diff --git a/patroni/api.py b/patroni/api.py index c82f679c..557de01a 100644 --- a/patroni/api.py +++ b/patroni/api.py @@ -8,7 +8,8 @@ import dateutil.parser import datetime import os -from patroni.postgresql import PostgresConnectionException, PostgresException, Postgresql +from patroni.postgresql import PostgresConnectionException +from patroni.postgresql.misc import postgres_version_to_int, PostgresException from patroni.utils import deep_compare, parse_bool, patch_config, Retry, \ RetryFailedError, parse_int, split_host_port, tzutc from six.moves.BaseHTTPServer import BaseHTTPRequestHandler, HTTPServer @@ -232,7 +233,7 @@ class RestApiHandler(BaseHTTPRequestHandler): break elif k == 'postgres_version': try: - Postgresql.postgres_version_to_int(request[k]) + postgres_version_to_int(request[k]) except PostgresException as e: status_code = 400 data = e.value diff --git a/patroni/async_executor.py b/patroni/async_executor.py index 730ebdeb..fec5506a 100644 --- a/patroni/async_executor.py +++ b/patroni/async_executor.py @@ -52,8 +52,8 @@ class CriticalTask(object): class AsyncExecutor(object): - def __init__(self, state_handler, ha_wakeup): - self.state_handler = state_handler + def __init__(self, cancellable, ha_wakeup): + self._cancellable = cancellable self._ha_wakeup = ha_wakeup self._thread_lock = RLock() self._scheduled_action = None @@ -92,7 +92,7 @@ class AsyncExecutor(object): return self._finish_event.clear() - self.state_handler.reset_is_cancelled() + self._cancellable.reset_is_cancelled() # if the func returned something (not None) - wake up main HA loop wakeup = func(*args) if args else func() return wakeup @@ -118,7 +118,7 @@ class AsyncExecutor(object): logger.warning('Cancelling long running task %s', self._scheduled_action) self._is_cancelled = True - self.state_handler.cancel() + self._cancellable.cancel() self._finish_event.wait() with self: diff --git a/patroni/config.py b/patroni/config.py index b99df660..c982637e 100644 --- a/patroni/config.py +++ b/patroni/config.py @@ -9,7 +9,7 @@ import yaml from collections import defaultdict from copy import deepcopy from patroni.dcs import ClusterConfig -from patroni.postgresql import Postgresql +from patroni.postgresql.config import ConfigHandler from patroni.utils import deep_compare, parse_bool, parse_int, patch_config from requests.structures import CaseInsensitiveDict @@ -59,7 +59,7 @@ class Config(object): 'postgresql': { 'bin_dir': '', 'use_slots': True, - 'parameters': CaseInsensitiveDict({p: v[0] for p, v in Postgresql.CMDLINE_OPTIONS.items()}) + 'parameters': CaseInsensitiveDict({p: v[0] for p, v in ConfigHandler.CMDLINE_OPTIONS.items()}) }, 'watchdog': { 'mode': 'automatic', @@ -178,11 +178,9 @@ class Config(object): @staticmethod def _process_postgresql_parameters(parameters, is_local=False): - ret = {} - for name, value in (parameters or {}).items(): - if name not in Postgresql.CMDLINE_OPTIONS or not is_local and Postgresql.CMDLINE_OPTIONS[name][1](value): - ret[name] = value - return ret + return {name: value for name, value in (parameters or {}).items() + if name not in ConfigHandler.CMDLINE_OPTIONS or + not is_local and ConfigHandler.CMDLINE_OPTIONS[name][1](value)} def _safe_copy_dynamic_configuration(self, dynamic_configuration): config = deepcopy(self.__DEFAULT_CONFIG) diff --git a/patroni/ctl.py b/patroni/ctl.py index bded4299..977ddba4 100644 --- a/patroni/ctl.py +++ b/patroni/ctl.py @@ -29,6 +29,7 @@ from patroni.config import Config from patroni.dcs import get_dcs as _get_dcs from patroni.exceptions import PatroniException from patroni.postgresql import Postgresql +from patroni.postgresql.misc import postgres_version_to_int from patroni.utils import patch_config, polling_loop from patroni.version import __version__ from prettytable import PrettyTable @@ -525,7 +526,7 @@ def restart(obj, cluster_name, member_names, force, role, p_any, scheduled, vers if version: try: - Postgresql.postgres_version_to_int(version) + postgres_version_to_int(version) except PatroniException as e: raise PatroniCtlException(e.value) diff --git a/patroni/ha.py b/patroni/ha.py index 647bb753..b1b745fe 100644 --- a/patroni/ha.py +++ b/patroni/ha.py @@ -13,6 +13,8 @@ from multiprocessing.pool import ThreadPool from patroni.async_executor import AsyncExecutor, CriticalTask from patroni.exceptions import DCSError, PostgresConnectionException, PatroniException from patroni.postgresql import ACTION_ON_START, ACTION_ON_ROLE_CHANGE +from patroni.postgresql.misc import postgres_version_to_int +from patroni.postgresql.rewind import Rewind from patroni.utils import polling_loop, tzutc from patroni.dcs import RemoteMember from threading import RLock @@ -59,6 +61,7 @@ class Ha(object): def __init__(self, patroni): self.patroni = patroni self.state_handler = patroni.postgresql + self._rewind = Rewind(self.state_handler) self.dcs = patroni.dcs self.cluster = None self.old_cluster = None @@ -71,7 +74,7 @@ class Ha(object): self._post_bootstrap_task = None self._crash_recovery_executed = False self._start_timeout = None - self._async_executor = AsyncExecutor(self.state_handler, self.wakeup) + self._async_executor = AsyncExecutor(self.state_handler.cancellable, self.wakeup) self.watchdog = patroni.watchdog # Each member publishes various pieces of information to the DCS using touch_member. This lock protects @@ -182,7 +185,7 @@ class Ha(object): if data['role'] == 'master' and not self.is_leader(): data['role'] = 'promoted' if self.is_leader(): - data['checkpoint_after_promote'] = self.state_handler.checkpoint_after_promote() + data['checkpoint_after_promote'] = self._rewind.checkpoint_after_promote() tags = self.get_effective_tags() if tags: data['tags'] = tags @@ -213,7 +216,8 @@ class Ha(object): if self.is_standby_cluster() and not isinstance(clone_member, RemoteMember): clone_member = self.get_remote_member(clone_member) - if self.state_handler.clone(clone_member): + self._rewind.reset_state() + if self.state_handler.bootstrap.clone(clone_member): logger.info('bootstrapped %s', msg) cluster = self.dcs.get_cluster() node_to_follow = self._get_node_to_follow(cluster) @@ -244,7 +248,7 @@ class Ha(object): else: self._async_executor.schedule('bootstrap') self._async_executor.run_async( - self.state_handler.bootstrap, + self.state_handler.bootstrap.bootstrap, args=(self.patroni.config['bootstrap'],) ) return 'trying to bootstrap a new cluster' @@ -276,12 +280,12 @@ class Ha(object): def _handle_rewind_or_reinitialize(self): leader = self.get_remote_master() if self.is_standby_cluster() else self.cluster.leader - if not self.state_handler.rewind_or_reinitialize_needed_and_possible(leader): + if not self._rewind.rewind_or_reinitialize_needed_and_possible(leader): return None - if self.state_handler.can_rewind: + if self._rewind.can_rewind: self._async_executor.schedule('running pg_rewind from ' + leader.name) - self._async_executor.run_async(self.state_handler.rewind, (leader,)) + self._async_executor.run_async(self._rewind.execute, (leader,)) return True # remove_data_directory_on_diverged_timelines is set @@ -311,8 +315,9 @@ class Ha(object): data = self.state_handler.controldata() logger.info('pg_controldata:\n%s\n', '\n'.join(' {0}: {1}'.format(k, v) for k, v in data.items())) - if data.get('Database cluster state') in ('in production', 'shutting down', 'in crash recovery') and \ - not self._crash_recovery_executed and (self.cluster.is_unlocked() or self.state_handler.can_rewind): + if data.get('Database cluster state') in ('in production', 'shutting down', 'in crash recovery') \ + and not self._crash_recovery_executed and \ + (self.cluster.is_unlocked() or self._rewind.can_rewind): self._crash_recovery_executed = True self._async_executor.schedule('doing crash recovery in a single user mode') self._async_executor.run_async(self.state_handler.fix_cluster_state) @@ -322,8 +327,8 @@ class Ha(object): role = 'replica' if self.is_standby_cluster() or not self.has_lock(): - if not self.state_handler.rewind_executed: - self.state_handler.trigger_check_diverged_lsn() + if not self._rewind.executed: + self._rewind.trigger_check_diverged_lsn() if self._handle_rewind_or_reinitialize(): return self._async_executor.scheduled_action @@ -368,7 +373,7 @@ class Ha(object): node_to_follow = self._get_node_to_follow(self.cluster) if self.is_paused(): - if not (self.state_handler.need_rewind and self.state_handler.can_rewind_or_reinitialize_allowed)\ + if not (self._rewind.is_needed and self._rewind.can_rewind_or_reinitialize_allowed)\ or self.cluster.is_unlocked(): self.state_handler.set_role('master' if is_leader else 'replica') if is_leader: @@ -387,7 +392,7 @@ class Ha(object): # In this case it is safe to continue running without changing recovery.conf if self.is_standby_cluster() and role == 'replica' and not (node_to_follow and node_to_follow.conn_url): return 'continue following the old known standby leader' - elif not self.state_handler.check_recovery_conf(node_to_follow): + elif not self.state_handler.config.check_recovery_conf(node_to_follow): self._async_executor.schedule('changing primary_conninfo and restarting') self._async_executor.run_async(self.state_handler.follow, (node_to_follow, role)) elif role == 'standby_leader' and self.state_handler.role != role: @@ -427,7 +432,7 @@ class Ha(object): logger.warning("No standbys available!") logger.info("Assigning synchronous standby status to %s", picked) - self.state_handler.set_synchronous_standby(picked) + self.state_handler.config.set_synchronous_standby(picked) if picked and picked != '*' and not allow_promote: # Wait for PostgreSQL to enable synchronous mode and see if we can immediately set sync_standby @@ -448,7 +453,7 @@ class Ha(object): else: if self.cluster.sync.leader and self.dcs.delete_sync_state(index=self.cluster.sync.index): logger.info("Disabled synchronous replication") - self.state_handler.set_synchronous_standby(None) + self.state_handler.config.set_synchronous_standby(None) def is_sync_standby(self, cluster): return cluster.leader and cluster.sync.leader == cluster.leader.name \ @@ -540,12 +545,17 @@ class Ha(object): # Somebody else updated sync state, it may be due to us losing the lock. To be safe, postpone # promotion until next cycle. TODO: trigger immediate retry of run_cycle return 'Postponing promotion because synchronous replication state was updated by somebody else' - self.state_handler.set_synchronous_standby('*' if self.is_synchronous_mode_strict() else None) + self.state_handler.config.set_synchronous_standby('*' if self.is_synchronous_mode_strict() else None) if self.state_handler.role != 'master': self.set_leader_access_is_restricted(self.cluster.has_permanent_logical_slots(self.state_handler.name)) + + def on_success(): + self._rewind.reset_state() + logger.info("cleared rewind state after becoming the leader") + self._async_executor.schedule('promote') self._async_executor.run_async(self.state_handler.promote, - args=(self.dcs.loop_wait, self._leader_access_is_restricted)) + args=(self.dcs.loop_wait, on_success, self._leader_access_is_restricted)) return promote_message @staticmethod @@ -744,7 +754,7 @@ class Ha(object): 'immediate-nolock': dict(stop='immediate', checkpoint=False, release=False, offline=False, async_req=True), }[mode] - self.state_handler.trigger_check_diverged_lsn() + self._rewind.trigger_check_diverged_lsn() self.state_handler.stop(mode_control['stop'], checkpoint=mode_control['checkpoint'], on_safepoint=self.watchdog.disable if self.watchdog.is_running else None) self.state_handler.set_role('demoted') @@ -768,8 +778,8 @@ class Ha(object): self._async_executor.run_async(self.state_handler.follow, (node_to_follow,)) else: if self.is_synchronous_mode(): - self.state_handler.set_synchronous_standby(None) - if self.state_handler.rewind_or_reinitialize_needed_and_possible(leader): + self.state_handler.config.set_synchronous_standby(None) + if self._rewind.rewind_or_reinitialize_needed_and_possible(leader): return False # do not start postgres, but run pg_rewind on the next iteration self.state_handler.follow(node_to_follow) @@ -885,7 +895,7 @@ class Ha(object): # when we are doing manual failover there is no guaranty that new leader is ahead of any other node # node tagged as nofailover can be ahead of the new leader either, but it is always excluded from elections if bool(self.cluster.failover) or self.patroni.nofailover: - self.state_handler.trigger_check_diverged_lsn() + self._rewind.trigger_check_diverged_lsn() time.sleep(2) # Give a time to somebody to take the leader lock if self.patroni.nofailover: @@ -904,7 +914,7 @@ class Ha(object): return 'removed leader lock because postgres is not running as master' if self.state_handler.is_leader() and self._leader_access_is_restricted: - self.state_handler.sync_replication_slots(self.cluster) + self.state_handler.slots_handler.sync_replication_slots(self.cluster) self.state_handler.call_nowait(ACTION_ON_ROLE_CHANGE) self.set_leader_access_is_restricted(False) @@ -914,7 +924,7 @@ class Ha(object): return msg # check if the node is ready to be used by pg_rewind - self.state_handler.check_for_checkpoint_after_promote() + self._rewind.check_for_checkpoint_after_promote() if self.is_standby_cluster(): # in case of standby cluster we don't really need to @@ -980,8 +990,7 @@ class Ha(object): if role and role != self.state_handler.role: reason_to_cancel = "host role mismatch" - if (postgres_version and - self.state_handler.postgres_version_to_int(postgres_version) <= int(self.state_handler.server_version)): + if postgres_version and postgres_version_to_int(postgres_version) <= int(self.state_handler.server_version): reason_to_cancel = "postgres version mismatch" if pending_restart and not self.state_handler.pending_restart: @@ -1145,7 +1154,7 @@ class Ha(object): self.state_handler.set_role('master') self._async_executor.schedule('post_bootstrap') - self._async_executor.run_async(self.state_handler.post_bootstrap, + self._async_executor.run_async(self.state_handler.bootstrap.post_bootstrap, args=(self.patroni.config['bootstrap'], self._post_bootstrap_task)) return 'running post_bootstrap' @@ -1154,7 +1163,7 @@ class Ha(object): if not self.watchdog.activate(): logger.error('Cancelling bootstrap because watchdog activation failed') self.cancel_initialization() - self.state_handler.sync_replication_slots(self.cluster) + self.state_handler.slots_handler.sync_replication_slots(self.cluster) self.dcs.take_leader() self.set_is_leader(True) self.state_handler.call_nowait(ACTION_ON_START) @@ -1243,7 +1252,7 @@ class Ha(object): if self.state_handler.bootstrapping: return self.post_bootstrap() - if self.recovering and not self.state_handler.need_rewind: + if self.recovering and not self._rewind.is_needed: self.recovering = False # Check if we tried to recover and failed msg = self.post_recover() @@ -1294,8 +1303,8 @@ class Ha(object): # the demote code follows through to starting Postgres right away, however, in the rewind case # it returns from demote and reaches this point to start PostgreSQL again after rewind. In that # case it makes no sense to continue to recover() unless rewind has finished successfully. - elif self.state_handler.rewind_failed or not self.state_handler.rewind_executed and not \ - (self.state_handler.need_rewind and self.state_handler.can_rewind_or_reinitialize_allowed): + elif self._rewind.failed or not self._rewind.executed and not \ + (self._rewind.is_needed and self._rewind.can_rewind_or_reinitialize_allowed): return 'postgres is not running' # try to start dead postgres @@ -1312,10 +1321,10 @@ class Ha(object): # stops PostgreSQL, therefore, we only reload replication slots if no # asynchronous processes are running (should be always the case for the master) if not self._async_executor.busy and not self.state_handler.is_starting(): - self.state_handler.sync_replication_slots(self.cluster) + self.state_handler.slots_handler.sync_replication_slots(self.cluster) if not self.state_handler.cb_called: if not self.state_handler.is_leader(): - self.state_handler.trigger_check_diverged_lsn() + self._rewind.trigger_check_diverged_lsn() self.state_handler.call_nowait(ACTION_ON_START) except DCSError: dcs_failed = True diff --git a/patroni/postgresql.py b/patroni/postgresql.py deleted file mode 100644 index 728e3dc2..00000000 --- a/patroni/postgresql.py +++ /dev/null @@ -1,2075 +0,0 @@ -import logging -import os -import psycopg2 -import re -import shlex -import shutil -import socket -import stat -import subprocess -import tempfile -import time - -from collections import defaultdict -from contextlib import contextmanager -from patroni.callback_executor import CallbackExecutor -from patroni.exceptions import PostgresConnectionException, PostgresException -from patroni.utils import compare_values, parse_bool, parse_int, Retry, RetryFailedError, polling_loop, split_host_port -from patroni.postmaster import PostmasterProcess -from patroni.dcs import slot_name_from_member_name, RemoteMember, Leader -from requests.structures import CaseInsensitiveDict -from six import string_types -from six.moves.urllib.parse import quote_plus -from threading import current_thread, Lock - - -logger = logging.getLogger(__name__) - -ACTION_ON_START = "on_start" -ACTION_ON_STOP = "on_stop" -ACTION_ON_RESTART = "on_restart" -ACTION_ON_RELOAD = "on_reload" -ACTION_ON_ROLE_CHANGE = "on_role_change" -ACTION_NOOP = "noop" - -STATE_RUNNING = 'running' -STATE_REJECT = 'rejecting connections' -STATE_NO_RESPONSE = 'not responding' -STATE_UNKNOWN = 'unknown' - -STOP_POLLING_INTERVAL = 1 -REWIND_STATUS = type('Enum', (), {'INITIAL': 0, 'CHECKPOINT': 1, 'CHECK': 2, 'NEED': 3, - 'NOT_NEED': 4, 'SUCCESS': 5, 'FAILED': 6}) -sync_standby_name_re = re.compile(r'^[A-Za-z_][A-Za-z_0-9\$]*$') - -cluster_info_query = ("SELECT CASE WHEN pg_catalog.pg_is_in_recovery() THEN 0 " - "ELSE ('x' || pg_catalog.substr(pg_catalog.pg_{0}file_name(" - "pg_catalog.pg_current_{0}_{1}()), 1, 8))::bit(32)::int END, " - "CASE WHEN pg_catalog.pg_is_in_recovery() THEN GREATEST(" - " pg_catalog.pg_{0}_{1}_diff(COALESCE(" - "pg_catalog.pg_last_{0}_receive_{1}(), '0/0'), '0/0')::bigint," - " pg_catalog.pg_{0}_{1}_diff(pg_catalog.pg_last_{0}_replay_{1}(), '0/0')::bigint)" - "ELSE pg_catalog.pg_{0}_{1}_diff(pg_catalog.pg_current_{0}_{1}(), '0/0')::bigint END") - - -def quote_ident(value): - """Very simplified version of quote_ident""" - return value if sync_standby_name_re.match(value) else '"' + value + '"' - - -@contextmanager -def null_context(): - yield - - -class Postgresql(object): - - # List of parameters which must be always passed to postmaster as command line options - # to make it not possible to change them with 'ALTER SYSTEM'. - # Some of these parameters have sane default value assigned and Patroni doesn't allow - # to decrease this value. E.g. 'wal_level' can't be lower then 'hot_standby' and so on. - # These parameters could be changed only globally, i.e. via DCS. - # P.S. 'listen_addresses' and 'port' are added here just for convenience, to mark them - # as a parameters which should always be passed through command line. - # - # Format: - # key - parameter name - # value - tuple(default_value, check_function, min_version) - # default_value -- some sane default value - # check_function -- if the new value is not correct must return `!False` - # min_version -- major version of PostgreSQL when parameter was introduced - CMDLINE_OPTIONS = CaseInsensitiveDict({ - 'listen_addresses': (None, lambda _: False, 90100), - 'port': (None, lambda _: False, 90100), - 'cluster_name': (None, lambda _: False, 90500), - 'wal_level': ('hot_standby', lambda v: v.lower() in ('hot_standby', 'replica', 'logical'), 90100), - 'hot_standby': ('on', lambda _: False, 90100), - 'max_connections': (100, lambda v: int(v) >= 100, 90100), - 'max_wal_senders': (10, lambda v: int(v) >= 10, 90100), - 'wal_keep_segments': (8, lambda v: int(v) >= 8, 90100), - 'max_prepared_transactions': (0, lambda v: int(v) >= 0, 90100), - 'max_locks_per_transaction': (64, lambda v: int(v) >= 64, 90100), - 'track_commit_timestamp': ('off', lambda v: parse_bool(v) is not None, 90500), - 'max_replication_slots': (10, lambda v: int(v) >= 10, 90400), - 'max_worker_processes': (8, lambda v: int(v) >= 8, 90400), - 'wal_log_hints': ('on', lambda _: False, 90400) - }) - - _CONFIG_WARNING_HEADER = '# Do not edit this file manually!\n# It will be overwritten by Patroni!\n' - - def __init__(self, config): - self.config = config - self.name = config['name'] - self.scope = config['scope'] - self._bin_dir = config.get('bin_dir') or '' - self._database = config.get('database', 'postgres') - self._data_dir = config['data_dir'] - self._config_dir = os.path.abspath(config.get('config_dir') or self._data_dir) - self._pending_restart = False - self.bootstrapping = False - self._running_custom_bootstrap = False - self.__thread_ident = current_thread().ident - - self._version_file = os.path.join(self._data_dir, 'PG_VERSION') - self._synchronous_standby_names = None - self._configure_server_parameters() - - self._connect_address = config.get('connect_address') - self._krbsrvname = config.get('krbsrvname') - - # for not so obvious connection attempts that may happen outside of pyscopg2 - if self._krbsrvname: - os.environ['PGKRBSRVNAME'] = self._krbsrvname - - self._superuser = config['authentication'].get('superuser', {}) - self.resolve_connection_addresses() - - self._rewind_state = REWIND_STATUS.INITIAL - self._use_slots = config.get('use_slots', True) - self._schedule_load_slots = self.use_slots - - self._pgpass = config.get('pgpass') or os.path.join(os.path.expanduser('~'), 'pgpass') - self._callback_executor = CallbackExecutor() - self.__cb_called = False - self.__cb_pending = None - config_base_name = config.get('config_base_name', 'postgresql') - self._postgresql_conf = os.path.join(self._config_dir, config_base_name + '.conf') - self._postgresql_base_conf_name = config_base_name + '.base.conf' - self._postgresql_base_conf = os.path.join(self._config_dir, self._postgresql_base_conf_name) - self._pg_hba_conf = os.path.join(self._config_dir, 'pg_hba.conf') - self._pg_ident_conf = os.path.join(self._config_dir, 'pg_ident.conf') - self._recovery_conf = os.path.join(self._data_dir, 'recovery.conf') - self._trigger_file = config.get('recovery_conf', {}).get('trigger_file') or 'promote' - self._trigger_file = os.path.abspath(os.path.join(self._data_dir, self._trigger_file)) - - self._is_cancelled = False - self._cancellable = None - self._cancellable_lock = Lock() - - self._connection_lock = Lock() - self._connection = None - self._cursor_holder = None - self._sysid = None - self._replication_slots = {} # already existing replication slots - self.retry = Retry(max_tries=-1, deadline=config['retry_timeout']/2.0, max_delay=1, - retry_exceptions=PostgresConnectionException) - - # Retry 'pg_is_in_recovery()' only once - self._is_leader_retry = Retry(max_tries=1, deadline=config['retry_timeout']/2.0, max_delay=1, - retry_exceptions=PostgresConnectionException) - - self._state_lock = Lock() - self.set_state('stopped') - self._role_lock = Lock() - self.set_role(self.get_postgres_role_from_data_directory()) - self._state_entry_timestamp = None - - self._cluster_info_state = {} - self._cached_replica_timeline = None - - # Last known running process - self._postmaster_proc = None - - if self.is_running(): - self.set_state('running') - self.set_role('master' if self.is_leader() else 'replica') - self._write_postgresql_conf() # we are "joining" already running postgres - if self._replace_pg_hba() or self._replace_pg_ident(): - self.reload() - elif self.role == 'master': - self.set_role('demoted') - - @property - def _create_replica_methods(self): - return (self.config.get('create_replica_methods', []) or - self.config.get('create_replica_method', [])) - - @property - def _configuration_to_save(self): - configuration = [os.path.basename(self._postgresql_conf)] - if 'custom_conf' not in self.config: - configuration.append(os.path.basename(self._postgresql_base_conf)) - if not self._server_parameters.get('hba_file'): - configuration.append('pg_hba.conf') - if not self._server_parameters.get('ident_file'): - configuration.append('pg_ident.conf') - return configuration - - @property - def use_slots(self): - return self._use_slots and self._major_version >= 90400 - - @property - def _replication(self): - return self.config['authentication']['replication'] - - @property - def _rewind_credentials(self): - return self.config['authentication'].get('rewind', self._superuser) \ - if self._major_version >= 110000 else self._superuser - - @property - def callback(self): - return self.config.get('callbacks') or {} - - @staticmethod - def _wal_name(version): - return 'wal' if version >= 100000 else 'xlog' - - @property - def wal_name(self): - return self._wal_name(self._major_version) - - @property - def lsn_name(self): - return 'lsn' if self._major_version >= 100000 else 'location' - - def _version_file_exists(self): - return not self.data_directory_empty() and os.path.isfile(self._version_file) - - def get_major_version(self): - if self._version_file_exists(): - try: - with open(self._version_file) as f: - return self.postgres_major_version_to_int(f.read().strip()) - except Exception: - logger.exception('Failed to read PG_VERSION from %s', self._data_dir) - return 0 - - def get_server_parameters(self, config): - parameters = config['parameters'].copy() - listen_addresses, port = split_host_port(config['listen'], 5432) - parameters.update({'cluster_name': self.scope, 'listen_addresses': listen_addresses, 'port': str(port)}) - if config.get('synchronous_mode', False): - if self._synchronous_standby_names is None: - if config.get('synchronous_mode_strict', False): - parameters['synchronous_standby_names'] = '*' - else: - parameters.pop('synchronous_standby_names', None) - else: - parameters['synchronous_standby_names'] = self._synchronous_standby_names - if self._major_version >= 90600 and parameters['wal_level'] == 'hot_standby': - parameters['wal_level'] = 'replica' - ret = CaseInsensitiveDict({k: v for k, v in parameters.items() if not self._major_version or - self._major_version >= self.CMDLINE_OPTIONS.get(k, (0, 1, 90100))[2]}) - ret.update({k: os.path.join(self._config_dir, ret[k]) for k in ('hba_file', 'ident_file') if k in ret}) - return ret - - def resolve_connection_addresses(self): - port = self._server_parameters['port'] - tcp_local_address = self._get_tcp_local_address() - - local_address = {'port': port} - if self.config.get('use_unix_socket'): - unix_socket_directories = self._server_parameters.get('unix_socket_directories') - if unix_socket_directories is not None: - # fallback to tcp if unix_socket_directories is set, but there are no sutable values - local_address['host'] = self._get_unix_local_address(unix_socket_directories) or tcp_local_address - - # if unix_socket_directories is not specified, but use_unix_socket is set to true - do our best - # to use default value, i.e. don't specify a host neither in connection url nor arguments - else: - local_address['host'] = tcp_local_address - - self._local_address = local_address - self._local_replication_address = {'host': tcp_local_address, 'port': port} - - self.connection_string = 'postgres://{0}/{1}'.format( - self._connect_address or tcp_local_address + ':' + port, self._database) - - def _pgcommand(self, cmd): - """Returns path to the specified PostgreSQL command""" - return os.path.join(self._bin_dir, cmd) - - def pg_ctl(self, cmd, *args, **kwargs): - """Builds and executes pg_ctl command - - :returns: `!True` when return_code == 0, otherwise `!False`""" - - pg_ctl = [self._pgcommand('pg_ctl'), cmd] - return subprocess.call(pg_ctl + ['-D', self._data_dir] + list(args), **kwargs) == 0 - - def pg_isready(self): - """Runs pg_isready to see if PostgreSQL is accepting connections. - - :returns: 'ok' if PostgreSQL is up, 'reject' if starting up, 'no_resopnse' if not up.""" - - cmd = [self._pgcommand('pg_isready'), '-p', self._local_address['port'], '-d', self._database] - - # Host is not set if we are connecting via default unix socket - if 'host' in self._local_address: - cmd.extend(['-h', self._local_address['host']]) - - # We only need the username because pg_isready does not try to authenticate - if 'username' in self._superuser: - cmd.extend(['-U', self._superuser['username']]) - - ret = subprocess.call(cmd) - return_codes = {0: STATE_RUNNING, - 1: STATE_REJECT, - 2: STATE_NO_RESPONSE, - 3: STATE_UNKNOWN} - return return_codes.get(ret, STATE_UNKNOWN) - - def reload_config(self, config): - self._superuser = config['authentication'].get('superuser', {}) - server_parameters = self.get_server_parameters(config) - - conf_changed = hba_changed = ident_changed = local_connection_address_changed = pending_restart = False - if self.state == 'running': - changes = CaseInsensitiveDict({p: v for p, v in server_parameters.items() if '.' not in p}) - changes.update({p: None for p in self._server_parameters.keys() if not ('.' in p or p in changes)}) - if changes: - if 'wal_segment_size' not in changes: - changes['wal_segment_size'] = '16384kB' - # XXX: query can raise an exception - for r in self.query("""SELECT name, setting, unit, vartype, context - FROM pg_catalog.pg_settings - WHERE pg_catalog.lower(name) IN (""" + ', '.join(['%s'] * len(changes)) + """) - ORDER BY 1 DESC""", *(k.lower() for k in changes.keys())): - if r[4] == 'internal': - if r[0] == 'wal_segment_size': - server_parameters.pop(r[0], None) - wal_segment_size = parse_int(r[2], 'kB') - if wal_segment_size is not None: - changes['wal_segment_size'] = '{0}kB'.format(int(r[1]) * wal_segment_size) - elif r[0] in changes: - unit = changes['wal_segment_size'] if r[0] in ('min_wal_size', 'max_wal_size') else r[2] - new_value = changes.pop(r[0]) - if new_value is None or not compare_values(r[3], unit, r[1], new_value): - if r[4] == 'postmaster': - pending_restart = True - logger.info('Changed %s from %s to %s (restart required)', r[0], r[1], new_value) - if config.get('use_unix_socket') and r[0] == 'unix_socket_directories'\ - or r[0] in ('listen_addresses', 'port'): - local_connection_address_changed = True - else: - logger.info('Changed %s from %s to %s', r[0], r[1], new_value) - conf_changed = True - for param in changes: - if param in server_parameters: - logger.warning('Removing invalid parameter `%s` from postgresql.parameters', param) - server_parameters.pop(param) - - # Check that user-defined-paramters have changed (parameters with period in name) - if not conf_changed: - for p, v in server_parameters.items(): - if '.' in p and (p not in self._server_parameters or str(v) != str(self._server_parameters[p])): - logger.info('Changed %s from %s to %s', p, self._server_parameters.get(p), v) - conf_changed = True - break - if not conf_changed: - for p, v in self._server_parameters.items(): - if '.' in p and (p not in server_parameters or str(v) != str(server_parameters[p])): - logger.info('Changed %s from %s to %s', p, v, server_parameters.get(p)) - conf_changed = True - break - - if not server_parameters.get('hba_file') and config.get('pg_hba'): - hba_changed = self.config.get('pg_hba', []) != config['pg_hba'] - - if not server_parameters.get('ident_file') and config.get('pg_ident'): - ident_changed = self.config.get('pg_ident', []) != config['pg_ident'] - - self.config = config - self._pending_restart = pending_restart - self._server_parameters = server_parameters - self._connect_address = config.get('connect_address') - self._krbsrvname = config.get('krbsrvname') - - # for not so obvious connection attempts that may happen outside of pyscopg2 - if self._krbsrvname: - os.environ['PGKRBSRVNAME'] = self._krbsrvname - - if not local_connection_address_changed: - self.resolve_connection_addresses() - - if conf_changed: - self._write_postgresql_conf() - - if hba_changed: - self._replace_pg_hba() - - if ident_changed: - self._replace_pg_ident() - - if conf_changed or hba_changed or ident_changed: - logger.info('PostgreSQL configuration items changed, reloading configuration.') - self.reload() - elif not pending_restart: - logger.info('No PostgreSQL configuration items changed, nothing to reload.') - - self._is_leader_retry.deadline = self.retry.deadline = config['retry_timeout']/2.0 - - @property - def pending_restart(self): - return self._pending_restart - - @staticmethod - def configuration_allows_rewind(data): - return data.get('wal_log_hints setting', 'off') == 'on' \ - or data.get('Data page checksum version', '0') != '0' - - @property - def can_rewind(self): - """ check if pg_rewind executable is there and that pg_controldata indicates - we have either wal_log_hints or checksums turned on - """ - # low-hanging fruit: check if pg_rewind configuration is there - if not self.config.get('use_pg_rewind'): - return False - - cmd = [self._pgcommand('pg_rewind'), '--help'] - try: - ret = subprocess.call(cmd, stdout=open(os.devnull, 'w'), stderr=subprocess.STDOUT) - if ret != 0: # pg_rewind is not there, close up the shop and go home - return False - except OSError: - return False - return self.configuration_allows_rewind(self.controldata()) - - @property - def can_rewind_or_reinitialize_allowed(self): - return self.config.get('remove_data_directory_on_diverged_timelines') or self.can_rewind - - @property - def sysid(self): - if not self._sysid and not self.bootstrapping: - data = self.controldata() - self._sysid = data.get('Database system identifier', "") - return self._sysid - - @staticmethod - def _get_unix_local_address(unix_socket_directories): - for d in unix_socket_directories.split(','): - d = d.strip() - if d.startswith('/'): # Only absolute path can be used to connect via unix-socket - return d - return '' - - def _get_tcp_local_address(self): - listen_addresses = self._server_parameters['listen_addresses'].split(',') - - for la in listen_addresses: - if la.strip().lower() in ('*', '0.0.0.0', '127.0.0.1', 'localhost'): # we are listening on '*' or localhost - return 'localhost' # connection via localhost is preferred - return listen_addresses[0].strip() # can't use localhost, take first address from listen_addresses - - def get_postgres_role_from_data_directory(self): - if self.data_directory_empty(): - return 'uninitialized' - elif os.path.exists(self._recovery_conf): - return 'replica' - else: - return 'master' - - @property - def _local_connect_kwargs(self): - ret = self._local_address.copy() - ret.update({'database': self._database, - 'fallback_application_name': 'Patroni', - 'connect_timeout': 3, - 'options': '-c statement_timeout=2000'}) - if 'username' in self._superuser: - ret['user'] = self._superuser['username'] - if 'password' in self._superuser: - ret['password'] = self._superuser['password'] - return ret - - def connection(self): - with self._connection_lock: - if not self._connection or self._connection.closed != 0: - self._connection = psycopg2.connect(**self._local_connect_kwargs) - self._connection.autocommit = True - self.server_version = self._connection.server_version - return self._connection - - def _cursor(self): - if not self._cursor_holder or self._cursor_holder.closed or self._cursor_holder.connection.closed != 0: - logger.info("establishing a new patroni connection to the postgres cluster") - self._cursor_holder = self.connection().cursor() - return self._cursor_holder - - def close_connection(self): - if self._connection and self._connection.closed == 0: - self._connection.close() - logger.info("closed patroni connection to the postgresql cluster") - self._cursor_holder = self._connection = None - - def _query(self, sql, *params): - """We are always using the same cursor, therefore this method is not thread-safe!!! - You can call it from different threads only if you are holding explicit `AsyncExecutor` lock, - because the main thread is always holding this lock when running HA cycle.""" - cursor = None - try: - cursor = self._cursor() - cursor.execute(sql, params) - return cursor - except psycopg2.Error as e: - if cursor and cursor.connection.closed == 0: - # When connected via unix socket, psycopg2 can't recoginze 'connection lost' - # and leaves `_cursor_holder.connection.closed == 0`, but psycopg2.OperationalError - # is still raised (what is correct). It doesn't make sense to continiue with existing - # connection and we will close it, to avoid its reuse by the `_cursor` method. - if isinstance(e, psycopg2.OperationalError): - self.close_connection() - else: - raise e - if self.state == 'restarting': - raise RetryFailedError('cluster is being restarted') - raise PostgresConnectionException('connection problems') - - def query(self, sql, *params): - try: - return self.retry(self._query, sql, *params) - except RetryFailedError as e: - raise PostgresConnectionException(str(e)) - - def data_directory_empty(self): - return not os.path.exists(self._data_dir) or os.listdir(self._data_dir) == [] - - @staticmethod - def process_user_options(tool, options, not_allowed_options, error_handler): - user_options = [] - - def option_is_allowed(name): - ret = name not in not_allowed_options - if not ret: - error_handler('{0} option for {1} is not allowed'.format(name, tool)) - return ret - - if isinstance(options, dict): - for k, v in options.items(): - if k and v: - user_options.append('--{0}={1}'.format(k, v)) - elif isinstance(options, list): - for opt in options: - if isinstance(opt, string_types) and option_is_allowed(opt): - user_options.append('--{0}'.format(opt)) - elif isinstance(opt, dict): - keys = list(opt.keys()) - if len(keys) != 1 or not isinstance(opt[keys[0]], string_types) or not option_is_allowed(keys[0]): - error_handler('Error when parsing {0} key-value option {1}: only one key-value is allowed' - ' and value should be a string'.format(tool, opt[keys[0]])) - user_options.append('--{0}={1}'.format(keys[0], opt[keys[0]])) - else: - error_handler('Error when parsing {0} option {1}: value should be string value' - ' or a single key-value pair'.format(tool, opt)) - else: - error_handler('{0} options must be list ot dict'.format(tool)) - return user_options - - def _initdb(self, config): - self.set_state('initalizing new cluster') - not_allowed_options = ('pgdata', 'nosync', 'pwfile', 'sync-only', 'version') - - def error_handler(e): - raise Exception(e) - - options = self.process_user_options('initdb', config.get('initdb') or [], not_allowed_options, error_handler) - pwfile = None - - if self._superuser: - if 'username' in self._superuser: - options.append('--username={0}'.format(self._superuser['username'])) - if 'password' in self._superuser: - (fd, pwfile) = tempfile.mkstemp() - os.write(fd, self._superuser['password'].encode('utf-8')) - os.close(fd) - options.append('--pwfile={0}'.format(pwfile)) - options = ['-o', ' '.join(options)] if options else [] - - ret = self.pg_ctl('initdb', *options) - if pwfile: - os.remove(pwfile) - if not ret: - self.set_state('initdb failed') - return ret - - def _custom_bootstrap(self, config): - self.set_state('running custom bootstrap script') - params = ['--scope=' + self.scope, '--datadir=' + self._data_dir] - try: - logger.info('Running custom bootstrap script: %s', config['command']) - if self.cancellable_subprocess_call(shlex.split(config['command']) + params) != 0: - self.set_state('custom bootstrap failed') - return False - except Exception: - logger.exception('Exception during custom bootstrap') - return False - self._post_restore() - - if 'recovery_conf' in config: - self.write_recovery_conf(config['recovery_conf']) - elif (os.path.isfile(self._recovery_conf) or os.path.islink(self._recovery_conf)) and \ - not config.get('keep_existing_recovery_conf'): - os.unlink(self._recovery_conf) - return True - - def run_bootstrap_post_init(self, config): - """ - runs a script after initdb or custom bootstrap script is called and waits until completion. - """ - cmd = config.get('post_bootstrap') or config.get('post_init') - if cmd: - r = self._local_connect_kwargs - - if 'host' in r: - # '/tmp' => '%2Ftmp' for unix socket path - host = quote_plus(r['host']) if r['host'].startswith('/') else r['host'] - else: - host = '' - - # https://www.postgresql.org/docs/current/static/libpq-pgpass.html - # A host name of localhost matches both TCP (host name localhost) and Unix domain socket - # (pghost empty or the default socket directory) connections coming from the local machine. - r['host'] = 'localhost' # set it to localhost to write into pgpass - - if 'user' in r: - user = r['user'] + '@' - else: - user = '' - if 'password' in r: - import getpass - r.setdefault('user', os.environ.get('PGUSER', getpass.getuser())) - - connstring = 'postgres://{0}{1}:{2}/{3}'.format(user, host, r['port'], r['database']) - env = self.write_pgpass(r) if 'password' in r else None - - try: - ret = self.cancellable_subprocess_call(shlex.split(cmd) + [connstring], env=env) - except OSError: - logger.error('post_init script %s failed', cmd) - return False - if ret != 0: - logger.error('post_init script %s returned non-zero code %d', cmd, ret) - return False - return True - - def delete_trigger_file(self): - if os.path.exists(self._trigger_file): - os.unlink(self._trigger_file) - - def write_pgpass(self, record): - if 'user' not in record or 'password' not in record: - return os.environ.copy() - - with open(self._pgpass, 'w') as f: - if os.name != 'nt': - os.fchmod(f.fileno(), 0o600) - f.write('{host}:{port}:*:{user}:{password}\n'.format(**record)) - - env = os.environ.copy() - env['PGPASSFILE'] = self._pgpass - return env - - def replica_method_can_work_without_replication_connection(self, method): - return method != 'basebackup' and self.config and self.config.get(method, {}).get('no_master') - - def can_create_replica_without_replication_connection(self, replica_methods=None): - """ go through the replication methods to see if there are ones - that does not require a working replication connection. - """ - if replica_methods is None: - replica_methods = self._create_replica_methods - return any(self.replica_method_can_work_without_replication_connection(method) for method in replica_methods) - - def create_replica(self, clone_member): - """ - create the replica according to the replica_method - defined by the user. this is a list, so we need to - loop through all methods the user supplies - """ - - self.set_state('creating replica') - self._sysid = None - - is_remote_master = isinstance(clone_member, RemoteMember) - - # get list of replica methods either from clone member or from - # the config. If there is no configuration key, or no value is - # specified, use basebackup - replica_methods = (clone_member.create_replica_methods if is_remote_master - else self._create_replica_methods) or ['basebackup'] - - if clone_member and clone_member.conn_url: - r = clone_member.conn_kwargs(self._replication) - connstring = 'postgres://{user}@{host}:{port}/{database}'.format(**r) - # add the credentials to connect to the replica origin to pgpass. - env = self.write_pgpass(r) - else: - connstring = '' - env = os.environ.copy() - # if we don't have any source, leave only replica methods that work without it - replica_methods = \ - [r for r in replica_methods if self.replica_method_can_work_without_replication_connection(r)] - - # go through them in priority order - ret = 1 - for replica_method in replica_methods: - with self._cancellable_lock: - if self._is_cancelled: - break - # if the method is basebackup, then use the built-in - if replica_method == "basebackup": - ret = self.basebackup(connstring, env, self.config.get(replica_method, {})) - if ret == 0: - logger.info("replica has been created using basebackup") - # if basebackup succeeds, exit with success - break - else: - if not self.data_directory_empty(): - if self.config.get(replica_method, {}).get('keep_data', False): - logger.info('Leaving data directory uncleaned') - else: - self.remove_data_directory() - - cmd = replica_method - method_config = {} - # user-defined method; check for configuration - # not required, actually - if self.config.get(replica_method, {}): - method_config = self.config[replica_method].copy() - # look to see if the user has supplied a full command path - # if not, use the method name as the command - cmd = method_config.pop('command', cmd) - - # add the default parameters - if not method_config.get('no_params', False): - method_config.update({"scope": self.scope, - "role": "replica", - "datadir": self._data_dir, - "connstring": connstring}) - else: - for param in ('no_params', 'no_master', 'keep_data'): - method_config.pop(param, None) - params = ["--{0}={1}".format(arg, val) for arg, val in method_config.items()] - try: - # call script with the full set of parameters - ret = self.cancellable_subprocess_call(shlex.split(cmd) + params, env=env) - # if we succeeded, stop - if ret == 0: - logger.info('replica has been created using %s', replica_method) - break - else: - logger.error('Error creating replica using method %s: %s exited with code=%s', - replica_method, cmd, ret) - except Exception: - logger.exception('Error creating replica using method %s', replica_method) - ret = 1 - - self.set_state('stopped') - return ret - - def reset_cluster_info_state(self): - self._cluster_info_state = {} - - def _cluster_info_state_get(self, name): - if not self._cluster_info_state: - stmt = cluster_info_query.format(self.wal_name, self.lsn_name) - try: - result = self._is_leader_retry(self._query, stmt).fetchone() - self._cluster_info_state = dict(zip(['timeline', 'wal_position'], result)) - except RetryFailedError as e: # SELECT failed two times - self._cluster_info_state = {'error': str(e)} - if not self.is_starting() and self.pg_isready() == STATE_REJECT: - self.set_state('starting') - - if 'error' in self._cluster_info_state: - raise PostgresConnectionException(self._cluster_info_state['error']) - - return self._cluster_info_state.get(name) - - def is_leader(self): - return bool(self._cluster_info_state_get('timeline')) - - def is_running(self): - """Returns PostmasterProcess if one is running on the data directory or None. If most recently seen process - is running updates the cached process based on pid file.""" - if self._postmaster_proc: - if self._postmaster_proc.is_running(): - return self._postmaster_proc - self._postmaster_proc = None - - # we noticed that postgres was restarted, force syncing of replication - self._schedule_load_slots = self.use_slots - - self._postmaster_proc = PostmasterProcess.from_pidfile(self._data_dir) - return self._postmaster_proc - - @property - def cb_called(self): - return self.__cb_called - - def call_nowait(self, cb_name): - """ pick a callback command and call it without waiting for it to finish """ - if self.bootstrapping: - return - if cb_name in (ACTION_ON_START, ACTION_ON_STOP, ACTION_ON_RESTART, ACTION_ON_ROLE_CHANGE): - self.__cb_called = True - - if self.callback and cb_name in self.callback: - cmd = self.callback[cb_name] - try: - cmd = shlex.split(self.callback[cb_name]) + [cb_name, self.role, self.scope] - self._callback_executor.call(cmd) - except Exception: - logger.exception('callback %s %s %s %s failed', cmd, cb_name, self.role, self.scope) - - @property - def role(self): - with self._role_lock: - return self._role - - def set_role(self, value): - with self._role_lock: - self._role = value - - @property - def state(self): - with self._state_lock: - return self._state - - def set_state(self, value): - with self._state_lock: - self._state = value - self._state_entry_timestamp = time.time() - - def time_in_state(self): - return time.time() - self._state_entry_timestamp - - def is_starting(self): - return self.state == 'starting' - - def wait_for_port_open(self, postmaster, timeout): - """Waits until PostgreSQL opens ports.""" - for _ in polling_loop(timeout): - with self._cancellable_lock: - if self._is_cancelled: - return False - - if not postmaster.is_running(): - logger.error('postmaster is not running') - self.set_state('start failed') - return False - - isready = self.pg_isready() - if isready != STATE_NO_RESPONSE: - if isready not in [STATE_REJECT, STATE_RUNNING]: - logger.warning("Can't determine PostgreSQL startup status, assuming running") - return True - - logger.warning("Timed out waiting for PostgreSQL to start") - return False - - def _build_effective_configuration(self): - """It might happen that the current value of one (or more) below parameters stored in - the controldata is higher than the value stored in the global cluster configuration. - - Example: max_connections in global configuration is 100, but in controldata - `Current max_connections setting: 200`. If we try to start postgres with - max_connections=100, it will immediately exit. - As a workaround we will start it with the values from controldata and set `pending_restart` - to true as an indicator that current values of parameters are not matching expectations.""" - - OPTIONS_MAPPING = { - 'max_connections': 'max_connections setting', - 'max_prepared_transactions': 'max_prepared_xacts setting', - 'max_locks_per_transaction': 'max_locks_per_xact setting' - } - - if self._major_version >= 90400: - OPTIONS_MAPPING['max_worker_processes'] = 'max_worker_processes setting' - - data = self.controldata() - effective_configuration = self._server_parameters.copy() - - for name, cname in OPTIONS_MAPPING.items(): - value = parse_int(effective_configuration[name]) - cvalue = parse_int(data[cname]) - if cvalue > value: - effective_configuration[name] = cvalue - self._pending_restart = True - return effective_configuration - - def start(self, timeout=None, task=None, block_callbacks=False, role=None): - """Start PostgreSQL - - Waits for postmaster to open ports or terminate so pg_isready can be used to check startup completion - or failure. - - :returns: True if start was initiated and postmaster ports are open, False if start failed""" - # make sure we close all connections established against - # the former node, otherwise, we might get a stalled one - # after kill -9, which would report incorrect data to - # patroni. - self.close_connection() - - if self.is_running(): - logger.error('Cannot start PostgreSQL because one is already running.') - self.set_state('starting') - return True - - if not block_callbacks: - self.__cb_pending = ACTION_ON_START - - self.set_role(role or self.get_postgres_role_from_data_directory()) - - self.set_state('starting') - self._pending_restart = False - - configuration = self._server_parameters if self.role == 'master' else self._build_effective_configuration() - self._write_postgresql_conf(configuration) - self.resolve_connection_addresses() - self._replace_pg_hba() - self._replace_pg_ident() - - options = ['--{0}={1}'.format(p, configuration[p]) for p in self.CMDLINE_OPTIONS - if p in configuration and p != 'wal_keep_segments'] - - with self._cancellable_lock: - if self._is_cancelled: - return False - - with task or null_context(): - if task and task.is_cancelled: - logger.info("PostgreSQL start cancelled.") - return False - - self._postmaster_proc = PostmasterProcess.start(self._pgcommand('postgres'), - self._data_dir, - self._postgresql_conf, - options) - - if task: - task.complete(self._postmaster_proc) - - start_timeout = timeout - if not start_timeout: - try: - start_timeout = float(self.config.get('pg_ctl_timeout', 60)) - except ValueError: - start_timeout = 60 - - # We want postmaster to open ports before we continue - if not self._postmaster_proc or not self.wait_for_port_open(self._postmaster_proc, start_timeout): - return False - - ret = self.wait_for_startup(start_timeout) - if ret is not None: - return ret - elif timeout is not None: - return False - else: - return None - - def checkpoint(self, connect_kwargs=None): - check_not_is_in_recovery = connect_kwargs is not None - connect_kwargs = connect_kwargs or self._local_connect_kwargs - for p in ['connect_timeout', 'options']: - connect_kwargs.pop(p, None) - try: - with self._get_connection_cursor(**connect_kwargs) as cur: - cur.execute("SET statement_timeout = 0") - if check_not_is_in_recovery: - cur.execute('SELECT pg_catalog.pg_is_in_recovery()') - if cur.fetchone()[0]: - return 'is_in_recovery=true' - return cur.execute('CHECKPOINT') - except psycopg2.Error: - logger.exception('Exception during CHECKPOINT') - return 'not accessible or not healty' - - def stop(self, mode='fast', block_callbacks=False, checkpoint=None, on_safepoint=None): - """Stop PostgreSQL - - Supports a callback when a safepoint is reached. A safepoint is when no user backend can return a successful - commit to users. Currently this means we wait for user backends to close. But in the future alternate mechanisms - could be added. - - :param on_safepoint: This callback is called when no user backends are running. - """ - if checkpoint is None: - checkpoint = False if mode == 'immediate' else True - - success, pg_signaled = self._do_stop(mode, block_callbacks, checkpoint, on_safepoint) - if success: - # block_callbacks is used during restart to avoid - # running start/stop callbacks in addition to restart ones - if not block_callbacks: - self.set_state('stopped') - if pg_signaled: - self.call_nowait(ACTION_ON_STOP) - else: - logger.warning('pg_ctl stop failed') - self.set_state('stop failed') - return success - - def _do_stop(self, mode, block_callbacks, checkpoint, on_safepoint): - postmaster = self.is_running() - if not postmaster: - if on_safepoint: - on_safepoint() - return True, False - - if checkpoint and not self.is_starting(): - self.checkpoint() - - if not block_callbacks: - self.set_state('stopping') - - # Send signal to postmaster to stop - success = postmaster.signal_stop(mode) - if success is not None: - if success and on_safepoint: - on_safepoint() - return success, True - - # We can skip safepoint detection if we don't have a callback - if on_safepoint: - # Wait for our connection to terminate so we can be sure that no new connections are being initiated - self._wait_for_connection_close(postmaster) - postmaster.wait_for_user_backends_to_close() - on_safepoint() - - postmaster.wait() - - return True, True - - @staticmethod - def terminate_starting_postmaster(postmaster): - """Terminates a postmaster that has not yet opened ports or possibly even written a pid file. Blocks - until the process goes away.""" - postmaster.signal_stop('immediate') - postmaster.wait() - - def _wait_for_connection_close(self, postmaster): - try: - with self.connection().cursor() as cur: - while postmaster.is_running(): # Need a timeout here? - cur.execute("SELECT 1") - time.sleep(STOP_POLLING_INTERVAL) - except psycopg2.Error: - pass - - def reload(self): - ret = self.pg_ctl('reload') - if ret: - self.call_nowait(ACTION_ON_RELOAD) - return ret - - def check_for_checkpoint_after_promote(self): - if self._rewind_state == REWIND_STATUS.INITIAL and self.is_leader(): - try: - timeline = int(self.controldata().get("Latest checkpoint's TimeLineID")) - if self.get_master_timeline() == timeline: - self._rewind_state = REWIND_STATUS.CHECKPOINT - except (TypeError, ValueError): - logger.exception('Failed to parse timeline from pg_controldata output') - - def checkpoint_after_promote(self): - return self._rewind_state == REWIND_STATUS.CHECKPOINT - - def check_for_startup(self): - """Checks PostgreSQL status and returns if PostgreSQL is in the middle of startup.""" - return self.is_starting() and not self.check_startup_state_changed() - - def check_startup_state_changed(self): - """Checks if PostgreSQL has completed starting up or failed or still starting. - - Should only be called when state == 'starting' - - :returns: True if state was changed from 'starting' - """ - ready = self.pg_isready() - - if ready == STATE_REJECT: - return False - elif ready == STATE_NO_RESPONSE: - self.set_state('start failed') - self._schedule_load_slots = False # TODO: can remove this? - if not self._running_custom_bootstrap: - self.save_configuration_files() # TODO: maybe remove this? - return True - else: - if ready != STATE_RUNNING: - # Bad configuration or unexpected OS error. No idea of PostgreSQL status. - # Let the main loop of run cycle clean up the mess. - logger.warning("%s status returned from pg_isready", - "Unknown" if ready == STATE_UNKNOWN else "Invalid") - self.set_state('running') - self._schedule_load_slots = self.use_slots - if not self._running_custom_bootstrap: - self.save_configuration_files() - # TODO: __cb_pending can be None here after PostgreSQL restarts on its own. Do we want to call the callback? - # Previously we didn't even notice. - action = self.__cb_pending or ACTION_ON_START - self.call_nowait(action) - self.__cb_pending = None - - return True - - def wait_for_startup(self, timeout=None): - """Waits for PostgreSQL startup to complete or fail. - - :returns: True if start was successful, False otherwise""" - if not self.is_starting(): - # Should not happen - logger.warning("wait_for_startup() called when not in starting state") - - while not self.check_startup_state_changed(): - with self._cancellable_lock: - if self._is_cancelled: - return None - if timeout and self.time_in_state() > timeout: - return None - time.sleep(1) - - return self.state == 'running' - - def restart(self, timeout=None, task=None, block_callbacks=False, role=None): - """Restarts PostgreSQL. - - When timeout parameter is set the call will block either until PostgreSQL has started, failed to start or - timeout arrives. - - :returns: True when restart was successful and timeout did not expire when waiting. - """ - self.set_state('restarting') - if not block_callbacks: - self.__cb_pending = ACTION_ON_RESTART - ret = self.stop(block_callbacks=True) and self.start(timeout, task, True, role) - if not ret and not self.is_starting(): - self.set_state('restart failed ({0})'.format(self.state)) - return ret - - def _write_postgresql_conf(self, configuration=None): - # rename the original configuration if it is necessary - if 'custom_conf' not in self.config and not os.path.exists(self._postgresql_base_conf): - os.rename(self._postgresql_conf, self._postgresql_base_conf) - - with open(self._postgresql_conf, 'w') as f: - f.write(self._CONFIG_WARNING_HEADER) - f.write("include '{0}'\n\n".format(self.config.get('custom_conf') or self._postgresql_base_conf_name)) - for name, value in sorted((configuration or self._server_parameters).items()): - if not self._running_custom_bootstrap or name != 'hba_file': - f.write("{0} = '{1}'\n".format(name, value)) - # when we are doing custom bootstrap we assume that we don't know superuser password - # and in order to be able to change it, we are opening trust access from a certain address - # therefore we need to make sure that hba_file is not overriden - # after changing superuser password we will "revert" all these "changes" - if self._running_custom_bootstrap or 'hba_file' not in self._server_parameters: - f.write("hba_file = '{0}'\n".format(self._pg_hba_conf.replace('\\', '\\\\'))) - if 'ident_file' not in self._server_parameters: - f.write("ident_file = '{0}'\n".format(self._pg_ident_conf.replace('\\', '\\\\'))) - - def is_healthy(self): - if not self.is_running(): - logger.warning('Postgresql is not running.') - return False - return True - - def append_pg_hba(self, config): - if not self._server_parameters.get('hba_file') and not self.config.get('pg_hba'): - with open(self._pg_hba_conf, 'a') as f: - f.write('\n{}\n'.format('\n'.join(config))) - return True - - def _replace_pg_hba(self): - """ - Replace pg_hba.conf content in the PGDATA if hba_file is not defined in the - `postgresql.parameters` and pg_hba is defined in `postgresql` configuration section. - - :returns: True if pg_hba.conf was rewritten. - """ - - # when we are doing custom bootstrap we assume that we don't know superuser password - # and in order to be able to change it, we are opening trust access from a certain address - if self._running_custom_bootstrap: - addresses = {'': 'local'} - if 'host' in self._local_address and not self._local_address['host'].startswith('/'): - for _, _, _, _, sa in socket.getaddrinfo(self._local_address['host'], self._local_address['port'], - 0, socket.SOCK_STREAM, socket.IPPROTO_TCP): - addresses[sa[0] + '/32'] = 'host' - - with open(self._pg_hba_conf, 'w') as f: - f.write(self._CONFIG_WARNING_HEADER) - for address, t in addresses.items(): - f.write(( - '{0}\treplication\t{1}\t{3}\ttrust\n' - '{0}\tall\t{2}\t{3}\ttrust\n' - ).format(t, self._replication['username'], self._superuser.get('username') or 'all', address)) - elif not self._server_parameters.get('hba_file') and self.config.get('pg_hba'): - with open(self._pg_hba_conf, 'w') as f: - f.write(self._CONFIG_WARNING_HEADER) - for line in self.config['pg_hba']: - f.write('{0}\n'.format(line)) - return True - - def _replace_pg_ident(self): - """ - Replace pg_ident.conf content in the PGDATA if ident_file is not defined in the - `postgresql.parameters` and pg_ident is defined in the `postgresql` section. - - :returns: True if pg_ident.conf was rewritten. - """ - - if not self._server_parameters.get('ident_file') and self.config.get('pg_ident'): - with open(self._pg_ident_conf, 'w') as f: - f.write(self._CONFIG_WARNING_HEADER) - for line in self.config['pg_ident']: - f.write('{0}\n'.format(line)) - return True - - def primary_conninfo(self, member): - if not (member and member.conn_url) or member.name == self.name: - return None - r = member.conn_kwargs(self._replication) - r.update(application_name=self.name, sslmode='prefer', sslcompression='1', krbsrvname=self._krbsrvname) - keywords = 'user password host port sslmode sslcompression application_name krbsrvname'.split() - return ' '.join('{0}={{{0}}}'.format(kw) for kw in keywords if r.get(kw)).format(**r) - - def check_recovery_conf(self, member): - # TODO: recovery.conf could be stale, would be nice to detect that. - primary_conninfo = self.primary_conninfo(member) - - if not os.path.isfile(self._recovery_conf): - return False - - with open(self._recovery_conf, 'r') as f: - for line in f: - if line.startswith('primary_conninfo'): - return primary_conninfo and (primary_conninfo in line) - return not primary_conninfo - - def write_recovery_conf(self, recovery_params): - with open(self._recovery_conf, 'w') as f: - os.chmod(self._recovery_conf, stat.S_IWRITE | stat.S_IREAD) - for name, value in recovery_params.items(): - f.write("{0} = '{1}'\n".format(name, value)) - - def pg_rewind(self, r): - # prepare pg_rewind connection - env = self.write_pgpass(r) - dsn_attrs = [ - ('user', r.get('user')), - ('host', r.get('host')), - ('port', r.get('port')), - ('dbname', r.get('database') or self._database), - ('sslmode', 'prefer'), - ('sslcompression', '1'), - ] - dsn = " ".join("{0}={1}".format(k, v) for k, v in dsn_attrs if v is not None) - logger.info('running pg_rewind from %s', dsn) - try: - return self.cancellable_subprocess_call([self._pgcommand('pg_rewind'), - '-D', self._data_dir, - '--source-server', dsn], env=env) == 0 - except OSError: - return False - - def controldata(self): - """ return the contents of pg_controldata, or non-True value if pg_controldata call failed """ - result = {} - # Don't try to call pg_controldata during backup restore - if self._version_file_exists() and self.state != 'creating replica': - try: - env = {'LANG': 'C', 'LC_ALL': 'C', 'PATH': os.getenv('PATH')} - if os.getenv('SYSTEMROOT') is not None: - env['SYSTEMROOT'] = os.getenv('SYSTEMROOT') - data = subprocess.check_output([self._pgcommand('pg_controldata'), self._data_dir], env=env) - if data: - data = data.decode('utf-8').splitlines() - # pg_controldata output depends on major verion. Some of parameters are prefixed by 'Current ' - result = {l.split(':')[0].replace('Current ', '', 1): l.split(':', 1)[1].strip() for l in data - if l and ':' in l} - except subprocess.CalledProcessError: - logger.exception("Error when calling pg_controldata") - return result - - @property - def need_rewind(self): - return self._rewind_state in (REWIND_STATUS.CHECK, REWIND_STATUS.NEED) - - @staticmethod - @contextmanager - def _get_connection_cursor(**kwargs): - with psycopg2.connect(**kwargs) as conn: - conn.autocommit = True - with conn.cursor() as cur: - yield cur - - @contextmanager - def _get_replication_connection_cursor(self, host='localhost', port=5432, database=None, **kwargs): - with self._get_connection_cursor(host=host, port=int(port), database=database or self._database, replication=1, - user=self._replication['username'], password=self._replication.get('password'), - connect_timeout=3, options='-c statement_timeout=2000') as cur: - yield cur - - def check_leader_is_not_in_recovery(self, **kwargs): - if not kwargs.get('database'): - kwargs['database'] = self._database - try: - with self._get_connection_cursor(connect_timeout=3, options='-c statement_timeout=2000', **kwargs) as cur: - cur.execute('SELECT pg_catalog.pg_is_in_recovery()') - if not cur.fetchone()[0]: - return True - logger.info('Leader is still in_recovery and therefore can\'t be used for rewind') - except Exception: - return logger.exception('Exception when working with leader') - - def _get_local_timeline_lsn_from_replication_connection(self): - timeline = lsn = None - try: - with self._get_replication_connection_cursor(**self._local_replication_address) as cur: - cur.execute('IDENTIFY_SYSTEM') - timeline, lsn = cur.fetchone()[1:3] - except Exception: - logger.exception('Can not fetch local timeline and lsn from replication connection') - return timeline, lsn - - def _get_local_timeline_lsn_from_controldata(self): - timeline = lsn = None - data = self.controldata() - try: - if data.get('Database cluster state') == 'shut down in recovery': - lsn = data.get('Minimum recovery ending location') - timeline = int(data.get("Min recovery ending loc's timeline")) - if lsn == '0/0' or timeline == 0: # it was a master when it crashed - data['Database cluster state'] = 'shut down' - if data.get('Database cluster state') == 'shut down': - lsn = data.get('Latest checkpoint location') - timeline = int(data.get("Latest checkpoint's TimeLineID")) - except (TypeError, ValueError): - logger.exception('Failed to get local timeline and lsn from pg_controldata output') - return timeline, lsn - - def _get_local_timeline_lsn(self): - if self.is_running(): # if postgres is running - get timeline and lsn from replication connection - timeline, lsn = self._get_local_timeline_lsn_from_replication_connection() - else: # otherwise analyze pg_controldata output - timeline, lsn = self._get_local_timeline_lsn_from_controldata() - logger.info('Local timeline=%s lsn=%s', timeline, lsn) - return timeline, lsn - - @staticmethod - def parse_lsn(lsn): - t = lsn.split('/') - return int(t[0], 16) * 0x100000000 + int(t[1], 16) - - @staticmethod - def parse_history(data): - for line in data.split('\n'): - values = line.strip().split('\t') - if len(values) == 3: - try: - values[0] = int(values[0]) - values[1] = Postgresql.parse_lsn(values[1]) - yield values - except (IndexError, ValueError): - logger.exception('Exception when parsing timeline history line "%s"', values) - - def _check_timeline_and_lsn(self, leader): - local_timeline, local_lsn = self._get_local_timeline_lsn() - if local_timeline is None or local_lsn is None: - return - - if isinstance(leader, Leader): - if leader.member.data.get('role') != 'master': - return - # standby cluster - elif not self.check_leader_is_not_in_recovery(**leader.conn_kwargs(self._replication)): - return - - history = need_rewind = None - try: - with self._get_replication_connection_cursor(**leader.conn_kwargs()) as cur: - cur.execute('IDENTIFY_SYSTEM') - master_timeline = cur.fetchone()[1] - logger.info('master_timeline=%s', master_timeline) - if local_timeline > master_timeline: # Not always supported by pg_rewind - need_rewind = True - elif master_timeline > 1: - cur.execute('TIMELINE_HISTORY %s', (master_timeline,)) - history = bytes(cur.fetchone()[1]).decode('utf-8') - logger.info('master: history=%s', history) - else: # local_timeline == master_timeline == 1 - need_rewind = False - except Exception: - return logger.exception('Exception when working with master via replication connection') - - if history is not None: - for parent_timeline, switchpoint, _ in self.parse_history(history): - if parent_timeline == local_timeline: - try: - need_rewind = self.parse_lsn(local_lsn) >= switchpoint - except (IndexError, ValueError): - logger.exception('Exception when parsing lsn') - break - elif parent_timeline > local_timeline: - break - - self._rewind_state = need_rewind and REWIND_STATUS.NEED or REWIND_STATUS.NOT_NEED - - def get_replica_timeline(self): - return self._get_local_timeline_lsn_from_replication_connection()[0] - - def replica_cached_timeline(self, master_timeline): - if not self._cached_replica_timeline or not master_timeline or self._cached_replica_timeline != master_timeline: - self._cached_replica_timeline = self.get_replica_timeline() - return self._cached_replica_timeline - - def get_master_timeline(self): - return self._cluster_info_state_get('timeline') - - def get_history(self, timeline): - history_path = 'pg_{0}/{1:08X}.history'.format(self.wal_name, timeline) - try: - cursor = self._cursor() - cursor.execute('SELECT isdir, modification FROM pg_catalog.pg_stat_file(%s)', (history_path,)) - isdir, modification = cursor.fetchone() - if not isdir: - cursor.execute('SELECT pg_catalog.pg_read_file(%s)', (history_path,)) - history = list(self.parse_history(cursor.fetchone()[0])) - if history[-1][0] == timeline - 1: - history[-1].append(modification.isoformat()) - return history - except Exception: - logger.exception('Failed to read and parse %s', (history_path,)) - - def rewind(self, leader): - if self.is_running() and not self.stop(checkpoint=False): - return logger.warning('Can not run pg_rewind because postgres is still running') - - # prepare pg_rewind connection - r = leader.conn_kwargs(self._rewind_credentials) - - # 1. make sure that we are really trying to rewind from the master - # 2. make sure that pg_control contains the new timeline by: - # running a checkpoint or - # waiting until Patroni on the master will expose checkpoint_after_promote=True - checkpoint_status = leader.checkpoint_after_promote if isinstance(leader, Leader) else None - if checkpoint_status is None: # master still runs the old Patroni - leader_status = self.checkpoint(leader.conn_kwargs(self._superuser)) - if leader_status: - return logger.warning('Can not use %s for rewind: %s', leader.name, leader_status) - elif not checkpoint_status: - return logger.info('Waiting for checkpoint on %s before rewind', leader.name) - elif not self.check_leader_is_not_in_recovery(**r): - return - - if self.pg_rewind(r): - self._rewind_state = REWIND_STATUS.SUCCESS - elif not self.check_leader_is_not_in_recovery(**r): - logger.warning('Failed to rewind because master %s become unreachable', leader.name) - else: - logger.error('Failed to rewind from healty master: %s', leader.name) - - for name in ('remove_data_directory_on_rewind_failure', 'remove_data_directory_on_diverged_timelines'): - if self.config.get(name): - logger.warning('%s is set. removing...', name) - self.remove_data_directory() - self._rewind_state = REWIND_STATUS.INITIAL - break - else: - self._rewind_state = REWIND_STATUS.FAILED - return False - - def trigger_check_diverged_lsn(self): - if self.can_rewind_or_reinitialize_allowed and self._rewind_state != REWIND_STATUS.NEED: - self._rewind_state = REWIND_STATUS.CHECK - - def rewind_or_reinitialize_needed_and_possible(self, leader): - if leader and leader.name != self.name and leader.conn_url and self._rewind_state == REWIND_STATUS.CHECK: - self._check_timeline_and_lsn(leader) - return leader and leader.conn_url and self._rewind_state == REWIND_STATUS.NEED - - @property - def rewind_executed(self): - return self._rewind_state > REWIND_STATUS.NOT_NEED - - @property - def rewind_failed(self): - return self._rewind_state == REWIND_STATUS.FAILED - - def follow(self, member, role='replica', timeout=None): - is_remote_master = isinstance(member, RemoteMember) - no_replication_slot = is_remote_master and member.no_replication_slot - restore_command = is_remote_master and member.restore_command - min_apply_delay = is_remote_master and member.recovery_min_apply_delay - archive_cleanup = is_remote_master and member.archive_cleanup_command - - primary_conninfo = self.primary_conninfo(member) - change_role = self.cb_called and (self.role in ('master', 'demoted') or - not {'standby_leader', 'replica'} - {self.role, role}) - - recovery_params = self.config.get('recovery_conf', {}).copy() - recovery_params.update({'standby_mode': 'on', 'recovery_target_timeline': 'latest'}) - if primary_conninfo: - recovery_params['primary_conninfo'] = primary_conninfo - if self.use_slots and not no_replication_slot: - required_name = is_remote_master and member.data.get('primary_slot_name') - name = required_name or slot_name_from_member_name(self.name) - recovery_params['primary_slot_name'] = name - if restore_command: - recovery_params['restore_command'] = restore_command - if min_apply_delay: - recovery_params['recovery_min_apply_delay'] = min_apply_delay - if archive_cleanup: - recovery_params['archive_cleanup_command'] = archive_cleanup - - self.write_recovery_conf(recovery_params) - - # When we demoting the master or standby_leader to replica or promoting replica to a standby_leader - # and we know for sure that postgres was already running before, we will only execute on_role_change - # callback and prevent execution of on_restart/on_start callback. - # If the role remains the same (replica or standby_leader), we will execute on_start or on_restart - if change_role: - self.__cb_pending = ACTION_NOOP - - if self.is_running(): - self.restart(block_callbacks=change_role, role=role) - else: - self.start(timeout=timeout, block_callbacks=change_role, role=role) - - if change_role: - # TODO: postpone this until start completes, or maybe do even earlier - self.call_nowait(ACTION_ON_ROLE_CHANGE) - return True - - def save_configuration_files(self): - """ - copy postgresql.conf to postgresql.conf.backup to be able to retrive configuration files - - originally stored as symlinks, those are normally skipped by pg_basebackup - - in case of WAL-E basebackup (see http://comments.gmane.org/gmane.comp.db.postgresql.wal-e/239) - """ - try: - for f in self._configuration_to_save: - config_file = os.path.join(self._config_dir, f) - backup_file = os.path.join(self._data_dir, f + '.backup') - if os.path.isfile(config_file): - shutil.copy(config_file, backup_file) - except IOError: - logger.exception('unable to create backup copies of configuration files') - return True - - def restore_configuration_files(self): - """ restore a previously saved postgresql.conf """ - try: - for f in self._configuration_to_save: - config_file = os.path.join(self._config_dir, f) - backup_file = os.path.join(self._data_dir, f + '.backup') - if not os.path.isfile(config_file): - if os.path.isfile(backup_file): - shutil.copy(backup_file, config_file) - # Previously we didn't backup pg_ident.conf, if file is missing just create empty - elif f == 'pg_ident.conf': - open(config_file, 'w').close() - except IOError: - logger.exception('unable to restore configuration files from backup') - - def _wait_promote(self, wait_seconds): - for _ in polling_loop(wait_seconds): - data = self.controldata() - if data.get('Database cluster state') == 'in production': - return True - - def promote(self, wait_seconds, access_is_restricted=False): - if self.role == 'master': - return True - ret = self.pg_ctl('promote', '-W') - if ret: - self.set_role('master') - logger.info("cleared rewind state after becoming the leader") - self._rewind_state = REWIND_STATUS.INITIAL - if not access_is_restricted: - self.call_nowait(ACTION_ON_ROLE_CHANGE) - ret = self._wait_promote(wait_seconds) - return ret - - def create_or_update_role(self, name, password, options): - options = list(map(str.upper, options)) - if 'NOLOGIN' not in options and 'LOGIN' not in options: - options.append('LOGIN') - - params = [name] - if password: - options.extend(['PASSWORD', '%s']) - params.extend([password, password]) - - sql = """DO $$ -BEGIN - SET local synchronous_commit = 'local'; - PERFORM * FROM pg_authid WHERE rolname = %s; - IF FOUND THEN - ALTER ROLE "{0}" WITH {1}; - ELSE - CREATE ROLE "{0}" WITH {1}; - END IF; -END;$$""".format(name, ' '.join(options)) - self.query(sql, *params) - - def timeline_wal_position(self): - # This method could be called from different threads (simultaneously with some other `_query` calls). - # If it is called not from main thread we will create a new cursor to execute statement. - if current_thread().ident == self.__thread_ident: - return self._cluster_info_state_get('timeline'), self._cluster_info_state_get('wal_position') - - with self.connection().cursor() as cursor: - cursor.execute(cluster_info_query.format(self.wal_name, self.lsn_name)) - return cursor.fetchone()[:2] - - def load_replication_slots(self): - if self.use_slots and self._schedule_load_slots: - replication_slots = {} - cursor = self._query('SELECT slot_name, slot_type, plugin, database FROM pg_catalog.pg_replication_slots') - for r in cursor: - value = {'type': r[1]} - if r[1] == 'logical': - value.update({'plugin': r[2], 'database': r[3]}) - replication_slots[r[0]] = value - self._replication_slots = replication_slots - self._schedule_load_slots = False - - def postmaster_start_time(self): - try: - cursor = self.query("SELECT pg_catalog.to_char(pg_catalog.pg_postmaster_start_time()," - " 'YYYY-MM-DD HH24:MI:SS.MS TZ')") - return cursor.fetchone()[0] - except psycopg2.Error: - return None - - def drop_replication_slot(self, name): - cursor = self._query(('SELECT pg_catalog.pg_drop_replication_slot(%s) WHERE EXISTS (SELECT 1 ' + - 'FROM pg_catalog.pg_replication_slots WHERE slot_name = %s AND NOT active)'), name, name) - # In normal situation rowcount should be 1, otherwise either slot doesn't exists or it is still active - return cursor.rowcount == 1 - - @staticmethod - def compare_slots(s1, s2): - return s1['type'] == s2['type'] and\ - (s1['type'] == 'physical' or s1['database'] == s2['database'] and s1['plugin'] == s2['plugin']) - - def sync_replication_slots(self, cluster): - if self.use_slots: - try: - self.load_replication_slots() - - slots = cluster.get_replication_slots(self.name, self.role) - - # drop old replication slots which are not presented in desired slots - for name in set(self._replication_slots) - set(slots): - if not self.drop_replication_slot(name): - logger.error("Failed to drop replication slot '%s'", name) - self._schedule_load_slots = True - - immediately_reserve = ', true' if self._major_version >= 90600 else '' - - logical_slots = defaultdict(dict) - for name, value in slots.items(): - if name in self._replication_slots and not self.compare_slots(value, self._replication_slots[name]): - logger.info("Trying to drop replication slot '%s' because value is changing from %s to %s", - name, self._replication_slots[name], value) - if not self.drop_replication_slot(name): - logger.error("Failed to drop replication slot '%s'", name) - self._schedule_load_slots = True - continue - self._replication_slots.pop(name) - if name not in self._replication_slots: - if value['type'] == 'physical': - try: - self._query(("SELECT pg_catalog.pg_create_physical_replication_slot(%s{0})" + - " WHERE NOT EXISTS (SELECT 1 FROM pg_catalog.pg_replication_slots" + - " WHERE slot_type = 'physical' AND slot_name = %s)").format( - immediately_reserve), name, name) - except Exception: - logger.exception("Failed to create physical replication slot '%s'", name) - self._schedule_load_slots = True - elif value['type'] == 'logical' and name not in self._replication_slots: - logical_slots[value['database']][name] = value - - # create new logical slots - for database, values in logical_slots.items(): - conn_kwargs = self._local_connect_kwargs - conn_kwargs['database'] = database - with self._get_connection_cursor(**conn_kwargs) as cur: - for name, value in values.items(): - try: - cur.execute("SELECT pg_catalog.pg_create_logical_replication_slot(%s, %s)" + - " WHERE NOT EXISTS (SELECT 1 FROM pg_catalog.pg_replication_slots" + - " WHERE slot_type = 'logical' AND slot_name = %s)", - (name, value['plugin'], name)) - except Exception: - logger.exception("Failed to create logical replication slot '%s' plugin='%s'", - name, value['plugin']) - self._schedule_load_slots = True - self._replication_slots = slots - except Exception: - logger.exception('Exception when changing replication slots') - self._schedule_load_slots = True - - def last_operation(self): - return str(self._cluster_info_state_get('wal_position')) - - def _post_restore(self): - self.delete_trigger_file() - self.restore_configuration_files() - - def _configure_server_parameters(self): - self._major_version = self.get_major_version() - self._server_parameters = self.get_server_parameters(self.config) - return True - - def clone(self, clone_member): - """ - - initialize the replica from an existing member (master or replica) - - initialize the replica using the replica creation method that - works without the replication connection (i.e. restore from on-disk - base backup) - """ - - self._rewind_state = REWIND_STATUS.INITIAL - ret = self.create_replica(clone_member) == 0 - if ret: - self._post_restore() - self._configure_server_parameters() - return ret - - def bootstrap(self, config): - """ Initialize a new node from scratch and start it. """ - pg_hba = config.get('pg_hba', []) - method = config.get('method') or 'initdb' - self._running_custom_bootstrap = method != 'initdb' and method in config and 'command' in config[method] - if self._running_custom_bootstrap: - do_initialize = self._custom_bootstrap - config = config[method] - else: - do_initialize = self._initdb - return do_initialize(config) and self.append_pg_hba(pg_hba) and self.save_configuration_files() \ - and self._configure_server_parameters() and self.start() - - def post_bootstrap(self, config, task): - try: - if 'username' in self._superuser and 'password' in self._superuser: - self.create_or_update_role(self._superuser['username'], self._superuser['password'], ['SUPERUSER']) - - task.complete(self.run_bootstrap_post_init(config)) - if task.result: - self.create_or_update_role(self._replication['username'], - self._replication.get('password'), ['REPLICATION']) - - if self._major_version >= 110000 and 'rewind' in self.config['authentication']: - rewind = self.config['authentication']['rewind'] - self.create_or_update_role(rewind['username'], rewind.get('password'), []) - for f in ('pg_ls_dir(text, boolean, boolean)', 'pg_stat_file(text, boolean)', - 'pg_read_binary_file(text)', 'pg_read_binary_file(text, bigint, bigint, boolean)'): - self.query('GRANT EXECUTE ON function pg_catalog.{0} TO "{1}"'.format(f, rewind['username'])) - - for name, value in (config.get('users') or {}).items(): - if name not in (self._superuser.get('username'), self._replication['username']): - self.create_or_update_role(name, value.get('password'), value.get('options', [])) - - # We were doing a custom bootstrap instead of running initdb, therefore we opened trust - # access from certain addresses to be able to reach cluster and change password - if self._running_custom_bootstrap: - self._running_custom_bootstrap = False - # If we don't have custom configuration for pg_hba.conf we need to restore original file - if not self.config.get('pg_hba'): - os.unlink(self._pg_hba_conf) - self.restore_configuration_files() - self._write_postgresql_conf() - self._replace_pg_ident() - # at this point there should be no recovery.conf - if os.path.isfile(self._recovery_conf) or os.path.islink(self._recovery_conf): - os.unlink(self._recovery_conf) - if self._server_parameters.get('hba_file') and \ - self._server_parameters['hba_file'] != self._pg_hba_conf: - self.restart() - else: - self._replace_pg_hba() - if self.pending_restart: - self.restart() - else: - self.reload() - time.sleep(1) # give a time to postgres to "reload" configuration files - self.close_connection() # close connection to reconnect with a new password - except Exception: - logger.exception('post_bootstrap') - task.complete(False) - return task.result - - def move_data_directory(self): - if os.path.isdir(self._data_dir) and not self.is_running(): - try: - new_name = '{0}_{1}'.format(self._data_dir, time.strftime('%Y-%m-%d-%H-%M-%S')) - logger.info('renaming data directory to %s', new_name) - os.rename(self._data_dir, new_name) - except OSError: - logger.exception("Could not rename data directory %s", self._data_dir) - - def remove_data_directory(self): - self.set_role('uninitialized') - logger.info('Removing data directory: %s', self._data_dir) - try: - if os.path.islink(self._data_dir): - os.unlink(self._data_dir) - elif not os.path.exists(self._data_dir): - return - elif os.path.isfile(self._data_dir): - os.remove(self._data_dir) - elif os.path.isdir(self._data_dir): - - # let's see if pg_xlog|pg_wal is a symlink, in this case we - # should clean the target - for pg_wal_dir in ('pg_xlog', 'pg_wal'): - pg_wal_path = os.path.join(self._data_dir, pg_wal_dir) - if os.path.exists(pg_wal_path) and os.path.islink(pg_wal_path): - pg_wal_realpath = os.path.realpath(pg_wal_path) - logger.info('Removing WAL directory: %s', pg_wal_realpath) - shutil.rmtree(pg_wal_realpath) - - shutil.rmtree(self._data_dir) - except (IOError, OSError): - logger.exception('Could not remove data directory %s', self._data_dir) - self.move_data_directory() - - def basebackup(self, conn_url, env, options): - # creates a replica data dir using pg_basebackup. - # this is the default, built-in create_replica_methods - # tries twice, then returns failure (as 1) - # uses "stream" as the xlog-method to avoid sync issues - # supports additional user-supplied options, those are not validated - maxfailures = 2 - ret = 1 - not_allowed_options = ('pgdata', 'format', 'wal-method', 'xlog-method', 'gzip', - 'version', 'compress', 'dbname', 'host', 'port', 'username', 'password') - user_options = self.process_user_options('basebackup', options, not_allowed_options, logger.error) - - for bbfailures in range(0, maxfailures): - with self._cancellable_lock: - if self._is_cancelled: - break - if not self.data_directory_empty(): - self.remove_data_directory() - try: - ret = self.cancellable_subprocess_call([self._pgcommand('pg_basebackup'), '--pgdata=' + self._data_dir, - '-X', 'stream', '--dbname=' + conn_url] + user_options, env=env) - if ret == 0: - break - else: - logger.error('Error when fetching backup: pg_basebackup exited with code=%s', ret) - - except Exception as e: - logger.error('Error when fetching backup with pg_basebackup: %s', e) - - if bbfailures < maxfailures - 1: - logger.warning('Trying again in 5 seconds') - time.sleep(5) - - return ret - - def pick_synchronous_standby(self, cluster): - """Finds the best candidate to be the synchronous standby. - - Current synchronous standby is always preferred, unless it has disconnected or does not want to be a - synchronous standby any longer. - - :returns tuple of candidate name or None, and bool showing if the member is the active synchronous standby. - """ - current = cluster.sync.sync_standby - current = current.lower() if current else current - members = {m.name.lower(): m for m in cluster.members} - candidates = [] - # Pick candidates based on who has flushed WAL farthest. - # TODO: for synchronous_commit = remote_write we actually want to order on write_location - for app_name, state, sync_state in self.query( - "SELECT pg_catalog.lower(application_name), state, sync_state" - " FROM pg_catalog.pg_stat_replication" - " ORDER BY flush_{0} DESC".format(self.lsn_name)): - member = members.get(app_name) - if state != 'streaming' or not member or member.tags.get('nosync', False): - continue - if sync_state == 'sync': - return app_name, True - if sync_state == 'potential' and app_name == current: - # Prefer current even if not the best one any more to avoid indecisivness and spurious swaps. - return current, False - if sync_state in ('async', 'potential'): - candidates.append(app_name) - - if candidates: - return candidates[0], False - return None, False - - def set_synchronous_standby(self, name): - """Sets a node to be synchronous standby and if changed does a reload for PostgreSQL.""" - if name and name != '*': - name = quote_ident(name) - if name != self._synchronous_standby_names: - if name is None: - self._server_parameters.pop('synchronous_standby_names', None) - else: - self._server_parameters['synchronous_standby_names'] = name - self._synchronous_standby_names = name - if self.state == 'running': - self._write_postgresql_conf() - self.reload() - - @staticmethod - def postgres_version_to_int(pg_version): - """Convert the server_version to integer - - >>> Postgresql.postgres_version_to_int('9.5.3') - 90503 - >>> Postgresql.postgres_version_to_int('9.3.13') - 90313 - >>> Postgresql.postgres_version_to_int('10.1') - 100001 - >>> Postgresql.postgres_version_to_int('10') # doctest: +IGNORE_EXCEPTION_DETAIL - Traceback (most recent call last): - ... - PostgresException: 'Invalid PostgreSQL version format: X.Y or X.Y.Z is accepted: 10' - >>> Postgresql.postgres_version_to_int('9.6') # doctest: +IGNORE_EXCEPTION_DETAIL - Traceback (most recent call last): - ... - PostgresException: 'Invalid PostgreSQL version format: X.Y or X.Y.Z is accepted: 9.6' - >>> Postgresql.postgres_version_to_int('a.b.c') # doctest: +IGNORE_EXCEPTION_DETAIL - Traceback (most recent call last): - ... - PostgresException: 'Invalid PostgreSQL version: a.b.c' - """ - - try: - components = list(map(int, pg_version.split('.'))) - except ValueError: - raise PostgresException('Invalid PostgreSQL version: {0}'.format(pg_version)) - - if len(components) < 2 or len(components) == 2 and components[0] < 10 or len(components) > 3: - raise PostgresException('Invalid PostgreSQL version format: X.Y or X.Y.Z is accepted: {0}' - .format(pg_version)) - - if len(components) == 2: - # new style verion numbers, i.e. 10.1 becomes 100001 - components.insert(1, 0) - - return int(''.join('{0:02d}'.format(c) for c in components)) - - @staticmethod - def postgres_major_version_to_int(pg_version): - """ - >>> Postgresql.postgres_major_version_to_int('10') - 100000 - >>> Postgresql.postgres_major_version_to_int('9.6') - 90600 - """ - return Postgresql.postgres_version_to_int(pg_version + '.0') - - def read_postmaster_opts(self): - """returns the list of option names/values from postgres.opts, Empty dict if read failed or no file""" - result = {} - try: - with open(os.path.join(self._data_dir, 'postmaster.opts')) as f: - data = f.read() - for opt in data.split('" "'): - if '=' in opt and opt.startswith('--'): - name, val = opt.split('=', 1) - result[name.strip('-')] = val.rstrip('"\n') - except IOError: - logger.exception('Error when reading postmaster.opts') - return result - - def single_user_mode(self, command=None, options=None): - """run a given command in a single-user mode. If the command is empty - then just start and stop""" - cmd = [self._pgcommand('postgres'), '--single', '-D', self._data_dir] - for opt, val in sorted((options or {}).items()): - cmd.extend(['-c', '{0}={1}'.format(opt, val)]) - # need a database name to connect - cmd.append(self._database) - return self.cancellable_subprocess_call(cmd, communicate_input=command) - - def cleanup_archive_status(self): - status_dir = os.path.join(self._data_dir, 'pg_' + self.wal_name, 'archive_status') - try: - for f in os.listdir(status_dir): - path = os.path.join(status_dir, f) - try: - if os.path.islink(path): - os.unlink(path) - elif os.path.isfile(path): - os.remove(path) - except OSError: - logger.exception('Unable to remove %s', path) - except OSError: - logger.exception('Unable to list %s', status_dir) - - def fix_cluster_state(self): - self.cleanup_archive_status() - - # Start in a single user mode and stop to produce a clean shutdown - opts = self.read_postmaster_opts() - opts.update({'archive_mode': 'on', 'archive_command': 'false'}) - if os.path.isfile(self._recovery_conf) or os.path.islink(self._recovery_conf): - os.unlink(self._recovery_conf) - return self.single_user_mode(options=opts) == 0 or None - - def cancellable_subprocess_call(self, *args, **kwargs): - for s in ('stdin', 'stdout', 'stderr'): - kwargs.pop(s, None) - - communicate_input = 'communicate_input' in kwargs - if communicate_input: - input_data = kwargs.pop('communicate_input', None) - if not isinstance(input_data, string_types): - input_data = '' - if input_data and input_data[-1] != '\n': - input_data += '\n' - kwargs['stdin'] = subprocess.PIPE - kwargs['stdout'] = open(os.devnull, 'w') - kwargs['stderr'] = subprocess.STDOUT - - try: - with self._cancellable_lock: - if self._is_cancelled: - raise PostgresException('cancelled') - - self._is_cancelled = False - self._cancellable = subprocess.Popen(*args, **kwargs) - - if communicate_input: - if input_data: - self._cancellable.communicate(input_data) - self._cancellable.stdin.close() - - return self._cancellable.wait() - finally: - with self._cancellable_lock: - self._cancellable = None - - def reset_is_cancelled(self): - with self._cancellable_lock: - self._is_cancelled = False - - def cancel(self): - with self._cancellable_lock: - self._is_cancelled = True - if self._cancellable is None or self._cancellable.returncode is not None: - return - self._cancellable.terminate() - - for _ in polling_loop(10): - with self._cancellable_lock: - if self._cancellable is None or self._cancellable.returncode is not None: - return - - with self._cancellable_lock: - if self._cancellable is not None and self._cancellable.returncode is None: - self._cancellable.kill() - - def schedule_sanity_checks_after_pause(self): - """ - After coming out of pause we have to: - 1. sync replication slots, because it might happen that slots were removed - 2. get new 'Database system identifier' to make sure that it wasn't changed - """ - self._schedule_load_slots = self.use_slots - self._sysid = None diff --git a/patroni/postgresql/__init__.py b/patroni/postgresql/__init__.py new file mode 100644 index 00000000..988577f6 --- /dev/null +++ b/patroni/postgresql/__init__.py @@ -0,0 +1,902 @@ +import logging +import os +import psycopg2 +import shlex +import shutil +import subprocess +import time + +from contextlib import contextmanager +from copy import deepcopy +from patroni.postgresql.callback_executor import CallbackExecutor +from patroni.postgresql.bootstrap import Bootstrap +from patroni.postgresql.cancellable import CancellableSubprocess +from patroni.postgresql.config import ConfigHandler +from patroni.postgresql.connection import Connection, get_connection_cursor +from patroni.postgresql.misc import parse_history, postgres_major_version_to_int +from patroni.postgresql.postmaster import PostmasterProcess +from patroni.postgresql.slots import SlotsHandler +from patroni.exceptions import PostgresConnectionException +from patroni.utils import Retry, RetryFailedError, polling_loop +from patroni.dcs import slot_name_from_member_name, RemoteMember +from threading import current_thread, Lock + + +logger = logging.getLogger(__name__) + +ACTION_ON_START = "on_start" +ACTION_ON_STOP = "on_stop" +ACTION_ON_RESTART = "on_restart" +ACTION_ON_RELOAD = "on_reload" +ACTION_ON_ROLE_CHANGE = "on_role_change" +ACTION_NOOP = "noop" + +STATE_RUNNING = 'running' +STATE_REJECT = 'rejecting connections' +STATE_NO_RESPONSE = 'not responding' +STATE_UNKNOWN = 'unknown' + +STOP_POLLING_INTERVAL = 1 + +cluster_info_query = ("SELECT CASE WHEN pg_catalog.pg_is_in_recovery() THEN 0 " + "ELSE ('x' || pg_catalog.substr(pg_catalog.pg_{0}file_name(" + "pg_catalog.pg_current_{0}_{1}()), 1, 8))::bit(32)::int END, " + "CASE WHEN pg_catalog.pg_is_in_recovery() THEN GREATEST(" + " pg_catalog.pg_{0}_{1}_diff(COALESCE(" + "pg_catalog.pg_last_{0}_receive_{1}(), '0/0'), '0/0')::bigint," + " pg_catalog.pg_{0}_{1}_diff(pg_catalog.pg_last_{0}_replay_{1}(), '0/0')::bigint)" + "ELSE pg_catalog.pg_{0}_{1}_diff(pg_catalog.pg_current_{0}_{1}(), '0/0')::bigint END") + + +@contextmanager +def null_context(): + yield + + +class Postgresql(object): + + def __init__(self, config): + self.name = config['name'] + self.scope = config['scope'] + self._data_dir = config['data_dir'] + self._database = config.get('database', 'postgres') + self._version_file = os.path.join(self._data_dir, 'PG_VERSION') + self._major_version = self.get_major_version() + + self._state_lock = Lock() + self.set_state('stopped') + + self._pending_restart = False + self._connection = Connection() + self.config = ConfigHandler(self, config) + + self._bin_dir = config.get('bin_dir') or '' + self.bootstrap = Bootstrap(self) + self.bootstrapping = False + self.__thread_ident = current_thread().ident + + self.slots_handler = SlotsHandler(self) + + self._pgpass = config.get('pgpass') or os.path.join(os.path.expanduser('~'), 'pgpass') + self._callback_executor = CallbackExecutor() + self.__cb_called = False + self.__cb_pending = None + + self.cancellable = CancellableSubprocess() + + self._sysid = None + self.retry = Retry(max_tries=-1, deadline=config['retry_timeout']/2.0, max_delay=1, + retry_exceptions=PostgresConnectionException) + + # Retry 'pg_is_in_recovery()' only once + self._is_leader_retry = Retry(max_tries=1, deadline=config['retry_timeout']/2.0, max_delay=1, + retry_exceptions=PostgresConnectionException) + + self._role_lock = Lock() + self.set_role(self.get_postgres_role_from_data_directory()) + self._state_entry_timestamp = None + + self._cluster_info_state = {} + self._cached_replica_timeline = None + + # Last known running process + self._postmaster_proc = None + + if self.is_running(): + self.set_state('running') + self.set_role('master' if self.is_leader() else 'replica') + self.config.write_postgresql_conf() # we are "joining" already running postgres + hba_saved = self.config.replace_pg_hba() + ident_saved = self.config.replace_pg_ident() + if hba_saved or ident_saved: + self.reload() + elif self.role == 'master': + self.set_role('demoted') + + @property + def create_replica_methods(self): + return self.config.get('create_replica_methods', []) or self.config.get('create_replica_method', []) + + @property + def major_version(self): + return self._major_version + + @property + def database(self): + return self._database + + @property + def data_dir(self): + return self._data_dir + + @property + def callback(self): + return self.config.get('callbacks') or {} + + @property + def wal_name(self): + return 'wal' if self._major_version >= 100000 else 'xlog' + + @property + def lsn_name(self): + return 'lsn' if self._major_version >= 100000 else 'location' + + def _version_file_exists(self): + return not self.data_directory_empty() and os.path.isfile(self._version_file) + + def get_major_version(self): + if self._version_file_exists(): + try: + with open(self._version_file) as f: + return postgres_major_version_to_int(f.read().strip()) + except Exception: + logger.exception('Failed to read PG_VERSION from %s', self._data_dir) + return 0 + + def pgcommand(self, cmd): + """Returns path to the specified PostgreSQL command""" + return os.path.join(self._bin_dir, cmd) + + def pg_ctl(self, cmd, *args, **kwargs): + """Builds and executes pg_ctl command + + :returns: `!True` when return_code == 0, otherwise `!False`""" + + pg_ctl = [self.pgcommand('pg_ctl'), cmd] + return subprocess.call(pg_ctl + ['-D', self._data_dir] + list(args), **kwargs) == 0 + + def pg_isready(self): + """Runs pg_isready to see if PostgreSQL is accepting connections. + + :returns: 'ok' if PostgreSQL is up, 'reject' if starting up, 'no_resopnse' if not up.""" + + r = self.config.local_connect_kwargs + cmd = [self.pgcommand('pg_isready'), '-p', r['port'], '-d', self._database] + + # Host is not set if we are connecting via default unix socket + if 'host' in r: + cmd.extend(['-h', r['host']]) + + # We only need the username because pg_isready does not try to authenticate + if 'user' in r: + cmd.extend(['-U', r['user']]) + + ret = subprocess.call(cmd) + return_codes = {0: STATE_RUNNING, + 1: STATE_REJECT, + 2: STATE_NO_RESPONSE, + 3: STATE_UNKNOWN} + return return_codes.get(ret, STATE_UNKNOWN) + + def reload_config(self, config): + self.config.reload_config(config) + self._is_leader_retry.deadline = self.retry.deadline = config['retry_timeout']/2.0 + + @property + def pending_restart(self): + return self._pending_restart + + def set_pending_restart(self, value): + self._pending_restart = value + + @property + def sysid(self): + if not self._sysid and not self.bootstrapping: + data = self.controldata() + self._sysid = data.get('Database system identifier', "") + return self._sysid + + def get_postgres_role_from_data_directory(self): + if self.data_directory_empty(): + return 'uninitialized' + elif self.config.recovery_conf_exists(): + return 'replica' + else: + return 'master' + + @property + def server_version(self): + return self._connection.server_version + + def connection(self): + return self._connection.get() + + def set_connection_kwargs(self, kwargs): + self._connection.set_conn_kwargs(kwargs) + + def _query(self, sql, *params): + """We are always using the same cursor, therefore this method is not thread-safe!!! + You can call it from different threads only if you are holding explicit `AsyncExecutor` lock, + because the main thread is always holding this lock when running HA cycle.""" + cursor = None + try: + cursor = self._connection.cursor() + cursor.execute(sql, params) + return cursor + except psycopg2.Error as e: + if cursor and cursor.connection.closed == 0: + # When connected via unix socket, psycopg2 can't recoginze 'connection lost' + # and leaves `_cursor_holder.connection.closed == 0`, but psycopg2.OperationalError + # is still raised (what is correct). It doesn't make sense to continiue with existing + # connection and we will close it, to avoid its reuse by the `cursor` method. + if isinstance(e, psycopg2.OperationalError): + self._connection.close() + else: + raise e + if self.state == 'restarting': + raise RetryFailedError('cluster is being restarted') + raise PostgresConnectionException('connection problems') + + def query(self, sql, *args, **kwargs): + if not kwargs.get('retry', True): + return self._query(sql, *args) + try: + return self.retry(self._query, sql, *args) + except RetryFailedError as e: + raise PostgresConnectionException(str(e)) + + def data_directory_empty(self): + return not os.path.exists(self._data_dir) or os.listdir(self._data_dir) == [] + + def write_pgpass(self, record): + if 'user' not in record or 'password' not in record: + return os.environ.copy() + + with open(self._pgpass, 'w') as f: + if os.name != 'nt': + os.fchmod(f.fileno(), 0o600) + f.write('{host}:{port}:*:{user}:{password}\n'.format(**record)) + + env = os.environ.copy() + env['PGPASSFILE'] = self._pgpass + return env + + def replica_method_options(self, method): + return deepcopy(self.config.get(method, {})) + + def replica_method_can_work_without_replication_connection(self, method): + return method != 'basebackup' and self.replica_method_options(method).get('no_master') + + def can_create_replica_without_replication_connection(self, replica_methods=None): + """ go through the replication methods to see if there are ones + that does not require a working replication connection. + """ + if replica_methods is None: + replica_methods = self.create_replica_methods + return any(self.replica_method_can_work_without_replication_connection(m) for m in replica_methods) + + def reset_cluster_info_state(self): + self._cluster_info_state = {} + + def _cluster_info_state_get(self, name): + if not self._cluster_info_state: + stmt = cluster_info_query.format(self.wal_name, self.lsn_name) + try: + result = self._is_leader_retry(self._query, stmt).fetchone() + self._cluster_info_state = dict(zip(['timeline', 'wal_position'], result)) + except RetryFailedError as e: # SELECT failed two times + self._cluster_info_state = {'error': str(e)} + if not self.is_starting() and self.pg_isready() == STATE_REJECT: + self.set_state('starting') + + if 'error' in self._cluster_info_state: + raise PostgresConnectionException(self._cluster_info_state['error']) + + return self._cluster_info_state.get(name) + + def is_leader(self): + return bool(self._cluster_info_state_get('timeline')) + + def is_running(self): + """Returns PostmasterProcess if one is running on the data directory or None. If most recently seen process + is running updates the cached process based on pid file.""" + if self._postmaster_proc: + if self._postmaster_proc.is_running(): + return self._postmaster_proc + self._postmaster_proc = None + + # we noticed that postgres was restarted, force syncing of replication + self.slots_handler.schedule() + + self._postmaster_proc = PostmasterProcess.from_pidfile(self._data_dir) + return self._postmaster_proc + + @property + def cb_called(self): + return self.__cb_called + + def call_nowait(self, cb_name): + """ pick a callback command and call it without waiting for it to finish """ + if self.bootstrapping: + return + if cb_name in (ACTION_ON_START, ACTION_ON_STOP, ACTION_ON_RESTART, ACTION_ON_ROLE_CHANGE): + self.__cb_called = True + + if self.callback and cb_name in self.callback: + cmd = self.callback[cb_name] + try: + cmd = shlex.split(self.callback[cb_name]) + [cb_name, self.role, self.scope] + self._callback_executor.call(cmd) + except Exception: + logger.exception('callback %s %s %s %s failed', cmd, cb_name, self.role, self.scope) + + @property + def role(self): + with self._role_lock: + return self._role + + def set_role(self, value): + with self._role_lock: + self._role = value + + @property + def state(self): + with self._state_lock: + return self._state + + def set_state(self, value): + with self._state_lock: + self._state = value + self._state_entry_timestamp = time.time() + + def time_in_state(self): + return time.time() - self._state_entry_timestamp + + def is_starting(self): + return self.state == 'starting' + + def wait_for_port_open(self, postmaster, timeout): + """Waits until PostgreSQL opens ports.""" + for _ in polling_loop(timeout): + if self.cancellable.is_cancelled: + return False + + if not postmaster.is_running(): + logger.error('postmaster is not running') + self.set_state('start failed') + return False + + isready = self.pg_isready() + if isready != STATE_NO_RESPONSE: + if isready not in [STATE_REJECT, STATE_RUNNING]: + logger.warning("Can't determine PostgreSQL startup status, assuming running") + return True + + logger.warning("Timed out waiting for PostgreSQL to start") + return False + + def start(self, timeout=None, task=None, block_callbacks=False, role=None): + """Start PostgreSQL + + Waits for postmaster to open ports or terminate so pg_isready can be used to check startup completion + or failure. + + :returns: True if start was initiated and postmaster ports are open, False if start failed""" + # make sure we close all connections established against + # the former node, otherwise, we might get a stalled one + # after kill -9, which would report incorrect data to + # patroni. + self._connection.close() + + if self.is_running(): + logger.error('Cannot start PostgreSQL because one is already running.') + self.set_state('starting') + return True + + if not block_callbacks: + self.__cb_pending = ACTION_ON_START + + self.set_role(role or self.get_postgres_role_from_data_directory()) + + self.set_state('starting') + self._pending_restart = False + + configuration = self.config.effective_configuration + self.config.write_postgresql_conf(configuration) + self.config.resolve_connection_addresses() + self.config.replace_pg_hba() + self.config.replace_pg_ident() + + options = ['--{0}={1}'.format(p, configuration[p]) for p in self.config.CMDLINE_OPTIONS + if p in configuration and p != 'wal_keep_segments'] + + if self.cancellable.is_cancelled: + return False + + with task or null_context(): + if task and task.is_cancelled: + logger.info("PostgreSQL start cancelled.") + return False + + self._postmaster_proc = PostmasterProcess.start(self.pgcommand('postgres'), + self._data_dir, + self.config.postgresql_conf, + options) + + if task: + task.complete(self._postmaster_proc) + + start_timeout = timeout + if not start_timeout: + try: + start_timeout = float(self.config.get('pg_ctl_timeout', 60)) + except ValueError: + start_timeout = 60 + + # We want postmaster to open ports before we continue + if not self._postmaster_proc or not self.wait_for_port_open(self._postmaster_proc, start_timeout): + return False + + ret = self.wait_for_startup(start_timeout) + if ret is not None: + return ret + elif timeout is not None: + return False + else: + return None + + def checkpoint(self, connect_kwargs=None): + check_not_is_in_recovery = connect_kwargs is not None + connect_kwargs = connect_kwargs or self.config.local_connect_kwargs + for p in ['connect_timeout', 'options']: + connect_kwargs.pop(p, None) + try: + with get_connection_cursor(**connect_kwargs) as cur: + cur.execute("SET statement_timeout = 0") + if check_not_is_in_recovery: + cur.execute('SELECT pg_catalog.pg_is_in_recovery()') + if cur.fetchone()[0]: + return 'is_in_recovery=true' + return cur.execute('CHECKPOINT') + except psycopg2.Error: + logger.exception('Exception during CHECKPOINT') + return 'not accessible or not healty' + + def stop(self, mode='fast', block_callbacks=False, checkpoint=None, on_safepoint=None): + """Stop PostgreSQL + + Supports a callback when a safepoint is reached. A safepoint is when no user backend can return a successful + commit to users. Currently this means we wait for user backends to close. But in the future alternate mechanisms + could be added. + + :param on_safepoint: This callback is called when no user backends are running. + """ + if checkpoint is None: + checkpoint = False if mode == 'immediate' else True + + success, pg_signaled = self._do_stop(mode, block_callbacks, checkpoint, on_safepoint) + if success: + # block_callbacks is used during restart to avoid + # running start/stop callbacks in addition to restart ones + if not block_callbacks: + self.set_state('stopped') + if pg_signaled: + self.call_nowait(ACTION_ON_STOP) + else: + logger.warning('pg_ctl stop failed') + self.set_state('stop failed') + return success + + def _do_stop(self, mode, block_callbacks, checkpoint, on_safepoint): + postmaster = self.is_running() + if not postmaster: + if on_safepoint: + on_safepoint() + return True, False + + if checkpoint and not self.is_starting(): + self.checkpoint() + + if not block_callbacks: + self.set_state('stopping') + + # Send signal to postmaster to stop + success = postmaster.signal_stop(mode) + if success is not None: + if success and on_safepoint: + on_safepoint() + return success, True + + # We can skip safepoint detection if we don't have a callback + if on_safepoint: + # Wait for our connection to terminate so we can be sure that no new connections are being initiated + self._wait_for_connection_close(postmaster) + postmaster.wait_for_user_backends_to_close() + on_safepoint() + + postmaster.wait() + + return True, True + + @staticmethod + def terminate_starting_postmaster(postmaster): + """Terminates a postmaster that has not yet opened ports or possibly even written a pid file. Blocks + until the process goes away.""" + postmaster.signal_stop('immediate') + postmaster.wait() + + def _wait_for_connection_close(self, postmaster): + try: + with self.connection().cursor() as cur: + while postmaster.is_running(): # Need a timeout here? + cur.execute("SELECT 1") + time.sleep(STOP_POLLING_INTERVAL) + except psycopg2.Error: + pass + + def reload(self): + ret = self.pg_ctl('reload') + if ret: + self.call_nowait(ACTION_ON_RELOAD) + return ret + + def check_for_startup(self): + """Checks PostgreSQL status and returns if PostgreSQL is in the middle of startup.""" + return self.is_starting() and not self.check_startup_state_changed() + + def check_startup_state_changed(self): + """Checks if PostgreSQL has completed starting up or failed or still starting. + + Should only be called when state == 'starting' + + :returns: True if state was changed from 'starting' + """ + ready = self.pg_isready() + + if ready == STATE_REJECT: + return False + elif ready == STATE_NO_RESPONSE: + self.set_state('start failed') + self.slots_handler.schedule(False) # TODO: can remove this? + self.config.save_configuration_files(True) # TODO: maybe remove this? + return True + else: + if ready != STATE_RUNNING: + # Bad configuration or unexpected OS error. No idea of PostgreSQL status. + # Let the main loop of run cycle clean up the mess. + logger.warning("%s status returned from pg_isready", + "Unknown" if ready == STATE_UNKNOWN else "Invalid") + self.set_state('running') + self.slots_handler.schedule() + self.config.save_configuration_files(True) + # TODO: __cb_pending can be None here after PostgreSQL restarts on its own. Do we want to call the callback? + # Previously we didn't even notice. + action = self.__cb_pending or ACTION_ON_START + self.call_nowait(action) + self.__cb_pending = None + + return True + + def wait_for_startup(self, timeout=None): + """Waits for PostgreSQL startup to complete or fail. + + :returns: True if start was successful, False otherwise""" + if not self.is_starting(): + # Should not happen + logger.warning("wait_for_startup() called when not in starting state") + + while not self.check_startup_state_changed(): + if self.cancellable.is_cancelled or timeout and self.time_in_state() > timeout: + return None + time.sleep(1) + + return self.state == 'running' + + def restart(self, timeout=None, task=None, block_callbacks=False, role=None): + """Restarts PostgreSQL. + + When timeout parameter is set the call will block either until PostgreSQL has started, failed to start or + timeout arrives. + + :returns: True when restart was successful and timeout did not expire when waiting. + """ + self.set_state('restarting') + if not block_callbacks: + self.__cb_pending = ACTION_ON_RESTART + ret = self.stop(block_callbacks=True) and self.start(timeout, task, True, role) + if not ret and not self.is_starting(): + self.set_state('restart failed ({0})'.format(self.state)) + return ret + + def is_healthy(self): + if not self.is_running(): + logger.warning('Postgresql is not running.') + return False + return True + + def controldata(self): + """ return the contents of pg_controldata, or non-True value if pg_controldata call failed """ + result = {} + # Don't try to call pg_controldata during backup restore + if self._version_file_exists() and self.state != 'creating replica': + try: + env = {'LANG': 'C', 'LC_ALL': 'C', 'PATH': os.getenv('PATH')} + if os.getenv('SYSTEMROOT') is not None: + env['SYSTEMROOT'] = os.getenv('SYSTEMROOT') + data = subprocess.check_output([self.pgcommand('pg_controldata'), self._data_dir], env=env) + if data: + data = data.decode('utf-8').splitlines() + # pg_controldata output depends on major verion. Some of parameters are prefixed by 'Current ' + result = {l.split(':')[0].replace('Current ', '', 1): l.split(':', 1)[1].strip() for l in data + if l and ':' in l} + except subprocess.CalledProcessError: + logger.exception("Error when calling pg_controldata") + return result + + @contextmanager + def get_replication_connection_cursor(self, host='localhost', port=5432, database=None, **kwargs): + replication = self.config.replication + with get_connection_cursor(host=host, port=int(port), database=database or self._database, replication=1, + user=replication['username'], password=replication.get('password'), + connect_timeout=3, options='-c statement_timeout=2000') as cur: + yield cur + + def get_local_timeline_lsn_from_replication_connection(self): + timeline = lsn = None + try: + with self.get_replication_connection_cursor(**self.config.local_replication_address) as cur: + cur.execute('IDENTIFY_SYSTEM') + timeline, lsn = cur.fetchone()[1:3] + except Exception: + logger.exception('Can not fetch local timeline and lsn from replication connection') + return timeline, lsn + + def get_replica_timeline(self): + return self.get_local_timeline_lsn_from_replication_connection()[0] + + def replica_cached_timeline(self, master_timeline): + if not self._cached_replica_timeline or not master_timeline or self._cached_replica_timeline != master_timeline: + self._cached_replica_timeline = self.get_replica_timeline() + return self._cached_replica_timeline + + def get_master_timeline(self): + return self._cluster_info_state_get('timeline') + + def get_history(self, timeline): + history_path = 'pg_{0}/{1:08X}.history'.format(self.wal_name, timeline) + try: + cursor = self._connection.cursor() + cursor.execute('SELECT isdir, modification FROM pg_catalog.pg_stat_file(%s)', (history_path,)) + isdir, modification = cursor.fetchone() + if not isdir: + cursor.execute('SELECT pg_catalog.pg_read_file(%s)', (history_path,)) + history = list(parse_history(cursor.fetchone()[0])) + if history[-1][0] == timeline - 1: + history[-1].append(modification.isoformat()) + return history + except Exception: + logger.exception('Failed to read and parse %s', (history_path,)) + + def follow(self, member, role='replica', timeout=None): + is_remote_master = isinstance(member, RemoteMember) + no_replication_slot = is_remote_master and member.no_replication_slot + restore_command = is_remote_master and member.restore_command + min_apply_delay = is_remote_master and member.recovery_min_apply_delay + archive_cleanup = is_remote_master and member.archive_cleanup_command + + primary_conninfo = self.config.primary_conninfo(member) + change_role = self.cb_called and (self.role in ('master', 'demoted') or + not {'standby_leader', 'replica'} - {self.role, role}) + + recovery_params = self.config.get('recovery_conf', {}).copy() + recovery_params.update({'standby_mode': 'on', 'recovery_target_timeline': 'latest'}) + if primary_conninfo: + recovery_params['primary_conninfo'] = primary_conninfo + if self.slots_handler.use_slots and not no_replication_slot: + required_name = is_remote_master and member.data.get('primary_slot_name') + name = required_name or slot_name_from_member_name(self.name) + recovery_params['primary_slot_name'] = name + if restore_command: + recovery_params['restore_command'] = restore_command + if min_apply_delay: + recovery_params['recovery_min_apply_delay'] = min_apply_delay + if archive_cleanup: + recovery_params['archive_cleanup_command'] = archive_cleanup + + self.config.write_recovery_conf(recovery_params) + + # When we demoting the master or standby_leader to replica or promoting replica to a standby_leader + # and we know for sure that postgres was already running before, we will only execute on_role_change + # callback and prevent execution of on_restart/on_start callback. + # If the role remains the same (replica or standby_leader), we will execute on_start or on_restart + if change_role: + self.__cb_pending = ACTION_NOOP + + if self.is_running(): + self.restart(block_callbacks=change_role, role=role) + else: + self.start(timeout=timeout, block_callbacks=change_role, role=role) + + if change_role: + # TODO: postpone this until start completes, or maybe do even earlier + self.call_nowait(ACTION_ON_ROLE_CHANGE) + return True + + def _wait_promote(self, wait_seconds): + for _ in polling_loop(wait_seconds): + data = self.controldata() + if data.get('Database cluster state') == 'in production': + return True + + def promote(self, wait_seconds, on_success=None, access_is_restricted=False): + if self.role == 'master': + return True + ret = self.pg_ctl('promote', '-W') + if ret: + self.set_role('master') + if on_success is not None: + on_success() + if not access_is_restricted: + self.call_nowait(ACTION_ON_ROLE_CHANGE) + ret = self._wait_promote(wait_seconds) + return ret + + def timeline_wal_position(self): + # This method could be called from different threads (simultaneously with some other `_query` calls). + # If it is called not from main thread we will create a new cursor to execute statement. + if current_thread().ident == self.__thread_ident: + return self._cluster_info_state_get('timeline'), self._cluster_info_state_get('wal_position') + + with self.connection().cursor() as cursor: + cursor.execute(cluster_info_query.format(self.wal_name, self.lsn_name)) + return cursor.fetchone()[:2] + + def postmaster_start_time(self): + try: + cursor = self.query("SELECT pg_catalog.to_char(pg_catalog.pg_postmaster_start_time()," + " 'YYYY-MM-DD HH24:MI:SS.MS TZ')") + return cursor.fetchone()[0] + except psycopg2.Error: + return None + + def last_operation(self): + return str(self._cluster_info_state_get('wal_position')) + + def configure_server_parameters(self): + self._major_version = self.get_major_version() + self.config.setup_server_parameters() + return True + + def move_data_directory(self): + if os.path.isdir(self._data_dir) and not self.is_running(): + try: + new_name = '{0}_{1}'.format(self._data_dir, time.strftime('%Y-%m-%d-%H-%M-%S')) + logger.info('renaming data directory to %s', new_name) + os.rename(self._data_dir, new_name) + except OSError: + logger.exception("Could not rename data directory %s", self._data_dir) + + def remove_data_directory(self): + self.set_role('uninitialized') + logger.info('Removing data directory: %s', self._data_dir) + try: + if os.path.islink(self._data_dir): + os.unlink(self._data_dir) + elif not os.path.exists(self._data_dir): + return + elif os.path.isfile(self._data_dir): + os.remove(self._data_dir) + elif os.path.isdir(self._data_dir): + + # let's see if pg_xlog|pg_wal is a symlink, in this case we + # should clean the target + for pg_wal_dir in ('pg_xlog', 'pg_wal'): + pg_wal_path = os.path.join(self._data_dir, pg_wal_dir) + if os.path.exists(pg_wal_path) and os.path.islink(pg_wal_path): + pg_wal_realpath = os.path.realpath(pg_wal_path) + logger.info('Removing WAL directory: %s', pg_wal_realpath) + shutil.rmtree(pg_wal_realpath) + + shutil.rmtree(self._data_dir) + except (IOError, OSError): + logger.exception('Could not remove data directory %s', self._data_dir) + self.move_data_directory() + + def pick_synchronous_standby(self, cluster): + """Finds the best candidate to be the synchronous standby. + + Current synchronous standby is always preferred, unless it has disconnected or does not want to be a + synchronous standby any longer. + + :returns tuple of candidate name or None, and bool showing if the member is the active synchronous standby. + """ + current = cluster.sync.sync_standby + current = current.lower() if current else current + members = {m.name.lower(): m for m in cluster.members} + candidates = [] + # Pick candidates based on who has flushed WAL farthest. + # TODO: for synchronous_commit = remote_write we actually want to order on write_location + for app_name, state, sync_state in self.query( + "SELECT pg_catalog.lower(application_name), state, sync_state" + " FROM pg_catalog.pg_stat_replication" + " ORDER BY flush_{0} DESC".format(self.lsn_name)): + member = members.get(app_name) + if state != 'streaming' or not member or member.tags.get('nosync', False): + continue + if sync_state == 'sync': + return app_name, True + if sync_state == 'potential' and app_name == current: + # Prefer current even if not the best one any more to avoid indecisivness and spurious swaps. + return current, False + if sync_state in ('async', 'potential'): + candidates.append(app_name) + + if candidates: + return candidates[0], False + return None, False + + def read_postmaster_opts(self): + """returns the list of option names/values from postgres.opts, Empty dict if read failed or no file""" + result = {} + try: + with open(os.path.join(self._data_dir, 'postmaster.opts')) as f: + data = f.read() + for opt in data.split('" "'): + if '=' in opt and opt.startswith('--'): + name, val = opt.split('=', 1) + result[name.strip('-')] = val.rstrip('"\n') + except IOError: + logger.exception('Error when reading postmaster.opts') + return result + + def single_user_mode(self, command=None, options=None): + """run a given command in a single-user mode. If the command is empty - then just start and stop""" + cmd = [self.pgcommand('postgres'), '--single', '-D', self._data_dir] + for opt, val in sorted((options or {}).items()): + cmd.extend(['-c', '{0}={1}'.format(opt, val)]) + # need a database name to connect + cmd.append(self._database) + return self.cancellable.call(cmd, communicate_input=command) + + def cleanup_archive_status(self): + status_dir = os.path.join(self._data_dir, 'pg_' + self.wal_name, 'archive_status') + try: + for f in os.listdir(status_dir): + path = os.path.join(status_dir, f) + try: + if os.path.islink(path): + os.unlink(path) + elif os.path.isfile(path): + os.remove(path) + except OSError: + logger.exception('Unable to remove %s', path) + except OSError: + logger.exception('Unable to list %s', status_dir) + + def fix_cluster_state(self): + self.cleanup_archive_status() + + # Start in a single user mode and stop to produce a clean shutdown + opts = self.read_postmaster_opts() + opts.update({'archive_mode': 'on', 'archive_command': 'false'}) + self.config.remove_recovery_conf() + return self.single_user_mode(options=opts) == 0 or None + + def schedule_sanity_checks_after_pause(self): + """ + After coming out of pause we have to: + 1. sync replication slots, because it might happen that slots were removed + 2. get new 'Database system identifier' to make sure that it wasn't changed + """ + self.slots_handler.schedule() + self._sysid = None diff --git a/patroni/postgresql/bootstrap.py b/patroni/postgresql/bootstrap.py new file mode 100644 index 00000000..a0f3b9ba --- /dev/null +++ b/patroni/postgresql/bootstrap.py @@ -0,0 +1,373 @@ +import logging +import os +import shlex +import tempfile +import time + +from patroni.dcs import RemoteMember +from patroni.utils import deep_compare +from six import string_types +from six.moves.urllib.parse import quote_plus + +logger = logging.getLogger(__name__) + + +class Bootstrap(object): + + def __init__(self, postgresql): + self._postgresql = postgresql + self._running_custom_bootstrap = False + + @property + def running_custom_bootstrap(self): + return self._running_custom_bootstrap + + @staticmethod + def process_user_options(tool, options, not_allowed_options, error_handler): + user_options = [] + + def option_is_allowed(name): + ret = name not in not_allowed_options + if not ret: + error_handler('{0} option for {1} is not allowed'.format(name, tool)) + return ret + + if isinstance(options, dict): + for k, v in options.items(): + if k and v: + user_options.append('--{0}={1}'.format(k, v)) + elif isinstance(options, list): + for opt in options: + if isinstance(opt, string_types) and option_is_allowed(opt): + user_options.append('--{0}'.format(opt)) + elif isinstance(opt, dict): + keys = list(opt.keys()) + if len(keys) != 1 or not isinstance(opt[keys[0]], string_types) or not option_is_allowed(keys[0]): + error_handler('Error when parsing {0} key-value option {1}: only one key-value is allowed' + ' and value should be a string'.format(tool, opt[keys[0]])) + user_options.append('--{0}={1}'.format(keys[0], opt[keys[0]])) + else: + error_handler('Error when parsing {0} option {1}: value should be string value' + ' or a single key-value pair'.format(tool, opt)) + else: + error_handler('{0} options must be list ot dict'.format(tool)) + return user_options + + def _initdb(self, config): + self._postgresql.set_state('initalizing new cluster') + not_allowed_options = ('pgdata', 'nosync', 'pwfile', 'sync-only', 'version') + + def error_handler(e): + raise Exception(e) + + options = self.process_user_options('initdb', config.get('initdb') or [], not_allowed_options, error_handler) + pwfile = None + + if self._postgresql.config.superuser: + if 'username' in self._postgresql.config.superuser: + options.append('--username={0}'.format(self._postgresql.config.superuser['username'])) + if 'password' in self._postgresql.config.superuser: + (fd, pwfile) = tempfile.mkstemp() + os.write(fd, self._postgresql.config.superuser['password'].encode('utf-8')) + os.close(fd) + options.append('--pwfile={0}'.format(pwfile)) + options = ['-o', ' '.join(options)] if options else [] + + ret = self._postgresql.pg_ctl('initdb', *options) + if pwfile: + os.remove(pwfile) + if not ret: + self._postgresql.set_state('initdb failed') + return ret + + def _post_restore(self): + # make sure there is no trigger file or postgres will be automatically promoted + trigger_file = self._postgresql.config.get('recovery_conf', {}).get('trigger_file') or 'promote' + trigger_file = os.path.abspath(os.path.join(self._postgresql.data_dir, trigger_file)) + if os.path.exists(trigger_file): + os.unlink(trigger_file) + self._postgresql.config.restore_configuration_files() + + def _custom_bootstrap(self, config): + self._postgresql.set_state('running custom bootstrap script') + params = ['--scope=' + self._postgresql.scope, '--datadir=' + self._postgresql.data_dir] + try: + logger.info('Running custom bootstrap script: %s', config['command']) + if self._postgresql.cancellable.call(shlex.split(config['command']) + params) != 0: + self._postgresql.set_state('custom bootstrap failed') + return False + except Exception: + logger.exception('Exception during custom bootstrap') + return False + self._post_restore() + + if 'recovery_conf' in config: + self._postgresql.config.write_recovery_conf(config['recovery_conf']) + elif not config.get('keep_existing_recovery_conf'): + self._postgresql.config.remove_recovery_conf() + return True + + def call_post_bootstrap(self, config): + """ + runs a script after initdb or custom bootstrap script is called and waits until completion. + """ + cmd = config.get('post_bootstrap') or config.get('post_init') + if cmd: + r = self._postgresql.config.local_connect_kwargs + + if 'host' in r: + # '/tmp' => '%2Ftmp' for unix socket path + host = quote_plus(r['host']) if r['host'].startswith('/') else r['host'] + else: + host = '' + + # https://www.postgresql.org/docs/current/static/libpq-pgpass.html + # A host name of localhost matches both TCP (host name localhost) and Unix domain socket + # (pghost empty or the default socket directory) connections coming from the local machine. + r['host'] = 'localhost' # set it to localhost to write into pgpass + + if 'user' in r: + user = r['user'] + '@' + else: + user = '' + if 'password' in r: + import getpass + r.setdefault('user', os.environ.get('PGUSER', getpass.getuser())) + + connstring = 'postgres://{0}{1}:{2}/{3}'.format(user, host, r['port'], r['database']) + env = self._postgresql.write_pgpass(r) if 'password' in r else None + + try: + ret = self._postgresql.cancellable.call(shlex.split(cmd) + [connstring], env=env) + except OSError: + logger.error('post_init script %s failed', cmd) + return False + if ret != 0: + logger.error('post_init script %s returned non-zero code %d', cmd, ret) + return False + return True + + def create_replica(self, clone_member): + """ + create the replica according to the replica_method + defined by the user. this is a list, so we need to + loop through all methods the user supplies + """ + + self._postgresql.set_state('creating replica') + self._postgresql.schedule_sanity_checks_after_pause() + + is_remote_master = isinstance(clone_member, RemoteMember) + + # get list of replica methods either from clone member or from + # the config. If there is no configuration key, or no value is + # specified, use basebackup + replica_methods = (clone_member.create_replica_methods if is_remote_master + else self._postgresql.create_replica_methods) or ['basebackup'] + + if clone_member and clone_member.conn_url: + r = clone_member.conn_kwargs(self._postgresql.config.replication) + connstring = 'postgres://{user}@{host}:{port}/{database}'.format(**r) + # add the credentials to connect to the replica origin to pgpass. + env = self._postgresql.write_pgpass(r) + else: + connstring = '' + env = os.environ.copy() + # if we don't have any source, leave only replica methods that work without it + replica_methods = [r for r in replica_methods + if self._postgresql.replica_method_can_work_without_replication_connection(r)] + + # go through them in priority order + ret = 1 + for replica_method in replica_methods: + if self._postgresql.cancellable.is_cancelled: + break + + method_config = self._postgresql.replica_method_options(replica_method) + + # if the method is basebackup, then use the built-in + if replica_method == "basebackup": + ret = self.basebackup(connstring, env, method_config) + if ret == 0: + logger.info("replica has been created using basebackup") + # if basebackup succeeds, exit with success + break + else: + if not self._postgresql.data_directory_empty(): + if method_config.get('keep_data', False): + logger.info('Leaving data directory uncleaned') + else: + self._postgresql.remove_data_directory() + + cmd = replica_method + # user-defined method; check for configuration + # not required, actually + if method_config: + # look to see if the user has supplied a full command path + # if not, use the method name as the command + cmd = method_config.pop('command', cmd) + + # add the default parameters + if not method_config.get('no_params', False): + method_config.update({"scope": self._postgresql.scope, + "role": "replica", + "datadir": self._postgresql.data_dir, + "connstring": connstring}) + else: + for param in ('no_params', 'no_master', 'keep_data'): + method_config.pop(param, None) + params = ["--{0}={1}".format(arg, val) for arg, val in method_config.items()] + try: + # call script with the full set of parameters + ret = self._postgresql.cancellable.call(shlex.split(cmd) + params, env=env) + # if we succeeded, stop + if ret == 0: + logger.info('replica has been created using %s', replica_method) + break + else: + logger.error('Error creating replica using method %s: %s exited with code=%s', + replica_method, cmd, ret) + except Exception: + logger.exception('Error creating replica using method %s', replica_method) + ret = 1 + + self._postgresql.set_state('stopped') + return ret + + def basebackup(self, conn_url, env, options): + # creates a replica data dir using pg_basebackup. + # this is the default, built-in create_replica_methods + # tries twice, then returns failure (as 1) + # uses "stream" as the xlog-method to avoid sync issues + # supports additional user-supplied options, those are not validated + maxfailures = 2 + ret = 1 + not_allowed_options = ('pgdata', 'format', 'wal-method', 'xlog-method', 'gzip', + 'version', 'compress', 'dbname', 'host', 'port', 'username', 'password') + user_options = self.process_user_options('basebackup', options, not_allowed_options, logger.error) + + for bbfailures in range(0, maxfailures): + if self._postgresql.cancellable.is_cancelled: + break + if not self._postgresql.data_directory_empty(): + self._postgresql.remove_data_directory() + try: + ret = self._postgresql.cancellable.call([self._postgresql.pgcommand('pg_basebackup'), + '--pgdata=' + self._postgresql.data_dir, '-X', 'stream', + '--dbname=' + conn_url] + user_options, env=env) + if ret == 0: + break + else: + logger.error('Error when fetching backup: pg_basebackup exited with code=%s', ret) + + except Exception as e: + logger.error('Error when fetching backup with pg_basebackup: %s', e) + + if bbfailures < maxfailures - 1: + logger.warning('Trying again in 5 seconds') + time.sleep(5) + + return ret + + def clone(self, clone_member): + """ + - initialize the replica from an existing member (master or replica) + - initialize the replica using the replica creation method that + works without the replication connection (i.e. restore from on-disk + base backup) + """ + + ret = self.create_replica(clone_member) == 0 + if ret: + self._post_restore() + self._postgresql.configure_server_parameters() + return ret + + def bootstrap(self, config): + """ Initialize a new node from scratch and start it. """ + pg_hba = config.get('pg_hba', []) + method = config.get('method') or 'initdb' + self._running_custom_bootstrap = method != 'initdb' and method in config and 'command' in config[method] + if self._running_custom_bootstrap: + do_initialize = self._custom_bootstrap + config = config[method] + else: + do_initialize = self._initdb + return do_initialize(config) and self._postgresql.config.append_pg_hba(pg_hba) \ + and self._postgresql.config.save_configuration_files() \ + and self._postgresql.configure_server_parameters() and self._postgresql.start() + + def create_or_update_role(self, name, password, options): + options = list(map(str.upper, options)) + if 'NOLOGIN' not in options and 'LOGIN' not in options: + options.append('LOGIN') + + params = [name] + if password: + options.extend(['PASSWORD', '%s']) + params.extend([password, password]) + + sql = """DO $$ +BEGIN + SET local synchronous_commit = 'local'; + PERFORM * FROM pg_authid WHERE rolname = %s; + IF FOUND THEN + ALTER ROLE "{0}" WITH {1}; + ELSE + CREATE ROLE "{0}" WITH {1}; + END IF; +END;$$""".format(name, ' '.join(options)) + self._postgresql.query(sql, *params) + + def post_bootstrap(self, config, task): + try: + postgresql = self._postgresql + superuser = postgresql.config.superuser + if 'username' in superuser and 'password' in superuser: + self.create_or_update_role(superuser['username'], superuser['password'], ['SUPERUSER']) + + task.complete(self.call_post_bootstrap(config)) + if task.result: + replication = postgresql.config.replication + self.create_or_update_role(replication['username'], replication.get('password'), ['REPLICATION']) + + rewind = postgresql.config.rewind_credentials + if not deep_compare(rewind, superuser): + self.create_or_update_role(rewind['username'], rewind.get('password'), []) + for f in ('pg_ls_dir(text, boolean, boolean)', 'pg_stat_file(text, boolean)', + 'pg_read_binary_file(text)', 'pg_read_binary_file(text, bigint, bigint, boolean)'): + postgresql.query('GRANT EXECUTE ON function pg_catalog.{0} TO "{1}"' + .format(f, rewind['username'])) + + for name, value in (config.get('users') or {}).items(): + if all(name != a.get('username') for a in (superuser, replication, rewind)): + self.create_or_update_role(name, value.get('password'), value.get('options', [])) + + # We were doing a custom bootstrap instead of running initdb, therefore we opened trust + # access from certain addresses to be able to reach cluster and change password + if self._running_custom_bootstrap: + self._running_custom_bootstrap = False + # If we don't have custom configuration for pg_hba.conf we need to restore original file + if not postgresql.config.get('pg_hba'): + os.unlink(postgresql.config.pg_hba_conf) + postgresql.config.restore_configuration_files() + postgresql.config.write_postgresql_conf() + postgresql.config.replace_pg_ident() + + # at this point there should be no recovery.conf + postgresql.config.remove_recovery_conf() + + if postgresql.config.hba_file and postgresql.config.hba_file != postgresql.config.pg_hba_conf: + postgresql.restart() + else: + postgresql.config.replace_pg_hba() + if postgresql.pending_restart: + postgresql.restart() + else: + postgresql.reload() + time.sleep(1) # give a time to postgres to "reload" configuration files + postgresql.connection().close() # close connection to reconnect with a new password + except Exception: + logger.exception('post_bootstrap') + task.complete(False) + return task.result diff --git a/patroni/callback_executor.py b/patroni/postgresql/callback_executor.py similarity index 100% rename from patroni/callback_executor.py rename to patroni/postgresql/callback_executor.py diff --git a/patroni/postgresql/cancellable.py b/patroni/postgresql/cancellable.py new file mode 100644 index 00000000..866fa194 --- /dev/null +++ b/patroni/postgresql/cancellable.py @@ -0,0 +1,76 @@ +import logging +import os +import subprocess + +from patroni.exceptions import PostgresException +from patroni.utils import polling_loop +from six import string_types +from threading import Lock + +logger = logging.getLogger(__name__) + + +class CancellableSubprocess(object): + + def __init__(self): + self._is_cancelled = False + self._process = None + self._lock = Lock() + + def call(self, *args, **kwargs): + for s in ('stdin', 'stdout', 'stderr'): + kwargs.pop(s, None) + + communicate_input = 'communicate_input' in kwargs + if communicate_input: + input_data = kwargs.pop('communicate_input', None) + if not isinstance(input_data, string_types): + input_data = '' + if input_data and input_data[-1] != '\n': + input_data += '\n' + kwargs['stdin'] = subprocess.PIPE + kwargs['stdout'] = open(os.devnull, 'w') + kwargs['stderr'] = subprocess.STDOUT + + try: + with self._lock: + if self._is_cancelled: + raise PostgresException('cancelled') + + self._is_cancelled = False + self._process = subprocess.Popen(*args, **kwargs) + + if communicate_input: + if input_data: + self._process.communicate(input_data) + self._process.stdin.close() + + return self._process.wait() + finally: + with self._lock: + self._process = None + + def reset_is_cancelled(self): + with self._lock: + self._is_cancelled = False + + @property + def is_cancelled(self): + with self._lock: + return self._is_cancelled + + def cancel(self): + with self._lock: + self._is_cancelled = True + if self._process is None or self._process.returncode is not None: + return + self._process.terminate() + + for _ in polling_loop(10): + with self._lock: + if self._process is None or self._process.returncode is not None: + return + + with self._lock: + if self._process is not None and self._process.returncode is None: + self._process.kill() diff --git a/patroni/postgresql/config.py b/patroni/postgresql/config.py new file mode 100644 index 00000000..27b048fb --- /dev/null +++ b/patroni/postgresql/config.py @@ -0,0 +1,459 @@ +import logging +import os +import re +import shutil +import socket +import stat + +from requests.structures import CaseInsensitiveDict + +from ..utils import compare_values, parse_bool, parse_int, split_host_port + +logger = logging.getLogger(__name__) + +SYNC_STANDBY_NAME_RE = re.compile(r'^[A-Za-z_][A-Za-z_0-9\$]*$') + + +def quote_ident(value): + """Very simplified version of quote_ident""" + return value if SYNC_STANDBY_NAME_RE.match(value) else '"' + value + '"' + + +class ConfigHandler(object): + + # List of parameters which must be always passed to postmaster as command line options + # to make it not possible to change them with 'ALTER SYSTEM'. + # Some of these parameters have sane default value assigned and Patroni doesn't allow + # to decrease this value. E.g. 'wal_level' can't be lower then 'hot_standby' and so on. + # These parameters could be changed only globally, i.e. via DCS. + # P.S. 'listen_addresses' and 'port' are added here just for convenience, to mark them + # as a parameters which should always be passed through command line. + # + # Format: + # key - parameter name + # value - tuple(default_value, check_function, min_version) + # default_value -- some sane default value + # check_function -- if the new value is not correct must return `!False` + # min_version -- major version of PostgreSQL when parameter was introduced + CMDLINE_OPTIONS = CaseInsensitiveDict({ + 'listen_addresses': (None, lambda _: False, 90100), + 'port': (None, lambda _: False, 90100), + 'cluster_name': (None, lambda _: False, 90500), + 'wal_level': ('hot_standby', lambda v: v.lower() in ('hot_standby', 'replica', 'logical'), 90100), + 'hot_standby': ('on', lambda _: False, 90100), + 'max_connections': (100, lambda v: int(v) >= 100, 90100), + 'max_wal_senders': (10, lambda v: int(v) >= 10, 90100), + 'wal_keep_segments': (8, lambda v: int(v) >= 8, 90100), + 'max_prepared_transactions': (0, lambda v: int(v) >= 0, 90100), + 'max_locks_per_transaction': (64, lambda v: int(v) >= 64, 90100), + 'track_commit_timestamp': ('off', lambda v: parse_bool(v) is not None, 90500), + 'max_replication_slots': (10, lambda v: int(v) >= 10, 90400), + 'max_worker_processes': (8, lambda v: int(v) >= 8, 90400), + 'wal_log_hints': ('on', lambda _: False, 90400) + }) + + _CONFIG_WARNING_HEADER = '# Do not edit this file manually!\n# It will be overwritten by Patroni!\n' + + def __init__(self, postgresql, config): + self._postgresql = postgresql + self._config_dir = os.path.abspath(config.get('config_dir') or postgresql.data_dir) + config_base_name = config.get('config_base_name', 'postgresql') + self._postgresql_conf = os.path.join(self._config_dir, config_base_name + '.conf') + self._postgresql_base_conf_name = config_base_name + '.base.conf' + self._postgresql_base_conf = os.path.join(self._config_dir, self._postgresql_base_conf_name) + self._pg_hba_conf = os.path.join(self._config_dir, 'pg_hba.conf') + self._pg_ident_conf = os.path.join(self._config_dir, 'pg_ident.conf') + self._recovery_conf = os.path.join(postgresql.data_dir, 'recovery.conf') + self._synchronous_standby_names = None + self._config = {} + self.reload_config(config) + + def setup_server_parameters(self): + self._server_parameters = self.get_server_parameters(self._config) + + @property + def _configuration_to_save(self): + configuration = [os.path.basename(self._postgresql_conf)] + if 'custom_conf' not in self._config: + configuration.append(os.path.basename(self._postgresql_base_conf_name)) + if not self.hba_file: + configuration.append('pg_hba.conf') + if not self._server_parameters.get('ident_file'): + configuration.append('pg_ident.conf') + return configuration + + def save_configuration_files(self, check_custom_bootstrap=False): + """ + copy postgresql.conf to postgresql.conf.backup to be able to retrive configuration files + - originally stored as symlinks, those are normally skipped by pg_basebackup + - in case of WAL-E basebackup (see http://comments.gmane.org/gmane.comp.db.postgresql.wal-e/239) + """ + if not (check_custom_bootstrap and self._postgresql.bootstrap.running_custom_bootstrap): + try: + for f in self._configuration_to_save: + config_file = os.path.join(self._config_dir, f) + backup_file = os.path.join(self._postgresql.data_dir, f + '.backup') + if os.path.isfile(config_file): + shutil.copy(config_file, backup_file) + except IOError: + logger.exception('unable to create backup copies of configuration files') + return True + + def restore_configuration_files(self): + """ restore a previously saved postgresql.conf """ + try: + for f in self._configuration_to_save: + config_file = os.path.join(self._config_dir, f) + backup_file = os.path.join(self._postgresql.data_dir, f + '.backup') + if not os.path.isfile(config_file): + if os.path.isfile(backup_file): + shutil.copy(backup_file, config_file) + # Previously we didn't backup pg_ident.conf, if file is missing just create empty + elif f == 'pg_ident.conf': + open(config_file, 'w').close() + except IOError: + logger.exception('unable to restore configuration files from backup') + + def write_postgresql_conf(self, configuration=None): + # rename the original configuration if it is necessary + if 'custom_conf' not in self._config and not os.path.exists(self._postgresql_base_conf): + os.rename(self._postgresql_conf, self._postgresql_base_conf) + + with open(self._postgresql_conf, 'w') as f: + f.write(self._CONFIG_WARNING_HEADER) + f.write("include '{0}'\n\n".format(self._config.get('custom_conf') or self._postgresql_base_conf_name)) + for name, value in sorted((configuration or self._server_parameters).items()): + if not self._postgresql.bootstrap.running_custom_bootstrap or name != 'hba_file': + f.write("{0} = '{1}'\n".format(name, value)) + # when we are doing custom bootstrap we assume that we don't know superuser password + # and in order to be able to change it, we are opening trust access from a certain address + # therefore we need to make sure that hba_file is not overriden + # after changing superuser password we will "revert" all these "changes" + if self._postgresql.bootstrap.running_custom_bootstrap or 'hba_file' not in self._server_parameters: + f.write("hba_file = '{0}'\n".format(self._pg_hba_conf.replace('\\', '\\\\'))) + if 'ident_file' not in self._server_parameters: + f.write("ident_file = '{0}'\n".format(self._pg_ident_conf.replace('\\', '\\\\'))) + + def append_pg_hba(self, config): + if not self.hba_file and not self._config.get('pg_hba'): + with open(self._pg_hba_conf, 'a') as f: + f.write('\n{}\n'.format('\n'.join(config))) + return True + + def replace_pg_hba(self): + """ + Replace pg_hba.conf content in the PGDATA if hba_file is not defined in the + `postgresql.parameters` and pg_hba is defined in `postgresql` configuration section. + + :returns: True if pg_hba.conf was rewritten. + """ + + # when we are doing custom bootstrap we assume that we don't know superuser password + # and in order to be able to change it, we are opening trust access from a certain address + if self._postgresql.bootstrap.running_custom_bootstrap: + addresses = {'': 'local'} + if 'host' in self._local_address and not self._local_address['host'].startswith('/'): + for _, _, _, _, sa in socket.getaddrinfo(self._local_address['host'], self._local_address['port'], + 0, socket.SOCK_STREAM, socket.IPPROTO_TCP): + addresses[sa[0] + '/32'] = 'host' + + with open(self._pg_hba_conf, 'w') as f: + f.write(self._CONFIG_WARNING_HEADER) + for address, t in addresses.items(): + f.write(( + '{0}\treplication\t{1}\t{3}\ttrust\n' + '{0}\tall\t{2}\t{3}\ttrust\n' + ).format(t, self.replication['username'], self._superuser.get('username') or 'all', address)) + elif not self.hba_file and self._config.get('pg_hba'): + with open(self._pg_hba_conf, 'w') as f: + f.write(self._CONFIG_WARNING_HEADER) + for line in self._config['pg_hba']: + f.write('{0}\n'.format(line)) + return True + + def replace_pg_ident(self): + """ + Replace pg_ident.conf content in the PGDATA if ident_file is not defined in the + `postgresql.parameters` and pg_ident is defined in the `postgresql` section. + + :returns: True if pg_ident.conf was rewritten. + """ + + if not self._server_parameters.get('ident_file') and self._config.get('pg_ident'): + with open(self._pg_ident_conf, 'w') as f: + f.write(self._CONFIG_WARNING_HEADER) + for line in self._config['pg_ident']: + f.write('{0}\n'.format(line)) + return True + + def primary_conninfo(self, member): + name = self._postgresql.name + if not (member and member.conn_url) or member.name == name: + return None + r = member.conn_kwargs(self.replication) + r.update(application_name=name, sslmode='prefer', sslcompression='1', krbsrvname=self._krbsrvname) + keywords = 'user password host port sslmode sslcompression application_name krbsrvname'.split() + return ' '.join('{0}={{{0}}}'.format(kw) for kw in keywords if r.get(kw)).format(**r) + + def recovery_conf_exists(self): + return os.path.exists(self._recovery_conf) + + def check_recovery_conf(self, member): + # TODO: recovery.conf could be stale, would be nice to detect that. + primary_conninfo = self.primary_conninfo(member) + + if not self.recovery_conf_exists(): + return False + + with open(self._recovery_conf, 'r') as f: + for line in f: + if line.startswith('primary_conninfo'): + return primary_conninfo and (primary_conninfo in line) + return not primary_conninfo + + def write_recovery_conf(self, recovery_params): + with open(self._recovery_conf, 'w') as f: + os.chmod(self._recovery_conf, stat.S_IWRITE | stat.S_IREAD) + for name, value in recovery_params.items(): + f.write("{0} = '{1}'\n".format(name, value)) + + def remove_recovery_conf(self): + if os.path.isfile(self._recovery_conf) or os.path.islink(self._recovery_conf): + os.unlink(self._recovery_conf) + + def get_server_parameters(self, config): + parameters = config['parameters'].copy() + listen_addresses, port = split_host_port(config['listen'], 5432) + parameters.update(cluster_name=self._postgresql.scope, listen_addresses=listen_addresses, port=str(port)) + if config.get('synchronous_mode', False): + if self._synchronous_standby_names is None: + if config.get('synchronous_mode_strict', False): + parameters['synchronous_standby_names'] = '*' + else: + parameters.pop('synchronous_standby_names', None) + else: + parameters['synchronous_standby_names'] = self._synchronous_standby_names + if self._postgresql.major_version >= 90600 and parameters['wal_level'] == 'hot_standby': + parameters['wal_level'] = 'replica' + ret = CaseInsensitiveDict({k: v for k, v in parameters.items() if not self._postgresql.major_version or + self._postgresql.major_version >= self.CMDLINE_OPTIONS.get(k, (0, 1, 90100))[2]}) + ret.update({k: os.path.join(self._config_dir, ret[k]) for k in ('hba_file', 'ident_file') if k in ret}) + return ret + + @staticmethod + def _get_unix_local_address(unix_socket_directories): + for d in unix_socket_directories.split(','): + d = d.strip() + if d.startswith('/'): # Only absolute path can be used to connect via unix-socket + return d + return '' + + def _get_tcp_local_address(self): + listen_addresses = self._server_parameters['listen_addresses'].split(',') + + for la in listen_addresses: + if la.strip().lower() in ('*', '0.0.0.0', '127.0.0.1', 'localhost'): # we are listening on '*' or localhost + return 'localhost' # connection via localhost is preferred + return listen_addresses[0].strip() # can't use localhost, take first address from listen_addresses + + @property + def local_connect_kwargs(self): + ret = self._local_address.copy() + ret.update({'database': self._postgresql.database, + 'fallback_application_name': 'Patroni', + 'connect_timeout': 3, + 'options': '-c statement_timeout=2000'}) + if 'username' in self._superuser: + ret['user'] = self._superuser['username'] + if 'password' in self._superuser: + ret['password'] = self._superuser['password'] + return ret + + def resolve_connection_addresses(self): + port = self._server_parameters['port'] + tcp_local_address = self._get_tcp_local_address() + + local_address = {'port': port} + if self._config.get('use_unix_socket'): + unix_socket_directories = self._server_parameters.get('unix_socket_directories') + if unix_socket_directories is not None: + # fallback to tcp if unix_socket_directories is set, but there are no sutable values + local_address['host'] = self._get_unix_local_address(unix_socket_directories) or tcp_local_address + + # if unix_socket_directories is not specified, but use_unix_socket is set to true - do our best + # to use default value, i.e. don't specify a host neither in connection url nor arguments + else: + local_address['host'] = tcp_local_address + + self._local_address = local_address + self.local_replication_address = {'host': tcp_local_address, 'port': port} + + self._postgresql.connection_string = 'postgres://{0}/{1}'.format( + self._config.get('connect_address') or tcp_local_address + ':' + port, self._postgresql.database) + + self._postgresql.set_connection_kwargs(self.local_connect_kwargs) + + def reload_config(self, config): + self._superuser = config['authentication'].get('superuser', {}) + server_parameters = self.get_server_parameters(config) + + conf_changed = hba_changed = ident_changed = local_connection_address_changed = pending_restart = False + if self._postgresql.state == 'running': + changes = CaseInsensitiveDict({p: v for p, v in server_parameters.items() if '.' not in p}) + changes.update({p: None for p in self._server_parameters.keys() if not ('.' in p or p in changes)}) + if changes: + if 'wal_segment_size' not in changes: + changes['wal_segment_size'] = '16384kB' + # XXX: query can raise an exception + for r in self._postgresql.query(('SELECT name, setting, unit, vartype, context ' + + 'FROM pg_catalog.pg_settings ' + + ' WHERE pg_catalog.lower(name) IN (' + + ', '.join(['%s'] * len(changes)) + + ') ORDER BY 1 DESC'), *(k.lower() for k in changes.keys())): + if r[4] == 'internal': + if r[0] == 'wal_segment_size': + server_parameters.pop(r[0], None) + wal_segment_size = parse_int(r[2], 'kB') + if wal_segment_size is not None: + changes['wal_segment_size'] = '{0}kB'.format(int(r[1]) * wal_segment_size) + elif r[0] in changes: + unit = changes['wal_segment_size'] if r[0] in ('min_wal_size', 'max_wal_size') else r[2] + new_value = changes.pop(r[0]) + if new_value is None or not compare_values(r[3], unit, r[1], new_value): + if r[4] == 'postmaster': + pending_restart = True + logger.info('Changed %s from %s to %s (restart required)', r[0], r[1], new_value) + if config.get('use_unix_socket') and r[0] == 'unix_socket_directories'\ + or r[0] in ('listen_addresses', 'port'): + local_connection_address_changed = True + else: + logger.info('Changed %s from %s to %s', r[0], r[1], new_value) + conf_changed = True + for param in changes: + if param in server_parameters: + logger.warning('Removing invalid parameter `%s` from postgresql.parameters', param) + server_parameters.pop(param) + + # Check that user-defined-paramters have changed (parameters with period in name) + if not conf_changed: + for p, v in server_parameters.items(): + if '.' in p and (p not in self._server_parameters or str(v) != str(self._server_parameters[p])): + logger.info('Changed %s from %s to %s', p, self._server_parameters.get(p), v) + conf_changed = True + break + if not conf_changed: + for p, v in self._server_parameters.items(): + if '.' in p and (p not in server_parameters or str(v) != str(server_parameters[p])): + logger.info('Changed %s from %s to %s', p, v, server_parameters.get(p)) + conf_changed = True + break + + if not server_parameters.get('hba_file') and config.get('pg_hba'): + hba_changed = self._config.get('pg_hba', []) != config['pg_hba'] + + if not server_parameters.get('ident_file') and config.get('pg_ident'): + ident_changed = self._config.get('pg_ident', []) != config['pg_ident'] + + self._config = config + self._postgresql.set_pending_restart(pending_restart) + self._server_parameters = server_parameters + self._connect_address = config.get('connect_address') + self._krbsrvname = config.get('krbsrvname') + + # for not so obvious connection attempts that may happen outside of pyscopg2 + if self._krbsrvname: + os.environ['PGKRBSRVNAME'] = self._krbsrvname + + if not local_connection_address_changed: + self.resolve_connection_addresses() + + if conf_changed: + self.write_postgresql_conf() + + if hba_changed: + self.replace_pg_hba() + + if ident_changed: + self.replace_pg_ident() + + if conf_changed or hba_changed or ident_changed: + logger.info('PostgreSQL configuration items changed, reloading configuration.') + self._postgresql.reload() + elif not pending_restart: + logger.info('No PostgreSQL configuration items changed, nothing to reload.') + + def set_synchronous_standby(self, name): + """Sets a node to be synchronous standby and if changed does a reload for PostgreSQL.""" + if name and name != '*': + name = quote_ident(name) + if name != self._synchronous_standby_names: + if name is None: + self._server_parameters.pop('synchronous_standby_names', None) + else: + self._server_parameters['synchronous_standby_names'] = name + self._synchronous_standby_names = name + if self._postgresql.state == 'running': + self.write_postgresql_conf() + self._postgresql.reload() + + @property + def effective_configuration(self): + """It might happen that the current value of one (or more) below parameters stored in + the controldata is higher than the value stored in the global cluster configuration. + + Example: max_connections in global configuration is 100, but in controldata + `Current max_connections setting: 200`. If we try to start postgres with + max_connections=100, it will immediately exit. + As a workaround we will start it with the values from controldata and set `pending_restart` + to true as an indicator that current values of parameters are not matching expectations.""" + + if self._postgresql.role == 'master': + return self._server_parameters + + options_mapping = { + 'max_connections': 'max_connections setting', + 'max_prepared_transactions': 'max_prepared_xacts setting', + 'max_locks_per_transaction': 'max_locks_per_xact setting' + } + + if self._postgresql.major_version >= 90400: + options_mapping['max_worker_processes'] = 'max_worker_processes setting' + + data = self._postgresql.controldata() + effective_configuration = self._server_parameters.copy() + + for name, cname in options_mapping.items(): + value = parse_int(effective_configuration[name]) + cvalue = parse_int(data[cname]) + if cvalue > value: + effective_configuration[name] = cvalue + self._postgresql.set_pending_restart(True) + return effective_configuration + + @property + def replication(self): + return self._config['authentication']['replication'] + + @property + def superuser(self): + return self._superuser + + @property + def rewind_credentials(self): + return self._config['authentication'].get('rewind', self._superuser) \ + if self._postgresql.major_version >= 110000 else self._superuser + + @property + def hba_file(self): + return self._server_parameters.get('hba_file') + + @property + def pg_hba_conf(self): + return self._pg_hba_conf + + @property + def postgresql_conf(self): + return self._postgresql_conf + + def get(self, key, default=None): + return self._config.get(key, default) diff --git a/patroni/postgresql/connection.py b/patroni/postgresql/connection.py new file mode 100644 index 00000000..933ee89c --- /dev/null +++ b/patroni/postgresql/connection.py @@ -0,0 +1,46 @@ +import logging +import psycopg2 + +from contextlib import contextmanager +from threading import Lock + +logger = logging.getLogger(__name__) + + +class Connection(object): + + def __init__(self): + self._lock = Lock() + self._connection = None + self._cursor_holder = None + + def set_conn_kwargs(self, conn_kwargs): + self._conn_kwargs = conn_kwargs + + def get(self): + with self._lock: + if not self._connection or self._connection.closed != 0: + self._connection = psycopg2.connect(**self._conn_kwargs) + self._connection.autocommit = True + self.server_version = self._connection.server_version + return self._connection + + def cursor(self): + if not self._cursor_holder or self._cursor_holder.closed or self._cursor_holder.connection.closed != 0: + logger.info("establishing a new patroni connection to the postgres cluster") + self._cursor_holder = self.get().cursor() + return self._cursor_holder + + def close(self): + if self._connection and self._connection.closed == 0: + self._connection.close() + logger.info("closed patroni connection to the postgresql cluster") + self._cursor_holder = self._connection = None + + +@contextmanager +def get_connection_cursor(**kwargs): + with psycopg2.connect(**kwargs) as conn: + conn.autocommit = True + with conn.cursor() as cur: + yield cur diff --git a/patroni/postgresql/misc.py b/patroni/postgresql/misc.py new file mode 100644 index 00000000..f5caffbd --- /dev/null +++ b/patroni/postgresql/misc.py @@ -0,0 +1,70 @@ +import logging + +from patroni.exceptions import PostgresException + +logger = logging.getLogger(__name__) + + +def postgres_version_to_int(pg_version): + """Convert the server_version to integer + + >>> postgres_version_to_int('9.5.3') + 90503 + >>> postgres_version_to_int('9.3.13') + 90313 + >>> postgres_version_to_int('10.1') + 100001 + >>> postgres_version_to_int('10') # doctest: +IGNORE_EXCEPTION_DETAIL + Traceback (most recent call last): + ... + PostgresException: 'Invalid PostgreSQL version format: X.Y or X.Y.Z is accepted: 10' + >>> postgres_version_to_int('9.6') # doctest: +IGNORE_EXCEPTION_DETAIL + Traceback (most recent call last): + ... + PostgresException: 'Invalid PostgreSQL version format: X.Y or X.Y.Z is accepted: 9.6' + >>> postgres_version_to_int('a.b.c') # doctest: +IGNORE_EXCEPTION_DETAIL + Traceback (most recent call last): + ... + PostgresException: 'Invalid PostgreSQL version: a.b.c' + """ + + try: + components = list(map(int, pg_version.split('.'))) + except ValueError: + raise PostgresException('Invalid PostgreSQL version: {0}'.format(pg_version)) + + if len(components) < 2 or len(components) == 2 and components[0] < 10 or len(components) > 3: + raise PostgresException('Invalid PostgreSQL version format: X.Y or X.Y.Z is accepted: {0}'.format(pg_version)) + + if len(components) == 2: + # new style verion numbers, i.e. 10.1 becomes 100001 + components.insert(1, 0) + + return int(''.join('{0:02d}'.format(c) for c in components)) + + +def postgres_major_version_to_int(pg_version): + """ + >>> postgres_major_version_to_int('10') + 100000 + >>> postgres_major_version_to_int('9.6') + 90600 + """ + return postgres_version_to_int(pg_version + '.0') + + +def parse_lsn(lsn): + t = lsn.split('/') + return int(t[0], 16) * 0x100000000 + int(t[1], 16) + + +def parse_history(data): + for line in data.split('\n'): + values = line.strip().split('\t') + if len(values) == 3: + try: + values[0] = int(values[0]) + values[1] = parse_lsn(values[1]) + yield values + except (IndexError, ValueError): + logger.exception('Exception when parsing timeline history line "%s"', values) diff --git a/patroni/postmaster.py b/patroni/postgresql/postmaster.py similarity index 100% rename from patroni/postmaster.py rename to patroni/postgresql/postmaster.py diff --git a/patroni/postgresql/rewind.py b/patroni/postgresql/rewind.py new file mode 100644 index 00000000..d667235a --- /dev/null +++ b/patroni/postgresql/rewind.py @@ -0,0 +1,216 @@ +import logging +import os +import subprocess + +from patroni.dcs import Leader +from patroni.postgresql.connection import get_connection_cursor +from patroni.postgresql.misc import parse_history, parse_lsn + +logger = logging.getLogger(__name__) + +REWIND_STATUS = type('Enum', (), {'INITIAL': 0, 'CHECKPOINT': 1, 'CHECK': 1, 'NEED': 2, + 'NOT_NEED': 3, 'SUCCESS': 4, 'FAILED': 5}) + + +class Rewind(object): + + def __init__(self, postgresql): + self._postgresql = postgresql + self.reset_state() + + @staticmethod + def configuration_allows_rewind(data): + return data.get('wal_log_hints setting', 'off') == 'on' or data.get('Data page checksum version', '0') != '0' + + @property + def can_rewind(self): + """ check if pg_rewind executable is there and that pg_controldata indicates + we have either wal_log_hints or checksums turned on + """ + # low-hanging fruit: check if pg_rewind configuration is there + if not self._postgresql.config.get('use_pg_rewind'): + return False + + cmd = [self._postgresql.pgcommand('pg_rewind'), '--help'] + try: + ret = subprocess.call(cmd, stdout=open(os.devnull, 'w'), stderr=subprocess.STDOUT) + if ret != 0: # pg_rewind is not there, close up the shop and go home + return False + except OSError: + return False + return self.configuration_allows_rewind(self._postgresql.controldata()) + + @property + def can_rewind_or_reinitialize_allowed(self): + return self._postgresql.config.get('remove_data_directory_on_diverged_timelines') or self.can_rewind + + def trigger_check_diverged_lsn(self): + if self.can_rewind_or_reinitialize_allowed and self._state != REWIND_STATUS.NEED: + self._state = REWIND_STATUS.CHECK + + def check_leader_is_not_in_recovery(self, **kwargs): + if not kwargs.get('database'): + kwargs['database'] = self._postgresql.database + try: + with get_connection_cursor(connect_timeout=3, options='-c statement_timeout=2000', **kwargs) as cur: + cur.execute('SELECT pg_catalog.pg_is_in_recovery()') + if not cur.fetchone()[0]: + return True + logger.info('Leader is still in_recovery and therefore can\'t be used for rewind') + except Exception: + return logger.exception('Exception when working with leader') + + def _get_local_timeline_lsn_from_controldata(self): + timeline = lsn = None + data = self._postgresql.controldata() + try: + if data.get('Database cluster state') == 'shut down in recovery': + lsn = data.get('Minimum recovery ending location') + timeline = int(data.get("Min recovery ending loc's timeline")) + if lsn == '0/0' or timeline == 0: # it was a master when it crashed + data['Database cluster state'] = 'shut down' + if data.get('Database cluster state') == 'shut down': + lsn = data.get('Latest checkpoint location') + timeline = int(data.get("Latest checkpoint's TimeLineID")) + except (TypeError, ValueError): + logger.exception('Failed to get local timeline and lsn from pg_controldata output') + return timeline, lsn + + def _get_local_timeline_lsn(self): + if self._postgresql.is_running(): # if postgres is running - get timeline and lsn from replication connection + timeline, lsn = self._postgresql.get_local_timeline_lsn_from_replication_connection() + else: # otherwise analyze pg_controldata output + timeline, lsn = self._get_local_timeline_lsn_from_controldata() + logger.info('Local timeline=%s lsn=%s', timeline, lsn) + return timeline, lsn + + def _check_timeline_and_lsn(self, leader): + local_timeline, local_lsn = self._get_local_timeline_lsn() + if local_timeline is None or local_lsn is None: + return + + if isinstance(leader, Leader): + if leader.member.data.get('role') != 'master': + return + # standby cluster + elif not self.check_leader_is_not_in_recovery(**leader.conn_kwargs(self._postgresql.config.replication)): + return + + history = need_rewind = None + try: + with self._postgresql.get_replication_connection_cursor(**leader.conn_kwargs()) as cur: + cur.execute('IDENTIFY_SYSTEM') + master_timeline = cur.fetchone()[1] + logger.info('master_timeline=%s', master_timeline) + if local_timeline > master_timeline: # Not always supported by pg_rewind + need_rewind = True + elif master_timeline > 1: + cur.execute('TIMELINE_HISTORY %s', (master_timeline,)) + history = bytes(cur.fetchone()[1]).decode('utf-8') + logger.info('master: history=%s', history) + else: # local_timeline == master_timeline == 1 + need_rewind = False + except Exception: + return logger.exception('Exception when working with master via replication connection') + + if history is not None: + for parent_timeline, switchpoint, _ in parse_history(history): + if parent_timeline == local_timeline: + try: + need_rewind = parse_lsn(local_lsn) >= switchpoint + except (IndexError, ValueError): + logger.exception('Exception when parsing lsn') + break + elif parent_timeline > local_timeline: + break + + self._state = need_rewind and REWIND_STATUS.NEED or REWIND_STATUS.NOT_NEED + + def rewind_or_reinitialize_needed_and_possible(self, leader): + if leader and leader.name != self._postgresql.name and leader.conn_url and self._state == REWIND_STATUS.CHECK: + self._check_timeline_and_lsn(leader) + return leader and leader.conn_url and self._state == REWIND_STATUS.NEED + + def check_for_checkpoint_after_promote(self): + if self._state == REWIND_STATUS.INITIAL and self._postgresql.is_leader(): + try: + timeline = int(self._postgresql.controldata().get("Latest checkpoint's TimeLineID")) + if self._postgresql.get_master_timeline() == timeline: + self._state = REWIND_STATUS.CHECKPOINT + except (TypeError, ValueError): + logger.exception('Failed to parse timeline from pg_controldata output') + + def checkpoint_after_promote(self): + return self._state == REWIND_STATUS.CHECKPOINT + + def pg_rewind(self, r): + # prepare pg_rewind connection + env = self._postgresql.write_pgpass(r) + dsn_attrs = [ + ('user', r.get('user')), + ('host', r.get('host')), + ('port', r.get('port')), + ('dbname', r.get('database') or self._postgresql.database), + ('sslmode', 'prefer'), + ('sslcompression', '1'), + ] + dsn = " ".join("{0}={1}".format(k, v) for k, v in dsn_attrs if v is not None) + logger.info('running pg_rewind from %s', dsn) + try: + return self._postgresql.cancellable.call([self._postgresql.pgcommand('pg_rewind'), '-D', + self._postgresql.data_dir, '--source-server', dsn], env=env) == 0 + except OSError: + return False + + def execute(self, leader): + if self._postgresql.is_running() and not self._postgresql.stop(checkpoint=False): + return logger.warning('Can not run pg_rewind because postgres is still running') + + # prepare pg_rewind connection + r = leader.conn_kwargs(self._postgresql.config.rewind_credentials) + + # 1. make sure that we are really trying to rewind from the master + # 2. make sure that pg_control contains the new timeline by: + # running a checkpoint or + # waiting until Patroni on the master will expose checkpoint_after_promote=True + checkpoint_status = leader.checkpoint_after_promote if isinstance(leader, Leader) else None + if checkpoint_status is None: # master still runs the old Patroni + leader_status = self._postgresql.checkpoint(leader.conn_kwargs(self._postgresql.config.superuser)) + if leader_status: + return logger.warning('Can not use %s for rewind: %s', leader.name, leader_status) + elif not checkpoint_status: + return logger.info('Waiting for checkpoint on %s before rewind', leader.name) + elif not self.check_leader_is_not_in_recovery(**r): + return + + if self.pg_rewind(r): + self._state = REWIND_STATUS.SUCCESS + elif not self.check_leader_is_not_in_recovery(**r): + logger.warning('Failed to rewind because master %s become unreachable', leader.name) + else: + logger.error('Failed to rewind from healty master: %s', leader.name) + + for name in ('remove_data_directory_on_rewind_failure', 'remove_data_directory_on_diverged_timelines'): + if self._postgresql.config.get(name): + logger.warning('%s is set. removing...', name) + self._postgresql.remove_data_directory() + self._state = REWIND_STATUS.INITIAL + break + else: + self._state = REWIND_STATUS.FAILED + return False + + def reset_state(self): + self._state = REWIND_STATUS.INITIAL + + @property + def is_needed(self): + return self._state in (REWIND_STATUS.CHECK, REWIND_STATUS.NEED) + + @property + def executed(self): + return self._state > REWIND_STATUS.NOT_NEED + + @property + def failed(self): + return self._state == REWIND_STATUS.FAILED diff --git a/patroni/postgresql/slots.py b/patroni/postgresql/slots.py new file mode 100644 index 00000000..91d6b743 --- /dev/null +++ b/patroni/postgresql/slots.py @@ -0,0 +1,108 @@ +import logging + +from patroni.postgresql.connection import get_connection_cursor +from collections import defaultdict + +logger = logging.getLogger(__name__) + + +def compare_slots(s1, s2): + return s1['type'] == s2['type'] and (s1['type'] == 'physical' or + s1['database'] == s2['database'] and s1['plugin'] == s2['plugin']) + + +class SlotsHandler(object): + + def __init__(self, postgresql): + self._postgresql = postgresql + self._use_slots = postgresql.config.get('use_slots', True) + self._replication_slots = {} # already existing replication slots + self.schedule() + + @property + def use_slots(self): + return self._use_slots and self._postgresql.major_version >= 90400 + + def _query(self, sql, *params): + return self._postgresql.query(sql, *params, retry=False) + + def load_replication_slots(self): + if self.use_slots and self._schedule_load_slots: + replication_slots = {} + cursor = self._query('SELECT slot_name, slot_type, plugin, database FROM pg_catalog.pg_replication_slots') + for r in cursor: + value = {'type': r[1]} + if r[1] == 'logical': + value.update({'plugin': r[2], 'database': r[3]}) + replication_slots[r[0]] = value + self._replication_slots = replication_slots + self._schedule_load_slots = False + + def drop_replication_slot(self, name): + cursor = self._query(('SELECT pg_catalog.pg_drop_replication_slot(%s) WHERE EXISTS (SELECT 1 ' + + 'FROM pg_catalog.pg_replication_slots WHERE slot_name = %s AND NOT active)'), name, name) + # In normal situation rowcount should be 1, otherwise either slot doesn't exists or it is still active + return cursor.rowcount == 1 + + def sync_replication_slots(self, cluster): + if self.use_slots: + try: + self.load_replication_slots() + + slots = cluster.get_replication_slots(self._postgresql.name, self._postgresql.role) + + # drop old replication slots which are not presented in desired slots + for name in set(self._replication_slots) - set(slots): + if not self.drop_replication_slot(name): + logger.error("Failed to drop replication slot '%s'", name) + self._schedule_load_slots = True + + immediately_reserve = ', true' if self._postgresql.major_version >= 90600 else '' + + logical_slots = defaultdict(dict) + for name, value in slots.items(): + if name in self._replication_slots and not compare_slots(value, self._replication_slots[name]): + logger.info("Trying to drop replication slot '%s' because value is changing from %s to %s", + name, self._replication_slots[name], value) + if not self.drop_replication_slot(name): + logger.error("Failed to drop replication slot '%s'", name) + self._schedule_load_slots = True + continue + self._replication_slots.pop(name) + if name not in self._replication_slots: + if value['type'] == 'physical': + try: + self._query(("SELECT pg_catalog.pg_create_physical_replication_slot(%s{0})" + + " WHERE NOT EXISTS (SELECT 1 FROM pg_catalog.pg_replication_slots" + + " WHERE slot_type = 'physical' AND slot_name = %s)").format( + immediately_reserve), name, name) + except Exception: + logger.exception("Failed to create physical replication slot '%s'", name) + self._schedule_load_slots = True + elif value['type'] == 'logical' and name not in self._replication_slots: + logical_slots[value['database']][name] = value + + # create new logical slots + for database, values in logical_slots.items(): + conn_kwargs = self._postgresql.config.local_connect_kwargs + conn_kwargs['database'] = database + with get_connection_cursor(**conn_kwargs) as cur: + for name, value in values.items(): + try: + cur.execute("SELECT pg_catalog.pg_create_logical_replication_slot(%s, %s)" + + " WHERE NOT EXISTS (SELECT 1 FROM pg_catalog.pg_replication_slots" + + " WHERE slot_type = 'logical' AND slot_name = %s)", + (name, value['plugin'], name)) + except Exception: + logger.exception("Failed to create logical replication slot '%s' plugin='%s'", + name, value['plugin']) + self._schedule_load_slots = True + self._replication_slots = slots + except Exception: + logger.exception('Exception when changing replication slots') + self._schedule_load_slots = True + + def schedule(self, value=None): + if value is None: + value = self.use_slots + self._schedule_load_slots = value diff --git a/tests/__init__.py b/tests/__init__.py new file mode 100644 index 00000000..7d6f4be0 --- /dev/null +++ b/tests/__init__.py @@ -0,0 +1,216 @@ +import datetime +import json +import os +import shutil +import unittest + +from mock import Mock, patch +from tempfile import gettempdir + +import psycopg2 +import requests + +from patroni.dcs import Leader, Member +from patroni.postgresql import Postgresql +from patroni.postgresql.config import ConfigHandler +from patroni.utils import RetryFailedError + + +class SleepException(Exception): + pass + + +class MockResponse(object): + + def __init__(self, status_code=200): + self.status_code = status_code + self.content = '{}' + self.ok = True + + def json(self): + return json.loads(self.content) + + @property + def data(self): + return self.content.encode('utf-8') + + @property + def text(self): + return self.content + + @property + def status(self): + return self.status_code + + @staticmethod + def getheader(*args): + return '' + + +def requests_get(url, **kwargs): + members = '[{"id":14855829450254237642,"peerURLs":["http://localhost:2380","http://localhost:7001"],' +\ + '"name":"default","clientURLs":["http://localhost:2379","http://localhost:4001"]}]' + response = MockResponse() + if url.startswith('http://local'): + raise requests.exceptions.RequestException() + elif ':8011/patroni' in url: + response.content = '{"role": "replica", "xlog": {"received_location": 0}, "tags": {}}' + elif url.endswith('/members'): + response.content = '[{}]' if url.startswith('http://error') else members + elif url.startswith('http://exhibitor'): + response.content = '{"servers":["127.0.0.1","127.0.0.2","127.0.0.3"],"port":2181}' + elif url.endswith(':8011/reinitialize'): + data = kwargs.get('data', '') + if ' false}' in data: + response.status_code = 503 + response.ok = False + response.content = 'restarting after failure already in progress' + else: + response.status_code = 404 + response.ok = False + return response + + +class MockPostmaster(object): + def __init__(self, is_running=True, is_single_master=False): + self.is_running = Mock(return_value=is_running) + self.is_single_master = Mock(return_value=is_single_master) + self.wait_for_user_backends_to_close = Mock() + self.signal_stop = Mock(return_value=None) + self.wait = Mock() + + +class MockCursor(object): + + def __init__(self, connection): + self.connection = connection + self.closed = False + self.rowcount = 0 + self.results = [] + + def execute(self, sql, *params): + if sql.startswith('blabla'): + raise psycopg2.ProgrammingError() + elif sql == 'CHECKPOINT' or sql.startswith('SELECT pg_catalog.pg_create_'): + raise psycopg2.OperationalError() + elif sql.startswith('RetryFailedError'): + raise RetryFailedError('retry') + elif sql.startswith('SELECT slot_name'): + self.results = [('blabla', 'physical'), ('foobar', 'physical'), ('ls', 'logical', 'a', 'b')] + elif sql.startswith('SELECT CASE WHEN pg_catalog.pg_is_in_recovery()'): + self.results = [(1, 2)] + elif sql.startswith('SELECT pg_catalog.pg_is_in_recovery()'): + self.results = [(False, 2)] + elif sql.startswith('WITH replication_info AS ('): + replication_info = '[{"application_name":"walreceiver","client_addr":"1.2.3.4",' +\ + '"state":"streaming","sync_state":"async","sync_priority":0}]' + self.results = [('', 0, '', '', '', '', False, replication_info)] + elif sql.startswith('SELECT name, setting'): + self.results = [('wal_segment_size', '2048', '8kB', 'integer', 'internal'), + ('search_path', 'public', None, 'string', 'user'), + ('port', '5433', None, 'integer', 'postmaster'), + ('listen_addresses', '*', None, 'string', 'postmaster'), + ('autovacuum', 'on', None, 'bool', 'sighup'), + ('unix_socket_directories', '/tmp', None, 'string', 'postmaster')] + elif sql.startswith('IDENTIFY_SYSTEM'): + self.results = [('1', 2, '0/402EEC0', '')] + elif sql.startswith('SELECT isdir, modification'): + self.results = [(False, datetime.datetime.now())] + elif sql.startswith('SELECT pg_catalog.pg_read_file'): + self.results = [('1\t0/40159C0\tno recovery target specified\n\n' + '2\t1/40159C0\tno recovery target specified\n',)] + elif sql.startswith('TIMELINE_HISTORY '): + self.results = [('', b'x\t0/40159C0\tno recovery target specified\n\n' + b'1\t0/40159C0\tno recovery target specified\n\n' + b'2\t0/402DD98\tno recovery target specified\n\n' + b'3\t0/403DD98\tno recovery target specified\n')] + else: + self.results = [(None, None, None, None, None, None, None, None, None, None)] + + def fetchone(self): + return self.results[0] + + def fetchall(self): + return self.results + + def __iter__(self): + for i in self.results: + yield i + + def __enter__(self): + return self + + def __exit__(self, *args): + pass + + +class MockConnect(object): + + server_version = 99999 + autocommit = False + closed = 0 + + def cursor(self): + return MockCursor(self) + + def __enter__(self): + return self + + def __exit__(self, *args): + pass + + @staticmethod + def close(): + pass + + +def psycopg2_connect(*args, **kwargs): + return MockConnect() + + +class PostgresInit(unittest.TestCase): + _PARAMETERS = {'wal_level': 'hot_standby', 'max_replication_slots': 5, 'f.oo': 'bar', + 'search_path': 'public', 'hot_standby': 'on', 'max_wal_senders': 5, + 'wal_keep_segments': 8, 'wal_log_hints': 'on', 'max_locks_per_transaction': 64, + 'max_worker_processes': 8, 'max_connections': 100, 'max_prepared_transactions': 0, + 'track_commit_timestamp': 'off', 'unix_socket_directories': '/tmp'} + + @patch('psycopg2.connect', psycopg2_connect) + @patch.object(ConfigHandler, 'write_postgresql_conf', Mock()) + @patch.object(ConfigHandler, 'replace_pg_hba', Mock()) + @patch.object(ConfigHandler, 'replace_pg_ident', Mock()) + def setUp(self): + data_dir = 'data/test0' + self.p = Postgresql({'name': 'postgresql0', 'scope': 'batman', 'data_dir': data_dir, + 'config_dir': data_dir, 'retry_timeout': 10, + 'krbsrvname': 'postgres', 'pgpass': os.path.join(gettempdir(), 'pgpass0'), + 'listen': '127.0.0.2, 127.0.0.3:5432', 'connect_address': '127.0.0.2:5432', + 'authentication': {'superuser': {'username': 'foo', 'password': 'test'}, + 'replication': {'username': '', 'password': 'rep-pass'}}, + 'remove_data_directory_on_rewind_failure': True, + 'use_pg_rewind': True, 'pg_ctl_timeout': 'bla', + 'parameters': self._PARAMETERS, + 'recovery_conf': {'foo': 'bar'}, + 'pg_hba': ['host all all 0.0.0.0/0 md5'], + 'pg_ident': ['krb realm postgres'], + 'callbacks': {'on_start': 'true', 'on_stop': 'true', 'on_reload': 'true', + 'on_restart': 'true', 'on_role_change': 'true'}}) + + +class BaseTestPostgresql(PostgresInit): + + def setUp(self): + super(BaseTestPostgresql, self).setUp() + + if not os.path.exists(self.p.data_dir): + os.makedirs(self.p.data_dir) + + self.leadermem = Member(0, 'leader', 28, {'conn_url': 'postgres://replicator:rep-pass@127.0.0.1:5435/postgres'}) + self.leader = Leader(-1, 28, self.leadermem) + self.other = Member(0, 'test-1', 28, {'conn_url': 'postgres://replicator:rep-pass@127.0.0.1:5433/postgres', + 'tags': {'replicatefrom': 'leader'}}) + self.me = Member(0, 'test0', 28, {'conn_url': 'postgres://replicator:rep-pass@127.0.0.1:5434/postgres'}) + + def tearDown(self): + if os.path.exists(self.p.data_dir): + shutil.rmtree(self.p.data_dir) diff --git a/tests/test_api.py b/tests/test_api.py index a3bf1966..8eb5c3ce 100644 --- a/tests/test_api.py +++ b/tests/test_api.py @@ -10,7 +10,7 @@ from patroni.ha import _MemberStatus from patroni.utils import tzutc from six import BytesIO as IO from six.moves import BaseHTTPServer -from test_postgresql import psycopg2_connect, MockCursor +from . import psycopg2_connect, MockCursor future_restart_time = datetime.datetime.now(tzutc) + datetime.timedelta(days=5) @@ -60,7 +60,7 @@ class MockHa(object): return 'reinitialize' @staticmethod - def restart(): + def restart(*args, **kwargs): return (True, '') @staticmethod diff --git a/tests/test_bootstrap.py b/tests/test_bootstrap.py new file mode 100644 index 00000000..21618b16 --- /dev/null +++ b/tests/test_bootstrap.py @@ -0,0 +1,219 @@ +import os + +from mock import Mock, PropertyMock, patch + +from patroni.async_executor import CriticalTask +from patroni.postgresql import Postgresql +from patroni.postgresql.bootstrap import Bootstrap +from patroni.postgresql.cancellable import CancellableSubprocess +from patroni.postgresql.config import ConfigHandler + +from . import psycopg2_connect, BaseTestPostgresql + + +@patch('subprocess.call', Mock(return_value=0)) +@patch('psycopg2.connect', psycopg2_connect) +@patch('os.rename', Mock()) +class TestBootstrap(BaseTestPostgresql): + + def setUp(self): + super(TestBootstrap, self).setUp() + self.b = self.p.bootstrap + + @patch('time.sleep', Mock()) + @patch.object(CancellableSubprocess, 'call') + @patch.object(Postgresql, 'remove_data_directory', Mock(return_value=True)) + @patch.object(Postgresql, 'data_directory_empty', Mock(return_value=False)) + @patch.object(Bootstrap, '_post_restore', Mock(side_effect=OSError)) + def test_create_replica(self, mock_cancellable_subprocess_call): + self.p.config._config['create_replica_methods'] = ['pgBackRest'] + self.p.config._config['pgBackRest'] = {'command': 'pgBackRest', 'keep_data': True, 'no_params': True} + mock_cancellable_subprocess_call.return_value = 0 + self.assertEqual(self.b.create_replica(self.leader), 0) + + self.p.config._config['create_replica_methods'] = ['basebackup'] + self.p.config._config['basebackup'] = [{'max_rate': '100M'}, 'no-sync'] + self.assertEqual(self.b.create_replica(self.leader), 0) + + self.p.config._config['basebackup'] = [{'max_rate': '100M', 'compress': '9'}] + with patch('patroni.postgresql.bootstrap.logger.error', new_callable=Mock()) as mock_logger: + self.b.create_replica(self.leader) + mock_logger.assert_called_once() + self.assertTrue("only one key-value is allowed and value should be a string" in mock_logger.call_args[0][0], + "not matching {0}".format(mock_logger.call_args[0][0])) + + self.p.config._config['basebackup'] = [42] + with patch('patroni.postgresql.bootstrap.logger.error', new_callable=Mock()) as mock_logger: + self.b.create_replica(self.leader) + mock_logger.assert_called_once() + self.assertTrue("value should be string value or a single key-value pair" in mock_logger.call_args[0][0], + "not matching {0}".format(mock_logger.call_args[0][0])) + + self.p.config._config['basebackup'] = {"foo": "bar"} + self.assertEqual(self.b.create_replica(self.leader), 0) + + self.p.config._config['create_replica_methods'] = ['wale', 'basebackup'] + del self.p.config._config['basebackup'] + mock_cancellable_subprocess_call.return_value = 1 + self.assertEqual(self.b.create_replica(self.leader), 1) + + mock_cancellable_subprocess_call.side_effect = Exception('foo') + self.assertEqual(self.b.create_replica(self.leader), 1) + + mock_cancellable_subprocess_call.side_effect = [1, 0] + self.assertEqual(self.b.create_replica(self.leader), 0) + + mock_cancellable_subprocess_call.side_effect = [Exception(), 0] + self.assertEqual(self.b.create_replica(self.leader), 0) + + self.p.cancellable.cancel() + self.assertEqual(self.b.create_replica(self.leader), 1) + + @patch('time.sleep', Mock()) + @patch.object(CancellableSubprocess, 'call') + @patch.object(Postgresql, 'remove_data_directory', Mock(return_value=True)) + @patch.object(Bootstrap, '_post_restore', Mock(side_effect=OSError)) + def test_create_replica_old_format(self, mock_cancellable_subprocess_call): + """ The same test as before but with old 'create_replica_method' + to test backward compatibility + """ + self.p.config._config['create_replica_method'] = ['wale', 'basebackup'] + self.p.config._config['wale'] = {'command': 'foo'} + mock_cancellable_subprocess_call.return_value = 0 + self.assertEqual(self.b.create_replica(self.leader), 0) + del self.p.config._config['wale'] + self.assertEqual(self.b.create_replica(self.leader), 0) + + self.p.config._config['create_replica_method'] = ['wale'] + mock_cancellable_subprocess_call.return_value = 1 + self.assertEqual(self.b.create_replica(self.leader), 1) + + def test_basebackup(self): + self.p.cancellable.cancel() + self.b.basebackup(None, None, {'foo': 'bar'}) + + def test__initdb(self): + self.assertRaises(Exception, self.b.bootstrap, {'initdb': [{'pgdata': 'bar'}]}) + self.assertRaises(Exception, self.b.bootstrap, {'initdb': [{'foo': 'bar', 1: 2}]}) + self.assertRaises(Exception, self.b.bootstrap, {'initdb': [1]}) + self.assertRaises(Exception, self.b.bootstrap, {'initdb': 1}) + + @patch.object(CancellableSubprocess, 'call', Mock()) + @patch.object(Postgresql, 'is_running', Mock(return_value=True)) + @patch.object(Postgresql, 'data_directory_empty', Mock(return_value=False)) + def test_bootstrap(self): + with patch('subprocess.call', Mock(return_value=1)): + self.assertFalse(self.b.bootstrap({})) + + config = {'users': {'replicator': {'password': 'rep-pass', 'options': ['replication']}}} + + with patch.object(Postgresql, 'is_running', Mock(return_value=False)),\ + patch('multiprocessing.Process', Mock(side_effect=Exception)): + self.assertRaises(Exception, self.b.bootstrap, config) + with open(os.path.join(self.p.data_dir, 'pg_hba.conf')) as f: + lines = f.readlines() + self.assertTrue('host all all 0.0.0.0/0 md5\n' in lines) + + self.p.config._config.pop('pg_hba') + config.update({'post_init': '/bin/false', + 'pg_hba': ['host replication replicator 127.0.0.1/32 md5', + 'hostssl all all 0.0.0.0/0 md5', + 'host all all 0.0.0.0/0 md5']}) + self.b.bootstrap(config) + with open(os.path.join(self.p.data_dir, 'pg_hba.conf')) as f: + lines = f.readlines() + self.assertTrue('host replication replicator 127.0.0.1/32 md5\n' in lines) + + @patch.object(CancellableSubprocess, 'call') + @patch.object(Postgresql, 'get_major_version', Mock(return_value=90600)) + def test_custom_bootstrap(self, mock_cancellable_subprocess_call): + self.p.config._config.pop('pg_hba') + config = {'method': 'foo', 'foo': {'command': 'bar'}} + + mock_cancellable_subprocess_call.return_value = 1 + self.assertFalse(self.b.bootstrap(config)) + + mock_cancellable_subprocess_call.return_value = 0 + with patch('multiprocessing.Process', Mock(side_effect=Exception("42"))),\ + patch('os.path.isfile', Mock(return_value=True)),\ + patch('os.unlink', Mock()),\ + patch.object(ConfigHandler, 'save_configuration_files', Mock()),\ + patch.object(ConfigHandler, 'restore_configuration_files', Mock()),\ + patch.object(ConfigHandler, 'write_recovery_conf', Mock()): + with self.assertRaises(Exception) as e: + self.b.bootstrap(config) + self.assertEqual(str(e.exception), '42') + + config['foo']['recovery_conf'] = {'foo': 'bar'} + + with self.assertRaises(Exception) as e: + self.b.bootstrap(config) + self.assertEqual(str(e.exception), '42') + + mock_cancellable_subprocess_call.side_effect = Exception + self.assertFalse(self.b.bootstrap(config)) + + @patch('time.sleep', Mock()) + @patch('os.unlink', Mock()) + @patch('shutil.copy', Mock()) + @patch('os.path.isfile', Mock(return_value=True)) + @patch.object(Bootstrap, 'call_post_bootstrap', Mock(return_value=True)) + @patch.object(Bootstrap, '_custom_bootstrap', Mock(return_value=True)) + @patch.object(Postgresql, 'start', Mock(return_value=True)) + @patch.object(Postgresql, 'get_major_version', Mock(return_value=110000)) + def test_post_bootstrap(self): + config = {'method': 'foo', 'foo': {'command': 'bar'}} + self.b.bootstrap(config) + + task = CriticalTask() + with patch.object(Bootstrap, 'create_or_update_role', Mock(side_effect=Exception)): + self.b.post_bootstrap({}, task) + self.assertFalse(task.result) + + self.p.config._config.pop('pg_hba') + self.b.post_bootstrap({}, task) + self.assertTrue(task.result) + + self.b.bootstrap(config) + with patch.object(Postgresql, 'pending_restart', PropertyMock(return_value=True)), \ + patch.object(Postgresql, 'restart', Mock()) as mock_restart: + self.b.post_bootstrap({}, task) + mock_restart.assert_called_once() + + self.b.bootstrap(config) + self.p.set_state('stopped') + self.p.reload_config({'authentication': {'superuser': {'username': 'p', 'password': 'p'}, + 'replication': {'username': 'r', 'password': 'r'}, + 'rewind': {'username': 'rw', 'password': 'rw'}}, + 'listen': '*', 'retry_timeout': 10, 'parameters': {'wal_level': '', 'hba_file': 'foo'}}) + with patch.object(Postgresql, 'restart', Mock()) as mock_restart: + self.b.post_bootstrap({}, task) + mock_restart.assert_called_once() + + @patch.object(CancellableSubprocess, 'call') + def test_call_post_bootstrap(self, mock_cancellable_subprocess_call): + mock_cancellable_subprocess_call.return_value = 1 + self.assertFalse(self.b.call_post_bootstrap({'post_init': '/bin/false'})) + + mock_cancellable_subprocess_call.return_value = 0 + self.p.config.superuser.pop('username') + self.assertTrue(self.b.call_post_bootstrap({'post_init': '/bin/false'})) + mock_cancellable_subprocess_call.assert_called() + args, kwargs = mock_cancellable_subprocess_call.call_args + self.assertTrue('PGPASSFILE' in kwargs['env']) + self.assertEqual(args[0], ['/bin/false', 'postgres://127.0.0.2:5432/postgres']) + + mock_cancellable_subprocess_call.reset_mock() + self.p.config._local_address.pop('host') + self.assertTrue(self.b.call_post_bootstrap({'post_init': '/bin/false'})) + mock_cancellable_subprocess_call.assert_called() + self.assertEqual(mock_cancellable_subprocess_call.call_args[0][0], ['/bin/false', 'postgres://:5432/postgres']) + + mock_cancellable_subprocess_call.side_effect = OSError + self.assertFalse(self.b.call_post_bootstrap({'post_init': '/bin/false'})) + + @patch('os.path.exists', Mock(return_value=True)) + @patch('os.unlink', Mock()) + @patch.object(Bootstrap, 'create_replica', Mock(return_value=0)) + def test_clone(self): + self.b.clone(self.leader) diff --git a/tests/test_callback_executor.py b/tests/test_callback_executor.py index 1b9a01a5..92f51782 100644 --- a/tests/test_callback_executor.py +++ b/tests/test_callback_executor.py @@ -1,7 +1,7 @@ import unittest from mock import Mock, patch -from patroni.callback_executor import CallbackExecutor +from patroni.postgresql.callback_executor import CallbackExecutor class TestCallbackExecutor(unittest.TestCase): diff --git a/tests/test_cancellable.py b/tests/test_cancellable.py new file mode 100644 index 00000000..33b72e2a --- /dev/null +++ b/tests/test_cancellable.py @@ -0,0 +1,23 @@ +import unittest + +from mock import Mock, PropertyMock, patch +from patroni.exceptions import PostgresException +from patroni.postgresql.cancellable import CancellableSubprocess + + +class TestCancellableSubprocess(unittest.TestCase): + + def setUp(self): + self.c = CancellableSubprocess() + + def test_call(self): + self.c.cancel() + self.assertRaises(PostgresException, self.c.call, communicate_input=None) + + @patch('patroni.postgresql.cancellable.polling_loop', Mock(return_value=[0, 0])) + def test_cancel(self): + self.c._process = Mock() + self.c._process.returncode = None + self.c.cancel() + type(self.c._process).returncode = PropertyMock(side_effect=[None, -15]) + self.c.cancel() diff --git a/tests/test_consul.py b/tests/test_consul.py index 15d1c707..c6571214 100644 --- a/tests/test_consul.py +++ b/tests/test_consul.py @@ -5,7 +5,7 @@ from consul import ConsulException, NotFound from mock import Mock, patch from patroni.dcs.consul import AbstractDCS, Cluster, Consul, ConsulInternalError, \ ConsulError, HTTPClient, InvalidSessionTTL, InvalidSession -from test_etcd import SleepException +from . import SleepException def kv_get(self, key, **kwargs): diff --git a/tests/test_ctl.py b/tests/test_ctl.py index 52681ffc..e4e47e94 100644 --- a/tests/test_ctl.py +++ b/tests/test_ctl.py @@ -13,10 +13,11 @@ from patroni.ctl import ctl, store_config, load_config, output_members, request_ from patroni.dcs.etcd import Client, Failover from patroni.utils import tzutc from psycopg2 import OperationalError -from test_etcd import etcd_read, requests_get, socket_getaddrinfo, MockResponse -from test_ha import get_cluster_initialized_without_leader, get_cluster_initialized_with_leader, \ + +from . import MockConnect, MockCursor, MockResponse, psycopg2_connect, requests_get +from .test_etcd import etcd_read, socket_getaddrinfo +from .test_ha import get_cluster_initialized_without_leader, get_cluster_initialized_with_leader, \ get_cluster_initialized_with_only_leader, get_cluster_not_initialized_without_leader, get_cluster, Member -from test_postgresql import MockConnect, psycopg2_connect CONFIG_FILE_PATH = './test-ctl.yaml' @@ -198,7 +199,7 @@ class TestCtl(unittest.TestCase): rows = query_member(None, None, None, 'replica', 'SELECT pg_catalog.pg_is_in_recovery()', {}) self.assertEqual(rows, (None, None)) - with patch('test_postgresql.MockCursor.execute', Mock(side_effect=OperationalError('bla'))): + with patch.object(MockCursor, 'execute', Mock(side_effect=OperationalError('bla'))): rows = query_member(None, None, None, 'replica', 'SELECT pg_catalog.pg_is_in_recovery()', {}) with patch('patroni.ctl.get_cursor', Mock(return_value=None)): @@ -289,14 +290,14 @@ class TestCtl(unittest.TestCase): with patch('requests.post', Mock(return_value=MockResponse())): # normal restart, the schedule is actually parsed, but not validated in patronictl result = self.runner.invoke(ctl, ['restart', 'alpha', '--pg-version', '42.0.0', - '--scheduled', '2300-10-01T14:30'], input='y') + '--scheduled', '2300-10-01T14:30'], input='y') assert result.exit_code == 0 with patch('requests.post', Mock(return_value=MockResponse(204))): # get restart with the non-200 return code # normal restart, the schedule is actually parsed, but not validated in patronictl result = self.runner.invoke(ctl, ['restart', 'alpha', '--pg-version', '42.0', - '--scheduled', '2300-10-01T14:30'], input='y') + '--scheduled', '2300-10-01T14:30'], input='y') assert result.exit_code == 0 # force restart with restart already present diff --git a/tests/test_etcd.py b/tests/test_etcd.py index fb66d139..932a117f 100644 --- a/tests/test_etcd.py +++ b/tests/test_etcd.py @@ -1,7 +1,5 @@ import etcd -import json import urllib3.util.connection -import requests import socket import unittest @@ -11,56 +9,7 @@ from patroni.dcs.etcd import AbstractDCS, Client, Cluster, Etcd, EtcdError, DnsC from patroni.exceptions import DCSError from urllib3.exceptions import ReadTimeoutError - -class MockResponse(object): - - def __init__(self, status_code=200): - self.status_code = status_code - self.content = '{}' - self.ok = True - - def json(self): - return json.loads(self.content) - - @property - def data(self): - return self.content.encode('utf-8') - - @property - def text(self): - return self.content - - @property - def status(self): - return self.status_code - - @staticmethod - def getheader(*args): - return '' - - -def requests_get(url, **kwargs): - members = '[{"id":14855829450254237642,"peerURLs":["http://localhost:2380","http://localhost:7001"],' +\ - '"name":"default","clientURLs":["http://localhost:2379","http://localhost:4001"]}]' - response = MockResponse() - if url.startswith('http://local'): - raise requests.exceptions.RequestException() - elif ':8011/patroni' in url: - response.content = '{"role": "replica", "xlog": {"received_location": 0}, "tags": {}}' - elif url.endswith('/members'): - response.content = '[{}]' if url.startswith('http://error') else members - elif url.startswith('http://exhibitor'): - response.content = '{"servers":["127.0.0.1","127.0.0.2","127.0.0.3"],"port":2181}' - elif url.endswith(':8011/reinitialize'): - data = kwargs.get('data', '') - if ' false}' in data: - response.status_code = 503 - response.ok = False - response.content = 'restarting after failure already in progress' - else: - response.status_code = 404 - response.ok = False - return response +from . import SleepException, MockResponse, requests_get def etcd_watch(self, key, index=None, timeout=None, recursive=None): @@ -122,10 +71,6 @@ def etcd_read(self, key, **kwargs): return result -class SleepException(Exception): - pass - - def dns_query(name, _): if '-server' not in name or '-ssl' in name: return [] diff --git a/tests/test_exhibitor.py b/tests/test_exhibitor.py index d9c0d6d3..b0d51196 100644 --- a/tests/test_exhibitor.py +++ b/tests/test_exhibitor.py @@ -3,8 +3,9 @@ import unittest from mock import Mock, patch from patroni.dcs.exhibitor import ExhibitorEnsembleProvider, Exhibitor from patroni.dcs.zookeeper import ZooKeeperError -from test_etcd import SleepException, requests_get -from test_zookeeper import MockKazooClient + +from . import SleepException, requests_get +from .test_zookeeper import MockKazooClient @patch('requests.get', requests_get) diff --git a/tests/test_ha.py b/tests/test_ha.py index 2a31cb8d..367c491a 100644 --- a/tests/test_ha.py +++ b/tests/test_ha.py @@ -1,7 +1,6 @@ import datetime import etcd import os -import unittest import sys from mock import Mock, MagicMock, PropertyMock, patch @@ -11,10 +10,16 @@ from patroni.dcs.etcd import Client from patroni.exceptions import DCSError, PostgresConnectionException, PatroniException from patroni.ha import Ha, _MemberStatus from patroni.postgresql import Postgresql -from patroni.watchdog import Watchdog +from patroni.postgresql.bootstrap import Bootstrap +from patroni.postgresql.cancellable import CancellableSubprocess +from patroni.postgresql.config import ConfigHandler +from patroni.postgresql.rewind import Rewind +from patroni.postgresql.slots import SlotsHandler from patroni.utils import tzutc -from test_etcd import socket_getaddrinfo, etcd_read, etcd_write, requests_get -from test_postgresql import psycopg2_connect, MockPostmaster +from patroni.watchdog import Watchdog + +from . import PostgresInit, MockPostmaster, psycopg2_connect, requests_get +from .test_etcd import socket_getaddrinfo, etcd_read, etcd_write SYSID = '12345678901' @@ -145,16 +150,16 @@ def run_async(self, func, args=()): @patch.object(Postgresql, 'call_nowait', Mock(return_value=True)) @patch.object(Postgresql, 'data_directory_empty', Mock(return_value=False)) @patch.object(Postgresql, 'controldata', Mock(return_value={'Database system identifier': SYSID})) -@patch.object(Postgresql, 'sync_replication_slots', Mock()) -@patch.object(Postgresql, 'append_pg_hba', Mock()) +@patch.object(SlotsHandler, 'sync_replication_slots', Mock()) +@patch.object(ConfigHandler, 'append_pg_hba', Mock()) @patch.object(Postgresql, 'write_pgpass', Mock(return_value={})) -@patch.object(Postgresql, 'write_recovery_conf', Mock()) +@patch.object(ConfigHandler, 'write_recovery_conf', Mock()) @patch.object(Postgresql, 'query', Mock()) @patch.object(Postgresql, 'checkpoint', Mock()) -@patch.object(Postgresql, 'cancellable_subprocess_call', Mock(return_value=0)) -@patch.object(Postgresql, '_get_local_timeline_lsn_from_replication_connection', Mock(return_value=[2, 10])) +@patch.object(CancellableSubprocess, 'call', Mock(return_value=0)) +@patch.object(Postgresql, 'get_local_timeline_lsn_from_replication_connection', Mock(return_value=[2, 10])) @patch.object(Postgresql, 'get_master_timeline', Mock(return_value=2)) -@patch.object(Postgresql, 'restore_configuration_files', Mock()) +@patch.object(ConfigHandler, 'restore_configuration_files', Mock()) @patch.object(etcd.Client, 'write', etcd_write) @patch.object(etcd.Client, 'read', etcd_read) @patch.object(etcd.Client, 'delete', Mock(side_effect=etcd.EtcdException)) @@ -163,21 +168,15 @@ def run_async(self, func, args=()): @patch('patroni.async_executor.AsyncExecutor.run_async', run_async) @patch('subprocess.call', Mock(return_value=0)) @patch('time.sleep', Mock()) -class TestHa(unittest.TestCase): +class TestHa(PostgresInit): @patch('socket.getaddrinfo', socket_getaddrinfo) - @patch('psycopg2.connect', psycopg2_connect) @patch('patroni.dcs.dcs_modules', Mock(return_value=['patroni.dcs.foo', 'patroni.dcs.etcd'])) @patch.object(etcd.Client, 'read', etcd_read) def setUp(self): + super(TestHa, self).setUp() with patch.object(Client, 'machines') as mock_machines: mock_machines.__get__ = Mock(return_value=['http://remotehost:2379']) - self.p = Postgresql({'name': 'postgresql0', 'scope': 'dummy', 'listen': '127.0.0.1:5432', - 'data_dir': 'data/postgresql0', 'retry_timeout': 10, - 'authentication': {'superuser': {'username': 'foo', 'password': 'bar'}, - 'replication': {'username': '', 'password': ''}}, - 'parameters': {'wal_level': 'hot_standby', 'max_replication_slots': 5, 'foo': 'bar', - 'hot_standby': 'on', 'max_wal_senders': 5, 'wal_keep_segments': 8}}) self.p.set_state('running') self.p.set_role('replica') self.p.postmaster_start_time = MagicMock(return_value=str(postmaster_start_time)) @@ -222,7 +221,7 @@ class TestHa(unittest.TestCase): @patch.object(Cluster, 'get_clone_member', Mock(return_value=Member(0, 'test', 1, {'api_url': 'http://127.0.0.1:8011/patroni', 'conn_url': 'postgres://127.0.0.1:5432/postgres'}))) - @patch.object(Postgresql, 'create_replica', Mock(return_value=0)) + @patch.object(Bootstrap, 'create_replica', Mock(return_value=0)) def test_start_as_cascade_replica_in_standby_cluster(self): self.p.data_directory_empty = true self.ha.cluster = get_standby_cluster_initialized_with_only_leader() @@ -251,15 +250,15 @@ class TestHa(unittest.TestCase): self.p.controldata = lambda: {'Database cluster state': 'in production', 'Database system identifier': SYSID} self.assertEqual(self.ha.run_cycle(), 'doing crash recovery in a single user mode') - @patch.object(Postgresql, 'rewind_or_reinitialize_needed_and_possible', Mock(return_value=True)) - @patch.object(Postgresql, 'can_rewind', PropertyMock(return_value=True)) + @patch.object(Rewind, 'rewind_or_reinitialize_needed_and_possible', Mock(return_value=True)) + @patch.object(Rewind, 'can_rewind', PropertyMock(return_value=True)) def test_recover_with_rewind(self): self.p.is_running = false self.ha.cluster = get_cluster_initialized_with_leader() self.assertEqual(self.ha.run_cycle(), 'running pg_rewind from leader') - @patch.object(Postgresql, 'rewind_or_reinitialize_needed_and_possible', Mock(return_value=True)) - @patch.object(Postgresql, 'create_replica', Mock(return_value=1)) + @patch.object(Rewind, 'rewind_or_reinitialize_needed_and_possible', Mock(return_value=True)) + @patch.object(Bootstrap, 'create_replica', Mock(return_value=1)) def test_recover_with_reinitialize(self): self.p.is_running = false self.ha.cluster = get_cluster_initialized_with_leader() @@ -363,11 +362,11 @@ class TestHa(unittest.TestCase): self.p.is_leader = false self.assertEqual(self.ha.run_cycle(), 'PAUSE: no action') - @patch.object(Postgresql, 'rewind_or_reinitialize_needed_and_possible', Mock(return_value=True)) - @patch.object(Postgresql, 'can_rewind', PropertyMock(return_value=True)) + @patch.object(Rewind, 'rewind_or_reinitialize_needed_and_possible', Mock(return_value=True)) + @patch.object(Rewind, 'can_rewind', PropertyMock(return_value=True)) def test_follow_triggers_rewind(self): self.p.is_leader = false - self.p.trigger_check_diverged_lsn() + self.ha._rewind.trigger_check_diverged_lsn() self.ha.cluster = get_cluster_initialized_with_leader() self.assertEqual(self.ha.run_cycle(), 'running pg_rewind from leader') @@ -478,7 +477,7 @@ class TestHa(unittest.TestCase): f = Failover(0, self.p.name, '', None) self.ha.cluster = get_cluster_initialized_with_leader(f) self.assertEqual(self.ha.run_cycle(), 'manual failover: demoting myself') - self.p.rewind_or_reinitialize_needed_and_possible = true + self.ha._rewind.rewind_or_reinitialize_needed_and_possible = true self.assertEqual(self.ha.run_cycle(), 'manual failover: demoting myself') self.ha.fetch_node_status = get_node_status(nofailover=True) self.assertEqual(self.ha.run_cycle(), 'no action. i am the leader with the lock') @@ -664,7 +663,7 @@ class TestHa(unittest.TestCase): def test_restart_matches(self): self.p._role = 'replica' - self.p.server_version = 90500 + self.p._connection.server_version = 90500 self.p._pending_restart = True self.assertFalse(self.ha.restart_matches("master", "9.5.0", True)) self.assertFalse(self.ha.restart_matches("replica", "9.4.3", True)) @@ -685,7 +684,7 @@ class TestHa(unittest.TestCase): self.p.is_leader = false self.p.name = 'leader' self.ha.cluster = get_standby_cluster_initialized_with_only_leader() - self.p.check_recovery_conf = true + self.p.config.check_recovery_conf = true self.assertEqual(self.ha.run_cycle(), 'promoted self to a standby leader because i had the session lock') self.assertEqual(self.ha.run_cycle(), 'no action. i am the standby leader with the lock') @@ -707,8 +706,8 @@ class TestHa(unittest.TestCase): self.p._sysid = True self.assertEqual(self.ha.run_cycle(), 'promoted self to a standby leader by acquiring session lock') - @patch.object(Postgresql, 'rewind_or_reinitialize_needed_and_possible', Mock(return_value=True)) - @patch.object(Postgresql, 'can_rewind', PropertyMock(return_value=True)) + @patch.object(Rewind, 'rewind_or_reinitialize_needed_and_possible', Mock(return_value=True)) + @patch.object(Rewind, 'can_rewind', PropertyMock(return_value=True)) def test_process_unhealthy_standby_cluster_as_cascade_replica(self): self.p.is_leader = false self.p.name = 'replica' @@ -819,7 +818,7 @@ class TestHa(unittest.TestCase): def test_process_sync_replication(self): self.ha.has_lock = true - mock_set_sync = self.p.set_synchronous_standby = Mock() + mock_set_sync = self.p.config.set_synchronous_standby = Mock() self.p.name = 'leader' # Test sync key removed when sync mode disabled @@ -896,7 +895,7 @@ class TestHa(unittest.TestCase): def test_sync_replication_become_master(self): self.ha.is_synchronous_mode = true - mock_set_sync = self.p.set_synchronous_standby = Mock() + mock_set_sync = self.p.config.set_synchronous_standby = Mock() self.p.is_leader = false self.p.set_role('replica') self.ha.has_lock = true diff --git a/tests/test_log.py b/tests/test_log.py index 5363e782..8cee50bf 100644 --- a/tests/test_log.py +++ b/tests/test_log.py @@ -1,3 +1,4 @@ +import logging import os import sys import unittest @@ -10,6 +11,12 @@ from patroni.log import PatroniLogger class TestPatroniLogger(unittest.TestCase): + def setUp(self): + self._handlers = logging.getLogger().handlers[:] + + def tearDown(self): + logging.getLogger().handlers[:] = self._handlers + @patch('logging.FileHandler._open', Mock()) def test_patroni_logger(self): config = { diff --git a/tests/test_patroni.py b/tests/test_patroni.py index b6ce61e9..b6fdf1e2 100644 --- a/tests/test_patroni.py +++ b/tests/test_patroni.py @@ -1,5 +1,5 @@ import etcd -import psycopg2 +import logging import signal import sys import time @@ -10,10 +10,14 @@ from patroni.api import RestApiServer from patroni.async_executor import AsyncExecutor from patroni.dcs.etcd import Client from patroni.exceptions import DCSError +from patroni.postgresql import Postgresql +from patroni.postgresql.config import ConfigHandler from patroni import Patroni, main as _main, patroni_main, check_psycopg2 from six.moves import BaseHTTPServer, builtins -from test_etcd import SleepException, etcd_read, etcd_write -from test_postgresql import Postgresql, psycopg2_connect, MockPostmaster + +from . import psycopg2_connect, SleepException +from .test_etcd import etcd_read, etcd_write +from .test_postgresql import MockPostmaster class MockFrozenImporter(object): @@ -24,9 +28,9 @@ class MockFrozenImporter(object): @patch('time.sleep', Mock()) @patch('subprocess.call', Mock(return_value=0)) @patch('psycopg2.connect', psycopg2_connect) -@patch.object(Postgresql, 'append_pg_hba', Mock()) -@patch.object(Postgresql, '_write_postgresql_conf', Mock()) -@patch.object(Postgresql, 'write_recovery_conf', Mock()) +@patch.object(ConfigHandler, 'append_pg_hba', Mock()) +@patch.object(ConfigHandler, 'write_postgresql_conf', Mock()) +@patch.object(ConfigHandler, 'write_recovery_conf', Mock()) @patch.object(Postgresql, 'is_running', Mock(return_value=MockPostmaster())) @patch.object(Postgresql, 'call_nowait', Mock()) @patch.object(BaseHTTPServer.HTTPServer, '__init__', Mock()) @@ -40,6 +44,7 @@ class TestPatroni(unittest.TestCase): @patch.object(BaseHTTPServer.HTTPServer, '__init__', Mock()) @patch.object(etcd.Client, 'read', etcd_read) def setUp(self): + self._handlers = logging.getLogger().handlers[:] RestApiServer._BaseServer__is_shut_down = Mock() RestApiServer._BaseServer__shutdown_request = True RestApiServer.socket = 0 @@ -48,6 +53,9 @@ class TestPatroni(unittest.TestCase): sys.argv = ['patroni.py', 'postgres0.yml'] self.p = Patroni() + def tearDown(self): + logging.getLogger().handlers[:] = self._handlers + @patch('patroni.dcs.AbstractDCS.get_cluster', Mock(side_effect=[None, DCSError('foo'), None])) def test_load_dynamic_configuration(self): self.p.config._dynamic_configuration = {} @@ -159,5 +167,5 @@ class TestPatroni(unittest.TestCase): def test_check_psycopg2(self): with patch.object(builtins, '__import__', Mock(side_effect=ImportError)): self.assertRaises(SystemExit, check_psycopg2) - with patch.object(psycopg2, '__version__', return_value='2.5.3.dev1 a b c'): + with patch('psycopg2.__version__', '2.5.3.dev1 a b c'): self.assertRaises(SystemExit, check_psycopg2) diff --git a/tests/test_postgresql.py b/tests/test_postgresql.py index 85ef2f48..c58ffee1 100644 --- a/tests/test_postgresql.py +++ b/tests/test_postgresql.py @@ -1,114 +1,20 @@ -import datetime import mock # for the mock.call method, importing it without a namespace breaks python3 import os import psycopg2 -import shutil import subprocess -import unittest from mock import Mock, MagicMock, PropertyMock, patch, mock_open from patroni.async_executor import CriticalTask -from patroni.dcs import Cluster, ClusterConfig, Leader, Member, RemoteMember, SyncState -from patroni.exceptions import PostgresConnectionException, PostgresException +from patroni.dcs import Cluster, ClusterConfig, Member, RemoteMember, SyncState +from patroni.exceptions import PostgresConnectionException from patroni.postgresql import Postgresql, STATE_REJECT, STATE_NO_RESPONSE -from patroni.postmaster import PostmasterProcess +from patroni.postgresql.postmaster import PostmasterProcess +from patroni.postgresql.slots import SlotsHandler from patroni.utils import RetryFailedError from six.moves import builtins from threading import Thread, current_thread -from tempfile import gettempdir - -class MockCursor(object): - - def __init__(self, connection): - self.connection = connection - self.closed = False - self.rowcount = 0 - self.results = [] - - def execute(self, sql, *params): - if sql.startswith('blabla'): - raise psycopg2.ProgrammingError() - elif sql == 'CHECKPOINT' or sql.startswith('SELECT pg_catalog.pg_create_'): - raise psycopg2.OperationalError() - elif sql.startswith('RetryFailedError'): - raise RetryFailedError('retry') - elif sql.startswith('SELECT slot_name'): - self.results = [('blabla', 'physical'), ('foobar', 'physical'), ('ls', 'logical', 'a', 'b')] - elif sql.startswith('SELECT CASE WHEN pg_catalog.pg_is_in_recovery()'): - self.results = [(1, 2)] - elif sql.startswith('SELECT pg_catalog.pg_is_in_recovery()'): - self.results = [(False, 2)] - elif sql.startswith('WITH replication_info AS ('): - replication_info = '[{"application_name":"walreceiver","client_addr":"1.2.3.4",' +\ - '"state":"streaming","sync_state":"async","sync_priority":0}]' - self.results = [('', 0, '', '', '', '', False, replication_info)] - elif sql.startswith('SELECT name, setting'): - self.results = [('wal_segment_size', '2048', '8kB', 'integer', 'internal'), - ('search_path', 'public', None, 'string', 'user'), - ('port', '5433', None, 'integer', 'postmaster'), - ('listen_addresses', '*', None, 'string', 'postmaster'), - ('autovacuum', 'on', None, 'bool', 'sighup'), - ('unix_socket_directories', '/tmp', None, 'string', 'postmaster')] - elif sql.startswith('IDENTIFY_SYSTEM'): - self.results = [('1', 2, '0/402EEC0', '')] - elif sql.startswith('SELECT isdir, modification'): - self.results = [(False, datetime.datetime.now())] - elif sql.startswith('SELECT pg_catalog.pg_read_file'): - self.results = [('1\t0/40159C0\tno recovery target specified\n\n' + - '2\t1/40159C0\tno recovery target specified\n',)] - elif sql.startswith('TIMELINE_HISTORY '): - self.results = [('', b'x\t0/40159C0\tno recovery target specified\n\n' + - b'1\t0/40159C0\tno recovery target specified\n\n' + - b'2\t0/402DD98\tno recovery target specified\n\n' + - b'3\t0/403DD98\tno recovery target specified\n')] - else: - self.results = [(None, None, None, None, None, None, None, None, None, None)] - - def fetchone(self): - return self.results[0] - - def fetchall(self): - return self.results - - def __iter__(self): - for i in self.results: - yield i - - def __enter__(self): - return self - - def __exit__(self, *args): - pass - - -class MockConnect(object): - - server_version = 99999 - autocommit = False - closed = 0 - - def cursor(self): - return MockCursor(self) - - def __enter__(self): - return self - - def __exit__(self, *args): - pass - - @staticmethod - def close(): - pass - - -class MockPostmaster(object): - def __init__(self, is_running=True, is_single_master=False): - self.is_running = Mock(return_value=is_running) - self.is_single_master = Mock(return_value=is_single_master) - self.wait_for_user_backends_to_close = Mock() - self.signal_stop = Mock(return_value=None) - self.wait = Mock() +from . import BaseTestPostgresql, MockCursor, MockPostmaster, psycopg2_connect def pg_controldata_string(*args, **kwargs): @@ -166,63 +72,18 @@ Data page checksum version: 0 """ -def psycopg2_connect(*args, **kwargs): - return MockConnect() - - @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, 'f.oo': 'bar', - 'search_path': 'public', 'hot_standby': 'on', 'max_wal_senders': 5, - 'wal_keep_segments': 8, 'wal_log_hints': 'on', 'max_locks_per_transaction': 64, - 'max_worker_processes': 8, 'max_connections': 100, 'max_prepared_transactions': 0, - 'track_commit_timestamp': 'off', 'unix_socket_directories': '/tmp'} +class TestPostgresql(BaseTestPostgresql): @patch('subprocess.call', Mock(return_value=0)) - @patch('psycopg2.connect', psycopg2_connect) @patch('os.rename', Mock()) @patch.object(Postgresql, 'get_major_version', Mock(return_value=90600)) @patch.object(Postgresql, 'is_running', Mock(return_value=True)) def setUp(self): - self.data_dir = 'data/test0' - self.config_dir = self.data_dir - if not os.path.exists(self.data_dir): - os.makedirs(self.data_dir) - self.p = Postgresql({'name': 'test0', 'scope': 'batman', 'data_dir': self.data_dir, - 'config_dir': self.config_dir, 'retry_timeout': 10, - 'krbsrvname': 'postgres', 'pgpass': os.path.join(gettempdir(), 'pgpass0'), - 'listen': '127.0.0.2, 127.0.0.3:5432', 'connect_address': '127.0.0.2:5432', - 'authentication': {'superuser': {'username': 'test', 'password': 'test'}, - 'replication': {'username': 'replicator', 'password': 'rep-pass'}}, - 'remove_data_directory_on_rewind_failure': True, - 'use_pg_rewind': True, 'pg_ctl_timeout': 'bla', - 'parameters': self._PARAMETERS, - 'recovery_conf': {'foo': 'bar'}, - 'pg_hba': ['host all all 0.0.0.0/0 md5'], - 'pg_ident': ['krb realm postgres'], - 'callbacks': {'on_start': 'true', 'on_stop': 'true', 'on_reload': 'true', - 'on_restart': 'true', 'on_role_change': 'true'}}) + super(TestPostgresql, self).setUp() + self.p.config.write_postgresql_conf() self.p._callback_executor = Mock() - self.leadermem = Member(0, 'leader', 28, {'conn_url': 'postgres://replicator:rep-pass@127.0.0.1:5435/postgres'}) - self.leader = Leader(-1, 28, self.leadermem) - self.other = Member(0, 'test-1', 28, {'conn_url': 'postgres://replicator:rep-pass@127.0.0.1:5433/postgres', - 'tags': {'replicatefrom': 'leader'}}) - self.me = Member(0, 'test0', 28, {'conn_url': 'postgres://replicator:rep-pass@127.0.0.1:5434/postgres'}) - - def tearDown(self): - shutil.rmtree('data') - - def test__initdb(self): - self.assertRaises(Exception, self.p.bootstrap, {'initdb': [{'pgdata': 'bar'}]}) - self.assertRaises(Exception, self.p.bootstrap, {'initdb': [{'foo': 'bar', 1: 2}]}) - self.assertRaises(Exception, self.p.bootstrap, {'initdb': [1]}) - self.assertRaises(Exception, self.p.bootstrap, {'initdb': 1}) - - @patch('os.path.exists', Mock(return_value=True)) - @patch('os.unlink', Mock()) - def test_delete_trigger_file(self): - self.p.delete_trigger_file() @patch('subprocess.Popen') @patch.object(Postgresql, 'wait_for_startup') @@ -238,7 +99,7 @@ class TestPostgresql(unittest.TestCase): mock_postmaster = MockPostmaster() with patch.object(PostmasterProcess, 'start', return_value=mock_postmaster): - pg_conf = os.path.join(self.data_dir, 'postgresql.conf') + pg_conf = os.path.join(self.p.data_dir, 'postgresql.conf') open(pg_conf, 'w').close() self.assertFalse(self.p.start(task=CriticalTask())) @@ -256,7 +117,7 @@ class TestPostgresql(unittest.TestCase): task.cancel() self.assertFalse(self.p.start(task=task)) - self.p.cancel() + self.p.cancellable.cancel() self.assertFalse(self.p.start()) @patch.object(Postgresql, 'pg_isready') @@ -278,7 +139,7 @@ class TestPostgresql(unittest.TestCase): self.assertTrue(self.p.wait_for_port_open(mock_postmaster, 1)) # cancelled - self.p.cancel() + self.p.cancellable.cancel() self.assertFalse(self.p.wait_for_port_open(mock_postmaster, 1)) @patch('time.sleep', Mock()) @@ -325,91 +186,11 @@ class TestPostgresql(unittest.TestCase): self.assertIsNone(self.p.checkpoint()) self.assertEqual(self.p.checkpoint(), 'not accessible or not healty') - @patch.object(Postgresql, 'cancellable_subprocess_call') - @patch('patroni.postgresql.Postgresql.write_pgpass', MagicMock(return_value=dict())) - def test_pg_rewind(self, mock_cancellable_subprocess_call): - r = {'user': '', 'host': '', 'port': '', 'database': '', 'password': ''} - mock_cancellable_subprocess_call.return_value = 0 - self.assertTrue(self.p.pg_rewind(r)) - mock_cancellable_subprocess_call.side_effect = OSError - self.assertFalse(self.p.pg_rewind(r)) - def test_check_recovery_conf(self): - self.p.write_recovery_conf({'primary_conninfo': 'foo'}) - self.assertFalse(self.p.check_recovery_conf(None)) - self.p.write_recovery_conf({}) - self.assertTrue(self.p.check_recovery_conf(None)) - - @patch.object(Postgresql, 'start', Mock()) - @patch.object(Postgresql, 'can_rewind', PropertyMock(return_value=True)) - def test__get_local_timeline_lsn(self): - self.p.trigger_check_diverged_lsn() - with patch.object(Postgresql, 'controldata', - Mock(return_value={'Database cluster state': 'shut down in recovery', - 'Minimum recovery ending location': '0/0', - "Min recovery ending loc's timeline": '0'})): - self.p.rewind_or_reinitialize_needed_and_possible(self.leader) - with patch.object(Postgresql, 'is_running', Mock(return_value=True)): - with patch.object(MockCursor, 'fetchone', Mock(side_effect=[(False, ), Exception])): - self.p.rewind_or_reinitialize_needed_and_possible(self.leader) - - @patch.object(Postgresql, 'start', Mock()) - @patch.object(Postgresql, 'can_rewind', PropertyMock(return_value=True)) - @patch.object(Postgresql, '_get_local_timeline_lsn', Mock(return_value=(2, '40159C1'))) - @patch.object(Postgresql, 'check_leader_is_not_in_recovery') - def test__check_timeline_and_lsn(self, mock_check_leader_is_not_in_recovery): - mock_check_leader_is_not_in_recovery.return_value = False - self.p.trigger_check_diverged_lsn() - self.assertFalse(self.p.rewind_or_reinitialize_needed_and_possible(self.leader)) - self.leader = self.leader.member - self.assertFalse(self.p.rewind_or_reinitialize_needed_and_possible(self.leader)) - mock_check_leader_is_not_in_recovery.return_value = True - self.assertFalse(self.p.rewind_or_reinitialize_needed_and_possible(self.leader)) - self.p.trigger_check_diverged_lsn() - with patch('psycopg2.connect', Mock(side_effect=Exception)): - self.assertFalse(self.p.rewind_or_reinitialize_needed_and_possible(self.leader)) - self.p.trigger_check_diverged_lsn() - with patch.object(MockCursor, 'fetchone', Mock(side_effect=[('', 2, '0/0'), ('', b'3\t0/40159C0\tn\n')])): - self.assertFalse(self.p.rewind_or_reinitialize_needed_and_possible(self.leader)) - self.p.trigger_check_diverged_lsn() - with patch.object(MockCursor, 'fetchone', Mock(return_value=('', 1, '0/0'))): - with patch.object(Postgresql, '_get_local_timeline_lsn', Mock(return_value=(1, '0/0'))): - self.assertFalse(self.p.rewind_or_reinitialize_needed_and_possible(self.leader)) - self.p.trigger_check_diverged_lsn() - self.assertTrue(self.p.rewind_or_reinitialize_needed_and_possible(self.leader)) - - @patch.object(MockCursor, 'fetchone', Mock(side_effect=[(True,), Exception])) - def test_check_leader_is_not_in_recovery(self): - self.p.check_leader_is_not_in_recovery() - self.p.check_leader_is_not_in_recovery() - - @patch.object(Postgresql, 'cancellable_subprocess_call', Mock(return_value=0)) - @patch.object(Postgresql, 'checkpoint', side_effect=['', '1']) - @patch.object(Postgresql, 'stop', Mock(return_value=False)) - @patch.object(Postgresql, 'start', Mock()) - def test_rewind(self, mock_checkpoint): - self.p.rewind(self.leader) - with patch.object(Postgresql, 'pg_rewind', Mock(return_value=False)): - mock_checkpoint.side_effect = ['1', '', '', ''] - self.p.rewind(self.leader) - self.p.rewind(self.leader) - with patch.object(Postgresql, 'check_leader_is_not_in_recovery', Mock(return_value=False)): - self.p.rewind(self.leader) - self.p.config['remove_data_directory_on_rewind_failure'] = False - self.p.trigger_check_diverged_lsn() - self.p.rewind(self.leader) - - self.leader.member.data.update(version='1.5.7', checkpoint_after_promote=False) - self.assertIsNone(self.p.rewind(self.leader)) - - self.leader.member.data['checkpoint_after_promote'] = True - with patch.object(Postgresql, 'check_leader_is_not_in_recovery', Mock(return_value=False)): - self.assertIsNone(self.p.rewind(self.leader)) - - with patch.object(Postgresql, 'is_running', Mock(return_value=True)): - self.p.rewind(self.leader) - self.p.is_leader = Mock(return_value=False) - self.p.rewind(self.leader) + self.p.config.write_recovery_conf({'primary_conninfo': 'foo'}) + self.assertFalse(self.p.config.check_recovery_conf(None)) + self.p.config.write_recovery_conf({}) + self.assertTrue(self.p.config.check_recovery_conf(None)) @patch.object(Postgresql, 'is_running', Mock(return_value=False)) @patch.object(Postgresql, 'start', Mock()) @@ -418,101 +199,6 @@ class TestPostgresql(unittest.TestCase): m = RemoteMember('1', {'restore_command': '2', 'recovery_min_apply_delay': 3, 'archive_cleanup_command': '4'}) self.p.follow(m) - @patch('subprocess.check_output', Mock(return_value=0, side_effect=pg_controldata_string)) - def test_can_rewind(self): - with patch('subprocess.call', MagicMock(return_value=1)): - self.assertFalse(self.p.can_rewind) - with patch('subprocess.call', side_effect=OSError): - self.assertFalse(self.p.can_rewind) - with patch.object(Postgresql, 'controldata', Mock(return_value={'wal_log_hints setting': 'on'})): - self.assertTrue(self.p.can_rewind) - self.p.config['use_pg_rewind'] = False - self.assertFalse(self.p.can_rewind) - - @patch('time.sleep', Mock()) - @patch.object(Postgresql, 'cancellable_subprocess_call') - @patch.object(Postgresql, 'remove_data_directory', Mock(return_value=True)) - def test_create_replica(self, mock_cancellable_subprocess_call): - self.p.delete_trigger_file = Mock(side_effect=OSError) - - self.p.config['create_replica_methods'] = ['pgBackRest'] - self.p.config['pgBackRest'] = {'command': 'pgBackRest', 'keep_data': True, 'no_params': True} - mock_cancellable_subprocess_call.return_value = 0 - self.assertEqual(self.p.create_replica(self.leader), 0) - - self.p.config['create_replica_methods'] = ['wale', 'basebackup'] - self.p.config['wale'] = {'command': 'foo'} - self.assertEqual(self.p.create_replica(self.leader), 0) - del self.p.config['wale'] - self.assertEqual(self.p.create_replica(self.leader), 0) - - self.p.config['create_replica_methods'] = ['basebackup'] - self.p.config['basebackup'] = [{'max_rate': '100M'}, 'no-sync'] - self.assertEqual(self.p.create_replica(self.leader), 0) - - self.p.config['basebackup'] = [{'max_rate': '100M', 'compress': '9'}] - with mock.patch('patroni.postgresql.logger.error', new_callable=Mock()) as mock_logger: - self.p.create_replica(self.leader) - mock_logger.assert_called_once() - self.assertTrue("only one key-value is allowed and value should be a string" in mock_logger.call_args[0][0], - "not matching {0}".format(mock_logger.call_args[0][0])) - - self.p.config['basebackup'] = [42] - with mock.patch('patroni.postgresql.logger.error', new_callable=Mock()) as mock_logger: - self.p.create_replica(self.leader) - mock_logger.assert_called_once() - self.assertTrue("value should be string value or a single key-value pair" in mock_logger.call_args[0][0], - "not matching {0}".format(mock_logger.call_args[0][0])) - - self.p.config['basebackup'] = {"foo": "bar"} - self.assertEqual(self.p.create_replica(self.leader), 0) - - self.p.config['create_replica_methods'] = ['wale', 'basebackup'] - del self.p.config['basebackup'] - mock_cancellable_subprocess_call.return_value = 1 - self.assertEqual(self.p.create_replica(self.leader), 1) - - mock_cancellable_subprocess_call.side_effect = Exception('foo') - self.assertEqual(self.p.create_replica(self.leader), 1) - - mock_cancellable_subprocess_call.side_effect = [1, 0] - self.assertEqual(self.p.create_replica(self.leader), 0) - - mock_cancellable_subprocess_call.side_effect = [Exception(), 0] - self.assertEqual(self.p.create_replica(self.leader), 0) - - self.p.cancel() - self.assertEqual(self.p.create_replica(self.leader), 1) - - @patch('time.sleep', Mock()) - @patch.object(Postgresql, 'cancellable_subprocess_call') - @patch.object(Postgresql, 'remove_data_directory', Mock(return_value=True)) - def test_create_replica_old_format(self, mock_cancellable_subprocess_call): - """ The same test as before but with old 'create_replica_method' - to test backward compatibility - """ - self.p.delete_trigger_file = Mock(side_effect=OSError) - - self.p.config['create_replica_method'] = ['wale', 'basebackup'] - self.p.config['wale'] = {'command': 'foo'} - mock_cancellable_subprocess_call.return_value = 0 - self.assertEqual(self.p.create_replica(self.leader), 0) - del self.p.config['wale'] - self.assertEqual(self.p.create_replica(self.leader), 0) - - self.p.config['create_replica_method'] = ['basebackup'] - self.p.config['basebackup'] = [{'max_rate': '100M'}, 'no-sync'] - self.assertEqual(self.p.create_replica(self.leader), 0) - - self.p.config['create_replica_method'] = ['wale', 'basebackup'] - del self.p.config['basebackup'] - mock_cancellable_subprocess_call.return_value = 1 - self.assertEqual(self.p.create_replica(self.leader), 1) - - def test_basebackup(self): - self.p.cancel() - self.p.basebackup(None, None, {'foo': 'bar'}) - @patch.object(Postgresql, 'is_running', Mock(return_value=True)) def test_sync_replication_slots(self): self.p.start() @@ -520,17 +206,16 @@ class TestPostgresql(unittest.TestCase): 'A': 0, 'test_3': 0, 'b': {'type': 'logical', 'plugin': '1'}}}, 1) cluster = Cluster(True, config, self.leader, 0, [self.me, self.other, self.leadermem], None, None, None) with mock.patch('patroni.postgresql.Postgresql._query', Mock(side_effect=psycopg2.OperationalError)): - self.p.sync_replication_slots(cluster) - self.p.sync_replication_slots(cluster) + self.p.slots_handler.sync_replication_slots(cluster) + self.p.slots_handler.sync_replication_slots(cluster) with mock.patch('patroni.postgresql.Postgresql.role', new_callable=PropertyMock(return_value='replica')): - self.p.sync_replication_slots(cluster) - with patch.object(Postgresql, 'drop_replication_slot', Mock(return_value=True)),\ + self.p.slots_handler.sync_replication_slots(cluster) + with patch.object(SlotsHandler, 'drop_replication_slot', Mock(return_value=True)),\ patch('patroni.dcs.logger.error', new_callable=Mock()) as errorlog_mock: - self.p.query = Mock() alias1 = Member(0, 'test-3', 28, {'conn_url': 'postgres://replicator:rep-pass@127.0.0.1:5436/postgres'}) alias2 = Member(0, 'test.3', 28, {'conn_url': 'postgres://replicator:rep-pass@127.0.0.1:5436/postgres'}) cluster.members.extend([alias1, alias2]) - self.p.sync_replication_slots(cluster) + self.p.slots_handler.sync_replication_slots(cluster) self.assertEqual(errorlog_mock.call_count, 5) ca = errorlog_mock.call_args_list[0][0][1] self.assertTrue("test-3" in ca, "non matching {0}".format(ca)) @@ -613,120 +298,6 @@ class TestPostgresql(unittest.TestCase): with patch('os.rename', Mock(side_effect=OSError)): self.p.move_data_directory() - @patch.object(Postgresql, 'is_running', Mock(return_value=True)) - def test_bootstrap(self): - with patch('subprocess.call', Mock(return_value=1)): - self.assertFalse(self.p.bootstrap({})) - - config = {'users': {'replicator': {'password': 'rep-pass', 'options': ['replication']}}} - - self.p.bootstrap(config) - with open(os.path.join(self.config_dir, 'pg_hba.conf')) as f: - lines = f.readlines() - self.assertTrue('host all all 0.0.0.0/0 md5\n' in lines) - - self.p.config.pop('pg_hba') - config.update({'post_init': '/bin/false', - 'pg_hba': ['host replication replicator 127.0.0.1/32 md5', - 'hostssl all all 0.0.0.0/0 md5', - 'host all all 0.0.0.0/0 md5']}) - self.p.bootstrap(config) - with open(os.path.join(self.data_dir, 'pg_hba.conf')) as f: - lines = f.readlines() - self.assertTrue('host replication replicator 127.0.0.1/32 md5\n' in lines) - - @patch.object(Postgresql, 'cancellable_subprocess_call') - @patch.object(Postgresql, 'get_major_version', Mock(return_value=90600)) - def test_custom_bootstrap(self, mock_cancellable_subprocess_call): - self.p.config.pop('pg_hba') - config = {'method': 'foo', 'foo': {'command': 'bar'}} - - mock_cancellable_subprocess_call.return_value = 1 - self.assertFalse(self.p.bootstrap(config)) - - mock_cancellable_subprocess_call.return_value = 0 - with patch('multiprocessing.Process', Mock(side_effect=Exception("42"))),\ - patch('os.path.isfile', Mock(return_value=True)),\ - patch('os.unlink', Mock()),\ - patch.object(Postgresql, 'save_configuration_files', Mock()),\ - patch.object(Postgresql, 'restore_configuration_files', Mock()),\ - patch.object(Postgresql, 'write_recovery_conf', Mock()): - with self.assertRaises(Exception) as e: - self.p.bootstrap(config) - self.assertEqual(str(e.exception), '42') - - config['foo']['recovery_conf'] = {'foo': 'bar'} - - with self.assertRaises(Exception) as e: - self.p.bootstrap(config) - self.assertEqual(str(e.exception), '42') - - mock_cancellable_subprocess_call.side_effect = Exception - self.assertFalse(self.p.bootstrap(config)) - - @patch('time.sleep', Mock()) - @patch('os.unlink', Mock()) - @patch('shutil.copy', Mock()) - @patch('os.path.isfile', Mock(return_value=True)) - @patch.object(Postgresql, 'run_bootstrap_post_init', Mock(return_value=True)) - @patch.object(Postgresql, '_custom_bootstrap', Mock(return_value=True)) - @patch.object(Postgresql, 'start', Mock(return_value=True)) - @patch.object(Postgresql, 'get_major_version', Mock(return_value=110000)) - def test_post_bootstrap(self): - config = {'method': 'foo', 'foo': {'command': 'bar'}} - self.p.bootstrap(config) - - task = CriticalTask() - with patch.object(Postgresql, 'create_or_update_role', Mock(side_effect=Exception)): - self.p.post_bootstrap({}, task) - self.assertFalse(task.result) - - self.p.config.pop('pg_hba') - self.p.post_bootstrap({}, task) - self.assertTrue(task.result) - - self.p.bootstrap(config) - with patch.object(Postgresql, 'pending_restart', PropertyMock(return_value=True)), \ - patch.object(Postgresql, 'restart', Mock()) as mock_restart: - self.p.post_bootstrap({}, task) - mock_restart.assert_called_once() - - self.p.bootstrap(config) - self.p.set_state('stopped') - self.p.reload_config({'authentication': {'superuser': {'username': 'p', 'password': 'p'}, - 'replication': {'username': 'r', 'password': 'r'}, - 'rewind': {'username': 'rw', 'password': 'rw'}}, - 'listen': '*', 'retry_timeout': 10, 'parameters': {'wal_level': '', 'hba_file': 'foo'}}) - with patch.object(Postgresql, 'restart', Mock()) as mock_restart: - self.p.post_bootstrap({}, task) - mock_restart.assert_called_once() - - @patch.object(Postgresql, 'cancellable_subprocess_call') - def test_run_bootstrap_post_init(self, mock_cancellable_subprocess_call): - mock_cancellable_subprocess_call.return_value = 1 - self.assertFalse(self.p.run_bootstrap_post_init({'post_init': '/bin/false'})) - - mock_cancellable_subprocess_call.return_value = 0 - self.p._superuser.pop('username') - self.assertTrue(self.p.run_bootstrap_post_init({'post_init': '/bin/false'})) - mock_cancellable_subprocess_call.assert_called() - args, kwargs = mock_cancellable_subprocess_call.call_args - self.assertTrue('PGPASSFILE' in kwargs['env']) - self.assertEqual(args[0], ['/bin/false', 'postgres://127.0.0.2:5432/postgres']) - - mock_cancellable_subprocess_call.reset_mock() - self.p._local_address.pop('host') - self.assertTrue(self.p.run_bootstrap_post_init({'post_init': '/bin/false'})) - mock_cancellable_subprocess_call.assert_called() - self.assertEqual(mock_cancellable_subprocess_call.call_args[0][0], ['/bin/false', 'postgres://:5432/postgres']) - - mock_cancellable_subprocess_call.side_effect = OSError - self.assertFalse(self.p.run_bootstrap_post_init({'post_init': '/bin/false'})) - - @patch('patroni.postgresql.Postgresql.create_replica', Mock(return_value=0)) - def test_clone(self): - self.p.clone(self.leader) - @patch('os.listdir', Mock(return_value=['recovery.conf'])) @patch('os.path.exists', Mock(return_value=True)) def test_get_postgres_role_from_data_directory(self): @@ -739,12 +310,12 @@ class TestPostgresql(unittest.TestCase): except OSError: if os.name == 'nt': # os.symlink under Windows needs admin rights skip it pass - os.makedirs(os.path.join(self.data_dir, 'foo')) - _symlink('foo', os.path.join(self.data_dir, 'pg_wal')) + os.makedirs(os.path.join(self.p.data_dir, 'foo')) + _symlink('foo', os.path.join(self.p.data_dir, 'pg_wal')) self.p.remove_data_directory() - open(self.data_dir, 'w').close() + open(self.p.data_dir, 'w').close() self.p.remove_data_directory() - _symlink('unexisting', self.data_dir) + _symlink('unexisting', self.p.data_dir) with patch('os.unlink', Mock(side_effect=OSError)): self.p.remove_data_directory() self.p.remove_data_directory() @@ -769,26 +340,26 @@ class TestPostgresql(unittest.TestCase): @patch('os.path.isfile', Mock(return_value=True)) @patch('shutil.copy', Mock(side_effect=IOError)) def test_save_configuration_files(self): - self.p.save_configuration_files() + self.p.config.save_configuration_files() @patch('os.path.isfile', Mock(side_effect=[False, True])) @patch('shutil.copy', Mock(side_effect=IOError)) def test_restore_configuration_files(self): - self.p.restore_configuration_files() + self.p.config.restore_configuration_files() def test_can_create_replica_without_replication_connection(self): - self.p.config['create_replica_method'] = [] + self.p.config._config['create_replica_method'] = [] self.assertFalse(self.p.can_create_replica_without_replication_connection()) - self.p.config['create_replica_method'] = ['wale', 'basebackup'] - self.p.config['wale'] = {'command': 'foo', 'no_master': 1} + self.p.config._config['create_replica_method'] = ['wale', 'basebackup'] + self.p.config._config['wale'] = {'command': 'foo', 'no_master': 1} self.assertTrue(self.p.can_create_replica_without_replication_connection()) def test_replica_method_can_work_without_replication_connection(self): self.assertFalse(self.p.replica_method_can_work_without_replication_connection('basebackup')) self.assertFalse(self.p.replica_method_can_work_without_replication_connection('foobar')) - self.p.config['foo'] = {'command': 'bar', 'no_master': 1} + self.p.config._config['foo'] = {'command': 'bar', 'no_master': 1} self.assertTrue(self.p.replica_method_can_work_without_replication_connection('foo')) - self.p.config['foo'] = {'command': 'bar'} + self.p.config._config['foo'] = {'command': 'bar'} self.assertFalse(self.p.replica_method_can_work_without_replication_connection('foo')) @patch.object(Postgresql, 'is_running', Mock(return_value=True)) @@ -808,7 +379,7 @@ class TestPostgresql(unittest.TestCase): self.p.reload_config(config) parameters['unix_socket_directories'] = '.' self.p.reload_config(config) - self.p.resolve_connection_addresses() + self.p.config.resolve_connection_addresses() @patch.object(Postgresql, '_version_file_exists', Mock(return_value=True)) def test_get_major_version(self): @@ -895,7 +466,7 @@ class TestPostgresql(unittest.TestCase): self.assertEqual(state['sleeps'], 3) with patch.object(Postgresql, 'check_startup_state_changed', Mock(return_value=False)): - self.p.cancel() + self.p.cancellable.cancel() self.p._state = 'starting' self.assertIsNone(self.p.wait_for_startup()) @@ -935,37 +506,37 @@ class TestPostgresql(unittest.TestCase): def test_set_sync_standby(self): def value_in_conf(): - with open(os.path.join(self.data_dir, 'postgresql.conf')) as f: + with open(os.path.join(self.p.data_dir, 'postgresql.conf')) as f: for line in f: if line.startswith('synchronous_standby_names'): return line.strip() mock_reload = self.p.reload = Mock() - self.p.set_synchronous_standby('n1') + self.p.config.set_synchronous_standby('n1') self.assertEqual(value_in_conf(), "synchronous_standby_names = 'n1'") mock_reload.assert_called() mock_reload.reset_mock() - self.p.set_synchronous_standby('n1') + self.p.config.set_synchronous_standby('n1') mock_reload.assert_not_called() self.assertEqual(value_in_conf(), "synchronous_standby_names = 'n1'") - self.p.set_synchronous_standby('n2') + self.p.config.set_synchronous_standby('n2') mock_reload.assert_called() self.assertEqual(value_in_conf(), "synchronous_standby_names = 'n2'") mock_reload.reset_mock() - self.p.set_synchronous_standby(None) + self.p.config.set_synchronous_standby(None) mock_reload.assert_called() self.assertEqual(value_in_conf(), None) def test_get_server_parameters(self): config = {'synchronous_mode': True, 'parameters': {'wal_level': 'hot_standby'}, 'listen': '0'} - self.p.get_server_parameters(config) + self.p.config.get_server_parameters(config) config['synchronous_mode_strict'] = True - self.p.get_server_parameters(config) - self.p.set_synchronous_standby('foo') - self.p.get_server_parameters(config) + self.p.config.get_server_parameters(config) + self.p.config.set_synchronous_standby('foo') + self.p.config.get_server_parameters(config) @patch('time.sleep', Mock()) def test__wait_for_connection_close(self): @@ -1027,18 +598,6 @@ class TestPostgresql(unittest.TestCase): def test_get_master_timeline(self): self.assertEqual(self.p.get_master_timeline(), 1) - def test_cancellable_subprocess_call(self): - self.p.cancel() - self.assertRaises(PostgresException, self.p.cancellable_subprocess_call, communicate_input=None) - - @patch('patroni.postgresql.polling_loop', Mock(return_value=[0, 0])) - def test_cancel(self): - self.p._cancellable = Mock() - self.p._cancellable.returncode = None - self.p.cancel() - type(self.p._cancellable).returncode = PropertyMock(side_effect=[None, -15]) - self.p.cancel() - @patch.object(Postgresql, 'get_postgres_role_from_data_directory', Mock(return_value='replica')) def test__build_effective_configuration(self): with patch.object(Postgresql, 'controldata', @@ -1046,10 +605,6 @@ class TestPostgresql(unittest.TestCase): 'max_worker_processes setting': '20', 'max_prepared_xacts setting': '100', 'max_locks_per_xact setting': '100'})): - self.p.cancel() + self.p.cancellable.cancel() self.assertFalse(self.p.start()) self.assertTrue(self.p.pending_restart) - - @patch.object(Postgresql, 'controldata', Mock(return_value={"Latest checkpoint's TimeLineID": 1})) - def test_check_for_checkpoint_after_promote(self): - self.p.check_for_checkpoint_after_promote() diff --git a/tests/test_postmaster.py b/tests/test_postmaster.py index 6929f55e..75dff566 100644 --- a/tests/test_postmaster.py +++ b/tests/test_postmaster.py @@ -2,7 +2,7 @@ import psutil import unittest from mock import Mock, patch, mock_open -from patroni.postmaster import PostmasterProcess +from patroni.postgresql.postmaster import PostmasterProcess from six.moves import builtins @@ -26,7 +26,7 @@ class TestPostmasterProcess(unittest.TestCase): @patch('psutil.Process.create_time') @patch('psutil.Process.__init__') - @patch('patroni.postmaster.PostmasterProcess._read_postmaster_pidfile') + @patch.object(PostmasterProcess, '_read_postmaster_pidfile') def test_from_pidfile(self, mock_read, mock_init, mock_create_time): mock_init.side_effect = psutil.NoSuchProcess(123) mock_read.return_value = {} diff --git a/tests/test_rewind.py b/tests/test_rewind.py new file mode 100644 index 00000000..90a459f7 --- /dev/null +++ b/tests/test_rewind.py @@ -0,0 +1,107 @@ +from mock import Mock, PropertyMock, patch + +from patroni.postgresql import Postgresql +from patroni.postgresql.cancellable import CancellableSubprocess +from patroni.postgresql.rewind import Rewind + +from . import BaseTestPostgresql, MockCursor, psycopg2_connect + + +@patch('subprocess.call', Mock(return_value=0)) +@patch('psycopg2.connect', psycopg2_connect) +class TestRewind(BaseTestPostgresql): + + def setUp(self): + super(TestRewind, self).setUp() + self.r = Rewind(self.p) + + def test_can_rewind(self): + with patch.object(Postgresql, 'controldata', Mock(return_value={'wal_log_hints setting': 'on'})): + self.assertTrue(self.r.can_rewind) + with patch('subprocess.call', Mock(return_value=1)): + self.assertFalse(self.r.can_rewind) + with patch('subprocess.call', side_effect=OSError): + self.assertFalse(self.r.can_rewind) + self.p.config._config['use_pg_rewind'] = False + self.assertFalse(self.r.can_rewind) + + @patch.object(CancellableSubprocess, 'call') + def test_pg_rewind(self, mock_cancellable_subprocess_call): + r = {'user': '', 'host': '', 'port': '', 'database': '', 'password': ''} + mock_cancellable_subprocess_call.return_value = 0 + self.assertTrue(self.r.pg_rewind(r)) + mock_cancellable_subprocess_call.side_effect = OSError + self.assertFalse(self.r.pg_rewind(r)) + + @patch.object(Rewind, 'can_rewind', PropertyMock(return_value=True)) + def test__get_local_timeline_lsn(self): + self.r.trigger_check_diverged_lsn() + with patch.object(Postgresql, 'controldata', + Mock(return_value={'Database cluster state': 'shut down in recovery', + 'Minimum recovery ending location': '0/0', + "Min recovery ending loc's timeline": '0'})): + self.r.rewind_or_reinitialize_needed_and_possible(self.leader) + + with patch.object(Postgresql, 'is_running', Mock(return_value=True)): + with patch.object(MockCursor, 'fetchone', Mock(side_effect=[(False, ), Exception])): + self.r.rewind_or_reinitialize_needed_and_possible(self.leader) + + @patch.object(CancellableSubprocess, 'call', Mock(return_value=0)) + @patch.object(Postgresql, 'checkpoint', side_effect=['', '1'],) + @patch.object(Postgresql, 'stop', Mock(return_value=False)) + @patch.object(Postgresql, 'start', Mock()) + def test_execute(self, mock_checkpoint): + self.r.execute(self.leader) + with patch.object(Rewind, 'pg_rewind', Mock(return_value=False)): + mock_checkpoint.side_effect = ['1', '', '', ''] + self.r.execute(self.leader) + self.r.execute(self.leader) + with patch.object(Rewind, 'check_leader_is_not_in_recovery', Mock(return_value=False)): + self.r.execute(self.leader) + self.p.config._config['remove_data_directory_on_rewind_failure'] = False + self.r.trigger_check_diverged_lsn() + self.r.execute(self.leader) + + self.leader.member.data.update(version='1.5.7', checkpoint_after_promote=False) + self.assertIsNone(self.r.execute(self.leader)) + + self.leader.member.data['checkpoint_after_promote'] = True + with patch.object(Rewind, 'check_leader_is_not_in_recovery', Mock(return_value=False)): + self.assertIsNone(self.r.execute(self.leader)) + + with patch.object(Postgresql, 'is_running', Mock(return_value=True)): + self.r.execute(self.leader) + + @patch.object(Postgresql, 'start', Mock()) + @patch.object(Rewind, 'can_rewind', PropertyMock(return_value=True)) + @patch.object(Rewind, '_get_local_timeline_lsn', Mock(return_value=(2, '40159C1'))) + @patch.object(Rewind, 'check_leader_is_not_in_recovery') + def test__check_timeline_and_lsn(self, mock_check_leader_is_not_in_recovery): + mock_check_leader_is_not_in_recovery.return_value = False + self.r.trigger_check_diverged_lsn() + self.assertFalse(self.r.rewind_or_reinitialize_needed_and_possible(self.leader)) + self.leader = self.leader.member + self.assertFalse(self.r.rewind_or_reinitialize_needed_and_possible(self.leader)) + mock_check_leader_is_not_in_recovery.return_value = True + self.assertFalse(self.r.rewind_or_reinitialize_needed_and_possible(self.leader)) + self.r.trigger_check_diverged_lsn() + with patch('psycopg2.connect', Mock(side_effect=Exception)): + self.assertFalse(self.r.rewind_or_reinitialize_needed_and_possible(self.leader)) + self.r.trigger_check_diverged_lsn() + with patch.object(MockCursor, 'fetchone', Mock(side_effect=[('', 2, '0/0'), ('', b'3\t0/40159C0\tn\n')])): + self.assertFalse(self.r.rewind_or_reinitialize_needed_and_possible(self.leader)) + self.r.trigger_check_diverged_lsn() + with patch.object(MockCursor, 'fetchone', Mock(return_value=('', 1, '0/0'))): + with patch.object(Rewind, '_get_local_timeline_lsn', Mock(return_value=(1, '0/0'))): + self.assertFalse(self.r.rewind_or_reinitialize_needed_and_possible(self.leader)) + self.r.trigger_check_diverged_lsn() + self.assertTrue(self.r.rewind_or_reinitialize_needed_and_possible(self.leader)) + + @patch.object(MockCursor, 'fetchone', Mock(side_effect=[(True,), Exception])) + def test_check_leader_is_not_in_recovery(self): + self.r.check_leader_is_not_in_recovery() + self.r.check_leader_is_not_in_recovery() + + @patch.object(Postgresql, 'controldata', Mock(return_value={"Latest checkpoint's TimeLineID": 1})) + def test_check_for_checkpoint_after_promote(self): + self.r.check_for_checkpoint_after_promote() diff --git a/tests/test_wale_restore.py b/tests/test_wale_restore.py index 057cfd96..0b450b33 100644 --- a/tests/test_wale_restore.py +++ b/tests/test_wale_restore.py @@ -6,9 +6,9 @@ from mock import Mock, PropertyMock, patch, mock_open from patroni.scripts import wale_restore from patroni.scripts.wale_restore import WALERestore, main as _main, get_major_version from six.moves import builtins -from test_postgresql import MockConnect, psycopg2_connect from threading import current_thread +from . import MockConnect, psycopg2_connect wale_output_header = ( b'name\tlast_modified\t'