Improve quality of code by resolving issues found by quantifiedcode and codacy

This commit is contained in:
Alexander Kukushkin
2016-02-12 12:23:49 +01:00
parent b77908ce58
commit df9b8fed2e
31 changed files with 762 additions and 437 deletions
+2 -1
View File
@@ -6,9 +6,10 @@ python:
install:
- if [[ $TRAVIS_PYTHON_VERSION == 2* ]]; then pip install -r requirements-py2.txt --use-mirrors; fi
- if [[ $TRAVIS_PYTHON_VERSION == 3* ]]; then pip install -r requirements-py3.txt; fi
- pip install coveralls
- pip install coveralls codacy-coverage
script:
- python setup.py test
- python setup.py flake8
after_success:
- coveralls
- python-codacy-coverage -r coverage.xml
+10 -5
View File
@@ -12,17 +12,22 @@ RUN curl https://www.postgresql.org/media/keys/ACCC4CF8.asc | apt-key add -
RUN apt-get update -y
RUN apt-get upgrade -y
ENV PGVERSION 9.4
RUN apt-get install python python-yaml python-requests python-boto postgresql-${PGVERSION} python-dnspython python-kazoo python-pip -y
RUN apt-get install python-dev postgresql-server-dev-${PGVERSION} -y
RUN pip install python-etcd psycopg2
ENV PGVERSION 9.5
RUN apt-get install postgresql-${PGVERSION} postgresql-server-dev-${PGVERSION} -y
RUN apt-get install python python-dev python-pip -y
ADD requirements-py2.txt /requirements-py2.txt
RUN pip install -r /requirements-py2.txt
ENV PATH /usr/lib/postgresql/${PGVERSION}/bin:$PATH
ADD patroni.py /patroni.py
ADD patronictl.py /patronictl.py
ADD patroni/ /patroni
ENV ETCDVERSION 2.0.13
RUN ln -s /patroni.py /usr/local/bin/patroni
RUN ln -s /patronictl.py /usr/local/bin/patronictl
ENV ETCDVERSION 2.2.5
RUN curl -L https://github.com/coreos/etcd/releases/download/v${ETCDVERSION}/etcd-v${ETCDVERSION}-linux-amd64.tar.gz | tar xz -C /bin --strip=1 --wildcards --no-anchored etcd etcdctl
### Setting up a simple script that will serve as an entrypoint
+3
View File
@@ -0,0 +1,3 @@
Alexander Kukushkin <[email protected]>
Feike Steenbergen <[email protected]>
Oleksii Kliukin <[email protected]>
+2 -1
View File
@@ -77,7 +77,8 @@ For an example file, see ``postgres0.yml``. Regarding settings:
- *postgresql*:
- *name*: the name of the Postgres host. Must be unique for the cluster.
- *listen*: IP address + port that Postgres listens to; must be accessible from other nodes in the cluster, if you're using streaming replication.
- *listen*: IP address + port that Postgres listens to; must be accessible from other nodes in the cluster, if you're using streaming replication. Multiple comma-separated addresses are permitted, as long as the port component is appended after to the last one with a colon, i.e. ``listen: 127.0.0.1,127.0.0.2:5432``. The first address from this list will be used by Patroni to establish local connections to the PostgreSQL node.
- *connect\_address*: IP address + port through which Postgres is accessible from other nodes and applications.
- *data\_dir*: file path to initialize and store Postgres data files.
- *maximum\_lag\_on\_failover*: the maximum bytes a follower may lag.
+10 -4
View File
@@ -79,7 +79,12 @@ then
ETCD_CLUSTER="127.0.0.1:4001"
fi
cat > /patroni/postgres.yml <<__EOF__
mkdir -p ~postgres/.config/patroni
cat > ~postgres/.config/patroni/patronictl.yaml <<__EOF__
{dcs_api: 'etcd://${ETCD_CLUSTER}', namespace: /service/}
__EOF__
cat > /patroni/postgres.yaml <<__EOF__
ttl: &ttl 30
loop_wait: &loop_wait 10
@@ -119,14 +124,15 @@ postgresql:
archive_command: 'true'
max_wal_senders: 20
listen_addresses: 0.0.0.0
checkpoint_segments: 64
max_wal_size: 1GB
min_wal_size: 128MB
wal_keep_segments: 64
archive_timeout: 1800s
max_replication_slots: 20
hot_standby: "on"
__EOF__
cat /patroni/postgres.yml
cat /patroni/postgres.yaml
if [ ! -z $CHEAT ]
then
@@ -135,5 +141,5 @@ then
sleep 60
done
else
exec python /patroni.py /patroni/postgres.yml
exec python /patroni.py /patroni/postgres.yaml
fi
+8 -2
View File
@@ -10,17 +10,19 @@ from patroni.ha import Ha
from patroni.postgresql import Postgresql
from patroni.utils import setup_signal_handlers, reap_children
from patroni.zookeeper import ZooKeeper
from .version import __version__
logger = logging.getLogger(__name__)
class Patroni:
class Patroni(object):
def __init__(self, config):
self.nap_time = config['loop_wait']
self.tags = config.get('tags', dict())
self.postgresql = Postgresql(config['postgresql'])
self.dcs = self.get_dcs(self.postgresql.name, config)
self.version = __version__
self.api = RestApiServer(self, config['restapi'])
self.ha = Ha(self)
self.next_run = time.time()
@@ -29,6 +31,10 @@ class Patroni:
def nofailover(self):
return self.tags.get('nofailover', False)
@property
def replicatefrom(self):
return self.tags.get('replicatefrom')
@staticmethod
def get_dcs(name, config):
if 'etcd' in config:
@@ -62,7 +68,7 @@ def main():
setup_signal_handlers()
if len(sys.argv) < 2 or not os.path.isfile(sys.argv[1]):
print('Usage: {} config.yml'.format(sys.argv[0]))
print('Usage: {0} config.yml'.format(sys.argv[0]))
return
with open(sys.argv[1], 'r') as f:
+8 -6
View File
@@ -92,6 +92,7 @@ class RestApiHandler(BaseHTTPRequestHandler):
def do_GET_patroni(self):
response = self.get_postgresql_status(True)
response.update(self.get_tags())
response['patroni'] = {'version': self.server.patroni.version, 'scope': self.server.patroni.postgresql.scope}
self.send_response(200)
self.send_header('Content-Type', 'application/json')
@@ -171,8 +172,8 @@ class RestApiHandler(BaseHTTPRequestHandler):
def do_POST_failover(self):
content_length = int(self.headers.get('content-length', 0))
request = json.loads(self.rfile.read(content_length).decode('utf-8'))
leader = request.get('leader', None)
member = request.get('member', None)
leader = request.get('leader')
member = request.get('member')
cluster = self.server.patroni.ha.dcs.get_cluster()
status_code = 503
data = self.is_failover_possible(cluster, leader, member)
@@ -233,6 +234,7 @@ class RestApiHandler(BaseHTTPRequestHandler):
'state': self.server.patroni.postgresql.state,
'postmaster_start_time': row[0],
'role': 'replica' if row[1] else 'master',
'server_version': self.server.patroni.postgresql.server_version,
'xlog': ({
'received_location': row[3],
'replayed_location': row[4],
@@ -250,8 +252,8 @@ class RestApiHandler(BaseHTTPRequestHandler):
def get_tags(self):
return {'tags': self.server.patroni.tags}
def log_message(self, format, *args):
logger.debug("API thread: " + format % args)
def log_message(self, fmt, *args):
logger.debug("API thread: " + fmt % args)
class RestApiServer(ThreadingMixIn, HTTPServer, Thread):
@@ -268,12 +270,12 @@ class RestApiServer(ThreadingMixIn, HTTPServer, Thread):
# wrap socket with ssl if 'certfile' is defined in a config.yaml
# Sometime it's also needed to pass reference to a 'keyfile'.
options = {option: config[option] for option in ['certfile', 'keyfile'] if option in config}
if options.get('certfile', None):
if options.get('certfile'):
import ssl
self.socket = ssl.wrap_socket(self.socket, server_side=True, **options)
protocol = 'https'
self.connection_string = '{}://{}/patroni'.format(protocol, config.get('connect_address', config['listen']))
self.connection_string = '{0}://{1}/patroni'.format(protocol, config.get('connect_address', config['listen']))
self.patroni = patroni
self.daemon = True
+2 -3
View File
@@ -4,10 +4,9 @@ from threading import Lock, Thread
logger = logging.getLogger(__name__)
class AsyncExecutor:
class AsyncExecutor(object):
def __init__(self):
Lock.__init__(self)
self._busy = False
self._thread_lock = Lock()
self._scheduled_action = None
@@ -51,5 +50,5 @@ class AsyncExecutor:
def __enter__(self):
self._thread_lock.acquire()
def __exit__(self, type, value, traceback):
def __exit__(self, exc_type, exc_value, exc_traceback):
self._thread_lock.release()
+59 -40
View File
@@ -56,7 +56,7 @@ def parse_dcs(dcs):
def load_config(path, dcs):
logging.debug('Loading configuration from file {}'.format(path))
logging.debug('Loading configuration from file %s', path)
config = dict()
try:
with open(path, 'rb') as fd:
@@ -74,8 +74,7 @@ def load_config(path, dcs):
def store_config(config, path):
dir_path = os.path.dirname(path)
if dir_path:
if not os.path.isdir(dir_path):
if dir_path and not os.path.isdir(dir_path):
os.makedirs(dir_path)
with open(path, 'w') as fd:
yaml.dump(config, fd)
@@ -102,19 +101,21 @@ def get_dcs(config, scope):
scheme, hostname, port = map(config.get('dcs', {}).get, ('scheme', 'hostname', 'port'))
if scheme == 'etcd':
return Etcd(name=scope, config={'scope': scope, 'host': '{}:{}'.format(hostname, port)})
return Etcd(name=scope, config={'scope': scope, 'host': '{0}:{1}'.format(hostname, port)})
raise PatroniCtlException('Can not find suitable configuration of distributed configuration store')
def post_patroni(member, endpoint, content, headers={'Content-Type': 'application/json'}):
def post_patroni(member, endpoint, content, headers=None):
url = urlparse(member.api_url)
logging.debug(url)
return requests.post('{}://{}/{}'.format(url.scheme, url.netloc, endpoint), headers=headers,
return requests.post('{0}://{1}/{2}'.format(url.scheme, url.netloc, endpoint),
headers=headers or {'Content-Type': 'application/json'},
data=json.dumps(content), timeout=60)
def print_output(columns, rows=[], alignment=None, format='pretty', header=True, delimiter='\t'):
def print_output(columns, rows=None, alignment=None, format='pretty', header=True, delimiter='\t'):
rows = rows or []
if format == 'pretty':
t = PrettyTable(columns)
for k, v in (alignment or {}).items():
@@ -135,7 +136,7 @@ def print_output(columns, rows=[], alignment=None, format='pretty', header=True,
if columns is not None and header:
click.echo(delimiter.join(columns) + '\n')
for r in rows or []:
for r in rows:
c = [str(c) for c in r]
click.echo(delimiter.join(c))
@@ -168,8 +169,8 @@ def watching(w, watch, max_count=None, clear=True):
yield 0
def build_connect_parameters(conn_url, connect_parameters={}):
params = connect_parameters.copy()
def build_connect_parameters(conn_url, connect_parameters=None):
params = (connect_parameters or {}).copy()
parsed = parseurl(conn_url)
params['host'] = parsed['host']
params['port'] = parsed['port']
@@ -200,12 +201,12 @@ def get_any_member(cluster, role='master', member=None):
return None
def get_cursor(cluster, role='master', member=None, connect_parameters={}):
def get_cursor(cluster, role='master', member=None, connect_parameters=None):
member = get_any_member(cluster=cluster, role=role, member=member)
if member is None:
return None
params = build_connect_parameters(member.conn_url, connect_parameters=connect_parameters)
params = build_connect_parameters(member.conn_url, connect_parameters)
conn = psycopg2.connect(**params)
conn.autocommit = True
@@ -243,7 +244,7 @@ def dsn(cluster_name, config_file, dcs, role, member):
raise PatroniCtlException('Can not find a suitable member')
params = build_connect_parameters(m.conn_url)
click.echo('host={} port={}'.format(params['host'], params['port']))
click.echo('host={host} port={port}'.format(**params))
@ctl.command('query', help='Query a Patroni PostgreSQL member')
@@ -252,6 +253,8 @@ def dsn(cluster_name, config_file, dcs, role, member):
@option_format
@click.option('--format', help='Output format (pretty, json)', default='tsv')
@click.option('--file', '-f', help='Execute the SQL commands from this file', type=click.File('rb'))
@click.option('--password', help='force password prompt', is_flag=True)
@click.option('-U', '--username', help='database user name', type=str)
@option_dcs
@option_watch
@option_watchrefresh
@@ -260,6 +263,7 @@ def dsn(cluster_name, config_file, dcs, role, member):
@click.option('--member', '-m', help='Query a specific member', type=str)
@click.option('--delimiter', help='The column delimiter', default='\t')
@click.option('--command', '-c', help='The SQL commands to execute')
@click.option('-d', '--dbname', help='database name to connect to', type=str)
def query(
cluster_name,
config_file,
@@ -271,6 +275,9 @@ def query(
delimiter,
command,
file,
password,
username,
dbname,
format='tsv',
):
if role is not None and member is not None:
@@ -281,6 +288,17 @@ def query(
if file is not None and command is not None:
raise PatroniCtlException('--file and --command are mutually exclusive options')
if file is None and command is None:
raise PatroniCtlException('You need to specify either --command or --file')
connect_parameters = dict()
if username:
connect_parameters['user'] = username
if password:
connect_parameters['password'] = click.prompt('Password', hide_input=True, type=str)
if dbname:
connect_parameters['database'] = dbname
if file is not None:
command = file.read()
@@ -289,23 +307,24 @@ def query(
cursor = None
for _ in watching(w, watch, clear=False):
output, cursor = query_member(cluster=cluster, cursor=cursor, member=member, role=role, command=command)
output, cursor = query_member(cluster=cluster, cursor=cursor, member=member, role=role, command=command,
connect_parameters=connect_parameters)
print_output(None, output, format=format, delimiter=delimiter)
if cursor is None:
cluster = dcs.get_cluster()
def query_member(cluster, cursor, member, role, command):
def query_member(cluster, cursor, member, role, command, connect_parameters=None):
try:
if cursor is None:
cursor = get_cursor(cluster, role=role, member=member)
cursor = get_cursor(cluster, role=role, member=member, connect_parameters=connect_parameters)
if cursor is None:
if role is None:
message = 'No connection to member {} is available'.format(member)
message = 'No connection to member {0} is available'.format(member)
else:
message = 'No connection to role={} is available'.format(role)
message = 'No connection to role={0} is available'.format(role)
logging.debug(message)
return [[timestamp(0), message]], None
@@ -324,7 +343,7 @@ def query_member(cluster, cursor, member, role, command):
cursor.connection.close()
message = oe.pgcode or oe.pgerror or str(oe)
message = message.replace('\n', ' ')
return [[timestamp(0), 'ERROR, SQLSTATE: {}'.format(message)]], None
return [[timestamp(0), 'ERROR, SQLSTATE: {0}'.format(message)]], None
@ctl.command('remove', help='Remove cluster from DCS')
@@ -336,7 +355,7 @@ def remove(config_file, cluster_name, format, dcs):
config, dcs, cluster = ctl_load_config(cluster_name, config_file, dcs)
if not isinstance(dcs, Etcd):
raise PatroniCtlException('We have not implemented this for DCS of type {}'.format(type(dcs)))
raise PatroniCtlException('We have not implemented this for DCS of type {0}'.format(type(dcs)))
output_members(cluster, format=format)
@@ -346,17 +365,17 @@ def remove(config_file, cluster_name, format, dcs):
message = 'Yes I am aware'
confirm = \
click.prompt('You are about to remove all information in DCS for {}, please type: "{}"'.format(cluster_name,
click.prompt('You are about to remove all information in DCS for {0}, please type: "{1}"'.format(cluster_name,
message), type=str)
if message != confirm:
raise PatroniCtlException('You did not exactly type "{}"'.format(message))
raise PatroniCtlException('You did not exactly type "{0}"'.format(message))
if cluster.leader:
confirm = click.prompt('This cluster currently is healthy. Please specify the master name to continue')
if confirm != cluster.leader.name:
raise PatroniCtlException('You did not specify the current master of the cluster')
dcs.client.delete(dcs._base_path, recursive=True)
dcs.client.delete(dcs.client_path(''), recursive=True)
def wait_for_leader(dcs, timeout=30):
@@ -378,25 +397,25 @@ def empty_post_to_members(cluster, member_names, force, endpoint):
for m in cluster.members:
candidates[m.name] = m
if len(member_names) == 0:
member_names = [click.prompt('Which member do you want to {} [{}]?'.format(endpoint,
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.keys():
raise PatroniCtlException('{} is not a member of cluster'.format(mn))
raise PatroniCtlException('{0} is not a member of cluster'.format(mn))
if not force:
confirm = click.confirm('Are you sure you want to {} members {}?'.format(endpoint, ', '.join(member_names)))
confirm = click.confirm('Are you sure you want to {0} members {1}?'.format(endpoint, ', '.join(member_names)))
if not confirm:
raise PatroniCtlException('Aborted {}'.format(endpoint))
raise PatroniCtlException('Aborted {0}'.format(endpoint))
for mn in member_names:
r = post_patroni(candidates[mn], endpoint, '')
if r.status_code != 200:
click.echo('{} failed for member {}, status code={}, ({})'.format(endpoint, mn, r.status_code, r.text))
click.echo('{0} failed for member {1}, status code={2}, ({3})'.format(endpoint, mn, r.status_code, r.text))
else:
click.echo('Succesful {} on member {}'.format(endpoint, mn))
click.echo('Succesful {0} on member {1}'.format(endpoint, mn))
def ctl_load_config(cluster_name, config_file, dcs):
@@ -421,7 +440,7 @@ def restart(cluster_name, member_names, config_file, dcs, force, role, any):
role_names = [m.name for m in get_all_members(cluster=cluster, role=role)]
if len(member_names) > 0:
if member_names:
member_names = list(set(member_names) & set(role_names))
else:
member_names = role_names
@@ -472,13 +491,13 @@ def failover(config_file, cluster_name, master, candidate, force, dcs):
master = click.prompt('Master', type=str, default=cluster.leader.member.name)
if cluster.leader.member.name != master:
raise PatroniCtlException('Member {} is not the leader of cluster {}'.format(master, cluster_name))
raise PatroniCtlException('Member {0} is not the leader of cluster {1}'.format(master, cluster_name))
candidate_names = [str(m.name) for m in cluster.members if m.name != master]
# We sort the names for consistent output to the client
candidate_names.sort()
if len(candidate_names) == 0:
if not candidate_names:
raise PatroniCtlException('No candidates found to failover to')
if candidate is None and not force:
@@ -488,7 +507,7 @@ def failover(config_file, cluster_name, master, candidate, force, dcs):
raise PatroniCtlException('Failover target and source are the same.')
if candidate and candidate not in candidate_names:
raise PatroniCtlException('Member {} does not exist in cluster {}'.format(candidate, cluster_name))
raise PatroniCtlException('Member {0} does not exist in cluster {1}'.format(candidate, cluster_name))
# By now we have established that the leader exists and the candidate exists
click.echo('Current cluster topology')
@@ -496,12 +515,12 @@ def failover(config_file, cluster_name, master, candidate, force, dcs):
if not force:
a = \
click.confirm('Are you sure you want to failover cluster {}, demoting current master {}?'.format(
click.confirm('Are you sure you want to failover cluster {0}, demoting current master {1}?'.format(
cluster_name, master))
if not a:
raise PatroniCtlException('Aborting failover')
failover_value = '{}:{}'.format(master, candidate or '')
failover_value = '{0}:{1}'.format(master, candidate or '')
t_started = time.time()
r = None
@@ -511,16 +530,16 @@ def failover(config_file, cluster_name, master, candidate, force, dcs):
logging.debug(r)
logging.debug(r.text)
cluster = dcs.get_cluster()
click.echo(timestamp() + ' Failing over to new leader: {}'.format(cluster.leader.member.name))
click.echo(timestamp() + ' Failing over to new leader: {0}'.format(cluster.leader.member.name))
else:
click.echo('Failover failed, details: {}, {}'.format(r.status_code, r.text))
click.echo('Failover failed, details: {0}, {1}'.format(r.status_code, r.text))
return
except:
logging.exception(r)
logging.warning('Failing over to DCS')
click.echo(timestamp() + ' Could not failover using Patroni api, falling back to DCS')
dcs.set_failover_value(failover_value)
click.echo(timestamp() + ' Initialized failover from master {}'.format(master))
click.echo(timestamp() + ' Initialized failover from master {0}'.format(master))
# The failover process should within a minute update the failover key, we will keep watching it until it changes
# or we timeout
cluster = wait_for_leader(dcs, timeout=60)
@@ -589,7 +608,7 @@ def output_members(cluster, name=None, format='pretty'):
@option_watchrefresh
@option_dcs
def members(config_file, cluster_names, format, watch, w, dcs):
if len(cluster_names) == 0:
if not cluster_names:
logging.warning('Listing members: No cluster names were provided')
return
+12 -5
View File
@@ -51,22 +51,26 @@ class Member(namedtuple('Member', 'index,name,session,data')):
else:
try:
data = json.loads(data)
except:
except (TypeError, ValueError):
data = {}
return Member(index, name, session, data)
@property
def conn_url(self):
return self.data.get('conn_url', None)
return self.data.get('conn_url')
@property
def api_url(self):
return self.data.get('api_url', None)
return self.data.get('api_url')
@property
def nofailover(self):
return self.data.get('tags', {}).get('nofailover', False)
@property
def replicatefrom(self):
return self.data.get('tags', {}).get('replicatefrom')
class Leader(namedtuple('Leader', 'index,session,member')):
@@ -107,8 +111,11 @@ class Cluster(namedtuple('Cluster', 'initialize,leader,last_leader_operation,mem
def is_unlocked(self):
return not (self.leader and self.leader.name)
def has_member(self, member_name):
return any(m for m in self.members if m.name == member_name)
class AbstractDCS:
class AbstractDCS(object):
__metaclass__ = abc.ABCMeta
@@ -126,7 +133,7 @@ class AbstractDCS:
i.e.: `zookeeper` for zookeeper, `etcd` for etcd, etc...
"""
self._name = name
self._namespace = '/{}'.format(config.get('namespace', '/service/').strip('/'))
self._namespace = '/{0}'.format(config.get('namespace', '/service/').strip('/'))
self._base_path = '/'.join([self._namespace, config['scope']])
self._cluster = None
+16 -12
View File
@@ -51,7 +51,8 @@ class Client(etcd.Client):
def api_execute(self, path, method, **kwargs):
# Update machines_cache if previous attempt of update has failed
self._update_machines_cache and self._load_machines_cache()
if self._update_machines_cache:
self._load_machines_cache()
try:
return super(Client, self).api_execute(path, method, **kwargs)
except etcd.EtcdConnectionFailed:
@@ -73,7 +74,7 @@ class Client(etcd.Client):
except urllib3.exceptions.TimeoutError:
raise
except Exception as e:
raise etcd.EtcdException('Unable to decode server response: %s' % e)
raise etcd.EtcdException('Unable to decode server response: {0}'.format(e))
return super(Client, self)._result_from_response(response)
def _get_machines_cache_from_srv(self, discovery_srv):
@@ -83,7 +84,7 @@ class Client(etcd.Client):
ret = []
for host, port in self.get_srv_record(discovery_srv):
url = '{}://{}:{}/members'.format(self._protocol, host, port)
url = '{0}://{1}:{2}/members'.format(self._protocol, host, port)
try:
response = requests.get(url, timeout=5)
if response.ok:
@@ -101,10 +102,10 @@ class Client(etcd.Client):
host, port = addr.split(':')
try:
for r in set(socket.getaddrinfo(host, port, socket.AF_INET, socket.SOCK_STREAM, socket.IPPROTO_TCP)):
ret.append('{}://{}:{}'.format(self._protocol, r[4][0], r[4][1]))
ret.append('{0}://{1}:{2}'.format(self._protocol, r[4][0], r[4][1]))
except socket.error:
logger.exception('Can not resolve %s', host)
return list(set(ret)) if ret else ['{}://{}:{}'.format(self._protocol, host, port)]
return list(set(ret)) if ret else ['{0}://{1}:{2}'.format(self._protocol, host, port)]
def _load_machines_cache(self):
"""This method should fill up `_machines_cache` from scratch.
@@ -132,7 +133,9 @@ class Client(etcd.Client):
# After filling up initial list of machines_cache we should ask etcd-cluster about actual list
self._base_uri = self._machines_cache.pop(0)
self._machines_cache = self.machines
self._base_uri in self._machines_cache and self._machines_cache.remove(self._base_uri)
if self._base_uri in self._machines_cache:
self._machines_cache.remove(self._base_uri)
self._update_machines_cache = False
@@ -140,7 +143,7 @@ class Client(etcd.Client):
def catch_etcd_errors(func):
def wrapper(*args, **kwargs):
try:
return not func(*args, **kwargs) is None
return func(*args, **kwargs) is not None
except (RetryFailedError, etcd.EtcdException):
return False
except:
@@ -165,7 +168,8 @@ class Etcd(AbstractDCS):
def retry(self, *args, **kwargs):
return self._retry.copy()(*args, **kwargs)
def get_etcd_client(self, config):
@staticmethod
def get_etcd_client(config):
client = None
while not client:
try:
@@ -185,25 +189,25 @@ class Etcd(AbstractDCS):
nodes = {os.path.relpath(node.key, result.key): node for node in result.leaves}
# get initialize flag
initialize = nodes.get(self._INITIALIZE, None)
initialize = nodes.get(self._INITIALIZE)
initialize = initialize and initialize.value
# get last leader operation
last_leader_operation = nodes.get(self._LEADER_OPTIME, None)
last_leader_operation = nodes.get(self._LEADER_OPTIME)
last_leader_operation = 0 if last_leader_operation is None else int(last_leader_operation.value)
# get list of members
members = [self.member(n) for k, n in nodes.items() if k.startswith(self._MEMBERS) and k.count('/') == 1]
# get leader
leader = nodes.get(self._LEADER, None)
leader = nodes.get(self._LEADER)
if leader:
member = Member(-1, leader.value, None, {})
member = ([m for m in members if m.name == leader.value] or [member])[0]
leader = Leader(leader.modifiedIndex, leader.ttl, member)
# failover key
failover = nodes.get(self._FAILOVER, None)
failover = nodes.get(self._FAILOVER)
if failover:
failover = Failover.from_node(failover.modifiedIndex, failover.value)
+4 -1
View File
@@ -1,3 +1,6 @@
from click import ClickException
class PatroniException(Exception):
"""Parent class for all kind of exceptions related to selected distributed configuration store"""
@@ -13,7 +16,7 @@ class PatroniException(Exception):
return repr(self.value)
class PatroniCtlException(Exception):
class PatroniCtlException(ClickException):
pass
+59 -39
View File
@@ -11,7 +11,7 @@ from multiprocessing.pool import ThreadPool
logger = logging.getLogger(__name__)
class Ha:
class Ha(object):
def __init__(self, patroni):
self.patroni = patroni
@@ -19,12 +19,13 @@ class Ha:
self.dcs = patroni.dcs
self.cluster = None
self.old_cluster = None
self.recovering = False
self._async_executor = AsyncExecutor()
def load_cluster_from_dcs(self):
cluster = self.dcs.get_cluster()
# We want to keep the state of cluster when it was healhy
# We want to keep the state of cluster when it was healthy
if not cluster.is_unlocked() or not self.old_cluster:
self.old_cluster = cluster
self.cluster = cluster
@@ -61,18 +62,18 @@ class Ha:
pass
self.dcs.touch_member(json.dumps(data, separators=(',', ':')))
def copy_backup_from_leader(self, leader):
if self.state_handler.bootstrap(leader):
logger.info('bootstrapped from leader')
def clone(self, leader):
if self.state_handler.bootstrap(cluster_initialized=True, current_leader=leader):
logger.info('bootstrapped from leader' if leader else 'bootstrapped without leader')
else:
self.state_handler.stop('immediate')
self.state_handler.remove_data_directory()
logger.error('failed to bootstrap from leader')
logger.error('failed to bootstrap from leader' if leader else 'failed to bootstrap (without leader)')
def bootstrap(self):
if not self.cluster.is_unlocked(): # cluster already has leader
self._async_executor.schedule('bootstrap from leader')
self._async_executor.run_async(self.copy_backup_from_leader, args=(self.cluster.leader, ))
self._async_executor.run_async(self.clone, args=(self.cluster.leader, ))
return 'trying to bootstrap from leader'
elif not self.cluster.initialize and not self.patroni.nofailover: # no initialize key
if self.dcs.initialize(create_new=True): # race for initialization
@@ -91,42 +92,40 @@ class Ha:
else:
return 'failed to acquire initialize lock'
else:
if self.state_handler.can_create_replica_without_leader():
self._async_executor.run_async(self.clone, args=(None, ))
return "trying to bootstrap without leader"
return 'waiting for leader to bootstrap'
def recover(self):
has_lock = self.has_lock()
# try to see if we are the former master that crashed. If so - we likely need to run pg_rewind
# in order to join the former standby being promoted.
pg_controldata = self.state_handler.controldata()
if not has_lock and pg_controldata and\
if (self.state_handler.role == 'master') and pg_controldata and\
pg_controldata.get('Database cluster state', '') == 'in production': # crashed master
self.state_handler.require_rewind()
self.recovering = True
return self.follow("started as readonly because i had the session lock",
"started as a secondary",
refresh=True, recovery=True)
# XXX: follow the leader calls stop, which might take quite some time.
# perhaps we should run sync asynchronously
# (we still need the exit code from follow_the_leader)
ret = self.state_handler.follow_the_leader(None if has_lock else self.cluster.leader, recovery=True)
if not ret:
if not has_lock:
return 'failed to start postgres'
self.dcs.delete_leader()
self.dcs.reset_cluster()
return 'removed leader key after trying and failing to start postgres'
if not has_lock:
return 'started as a secondary'
logger.info('started as readonly because i had the session lock')
def follow(self, demote_reason, follow_reason, refresh=True, recovery=False):
if refresh:
self.load_cluster_from_dcs()
def follow_the_leader(self, demote_reason, follow_reason, refresh=True):
refresh and self.load_cluster_from_dcs()
ret = demote_reason if self.state_handler.is_leader() else follow_reason
leader = self.cluster.leader
leader = None if (leader and leader.name) == self.state_handler.name else leader
if not self.state_handler.check_recovery_conf(leader):
# determine the node to follow. If replicatefrom tag is set,
# try to follow the node mentioned there, otherwise, follow the leader.
if self.patroni.replicatefrom:
node_to_follow = [m for m in self.cluster.members if m.name == self.patroni.replicatefrom]
node_to_follow = node_to_follow[0] if node_to_follow else self.cluster.leader
else:
node_to_follow = self.cluster.leader
node_to_follow = None if node_to_follow and node_to_follow.name == self.state_handler.name else node_to_follow
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_the_leader, (leader, ))
return ret
self._async_executor.run_async(self.state_handler.follow, (node_to_follow, recovery))
if not recovery and self.state_handler.is_leader() or recovery and self.state_handler.role == 'master':
return demote_reason
return follow_reason
def enforce_master_role(self, message, promote_message):
if self.state_handler.is_leader() or self.state_handler.role == 'master':
@@ -298,13 +297,13 @@ class Ha:
return self.enforce_master_role('acquired session lock as a leader',
'promoted self to leader by acquiring session lock')
else:
return self.follow_the_leader('demoted self due after trying and failing to obtain lock',
return self.follow('demoted self after trying and failing to obtain lock',
'following new leader after trying and failing to obtain lock')
else:
if self.patroni.nofailover:
return self.follow_the_leader('demoting self because I am not allowed to become master',
return self.follow('demoting self because I am not allowed to become master',
'following a different leader because I am not allowed to promote')
return self.follow_the_leader('demoting self because i am not the healthiest node',
return self.follow('demoting self because i am not the healthiest node',
'following a different leader because i am not the healthiest node')
def process_healthy_cluster(self):
@@ -323,7 +322,7 @@ class Ha:
self.load_cluster_from_dcs()
else:
logger.info('does not have lock')
return self.follow_the_leader('demoting self because i do not have the lock and i was a leader',
return self.follow('demoting self because i do not have the lock and i was a leader',
'no action. i am a secondary and i am following a leader', False)
def schedule(self, action):
@@ -352,7 +351,7 @@ class Ha:
def reinitialize(self, cluster):
self.state_handler.stop('immediate')
self.state_handler.remove_data_directory()
self.copy_backup_from_leader(cluster.leader)
self.clone(cluster.leader)
def process_scheduled_action(self):
if self.reinitialize_scheduled():
@@ -377,11 +376,21 @@ class Ha:
else:
return self._async_executor.scheduled_action + ' in progress'
def sysid_valid(self, sysid):
@staticmethod
def sysid_valid(sysid):
# sysid does tv_sec << 32, where tv_sec is the number of seconds sine 1970,
# so even 1 << 32 would have 10 digits.
return str(sysid) and len(str(sysid)) >= 10 and str(sysid).isdigit()
def post_recover(self):
if not self.state_handler.is_running():
if self.has_lock():
self.dcs.delete_leader()
self.dcs.reset_cluster()
return 'removed leader key after trying and failing to start postgres'
return 'failed to start postgres'
return None
def _run_cycle(self):
try:
self.load_cluster_from_dcs()
@@ -395,6 +404,13 @@ class Ha:
if self._async_executor.busy:
return self.handle_long_action_in_progress()
# we've got here, so any async action has finished. Check if we tried to recover and failed
if self.recovering:
self.recovering = False
msg = self.post_recover()
if msg is not None:
return msg
# currently it can trigger only reinitialize
msg = self.process_scheduled_action()
if msg is not None:
@@ -425,6 +441,10 @@ class Ha:
else:
return self.process_healthy_cluster()
finally:
# we might not have a valid PostgreSQL connection here if another thread
# stops PostgreSQL, therefore, we only reload replication slots if no
# asynchronous processes are running (should be always the case for the master)
if not self._async_executor.busy:
self.state_handler.sync_replication_slots(self.cluster)
except DCSError:
logger.error('Error communicating with DCS')
@@ -432,7 +452,7 @@ class Ha:
self.demote(delete_leader=False)
return 'demoted self because DCS is not accessible and i was a leader'
except (psycopg2.Error, PostgresConnectionException):
logger.exception('Error communicating with Postgresql. Will try again later')
logger.exception('Error communicating with PostgreSQL. Will try again later')
def run_cycle(self):
with self._async_executor:
+137 -82
View File
@@ -39,7 +39,7 @@ def parseurl(url):
return ret
class Postgresql:
class Postgresql(object):
def __init__(self, config):
self.config = config
@@ -52,7 +52,7 @@ class Postgresql:
self.superuser = config['superuser']
self.admin = config['admin']
self.initdb_options = config.get('initdb', [])
self.pgpass = config.get('pgpass', None) or os.path.join(os.path.expanduser('~'), 'pgpass')
self.pgpass = config.get('pgpass') or os.path.join(os.path.expanduser('~'), 'pgpass')
self.pg_rewind = config.get('pg_rewind', {})
self.callback = config.get('callbacks', {})
self.use_slots = config.get('use_slots', True)
@@ -61,13 +61,13 @@ class Postgresql:
self.configuration_to_save = (os.path.join(self.data_dir, 'pg_hba.conf'),
os.path.join(self.data_dir, 'postgresql.conf'))
self.postmaster_pid = os.path.join(self.data_dir, 'postmaster.pid')
self.trigger_file = config.get('recovery_conf', {}).get('trigger_file', None) or 'promote'
self.trigger_file = config.get('recovery_conf', {}).get('trigger_file') or 'promote'
self.trigger_file = os.path.abspath(os.path.join(self.data_dir, self.trigger_file))
self._pg_ctl = ['pg_ctl', '-w', '-D', self.data_dir]
self.local_address = self.get_local_address()
connect_address = config.get('connect_address', None) or self.local_address
connect_address = config.get('connect_address') or self.local_address
self.connection_string = 'postgres://{username}:{password}@{connect_address}/postgres'.format(
connect_address=connect_address, **self.replication)
@@ -128,16 +128,25 @@ class Postgresql:
break
return local_address + ':' + self.port
@property
def _connect_kwargs(self):
r = parseurl('postgres://{0}/postgres'.format(self.local_address))
if 'username' in self.superuser:
r['user'] = self.superuser['username']
if 'password' in self.superuser:
r['password'] = self.superuser['password']
return r
def connection(self):
if not self._connection or self._connection.closed != 0:
r = parseurl('postgres://{}/postgres'.format(self.local_address))
self._connection = psycopg2.connect(**r)
self._connection = psycopg2.connect(**self._connect_kwargs)
self._connection.autocommit = True
self.server_version = self._connection.server_version
return self._connection
def _cursor(self):
if not self._cursor_holder or self._cursor_holder.closed or self._cursor_holder.connection.closed != 0:
logger.info("established a new patroni connection to the postgres cluster")
logger.info("establishing a new patroni connection to the postgres cluster")
self._cursor_holder = self.connection().cursor()
return self._cursor_holder
@@ -171,32 +180,36 @@ class Postgresql:
@staticmethod
def initdb_allowed_option(name):
if name in ['pgdata', 'nosync', 'pwfile', 'sync-only']:
raise Exception('{} option for initdb is not allowed'.format(name))
raise Exception('{0} option for initdb is not allowed'.format(name))
return True
def get_initdb_options(self):
options = []
for o in self.initdb_options:
if isinstance(o, string_types) and self.initdb_allowed_option(o):
options.append('--{}'.format(o))
options.append('--{0}'.format(o))
elif isinstance(o, dict):
keys = list(o.keys())
if len(keys) != 1 or not isinstance(keys[0], string_types) or not self.initdb_allowed_option(keys[0]):
raise Exception('Invalid option: {}'.format(o))
options.append('--{}={}'.format(keys[0], o[keys[0]]))
raise Exception('Invalid option: {0}'.format(o))
options.append('--{0}={1}'.format(keys[0], o[keys[0]]))
else:
raise Exception('Unknown type of initdb option: {}'.format(o))
raise Exception('Unknown type of initdb option: {0}'.format(o))
return options
def initialize(self):
self.set_state('initalizing new cluster')
options = self.get_initdb_options()
pwfile = None
if self.superuser and 'username' not in self.superuser and 'password' in self.superuser:
if self.superuser:
if 'username' in self.superuser:
options.append('--username={0}'.format(self.superuser['username']))
if 'password' in self.superuser:
(fd, pwfile) = tempfile.mkstemp()
os.write(fd, self.superuser['password'].encode())
os.close(fd)
options.append('--pwfile={}'.format(pwfile))
options.append('--pwfile={0}'.format(pwfile))
ret = subprocess.call(self._pg_ctl + ['initdb'] + (['-o', ' '.join(options)] if options else [])) == 0
if pwfile:
@@ -208,7 +221,8 @@ class Postgresql:
return ret
def delete_trigger_file(self):
os.path.exists(self.trigger_file) and os.unlink(self.trigger_file)
if os.path.exists(self.trigger_file):
os.unlink(self.trigger_file)
def write_pgpass(self, record):
with open(self.pgpass, 'w') as f:
@@ -219,13 +233,12 @@ class Postgresql:
env['PGPASSFILE'] = self.pgpass
return env
def sync_from_leader(self, leader):
r = parseurl(leader.conn_url)
env = self.write_pgpass(r)
ret = self.create_replica(leader, env) == 0
ret and self.delete_trigger_file()
return ret
def sync_replica(self, leader):
env = self.write_pgpass(parseurl(leader.conn_url)) if leader else os.environ.copy()
if self.create_replica(leader, env) == 0:
self.delete_trigger_file()
return True
return False
@staticmethod
def build_connstring(conn):
@@ -233,16 +246,32 @@ class Postgresql:
>>> Postgresql.build_connstring({'host': '127.0.0.1', 'port': '5432'}) == 'host=127.0.0.1 port=5432'
True
"""
return ' '.join('{}={}'.format(param, val) for param, val in sorted(conn.items()))
return ' '.join('{0}={1}'.format(param, val) for param, val in sorted(conn.items()))
def replica_method_can_work_without_leader(self, method):
return method != 'basebackup' and self.config and self.config.get(method, {}).get('no_master')
def can_create_replica_without_leader(self):
""" go through the replication methods to see if there are ones
that does not require a running leader to create the replica.
"""
replica_methods = self.config.get('create_replica_method', [])
for replica_method in replica_methods:
if self.replica_method_can_work_without_leader(replica_method):
return True
return False
def create_replica(self, leader, env):
# create the replica according to the replica_method
# defined by the user. this is a list, so we need to
# loop through all methods the user supplies
connstring = leader.conn_url
connstring = leader.conn_url if leader else ""
# get list of replica methods from config.
# If there is no configuration key, or no value is specified, use basebackup
replica_methods = self.config.get('create_replica_method') or ['basebackup']
# if we don't have any leader, leave only replica methods that work without it
replica_methods = [r for r in replica_methods if self.replica_method_can_work_without_leader(r)] if not leader \
else replica_methods
# go through them in priority order
ret = 1
for replica_method in replica_methods:
@@ -295,7 +324,7 @@ class Postgresql:
cmd = self.callback[cb_name]
try:
subprocess.Popen(shlex.split(cmd) + [cb_name, self.role, self.scope])
except:
except OSError:
logger.exception('callback %s %s %s %s failed', cmd, cb_name, self.role, self.scope)
return False
return True
@@ -331,7 +360,10 @@ class Postgresql:
if not block_callbacks:
self.set_state('starting')
ret = subprocess.call(self._pg_ctl + ['start', '-o', self.server_options()]) == 0
env = os.environ.copy()
if 'username' in self.superuser:
env['PGUSER'] = self.superuser['username']
ret = subprocess.call(self._pg_ctl + ['start', '-o', self.server_options()], env=env, preexec_fn=os.setsid) == 0
self.set_state('running' if ret else 'start failed')
@@ -339,18 +371,21 @@ class Postgresql:
self.save_configuration_files()
# block_callbacks is used during restart to avoid
# running start/stop callbacks in addition to restart ones
ret and not block_callbacks and self.call_nowait(ACTION_ON_START)
if ret and not block_callbacks:
self.call_nowait(ACTION_ON_START)
return ret
def checkpoint(self, connstring=None):
def checkpoint(self, connect_kwargs=None):
connect_kwargs = connect_kwargs or self._connect_kwargs
for p in ['connect_timeout', 'options']:
connect_kwargs.pop(p, None)
try:
connstring = connstring or 'postgres://{}/postgres'.format(self.local_address)
with psycopg2.connect(connstring) as conn:
with psycopg2.connect(**connect_kwargs) as conn:
conn.autocommit = True
with conn.cursor() as cur:
cur.execute("SET statement_timeout = 0")
cur.execute('CHECKPOINT')
except:
except psycopg2.Error:
logging.exception('Exception during CHECKPOINT')
def stop(self, mode='fast', block_callbacks=False):
@@ -360,8 +395,7 @@ class Postgresql:
# patroni.
self.close_connection()
if not self.is_running():
if not block_callbacks:
if not self.is_running() and not block_callbacks:
self.set_state('stopped')
return True
@@ -382,7 +416,8 @@ class Postgresql:
def reload(self):
ret = subprocess.call(self._pg_ctl + ['reload']) == 0
ret and self.call_nowait(ACTION_ON_RELOAD)
if ret:
self.call_nowait(ACTION_ON_RELOAD)
return ret
def restart(self):
@@ -391,13 +426,13 @@ class Postgresql:
if ret:
self.call_nowait(ACTION_ON_RESTART)
else:
self.set_state('restart failed ({})'.format(self.state))
self.set_state('restart failed ({0})'.format(self.state))
return ret
def server_options(self):
options = "--listen_addresses='{}' --port={}".format(self.listen_addresses, self.port)
options = "--listen_addresses='{0}' --port={1}".format(self.listen_addresses, self.port)
for setting, value in self.server_parameters.items():
options += " --{}='{}'".format(setting, value)
options += " --{0}='{1}'".format(setting, value)
return options
def is_healthy(self):
@@ -407,8 +442,7 @@ class Postgresql:
return True
def check_replication_lag(self, last_leader_operation):
return (last_leader_operation if last_leader_operation else 0) - self.xlog_position() <=\
self.config.get('maximum_lag_on_failover', 0)
return (last_leader_operation or 0) - self.xlog_position() <= self.config.get('maximum_lag_on_failover', 0)
def write_pg_hba(self):
with open(os.path.join(self.data_dir, 'pg_hba.conf'), 'a') as f:
@@ -435,33 +469,34 @@ class Postgresql:
return pattern and (pattern in line)
return not pattern
def write_recovery_conf(self, leader):
def write_recovery_conf(self, leader, bootstrap=False):
with open(self.recovery_conf, 'w') as f:
f.write("""standby_mode = 'on'
recovery_target_timeline = 'latest'
""")
if leader and leader.conn_url:
f.write("""primary_conninfo = '{}'\n""".format(self.primary_conninfo(leader.conn_url)))
f.write("""primary_conninfo = '{0}'\n""".format(self.primary_conninfo(leader.conn_url)))
if self.use_slots:
f.write("""primary_slot_name = '{}'\n""".format(self.name))
f.write("""primary_slot_name = '{0}'\n""".format(self.name))
if (leader and leader.conn_url) or bootstrap:
for name, value in self.config.get('recovery_conf', {}).items():
f.write("{} = '{}'\n".format(name, value))
f.write("{0} = '{1}'\n".format(name, value))
def rewind(self, leader):
# prepare pg_rewind connection
r = parseurl(leader.conn_url)
r.update(self.pg_rewind)
r['user'] = r['username']
r['user'] = r.pop('username')
env = self.write_pgpass(r)
pc = "user={user} host={host} port={port} dbname=postgres sslmode=prefer sslcompression=1".format(**r)
# first run a checkpoint on a promoted master in order
# to make it store the new timeline ([email protected])
self.checkpoint(pc)
logger.info("running pg_rewind from {}".format(pc))
self.checkpoint(r)
logger.info("running pg_rewind from {0}".format(pc))
pg_rewind = ['pg_rewind', '-D', self.data_dir, '--source-server', pc]
try:
ret = (subprocess.call(pg_rewind, env=env) == 0)
except:
ret = subprocess.call(pg_rewind, env=env) == 0
except OSError:
ret = False
if ret:
self.write_recovery_conf(leader)
@@ -497,16 +532,17 @@ recovery_target_timeline = 'latest'
finally:
return result
def single_user_mode(self, command=None, options={}):
def single_user_mode(self, command=None, options=None):
""" run a given command in a single-user mode. If the command is empty - then just start and stop """
cmd = ['postgres', '--single', '-D', self.data_dir]
for opt in sorted(options):
cmd.extend(['-c', '{0}={1}'.format(opt, options[opt])])
for opt, val in sorted((options or {}).items()):
cmd.extend(['-c', '{0}={1}'.format(opt, val)])
# need a database name to connect
cmd.append('postgres')
p = subprocess.Popen(cmd, stdin=subprocess.PIPE, stdout=open(os.devnull, 'w'), stderr=subprocess.STDOUT)
if p:
command and p.communicate('{}\n'.format(command))
if command:
p.communicate('{0}\n'.format(command))
p.stdin.close()
return p.wait()
return 1
@@ -521,10 +557,10 @@ recovery_target_timeline = 'latest'
os.unlink(path)
elif os.path.isfile(path):
os.remove(path)
except:
logger.exception("Unable to remove {}".format(path))
except OSError:
logger.exception("Unable to remove %s", path)
def follow_the_leader(self, leader, recovery=False):
def follow(self, leader, recovery=False):
if not self.check_recovery_conf(leader) or recovery:
change_role = (self.role == 'master')
@@ -560,7 +596,8 @@ recovery_target_timeline = 'latest'
self.remove_data_directory()
ret = True
self._need_rewind = False
change_role and ret and self.call_nowait(ACTION_ON_ROLE_CHANGE)
if change_role and ret:
self.call_nowait(ACTION_ON_ROLE_CHANGE)
return ret
else:
return True
@@ -573,16 +610,18 @@ recovery_target_timeline = 'latest'
"""
try:
for f in self.configuration_to_save:
os.path.isfile(f) and shutil.copy(f, f + '.backup')
except:
if os.path.isfile(f):
shutil.copy(f, f + '.backup')
except IOError:
logger.exception('unable to create backup copies of configuration files')
def restore_configuration_files(self):
""" restore a previously saved postgresql.conf """
try:
for f in self.configuration_to_save:
not os.path.isfile(f) and os.path.isfile(f + '.backup') and shutil.copy(f + '.backup', f)
except:
if not os.path.isfile(f) and os.path.isfile(f + '.backup'):
shutil.copy(f + '.backup', f)
except IOError:
logger.exception('unable to restore configuration files from backup')
def promote(self):
@@ -596,9 +635,6 @@ recovery_target_timeline = 'latest'
self.call_nowait(ACTION_ON_ROLE_CHANGE)
return ret
def demote(self):
self.follow_the_leader(None)
def create_or_update_role(self, name, password, options):
self.query("""DO $$
BEGIN
@@ -615,9 +651,7 @@ $$""".format(name, options), name, password, password)
def create_replication_user(self):
self.create_or_update_role(self.replication['username'], self.replication['password'], 'REPLICATION')
def create_connection_users(self):
if 'username' in self.superuser:
self.create_or_update_role(self.superuser['username'], self.superuser['password'], 'SUPERUSER')
def create_connection_user(self):
if self.admin:
self.create_or_update_role(self.admin['username'], self.admin['password'], 'CREATEDB CREATEROLE')
@@ -637,7 +671,18 @@ $$""".format(name, options), name, password, password)
if self.use_slots:
try:
self.load_replication_slots()
slots = [m.name for m in cluster.members if m.name != self.name] if self.role == 'master' else []
# if the replicatefrom tag is set on the member - we should not create the replication slot for it on
# the current master, because that member would replicate from elsewhere. We still create the slot if
# the replicatefrom destination member is currently not a member of the cluster (fallback to the
# master), or if replicatefrom destination member happens to be the current master
if self.role == 'master':
slots = [m.name for m in cluster.members if m.name != self.name and
(m.replicatefrom is None or m.replicatefrom == self.name or
not cluster.has_member(m.replicatefrom))]
else:
# only manage slots for replicas that replicate from this one, except for the leader among them
slots = [m.name for m in cluster.members if m.replicatefrom == self.name and
m.name != cluster.leader.name]
# drop unused slots
for slot in set(self.replication_slots) - set(slots):
self.query("""SELECT pg_drop_replication_slot(%s)
@@ -651,34 +696,44 @@ $$""".format(name, options), name, password, password)
WHERE slot_name = %s)""", slot, slot)
self.replication_slots = slots
except:
except psycopg2.Error:
logger.exception('Exception when changing replication slots')
def last_operation(self):
return str(self.xlog_position())
def bootstrap(self, current_leader=None):
def bootstrap(self, cluster_initialized=False, current_leader=None):
"""
Initially bootstrap PostgreSQL, either by creating a data
directory with initdb, or by initalizing a replica from an
exiting leader. Failure in the first case always leads to
exception, since there is no point in continuing if initdb failed.
In the second case, however, a False is returned on failure, since
it is normal for the replica to retry a failed attempt to initialize
from the master.
Populate PostgreSQL data directory by doing one of the following:
- create with initdb if there is no master.
- initialize the replica from an existing master
- initialize the replica using the replica creation method that
works without the master (i.e. restore from on-disk base backup)
The choice between the last 2 is triggered by the initialize flag.
We should never try to initdb an already initialized cluster, nor
try to bootstrap the cluster that lacks the initialize key from from
the master-less replica creation method (in the latter case, there is
no clear inidicator of the moment we should abandon our attempts and
swich to initdb).
Failure during initdb always leads to an exception, since there is
no point in continuing if initdb fails. For the rest of the cases,
the function returns False in order to inidicate a failed attempt
that should be retried in the future.
"""
ret = False
if not current_leader:
if not (cluster_initialized or current_leader):
ret = self.initialize() and self.start()
if ret:
self.create_replication_user()
self.create_connection_users()
self.create_connection_user()
else:
raise PostgresException("Could not bootstrap master PostgreSQL")
else:
if self.sync_from_leader(current_leader):
if self.sync_replica(current_leader):
self.restore_configuration_files()
self.write_recovery_conf(current_leader)
self.write_recovery_conf(current_leader, True)
ret = self.start()
return ret
@@ -688,7 +743,7 @@ $$""".format(name, options), name, password, password)
new_name = '{0}_{1}'.format(self.data_dir, time.strftime('%Y-%m-%d-%H-%M-%S'))
logger.info('renaming data directory to %s', new_name)
os.rename(self.data_dir, new_name)
except:
except OSError:
logger.exception("Could not rename data directory %s", self.data_dir)
def remove_data_directory(self):
@@ -702,7 +757,7 @@ $$""".format(name, options), name, password, password)
os.remove(self.data_dir)
elif os.path.isdir(self.data_dir):
shutil.rmtree(self.data_dir)
except:
except (IOError, OSError):
logger.exception('Could not remove data directory %s', self.data_dir)
self.move_data_directory()
+2 -2
View File
@@ -9,7 +9,7 @@ import boto.ec2
logger = logging.getLogger(__name__)
class AWSConnection:
class AWSConnection(object):
def __init__(self, cluster_name):
self.available = False
self.cluster_name = cluster_name if cluster_name is not None else 'unknown'
@@ -56,7 +56,7 @@ class AWSConnection:
conn = boto.ec2.connect_to_region(self.region)
conn.create_tags([self.instance_id], tags)
except Exception as e:
logger.info("could not set tags for EC2 instance {}: {}".format(self.instance_id, e))
logger.info("could not set tags for EC2 instance %s: %s", self.instance_id, e)
return False
return True
+10 -3
View File
@@ -42,7 +42,7 @@ logger = logging.getLogger(__name__)
class WALERestore(object):
def __init__(self, scope, datadir, connstring, env_dir, threshold_mb, threshold_pct, use_iam):
def __init__(self, scope, datadir, connstring, env_dir, threshold_mb, threshold_pct, use_iam, no_master):
self.scope = scope
self.master_connection = connstring
self.data_dir = datadir
@@ -51,6 +51,7 @@ class WALERestore(object):
self.wal_e.threshold_mb = threshold_mb
self.wal_e.threshold_pct = threshold_pct
self.wal_e.iam_string = ' --aws-instance-profile ' if use_iam == 1 else ''
self.no_master = no_master
self.wal_e.cmd = 'envdir {0} wal-e {1} '.format(self.wal_e.dir, self.wal_e.iam_string)
self.init_error = (not os.path.exists(self.wal_e.dir))
@@ -104,11 +105,12 @@ class WALERestore(object):
lsn_offset = hex((long(backup_start_segment[16:32], 16) << 24) + long(backup_start_offset))[2:-1]
# construct the LSN from the segment and offset
backup_start_lsn = '{}/{}'.format(lsn_segment, lsn_offset)
backup_start_lsn = '{0}/{1}'.format(lsn_segment, lsn_offset)
conn = None
cursor = None
diff_in_bytes = long(backup_size)
if not self.no_master:
try:
# get the difference in bytes between the current WAL location and the backup start offset
conn = psycopg2.connect(self.master_connection)
@@ -122,6 +124,9 @@ class WALERestore(object):
finally:
cursor and cursor.close()
conn and conn.close()
else:
# always try to use WAL-E if base backup is available
diff_in_bytes = 0
# if the size of the accumulated WAL segments is more than a certan percentage of the backup size
# or exceeds the pre-determined size - pg_basebackup is chosen instead.
@@ -150,13 +155,15 @@ def main():
parser.add_argument('--threshold_megabytes', type=int, default=10240)
parser.add_argument('--threshold_backup_size_percentage', type=int, default=30)
parser.add_argument('--use_iam', type=int, default=0)
parser.add_argument('--no_master', type=int, default=0)
args = parser.parse_args()
# retry cloning in a loop
for retry in range(0, args.retries + 1):
restore = WALERestore(scope=args.scope, datadir=args.datadir, connstring=args.connstring,
env_dir=args.envdir, threshold_mb=args.threshold_megabytes,
threshold_pct=args.threshold_backup_size_percentage, use_iam=args.use_iam)
threshold_pct=args.threshold_backup_size_percentage, use_iam=args.use_iam,
no_master=args.no_master)
ret = restore.run()
if ret == 0:
break
+16 -16
View File
@@ -8,9 +8,9 @@ import time
from patroni.exceptions import PatroniException
ignore_sigterm = False
interrupted_sleep = False
reap_children = False
__ignore_sigterm = False
__interrupted_sleep = False
__reap_children = False
_DATE_TIME_RE = re.compile(r'''^
(?P<year>\d{4})\-(?P<month>\d{2})\-(?P<day>\d{2}) # date
@@ -49,28 +49,28 @@ def calculate_ttl(expiration):
def sigterm_handler(signo, stack_frame):
global ignore_sigterm
if not ignore_sigterm:
ignore_sigterm = True
global __ignore_sigterm
if not __ignore_sigterm:
__ignore_sigterm = True
sys.exit()
def sigchld_handler(signo, stack_frame):
global interrupted_sleep, reap_children
reap_children = interrupted_sleep = True
global __interrupted_sleep, __reap_children
__reap_children = __interrupted_sleep = True
def sleep(interval):
global interrupted_sleep
global __interrupted_sleep
current_time = time.time()
end_time = current_time + interval
while current_time < end_time:
interrupted_sleep = False
__interrupted_sleep = False
time.sleep(end_time - current_time)
if not interrupted_sleep: # we will ignore only sigchld
if not __interrupted_sleep: # we will ignore only sigchld
break
current_time = time.time()
interrupted_sleep = False
__interrupted_sleep = False
def setup_signal_handlers():
@@ -79,8 +79,8 @@ def setup_signal_handlers():
def reap_children():
global reap_children
if reap_children:
global __reap_children
if __reap_children:
try:
while True:
ret = os.waitpid(-1, os.WNOHANG)
@@ -89,7 +89,7 @@ def reap_children():
except OSError:
pass
finally:
reap_children = False
__reap_children = False
class RetryFailedError(PatroniException):
@@ -97,7 +97,7 @@ class RetryFailedError(PatroniException):
"""Raised when retrying an operation ultimately failed, after retrying the maximum number of attempts."""
class Retry:
class Retry(object):
"""Helper for retrying a method in the face of retry-able exceptions"""
+6 -5
View File
@@ -17,7 +17,7 @@ class ZooKeeperError(DCSError):
pass
class ExhibitorEnsembleProvider:
class ExhibitorEnsembleProvider(object):
TIMEOUT = 3.1
@@ -54,7 +54,7 @@ class ExhibitorEnsembleProvider:
def _query_exhibitors(self, exhibitors):
random.shuffle(exhibitors)
for host in exhibitors:
uri = 'http://{}:{}{}'.format(host, self._exhibitor_port, self._uri_path)
uri = 'http://{0}:{1}{2}'.format(host, self._exhibitor_port, self._uri_path)
try:
response = requests.get(uri, timeout=self.TIMEOUT)
return response.json()
@@ -84,9 +84,9 @@ class ZooKeeper(AbstractDCS):
hosts = self.exhibitor.zookeeper_hosts
self.client = KazooClient(hosts=hosts,
timeout=(config.get('session_timeout', None) or 30),
timeout=(config.get('session_timeout') or 30),
command_retry={
'deadline': (config.get('reconnect_timeout', None) or 10),
'deadline': (config.get('reconnect_timeout') or 10),
'max_delay': 1,
'max_tries': -1},
connection_retry={'max_delay': 1, 'max_tries': -1})
@@ -190,7 +190,8 @@ class ZooKeeper(AbstractDCS):
def attempt_to_acquire_leader(self):
ret = self._create(self.leader_path, self._name, makepath=True, ephemeral=True)
ret or logger.info('Could not take out TTL lock')
if ret:
logger.info('Could not take out TTL lock')
return ret
def set_failover_value(self, value, index=None):
+3 -5
View File
@@ -4,7 +4,7 @@ scope: &scope batman
restapi:
listen: 127.0.0.1:8008
connect_address: 127.0.0.1:8008
auth: 'username:password'
# auth: 'username:password'
# certfile: /etc/ssl/certs/ssl-cert-snakeoil.pem
# keyfile: /etc/ssl/private/ssl-cert-snakeoil.key
etcd:
@@ -72,7 +72,6 @@ postgresql:
- basebackup
# - wal_e
# commented-out example for wal-e provisioning
#create_replica_method: wal_e, basebackup
#wal_e:
#command: /patroni/scripts/wale_restore.py
#env_dir: /etc/wal-e.d/env
@@ -88,14 +87,13 @@ postgresql:
archive_mode: "on"
wal_level: hot_standby
archive_command: mkdir -p ../wal_archive && test ! -f ../wal_archive/%f && cp %p ../wal_archive/%f
max_wal_senders: 5
max_wal_senders: 10
wal_keep_segments: 8
archive_timeout: 1800s
max_replication_slots: 5
max_replication_slots: 10
hot_standby: "on"
wal_log_hints: "on"
tags:
nofailover: False
noloadbalance: False
clonefrom: False
replicatefrom: 127.0.0.1
+4 -5
View File
@@ -4,7 +4,7 @@ scope: &scope batman
restapi:
listen: 127.0.0.1:8009
connect_address: 127.0.0.1:8009
auth: 'username:password'
# auth: 'username:password'
# certfile: /etc/ssl/certs/ssl-cert-snakeoil.pem
# keyfile: /etc/ssl/private/ssl-cert-snakeoil.key
etcd:
@@ -63,7 +63,7 @@ postgresql:
password: rep-pass
network: 127.0.0.1/32
superuser:
user: postgres
username: postgres
password: zalando
admin:
username: admin
@@ -88,14 +88,13 @@ postgresql:
archive_mode: "on"
wal_level: hot_standby
archive_command: mkdir -p ../wal_archive && test ! -f ../wal_archive/%f && cp %p ../wal_archive/%f
max_wal_senders: 5
max_wal_senders: 10
wal_keep_segments: 8
archive_timeout: 1800s
max_replication_slots: 5
max_replication_slots: 10
hot_standby: "on"
wal_log_hints: "on"
tags:
nofailover: False
noloadbalance: False
clonefrom: False
replicatefrom: 127.0.0.1
+101
View File
@@ -0,0 +1,101 @@
ttl: &ttl 30
loop_wait: &loop_wait 10
scope: &scope batman
restapi:
listen: 127.0.0.1:8010
connect_address: 127.0.0.1:8010
auth: 'username:password'
# certfile: /etc/ssl/certs/ssl-cert-snakeoil.pem
# keyfile: /etc/ssl/private/ssl-cert-snakeoil.key
etcd:
scope: *scope
ttl: *ttl
host: 127.0.0.1:4001
#discovery_srv: my-etcd.domain
#zookeeper:
# scope: *scope
# session_timeout: *ttl
# reconnect_timeout: *loop_wait
# hosts:
# - 127.0.0.1:2181
# - 127.0.0.2:2181
# exhibitor:
# poll_interval: 300
# port: 8181
# hosts:
# - host1
# - host2
# - host3
postgresql:
name: postgresql2
scope: *scope
listen: 127.0.0.1:5434
connect_address: 127.0.0.1:5434
data_dir: data/postgresql2
maximum_lag_on_failover: 1048576 # 1 megabyte in bytes
use_slots: True
pgpass: /tmp/pgpass2
initdb: ## We allow the following options to be passed on to initdb
# - auth: authmethod
# - auth-host: authmethod
# - auth-local: authmethod
- encoding: UTF8
# - data-checksums # When pg_rewind is needed on 9.3, this needs to be enabled
# - locale: locale
# - lc-collate: locale
# - lc-ctype: locale
# - lc-messages: locale
# - lc-monetary: locale
# - lc-numeric: locale
# - lc-time: locale
# - text-search-config: CFG
# - xlogdir: directory
# - debug
# - noclean
pg_rewind:
username: postgres
password: zalando
pg_hba:
- host all all 0.0.0.0/0 md5
- hostssl all all 0.0.0.0/0 md5
replication:
username: replicator
password: rep-pass
network: 127.0.0.1/32
superuser:
username: postgres
password: zalando
admin:
username: admin
password: admin
# commented-out example for wal-e provisioning
create_replica_method:
- basebackup
# - wal_e
# commented-out example for wal-e provisioning
#wal_e:
#command: /patroni/scripts/wale_restore.py
#env_dir: /home/postgres/etc/wal-e.d/env
#threshold_megabytes: 10240
#threshold_backup_size_percentage: 30
#retries: 2
#use_iam: 1
#recovery_conf:
#restore_command: envdir /etc/wal-e.d/env wal-e wal-fetch "%f" "%p" -p 1
recovery_conf:
restore_command: cp ../wal_archive/%f %p
parameters:
archive_mode: "on"
wal_level: hot_standby
archive_command: mkdir -p ../wal_archive && test ! -f ../wal_archive/%f && cp %p ../wal_archive/%f
max_wal_senders: 10
wal_keep_segments: 8
archive_timeout: 1800s
max_replication_slots: 10
hot_standby: "on"
wal_log_hints: "on"
tags:
nofailover: False
noloadbalance: False
clonefrom: False
replicatefrom: postgresql1
+19 -9
View File
@@ -16,11 +16,15 @@ class MockPostgresql(Mock):
name = 'test'
state = 'running'
role = 'master'
server_version = '999999'
scope = 'dummy'
def connection(self):
@staticmethod
def connection():
return psycopg2_connect()
def is_running(self):
@staticmethod
def is_running():
return True
@@ -29,31 +33,37 @@ class MockHa(Mock):
dcs = Mock()
state_handler = MockPostgresql()
def schedule_restart(self):
@staticmethod
def schedule_restart():
return 'restart'
def schedule_reinitialize(self):
@staticmethod
def schedule_reinitialize():
return 'reinitialize'
def restart(self):
@staticmethod
def restart():
return (True, '')
def restart_scheduled(self):
@staticmethod
def restart_scheduled():
return False
def fetch_nodes_statuses(self, members):
@staticmethod
def fetch_nodes_statuses(members):
return [[None, True, None, None, {}]]
class MockPatroni:
class MockPatroni(Mock):
postgresql = MockPostgresql()
ha = MockHa()
dcs = Mock()
tags = {}
version = '0.00'
class MockRequest:
class MockRequest(object):
def __init__(self, path):
self.path = path
+7 -15
View File
@@ -1,12 +1,13 @@
import unittest
import requests
import boto.ec2
from collections import namedtuple
from patroni.scripts.aws import AWSConnection
from requests.exceptions import RequestException
class MockEc2Connection:
class MockEc2Connection(object):
def __init__(self, error=False):
self.error = error
@@ -23,7 +24,7 @@ class MockEc2Connection:
return True
class MockResponse:
class MockResponse(object):
def __init__(self, content):
self.content = content
@@ -35,15 +36,6 @@ class MockResponse:
class TestAWSConnection(unittest.TestCase):
def __init__(self, method_name='runTest'):
super(TestAWSConnection, self).__init__(method_name)
def set_error(self):
self.error = True
def set_json_error(self):
self.json_error = True
def boto_ec2_connect_to_region(self, region):
return MockEc2Connection(self.error)
@@ -74,21 +66,21 @@ class TestAWSConnection(unittest.TestCase):
self.assertTrue(self.conn.on_role_change('master'))
def test_non_aws(self):
self.set_error()
self.error = True
conn = AWSConnection('test')
self.assertFalse(conn.aws_available())
self.assertFalse(conn._tag_ebs('master'))
self.assertFalse(conn._tag_ec2('master'))
def test_aws_bizare_response(self):
self.set_json_error()
self.json_error = True
conn = AWSConnection('test')
self.assertFalse(conn.aws_available())
def test_aws_tag_ebs_error(self):
self.set_error()
self.error = True
self.assertFalse(self.conn._tag_ebs("master"))
def test_aws_tag_ec2_error(self):
self.set_error()
self.error = True
self.assertFalse(self.conn._tag_ec2("master"))
+36 -25
View File
@@ -8,7 +8,8 @@ import psycopg2
import requests
import patroni.exceptions
import etcd
from mock import patch, Mock
from mock import patch, Mock, MagicMock
from click.testing import CliRunner
from patroni.ctl import ctl, members, store_config, load_config, output_members, post_patroni, get_dcs, \
@@ -23,6 +24,7 @@ from test_postgresql import MockConnect, psycopg2_connect
CONFIG_FILE_PATH = './test-ctl.yaml'
def test_rw_config():
runner = CliRunner()
config = {'a': 'b'}
@@ -45,12 +47,13 @@ def test_rw_config():
load_config(CONFIG_FILE_PATH, None)
load_config(CONFIG_FILE_PATH, '0.0.0.0')
@patch('patroni.ctl.load_config', Mock(return_value={'dcs': {'scheme': 'etcd', 'hostname': 'localhost', 'port': 4001}}))
class TestCtl(unittest.TestCase):
@patch('socket.getaddrinfo', socket_getaddrinfo)
@patch.object(Client, 'machines')
def setUp(self, mock_machines):
def setUp(self):
with patch.object(Client, 'machines') as mock_machines:
mock_machines.__get__ = Mock(return_value=['http://remotehost:2379'])
self.p = MockPostgresql()
self.e = Etcd('foo', {'ttl': 30, 'host': 'ok:2379', 'scope': 'test'})
@@ -103,35 +106,35 @@ y''')
result = runner.invoke(ctl, ['failover', 'dummy', '--dcs', '8.8.8.8'], input='''leader
other
N''')
assert 'Aborting failover' in str(result.exception)
assert 'Aborting failover' in str(result.output)
result = runner.invoke(ctl, ['failover', 'dummy', '--dcs', '8.8.8.8'], input='''leader
leader
y''')
assert 'target and source are the same' in str(result.exception)
assert 'target and source are the same' in str(result.output)
result = runner.invoke(ctl, ['failover', 'dummy', '--dcs', '8.8.8.8'], input='''leader
Reality
y''')
assert 'Reality does not exist' in str(result.exception)
assert 'Reality does not exist' in str(result.output)
result = runner.invoke(ctl, ['failover', 'dummy', '--force'])
assert 'Failing over to new leader' in result.output
result = runner.invoke(ctl, ['failover', 'dummy', '--dcs', '8.8.8.8'], input='dummy')
assert 'is not the leader of cluster' in str(result.exception)
assert 'is not the leader of cluster' in str(result.output)
with patch('patroni.etcd.Etcd.get_cluster', Mock(return_value=get_cluster_initialized_with_only_leader())):
result = runner.invoke(ctl, ['failover', 'dummy', '--dcs', '8.8.8.8'], input='''leader
other
y''')
assert 'No candidates found to failover to' in str(result.exception)
assert 'No candidates found to failover to' in str(result.output)
with patch('patroni.etcd.Etcd.get_cluster', Mock(return_value=get_cluster_initialized_without_leader())):
result = runner.invoke(ctl, ['failover', 'dummy', '--dcs', '8.8.8.8'], input='''leader
other
y''')
assert 'This cluster has no master' in str(result.exception)
assert 'This cluster has no master' in str(result.output)
with patch('patroni.ctl.post_patroni', Mock(side_effect=Exception())):
result = runner.invoke(ctl, ['failover', 'dummy', '--dcs', '8.8.8.8'], input='''leader
@@ -150,13 +153,13 @@ y''')
# with patch('patroni.dcs.AbstractDCS.get_cluster', Mock(return_value=get_cluster_initialized_with_leader())):
# result = runner.invoke(ctl, ['failover', 'alpha', '--dcs', '8.8.8.8'], input='nonsense')
# assert 'is not the leader of cluster' in str(result.exception)
# assert 'is not the leader of cluster' in str(result.output)
# result = runner.invoke(ctl, ['failover', 'alpha', '--dcs', '8.8.8.8', '--master', 'nonsense'])
# assert 'is not the leader of cluster' in str(result.exception)
# assert 'is not the leader of cluster' in str(result.output)
# result = runner.invoke(ctl, ['failover', 'alpha', '--dcs', '8.8.8.8'], input='leader\nother\nn')
# assert 'Aborting failover' in str(result.exception)
# assert 'Aborting failover' in str(result.output)
# with patch('patroni.ctl.wait_for_leader', Mock(return_value = get_cluster_initialized_with_leader())):
# result = runner.invoke(ctl, ['failover', 'alpha', '--dcs', '8.8.8.8'], input='leader\nother\nY')
@@ -182,12 +185,17 @@ y''')
'--role',
'master',
])
assert 'mutually exclusive' in str(result.exception)
assert 'mutually exclusive' in str(result.output)
with runner.isolated_filesystem():
dummy_file = open('dummy', 'w')
with open('dummy', 'w') as dummy_file:
dummy_file.write('SELECT 1')
dummy_file.close()
result = runner.invoke(ctl, [
'query',
'alpha'
])
assert 'You need to specify' in str(result.output)
result = runner.invoke(ctl, [
'query',
@@ -197,7 +205,7 @@ y''')
'--command',
'dummy',
])
assert 'mutually exclusive' in str(result.exception)
assert 'mutually exclusive' in str(result.output)
result = runner.invoke(ctl, ['query', 'alpha', '--file', 'dummy'])
@@ -206,6 +214,10 @@ y''')
result = runner.invoke(ctl, ['query', 'alpha', '--command', 'SELECT 1'])
assert 'mock column' in result.output
result = runner.invoke(ctl, ['query', 'alpha', '--command', 'SELECT 1', '--dbname', 'dummy',
'--password', '--username', 'dummy'], input='password\n')
assert 'mock column' in result.output
@patch('patroni.ctl.get_cursor', Mock(return_value=MockConnect().cursor()))
def test_query_member(self):
rows = query_member(None, None, None, 'master', 'SELECT pg_is_in_recovery()')
@@ -243,10 +255,10 @@ y''')
'--member',
'dummy',
])
assert 'mutually exclusive' in str(result.exception)
assert 'mutually exclusive' in str(result.output)
result = runner.invoke(ctl, ['dsn', 'alpha', '--member', 'dummy'])
assert 'Can not find' in str(result.exception)
assert 'Can not find' in str(result.output)
# result = runner.invoke(ctl, ['dsn', 'alpha', '--dcs', '8.8.8.8', '--role', 'replica'])
# assert 'host=127.0.0.1 port=5436' in result.output
@@ -270,7 +282,7 @@ y''')
'dummy',
'--any',
], input='y')
assert 'not a member' in str(result.exception)
assert 'not a member' in str(result.output)
with patch('requests.post', Mock(return_value=MockResponse())):
result = runner.invoke(ctl, ['restart', 'alpha', '--dcs', '8.8.8.8'], input='y')
@@ -283,15 +295,15 @@ y''')
result = runner.invoke(ctl, ['remove', 'alpha', '--dcs', '8.8.8.8'], input='alpha\nslave')
assert 'Please confirm' in result.output
assert 'You are about to remove all' in result.output
assert 'You did not exactly type' in str(result.exception)
assert 'You did not exactly type' in str(result.output)
result = runner.invoke(ctl, ['remove', 'alpha', '--dcs', '8.8.8.8'], input='''alpha
Yes I am aware
slave''')
assert 'You did not specify the current master of the cluster' in str(result.exception)
assert 'You did not specify the current master of the cluster' in str(result.output)
result = runner.invoke(ctl, ['remove', 'alpha', '--dcs', '8.8.8.8'], input='beta\nleader')
assert 'Cluster names specified do not match' in str(result.exception)
assert 'Cluster names specified do not match' in str(result.output)
with patch('patroni.etcd.Etcd.get_cluster', get_cluster_initialized_with_leader):
result = runner.invoke(ctl, ['remove', 'alpha', '--dcs', '8.8.8.8'],
@@ -305,7 +317,7 @@ leader''')
input='''alpha
Yes I am aware
leader''')
assert 'We have not implemented this for DCS of type' in str(result.exception)
assert 'We have not implemented this for DCS of type' in str(result.output)
@patch('patroni.etcd.Etcd.watch', Mock(return_value=None))
@patch('patroni.etcd.Etcd.get_cluster', Mock(return_value=get_cluster_initialized_with_leader()))
@@ -317,6 +329,7 @@ leader''')
assert cluster.leader.member.name == 'leader'
def test_post_patroni(self):
with patch('requests.post', MagicMock(side_effect=requests.exceptions.ConnectionError('foo'))):
member = get_cluster_initialized_with_leader().leader.member
self.assertRaises(requests.exceptions.ConnectionError, post_patroni, member, 'dummy', {})
@@ -373,5 +386,3 @@ leader''')
])
assert result.exit_code == 0
+12 -6
View File
@@ -11,7 +11,7 @@ from patroni.dcs import Cluster, DCSError, Leader
from patroni.etcd import Client, Etcd, EtcdError
class MockResponse:
class MockResponse(object):
def __init__(self):
self.status_code = 200
@@ -34,13 +34,18 @@ class MockResponse:
def status(self):
return self.status_code
@staticmethod
def getheader(*args):
return ''
class MockPostgresql(Mock):
def last_operation(self):
server_version = '999999'
scope = 'dummy'
@staticmethod
def last_operation():
return '0'
@@ -81,8 +86,8 @@ def etcd_watch(key, index=None, timeout=None, recursive=None):
def etcd_write(key, value, **kwargs):
if key == '/service/exists/leader':
raise etcd.EtcdAlreadyExist
if key == '/service/test/leader' or key == '/patroni/test/leader':
if kwargs.get('prevValue', None) == 'foo' or not kwargs.get('prevExist', True):
if key in ['/service/test/leader', '/patroni/test/leader'] and \
(kwargs.get('prevValue') == 'foo' or not kwargs.get('prevExist', True)):
return True
raise etcd.EtcdException
@@ -124,12 +129,12 @@ class SleepException(Exception):
pass
class MockSRV:
class MockSRV(object):
port = 2380
target = '127.0.0.1'
def dns_query(name, type):
def dns_query(name, _):
if name == '_etcd-server._tcp.blabla':
return []
elif name == '_etcd-server._tcp.exception':
@@ -171,6 +176,7 @@ class TestClient(unittest.TestCase):
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')
def test_get_srv_record(self):
self.assertEquals(self.client.get_srv_record('blabla'), [])
+58 -20
View File
@@ -18,7 +18,7 @@ def false(*args, **kwargs):
def get_cluster(initialize, leader, members, failover):
return Cluster(initialize, leader, None, members, failover)
return Cluster(initialize, leader, 10, members, failover)
def get_cluster_not_initialized_without_leader():
@@ -37,6 +37,7 @@ def get_cluster_initialized_without_leader(leader=False, failover=None):
def get_cluster_initialized_with_leader(failover=None):
return get_cluster_initialized_without_leader(leader=True, failover=failover)
def get_cluster_initialized_with_only_leader(failover=None):
l = get_cluster_initialized_without_leader(leader=True, failover=failover).leader
return get_cluster(True, l, [l], failover)
@@ -48,39 +49,51 @@ class MockPostgresql(Mock):
role = 'replica'
state = 'running'
connection_string = 'postgres://foo@bar/postgres'
server_version = '999999'
scope = 'dummy'
def is_healthy(self):
@staticmethod
def is_healthy():
return True
def start(self):
@staticmethod
def start():
return True
def is_healthiest_node(self, members):
@staticmethod
def is_healthiest_node(members):
return True
def is_leader(self):
@staticmethod
def is_leader():
return True
def xlog_position(self):
@staticmethod
def xlog_position():
return 0
def last_operation(self):
@staticmethod
def last_operation():
return 0
def data_directory_empty(self):
@staticmethod
def data_directory_empty():
return False
def bootstrap(self, *args, **kwargs):
@staticmethod
def bootstrap(*args, **kwargs):
return True
def check_replication_lag(self, last_leader_operation):
@staticmethod
def check_replication_lag(last_leader_operation):
return True
def check_recovery_conf(self, leader):
@staticmethod
def check_recovery_conf(leader):
return False
class MockPatroni:
class MockPatroni(object):
def __init__(self, p, d):
self.postgresql = p
@@ -88,20 +101,22 @@ class MockPatroni:
self.api = Mock()
self.tags = {}
self.nofailover = None
self.replicatefrom = None
self.api.connection_string = 'http://127.0.0.1:8008'
def run_async(func, args=()):
func(*args) if args else func()
return func(*args) if args else func()
class TestHa(unittest.TestCase):
@patch('socket.getaddrinfo', socket_getaddrinfo)
@patch.object(Client, 'machines')
def setUp(self, mock_machines):
def setUp(self):
with patch.object(Client, 'machines') as mock_machines:
mock_machines.__get__ = Mock(return_value=['http://remotehost:2379'])
self.p = MockPostgresql()
self.p.can_create_replica_without_leader = MagicMock(return_value=False)
self.e = Etcd('foo', {'ttl': 30, 'host': 'ok:2379', 'scope': 'test'})
self.e.client.read = etcd_read
self.e.client.write = etcd_write
@@ -127,13 +142,19 @@ class TestHa(unittest.TestCase):
def test_recover_replica_failed(self):
self.p.controldata = lambda: {'Database cluster state': 'in production'}
self.p.is_healthy = false
self.p.follow_the_leader = false
self.p.is_running = false
self.p.follow = false
self.assertEquals(self.ha.run_cycle(), 'started as a secondary')
self.assertEquals(self.ha.run_cycle(), 'failed to start postgres')
def test_recover_master_failed(self):
self.p.follow_the_leader = false
self.p.follow = false
self.p.is_healthy = false
self.p.is_running = false
self.ha.has_lock = true
self.p.role = 'master'
self.p.controldata = lambda: {'Database cluster state': 'in production'}
self.assertEquals(self.ha.run_cycle(), 'started as readonly because i had the session lock')
self.assertEquals(self.ha.run_cycle(), 'removed leader key after trying and failing to start postgres')
@patch('sys.exit', return_value=1)
@@ -144,7 +165,8 @@ class TestHa(unittest.TestCase):
@patch.object(Cluster, 'is_unlocked', Mock(return_value=False))
def test_start_as_readonly(self):
self.p.is_leader = self.p.is_healthy = false
self.p.is_leader = false
self.p.is_healthy = true
self.ha.has_lock = true
self.assertEquals(self.ha.run_cycle(), 'promoted self to leader because i had the session lock')
@@ -158,7 +180,7 @@ class TestHa(unittest.TestCase):
def test_demote_after_failing_to_obtain_lock(self):
self.ha.acquire_lock = false
self.assertEquals(self.ha.run_cycle(), 'demoted self due after trying and failing to obtain lock')
self.assertEquals(self.ha.run_cycle(), 'demoted self after trying and failing to obtain lock')
def test_follow_new_leader_after_failing_to_obtain_lock(self):
self.ha.is_healthiest_node = true
@@ -196,10 +218,12 @@ class TestHa(unittest.TestCase):
self.ha.update_lock = false
self.assertEquals(self.ha.run_cycle(), 'demoting self because i do not have the lock and i was a leader')
def test_follow_the_leader(self):
def test_follow(self):
self.ha.cluster.is_unlocked = false
self.p.is_leader = false
self.assertEquals(self.ha.run_cycle(), 'no action. i am a secondary and i am following a leader')
self.ha.patroni.replicatefrom = "foo"
self.assertEquals(self.ha.run_cycle(), 'no action. i am a secondary and i am following a leader')
def test_no_etcd_connection_master_demote(self):
self.ha.load_cluster_from_dcs = Mock(side_effect=DCSError('Etcd is not responding properly'))
@@ -214,6 +238,11 @@ class TestHa(unittest.TestCase):
self.ha.cluster = get_cluster_initialized_without_leader()
self.assertEquals(self.ha.bootstrap(), 'waiting for leader to bootstrap')
def test_bootstrap_without_leader(self):
self.ha.cluster = get_cluster_initialized_without_leader()
self.p.can_create_replica_without_leader = MagicMock(return_value=True)
self.assertEquals(self.ha.bootstrap(), "trying to bootstrap without leader")
def test_bootstrap_initialize_lock_failed(self):
self.ha.cluster = get_cluster_not_initialized_without_leader()
self.assertEquals(self.ha.bootstrap(), 'failed to acquire initialize lock')
@@ -333,3 +362,12 @@ class TestHa(unittest.TestCase):
self.ha.fetch_node_status(member)
member = Member(0, 'test', 1, {'api_url': 'http://localhost:8011/patroni'})
self.ha.fetch_node_status(member)
def test_post_recover(self):
self.p.is_running = false
self.ha.has_lock = true
self.assertEqual(self.ha.post_recover(), 'removed leader key after trying and failing to start postgres')
self.ha.has_lock = false
self.assertEqual(self.ha.post_recover(), 'failed to start postgres')
self.p.is_running = true
self.assertIsNone(self.ha.post_recover())
+7 -2
View File
@@ -28,8 +28,8 @@ def time_sleep(*args):
@patch.object(AsyncExecutor, 'run', Mock())
class TestPatroni(unittest.TestCase):
@patch.object(Client, 'machines')
def setUp(self, mock_machines):
def setUp(self):
with patch.object(Client, 'machines') as mock_machines:
mock_machines.__get__ = Mock(return_value=['http://remotehost:2379'])
self.touched = False
self.init_cancelled = False
@@ -80,3 +80,8 @@ class TestPatroni(unittest.TestCase):
self.assertTrue(self.p.nofailover)
self.p.tags['nofailover'] = None
self.assertFalse(self.p.nofailover)
def test_replicatefrom(self):
self.assertIsNone(self.p.replicatefrom)
self.p.tags['replicatefrom'] = 'foo'
self.assertEqual(self.p.replicatefrom, 'foo')
+58 -40
View File
@@ -19,7 +19,7 @@ def is_file_raise_on_backup(*args, **kwargs):
raise Exception("foo")
class MockCursor:
class MockCursor(object):
def __init__(self, connection):
self.connection = connection
@@ -59,7 +59,8 @@ class MockCursor:
def fetchall(self):
return self.results
def close(self):
@staticmethod
def close():
pass
def __iter__(self):
@@ -154,9 +155,12 @@ def psycopg2_connect(*args, **kwargs):
return MockConnect()
def fake_listdir(path):
return ["a", "b", "c"] if path.endswith('pg_xlog/archive_status') else []
@patch('subprocess.call', Mock(return_value=0))
@patch('psycopg2.connect', psycopg2_connect)
@patch('shutil.copy', Mock())
class TestPostgresql(unittest.TestCase):
@patch('subprocess.call', Mock(return_value=0))
@@ -165,7 +169,7 @@ class TestPostgresql(unittest.TestCase):
self.p = Postgresql({'name': 'test0', 'scope': 'batman', 'data_dir': 'data/test0',
'listen': '127.0.0.1, *:5432', 'connect_address': '127.0.0.2:5432',
'pg_hba': ['hostssl all all 0.0.0.0/0 md5', 'host all all 0.0.0.0/0 md5'],
'superuser': {'password': 'test'},
'superuser': {'username': 'test', 'password': 'test'},
'admin': {'username': 'admin', 'password': 'admin'},
'pg_rewind': {'username': 'admin', 'password': 'admin'},
'replication': {'username': 'replicator',
@@ -181,7 +185,8 @@ class TestPostgresql(unittest.TestCase):
os.makedirs(self.p.data_dir)
self.leadermem = Member(0, 'leader', 28, {'conn_url': 'postgres://replicator:[email protected]:5435/postgres'})
self.leader = Leader(-1, 28, self.leadermem)
self.other = Member(0, 'test1', 28, {'conn_url': 'postgres://replicator:[email protected]:5433/postgres'})
self.other = Member(0, 'test1', 28, {'conn_url': 'postgres://replicator:[email protected]:5433/postgres',
'tags': {'replicatefrom': 'leader'}})
self.me = Member(0, 'test0', 28, {'conn_url': 'postgres://replicator:[email protected]:5434/postgres'})
def tearDown(self):
@@ -204,6 +209,11 @@ class TestPostgresql(unittest.TestCase):
self.assertTrue(self.p.initialize())
self.assertTrue(os.path.exists(os.path.join(self.p.data_dir, 'pg_hba.conf')))
@patch('os.path.exists', Mock(return_value=True))
@patch('os.unlink', Mock())
def test_delete_trigger_file(self):
self.p.delete_trigger_file()
def test_start(self):
self.assertTrue(self.p.start())
self.p.is_running = false
@@ -228,10 +238,12 @@ class TestPostgresql(unittest.TestCase):
self.p.write_pgpass({'host': 'localhost', 'port': '5432', 'user': 'foo', 'password': 'bar'})
@patch('patroni.postgresql.Postgresql.write_pgpass', MagicMock(return_value=dict()))
def test_sync_from_leader(self):
self.assertTrue(self.p.sync_from_leader(self.leader))
def test_sync_replica(self):
self.assertTrue(self.p.sync_replica(self.leader))
self.p.create_replica = Mock(return_value=1)
self.assertFalse(self.p.sync_replica(self.leader))
@patch('subprocess.call', side_effect=Exception("Test"))
@patch('subprocess.call', side_effect=OSError)
@patch('patroni.postgresql.Postgresql.write_pgpass', MagicMock(return_value=dict()))
def test_pg_rewind(self, mock_call):
self.assertTrue(self.p.rewind(self.leader))
@@ -242,25 +254,25 @@ class TestPostgresql(unittest.TestCase):
@patch('patroni.postgresql.Postgresql.remove_data_directory', MagicMock(return_value=True))
@patch('patroni.postgresql.Postgresql.single_user_mode', MagicMock(return_value=1))
@patch('patroni.postgresql.Postgresql.write_pgpass', MagicMock(return_value=dict()))
def test_follow_the_leader(self, mock_pg_rewind):
self.p.demote()
self.p.follow_the_leader(None)
self.p.demote()
self.p.follow_the_leader(self.leader)
self.p.follow_the_leader(Leader(-1, 28, self.other))
def test_follow(self, mock_pg_rewind):
self.p.follow(None)
self.p.follow(self.leader)
self.p.follow(Leader(-1, 28, self.other))
self.p.rewind = mock_pg_rewind
self.p.follow_the_leader(self.leader)
self.p.follow(self.leader)
self.p.require_rewind()
with mock.patch('os.path.islink', MagicMock(return_value=True)):
with mock.patch('patroni.postgresql.Postgresql.can_rewind', new_callable=PropertyMock(return_value=True)):
with mock.patch('os.unlink', MagicMock(return_value=True)):
self.p.follow_the_leader(self.leader, recovery=True)
self.p.follow(self.leader, recovery=True)
self.p.require_rewind()
with mock.patch('patroni.postgresql.Postgresql.can_rewind', new_callable=PropertyMock(return_value=True)):
self.p.rewind.return_value = True
self.p.follow_the_leader(self.leader, recovery=True)
self.p.follow(self.leader, recovery=True)
self.p.rewind.return_value = False
self.p.follow_the_leader(self.leader, recovery=True)
self.p.follow(self.leader, recovery=True)
with mock.patch('patroni.postgresql.Postgresql.check_recovery_conf', MagicMock(return_value=True)):
self.assertTrue(self.p.follow(None))
def test_can_rewind(self):
tmp = self.p.pg_rewind
@@ -294,12 +306,6 @@ class TestPostgresql(unittest.TestCase):
with patch('subprocess.call', Mock(side_effect=Exception("foo"))):
self.assertEquals(self.p.create_replica(self.leader, ''), 1)
def test_create_connection_users(self):
cfg = self.p.config
cfg['superuser']['username'] = 'test'
p = Postgresql(cfg)
p.create_connection_users()
def test_sync_replication_slots(self):
self.p.start()
cluster = Cluster(True, self.leader, 0, [self.me, self.other, self.leadermem], None)
@@ -307,6 +313,9 @@ class TestPostgresql(unittest.TestCase):
self.p.query = Mock(side_effect=psycopg2.OperationalError)
self.p.schedule_load_slots = True
self.p.sync_replication_slots(cluster)
self.p.schedule_load_slots = False
with mock.patch('patroni.postgresql.Postgresql.role', new_callable=PropertyMock(return_value='replica')):
self.p.sync_replication_slots(cluster)
@patch.object(MockConnect, 'closed', 2)
def test__query(self):
@@ -366,6 +375,7 @@ class TestPostgresql(unittest.TestCase):
with patch('subprocess.call', Mock(return_value=1)):
self.assertRaises(PostgresException, self.p.bootstrap)
self.p.bootstrap()
with patch('patroni.postgresql.Postgresql.sync_replica', MagicMock(return_value=True)):
self.p.bootstrap(self.leader)
def test_remove_data_directory(self):
@@ -376,7 +386,7 @@ class TestPostgresql(unittest.TestCase):
open(self.p.data_dir, 'w').close()
self.p.remove_data_directory()
os.symlink('unexisting', self.p.data_dir)
with patch('os.unlink', Mock(side_effect=Exception)):
with patch('os.unlink', Mock(side_effect=OSError)):
self.p.remove_data_directory()
self.p.remove_data_directory()
@@ -431,11 +441,6 @@ class TestPostgresql(unittest.TestCase):
subprocess_popen_mock.return_value = None
self.assertEquals(self.p.single_user_mode(), 1)
def fake_listdir(path):
if path.endswith(os.path.join('pg_xlog', 'archive_status')):
return ["a", "b", "c"]
return []
@patch('os.listdir', MagicMock(side_effect=fake_listdir))
@patch('os.path.isdir', MagicMock(return_value=True))
@patch('os.unlink', return_value=True)
@@ -459,8 +464,8 @@ class TestPostgresql(unittest.TestCase):
mock_unlink.reset_mock()
mock_remove.reset_mock()
mock_file.side_effect = Exception("foo")
mock_link.side_effect = Exception("foo")
mock_file.side_effect = OSError
mock_link.side_effect = OSError
self.p.cleanup_archive_status()
mock_unlink.assert_not_called()
mock_remove.assert_not_called()
@@ -469,14 +474,27 @@ class TestPostgresql(unittest.TestCase):
def test_sysid(self):
self.assertEqual(self.p.sysid, "6200971513092291716")
@patch('os.path.isfile', MagicMock(return_value=True))
@patch('shutil.copy', side_effect=Exception)
def test_save_configuration_files(self, mock_copy):
shutil.copy = mock_copy
@patch('os.path.isfile', Mock(return_value=True))
@patch('shutil.copy', Mock(side_effect=IOError))
def test_save_configuration_files(self):
self.p.save_configuration_files()
@patch('os.path.isfile', MagicMock(side_effect=is_file_raise_on_backup))
@patch('shutil.copy', side_effect=Exception)
def test_restore_configuration_files(self, mock_copy):
shutil.copy = mock_copy
@patch('os.path.isfile', Mock(side_effect=[False, True]))
@patch('shutil.copy', Mock(side_effect=IOError))
def test_restore_configuration_files(self):
self.p.restore_configuration_files()
def test_can_create_replica_without_leader(self):
self.p.config['create_replica_method'] = []
self.assertFalse(self.p.can_create_replica_without_leader())
self.p.config['create_replica_method'] = ['wale', 'basebackup']
self.p.config['wale'] = {'command': 'foo', 'no_master': 1}
self.assertTrue(self.p.can_create_replica_without_leader())
def test_replica_method_can_work_without_leader(self):
self.assertFalse(self.p.replica_method_can_work_without_leader('basebackup'))
self.assertFalse(self.p.replica_method_can_work_without_leader('foobar'))
self.p.config['foo'] = {'command': 'bar', 'no_master': 1}
self.assertTrue(self.p.replica_method_can_work_without_leader('foo'))
self.p.config['foo'] = {'command': 'bar'}
self.assertFalse(self.p.replica_method_can_work_without_leader('foo'))
+8 -10
View File
@@ -29,7 +29,8 @@ class TestUtils(unittest.TestCase):
@patch('time.sleep', Mock())
class TestRetrySleeper(unittest.TestCase):
def _fail(self, times=1):
@staticmethod
def _fail(times=1):
scope = dict(times=0)
def inner():
@@ -40,36 +41,33 @@ class TestRetrySleeper(unittest.TestCase):
raise PatroniException('Failed!')
return inner
def _makeOne(self, *args, **kwargs):
return Retry(*args, **kwargs)
def test_reset(self):
retry = self._makeOne(delay=0, max_tries=2)
retry = Retry(delay=0, max_tries=2)
retry(self._fail())
self.assertEquals(retry._attempts, 1)
retry.reset()
self.assertEquals(retry._attempts, 0)
def test_too_many_tries(self):
retry = self._makeOne(delay=0)
retry = Retry(delay=0)
self.assertRaises(RetryFailedError, retry, self._fail(times=999))
self.assertEquals(retry._attempts, 1)
def test_maximum_delay(self):
retry = self._makeOne(delay=10, max_tries=100)
retry = Retry(delay=10, max_tries=100)
retry(self._fail(times=10))
self.assertTrue(retry._cur_delay < 4000, retry._cur_delay)
# gevent's sleep function is picky about the type
self.assertEquals(type(retry._cur_delay), float)
def test_deadline(self):
retry = self._makeOne(deadline=0.0001)
retry = Retry(deadline=0.0001)
self.assertRaises(RetryFailedError, retry, self._fail(times=100))
def test_copy(self):
def _sleep(t):
None
pass
retry = self._makeOne(sleep_func=_sleep)
retry = Retry(sleep_func=_sleep)
rcopy = retry.copy()
self.assertTrue(rcopy.sleep_func is _sleep)
+11 -3
View File
@@ -1,9 +1,8 @@
import unittest
from mock import MagicMock, patch, PropertyMock
import os
import psycopg2
import subprocess
from patroni.scripts.wale_restore import WALERestore
from patroni.scripts.wale_restore import WALERestore, main
def fake_cursor_fetchone(*args, **kwargs):
@@ -28,16 +27,19 @@ def fake_backup_data(self, *args, **kwargs):
base_00000001000000000000007F_00000040 2015-05-18T10:13:25.000Z 167772160 00000001000000000000007F 00000040 00000001000000000000007F 00000240
"""
def fake_backup_data_2(self, *args, **kwargs):
""" return the fake result of WAL-E backup-list"""
return """name last_modified expanded_size_bytes wal_segment_backup_start wal_segment_offset_backup_start wal_segment_backup_stop wal_segment_offset_backup_stop """
def fake_backup_data_3(self, *args, **kwargs):
""" return the fake result of WAL-E backup-list"""
return """name last_modified expanded_size_bytes wal_segment_backup_start wal_segment_offset_backup_start wal_segment_backup_stop
base_00000001000000000000007F_00000040 2015-05-18T10:13:25.000Z 167772160 00000001000000000000007F 00000040 00000001000000000000007F 00000240
"""
def fake_backup_data_4(self, *args, **kwargs):
""" return the fake result of WAL-E backup-list"""
return """name last_modified expanded_size_foo wal_segment_backup_start wal_segment_offset_backup_start wal_segment_backup_stop wal_segment_offset_backup_stop
@@ -58,7 +60,7 @@ class TestWALERestore(unittest.TestCase):
def setUp(self):
self.wale_restore = WALERestore("batman", "/data",
"host=batman port=5432 user=batman", "/etc", 100, 100, 1)
"host=batman port=5432 user=batman", "/etc", 100, 100, 1, 0)
def tearDown(self):
pass
@@ -76,6 +78,8 @@ class TestWALERestore(unittest.TestCase):
self.assertFalse(self.wale_restore.should_use_s3_to_create_replica())
self.wale_restore.should_use_s3_to_create_replica()
self.wale_restore.no_master = 1
self.assertTrue(self.wale_restore.should_use_s3_to_create_replica())
def test_create_replica_with_s3(self):
with patch('subprocess.call', MagicMock(return_value=0)):
@@ -89,3 +93,7 @@ class TestWALERestore(unittest.TestCase):
with patch.object(self.wale_restore, 'should_use_s3_to_create_replica', MagicMock(return_value=True)):
with patch.object(self.wale_restore, 'create_replica_with_s3', MagicMock(return_value=0)):
self.assertEqual(self.wale_restore.run(), 0)
def test_main(self):
with patch('sys.exit', MagicMock(return_value=0)):
self.assertEqual(main(), None)
+7 -5
View File
@@ -20,7 +20,8 @@ class MockKazooClient(Mock):
def client_id(self):
return (-1, '')
def retry(self, func, *args, **kwargs):
@staticmethod
def retry(func, *args, **kwargs):
func(*args, **kwargs)
def get(self, path, watch=None):
@@ -43,7 +44,8 @@ class MockKazooClient(Mock):
return (b'foo', ZnodeStat(0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0))
return (b'', ZnodeStat(0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0))
def get_children(self, path, watch=None, include_data=False):
@staticmethod
def get_children(path, watch=None, include_data=False):
if not isinstance(path, six.string_types):
raise TypeError("Invalid type for 'path' (string expected)")
if path.startswith('/no_node'):
@@ -62,15 +64,15 @@ class MockKazooClient(Mock):
elif value == b'retry' or (value == b'exists' and self.exists):
raise NodeExistsError
def set(self, path, value, version=-1):
@staticmethod
def set(path, value, version=-1):
if not isinstance(path, six.string_types):
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 path == '/service/bla/optime/leader':
raise Exception
if path == '/service/test/members/bar':
if value == b'retry':
if path == '/service/test/members/bar' and value == b'retry':
return
if path == '/service/test/failover':
if value == b'Exception':