Merge branch 'master' into feature/ctl_scaffolding

This commit is contained in:
Oleksii Kliukin
2016-08-10 11:49:08 +02:00
17 changed files with 606 additions and 280 deletions
+3 -3
View File
@@ -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
+2 -5
View File
@@ -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:
+13 -14
View File
@@ -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",
+177 -77
View File
@@ -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))
+43 -7
View File
@@ -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
+58 -32
View File
@@ -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):
+64 -27
View File
@@ -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
+6 -5
View File
@@ -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
+43 -44
View File
@@ -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')
+5
View File
@@ -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."""
+32 -22
View File
@@ -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)
+87 -12
View File
@@ -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
+34 -15
View File
@@ -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})
+8 -5
View File
@@ -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:[email protected]: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()
+4 -3
View File
@@ -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):
+1 -1
View File
@@ -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):
+26 -8
View File
@@ -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)