Don't expose replication user/passwd in DCS

This commit is contained in:
Alexander Kukushkin
2016-06-15 09:34:04 +02:00
parent 25f20ca7d7
commit 57807ff337
4 changed files with 47 additions and 52 deletions
+14 -15
View File
@@ -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 = ''
+21 -27
View File
@@ -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:
+10 -10
View File
@@ -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):
+2
View File
@@ -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")