diff --git a/cassandra/cluster.py b/cassandra/cluster.py index 88c8d2707a..e959ef9cdb 100644 --- a/cassandra/cluster.py +++ b/cassandra/cluster.py @@ -5033,7 +5033,10 @@ def _query(self, host, message=None, cb=None): try: # TODO get connectTimeout from cluster settings if self.query: - connection, request_id = pool.borrow_connection(timeout=2.0, routing_key=self.query.routing_key, keyspace=self.query.keyspace, table=self.query.table) + connection, request_id = pool.borrow_connection( + timeout=2.0, routing_key=self.query.routing_key, + keyspace=self.query.keyspace, table=self.query.table, + tablet=getattr(self.query, '_tablet', None)) else: connection, request_id = pool.borrow_connection(timeout=2.0) self._connection = connection diff --git a/cassandra/policies.py b/cassandra/policies.py index 89702e8c89..bb32ea1b82 100644 --- a/cassandra/policies.py +++ b/cassandra/policies.py @@ -498,6 +498,11 @@ def make_query_plan(self, working_keyspace=None, query=None): child = self._child_policy if query is None or query.routing_key is None or keyspace is None: + if query is not None: + # A Statement (e.g. BoundStatement) can be rebound and + # re-executed by the caller; make sure a tablet stashed by + # an earlier, unrelated execution isn't picked up below. + query._tablet = None for host in child.make_query_plan(keyspace, query): yield host return @@ -507,12 +512,22 @@ def make_query_plan(self, working_keyspace=None, query=None): keyspace, query.table, self._cluster_metadata.token_map.token_class.from_key(query.routing_key)) if tablet is not None: - replicas_mapped = set(map(lambda r: r[0], tablet.replicas)) + replica_dict = tablet._replica_dict child_plan = child.make_query_plan(keyspace, query) - replicas = [host for host in child_plan if host.host_id in replicas_mapped] + replicas = [host for host in child_plan if host.host_id in replica_dict] + # Stash the tablet so that downstream shard-aware + # connection selection can reuse it instead of + # repeating the bisect lookup. + query._tablet = tablet else: replicas = self._cluster_metadata.get_replicas(keyspace, query.routing_key) + # Clear any tablet stashed by a previous execution of this same + # query object (statements may be rebound and reused, e.g. via + # BoundStatement.bind()) so a stale tablet -- for a different + # routing key -- isn't reused for shard-aware connection + # selection below. + query._tablet = None if self.shuffle_replicas and not query.is_lwt() and not ConsistencyLevel.is_serial(query.consistency_level): shuffle(replicas) diff --git a/cassandra/pool.py b/cassandra/pool.py index 176751f60a..87f709712a 100644 --- a/cassandra/pool.py +++ b/cassandra/pool.py @@ -440,7 +440,7 @@ def __init__(self, host, host_distance, session): log.debug("Finished initializing connection for host %s", self.host) - def _get_connection_for_routing_key(self, routing_key=None, keyspace=None, table=None): + def _get_connection_for_routing_key(self, routing_key=None, keyspace=None, table=None, tablet=None): if self.is_shutdown: raise ConnectionException( "Pool for %s is shutdown" % (self.host,), self.host) @@ -454,16 +454,18 @@ def _get_connection_for_routing_key(self, routing_key=None, keyspace=None, table shard_id = None if self.tablets_routing_v1 and table is not None: - if keyspace is None: - keyspace = self._keyspace + # Reuse tablet from query planning if available, avoiding + # a redundant bisect lookup in the tablet map. + if tablet is not None: + shard_id = tablet._replica_dict.get(self.host.host_id) + else: + if keyspace is None: + keyspace = self._keyspace - tablet = self._session.cluster.metadata._tablets.get_tablet_for_key(keyspace, table, t) + tablet = self._session.cluster.metadata._tablets.get_tablet_for_key(keyspace, table, t) - if tablet is not None: - for replica in tablet.replicas: - if replica[0] == self.host.host_id: - shard_id = replica[1] - break + if tablet is not None: + shard_id = tablet._replica_dict.get(self.host.host_id) if shard_id is None: shard_id = self.host.sharding_info.shard_id_from_token(t.value) @@ -506,15 +508,15 @@ def _get_connection_for_routing_key(self, routing_key=None, keyspace=None, table return random.choice(active_connections) return random.choice(list(self._connections.values())) - def borrow_connection(self, timeout, routing_key=None, keyspace=None, table=None): - conn = self._get_connection_for_routing_key(routing_key, keyspace, table) + def borrow_connection(self, timeout, routing_key=None, keyspace=None, table=None, tablet=None): + conn = self._get_connection_for_routing_key(routing_key, keyspace, table, tablet) start = time.time() remaining = timeout last_retry = False while True: if conn.is_closed: # The connection might have been closed in the meantime - if so, try again - conn = self._get_connection_for_routing_key(routing_key, keyspace, table) + conn = self._get_connection_for_routing_key(routing_key, keyspace, table, tablet) with conn.lock: if (not conn.is_closed or last_retry) and conn.in_flight < conn.max_request_id: # On last retry we ignore connection status, since it is better to return closed connection than diff --git a/cassandra/tablets.py b/cassandra/tablets.py index 96e61a50c2..eb7d2025b5 100644 --- a/cassandra/tablets.py +++ b/cassandra/tablets.py @@ -1,13 +1,8 @@ from bisect import bisect_left -from operator import attrgetter from threading import Lock from typing import Optional from uuid import UUID -# C-accelerated attrgetter avoids per-call lambda allocation overhead -_get_first_token = attrgetter("first_token") -_get_last_token = attrgetter("last_token") - class Tablet(object): """ @@ -15,63 +10,83 @@ class Tablet(object): It stores information about each replica, its host and shard, and the token interval in the format (first_token, last_token]. """ - first_token = 0 - last_token = 0 - replicas = None + __slots__ = ('first_token', 'last_token', 'replicas', '_replica_dict') def __init__(self, first_token=0, last_token=0, replicas=None): self.first_token = first_token self.last_token = last_token - self.replicas = replicas + if replicas is not None: + replicas_tuple = tuple(replicas) + self.replicas = replicas_tuple + self._replica_dict = {r[0]: r[1] for r in replicas_tuple} + else: + self.replicas = None + self._replica_dict = {} def __str__(self): return "" \ % (self.first_token, self.last_token, self.replicas) __repr__ = __str__ - @staticmethod - def _is_valid_tablet(replicas): - return replicas is not None and len(replicas) != 0 - @staticmethod def from_row(first_token, last_token, replicas): - if Tablet._is_valid_tablet(replicas): - tablet = Tablet(first_token, last_token, replicas) - return tablet - return None + # Materialize once: `replicas` may be a one-shot iterator (e.g. a + # generator), and a plain `if not replicas` truthiness check would + # always be False for such an object even when it yields nothing, + # since iterators have no __len__/__bool__ and are always truthy. + replicas_tuple = tuple(replicas) if replicas is not None else () + if not replicas_tuple: + return None + return Tablet(first_token, last_token, replicas_tuple) def replica_contains_host_id(self, uuid: UUID) -> bool: - for replica in self.replicas: - if replica[0] == uuid: - return True - return False + return uuid in self._replica_dict + def get_replica_shard_id(self, uuid: UUID) -> Optional[int]: + return self._replica_dict.get(uuid) -class Tablets(object): - _lock = None - _tablets = {} +class Tablets(object): def __init__(self, tablets): - self._tablets = tablets + # NOTE: these are intentionally instance attributes only (not class + # attributes) to avoid mutable class-level dicts being shared across + # instances, e.g. if a future alternative constructor were to bypass + # __init__. self._lock = Lock() + self._tablets = tablets + # Build parallel token index lists from any pre-populated data + # (keyspace, table) -> list[int] for both _first_tokens/_last_tokens + self._first_tokens = { + key: [t.first_token for t in tlist] + for key, tlist in tablets.items() + } + self._last_tokens = { + key: [t.last_token for t in tlist] + for key, tlist in tablets.items() + } def table_has_tablets(self, keyspace, table) -> bool: return bool(self._tablets.get((keyspace, table), [])) def get_tablet_for_key(self, keyspace, table, t): - tablet = self._tablets.get((keyspace, table), []) - if not tablet: + key = (keyspace, table) + last_tokens = self._last_tokens.get(key) + if not last_tokens: return None - id = bisect_left(tablet, t.value, key=_get_last_token) - if id < len(tablet) and t.value > tablet[id].first_token: - return tablet[id] + token_value = t.value + id = bisect_left(last_tokens, token_value) + if id < len(last_tokens) and token_value > self._first_tokens[key][id]: + return self._tablets[key][id] return None def drop_tablets(self, keyspace: str, table: Optional[str] = None): with self._lock: if table is not None: - self._tablets.pop((keyspace, table), None) + key = (keyspace, table) + self._tablets.pop(key, None) + self._first_tokens.pop(key, None) + self._last_tokens.pop(key, None) return to_be_deleted = [] @@ -81,36 +96,48 @@ def drop_tablets(self, keyspace: str, table: Optional[str] = None): for key in to_be_deleted: del self._tablets[key] + self._first_tokens.pop(key, None) + self._last_tokens.pop(key, None) def drop_tablets_by_host_id(self, host_id: Optional[UUID]): if host_id is None: return with self._lock: for key, tablets in self._tablets.items(): - to_be_deleted = [] - for tablet_id, tablet in enumerate(tablets): - if tablet.replica_contains_host_id(host_id): - to_be_deleted.append(tablet_id) - - for tablet_id in reversed(to_be_deleted): - tablets.pop(tablet_id) + # Filter in one pass instead of popping one-by-one (O(n) vs O(k*n)) + keep = [i for i, t in enumerate(tablets) + if not t.replica_contains_host_id(host_id)] + if len(keep) == len(tablets): + continue # nothing to drop + self._tablets[key] = [tablets[i] for i in keep] + first = self._first_tokens[key] + last = self._last_tokens[key] + self._first_tokens[key] = [first[i] for i in keep] + self._last_tokens[key] = [last[i] for i in keep] def add_tablet(self, keyspace, table, tablet): with self._lock: - tablets_for_table = self._tablets.setdefault((keyspace, table), []) + key = (keyspace, table) + tablets_for_table = self._tablets.setdefault(key, []) + first_tokens = self._first_tokens.setdefault(key, []) + last_tokens = self._last_tokens.setdefault(key, []) # find first overlapping range - start = bisect_left(tablets_for_table, tablet.first_token, key=_get_first_token) - if start > 0 and tablets_for_table[start - 1].last_token > tablet.first_token: + start = bisect_left(first_tokens, tablet.first_token) + if start > 0 and last_tokens[start - 1] > tablet.first_token: start = start - 1 # find last overlapping range - end = bisect_left(tablets_for_table, tablet.last_token, key=_get_last_token) - if end < len(tablets_for_table) and tablets_for_table[end].first_token >= tablet.last_token: + end = bisect_left(last_tokens, tablet.last_token) + if end < len(last_tokens) and first_tokens[end] >= tablet.last_token: end = end - 1 if start <= end: del tablets_for_table[start:end + 1] + del first_tokens[start:end + 1] + del last_tokens[start:end + 1] tablets_for_table.insert(start, tablet) + first_tokens.insert(start, tablet.first_token) + last_tokens.insert(start, tablet.last_token) diff --git a/tests/unit/test_policies.py b/tests/unit/test_policies.py index 63a3c3d12d..a3db9c890b 100644 --- a/tests/unit/test_policies.py +++ b/tests/unit/test_policies.py @@ -972,6 +972,64 @@ def test_no_shuffle_for_serial_consistency(self, patched_shuffle): assert patched_shuffle.call_count == 0, \ "shuffle should not be called for consistency level %s" % cl + def test_stale_tablet_not_reused_across_query_plans(self): + """ + A Statement (e.g. a BoundStatement) may be rebound and re-executed by + the caller, so the same query object can be passed to + make_query_plan() multiple times with a different routing key each + time. Verify that a tablet stashed on the query object for shard-aware + connection selection (query._tablet) from one call doesn't leak into + a later call for which no tablet is found -- otherwise downstream + shard selection could pick a shard belonging to an unrelated, + previously-looked-up tablet. + """ + cluster = self._prepare_cluster_with_tablets() + hosts = cluster.metadata.all_hosts() + tablet = cluster.metadata._tablets.get_tablet_for_key.return_value + + child_policy = Mock() + child_policy.make_query_plan.return_value = hosts + child_policy.distance.return_value = HostDistance.LOCAL + + policy = TokenAwarePolicy(child_policy, shuffle_replicas=False) + policy.populate(cluster, hosts) + + query = Statement(routing_key='routing_key', keyspace='keyspace') + list(policy.make_query_plan('keyspace', query)) + self.assertIs(query._tablet, tablet) + + # Same (reused) query object, but this time no tablet is found for + # the (new) routing key -- e.g. it hasn't been discovered yet, or + # the table isn't tablets-based. + cluster.metadata._tablets.get_tablet_for_key.return_value = None + list(policy.make_query_plan('keyspace', query)) + self.assertIsNone(query._tablet) + + def test_stale_tablet_not_reused_when_no_routing_key(self): + """ + Same as above, but covers the early-return path (no routing key / + no keyspace), which must also clear any previously stashed tablet. + """ + cluster = self._prepare_cluster_with_tablets() + hosts = cluster.metadata.all_hosts() + tablet = cluster.metadata._tablets.get_tablet_for_key.return_value + + child_policy = Mock() + child_policy.make_query_plan.return_value = hosts + child_policy.distance.return_value = HostDistance.LOCAL + + policy = TokenAwarePolicy(child_policy, shuffle_replicas=False) + policy.populate(cluster, hosts) + + query = Statement(routing_key='routing_key', keyspace='keyspace') + list(policy.make_query_plan('keyspace', query)) + self.assertIs(query._tablet, tablet) + + # Reuse the same statement without a routing key this time. + query.routing_key = None + list(policy.make_query_plan('keyspace', query)) + self.assertIsNone(query._tablet) + class ConvictionPolicyTest(unittest.TestCase): def test_not_implemented(self): diff --git a/tests/unit/test_response_future.py b/tests/unit/test_response_future.py index cf1194a91f..c08d7a8b11 100644 --- a/tests/unit/test_response_future.py +++ b/tests/unit/test_response_future.py @@ -94,7 +94,7 @@ def test_result_message(self): rf.send_request() rf.session._pools.get.assert_called_once_with('ip1') - pool.borrow_connection.assert_called_once_with(timeout=ANY, routing_key=ANY, keyspace=ANY, table=ANY) + pool.borrow_connection.assert_called_once_with(timeout=ANY, routing_key=ANY, keyspace=ANY, table=ANY, tablet=ANY) connection.send_msg.assert_called_once_with(rf.message, 1, cb=ANY, encoder=ProtocolHandler.encode_message, decoder=ProtocolHandler.decode_message, result_metadata=[]) @@ -285,7 +285,7 @@ def test_retry_policy_says_retry(self): rf.send_request() rf.session._pools.get.assert_called_once_with('ip1') - pool.borrow_connection.assert_called_once_with(timeout=ANY, routing_key=ANY, keyspace=ANY, table=ANY) + pool.borrow_connection.assert_called_once_with(timeout=ANY, routing_key=ANY, keyspace=ANY, table=ANY, tablet=ANY) connection.send_msg.assert_called_once_with(rf.message, 1, cb=ANY, encoder=ProtocolHandler.encode_message, decoder=ProtocolHandler.decode_message, result_metadata=[]) result = Mock(spec=UnavailableErrorMessage, info={}) @@ -304,7 +304,7 @@ def test_retry_policy_says_retry(self): # it should try again with the same host since this was # an UnavailableException rf.session._pools.get.assert_called_with(host) - pool.borrow_connection.assert_called_with(timeout=ANY, routing_key=ANY, keyspace=ANY, table=ANY) + pool.borrow_connection.assert_called_with(timeout=ANY, routing_key=ANY, keyspace=ANY, table=ANY, tablet=ANY) connection.send_msg.assert_called_with(rf.message, 2, cb=ANY, encoder=ProtocolHandler.encode_message, decoder=ProtocolHandler.decode_message, result_metadata=[]) def test_retry_with_different_host(self): @@ -319,7 +319,7 @@ def test_retry_with_different_host(self): rf.send_request() rf.session._pools.get.assert_called_once_with('ip1') - pool.borrow_connection.assert_called_once_with(timeout=ANY, routing_key=ANY, keyspace=ANY, table=ANY) + pool.borrow_connection.assert_called_once_with(timeout=ANY, routing_key=ANY, keyspace=ANY, table=ANY, tablet=ANY) connection.send_msg.assert_called_once_with(rf.message, 1, cb=ANY, encoder=ProtocolHandler.encode_message, decoder=ProtocolHandler.decode_message, result_metadata=[]) assert ConsistencyLevel.QUORUM == rf.message.consistency_level @@ -338,7 +338,7 @@ def test_retry_with_different_host(self): # it should try with a different host rf.session._pools.get.assert_called_with('ip2') - pool.borrow_connection.assert_called_with(timeout=ANY, routing_key=ANY, keyspace=ANY, table=ANY) + pool.borrow_connection.assert_called_with(timeout=ANY, routing_key=ANY, keyspace=ANY, table=ANY, tablet=ANY) connection.send_msg.assert_called_with(rf.message, 2, cb=ANY, encoder=ProtocolHandler.encode_message, decoder=ProtocolHandler.decode_message, result_metadata=[]) # the consistency level should be the same @@ -1055,7 +1055,7 @@ def test_single_host_query_plan_exhausted_after_one_retry(self): # Verify initial request was sent rf.session._pools.get.assert_called_once_with(specific_host) - pool.borrow_connection.assert_called_once_with(timeout=ANY, routing_key=ANY, keyspace=ANY, table=ANY) + pool.borrow_connection.assert_called_once_with(timeout=ANY, routing_key=ANY, keyspace=ANY, table=ANY, tablet=ANY) connection.send_msg.assert_called_once_with(rf.message, 1, cb=ANY, encoder=ProtocolHandler.encode_message, decoder=ProtocolHandler.decode_message, result_metadata=[]) # Simulate a ServerError response (which triggers RETRY_NEXT_HOST by default) diff --git a/tests/unit/test_tablets.py b/tests/unit/test_tablets.py index 7a40e7de4d..27496361ef 100644 --- a/tests/unit/test_tablets.py +++ b/tests/unit/test_tablets.py @@ -1,4 +1,5 @@ import unittest +from uuid import UUID from cassandra.tablets import Tablets, Tablet @@ -88,6 +89,31 @@ def test_add_tablet_intersecting_with_last(self): (-5011686018427387905, -2987529027641081857)]) +class TabletsInstanceStateTest(unittest.TestCase): + """Tests that Tablets' internal dicts are per-instance state, not + shared mutable class attributes (a well-known Python footgun).""" + + def test_internal_dicts_are_not_class_attributes(self): + self.assertNotIn('_tablets', vars(Tablets)) + self.assertNotIn('_first_tokens', vars(Tablets)) + self.assertNotIn('_last_tokens', vars(Tablets)) + + def test_instances_do_not_share_internal_dicts(self): + a = Tablets({}) + b = Tablets({}) + self.assertIsNot(a._tablets, b._tablets) + self.assertIsNot(a._first_tokens, b._first_tokens) + self.assertIsNot(a._last_tokens, b._last_tokens) + + t1 = Tablet(0, 100, [("host1", 0)]) + a.add_tablet("ks", "tb", t1) + # Mutating `a` must not be visible through `b`. + self.assertFalse(b.table_has_tablets("ks", "tb")) + self.assertEqual(b._tablets, {}) + self.assertEqual(b._first_tokens, {}) + self.assertEqual(b._last_tokens, {}) + + class GetTabletForKeyTest(unittest.TestCase): """Tests for Tablets.get_tablet_for_key.""" @@ -124,3 +150,147 @@ def __init__(self, v): # Token value 50 is not > first_token (100) of the tablet whose # last_token (200) is >= 50, so no match. self.assertIsNone(tablets.get_tablet_for_key("ks", "tb", Token(50))) + + +class TabletFromRowTest(unittest.TestCase): + """Tests for Tablet.from_row, in particular that emptiness is detected + correctly regardless of whether `replicas` is a reusable sequence or a + one-shot iterator/generator.""" + + def test_empty_list_returns_none(self): + self.assertIsNone(Tablet.from_row(0, 100, [])) + + def test_empty_generator_returns_none(self): + # A generator is always truthy, even when empty, so a naive + # `if not replicas` check would fail to detect this case. + self.assertIsNone(Tablet.from_row(0, 100, (x for x in []))) + + def test_none_returns_none(self): + self.assertIsNone(Tablet.from_row(0, 100, None)) + + def test_non_empty_list_builds_tablet(self): + u1 = UUID('12345678-1234-5678-1234-567812345678') + u2 = UUID('87654321-4321-8765-4321-876543218765') + tablet = Tablet.from_row(0, 100, [(u1, 3), (u2, 7)]) + self.assertIsNotNone(tablet) + self.assertEqual(tablet.replicas, ((u1, 3), (u2, 7))) + self.assertTrue(tablet.replica_contains_host_id(u1)) + self.assertEqual(tablet.get_replica_shard_id(u2), 7) + + def test_non_empty_generator_builds_tablet(self): + # Generators are single-use: confirm the fix materializes the + # replicas exactly once and doesn't lose data by iterating twice. + u1 = UUID('12345678-1234-5678-1234-567812345678') + u2 = UUID('87654321-4321-8765-4321-876543218765') + + def gen(): + yield (u1, 3) + yield (u2, 7) + + tablet = Tablet.from_row(0, 100, gen()) + self.assertIsNotNone(tablet) + self.assertEqual(tablet.replicas, ((u1, 3), (u2, 7))) + self.assertTrue(tablet.replica_contains_host_id(u1)) + self.assertTrue(tablet.replica_contains_host_id(u2)) + self.assertEqual(tablet.get_replica_shard_id(u1), 3) + self.assertEqual(tablet.get_replica_shard_id(u2), 7) + + +class TabletReplicaDictTest(unittest.TestCase): + """Tests for Tablet's replica/shard lookup behavior, backed internally + by a cached _replica_dict for O(1) host/shard lookup. + + Most of these tests go through the public API (replica_contains_host_id + and get_replica_shard_id) so they keep working across internal + refactors of the cache; see test_replica_dict_populated_as_expected + for the one targeted check of the internal structure itself. + """ + + def test_replica_contains_host_id(self): + u1 = UUID('12345678-1234-5678-1234-567812345678') + u2 = UUID('87654321-4321-8765-4321-876543218765') + u3 = UUID('aaaaaaaa-bbbb-cccc-dddd-eeeeeeeeeeee') + t = Tablet(0, 100, [(u1, 3), (u2, 7)]) + self.assertTrue(t.replica_contains_host_id(u1)) + self.assertTrue(t.replica_contains_host_id(u2)) + self.assertFalse(t.replica_contains_host_id(u3)) + + def test_replica_contains_host_id_false_when_no_replicas(self): + u1 = UUID('12345678-1234-5678-1234-567812345678') + t = Tablet(0, 100, None) + self.assertFalse(t.replica_contains_host_id(u1)) + + def test_get_replica_shard_id(self): + u1 = UUID('12345678-1234-5678-1234-567812345678') + u2 = UUID('87654321-4321-8765-4321-876543218765') + u3 = UUID('aaaaaaaa-bbbb-cccc-dddd-eeeeeeeeeeee') + t = Tablet(0, 100, [(u1, 3), (u2, 7)]) + self.assertEqual(t.get_replica_shard_id(u1), 3) + self.assertEqual(t.get_replica_shard_id(u2), 7) + self.assertIsNone(t.get_replica_shard_id(u3)) + + def test_replicas_stored_as_tuple(self): + t = Tablet(0, 100, [("host1", 0), ("host2", 1)]) + self.assertIsInstance(t.replicas, tuple) + + def test_replica_lookup_from_iterator(self): + """Ensure replica lookups work correctly even when replicas is a + one-shot iterator (generator), not a reusable list.""" + u1 = UUID('12345678-1234-5678-1234-567812345678') + u2 = UUID('87654321-4321-8765-4321-876543218765') + + def gen(): + yield (u1, 3) + yield (u2, 7) + + t = Tablet(0, 100, gen()) + self.assertEqual(t.replicas, ((u1, 3), (u2, 7))) + self.assertTrue(t.replica_contains_host_id(u1)) + self.assertTrue(t.replica_contains_host_id(u2)) + self.assertEqual(t.get_replica_shard_id(u1), 3) + self.assertEqual(t.get_replica_shard_id(u2), 7) + + def test_replica_dict_populated_as_expected(self): + """Minimal targeted regression test for the internal _replica_dict + cache: confirms the O(1)-lookup structure this optimization relies + on is actually populated as {host_id: shard_id}, which the public + API alone does not prove.""" + u1 = UUID('12345678-1234-5678-1234-567812345678') + u2 = UUID('87654321-4321-8765-4321-876543218765') + t = Tablet(0, 100, [(u1, 3), (u2, 7)]) + self.assertEqual(t._replica_dict, {u1: 3, u2: 7}) + + +class DropTabletsByHostIdTest(unittest.TestCase): + """Tests for Tablets.drop_tablets_by_host_id batch-filter path.""" + + def test_drop_removes_matching_tablets(self): + u1 = UUID('12345678-1234-5678-1234-567812345678') + u2 = UUID('87654321-4321-8765-4321-876543218765') + t1 = Tablet(0, 100, [(u1, 0)]) + t2 = Tablet(100, 200, [(u2, 0)]) + t3 = Tablet(200, 300, [(u1, 1), (u2, 1)]) + tablets = Tablets({("ks", "tb"): [t1, t2, t3]}) + + tablets.drop_tablets_by_host_id(u1) + + remaining = tablets._tablets[("ks", "tb")] + self.assertEqual(len(remaining), 1) + self.assertIs(remaining[0], t2) + # Verify token index lists are in sync + self.assertEqual(tablets._first_tokens[("ks", "tb")], [100]) + self.assertEqual(tablets._last_tokens[("ks", "tb")], [200]) + + def test_drop_none_host_id_is_noop(self): + t1 = Tablet(0, 100, [("host1", 0)]) + tablets = Tablets({("ks", "tb"): [t1]}) + tablets.drop_tablets_by_host_id(None) + self.assertEqual(len(tablets._tablets[("ks", "tb")]), 1) + + def test_drop_nonexistent_host_id_is_noop(self): + u1 = UUID('12345678-1234-5678-1234-567812345678') + u_missing = UUID('aaaaaaaa-bbbb-cccc-dddd-eeeeeeeeeeee') + t1 = Tablet(0, 100, [(u1, 0)]) + tablets = Tablets({("ks", "tb"): [t1]}) + tablets.drop_tablets_by_host_id(u_missing) + self.assertEqual(len(tablets._tablets[("ks", "tb")]), 1)