From 8ef7178ddf71b09bab674bcb3921b41677bd7225 Mon Sep 17 00:00:00 2001 From: Alexander Kukushkin Date: Wed, 10 Aug 2016 10:19:52 +0200 Subject: [PATCH] Refactor code dealing with database connection string/params (#255) In the original code we were parsing/deparsing url-style connection strings back and forth. That was not really resource greedy but rather annoying. Also it was not really obvious how to switch all local connections to unix-sockets (preferably). This commit isolates different use-cases of working with connection strings and minimizes amount of code parsing and deparsing them. Also it introduces one new helper method in the `Member` object - `conn_kwargs`. This method can accept as a parameter dict object with credentials (username and password). As a result it returns dict object which could be used by `psycopg2.connect` or for building connection urls for pg_rewind, pg_basebackup or some other replica creation methods. Params for local connection are builded in the `_local_connect_kwargs` method and could be changed to unix-socket later easily. --- patroni/ctl.py | 22 ++++------ patroni/dcs/__init__.py | 25 +++++++++++- patroni/ha.py | 5 +-- patroni/postgresql.py | 87 ++++++++++++++++++++-------------------- tests/test_ctl.py | 6 ++- tests/test_postgresql.py | 2 +- 6 files changed, 82 insertions(+), 65 deletions(-) diff --git a/patroni/ctl.py b/patroni/ctl.py index 0f590457..cb6978b9 100644 --- a/patroni/ctl.py +++ b/patroni/ctl.py @@ -21,7 +21,6 @@ from click import ClickException from patroni.config import Config from patroni.dcs import get_dcs as _get_dcs from patroni.exceptions import PatroniException -from patroni.postgresql import get_conn_kwargs from patroni.utils import is_valid_pg_version from prettytable import PrettyTable from six.moves.urllib_parse import urlparse @@ -185,16 +184,6 @@ def watching(w, watch, max_count=None, clear=True): yield 0 -def build_connect_parameters(conn_url, connect_parameters): - params = get_conn_kwargs(conn_url, connect_parameters) - params.update({'fallback_application_name': 'Patroni ctl', 'connect_timeout': '5'}) - if 'database' in connect_parameters: - params['database'] = connect_parameters['database'] - else: - params.pop('database') - return params - - def get_all_members(cluster, role='master'): if role == 'master': if cluster.leader is not None: @@ -219,7 +208,12 @@ def get_cursor(cluster, connect_parameters, role='master', member=None): if member is None: return None - params = build_connect_parameters(member.conn_url, connect_parameters) + params = member.conn_kwargs(connect_parameters) + params.update({'fallback_application_name': 'Patroni ctl', 'connect_timeout': '5'}) + if 'database' in connect_parameters: + params['database'] = connect_parameters['database'] + else: + params.pop('database') conn = psycopg2.connect(**params) conn.autocommit = True @@ -287,7 +281,7 @@ def dsn(cluster_name, config_file, dcs, role, member): if m is None: raise PatroniCtlException('Can not find a suitable member') - params = get_conn_kwargs(m.conn_url) + params = m.conn_kwargs() click.echo('host={host} port={port}'.format(**params)) @@ -650,7 +644,7 @@ def output_members(cluster, name, extended=False, fmt='pretty'): if m.name == leader_name: leader = '*' - host = get_conn_kwargs(m.conn_url)['host'] + host = m.conn_kwargs()['host'] xlog_location = m.data.get('xlog_location') or 0 lag = '' diff --git a/patroni/dcs/__init__.py b/patroni/dcs/__init__.py index 28fda908..7b291d55 100644 --- a/patroni/dcs/__init__.py +++ b/patroni/dcs/__init__.py @@ -86,6 +86,26 @@ class Member(namedtuple('Member', 'index,name,session,data')): def conn_url(self): return self.data.get('conn_url') + def conn_kwargs(self, auth=None): + ret = self.data.get('conn_kwargs') + if ret: + ret = ret.copy() + else: + r = urlparse(self.conn_url) + ret = { + 'host': r.hostname, + 'port': r.port or 5432, + 'database': r.path[1:] + } + self.data['conn_kwargs'] = ret.copy() + + if auth and isinstance(auth, dict): + if 'username' in auth: + ret['user'] = auth['username'] + if 'password' in auth: + ret['password'] = auth['password'] + return ret + @property def api_url(self): return self.data.get('api_url') @@ -104,7 +124,7 @@ class Member(namedtuple('Member', 'index,name,session,data')): @property def clonefrom(self): - return self.tags.get('clonefrom', False) + return self.tags.get('clonefrom', False) and bool(self.conn_url) class Leader(namedtuple('Leader', 'index,session,member')): @@ -119,6 +139,9 @@ class Leader(namedtuple('Leader', 'index,session,member')): def name(self): return self.member.name + def conn_kwargs(self, auth=None): + return self.member.conn_kwargs(auth) + @property def conn_url(self): return self.member.conn_url diff --git a/patroni/ha.py b/patroni/ha.py index 79cc8b04..8c044ba8 100644 --- a/patroni/ha.py +++ b/patroni/ha.py @@ -141,9 +141,8 @@ class Ha(object): node_to_follow = self._get_node_to_follow(self.cluster) - if not self.state_handler.check_recovery_conf(node_to_follow) or recovery: - self._async_executor.schedule('changing primary_conninfo and restarting') - self._async_executor.run_async(self.state_handler.follow, (node_to_follow, self.cluster.leader, recovery)) + self.state_handler.follow(node_to_follow, self.cluster.leader, recovery, self._async_executor) + return ret def enforce_master_role(self, message, promote_message): diff --git a/patroni/postgresql.py b/patroni/postgresql.py index d8bb1b60..7dff8b32 100644 --- a/patroni/postgresql.py +++ b/patroni/postgresql.py @@ -10,7 +10,6 @@ import time from patroni.exceptions import PostgresConnectionException, PostgresException from patroni.utils import compare_values, parse_bool, parse_int, Retry, RetryFailedError from six import string_types -from six.moves.urllib_parse import urlparse from threading import Lock logger = logging.getLogger(__name__) @@ -22,24 +21,6 @@ ACTION_ON_RELOAD = "on_reload" ACTION_ON_ROLE_CHANGE = "on_role_change" -def get_conn_kwargs(url, auth=None): - r = urlparse(url) - ret = { - 'host': r.hostname, - 'port': r.port or 5432, - 'database': r.path[1:], - 'fallback_application_name': 'Patroni', - 'connect_timeout': 3, - 'options': '-c statement_timeout=2000', - } - if auth and isinstance(auth, dict): - if 'username' in auth: - ret['user'] = auth['username'] - if 'password' in auth: - ret['password'] = auth['password'] - return ret - - class Postgresql(object): # List of parameters which must be always passed to postmaster as command line options @@ -147,8 +128,8 @@ class Postgresql(object): def resolve_connection_addresses(self): self._local_address = self.get_local_address() - self.connection_string = 'postgres://{connect_address}/{database}'.format( - connect_address=self._connect_address or self._local_address, database=self._database) + self.connection_string = 'postgres://{0}/{1}'.format( + self._connect_address or self._local_address['host'] + ':' + self._local_address['port'], self._database) def pg_ctl(self, cmd, *args, **kwargs): """Builds and executes pg_ctl command @@ -266,7 +247,7 @@ class Postgresql(object): if la.strip().lower() in ('*', '0.0.0.0', '127.0.0.1', 'localhost'): # we are listening on '*' or localhost local_address = 'localhost' # connection via localhost is preferred break - return local_address + ':' + self._server_parameters['port'] + return {'host': local_address, 'port': self._server_parameters['port']} def get_postgres_role_from_data_directory(self): if self.data_directory_empty(): @@ -277,12 +258,21 @@ class Postgresql(object): return 'master' @property - def _connect_kwargs(self): - return get_conn_kwargs('postgres://{0}/{1}'.format(self._local_address, self._database), self._superuser) + def _local_connect_kwargs(self): + ret = self._local_address.copy() + ret.update({'database': self._database, + 'fallback_application_name': 'Patroni', + 'connect_timeout': 3, + 'options': '-c statement_timeout=2000'}) + if 'username' in self._superuser: + ret['user'] = self._superuser['username'] + if 'password' in self._superuser: + ret['password'] = self._superuser['password'] + return ret def connection(self): if not self._connection or self._connection.closed != 0: - self._connection = psycopg2.connect(**self._connect_kwargs) + self._connection = psycopg2.connect(**self._local_connect_kwargs) self._connection.autocommit = True self.server_version = self._connection.server_version return self._connection @@ -403,7 +393,7 @@ class Postgresql(object): replica_methods = self.config.get('create_replica_method') or ['basebackup'] if clone_member: - r = get_conn_kwargs(clone_member.conn_url, self._replication) + r = clone_member.conn_kwargs(self._replication) connstring = 'postgres://{user}@{host}:{port}/{database}'.format(**r) # add the credentials to connect to the replica origin to pgpass. env = self.write_pgpass(r) @@ -538,7 +528,7 @@ class Postgresql(object): def checkpoint(self, connect_kwargs=None): check_not_is_in_recovery = connect_kwargs is not None - connect_kwargs = connect_kwargs or self._connect_kwargs + connect_kwargs = connect_kwargs or self._local_connect_kwargs for p in ['connect_timeout', 'options']: connect_kwargs.pop(p, None) try: @@ -624,29 +614,29 @@ class Postgresql(object): with open(os.path.join(self._data_dir, 'pg_hba.conf'), 'a') as f: f.write('\n{}\n'.format('\n'.join(config))) - def primary_conninfo(self, node_to_follow_url): - r = get_conn_kwargs(node_to_follow_url, self._replication) + def primary_conninfo(self, member): + if not (member and member.conn_url): + return None + r = member.conn_kwargs(self._replication) r.update({'application_name': self.name, 'sslmode': 'prefer', 'sslcompression': '1'}) keywords = 'user password host port sslmode sslcompression application_name'.split() return ' '.join('{0}={{{0}}}'.format(kw) for kw in keywords).format(**r) - def check_recovery_conf(self, node_to_follow): + def check_recovery_conf(self, primary_conninfo): if not os.path.isfile(self._recovery_conf): return False - pattern = node_to_follow and node_to_follow.conn_url and self.primary_conninfo(node_to_follow.conn_url) - with open(self._recovery_conf, 'r') as f: for line in f: if line.startswith('primary_conninfo'): - return pattern and (pattern in line) - return not pattern + return primary_conninfo and (primary_conninfo in line) + return not primary_conninfo - def write_recovery_conf(self, node_to_follow): + def write_recovery_conf(self, primary_conninfo): with open(self._recovery_conf, 'w') as f: f.write("standby_mode = 'on'\nrecovery_target_timeline = 'latest'\n") - if node_to_follow and node_to_follow.conn_url: - f.write("primary_conninfo = '{0}'\n".format(self.primary_conninfo(node_to_follow.conn_url))) + if primary_conninfo: + f.write("primary_conninfo = '{0}'\n".format(primary_conninfo)) if self.use_slots: f.write("primary_slot_name = '{0}'\n".format(self.name)) for name, value in self.config.get('recovery_conf', {}).items(): @@ -723,24 +713,33 @@ class Postgresql(object): except OSError: logger.exception("Unable to list %s", status_dir) - def follow(self, member, leader, recovery=False): - if self.check_recovery_conf(member) and not recovery: + def follow(self, member, leader, recovery=False, async_executor=None): + primary_conninfo = self.primary_conninfo(member) + + if self.check_recovery_conf(primary_conninfo) and not recovery: return True + if async_executor: + async_executor.schedule('changing primary_conninfo and restarting') + async_executor.run_async(self._do_follow, (primary_conninfo, leader, recovery)) + else: + self._do_follow(primary_conninfo, leader, recovery) + + def _do_follow(self, primary_conninfo, leader, recovery=False): change_role = self.role == 'master' if change_role: if leader: if leader.name == self.name: self._need_rewind = False - member = None + primary_conninfo = None if self.is_running(): return else: self._need_rewind = bool(leader.conn_url) and self.can_rewind else: self._need_rewind = False - member = None + primary_conninfo = None if self._need_rewind: logger.info("set the rewind flag after demote") @@ -753,7 +752,7 @@ class Postgresql(object): return logger.info('Leader unknown, can not rewind') # prepare pg_rewind connection - r = get_conn_kwargs(leader.conn_url, self._superuser) + r = leader.conn_kwargs(self._superuser) # first make sure that we are really trying to rewind # from the master and run a checkpoint on a t in order to @@ -779,7 +778,7 @@ class Postgresql(object): self.single_user_mode(options=opts) if self.rewind(r) or not self.config.get('remove_data_directory_on_rewind_failure', False): - self.write_recovery_conf(member) + self.write_recovery_conf(primary_conninfo) ret = self.start() else: logger.error('unable to rewind the former master') @@ -788,7 +787,7 @@ class Postgresql(object): ret = True self._need_rewind = False else: - self.write_recovery_conf(member) + self.write_recovery_conf(primary_conninfo) ret = self.restart() self.set_role('replica') diff --git a/tests/test_ctl.py b/tests/test_ctl.py index ccbc0b37..0e4ea51b 100644 --- a/tests/test_ctl.py +++ b/tests/test_ctl.py @@ -235,7 +235,8 @@ class TestCtl(unittest.TestCase): with patch('requests.delete', Mock(return_value=MockResponse(500))): # normal restart, the schedule is actually parsed, but not validated in patronictl - result = self.runner.invoke(ctl, ['restart', 'alpha', 'other', '--force', '--scheduled', '2300-10-01T14:30']) + result = self.runner.invoke(ctl, ['restart', 'alpha', 'other', '--force', + '--scheduled', '2300-10-01T14:30']) assert 'Failed: flush scheduled restart' in result.output with patch('requests.post', Mock(return_value=MockResponse())): @@ -253,7 +254,8 @@ class TestCtl(unittest.TestCase): # force restart with restart already present with patch('patroni.ctl.request_patroni', Mock(return_value=MockResponse(204))): - result = self.runner.invoke(ctl, ['restart', 'alpha', 'other', '--force', '--scheduled', '2300-10-01T14:30']) + result = self.runner.invoke(ctl, ['restart', 'alpha', 'other', '--force', + '--scheduled', '2300-10-01T14:30']) assert result.exit_code == 0 with patch('requests.post', Mock(return_value=MockResponse(202))): diff --git a/tests/test_postgresql.py b/tests/test_postgresql.py index bb3b17c6..2ab29124 100644 --- a/tests/test_postgresql.py +++ b/tests/test_postgresql.py @@ -278,7 +278,7 @@ class TestPostgresql(unittest.TestCase): mock_pg_rewind.return_value = True self.p.follow(self.leader, self.leader) - self.assertTrue(self.p.follow(None, None)) # check_recovery_conf... + self.p.follow(None, None) # check_recovery_conf... @patch('subprocess.check_output', Mock(return_value=0, side_effect=pg_controldata_string)) def test_can_rewind(self):