diff --git a/CHANGELOG.rst b/CHANGELOG.rst index 812d46960a..8f0c0928f5 100644 --- a/CHANGELOG.rst +++ b/CHANGELOG.rst @@ -25,6 +25,9 @@ Bug Fixes replace, close, or misconfigure a live pool. Closing a pool with requests in flight could retry an already executed request, causing it to execute twice (#317). +* Tablet updates no longer mutate a table's tablet list in place, so a + concurrent lookup cannot fail with ``IndexError``, miss its tablet, or route + to the wrong one (#1086). Behavior Changes ---------------- diff --git a/cassandra/cluster.py b/cassandra/cluster.py index da4c79b8c2..f281c22eed 100644 --- a/cassandra/cluster.py +++ b/cassandra/cluster.py @@ -83,7 +83,7 @@ RetryPolicy, IdentityTranslator, NoSpeculativeExecutionPlan, NoSpeculativeExecutionPolicy, DefaultLoadBalancingPolicy, NeverRetryPolicy) -from cassandra.pool import (Host, _ReconnectionHandler, _HostReconnectionHandler, +from cassandra.pool import (_TABLET_NOT_LOOKED_UP, Host, _ReconnectionHandler, _HostReconnectionHandler, HostConnection, NoConnectionsAvailable) from cassandra.query import (SimpleStatement, PreparedStatement, BoundStatement, @@ -3855,12 +3855,16 @@ def _create_response_future(self, query, parameters, trace, custom_payload, # balancing policy drops token awareness in that case. Without the check # this path would raise on every prepared-statement execution instead. routing_token = None + routing_tablet = None routing_key = query.routing_key if routing_key is not None: metadata = self.cluster.metadata token_map = metadata.token_map if token_map is not None and metadata.can_support_partitioner(): routing_token = token_map.token_class.from_key(routing_key) + # One tablet lookup per request, shared with the version block and the pool. + routing_tablet = metadata._tablets.get_tablet_for_key( + query.keyspace or self.keyspace, query.table, routing_token) if isinstance(query, SimpleStatement): query_string = query.query_string @@ -3899,7 +3903,7 @@ def _create_response_future(self, query, parameters, trace, custom_payload, and continuous_paging_options is None, continuous_paging_options=continuous_paging_options, result_metadata_id=result_metadata_id, - tablet_version_block=self._compute_tablet_version_block(query, routing_key, routing_token)) + tablet_version_block=self._compute_tablet_version_block(query, routing_key, routing_token, routing_tablet)) elif isinstance(query, BatchStatement): if self._protocol_version < 2: raise UnsupportedOperation( @@ -3927,10 +3931,10 @@ def _create_response_future(self, query, parameters, trace, custom_payload, prepared_statement=prepared_statement, retry_policy=retry_policy, row_factory=row_factory, load_balancer=load_balancing_policy, start_time=start_time, speculative_execution_plan=spec_exec_plan, continuous_paging_state=None, host=host, bound_result_metadata=bound_result_metadata, - routing_token=routing_token) + routing_token=routing_token, routing_tablet=routing_tablet) def _compute_tablet_version_block(self, query, routing_key: Optional[bytes], - routing_token: Optional[Token]) -> int: + routing_token: Optional[Token], tablet: Optional[Tablet] = None) -> int: """ Compute the tablet_version_block byte for a BoundStatement. @@ -3972,11 +3976,8 @@ def _compute_tablet_version_block(self, query, routing_key: Optional[bytes], # useful. Make it possible to obtain it. return random_tablet_version_block() - # A single lookup: get_tablet_for_key already reports a table with no - # cached tablets (a vnode table, or a tablet table on cold start) as - # None, and going through the mutable cache twice would leave a window - # for the tablet to disappear between the checks. - tablet = self.cluster.metadata._tablets.get_tablet_for_key(keyspace, table, routing_token) + # ``tablet`` is the caller's single lookup; None covers a vnode table and + # a tablet table on cold start alike. if tablet is None or tablet.tablet_version is None: # A version miss on the server, which replies with fresh routing info. return random_tablet_version_block() @@ -6729,7 +6730,8 @@ class ResponseFuture(object): def __init__(self, session, message, query, timeout, metrics=None, prepared_statement=None, retry_policy=RetryPolicy(), row_factory=None, load_balancer=None, start_time=None, speculative_execution_plan=None, continuous_paging_state=None, host=None, - bound_result_metadata=_NOT_SET, routing_token=None): + bound_result_metadata=_NOT_SET, routing_token=None, + routing_tablet=_TABLET_NOT_LOOKED_UP): self.session = session # TODO: normalize handling of retry policy and row factory self.row_factory = row_factory or session.row_factory @@ -6750,6 +6752,7 @@ def __init__(self, session, message, query, timeout, metrics=None, prepared_stat self._start_time = start_time or time.time() self._host = host self._routing_token = routing_token + self._routing_tablet = routing_tablet self._control_connection_query_attempted = False self._page_generation = 0 self._retry_aborted = False @@ -7293,7 +7296,7 @@ def _query(self, host, message=None, cb=None): connection, request_id = pool.borrow_connection( timeout=2.0, routing_key=self.query.routing_key, keyspace=self.query.keyspace, table=self.query.table, - routing_token=self._routing_token) + routing_token=self._routing_token, routing_tablet=self._routing_tablet) else: connection, request_id = pool.borrow_connection(timeout=2.0) self._connection = connection diff --git a/cassandra/policies.py b/cassandra/policies.py index 812603616f..9efbe0eb32 100644 --- a/cassandra/policies.py +++ b/cassandra/policies.py @@ -574,15 +574,16 @@ def make_query_plan(self, working_keyspace=None, query=None): tablet = self._cluster_metadata._tablets.get_tablet_for_key(keyspace, query.table, token) if tablet is not None: + replica_dict = tablet._replica_dict if keep_order: # The child plan is round-robin rotated, so it cannot provide a stable order. replicas = [host for host in (self._cluster_metadata.get_host_by_host_id(host_id) for host_id, _ in tablet.replicas) if host is not None] else: - replicas_mapped = set(map(lambda r: r[0], tablet.replicas)) 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 is not None and host.host_id.int in replica_dict] # The leader concept only exists for strongly-consistent keyspaces, # which today means exactly the keyspaces whose consistency mode is diff --git a/cassandra/pool.py b/cassandra/pool.py index 26b5670e09..3d3cff4e26 100644 --- a/cassandra/pool.py +++ b/cassandra/pool.py @@ -34,6 +34,9 @@ DefaultEndPoint, UnixSocketEndPoint) from cassandra.policies import HostDistance +# Default for routing_tablet: "not looked up" (None means looked up, no tablet). +_TABLET_NOT_LOOKED_UP = object() + log = logging.getLogger(__name__) @@ -510,7 +513,8 @@ 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, routing_token=None): + def _get_connection_for_routing_key(self, routing_key=None, keyspace=None, table=None, routing_token=None, + routing_tablet=_TABLET_NOT_LOOKED_UP): if self.is_shutdown: raise ConnectionException( "Pool for %s is shutdown" % (self.host,), self.host) @@ -531,19 +535,19 @@ def _get_connection_for_routing_key(self, routing_key=None, keyspace=None, table if t is None and metadata.token_map is not None and metadata.can_support_partitioner(): t = metadata.token_map.token_class.from_key(routing_key) if t is not None and self.supports_tablet_routing and table is not None: - if keyspace is None: - keyspace = self._keyspace - - tablet = self._session.cluster.metadata._tablets.get_tablet_for_key(keyspace, table, t) + if routing_tablet is not _TABLET_NOT_LOOKED_UP: + # The caller already did the lookup; None means no tablet. + tablet = routing_tablet + else: + if keyspace is None: + keyspace = self._keyspace + tablet = self._session.cluster.metadata._tablets.get_tablet_for_key(keyspace, table, t) # In both V1 and V2 the request is sent to this host, so we pick # the shard that this host owns for the tablet. Leader-aware host # selection (V2) happens earlier, in the load balancing policy. if tablet is not None: - for replica in tablet.replicas: - if replica[0] == self.host.host_id: - shard_id = replica[1] - break + shard_id = tablet.get_replica_shard_id(self.host.host_id) if shard_id is None and t is not None: shard_id = self.host.sharding_info.shard_id_from_token(t.value) @@ -586,15 +590,16 @@ 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, routing_token=None): - conn = self._get_connection_for_routing_key(routing_key, keyspace, table, routing_token) + def borrow_connection(self, timeout, routing_key=None, keyspace=None, table=None, routing_token=None, + routing_tablet=_TABLET_NOT_LOOKED_UP): + conn = self._get_connection_for_routing_key(routing_key, keyspace, table, routing_token, routing_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, routing_token) + conn = self._get_connection_for_routing_key(routing_key, keyspace, table, routing_token, routing_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 b386d1a372..78ed921e50 100644 --- a/cassandra/tablets.py +++ b/cassandra/tablets.py @@ -1,14 +1,9 @@ -from bisect import bisect_left -from operator import attrgetter +from bisect import bisect_left, bisect_right from random import getrandbits 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") - def choose_tablet_version_block(tablet_version: int) -> int: """ @@ -36,23 +31,27 @@ def random_tablet_version_block() -> int: return getrandbits(8) +# host_id.int -> one shared UUID per host; tablets key by its .int so all share one int per host. +# ponytail: never pruned; grows with every host ever seen, fine unless hosts churn by thousands. +_host_ids = {} + + class Tablet(object): """ Represents a single ScyllaDB tablet. 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 - # uint64 hash; None means unknown -- a cold start, or a tablet learned over - # TABLETS_ROUTING_V1, which does not report a version. - tablet_version = None + __slots__ = ('first_token', 'last_token', 'tablet_version', '_replica_dict') def __init__(self, first_token=0, last_token=0, replicas=None, tablet_version=None): self.first_token = first_token self.last_token = last_token - self.replicas = replicas + # The only replica storage: {host_id.int: shard}, in wire order so the leader is first. + # Int keys: UUID.__hash__ is pure Python, and the wire's UUID objects can be freed. + intern = _host_ids.setdefault + self._replica_dict = {intern(u.int, u).int: s for u, s in replicas} if replicas is not None else {} + # uint64 hash; None = unknown (cold start, or learned over TABLETS_ROUTING_V1). self.tablet_version = tablet_version def __str__(self): @@ -60,21 +59,21 @@ def __str__(self): % (self.first_token, self.last_token, self.replicas, self.tablet_version) __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, tablet_version=None): - if Tablet._is_valid_tablet(replicas): - if tablet_version is not None: - # tablet_version is an unsigned 64-bit value, but it is - # deserialized from the wire as a signed LongType; normalize it - # back to unsigned so it matches the server's representation. - tablet_version &= 0xFFFFFFFFFFFFFFFF - tablet = Tablet(first_token, last_token, replicas, tablet_version) - return tablet - return None + if tablet_version is not None: + # tablet_version is an unsigned 64-bit value, but it is + # deserialized from the wire as a signed LongType; normalize it + # back to unsigned so it matches the server's representation. + tablet_version &= 0xFFFFFFFFFFFFFFFF + # __init__ consumes replicas once, so empty generators are caught too. + tablet = Tablet(first_token, last_token, replicas, tablet_version) + return tablet if tablet._replica_dict else None + + @property + def replicas(self): + # Rebuilt on each access; kept for compatibility, not used on the query path. + return tuple([(_host_ids[k], s) for k, s in self._replica_dict.items()]) @property def leader(self) -> Optional[UUID]: @@ -99,36 +98,47 @@ def leader(self) -> Optional[UUID]: Returns ``None`` for a tablet with no replicas rather than raising, so callers do not have to guard the lookup themselves. """ - if not self.replicas: - return None - return self.replicas[0][0] + for key in self._replica_dict: + return _host_ids[key] + return None - def replica_contains_host_id(self, uuid: UUID) -> bool: - for replica in self.replicas: - if replica[0] == uuid: - return True - return False + def replica_contains_host_id(self, uuid: Optional[UUID]) -> bool: + # A host whose id is not yet known (discovery/metadata transitions) is + # not a replica; treat it as a non-match rather than raising. + if uuid is None: + return False + return uuid.int in self._replica_dict + def get_replica_shard_id(self, uuid: Optional[UUID]) -> Optional[int]: + if uuid is None: + return None + return self._replica_dict.get(uuid.int) -class Tablets(object): - _lock = None - _tablets = {} +class Tablets(object): def __init__(self, tablets): - self._tablets = tablets + # Instance-only: mutable class-level dicts would be shared across instances. self._lock = Lock() + # (keyspace, table) -> (tablets, last_tokens): one snapshot, replaced whole under _lock + # and never mutated, so lock-free readers see a matching pair. last_tokens lets bisect skip key=. + self._tablets = {key: (tlist, [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), [])) + entry = self._tablets.get((keyspace, table)) + return entry is not None and bool(entry[0]) def get_tablet_for_key(self, keyspace, table, t): - tablet = self._tablets.get((keyspace, table), []) - if not tablet: + entry = self._tablets.get((keyspace, table)) + if entry is None: 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] + tablets, last_tokens = entry + token_value = t.value + try: + tablet = tablets[bisect_left(last_tokens, token_value)] + except IndexError: + return None + if token_value > tablet.first_token: + return tablet return None def drop_tablets(self, keyspace: str, table: Optional[str] = None): @@ -149,31 +159,31 @@ 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) + for key, (tablets, _) in list(self._tablets.items()): + kept = [tablet for tablet in tablets if not tablet.replica_contains_host_id(host_id)] + if not kept: + # Don't leave an empty entry for a table with no tablets left. + del self._tablets[key] + elif len(kept) != len(tablets): + self._tablets[key] = (kept, [t.last_token for t in kept]) def add_tablet(self, keyspace, table, tablet): with self._lock: - tablets_for_table = self._tablets.setdefault((keyspace, table), []) + key = (keyspace, table) + # Copy-on-write: lock-free readers in get_tablet_for_key may hold the old snapshot. + tablets_for_table, last_tokens = self._tablets.get(key, ((), ())) + tablets_for_table = list(tablets_for_table) + last_tokens = list(last_tokens) # 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 = start - 1 + start = bisect_right(last_tokens, tablet.first_token) # 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 tablets_for_table[end].first_token >= tablet.last_token: end = end - 1 - if start <= end: - del tablets_for_table[start:end + 1] - - tablets_for_table.insert(start, tablet) - + # Slice assignment replaces the overlap, or inserts when start > end. + tablets_for_table[start:end + 1] = (tablet,) + last_tokens[start:end + 1] = (tablet.last_token,) + self._tablets[key] = (tablets_for_table, last_tokens) diff --git a/tests/unit/test_host_connection_pool.py b/tests/unit/test_host_connection_pool.py index ca1cb51573..09e9baafa3 100644 --- a/tests/unit/test_host_connection_pool.py +++ b/tests/unit/test_host_connection_pool.py @@ -338,3 +338,32 @@ def mock_connection_factory(self, *args, **kwargs): # Cleanup executor with proper wait session.cluster.executor.shutdown(wait=True) + + +def test_pool_reuses_session_tablet_without_second_lookup(): + # The session looks the tablet up once; the pool must not repeat it. + tablet = Mock() + tablet.get_replica_shard_id.return_value = 3 + pool = Mock(spec=HostConnection, is_shutdown=False, supports_tablet_routing=True, _connections={3: Mock(orphaned_threshold_reached=False)}) + pool._session.cluster.shard_aware_options.disable = False + pool.host.sharding_info = Mock() + pool._session.cluster.metadata._tablets.get_tablet_for_key.side_effect = AssertionError("second lookup") + HostConnection._get_connection_for_routing_key( + pool, b'k', 'ks', 'tb', routing_token=Mock(), routing_tablet=tablet) + tablet.get_replica_shard_id.assert_called_once() + + +def test_pool_looks_up_tablet_when_only_token_given(): + # A caller passing routing_token without routing_tablet must still get tablet routing. + tablet = Mock() + tablet.get_replica_shard_id.return_value = 3 + conn = Mock(orphaned_threshold_reached=False) + pool = Mock(spec=HostConnection, is_shutdown=False, supports_tablet_routing=True, _connections={3: conn}) + pool._session.cluster.shard_aware_options.disable = False + pool.host.sharding_info = Mock() + lookup = pool._session.cluster.metadata._tablets.get_tablet_for_key + lookup.return_value = tablet + assert HostConnection._get_connection_for_routing_key( + pool, b'k', 'ks', 'tb', routing_token=Mock()) is conn + lookup.assert_called_once() + tablet.get_replica_shard_id.assert_called_once() diff --git a/tests/unit/test_metadata.py b/tests/unit/test_metadata.py index dbd71c40aa..bb2736aa16 100644 --- a/tests/unit/test_metadata.py +++ b/tests/unit/test_metadata.py @@ -459,7 +459,7 @@ def setUp(self): keyspace = KeyspaceMetadata("ks", True, "NetworkTopologyStrategy", {"dc1": "1"}) keyspace.tables["tb"] = TableMetadata("ks", "tb") self.metadata.keyspaces["ks"] = keyspace - self.metadata._tablets.add_tablet("ks", "tb", Tablet(0, 100, [("host1", 0)])) + self.metadata._tablets.add_tablet("ks", "tb", Tablet(0, 100, [(uuid.uuid4(), 0)])) def test_drop_table_invalidates_tablets(self): """Dropping a known table removes its tablet and table metadata.""" @@ -470,7 +470,7 @@ def test_drop_table_invalidates_tablets(self): def test_drop_table_invalidates_tablets_for_unknown_keyspace(self): """Dropping a table in an unknown keyspace still removes its tablet metadata.""" - self.metadata._tablets.add_tablet("unknown", "tb", Tablet(0, 100, [("host1", 0)])) + self.metadata._tablets.add_tablet("unknown", "tb", Tablet(0, 100, [(uuid.uuid4(), 0)])) self.metadata._drop_table("unknown", "tb") assert self.metadata._tablets.table_has_tablets("unknown", "tb") is False diff --git a/tests/unit/test_policies.py b/tests/unit/test_policies.py index d1d86f0f5f..5899c83c32 100644 --- a/tests/unit/test_policies.py +++ b/tests/unit/test_policies.py @@ -1013,6 +1013,44 @@ def test_lwt_tablet_keeps_natural_replica_order(self): for _ in range(len(hosts)): assert list(policy.make_query_plan(None, query))[:3] == order + def test_tablet_routing_skips_host_with_unknown_host_id(self): + """A host whose host_id is not yet known must be treated as a + non-replica during tablet routing, not raise AttributeError. + + Host.__init__ rejects host_id=None, but Host.host_id carries a + class-level None default, and a child policy can yield a host whose id + is not populated yet.""" + hosts = [Host(DefaultEndPoint(str(i)), SimpleConvictionPolicy, host_id=uuid.uuid4()) + for i in range(3)] + for h in hosts: + h.set_up() + h.set_location_info("dc1", "rack1") + unknown = Mock() + unknown.host_id = None + unknown.is_up = True + all_hosts = hosts + [unknown] + + cluster = Mock(spec=Cluster) + cluster.metadata = Mock(spec=Metadata) + cluster.metadata._tablets = Mock(spec=Tablets) + cluster.metadata.all_hosts.return_value = all_hosts + cluster.metadata._tablets.get_tablet_for_key.return_value = Tablet( + replicas=[(h.host_id, 0) for h in hosts[:2]]) + + child_policy = Mock() + child_policy.make_query_plan.return_value = all_hosts + child_policy.distance.return_value = HostDistance.LOCAL + + policy = TokenAwarePolicy(child_policy, shuffle_replicas=False) + policy.populate(cluster, all_hosts) + + query = Statement(routing_key=b"key", keyspace="ks") + plan = list(policy.make_query_plan(None, query)) + # The two known replicas come first; the id-less host is skipped from + # replica selection and only surfaces later as a non-replica fallback. + assert hosts[:2] == plan[:2] + assert unknown in plan[2:] + def test_leader_aware_routing_with_tablet_version(self): """ For a strongly-consistent keyspace, the leader (first replica in the diff --git a/tests/unit/test_response_future.py b/tests/unit/test_response_future.py index ba1f179748..8f28f75245 100644 --- a/tests/unit/test_response_future.py +++ b/tests/unit/test_response_future.py @@ -111,7 +111,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, routing_token=ANY) + pool.borrow_connection.assert_called_once_with(timeout=ANY, routing_key=ANY, keyspace=ANY, table=ANY, routing_token=ANY, routing_tablet=ANY) connection.send_msg.assert_called_once_with(rf.message, 1, cb=ANY, encoder=ProtocolHandler.encode_message, decoder=ProtocolHandler.decode_message, result_metadata=[]) @@ -304,7 +304,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, routing_token=ANY) + pool.borrow_connection.assert_called_once_with(timeout=ANY, routing_key=ANY, keyspace=ANY, table=ANY, routing_token=ANY, routing_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={}) @@ -324,7 +324,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, routing_token=ANY) + pool.borrow_connection.assert_called_with(timeout=ANY, routing_key=ANY, keyspace=ANY, table=ANY, routing_token=ANY, routing_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): @@ -339,7 +339,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, routing_token=ANY) + pool.borrow_connection.assert_called_once_with(timeout=ANY, routing_key=ANY, keyspace=ANY, table=ANY, routing_token=ANY, routing_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 @@ -359,7 +359,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, routing_token=ANY) + pool.borrow_connection.assert_called_with(timeout=ANY, routing_key=ANY, keyspace=ANY, table=ANY, routing_token=ANY, routing_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 @@ -2013,7 +2013,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, routing_token=ANY) + pool.borrow_connection.assert_called_once_with(timeout=ANY, routing_key=ANY, keyspace=ANY, table=ANY, routing_token=ANY, routing_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 656ae42da7..077a158a9f 100644 --- a/tests/unit/test_tablets.py +++ b/tests/unit/test_tablets.py @@ -1,6 +1,6 @@ import unittest from io import BytesIO -from uuid import uuid4 +from uuid import UUID, uuid4 from cassandra import ConsistencyLevel, ProtocolVersion from cassandra.protocol import ExecuteMessage @@ -20,7 +20,7 @@ def test_add_tablet_to_empty_tablets(self): tablets.add_tablet("test_ks", "test_tb", Tablet(-6917529027641081857, -4611686018427387905, None)) - tablets_list = tablets._tablets.get(("test_ks", "test_tb")) + tablets_list = tablets._tablets.get(("test_ks", "test_tb"))[0] self.compare_ranges(tablets_list, [(-6917529027641081857, -4611686018427387905)]) @@ -29,7 +29,7 @@ def test_add_tablet_at_the_beggining(self): tablets.add_tablet("test_ks", "test_tb", Tablet(-8611686018427387905, -7917529027641081857, None)) - tablets_list = tablets._tablets.get(("test_ks", "test_tb")) + tablets_list = tablets._tablets.get(("test_ks", "test_tb"))[0] self.compare_ranges(tablets_list, [(-8611686018427387905, -7917529027641081857), (-6917529027641081857, -4611686018427387905)]) @@ -39,7 +39,7 @@ def test_add_tablet_at_the_end(self): tablets.add_tablet("test_ks", "test_tb", Tablet(-1, 2305843009213693951, None)) - tablets_list = tablets._tablets.get(("test_ks", "test_tb")) + tablets_list = tablets._tablets.get(("test_ks", "test_tb"))[0] self.compare_ranges(tablets_list, [(-6917529027641081857, -4611686018427387905), (-1, 2305843009213693951)]) @@ -50,7 +50,7 @@ def test_add_tablet_in_the_middle(self): tablets.add_tablet("test_ks", "test_tb", Tablet(-4611686018427387905, -2305843009213693953, None)) - tablets_list = tablets._tablets.get(("test_ks", "test_tb")) + tablets_list = tablets._tablets.get(("test_ks", "test_tb"))[0] self.compare_ranges(tablets_list, [(-6917529027641081857, -4611686018427387905), (-4611686018427387905, -2305843009213693953), @@ -64,7 +64,7 @@ def test_add_tablet_intersecting(self): tablets.add_tablet("test_ks", "test_tb", Tablet(-3611686018427387905, -6, None)) - tablets_list = tablets._tablets.get(("test_ks", "test_tb")) + tablets_list = tablets._tablets.get(("test_ks", "test_tb"))[0] self.compare_ranges(tablets_list, [(-6917529027641081857, -4611686018427387905), (-3611686018427387905, -6), @@ -76,7 +76,7 @@ def test_add_tablet_intersecting_with_first(self): tablets.add_tablet("test_ks", "test_tb", Tablet(-8011686018427387905, -7987529027641081857, None)) - tablets_list = tablets._tablets.get(("test_ks", "test_tb")) + tablets_list = tablets._tablets.get(("test_ks", "test_tb"))[0] self.compare_ranges(tablets_list, [(-8011686018427387905, -7987529027641081857), (-6917529027641081857, -4611686018427387905)]) @@ -87,7 +87,7 @@ def test_add_tablet_intersecting_with_last(self): tablets.add_tablet("test_ks", "test_tb", Tablet(-5011686018427387905, -2987529027641081857, None)) - tablets_list = tablets._tablets.get(("test_ks", "test_tb")) + tablets_list = tablets._tablets.get(("test_ks", "test_tb"))[0] self.compare_ranges(tablets_list, [(-8611686018427387905, -7917529027641081857), (-5011686018427387905, -2987529027641081857)]) @@ -97,9 +97,9 @@ class GetTabletForKeyTest(unittest.TestCase): """Tests for Tablets.get_tablet_for_key.""" def test_found(self): - t1 = Tablet(0, 100, [("host1", 0)]) - t2 = Tablet(100, 200, [("host2", 0)]) - t3 = Tablet(200, 300, [("host3", 0)]) + t1 = Tablet(0, 100, [(uuid4(), 0)]) + t2 = Tablet(100, 200, [(uuid4(), 0)]) + t3 = Tablet(200, 300, [(uuid4(), 0)]) tablets = Tablets({("ks", "tb"): [t1, t2, t3]}) class Token: @@ -119,7 +119,7 @@ def __init__(self, v): self.assertIsNone(tablets.get_tablet_for_key("ks", "tb", Token(50))) def test_not_found_outside_range(self): - t1 = Tablet(100, 200, [("host1", 0)]) + t1 = Tablet(100, 200, [(uuid4(), 0)]) tablets = Tablets({("ks", "tb"): [t1]}) class Token: @@ -131,6 +131,94 @@ def __init__(self, v): self.assertIsNone(tablets.get_tablet_for_key("ks", "tb", Token(50))) +class _Token: + def __init__(self, v): + self.value = v + + +class TabletsCopyOnWriteTest(unittest.TestCase): + """Writers must publish a new (tablets, last_tokens) snapshot, never mutate one a lock-free reader may hold (#1086).""" + + def _ranges(self, lst): + return [(t.first_token, t.last_token) for t in lst] + + def _make(self): + h1, h2 = uuid4(), uuid4() + tablets = [Tablet(0, 100, [(h1, 0)]), Tablet(100, 200, [(h2, 0)]), + Tablet(200, 300, [(h1, 0)]), Tablet(300, 400, [(h2, 0)])] + return Tablets({("ks", "tb"): tablets}), h1, h2 + + def _assert_consistent(self, entry): + tablets, last_tokens = entry + self.assertEqual([t.last_token for t in tablets], last_tokens) + + def _assert_snapshot_unchanged(self, mutate, expected_after): + tablets, h1, h2 = self._make() + snapshot = tablets._tablets[("ks", "tb")] + before = (list(snapshot[0]), list(snapshot[1])) + mutate(tablets, h1, h2) + self.assertEqual(snapshot, before) + self._assert_consistent(tablets._tablets[("ks", "tb")]) + self.assertEqual(self._ranges(tablets._tablets[("ks", "tb")][0]), expected_after) + + def test_add_overlapping_tablet_keeps_snapshot(self): + self._assert_snapshot_unchanged( + lambda t, h1, h2: t.add_tablet("ks", "tb", Tablet(50, 350, [(h1, 0)])), + [(50, 350)]) + + def test_add_non_overlapping_tablet_keeps_snapshot(self): + self._assert_snapshot_unchanged( + lambda t, h1, h2: t.add_tablet("ks", "tb", Tablet(400, 500, [(h1, 0)])), + [(0, 100), (100, 200), (200, 300), (300, 400), (400, 500)]) + + def test_drop_tablets_by_host_id_keeps_snapshot(self): + self._assert_snapshot_unchanged( + lambda t, h1, h2: t.drop_tablets_by_host_id(h1), + [(100, 200), (300, 400)]) + + def test_drop_tablets_by_host_id(self): + tablets, h1, h2 = self._make() + tablets._tablets[("ks", "other")] = ([Tablet(0, 10, [(h2, 0)])], [10]) + tablets.drop_tablets_by_host_id(h2) + self.assertEqual(self._ranges(tablets._tablets[("ks", "tb")][0]), [(0, 100), (200, 300)]) + self._assert_consistent(tablets._tablets[("ks", "tb")]) + self.assertNotIn(("ks", "other"), tablets._tablets) + self.assertIsNone(tablets.get_tablet_for_key("ks", "tb", _Token(150))) + self.assertEqual(tablets.get_tablet_for_key("ks", "tb", _Token(250)).first_token, 200) + + def _lookup_racing(self, write): + # The reader takes its snapshot, then reads t.value; the write lands in between. + tablets, h1, h2 = self._make() + snapshot = tablets._tablets[("ks", "tb")] + before = (list(snapshot[0]), list(snapshot[1])) + last = snapshot[0][3] + + class RacingToken: + reads = 0 + + @property + def value(self): + RacingToken.reads += 1 + write(tablets, h1, h2) + return 350 + + self.assertIs(tablets.get_tablet_for_key("ks", "tb", RacingToken()), last) + self.assertEqual(RacingToken.reads, 1) + self.assertEqual(snapshot, before) + self._assert_consistent(tablets._tablets[("ks", "tb")]) + return tablets + + def test_add_tablet_during_lookup(self): + tablets = self._lookup_racing( + lambda t, h1, h2: t.add_tablet("ks", "tb", Tablet(-1, 1000, [(h1, 0)]))) + self.assertEqual(self._ranges(tablets._tablets[("ks", "tb")][0]), [(-1, 1000)]) + + def test_drop_tablets_by_host_id_during_lookup(self): + # Drops the tablet the reader is about to pick and shifts every later index. + tablets = self._lookup_racing(lambda t, h1, h2: t.drop_tablets_by_host_id(h2)) + self.assertEqual(self._ranges(tablets._tablets[("ks", "tb")][0]), [(0, 100), (200, 300)]) + + class TabletLeaderTest(unittest.TestCase): """Tests for Tablet.leader, the leader-first replica ordering V2 provides.""" @@ -214,7 +302,7 @@ def test_random_tablet_version_block_returns_byte(self): def test_from_row_stores_tablet_version(self): """Tablet.from_row stores the tablet_version it is given (the V2 payload field).""" version = 0xDEADBEEFCAFEBABE - tablet = Tablet.from_row(-100, 100, [("host1", 0), ("host2", 1)], tablet_version=version) + tablet = Tablet.from_row(-100, 100, [(uuid4(), 0), (uuid4(), 1)], tablet_version=version) self.assertIsNotNone(tablet) self.assertEqual(tablet.tablet_version, version) self.assertEqual(tablet.first_token, -100) @@ -279,3 +367,131 @@ def test_same_message_encodes_consistently_across_connections(self): first_again = self._encode_body(message, ProtocolFeatures(tablets_routing_v2=True)) self.assertEqual(first, first_again) self.assertEqual(first, second_plain + bytes([0x3C])) + +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) + + +class TabletReplicaDictTest(unittest.TestCase): + """replica_contains_host_id / get_replica_shard_id lookups.""" + + 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_none_host_id_is_not_a_replica(self): + # A host whose id is still unknown must be treated as a non-replica + # rather than raising (discovery/metadata transitions). + u1 = UUID('12345678-1234-5678-1234-567812345678') + t = Tablet(0, 100, [(u1, 3)]) + self.assertFalse(t.replica_contains_host_id(None)) + self.assertIsNone(t.get_replica_shard_id(None)) + + 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.assertEqual(t.get_replica_shard_id(u2), 7) + + +class TabletHostIdInternTest(unittest.TestCase): + def test_equal_host_ids_share_one_object(self): + u = uuid4() + t1 = Tablet(0, 100, [(UUID(int=u.int), 0)]) + t2 = Tablet(100, 200, [(UUID(int=u.int), 1)]) + self.assertIs(t1.leader, t2.leader) + self.assertEqual(t1.leader, u) + self.assertIs(next(iter(t1._replica_dict)), next(iter(t2._replica_dict))) + + +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, last_tokens = tablets._tablets[("ks", "tb")] + self.assertEqual(len(remaining), 1) + self.assertIs(remaining[0], t2) + self.assertEqual(last_tokens, [200]) + + def test_drop_none_host_id_is_noop(self): + t1 = Tablet(0, 100, [(uuid4(), 0)]) + tablets = Tablets({("ks", "tb"): [t1]}) + tablets.drop_tablets_by_host_id(None) + self.assertEqual(len(tablets._tablets[("ks", "tb")][0]), 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")][0]), 1) + + def test_drop_last_tablet_removes_table_keys(self): + # Dropping the only tablet of a table must not leave an empty entry (PR #651 cleanup). + u1 = UUID('12345678-1234-5678-1234-567812345678') + t1 = Tablet(0, 100, [(u1, 0)]) + tablets = Tablets({("ks", "tb"): [t1]}) + + tablets.drop_tablets_by_host_id(u1) + + self.assertNotIn(("ks", "tb"), tablets._tablets) + self.assertFalse(tablets.table_has_tablets("ks", "tb"))