From a92a76461e2b0bef6bf0830fd3bbecb8854a75d0 Mon Sep 17 00:00:00 2001 From: Dmitry Kropachev Date: Fri, 11 Sep 2026 16:02:57 -0400 Subject: [PATCH 1/5] host: define identity by immutable host ID --- CHANGELOG.rst | 5 ++ cassandra/metrics.py | 4 +- cassandra/pool.py | 70 +++++++++++++------ tests/integration/simulacron/test_cluster.py | 21 ++++-- .../integration/simulacron/test_connection.py | 2 +- .../integration/standard/test_shard_aware.py | 3 +- .../standard/test_tablets_routing_v2.py | 12 ++-- tests/unit/advanced/test_policies.py | 4 +- tests/unit/test_control_connection.py | 4 +- tests/unit/test_host_connection_pool.py | 66 +++++++++++++++-- tests/unit/test_policies.py | 10 +-- tests/unit/test_types.py | 14 ++-- 12 files changed, 160 insertions(+), 55 deletions(-) diff --git a/CHANGELOG.rst b/CHANGELOG.rst index bd2f2b28be..91560dc942 100644 --- a/CHANGELOG.rst +++ b/CHANGELOG.rst @@ -31,6 +31,11 @@ Features Others ------ +* ``Host.host_id`` is now the immutable identity of a node. Constructing a + ``Host`` requires a non-nil ``uuid.UUID``; equality, hashing, and ordering use + only that ID, and comparing a ``Host`` with an address no longer reports them + as equal. Because hashing is stable, existing Host-keyed session pool + lookups and removals remain valid across endpoint changes (issue #867). * ``DCAwareRoundRobinPolicy.local_dc`` is now read-only. It is set by the constructor, and filled in by the policy itself when the constructor was given none, from the first host to come up. Assigning it afterwards was indistinguishable from that inference, diff --git a/cassandra/metrics.py b/cassandra/metrics.py index 7ff44107af..9a4e32f417 100644 --- a/cassandra/metrics.py +++ b/cassandra/metrics.py @@ -441,7 +441,9 @@ def __init__(self, cluster_proxy): Stat('known_hosts', lambda: len(cluster_proxy.metadata.all_hosts())), Stat('connected_to', - lambda: len(set(chain.from_iterable(list(s._pools.keys()) for s in cluster_proxy.sessions)))), + lambda: len(set(chain.from_iterable( + (pool.host for pool in list(s._pools.values())) + for s in cluster_proxy.sessions)))), Stat('open_connections', lambda: sum(sum(p.open_count for p in list(s._pools.values())) for s in cluster_proxy.sessions))) diff --git a/cassandra/pool.py b/cassandra/pool.py index 14829ffa26..3e8bfcbd66 100644 --- a/cassandra/pool.py +++ b/cassandra/pool.py @@ -21,6 +21,7 @@ import time import random import copy +import uuid from threading import Lock, RLock, Condition import weakref try: @@ -129,11 +130,6 @@ class Host(object): release_version as queried from the control connection system tables """ - host_id = None - """ - The unique identifier of the cassandra node - """ - dse_version = None """ dse_version as queried from the control connection system tables. Only populated when connecting to @@ -155,6 +151,7 @@ class Host(object): Not queried if :attr:`~.Cluster.token_metadata_enabled` is ``False``. """ + _host_id = None _datacenter = None _rack = None _reconnection_handler = None @@ -170,11 +167,15 @@ def __init__(self, endpoint, conviction_policy_factory, datacenter=None, rack=No if conviction_policy_factory is None: raise ValueError("conviction_policy_factory may not be None") + if not isinstance(host_id, uuid.UUID): + raise TypeError("host_id must be a uuid.UUID") + if host_id.int == 0: + raise ValueError("host_id may not be the nil UUID") + self.endpoint = endpoint if isinstance(endpoint, EndPoint) else DefaultEndPoint(endpoint) + self._host_id = host_id + self._is_removed = False self.conviction_policy = conviction_policy_factory(self) - if not host_id: - raise ValueError("host_id may not be None") - self.host_id = host_id self.set_location_info(datacenter, rack) self.lock = RLock() @@ -186,6 +187,13 @@ def address(self): # backward compatibility return self.endpoint.address + @property + def host_id(self): + """ + The immutable unique identifier of the Cassandra node. + """ + return self._host_id + @property def datacenter(self): """ The datacenter the node is in. """ @@ -230,23 +238,25 @@ def get_and_set_reconnection_handler(self, new_handler): self._reconnection_handler = new_handler return old + def _clear_reconnection_handler(self, handler): + with self.lock: + if self._reconnection_handler is not handler: + return False + self._reconnection_handler = None + return True + def __eq__(self, other): - if isinstance(other, Host): - return self.endpoint == other.endpoint - else: # TODO Backward compatibility, remove next major - return self.endpoint.address == other + if not isinstance(other, Host): + return NotImplemented + return self.host_id == other.host_id def __hash__(self): - return hash(self.endpoint) + return hash(self.host_id) def __lt__(self, other): - self_is_unix = isinstance(self.endpoint, UnixSocketEndPoint) - other_is_unix = isinstance(other.endpoint, UnixSocketEndPoint) - if self_is_unix != other_is_unix: - # Endpoint comparators assume same-kind operands, so partition - # Unix and network Hosts before delegating their ordering. - return self_is_unix - return self.endpoint < other.endpoint + if not isinstance(other, Host): + return NotImplemented + return self.host_id < other.host_id def __str__(self): return str(self.endpoint) @@ -263,6 +273,7 @@ class _ReconnectionHandler(object): """ _cancelled = False + _clear_handler_before_reconnection = False def __init__(self, scheduler, schedule, callback, *callback_args, **callback_kwargs): self.scheduler = scheduler @@ -303,8 +314,12 @@ def run(self): self.scheduler.schedule(next_delay, self.run) else: if not self._cancelled: + if (self._clear_handler_before_reconnection and + not self._release_reconnection_handler()): + return self.on_reconnection(conn) - self.callback(*(self.callback_args), **(self.callback_kwargs)) + if not self._clear_handler_before_reconnection: + self._release_reconnection_handler() finally: if conn: conn.close() @@ -312,6 +327,10 @@ def run(self): def cancel(self): self._cancelled = True + def _release_reconnection_handler(self): + self.callback(*(self.callback_args), **(self.callback_kwargs)) + return True + def try_reconnect(self): """ Subclasses must implement this method. It should attempt to @@ -347,6 +366,12 @@ def on_exception(self, exc, next_delay): class _HostReconnectionHandler(_ReconnectionHandler): + # Host reconnection callbacks can synchronously start another reconnector + # when rebuilding pools fails. Clear this handler first so that failure is + # not suppressed as already reconnecting and post-callback cleanup cannot + # clear the successor. + _clear_handler_before_reconnection = True + def __init__(self, host, connection_factory, is_host_addition, on_add, on_up, *args, **kwargs): _ReconnectionHandler.__init__(self, *args, **kwargs) self.is_host_addition = is_host_addition @@ -358,6 +383,9 @@ def __init__(self, host, connection_factory, is_host_addition, on_add, on_up, *a def try_reconnect(self): return self.connection_factory() + def _release_reconnection_handler(self): + return self.host._clear_reconnection_handler(self) + def on_reconnection(self, connection): log.info("Successful reconnection to %s, marking node up if it isn't already", self.host) if self.is_host_addition: diff --git a/tests/integration/simulacron/test_cluster.py b/tests/integration/simulacron/test_cluster.py index b8b908e3bb..20378115b1 100644 --- a/tests/integration/simulacron/test_cluster.py +++ b/tests/integration/simulacron/test_cluster.py @@ -89,12 +89,23 @@ def test_duplicate(self): with MockLoggingHandler().set_module_name(cassandra.cluster.__name__) as mock_handler: address_column = "rpc_address" rows = [ - {"peer": "127.0.0.1", "data_center": "dc", "host_id": "dontcare1", "rack": "rack1", - "release_version": "3.11.4", address_column: "127.0.0.1", "schema_version": "dontcare", "tokens": "1"}, - {"peer": "127.0.0.2", "data_center": "dc", "host_id": "dontcare2", "rack": "rack1", - "release_version": "3.11.4", address_column: "127.0.0.2", "schema_version": "dontcare", "tokens": "2"}, + {"peer": "127.0.0.1", "data_center": "dc", "host_id": "00000000-0000-0000-0000-000000000001", "rack": "rack1", + "release_version": "3.11.4", address_column: "127.0.0.1", "schema_version": "00000000-0000-0000-0000-000000000011", "tokens": ["1"]}, + {"peer": "127.0.0.2", "data_center": "dc", "host_id": "00000000-0000-0000-0000-000000000002", "rack": "rack1", + "release_version": "3.11.4", address_column: "127.0.0.2", "schema_version": "00000000-0000-0000-0000-000000000012", "tokens": ["2"]}, ] - prime_query(ControlConnection._SELECT_PEERS, rows=rows) + prime_query( + ControlConnection._SELECT_PEERS, rows=rows, + column_types={ + "peer": "inet", + "data_center": "varchar", + "host_id": "uuid", + "rack": "varchar", + "release_version": "varchar", + address_column: "inet", + "schema_version": "uuid", + "tokens": "set", + }) cluster = Cluster(protocol_version=PROTOCOL_VERSION, compression=False) session = cluster.connect(wait_for_all_pools=True) diff --git a/tests/integration/simulacron/test_connection.py b/tests/integration/simulacron/test_connection.py index 574f153edf..d51c6955d2 100644 --- a/tests/integration/simulacron/test_connection.py +++ b/tests/integration/simulacron/test_connection.py @@ -224,7 +224,7 @@ class PatchedRoundRobinPolicy(RoundRobinPolicy): # Send always to same host def make_query_plan(self, working_keyspace=None, query=None): if query and query.query_string == query_to_prime: - return filter(lambda h: h == query_host, self._live_hosts) + return filter(lambda h: h.address == query_host, self._live_hosts) else: return super(PatchedRoundRobinPolicy, self).make_query_plan() diff --git a/tests/integration/standard/test_shard_aware.py b/tests/integration/standard/test_shard_aware.py index 6daba6e26f..2fc36af577 100644 --- a/tests/integration/standard/test_shard_aware.py +++ b/tests/integration/standard/test_shard_aware.py @@ -152,7 +152,8 @@ def _assert_blocked_node_disconnected(self, node_ip_address, node_port): assert active_control_connection.is_closed or active_control_connection.is_defunct pools = getattr(self.session, '_pools', None) or {} - for host, pool in pools.items(): + for pool in pools.values(): + host = pool.host if host.endpoint.address != node_ip_address or host.endpoint.port != node_port: continue diff --git a/tests/integration/standard/test_tablets_routing_v2.py b/tests/integration/standard/test_tablets_routing_v2.py index 9edcdcd64e..d56a748503 100644 --- a/tests/integration/standard/test_tablets_routing_v2.py +++ b/tests/integration/standard/test_tablets_routing_v2.py @@ -276,7 +276,8 @@ def _any_connection(self): @staticmethod def _all_shard_connections(session): - for host, pool in session._pools.items(): + for pool in session._pools.values(): + host = pool.host for shard, conn in pool._connections.items(): yield host, shard, conn @@ -287,9 +288,9 @@ def _wait_for_shard_connections(session, timeout=15): deadline = time.time() + timeout while time.time() < deadline: if all( - len(pool._connections) >= (min(host.sharding_info.shards_count, 2) - if host.sharding_info else 1) - for host, pool in session._pools.items() + len(pool._connections) >= (min(pool.host.sharding_info.shards_count, 2) + if pool.host.sharding_info else 1) + for pool in session._pools.values() ): return time.sleep(0.05) @@ -329,7 +330,8 @@ def _find_replica_wrong_shard(session, tablet): (host, owner_shard, wrong_shard, conn) or None if no host has >=2 shards. """ replica_shard = {host_id: shard for host_id, shard in tablet.replicas} - for host, pool in session._pools.items(): + for pool in session._pools.values(): + host = pool.host owner = replica_shard.get(host.host_id) if owner is None: continue diff --git a/tests/unit/advanced/test_policies.py b/tests/unit/advanced/test_policies.py index 75cfd3fbf9..a776aa8b9f 100644 --- a/tests/unit/advanced/test_policies.py +++ b/tests/unit/advanced/test_policies.py @@ -79,7 +79,7 @@ def test_target_host_down(self): policy = DSELoadBalancingPolicy(RoundRobinPolicy()) policy.populate(Mock(metadata=ClusterMetaMock({'127.0.0.1': target_host})), hosts) query_plan = list(policy.make_query_plan(None, Mock(target_host='127.0.0.1'))) - assert sorted(query_plan) == hosts + assert sorted(query_plan) == sorted(hosts) target_host.is_up = False policy.on_down(target_host) @@ -96,5 +96,5 @@ def test_target_host_nominal(self): policy.populate(Mock(metadata=ClusterMetaMock({'127.0.0.1': target_host})), hosts) for _ in range(10): query_plan = list(policy.make_query_plan(None, Mock(target_host='127.0.0.1'))) - assert sorted(query_plan) == hosts + assert sorted(query_plan) == sorted(hosts) assert query_plan[0] == target_host diff --git a/tests/unit/test_control_connection.py b/tests/unit/test_control_connection.py index 13a68c7e55..fb5abdb3f6 100644 --- a/tests/unit/test_control_connection.py +++ b/tests/unit/test_control_connection.py @@ -359,7 +359,7 @@ def test_refresh_sets_local_listen_address_when_rpc_address_changes(self): self.control_connection.refresh_node_list_and_token_map() - local_host = self.cluster.metadata.get_host_by_host_id('uuid1') + local_host = self.cluster.metadata.get_host_by_host_id(HOST_ID_1) assert local_host.endpoint == DefaultEndPoint('192.168.1.4') assert local_host.listen_address == '192.168.1.0' @@ -380,7 +380,7 @@ def test_refresh_sets_local_addresses_without_token_metadata(self): assert 'listen_address' in local_projection assert 'broadcast_address' in local_projection assert 'tokens' not in local_projection - local_host = self.cluster.metadata.get_host_by_host_id('uuid1') + local_host = self.cluster.metadata.get_host_by_host_id(HOST_ID_1) assert local_host.listen_address == '192.168.1.0' assert local_host.broadcast_address == '10.0.0.1' diff --git a/tests/unit/test_host_connection_pool.py b/tests/unit/test_host_connection_pool.py index 8bb57d0dc0..ca1cb51573 100644 --- a/tests/unit/test_host_connection_pool.py +++ b/tests/unit/test_host_connection_pool.py @@ -23,7 +23,7 @@ from unittest.mock import Mock, NonCallableMagicMock, MagicMock from cassandra.cluster import Session, ShardAwareOptions -from cassandra.connection import Connection +from cassandra.connection import Connection, DefaultEndPoint from cassandra.pool import HostConnection from cassandra.pool import Host, NoConnectionsAvailable from cassandra.policies import HostDistance, SimpleConvictionPolicy @@ -211,18 +211,70 @@ def test_host_instantiations(self): with pytest.raises(ValueError): Host(None, SimpleConvictionPolicy, host_id=uuid.uuid4()) + for host_id in (None, "not-a-uuid", object()): + with pytest.raises( + TypeError, match=r"^host_id must be a uuid\.UUID$"): + Host('127.0.0.1', SimpleConvictionPolicy, host_id=host_id) + + with pytest.raises(ValueError): + Host('127.0.0.1', SimpleConvictionPolicy, host_id=uuid.UUID(int=0)) + + with pytest.raises( + TypeError, match=r"^host_id must be a uuid\.UUID$"): + Host('127.0.0.1', SimpleConvictionPolicy) + def test_host_equality(self): """ Test host equality has correct logic """ - a = Host('127.0.0.1', SimpleConvictionPolicy, host_id=uuid.uuid4()) - b = Host('127.0.0.1', SimpleConvictionPolicy, host_id=uuid.uuid4()) - c = Host('127.0.0.2', SimpleConvictionPolicy, host_id=uuid.uuid4()) + shared_id = uuid.uuid4() + a = Host('127.0.0.1', SimpleConvictionPolicy, host_id=shared_id) + b = Host('127.0.0.2', SimpleConvictionPolicy, host_id=shared_id) + c = Host('127.0.0.1', SimpleConvictionPolicy, host_id=uuid.uuid4()) + + assert a == b, 'Two Host instances with the same host ID should be equal.' + assert a != c, 'Two Host instances with different host IDs should not be equal.' + assert a != a.address, 'A Host should not compare equal to its address.' + + def test_host_id_is_read_only(self): + host_id = uuid.uuid4() + host = Host('127.0.0.1', SimpleConvictionPolicy, host_id=host_id) + + with pytest.raises(AttributeError): + host.host_id = uuid.uuid4() + + assert host.host_id == host_id + + def test_host_hash_is_stable_when_endpoint_changes(self): + host = Host('127.0.0.1', SimpleConvictionPolicy, host_id=uuid.uuid4()) + hosts_by_id = {host: "pool"} + hosts = {host} + + host.endpoint = DefaultEndPoint('127.0.0.2') + + assert hosts_by_id[host] == "pool" + assert host in hosts + + def test_host_ordering_uses_host_id(self): + first = Host('127.0.0.2', SimpleConvictionPolicy, host_id=uuid.UUID(int=1)) + second = Host('127.0.0.1', SimpleConvictionPolicy, host_id=uuid.UUID(int=2)) + + assert first < second + + def test_host_id_is_set_before_conviction_policy_is_created(self): + host_id = uuid.UUID(int=1) + hosts = set() + + def conviction_policy_factory(host): + assert host.host_id == host_id + hosts.add(host) + return SimpleConvictionPolicy(host) + + host = Host( + '127.0.0.1', conviction_policy_factory, host_id=host_id) - assert a == b, 'Two Host instances should be equal when sharing.' - assert a != c, 'Two Host instances should NOT be equal when using two different addresses.' - assert b != c, 'Two Host instances should NOT be equal when using two different addresses.' + assert host in hosts class HostConnectionTests(_PoolTests): diff --git a/tests/unit/test_policies.py b/tests/unit/test_policies.py index 2fc31a31df..31ded0f45d 100644 --- a/tests/unit/test_policies.py +++ b/tests/unit/test_policies.py @@ -2174,9 +2174,9 @@ def get_replicas(keyspace, packed_key): query_plan = hfp.make_query_plan("keyspace", mocked_query) # First the not filtered replica, and then the rest of the allowed hosts ordered query_plan = list(query_plan) - assert query_plan[0] == Host(DefaultEndPoint("127.0.0.2"), SimpleConvictionPolicy, host_id=uuid.uuid4()) - assert set(query_plan[1:]) == {Host(DefaultEndPoint("127.0.0.3"), SimpleConvictionPolicy, host_id=uuid.uuid4()), - Host(DefaultEndPoint("127.0.0.5"), SimpleConvictionPolicy, host_id=uuid.uuid4())} + assert query_plan[0].endpoint == DefaultEndPoint("127.0.0.2") + assert {host.endpoint for host in query_plan[1:]} == { + DefaultEndPoint("127.0.0.3"), DefaultEndPoint("127.0.0.5")} def test_create_whitelist(self): cluster = Mock(spec=Cluster) @@ -2198,5 +2198,5 @@ def test_create_whitelist(self): mocked_query = Mock() query_plan = hfp.make_query_plan("keyspace", mocked_query) # Only the filtered replicas should be allowed - assert set(query_plan) == {Host(DefaultEndPoint("127.0.0.1"), SimpleConvictionPolicy, host_id=uuid.uuid4()), - Host(DefaultEndPoint("127.0.0.4"), SimpleConvictionPolicy, host_id=uuid.uuid4())} + assert {host.endpoint for host in query_plan} == { + DefaultEndPoint("127.0.0.1"), DefaultEndPoint("127.0.0.4")} diff --git a/tests/unit/test_types.py b/tests/unit/test_types.py index 11aab2748d..654d16fe56 100644 --- a/tests/unit/test_types.py +++ b/tests/unit/test_types.py @@ -1018,11 +1018,15 @@ def test_host_order(self): @test_category data_types """ - hosts = [Host(addr, SimpleConvictionPolicy, host_id=uuid.uuid4()) for addr in - ("127.0.0.1", "127.0.0.2", "127.0.0.3", "127.0.0.4")] - hosts_equal = [Host(addr, SimpleConvictionPolicy, host_id=uuid.uuid4()) for addr in - ("127.0.0.1", "127.0.0.1")] - hosts_equal_conviction = [Host("127.0.0.1", SimpleConvictionPolicy, host_id=uuid.uuid4()), Host("127.0.0.1", ConvictionPolicy, host_id=uuid.uuid4())] + hosts = [Host(addr, SimpleConvictionPolicy, host_id=uuid.UUID(int=index)) + for index, addr in enumerate( + ("127.0.0.4", "127.0.0.3", "127.0.0.2", "127.0.0.1"), + start=1)] + shared_id = uuid.uuid4() + hosts_equal = [Host(addr, SimpleConvictionPolicy, host_id=shared_id) for addr in + ("127.0.0.1", "127.0.0.2")] + hosts_equal_conviction = [Host("127.0.0.1", SimpleConvictionPolicy, host_id=shared_id), + Host("127.0.0.2", ConvictionPolicy, host_id=shared_id)] check_sequence_consistency(hosts) check_sequence_consistency(hosts_equal, equal=True) check_sequence_consistency(hosts_equal_conviction, equal=True) From 6fbc7672ba0dc060e8f9b89c1137fbf4513ba08a Mon Sep 17 00:00:00 2001 From: Dmitry Kropachev Date: Mon, 14 Sep 2026 16:05:19 -0400 Subject: [PATCH 2/5] host: preserve pool lifecycle across replacements Reconcile reused endpoints through the existing host lifecycle without overlapping control reconnects. Fence pool publication after lifecycle removal, preserve replacement recovery across partial failures and reconnector handoffs, and keep control-timeout cleanup separate from data pools. --- cassandra/cluster.py | 356 ++++++++++++--- cassandra/pool.py | 2 +- tests/unit/test_cluster.py | 605 +++++++++++++++++++++++++- tests/unit/test_control_connection.py | 371 +++++++++++++--- tests/unit/test_response_future.py | 17 +- 5 files changed, 1211 insertions(+), 140 deletions(-) diff --git a/cassandra/cluster.py b/cassandra/cluster.py index a87b2d2082..900c4c2ee3 100644 --- a/cassandra/cluster.py +++ b/cassandra/cluster.py @@ -1986,8 +1986,10 @@ def on_up(self, host): # for testing purposes return futures - def _start_reconnector(self, host, is_host_addition): - if self.profile_manager.distance(host) == HostDistance.IGNORED: + def _start_reconnector(self, host, is_host_addition, + on_add_reconnection=None, start=True): + if (on_add_reconnection is None and + self.profile_manager.distance(host) == HostDistance.IGNORED): return schedule = self.reconnection_policy.new_schedule() @@ -1996,22 +1998,33 @@ def _start_reconnector(self, host, is_host_addition): # proper shutdown when the program ends, we'll just make a closure # of the current Cluster attributes to create new Connections with conn_factory = self._make_connection_factory(host) + on_add = (self.on_add if on_add_reconnection is None + else on_add_reconnection) - reconnector = _HostReconnectionHandler( - host, conn_factory, is_host_addition, self.on_add, self.on_up, - self.scheduler, schedule, host.get_and_set_reconnection_handler, - new_handler=None) - - old_reconnector = host.get_and_set_reconnection_handler(reconnector) + # Pair installation with lifecycle removal. If installation wins, + # on_remove() will clear and cancel this handler; if removal wins, do + # not leave retry work attached to a terminal Host object. + with host.lock: + if host._is_removed: + return + reconnector = _HostReconnectionHandler( + host, conn_factory, is_host_addition, on_add, self.on_up, + self.scheduler, schedule, + host.get_and_set_reconnection_handler, new_handler=None) + old_reconnector = host._reconnection_handler + host._reconnection_handler = reconnector if old_reconnector: log.debug("Old host reconnector found for %s, cancelling", host) old_reconnector.cancel() - log.debug("Starting reconnector for host %s", host) - reconnector.start() + if start: + log.debug("Starting reconnector for host %s", host) + reconnector.start() + return reconnector @run_in_executor - def on_down_potentially_blocking(self, host, is_host_addition): + def on_down_potentially_blocking(self, host, is_host_addition, + on_add_reconnection=None): self.profile_manager.on_down(host) self.control_connection.on_down(host) for session in tuple(self.sessions): @@ -2020,9 +2033,15 @@ def on_down_potentially_blocking(self, host, is_host_addition): for listener in self.listeners: listener.on_down(host) - self._start_reconnector(host, is_host_addition) + if on_add_reconnection is None: + self._start_reconnector(host, is_host_addition) + else: + self._start_reconnector( + host, is_host_addition, + on_add_reconnection=on_add_reconnection) - def on_down(self, host, is_host_addition, expect_host_to_be_down=False): + def on_down(self, host, is_host_addition, expect_host_to_be_down=False, + on_add_reconnection=None): """ Intended for internal use only. """ @@ -2049,9 +2068,15 @@ def on_down(self, host, is_host_addition, expect_host_to_be_down=False): return log.warning("Host %s has been marked down", host) - self.on_down_potentially_blocking(host, is_host_addition) + if on_add_reconnection is None: + self.on_down_potentially_blocking(host, is_host_addition) + else: + self.on_down_potentially_blocking( + host, is_host_addition, + on_add_reconnection=on_add_reconnection) - def on_add(self, host, refresh_nodes=True): + def on_add(self, host, refresh_nodes=True, + reconcile_pools_on_failure=False): if self.is_shutdown: return @@ -2074,6 +2099,11 @@ def on_add(self, host, refresh_nodes=True): futures_lock = Lock() futures_results = [] futures = set() + on_add_reconnection = None + if reconcile_pools_on_failure: + on_add_reconnection = partial( + self.on_add, refresh_nodes=refresh_nodes, + reconcile_pools_on_failure=True) def future_completed(future): with futures_lock: @@ -2089,26 +2119,104 @@ def future_completed(future): log.debug('All futures have completed for added host %s', host) - for exc in [f for f in futures_results if isinstance(f, Exception)]: - log.error("Unexpected failure while adding node %s, will not mark up:", host, exc_info=exc) - return - - if not all(futures_results): + failures = [f for f in futures_results if isinstance(f, Exception)] + if failures: + for exc in failures: + log.error("Unexpected failure while adding node %s, will not mark up:", host, exc_info=exc) + elif not all(futures_results): log.warning("Connection pool could not be created, not marking node %s up", host) + else: + self._finalize_add(host) return - self._finalize_add(host) + # A same-endpoint replacement removes its predecessor without + # reconciling pools because doing so can race this addition. If + # the replacement fails, _finalize_add() cannot provide that + # reconciliation either, so restore it without immediately + # recreating a pool for the host whose addition just failed. + if reconcile_pools_on_failure and not self.is_shutdown: + # A pool created by another session can cause the failure's + # DOWN signal to be discounted, and not every failure convicts + # the host. Aggregate replacement handling therefore owns + # cleanup and makes sure retry work exists. + with host.lock: + if host._is_removed: + return + host.set_down() + sessions = tuple(self.sessions) + self.profile_manager.on_down(host) + for session in sessions: + session.remove_pool(host) + + authentication_failed = any( + result is None for result in futures_results) + pending_reconnector = None + if authentication_failed: + old_reconnector = \ + host.get_and_set_reconnection_handler(None) + if old_reconnector: + old_reconnector.cancel() + else: + pending_reconnector = self._start_reconnector( + host, is_host_addition=True, + on_add_reconnection=on_add_reconnection, + start=False) + + # Reconcile only after every partial replacement pool is gone + # and retry ownership has been established. Freeze the host + # set while removal of this object is excluded, so a stale + # callback cannot discover and duplicate-create a newer + # same-endpoint replacement. + with host.lock: + if host._is_removed: + if pending_reconnector: + host._clear_reconnection_handler( + pending_reconnector) + pending_reconnector.cancel() + return + reconciliation_hosts = tuple( + self.metadata.all_hosts()) + + try: + for session in sessions: + session.update_created_pools( + excluded_host=host, + hosts=reconciliation_hosts) + finally: + if pending_reconnector: + with host.lock: + active = ( + not host._is_removed and + host._reconnection_handler is + pending_reconnector) + if active: + log.debug( + "Starting reconnector for host %s", host) + pending_reconnector.start() + else: + pending_reconnector.cancel() - have_future = False for session in tuple(self.sessions): - future = session.add_or_renew_pool(host, is_host_addition=True) + if on_add_reconnection is None: + future = session.add_or_renew_pool( + host, is_host_addition=True) + else: + future = session.add_or_renew_pool( + host, is_host_addition=True, + on_add_reconnection=on_add_reconnection) if future is not None: - have_future = True futures.add(future) - future.add_done_callback(future_completed) - if not have_future: + if not futures: self._finalize_add(host) + return + + # Register callbacks only after every session has had its pool + # scheduled. Future.add_done_callback() invokes the callback inline + # when the future is already complete, so registering inside the loop + # can finalize the host while later sessions are still unscheduled. + for future in tuple(futures): + future.add_done_callback(future_completed) def _finalize_add(self, host, set_up=True): if set_up: @@ -2121,35 +2229,57 @@ def _finalize_add(self, host, set_up=True): for session in tuple(self.sessions): session.update_created_pools() - def on_remove(self, host): + def on_remove(self, host, trigger_reconciliation=True): if self.is_shutdown: return log.debug("[cluster] Removing host %s", host) - host.set_down() + # A Host object is never re-added after lifecycle removal; a later + # discovery creates a new object, even when it carries the same ID. + # Mark this instance before removing its pools so pool creation + # already in flight cannot publish after removal has passed it. + with host.lock: + host._is_removed = True + host.set_down() self.profile_manager.on_remove(host) for session in tuple(self.sessions): - session.on_remove(host) + session.on_remove( + host, trigger_reconciliation=trigger_reconciliation) for listener in self.listeners: listener.on_remove(host) - self.control_connection.on_remove(host) + self.control_connection.on_remove( + host, trigger_reconciliation=trigger_reconciliation) reconnection_handler = host.get_and_set_reconnection_handler(None) if reconnection_handler: reconnection_handler.cancel() - def signal_connection_failure(self, host, connection_exc, is_host_addition, expect_host_to_be_down=False): + def signal_connection_failure(self, host, connection_exc, + is_host_addition, + expect_host_to_be_down=False, + on_add_reconnection=None): is_down = host.signal_connection_failure(connection_exc) if is_down: - self.on_down(host, is_host_addition, expect_host_to_be_down) + if on_add_reconnection is None: + self.on_down( + host, is_host_addition, expect_host_to_be_down) + else: + self.on_down( + host, is_host_addition, expect_host_to_be_down, + on_add_reconnection=on_add_reconnection) return is_down - def add_host(self, endpoint, datacenter=None, rack=None, signal=True, refresh_nodes=True, host_id=None): + def add_host(self, endpoint, datacenter=None, rack=None, signal=True, + refresh_nodes=True, host_id=None, + reconcile_pools_on_failure=False): """ Called when adding initial contact points and when the control connection subsequently discovers a new node. Returns a Host instance, and a flag indicating whether it was new in the metadata. + + ``reconcile_pools_on_failure`` restores reconciliation deferred by a + same-endpoint replacement if creating the replacement pool fails. Intended for internal use only. """ with self.metadata._hosts_lock: @@ -2158,18 +2288,24 @@ def add_host(self, endpoint, datacenter=None, rack=None, signal=True, refresh_no host, new = self.metadata.add_or_return_host(Host(endpoint, self.conviction_policy_factory, datacenter, rack, host_id=host_id)) if new and signal: log.info("New Cassandra host %r discovered", host) - self.on_add(host, refresh_nodes) + self.on_add( + host, refresh_nodes, + reconcile_pools_on_failure=reconcile_pools_on_failure) return host, new - def remove_host(self, host): + def remove_host(self, host, trigger_reconciliation=True): """ Called when the control connection observes that a node has left the ring. Intended for internal use only. + + ``trigger_reconciliation`` may be disabled when an immediately + following host addition owns pool and control-connection follow-up. """ if host and self.metadata.remove_host(host): log.info("Cassandra host %s removed", host) - self.on_remove(host) + self.on_remove( + host, trigger_reconciliation=trigger_reconciliation) def register_listener(self, listener): """ @@ -3263,7 +3399,8 @@ def prepare_on_all_hosts(self, query, excluded_host, keyspace=None): Intended for internal use only. """ futures = [] - for host in tuple(self._pools.keys()): + for pool in tuple(self._pools.values()): + host = pool.host if host != excluded_host and host.is_up: future = ResponseFuture(self, PrepareMessage(query=query, keyspace=keyspace), None, self.default_timeout) @@ -3327,7 +3464,8 @@ def __del__(self): # when cluster.shutdown() is called explicitly. pass - def add_or_renew_pool(self, host, is_host_addition): + def add_or_renew_pool(self, host, is_host_addition, + on_add_reconnection=None): """ For internal use only. """ @@ -3343,18 +3481,35 @@ def run_add_or_renew_pool(): new_pool = HostConnection(host, distance, self) except AuthenticationFailed as auth_exc: conn_exc = ConnectionException(str(auth_exc), endpoint=host) - self.cluster.signal_connection_failure(host, conn_exc, is_host_addition) - return False + if on_add_reconnection is None: + self.cluster.signal_connection_failure( + host, conn_exc, is_host_addition) + else: + self.cluster.signal_connection_failure( + host, conn_exc, is_host_addition, + on_add_reconnection=on_add_reconnection) + # Replacement aggregation uses None to distinguish a + # non-retryable authentication failure from other falsey + # pool-creation results. + return (None if on_add_reconnection is not None else False) except Exception as conn_exc: log.warning("Failed to create connection pool for new host %s:", host, exc_info=conn_exc) # the host itself will still be marked down, so we need to pass # a special flag to make sure the reconnector is created - self.cluster.signal_connection_failure( - host, conn_exc, is_host_addition, expect_host_to_be_down=True) + if on_add_reconnection is None: + self.cluster.signal_connection_failure( + host, conn_exc, is_host_addition, + expect_host_to_be_down=True) + else: + self.cluster.signal_connection_failure( + host, conn_exc, is_host_addition, + expect_host_to_be_down=True, + on_add_reconnection=on_add_reconnection) return False previous = self._pools.get(host) + publish_pool = True with self._lock: while new_pool._keyspace != self.keyspace: self._lock.release() @@ -3369,12 +3524,30 @@ def callback(pool, errors): set_keyspace_event.wait(self.cluster.connect_timeout) if not set_keyspace_event.is_set() or errors_returned: log.warning("Failed setting keyspace for pool after keyspace changed during connect: %s", errors_returned) - self.cluster.on_down(host, is_host_addition) + if on_add_reconnection is None: + self.cluster.on_down(host, is_host_addition) + else: + self.cluster.on_down( + host, is_host_addition, + expect_host_to_be_down=True, + on_add_reconnection=on_add_reconnection) new_pool.shutdown() self._lock.acquire() return False self._lock.acquire() - self._pools[host] = new_pool + # Pair this check with Cluster.on_remove() marking the Host + # under the same lock. If publication wins, removal will pop + # the pool; if removal wins, discard the obsolete result. + with host.lock: + if host._is_removed: + publish_pool = False + else: + self._pools[host] = new_pool + + if not publish_pool: + log.debug("Discarding pool created for removed host %s", host) + new_pool.shutdown() + return False log.debug("Added pool for host %s to session", host) if previous: @@ -3392,7 +3565,7 @@ def remove_pool(self, host): else: return None - def update_created_pools(self): + def update_created_pools(self, excluded_host=None, hosts=None): """ When the set of live nodes change, the loadbalancer will change its mind on host distances. It might change it on the node that came/left @@ -3402,13 +3575,22 @@ def update_created_pools(self): This method ensures that all hosts for which a pool should exist have one, and hosts that shouldn't don't. + ``excluded_host`` suppresses a new attempt for a host whose pool + creation has just failed while still reconciling all other hosts. + ``hosts`` may freeze the metadata snapshot used for reconciliation. + For internal use only. """ if self.cluster.allow_control_connection_query_fallback is ControlConnectionQueryFallback.SkipPoolCreation: return set() + if hosts is None: + hosts = self.cluster.metadata.all_hosts() + futures = set() - for host in self.cluster.metadata.all_hosts(): + for host in hosts: + if excluded_host is not None and host == excluded_host: + continue distance = self._profile_manager.distance(host) pool = self._pools.get(host) future = None @@ -3437,9 +3619,12 @@ def on_down(self, host): if future: future.add_done_callback(lambda f: self.update_created_pools()) - def on_remove(self, host): - """ Internal """ - self.on_down(host) + def on_remove(self, host, trigger_reconciliation=True): + """Remove this host's pool and optionally reconcile remaining pools.""" + if trigger_reconciliation: + self.on_down(host) + else: + self.remove_pool(host) def set_keyspace(self, keyspace): """ @@ -3604,11 +3789,14 @@ def _get_schema_agreement_hosts(self, scope: SchemaAgreementScope) -> Tuple[Host else: allowed_distances = (HostDistance.LOCAL_RACK, HostDistance.LOCAL, HostDistance.REMOTE) - return tuple( - host for host, pool in tuple(self._pools.items()) - if host.is_up - and not pool.is_shutdown - and self._profile_manager.distance(host) in allowed_distances) + hosts = [] + for pool in tuple(self._pools.values()): + host = pool.host + if (host.is_up + and not pool.is_shutdown + and self._profile_manager.distance(host) in allowed_distances): + hosts.append(host) + return tuple(hosts) def _query_local_schema_version(self, host: Host, query: str, deadline: float) -> Future: remaining = max(0.0, deadline - time.time()) @@ -3693,7 +3881,7 @@ def submit(self, fn, *args, **kwargs): return self.cluster.executor.submit(fn, *args, **kwargs) def get_pool_state(self): - return dict((host, pool.get_state()) for host, pool in tuple(self._pools.items())) + return dict((pool.host, pool.get_state()) for pool in tuple(self._pools.values())) def get_pools(self): return self._pools.values() @@ -4171,9 +4359,8 @@ def _refresh_node_list_and_token_map(self, connection, preloaded_results=None, found_endpoints.add(factory_endpoint) existing_host = self._cluster.metadata.get_host_by_host_id(host_id) - # Host hashes depend on their endpoint, so never replace the route - # of an existing Host with or from a Unix socket. A newly discovered - # local Host keeps the socket which actually reached the node. + # Preserve an existing Host's Unix route. A newly discovered local + # Host keeps the socket which actually reached the node. if (existing_host is not None and isinstance(existing_host.endpoint, UnixSocketEndPoint)): endpoint = existing_host.endpoint @@ -4187,6 +4374,22 @@ def _refresh_node_list_and_token_map(self, connection, preloaded_results=None, host = self._cluster.metadata.get_host(endpoint) datacenter = row.get("data_center") rack = row.get("rack") + reconcile_pools_on_failure = False + + # host_id is immutable. Preserve the existing replacement + # behavior without folding endpoint-transition orchestration into + # this identity change; coordinated moves belong to #923. + if host is not None and host.host_id != host_id: + log.debug( + "[control connection] Replacing host %s at %s with host %s", + host.host_id, endpoint, host_id) + # The addition below owns replacement pool creation. Avoid the + # removal callback racing it with a second creation attempt. + self._cluster.remove_host( + host, trigger_reconciliation=False) + should_rebuild_token_map = True + host = None + reconcile_pools_on_failure = True if host is None: host = existing_host @@ -4204,12 +4407,14 @@ def _refresh_node_list_and_token_map(self, connection, preloaded_results=None, if host is None: log.debug("[control connection] Found new host to connect to: %s", endpoint) - host, _ = self._cluster.add_host(endpoint, datacenter=datacenter, rack=rack, signal=True, refresh_nodes=False, host_id=host_id) + host, _ = self._cluster.add_host( + endpoint, datacenter=datacenter, rack=rack, signal=True, + refresh_nodes=False, host_id=host_id, + reconcile_pools_on_failure=reconcile_pools_on_failure) should_rebuild_token_map = True else: should_rebuild_token_map |= self._update_location_info(host, datacenter, rack) - host.host_id = host_id host.broadcast_address = _NodeInfo.get_broadcast_address(row) host.broadcast_port = _NodeInfo.get_broadcast_port(row) host.broadcast_rpc_address = _NodeInfo.get_broadcast_rpc_address(row) @@ -4256,6 +4461,12 @@ def _is_valid_peer(row): broadcast_rpc) return False + if not isinstance(host_id, uuid.UUID) or host_id.int == 0: + log.warning( + "Found an invalid row for peer - invalid host_id %r (broadcast_rpc: %s). Ignoring host." % + (host_id, broadcast_rpc)) + return False + if not row.get("data_center"): log.warning( "Found an invalid row for peer - missing data_center (broadcast_rpc: %s, host_id: %s). Ignoring host." % @@ -4649,7 +4860,14 @@ def on_add(self, host, refresh_nodes=True): if refresh_nodes: self.refresh_node_list_and_token_map(force_token_rebuild=True) - def on_remove(self, host): + def on_remove(self, host, trigger_reconciliation=True): + if not trigger_reconciliation: + # Same-endpoint replacement reconciliation is already running on + # either the current connection or a candidate that its caller + # owns. Starting another reconnect here can race that candidate's + # publication and close the newly published connection. + return + c = self._connection if self._connection_matches_host(c, host): log.debug("[control connection] Control connection host (%s) is being removed. Reconnecting", host) @@ -4934,7 +5152,15 @@ def _on_timeout(self, _attempts=0): # Capture connection stats before pool.return_connection() can alter state conn_in_flight = self._connection.in_flight - pool = self.session._pools.get(self._current_host) + pool = None + if self._connection.is_control_connection: + with self._connection.lock: + self._connection.orphaned_request_ids.add(self._req_id) + if len(self._connection.orphaned_request_ids) >= self._connection.orphaned_threshold: + self._connection.orphaned_threshold_reached = True + else: + pool = self.session._pools.get(self._current_host) + if pool and not pool.is_shutdown: # Do not return the stream ID to the pool yet. We cannot reuse it # because the node might still be processing the query and will @@ -4947,12 +5173,6 @@ def _on_timeout(self, _attempts=0): self._connection.orphaned_threshold_reached = True pool.return_connection(self._connection, stream_was_orphaned=True) - elif self._connection.is_control_connection: - with self._connection.lock: - self._connection.orphaned_request_ids.add(self._req_id) - if len(self._connection.orphaned_request_ids) >= self._connection.orphaned_threshold: - self._connection.orphaned_threshold_reached = True - errors = self._errors if not errors: if self.is_schema_agreed: diff --git a/cassandra/pool.py b/cassandra/pool.py index 3e8bfcbd66..bf87b714c1 100644 --- a/cassandra/pool.py +++ b/cassandra/pool.py @@ -243,7 +243,7 @@ def _clear_reconnection_handler(self, handler): if self._reconnection_handler is not handler: return False self._reconnection_handler = None - return True + return not self._is_removed def __eq__(self, other): if not isinstance(other, Host): diff --git a/tests/unit/test_cluster.py b/tests/unit/test_cluster.py index 74ed346c68..c591491503 100644 --- a/tests/unit/test_cluster.py +++ b/tests/unit/test_cluster.py @@ -16,16 +16,17 @@ from concurrent.futures import Future import logging import socket +from threading import RLock from types import SimpleNamespace -from unittest.mock import patch, Mock +from unittest.mock import ANY, Mock, call, patch import uuid from cassandra import ConsistencyLevel, DriverException, Timeout, Unavailable, RequestExecutionException, ReadTimeout, WriteTimeout, CoordinationFailure, ReadFailure, WriteFailure, FunctionFailure, AlreadyExists,\ InvalidRequest, Unauthorized, AuthenticationFailed, OperationTimedOut, UnsupportedOperation, RequestValidationException, ConfigurationException, ProtocolVersion from cassandra.cluster import _Scheduler, Session, Cluster, ResultSet, SchemaAgreementScope, ControlConnectionQueryFallback, default_lbp_factory, \ ExecutionProfile, _ConfigMode, EXEC_PROFILE_DEFAULT -from cassandra.connection import ConnectionBusy, ConnectionException +from cassandra.connection import ConnectionBusy, ConnectionException, DefaultEndPoint from cassandra.driver_config import DriverConfigReporter from cassandra.pool import Host from cassandra.policies import HostDistance, RetryPolicy, RoundRobinPolicy, DowngradingConsistencyRetryPolicy, SimpleConvictionPolicy @@ -240,6 +241,551 @@ def test_control_connection_query_fallback_fallback_tolerates_empty_initial_pool assert session._initial_connect_futures == {future} assert session._pools == {} + def test_on_add_waits_until_every_session_pool_is_scheduled(self): + cluster = Cluster() + self.addCleanup(cluster.shutdown) + cluster.profile_manager = Mock() + cluster.profile_manager.distance.return_value = HostDistance.LOCAL + cluster.control_connection.on_add = Mock() + cluster._prepare_all_queries = Mock() + + host = Host( + "127.0.0.1", SimpleConvictionPolicy, host_id=uuid.uuid4()) + cluster.metadata.add_or_return_host(host) + + def make_session(): + session = Session.__new__(Session) + session.cluster = cluster + session._profile_manager = cluster.profile_manager + session._pools = {} + session.shutdown = Mock() + session.update_created_pools = Mock( + wraps=session.update_created_pools) + return session + + first_session = make_session() + second_session = make_session() + first_pool = Mock( + host=host, is_shutdown=False, host_distance=HostDistance.LOCAL) + second_pool = Mock( + host=host, is_shutdown=False, host_distance=HostDistance.LOCAL) + + first_future = Future() + first_future.set_result(True) + second_future = Future() + reconciliation_future = Future() + + def add_first_pool(pool_host, is_host_addition): + first_session._pools[pool_host] = first_pool + return first_future + + def add_second_pool(pool_host, is_host_addition): + return second_future if is_host_addition else reconciliation_future + + first_session.add_or_renew_pool = Mock(side_effect=add_first_pool) + second_session.add_or_renew_pool = Mock(side_effect=add_second_pool) + cluster.sessions = (first_session, second_session) + + listener = Mock() + cluster.register_listener(listener) + + cluster.on_add(host, refresh_nodes=False) + + first_session.add_or_renew_pool.assert_called_once_with( + host, is_host_addition=True) + second_session.add_or_renew_pool.assert_called_once_with( + host, is_host_addition=True) + listener.on_add.assert_not_called() + assert host.is_up is None + + # Model the pending task publishing its pool before completing. + second_session._pools[host] = second_pool + second_future.set_result(True) + + listener.on_add.assert_called_once_with(host) + assert host.is_up is True + first_session.update_created_pools.assert_called_once_with() + second_session.update_created_pools.assert_called_once_with() + + def test_failed_replacement_add_removes_partial_pools_and_reconnects(self): + cluster = Cluster() + self.addCleanup(cluster.shutdown) + cluster.profile_manager = Mock() + cluster.profile_manager.distance.return_value = HostDistance.LOCAL + cluster.control_connection.on_add = Mock() + cluster._prepare_all_queries = Mock() + cluster._start_reconnector = Mock() + + host = Host( + "127.0.0.1", SimpleConvictionPolicy, host_id=uuid.uuid4()) + cluster.metadata.add_or_return_host(host) + + def make_session(pool_future): + session = Session.__new__(Session) + session.cluster = cluster + session._profile_manager = cluster.profile_manager + session._pools = {} + session.is_shutdown = False + session.submit = lambda fn, *args, **kwargs: fn(*args, **kwargs) + session.add_or_renew_pool = Mock(return_value=pool_future) + session.update_created_pools = Mock( + wraps=session.update_created_pools) + session.shutdown = Mock() + return session + + successful_future = Future() + failed_future = Future() + successful_session = make_session(successful_future) + failed_session = make_session(failed_future) + cluster.sessions = (successful_session, failed_session) + + listener = Mock() + cluster.register_listener(listener) + + cluster.on_add( + host, refresh_nodes=False, reconcile_pools_on_failure=True) + + partial_pool = Mock(host=host, is_shutdown=False) + partial_pool.get_state.return_value = {"open_count": 1} + successful_session._pools[host] = partial_pool + successful_future.set_result(True) + + # The healthy pool makes the failing session's DOWN signal look like + # an isolated connection failure, so Cluster.on_down() discounts it. + cluster.signal_connection_failure( + host, ConnectionException("pool creation failed"), + is_host_addition=True, expect_host_to_be_down=True) + assert host.is_up is None + cluster._start_reconnector.assert_not_called() + + failed_future.set_result(False) + + assert host.is_up is False + assert successful_session._pools == {} + partial_pool.shutdown.assert_called_once_with() + successful_session.update_created_pools.assert_called_once_with( + excluded_host=host, hosts=ANY) + failed_session.update_created_pools.assert_called_once_with( + excluded_host=host, hosts=ANY) + cluster._start_reconnector.assert_called_once_with( + host, is_host_addition=True, + on_add_reconnection=ANY, start=False) + listener.on_add.assert_not_called() + + def test_replacement_reconnector_preserves_recovery_context(self): + cluster = Cluster() + self.addCleanup(cluster.shutdown) + cluster.profile_manager = Mock() + cluster.profile_manager.distance.return_value = HostDistance.LOCAL + cluster._prepare_all_queries = Mock() + cluster.control_connection.refresh_node_list_and_token_map = Mock() + cluster.control_connection.on_add = Mock( + wraps=cluster.control_connection.on_add) + cluster.scheduler.schedule = Mock() + + probe_connections = [Mock(), Mock()] + connection_factory = Mock(side_effect=probe_connections) + cluster._make_connection_factory = Mock( + return_value=connection_factory) + + host = Host( + "127.0.0.1", SimpleConvictionPolicy, host_id=uuid.uuid4()) + cluster.metadata.add_or_return_host(host) + + def submit(fn, *args, **kwargs): + future = Future() + try: + future.set_result(fn(*args, **kwargs)) + except Exception as exc: + future.set_exception(exc) + return future + + def make_session(): + session = Session.__new__(Session) + session.cluster = cluster + session._profile_manager = cluster.profile_manager + session._pools = {} + session._lock = RLock() + session.keyspace = None + session.is_shutdown = False + session.submit = submit + session.update_created_pools = Mock( + wraps=session.update_created_pools) + session.shutdown = Mock() + return session + + successful_session = make_session() + failing_session = make_session() + cluster.sessions = (successful_session, failing_session) + + created_pools = [] + + def create_pool(pool_host, distance, session): + if session is failing_session: + raise ConnectionException("pool creation failed") + + pool = Mock( + host=pool_host, host_distance=distance, + is_shutdown=False, _keyspace=None) + pool.get_state.return_value = {"open_count": 1} + created_pools.append(pool) + return pool + + listener = Mock() + cluster.register_listener(listener) + + with patch('cassandra.cluster.HostConnection', side_effect=create_pool): + cluster.on_add( + host, refresh_nodes=False, + reconcile_pools_on_failure=True) + + # Exercise two complete reconnector handoffs. On each attempt, the + # first session's open pool makes the second session's DOWN signal + # get discounted, so replacement-aware aggregate recovery is the + # only owner that can install the successor. + cluster.scheduler.schedule.call_args_list[0].args[1]() + cluster.scheduler.schedule.call_args_list[1].args[1]() + + assert cluster.scheduler.schedule.call_count == 3 + pending_run = cluster.scheduler.schedule.call_args_list[2].args[1] + assert host._reconnection_handler is pending_run.__self__ + assert host.is_up is False + assert len(created_pools) == 3 + for pool in created_pools: + pool.shutdown.assert_called_once_with() + expected_reconciliation = [ + call(excluded_host=host, hosts=ANY)] * 3 + assert successful_session.update_created_pools.call_args_list == \ + expected_reconciliation + assert failing_session.update_created_pools.call_args_list == \ + expected_reconciliation + assert cluster.control_connection.on_add.call_args_list == \ + [call(host, False)] * 3 + cluster.control_connection.refresh_node_list_and_token_map \ + .assert_not_called() + listener.on_add.assert_not_called() + + def test_unconvicted_replacement_retry_keeps_reconnecting(self): + cluster = Cluster() + self.addCleanup(cluster.shutdown) + cluster.profile_manager = Mock() + cluster.profile_manager.distance.return_value = HostDistance.LOCAL + cluster._prepare_all_queries = Mock() + cluster.control_connection.on_add = Mock() + cluster.scheduler.schedule = Mock() + probe = Mock() + cluster._make_connection_factory = Mock(return_value=Mock( + return_value=probe)) + + host = Host( + "127.0.0.1", SimpleConvictionPolicy, host_id=uuid.uuid4()) + cluster.metadata.add_or_return_host(host) + + session = Session.__new__(Session) + session.cluster = cluster + session._profile_manager = cluster.profile_manager + session._pools = {} + session._lock = RLock() + session.keyspace = None + session.is_shutdown = False + session.shutdown = Mock() + session.update_created_pools = Mock(return_value=set()) + + def submit(fn, *args, **kwargs): + future = Future() + try: + future.set_result(fn(*args, **kwargs)) + except Exception as exc: + future.set_exception(exc) + return future + + session.submit = submit + cluster.sessions.add(session) + + with patch('cassandra.cluster.HostConnection', + side_effect=OperationTimedOut()): + cluster.on_add( + host, refresh_nodes=False, + reconcile_pools_on_failure=True) + first_handler = host._reconnection_handler + first_handler.run() + + assert cluster.scheduler.schedule.call_count == 2 + assert host.is_currently_reconnecting() + assert host._reconnection_handler is not first_handler + probe.close.assert_called_once_with() + + def test_failed_replacement_keyspace_sync_starts_reconnector(self): + cluster = Cluster() + self.addCleanup(cluster.shutdown) + cluster.profile_manager = Mock() + cluster.profile_manager.distance.return_value = HostDistance.LOCAL + cluster.control_connection.on_add = Mock() + cluster._prepare_all_queries = Mock() + cluster._start_reconnector = Mock() + + existing_reconnector = Mock() + aggregate_reconnector = Mock() + + def install_reconnector(host, is_host_addition, + on_add_reconnection=None, start=True): + old_reconnector = host.get_and_set_reconnection_handler( + aggregate_reconnector) + if old_reconnector: + old_reconnector.cancel() + return aggregate_reconnector + + cluster._start_reconnector.side_effect = install_reconnector + + def handle_down(host, is_host_addition, + on_add_reconnection=None): + host.get_and_set_reconnection_handler(existing_reconnector) + + cluster.on_down_potentially_blocking = Mock( + side_effect=handle_down) + + host = Host( + "127.0.0.1", SimpleConvictionPolicy, host_id=uuid.uuid4()) + cluster.metadata.add_or_return_host(host) + + session = Session.__new__(Session) + session.cluster = cluster + session._profile_manager = cluster.profile_manager + session._pools = {} + session._lock = RLock() + session.keyspace = "new_keyspace" + session.is_shutdown = False + session.update_created_pools = Mock(return_value=set()) + session.shutdown = Mock() + + def submit(fn, *args, **kwargs): + future = Future() + try: + future.set_result(fn(*args, **kwargs)) + except Exception as exc: + future.set_exception(exc) + return future + + session.submit = submit + cluster.sessions.add(session) + + new_pool = Mock( + host=host, host_distance=HostDistance.LOCAL, + is_shutdown=False, _keyspace="old_keyspace") + keyspace_error = ConnectionException("keyspace update failed") + + def fail_keyspace_update(_keyspace, callback): + callback(new_pool, [keyspace_error]) + + new_pool._set_keyspace_for_all_conns.side_effect = \ + fail_keyspace_update + + with patch('cassandra.cluster.HostConnection', return_value=new_pool): + cluster.on_add( + host, refresh_nodes=False, + reconcile_pools_on_failure=True) + + assert host.is_up is False + assert session._pools == {} + new_pool.shutdown.assert_called_once_with() + new_pool._set_keyspace_for_all_conns.assert_called_once_with( + "new_keyspace", ANY) + cluster.on_down_potentially_blocking.assert_called_once_with( + host, True, on_add_reconnection=ANY) + cluster._start_reconnector.assert_called_once_with( + host, is_host_addition=True, + on_add_reconnection=ANY, start=False) + existing_reconnector.cancel.assert_called_once_with() + aggregate_reconnector.start.assert_called_once_with() + assert host._reconnection_handler is aggregate_reconnector + + def test_unconvicted_replacement_failure_starts_reconnector(self): + cluster = Cluster() + self.addCleanup(cluster.shutdown) + cluster.profile_manager = Mock() + cluster.profile_manager.distance.return_value = HostDistance.LOCAL + cluster.control_connection.on_add = Mock() + cluster._prepare_all_queries = Mock() + cluster._start_reconnector = Mock() + + host = Host( + "127.0.0.1", SimpleConvictionPolicy, host_id=uuid.uuid4()) + cluster.metadata.add_or_return_host(host) + + session = Session.__new__(Session) + session.cluster = cluster + session._profile_manager = cluster.profile_manager + session._pools = {} + session._lock = RLock() + session.keyspace = None + session.is_shutdown = False + session.shutdown = Mock() + session.update_created_pools = Mock(return_value=set()) + + def submit(fn, *args, **kwargs): + future = Future() + try: + future.set_result(fn(*args, **kwargs)) + except Exception as exc: + future.set_exception(exc) + return future + + session.submit = submit + cluster.sessions.add(session) + + # SimpleConvictionPolicy deliberately does not convict a host for an + # OperationTimedOut. Aggregate replacement handling must still own the + # retry after pool creation fails. + with patch('cassandra.cluster.HostConnection', + side_effect=OperationTimedOut()): + cluster.on_add( + host, refresh_nodes=False, + reconcile_pools_on_failure=True) + + assert host.is_up is False + assert session._pools == {} + session.update_created_pools.assert_called_once_with( + excluded_host=host, hosts=ANY) + cluster._start_reconnector.assert_called_once_with( + host, is_host_addition=True, + on_add_reconnection=ANY, start=False) + + def test_pool_creation_cannot_publish_after_host_removal(self): + cluster = Cluster() + self.addCleanup(cluster.shutdown) + cluster.profile_manager = Mock() + cluster.profile_manager.distance.return_value = HostDistance.LOCAL + cluster.control_connection.on_remove = Mock() + + host = Host( + "127.0.0.1", SimpleConvictionPolicy, host_id=uuid.uuid4()) + cluster.metadata.add_or_return_host(host) + + session = Session.__new__(Session) + session.cluster = cluster + session._profile_manager = cluster.profile_manager + session._pools = {} + session._lock = RLock() + session.keyspace = None + session.is_shutdown = False + session.shutdown = Mock() + + def submit(fn, *args, **kwargs): + future = Future() + try: + future.set_result(fn(*args, **kwargs)) + except Exception as exc: + future.set_exception(exc) + return future + + session.submit = submit + cluster.sessions.add(session) + + new_pool = Mock( + host=host, host_distance=HostDistance.LOCAL, + is_shutdown=False, _keyspace=None) + + def create_pool(*args, **kwargs): + cluster.remove_host(host, trigger_reconciliation=False) + return new_pool + + with patch('cassandra.cluster.HostConnection', side_effect=create_pool): + future = session.add_or_renew_pool( + host, is_host_addition=True) + + assert future.result() is False + assert session._pools == {} + new_pool.shutdown.assert_called_once_with() + + def test_removed_replacement_failure_does_not_reconcile(self): + cluster = Cluster() + self.addCleanup(cluster.shutdown) + cluster.profile_manager = Mock() + cluster.profile_manager.distance.return_value = HostDistance.LOCAL + cluster.control_connection.on_add = Mock() + cluster._prepare_all_queries = Mock() + cluster._start_reconnector = Mock() + + host = Host( + "127.0.0.1", SimpleConvictionPolicy, host_id=uuid.uuid4()) + cluster.metadata.add_or_return_host(host) + pool_future = Future() + session = Mock() + session.add_or_renew_pool.return_value = pool_future + cluster.sessions = (session,) + + cluster.on_add( + host, refresh_nodes=False, reconcile_pools_on_failure=True) + cluster.remove_host(host, trigger_reconciliation=False) + session.reset_mock() + + pool_future.set_result(False) + + session.remove_pool.assert_not_called() + session.update_created_pools.assert_not_called() + cluster._start_reconnector.assert_not_called() + + def test_replacement_removed_during_failure_cleanup_does_not_reconcile(self): + cluster = Cluster() + self.addCleanup(cluster.shutdown) + cluster.profile_manager = Mock() + cluster.profile_manager.distance.return_value = HostDistance.LOCAL + cluster.control_connection.on_add = Mock() + cluster.control_connection.on_remove = Mock() + cluster._prepare_all_queries = Mock() + + host = Host( + "127.0.0.1", SimpleConvictionPolicy, host_id=uuid.uuid4()) + cluster.metadata.add_or_return_host(host) + pool_future = Future() + session = Mock() + session.add_or_renew_pool.return_value = pool_future + cluster.sessions = (session,) + + cluster.on_add( + host, refresh_nodes=False, reconcile_pools_on_failure=True) + + replacement_id = uuid.uuid4() + + def remove_and_replace(_host): + cluster.remove_host(host, trigger_reconciliation=False) + cluster.add_host( + host.endpoint, signal=False, host_id=replacement_id) + + session.remove_pool.side_effect = remove_and_replace + pool_future.set_result(False) + + replacement = cluster.metadata.get_host_by_host_id(replacement_id) + assert replacement is not None + assert replacement.endpoint == host.endpoint + session.update_created_pools.assert_not_called() + assert not host.is_currently_reconnecting() + + def test_reconnector_cannot_be_installed_after_host_removal(self): + cluster = Cluster() + self.addCleanup(cluster.shutdown) + cluster.profile_manager = Mock() + cluster.profile_manager.distance.return_value = HostDistance.LOCAL + + host = Host( + "127.0.0.1", SimpleConvictionPolicy, host_id=uuid.uuid4()) + cluster.metadata.add_or_return_host(host) + + def remove_during_setup(_host): + cluster.remove_host(host, trigger_reconciliation=False) + return Mock() + + cluster._make_connection_factory = Mock( + side_effect=remove_during_setup) + + result = cluster._start_reconnector( + host, is_host_addition=True, + on_add_reconnection=Mock()) + + assert result is None + assert host._is_removed + assert not host.is_currently_reconnecting() + def test_compression_autodisabled_without_libraries(self): with patch.dict('cassandra.cluster.locally_supported_compressions', {}, clear=True): with patch('cassandra.cluster.log') as patched_logger: @@ -771,6 +1317,61 @@ def test_set_keyspace_for_all_pools_reports_all_errors(self, *_): callback.assert_called_once() assert callback.call_args.args[0] == {'host1': [keyspace_error]} + def test_remove_pool_after_host_endpoint_changes(self): + session = Session.__new__(Session) + host = Host("127.0.0.1", SimpleConvictionPolicy, host_id=uuid.uuid4()) + pool = Mock(host=host) + shutdown_future = Future() + session._pools = {host: pool} + session.cluster = Mock() + session.cluster.executor.submit.return_value = shutdown_future + session.is_shutdown = False + + host.endpoint = DefaultEndPoint("127.0.0.2") + same_host_at_another_endpoint = Host( + "127.0.0.3", SimpleConvictionPolicy, host_id=host.host_id) + + assert session._pools[host] is pool + assert session._pools[same_host_at_another_endpoint] is pool + assert session.remove_pool(same_host_at_another_endpoint) is shutdown_future + assert session._pools == {} + session.cluster.executor.submit.assert_called_once_with(pool.shutdown) + + def test_pool_renewal_uses_pool_host_not_retained_dict_key(self): + session = Session.__new__(Session) + host_id = uuid.uuid4() + original_host = Host( + "127.0.0.1", SimpleConvictionPolicy, host_id=host_id) + current_host = Host( + "127.0.0.2", SimpleConvictionPolicy, host_id=host_id) + old_pool = Mock(host=original_host) + pool_state = {"open_count": 1} + new_pool = Mock(host=current_host, _keyspace=None) + new_pool.get_state.return_value = pool_state + session._pools = {original_host: old_pool} + session._lock = RLock() + session.keyspace = None + session.is_shutdown = False + session.cluster = Mock( + allow_control_connection_query_fallback=ControlConnectionQueryFallback.Disabled) + session._profile_manager = Mock() + session._profile_manager.distance.return_value = HostDistance.LOCAL + session.submit = lambda fn, *args, **kwargs: fn(*args, **kwargs) + + with patch('cassandra.cluster.HostConnection', return_value=new_pool): + assert session.add_or_renew_pool( + current_host, is_host_addition=False) + + old_pool.shutdown.assert_called_once_with() + assert len(session._pools) == 1 + assert session._pools[current_host] is new_pool + # Assigning through the equal current Host retains the original key + # object, so public state must take its Host from the replacement pool. + assert next(iter(session._pools)) is original_host + state = session.get_pool_state() + assert next(iter(state)) is current_host + assert state[current_host] == pool_state + class ProtocolVersionTests(unittest.TestCase): def test_protocol_downgrade_test(self): diff --git a/tests/unit/test_control_connection.py b/tests/unit/test_control_connection.py index fb5abdb3f6..b90a572987 100644 --- a/tests/unit/test_control_connection.py +++ b/tests/unit/test_control_connection.py @@ -13,36 +13,46 @@ # limitations under the License. import unittest +import uuid -from concurrent.futures import ThreadPoolExecutor +from concurrent.futures import Future, ThreadPoolExecutor from unittest.mock import Mock, ANY, call, patch from cassandra import OperationTimedOut, SchemaTargetType, SchemaChangeType from cassandra.protocol import ResultMessage, RESULT_KIND_ROWS -from cassandra.cluster import (Cluster, ControlConnection, _Scheduler, +from cassandra.cluster import (Cluster, ControlConnection, Session, _Scheduler, ProfileManager, EXEC_PROFILE_DEFAULT, ExecutionProfile) from cassandra.pool import Host -from cassandra.connection import (ConnectionException, EndPoint, DefaultEndPoint, - DefaultEndPointFactory, UnixSocketEndPoint) -from cassandra.policies import (SimpleConvictionPolicy, RoundRobinPolicy, +from cassandra.connection import (ConnectionException, EndPoint, + DefaultEndPoint, DefaultEndPointFactory, + UnixSocketEndPoint) +from cassandra.policies import (DCAwareRoundRobinPolicy, HostDistance, + SimpleConvictionPolicy, RoundRobinPolicy, ConstantReconnectionPolicy, IdentityTranslator) PEER_IP = "foobar" +HOST_ID_1 = uuid.UUID(int=1) +HOST_ID_2 = uuid.UUID(int=2) +HOST_ID_3 = uuid.UUID(int=3) +HOST_ID_4 = uuid.UUID(int=4) +HOST_ID_6 = uuid.UUID(int=6) +HOST_ID_7 = uuid.UUID(int=7) + class MockMetadata(object): def __init__(self): self.hosts = { - 'uuid1': Host(endpoint=DefaultEndPoint("192.168.1.0"), conviction_policy_factory=SimpleConvictionPolicy, host_id='uuid1'), - 'uuid2': Host(endpoint=DefaultEndPoint("192.168.1.1"), conviction_policy_factory=SimpleConvictionPolicy, host_id='uuid2'), - 'uuid3': Host(endpoint=DefaultEndPoint("192.168.1.2"), conviction_policy_factory=SimpleConvictionPolicy, host_id='uuid3') + HOST_ID_1: Host(endpoint=DefaultEndPoint("192.168.1.0"), conviction_policy_factory=SimpleConvictionPolicy, host_id=HOST_ID_1), + HOST_ID_2: Host(endpoint=DefaultEndPoint("192.168.1.1"), conviction_policy_factory=SimpleConvictionPolicy, host_id=HOST_ID_2), + HOST_ID_3: Host(endpoint=DefaultEndPoint("192.168.1.2"), conviction_policy_factory=SimpleConvictionPolicy, host_id=HOST_ID_3) } self._host_id_by_endpoint = { - DefaultEndPoint("192.168.1.0"): 'uuid1', - DefaultEndPoint("192.168.1.1"): 'uuid2', - DefaultEndPoint("192.168.1.2"): 'uuid3', + DefaultEndPoint("192.168.1.0"): HOST_ID_1, + DefaultEndPoint("192.168.1.1"): HOST_ID_2, + DefaultEndPoint("192.168.1.2"): HOST_ID_3, } for host in self.hosts.values(): host.set_up() @@ -115,13 +125,15 @@ def __init__(self): self.endpoint_factory = DefaultEndPointFactory().configure(self) self.ssl_options = None - def add_host(self, endpoint, datacenter, rack, signal=False, refresh_nodes=True, host_id=None): + def add_host(self, endpoint, datacenter, rack, signal=False, + refresh_nodes=True, host_id=None, + reconcile_pools_on_failure=False): host = Host(endpoint, SimpleConvictionPolicy, datacenter, rack, host_id=host_id) host, _ = self.metadata.add_or_return_host(host) self.added_hosts.append(host) return host, True - def remove_host(self, host): + def remove_host(self, host, trigger_reconciliation=True): pass def on_up(self, host): @@ -149,28 +161,32 @@ def _node_meta_results(local_results, peer_results): class MockConnection(object): is_defunct = False + is_closed = False def __init__(self): self.endpoint = DefaultEndPoint("192.168.1.0") self.original_endpoint = self.endpoint self.local_results = [ ["rpc_address", "schema_version", "cluster_name", "data_center", "rack", "partitioner", "release_version", "tokens", "host_id", "listen_address"], - [["192.168.1.0", "a", "foocluster", "dc1", "rack1", "Murmur3Partitioner", "2.2.0", ["0", "100", "200"], "uuid1", "192.168.1.0"]] + [["192.168.1.0", "a", "foocluster", "dc1", "rack1", "Murmur3Partitioner", "2.2.0", ["0", "100", "200"], HOST_ID_1, "192.168.1.0"]] ] self.peer_results = [ ["rpc_address", "peer", "schema_version", "data_center", "rack", "tokens", "host_id"], - [["192.168.1.1", "10.0.0.1", "a", "dc1", "rack1", ["1", "101", "201"], "uuid2"], - ["192.168.1.2", "10.0.0.2", "a", "dc1", "rack1", ["2", "102", "202"], "uuid3"]] + [["192.168.1.1", "10.0.0.1", "a", "dc1", "rack1", ["1", "101", "201"], HOST_ID_2], + ["192.168.1.2", "10.0.0.2", "a", "dc1", "rack1", ["2", "102", "202"], HOST_ID_3]] ] self.peer_results_v2 = [ ["native_address", "native_port", "peer", "peer_port", "schema_version", "data_center", "rack", "tokens", "host_id"], - [["192.168.1.1", 9042, "10.0.0.1", 7042, "a", "dc1", "rack1", ["1", "101", "201"], "uuid2"], - ["192.168.1.2", 9042, "10.0.0.2", 7040, "a", "dc1", "rack1", ["2", "102", "202"], "uuid3"]] + [["192.168.1.1", 9042, "10.0.0.1", 7042, "a", "dc1", "rack1", ["1", "101", "201"], HOST_ID_2], + ["192.168.1.2", 9042, "10.0.0.2", 7040, "a", "dc1", "rack1", ["2", "102", "202"], HOST_ID_3]] ] self.wait_for_responses = Mock(return_value=_node_meta_results(self.local_results, self.peer_results)) + def close(self): + self.is_closed = True + class FakeTime(object): @@ -188,17 +204,17 @@ class ControlConnectionTest(unittest.TestCase): _matching_schema_preloaded_results = _node_meta_results( local_results=(["rpc_address", "schema_version", "cluster_name", "data_center", "rack", "partitioner", "release_version", "tokens", "host_id", "listen_address"], - [["192.168.1.0", "a", "foocluster", "dc1", "rack1", "Murmur3Partitioner", "2.2.0", ["0", "100", "200"], "uuid1", "192.168.1.0"]]), + [["192.168.1.0", "a", "foocluster", "dc1", "rack1", "Murmur3Partitioner", "2.2.0", ["0", "100", "200"], HOST_ID_1, "192.168.1.0"]]), peer_results=(["rpc_address", "peer", "schema_version", "data_center", "rack", "tokens", "host_id"], - [["192.168.1.1", "10.0.0.1", "a", "dc1", "rack1", ["1", "101", "201"], "uuid2"], - ["192.168.1.2", "10.0.0.2", "a", "dc1", "rack1", ["2", "102", "202"], "uuid3"]])) + [["192.168.1.1", "10.0.0.1", "a", "dc1", "rack1", ["1", "101", "201"], HOST_ID_2], + ["192.168.1.2", "10.0.0.2", "a", "dc1", "rack1", ["2", "102", "202"], HOST_ID_3]])) _nonmatching_schema_preloaded_results = _node_meta_results( local_results=(["rpc_address", "schema_version", "cluster_name", "data_center", "rack", "partitioner", "release_version", "tokens", "host_id", "listen_address"], - [["192.168.1.0", "a", "foocluster", "dc1", "rack1", "Murmur3Partitioner", "2.2.0", ["0", "100", "200"], "uuid1", "192.168.1.0"]]), + [["192.168.1.0", "a", "foocluster", "dc1", "rack1", "Murmur3Partitioner", "2.2.0", ["0", "100", "200"], HOST_ID_1, "192.168.1.0"]]), peer_results=(["rpc_address", "peer", "schema_version", "data_center", "rack", "tokens", "host_id"], - [["192.168.1.1", "10.0.0.1", "a", "dc1", "rack1", ["1", "101", "201"], "uuid2"], - ["192.168.1.2", "10.0.0.2", "b", "dc1", "rack1", ["2", "102", "202"], "uuid3"]])) + [["192.168.1.1", "10.0.0.1", "a", "dc1", "rack1", ["1", "101", "201"], HOST_ID_2], + ["192.168.1.2", "10.0.0.2", "b", "dc1", "rack1", ["2", "102", "202"], HOST_ID_3]])) def setUp(self): self.cluster = MockCluster() @@ -213,7 +229,7 @@ def setUp(self): def _forget_local_host(self): endpoint = DefaultEndPoint('192.168.1.0') self.cluster.metadata._host_id_by_endpoint.pop(endpoint) - self.cluster.metadata.hosts.pop('uuid1') + self.cluster.metadata.hosts.pop(HOST_ID_1) def _discover_local_host_over_unix(self): maintenance_endpoint = UnixSocketEndPoint('/tmp/maintenance.sock') @@ -221,7 +237,7 @@ def _discover_local_host_over_unix(self): self.connection.endpoint = maintenance_endpoint self.connection.original_endpoint = maintenance_endpoint self.control_connection.refresh_node_list_and_token_map() - local_host = self.cluster.metadata.get_host_by_host_id('uuid1') + local_host = self.cluster.metadata.get_host_by_host_id(HOST_ID_1) local_host.set_up() return maintenance_endpoint, local_host @@ -305,9 +321,9 @@ def test_wait_for_schema_agreement_rpc_lookup(self): If the rpc_address is 0.0.0.0, the "peer" column should be used instead. """ self.connection.peer_results[1].append( - ["0.0.0.0", PEER_IP, "b", "dc1", "rack1", ["3", "103", "203"], "uuid6"] + ["0.0.0.0", PEER_IP, "b", "dc1", "rack1", ["3", "103", "203"], HOST_ID_6] ) - host = Host(DefaultEndPoint("0.0.0.0"), SimpleConvictionPolicy, host_id='uuid6') + host = Host(DefaultEndPoint("0.0.0.0"), SimpleConvictionPolicy, host_id=HOST_ID_6) self.cluster.metadata.hosts[host.host_id] = host self.cluster.metadata._host_id_by_endpoint[DefaultEndPoint(PEER_IP)] = host.host_id host.is_up = False @@ -392,10 +408,10 @@ def test_refresh_uses_control_endpoint_for_local_unix_host(self): self.control_connection.refresh_node_list_and_token_map() - local_host = self.cluster.metadata.get_host_by_host_id('uuid1') + local_host = self.cluster.metadata.get_host_by_host_id(HOST_ID_1) assert local_host.endpoint == maintenance_endpoint assert local_host.broadcast_rpc_address == '192.168.1.0' - peer_host = self.cluster.metadata.get_host_by_host_id('uuid2') + peer_host = self.cluster.metadata.get_host_by_host_id(HOST_ID_2) assert peer_host.endpoint == DefaultEndPoint('192.168.1.1') assert sorted([local_host, peer_host]) == \ sorted([peer_host, local_host]) @@ -407,11 +423,11 @@ def test_refresh_checks_unix_local_advertised_endpoint_for_duplicates(self): UnixSocketEndPoint('/tmp/maintenance.sock') self.connection.peer_results[1].append([ '192.168.1.0', '10.0.0.4', 'a', 'dc1', 'rack1', - ['4', '104', '204'], 'uuid4']) + ['4', '104', '204'], HOST_ID_4]) self.control_connection.refresh_node_list_and_token_map() - assert self.cluster.metadata.get_host_by_host_id('uuid4') is None + assert self.cluster.metadata.get_host_by_host_id(HOST_ID_4) is None def test_refresh_preserves_known_unix_endpoint_when_host_becomes_peer(self): maintenance_endpoint = UnixSocketEndPoint('/tmp/maintenance.sock') @@ -424,13 +440,13 @@ def test_refresh_preserves_known_unix_endpoint_when_host_becomes_peer(self): self.connection.local_results[0], [['192.168.1.1', 'a', 'foocluster', 'dc1', 'rack1', 'Murmur3Partitioner', '2.2.0', ['1', '101', '201'], - 'uuid2', '192.168.1.1']]) + HOST_ID_2, '192.168.1.1']]) peer_results = ( self.connection.peer_results[0], [['192.168.1.0', '10.0.0.1', 'a', 'dc1', 'rack1', - ['0', '100', '200'], 'uuid1'], + ['0', '100', '200'], HOST_ID_1], ['192.168.1.2', '10.0.0.2', 'a', 'dc1', 'rack1', - ['2', '102', '202'], 'uuid3']]) + ['2', '102', '202'], HOST_ID_3]]) self.connection.endpoint = DefaultEndPoint('192.168.1.1') self.connection.original_endpoint = self.connection.endpoint @@ -438,7 +454,7 @@ def test_refresh_preserves_known_unix_endpoint_when_host_becomes_peer(self): self.connection, preloaded_results=_node_meta_results(local_results, peer_results)) - local_host = self.cluster.metadata.get_host_by_host_id('uuid1') + local_host = self.cluster.metadata.get_host_by_host_id(HOST_ID_1) assert local_host.endpoint == maintenance_endpoint peer_results[1][0][2] = 'b' @@ -453,11 +469,11 @@ def test_refresh_uses_factory_for_local_network_host(self): self.control_connection.refresh_node_list_and_token_map() - local_host = self.cluster.metadata.get_host_by_host_id('uuid1') + local_host = self.cluster.metadata.get_host_by_host_id(HOST_ID_1) assert local_host.endpoint == DefaultEndPoint('192.168.1.0') def test_schema_query_uses_shard_aware_connection_original_endpoint(self): - host = self.cluster.metadata.get_host_by_host_id('uuid1') + host = self.cluster.metadata.get_host_by_host_id(HOST_ID_1) self.connection.endpoint = DefaultEndPoint('192.168.1.0', 19042) self.connection.original_endpoint = host.endpoint self.control_connection._uses_peers_v2 = False @@ -476,7 +492,7 @@ def test_refresh_network_local_preserves_known_unix_endpoint(self): self._refresh_control_connection_over_network() - assert self.cluster.metadata.get_host_by_host_id('uuid1') is local_host + assert self.cluster.metadata.get_host_by_host_id(HOST_ID_1) is local_host assert local_host.endpoint == maintenance_endpoint assert host_index[local_host] is not None assert Cluster.get_control_connection_host(self.cluster) is local_host @@ -513,7 +529,7 @@ def test_unix_signal_error_reconnects_if_down_notification_suppressed(self): def test_tcp_route_mismatch_reconnects_if_down_notification_suppressed(self): self.control_connection.refresh_node_list_and_token_map() - local_host = self.cluster.metadata.get_host_by_host_id('uuid1') + local_host = self.cluster.metadata.get_host_by_host_id(HOST_ID_1) self.connection.endpoint = DefaultEndPoint('192.168.1.0', 19042) self.connection.original_endpoint = local_host.endpoint connection_error = ConnectionException('control connection failed') @@ -603,10 +619,10 @@ def test_remove_matches_control_connection_by_host_id(self): self.connection.endpoint = maintenance_endpoint self.connection.original_endpoint = maintenance_endpoint self.control_connection.refresh_node_list_and_token_map() - local_host = self.cluster.metadata.get_host_by_host_id('uuid1') + local_host = self.cluster.metadata.get_host_by_host_id(HOST_ID_1) self.connection.endpoint = DefaultEndPoint('192.168.1.0') - self.cluster.metadata.hosts.pop('uuid1') + self.cluster.metadata.hosts.pop(HOST_ID_1) self.cluster.metadata._host_id_by_endpoint.pop(maintenance_endpoint) self.cluster.executor.reset_mock() @@ -617,16 +633,16 @@ def test_remove_matches_control_connection_by_host_id(self): def test_down_matches_replacement_at_stale_control_endpoint(self): self.control_connection.refresh_node_list_and_token_map() - old_host = self.cluster.metadata.get_host_by_host_id('uuid1') + old_host = self.cluster.metadata.get_host_by_host_id(HOST_ID_1) endpoint = old_host.endpoint - self.cluster.metadata.hosts.pop('uuid1') + self.cluster.metadata.hosts.pop(HOST_ID_1) replacement_host = Host( - endpoint, SimpleConvictionPolicy, host_id='replacement-id') + endpoint, SimpleConvictionPolicy, host_id=HOST_ID_4) replacement_host.set_up() - self.cluster.metadata.hosts['replacement-id'] = replacement_host + self.cluster.metadata.hosts[HOST_ID_4] = replacement_host self.cluster.metadata._host_id_by_endpoint[endpoint] = \ - 'replacement-id' + HOST_ID_4 connection_error = ConnectionException('old control failed') self.connection.is_defunct = True @@ -647,14 +663,14 @@ def test_down_matches_replacement_at_stale_control_endpoint(self): def test_refresh_unix_local_preserves_known_network_endpoint(self): maintenance_endpoint = UnixSocketEndPoint('/tmp/maintenance.sock') - local_host = self.cluster.metadata.get_host_by_host_id('uuid1') + local_host = self.cluster.metadata.get_host_by_host_id(HOST_ID_1) host_index = {local_host: object()} self.connection.endpoint = maintenance_endpoint self.connection.original_endpoint = maintenance_endpoint self.control_connection.refresh_node_list_and_token_map() - assert self.cluster.metadata.get_host_by_host_id('uuid1') is local_host + assert self.cluster.metadata.get_host_by_host_id(HOST_ID_1) is local_host assert local_host.endpoint == DefaultEndPoint('192.168.1.0') assert host_index[local_host] is not None @@ -669,12 +685,14 @@ def refresh_and_validate_added_hosts(): del self.connection.peer_results[:] self.connection.peer_results.extend([ ["rpc_address", "peer", "schema_version", "data_center", "rack", "tokens", "host_id"], - [["192.168.1.3", "10.0.0.1", "a", "dc1", "rack1", ["1", "101", "201"], 'uuid6'], + [["192.168.1.3", "10.0.0.1", "a", "dc1", "rack1", ["1", "101", "201"], HOST_ID_6], # all others are invalid - [None, None, "a", "dc1", "rack1", ["1", "101", "201"], 'uuid1'], - ["192.168.1.7", "10.0.0.1", "a", None, "rack1", ["1", "101", "201"], 'uuid2'], - ["192.168.1.6", "10.0.0.1", "a", "dc1", None, ["1", "101", "201"], 'uuid3'], - ["192.168.1.5", "10.0.0.1", "a", "dc1", "rack1", None, 'uuid4'], + [None, None, "a", "dc1", "rack1", ["1", "101", "201"], HOST_ID_1], + ["192.168.1.7", "10.0.0.1", "a", None, "rack1", ["1", "101", "201"], HOST_ID_2], + ["192.168.1.6", "10.0.0.1", "a", "dc1", None, ["1", "101", "201"], HOST_ID_3], + ["192.168.1.5", "10.0.0.1", "a", "dc1", "rack1", None, HOST_ID_4], + ["192.168.1.8", "10.0.0.1", "a", "dc1", "rack1", ["1", "101", "201"], "not-a-uuid"], + ["192.168.1.9", "10.0.0.1", "a", "dc1", "rack1", ["1", "101", "201"], uuid.UUID(int=0)], ["192.168.1.4", "10.0.0.1", "a", "dc1", "rack1", ["1", "101", "201"], None]]]) refresh_and_validate_added_hosts() @@ -683,12 +701,12 @@ def refresh_and_validate_added_hosts(): del self.connection.peer_results[:] self.connection.peer_results.extend([ ["native_address", "native_port", "peer", "peer_port", "schema_version", "data_center", "rack", "tokens", "host_id"], - [["192.168.1.4", 9042, "10.0.0.1", 7042, "a", "dc1", "rack1", ["1", "101", "201"], "uuid7"], + [["192.168.1.4", 9042, "10.0.0.1", 7042, "a", "dc1", "rack1", ["1", "101", "201"], HOST_ID_7], # all others are invalid - [None, 9042, None, 7040, "a", "dc1", "rack1", ["2", "102", "202"], "uuid2"], - ["192.168.1.5", 9042, "10.0.0.2", 7040, "a", None, "rack1", ["2", "102", "202"], "uuid2"], - ["192.168.1.5", 9042, "10.0.0.2", 7040, "a", "dc1", None, ["2", "102", "202"], "uuid2"], - ["192.168.1.5", 9042, "10.0.0.2", 7040, "a", "dc1", "rack1", None, "uuid2"], + [None, 9042, None, 7040, "a", "dc1", "rack1", ["2", "102", "202"], HOST_ID_2], + ["192.168.1.5", 9042, "10.0.0.2", 7040, "a", None, "rack1", ["2", "102", "202"], HOST_ID_2], + ["192.168.1.5", 9042, "10.0.0.2", 7040, "a", "dc1", None, ["2", "102", "202"], HOST_ID_2], + ["192.168.1.5", 9042, "10.0.0.2", 7040, "a", "dc1", "rack1", None, HOST_ID_2], ["192.168.1.5", 9042, "10.0.0.2", 7040, "a", "dc1", "rack1", ["2", "102", "202"], None]]]) refresh_and_validate_added_hosts() @@ -703,8 +721,8 @@ def test_change_ip(self): self.connection.peer_results.extend([ ["rpc_address", "peer", "schema_version", "data_center", "rack", "tokens", "host_id"], - [["192.168.1.5", "10.0.0.5", "a", "dc1", "rack1", ["2", "102", "202"], 'uuid2'], - ["192.168.1.6", "10.0.0.6", "a", "dc1", "rack1", ["3", "103", "203"], 'uuid3']]]) + [["192.168.1.5", "10.0.0.5", "a", "dc1", "rack1", ["2", "102", "202"], HOST_ID_2], + ["192.168.1.6", "10.0.0.6", "a", "dc1", "rack1", ["3", "103", "203"], HOST_ID_3]]]) self.connection.wait_for_responses = Mock( return_value=_node_meta_results( self.connection.local_results, self.connection.peer_results)) @@ -717,6 +735,229 @@ def test_change_ip(self): assert 3 == len(self.cluster.metadata.all_hosts()) + def test_same_endpoint_with_new_host_id_removes_old_session_pool(self): + cluster = Cluster() + self.addCleanup(cluster.shutdown) + cluster.control_connection.shutdown() + + hosts = [] + for host_id, address in ( + (HOST_ID_1, "192.168.1.0"), + (HOST_ID_2, "192.168.1.1"), + (HOST_ID_3, "192.168.1.2")): + host, _ = cluster.add_host( + DefaultEndPoint(address), datacenter="dc1", rack="rack1", + signal=False, host_id=host_id) + host.set_up() + hosts.append(host) + + old_host = hosts[1] + old_pool = Mock(host=old_host, is_shutdown=False) + retained_pools = { + host: Mock( + host=host, is_shutdown=False, + host_distance=HostDistance.LOCAL) + for host in (hosts[0], hosts[2]) + } + removal_future = Future() + addition_future = Future() + session = Session.__new__(Session) + session.cluster = cluster + session._pools = dict(retained_pools) + session._pools[old_host] = old_pool + session.is_shutdown = False + + def submit(fn, *args, **kwargs): + fn(*args, **kwargs) + return removal_future + + session.submit = submit + session._profile_manager = Mock() + session._profile_manager.distance.return_value = HostDistance.LOCAL + session.add_or_renew_pool = Mock(return_value=addition_future) + session.shutdown = Mock() + cluster.sessions.add(session) + + connection = MockConnection() + connection.peer_results[1][0][-1] = HOST_ID_4 + connection.wait_for_responses = Mock( + return_value=_node_meta_results( + connection.local_results, connection.peer_results)) + control_connection = ControlConnection(cluster, 1, 2, 0, 0) + control_connection._connection = connection + cluster.control_connection = control_connection + control_connection.refresh_node_list_and_token_map() + + assert connection.wait_for_responses.call_count == 1 + assert old_host not in session._pools + old_pool.shutdown.assert_called_once_with() + assert cluster.metadata.get_host_by_host_id(HOST_ID_2) is None + replacement = cluster.metadata.get_host_by_host_id(HOST_ID_4) + assert replacement is not None + assert replacement.endpoint == old_host.endpoint + session.add_or_renew_pool.assert_called_once_with( + replacement, is_host_addition=True, + on_add_reconnection=ANY) + + removal_future.set_result(None) + session.add_or_renew_pool.assert_called_once_with( + replacement, is_host_addition=True, + on_add_reconnection=ANY) + + replacement_pool = Mock( + host=replacement, is_shutdown=False, + host_distance=HostDistance.LOCAL) + session._pools[replacement] = replacement_pool + addition_future.set_result(True) + + assert session._pools[replacement] is replacement_pool + + def test_failed_same_endpoint_replacement_reconciles_surviving_pools(self): + class ReplacementFirstDCAwareRoundRobinPolicy( + DCAwareRoundRobinPolicy): + def on_add(self, host): + super().on_add(host) + dc = self._dc(host) + with self._hosts_lock: + current_hosts = self._dc_live_hosts[dc] + self._dc_live_hosts[dc] = (host,) + tuple( + current_host for current_host in current_hosts + if current_host != host) + + cluster = Cluster( + load_balancing_policy=ReplacementFirstDCAwareRoundRobinPolicy( + local_dc="dc1", used_hosts_per_remote_dc=1)) + self.addCleanup(cluster.shutdown) + cluster.control_connection.shutdown() + + hosts = [] + for host_id, address, datacenter in ( + (HOST_ID_1, "192.168.1.0", "dc1"), + (HOST_ID_2, "192.168.1.1", "dc2"), + (HOST_ID_3, "192.168.1.2", "dc2")): + host, _ = cluster.add_host( + DefaultEndPoint(address), datacenter=datacenter, rack="rack1", + signal=False, host_id=host_id) + host.set_up() + hosts.append(host) + + local_host, old_host, promoted_host = hosts + # The custom policy models a policy whose newest host preempts an + # existing eligible host. Seed the old host last so it starts as the + # only eligible remote host. + for host in (local_host, promoted_host, old_host): + cluster.profile_manager.on_add(host) + assert cluster.profile_manager.distance(old_host) == HostDistance.REMOTE + assert cluster.profile_manager.distance(promoted_host) == HostDistance.IGNORED + + old_pool = Mock( + host=old_host, is_shutdown=False, + host_distance=HostDistance.REMOTE) + local_pool = Mock( + host=local_host, is_shutdown=False, + host_distance=HostDistance.LOCAL) + removal_future = Future() + replacement_future = Future() + promoted_future = Future() + session = Session.__new__(Session) + session.cluster = cluster + session._pools = {local_host: local_pool, old_host: old_pool} + session.is_shutdown = False + + def submit(fn, *args, **kwargs): + fn(*args, **kwargs) + return removal_future + + session.submit = submit + session._profile_manager = cluster.profile_manager + def add_or_renew_pool(host, is_host_addition, + on_add_reconnection=None): + if is_host_addition: + assert on_add_reconnection is not None + return replacement_future + assert host is promoted_host + assert cluster.profile_manager.distance(host) == HostDistance.REMOTE + return promoted_future + + session.add_or_renew_pool = Mock(side_effect=add_or_renew_pool) + session.shutdown = Mock() + cluster.sessions.add(session) + + listener = Mock() + cluster.register_listener(listener) + + connection = MockConnection() + connection.peer_results[1][0][3] = "dc2" + connection.peer_results[1][0][-1] = HOST_ID_4 + connection.peer_results[1][1][3] = "dc2" + connection.wait_for_responses = Mock( + return_value=_node_meta_results( + connection.local_results, connection.peer_results)) + control_connection = ControlConnection(cluster, 1, 2, 0, 0) + control_connection._connection = connection + cluster.control_connection = control_connection + + control_connection.refresh_node_list_and_token_map() + + replacement = cluster.metadata.get_host_by_host_id(HOST_ID_4) + assert replacement is not None + assert cluster.profile_manager.distance(replacement) == HostDistance.REMOTE + assert cluster.profile_manager.distance(promoted_host) == HostDistance.IGNORED + session.add_or_renew_pool.assert_called_once_with( + replacement, is_host_addition=True, + on_add_reconnection=ANY) + + # Completing removal must not race the replacement with another pool + # creation attempt. + removal_future.set_result(None) + session.add_or_renew_pool.assert_called_once_with( + replacement, is_host_addition=True, + on_add_reconnection=ANY) + + replacement_future.set_result(False) + + assert cluster.profile_manager.distance(replacement) == HostDistance.IGNORED + assert cluster.profile_manager.distance(promoted_host) == HostDistance.REMOTE + assert replacement.is_currently_reconnecting() + assert session.add_or_renew_pool.call_args_list == [ + call(replacement, is_host_addition=True, + on_add_reconnection=ANY), + call(promoted_host, False), + ] + listener.on_add.assert_not_called() + + def test_same_control_endpoint_with_new_host_id_does_not_reconnect(self): + cluster = Cluster() + self.addCleanup(cluster.shutdown) + cluster.control_connection.shutdown() + + old_host, _ = cluster.add_host( + DefaultEndPoint("192.168.1.0"), datacenter="dc1", rack="rack1", + signal=False, host_id=HOST_ID_1) + + published_connection = MockConnection() + candidate_connection = MockConnection() + host_id_index = candidate_connection.local_results[0].index('host_id') + candidate_connection.local_results[1][0][host_id_index] = HOST_ID_4 + candidate_connection.wait_for_responses = Mock( + return_value=_node_meta_results( + candidate_connection.local_results, + candidate_connection.peer_results)) + + control_connection = ControlConnection(cluster, 1, 2, 0, 0) + control_connection._connection = published_connection + control_connection.reconnect = Mock() + cluster.control_connection = control_connection + control_connection._refresh_node_list_and_token_map( + candidate_connection) + + control_connection.reconnect.assert_not_called() + assert candidate_connection.wait_for_responses.call_count == 1 + assert control_connection._connection is published_connection + assert cluster.metadata.get_host_by_host_id(HOST_ID_1) is None + replacement = cluster.metadata.get_host_by_host_id(HOST_ID_4) + assert replacement is not None + assert replacement.endpoint == old_host.endpoint def test_refresh_nodes_and_tokens_uses_preloaded_results_if_given(self): """ @@ -754,7 +995,7 @@ def test_refresh_nodes_and_tokens_no_partitioner(self): def test_refresh_nodes_and_tokens_add_host(self): self.connection.peer_results[1].append( - ["192.168.1.3", "10.0.0.3", "a", "dc1", "rack1", ["3", "103", "203"], "uuid4"] + ["192.168.1.3", "10.0.0.3", "a", "dc1", "rack1", ["3", "103", "203"], HOST_ID_4] ) self.cluster.scheduler.schedule = lambda delay, f, *args, **kwargs: f(*args, **kwargs) self.control_connection.refresh_node_list_and_token_map() @@ -762,7 +1003,7 @@ def test_refresh_nodes_and_tokens_add_host(self): assert self.cluster.added_hosts[0].address == "192.168.1.3" assert self.cluster.added_hosts[0].datacenter == "dc1" assert self.cluster.added_hosts[0].rack == "rack1" - assert self.cluster.added_hosts[0].host_id == "uuid4" + assert self.cluster.added_hosts[0].host_id == HOST_ID_4 def test_refresh_nodes_and_tokens_remove_host(self): del self.connection.peer_results[1][1] @@ -928,7 +1169,7 @@ def test_refresh_nodes_and_tokens_add_host_detects_port(self): del self.connection.peer_results[:] self.connection.peer_results.extend(self.connection.peer_results_v2) self.connection.peer_results[1].append( - ["192.168.1.3", 555, "10.0.0.3", 666, "a", "dc1", "rack1", ["3", "103", "203"], "uuid4"] + ["192.168.1.3", 555, "10.0.0.3", 666, "a", "dc1", "rack1", ["3", "103", "203"], HOST_ID_4] ) self.connection.wait_for_responses = Mock(return_value=_node_meta_results( self.connection.local_results, self.connection.peer_results)) @@ -948,7 +1189,7 @@ def test_refresh_nodes_and_tokens_add_host_detects_invalid_port(self): del self.connection.peer_results[:] self.connection.peer_results.extend(self.connection.peer_results_v2) self.connection.peer_results[1].append( - ["192.168.1.3", -1, "10.0.0.3", 0, "a", "dc1", "rack1", ["3", "103", "203"], "uuid4"] + ["192.168.1.3", -1, "10.0.0.3", 0, "a", "dc1", "rack1", ["3", "103", "203"], HOST_ID_4] ) self.connection.wait_for_responses = Mock(return_value=_node_meta_results( self.connection.local_results, self.connection.peer_results)) diff --git a/tests/unit/test_response_future.py b/tests/unit/test_response_future.py index d71943ec04..a4290494b5 100644 --- a/tests/unit/test_response_future.py +++ b/tests/unit/test_response_future.py @@ -13,6 +13,7 @@ # limitations under the License. import unittest +import uuid from collections import deque from threading import RLock @@ -29,8 +30,8 @@ RESULT_KIND_ROWS, RESULT_KIND_SET_KEYSPACE, RESULT_KIND_SCHEMA_CHANGE, RESULT_KIND_PREPARED, ProtocolHandler) -from cassandra.policies import RetryPolicy, ExponentialBackoffRetryPolicy -from cassandra.pool import NoConnectionsAvailable +from cassandra.policies import RetryPolicy, ExponentialBackoffRetryPolicy, SimpleConvictionPolicy +from cassandra.pool import Host, NoConnectionsAvailable from cassandra.query import SimpleStatement, PreparedStatement, BoundStatement from tests.util import assertEqual, assertIsInstance import pytest @@ -633,9 +634,13 @@ def test_control_connection_fallback_orphans_stream_on_timeout(self): session = self.make_basic_session() session.cluster.allow_control_connection_query_fallback = ControlConnectionQueryFallback.Fallback session.cluster._default_load_balancing_policy.make_query_plan.return_value = ['ip1'] - session._pools = {} connection = self.make_control_connection() session.cluster.control_connection._connection = connection + control_host = Host( + connection.endpoint, SimpleConvictionPolicy, host_id=uuid.uuid4()) + session.cluster.get_control_connection_host.return_value = control_host + data_pool = Mock(is_shutdown=False) + session._pools = {} def send_msg(message, request_id, cb, **kwargs): connection._requests[request_id] = (cb, kwargs.get('decoder'), kwargs.get('result_metadata')) @@ -645,10 +650,13 @@ def send_msg(message, request_id, cb, **kwargs): rf = self.make_response_future(session) rf.send_request() + # Model a node pool appearing while the control request is in flight. + session._pools = {control_host: data_pool} rf._on_timeout() assert 7 in connection.orphaned_request_ids assert connection.in_flight == 1 + data_pool.return_connection.assert_not_called() with pytest.raises(OperationTimedOut): rf.result() @@ -1016,7 +1024,8 @@ def test_timeout_does_not_release_stream_id(self): pool = self.make_pool() session._pools.get.return_value = pool connection = Mock(spec=Connection, lock=RLock(), _requests={}, request_ids=deque(), - orphaned_request_ids=set(), orphaned_threshold=256, in_flight=3) + orphaned_request_ids=set(), orphaned_threshold=256, in_flight=3, + is_control_connection=False) pool.borrow_connection.return_value = (connection, 1) rf = self.make_response_future(session) From 8c274602f9977091cdd7276a2e58f0b85157d02a Mon Sep 17 00:00:00 2001 From: Dmitry Kropachev Date: Mon, 14 Sep 2026 10:45:45 -0400 Subject: [PATCH 3/5] host: fix replacement retry recovery --- cassandra/cluster.py | 7 +++++-- tests/unit/test_cluster.py | 29 +++++++++++++++++++++++++---- 2 files changed, 30 insertions(+), 6 deletions(-) diff --git a/cassandra/cluster.py b/cassandra/cluster.py index 900c4c2ee3..8c965de6cc 100644 --- a/cassandra/cluster.py +++ b/cassandra/cluster.py @@ -2101,8 +2101,11 @@ def on_add(self, host, refresh_nodes=True, futures = set() on_add_reconnection = None if reconcile_pools_on_failure: + # The initial addition may be running inside a topology refresh, + # but a later reconnection is outside that refresh and must fetch + # current node metadata. on_add_reconnection = partial( - self.on_add, refresh_nodes=refresh_nodes, + self.on_add, refresh_nodes=True, reconcile_pools_on_failure=True) def future_completed(future): @@ -3589,7 +3592,7 @@ def update_created_pools(self, excluded_host=None, hosts=None): futures = set() for host in hosts: - if excluded_host is not None and host == excluded_host: + if excluded_host is not None and host is excluded_host: continue distance = self._profile_manager.distance(host) pool = self._pools.get(host) diff --git a/tests/unit/test_cluster.py b/tests/unit/test_cluster.py index c591491503..ecf2c7e648 100644 --- a/tests/unit/test_cluster.py +++ b/tests/unit/test_cluster.py @@ -372,7 +372,7 @@ def make_session(pool_future): on_add_reconnection=ANY, start=False) listener.on_add.assert_not_called() - def test_replacement_reconnector_preserves_recovery_context(self): + def test_replacement_reconnector_refreshes_nodes(self): cluster = Cluster() self.addCleanup(cluster.shutdown) cluster.profile_manager = Mock() @@ -460,9 +460,9 @@ def create_pool(pool_host, distance, session): assert failing_session.update_created_pools.call_args_list == \ expected_reconciliation assert cluster.control_connection.on_add.call_args_list == \ - [call(host, False)] * 3 - cluster.control_connection.refresh_node_list_and_token_map \ - .assert_not_called() + [call(host, False), call(host, True), call(host, True)] + assert cluster.control_connection.refresh_node_list_and_token_map \ + .call_args_list == [call(force_token_rebuild=True)] * 2 listener.on_add.assert_not_called() def test_unconvicted_replacement_retry_keeps_reconnecting(self): @@ -1337,6 +1337,27 @@ def test_remove_pool_after_host_endpoint_changes(self): assert session._pools == {} session.cluster.executor.submit.assert_called_once_with(pool.shutdown) + def test_update_created_pools_excludes_only_same_host_object(self): + session = Session.__new__(Session) + host_id = uuid.uuid4() + excluded_host = Host( + "127.0.0.1", SimpleConvictionPolicy, host_id=host_id) + replacement_host = Host( + "127.0.0.2", SimpleConvictionPolicy, host_id=host_id) + replacement_host.set_up() + session._pools = {} + session.cluster = Mock( + allow_control_connection_query_fallback=ControlConnectionQueryFallback.Disabled) + session._profile_manager = Mock() + session._profile_manager.distance.return_value = HostDistance.LOCAL + session.add_or_renew_pool = Mock(return_value=None) + + assert session.update_created_pools( + excluded_host=excluded_host, hosts=(replacement_host,)) == set() + + session.add_or_renew_pool.assert_called_once_with( + replacement_host, False) + def test_pool_renewal_uses_pool_host_not_retained_dict_key(self): session = Session.__new__(Session) host_id = uuid.uuid4() From f4db94706999b9d28f680a56bc16dddeb8b34e9f Mon Sep 17 00:00:00 2001 From: Dmitry Kropachev Date: Mon, 14 Sep 2026 12:15:21 -0400 Subject: [PATCH 4/5] host: fence reconnect after authentication failure --- cassandra/cluster.py | 11 +++-- cassandra/pool.py | 4 +- tests/unit/test_cluster.py | 60 ++++++++++++++++++++++++- tests/unit/test_host_connection_pool.py | 10 +++++ 4 files changed, 80 insertions(+), 5 deletions(-) diff --git a/cassandra/cluster.py b/cassandra/cluster.py index 8c965de6cc..260dba524c 100644 --- a/cassandra/cluster.py +++ b/cassandra/cluster.py @@ -2005,7 +2005,7 @@ def _start_reconnector(self, host, is_host_addition, # on_remove() will clear and cancel this handler; if removal wins, do # not leave retry work attached to a terminal Host object. with host.lock: - if host._is_removed: + if host._is_removed or host._reconnection_disabled: return reconnector = _HostReconnectionHandler( host, conn_factory, is_host_addition, on_add, self.on_up, @@ -2155,8 +2155,13 @@ def future_completed(future): result is None for result in futures_results) pending_reconnector = None if authentication_failed: - old_reconnector = \ - host.get_and_set_reconnection_handler(None) + # Authentication is terminal for this Host lifecycle. + # Fence DOWN work already queued by another failed pool + # before detaching any handler it installed. + with host.lock: + host._reconnection_disabled = True + old_reconnector = host._reconnection_handler + host._reconnection_handler = None if old_reconnector: old_reconnector.cancel() else: diff --git a/cassandra/pool.py b/cassandra/pool.py index bf87b714c1..6ce9eeb30f 100644 --- a/cassandra/pool.py +++ b/cassandra/pool.py @@ -155,6 +155,7 @@ class Host(object): _datacenter = None _rack = None _reconnection_handler = None + _reconnection_disabled = False lock = None _currently_handling_node_up = False @@ -175,6 +176,7 @@ def __init__(self, endpoint, conviction_policy_factory, datacenter=None, rack=No self.endpoint = endpoint if isinstance(endpoint, EndPoint) else DefaultEndPoint(endpoint) self._host_id = host_id self._is_removed = False + self._reconnection_disabled = False self.conviction_policy = conviction_policy_factory(self) self.set_location_info(datacenter, rack) self.lock = RLock() @@ -243,7 +245,7 @@ def _clear_reconnection_handler(self, handler): if self._reconnection_handler is not handler: return False self._reconnection_handler = None - return not self._is_removed + return not self._is_removed and not self._reconnection_disabled def __eq__(self, other): if not isinstance(other, Host): diff --git a/tests/unit/test_cluster.py b/tests/unit/test_cluster.py index ecf2c7e648..377c449d78 100644 --- a/tests/unit/test_cluster.py +++ b/tests/unit/test_cluster.py @@ -16,7 +16,7 @@ from concurrent.futures import Future import logging import socket -from threading import RLock +from threading import Event, RLock from types import SimpleNamespace from unittest.mock import ANY, Mock, call, patch @@ -372,6 +372,64 @@ def make_session(pool_future): on_add_reconnection=ANY, start=False) listener.on_add.assert_not_called() + def test_authentication_failure_fences_queued_replacement_down(self): + cluster = Cluster(executor_threads=1) + self.addCleanup(cluster.shutdown) + cluster.profile_manager = Mock() + cluster.profile_manager.distance.return_value = HostDistance.LOCAL + cluster.control_connection.on_add = Mock() + cluster.control_connection.on_down = Mock() + cluster._prepare_all_queries = Mock() + cluster._make_connection_factory = Mock(return_value=Mock()) + cluster.scheduler.schedule = Mock() + + worker_started = Event() + release_worker = Event() + self.addCleanup(release_worker.set) + + def block_worker(): + worker_started.set() + release_worker.wait() + + blocker = cluster.executor.submit(block_worker) + assert worker_started.wait(5) + + host = Host( + "127.0.0.1", SimpleConvictionPolicy, host_id=uuid.uuid4()) + cluster.metadata.add_or_return_host(host) + + retryable_future = Future() + authentication_future = Future() + retryable_session = Mock() + retryable_session.add_or_renew_pool.return_value = retryable_future + retryable_session.get_pool_state.return_value = {} + authentication_session = Mock() + authentication_session.add_or_renew_pool.return_value = \ + authentication_future + authentication_session.get_pool_state.return_value = {} + cluster.sessions = (retryable_session, authentication_session) + + cluster.on_add( + host, refresh_nodes=False, reconcile_pools_on_failure=True) + on_add_reconnection = retryable_session.add_or_renew_pool \ + .call_args.kwargs["on_add_reconnection"] + + # Model retryable pool failure queuing DOWN work while the sole worker + # is busy, followed by a non-retryable authentication result. + cluster.on_down( + host, is_host_addition=True, expect_host_to_be_down=True, + on_add_reconnection=on_add_reconnection) + retryable_future.set_result(False) + authentication_future.set_result(None) + + release_worker.set() + blocker.result(timeout=5) + cluster.executor.submit(lambda: None).result(timeout=5) + + assert host._reconnection_disabled + assert not host.is_currently_reconnecting() + cluster.scheduler.schedule.assert_not_called() + def test_replacement_reconnector_refreshes_nodes(self): cluster = Cluster() self.addCleanup(cluster.shutdown) diff --git a/tests/unit/test_host_connection_pool.py b/tests/unit/test_host_connection_pool.py index ca1cb51573..6e8b258aef 100644 --- a/tests/unit/test_host_connection_pool.py +++ b/tests/unit/test_host_connection_pool.py @@ -246,6 +246,16 @@ def test_host_id_is_read_only(self): assert host.host_id == host_id + def test_disabled_reconnection_handler_cannot_resume_recovery(self): + host = Host( + '127.0.0.1', SimpleConvictionPolicy, host_id=uuid.uuid4()) + handler = Mock() + host.get_and_set_reconnection_handler(handler) + host._reconnection_disabled = True + + assert not host._clear_reconnection_handler(handler) + assert not host.is_currently_reconnecting() + def test_host_hash_is_stable_when_endpoint_changes(self): host = Host('127.0.0.1', SimpleConvictionPolicy, host_id=uuid.uuid4()) hosts_by_id = {host: "pool"} From 72adbefb1c35ea818db70205ebefc1c114f6b041 Mon Sep 17 00:00:00 2001 From: Dmitry Kropachev Date: Mon, 14 Sep 2026 13:16:38 -0400 Subject: [PATCH 5/5] host: restore ignored replacement eligibility --- cassandra/cluster.py | 8 ++++-- tests/unit/test_cluster.py | 56 ++++++++++++++++++++++++++++++++++++++ 2 files changed, 61 insertions(+), 3 deletions(-) diff --git a/cassandra/cluster.py b/cassandra/cluster.py index 260dba524c..26d270ddf0 100644 --- a/cassandra/cluster.py +++ b/cassandra/cluster.py @@ -2076,7 +2076,8 @@ def on_down(self, host, is_host_addition, expect_host_to_be_down=False, on_add_reconnection=on_add_reconnection) def on_add(self, host, refresh_nodes=True, - reconcile_pools_on_failure=False): + reconcile_pools_on_failure=False, + set_up_if_ignored=False): if self.is_shutdown: return @@ -2093,7 +2094,7 @@ def on_add(self, host, refresh_nodes=True, if distance == HostDistance.IGNORED: log.debug("Not adding connection pool for new host %r because the " "load balancing policy has marked it as IGNORED", host) - self._finalize_add(host, set_up=False) + self._finalize_add(host, set_up=set_up_if_ignored) return futures_lock = Lock() @@ -2106,7 +2107,8 @@ def on_add(self, host, refresh_nodes=True, # current node metadata. on_add_reconnection = partial( self.on_add, refresh_nodes=True, - reconcile_pools_on_failure=True) + reconcile_pools_on_failure=True, + set_up_if_ignored=True) def future_completed(future): with futures_lock: diff --git a/tests/unit/test_cluster.py b/tests/unit/test_cluster.py index 377c449d78..f1c0058f7e 100644 --- a/tests/unit/test_cluster.py +++ b/tests/unit/test_cluster.py @@ -523,6 +523,62 @@ def create_pool(pool_host, distance, session): .call_args_list == [call(force_token_rebuild=True)] * 2 listener.on_add.assert_not_called() + def test_ignored_replacement_reconnect_restores_pool_eligibility(self): + cluster = Cluster() + self.addCleanup(cluster.shutdown) + distance = [HostDistance.LOCAL] + cluster.profile_manager = Mock() + cluster.profile_manager.distance.side_effect = lambda host: distance[0] + cluster.control_connection.on_add = Mock() + cluster._prepare_all_queries = Mock() + cluster.scheduler.schedule = Mock() + probe = Mock() + cluster._make_connection_factory = Mock(return_value=Mock( + return_value=probe)) + + host = Host( + "127.0.0.1", SimpleConvictionPolicy, host_id=uuid.uuid4()) + cluster.metadata.add_or_return_host(host) + + failed_future = Future() + eligible_future = Future() + session = Session.__new__(Session) + session.cluster = cluster + session._profile_manager = cluster.profile_manager + session._pools = {} + session.is_shutdown = False + session.add_or_renew_pool = Mock(return_value=failed_future) + session.shutdown = Mock() + cluster.sessions.add(session) + + cluster.on_add( + host, refresh_nodes=False, reconcile_pools_on_failure=True) + failed_future.set_result(False) + + assert host.is_up is False + assert host.is_currently_reconnecting() + + # Another remote host can consume the policy's eligible slot before + # the replacement probe succeeds. The probe still establishes that + # this host is reachable, even though it does not currently need a + # pool. + distance[0] = HostDistance.IGNORED + cluster.scheduler.schedule.call_args.args[1]() + + assert host.is_up is True + assert not host.is_currently_reconnecting() + probe.close.assert_called_once_with() + + # Once the slot opens, update_created_pools must be able to create a + # pool without waiting for an unrelated UP event. + session.add_or_renew_pool.reset_mock(return_value=True) + session.add_or_renew_pool.return_value = eligible_future + distance[0] = HostDistance.REMOTE + futures = session.update_created_pools() + + assert futures == {eligible_future} + session.add_or_renew_pool.assert_called_once_with(host, False) + def test_unconvicted_replacement_retry_keeps_reconnecting(self): cluster = Cluster() self.addCleanup(cluster.shutdown)