Introduce patroni.collections (#2629)

For now it implements:
- CaseInsensitiveDict()
- CaseInsensitiveSet()

Update `patroni.postgresql.sync.parse_sync_standby_names()` to use `CaseInsensitiveSet()` instead of `CaseInsensitiveDict()`
This commit is contained in:
Alexander Kukushkin
2023-04-03 11:19:08 +02:00
committed by GitHub
parent 3fe2a7868a
commit 2c7b547a29
6 changed files with 94 additions and 34 deletions
+74
View File
@@ -0,0 +1,74 @@
from collections import OrderedDict
from collections.abc import MutableMapping, MutableSet
from typing import Any, Collection, Dict, Iterable, Iterator, Optional, Tuple, Union
class CaseInsensitiveSet(MutableSet):
"""A case-insensitive ``set``-like object.
Implements all methods and operations of :class:``MutableSet``. All values are expected to be strings.
The structure remembers the case of the last value set, however, contains testing is case insensitive.
"""
def __init__(self, values: Optional[Collection[str]] = None) -> None:
self._values = {}
for v in values or ():
self.add(v)
def __repr__(self) -> str:
return '<{0}{1} at {2:x}>'.format(type(self).__name__, tuple(self._values.values()), id(self))
def __str__(self) -> str:
return str(set(self._values.values()))
def __contains__(self, value: str) -> bool:
return value.lower() in self._values
def __iter__(self) -> Iterator[str]:
return iter(self._values.values())
def __len__(self) -> int:
return len(self._values)
def add(self, value: str) -> None:
self._values[value.lower()] = value
def discard(self, value: str) -> None:
self._values.pop(value.lower(), None)
def issubset(self, other: 'CaseInsensitiveSet'):
return self <= other
class CaseInsensitiveDict(MutableMapping):
"""A case-insensitive ``dict``-like object.
Implements all methods and operations of :class:``MutableMapping`` as well as dict's :func:``copy``.
All keys are expected to be strings. The structure remembers the case of the last key to be set,
and ``iter(instance)``, ``keys()``, ``items()``, ``iterkeys()``, and ``iteritems()`` will contain
case-sensitive keys. However, querying and contains testing is case insensitive.
"""
def __init__(self, data: Optional[Union[Dict[str, Any], Iterable[Tuple[str, Any]]]] = None) -> None:
self._values = OrderedDict()
self.update(data or {})
def __setitem__(self, key: str, value: Any) -> None:
# Use the lowercase key for lookups, but store the actual key alongside the value.
self._values[key.lower()] = (key, value)
def __getitem__(self, key: str) -> Any:
return self._values[key.lower()][1]
def __delitem__(self, key: str) -> Any:
del self._values[key.lower()]
def __iter__(self) -> Iterator[str]:
return iter(key for key, _ in self._values.values())
def __len__(self) -> int:
return len(self._values)
def copy(self) -> 'CaseInsensitiveDict':
return CaseInsensitiveDict(self._values.values())
def __repr__(self) -> str:
return '<{0}{1} at {2:x}>'.format(type(self).__name__, dict(self.items()), id(self))
+3 -2
View File
@@ -10,9 +10,10 @@ from copy import deepcopy
from typing import Any, Dict, Optional, Union
from . import PATRONI_ENV_PREFIX
from .exceptions import ConfigParseError
from .collections import CaseInsensitiveDict
from .dcs import ClusterConfig, Cluster
from .postgresql.config import CaseInsensitiveDict, ConfigHandler
from .exceptions import ConfigParseError
from .postgresql.config import ConfigHandler
from .utils import deep_compare, parse_bool, parse_int, patch_config
logger = logging.getLogger(__name__)
+2 -2
View File
@@ -8,8 +8,8 @@ import time
from urllib.parse import urlparse, parse_qsl, unquote
from .validator import CaseInsensitiveDict, recovery_parameters,\
transform_postgresql_parameter_value, transform_recovery_parameter_value
from .validator import recovery_parameters, transform_postgresql_parameter_value, transform_recovery_parameter_value
from ..collections import CaseInsensitiveDict
from ..dcs import RemoteMember, slot_name_from_member_name
from ..exceptions import PatroniFatalException
from ..utils import compare_values, parse_bool, parse_int, split_host_port, uri, \
+13 -13
View File
@@ -4,7 +4,7 @@ import time
from copy import deepcopy
from .validator import CaseInsensitiveDict
from ..collections import CaseInsensitiveDict, CaseInsensitiveSet
from ..psycopg import quote_ident as _quote_ident
logger = logging.getLogger(__name__)
@@ -23,7 +23,7 @@ SYNC_REP_PARSER_RE = re.compile(r"""
| (?P<parenend> \) )
| (?P<JUNK> . )
""", re.X)
_EMPTY_SSN = {'type': 'off', 'num': 0, 'members': CaseInsensitiveDict({})}
_EMPTY_SSN = {'type': 'off', 'num': 0, 'members': CaseInsensitiveSet()}
def quote_ident(value):
@@ -36,7 +36,7 @@ def parse_sync_standby_names(value):
Returns dict with the following keys:
* type: 'quorum'|'priority'
* num: int
* members: CaseInsensitiveDict, with names as keys
* members: CaseInsensitiveSet, with name
* has_star: bool - Present if true
If the configuration value can not be parsed, raises a ValueError.
@@ -46,14 +46,14 @@ def parse_sync_standby_names(value):
>>> parse_sync_standby_names('FiRsT')['type']
'priority'
>>> parse_sync_standby_names('FiRsT')['members']
{'FiRsT': True}
>>> 'first' in parse_sync_standby_names('FiRsT')['members']
True
>>> parse_sync_standby_names('"1"')['members']
{'1': True}
>>> set(parse_sync_standby_names('"1"')['members'])
{'1'}
>>> parse_sync_standby_names(' a , b ')['members']
{'a': True, 'b': True}
>>> parse_sync_standby_names(' a , b ')['members'] == {'a', 'b'}
True
>>> parse_sync_standby_names(' a , b ')['num']
1
@@ -107,7 +107,7 @@ def parse_sync_standby_names(value):
else:
result = {'type': 'priority', 'num': 1}
synclist = tokens
result['members'] = CaseInsensitiveDict({})
result['members'] = CaseInsensitiveSet()
for i, (a_type, a_value, a_pos) in enumerate(synclist):
if i % 2 == 1: # odd elements are supposed to be commas
if len(synclist) == i + 1: # except the last token
@@ -117,12 +117,12 @@ def parse_sync_standby_names(value):
raise ValueError("Unparseable synchronous_standby_names value %r: ""Got token %s %r while"
" expecting comma at %d" % (value, a_type, a_value, a_pos))
elif a_type in {'ident', 'first', 'any'}:
result['members'][a_value] = True
result['members'].add(a_value)
elif a_type == 'star':
result['members'][a_value] = True
result['members'].add(a_value)
result['has_star'] = True
elif a_type == 'dquot':
result['members'][a_value[1:-1].replace('""', '"')] = True
result['members'].add(a_value[1:-1].replace('""', '"'))
else:
raise ValueError("Unparseable synchronous_standby_names value %r: Unexpected token %s %r at %d" %
(value, a_type, a_value, a_pos))
+1 -16
View File
@@ -2,28 +2,13 @@ import abc
import logging
from collections import namedtuple
from urllib3.response import HTTPHeaderDict
from ..collections import CaseInsensitiveDict
from ..utils import parse_bool, parse_int, parse_real
logger = logging.getLogger(__name__)
class CaseInsensitiveDict(HTTPHeaderDict):
def add(self, key, val):
self[key] = val
def __getitem__(self, key):
return self._container[key.lower()][1]
def __repr__(self):
return str(dict(self.items()))
def copy(self):
return CaseInsensitiveDict(self._container.values())
class Bool(namedtuple('Bool', 'version_from,version_till')):
@staticmethod
+1 -1
View File
@@ -651,7 +651,7 @@ class TestPostgresql(BaseTestPostgresql):
self.p._global_config = GlobalConfig({'synchronous_mode': True, 'synchronous_mode_strict': True})
self.p.config.get_server_parameters(config)
self.p.config.set_synchronous_standby_names('foo')
self.assertTrue(str(self.p.config.get_server_parameters(config)).startswith('{'))
self.assertTrue(str(self.p.config.get_server_parameters(config)).startswith('<CaseInsensitiveDict'))
@patch('time.sleep', Mock())
def test__wait_for_connection_close(self):