diff --git a/patroni/postgresql.py b/patroni/postgresql.py index 51c3db13..5d0e27da 100644 --- a/patroni/postgresql.py +++ b/patroni/postgresql.py @@ -47,7 +47,7 @@ class Postgresql: self.replication = config['replication'] self.superuser = config['superuser'] self.admin = config['admin'] - self.pg_rewind = config.get('pg_rewind', {}) + self._pg_rewind = config.get('pg_rewind', {}) self.callback = config.get('callbacks', {}) self.use_slots = config.get('use_slots', True) self.schedule_load_slots = self.use_slots @@ -70,16 +70,19 @@ class Postgresql: self._cursor_holder = None self.members = [] # list of already existing replication slots self.retry = Retry(max_tries=-1, deadline=10, max_delay=1, retry_exceptions=PostgresConnectionException) + self.init_pg_rewind() + + def init_pg_rewind(self): try: - self._pg_rewind_present = ('username' in self.pg_rewind and + self._pg_rewind_present = ('username' in self._pg_rewind and ('wal_log_hints' in self.config['parameters'] or 'data_checksums' in self.config['parameters']) and os.system("pg_rewind --version >/dev/null 2>&1") == 0) if self._pg_rewind_present: - self.pg_rewind['user'] = self.pg_rewind['username'] + self._pg_rewind['user'] = self._pg_rewind['username'] except: self._pg_rewind_present = False - if self.pg_rewind and not self._pg_rewind_present: + if self._pg_rewind and not self._pg_rewind_present: logger.warning("pg_rewind support is disabled") def get_local_address(self): @@ -310,6 +313,18 @@ recovery_target_timeline = 'latest' for name, value in self.config.get('recovery_conf', {}).items(): f.write("{} = '{}'\n".format(name, value)) + def pg_rewind(self, leader): + pc = self.primary_conninfo(leader.conn_url, self._pg_rewind) + ' dbname=postgres' + logger.info("running pg_rewind from {}".format(pc)) + pg_rewind = ['pg_rewind', '-D', self.data_dir, '--source-server', pc] + try: + ret = (subprocess.call(pg_rewind) == 0) + except: + ret = False + if ret: + self.write_recovery_conf(leader) + return ret + def follow_the_leader(self, leader): if not self.check_recovery_conf(leader): self.write_recovery_conf(leader) @@ -317,20 +332,11 @@ recovery_target_timeline = 'latest' if leader and change_role and self._pg_rewind_present: self.stop() - pc = self.primary_conninfo(leader.conn_url, - self.pg_rewind) + ' dbname=postgres' - logger.info("running pg_rewind from {}".format(pc)) - pg_rewind = ['pg_rewind', '-D', self.data_dir, '--source-server', pc] - try: - ret = (subprocess.call(pg_rewind) == 0) - except: - ret = False - # pg_rewind removes recovery.conf, we have to reinstate it. - if ret: - self.write_recovery_conf(leader) + if self.pg_rewind(leader): ret = self.start() else: - self.remove_data_directory() + ret = False + self.move_data_directory() logger.error("unable to rewind the former leader") else: ret = self.restart() diff --git a/tests/test_postgresql.py b/tests/test_postgresql.py index eaf84943..35019780 100644 --- a/tests/test_postgresql.py +++ b/tests/test_postgresql.py @@ -9,6 +9,7 @@ from patroni.exceptions import PostgresConnectionException from patroni.postgresql import Postgresql from patroni.utils import RetryFailedError from test_ha import false +import subprocess class MockCursor: @@ -132,12 +133,32 @@ class TestPostgresql(unittest.TestCase): def test_sync_from_leader(self): self.assertTrue(self.p.sync_from_leader(self.leader)) - def test_follow_the_leader(self): + @patch('os.system', side_effect=Exception("Test")) + def test_init_pg_rewind(self, mock_system): + self.p.init_pg_rewind() + # prepare parameters for pg_rewind + self.p._pg_rewind = {'username': 'foo'} + self.p.config['parameters']['data_checksums'] = 1 + os.system = mock_system + self.p.init_pg_rewind() + + @patch('subprocess.call', side_effect=Exception("Test")) + def test_pg_rewind(self, mock_call): + self.assertTrue(self.p.pg_rewind(self.leader)) + self.p + subprocess.call = mock_call + self.assertFalse(self.p.pg_rewind(self.leader)) + + @patch('patroni.postgresql.Postgresql.pg_rewind', return_value=False) + def test_follow_the_leader(self, mock_pg_rewind): self.p.demote(self.leader) self.p.follow_the_leader(None) + self.p._pg_rewind_present = True self.p.demote(self.leader) self.p.follow_the_leader(self.leader) self.p.follow_the_leader(Leader(-1, None, 28, self.other)) + self.p.pg_rewind = mock_pg_rewind + self.p.follow_the_leader(self.leader) def test_create_replica(self): self.p.delete_trigger_file = Mock(side_effect=OSError())