From a47a2bceffd1081525b8ca452a24762021e1d599 Mon Sep 17 00:00:00 2001 From: Murat Kabilov Date: Tue, 9 Aug 2016 12:54:48 +0200 Subject: [PATCH 1/5] Manage scheduled restarts using patronictl (#248) Manage scheduled restarts using patronictl --- features/patroni_api.feature | 6 +- patroni/api.py | 5 +- patroni/ctl.py | 231 +++++++++++++++++++++++++---------- patroni/ha.py | 2 + patroni/utils.py | 5 + tests/test_ctl.py | 91 ++++++++++++-- tests/test_etcd.py | 4 +- tests/test_ha.py | 9 +- 8 files changed, 270 insertions(+), 83 deletions(-) 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/api.py b/patroni/api.py index 756031f4..0029d8de 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 @@ -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 diff --git a/patroni/ctl.py b/patroni/ctl.py index e45a77a4..0f590457 100644 --- a/patroni/ctl.py +++ b/patroni/ctl.py @@ -22,6 +22,7 @@ from patroni.config import Config from patroni.dcs import get_dcs as _get_dcs from patroni.exceptions import PatroniException from patroni.postgresql import get_conn_kwargs +from patroni.utils import is_valid_pg_version from prettytable import PrettyTable from six.moves.urllib_parse import urlparse @@ -119,15 +120,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'): @@ -235,6 +238,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) @@ -363,7 +397,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: @@ -398,30 +432,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) @@ -430,31 +440,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') @@ -465,7 +534,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') @@ -474,7 +547,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): @@ -519,16 +592,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} @@ -547,7 +613,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() @@ -566,7 +632,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 @@ -588,17 +654,28 @@ def output_members(cluster, name, fmt='pretty'): 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', @@ -610,17 +687,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 @@ -630,7 +712,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): @@ -646,3 +729,25 @@ def configure(config_file, dcs, namespace): config['dcs_api'] = str(dcs) config['namespace'] = str(namespace) store_config(config, config_file) + + +@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/ha.py b/patroni/ha.py index 4db0ef3f..cf3c4876 100644 --- a/patroni/ha.py +++ b/patroni/ha.py @@ -431,6 +431,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 +440,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/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_ctl.py b/tests/test_ctl.py index 42929a43..71f2ba8f 100644 --- a/tests/test_ctl.py +++ b/tests/test_ctl.py @@ -6,7 +6,7 @@ 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 psycopg2 import OperationalError from test_etcd import etcd_read, requests_get, socket_getaddrinfo, MockResponse @@ -66,7 +66,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 @@ -109,12 +109,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 @@ -205,23 +205,72 @@ 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') @@ -253,9 +302,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']) @@ -292,3 +341,27 @@ class TestCtl(unittest.TestCase): def test_configure(self): result = self.runner.invoke(configure, ['--dcs', 'abc', '-c', 'dummy', '-n', 'bla']) assert result.exit_code == 0 + + @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..5b8c74fd 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 = '' diff --git a/tests/test_ha.py b/tests/test_ha.py index 6a788e5b..0c3c2d0f 100644 --- a/tests/test_ha.py +++ b/tests/test_ha.py @@ -34,7 +34,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) @@ -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() From 5fe74bec3b2edfbd9bf57f78c67fdf4b8809fd68 Mon Sep 17 00:00:00 2001 From: Alexander Kukushkin Date: Wed, 10 Aug 2016 10:15:09 +0200 Subject: [PATCH 2/5] Make different kazoo timeouts depend on loop_wait (#243) * Make different kazoo timeouts dependant on loop_wait ping timeout ~ 1/2 * loop_wait connect_timeout ~ 1/2 * loop_wait Originally these values were calculated from negotiated session timeout and didn't worked very well, because it was taking significant time to figure out that connection is dead and reconnect (up to session timeout) and not giving us time to retry. * Address the code review --- patroni/__init__.py | 7 +--- patroni/api.py | 22 +++++----- patroni/dcs/__init__.py | 17 +++++++- patroni/dcs/zookeeper.py | 91 ++++++++++++++++++++++++++++------------ patroni/ha.py | 4 +- tests/test_api.py | 54 ++++++++++++++---------- tests/test_ha.py | 1 - tests/test_patroni.py | 2 +- tests/test_zookeeper.py | 34 +++++++++++---- 9 files changed, 153 insertions(+), 79 deletions(-) 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 0029d8de..6f3a6ae9 100644 --- a/patroni/api.py +++ b/patroni/api.py @@ -110,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: @@ -134,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) @@ -146,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) @@ -249,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' @@ -268,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: @@ -309,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/dcs/__init__.py b/patroni/dcs/__init__.py index 39049792..e7f850b3 100644 --- a/patroni/dcs/__init__.py +++ b/patroni/dcs/__init__.py @@ -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)) @@ -225,6 +225,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 +270,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/zookeeper.py b/patroni/dcs/zookeeper.py index be6fe45a..9f8c9e12 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): 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=True) - 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=True).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 cf3c4876..79cc8b04 100644 --- a/patroni/ha.py +++ b/patroni/ha.py @@ -296,11 +296,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() 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_ha.py b/tests/test_ha.py index 0c3c2d0f..1b7f937e 100644 --- a/tests/test_ha.py +++ b/tests/test_ha.py @@ -82,7 +82,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 diff --git a/tests/test_patroni.py b/tests/test_patroni.py index af8b36b0..30c1b585 100644 --- a/tests/test_patroni.py +++ b/tests/test_patroni.py @@ -74,7 +74,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_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) From 702ab261a25f55e9db8651e9b9e1013a1d45bc5a Mon Sep 17 00:00:00 2001 From: Alexander Kukushkin Date: Wed, 10 Aug 2016 10:15:55 +0200 Subject: [PATCH 3/5] Use pgkutil to find dcs modules (#253) --- patroni/dcs/__init__.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/patroni/dcs/__init__.py b/patroni/dcs/__init__.py index e7f850b3..28fda908 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() From 413a84836bbc40b8a9525920c68aff2650e553ab Mon Sep 17 00:00:00 2001 From: Alexander Kukushkin Date: Wed, 10 Aug 2016 10:17:37 +0200 Subject: [PATCH 4/5] Update etcd topology only after original request succeed (#254) There is no point to try to update topology until original request is not performed. Also for us it is more important to execute original request rather then keep topology of etcd cluster in sync. In addition to that implement the same retry-timeout logic in the `machines` property which already is used in `api_execute` method. --- patroni/dcs/etcd.py | 90 ++++++++++++++++++++++++++++--------------- tests/test_ctl.py | 5 ++- tests/test_etcd.py | 45 +++++++++++++++------- tests/test_ha.py | 3 +- tests/test_patroni.py | 5 ++- 5 files changed, 98 insertions(+), 50 deletions(-) diff --git a/patroni/dcs/etcd.py b/patroni/dcs/etcd.py index c376ce8b..6f96beb4 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/tests/test_ctl.py b/tests/test_ctl.py index 71f2ba8f..ccbc0b37 100644 --- a/tests/test_ctl.py +++ b/tests/test_ctl.py @@ -8,6 +8,7 @@ from click.testing import CliRunner from mock import patch, Mock 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, \ @@ -33,9 +34,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) diff --git a/tests/test_etcd.py b/tests/test_etcd.py index 5b8c74fd..f3a32a41 100644 --- a/tests/test_etcd.py +++ b/tests/test_etcd.py @@ -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 1b7f937e..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 @@ -114,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, diff --git a/tests/test_patroni.py b/tests/test_patroni.py index 30c1b585..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'] From 8ef7178ddf71b09bab674bcb3921b41677bd7225 Mon Sep 17 00:00:00 2001 From: Alexander Kukushkin Date: Wed, 10 Aug 2016 10:19:52 +0200 Subject: [PATCH 5/5] Refactor code dealing with database connection string/params (#255) In the original code we were parsing/deparsing url-style connection strings back and forth. That was not really resource greedy but rather annoying. Also it was not really obvious how to switch all local connections to unix-sockets (preferably). This commit isolates different use-cases of working with connection strings and minimizes amount of code parsing and deparsing them. Also it introduces one new helper method in the `Member` object - `conn_kwargs`. This method can accept as a parameter dict object with credentials (username and password). As a result it returns dict object which could be used by `psycopg2.connect` or for building connection urls for pg_rewind, pg_basebackup or some other replica creation methods. Params for local connection are builded in the `_local_connect_kwargs` method and could be changed to unix-socket later easily. --- patroni/ctl.py | 22 ++++------ patroni/dcs/__init__.py | 25 +++++++++++- patroni/ha.py | 5 +-- patroni/postgresql.py | 87 ++++++++++++++++++++-------------------- tests/test_ctl.py | 6 ++- tests/test_postgresql.py | 2 +- 6 files changed, 82 insertions(+), 65 deletions(-) diff --git a/patroni/ctl.py b/patroni/ctl.py index 0f590457..cb6978b9 100644 --- a/patroni/ctl.py +++ b/patroni/ctl.py @@ -21,7 +21,6 @@ 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 get_conn_kwargs from patroni.utils import is_valid_pg_version from prettytable import PrettyTable from six.moves.urllib_parse import urlparse @@ -185,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: @@ -219,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 @@ -287,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)) @@ -650,7 +644,7 @@ def output_members(cluster, name, extended=False, 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 = '' diff --git a/patroni/dcs/__init__.py b/patroni/dcs/__init__.py index 28fda908..7b291d55 100644 --- a/patroni/dcs/__init__.py +++ b/patroni/dcs/__init__.py @@ -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 diff --git a/patroni/ha.py b/patroni/ha.py index 79cc8b04..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): 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/tests/test_ctl.py b/tests/test_ctl.py index ccbc0b37..0e4ea51b 100644 --- a/tests/test_ctl.py +++ b/tests/test_ctl.py @@ -235,7 +235,8 @@ class TestCtl(unittest.TestCase): 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']) + 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())): @@ -253,7 +254,8 @@ class TestCtl(unittest.TestCase): # 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']) + 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))): 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):