diff --git a/features/patroni_api.feature b/features/patroni_api.feature index bf1e2fa9..b9199b00 100644 --- a/features/patroni_api.feature +++ b/features/patroni_api.feature @@ -12,7 +12,7 @@ Scenario: check API requests on a stand-alone server Then I receive a response code 503 When I run patronictl.py reinit batman postgres0 --force Then I receive a response returncode 0 - And I receive a response output "reinitialize failed for member postgres0, status code=503, (I am the leader, can not reinitialize)" + And I receive a response output "Failed: reinitialize for member postgres0, status code=503, (I am the leader, can not reinitialize)" When I run patronictl.py failover batman --master postgres0 --force Then I receive a response returncode 1 And I receive a response output "Error: No candidates found to failover to" @@ -54,10 +54,10 @@ Scenario: check API requests for the primary-replica pair And I receive a response role replica When I run patronictl.py reinit batman postgres1 --force Then I receive a response returncode 0 - And I receive a response output "Succesful reinitialize on member postgres1" + And I receive a response output "Success: reinitialize for member postgres1" When I run patronictl.py restart batman postgres0 --force Then I receive a response returncode 0 - And I receive a response output "Succesful restart on member postgres0" + And I receive a response output "Success: restart on member postgres0" And postgres0 role is the primary after 5 seconds When I sleep for 10 seconds Then postgres1 role is the secondary after 15 seconds diff --git a/patroni/__init__.py b/patroni/__init__.py index 1403dacf..37a8e848 100644 --- a/patroni/__init__.py +++ b/patroni/__init__.py @@ -30,7 +30,6 @@ class Patroni(object): self.ha = Ha(self) self.tags = self.get_tags() - self.nap_time = self.config['loop_wait'] self.next_run = time.time() self.scheduled_restart = {} @@ -57,9 +56,7 @@ class Patroni(object): def reload_config(self): try: self.tags = self.get_tags() - self.nap_time = self.config['loop_wait'] - self.dcs.set_ttl(self.config.get('ttl') or 30) - self.dcs.set_retry_timeout(self.config.get('retry_timeout') or self.nap_time) + self.dcs.reload_config(self.config) self.api.reload_config(self.config['restapi']) self.postgresql.reload_config(self.config['postgresql']) except Exception: @@ -82,7 +79,7 @@ class Patroni(object): return self.tags.get('noloadbalance', False) def schedule_next_run(self): - self.next_run += self.nap_time + self.next_run += self.dcs.loop_wait current_time = time.time() nap_time = self.next_run - current_time if nap_time <= 0: diff --git a/patroni/api.py b/patroni/api.py index 756031f4..6f3a6ae9 100644 --- a/patroni/api.py +++ b/patroni/api.py @@ -7,10 +7,9 @@ import time import dateutil.parser import datetime import pytz -import re from patroni.exceptions import PostgresConnectionException -from patroni.utils import deep_compare, patch_config, Retry, RetryFailedError +from patroni.utils import deep_compare, patch_config, Retry, RetryFailedError, is_valid_pg_version from six.moves.BaseHTTPServer import BaseHTTPRequestHandler, HTTPServer from six.moves.socketserver import ThreadingMixIn from threading import Thread @@ -111,7 +110,7 @@ class RestApiHandler(BaseHTTPRequestHandler): self._write_status_response(200, response) def do_GET_config(self): - cluster = self.server.patroni.ha.dcs.cluster or self.server.patroni.ha.dcs.get_cluster() + cluster = self.server.patroni.dcs.cluster or self.server.patroni.dcs.get_cluster() if cluster.config: self._write_json_response(200, cluster.config.data) else: @@ -135,11 +134,11 @@ class RestApiHandler(BaseHTTPRequestHandler): def do_PATCH_config(self): request = self._read_json_content() if request: - cluster = self.server.patroni.ha.dcs.get_cluster() + cluster = self.server.patroni.dcs.get_cluster() data = cluster.config.data.copy() if patch_config(data, request): value = json.dumps(data, separators=(',', ':')) - if not self.server.patroni.ha.dcs.set_config_value(value, cluster.config.index): + if not self.server.patroni.dcs.set_config_value(value, cluster.config.index): return self.send_error(409) self._write_json_response(200, data) @@ -147,10 +146,10 @@ class RestApiHandler(BaseHTTPRequestHandler): def do_PUT_config(self): request = self._read_json_content() if request: - cluster = self.server.patroni.ha.dcs.get_cluster() + cluster = self.server.patroni.dcs.get_cluster() if not deep_compare(request, cluster.config.data): value = json.dumps(request, separators=(',', ':')) - if not self.server.patroni.ha.dcs.set_config_value(value): + if not self.server.patroni.dcs.set_config_value(value): return self.send_error(502) self._write_json_response(200, request) @@ -212,7 +211,7 @@ class RestApiHandler(BaseHTTPRequestHandler): data = "PostgreSQL role should be either master or replica" break elif k == 'postgres_version': - if not re.match(r'[1-9][0-9]?(\.(0|([1-9][0-9]?))){2}$', request[k]): + if not is_valid_pg_version(request[k]): status_code = 400 data = "PostgreSQL version should be in the first.major.minor format" break @@ -250,16 +249,16 @@ class RestApiHandler(BaseHTTPRequestHandler): @check_auth def do_POST_reinitialize(self): - ha = self.server.patroni.ha - cluster = ha.dcs.get_cluster() + patroni = self.server.patroni + cluster = patroni.dcs.get_cluster() if cluster.is_unlocked(): status_code = 503 data = 'Cluster has no leader, can not reinitialize' - elif cluster.leader.name == ha.state_handler.name: + elif cluster.leader.name == patroni.ha.state_handler.name: status_code = 503 data = 'I am the leader, can not reinitialize' else: - action = ha.schedule_reinitialize() + action = patroni.ha.schedule_reinitialize() if action is not None: status_code = 503 data = action + ' already in progress' @@ -269,7 +268,7 @@ class RestApiHandler(BaseHTTPRequestHandler): self._write_response(status_code, data) def poll_failover_result(self, leader, candidate): - timeout = 10 if self.server.patroni.nap_time < 10 else self.server.patroni.nap_time + timeout = max(10, self.server.patroni.dcs.loop_wait) for _ in range(0, timeout*2): time.sleep(1) try: @@ -310,7 +309,7 @@ class RestApiHandler(BaseHTTPRequestHandler): leader = request.get('leader') candidate = request.get('candidate') or request.get('member') scheduled_at = request.get('scheduled_at') - cluster = self.server.patroni.ha.dcs.get_cluster() + cluster = self.server.patroni.dcs.get_cluster() status_code = 500 logger.info("received failover request with leader=%s candidate=%s scheduled_at=%s", diff --git a/patroni/ctl.py b/patroni/ctl.py index 5c98cf28..b272d5ca 100644 --- a/patroni/ctl.py +++ b/patroni/ctl.py @@ -21,7 +21,8 @@ from click import ClickException from patroni.config import Config from patroni.dcs import get_dcs as _get_dcs from patroni.exceptions import PatroniException -from patroni.postgresql import Postgresql, get_conn_kwargs +from patroni.postgresql import Postgresql +from patroni.utils import is_valid_pg_version from prettytable import PrettyTable from six.moves.urllib_parse import urlparse @@ -118,15 +119,17 @@ def auth_header(config): return {'Authorization': 'Basic ' + base64.b64encode(config['restapi']['auth'].encode('utf-8')).decode('utf-8')} -def post_patroni(member, endpoint, content, headers=None): +def request_patroni(member, request_type, endpoint, content=None, headers=None): headers = headers or {} - url = urlparse(member.api_url) - logging.debug(url) + url_parts = urlparse(member.api_url) + logging.debug(url_parts) if 'Content-Type' not in headers: headers['Content-Type'] = 'application/json' - return requests.post('{0}://{1}/{2}'.format(url.scheme, url.netloc, endpoint), - headers=headers, - data=json.dumps(content) if content else None, timeout=60) + + url = '{0}://{1}/{2}'.format(url_parts.scheme, url_parts.netloc, endpoint) + + return getattr(requests, request_type)(url, headers=headers, + data=json.dumps(content) if content else None, timeout=60) def print_output(columns, rows=None, alignment=None, fmt='pretty', header=True, delimiter='\t'): @@ -181,16 +184,6 @@ def watching(w, watch, max_count=None, clear=True): yield 0 -def build_connect_parameters(conn_url, connect_parameters): - params = get_conn_kwargs(conn_url, connect_parameters) - params.update({'fallback_application_name': 'Patroni ctl', 'connect_timeout': '5'}) - if 'database' in connect_parameters: - params['database'] = connect_parameters['database'] - else: - params.pop('database') - return params - - def get_all_members(cluster, role='master'): if role == 'master': if cluster.leader is not None: @@ -215,7 +208,12 @@ def get_cursor(cluster, connect_parameters, role='master', member=None): if member is None: return None - params = build_connect_parameters(member.conn_url, connect_parameters) + params = member.conn_kwargs(connect_parameters) + params.update({'fallback_application_name': 'Patroni ctl', 'connect_timeout': '5'}) + if 'database' in connect_parameters: + params['database'] = connect_parameters['database'] + else: + params.pop('database') conn = psycopg2.connect(**params) conn.autocommit = True @@ -234,6 +232,37 @@ def get_cursor(cluster, connect_parameters, role='master', member=None): return None +def get_members(cluster, cluster_name, member_names, role, force, action): + candidates = {m.name: m for m in cluster.members} + + if not force or role: + output_members(cluster, cluster_name) + + if role: + role_names = [m.name for m in get_all_members(cluster, role)] + if member_names: + member_names = list(set(member_names) & set(role_names)) + if not member_names: + raise PatroniCtlException('No {0} among provided members'.format(role)) + else: + member_names = role_names + + if not member_names and not force: + member_names = [click.prompt('Which member do you want to {0} [{1}]?'.format(action, + ', '.join(candidates.keys())), type=str, default='')] + + for mn in member_names: + if mn not in candidates: + raise PatroniCtlException('{0} is not a member of cluster'.format(mn)) + + if not force: + confirm = click.confirm('Are you sure you want to {0} members {1}?'.format(action, ', '.join(member_names))) + if not confirm: + raise PatroniCtlException('Aborted {0}'.format(action)) + + return [candidates[n] for n in member_names] + + @ctl.command('dsn', help='Generate a dsn for the provided member, defaults to a dsn of the master') @click.option('--role', '-r', help='Give a dsn of any member with this role', type=click.Choice(['master', 'replica', 'any']), default=None) @@ -252,7 +281,7 @@ def dsn(cluster_name, config_file, dcs, role, member): if m is None: raise PatroniCtlException('Can not find a suitable member') - params = get_conn_kwargs(m.conn_url) + params = m.conn_kwargs() click.echo('host={host} port={port}'.format(**params)) @@ -362,7 +391,7 @@ def query_member(cluster, cursor, member, role, command, connect_parameters): def remove(config_file, cluster_name, fmt, dcs): _, dcs, cluster = ctl_load_config(cluster_name, config_file, dcs) - output_members(cluster, cluster_name, fmt) + output_members(cluster, cluster_name, fmt=fmt) confirm = click.prompt('Please confirm the cluster name to remove', type=str) if confirm != cluster_name: @@ -397,30 +426,6 @@ def wait_for_leader(dcs, timeout=30): raise PatroniCtlException('Timeout occured') -def empty_post_to_members(cluster, member_names, force, endpoint, headers=None): - candidates = {m.name: m for m in cluster.members} - - if not member_names: - member_names = [click.prompt('Which member do you want to {0} [{1}]?'.format(endpoint, - ', '.join(candidates.keys())), type=str, default='')] - - for mn in member_names: - if mn not in candidates: - raise PatroniCtlException('{0} is not a member of cluster'.format(mn)) - - if not force: - confirm = click.confirm('Are you sure you want to {0} members {1}?'.format(endpoint, ', '.join(member_names))) - if not confirm: - raise PatroniCtlException('Aborted {0}'.format(endpoint)) - - for mn in member_names: - r = post_patroni(candidates[mn], endpoint, '', headers) - if r.status_code != 200: - click.echo('{0} failed for member {1}, status code={2}, ({3})'.format(endpoint, mn, r.status_code, r.text)) - else: - click.echo('Succesful {0} on member {1}'.format(endpoint, mn)) - - def ctl_load_config(cluster_name, config_file, dcs): config = load_config(config_file, dcs) dcs = get_dcs(config, cluster_name) @@ -429,31 +434,90 @@ def ctl_load_config(cluster_name, config_file, dcs): return config, dcs, cluster +def check_response(response, member_name, action_name, silent_success=False): + if response.status_code >= 400: + click.echo('Failed: {0} for member {1}, status code={2}, ({3})'.format( + action_name, member_name, response.status_code, response.text + )) + elif not silent_success: + click.echo('Success: {0} for member {1}'.format(action_name, member_name)) + + +def parse_scheduled(scheduled): + if (scheduled or 'now') != 'now': + try: + scheduled_at = dateutil.parser.parse(scheduled) + if scheduled_at.tzinfo is None: + scheduled_at = tzlocal.get_localzone().localize(scheduled_at) + except (ValueError, TypeError): + message = 'Unable to parse scheduled timestamp ({0}). It should be in an unambiguous format (e.g. ISO 8601)' + raise PatroniCtlException(message.format(scheduled)) + return scheduled_at + + return None + + @ctl.command('restart', help='Restart cluster member') @click.argument('cluster_name') @click.argument('member_names', nargs=-1) @click.option('--role', '-r', help='Restart only members with this role', default='any', type=click.Choice(['master', 'replica', 'any'])) @click.option('--any', 'p_any', help='Restart a single member only', is_flag=True) +@click.option('--scheduled', help='Timestamp of a scheduled restart in unambiguous format (e.g. ISO 8601)', + default=None) +@click.option('--pg-version', 'version', help='Restart if the PostgreSQL version is less than provided (e.g. 9.5.2)', + default=None) +@click.option('--pending', help='Restart if pending', is_flag=True) @option_config_file @option_force @option_dcs -def restart(cluster_name, member_names, config_file, dcs, force, role, p_any): +def restart(cluster_name, member_names, config_file, dcs, force, role, p_any, scheduled, version, pending): config, dcs, cluster = ctl_load_config(cluster_name, config_file, dcs) - role_names = [m.name for m in get_all_members(cluster, role)] - - if member_names: - member_names = list(set(member_names) & set(role_names)) - else: - member_names = role_names - + members = get_members(cluster, cluster_name, member_names, role, force, 'restart') if p_any: - random.shuffle(member_names) - member_names = member_names[:1] + random.shuffle(members) + members = members[:1] - output_members(cluster, cluster_name) - empty_post_to_members(cluster, member_names, force, 'restart', auth_header(config)) + if version is None and not force: + version = click.prompt('Restart if the PostgreSQL version is less than provided (e.g. 9.5.2) ', + type=str, default='') + + content = {} + if pending: + content['restart_pending'] = True + + if version: + if not is_valid_pg_version(version): + message = 'PostgreSQL version should be in the first.major.minor format' + raise PatroniCtlException(message) + else: + content['postgres_version'] = version + + if scheduled is None and not force: + scheduled = click.prompt('When should the restart take place (e.g. 2015-10-01T14:30) ', type=str, default='now') + + scheduled_at = parse_scheduled(scheduled) + if scheduled_at: + content['schedule'] = scheduled_at.isoformat() + + for member in members: + if 'schedule' in content: + if force and member.data.get('scheduled_restart'): + r = request_patroni(member, 'delete', 'restart', headers=auth_header(config)) + check_response(r, member.name, 'flush scheduled restart', True) + + r = request_patroni(member, 'post', 'restart', content, auth_header(config)) + if r.status_code == 200: + click.echo('Success: restart on member {0}'.format(member.name)) + elif r.status_code == 202: + click.echo('Success: restart scheduled on member {0}'.format(member.name)) + elif r.status_code == 409: + click.echo('Failed: another restart is already scheduled on member {0}'.format(member.name)) + else: + click.echo('Failed: restart for member {0}, status code={1}, ({2})'.format( + member.name, r.status_code, r.text) + ) @ctl.command('reinit', help='Reinitialize cluster member') @@ -464,7 +528,11 @@ def restart(cluster_name, member_names, config_file, dcs, force, role, p_any): @option_dcs def reinit(cluster_name, member_names, config_file, dcs, force): config, dcs, cluster = ctl_load_config(cluster_name, config_file, dcs) - empty_post_to_members(cluster, member_names, force, 'reinitialize', auth_header(config)) + members = get_members(cluster, cluster_name, member_names, None, force, 'reinitialize') + + for member in members: + r = request_patroni(member, 'post', 'reinitialize', headers=auth_header(config)) + check_response(r, member.name, 'reinitialize') @ctl.command('failover', help='Failover to a replica') @@ -473,7 +541,7 @@ def reinit(cluster_name, member_names, config_file, dcs, force): @click.option('--candidate', help='The name of the candidate', default=None) @click.option('--scheduled', help='Timestamp of a scheduled failover in unambiguous format (e.g. ISO 8601)', default=None) -@click.option('--force', is_flag=True) +@option_force @option_config_file @option_dcs def failover(config_file, cluster_name, master, candidate, force, dcs, scheduled): @@ -518,16 +586,9 @@ def failover(config_file, cluster_name, master, candidate, force, dcs, scheduled scheduled = click.prompt('When should the failover take place (e.g. 2015-10-01T14:30) ', type=str, default='now') - if (scheduled or 'now') == 'now': - scheduled_at = None - else: - try: - scheduled_at = dateutil.parser.parse(scheduled) - if scheduled_at.tzinfo is None: - scheduled_at = tzlocal.get_localzone().localize(scheduled_at) - except (ValueError, TypeError): - message = 'Unable to parse scheduled timestamp ({0}). It should be in an unambiguous format (e.g. ISO 8601)' - raise PatroniCtlException(message.format(scheduled)) + scheduled_at = parse_scheduled(scheduled) + + if scheduled_at: scheduled_at = scheduled_at.isoformat() failover_value = {'leader': master, 'candidate': candidate, 'scheduled_at': scheduled_at} @@ -546,7 +607,7 @@ def failover(config_file, cluster_name, master, candidate, force, dcs, scheduled r = None try: - r = post_patroni(cluster.leader.member, 'failover', failover_value, auth_header(config)) + r = request_patroni(cluster.leader.member, 'post', 'failover', failover_value, auth_header(config)) if r.status_code in (200, 202): logging.debug(r) cluster = dcs.get_cluster() @@ -565,7 +626,7 @@ def failover(config_file, cluster_name, master, candidate, force, dcs, scheduled output_members(cluster, cluster_name) -def output_members(cluster, name, fmt='pretty'): +def output_members(cluster, name, extended=False, fmt='pretty'): rows = [] logging.debug(cluster) leader_name = None @@ -583,21 +644,32 @@ def output_members(cluster, name, fmt='pretty'): if m.name == leader_name: leader = '*' - host = get_conn_kwargs(m.conn_url)['host'] + host = m.conn_kwargs()['host'] xlog_location = m.data.get('xlog_location') or 0 lag = '' - if (xlog_location_cluster >= xlog_location): + if xlog_location_cluster >= xlog_location: lag = round((xlog_location_cluster - xlog_location)/1024/1024) - rows.append([ + row = [ name, m.name, host, leader, m.data.get('state', ''), - lag - ]) + lag, + ] + if extended: + value = '' + scheduled_restart = m.data.get('scheduled_restart') + if scheduled_restart: + value = scheduled_restart['schedule'] + if 'postgres_version' in scheduled_restart: + value += ' if version < {0}'.format(scheduled_restart['postgres_version']) + + row.append(value) + + rows.append(row) columns = [ 'Cluster', @@ -609,17 +681,22 @@ def output_members(cluster, name, fmt='pretty'): ] alignment = {'Cluster': 'l', 'Member': 'l', 'Host': 'l', 'Lag in MB': 'r'} + if extended: + columns.append('Scheduled restart') + alignment['Scheduled restart'] = 'l' + print_output(columns, rows, alignment, fmt) @ctl.command('list', help='List the Patroni members for a given Patroni') @click.argument('cluster_names', nargs=-1) +@click.option('--extended', '-e', help='Show some extra information', is_flag=True) @option_config_file @option_format @option_watch @option_watchrefresh @option_dcs -def members(config_file, cluster_names, fmt, watch, w, dcs): +def members(config_file, cluster_names, fmt, watch, w, dcs, extended): if not cluster_names: logging.warning('Listing members: No cluster names were provided') return @@ -629,7 +706,8 @@ def members(config_file, cluster_names, fmt, watch, w, dcs): dcs = get_dcs(config, cluster_name) for _ in watching(w, watch): - output_members(dcs.get_cluster(), cluster_name, fmt) + cluster = dcs.get_cluster() + output_members(cluster, cluster_name, extended, fmt) def timestamp(precision=6): @@ -700,3 +778,25 @@ def scaffold(cluster_name, config_file, dcs, sysid): dcs.delete_cluster() raise PatroniCtlException("Unable to install permanent leader for cluster {0}".format(cluster_name)) click.echo("Cluster {0} has been created successfully".format(cluster_name)) + + +@ctl.command('flush', help='Flush scheduled events') +@click.argument('cluster_name') +@click.argument('member_names', nargs=-1) +@click.argument('target', type=click.Choice(['restart'])) +@click.option('--role', '-r', help='Flush only members with this role', default='any', + type=click.Choice(['master', 'replica', 'any'])) +@option_config_file +@option_force +@option_dcs +def flush(cluster_name, member_names, config_file, dcs, force, role, target): + config, dcs, cluster = ctl_load_config(cluster_name, config_file, dcs) + + members = get_members(cluster, cluster_name, member_names, role, force, 'flush') + for member in members: + if target == 'restart': + if member.data.get('scheduled_restart'): + r = request_patroni(member, 'delete', 'restart', None, auth_header(config)) + check_response(r, member.name, 'flush scheduled restart') + else: + click.echo('No scheduled restart for member {0}'.format(member.name)) diff --git a/patroni/dcs/__init__.py b/patroni/dcs/__init__.py index b26162f9..d7a69997 100644 --- a/patroni/dcs/__init__.py +++ b/patroni/dcs/__init__.py @@ -4,6 +4,7 @@ import importlib import inspect import json import os +import pkgutil import six from collections import namedtuple @@ -32,10 +33,9 @@ def parse_connection_string(value): def get_dcs(config): available_implementations = set() - for module in os.listdir(os.path.dirname(__file__)): - if module.endswith('.py') and not module.startswith('__'): # find module - module_name = module[:-3].lower() - module = importlib.import_module(__package__ + '.' + module[:-3]) + for _, module_name, is_pkg in pkgutil.iter_modules([os.path.dirname(__file__)]): + if not is_pkg: + module = importlib.import_module(__package__ + '.' + module_name) for name in filter(lambda name: not name.startswith('__'), dir(module)): # iterate through module content value = getattr(module, name) name = name.lower() @@ -44,8 +44,8 @@ def get_dcs(config): available_implementations.add(name) if name in config: # which has configuration section in the config file # propagate some parameters - config[name].update({p: config[p] for p in ('namespace', 'name', - 'scope', 'ttl', 'retry_timeout') if p in config}) + config[name].update({p: config[p] for p in ('namespace', 'name', 'scope', + 'loop_wait', 'ttl', 'retry_timeout') if p in config}) return value(config[name]) raise PatroniException("""Can not find suitable configuration of distributed configuration store Available implementations: """ + ', '.join(available_implementations)) @@ -86,6 +86,26 @@ class Member(namedtuple('Member', 'index,name,session,data')): def conn_url(self): return self.data.get('conn_url') + def conn_kwargs(self, auth=None): + ret = self.data.get('conn_kwargs') + if ret: + ret = ret.copy() + else: + r = urlparse(self.conn_url) + ret = { + 'host': r.hostname, + 'port': r.port or 5432, + 'database': r.path[1:] + } + self.data['conn_kwargs'] = ret.copy() + + if auth and isinstance(auth, dict): + if 'username' in auth: + ret['user'] = auth['username'] + if 'password' in auth: + ret['password'] = auth['password'] + return ret + @property def api_url(self): return self.data.get('api_url') @@ -104,7 +124,7 @@ class Member(namedtuple('Member', 'index,name,session,data')): @property def clonefrom(self): - return self.tags.get('clonefrom', False) + return self.tags.get('clonefrom', False) and bool(self.conn_url) class Leader(namedtuple('Leader', 'index,session,member')): @@ -119,6 +139,9 @@ class Leader(namedtuple('Leader', 'index,session,member')): def name(self): return self.member.name + def conn_kwargs(self, auth=None): + return self.member.conn_kwargs(auth) + @property def conn_url(self): return self.member.conn_url @@ -225,6 +248,7 @@ class AbstractDCS(object): self._name = config['name'] self._namespace = '/{0}'.format(config.get('namespace', '/service/').strip('/')) self._base_path = '/'.join([self._namespace, config['scope']]) + self._set_loop_wait(config.get('loop_wait', 10)) self._cluster = None self._cluster_thread_lock = Lock() @@ -269,6 +293,18 @@ class AbstractDCS(object): def set_retry_timeout(self, retry_timeout): """Set the new value for retry_timeout""" + def _set_loop_wait(self, loop_wait): + self._loop_wait = loop_wait + + def reload_config(self, config): + self._set_loop_wait(config['loop_wait']) + self.set_ttl(config['ttl']) + self.set_retry_timeout(config['retry_timeout']) + + @property + def loop_wait(self): + return self._loop_wait + @abc.abstractmethod def _load_cluster(self): """Internally this method should build `Cluster` object which diff --git a/patroni/dcs/etcd.py b/patroni/dcs/etcd.py index 0116aeab..d81cf968 100644 --- a/patroni/dcs/etcd.py +++ b/patroni/dcs/etcd.py @@ -31,24 +31,57 @@ class Client(etcd.Client): self._load_machines_cache() self._allow_reconnect = True + def _build_request_parameters(self): + kwargs = {'headers': self._get_headers(), 'redirect': self.allow_redirect} + + # calculate the number of retries and timeout *per node* + # actual number of retries depends on the number of nodes + etcd_nodes = len(self._machines_cache) + 1 + kwargs['retries'] = 0 if etcd_nodes > 3 else (1 if etcd_nodes > 1 else 2) + + # if etcd_nodes > 3: + # kwargs.update({'retries': 0, 'timeout': float(self.read_timeout)/etcd_nodes}) + # elif etcd_nodes > 1: + # kwargs.update({'retries': 1, 'timeout': self.read_timeout/2.0/etcd_nodes}) + # else: + # kwargs.update({'retries': 2, 'timeout': self.read_timeout/3.0}) + kwargs['timeout'] = self.read_timeout/float(kwargs['retries'] + 1)/etcd_nodes + return kwargs + @property def machines(self): """Original `machines` method(property) of `etcd.Client` class raise exception when it failed to get list of etcd cluster members. This method is being called only when request failed on one of the etcd members during `api_execute` call. - For us it's more important to execute original request rather then get new - topology of etcd cluster. So we will catch this exception and return valid list - of machines with setting flag `self._update_machines_cache` to `!True`. - Later, during next `api_execute` call we will forcefully update machines_cache""" - try: - ret = super(Client, self).machines - random.shuffle(ret) - return ret - except etcd.EtcdException: - if self._update_machines_cache: # We are updating machines_cache - raise # This exception is fatal, we should re-raise it. - self._update_machines_cache = True - return [self._base_uri] + For us it's more important to execute original request rather then get new topology + of etcd cluster. So we will catch this exception and return empty list of machines. + Later, during next `api_execute` call we will forcefully update machines_cache. + + Also this method implements the same timeout-retry logic as `api_execute`, because + the original method was retrying 2 times with the `read_timeout` on each node.""" + + kwargs = self._build_request_parameters() + + while True: + try: + response = self.http.request(self._MGET, self._base_uri + self.version_prefix + '/machines', **kwargs) + machines = [n.strip() for n in self._handle_server_response(response).data.decode('utf-8').split(',')] + logger.debug("Retrieved list of machines: %s", machines) + random.shuffle(machines) + return machines + except Exception as e: + # We can't get the list of machines, if one server is in the + # machines cache, try on it + logger.error("Failed to get list of machines from %s%s: %r", self._base_uri, self.version_prefix, e) + if self._machines_cache: + self._base_uri = self._machines_cache.pop(0) + logger.info("Retrying on %s", self._base_uri) + elif self._update_machines_cache: + raise etcd.EtcdException("Could not get the list of servers, " + "maybe you provided the wrong " + "host(s) to connect to?") + else: + return [] def set_read_timeout(self, timeout): self._read_timeout = timeout @@ -73,8 +106,7 @@ class Client(etcd.Client): if not path.startswith('/'): raise ValueError('Path does not start with /') - kwargs = {'fields': params, 'redirect': self.allow_redirect, - 'headers': self._get_headers(), 'preload_content': False} + kwargs = {'fields': params, 'preload_content': False} if method in [self._MGET, self._MDELETE]: request_executor = self.http.request @@ -88,35 +120,29 @@ class Client(etcd.Client): if self._update_machines_cache: self._load_machines_cache() - if timeout is None: - # calculate the number of retries and timeout *per node* - # actual number of retries depends on the number of nodes - etcd_nodes = len(self._machines_cache) + 1 - kwargs['retries'] = 0 if etcd_nodes > 3 else (1 if etcd_nodes > 1 else 2) + kwargs.update(self._build_request_parameters()) - # if etcd_nodes > 3: - # kwargs.update({'retries': 0, 'timeout': float(self.read_timeout)/etcd_nodes}) - # elif etcd_nodes > 1: - # kwargs.update({'retries': 1, 'timeout': self.read_timeout/2.0/etcd_nodes}) - # else: - # kwargs.update({'retries': 2, 'timeout': self.read_timeout/3.0}) - kwargs['timeout'] = self.read_timeout/float(kwargs['retries'] + 1)/etcd_nodes - else: + if timeout is not None: kwargs.update({'retries': 0, 'timeout': timeout}) response = False try: + some_request_failed = False while not response: response = self._do_http_request(request_executor, method, self._base_uri + path, **kwargs) - if response is False and not self._use_proxies: - self._machines_cache = self.machines + if response is False: + some_request_failed = True + if some_request_failed and not self._use_proxies: + self._machines_cache = self.machines + if self._base_uri in self._machines_cache: self._machines_cache.remove(self._base_uri) - return self._handle_server_response(response) except etcd.EtcdConnectionFailed: self._update_machines_cache = True - raise + if not response: + raise + return self._handle_server_response(response) @staticmethod def get_srv_record(host): diff --git a/patroni/dcs/zookeeper.py b/patroni/dcs/zookeeper.py index 9e018816..68f106a3 100644 --- a/patroni/dcs/zookeeper.py +++ b/patroni/dcs/zookeeper.py @@ -20,7 +20,7 @@ class PatroniSequentialThreadingHandler(SequentialThreadingHandler): self.set_connect_timeout(connect_timeout) def set_connect_timeout(self, connect_timeout): - self._connect_timeout = max(1.0, connect_timeout/4.0) + self._connect_timeout = max(1.0, connect_timeout/2.0) # try to connect to zookeeper node during loop_wait/2 def create_connection(self, *args, **kwargs): """This method is trying to establish connection with one of the zookeeper nodes. @@ -59,8 +59,27 @@ class ZooKeeper(AbstractDCS): self._fetch_cluster = True self._last_leader_operation = 0 + self._orig_kazoo_connect = self._client._connection._connect + self._client._connection._connect = self._kazoo_connect + self._client.start() + def _kazoo_connect(self, host, port): + + """Kazoo is using Ping's to determine health of connection to zookeeper. If there is no + response on Ping after Ping interval (1/2 from read_timeout) it will consider current + connection dead and try to connect to another node. Without this "magic" it was taking + up to 2/3 from session timeout (ttl) to figure out that connection was dead and we had + only small time for reconnect and retry. + + This method is needed to return different value of read_timeout, which is not calculated + from negotiated session timeout but from value of `loop_wait`. And it is 2 sec smaller + than loop_wait, because we can spend up to 2 seconds when calling `touch_member()` and + `write_leader_optime()` methods, which also may hang...""" + + ret = self._orig_kazoo_connect(host, port) + return max(self.loop_wait - 2, 2)*1000, ret[1] + def session_listener(self, state): if state in [KazooState.SUSPENDED, KazooState.LOST]: self.cluster_watcher(None) @@ -69,15 +88,34 @@ class ZooKeeper(AbstractDCS): self._fetch_cluster = True self.event.set() + def reload_config(self, config): + self.set_retry_timeout(config['retry_timeout']) + + loop_wait = config['loop_wait'] + + loop_wait_changed = self._loop_wait != loop_wait + self._loop_wait = loop_wait + self._client.handler.set_connect_timeout(loop_wait) + + # We need to reestablish connection to zookeeper if we want to change + # read_timeout (and Ping interval respectively), because read_timeout + # is calculated in `_kazoo_connect` method. If we are changing ttl at + # the same time, set_ttl method will reestablish connection and return + # `!True`, otherwise we will close existing connection and let kazoo + # open the new one. + if not self.set_ttl(int(config['ttl'] * 1000)) and loop_wait_changed: + self._client._connection._socket.close() + def set_ttl(self, ttl): - ttl = int(ttl * 1000) - # I know, it's weird to access private attributes + """It is not possible to change ttl (session_timeout) in zookeeper without + destroying old session and creating the new one. This method returns `!True` + if session_timeout has been changed (`restart()` has been called).""" if self._client._session_timeout != ttl: self._client._session_timeout = ttl self._client.restart() + return True def set_retry_timeout(self, retry_timeout): - self._client.handler.set_connect_timeout(retry_timeout) self._client._retry.deadline = retry_timeout def get_node(self, key, watch=None): @@ -150,7 +188,7 @@ class ZooKeeper(AbstractDCS): if self._fetch_cluster or self._cluster is None: try: self._client.retry(self._inner_load_cluster) - except: + except Exception: logger.exception('get_cluster') self.cluster_watcher(None) raise ZooKeeperError('ZooKeeper in not responding properly') @@ -195,36 +233,35 @@ class ZooKeeper(AbstractDCS): def touch_member(self, data, ttl=None, permanent=False): cluster = self.cluster member = cluster and ([m for m in cluster.members if m.name == self._name] or [None])[0] - path = self.member_path data = data.encode('utf-8') if member and self._client.client_id is not None and member.session != self._client.client_id[0]: try: - self._client.retry(self._client.delete, path) + self._client.delete_async(self.member_path).get(timeout=1) except NoNodeError: pass except: return False member = None - if member and data == self._my_member_data: - return True - - try: - if member: - self._client.retry(self._client.set, path, data) - else: - self._client.retry(self._client.create, path, data, makepath=True, ephemeral=not permanent) - self._my_member_data = data - return True - except NodeExistsError: + if member: + if data == self._my_member_data: + return True + else: try: - self._client.retry(self._client.set, path, data) + self._client.create_async(self.member_path, data, makepath=True, ephemeral=not permanent).get(timeout=1) self._my_member_data = data return True - except: - logger.exception('touch_member') + except Exception as e: + if not isinstance(e, NodeExistsError): + logger.exception('touch_member') + return False + try: + self._client.set_async(self.member_path, data).get(timeout=1) + self._my_member_data = data + return True except: logger.exception('touch_member') + return False def take_leader(self): @@ -233,17 +270,17 @@ class ZooKeeper(AbstractDCS): def write_leader_optime(self, last_operation): last_operation = last_operation.encode('utf-8') if last_operation != self._last_leader_operation: - self._last_leader_operation = last_operation - path = self.leader_optime_path try: - self._client.retry(self._client.set, path, last_operation) + self._client.set_async(self.leader_optime_path, last_operation).get(timeout=1) + self._last_leader_operation = last_operation except NoNodeError: try: - self._client.retry(self._client.create, path, last_operation, makepath=True) + self._client.create_async(self.leader_optime_path, last_operation, makepath=True).get(timeout=1) + self._last_leader_operation = last_operation except: - logger.exception('Failed to create %s', path) + logger.exception('Failed to create %s', self.leader_optime_path) except: - logger.exception('Failed to update %s', path) + logger.exception('Failed to update %s', self.leader_optime_path) def update_leader(self): return True diff --git a/patroni/ha.py b/patroni/ha.py index 4db0ef3f..8c044ba8 100644 --- a/patroni/ha.py +++ b/patroni/ha.py @@ -141,9 +141,8 @@ class Ha(object): node_to_follow = self._get_node_to_follow(self.cluster) - if not self.state_handler.check_recovery_conf(node_to_follow) or recovery: - self._async_executor.schedule('changing primary_conninfo and restarting') - self._async_executor.run_async(self.state_handler.follow, (node_to_follow, self.cluster.leader, recovery)) + self.state_handler.follow(node_to_follow, self.cluster.leader, recovery, self._async_executor) + return ret def enforce_master_role(self, message, promote_message): @@ -296,11 +295,11 @@ class Ha(object): try: delta = (scheduled_at - now).total_seconds() - if delta > self.patroni.nap_time: + if delta > self.dcs.loop_wait: logger.info('Awaiting %s at %s (in %.0f seconds)', action_name, scheduled_at.isoformat(), delta) return False - elif delta < - int(self.patroni.nap_time * 1.5): + elif delta < - int(self.dcs.loop_wait * 1.5): logger.warning('Found a stale %s value, cleaning up: %s', action_name, scheduled_at.isoformat()) cleanup_fn() @@ -431,6 +430,7 @@ class Ha(object): with self._async_executor: if not self.patroni.scheduled_restart: self.patroni.scheduled_restart = restart_data + self.touch_member() return True return False @@ -439,6 +439,7 @@ class Ha(object): with self._async_executor: if self.patroni.scheduled_restart: self.patroni.scheduled_restart = {} + self.touch_member() ret = True return ret diff --git a/patroni/postgresql.py b/patroni/postgresql.py index d8bb1b60..7dff8b32 100644 --- a/patroni/postgresql.py +++ b/patroni/postgresql.py @@ -10,7 +10,6 @@ import time from patroni.exceptions import PostgresConnectionException, PostgresException from patroni.utils import compare_values, parse_bool, parse_int, Retry, RetryFailedError from six import string_types -from six.moves.urllib_parse import urlparse from threading import Lock logger = logging.getLogger(__name__) @@ -22,24 +21,6 @@ ACTION_ON_RELOAD = "on_reload" ACTION_ON_ROLE_CHANGE = "on_role_change" -def get_conn_kwargs(url, auth=None): - r = urlparse(url) - ret = { - 'host': r.hostname, - 'port': r.port or 5432, - 'database': r.path[1:], - 'fallback_application_name': 'Patroni', - 'connect_timeout': 3, - 'options': '-c statement_timeout=2000', - } - if auth and isinstance(auth, dict): - if 'username' in auth: - ret['user'] = auth['username'] - if 'password' in auth: - ret['password'] = auth['password'] - return ret - - class Postgresql(object): # List of parameters which must be always passed to postmaster as command line options @@ -147,8 +128,8 @@ class Postgresql(object): def resolve_connection_addresses(self): self._local_address = self.get_local_address() - self.connection_string = 'postgres://{connect_address}/{database}'.format( - connect_address=self._connect_address or self._local_address, database=self._database) + self.connection_string = 'postgres://{0}/{1}'.format( + self._connect_address or self._local_address['host'] + ':' + self._local_address['port'], self._database) def pg_ctl(self, cmd, *args, **kwargs): """Builds and executes pg_ctl command @@ -266,7 +247,7 @@ class Postgresql(object): if la.strip().lower() in ('*', '0.0.0.0', '127.0.0.1', 'localhost'): # we are listening on '*' or localhost local_address = 'localhost' # connection via localhost is preferred break - return local_address + ':' + self._server_parameters['port'] + return {'host': local_address, 'port': self._server_parameters['port']} def get_postgres_role_from_data_directory(self): if self.data_directory_empty(): @@ -277,12 +258,21 @@ class Postgresql(object): return 'master' @property - def _connect_kwargs(self): - return get_conn_kwargs('postgres://{0}/{1}'.format(self._local_address, self._database), self._superuser) + 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): if not self._connection or self._connection.closed != 0: - self._connection = psycopg2.connect(**self._connect_kwargs) + self._connection = psycopg2.connect(**self._local_connect_kwargs) self._connection.autocommit = True self.server_version = self._connection.server_version return self._connection @@ -403,7 +393,7 @@ class Postgresql(object): replica_methods = self.config.get('create_replica_method') or ['basebackup'] if clone_member: - r = get_conn_kwargs(clone_member.conn_url, self._replication) + 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) @@ -538,7 +528,7 @@ class Postgresql(object): def checkpoint(self, connect_kwargs=None): check_not_is_in_recovery = connect_kwargs is not None - connect_kwargs = connect_kwargs or self._connect_kwargs + connect_kwargs = connect_kwargs or self._local_connect_kwargs for p in ['connect_timeout', 'options']: connect_kwargs.pop(p, None) try: @@ -624,29 +614,29 @@ class Postgresql(object): with open(os.path.join(self._data_dir, 'pg_hba.conf'), 'a') as f: f.write('\n{}\n'.format('\n'.join(config))) - def primary_conninfo(self, node_to_follow_url): - r = get_conn_kwargs(node_to_follow_url, self._replication) + def primary_conninfo(self, member): + if not (member and member.conn_url): + return None + r = member.conn_kwargs(self._replication) r.update({'application_name': self.name, 'sslmode': 'prefer', 'sslcompression': '1'}) keywords = 'user password host port sslmode sslcompression application_name'.split() return ' '.join('{0}={{{0}}}'.format(kw) for kw in keywords).format(**r) - def check_recovery_conf(self, node_to_follow): + def check_recovery_conf(self, primary_conninfo): if not os.path.isfile(self._recovery_conf): return False - pattern = node_to_follow and node_to_follow.conn_url and self.primary_conninfo(node_to_follow.conn_url) - with open(self._recovery_conf, 'r') as f: for line in f: if line.startswith('primary_conninfo'): - return pattern and (pattern in line) - return not pattern + return primary_conninfo and (primary_conninfo in line) + return not primary_conninfo - def write_recovery_conf(self, node_to_follow): + def write_recovery_conf(self, primary_conninfo): with open(self._recovery_conf, 'w') as f: f.write("standby_mode = 'on'\nrecovery_target_timeline = 'latest'\n") - if node_to_follow and node_to_follow.conn_url: - f.write("primary_conninfo = '{0}'\n".format(self.primary_conninfo(node_to_follow.conn_url))) + if primary_conninfo: + f.write("primary_conninfo = '{0}'\n".format(primary_conninfo)) if self.use_slots: f.write("primary_slot_name = '{0}'\n".format(self.name)) for name, value in self.config.get('recovery_conf', {}).items(): @@ -723,24 +713,33 @@ class Postgresql(object): except OSError: logger.exception("Unable to list %s", status_dir) - def follow(self, member, leader, recovery=False): - if self.check_recovery_conf(member) and not recovery: + def follow(self, member, leader, recovery=False, async_executor=None): + primary_conninfo = self.primary_conninfo(member) + + if self.check_recovery_conf(primary_conninfo) and not recovery: return True + if async_executor: + async_executor.schedule('changing primary_conninfo and restarting') + async_executor.run_async(self._do_follow, (primary_conninfo, leader, recovery)) + else: + self._do_follow(primary_conninfo, leader, recovery) + + def _do_follow(self, primary_conninfo, leader, recovery=False): change_role = self.role == 'master' if change_role: if leader: if leader.name == self.name: self._need_rewind = False - member = None + primary_conninfo = None if self.is_running(): return else: self._need_rewind = bool(leader.conn_url) and self.can_rewind else: self._need_rewind = False - member = None + primary_conninfo = None if self._need_rewind: logger.info("set the rewind flag after demote") @@ -753,7 +752,7 @@ class Postgresql(object): return logger.info('Leader unknown, can not rewind') # prepare pg_rewind connection - r = get_conn_kwargs(leader.conn_url, self._superuser) + r = leader.conn_kwargs(self._superuser) # first make sure that we are really trying to rewind # from the master and run a checkpoint on a t in order to @@ -779,7 +778,7 @@ class Postgresql(object): self.single_user_mode(options=opts) if self.rewind(r) or not self.config.get('remove_data_directory_on_rewind_failure', False): - self.write_recovery_conf(member) + self.write_recovery_conf(primary_conninfo) ret = self.start() else: logger.error('unable to rewind the former master') @@ -788,7 +787,7 @@ class Postgresql(object): ret = True self._need_rewind = False else: - self.write_recovery_conf(member) + self.write_recovery_conf(primary_conninfo) ret = self.restart() self.set_role('replica') diff --git a/patroni/utils.py b/patroni/utils.py index 8b2389f7..3072b42e 100644 --- a/patroni/utils.py +++ b/patroni/utils.py @@ -2,6 +2,7 @@ import os import random import sys import time +import re from patroni.exceptions import PatroniException @@ -226,6 +227,10 @@ def reap_children(): __reap_children = False +def is_valid_pg_version(version): + return re.match(r'[1-9][0-9]?(\.(0|([1-9][0-9]?))){2}$', version) + + class RetryFailedError(PatroniException): """Raised when retrying an operation ultimately failed, after retrying the maximum number of attempts.""" diff --git a/tests/test_api.py b/tests/test_api.py index 51bdf9b3..ea68486c 100644 --- a/tests/test_api.py +++ b/tests/test_api.py @@ -37,7 +37,6 @@ class MockPostgresql(object): class MockHa(object): - dcs = Mock() state_handler = MockPostgresql() @staticmethod @@ -67,10 +66,9 @@ class MockHa(object): class MockPatroni(object): - nap_time = 10 - config = Mock() - postgresql = MockPostgresql() ha = MockHa() + config = Mock() + postgresql = ha.state_handler dcs = Mock() tags = {} version = '0.00' @@ -138,14 +136,14 @@ class TestRestApiHandler(unittest.TestCase): self.assertIsNotNone(MockRestApiServer(RestApiHandler, 'POST /restart HTTP/1.0')) MockRestApiServer(RestApiHandler, 'POST /restart HTTP/1.0\nAuthorization:') - @patch.object(MockHa, 'dcs') + @patch.object(MockPatroni, 'dcs') def test_do_GET_config(self, mock_dcs): mock_dcs.cluster.config.data = {} self.assertIsNotNone(MockRestApiServer(RestApiHandler, 'GET /config')) mock_dcs.cluster.config = None self.assertIsNotNone(MockRestApiServer(RestApiHandler, 'GET /config')) - @patch.object(MockHa, 'dcs') + @patch.object(MockPatroni, 'dcs') def test_do_PATCH_config(self, mock_dcs): config = {'postgresql': {'use_slots': False, 'use_pg_rewind': True, 'parameters': {'wal_level': 'logical'}}} mock_dcs.get_cluster.return_value.config = ClusterConfig.from_node(1, json.dumps(config)) @@ -161,7 +159,7 @@ class TestRestApiHandler(unittest.TestCase): mock_dcs.set_config_value.return_value = False MockRestApiServer(RestApiHandler, request) - @patch.object(MockHa, 'dcs') + @patch.object(MockPatroni, 'dcs') def test_do_PUT_config(self, mock_dcs): mock_dcs.get_cluster.return_value.config = ClusterConfig.from_node(1, '{}') request = 'PUT /config HTTP/1.0' + self._authorization + '\nContent-Length: ' @@ -181,9 +179,11 @@ class TestRestApiHandler(unittest.TestCase): MockRestApiServer(RestApiHandler, 'POST /reload HTTP/1.0' + self._authorization) self.assertIsNotNone(MockRestApiServer(RestApiHandler, 'POST /reload HTTP/1.0' + self._authorization)) + #@patch.object(MockPatroni, 'dcs') def test_do_POST_restart(self): request = 'POST /restart HTTP/1.0' + self._authorization self.assertIsNotNone(MockRestApiServer(RestApiHandler, request)) + with patch.object(MockHa, 'restart', Mock(side_effect=Exception)): MockRestApiServer(RestApiHandler, request) @@ -221,13 +221,14 @@ class TestRestApiHandler(unittest.TestCase): request = make_request('{"role": "master", "postgres_version": "9.5.2"}') MockRestApiServer(RestApiHandler, request) + #@patch.object(MockPatroni, 'dcs') def test_do_DELETE_restart(self): for retval in (True, False): with patch.object(MockHa, 'delete_future_restart', Mock(return_value=retval)): request = 'DELETE /restart HTTP/1.0' + self._authorization self.assertIsNotNone(MockRestApiServer(RestApiHandler, request)) - @patch.object(MockHa, 'dcs') + @patch.object(MockPatroni, 'dcs') def test_do_POST_reinitialize(self, dcs): cluster = dcs.get_cluster.return_value request = 'POST /reinitialize HTTP/1.0' + self._authorization @@ -247,8 +248,9 @@ class TestRestApiHandler(unittest.TestCase): self.assertIsNotNone(MockRestApiServer(RestApiHandler, 'GET /patroni')) @patch('time.sleep', Mock()) - @patch.object(MockHa, 'dcs') + @patch.object(MockPatroni, 'dcs') def test_do_POST_failover(self, dcs): + dcs.loop_wait = 10 cluster = dcs.get_cluster.return_value post = 'POST /failover HTTP/1.0' + self._authorization + '\nContent-Length: ' @@ -273,19 +275,27 @@ class TestRestApiHandler(unittest.TestCase): cluster.members = [Member(0, 'postgresql0', 30, {'api_url': 'http'}), Member(0, 'postgresql2', 30, {'api_url': 'http'})] MockRestApiServer(RestApiHandler, request) - with patch.object(MockPatroni, 'dcs') as d: - cluster = d.get_cluster.return_value - cluster.leader.name = 'postgresql0' - MockRestApiServer(RestApiHandler, request) - cluster.leader.name = 'postgresql2' - MockRestApiServer(RestApiHandler, request) - cluster.leader.name = 'postgresql1' - cluster.failover = None - MockRestApiServer(RestApiHandler, request) - d.get_cluster = Mock(side_effect=Exception) - MockRestApiServer(RestApiHandler, request) - d.manual_failover.return_value = False - MockRestApiServer(RestApiHandler, request) + + cluster.failover = None + MockRestApiServer(RestApiHandler, request) + + dcs.get_cluster.side_effect = [cluster] + MockRestApiServer(RestApiHandler, request) + + cluster2 = cluster.copy() + cluster2.leader.name = 'postgresql0' + dcs.get_cluster.side_effect = [cluster, cluster2] + MockRestApiServer(RestApiHandler, request) + + cluster2.leader.name = 'postgresql2' + dcs.get_cluster.side_effect = [cluster, cluster2] + MockRestApiServer(RestApiHandler, request) + + dcs.get_cluster.side_effect = None + dcs.manual_failover.return_value = False + MockRestApiServer(RestApiHandler, request) + dcs.manual_failover.return_value = True + with patch.object(MockHa, 'fetch_nodes_statuses', Mock(return_value=[])): MockRestApiServer(RestApiHandler, request) diff --git a/tests/test_ctl.py b/tests/test_ctl.py index 4d54fc0f..0bf6eaf8 100644 --- a/tests/test_ctl.py +++ b/tests/test_ctl.py @@ -6,9 +6,9 @@ import unittest from click.testing import CliRunner from mock import patch, Mock -from patroni.ctl import ctl, members, store_config, load_config, output_members, post_patroni, get_dcs, parse_dcs, \ +from patroni.ctl import ctl, members, store_config, load_config, output_members, request_patroni, get_dcs, parse_dcs, \ wait_for_leader, get_all_members, get_any_member, get_cursor, query_member, configure, PatroniCtlException - +from patroni.dcs.etcd import Client 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, \ @@ -36,9 +36,9 @@ class TestCtl(unittest.TestCase): @patch('socket.getaddrinfo', socket_getaddrinfo) def setUp(self): - self.runner = CliRunner() - with patch.object(etcd.Client, 'machines') as mock_machines: + with patch.object(Client, 'machines') as mock_machines: mock_machines.__get__ = Mock(return_value=['http://remotehost:2379']) + self.runner = CliRunner() self.e = get_dcs({'etcd': {'ttl': 30, 'host': 'ok:2379', 'retry_timeout': 10}}, 'foo') @patch('psycopg2.connect', psycopg2_connect) @@ -69,7 +69,7 @@ class TestCtl(unittest.TestCase): self.assertIsNone(output_members(cluster, name='abc', fmt='tsv')) @patch('patroni.ctl.get_dcs') - @patch('patroni.ctl.post_patroni', Mock(return_value=MockResponse())) + @patch('patroni.ctl.request_patroni', Mock(return_value=MockResponse())) def test_failover(self, mock_get_dcs): mock_get_dcs.return_value = self.e mock_get_dcs.return_value.get_cluster = get_cluster_initialized_with_leader @@ -112,12 +112,12 @@ class TestCtl(unittest.TestCase): result = self.runner.invoke(ctl, ['failover', 'dummy'], input='dummy') assert result.exit_code == 1 - with patch('patroni.ctl.post_patroni', Mock(side_effect=Exception)): + with patch('patroni.ctl.request_patroni', Mock(side_effect=Exception)): # Non-responding patroni result = self.runner.invoke(ctl, ['failover', 'dummy'], input='leader\nother\n\ny') assert 'falling back to DCS' in result.output - with patch('patroni.ctl.post_patroni') as mocked: + with patch('patroni.ctl.request_patroni') as mocked: mocked.return_value.status_code = 500 result = self.runner.invoke(ctl, ['failover', 'dummy'], input='leader\nother\n\ny') assert 'Failover failed' in result.output @@ -208,23 +208,74 @@ class TestCtl(unittest.TestCase): @patch('patroni.ctl.get_dcs') def test_restart_reinit(self, mock_get_dcs): mock_get_dcs.return_value.get_cluster = get_cluster_initialized_with_leader - result = self.runner.invoke(ctl, ['restart', 'alpha'], input='y') - assert 'restart failed for' in result.output + result = self.runner.invoke(ctl, ['restart', 'alpha'], input='y\n\nnow') + assert 'Failed: restart for' in result.output assert result.exit_code == 0 result = self.runner.invoke(ctl, ['reinit', 'alpha'], input='y') assert result.exit_code == 1 + # successful reinit + result = self.runner.invoke(ctl, ['reinit', 'alpha', 'other'], input='y') + assert result.exit_code == 0 + # Aborted restart result = self.runner.invoke(ctl, ['restart', 'alpha'], input='N') assert result.exit_code == 1 + result = self.runner.invoke(ctl, ['restart', 'alpha', '--pending', '--force']) + assert result.exit_code == 0 + # Not a member result = self.runner.invoke(ctl, ['restart', 'alpha', 'dummy', '--any'], input='y') assert result.exit_code == 1 + # Wrong pg version + result = self.runner.invoke(ctl, ['restart', 'alpha', '--any', '--pg-version', '9.1'], input='y') + assert 'Error: PostgreSQL version' in result.output + assert result.exit_code == 1 + + with patch('requests.delete', Mock(return_value=MockResponse(500))): + # normal restart, the schedule is actually parsed, but not validated in patronictl + result = self.runner.invoke(ctl, ['restart', 'alpha', 'other', '--force', + '--scheduled', '2300-10-01T14:30']) + assert 'Failed: flush scheduled restart' in result.output + with patch('requests.post', Mock(return_value=MockResponse())): - result = self.runner.invoke(ctl, ['restart', 'alpha'], input='y') + # 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') + 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.0', + '--scheduled', '2300-10-01T14:30'], input='y') + assert result.exit_code == 0 + + # force restart with restart already present + with patch('patroni.ctl.request_patroni', Mock(return_value=MockResponse(204))): + result = self.runner.invoke(ctl, ['restart', 'alpha', 'other', '--force', + '--scheduled', '2300-10-01T14:30']) + assert result.exit_code == 0 + + with patch('requests.post', Mock(return_value=MockResponse(202))): + # 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', '99.0.0', '--scheduled', '2300-10-01T14:30'], input='y' + ) + assert 'Success: restart scheduled' in result.output + assert result.exit_code == 0 + + with patch('requests.post', Mock(return_value=MockResponse(409))): + # 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', '99.0.0', '--scheduled', '2300-10-01T14:30'], input='y' + ) + assert 'Failed: another restart is already' in result.output assert result.exit_code == 0 @patch('patroni.ctl.get_dcs') @@ -256,9 +307,9 @@ class TestCtl(unittest.TestCase): assert cluster.leader.member.name == 'leader' @patch('requests.post', Mock(side_effect=requests.exceptions.ConnectionError('foo'))) - def test_post_patroni(self): + def test_request_patroni(self): member = get_cluster_initialized_with_leader().leader.member - self.assertRaises(requests.exceptions.ConnectionError, post_patroni, member, 'dummy', {}) + self.assertRaises(requests.exceptions.ConnectionError, request_patroni, member, 'post', 'dummy', {}) def test_ctl(self): self.runner.invoke(ctl, ['list']) @@ -318,3 +369,27 @@ class TestCtl(unittest.TestCase): mock_get_dcs.return_value.get_cluster = get_cluster_initialized_with_leader result = self.runner.invoke(ctl, ['scaffold', 'alpha']) assert result.exception + + @patch('patroni.ctl.get_dcs') + def test_list_extended(self, mock_get_dcs): + mock_get_dcs.return_value = self.e + mock_get_dcs.return_value.get_cluster = get_cluster_initialized_with_leader + + result = self.runner.invoke(ctl, ['list', 'dummy', '--extended']) + assert '2100' in result.output + assert 'Scheduled restart' in result.output + + @patch('patroni.ctl.get_dcs') + @patch('requests.delete', Mock(return_value=MockResponse())) + def test_flush(self, mock_get_dcs): + mock_get_dcs.return_value = self.e + mock_get_dcs.return_value.get_cluster = get_cluster_initialized_with_leader + + result = self.runner.invoke(ctl, ['flush', 'dummy', 'restart', '-r', 'master'], input='y') + assert 'No scheduled restart' in result.output + + result = self.runner.invoke(ctl, ['flush', 'dummy', 'restart', '--force']) + assert 'Success: flush scheduled restart' in result.output + with patch.object(requests, 'delete', return_value=MockResponse(404)): + result = self.runner.invoke(ctl, ['flush', 'dummy', 'restart', '--force']) + assert 'Failed: flush scheduled restart' in result.output diff --git a/tests/test_etcd.py b/tests/test_etcd.py index 0fb6ca9c..f3a32a41 100644 --- a/tests/test_etcd.py +++ b/tests/test_etcd.py @@ -13,8 +13,8 @@ from urllib3.exceptions import ReadTimeoutError class MockResponse(object): - def __init__(self): - self.status_code = 200 + def __init__(self, status_code=200): + self.status_code = status_code self.content = '{}' self.ok = True self.text = '' @@ -132,6 +132,10 @@ def socket_getaddrinfo(*args): def http_request(method, url, **kwargs): if url == 'http://localhost:2379/timeout': raise ReadTimeoutError(None, None, None) + if url == 'http://localhost:2379/v2/machines': + ret = MockResponse() + ret.content = 'http://localhost:2379,http://localhost:4001' + return ret if url == 'http://localhost:2379/': return MockResponse() raise socket.error @@ -145,26 +149,39 @@ class TestClient(unittest.TestCase): @patch('dns.resolver.query', dns_query) @patch('requests.get', requests_get) def setUp(self): - with patch.object(etcd.Client, 'machines') as mock_machines: + with patch.object(Client, 'machines') as mock_machines: mock_machines.__get__ = Mock(return_value=['http://localhost:2379', 'http://localhost:4001']) self.client = Client({'discovery_srv': 'test', 'retry_timeout': 3}) self.client.http.request = http_request self.client.http.request_encode_body = http_request - def test_api_execute(self): + def test_machines(self): self.client._base_uri = 'http://localhost:4001' self.client._machines_cache = ['http://localhost:2379'] - self.assertRaises(etcd.EtcdWatchTimedOut, self.client.api_execute, '/timeout', 'POST', params={'wait': 'true'}) - self.client._update_machines_cache = False - self.client.api_execute('/', 'POST', timeout=0) - self.client._update_machines_cache = False + self.assertIsNotNone(self.client.machines) self.client._base_uri = 'http://localhost:4001' self.client._machines_cache = [] - self.assertRaises(etcd.EtcdConnectionFailed, self.client.api_execute, '/', 'GET') - self.assertTrue(self.client._update_machines_cache) - self.assertRaises(etcd.EtcdException, self.client.api_execute, '/', 'GET') - self.assertRaises(etcd.EtcdException, self.client.api_execute, '/', '') + self.assertIsNotNone(self.client.machines) + self.client._update_machines_cache = True + machines = None + try: + machines = self.client.machines + self.assertFail() + except Exception: + self.assertIsNone(machines) + + @patch.object(Client, 'machines') + def test_api_execute(self, mock_machines): + mock_machines.__get__ = Mock(return_value=['http://localhost:2379']) self.assertRaises(ValueError, self.client.api_execute, '', '') + self.client._base_uri = 'http://localhost:4001' + self.client._machines_cache = ['http://localhost:2379'] + self.client.api_execute('/', 'POST', timeout=0) + self.assertRaises(etcd.EtcdWatchTimedOut, self.client.api_execute, '/timeout', 'POST', params={'wait': 'true'}) + self.assertRaises(etcd.EtcdException, self.client.api_execute, '/', '') + self.client._update_machines_cache = True + with patch.object(Client, '_load_machines_cache', Mock(side_effect=etcd.EtcdException)): + self.assertRaises(etcd.EtcdException, self.client.api_execute, '/', 'GET') def test_get_srv_record(self): self.assertEquals(self.client.get_srv_record('blabla'), []) @@ -177,7 +194,9 @@ class TestClient(unittest.TestCase): def test__get_machines_cache_from_dns(self): self.client._get_machines_cache_from_dns('error:2379') - def test__load_machines_cache(self): + @patch.object(Client, 'machines') + def test__load_machines_cache(self, mock_machines): + mock_machines.__get__ = Mock(return_value=['http://localhost:2379']) self.client._config = {} self.assertRaises(Exception, self.client._load_machines_cache) self.client._config = {'discovery_srv': 'blabla'} @@ -201,9 +220,9 @@ class TestEtcd(unittest.TestCase): @patch('dns.resolver.query', dns_query) def test_get_etcd_client(self): - with patch.object(etcd.Client, 'machines') as mock_machines: + with patch.object(Client, 'machines') as mock_machines: mock_machines.__get__ = Mock(side_effect=etcd.EtcdException) - with patch('time.sleep', Mock(side_effect=SleepException())): + with patch('time.sleep', Mock(side_effect=SleepException)): self.assertRaises(SleepException, self.etcd.get_etcd_client, {'discovery_srv': 'test', 'retry_timeout': 10}) diff --git a/tests/test_ha.py b/tests/test_ha.py index 6a788e5b..de458b60 100644 --- a/tests/test_ha.py +++ b/tests/test_ha.py @@ -7,6 +7,7 @@ import unittest from mock import Mock, MagicMock, patch from patroni.config import Config from patroni.dcs import Cluster, Failover, Leader, Member, get_dcs +from patroni.dcs.etcd import Client from patroni.exceptions import DCSError, PostgresException from patroni.ha import Ha from patroni.postgresql import Postgresql @@ -34,7 +35,10 @@ def get_cluster_initialized_without_leader(leader=False, failover=None): 'api_url': 'http://127.0.0.1:8008/patroni', 'xlog_location': 4}) l = Leader(0, 0, m1) if leader else None m2 = Member(0, 'other', 28, {'conn_url': 'postgres://replicator:rep-pass@127.0.0.1:5436/postgres', - 'api_url': 'http://127.0.0.1:8011/patroni', 'tags': {'clonefrom': True}}) + 'api_url': 'http://127.0.0.1:8011/patroni', + 'tags': {'clonefrom': True}, + 'scheduled_restart': {'schedule': "2100-01-01 10:53:07.560445+00:00", + 'postgres_version': '99.0.0'}}) return get_cluster(True, l, [m1, m2], failover) @@ -79,7 +83,6 @@ zookeeper: self.api = Mock() self.tags = {'foo': 'bar'} self.nofailover = None - self.nap_time = 10 self.replicatefrom = None self.api.connection_string = 'http://127.0.0.1:8008' self.clonefrom = None @@ -112,7 +115,7 @@ class TestHa(unittest.TestCase): @patch('socket.getaddrinfo', socket_getaddrinfo) @patch.object(etcd.Client, 'read', etcd_read) def setUp(self): - with patch.object(etcd.Client, 'machines') as mock_machines: + 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, @@ -407,8 +410,8 @@ class TestHa(unittest.TestCase): def test_schedule_future_restart(self): self.ha.patroni.scheduled_restart = {} # do the restart 2 times. The first one should succeed, the second one should fail - self.assertTrue(self.ha.schedule_future_restart({'schedule': str(future_restart_time)})) - self.assertFalse(self.ha.schedule_future_restart({'schedule': str(future_restart_time)})) + self.assertTrue(self.ha.schedule_future_restart({'schedule': future_restart_time})) + self.assertFalse(self.ha.schedule_future_restart({'schedule': future_restart_time})) def test_delete_future_restarts(self): self.ha.delete_future_restart() diff --git a/tests/test_patroni.py b/tests/test_patroni.py index af8b36b0..f9f135cc 100644 --- a/tests/test_patroni.py +++ b/tests/test_patroni.py @@ -6,6 +6,7 @@ import unittest from mock import Mock, patch from patroni.api import RestApiServer from patroni.async_executor import AsyncExecutor +from patroni.dcs.etcd import Client from patroni.exceptions import DCSError from patroni import Patroni, main as _main from six.moves import BaseHTTPServer @@ -31,7 +32,7 @@ class TestPatroni(unittest.TestCase): RestApiServer._BaseServer__is_shut_down = Mock() RestApiServer._BaseServer__shutdown_request = True RestApiServer.socket = 0 - with patch.object(etcd.Client, 'machines') as mock_machines: + with patch.object(Client, 'machines') as mock_machines: mock_machines.__get__ = Mock(return_value=['http://remotehost:2379']) sys.argv = ['patroni.py', 'postgres0.yml'] self.p = Patroni() @@ -44,7 +45,7 @@ class TestPatroni(unittest.TestCase): @patch('time.sleep', Mock(side_effect=SleepException)) @patch.object(etcd.Client, 'delete', Mock()) - @patch.object(etcd.Client, 'machines') + @patch.object(Client, 'machines') def test_patroni_main(self, mock_machines): with patch('subprocess.call', Mock(return_value=1)): sys.argv = ['patroni.py', 'postgres0.yml'] @@ -74,7 +75,7 @@ class TestPatroni(unittest.TestCase): def test_schedule_next_run(self): self.p.ha.dcs.watch = Mock(return_value=True) self.p.schedule_next_run() - self.p.next_run = time.time() - self.p.nap_time - 1 + self.p.next_run = time.time() - self.p.dcs.loop_wait - 1 self.p.schedule_next_run() def test_noloadbalance(self): diff --git a/tests/test_postgresql.py b/tests/test_postgresql.py index bb3b17c6..2ab29124 100644 --- a/tests/test_postgresql.py +++ b/tests/test_postgresql.py @@ -278,7 +278,7 @@ class TestPostgresql(unittest.TestCase): mock_pg_rewind.return_value = True self.p.follow(self.leader, self.leader) - self.assertTrue(self.p.follow(None, None)) # check_recovery_conf... + self.p.follow(None, None) # check_recovery_conf... @patch('subprocess.check_output', Mock(return_value=0, side_effect=pg_controldata_string)) def test_can_rewind(self): diff --git a/tests/test_zookeeper.py b/tests/test_zookeeper.py index a406ec82..d7ee7ad5 100644 --- a/tests/test_zookeeper.py +++ b/tests/test_zookeeper.py @@ -58,11 +58,16 @@ class MockKazooClient(Mock): raise TypeError("Invalid type for 'path' (string expected)") if not isinstance(value, (six.binary_type,)): raise TypeError("Invalid type for 'value' (must be a byte string)") + if value == b'Exception': + raise Exception if path.endswith('/initialize') or path == '/service/test/optime/leader': raise Exception elif value == b'retry' or (value == b'exists' and self.exists): raise NodeExistsError + def create_async(self, path, value=b"", acl=None, ephemeral=False, sequence=False, makepath=False): + return self.create(path, value, acl, ephemeral, sequence, makepath) or Mock() + @staticmethod def set(path, value, version=-1): if not isinstance(path, six.string_types): @@ -80,6 +85,9 @@ class MockKazooClient(Mock): return raise NoNodeError + def set_async(self, path, value, version=-1): + return self.set(path, value, version) or Mock() + def delete(self, path, version=-1, recursive=False): if not isinstance(path, six.string_types): raise TypeError("Invalid type for 'path' (string expected)") @@ -92,6 +100,9 @@ class MockKazooClient(Mock): elif path.endswith('/') or path.endswith('/initialize') or path == '/service/test/members/bar': raise NoNodeError + def delete_async(self, path, version=-1, recursive=False): + return self.delete(path, version, recursive) or Mock() + class TestPatroniSequentialThreadingHandler(unittest.TestCase): @@ -109,16 +120,14 @@ class TestZooKeeper(unittest.TestCase): @patch('patroni.dcs.zookeeper.KazooClient', MockKazooClient) def setUp(self): self.zk = ZooKeeper({'hosts': ['localhost:2181'], 'scope': 'test', - 'name': 'foo', 'ttl': 30, 'retry_timeout': 10}) + 'name': 'foo', 'ttl': 30, 'retry_timeout': 10, 'loop_wait': 10}) def test_session_listener(self): self.zk.session_listener(KazooState.SUSPENDED) - def test_set_ttl(self): - self.zk.set_ttl(20) - - def test_set_retry_timeout(self): - self.zk.set_retry_timeout(10) + def test_reload_config(self): + self.zk.reload_config({'ttl': 20, 'retry_timeout': 10, 'loop_wait': 10}) + self.zk.reload_config({'ttl': 20, 'retry_timeout': 10, 'loop_wait': 5}) def test_get_node(self): self.assertIsNone(self.zk.get_node('/no_node')) @@ -165,7 +174,7 @@ class TestZooKeeper(unittest.TestCase): self.zk.touch_member('new') self.zk._name = 'na' self.zk._client.exists = 1 - self.zk.touch_member('exists') + self.zk.touch_member('Exception') self.zk._name = 'bar' self.zk.touch_member('retry') self.zk._fetch_cluster = True @@ -183,8 +192,12 @@ class TestZooKeeper(unittest.TestCase): def test_write_leader_optime(self): self.zk.last_leader_operation = '0' self.zk.write_leader_optime('1') + with patch.object(MockKazooClient, 'create_async', Mock()): + self.zk.write_leader_optime('1') + with patch.object(MockKazooClient, 'set_async', Mock()): + self.zk.write_leader_optime('2') self.zk._base_path = self.zk._base_path.replace('test', 'bla') - self.zk.write_leader_optime('2') + self.zk.write_leader_optime('3') def test_delete_cluster(self): self.assertTrue(self.zk.delete_cluster()) @@ -193,3 +206,8 @@ class TestZooKeeper(unittest.TestCase): self.zk.watch(0) self.zk.event.isSet = lambda: True self.zk.watch(0) + + def test__kazoo_connect(self): + self.zk._client._retry.deadline = 1 + self.zk._orig_kazoo_connect = Mock(return_value=(0, 0)) + self.zk._kazoo_connect(None, None)