diff --git a/patroni/ctl.py b/patroni/ctl.py index 499f6563..ad41dacd 100644 --- a/patroni/ctl.py +++ b/patroni/ctl.py @@ -18,7 +18,7 @@ import yaml from click import ClickException from patroni.dcs import get_dcs as _get_dcs from patroni.exceptions import PatroniException -from patroni.postgresql import parseurl +from patroni.postgresql import get_conn_kwargs from prettytable import PrettyTable from six.moves.urllib_parse import urlparse @@ -165,14 +165,13 @@ def watching(w, watch, max_count=None, clear=True): yield 0 -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'] - params['fallback_application_name'] = 'Patroni ctl' - params['connect_timeout'] = '5' - +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 @@ -195,7 +194,7 @@ def get_any_member(cluster, role='master', member=None): return m -def get_cursor(cluster, role='master', member=None, connect_parameters=None): +def get_cursor(cluster, connect_parameters, role='master', member=None): member = get_any_member(cluster, role=role, member=member) if member is None: return None @@ -237,7 +236,7 @@ def dsn(cluster_name, config_file, dcs, role, member): if m is None: raise PatroniCtlException('Can not find a suitable member') - params = build_connect_parameters(m.conn_url) + params = get_conn_kwargs(m.conn_url) click.echo('host={host} port={port}'.format(**params)) @@ -287,7 +286,7 @@ def query( connect_parameters = dict() if username: - connect_parameters['user'] = username + connect_parameters['username'] = username if password: connect_parameters['password'] = click.prompt('Password', hide_input=True, type=str) if dbname: @@ -308,10 +307,10 @@ def query( cluster = dcs.get_cluster() -def query_member(cluster, cursor, member, role, command, connect_parameters=None): +def query_member(cluster, cursor, member, role, command, connect_parameters): try: if cursor is None: - cursor = get_cursor(cluster, role=role, member=member, connect_parameters=connect_parameters) + cursor = get_cursor(cluster, connect_parameters, role=role, member=member) if cursor is None: if role is None: @@ -570,7 +569,7 @@ def output_members(cluster, name, fmt='pretty'): if m.name == leader_name: leader = '*' - host = build_connect_parameters(m.conn_url)['host'] + host = get_conn_kwargs(m.conn_url)['host'] xlog_location = m.data.get('xlog_location') or 0 lag = '' diff --git a/patroni/postgresql.py b/patroni/postgresql.py index aa3adf0c..4016c6d7 100644 --- a/patroni/postgresql.py +++ b/patroni/postgresql.py @@ -22,7 +22,7 @@ ACTION_ON_RELOAD = "on_reload" ACTION_ON_ROLE_CHANGE = "on_role_change" -def parseurl(url): +def get_conn_kwargs(url, auth=None): r = urlparse(url) ret = { 'host': r.hostname, @@ -32,10 +32,11 @@ def parseurl(url): 'connect_timeout': 3, 'options': '-c statement_timeout=2000', } - if r.username: - ret['user'] = r.username - if r.password: - ret['password'] = r.password + if auth and isinstance(auth, dict): + if 'username' in auth: + ret['user'] = auth['username'] + if 'password' in auth: + ret['password'] = auth['password'] return ret @@ -148,8 +149,8 @@ class Postgresql(object): def resolve_connection_addresses(self): self._local_address = self.get_local_address() - self.connection_string = 'postgres://{username}:{password}@{connect_address}/{database}'.format( - connect_address=self._connect_address or self._local_address, database=self._database, **self._replication) + self.connection_string = 'postgres://{connect_address}/{database}'.format( + connect_address=self._connect_address or self._local_address, database=self._database) def reload_config(self, config): server_parameters = self.get_server_parameters(config) @@ -248,7 +249,7 @@ class Postgresql(object): local_address = listen_addresses[0].strip() # take first address from listen_addresses for la in listen_addresses: - if la.strip() in ('*', '0.0.0.0', '127.0.0.1', 'localhost'): # we are listening on '*' or localhost + 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'] @@ -263,12 +264,7 @@ class Postgresql(object): @property def _connect_kwargs(self): - r = parseurl('postgres://{0}/{1}'.format(self._local_address, self._database)) - if 'username' in self._superuser: - r['user'] = self._superuser['username'] - if 'password' in self._superuser: - r['password'] = self._superuser['password'] - return r + return get_conn_kwargs('postgres://{0}/{1}'.format(self._local_address, self._database), self._superuser) def connection(self): if not self._connection or self._connection.closed != 0: @@ -392,7 +388,7 @@ class Postgresql(object): replica_methods = self.config.get('create_replica_method') or ['basebackup'] if clone_member: - r = parseurl(clone_member.conn_url) + r = get_conn_kwargs(clone_member.conn_url, 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) @@ -606,17 +602,17 @@ 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, leader_url): - r = parseurl(leader_url) + def primary_conninfo(self, node_to_follow_url): + r = get_conn_kwargs(node_to_follow_url, 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, leader): + def check_recovery_conf(self, node_to_follow): if not os.path.isfile(self._recovery_conf): return False - pattern = leader and leader.conn_url and self.primary_conninfo(leader.conn_url) + 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: @@ -624,11 +620,11 @@ class Postgresql(object): return pattern and (pattern in line) return not pattern - def write_recovery_conf(self, leader): + def write_recovery_conf(self, node_to_follow): with open(self._recovery_conf, 'w') as f: f.write("standby_mode = 'on'\nrecovery_target_timeline = 'latest'\n") - if leader and leader.conn_url: - f.write("primary_conninfo = '{0}'\n".format(self.primary_conninfo(leader.conn_url))) + 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 self.use_slots: f.write("primary_slot_name = '{0}'\n".format(self.name)) for name, value in self.config.get('recovery_conf', {}).items(): @@ -637,10 +633,7 @@ class Postgresql(object): def rewind(self, leader): # prepare pg_rewind connection - r = parseurl(leader.conn_url) - r.update(self._superuser) - r['user'] = r.pop('username') - r['database'] = self._database + r = get_conn_kwargs(leader.conn_url, self._superuser) env = self.write_pgpass(r) pc = "user={user} host={host} port={port} dbname={database} sslmode=prefer sslcompression=1".format(**r) # first run a checkpoint on a promoted master in order @@ -656,7 +649,8 @@ class Postgresql(object): def controldata(self): """ return the contents of pg_controldata, or non-True value if pg_controldata call failed """ result = {} - if self.state != 'creating replica': # Don't try to call pg_controldata during backup restore + # Don't try to call pg_controldata during backup restore + if self._version_file_exists() and self.state != 'creating replica': try: data = subprocess.check_output(['pg_controldata', self._data_dir]) if data: diff --git a/tests/test_ctl.py b/tests/test_ctl.py index eb63e153..7e4d8a03 100644 --- a/tests/test_ctl.py +++ b/tests/test_ctl.py @@ -55,14 +55,14 @@ class TestCtl(unittest.TestCase): @patch('psycopg2.connect', psycopg2_connect) def test_get_cursor(self): - self.assertIsNone(get_cursor(get_cluster_initialized_without_leader(), role='master')) + self.assertIsNone(get_cursor(get_cluster_initialized_without_leader(), {}, role='master')) - self.assertIsNotNone(get_cursor(get_cluster_initialized_with_leader(), role='master')) + self.assertIsNotNone(get_cursor(get_cluster_initialized_with_leader(), {}, role='master')) # MockCursor returns pg_is_in_recovery as false - self.assertIsNone(get_cursor(get_cluster_initialized_with_leader(), role='replica')) + self.assertIsNone(get_cursor(get_cluster_initialized_with_leader(), {}, role='replica')) - self.assertIsNotNone(get_cursor(get_cluster_initialized_with_leader(), role='any')) + self.assertIsNotNone(get_cursor(get_cluster_initialized_with_leader(), {'database': 'foo'}, role='any')) def test_parse_dcs(self): assert parse_dcs(None) is None @@ -183,24 +183,24 @@ class TestCtl(unittest.TestCase): def test_query_member(self): with patch('patroni.ctl.get_cursor', Mock(return_value=MockConnect().cursor())): - rows = query_member(None, None, None, 'master', 'SELECT pg_is_in_recovery()') + rows = query_member(None, None, None, 'master', 'SELECT pg_is_in_recovery()', {}) self.assertTrue('False' in str(rows)) - rows = query_member(None, None, None, 'replica', 'SELECT pg_is_in_recovery()') + rows = query_member(None, None, None, 'replica', 'SELECT pg_is_in_recovery()', {}) self.assertEquals(rows, (None, None)) with patch('test_postgresql.MockCursor.execute', Mock(side_effect=OperationalError('bla'))): - rows = query_member(None, None, None, 'replica', 'SELECT pg_is_in_recovery()') + rows = query_member(None, None, None, 'replica', 'SELECT pg_is_in_recovery()', {}) with patch('patroni.ctl.get_cursor', Mock(return_value=None)): - rows = query_member(None, None, None, None, 'SELECT pg_is_in_recovery()') + rows = query_member(None, None, None, None, 'SELECT pg_is_in_recovery()', {}) self.assertTrue('No connection to' in str(rows)) - rows = query_member(None, None, None, 'replica', 'SELECT pg_is_in_recovery()') + rows = query_member(None, None, None, 'replica', 'SELECT pg_is_in_recovery()', {}) self.assertTrue('No connection to' in str(rows)) with patch('patroni.ctl.get_cursor', Mock(side_effect=OperationalError('bla'))): - rows = query_member(None, None, None, 'replica', 'SELECT pg_is_in_recovery()') + rows = query_member(None, None, None, 'replica', 'SELECT pg_is_in_recovery()', {}) @patch('patroni.ctl.get_dcs') def test_dsn(self, mock_get_dcs): diff --git a/tests/test_postgresql.py b/tests/test_postgresql.py index 7fa7279a..ec0151a2 100644 --- a/tests/test_postgresql.py +++ b/tests/test_postgresql.py @@ -396,6 +396,7 @@ class TestPostgresql(unittest.TestCase): self.p.remove_data_directory() self.p.remove_data_directory() + @patch('patroni.postgresql.Postgresql._version_file_exists', Mock(return_value=True)) def test_controldata(self): with patch('subprocess.check_output', Mock(return_value=0, side_effect=pg_controldata_string)): data = self.p.controldata() @@ -466,6 +467,7 @@ class TestPostgresql(unittest.TestCase): mock_unlink.assert_not_called() mock_remove.assert_not_called() + @patch('patroni.postgresql.Postgresql._version_file_exists', Mock(return_value=True)) @patch('subprocess.check_output', MagicMock(return_value=0, side_effect=pg_controldata_string)) def test_sysid(self): self.assertEqual(self.p.sysid, "6200971513092291716")