diff --git a/patroni/postgresql.py b/patroni/postgresql.py index b61ae77b..d13b33f8 100644 --- a/patroni/postgresql.py +++ b/patroni/postgresql.py @@ -8,7 +8,7 @@ import tempfile import time from patroni.exceptions import PostgresConnectionException, PostgresException -from patroni.utils import Retry, RetryFailedError +from patroni.utils import compare_values, Retry, RetryFailedError from six import string_types from six.moves.urllib_parse import urlparse from threading import Lock @@ -51,6 +51,7 @@ class Postgresql(object): CMDLINE_OPTIONS = { 'listen_addresses': None, 'port': None, + 'config_file': None, 'wal_level': 'hot_standby', 'hot_standby': 'on', 'max_wal_senders': 5, @@ -136,21 +137,23 @@ class Postgresql(object): def reload_config(self, config): server_parameters = self.get_server_parameters(config) - listen_address_changed = reload_pending = False + listen_address_changed = pending_reload = False if self.is_healthy(): changes = server_parameters.copy() changes.update({p: None for p, v in self._server_parameters.items() if p not in server_parameters}) if changes: - for r in self.query("""SELECT name, setting, context + for r in self.query("""SELECT name, setting, unit, vartype, context FROM pg_settings WHERE name IN (""" + ', '.join('%s' for _ in changes.keys()) + ')', *(list(changes.keys()))): - if server_parameters[r[0]] is None or str(server_parameters[r[0]]) != str(r[1]): - reload_pending = True - if r[2] in ('internal', 'postmaster'): + unit = '16384kB' if r[0] in ('min_wal_size', 'max_wal_size') else r[2] + if server_parameters[r[0]] is None or not compare_values(r[3], unit, r[1], server_parameters[r[0]]): + if r[4] == 'postmaster': self._pending_restart = True if r[0] in ('listen_addresses', 'port'): listen_address_changed = True + elif r[4] != 'internal': + pending_reload = True self.config = config self._server_parameters = server_parameters self._connect_address = config.get('connect_address') @@ -158,7 +161,7 @@ class Postgresql(object): if not listen_address_changed: self.resolve_connection_addresses() - if reload_pending: + if pending_reload: self._write_postgresql_conf() self.reload() self.retry.deadline = config['retry_timeout']/2.0 @@ -451,7 +454,8 @@ class Postgresql(object): self.resolve_connection_addresses() options = ' '.join("--{0}='{1}'".format(p, self._server_parameters[p]) for p in self.CMDLINE_OPTIONS - if not (self._major_version < 9.4 and p in ('max_replication_slots', 'wal_log_hints'))) + if self._server_parameters[p] is not None and + not (self._major_version < 9.4 and p in ('max_replication_slots', 'wal_log_hints'))) ret = subprocess.call(self._pg_ctl + ['start', '-o', options], env=env, preexec_fn=os.setsid) == 0 self._pending_restart = False diff --git a/patroni/utils.py b/patroni/utils.py index d5e8bc41..f85ec56f 100644 --- a/patroni/utils.py +++ b/patroni/utils.py @@ -2,6 +2,7 @@ import datetime import os import random import signal +import six import sys import time import pytz @@ -59,6 +60,95 @@ def deep_compare(obj1, obj2): return True +def parse_bool(value): + """ + >>> parse_bool(1) + True + >>> parse_bool('off') + False + >>> parse_bool('foo') + """ + value = str(value).lower() + if value in ('on', 'true', 'yes', '1'): + return True + if value in ('off', 'false', 'no', '0'): + return False + + +def split_int_unit(value, strict=True): + value = str(value) + l = len(value) - 1 + while l >= 0 and not value[l].isdigit(): + l -= 1 + unit = value[l + 1:].strip() + try: + value = int(value[:l + 1], 0) if six.PY3 else long(value[:l + 1], 0) + except ValueError: + value = None if strict else 1 + return (value, unit) + + +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('1GB', 'MB') is None + True + """ + + convert = { + 'kB': {'kB': 1, 'MB': 1024, 'GB': 1024 * 1024, 'TB': 1024 * 1024 * 1024}, + 'ms': {'ms': 1, 's': 1000, 'min': 1000 * 60, 'h': 1000 * 60 * 60, 'd': 1000 * 60 * 60 * 24}, + 's': {'ms': -1000, 's': 1, 'min': 60, 'h': 60 * 60, 'd': 60 * 60 * 24}, + 'min': {'ms': -1000 * 60, 's': -60, 'min': 1, 'h': 60, 'd': 60 * 24} + } + + value, unit = split_int_unit(value) + if value is not None: + if not unit: + return value + + if base_unit and base_unit not in convert: + base_value, base_unit = split_int_unit(base_unit, False) + else: + base_value = 1 + if base_unit in convert and unit in convert[base_unit]: + multiplier = convert[base_unit][unit] + if multiplier < 0: + value /= -multiplier + else: + value *= multiplier + return int(value/base_value) + + +def compare_values(vartype, unit, old_value, new_value): + """ + >>> compare_values('enum', None, 'remote_write', 'REMOTE_WRITE') + True + >>> compare_values('real', None, '1.23', 1.23) + True + """ + + # if the integer or bool new_value is not correct this function will return False + if vartype == 'bool': + old_value = parse_bool(old_value) + new_value = parse_bool(new_value) + elif vartype == 'integer': + old_value = parse_int(old_value) + new_value = parse_int(new_value, unit) + elif vartype == 'enum': + return str(old_value).lower() == str(new_value).lower() + else: # ('string', 'real') + return str(old_value) == str(new_value) + return old_value is not None and new_value is not None and old_value == new_value + + def set_ignore_sigterm(value=True): global __ignore_sigterm __ignore_sigterm = value diff --git a/tests/test_postgresql.py b/tests/test_postgresql.py index 31ec0c60..3107ed9b 100644 --- a/tests/test_postgresql.py +++ b/tests/test_postgresql.py @@ -35,7 +35,8 @@ class MockCursor(object): elif sql.startswith('SELECT to_char(pg_postmaster_start_time'): self.results = [('', True, '', '', '', '', False)] elif sql.startswith('SELECT name, setting'): - self.results = [('port', '5433', 'postmaster')] + self.results = [('port', '5433', None, 'integer', 'postmaster'), + ('autovacuum', 'on', None, 'bool', 'sighup')] else: self.results = [(None, None, None, None, None, None, None, None, None, None)] @@ -149,7 +150,7 @@ def fake_listdir(path): @patch('subprocess.call', Mock(return_value=0)) @patch('psycopg2.connect', psycopg2_connect) class TestPostgresql(unittest.TestCase): - _PARAMETERS = {'wal_level': 'hot_standby', 'max_replication_slots': 5, 'foo': 'bar', + _PARAMETERS = {'wal_level': 'hot_standby', 'max_replication_slots': 5, 'foo': 'bar', 'config_file': None, 'hot_standby': 'on', 'max_wal_senders': 5, 'wal_keep_segments': 8, 'wal_log_hints': 'on'} @patch('subprocess.call', Mock(return_value=0)) @@ -475,8 +476,11 @@ class TestPostgresql(unittest.TestCase): self.assertFalse(self.p.replica_method_can_work_without_replication_connection('foo')) def test_reload_config(self): - self.p.reload_config({'retry_timeout': 10, 'listen': '*', 'parameters': self._PARAMETERS}) - self.p.reload_config({'retry_timeout': 10, 'listen': '*:5433', 'parameters': self._PARAMETERS}) + parameters = self._PARAMETERS.copy() + parameters['autovacuum'] = 'on' + self.p.reload_config({'retry_timeout': 10, 'listen': '*', 'parameters': parameters}) + parameters['autovacuum'] = 'off' + self.p.reload_config({'retry_timeout': 10, 'listen': '*:5433', 'parameters': parameters}) @patch.object(builtins, 'open', mock_open(read_data='9.4')) def test_get_major_version(self):