From 2c7b547a2980e47ff5f807c3ad8d4169b64f0355 Mon Sep 17 00:00:00 2001 From: Alexander Kukushkin Date: Mon, 3 Apr 2023 11:19:08 +0200 Subject: [PATCH] Introduce patroni.collections (#2629) For now it implements: - CaseInsensitiveDict() - CaseInsensitiveSet() Update `patroni.postgresql.sync.parse_sync_standby_names()` to use `CaseInsensitiveSet()` instead of `CaseInsensitiveDict()` --- patroni/collections.py | 74 +++++++++++++++++++++++++++++++++ patroni/config.py | 5 ++- patroni/postgresql/config.py | 4 +- patroni/postgresql/sync.py | 26 ++++++------ patroni/postgresql/validator.py | 17 +------- tests/test_postgresql.py | 2 +- 6 files changed, 94 insertions(+), 34 deletions(-) create mode 100644 patroni/collections.py diff --git a/patroni/collections.py b/patroni/collections.py new file mode 100644 index 00000000..a42060bb --- /dev/null +++ b/patroni/collections.py @@ -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)) diff --git a/patroni/config.py b/patroni/config.py index 1fca41b9..6b645725 100644 --- a/patroni/config.py +++ b/patroni/config.py @@ -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__) diff --git a/patroni/postgresql/config.py b/patroni/postgresql/config.py index d2a097bc..3f3d8060 100644 --- a/patroni/postgresql/config.py +++ b/patroni/postgresql/config.py @@ -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, \ diff --git a/patroni/postgresql/sync.py b/patroni/postgresql/sync.py index dd75c293..43c3d9f2 100644 --- a/patroni/postgresql/sync.py +++ b/patroni/postgresql/sync.py @@ -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 \) ) | (?P . ) """, 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)) diff --git a/patroni/postgresql/validator.py b/patroni/postgresql/validator.py index 61eb1bde..95162d3e 100644 --- a/patroni/postgresql/validator.py +++ b/patroni/postgresql/validator.py @@ -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 diff --git a/tests/test_postgresql.py b/tests/test_postgresql.py index bab65159..31a807f5 100644 --- a/tests/test_postgresql.py +++ b/tests/test_postgresql.py @@ -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('