From 7941c8677567e7e65db4f3f0a2d127758a591d23 Mon Sep 17 00:00:00 2001 From: Alexander Kukushkin Date: Thu, 11 May 2023 09:58:15 +0200 Subject: [PATCH] Refactor write_sync_state() (#2669) Make it return the new `SyncState` object in order to avoid reading the new cluster state in the Ha.process_sync_replication(). Now it is a small optimization, but it will become very handy in the quorum commit feature. --- patroni/dcs/__init__.py | 12 +++++++----- patroni/dcs/consul.py | 13 +++++++++++-- patroni/dcs/etcd.py | 6 +++--- patroni/dcs/etcd3.py | 8 +++++--- patroni/dcs/kubernetes.py | 19 ++++++++++--------- patroni/dcs/raft.py | 38 +++++++++++++++++++++----------------- patroni/dcs/zookeeper.py | 27 ++++++++++++++------------- patroni/ha.py | 24 ++++++++---------------- tests/test_consul.py | 8 +++++++- tests/test_etcd.py | 2 +- tests/test_etcd3.py | 9 ++++++--- tests/test_ha.py | 20 ++++++++++---------- tests/test_kubernetes.py | 8 ++++++-- tests/test_raft.py | 1 + 14 files changed, 110 insertions(+), 85 deletions(-) diff --git a/patroni/dcs/__init__.py b/patroni/dcs/__init__.py index 2c2d0187..f1429c6f 100644 --- a/patroni/dcs/__init__.py +++ b/patroni/dcs/__init__.py @@ -1087,28 +1087,30 @@ class AbstractDCS(abc.ABC): return {'leader': leader, 'sync_standby': ','.join(sorted(sync_standby)) if sync_standby else None} def write_sync_state(self, leader: Optional[str], sync_standby: Optional[Collection[str]], - index: Optional[Any] = None) -> bool: + index: Optional[Any] = None) -> Optional[SyncState]: """Write the new synchronous state to DCS. Calls :func:`sync_state` method to build a dict and than calls DCS specific :func:`set_sync_state_value` method. :param leader: name of the leader node that manages /sync key :param sync_standby: collection of currently known synchronous standby node names :param index: for conditional update of the key/object - :returns: `True` if /sync key was successfully updated + :returns: the new :class:`SyncState` object or None """ sync_value = self.sync_state(leader, sync_standby) - return self.set_sync_state_value(json.dumps(sync_value, separators=(',', ':')), index) + ret = self.set_sync_state_value(json.dumps(sync_value, separators=(',', ':')), index) + if not isinstance(ret, bool): + return SyncState.from_node(ret, sync_value) @abc.abstractmethod def set_history_value(self, value: str) -> bool: """""" @abc.abstractmethod - def set_sync_state_value(self, value: str, index: Optional[Any] = None) -> bool: + def set_sync_state_value(self, value: str, index: Optional[Any] = None) -> Union[Any, bool]: """Set synchronous state in DCS, should be implemented in the child class. :param value: the new value of /sync key :param index: for conditional update of the key/object - :returns: `True` if key/object was successfully updated + :returns: version of the new object or `False` in case of error """ @abc.abstractmethod diff --git a/patroni/dcs/consul.py b/patroni/dcs/consul.py index 5a2e5700..2d2d6708 100644 --- a/patroni/dcs/consul.py +++ b/patroni/dcs/consul.py @@ -658,8 +658,17 @@ class Consul(AbstractDCS): return True @catch_consul_errors - def set_sync_state_value(self, value: str, index: Optional[int] = None) -> bool: - return self.retry(self._client.kv.put, self.sync_path, value, cas=index) + def set_sync_state_value(self, value: str, index: Optional[int] = None) -> Union[int, bool]: + retry = self._retry.copy() + ret = retry(self._client.kv.put, self.sync_path, value, cas=index) + if ret: # We have no other choise, only read after write :( + retry.deadline = retry.stoptime - time.time() + if retry.deadline < 0.5: + return False + _, ret = self.retry(self._client.kv.get, self.sync_path) + if ret and (ret.get('Value') or b'').decode('utf-8') == value: + return ret['ModifyIndex'] + return False @catch_consul_errors def delete_sync_state(self, index: Optional[int] = None) -> bool: diff --git a/patroni/dcs/etcd.py b/patroni/dcs/etcd.py index c2a83dab..e53d8a67 100644 --- a/patroni/dcs/etcd.py +++ b/patroni/dcs/etcd.py @@ -625,7 +625,7 @@ class AbstractEtcd(AbstractDCS): def catch_etcd_errors(func: Callable[..., Any]) -> Any: def wrapper(self: AbstractEtcd, *args: Any, **kwargs: Any) -> Any: try: - retval = func(self, *args, **kwargs) is not None + retval = func(self, *args, **kwargs) self._has_failed = False return retval except (RetryFailedError, etcd.EtcdException) as e: @@ -817,8 +817,8 @@ class Etcd(AbstractEtcd): return bool(self._client.write(self.history_path, value)) @catch_etcd_errors - def set_sync_state_value(self, value: str, index: Optional[int] = None) -> bool: - return bool(self.retry(self._client.write, self.sync_path, value, prevIndex=index or 0)) + def set_sync_state_value(self, value: str, index: Optional[int] = None) -> Union[int, bool]: + return self.retry(self._client.write, self.sync_path, value, prevIndex=index or 0).modifiedIndex @catch_etcd_errors def delete_sync_state(self, index: Optional[int] = None) -> bool: diff --git a/patroni/dcs/etcd3.py b/patroni/dcs/etcd3.py index b526e1e2..bdf5cc33 100644 --- a/patroni/dcs/etcd3.py +++ b/patroni/dcs/etcd3.py @@ -340,7 +340,8 @@ class Etcd3Client(AbstractEtcdClientWithFailover): return self.call_rpc('/lease/keepalive', {'ID': ID}, retry).get('result', {}).get('TTL') def txn(self, compare: Dict[str, Any], success: Dict[str, Any], retry: Optional[Retry] = None) -> Dict[str, Any]: - return self.call_rpc('/kv/txn', {'compare': [compare], 'success': [success]}, retry).get('succeeded', {}) + ret = self.call_rpc('/kv/txn', {'compare': [compare], 'success': [success]}, retry) + return ret if ret.get('succeeded') else {} @_handle_auth_errors def put(self, key: str, value: str, lease: Optional[str] = None, create_revision: Optional[str] = None, @@ -908,8 +909,9 @@ class Etcd3(AbstractEtcd): return bool(self._client.put(self.history_path, value)) @catch_etcd_errors - def set_sync_state_value(self, value: str, index: Optional[str] = None) -> bool: - return self.retry(self._client.put, self.sync_path, value, mod_revision=index) + def set_sync_state_value(self, value: str, index: Optional[str] = None) -> Union[str, bool]: + return self.retry(self._client.put, self.sync_path, value, mod_revision=index)\ + .get('header', {}).get('revision', False) @catch_etcd_errors def delete_sync_state(self, index: Optional[str] = None) -> bool: diff --git a/patroni/dcs/kubernetes.py b/patroni/dcs/kubernetes.py index 61cebc10..e33a8509 100644 --- a/patroni/dcs/kubernetes.py +++ b/patroni/dcs/kubernetes.py @@ -1078,10 +1078,9 @@ class Kubernetes(AbstractDCS): @catch_kubernetes_errors def patch_or_create(self, name: str, annotations: Dict[str, Any], resource_version: Optional[str] = None, - patch: bool = False, retry: bool = True, ips: Optional[List[str]] = None) -> bool: + patch: bool = False, retry: bool = True, ips: Optional[List[str]] = None) -> K8sObject: try: - return bool(self._patch_or_create(name, annotations, resource_version, - patch, self.retry if retry else None, ips)) + return self._patch_or_create(name, annotations, resource_version, patch, self.retry if retry else None, ips) except k8s_client.rest.ApiException as e: if e.status == 409 and resource_version: # Conflict in resource_version # Terminate watchers, it could be a sign that K8s API is in a failed state @@ -1095,7 +1094,7 @@ class Kubernetes(AbstractDCS): if self._api.use_endpoints and not patch and not resource_version: self._should_create_config_service = True self._create_config_service() - return self.patch_or_create(self.config_path, annotations, resource_version, patch, retry) + return bool(self.patch_or_create(self.config_path, annotations, resource_version, patch, retry)) def _create_config_service(self) -> None: metadata = k8s_client.V1ObjectMeta(namespace=self._namespace, name=self.config_path, labels=self._labels) @@ -1242,7 +1241,7 @@ class Kubernetes(AbstractDCS): annotations = {'leader': leader or None, 'member': candidate or None, 'scheduled_at': scheduled_at and scheduled_at.isoformat()} patch = bool(self.cluster and isinstance(self.cluster.failover, Failover) and self.cluster.failover.index) - return self.patch_or_create(self.failover_path, annotations, index, bool(index or patch), False) + return bool(self.patch_or_create(self.failover_path, annotations, index, bool(index or patch), False)) @property def _config_resource_version(self) -> Optional[str]: @@ -1314,16 +1313,18 @@ class Kubernetes(AbstractDCS): raise NotImplementedError # pragma: no cover def write_sync_state(self, leader: Optional[str], sync_standby: Optional[Collection[str]], - index: Optional[str] = None) -> bool: + index: Optional[str] = None) -> Optional[SyncState]: """Prepare and write annotations to $SCOPE-sync Endpoint or ConfigMap. :param leader: name of the leader node that manages /sync key :param sync_standby: collection of currently known synchronous standby node names :param index: last known `resource_version` for conditional update of the object - :returns: `True` if update was successful + :returns: the new :class:`SyncState` object or None """ sync_state = self.sync_state(leader, sync_standby) - return self.patch_or_create(self.sync_path, sync_state, index, False) + ret = self.patch_or_create(self.sync_path, sync_state, index, False) + if not isinstance(ret, bool): + return SyncState.from_node(ret.metadata.resource_version, sync_state) def delete_sync_state(self, index: Optional[str] = None) -> bool: """Patch annotations of $SCOPE-sync Endpoint or ConfigMap with empty values. @@ -1332,7 +1333,7 @@ class Kubernetes(AbstractDCS): :param index: last known `resource_version` for conditional update of the object :returns: `True` if "delete" was successful """ - return self.write_sync_state(None, None, index=index) + return self.write_sync_state(None, None, index=index) is not None def watch(self, leader_index: Optional[str], timeout: float) -> bool: if self.__do_not_watch: diff --git a/patroni/dcs/raft.py b/patroni/dcs/raft.py index 1c563788..10113931 100644 --- a/patroni/dcs/raft.py +++ b/patroni/dcs/raft.py @@ -142,7 +142,7 @@ class KVStoreTTL(DynMemberSyncObj): def __check_requirements(old_value: Dict[str, Any], **kwargs: Any) -> bool: return bool(('prevExist' not in kwargs or bool(kwargs['prevExist']) == bool(old_value)) and ('prevValue' not in kwargs or old_value and old_value['value'] == kwargs['prevValue']) - and (not kwargs.get('prevIndex') or old_value and old_value['index'] == kwargs['prevIndex'])) + and (kwargs.get('prevIndex') is None or old_value and old_value['index'] == kwargs['prevIndex'])) def set_retry_timeout(self, retry_timeout: int) -> None: self.__retry_timeout = retry_timeout @@ -175,7 +175,7 @@ class KVStoreTTL(DynMemberSyncObj): return False @replicated - def _set(self, key: str, value: Dict[str, Any], **kwargs: Any) -> bool: + def _set(self, key: str, value: Dict[str, Any], **kwargs: Any) -> Union[bool, Dict[str, Any]]: old_value = self.__data.get(key, {}) if not self.__check_requirements(old_value, **kwargs): return False @@ -187,10 +187,10 @@ class KVStoreTTL(DynMemberSyncObj): self.__data[key] = value if self.__on_set: self.__on_set(key, value) - return True + return value def set(self, key: str, value: str, ttl: Optional[int] = None, - handle_raft_error: bool = True, **kwargs: Any) -> bool: + handle_raft_error: bool = True, **kwargs: Any) -> Union[bool, Dict[str, Any]]: old_value = self.__data.get(key, {}) if not self.__check_requirements(old_value, **kwargs): return False @@ -411,39 +411,40 @@ class Raft(AbstractDCS): return loader(path) def _write_leader_optime(self, last_lsn: str) -> bool: - return self._sync_obj.set(self.leader_optime_path, last_lsn, timeout=1) + return self._sync_obj.set(self.leader_optime_path, last_lsn, timeout=1) is not False def _write_status(self, value: str) -> bool: - return self._sync_obj.set(self.status_path, value, timeout=1) + return self._sync_obj.set(self.status_path, value, timeout=1) is not False def _write_failsafe(self, value: str) -> bool: - return self._sync_obj.set(self.failsafe_path, value, timeout=1) + return self._sync_obj.set(self.failsafe_path, value, timeout=1) is not False def _update_leader(self) -> bool: ret = self._sync_obj.set(self.leader_path, self._name, ttl=self._ttl, - handle_raft_error=False, prevValue=self._name) + handle_raft_error=False, prevValue=self._name) is not False if not ret and self._sync_obj.get(self.leader_path) is None: ret = self.attempt_to_acquire_leader() return ret def attempt_to_acquire_leader(self) -> bool: - return self._sync_obj.set(self.leader_path, self._name, ttl=self._ttl, handle_raft_error=False, prevExist=False) + return self._sync_obj.set(self.leader_path, self._name, ttl=self._ttl, + handle_raft_error=False, prevExist=False) is not False def set_failover_value(self, value: str, index: Optional[int] = None) -> bool: - return self._sync_obj.set(self.failover_path, value, prevIndex=index) + return self._sync_obj.set(self.failover_path, value, prevIndex=index) is not False def set_config_value(self, value: str, index: Optional[int] = None) -> bool: - return self._sync_obj.set(self.config_path, value, prevIndex=index) + return self._sync_obj.set(self.config_path, value, prevIndex=index) is not False def touch_member(self, data: Dict[str, Any]) -> bool: value = json.dumps(data, separators=(',', ':')) - return self._sync_obj.set(self.member_path, value, self._ttl, timeout=2) + return self._sync_obj.set(self.member_path, value, self._ttl, timeout=2) is not False def take_leader(self) -> bool: - return self._sync_obj.set(self.leader_path, self._name, ttl=self._ttl) + return self._sync_obj.set(self.leader_path, self._name, ttl=self._ttl) is not False def initialize(self, create_new: bool = True, sysid: str = '') -> bool: - return self._sync_obj.set(self.initialize_path, sysid, prevExist=(not create_new)) + return self._sync_obj.set(self.initialize_path, sysid, prevExist=(not create_new)) is not False def _delete_leader(self) -> bool: return self._sync_obj.delete(self.leader_path, prevValue=self._name, timeout=1) @@ -455,10 +456,13 @@ class Raft(AbstractDCS): return self._sync_obj.delete(self.client_path(''), recursive=True) def set_history_value(self, value: str) -> bool: - return self._sync_obj.set(self.history_path, value) + return self._sync_obj.set(self.history_path, value) is not False - def set_sync_state_value(self, value: str, index: Optional[int] = None) -> bool: - return self._sync_obj.set(self.sync_path, value, prevIndex=index) + def set_sync_state_value(self, value: str, index: Optional[int] = None) -> Union[int, bool]: + ret = self._sync_obj.set(self.sync_path, value, prevIndex=index) + if isinstance(ret, dict): + return ret['index'] + return ret def delete_sync_state(self, index: Optional[int] = None) -> bool: return self._sync_obj.delete(self.sync_path, prevIndex=index) diff --git a/patroni/dcs/zookeeper.py b/patroni/dcs/zookeeper.py index eff6aceb..3d5f7b4b 100644 --- a/patroni/dcs/zookeeper.py +++ b/patroni/dcs/zookeeper.py @@ -358,19 +358,20 @@ class ZooKeeper(AbstractDCS): return False def _set_or_create(self, key: str, value: str, index: Optional[int] = None, - retry: bool = False, do_not_create_empty: bool = False) -> bool: + retry: bool = False, do_not_create_empty: bool = False) -> Union[int, bool]: value_bytes = value.encode('utf-8') try: if retry: - self._client.retry(self._client.set, key, value_bytes, version=index or -1) + ret = self._client.retry(self._client.set, key, value_bytes, version=index or -1) else: - self._client.set_async(key, value_bytes, version=index or -1).get(timeout=1) - return True + ret = self._client.set_async(key, value_bytes, version=index or -1).get(timeout=1) + return ret.version except NoNodeError: if do_not_create_empty and not value_bytes: return True elif index is None: - return self._create(key, value_bytes, retry) + if self._create(key, value_bytes, retry): + return 0 else: return False except Exception: @@ -378,10 +379,10 @@ class ZooKeeper(AbstractDCS): return False def set_failover_value(self, value: str, index: Optional[int] = None) -> bool: - return self._set_or_create(self.failover_path, value, index) + return self._set_or_create(self.failover_path, value, index) is not False def set_config_value(self, value: str, index: Optional[int] = None) -> bool: - return self._set_or_create(self.config_path, value, index, retry=True) + return self._set_or_create(self.config_path, value, index, retry=True) is not False def initialize(self, create_new: bool = True, sysid: str = "") -> bool: sysid_bytes = sysid.encode('utf-8') @@ -434,13 +435,13 @@ class ZooKeeper(AbstractDCS): return self.attempt_to_acquire_leader() def _write_leader_optime(self, last_lsn: str) -> bool: - return self._set_or_create(self.leader_optime_path, last_lsn) + return self._set_or_create(self.leader_optime_path, last_lsn) is not False def _write_status(self, value: str) -> bool: - return self._set_or_create(self.status_path, value) + return self._set_or_create(self.status_path, value) is not False def _write_failsafe(self, value: str) -> bool: - return self._set_or_create(self.failsafe_path, value) + return self._set_or_create(self.failsafe_path, value) is not False def _update_leader(self) -> bool: cluster = self.cluster @@ -491,13 +492,13 @@ class ZooKeeper(AbstractDCS): return True def set_history_value(self, value: str) -> bool: - return self._set_or_create(self.history_path, value) + return self._set_or_create(self.history_path, value) is not False - def set_sync_state_value(self, value: str, index: Optional[int] = None) -> bool: + def set_sync_state_value(self, value: str, index: Optional[int] = None) -> Union[int, bool]: return self._set_or_create(self.sync_path, value, index, retry=True, do_not_create_empty=True) def delete_sync_state(self, index: Optional[int] = None) -> bool: - return self.set_sync_state_value("{}", index) + return self.set_sync_state_value("{}", index) is not False def watch(self, leader_index: Optional[int], timeout: float) -> bool: ret = super(ZooKeeper, self).watch(leader_index, timeout + 0.5) diff --git a/patroni/ha.py b/patroni/ha.py index b994ad9a..e1559a27 100644 --- a/patroni/ha.py +++ b/patroni/ha.py @@ -590,15 +590,15 @@ class Ha(object): picked, allow_promote = self.state_handler.sync_handler.current_state(self.cluster) if picked != current: + sync = self.cluster.sync # update synchronous standby list in dcs temporarily to point to common nodes in current and picked sync_common = current & allow_promote if sync_common != current: logger.info("Updating synchronous privilege temporarily from %s to %s", list(current), list(sync_common)) - if not self.dcs.write_sync_state(self.state_handler.name, sync_common, - index=self.cluster.sync.index): - logger.info('Synchronous replication key updated by someone else.') - return + sync = self.dcs.write_sync_state(self.state_handler.name, sync_common, index=sync.index) + if not sync: + return logger.info('Synchronous replication key updated by someone else.') # When strict mode and no suitable replication connections put "*" to synchronous_standby_names if self.global_config.is_synchronous_mode_strict and not picked: @@ -606,24 +606,16 @@ class Ha(object): logger.warning("No standbys available!") # Update postgresql.conf and wait 2 secs for changes to become active - logger.info("Assigning synchronous standby status to %s", picked) + logger.info("Assigning synchronous standby status to %s", list(picked)) self.state_handler.sync_handler.set_synchronous_standby_names(picked) - if picked and picked != CaseInsensitiveSet('*') and allow_promote != picked and not allow_promote: + if picked and picked != CaseInsensitiveSet('*') and allow_promote != picked: # Wait for PostgreSQL to enable synchronous mode and see if we can immediately set sync_standby time.sleep(2) _, allow_promote = self.state_handler.sync_handler.current_state(self.cluster) if allow_promote and allow_promote != sync_common: - try: - cluster = self.dcs.get_cluster() - except DCSError: - return logger.warning("Could not get cluster state from DCS during process_sync_replication()") - if not cluster.sync.is_empty and not cluster.sync.leader_matches(self.state_handler.name): - logger.info("Synchronous replication key updated by someone else") - return - if not self.dcs.write_sync_state(self.state_handler.name, allow_promote, index=cluster.sync.index): - logger.info("Synchronous replication key updated by someone else") - return + if not self.dcs.write_sync_state(self.state_handler.name, allow_promote, index=sync.index): + return logger.info("Synchronous replication key updated by someone else") logger.info("Synchronous standby status assigned to %s", list(allow_promote)) else: if not self.cluster.sync.is_empty and self.dcs.delete_sync_state(index=self.cluster.sync.index): diff --git a/tests/test_consul.py b/tests/test_consul.py index f2bb6bc5..165d2382 100644 --- a/tests/test_consul.py +++ b/tests/test_consul.py @@ -15,6 +15,8 @@ def kv_get(self, key, **kwargs): return None, None if key == 'service/good/leader': return '1', None + if key == 'service/good/sync': + return '1', {'ModifyIndex': 1, 'Value': b'{}'} good_cls = ('6429', [{'CreateIndex': 1334, 'Flags': 0, 'Key': key + 'failover', 'LockIndex': 0, 'ModifyIndex': 1334, 'Value': b''}, @@ -224,7 +226,11 @@ class TestConsul(unittest.TestCase): @patch.object(consul.Consul.KV, 'delete', Mock(return_value=True)) @patch.object(consul.Consul.KV, 'put', Mock(return_value=True)) def test_sync_state(self): - self.assertTrue(self.c.set_sync_state_value('{}')) + self.assertEqual(self.c.set_sync_state_value('{}'), 1) + with patch('time.time', Mock(side_effect=[1, 100, 1000])): + self.assertFalse(self.c.set_sync_state_value('{}')) + with patch.object(consul.Consul.KV, 'put', Mock(return_value=False)): + self.assertFalse(self.c.set_sync_state_value('{}')) self.assertTrue(self.c.delete_sync_state()) @patch.object(consul.Consul.KV, 'put', Mock(return_value=True)) diff --git a/tests/test_etcd.py b/tests/test_etcd.py index 914ddeb8..6699127a 100644 --- a/tests/test_etcd.py +++ b/tests/test_etcd.py @@ -338,7 +338,7 @@ class TestEtcd(unittest.TestCase): self.assertTrue(self.etcd.watch(None, 1)) def test_sync_state(self): - self.assertFalse(self.etcd.write_sync_state('leader', None)) + self.assertIsNone(self.etcd.write_sync_state('leader', None)) self.assertFalse(self.etcd.delete_sync_state()) def test_set_history_value(self): diff --git a/tests/test_etcd3.py b/tests/test_etcd3.py index c0acb67f..85c2173a 100644 --- a/tests/test_etcd3.py +++ b/tests/test_etcd3.py @@ -54,8 +54,11 @@ def mock_urlopen(self, method, url, **kwargs): ]} })[:-1].encode('utf-8'), b'}{"error":{"grpc_code":14,"message":"","http_code":503}}']) elif url.endswith('/kv/put') or url.endswith('/kv/txn'): - ret.status_code = 400 - ret.content = '{"code":5,"error":"etcdserver: requested lease not found"}' + if base64_encode('/patroni/test/sync') in kwargs['body']: + ret.content = '{"header":{"revision":"1"},"succeeded":true}' + else: + ret.status_code = 400 + ret.content = '{"code":5,"error":"etcdserver: requested lease not found"}' elif not url.endswith('/kv/deleterange'): raise Exception('Unexpected url: {0} {1} {2}'.format(method, url, kwargs)) return ret @@ -295,7 +298,7 @@ class TestEtcd3(BaseTestEtcd3): self.etcd3.set_history_value('') def test_set_sync_state_value(self): - self.etcd3.set_sync_state_value('') + self.etcd3.set_sync_state_value('', 1) def test_delete_sync_state(self): self.etcd3.delete_sync_state() diff --git a/tests/test_ha.py b/tests/test_ha.py index 98e9966e..c9734517 100644 --- a/tests/test_ha.py +++ b/tests/test_ha.py @@ -793,7 +793,7 @@ class TestHa(PostgresInit): self.ha.cluster = get_cluster_initialized_without_leader(failover=Failover(0, '', 'other', None), sync=('leader1', 'postgresql0')) self.p.sync_handler.current_state = Mock(return_value=(CaseInsensitiveSet(), CaseInsensitiveSet())) - self.ha.dcs.write_sync_state = true + self.ha.dcs.write_sync_state = Mock(return_value=SyncState.empty()) self.assertEqual(self.ha.run_cycle(), 'promoted self to leader by acquiring session lock') # manual failover to our node (postgresql0), @@ -1136,46 +1136,46 @@ class TestHa(PostgresInit): # Test sync standby is replaced when switching standbys self.p.sync_handler.current_state = Mock(return_value=(CaseInsensitiveSet(['other2']), CaseInsensitiveSet())) - self.ha.dcs.write_sync_state = Mock(return_value=True) + self.ha.dcs.write_sync_state = Mock(return_value=SyncState.empty()) self.ha.run_cycle() mock_set_sync.assert_called_once_with(CaseInsensitiveSet(['other2'])) # Test sync standby is replaced when new standby is joined self.p.sync_handler.current_state = Mock(return_value=(CaseInsensitiveSet(['other2', 'other3']), CaseInsensitiveSet(['other2']))) - self.ha.dcs.write_sync_state = Mock(return_value=True) + self.ha.dcs.write_sync_state = Mock(return_value=SyncState.empty()) self.ha.run_cycle() self.assertEqual(mock_set_sync.call_args_list[0][0], (CaseInsensitiveSet(['other2']),)) self.assertEqual(mock_set_sync.call_args_list[1][0], (CaseInsensitiveSet(['other2', 'other3']),)) mock_set_sync.reset_mock() # Test sync standby is not disabled when updating dcs fails - self.ha.dcs.write_sync_state = Mock(return_value=False) + self.ha.dcs.write_sync_state = Mock(return_value=None) self.ha.run_cycle() mock_set_sync.assert_not_called() mock_set_sync.reset_mock() # Test changing sync standby - self.ha.dcs.write_sync_state = Mock(return_value=True) + self.ha.dcs.write_sync_state = Mock(return_value=SyncState.empty()) self.ha.dcs.get_cluster = Mock(return_value=get_cluster_initialized_with_leader(sync=('leader', 'other'))) # self.ha.cluster = get_cluster_initialized_with_leader(sync=('leader', 'other')) self.p.sync_handler.current_state = Mock(return_value=(CaseInsensitiveSet(['other2']), CaseInsensitiveSet(['other2']))) self.ha.run_cycle() - self.ha.dcs.get_cluster.assert_called_once() self.assertEqual(self.ha.dcs.write_sync_state.call_count, 2) # Test updating sync standby key failed due to race - self.ha.dcs.write_sync_state = Mock(side_effect=[True, False]) + self.ha.dcs.write_sync_state = Mock(side_effect=[SyncState.empty(), None]) self.ha.run_cycle() self.assertEqual(self.ha.dcs.write_sync_state.call_count, 2) # Test updating sync standby key failed due to DCS being not accessible - self.ha.dcs.write_sync_state = Mock(return_value=True) + self.ha.dcs.write_sync_state = Mock(return_value=SyncState.empty()) self.ha.dcs.get_cluster = Mock(side_effect=DCSError('foo')) self.ha.run_cycle() # Test changing sync standby failed due to race + self.ha.dcs.write_sync_state = Mock(return_value=SyncState.empty()) self.ha.dcs.get_cluster = Mock(return_value=get_cluster_initialized_with_leader(sync=('somebodyelse', None))) self.ha.run_cycle() self.assertEqual(self.ha.dcs.write_sync_state.call_count, 2) @@ -1194,7 +1194,7 @@ class TestHa(PostgresInit): self.p.is_leader = false self.p.set_role('replica') self.ha.has_lock = true - mock_write_sync = self.ha.dcs.write_sync_state = Mock(return_value=True) + mock_write_sync = self.ha.dcs.write_sync_state = Mock(return_value=SyncState.empty()) self.p.name = 'leader' self.ha.cluster = get_cluster_initialized_with_leader(sync=('other', None)) @@ -1218,7 +1218,7 @@ class TestHa(PostgresInit): self.p.set_role('replica') self.p.name = 'other' self.ha.cluster = get_cluster_initialized_without_leader(sync=('leader', 'other2')) - mock_write_sync = self.ha.dcs.write_sync_state = Mock(return_value=True) + mock_write_sync = self.ha.dcs.write_sync_state = Mock(return_value=SyncState.empty()) mock_acquire = self.ha.acquire_lock = Mock(return_value=True) mock_follow = self.p.follow = Mock() mock_promote = self.p.promote = Mock() diff --git a/tests/test_kubernetes.py b/tests/test_kubernetes.py index fd2185e0..bbe55157 100644 --- a/tests/test_kubernetes.py +++ b/tests/test_kubernetes.py @@ -370,11 +370,15 @@ class TestKubernetesEndpoints(BaseTestKubernetes): mock_read.side_effect = Exception self.assertFalse(self.k.update_leader('123')) - @patch.object(k8s_client.CoreV1Api, 'create_namespaced_endpoints', + @patch.object(k8s_client.CoreV1Api, 'patch_namespaced_endpoints', Mock(side_effect=[k8s_client.rest.ApiException(500, ''), k8s_client.rest.ApiException(502, '')]), create=True) def test_delete_sync_state(self): - self.assertFalse(self.k.delete_sync_state()) + self.assertFalse(self.k.delete_sync_state(1)) + + @patch.object(k8s_client.CoreV1Api, 'patch_namespaced_endpoints', mock_namespaced_kind, create=True) + def test_write_sync_state(self): + self.assertIsNotNone(self.k.write_sync_state('a', ['b'], 1)) @patch.object(k8s_client.CoreV1Api, 'patch_namespaced_pod', mock_namespaced_kind, create=True) @patch.object(k8s_client.CoreV1Api, 'create_namespaced_endpoints', mock_namespaced_kind, create=True) diff --git a/tests/test_raft.py b/tests/test_raft.py index 4f8ddb2c..2029a61e 100644 --- a/tests/test_raft.py +++ b/tests/test_raft.py @@ -138,6 +138,7 @@ class TestRaft(unittest.TestCase): self.assertTrue(raft.cancel_initialization()) self.assertTrue(raft.set_config_value('{}')) self.assertTrue(raft.write_sync_state('foo', 'bar')) + self.assertFalse(raft.write_sync_state('foo', 'bar', 1)) raft._citus_group = '1' self.assertTrue(raft.manual_failover('foo', 'bar')) raft._citus_group = '0'