diff --git a/patroni/utils.py b/patroni/utils.py index a52a69e2..fae246dc 100644 --- a/patroni/utils.py +++ b/patroni/utils.py @@ -1,3 +1,13 @@ +"""Utilitary objects and functions that can be used throughout Patroni code. + +:var tzutc: UTC time zone info object. +:var logger: logger of this module. +:var USER_AGENT: identifies the Patroni version, Python version, and the underlying platform. +:var OCT_RE: regular expression to match octal numbers, signed or unsigned. +:var DEC_RE: regular expression to match decimal numbers, signed or unsigned. +:var HEX_RE: regular expression to match hex strings, signed or unsigned. +:var DBL_RE: regular expression to match double precision numbers, signed or unsigned. Matches scientific notation too. +""" import errno import json.decoder as json_decoder import logging @@ -10,11 +20,17 @@ import sys import tempfile import time +from typing import Any, Callable, Dict, Iterator, List, Optional, Tuple, Union, TYPE_CHECKING + from dateutil import tz +from urllib3.response import HTTPResponse from .exceptions import PatroniException from .version import __version__ +if TYPE_CHECKING: + from .dcs import Cluster + tzutc = tz.tzutc() logger = logging.getLogger(__name__) @@ -26,20 +42,34 @@ HEX_RE = re.compile(r'^[-+]?0x[0-9a-fA-F]+') DBL_RE = re.compile(r'^[-+]?[0-9]*\.?[0-9]+([eE][-+]?[0-9]+)?') -def deep_compare(obj1, obj2): - """ - >>> deep_compare({'1': None}, {}) - False - >>> deep_compare({'1': {}}, {'1': None}) - False - >>> deep_compare({'1': [1]}, {'1': [2]}) - False - >>> deep_compare({'1': 2}, {'1': '2'}) - True - >>> deep_compare({'1': {'2': [3, 4]}}, {'1': {'2': [3, 4]}}) - True - """ +def deep_compare(obj1: Dict, obj2: Dict) -> bool: + """Recursively compare two dictionaries to check if they are equal in terms of keys and values. + .. note:: + Values are compared based on their string representation. + + :param obj1: dictionary to be compared with *obj2*. + :param obj2: dictionary to be compared with *obj1*. + + :returns: ``True`` if all keys and values match between the two dictionaries. + + :Example: + + >>> deep_compare({'1': None}, {}) + False + + >>> deep_compare({'1': {}}, {'1': None}) + False + + >>> deep_compare({'1': [1]}, {'1': [2]}) + False + + >>> deep_compare({'1': 2}, {'1': '2'}) + True + + >>> deep_compare({'1': {'2': [3, 4]}}, {'1': {'2': [3, 4]}}) + True + """ if set(list(obj1.keys())) != set(list(obj2.keys())): # Objects have different sets of keys return False @@ -52,9 +82,21 @@ def deep_compare(obj1, obj2): return True -def patch_config(config, data): - """recursively 'patch' `config` with `data` - :returns: `!True` if the `config` was changed""" +def patch_config(config: Dict, data: Dict) -> bool: + """Update and append to dictionary *config* from overrides in *data*. + + .. note:: + + * If the value of a given key in *data* is ``None``, then the key is removed from *config*; + * If a key is present in *data* but not in *config*, the key with the corresponding value is added to *config* + * For keys that are present on both sides it will compare the string representation of the corresponding values, + if the comparison doesn't match override the value + + :param config: configuration to be patched. + :param data: new configuration values to patch *config* with. + + :returns: ``True`` if *config* was changed. + """ is_changed = False for name, value in data.items(): if value is None: @@ -77,13 +119,28 @@ def patch_config(config, data): return is_changed -def parse_bool(value): - """ - >>> parse_bool(1) - True - >>> parse_bool('off') - False - >>> parse_bool('foo') +def parse_bool(value: Any) -> Union[bool, None]: + """Parse a given value to a :class:`bool` object. + + .. note:: + + The parsing is case-insensitive, and takes into consideration these values: + * ``on``, ``true``, ``yes``, and ``1`` as ``True``. + * ``off``, ``false``, ``no``, and ``0`` as ``False``. + + :param value: value to be parsed to :class:`bool`. + + :returns: the parsed value. If not able to parse, returns ``None``. + + :Example: + + >>> parse_bool(1) + True + + >>> parse_bool('off') + False + + >>> parse_bool('foo') """ value = str(value).lower() if value in ('on', 'true', 'yes', '1'): @@ -92,27 +149,50 @@ def parse_bool(value): return False -def strtol(value, strict=True): - """As most as possible close equivalent of strtol(3) function (with base=0), - used by postgres to parse parameter values. - >>> strtol(0) == (0, '') - True - >>> strtol(1) == (1, '') - True - >>> strtol(9) == (9, '') - True - >>> strtol(' +0x400MB') == (1024, 'MB') - True - >>> strtol(' -070d') == (-56, 'd') - True - >>> strtol(' d ') == (None, 'd') - True - >>> strtol(' 1 d ') == (1, ' d') - True - >>> strtol('9s', False) == (9, 's') - True - >>> strtol(' s ', False) == (1, 's') - True +def strtol(value: Any, strict: Optional[bool] = True) -> Tuple[Union[int, None], str]: + """Extract the long integer part from the beginning of a string that represents a configuration value. + + As most as possible close equivalent of ``strtol(3)`` C function (with base=0), which is used by postgres to parse + parameter values. + + Takes into consideration numbers represented either as hex, octal or decimal formats. + + :param value: any value from which we want to extract a long integer. + :param strict: dictates how the first item in the returning tuple is set when :func:`strtol` is not able to find a + long integer in *value*. If *strict* is ``True``, then the first item will be ``None``, else it will be ``1``. + + :returns: the first item is the extracted long integer from *value*, and the second item is the remaining string of + *value*. If not able to match a long integer in *value*, then the first item will be either ``None`` or ``1`` + (depending on *strict* argument), and the second item will be the original *value*. + + :Example: + + >>> strtol(0) == (0, '') + True + + >>> strtol(1) == (1, '') + True + + >>> strtol(9) == (9, '') + True + + >>> strtol(' +0x400MB') == (1024, 'MB') + True + + >>> strtol(' -070d') == (-56, 'd') + True + + >>> strtol(' d ') == (None, 'd') + True + + >>> strtol(' 1 d ') == (1, ' d') + True + + >>> strtol('9s', False) == (9, 's') + True + + >>> strtol(' s ', False) == (1, 's') + True """ value = str(value).strip() for regex, base in ((HEX_RE, 16), (OCT_RE, 8), (DEC_RE, 10)): @@ -123,10 +203,31 @@ def strtol(value, strict=True): return (None if strict else 1), value -def strtod(value): - """As most as possible close equivalent of strtod(3) function used by postgres to parse parameter values. - >>> strtod(' A ') == (None, 'A') - True +def strtod(value: Any) -> Tuple[Union[float, None], str]: + """Extract the double precision part from the beginning of a string that reprensents a configuration value. + + As most as possible close equivalent of ``strtod(3)`` C function, which is used by postgres to parse parameter + values. + + :param value: any value from which we want to extract a double precision. + + :returns: the first item is the extracted double precision from *value*, and the second item is the remaining + string of *value*. If not able to match a double precision in *value*, then the first item will be ``None``, + and the second item will be the original *value*. + + :Example: + + >>> strtod(' A ') == (None, 'A') + True + + >>> strtod('1 A ') == (1.0, ' A') + True + + >>> strtod('1.5A') == (1.5, 'A') + True + + >>> strtod('8.325e-10A B C') == (8.325e-10, 'A B C') + True """ value = str(value).strip() match = DBL_RE.match(value) @@ -136,21 +237,35 @@ def strtod(value): return None, value -def rint(value): +def convert_to_base_unit(value: Union[int, float], unit: str, base_unit: str) -> Union[int, float, None]: + """Convert *value* as a *unit* of compute information or time to *base_unit*. + + :param value: value to be converted to the base unit. + :param unit: unit of *value*. Accepts these units (case sensitive) + * For space: ``B``, ``kB``, ``MB``, ``GB``, or ``TB``; + * For time: ``d``, ``h``, ``min``, ``s``, ``ms``, or ``us``. + + :param base_unit: target unit in the conversion. May contain the target unit with an associated value, e.g + ``512MB``. Accepts these units (case sensitive) + * For space: ``B``, ``kB``, or ``MB``; + * For time: ``ms``, ``s``, or ``min``. + + :returns: *value* in *unit* converted to *base_unit*. Returns ``None`` if *unit* or *base_unit* is invalid. + + :Example: + + >>> convert_to_base_unit(1, 'GB', '256MB') + 4 + + >>> convert_to_base_unit(1, 'GB', 'MB') + 1024 + + >>> convert_to_base_unit(1, 'gB', '512MB') is None + True + + >>> convert_to_base_unit(1, 'GB', '512 MB') is None + True """ - >>> rint(0.5) == 0 - True - >>> rint(0.501) == 1 - True - >>> rint(1.5) == 2 - True - """ - - ret = round(value) - return 2.0 * round(value / 2.0) if abs(ret - value) == 0.5 else ret - - -def convert_to_base_unit(value, unit, base_unit): convert = { 'B': {'B': 1, 'kB': 1024, 'MB': 1024 * 1024, 'GB': 1024 * 1024 * 1024, 'TB': 1024 * 1024 * 1024 * 1024}, 'kB': {'B': 1.0 / 1024, 'kB': 1, 'MB': 1024, 'GB': 1024 * 1024, 'TB': 1024 * 1024 * 1024}, @@ -175,33 +290,50 @@ def convert_to_base_unit(value, unit, base_unit): if unit in round_order: multiplier = convert[base_unit][round_order[unit]] - value = rint(value / float(multiplier)) * multiplier + value = round(value / float(multiplier)) * multiplier return value -def parse_int(value, base_unit=None): - """ - >>> parse_int('1') == 1 - True - >>> parse_int(' 0x400 MB ', '16384kB') == 64 - True - >>> parse_int('1MB', 'kB') == 1024 - True - >>> parse_int('1000 ms', 's') == 1 - True - >>> parse_int('1TB', 'GB') is None - True - >>> parse_int(0) == 0 - True - >>> parse_int('6GB', '16MB') == 384 - True - >>> parse_int('4097.4kB', 'kB') == 4097 - True - >>> parse_int('4097.5kB', 'kB') == 4098 - True - """ +def parse_int(value: Any, base_unit: Optional[str] = None) -> Union[int, None]: + """Parse *value* as an :class:`int`. + :param value: any value that can be handled either by :func:`strtol` or :func:`strtod`. If *value* contains a + unit, then *base_unit* must be given. + :param base_unit: an optional base unit to convert *value* through :func:`convert_to_base_unit`. Not used if + *value* does not contain a unit. + + :returns: the parsed value, if able to parse. Otherwise returns ``None``. + + :Example: + + >>> parse_int('1') == 1 + True + + >>> parse_int(' 0x400 MB ', '16384kB') == 64 + True + + >>> parse_int('1MB', 'kB') == 1024 + True + + >>> parse_int('1000 ms', 's') == 1 + True + + >>> parse_int('1TB', 'GB') is None + True + + >>> parse_int(0) == 0 + True + + >>> parse_int('6GB', '16MB') == 384 + True + + >>> parse_int('4097.4kB', 'kB') == 4097 + True + + >>> parse_int('4097.5kB', 'kB') == 4098 + True + """ val, unit = strtol(value) if val is None and unit.startswith('.') or unit and unit[0] in ('.', 'e', 'E'): val, unit = strtod(value) @@ -209,21 +341,33 @@ def parse_int(value, base_unit=None): if val is not None: unit = unit.strip() if not unit: - return int(rint(val)) + return round(val) val = convert_to_base_unit(val, unit, base_unit) if val is not None: - return int(rint(val)) + return round(val) -def parse_real(value, base_unit=None): - """ - >>> parse_real(' +0.0005 ') == 0.0005 - True - >>> parse_real('0.0005ms', 'ms') == 0.0 - True - >>> parse_real('0.00051ms', 'ms') == 0.001 - True +def parse_real(value: Any, base_unit: Optional[str] = None) -> Union[float, None]: + """Parse *value* as a :class:`float`. + + :param value: any value that can be handled by :func:`strtod`. If *value* contains a unit, then *base_unit* must + be given. + :param base_unit: an optional base unit to convert *value* through :func:`convert_to_base_unit`. Not used if + *value* does not contain a unit. + + :returns: the parsed value, if able to parse. Otherwise returns ``None``. + + :Example: + + >>> parse_real(' +0.0005 ') == 0.0005 + True + + >>> parse_real('0.0005ms', 'ms') == 0.0 + True + + >>> parse_real('0.00051ms', 'ms') == 0.001 + True """ val, unit = strtod(value) @@ -235,14 +379,51 @@ def parse_real(value, base_unit=None): return convert_to_base_unit(val, unit, base_unit) -def compare_values(vartype, unit, old_value, new_value): - """ - >>> compare_values('enum', None, 'remote_write', 'REMOTE_WRITE') - True - >>> compare_values('real', None, '1e-06', 0.000001) - True - """ +def compare_values(vartype: str, unit: str, old_value: Any, new_value: Any) -> bool: + """Check if *old_value* and *new_value* are equivalent after parsing them as *vartype*. + :param vartpe: the target type to parse *old_value* and *new_value* before comparing them. Accepts any among of the + following (case sensitive) + * ``bool``: parse values using :func:`parse_bool`; or + * ``integer``: parse values using :func:`parse_int`; or + * ``real``: parse values using :func:`parse_real`; or + * ``enum``: parse values as lowercase strings; or + * ``string``: parse values as strings. This one is used by default if no valid value is passed as *vartype*. + :param unit: base unit to be used as argument when calling :func:`parse_int` or :func:`parse_real` for *new_value*. + :param old_value: value to be compared with *new_value*. + :param new_value: value to be compared with *old_value*. + + :returns: ``True`` if *old_value* is equivalent to *new_value* when both are parsed as *vartype*. + + :Example: + + >>> compare_values('enum', None, 'remote_write', 'REMOTE_WRITE') + True + + >>> compare_values('string', None, 'remote_write', 'REMOTE_WRITE') + False + + >>> compare_values('real', None, '1e-06', 0.000001) + True + + >>> compare_values('integer', 'MB', '6GB', '6GB') + False + + >>> compare_values('integer', None, '6GB', '6GB') + False + + >>> compare_values('integer', '16384kB', '64', ' 0x400 MB ') + True + + >>> compare_values('integer', '2MB', 524288, '1TB') + True + + >>> compare_values('integer', 'MB', 1048576, '1TB') + True + + >>> compare_values('integer', 'kB', 4098, '4097.5kB') + True + """ converters = { 'bool': lambda v1, v2: parse_bool(v1), 'integer': parse_int, @@ -258,30 +439,48 @@ def compare_values(vartype, unit, old_value, new_value): return old_value is not None and new_value is not None and old_value == new_value -def _sleep(interval): +def _sleep(interval: Union[int, float]) -> None: + """Wrap :func:`time.sleep`. + + :param interval: Delay execution for a given number of seconds. The argument may be a floating point number for + subsecond precision. + """ time.sleep(interval) class RetryFailedError(PatroniException): - - """Raised when retrying an operation ultimately failed, after retrying the maximum number of attempts.""" + """Maximum number of attempts exhausted in retry operation.""" class Retry(object): + """Helper for retrying a method in the face of retryable exceptions. - """Helper for retrying a method in the face of retry-able exceptions""" + :ivar max_tries: how many times to retry the command. + :ivar delay: initial delay between retry attempts. + :ivar backoff: backoff multiplier between retry attempts. + :ivar max_jitter: additional max jitter period to wait between retry attempts to avoid slamming the server. + :ivar max_delay: maximum delay in seconds, regardless of other backoff settings. + :ivar sleep_func: function used to introduce artificial delays. + :ivar deadline: timeout for operation retries. + :ivar retry_exceptions: single exception or tuple + """ - def __init__(self, max_tries=1, delay=0.1, backoff=2, max_jitter=0.8, max_delay=3600, - sleep_func=_sleep, deadline=None, retry_exceptions=PatroniException): - """Create a :class:`Retry` instance for retrying function calls - - :param max_tries: How many times to retry the command. -1 means infinite tries. - :param delay: Initial delay between retry attempts. - :param backoff: Backoff multiplier between retry attempts. Defaults to 2 for exponential backoff. - :param max_jitter: Additional max jitter period to wait between retry attempts to avoid slamming the server. - :param max_delay: Maximum delay in seconds, regardless of other backoff settings. Defaults to one hour. - :param retry_exceptions: single exception or tuple""" + def __init__(self, max_tries: Optional[int] = 1, delay: Optional[float] = 0.1, backoff: Optional[int] = 2, + max_jitter: Optional[float] = 0.8, max_delay: Optional[int] = 3600, + sleep_func: Optional[Callable[[Union[int, float]], None]] = _sleep, + deadline: Optional[Union[int, float]] = None, + retry_exceptions: Optional[Union[Exception, Tuple[Exception]]] = PatroniException) -> None: + """Create a :class:`Retry` instance for retrying function calls. + :param max_tries: how many times to retry the command. ``-1`` means infinite tries. + :param delay: initial delay between retry attempts. + :param backoff: backoff multiplier between retry attempts. Defaults to ``2`` for exponential backoff. + :param max_jitter: additional max jitter period to wait between retry attempts to avoid slamming the server. + :param max_delay: maximum delay in seconds, regardless of other backoff settings. + :param sleep_func: function used to introduce artificial delays. + :param deadline: timeout for operation retries. + :param retry_exceptions: single exception or tuple + """ self.max_tries = max_tries self.delay = delay self.backoff = backoff @@ -294,37 +493,59 @@ class Retry(object): self.sleep_func = sleep_func self.retry_exceptions = retry_exceptions - def reset(self): - """Reset the attempt counter""" + def reset(self) -> None: + """Reset the attempt counter, delay and stop time.""" self._attempts = 0 self._cur_delay = self.delay self._cur_stoptime = None - def copy(self): - """Return a clone of this retry manager""" + def copy(self) -> 'Retry': + """Return a clone of this retry manager.""" return Retry(max_tries=self.max_tries, delay=self.delay, backoff=self.backoff, max_jitter=self.max_jitter / 100.0, max_delay=self.max_delay, sleep_func=self.sleep_func, deadline=self.deadline, retry_exceptions=self.retry_exceptions) @property - def sleeptime(self): + def sleeptime(self) -> float: + """Get next cycle sleep time. + + It is based on the current delay plus a number up to ``max_jitter``. + """ return self._cur_delay + (random.randint(0, self.max_jitter) / 100.0) - def update_delay(self): + def update_delay(self) -> None: + """Set next cycle delay. + + It will be the minimum value between: + * current delay with ``backoff``; or + * ``max_delay``. + """ self._cur_delay = min(self._cur_delay * self.backoff, self.max_delay) @property - def stoptime(self): + def stoptime(self) -> Union[float, None]: + """Get the current stop time.""" return self._cur_stoptime - def __call__(self, func, *args, **kwargs): - """Call a function with arguments until it completes without throwing a `retry_exceptions` + def __call__(self, func: Callable, *args: Any, **kwargs: Any) -> Any: + """Call a function *func* with arguments ``*args`` and ``*kwargs`` in a loop. - :param func: Function to call - :param args: Positional arguments to call the function with - :params kwargs: Keyword arguments to call the function with + *func* will be called until one of the following conditions is met: + * It completes without throwing one of the configured ``retry_exceptions``; or + * ``max_retries`` is exceeded.; or + * ``deadline`` is exceeded. - The function will be called until it doesn't throw one of the retryable exceptions""" + .. note:: + * It will set loop stop time based on ``deadline`` attribute. + * It will adjust delay on each cycle. + + :param func: function to call. + :param args: positional arguments to call *func* with. + :params kwargs: keyword arguments to call *func* with. + :raises :class:`RetryFailedError` + * If ``max_tries`` is exceeded; or + * If ``deadline`` is exceeded. + """ self.reset() while True: @@ -348,8 +569,16 @@ class Retry(object): self.update_delay() -def polling_loop(timeout, interval=1): - """Returns an iterator that returns values until timeout has passed. Timeout is measured from start of iteration.""" +def polling_loop(timeout: Union[int, float], interval: Optional[Union[int, float]] = 1) -> Iterator[int]: + """Return an iterator that returns values every *interval* seconds until *timeout* has passed. + + .. note:: + Timeout is measured from start of iteration. + + :param timeout: for how long (in seconds) from now it should keep returning values. + :param interval: for how long to sleep before returning a new value. + :rtype: Iterator[:class:`int`] with current iteration counter, starting from ``0``. + """ start_time = time.time() iteration = 0 end_time = start_time + timeout @@ -359,16 +588,71 @@ def polling_loop(timeout, interval=1): time.sleep(interval) -def split_host_port(value, default_port): +def split_host_port(value: str, default_port: int) -> Tuple[str, int]: + """Extract host(s) and port from *value*. + + :param value: string from where host(s) and port will be extracted. Accepts either of these formats + * ``host:port``; or + * ``host1,host2,...,hostn:port``. + + Each ``host`` portion of *value* can be either: + * A FQDN; or + * An IPv4 address; or + * An IPv6 address, with or without square brackets. + + :param default_port: if no port can be found in *param*, use *default_port* instead. + + :returns: the first item is composed of a CSV list of hosts from *value*, and the second item is either the port + from *value* or *default_port*. + + :Example: + + >>> split_host_port('127.0.0.1', 5432) + ('127.0.0.1', 5432) + + >>> split_host_port('127.0.0.1:5400', 5432) + ('127.0.0.1', 5400) + + >>> split_host_port('127.0.0.1,192.168.0.101:5400', 5432) + ('127.0.0.1,192.168.0.101', 5400) + + >>> split_host_port('127.0.0.1,www.mydomain.com,[fe80:0:0:0:213:72ff:fe3c:21bf], 0:0:0:0:0:0:0:0:5400', 5432) + ('127.0.0.1,www.mydomain.com,fe80:0:0:0:213:72ff:fe3c:21bf,0:0:0:0:0:0:0:0', 5400) + """ t = value.rsplit(':', 1) + # If *value* contains ``:`` we consider it to be an IPv6 address, so we attempt to remove possible square brackets if ':' in t[0]: t[0] = ','.join([h.strip().strip('[]') for h in t[0].split(',')]) t.append(default_port) return t[0], int(t[1]) -def uri(proto, netloc, path='', user=None): +def uri(proto: str, netloc: Union[List, Tuple[str, int], str], path: Optional[str] = '', + user: Optional[str] = None) -> str: + """Construct URI from given arguments. + + :param proto: the URI protocol. + :param netloc: the URI host(s) and port. Can be specified in either way among + * A :class:`list` or :class:`tuple`. The second item should be a port, and the first item should be composed of + hosts in either of these formats: + * ``host``; or. + * ``host1,host2,...,hostn``. + * A :class:`str` in either of these formats: + * ``host:port``; or + * ``host1,host2,...,hostn:port``. + + In all cases, each ``host`` portion of *netloc* can be either: + * An FQDN; or + * An IPv4 address; or + * An IPv6 address, with or without square brackets. + + :param path: the URI path. + :param user: the authenticating user, if any. + + :returns: constructed URI. + """ host, port = netloc if isinstance(netloc, (list, tuple)) else split_host_port(netloc, 0) + # If ``host`` contains ``:`` we consider it to be an IPv6 address, so we add square brackets if they are missing if host and ':' in host and host[0] != '[' and host[-1] != ']': host = '[{0}]'.format(host) port = ':{0}'.format(port) if port else '' @@ -377,7 +661,12 @@ def uri(proto, netloc, path='', user=None): return '{0}://{1}{2}{3}{4}'.format(proto, user, host, port, path) -def iter_response_objects(response): +def iter_response_objects(response: HTTPResponse) -> Iterator[Dict[str, Any]]: + """Iterate over the chunks of a :class:`HTTPResponse` and yield each JSON document that is found along the way. + + :param response: the HTTP response from which JSON documents will be retrieved. + :rtype: Iterator[:class:`dict`] with current JSON document. + """ prev = '' decoder = json_decoder.JSONDecoder() for chunk in response.read_chunked(decode_content=False): @@ -386,24 +675,65 @@ def iter_response_objects(response): chunk = prev + chunk length = len(chunk) + # ``chunk`` is analyzed in parts. ``idx`` holds the position of the first character in the current part that is + # neither space nor tab nor line-break, or in other words, the position in the ``chunk`` where it is likely + # that a JSON document begins idx = json_decoder.WHITESPACE.match(chunk, 0).end() while idx < length: try: + # Get a JSON document from the chunk. ``message`` is a dictionary representing the JSON document, and + # ``idx`` becomes the position in the ``chunk`` where the retrieved JSON document ends message, idx = decoder.raw_decode(chunk, idx) except ValueError: # malformed or incomplete JSON, unlikely to happen break else: yield message idx = json_decoder.WHITESPACE.match(chunk, idx).end() + # It is not usual that a ``chunk`` would contain more than one JSON document, but we handle that just in case prev = chunk[idx:] -def is_standby_cluster(config): - # Check whether or not provided configuration describes a standby cluster +def is_standby_cluster(config: Union[Dict[str, Any], None]) -> bool: + """Check provided configuration describes a standby cluster. + + :param config: the configuration to be checked. It is expected to be the :class:`dict` that represents the value of + the ``standby_cluster`` key in the main Patroni configuration. ``None`` can be used if ``standby_cluster`` is + absent in the main Patroni configuration. + + :returns: ``True`` if configuration is a Patroni standby cluster. + """ return isinstance(config, dict) and (config.get('host') or config.get('port') or config.get('restore_command')) -def cluster_as_json(cluster): +def cluster_as_json(cluster: 'Cluster') -> Dict[str, Any]: + """Get a JSON representation of *cluster*. + + :param cluster: the :class:`Cluster` object to be parsed as JSON. + + :returns: JSON representation of *cluster*. + + These are the possible keys in the returning object depending on the available information in *cluster*: + + * ``members``: list of members in the cluster. Each value is a :class:`dict` that may have the following keys: + * ``name``: the name of the host (unique in the cluster). The ``members`` list is sorted by this key; + * ``role``: ``leader``, ``standby_leader``, ``sync_standby``, or ``replica``; + * ``state``: ``stopping``, ``stopped``, ``stop failed``, ``crashed``, ``running``, ``starting``, + ``start failed``, ``restarting``, ``restart failed``, ``initializing new cluster``, ``initdb failed``, + ``running custom bootstrap script``, ``custom bootstrap failed``, or ``creating replica``; + * ``api_url``: REST API URL based on ``restapi->connect_address`` configuration; + * ``host``: PostgreSQL host based on ``postgresql->connect_address``; + * ``port``: PostgreSQL port based on ``postgresql->connect_address``; + * ``timeline``: PostgreSQL current timeline; + * ``pending_restart``: ``True`` if PostgreSQL is pending to be restarted; + * ``scheduled_restart``: scheduled restart timestamp, if any; + * ``tags``: any tags that were set for this member; + * ``lag``: replication lag, if applicable; + * ``pause``: ``True`` if cluster is in maintenance mode; + * ``scheduled_switchover``: if a switchover has been scheduled, then it contains this entry with these keys: + * ``at``: timestamp when switchover was scheduled to occur; + * ``from``: name of the member to be demoted; + * ``to``: name of the member to be promoted. + """ leader_name = cluster.leader.name if cluster.leader else None cluster_lsn = cluster.last_lsn or 0 @@ -450,13 +780,40 @@ def cluster_as_json(cluster): return ret -def is_subpath(d1, d2): +def is_subpath(d1: str, d2: str) -> bool: + """Check if the file system path *d2* is contained within *d1* after resolving symbolic links. + + .. note:: + It will not check if the paths actually exist, it will only expand the paths and resolve any symbolic links + that happen to be found. + + :param d1: path to a directory. + :param d2: path to be checked if is within *d1*. + + :returns: ``True`` if *d1* is a subpath of *d2*. + """ real_d1 = os.path.realpath(d1) + os.path.sep real_d2 = os.path.realpath(os.path.join(real_d1, d2)) return os.path.commonprefix([real_d1, real_d2 + os.path.sep]) == real_d1 -def validate_directory(d, msg="{} {}"): +def validate_directory(d: str, msg: Optional[str] = "{} {}") -> None: + """Ensure directory exists and is writable. + + .. note:: + If the directory does not exist, :func:`validate_directory` will attempt to create it. + + :param d: the directory to be checked. + :param msg: a message to be thrown when raising :class:`PatroniException`, if any issue is faced. It must contain + 2 placeholders to be used by :func:`format`: + * The first placeholder will be replaced with path *d*; + * The second placeholder will be replaced with the error condition. + + :raises :class:`PatroniException`: if any issue is observed while validating *d*. Can be thrown in these situations + * *d* did not exist, and :func:`validate_directory` was not able to create it; or + * *d* is an existing directory, but Patroni is not able to write to that directory; or + * *d* is an existing file, not a directory. + """ if not os.path.exists(d): try: os.makedirs(d) @@ -475,17 +832,49 @@ def validate_directory(d, msg="{} {}"): raise PatroniException(msg.format(d, "is not a directory")) -def data_directory_is_empty(data_dir): +def data_directory_is_empty(data_dir: str) -> bool: + """Check if a PostgreSQL data directory is empty. + + .. note:: + In non-Windows environments *data_dir* is also considered empty if it only contains hidden files and/or + ``lost+found`` directory. + + :param data_dir: the PostgreSQL data directory to be checked. + + :returns: ``True`` if *data_dir* is empty. + """ if not os.path.exists(data_dir): return True return all(os.name != 'nt' and (n.startswith('.') or n == 'lost+found') for n in os.listdir(data_dir)) -def keepalive_intvl(timeout, idle, cnt=3): +def keepalive_intvl(timeout: int, idle: int, cnt: Optional[int] = 3) -> int: + """Calculate the value to be used as ``TCP_KEEPINTVL`` based on *timeout*, *idle*, and *cnt*. + + :param timeout: value for ``TCP_USER_TIMEOUT``. + :param idle: value for ``TCP_KEEPIDLE``. + :param cnt: value for ``TCP_KEEPCNT``. + + :returns: the value to be used as ``TCP_KEEPINTVL``. + """ return max(1, int(float(timeout - idle) / cnt)) -def keepalive_socket_options(timeout, idle, cnt=3): +def keepalive_socket_options(timeout: int, idle: int, cnt: Optional[int] = 3) -> Iterator[Tuple[int, int, int]]: + """Get all keepalive related options to be set in a socket. + + :param timeout: value for ``TCP_USER_TIMEOUT``. + :param idle: value for ``TCP_KEEPIDLE``. + :param cnt: value for ``TCP_KEEPCNT``. + + :rtype: Iterator[Tuple[:class:`int`, :class:`int`, :class:`int`]] of all keepalive related socket options to be + set. The first item in the tuple is the protocol, the second item is the option, and the third item is the + value to be used. The return values depend on the platform: + * ``Windows``: yield ``SO_KEEPALIVE``; + * ``Linux``: yield ``SO_KEEPALIVE``, ``TCP_USER_TIMEOUT``, ``TCP_KEEPIDLE`, ``TCP_KEEPINTVL``, and + ``TCP_KEEPCNT``; + * ``MacOS``: yield ``SO_KEEPALIVE``, ``TCP_KEEPIDLE`, ``TCP_KEEPINTVL``, and ``TCP_KEEPCNT`` + """ yield (socket.SOL_SOCKET, socket.SO_KEEPALIVE, 1) if sys.platform.startswith('linux'): @@ -506,7 +895,22 @@ def keepalive_socket_options(timeout, idle, cnt=3): yield (socket.IPPROTO_TCP, TCP_KEEPCNT, cnt) -def enable_keepalive(sock, timeout, idle, cnt=3): +def enable_keepalive(sock: socket.socket, timeout: int, idle: int, cnt: Optional[int] = 3) -> Union[int, None]: + """Enable keepalive for *sock*. + + Will set socket options depending on the platform, as per return of :func:`keepalive_socket_options`. + + .. note:: + Value for ``TCP_KEEPINTVL`` will be calculated through :func:`keepalive_intvl` based on *timeout*, *idle*, and + *cnt*. + + :param sock: the socket for which keepalive will be enabled. + :param timeout: value for ``TCP_USER_TIMEOUT``. + :param idle: value for ``TCP_KEEPIDLE``. + :param cnt: value for ``TCP_KEEPCNT``. + + :returns: output of :func:`socket.ioctl` if we are on Windows, nothing otherwise. + """ SIO_KEEPALIVE_VALS = getattr(socket, 'SIO_KEEPALIVE_VALS', None) if SIO_KEEPALIVE_VALS is not None: # Windows intvl = keepalive_intvl(timeout, idle, cnt)