From 08e8726084280620128d21d6e783cf300febe11a Mon Sep 17 00:00:00 2001 From: Dmitry Kropachev Date: Thu, 30 Jul 2026 20:03:55 -0400 Subject: [PATCH] connection: tolerate clean startup close --- cassandra/cluster.py | 3101 +++++++++++++++-- cassandra/connection.py | 440 ++- cassandra/io/twistedreactor.py | 75 +- cassandra/pool.py | 2028 +++++++++-- .../test_maintenance_mode_connection.py | 151 + tests/unit/io/test_twistedreactor.py | 95 +- tests/unit/test_cluster.py | 2686 +++++++++++++- tests/unit/test_connection.py | 680 +++- tests/unit/test_control_connection.py | 301 +- tests/unit/test_host_connection_pool.py | 2021 ++++++++++- tests/unit/test_response_future.py | 484 ++- tests/unit/test_shard_aware.py | 70 + 12 files changed, 11476 insertions(+), 656 deletions(-) create mode 100644 tests/integration/standard/test_maintenance_mode_connection.py diff --git a/cassandra/cluster.py b/cassandra/cluster.py index 88c8d2707a..3f17813cfd 100644 --- a/cassandra/cluster.py +++ b/cassandra/cluster.py @@ -39,7 +39,7 @@ import socket import sys import time -from threading import Lock, RLock, Thread, Event +from threading import Lock, RLock, Thread, Event, get_ident, local import uuid import weakref @@ -47,13 +47,17 @@ from cassandra import (ConsistencyLevel, AuthenticationFailed, OperationTimedOut, UnsupportedOperation, SchemaTargetType, DriverException, ProtocolVersion, - UnresolvableContactPoints, DependencyException) + UnresolvableContactPoints, DependencyException, + RequestValidationException) from cassandra.auth import _proxy_execute_key, PlainTextAuthProvider from cassandra.client_routes import ClientRoutesChangeType, ClientRoutesConfig, _ClientRoutesHandler from cassandra.connection import (ClientRoutesEndPointFactory, ConnectionException, ConnectionShutdown, - ConnectionHeartbeat, ProtocolVersionUnsupported, - EndPoint, DefaultEndPoint, DefaultEndPointFactory, - SniEndPointFactory, ConnectionBusy, locally_supported_compressions) + ConnectionHeartbeat, ProtocolVersionUnsupported, + EndPoint, DefaultEndPoint, DefaultEndPointFactory, + SniEndPointFactory, ConnectionBusy, locally_supported_compressions, + _ConnectionClosedDuringStartup, + _set_keyspace_blocking, + _startup_close_error) from cassandra.cqltypes import UserType import cassandra.cqltypes as types from cassandra.encoder import Encoder @@ -69,7 +73,9 @@ BatchMessage, RESULT_KIND_PREPARED, RESULT_KIND_SET_KEYSPACE, RESULT_KIND_ROWS, RESULT_KIND_SCHEMA_CHANGE, ProtocolHandler, - RESULT_KIND_VOID, ProtocolException) + RESULT_KIND_VOID, ProtocolException, + RequestValidationException as + ProtocolRequestValidationException) from cassandra.metadata import Metadata, protect_name, murmur3, _NodeInfo from cassandra.policies import (TokenAwarePolicy, DCAwareRoundRobinPolicy, SimpleConvictionPolicy, ExponentialReconnectionPolicy, HostDistance, @@ -78,7 +84,8 @@ NeverRetryPolicy) from cassandra.pool import (Host, _ReconnectionHandler, _HostReconnectionHandler, HostConnection, - NoConnectionsAvailable) + NoConnectionsAvailable, + _signal_connection_failure) from cassandra.query import (SimpleStatement, PreparedStatement, BoundStatement, BatchStatement, bind_params, QueryTrace, TraceUnavailable, named_tuple_factory, dict_factory, tuple_factory, FETCH_SIZE_UNSET, @@ -192,6 +199,68 @@ def _connection_reduce_fn(val,import_fn): log = logging.getLogger(__name__) + +def _is_use_statement(query): + """Return whether CQL ``query`` starts with a USE statement.""" + if not isinstance(query, str): + return False + + length = len(query) + position = 0 + while True: + while position < length and query[position].isspace(): + position += 1 + + if query.startswith('--', position) or query.startswith('//', position): + newline = query.find('\n', position + 2) + if newline < 0: + return False + position = newline + 1 + continue + + if query.startswith('/*', position): + comment_end = query.find('*/', position + 2) + if comment_end < 0: + return False + position = comment_end + 2 + continue + break + + if query[position:position + 3].upper() != 'USE': + return False + boundary = position + 3 + return bool( + boundary < length and ( + query[boundary].isspace() or + query.startswith('--', boundary) or + query.startswith('//', boundary) or + query.startswith('/*', boundary))) + + +class _HostTransitionResult(object): + """ + Work to run only after a host's serialized transition lane is released. + + External notifications to enqueue after this core transition. + + Notifications use a separate per-host FIFO so extension callbacks cannot + keep a later DOWN or REMOVE from performing driver cleanup or observe + events out of order. + """ + + def __init__(self, value=None, notifications=()): + self.value = value + self.notifications = tuple(notifications) + + +_HOST_TRANSITION_CONTEXT = local() +_HOST_TRANSITION_DEFERRED = object() + + +_REQUEST_VALIDATION_EXCEPTIONS = ( + RequestValidationException, + ProtocolRequestValidationException) + _GRAPH_PAGING_MIN_DSE_VERSION = Version('6.8.0') _NOT_SET = object() @@ -243,8 +312,37 @@ def new_f(self, *args, **kwargs): try: future = self.executor.submit(f, self, *args, **kwargs) future.add_done_callback(_future_completed) + return future except Exception: log.exception("Failed to submit task to executor") + if self.is_shutdown: + return None + + # A DOWN transition has already committed before this helper is + # called. Losing its cleanup task would leave the Host permanently + # down without pool teardown or a reconnector. A dedicated daemon + # is a last-resort path for an active Cluster whose executor rejects + # work; it also lets caller-held Cluster/Host/Session locks unwind. + def run_fallback(): + try: + f(self, *args, **kwargs) + except BaseException: + log.exception( + "Failed to run rejected executor task in fallback " + "thread") + + try: + fallback = Thread( + target=run_fallback, + name="cassandra-executor-fallback") + fallback.daemon = True + fallback.start() + return fallback + except BaseException: + log.exception( + "Failed to start fallback thread for rejected executor " + "task") + return None return new_f @@ -553,16 +651,40 @@ def on_up(self, host): def on_down(self, host): if not self.pools_allowed: return + first_error = None for p in self.profiles.values(): - p.load_balancing_policy.on_down(host) + try: + p.load_balancing_policy.on_down(host) + except BaseException as exc: + if first_error is None: + first_error = exc + else: + log.exception( + "Additional load-balancing policy failure while " + "marking host %s down", + host) + if first_error is not None: + raise first_error.with_traceback(first_error.__traceback__) def on_add(self, host): for p in self.profiles.values(): p.load_balancing_policy.on_add(host) def on_remove(self, host): + first_error = None for p in self.profiles.values(): - p.load_balancing_policy.on_remove(host) + try: + p.load_balancing_policy.on_remove(host) + except BaseException as exc: + if first_error is None: + first_error = exc + else: + log.exception( + "Additional load-balancing policy failure while " + "removing host %s", + host) + if first_error is not None: + raise first_error.with_traceback(first_error.__traceback__) @property def default(self): @@ -644,6 +766,9 @@ class ControlConnectionQueryFallback(enum.Enum): the control-connection fallback path for application queries. The fallback path is not used for requests targeted to an explicit host. + Continuous-paging and ``USE`` requests are not sent over the shared + control connection. A Session keyspace requires protocol v5 or DSE v2 so + it can be encoded independently on each request. """ Disabled = "Disabled" @@ -979,6 +1104,8 @@ def default_retry_policy(self, policy): ``Fallback`` enables control-connection fallback when no usable node pools exist. ``SkipPoolCreation`` skips node-pool creation and uses the control connection fallback path. This fallback is still not used for requests targeted to an explicit host. + It also excludes continuous paging and ``USE`` statements; Session + keyspaces require protocol v5 or DSE v2. """ idle_heartbeat_interval = 30 @@ -1743,17 +1870,33 @@ def add_execution_profile(self, name, profile, pool_wait_timeout=5): raise OperationTimedOut("Failed to create all new connection pools in the %ss timeout." % pool_wait_timeout, timeout=pool_wait_timeout) - def connection_factory(self, endpoint, host_conn = None, *args, **kwargs): + def connection_factory(self, endpoint, *args, **kwargs): """ Called to create a new connection with proper configuration. + + This preserves ``Connection.factory``'s return contract, including + a closed connection returned after a clean server close during startup. + Cluster-owned callers validate the result before adopting it. + Intended for internal use only. """ + host_conn = kwargs.pop('host_conn', None) kwargs = self._make_connection_kwargs(endpoint, kwargs) - return self.connection_class.factory(endpoint, self.connect_timeout, host_conn, *args, **kwargs) + if host_conn is not None: + kwargs['host_conn'] = host_conn + return self.connection_class.factory( + endpoint, self.connect_timeout, *args, **kwargs) def _make_connection_factory(self, host, *args, **kwargs): kwargs = self._make_connection_kwargs(host.endpoint, kwargs) - return partial(self.connection_class.factory, host.endpoint, self.connect_timeout, *args, **kwargs) + # Keep the raw factory result observable. The host reconnector owns the + # decision to park a clean startup close. + return partial( + self.connection_class.factory, + host.endpoint, + self.connect_timeout, + *args, + **kwargs) def _make_connection_kwargs(self, endpoint, kwargs_dict): if self._auth_provider_callable: @@ -1911,9 +2054,75 @@ def __exit__(self, *args): self.shutdown() def _new_session(self, keyspace): - session = Session(self, self.metadata.all_hosts(), keyspace) - self._session_register_user_types(session) - self.sessions.add(session) + session = None + try: + session = Session(self, self.metadata.all_hosts(), keyspace) + self._session_register_user_types(session) + except BaseException: + if session is not None: + session.shutdown() + raise + + stale_pools = [] + register_session = False + # Close the pool-publication -> Session-registration gap against + # Cluster.shutdown and DOWN. A DOWN which ran before this lock leaves + # host state for us to reconcile; one which runs after sees the newly + # registered Session. + with self._lock: + if not self.is_shutdown: + self.sessions.add(session) + register_session = True + for host in tuple(session.hosts): + with host.lock: + if host._is_removed or host.is_up is False: + with session._lock: + session._advance_pool_generation_locked(host) + stale_pools.extend( + session._pop_pools_locked(host)) + + if not register_session: + session.shutdown() + raise DriverException("Cluster is already shut down") + + try: + for pool in stale_pools: + pool.shutdown() + + # Hosts may also have become UP/ADD while the Session was being + # constructed and invisible to topology callbacks. Reconcile each + # initial host only when its own attempt completes; rescanning all + # hosts from every Future callback can fan out duplicate pool + # constructions while other initial attempts are still pending. + def reconcile_initial_host(host_id, _): + try: + session.update_created_pools( + only_host_ids=(host_id,)) + except BaseException: + # A done callback must not escape an executor thread. + # Later topology events retain the same reconciliation. + log.exception( + "Unable to reconcile pool after initial Session " + "connection attempt") + + initial_host_ids = { + id(host) + for host, _ in session._initial_connect_host_futures} + for host, future in session._initial_connect_host_futures: + future.add_done_callback( + partial(reconcile_initial_host, id(host))) + + # Initial hosts are covered by their per-Future callbacks, + # including Futures which were already complete when registered. + # This all-host pass is only for hosts which appeared while the + # Session was being constructed. + session.update_created_pools( + skip_host_ids=initial_host_ids) + except BaseException: + with self._lock: + self.sessions.discard(session) + session.shutdown() + raise return session def _session_register_user_types(self, session): @@ -1921,50 +2130,394 @@ def _session_register_user_types(self, session): for udt_name, klass in type_map.items(): session.user_type_registered(keyspace, udt_name, klass) - def _cleanup_failed_on_up_handling(self, host): - self.profile_manager.on_down(host) - self.control_connection.on_down(host) - for session in tuple(self.sessions): - session.remove_pool(host) + @staticmethod + def _reserve_host_transition_event(host): + """Return a monotonically ordered event token for one Host.""" + with host._transition_lock: + host._transition_event_sequence += 1 + return host._transition_event_sequence - self._start_reconnector(host, is_host_addition=False) + @staticmethod + def _reserve_endpoint_change_event(host): + """Reserve and publish the newest requested endpoint-change token.""" + with host._transition_lock: + host._transition_event_sequence += 1 + event_sequence = host._transition_event_sequence + host._latest_endpoint_change_sequence = event_sequence + return event_sequence + + def _run_host_transition(self, host, callback): + """ + Serialize topology side effects for one host without holding a driver + lock while callbacks execute. + + Reentrant and concurrent transitions enqueue and return. The active + runner drains them in event order, so the newest transition leaves + policies, listeners, and pools in the final host state. + """ + with host._transition_lock: + host._transition_queue.append(callback) + if host._transition_running: + return None + host._transition_running = True + host._transition_owner = get_ident() + + initial_result = None + initial_exception = None + run_notifications = False + while True: + with host._transition_lock: + if not host._transition_queue: + host._transition_running = False + host._transition_owner = None + if ( + host._transition_notification_queue and + not host._transition_notification_running): + host._transition_notification_running = True + run_notifications = True + break + work = host._transition_queue.popleft() + + owned_hosts = getattr( + _HOST_TRANSITION_CONTEXT, 'hosts', None) + if owned_hosts is None: + owned_hosts = [] + _HOST_TRANSITION_CONTEXT.hosts = owned_hosts + owned_hosts.append(host) + try: + result = work() + if isinstance(result, _HostTransitionResult): + if result.notifications: + with host._transition_lock: + host._transition_notification_queue.extend( + result.notifications) + result = result.value + if work is callback: + initial_result = result + except BaseException as exc: + if work is callback: + initial_exception = exc + else: + log.error( + "Failure in asynchronously queued host transition " + "for %s", + host, + exc_info=( + type(exc), + exc, + exc.__traceback__)) + finally: + owned_hosts.pop() + + if run_notifications: + self._drain_host_transition_notifications(host) + + if initial_exception is not None: + raise initial_exception.with_traceback( + initial_exception.__traceback__) + return initial_result + + @staticmethod + def _drain_host_transition_notifications(host): + """Drain ordered external notifications outside the core lane.""" + while True: + with host._transition_lock: + if not host._transition_notification_queue: + host._transition_notification_running = False + break + callback = \ + host._transition_notification_queue.popleft() + + try: + callback() + except BaseException as exc: + # A notification may be drained by a different transition's + # caller. Never misattribute extension errors across events. + log.error( + "Failure publishing queued host transition for %s", + host, + exc_info=( + type(exc), + exc, + exc.__traceback__)) + + def _run_host_transition_and_wait(self, host, callback): + """Run after previously queued host work before returning.""" + if getattr(_HOST_TRANSITION_CONTEXT, 'hosts', ()): + # Waiting while this thread owns any host lane can deadlock with a + # cross-host callback cycle. Queue the work in event order and + # report that it was accepted instead. + result = self._run_host_transition(host, callback) + return ( + _HOST_TRANSITION_DEFERRED + if result is None + else result) + + completed = Event() + outcome = {} + + def run_and_signal(): + try: + result = callback() + outcome['result'] = ( + result.value + if isinstance(result, _HostTransitionResult) + else result) + return result + except BaseException as exc: + outcome['error'] = exc + finally: + completed.set() + + self._run_host_transition(host, run_and_signal) + completed.wait() + error = outcome.get('error') + if error is not None: + raise error.with_traceback(error.__traceback__) + return outcome.get('result') + + def _run_host_callbacks_best_effort( + self, host, recovery_epoch, callbacks, action): + """Run every still-current extension callback, preserving one error.""" + first_error = None + for callback in callbacks: + if ( + recovery_epoch is not None and + not self._host_transition_is_current( + host, + recovery_epoch)): + break + try: + callback() + except BaseException as exc: + if first_error is None: + first_error = exc + else: + log.exception( + "Additional failure while %s host %s", + action, + host) - def _on_up_future_completed(self, host, futures, results, lock, finished_future): + if first_error is not None: + raise first_error.with_traceback(first_error.__traceback__) + return bool( + recovery_epoch is None or + self._host_transition_is_current(host, recovery_epoch)) + + @staticmethod + def _host_transition_is_current( + host, recovery_epoch, handling_attribute=None, + require_down=False): + with host.lock: + if ( + host._recovery_epoch != recovery_epoch or + host._is_removed): + return False + if ( + handling_attribute is not None and + not getattr(host, handling_attribute)): + return False + if require_down and host.is_up is not False: + return False + return True + + def _cleanup_failed_on_up_handling(self, host, recovery_epoch=None): + def is_current(): + return ( + recovery_epoch is None or + self._host_transition_is_current( + host, recovery_epoch, require_down=True)) + + first_error = None + + def run_cleanup(callback): + nonlocal first_error + if not is_current(): + return False + try: + callback() + except BaseException as exc: + if first_error is None: + first_error = exc + else: + log.exception( + "Additional failure cleaning up failed UP transition " + "for host %s", + host) + return is_current() + + continue_cleanup = run_cleanup( + lambda: self.profile_manager.on_down(host)) + if continue_cleanup: + continue_cleanup = run_cleanup( + lambda: self.control_connection.on_down(host)) + if continue_cleanup: + for session in tuple(self.sessions): + if not run_cleanup( + lambda session=session: session.remove_pool(host)): + break + + if is_current(): + run_cleanup(lambda: self._start_reconnector( + host, + is_host_addition=False, + recovery_epoch=recovery_epoch)) + + if first_error is not None: + raise first_error.with_traceback(first_error.__traceback__) + + def _on_up_future_completed( + self, host, recovery_epoch, futures, results, lock, + finished_future): with lock: futures.discard(finished_future) try: results.append(finished_future.result()) - except Exception as exc: + except BaseException as exc: + # A done callback must still finalize host recovery for + # cancellation-style BaseExceptions. The Future itself keeps + # the exception available to callers that inspect it. results.append(exc) if futures: return - try: - # all futures have completed at this point - for exc in [f for f in results if isinstance(f, Exception)]: - log.error("Unexpected failure while marking node %s up:", host, exc_info=exc) - self._cleanup_failed_on_up_handling(host) - return + completed_results = tuple(results) + + return self._run_host_transition( + host, + partial( + self._finish_on_up, + host, + recovery_epoch, + completed_results)) + + @staticmethod + def _publish_host_up( + host, recovery_epoch, handling_attribute): + """ + Reset conviction state without holding ``host.lock``, then publish UP + only if the same transition still owns the host. + + Conviction policies are application-extensible and may reenter + Cluster. Running reset under the host lock would invert the normal + Cluster -> Host lock order and could strand a failed transition if + reset raises. + """ + with host.lock: + if ( + host._recovery_epoch != recovery_epoch or + host._is_removed or + not getattr(host, handling_attribute)): + return False + + host.conviction_policy.reset() + + with host.lock: + if ( + host._recovery_epoch != recovery_epoch or + host._is_removed or + not getattr(host, handling_attribute)): + return False + setattr(host, handling_attribute, False) + host._pending_down_notification_epoch = None + if not host.is_up: + log.debug("Host %s is now marked up", host.endpoint) + host.is_up = True + return True - if not all(results): - log.debug("Connection pool could not be created, not marking node %s up", host) - self._cleanup_failed_on_up_handling(host) + def _finish_on_up( + self, host, recovery_epoch, results, + cleanup_on_failure=True): + # All futures have completed. Atomically reject a completion made stale + # by a newer DOWN transition before it can mark the host up. + failures = [result for result in results + if isinstance(result, BaseException)] + pools_created = not failures and all( + result is _SESSION_LOCAL_POOL_FAILURE or + result is _STALE_POOL_ATTEMPT or + result + for result in results) + if pools_created: + try: + published = self._publish_host_up( + host, + recovery_epoch, + '_currently_handling_node_up') + except BaseException: + if cleanup_on_failure: + with host.lock: + if ( + host._recovery_epoch == recovery_epoch and + not host._is_removed and + host._currently_handling_node_up): + host._currently_handling_node_up = False + self._cleanup_failed_on_up_handling( + host, recovery_epoch) + raise + if not published: return - log.info("Connection pools established for node %s", host) - # mark the host as up and notify all listeners - host.set_up() - for listener in self.listeners: - listener.on_up(host) - finally: + log.info( + "Connection pools established for node %s", + host) + + listener_callbacks = [ + partial(listener.on_up, host) + for listener in self.listeners] + session_callbacks = [ + session.update_created_pools + for session in tuple(self.sessions)] + # Keep deferred callbacks freshness-fenced so a queued DOWN or + # REMOVE cannot publish a stale UP event after its core transition. + # A listener which has already begun remains free to finish while + # later core cleanup proceeds on the released transition lane. + try: + self._run_host_callbacks_best_effort( + host, + recovery_epoch, + session_callbacks, + "reconciling pools after UP for") + except BaseException: + # Session repair is internal bookkeeping. A failure in one + # Session must not suppress the ordered external UP event. + log.exception( + "Unable to reconcile pools after UP for host %s", + host) + + return _HostTransitionResult( + notifications=( + partial( + self._run_host_callbacks_best_effort, + host, + recovery_epoch, + listener_callbacks, + "publishing UP for"), + )) + else: + if failures: + for exc in failures: + log.error( + "Unexpected failure while marking node %s up:", + host, + exc_info=exc) + else: + log.debug( + "Connection pool could not be created, not marking node " + "%s up", + host) + with host.lock: + if ( + host._recovery_epoch != recovery_epoch or + host._is_removed or + not host._currently_handling_node_up): + return host._currently_handling_node_up = False - - # see if there are any pools to add or remove now that the host is marked up - for session in tuple(self.sessions): - session.update_created_pools() + self._cleanup_failed_on_up_handling( + host, recovery_epoch) + return def on_up(self, host): """ @@ -1972,18 +2525,191 @@ def on_up(self, host): """ if self.is_shutdown or self.allow_control_connection_query_fallback == ControlConnectionQueryFallback.SkipPoolCreation: return - - log.debug("Waiting to acquire lock for handling up status of node %s", host) + forced_event = getattr( + _HOST_TRANSITION_CONTEXT, 'forced_event', None) + if forced_event is not None and forced_event[0] is host: + event_sequence = forced_event[1] + else: + event_sequence = self._reserve_host_transition_event(host) + return self._run_host_transition( + host, + partial(self._begin_up, host, event_sequence)) + + def _begin_up(self, host, event_sequence=None): + """Claim and execute UP from inside the host event lane.""" + log.debug( + "Waiting to acquire lock for handling up status of node %s", + host) + add_recovery_epoch = None with host.lock: - if host._currently_handling_node_up: - log.debug("Another thread is already handling up status of node %s", host) + if host._is_removed: return - - if host.is_up: + if ( + event_sequence is not None and + host._latest_preemptive_down_sequence > + event_sequence): + return + if host._currently_handling_node_add: + # NEW_NODE establishes topology membership before a STATUS + # event changes liveness. Do not invalidate an ADD while its + # policy/control callbacks are still in the serialized lane. + add_recovery_epoch = host._recovery_epoch + elif host._currently_handling_node_up: + log.debug( + "Another thread is already handling up status of node %s", + host) + return + elif host.is_up: log.debug("Host %s was already marked up", host) return + else: + # Until the UP attempt succeeds, an unknown host must have + # the same recoverable state as a known-down host. Failed-UP + # cleanup and reconnector installation both require an exact + # False value. + if host.is_up is None: + host.set_down() + + host._recovery_epoch += 1 + recovery_epoch = host._recovery_epoch + host._currently_handling_node_add = False + host._currently_handling_node_up = True + reconnector = host.get_and_set_reconnection_handler(None) + + if add_recovery_epoch is not None: + return self._complete_add_before_status( + host, + add_recovery_epoch, + partial(self._begin_up, host, event_sequence)) + + if reconnector: + log.debug( + "Now that host %s is up, cancelling the reconnection handler", + host) + try: + reconnector.cancel() + except BaseException: + log.exception( + "Unable to cancel old recovery while marking host %s up", + host) + + return self._on_up_locked(host, recovery_epoch) + + def _complete_add_before_status( + self, host, add_recovery_epoch, status_callback): + """ + Preserve NEW_NODE ordering when a STATUS event races pool creation. - host._currently_handling_node_up = True + The ADD's policy and control-connection work runs first. If its pool + Future is still pending, fence that attempt without claiming success, + publish the topology notification after the lane is released, and + immediately enqueue/apply the newer liveness transition. This keeps a + slow listener from blocking authoritative cleanup and prevents a + failed ADD pool attempt from leaving the host falsely marked UP. + """ + with host.lock: + if ( + host._recovery_epoch != add_recovery_epoch or + host._is_removed or + not host._currently_handling_node_add): + add_is_pending = False + else: + add_is_pending = True + host._currently_handling_node_add = False + host._recovery_epoch += 1 + + if not add_is_pending: + return status_callback() + + notifications = [ + partial( + self._run_host_callbacks_best_effort, + host, + None, + [ + partial(listener.on_add, host) + for listener in self.listeners + ], + "publishing superseded ADD for")] + try: + status_result = status_callback() + except BaseException as exc: + def raise_status_error(error=exc): + raise error.with_traceback(error.__traceback__) + + notifications.append(raise_status_error) + status_result = None + + if isinstance(status_result, _HostTransitionResult): + notifications.extend(status_result.notifications) + status_result = status_result.value + + return _HostTransitionResult( + status_result, + notifications) + + def _complete_add_before_down( + self, host, add_recovery_epoch, is_host_addition, + expect_host_to_be_down, force, event_sequence=None): + """ + Decide DOWN in the ADD lane, then leave teardown in the executor. + + An open-pool discount must leave the ADD attempt intact. An accepted + DOWN fences it and publishes ADD before DOWN through the independent + notification FIFO. + """ + with host.lock: + if ( + host._recovery_epoch != add_recovery_epoch or + host._is_removed or + not host._currently_handling_node_add): + add_is_pending = False + else: + add_is_pending = True + + if not add_is_pending: + return self._on_down( + host, + is_host_addition, + expect_host_to_be_down, + force, + event_sequence=event_sequence) + + with self._lock: + recovery_epoch = self._on_down_locked( + host, + is_host_addition, + expect_host_to_be_down, + force, + defer_cleanup=True, + event_sequence=event_sequence) + + if recovery_epoch is None: + return + + self.on_down_potentially_blocking( + host, + is_host_addition, + recovery_epoch) + return _HostTransitionResult( + recovery_epoch, + notifications=( + partial( + self._run_host_callbacks_best_effort, + host, + None, + [ + partial(listener.on_add, host) + for listener in self.listeners + ], + "publishing superseded ADD for"),)) + + def _on_up_locked(self, host, recovery_epoch): + if not self._host_transition_is_current( + host, + recovery_epoch, + '_currently_handling_node_up'): + return log.debug("Starting to handle up status of node %s", host) have_future = False @@ -1991,54 +2717,97 @@ def on_up(self, host): try: log.info("Host %s may be up; will prepare queries and open connection pool", host) - reconnector = host.get_and_set_reconnection_handler(None) - if reconnector: - log.debug("Now that host %s is up, cancelling the reconnection handler", host) - reconnector.cancel() - if self.profile_manager.distance(host) != HostDistance.IGNORED: self._prepare_all_queries(host) + if not self._host_transition_is_current( + host, + recovery_epoch, + '_currently_handling_node_up'): + return futures log.debug("Done preparing all queries for host %s, ", host) for session in tuple(self.sessions): + if not self._host_transition_is_current( + host, + recovery_epoch, + '_currently_handling_node_up'): + return futures session.remove_pool(host) + if not self._host_transition_is_current( + host, + recovery_epoch, + '_currently_handling_node_up'): + return futures log.debug("Signalling to load balancing policies that host %s is up", host) self.profile_manager.on_up(host) + if not self._host_transition_is_current( + host, + recovery_epoch, + '_currently_handling_node_up'): + return futures log.debug("Signalling to control connection that host %s is up", host) self.control_connection.on_up(host) + if not self._host_transition_is_current( + host, + recovery_epoch, + '_currently_handling_node_up'): + return futures log.debug("Attempting to open new connection pools for host %s", host) futures_lock = Lock() futures_results = [] - callback = partial(self._on_up_future_completed, host, futures, futures_results, futures_lock) + callback = partial( + self._on_up_future_completed, + host, + recovery_epoch, + futures, + futures_results, + futures_lock) for session in tuple(self.sessions): - future = session.add_or_renew_pool(host, is_host_addition=False) + if not self._host_transition_is_current( + host, + recovery_epoch, + '_currently_handling_node_up'): + return futures + future = session.add_or_renew_pool( + host, + is_host_addition=False, + recovery_epoch=recovery_epoch) if future is not None: have_future = True - future.add_done_callback(callback) futures.add(future) - except Exception: + for future in tuple(futures): + future.add_done_callback(callback) + if not have_future: + return self._finish_on_up( + host, + recovery_epoch, + (), + cleanup_on_failure=False) + except BaseException: log.exception("Unexpected failure handling node %s being marked up:", host) for future in futures: future.cancel() - self._cleanup_failed_on_up_handling(host) - with host.lock: + if ( + host._recovery_epoch != recovery_epoch or + host._is_removed): + raise host._currently_handling_node_up = False + self._cleanup_failed_on_up_handling( + host, recovery_epoch) raise - else: - if not have_future: - with host.lock: - host.set_up() - host._currently_handling_node_up = False # for testing purposes return futures - def _start_reconnector(self, host, is_host_addition): + def _start_reconnector( + self, host, is_host_addition, recovery_epoch=None): + if self.is_shutdown: + return if self.profile_manager.distance(host) == HostDistance.IGNORED: return @@ -2049,44 +2818,267 @@ def _start_reconnector(self, host, is_host_addition): # of the current Cluster attributes to create new Connections with conn_factory = self._make_connection_factory(host) + reconnector_holder = [] + + def clear_reconnector(): + if reconnector_holder: + host.clear_reconnection_handler( + reconnector_holder[0]) + 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) + self.scheduler, schedule, clear_reconnector) + reconnector_holder.append(reconnector) - old_reconnector = host.get_and_set_reconnection_handler(reconnector) + with host.lock: + if ( + self.is_shutdown or + host._is_removed or + host.is_up is not False or + (recovery_epoch is not None and + host._recovery_epoch != recovery_epoch)): + return + old_reconnector = host.get_and_set_reconnection_handler( + reconnector) if old_reconnector: log.debug("Old host reconnector found for %s, cancelling", host) - old_reconnector.cancel() + try: + old_reconnector.cancel() + except BaseException: + log.exception( + "Unable to cancel superseded recovery for host %s", + host) log.debug("Starting reconnector for host %s", host) - reconnector.start() + try: + reconnector.start() + except BaseException: + host.clear_reconnection_handler(reconnector) + raise @run_in_executor - def on_down_potentially_blocking(self, host, is_host_addition): - self.profile_manager.on_down(host) - self.control_connection.on_down(host) - for session in tuple(self.sessions): - session.on_down(host) - - for listener in self.listeners: - listener.on_down(host) + def on_down_potentially_blocking( + self, host, is_host_addition, recovery_epoch=None): + return self._run_host_transition( + host, + partial( + self._on_down_potentially_blocking_serialized, + host, + is_host_addition, + recovery_epoch)) + + def _on_down_potentially_blocking_serialized( + self, host, is_host_addition, recovery_epoch=None, + start_reconnector=True, + down_notification_epoch=_NOT_SET, + force_down_notification=False): + if down_notification_epoch is _NOT_SET: + down_notification_epoch = recovery_epoch + + def is_current(): + return ( + recovery_epoch is None or + self._host_transition_is_current( + host, recovery_epoch, require_down=True)) + + first_error = None + + def run_cleanup(callback): + nonlocal first_error + if not is_current(): + return False + try: + callback() + except BaseException as exc: + if first_error is None: + first_error = exc + else: + log.exception( + "Additional failure handling DOWN transition for " + "host %s", + host) + return is_current() + + continue_cleanup = run_cleanup( + lambda: self.profile_manager.on_down(host)) + if continue_cleanup: + continue_cleanup = run_cleanup( + lambda: self.control_connection.on_down(host)) + if continue_cleanup: + for session in tuple(self.sessions): + if not run_cleanup( + lambda session=session: session.on_down(host)): + break + if start_reconnector and is_current(): + run_cleanup(lambda: self._start_reconnector( + host, + is_host_addition, + recovery_epoch=recovery_epoch)) + + listener_callbacks = [ + partial(listener.on_down, host) + for listener in self.listeners] + if force_down_notification: + notifications = [ + partial( + self._publish_authoritative_down_notification, + host, + down_notification_epoch, + listener_callbacks, + )] + else: + notifications = [ + partial( + self._publish_pending_down_notification, + host, + down_notification_epoch, + listener_callbacks)] + if first_error is not None: + log.error( + "Failure handling DOWN transition for host %s", + host, + exc_info=( + type(first_error), + first_error, + first_error.__traceback__)) + return _HostTransitionResult( + value=first_error, + notifications=notifications) + + def _publish_pending_down_notification( + self, host, recovery_epoch, callbacks): + """ + Claim and publish the listener event for an accepted DOWN transition. + + A pending UP may supersede internal DOWN cleanup, but it does not + invalidate this external event until the host is actually published + UP. This lets a failed superseding UP retain exactly one DOWN + notification while successful recovery suppresses the stale event. + """ + with host.lock: + if ( + host._is_removed or + host.is_up is not False or + host._pending_down_notification_epoch != + recovery_epoch): + return False + host._pending_down_notification_epoch = None - self._start_reconnector(host, is_host_addition) + return self._run_host_callbacks_best_effort( + host, + None, + callbacks, + "publishing DOWN for") - def on_down(self, host, is_host_addition, expect_host_to_be_down=False): + def _publish_authoritative_down_notification( + self, host, recovery_epoch, callbacks): + """Publish a relocation DOWN even if its immediate UP already won.""" + with host.lock: + if ( + host._pending_down_notification_epoch == + recovery_epoch): + host._pending_down_notification_epoch = None + return self._run_host_callbacks_best_effort( + host, + None, + callbacks, + "publishing authoritative DOWN for") + + def on_down(self, host, is_host_addition, expect_host_to_be_down=False, + force=False): """ Intended for internal use only. """ - if self.is_shutdown or self.allow_control_connection_query_fallback == ControlConnectionQueryFallback.SkipPoolCreation: + forced_event = getattr( + _HOST_TRANSITION_CONTEXT, 'forced_event', None) + if forced_event is not None and forced_event[0] is host: + event_sequence = forced_event[1] + else: + event_sequence = self._reserve_host_transition_event(host) + if ( + forced_event is not None and + forced_event[0] is host and + getattr( + _HOST_TRANSITION_CONTEXT, + 'suppress_relocation_down_host', + None) is host): + # A custom relocation hook may delegate to ``super().on_down``. + # The relocation callback performs the authoritative base DOWN + # itself, so do not enqueue a second teardown behind it. + return + return self._run_host_transition( + host, + partial( + self._begin_down, + host, + is_host_addition, + expect_host_to_be_down, + force, + event_sequence)) + + def _begin_down( + self, host, is_host_addition, + expect_host_to_be_down=False, force=False, + event_sequence=None): + """Claim and schedule DOWN from inside the host event lane.""" + with host.lock: + add_recovery_epoch = ( + host._recovery_epoch + if host._currently_handling_node_add and not force + else None) + if add_recovery_epoch is not None: + return self._complete_add_before_down( + host, + add_recovery_epoch, + is_host_addition, + expect_host_to_be_down, + force, + event_sequence) + return self._on_down( + host, + is_host_addition, + expect_host_to_be_down, + force=force, + event_sequence=event_sequence) + + def _on_down(self, host, is_host_addition, expect_host_to_be_down=False, + force=False, event_sequence=None): + """ + ``force`` bypasses the open-pool discount when the caller has + authoritative evidence that the host cannot accept new CQL + connections. + """ + # Pool-construction failures hold their Session lock across this + # transition. Serializing transitions before taking host/session locks + # prevents two failures for different hosts from taking the session + # locks in opposite orders. + with self._lock: + return self._on_down_locked( + host, + is_host_addition, + expect_host_to_be_down, + force, + event_sequence=event_sequence) + + def _on_down_locked( + self, host, is_host_addition, + expect_host_to_be_down=False, force=False, + defer_cleanup=False, allow_when_skipping_pool_creation=False, + event_sequence=None): + if self.is_shutdown or ( + self.allow_control_connection_query_fallback == + ControlConnectionQueryFallback.SkipPoolCreation and + not allow_when_skipping_pool_creation): return - with host.lock: + if host._is_removed: + return was_up = host.is_up # ignore down signals if we have open pools to the host # this is to avoid closing pools when a control connection host became isolated - if self._discount_down_events and self.profile_manager.distance(host) != HostDistance.IGNORED: + if not force and self._discount_down_events and \ + self.profile_manager.distance(host) != HostDistance.IGNORED: connected = False for session in tuple(self.sessions): pool_states = session.get_pool_state() @@ -2096,32 +3088,415 @@ def on_down(self, host, is_host_addition, expect_host_to_be_down=False): if connected: return - host.set_down() - if (not was_up and not expect_host_to_be_down) or host.is_currently_reconnecting(): + recovery_in_progress = ( + host._currently_handling_node_up or + host._currently_handling_node_add) + + if event_sequence is None: + event_sequence = \ + self._reserve_host_transition_event(host) + host._latest_preemptive_down_sequence = max( + host._latest_preemptive_down_sequence, + event_sequence) + + # A duplicate DOWN must not invalidate the epoch of the worker + # already queued for the original transition. Recovery attempts + # are the exception: this DOWN supersedes and fences them. + if ( + (was_up is False and + not expect_host_to_be_down and + not recovery_in_progress) or + (host.is_currently_reconnecting() and + not recovery_in_progress)): return + + host.set_down() + host._recovery_epoch += 1 + host._currently_handling_node_up = False + host._currently_handling_node_add = False + recovery_epoch = host._recovery_epoch + host._pending_down_notification_epoch = recovery_epoch + + # Fence pool creations that started before this down transition + # before an on_up transition can start new ones. The blocking + # teardown still runs in the executor below. + for session in tuple(self.sessions): + try: + session._invalidate_pool_attempts(host) + except BaseException: + # The serialized session.on_down cleanup below repeats the + # invalidation. One broken Session must not prevent the + # cleanup/reconnector handoff for the Host. + log.exception( + "Unable to invalidate a pool attempt while marking " + "host %s down", + host) + log.warning("Host %s has been marked down", host) - self.on_down_potentially_blocking(host, is_host_addition) + if defer_cleanup: + return recovery_epoch + + self.on_down_potentially_blocking( + host, is_host_addition, recovery_epoch) + return recovery_epoch + + def _force_down_for_endpoint_change(self, host, endpoint): + """ + Remove every old-endpoint consumer before mutating ``Host.endpoint``. + + Host hashes by endpoint. Running the DOWN callbacks, endpoint mutation, + metadata reindex, and UP scheduling as one serialized transition + prevents hashed policy state from retaining an entry under the old + hash while the control connection publishes the new endpoint. + """ + with self._lock: + if self.is_shutdown: + return False + with host.lock: + if host._is_removed: + return False + + event_sequence = self._reserve_endpoint_change_event(host) + + def endpoint_change_is_current(): + with host._transition_lock: + return ( + host._latest_endpoint_change_sequence == + event_sequence) + + def relocate(): + if not endpoint_change_is_current(): + return False + + with self._lock: + if self.is_shutdown: + return False + with host.lock: + if host._is_removed: + return False + if host.endpoint == endpoint: + return True + + if type(self).on_down is not Cluster.on_down: + previous_forced_event = getattr( + _HOST_TRANSITION_CONTEXT, 'forced_event', None) + previous_suppressed_host = getattr( + _HOST_TRANSITION_CONTEXT, + 'suppress_relocation_down_host', + None) + _HOST_TRANSITION_CONTEXT.forced_event = \ + (host, event_sequence) + _HOST_TRANSITION_CONTEXT.suppress_relocation_down_host = host + try: + # Preserve the historical three-argument extension hook, + # but run it only after this endpoint event owns the lane. + self.on_down(host, False, True) + except BaseException: + # Endpoint identity is authoritative metadata. An extension + # failure cannot prevent the base cleanup and relocation. + log.exception( + "Custom DOWN hook failed while relocating host %s", + host) + finally: + _HOST_TRANSITION_CONTEXT.forced_event = \ + previous_forced_event + _HOST_TRANSITION_CONTEXT.suppress_relocation_down_host = \ + previous_suppressed_host + + if not endpoint_change_is_current(): + return False + + with self._lock: + if self.is_shutdown: + return False + with host.lock: + if host._is_removed: + return False + # Claim authoritative DOWN only after this relocation has + # entered the host lane. A later UP can no longer stale + # old-hash cleanup before the endpoint mutation. + reconnector = host.get_and_set_reconnection_handler(None) + recovery_epoch = Cluster._on_down_locked( + self, + host, + is_host_addition=False, + expect_host_to_be_down=True, + force=True, + defer_cleanup=True, + allow_when_skipping_pool_creation=True, + event_sequence=event_sequence) + + if reconnector: + try: + reconnector.cancel() + except BaseException: + log.exception( + "Unable to cancel old-endpoint recovery for host %s", + host) + + if recovery_epoch is None: + return False + + down_result = None + cleanup_error = None + try: + # Relocation owns the host lane and is authoritative. Do not + # freshness-fence cleanup or its DOWN notification against a + # concurrently announced newer status. + down_result = \ + self._on_down_potentially_blocking_serialized( + host, + False, + recovery_epoch=None, + start_reconnector=False, + down_notification_epoch=recovery_epoch, + force_down_notification=True) + if ( + isinstance(down_result, _HostTransitionResult) and + isinstance(down_result.value, BaseException)): + cleanup_error = down_result.value + except BaseException as exc: + cleanup_error = exc + # Endpoint identity is authoritative metadata. Extension + # cleanup failures cannot leave Host hashed under an endpoint + # metadata has already replaced. + log.exception( + "DOWN cleanup failed while relocating host %s", + host) + + notifications = list( + down_result.notifications + if isinstance(down_result, _HostTransitionResult) + else ()) + if cleanup_error is not None: + # Mutating a Host hash after any policy/session cleanup failed + # can make the old entry unreachable. Keep the old identity and + # restore an epoch-fenced recovery path; a later metadata + # refresh can retry the relocation safely. + try: + self._start_reconnector( + host, + is_host_addition=False, + recovery_epoch=recovery_epoch) + except BaseException: + log.exception( + "Unable to start recovery after endpoint cleanup " + "failed for host %s", + host) + return _HostTransitionResult( + False, + notifications=notifications) + + if not endpoint_change_is_current(): + return _HostTransitionResult( + False, + notifications=notifications) + + # Hold the Host and real Metadata host-map locks through + # reindexing. Metadata.remove_host() removes the map entry before + # Cluster.on_remove() can mark Host._is_removed; checking both + # closes that otherwise-resurrecting gap. + with self._lock: + if self.is_shutdown: + return False + with host.lock: + if host._is_removed: + return False + with host._transition_lock: + if ( + host._latest_endpoint_change_sequence != + event_sequence): + return _HostTransitionResult( + False, + notifications=notifications) + metadata_hosts = getattr( + self.metadata, '_hosts', None) + metadata_lock = getattr( + self.metadata, '_hosts_lock', None) + if ( + isinstance(metadata_hosts, dict) and + metadata_lock is not None): + with metadata_lock: + if metadata_hosts.get( + host.host_id) is not host: + return False + old_endpoint = host.endpoint + host.endpoint = endpoint + self.metadata.update_host(host, old_endpoint) + else: + old_endpoint = host.endpoint + host.endpoint = endpoint + self.metadata.update_host(host, old_endpoint) + + # Restore UP as part of the relocation event itself. A later + # queued DOWN must run after this work and leave the final state + # DOWN, rather than being overtaken by a synthetic queued UP. + if type(self).on_up is not Cluster.on_up: + previous_forced_event = getattr( + _HOST_TRANSITION_CONTEXT, 'forced_event', None) + _HOST_TRANSITION_CONTEXT.forced_event = \ + (host, event_sequence) + try: + self.on_up(host) + except BaseException: + log.exception( + "Custom UP hook failed after relocating host %s", + host) + finally: + _HOST_TRANSITION_CONTEXT.forced_event = \ + previous_forced_event + + try: + up_result = Cluster._begin_up( + self, + host, + event_sequence) + except BaseException: + log.exception( + "UP recovery failed after relocating host %s", + host) + up_result = None + + if isinstance(up_result, _HostTransitionResult): + notifications.extend(up_result.notifications) + return _HostTransitionResult( + True, + notifications=notifications) + + result = self._run_host_transition_and_wait(host, relocate) + return result def on_add(self, host, refresh_nodes=True): if self.is_shutdown: return + event_sequence = self._reserve_host_transition_event(host) + return self._run_host_transition( + host, + partial( + self._begin_add, + host, + refresh_nodes, + event_sequence)) + + def _begin_add( + self, host, refresh_nodes=True, event_sequence=None): + """Claim ADD state from inside the host's serialized event lane.""" + with host.lock: + if host._is_removed: + return + if ( + event_sequence is not None and + host._latest_preemptive_down_sequence > + event_sequence): + return + host._recovery_epoch += 1 + recovery_epoch = host._recovery_epoch + host._currently_handling_node_up = False + host._currently_handling_node_add = True + reconnector = host.get_and_set_reconnection_handler(None) + if reconnector: + try: + reconnector.cancel() + except BaseException: + log.exception( + "Unable to cancel old recovery while adding host %s", + host) + + return self._on_add_locked( + host, + refresh_nodes, + recovery_epoch) + + def _on_add_locked( + self, host, refresh_nodes=True, recovery_epoch=None): + try: + return self._on_add_locked_impl( + host, refresh_nodes, recovery_epoch) + except BaseException: + # on_add cancels any existing reconnector before this work is + # queued. If a policy, control callback, or pool setup then + # fails, transition authoritatively to DOWN so the host is not + # stranded in a half-added state without recovery. + self._recover_failed_add(host, recovery_epoch) + raise + + def _recover_failed_add(self, host, recovery_epoch): + """ + Convert a still-current failed ADD into authoritative DOWN recovery. + + Legacy ``on_down`` overrides are notified through their historical + positional API. A no-op or raising override cannot strand the host: + the base transition runs in ``finally`` if the ADD epoch remains + current. + """ + if not self._host_transition_is_current( + host, + recovery_epoch, + '_currently_handling_node_add'): + return False + + if type(self).on_down is Cluster.on_down: + self._on_down( + host, + is_host_addition=True, + expect_host_to_be_down=True, + force=True) + return True + + try: + self.on_down(host, True, True) + finally: + if self._host_transition_is_current( + host, + recovery_epoch, + '_currently_handling_node_add'): + Cluster._on_down( + self, + host, + is_host_addition=True, + expect_host_to_be_down=True, + force=True) + return True + + def _on_add_locked_impl( + self, host, refresh_nodes=True, recovery_epoch=None): + if not self._host_transition_is_current( + host, + recovery_epoch, + '_currently_handling_node_add'): + return log.debug("Handling new host %r and notifying listeners", host) self.profile_manager.on_add(host) + if not self._host_transition_is_current( + host, + recovery_epoch, + '_currently_handling_node_add'): + return self.control_connection.on_add(host, refresh_nodes) + if not self._host_transition_is_current( + host, + recovery_epoch, + '_currently_handling_node_add'): + return distance = self.profile_manager.distance(host) if distance != HostDistance.IGNORED: self._prepare_all_queries(host) + if not self._host_transition_is_current( + host, + recovery_epoch, + '_currently_handling_node_add'): + return log.debug("Done preparing queries for new host %r", host) 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) - return + return self._finalize_add( + host, set_up=False, recovery_epoch=recovery_epoch) futures_lock = Lock() futures_results = [] @@ -2133,69 +3508,272 @@ def future_completed(future): try: futures_results.append(future.result()) - except Exception as exc: + except BaseException as exc: + # Keep cleanup/reconnection alive for executor + # cancellations and other BaseException subclasses. futures_results.append(exc) if futures: return - log.debug('All futures have completed for added host %s', host) + completed_results = tuple(futures_results) - 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): - log.warning("Connection pool could not be created, not marking node %s up", host) - return - - self._finalize_add(host) + self._run_host_transition( + host, + partial( + self._finish_add, + host, + recovery_epoch, + completed_results)) have_future = False for session in tuple(self.sessions): - future = session.add_or_renew_pool(host, is_host_addition=True) + if not self._host_transition_is_current( + host, + recovery_epoch, + '_currently_handling_node_add'): + return futures + future = session.add_or_renew_pool( + host, + is_host_addition=True, + recovery_epoch=recovery_epoch) if future is not None: have_future = True futures.add(future) - future.add_done_callback(future_completed) + for future in tuple(futures): + future.add_done_callback(future_completed) if not have_future: - self._finalize_add(host) + return self._finalize_add( + host, recovery_epoch=recovery_epoch) - def _finalize_add(self, host, set_up=True): - if set_up: - host.set_up() + def _finish_add(self, host, recovery_epoch, results): + log.debug('All futures have completed for added host %s', host) - for listener in self.listeners: - listener.on_add(host) + if not self._host_transition_is_current( + host, + recovery_epoch, + '_currently_handling_node_add'): + return - # see if there are any pools to add or remove now that the host is marked up - for session in tuple(self.sessions): - session.update_created_pools() + failures = [ + result for result in results + if isinstance(result, BaseException)] + for exc in failures: + log.error( + "Unexpected failure while adding node %s, will not mark up:", + host, + exc_info=exc) + if failures: + self._recover_failed_add(host, recovery_epoch) + return + + if not all( + result is _SESSION_LOCAL_POOL_FAILURE or + result is _STALE_POOL_ATTEMPT or + result + for result in results): + log.warning( + "Connection pool could not be created, not marking node %s up", + host) + self._recover_failed_add(host, recovery_epoch) + return + + return self._finalize_add( + host, recovery_epoch=recovery_epoch) + + def _finalize_add(self, host, set_up=True, recovery_epoch=None): + if set_up: + try: + if not self._publish_host_up( + host, + recovery_epoch, + '_currently_handling_node_add'): + return False + except BaseException: + self._recover_failed_add(host, recovery_epoch) + raise + else: + with host.lock: + if ( + host._is_removed or + (recovery_epoch is not None and + host._recovery_epoch != recovery_epoch) or + not host._currently_handling_node_add): + return False + host._currently_handling_node_add = False + host._pending_down_notification_epoch = None + + listener_callbacks = [ + partial(listener.on_add, host) + for listener in self.listeners] + session_callbacks = [ + session.update_created_pools + for session in tuple(self.sessions)] + try: + self._run_host_callbacks_best_effort( + host, + recovery_epoch, + session_callbacks, + "reconciling pools after ADD for") + except BaseException: + log.exception( + "Unable to reconcile pools after ADD for host %s", + host) + + return _HostTransitionResult( + True, + notifications=( + partial( + self._run_host_callbacks_best_effort, + host, + None, + listener_callbacks, + "publishing ADD for"), + )) def on_remove(self, host): + if self.is_shutdown: + return + with host.lock: + if host._is_removed: + return + host._recovery_epoch += 1 + recovery_epoch = host._recovery_epoch + host._currently_handling_node_up = False + host._currently_handling_node_add = False + host._is_removed = True + host.set_down() + host._pending_down_notification_epoch = None + for session in tuple(self.sessions): + try: + session._invalidate_pool_attempts(host) + except BaseException: + # _on_remove_locked repeats Session cleanup. Never let one + # invalidation failure suppress removal from every other + # driver component. + log.exception( + "Unable to invalidate a pool attempt while removing " + "host %s", + host) + reconnection_handler = host.get_and_set_reconnection_handler(None) + if reconnection_handler: + try: + reconnection_handler.cancel() + except BaseException: + log.exception( + "Unable to cancel recovery while removing host %s", + host) + + return self._run_host_transition( + host, + partial(self._on_remove_locked, host, recovery_epoch)) + + def _on_remove_locked(self, host, recovery_epoch=None): + with host.lock: + if ( + recovery_epoch is not None and + host._recovery_epoch != recovery_epoch): + return + if not host._is_removed: + return + if self.is_shutdown: return log.debug("[cluster] Removing host %s", host) - host.set_down() - self.profile_manager.on_remove(host) - for session in tuple(self.sessions): - session.on_remove(host) - for listener in self.listeners: - listener.on_remove(host) - self.control_connection.on_remove(host) + first_error = None - reconnection_handler = host.get_and_set_reconnection_handler(None) - if reconnection_handler: - reconnection_handler.cancel() + def run_cleanup(callback): + nonlocal first_error + try: + callback() + except BaseException as exc: + if first_error is None: + first_error = exc + else: + log.exception( + "Additional failure cleaning up removed host %s", + host) - def signal_connection_failure(self, host, connection_exc, is_host_addition, expect_host_to_be_down=False): - is_down = host.signal_connection_failure(connection_exc) - if is_down: - self.on_down(host, is_host_addition, expect_host_to_be_down) + run_cleanup(lambda: self.profile_manager.on_remove(host)) + for session in tuple(self.sessions): + run_cleanup(lambda session=session: session.on_remove(host)) + run_cleanup(lambda: self.control_connection.on_remove(host)) + + notifications = [ + partial( + self._run_host_callbacks_best_effort, + host, + None, + [ + partial(listener.on_remove, host) + for listener in self.listeners + ], + "publishing REMOVE for")] + if first_error is not None: + log.error( + "Failure cleaning up removed host %s", + host, + exc_info=( + type(first_error), + first_error, + first_error.__traceback__)) + return _HostTransitionResult( + notifications=notifications) + + def signal_connection_failure(self, host, connection_exc, is_host_addition, + expect_host_to_be_down=False, force=False): + # ``force`` is reserved for authoritative failures that cannot leave a + # usable pool, and bypasses both conviction and open-pool discounting. + try: + is_down = host.signal_connection_failure(connection_exc) + except BaseException: + if not force: + raise + is_down = False + log.exception( + "Conviction policy failed for authoritative connection " + "failure on host %s; continuing host recovery", + host) + if force: + if type(self).on_down is Cluster.on_down: + if isinstance(host, Host): + self.on_down( + host, + is_host_addition, + expect_host_to_be_down, + force=True) + else: + # Preserve the long-standing duck-typed private hook used + # by tests and third-party Cluster stand-ins. + self._on_down( + host, + is_host_addition, + expect_host_to_be_down, + force=True) + else: + # Legacy subclasses commonly override the historical + # three-argument on_down hook. Preserve that notification + # without passing the newer force keyword. + self.on_down( + host, + is_host_addition, + expect_host_to_be_down) + elif is_down: + self.on_down( + host, + is_host_addition, + expect_host_to_be_down) return is_down + def _uses_default_failure_hooks(self): + """Whether private fenced DOWN commits preserve public hook behavior.""" + return bool( + type(self).signal_connection_failure is + Cluster.signal_connection_failure and + type(self).on_down is Cluster.on_down) + def add_host(self, endpoint, datacenter=None, rack=None, signal=True, refresh_nodes=True, host_id=None): """ Called when adding initial contact points and when the control @@ -2278,8 +3856,8 @@ def get_control_connection_host(self): Returns the control connection host metadata. """ connection = self.control_connection._connection - endpoint = connection.endpoint if connection else None - return self.metadata.get_host(endpoint) if endpoint else None + return self.control_connection._get_host_for_connection( + connection, require_current=True) def refresh_schema_metadata(self, max_schema_agreement_wait=None): """ @@ -2412,6 +3990,8 @@ def _prepare_all_queries(self, host): connection = None try: connection = self.connection_factory(host.endpoint) + if connection.is_closed: + raise _startup_close_error(connection, host.endpoint) statements = list(self._prepared_statements.values()) if ProtocolVersion.uses_keyspace_flag(self.protocol_version): # V5 protocol and higher, no need to set the keyspace @@ -2422,7 +4002,10 @@ def _prepare_all_queries(self, host): else: for keyspace, ks_statements in groupby(statements, lambda s: s.keyspace): if keyspace is not None: - connection.set_keyspace_blocking(keyspace) + _set_keyspace_blocking( + connection, + keyspace, + self.connect_timeout) # prepare 10 statements at a time ks_statements = list(ks_statements) @@ -2446,6 +4029,30 @@ def add_prepared(self, query_id, prepared_statement): with self._prepared_statement_lock: self._prepared_statements[query_id] = prepared_statement +class _SessionLocalPoolFailure(object): + """A pool failed for session-local state, not host health.""" + + def __bool__(self): + return False + + __nonzero__ = __bool__ + + +_SESSION_LOCAL_POOL_FAILURE = _SessionLocalPoolFailure() + + +class _StalePoolAttempt(object): + """A pool worker lost ownership without learning anything about the host.""" + + def __bool__(self): + return False + + __nonzero__ = __bool__ + + +_STALE_POOL_ATTEMPT = _StalePoolAttempt() + + class Session(object): """ A collection of connection pools for each host in the cluster. @@ -2662,6 +4269,7 @@ def default_serial_consistency_level(self, cl): _lock = None _pools = None + _pool_generations = None _profile_manager = None _metrics = None _request_init_callbacks = None @@ -2673,7 +4281,17 @@ def __init__(self, cluster, hosts, keyspace=None): self.keyspace = keyspace self._lock = RLock() + self._keyspace_dispatch_lock = RLock() + self._keyspace_completion_lock = Lock() + self._keyspace_completion_queue = [] + self._keyspace_completion_runner_active = False + self._keyspace_generation = 0 + self._pool_repair_schedule = None + self._pool_repair_scheduled = False self._pools = {} + # Host.__hash__ follows its mutable endpoint, so generation entries + # use weak identity keys maintained by the helpers below. + self._pool_generations = {} self._profile_manager = cluster.profile_manager self._metrics = cluster.metrics self._request_init_callbacks = [] @@ -2683,54 +4301,93 @@ def __init__(self, cluster, hosts, keyspace=None): # create connection pools in parallel self._initial_connect_futures = set() + self._initial_connect_host_futures = [] fallback_mode = self.cluster.allow_control_connection_query_fallback - if fallback_mode is not ControlConnectionQueryFallback.SkipPoolCreation: - for host in hosts: - future = self.add_or_renew_pool(host, is_host_addition=False) - if future: - self._initial_connect_futures.add(future) - - futures = wait_futures(self._initial_connect_futures, return_when=FIRST_COMPLETED) - while futures.not_done and not any(f.result() for f in futures.done): - futures = wait_futures(futures.not_done, return_when=FIRST_COMPLETED) - - # Only Disabled requires an initial pool to come up. - if not any(f.result() for f in self._initial_connect_futures) and \ - fallback_mode is ControlConnectionQueryFallback.Disabled: - msg = "Unable to connect to any servers" - if self.keyspace: - msg += " using keyspace '%s'" % self.keyspace - raise NoHostAvailable(msg, [h.address for h in hosts]) - - self.session_id = uuid.uuid4() - - if self.cluster.column_encryption_policy is not None: - try: - self.client_protocol_handler = type( - str(self.session_id) + "-ProtocolHandler", - (ProtocolHandler,), - {"column_encryption_policy": self.cluster.column_encryption_policy}) - except AttributeError: - log.info("Unable to set column encryption policy for session") - raise Exception( - "column_encryption_policy is temporary disabled, until https://github.com/scylladb/python-driver/issues/365 is sorted out") - - if self.cluster.monitor_reporting_enabled: - cc_host = self.cluster.get_control_connection_host() - valid_insights_version = (cc_host and version_supports_insights(cc_host.dse_version)) - if valid_insights_version: - self._monitor_reporter = MonitorReporter( - interval_sec=self.cluster.monitor_reporting_interval, - session=self, - ) - else: - if cc_host: - log.debug('Not starting MonitorReporter thread for Insights; ' - 'not supported by server version {v} on ' - 'ControlConnection host {c}'.format(v=cc_host.release_version, c=cc_host)) + try: + if fallback_mode is not \ + ControlConnectionQueryFallback.SkipPoolCreation: + for host in hosts: + future = self.add_or_renew_pool( + host, is_host_addition=False) + if future: + self._initial_connect_futures.add(future) + self._initial_connect_host_futures.append( + (host, future)) + + completed = wait_futures( + self._initial_connect_futures, + return_when=FIRST_COMPLETED) + pool_created = False + while True: + pool_created = any( + future.result() + for future in completed.done) + if pool_created or not completed.not_done: + break + completed = wait_futures( + completed.not_done, + return_when=FIRST_COMPLETED) + + # Only Disabled requires an initial pool to come up. + if not pool_created and \ + fallback_mode is \ + ControlConnectionQueryFallback.Disabled: + msg = "Unable to connect to any servers" + if self.keyspace: + msg += " using keyspace '%s'" % self.keyspace + raise NoHostAvailable( + msg, [h.address for h in hosts]) + except BaseException: + # The partially constructed Session is not yet registered in the + # Cluster, so it alone owns every published pool and in-progress + # candidate. Drain them deterministically before propagation. + self.shutdown() + raise + + try: + self.session_id = uuid.uuid4() + + if self.cluster.column_encryption_policy is not None: + try: + self.client_protocol_handler = type( + str(self.session_id) + "-ProtocolHandler", + (ProtocolHandler,), + {"column_encryption_policy": self.cluster.column_encryption_policy}) + except AttributeError: + log.info( + "Unable to set column encryption policy for session") + raise Exception( + "column_encryption_policy is temporary disabled, until " + "https://github.com/scylladb/python-driver/issues/365 is " + "sorted out") + + if self.cluster.monitor_reporting_enabled: + cc_host = self.cluster.get_control_connection_host() + valid_insights_version = ( + cc_host and + version_supports_insights(cc_host.dse_version)) + if valid_insights_version: + self._monitor_reporter = MonitorReporter( + interval_sec=self.cluster.monitor_reporting_interval, + session=self, + ) + elif cc_host: + log.debug( + 'Not starting MonitorReporter thread for Insights; ' + 'not supported by server version {v} on ' + 'ControlConnection host {c}'.format( + v=cc_host.release_version, + c=cc_host)) - log.debug('Started Session with client_id {} and session_id {}'.format(self.cluster.client_id, - self.session_id)) + log.debug( + 'Started Session with client_id {} and session_id {}'.format( + self.cluster.client_id, + self.session_id)) + except BaseException: + # Construction has not returned to Cluster._new_session yet, so + # only this object can release pools created above. + self.shutdown() + raise def execute(self, query, parameters=None, timeout=_NOT_SET, trace=False, custom_payload=None, execution_profile=EXEC_PROFILE_DEFAULT, @@ -3224,7 +4881,10 @@ def prepare(self, query, custom_payload=None, keyspace=None): log.exception("Error preparing query:") raise - prepared_keyspace = keyspace if keyspace else None + # Preserve the preparation context even when it came from the + # Session. Control-fallback reprepare must not drift to a keyspace + # selected on the Session after this statement was prepared. + prepared_keyspace = future._control_connection_keyspace prepared_statement = PreparedStatement.from_message( response.query_id, response.bind_metadata, response.pk_indexes, self.cluster.metadata, query, prepared_keyspace, self._protocol_version, response.column_metadata, response.result_metadata_id, response.is_lwt, self.cluster.column_encryption_policy) @@ -3235,7 +4895,15 @@ def prepare(self, query, custom_payload=None, keyspace=None): if self.cluster.prepare_on_all_hosts: host = future._current_host try: - self.prepare_on_all_hosts(prepared_statement.query_string, host, prepared_keyspace) + prepare_message_keyspace = ( + prepared_keyspace + if ProtocolVersion.uses_keyspace_flag( + self._protocol_version) + else None) + self.prepare_on_all_hosts( + prepared_statement.query_string, + host, + prepare_message_keyspace) except Exception: log.exception("Error preparing query on all hosts:") @@ -3285,6 +4953,10 @@ def shutdown(self): return else: self.is_shutdown = True + self._pool_repair_schedule = None + self._pool_repair_scheduled = False + pools = tuple(self._pools.values()) + self._pools.clear() # PYTHON-673. If shutdown was called shortly after session init, avoid # a race by cancelling any initial connection attempts haven't started, @@ -3296,7 +4968,7 @@ def shutdown(self): if self._monitor_reporter: self._monitor_reporter.stop() - for pool in tuple(self._pools.values()): + for pool in pools: pool.shutdown() def __enter__(self): @@ -3314,7 +4986,87 @@ def __del__(self): # when cluster.shutdown() is called explicitly. pass - def add_or_renew_pool(self, host, is_host_addition): + def _pool_generation_entry_locked(self, host): + key = id(host) + entry = self._pool_generations.get(key) + if entry is not None and entry[0]() is host: + return entry + + session_ref = weakref.ref(self) + + def remove_generation(host_ref): + session = session_ref() + if session is None: + return + with session._lock: + current = session._pool_generations.get(key) + if current is not None and current[0] is host_ref: + del session._pool_generations[key] + + entry = [weakref.ref(host, remove_generation), 0] + self._pool_generations[key] = entry + return entry + + def _get_pool_generation_locked(self, host): + return self._pool_generation_entry_locked(host)[1] + + def _advance_pool_generation_locked(self, host): + entry = self._pool_generation_entry_locked(host) + entry[1] += 1 + return entry[1] + + def _set_pool_generation_locked(self, host, generation): + self._pool_generation_entry_locked(host)[1] = generation + + def _get_pool(self, host): + """ + Return the pool for ``host`` even if its endpoint changed after it was + inserted into the dict (Host hashes by endpoint). + """ + pool = self._pools.get(host) + if pool is not None: + return pool + with self._lock: + pool = self._pools.get(host) + if pool is not None: + return pool + for pool_host, candidate in tuple(self._pools.items()): + if pool_host is host: + return candidate + return None + + def _pop_pools_locked(self, host): + """ + Remove all entries for the same Host identity. + + A dict can retain an unreachable entry after the key's endpoint (and + therefore hash) changes. Rebuilding also removes any duplicate entry + that may have been inserted under the new hash. + """ + entries = tuple(self._pools.items()) + pools = [pool for pool_host, pool in entries if pool_host is host] + if pools: + remaining = dict( + (pool_host, pool) + for pool_host, pool in entries + if pool_host is not host) + self._pools.clear() + self._pools.update(remaining) + else: + pool = self._pools.pop(host, None) + if pool is not None: + pools.append(pool) + + unique_pools = [] + seen = set() + for pool in pools: + if id(pool) not in seen: + seen.add(id(pool)) + unique_pools.append(pool) + return tuple(unique_pools) + + def add_or_renew_pool( + self, host, is_host_addition, recovery_epoch=None): """ For internal use only. """ @@ -3325,61 +5077,316 @@ def add_or_renew_pool(self, host, is_host_addition): if distance == HostDistance.IGNORED: return None + with host.lock: + if ( + host._is_removed or + ( + recovery_epoch is None and + host.is_up is False) or + (recovery_epoch is not None and + host._recovery_epoch != recovery_epoch)): + return None + attempt_recovery_epoch = host._recovery_epoch + with self._lock: + if self.cluster.is_shutdown or self.is_shutdown: + return None + generation = self._get_pool_generation_locked(host) + + def is_current_locked(): + return not ( + self.cluster.is_shutdown or + self.is_shutdown or + host._is_removed or + ( + recovery_epoch is None and + host.is_up is False) or + host._recovery_epoch != attempt_recovery_epoch or + self._get_pool_generation_locked(host) != generation) + + def is_current(): + with self.cluster._lock: + with host.lock: + with self._lock: + return is_current_locked() + + def commit_down_if_current( + expect_host_to_be_down=False, force=False): + with self.cluster._lock: + with host.lock: + with self._lock: + if not is_current_locked(): + return False + self.cluster._on_down_locked( + host, + is_host_addition, + expect_host_to_be_down, + force) + return True + + def signal_connection_failure_if_current( + connection_exc, expect_host_to_be_down=False, force=False): + if not is_current(): + return False + + uses_default_hooks = getattr( + self.cluster, '_uses_default_failure_hooks', None) + if ( + getattr(type(self.cluster), '_on_down_locked', None) is + None or + not callable(uses_default_hooks) or + uses_default_hooks() is not True): + # Third-party Cluster overrides are part of the driver's + # extension surface. They cannot participate in the private + # atomic commit, so preserve their public failure hook after + # the strongest available freshness check. + forced_commit = False + failure_epoch = getattr(host, '_recovery_epoch', None) + if not isinstance(failure_epoch, int): + failure_epoch = None + try: + _signal_connection_failure( + self.cluster, + host, + connection_exc, + is_host_addition, + expect_host_to_be_down=expect_host_to_be_down, + force=force) + finally: + # A pre-force override may accept the compatible call but + # return False without starting recovery. Authoritative + # startup-close handoffs must still commit DOWN/backoff if + # this pool attempt remains current. + recovery_was_started = bool( + failure_epoch is not None and + getattr( + host, + '_recovery_epoch', + failure_epoch) != failure_epoch) + if ( + force and + not recovery_was_started and + is_current()): + forced_commit = commit_down_if_current( + expect_host_to_be_down, + force=True) + return forced_commit or is_current() + + # Conviction policies are user-extensible and may reenter the + # Session. Never invoke one while Cluster/Host/Session locks are + # held; revalidate before committing the state transition. + try: + is_down = host.signal_connection_failure(connection_exc) + except BaseException: + if not force: + raise + is_down = False + log.exception( + "Conviction policy failed for authoritative pool startup " + "failure on host %s; continuing host recovery", + host) + if force or is_down: + return commit_down_if_current( + expect_host_to_be_down, + force) + return is_current() + + def mark_host_down_if_current(): + uses_default_hooks = getattr( + self.cluster, '_uses_default_failure_hooks', None) + if ( + getattr(type(self.cluster), '_on_down_locked', None) is + None or + not callable(uses_default_hooks) or + uses_default_hooks() is not True): + if not is_current(): + return False + self.cluster.on_down(host, is_host_addition) + return is_current() + return commit_down_if_current() + def run_add_or_renew_pool(): try: - new_pool = HostConnection(host, distance, self) + new_pool = HostConnection( + host, + distance, + self, + pool_generation=generation, + host_recovery_epoch=attempt_recovery_epoch) + except _ConnectionClosedDuringStartup as conn_exc: + if not signal_connection_failure_if_current( + conn_exc, + expect_host_to_be_down=True, + force=True): + return _STALE_POOL_ATTEMPT + log.info( + "Connection pool startup for host %s closed cleanly; " + "handing recovery to the host reconnector", + host) + return False except AuthenticationFailed as auth_exc: conn_exc = ConnectionException(str(auth_exc), endpoint=host) - self.cluster.signal_connection_failure(host, conn_exc, is_host_addition) + if not signal_connection_failure_if_current(conn_exc): + return _STALE_POOL_ATTEMPT return False + except _REQUEST_VALIDATION_EXCEPTIONS as validation_exc: + # A failed initial USE reflects session/keyspace state, not + # host health. HostConnection owns closing its partial pool. + log.warning( + "Failed to initialize connection pool for host %s: %s", + host, + validation_exc) + return _SESSION_LOCAL_POOL_FAILURE except Exception as conn_exc: + if not signal_connection_failure_if_current( + conn_exc, expect_host_to_be_down=True): + return _STALE_POOL_ATTEMPT 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) return False - previous = self._pools.get(host) - with self._lock: - while new_pool._keyspace != self.keyspace: - self._lock.release() + try: + while True: + with host.lock: + with self._lock: + if ( + self.cluster.is_shutdown or + self.is_shutdown or + host._is_removed or + ( + recovery_epoch is None and + host.is_up is False) or + host._recovery_epoch != + attempt_recovery_epoch or + self._get_pool_generation_locked(host) != + generation): + publish_pool = False + previous_pools = () + break + + keyspace = self.keyspace + if new_pool._keyspace == keyspace: + previous_pools = \ + self._pop_pools_locked(host) + self._pools[host] = new_pool + self._set_pool_generation_locked( + host, generation + 1) + new_pool._pool_generation = generation + 1 + publish_pool = True + break + set_keyspace_event = Event() errors_returned = [] - def callback(pool, errors): + def callback( + pool, + errors, + errors_returned=errors_returned, + set_keyspace_event=set_keyspace_event): errors_returned.extend(errors) set_keyspace_event.set() - new_pool._set_keyspace_for_all_conns(self.keyspace, callback) + new_pool._set_keyspace_for_all_conns( + keyspace, + callback) 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) - new_pool.shutdown() - self._lock.acquire() - return False - self._lock.acquire() - self._pools[host] = new_pool + 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) + validation_failure = bool( + errors_returned and all( + isinstance( + error, + _REQUEST_VALIDATION_EXCEPTIONS) + for error in errors_returned)) + if ( + not validation_failure and + not mark_host_down_if_current()): + return _STALE_POOL_ATTEMPT + return _SESSION_LOCAL_POOL_FAILURE \ + if validation_failure else False + + if not publish_pool: + return _STALE_POOL_ATTEMPT + + log.debug("Added pool for host %s to session", host) + for previous in previous_pools: + if previous is not new_pool: + previous.shutdown() - log.debug("Added pool for host %s to session", host) - if previous: - previous.shutdown() - - return True + return True + finally: + # Until publication, this worker is the only owner which can + # release the pool. This covers callback cancellation and + # custom hooks raising anywhere in post-connect setup. + with self._lock: + published = any( + pool is new_pool + for pool in self._pools.values()) + if not published: + new_pool.shutdown() return self.submit(run_add_or_renew_pool) + def _invalidate_pool_attempts(self, host): + with self._lock: + self._advance_pool_generation_locked(host) + def remove_pool(self, host): - pool = self._pools.pop(host, None) - if pool: + with self._lock: + self._advance_pool_generation_locked(host) + pools = self._pop_pools_locked(host) + + if pools: log.debug("Removed connection pool for %r", host) - return self.submit(pool.shutdown) + futures = [] + first_error = None + for pool in pools: + try: + future = self.submit(pool.shutdown) + except BaseException as exc: + future = None + if ( + not isinstance(exc, RuntimeError) and + first_error is None): + first_error = exc + if future is None: + try: + pool.shutdown() + except BaseException as exc: + if first_error is None: + first_error = exc + else: + if isinstance(future, Future): + def finish_cancelled( + completed_future, owned_pool=pool): + if completed_future.cancelled(): + owned_pool.shutdown() + + try: + future.add_done_callback(finish_cancelled) + except BaseException as exc: + try: + pool.shutdown() + except BaseException: + log.exception( + "Additional failure shutting down pool " + "for %r", + host) + if first_error is None: + first_error = exc + futures.append(future) + if first_error is not None: + raise first_error.with_traceback(first_error.__traceback__) + return futures[0] if futures else None else: return None - def update_created_pools(self): + def update_created_pools( + self, skip_host_ids=None, only_host_ids=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 @@ -3395,9 +5402,18 @@ def update_created_pools(self): return set() futures = set() + skip_host_ids = skip_host_ids or () + if only_host_ids is not None: + only_host_ids = set(only_host_ids) for host in self.cluster.metadata.all_hosts(): + host_id = id(host) + if ( + host_id in skip_host_ids or + (only_host_ids is not None and + host_id not in only_host_ids)): + continue distance = self._profile_manager.distance(host) - pool = self._pools.get(host) + pool = self._get_pool(host) future = None if not pool or pool.is_shutdown: # we don't eagerly set is_up on previously ignored hosts. None is included here @@ -3412,9 +5428,80 @@ def update_created_pools(self): else: pool.host_distance = distance if future: + self._track_session_local_pool_repair(future) futures.add(future) return futures + def _track_session_local_pool_repair(self, future): + """ + Keep retrying a missing pool whose creation failed only because of + Session-local state, such as a selected keyspace which is temporarily + absent. Transport failures continue through normal host recovery. + """ + add_done_callback = getattr(future, 'add_done_callback', None) + if not callable(add_done_callback): + return + + def pool_attempt_completed(completed_future): + if completed_future.cancelled(): + return + try: + result = completed_future.result() + except BaseException: + return + if result is _SESSION_LOCAL_POOL_FAILURE: + self._schedule_session_local_pool_repair() + else: + # A completed non-local outcome ends this Session's retry + # chain. A later independent keyspace failure starts with a + # fresh policy schedule rather than inheriting old backoff. + with self._lock: + self._pool_repair_schedule = None + + try: + add_done_callback(pool_attempt_completed) + except BaseException: + log.exception( + "Unable to track Session-local pool repair completion") + + def _schedule_session_local_pool_repair(self): + with self._lock: + if ( + self.cluster.is_shutdown or + self.is_shutdown or + self._pool_repair_scheduled): + return + if self._pool_repair_schedule is None: + self._pool_repair_schedule = \ + self.cluster.reconnection_policy.new_schedule() + try: + delay = next(self._pool_repair_schedule) + except StopIteration: + log.warning( + "Session-local pool repair schedule was exhausted") + self._pool_repair_schedule = None + return + self._pool_repair_scheduled = True + + try: + self.cluster.scheduler.schedule( + delay, self._retry_session_local_pools) + except BaseException: + with self._lock: + self._pool_repair_scheduled = False + log.exception("Unable to schedule Session-local pool repair") + + def _retry_session_local_pools(self): + with self._lock: + self._pool_repair_scheduled = False + if self.cluster.is_shutdown or self.is_shutdown: + return + + futures = self.update_created_pools() + if not futures: + with self._lock: + self._pool_repair_schedule = None + def on_down(self, host): """ Called by the parent Cluster instance when a node is marked down. @@ -3442,9 +5529,79 @@ def _set_keyspace_for_all_pools(self, keyspace, callback): called with a dictionary of all errors that occurred, keyed by the `Host` that they occurred against. """ - with self._lock: - self.keyspace = keyspace - remaining_callbacks = set(self._pools.values()) + completion = { + 'callback': callback, + 'dispatch_complete': False, + 'errors': None, + 'ready': False, + } + + def complete_after_dispatch(errors): + with self._keyspace_completion_lock: + if completion['ready']: + return + completion['errors'] = errors + completion['ready'] = True + should_drain = completion['dispatch_complete'] + if should_drain: + self._drain_keyspace_completions() + + with self._keyspace_dispatch_lock: + with self._keyspace_completion_lock: + self._keyspace_completion_queue.append(completion) + with self._lock: + self.keyspace = keyspace + self._keyspace_generation += 1 + pools = tuple(self._pools.values()) + self._dispatch_keyspace_update_locked( + pools, keyspace, complete_after_dispatch) + + with self._keyspace_completion_lock: + completion['dispatch_complete'] = True + self._drain_keyspace_completions() + + def _drain_keyspace_completions(self): + with self._keyspace_completion_lock: + if self._keyspace_completion_runner_active: + return + self._keyspace_completion_runner_active = True + + rerun = False + try: + while True: + with self._keyspace_completion_lock: + if not self._keyspace_completion_queue: + return + completion = self._keyspace_completion_queue[0] + if not ( + completion['dispatch_complete'] and + completion['ready']): + return + self._keyspace_completion_queue.pop(0) + + try: + completion['callback'](completion['errors']) + except Exception: + log.exception( + "Error completing Session keyspace update") + finally: + with self._keyspace_completion_lock: + self._keyspace_completion_runner_active = False + rerun = bool( + self._keyspace_completion_queue and + self._keyspace_completion_queue[0][ + 'dispatch_complete'] and + self._keyspace_completion_queue[0]['ready']) + if rerun: + self._drain_keyspace_completions() + + def _dispatch_keyspace_update_locked(self, pools, keyspace, callback): + """ + Enqueue one Session keyspace generation on every pool while holding + the dispatch lock, preserving the same A-before-B order cluster-wide. + """ + remaining_callbacks = set(pools) + callbacks_lock = Lock() errors = {} if not remaining_callbacks: @@ -3452,15 +5609,30 @@ def _set_keyspace_for_all_pools(self, keyspace, callback): return def pool_finished_setting_keyspace(pool, host_errors): - remaining_callbacks.remove(pool) - if host_errors: - errors[pool.host] = host_errors + callback_errors = None + with callbacks_lock: + # A pool should call back once, but ignoring a duplicate keeps + # completion exactly-once if an implementation violates that + # contract. + if pool not in remaining_callbacks: + return + + remaining_callbacks.remove(pool) + if host_errors: + errors[pool.host] = host_errors + + if not remaining_callbacks: + callback_errors = dict(errors) - if not remaining_callbacks: - callback(errors) + if callback_errors is not None: + callback(callback_errors) - for pool in tuple(self._pools.values()): - pool._set_keyspace_for_all_conns(keyspace, pool_finished_setting_keyspace) + for pool in pools: + try: + pool._set_keyspace_for_all_conns( + keyspace, pool_finished_setting_keyspace) + except BaseException as exc: + pool_finished_setting_keyspace(pool, [exc]) def wait_for_schema_agreement(self, wait_time: Optional[float] = None, scope: SchemaAgreementScope = SchemaAgreementScope.CLUSTER) -> bool: @@ -3714,7 +5886,11 @@ def try_reconnect(self): return self.control_connection._reconnect_internal() def on_reconnection(self, connection): - self.control_connection._set_new_connection(connection) + transferred = self.control_connection._set_new_connection(connection) + # The control connection now owns this socket. Host reconnection + # probes return the default falsey value and are still closed by the + # generic handler after they have verified host health. + return transferred def on_exception(self, exc, next_delay): # TODO only overridden to add logging, so add logging @@ -3820,27 +5996,108 @@ def __init__(self, cluster, timeout, self._event_schedule_times = {} + def _get_host_for_connection(self, connection, require_current=False): + if connection is None: + return None + connection_state = getattr(connection, '__dict__', None) + host = connection_state.get( + '_control_connection_host') if connection_state else None + if host is None: + host = self._cluster.metadata.get_host(connection.endpoint) + if ( + require_current and + host is not None and + self._cluster.metadata.get_host_by_host_id( + host.host_id) is not host): + return None + return host + def connect(self): if self._is_shutdown: return self._protocol_version = self._cluster.protocol_version - self._set_new_connection(self._reconnect_internal()) + connection = self._reconnect_internal() + self._set_new_connection(connection) - self._cluster.metadata.dbaas = self._connection._product_type == dscloud.DATASTAX_CLOUD_PRODUCT_TYPE + # shutdown() can reject the candidate in _set_new_connection or clear + # it immediately after publication. Snapshot the accepted connection + # under the same lock before reading any of its negotiated features. + with self._lock: + if ( + self._is_shutdown or + self._connection is not connection): + return + product_type = connection._product_type + self._cluster.metadata.dbaas = \ + product_type == dscloud.DATASTAX_CLOUD_PRODUCT_TYPE def _set_new_connection(self, conn): """ Replace existing connection (if there is one) and close it. """ with self._lock: - old = self._connection - self._connection = conn + reject_connection = self._is_shutdown + if reject_connection: + old = None + else: + old = self._connection + self._connection = conn + + if reject_connection: + # The caller transferred ownership by invoking this method, even + # though shutdown won before publication. + conn.close() + return True if old: log.debug("[control connection] Closing old connection %r, replacing with %r", old, conn) old.close() + # A successful control reconnection is authoritative evidence that its + # host can serve CQL again. This also resumes hosts whose regular + # reconnector was parked after a clean maintenance-mode startup close. + if old is not None: + host = self._get_host_for_connection(conn) + if host is not None: + parked_addition = None + with host.lock: + if host.is_up: + return True + reconnector = host._reconnection_handler + is_host_addition = bool( + reconnector and + getattr(reconnector, 'is_host_addition', False)) + if is_host_addition: + parked_addition = \ + host.get_and_set_reconnection_handler(None) + if parked_addition is not None: + parked_addition.cancel() + on_reconnection = self._cluster.on_add \ + if is_host_addition else self._cluster.on_up + try: + self._cluster.scheduler.schedule_unique( + 0, on_reconnection, host) + except BaseException: + # The control socket is already adopted. Never let an + # optional scheduler failure escape to the reconnector, + # which would close this now-active connection and retain + # a stale handler. Fall back to the same serialized + # transition directly. + log.exception( + "Unable to schedule host recovery after control " + "reconnection to %s; running it directly", + host) + if not self._cluster.is_shutdown: + try: + on_reconnection(host) + except BaseException: + log.exception( + "Host recovery after control reconnection " + "failed for %s", + host) + return True + def _try_connect_to_hosts(self): errors = {} @@ -3895,6 +6152,8 @@ def _try_connect(self, endpoint): if self._is_shutdown: connection.close() raise DriverException("Reconnecting during shutdown") + if connection.is_closed: + raise _startup_close_error(connection, endpoint) break except ProtocolVersionUnsupported as e: self._cluster.protocol_downgrade(endpoint, e.startup_version) @@ -3911,23 +6170,35 @@ def _try_connect(self, endpoint): "registering watchers and refreshing schema and topology", connection) - # Indirect way to determine if conencted to a ScyllaDB cluster, which does not support peers_v2 - # If sharding information is available, it's a ScyllaDB cluster, so do not use peers_v2 table. - if connection.features.sharding_info is not None: - self._uses_peers_v2 = False - - # Only ScyllaDB supports "USING TIMEOUT" - # Sharding information signals it is ScyllaDB - self._metadata_request_timeout = None if connection.features.sharding_info is None or not self._cluster.metadata_request_timeout \ - else datetime.timedelta(seconds=self._cluster.metadata_request_timeout) + try: + # Indirect way to determine if connected to a ScyllaDB cluster, + # which does not support peers_v2. If sharding information is + # available, do not use peers_v2. + if connection.features.sharding_info is not None: + self._uses_peers_v2 = False - self._tablets_routing_v1 = connection.features.tablets_routing_v1 + # Only ScyllaDB supports "USING TIMEOUT"; sharding information + # identifies a ScyllaDB connection. + self._metadata_request_timeout = ( + None + if ( + connection.features.sharding_info is None or + not self._cluster.metadata_request_timeout) + else datetime.timedelta( + seconds=self._cluster.metadata_request_timeout)) + + self._tablets_routing_v1 = \ + connection.features.tablets_routing_v1 + + # Use weak references in both directions. _clear_watcher runs + # when this ControlConnection is finalized. + self_weakref = weakref.ref( + self, + partial(_clear_watcher, weakref.proxy(connection))) + except BaseException: + connection.close() + raise - # use weak references in both directions - # _clear_watcher will be called when this ControlConnection is about to be finalized - # _watch_callback will get the actual callback from the Connection and relay it to - # this object (after a dereferencing a weakref) - self_weakref = weakref.ref(self, partial(_clear_watcher, weakref.proxy(connection))) try: watchers = { "TOPOLOGY_CHANGE": partial(_watch_callback, self_weakref, '_handle_topology_change'), @@ -3969,7 +6240,7 @@ def _try_connect(self, endpoint): shared_results = (peers_result, local_result) self._refresh_node_list_and_token_map(connection, preloaded_results=shared_results) self._refresh_schema(connection, preloaded_results=shared_results, schema_agreement_wait=-1) - except Exception: + except BaseException: connection.close() raise @@ -3998,11 +6269,19 @@ def _reconnect(self): # when a connection is successfully made, _set_new_connection # will be called with the new connection and then our # _reconnection_handler will be cleared out - self._reconnection_handler = _ControlReconnectionHandler( + handler_holder = [] + + def clear_reconnector(): + if handler_holder: + self._clear_reconnection_handler( + handler_holder[0]) + + handler = _ControlReconnectionHandler( self, self._cluster.scheduler, schedule, - self._get_and_set_reconnection_handler, - new_handler=None) - self._reconnection_handler.start() + clear_reconnector) + handler_holder.append(handler) + self._reconnection_handler = handler + handler.start() except Exception: log.debug("[control connection] error reconnecting", exc_info=True) raise @@ -4018,6 +6297,13 @@ def _get_and_set_reconnection_handler(self, new_handler): self._reconnection_handler = new_handler return old + def _clear_reconnection_handler(self, expected_handler): + with self._reconnection_lock: + if self._reconnection_handler is expected_handler: + self._reconnection_handler = None + return expected_handler + return None + def _submit(self, *args, **kwargs): try: if not self._cluster.is_shutdown: @@ -4120,10 +6406,12 @@ def _refresh_node_list_and_token_map(self, connection, preloaded_results=None, found_host_ids = set() found_endpoints = set() + connected_host_id = None if local_result.parsed_rows: local_rows = dict_factory(local_result.column_names, local_result.parsed_rows) local_row = local_rows[0] + connected_host_id = local_row.get("host_id") cluster_name = local_row["cluster_name"] self._cluster.metadata.cluster_name = cluster_name @@ -4164,12 +6452,38 @@ def _refresh_node_list_and_token_map(self, connection, preloaded_results=None, reconnector = host.get_and_set_reconnection_handler(None) if reconnector: reconnector.cancel() - self._cluster.on_down(host, is_host_addition=False, expect_host_to_be_down=True) - - old_endpoint = host.endpoint - host.endpoint = endpoint - self._cluster.metadata.update_host(host, old_endpoint) - self._cluster.on_up(host) + force_down = getattr( + self._cluster, + '_force_down_for_endpoint_change', + None) + if callable(force_down): + relocation_result = force_down(host, endpoint) + if relocation_result is \ + _HOST_TRANSITION_DEFERRED: + # This refresh is running from another Host's + # transition lane. Relocation is queued but has not + # mutated the endpoint yet, so using this row now + # would rebuild tokens under the stale key. + self._cluster.scheduler.schedule_unique( + 0, + self.refresh_node_list_and_token_map, + force_token_rebuild=True) + return + if not relocation_result: + return + else: + # Preserve the historical duck-typed Cluster surface + # used by integrations which do not subclass Cluster. + self._cluster.on_down( + host, + is_host_addition=False, + expect_host_to_be_down=True) + old_endpoint = host.endpoint + host.endpoint = endpoint + self._cluster.metadata.update_host( + host, + old_endpoint) + self._cluster.on_up(host) if host is None: log.debug("[control connection] Found new host to connect to: %s", endpoint) @@ -4193,6 +6507,14 @@ def _refresh_node_list_and_token_map(self, connection, preloaded_results=None, token_map[host] = tokens self._cluster.metadata.update_host(host, old_endpoint=endpoint) + if connected_host_id in found_host_ids: + # Keep the identity verified by system.local. The endpoint used to + # reach the control connection may be a contact-point alias, + # translated address, or shared ClientRoutes proxy and therefore + # cannot reliably key metadata. + connection._control_connection_host = \ + self._cluster.metadata.get_host_by_host_id(connected_host_id) + for old_host_id, old_host in self._cluster.metadata.all_hosts_items(): if old_host_id not in found_host_ids: should_rebuild_token_map = True @@ -4513,18 +6835,64 @@ def _signal_error(self): with self._lock: if self._is_shutdown: return - - # try just signaling the cluster, as this will trigger a reconnect - # as part of marking the host down - if self._connection and self._connection.is_defunct: - host = self._cluster.metadata.get_host(self._connection.endpoint) - # host may be None if it's already been removed, but that indicates - # that errors have already been reported, so we're fine - if host: + connection = self._connection + + # Conviction and Cluster transitions can reenter the control + # connection. Keep them outside the ControlConnection lock to avoid a + # Cluster-lock/control-lock inversion during initial connect. + if connection and connection.is_defunct: + host = self._get_host_for_connection( + connection, require_current=True) + # host may be None if it's already been removed, but that indicates + # that errors have already been reported, so we're fine + if host: + # ``_cluster`` is a weakref proxy, so inspecting its type + # yields ProxyType and would make the fenced production path + # permanently unreachable. Resolve the bound method through + # the proxy instead. + down_locked = getattr( + self._cluster, '_on_down_locked', None) + uses_default_hooks = getattr( + self._cluster, '_uses_default_failure_hooks', None) + if ( + down_locked is None or + not callable(uses_default_hooks) or + uses_default_hooks() is not True): + with self._lock: + if ( + self._connection is not connection or + self._is_shutdown): + return self._cluster.signal_connection_failure( - host, self._connection.last_error, is_host_addition=False) + host, + connection.last_error, + is_host_addition=False) return + # The policy is user-extensible, so run it lock-free. Commit + # only if this is still the active defunct control connection, + # using the global Cluster -> ControlConnection lock order. + is_down = host.signal_connection_failure( + connection.last_error) + if not is_down: + return + with self._cluster._lock: + with self._lock: + if ( + self._is_shutdown or + self._connection is not connection or + not connection.is_defunct or + self._get_host_for_connection( + connection, + require_current=True) is not host): + return + down_locked( + host, + is_host_addition=False, + expect_host_to_be_down=False, + force=False) + return + # if the connection is not defunct or the host already left, reconnect # manually self.reconnect() @@ -4535,7 +6903,9 @@ def on_up(self, host): def on_down(self, host): conn = self._connection - if conn and conn.endpoint == host.endpoint and \ + connection_host = self._get_host_for_connection(conn) + if conn and (connection_host is host or ( + connection_host is None and conn.endpoint == host.endpoint)) and \ self._reconnection_handler is None: log.debug("[control connection] Control connection host (%s) is " "considered down, starting reconnection", host) @@ -4548,7 +6918,9 @@ def on_add(self, host, refresh_nodes=True): def on_remove(self, host): c = self._connection - if c and c.endpoint == host.endpoint: + connection_host = self._get_host_for_connection(c) + if c and (connection_host is host or ( + connection_host is None and c.endpoint == host.endpoint)): log.debug("[control connection] Control connection host (%s) is being removed. Reconnecting", host) # refresh will be done on reconnect self.reconnect() @@ -4735,6 +7107,7 @@ class ResponseFuture(object): _continuous_paging_session = None _host = None _control_connection_query_attempted = False + _control_connection_keyspace = None _TABLET_ROUTING_CTYPE = None _bound_result_metadata = None @@ -4760,6 +7133,21 @@ def __init__(self, session, message, query, timeout, metrics=None, prepared_stat # even if a concurrent METADATA_CHANGED replaces the prepared statement's cache in # between. Defaults to [] for unprepared statements (no cached metadata). self._bound_result_metadata = [] if bound_result_metadata is _NOT_SET else bound_result_metadata + self._control_connection_keyspace = None + if not isinstance(query, GraphStatement): + keyspace_candidates = ( + getattr(message, 'keyspace', None), + getattr(prepared_statement, 'keyspace', None), + getattr(query, 'keyspace', None), + getattr(session, 'keyspace', None), + ) + self._control_connection_keyspace = next( + ( + keyspace + for keyspace in keyspace_candidates + if keyspace is not None + ), + None) self._callback_lock = Lock() self._start_time = start_time or time.time() self._host = host @@ -4829,7 +7217,8 @@ 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 = Session._get_pool( + self.session, 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 @@ -4927,11 +7316,44 @@ def send_request(self, error_no_hosts=True): def _has_usable_node_pool(self): try: - pools = tuple(self.session._pools.values()) + session_lock = getattr(self.session, '_lock', None) + if session_lock is None: + pools = tuple(self.session._pools.values()) + else: + with session_lock: + pools = tuple(self.session._pools.values()) except (AttributeError, TypeError): return False + except RuntimeError: + # A third-party Session stand-in without a shared lock may mutate + # its pool mapping concurrently. Treat an uncertain snapshot as + # usable rather than leaking the race or borrowing the control + # connection while a node pool may exist. + return True - return any(pool and not pool.is_shutdown for pool in pools) + for pool in pools: + if not pool or pool.is_shutdown: + continue + + try: + connections = tuple(pool.get_connections()) + except (AttributeError, TypeError, RuntimeError): + # Preserve compatibility with third-party pool implementations + # which predate per-connection quarantine. + return True + + if any( + not ( + getattr(connection, 'is_closed', False) or + getattr(connection, 'is_defunct', False) or + getattr(connection, '_pool_retired', False) is True or + getattr( + connection, + '_pool_keyspace_mismatch', + False) is True) + for connection in connections): + return True + return False def _fallback_to_control_connection(self): fallback_mode = self.session.cluster.allow_control_connection_query_fallback @@ -4939,6 +7361,11 @@ def _fallback_to_control_connection(self): return False if self._host or self._control_connection_query_attempted: return False + if getattr(self.message, 'continuous_paging_options', None): + self._errors['control connection'] = UnsupportedOperation( + "Continuous paging is not supported over the control " + "connection fallback") + return False if fallback_mode is ControlConnectionQueryFallback.SkipPoolCreation: return True return not self._has_usable_node_pool() @@ -4981,6 +7408,39 @@ def _query_control_connection(self, message=None, cb=None, connection=None, host request_id = None request_sent = False try: + query_text = getattr(message, 'query', None) + if _is_use_statement(query_text): + raise UnsupportedOperation( + "USE statements cannot be executed over the shared " + "control connection fallback") + + message_keyspace = getattr(message, 'keyspace', None) + effective_keyspace = ( + message_keyspace + if message_keyspace is not None + else self._control_connection_keyspace) + protocol_version = getattr( + connection, + 'protocol_version', + self.session.cluster.protocol_version) + if effective_keyspace is not None: + if not ProtocolVersion.uses_keyspace_flag(protocol_version): + raise UnsupportedOperation( + "Session keyspaces on the control connection " + "fallback require protocol version 5 or DSE_V2") + if not hasattr(message, 'keyspace'): + raise UnsupportedOperation( + "This request type cannot encode a keyspace for the " + "control connection fallback") + message.keyspace = effective_keyspace + elif getattr(connection, 'keyspace', None) is not None: + # Native protocol v5 can scope a request to a keyspace, but it + # cannot express "no keyspace". Never inherit mutable state + # left on this cluster-wide connection by another user. + raise UnsupportedOperation( + "An unscoped request cannot use a control connection " + "which already selected a keyspace") + request_id = self._borrow_control_connection(connection) self._connection = connection result_meta = self._bound_result_metadata @@ -5019,7 +7479,7 @@ def _query(self, host, message=None, cb=None): self._control_connection_query_attempted = False - pool = self.session._pools.get(host) + pool = Session._get_pool(self.session, host) if not pool: self._errors[host] = ConnectionException("Host has been marked down or removed") return None @@ -5031,11 +7491,28 @@ def _query(self, host, message=None, cb=None): connection = None try: + message_query = getattr(message, 'query', None) + allow_keyspace_mismatch = _is_use_statement(message_query) # 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) + borrow_kwargs = { + 'timeout': 2.0, + 'routing_key': self.query.routing_key, + 'keyspace': self.query.keyspace, + 'table': self.query.table, + } + if allow_keyspace_mismatch: + borrow_kwargs['allow_keyspace_mismatch'] = True + connection, request_id = pool.borrow_connection( + **borrow_kwargs) else: - connection, request_id = pool.borrow_connection(timeout=2.0) + if allow_keyspace_mismatch: + connection, request_id = pool.borrow_connection( + timeout=2.0, + allow_keyspace_mismatch=True) + else: + connection, request_id = pool.borrow_connection( + timeout=2.0) self._connection = connection result_meta = self._bound_result_metadata @@ -5145,8 +7622,36 @@ def _reprepare(self, prepare_message, host, connection, pool): self.send_request() def _set_result(self, host, connection, pool, response): + pending_continuous_paging_session = None + continuous_paging_session_handed_off = False try: self.coordinator_host = host + if ( + isinstance(response, ResultMessage) and + response.kind == RESULT_KIND_ROWS and + getattr(self.message, 'continuous_paging_options', None)): + # The initial request's pool return can retire and close this + # connection as soon as in_flight reaches zero. Register the + # stream first so retirement sees the paging session as an + # active use which still requires the socket. + pending_continuous_paging_session = \ + connection.new_continuous_paging_session( + response.stream_id, + self._protocol_handler.decode_message, + self.row_factory, + self._continuous_paging_state) + + if ( + isinstance(response, ResultMessage) and + response.kind == RESULT_KIND_SET_KEYSPACE and + connection is not None): + # Publish recovery of a quarantined socket before returning it + # to the pool. return_connection() wakes blocked borrowers, and + # they must not observe the old mismatch flag after USE + # succeeded. + connection.keyspace = response.new_keyspace + connection._pool_keyspace_mismatch = False + if pool and not pool.is_shutdown: pool.return_connection(connection) @@ -5178,8 +7683,6 @@ def _set_result(self, host, connection, pool, response): if isinstance(response, ResultMessage): if response.kind == RESULT_KIND_SET_KEYSPACE: session = getattr(self, 'session', None) - if connection is not None: - connection.keyspace = response.new_keyspace # since we're running on the event loop thread, we need to # use a non-blocking method for setting the keyspace on # all connections in this session, otherwise the event @@ -5230,7 +7733,11 @@ def _set_result(self, host, connection, pool, response): getattr(self.prepared_statement, 'query_id', None) ) if getattr(self.message, 'continuous_paging_options', None): - self._handle_continuous_paging_first_response(connection, response) + self._handle_continuous_paging_first_response( + connection, + response, + pending_continuous_paging_session) + continuous_paging_session_handed_off = True else: self._set_final_result(self.row_factory(response.column_names, response.parsed_rows)) elif response.kind == RESULT_KIND_VOID: @@ -5337,21 +7844,67 @@ def _set_result(self, host, connection, pool, response): self._connection.defunct(exc) self._set_final_exception(exc) except Exception as exc: + notify_continuous_paging_release = False + if ( + pending_continuous_paging_session is not None and + not continuous_paging_session_handed_off): + # Registration happened before the pool return, but response + # processing failed before the session was handed to the + # result consumer. Leave stream-id recycling to process_msg, + # which still owns the current response. + with connection.lock: + current_session = \ + connection._continuous_paging_sessions.get( + response.stream_id) + if current_session is pending_continuous_paging_session: + del connection._continuous_paging_sessions[ + response.stream_id] + notify_continuous_paging_release = True + if ( + self._continuous_paging_session is + pending_continuous_paging_session): + self._continuous_paging_session = None + if notify_continuous_paging_release and pool is not None: + on_connection_released = getattr( + pool, 'on_connection_released', None) + if on_connection_released is not None: + # HostConnection takes pool -> connection locks. + try: + on_connection_released(connection) + except Exception: + log.exception( + "Error releasing failed continuous-paging " + "reservation on connection (%s)", + id(connection)) # almost certainly caused by a bug, but we need to set something here log.exception("Unexpected exception while handling result in ResponseFuture:") self._set_final_exception(exc) - def _handle_continuous_paging_first_response(self, connection, response): - self._continuous_paging_session = connection.new_continuous_paging_session(response.stream_id, - self._protocol_handler.decode_message, - self.row_factory, - self._continuous_paging_state) + def _handle_continuous_paging_first_response( + self, connection, response, paging_session=None): + self._continuous_paging_session = ( + paging_session or + connection.new_continuous_paging_session( + response.stream_id, + self._protocol_handler.decode_message, + self.row_factory, + self._continuous_paging_state)) self._continuous_paging_session.on_message(response) self._set_final_result(self._continuous_paging_session.results()) def _set_keyspace_completed(self, errors): if not errors: self._set_final_result(None) + # A prior invalid/dropped keyspace can leave an otherwise-UP host + # without a Session-local pool. The successful new keyspace is + # now installed on existing pools, so retry any missing ones. + try: + self.session.submit(self.session.update_created_pools) + except BaseException: + # Repair is best-effort and must never change an already + # successful USE result or escape the reactor callback. + log.exception( + "Unable to schedule missing-pool repair after USE") else: self._set_final_exception(ConnectionException( "Failed to set keyspace on all hosts: %s" % (errors,))) diff --git a/cassandra/connection.py b/cassandra/connection.py index f238416b29..8d9c521d5b 100644 --- a/cassandra/connection.py +++ b/cassandra/connection.py @@ -18,6 +18,7 @@ from functools import wraps, partial, total_ordering from heapq import heappush, heappop import io +import inspect import logging import socket import struct @@ -40,12 +41,16 @@ else: from queue import Queue, Empty # noqa -from cassandra import ConsistencyLevel, AuthenticationFailed, OperationTimedOut, ProtocolVersion +from cassandra import (ConsistencyLevel, AuthenticationFailed, + OperationTimedOut, ProtocolVersion, + RequestValidationException) from cassandra.marshal import int32_pack from cassandra.protocol import (ReadyMessage, AuthenticateMessage, OptionsMessage, StartupMessage, ErrorMessage, CredentialsMessage, QueryMessage, ResultMessage, ProtocolHandler, - InvalidRequestException, SupportedMessage, + RequestValidationException as + ProtocolRequestValidationException, + SupportedMessage, AuthResponseMessage, AuthChallengeMessage, AuthSuccessMessage, ProtocolException, RegisterMessage, ReviseRequestMessage) @@ -511,6 +516,39 @@ def __str__(self): NONBLOCKING = (errno.EAGAIN, errno.EWOULDBLOCK) +class _ConnectionStartupEvent(Event): + """ + Track legacy reactors which publish READY by setting ``connected_event``. + + Modern startup paths publish ``_startup_completed`` explicitly. Historical + third-party reactors only set this event, so remember whether it was set + while the connection was still open. This preserves that publication if a + close wins the race before ``Connection.factory`` snapshots the state. + """ + + def __init__(self, connection): + Event.__init__(self) + self._connection_ref = weakref.ref(connection) + + def set(self): + connection = self._connection_ref() + if connection is None: + Event.set(self) + return + + # Publish the legacy READY marker and wake factory atomically with + # close/defunct and factory's state snapshot. Connection.lock is an + # RLock because modern startup and Twisted close already set this + # event while holding it. + with connection.lock: + if ( + not connection.is_closed and + not connection.is_defunct and + connection.last_error is None): + connection._startup_event_set_while_open = True + Event.set(self) + + class ConnectionException(Exception): """ An unrecoverable error was hit when attempting to use a connection, @@ -533,6 +571,84 @@ class ConnectionShutdown(ConnectionException): pass +class _ConnectionClosedDuringStartup(ConnectionShutdown): + """ + Internal owner-side signal for a clean close returned by + ``Connection.factory``. + """ + pass + + +def _startup_close_error(connection, endpoint=None): + """Build an owner-side error without changing the factory return contract.""" + connection_endpoint = getattr(connection, 'endpoint', None) + if connection_endpoint is not None: + endpoint = connection_endpoint + connection_lock = getattr(connection, 'lock', None) + if hasattr(connection_lock, '__enter__'): + with connection_lock: + startup_completed = ( + getattr(connection, '_startup_completed', False) is True) + last_error = getattr(connection, 'last_error', None) + else: + startup_completed = ( + getattr(connection, '_startup_completed', False) is True) + last_error = getattr(connection, 'last_error', None) + + if startup_completed: + if isinstance(last_error, BaseException): + return last_error + return ConnectionShutdown( + "Connection to %s was closed after startup" % (endpoint,), + endpoint) + return _ConnectionClosedDuringStartup( + "Connection to %s was closed during the startup handshake" % (endpoint,), + endpoint) + + +def _set_keyspace_blocking(connection, keyspace, timeout): + """ + Call the timeout-aware API while preserving custom Connection subclasses + which implemented the historical one-argument override. + """ + method = connection.set_keyspace_blocking + try: + parameters = tuple(inspect.signature(method).parameters.values()) + except (TypeError, ValueError): + # Uninspectable C/Cython callables are often historical one-argument + # overrides. Preserve that API rather than guessing at a keyword. + return method(keyspace) + + timeout_parameter = next( + ( + parameter for parameter in parameters + if parameter.name == 'timeout' + ), + None) + if ( + timeout_parameter is not None and + timeout_parameter.kind == inspect.Parameter.POSITIONAL_ONLY): + return method(keyspace, timeout) + + if ( + timeout_parameter is not None and + timeout_parameter.kind in ( + inspect.Parameter.POSITIONAL_OR_KEYWORD, + inspect.Parameter.KEYWORD_ONLY)): + return method(keyspace, timeout=timeout) + + if any( + parameter.kind == inspect.Parameter.VAR_KEYWORD + for parameter in parameters): + return method(keyspace, timeout=timeout) + + if any( + parameter.kind == inspect.Parameter.VAR_POSITIONAL + for parameter in parameters): + return method(keyspace, timeout) + return method(keyspace) + + class ProtocolVersionUnsupported(ConnectionException): """ Server rejected startup message due to unsupported protocol version @@ -841,6 +957,9 @@ class Connection(object): is_defunct = False is_closed = False + # Set by a pool before it deliberately retires this connection. Returns + # racing that close must not treat it as a transport failure. + _pool_retired = False lock = None user_type_map = None @@ -896,12 +1015,18 @@ def __init__(self, host='127.0.0.1', port=9042, authenticator=None, self.connect_timeout = connect_timeout self.allow_beta_protocol_version = allow_beta_protocol_version self.no_compact = no_compact + self._owning_pool = owning_pool self._push_watchers = defaultdict(set) self._requests = {} + # Per-request dispatch state used by set_keyspace_async to distinguish + # failures known to precede socket submission from ambiguous failures + # after push() has begun. + self._send_msg_states = {} self._io_buffer = _ConnectionIOBuffer(self) self._continuous_paging_sessions = {} self._socket_writable = True self.orphaned_request_ids = set() + self._pool_retired = False self._on_orphaned_stream_released = on_orphaned_stream_released self._application_info = application_info @@ -929,7 +1054,12 @@ def __init__(self, host='127.0.0.1', port=9042, authenticator=None, self.highest_request_id = initial_size - 1 self.lock = RLock() - self.connected_event = Event() + self._startup_event_set_while_open = False + self.connected_event = _ConnectionStartupEvent(self) + # ``connected_event`` is also set by close()/defunct() to wake factory + # waiters, so it cannot by itself prove that the startup handshake + # completed. Keep that state separately and publish it under ``lock``. + self._startup_completed = False self.features = ProtocolFeatures(shard_id=shard_id) self.total_shards = total_shards self.original_endpoint = self.endpoint @@ -962,34 +1092,146 @@ def handle_fork(cls): def create_timer(cls, timeout, callback): raise NotImplementedError() + def _is_clean_close_error(self, exc): + return not self.is_defunct and isinstance(exc, ConnectionShutdown) + @classmethod - def factory(cls, endpoint, timeout, host_conn = None, *args, **kwargs): + def factory(cls, endpoint, timeout, *args, **kwargs): """ - A factory function which returns connections which have - succeeded in connecting and are ready for service (or - raises an exception otherwise). + A factory function which returns a connection once startup has + completed, returns a closed connection if the server closes during + startup, or raises an exception otherwise. + + Pool callers pass ``host_conn`` as a keyword-only factory option so + the connection remains tracked from construction through caller + adoption and can be closed by pool shutdown throughout that handoff. + Positional arguments retain their historical meaning as constructor + arguments. """ start = time.time() + host_conn = kwargs.pop('host_conn', None) kwargs['connect_timeout'] = timeout conn = cls(endpoint, *args, **kwargs) - if host_conn is not None: - host_conn._pending_connections.append(conn) - if host_conn.is_shutdown: + if host_conn is not None and conn._owning_pool is None: + # ``host_conn`` is deliberately not forwarded to third-party + # Connection constructors, but the established owning_pool slot + # still needs to identify the pool for asynchronous release + # notifications after factory returns. + conn._owning_pool = host_conn + pending_registered = False + factory_completed = False + explicitly_closed = False + try: + if host_conn is not None: + register_pending = getattr( + type(host_conn), '_register_pending_connection', None) + if register_pending is not None: + pending_registered = register_pending(host_conn, conn) + host_conn_shutdown = not pending_registered + else: + pending_lock = getattr(host_conn, '_lock', None) + if hasattr(pending_lock, '__enter__'): + with pending_lock: + host_conn._pending_connections.append(conn) + pending_registered = True + host_conn_shutdown = host_conn.is_shutdown + else: + host_conn._pending_connections.append(conn) + pending_registered = True + host_conn_shutdown = host_conn.is_shutdown + + if host_conn_shutdown: + conn.close() + explicitly_closed = True + elapsed = time.time() - start + conn.connected_event.wait(timeout - elapsed) + # READY/auth success and close/defunct all wake connected_event. + # Snapshot their state under the connection lock so a defunct + # transition cannot expose is_defunct before last_error, and a + # post-startup close cannot be mistaken for a clean startup close. + with conn.lock: + last_error = conn.last_error + is_closed = conn.is_closed + startup_completed = conn._startup_completed + event_is_set = conn.connected_event.is_set() + is_unsupported_proto_version = conn.is_unsupported_proto_version + legacy_startup_completed = ( + event_is_set and + conn._startup_event_set_while_open) + if ( + event_is_set and + not startup_completed and + ( + legacy_startup_completed or + ( + not is_closed and + not conn.is_defunct and + last_error is None))): + # Preserve the historical reactor-subclass contract: + # before _mark_startup_completed existed, setting this + # event on an open/error-free connection signaled READY. + # The event records that ordering so a later close cannot + # erase a READY publication before this snapshot. + conn._startup_completed = True + startup_completed = True + clean_startup_close = ( + last_error is not None + and is_closed + and not startup_completed + and conn._is_clean_close_error(last_error)) + + if last_error: + if is_unsupported_proto_version: + raise ProtocolVersionUnsupported(endpoint, conn.protocol_version) + if clean_startup_close: + factory_completed = True + return conn + raise last_error + elif not event_is_set: conn.close() - elapsed = time.time() - start - conn.connected_event.wait(timeout - elapsed) - if conn.last_error: - if conn.is_unsupported_proto_version: - raise ProtocolVersionUnsupported(endpoint, conn.protocol_version) - raise conn.last_error - elif not conn.connected_event.is_set(): - conn.close() - raise OperationTimedOut("Timed out creating connection (%s seconds)" % timeout, - timeout=timeout) - elif conn.is_closed: - raise ConnectionShutdown("Connection to %s was closed by server" % conn.endpoint) - else: + explicitly_closed = True + raise OperationTimedOut("Timed out creating connection (%s seconds)" % timeout, + timeout=timeout) + elif is_closed and startup_completed: + raise ConnectionShutdown( + "Connection to %s was closed by server" % conn.endpoint, + conn.endpoint) + factory_completed = True return conn + finally: + # Event.wait can be interrupted by cancellation exceptions which + # do not derive from Exception. Until a result is returned the + # factory still owns the socket and must not leave it untracked. + if ( + not factory_completed and + not explicitly_closed and + not conn.is_closed): + conn.close() + # Failed or closed candidates remain factory-owned and are + # unregistered here. A successful open pool candidate stays + # registered until the caller atomically adopts it, closing the + # otherwise unowned return-to-caller window against shutdown. + if pending_registered and ( + not factory_completed or conn.is_closed): + unregister_pending = getattr( + type(host_conn), '_unregister_pending_connection', None) + if unregister_pending is not None: + unregister_pending(host_conn, conn) + else: + pending_lock = getattr(host_conn, '_lock', None) + if hasattr(pending_lock, '__enter__'): + with pending_lock: + pending_connections = host_conn._pending_connections + for i, pending in enumerate(pending_connections): + if pending is conn: + del pending_connections[i] + break + else: + pending_connections = host_conn._pending_connections + for i, pending in enumerate(pending_connections): + if pending is conn: + del pending_connections[i] + break def _build_ssl_context_from_options(self): @@ -1121,6 +1363,11 @@ def defunct(self, exc): if self.is_defunct or self.is_closed: return self.is_defunct = True + # Publish the cause atomically with the defunct state. In + # particular, factory may already be awake after READY/auth + # success and must never observe a defunct connection without its + # error. + self.last_error = exc exc_info = sys.exc_info() # if we are not handling an exception, just use the passed exception, and don't try to format exc_info with the message @@ -1131,13 +1378,26 @@ def defunct(self, exc): log.debug("Defuncting connection (%s) to %s: %s", id(self), self.endpoint, exc) - self.last_error = exc self.close() self.error_all_cp_sessions(exc) self.error_all_requests(exc) self.connected_event.set() return exc + def _mark_startup_completed(self): + """ + Atomically publish successful startup and wake factory waiters. + + A concurrent close which wins the lock remains a pre-startup close; + READY or AUTH_SUCCESS received after that close must not reclassify it. + """ + with self.lock: + if self.is_closed or self.is_defunct: + return False + self._startup_completed = True + self.connected_event.set() + return True + def error_all_cp_sessions(self, exc): stream_ids = list(self._continuous_paging_sessions.keys()) for stream_id in stream_ids: @@ -1205,6 +1465,10 @@ def handle_pushed(self, response): log.exception("Pushed event handler errored, ignoring:") def send_msg(self, msg, request_id, cb, encoder=ProtocolHandler.encode_message, decoder=ProtocolHandler.decode_message, result_metadata=None): + send_state = getattr(self, '_send_msg_states', {}).get(request_id) + if send_state is not None: + send_state['phase'] = 'before_registration' + if self.is_defunct: msg = "Connection to %s is defunct" % self.endpoint if self.last_error: @@ -1221,6 +1485,8 @@ def send_msg(self, msg, request_id, cb, encoder=ProtocolHandler.encode_message, # queue the decoder function with the request # this allows us to inject custom functions per request to encode, decode messages self._requests[request_id] = (cb, decoder, result_metadata) + if send_state is not None: + send_state['phase'] = 'registered' msg = encoder(msg, request_id, self.protocol_version, compressor=self.compressor, allow_beta_protocol_version=self.allow_beta_protocol_version, protocol_features=self.features) @@ -1230,7 +1496,14 @@ def send_msg(self, msg, request_id, cb, encoder=ProtocolHandler.encode_message, self._segment_codec.encode(buffer, msg) msg = buffer.getvalue() + if send_state is not None: + # Once push starts, an exception no longer proves that no bytes + # reached the reactor/socket. The request id must not be reused on + # this connection after such an ambiguous failure. + send_state['phase'] = 'push_started' self.push(msg) + if send_state is not None: + send_state['phase'] = 'complete' return len(msg) def wait_for_response(self, msg, timeout=None, **kwargs): @@ -1284,6 +1557,11 @@ def wait_for_responses(self, *msgs, **kwargs): return waiter.deliver(timeout) except OperationTimedOut: raise + except (RequestValidationException, + ProtocolRequestValidationException): + # Query validation reflects request/session state, not a broken + # transport. Keep the connection open for subsequent requests. + raise except Exception as exc: self.defunct(exc) raise @@ -1455,17 +1733,27 @@ def process_msg(self, header, body): def new_continuous_paging_session(self, stream_id, decoder, row_factory, state): session = ContinuousPagingSession(stream_id, decoder, row_factory, self, state) - self._continuous_paging_sessions[stream_id] = session + with self.lock: + self._continuous_paging_sessions[stream_id] = session return session def remove_continuous_paging_session(self, stream_id): - try: - self._continuous_paging_sessions.pop(stream_id) - with self.lock: + notify_owner = None + with self.lock: + try: + self._continuous_paging_sessions.pop(stream_id) log.debug("Returning cp session stream id %s", stream_id) self.request_ids.append(stream_id) - except KeyError: - pass + if not self._continuous_paging_sessions: + notify_owner = getattr( + self._owning_pool, 'on_connection_released', None) + except KeyError: + return + + # HostConnection retirement takes pool -> connection locks. Notify + # only after releasing the connection lock to preserve that order. + if notify_owner is not None: + notify_owner(self) @defunct_on_error def _send_options_message(self): @@ -1584,7 +1872,7 @@ def _handle_startup_response(self, startup_response, did_authenticate=False): if ProtocolVersion.has_checksumming_support(self.protocol_version): self._enable_checksumming() - self.connected_event.set() + self._mark_startup_completed() elif isinstance(startup_response, AuthenticateMessage): log.debug("Got AuthenticateMessage on new connection (%s) from %s: %s", id(self), self.endpoint, startup_response.authenticator) @@ -1640,7 +1928,7 @@ def _handle_auth_response(self, auth_response): self.authenticator.on_authentication_success(auth_response.token) if self._compressor: self.compressor = self._compressor - self.connected_event.set() + self._mark_startup_completed() elif isinstance(auth_response, AuthChallengeMessage): response = self.authenticator.evaluate_challenge(auth_response.challenge) msg = AuthResponseMessage("" if response is None else response) @@ -1660,7 +1948,7 @@ def _handle_auth_response(self, auth_response): log.error(msg, self.endpoint, auth_response) raise ProtocolError(msg % (self.endpoint, auth_response)) - def set_keyspace_blocking(self, keyspace): + def set_keyspace_blocking(self, keyspace, timeout=None): if not keyspace or keyspace == self.keyspace: return @@ -1668,10 +1956,11 @@ def set_keyspace_blocking(self, keyspace): query = QueryMessage(query='USE %s' % (escape_name(keyspace),), consistency_level=ConsistencyLevel.ONE) try: - result = self.wait_for_response(query) - except InvalidRequestException as ire: - # the keyspace probably doesn't exist - raise ire.to_exception() + result = self.wait_for_response(query, timeout=timeout) + except ProtocolRequestValidationException as validation_error: + raise validation_error.to_exception() + except RequestValidationException: + raise except Exception as exc: conn_exc = ConnectionException( "Problem while setting keyspace: %r" % (exc,), self.endpoint) @@ -1709,10 +1998,21 @@ def set_keyspace_async(self, keyspace, callback): # unlikely for a set_keyspace call # - it allows us to avoid signaling a condition every time a request completes while True: + closed_error = None with self.lock: - if self.in_flight < self.max_request_id: + if self.is_closed or self.is_defunct: + # Preserve the in_flight contract for callers which always + # balance this callback through pool.return_connection. + self.in_flight += 1 + closed_error = ConnectionShutdown( + "Connection to %s is closed" % self.endpoint, + self.endpoint) + elif self.in_flight < self.max_request_id: self.in_flight += 1 break + if closed_error is not None: + callback(self, closed_error) + return time.sleep(0.001) if not keyspace or keyspace == self.keyspace: @@ -1727,17 +2027,53 @@ def process_result(result): if isinstance(result, ResultMessage): self.keyspace = keyspace callback(self, None) - elif isinstance(result, InvalidRequestException): + elif isinstance(result, ProtocolRequestValidationException): callback(self, result.to_exception()) + elif isinstance(result, RequestValidationException): + callback(self, result) + elif isinstance(result, Exception): + callback(self, result) else: - callback(self, self.defunct(ConnectionException( - "Problem while setting keyspace: %r" % (result,), self.endpoint))) + conn_exc = ConnectionException( + "Problem while setting keyspace: %r" % (result,), + self.endpoint) + self.defunct(conn_exc) + callback(self, conn_exc) # We've incremented self.in_flight above, so we "have permission" to - # acquire a new request id - request_id = self.get_request_id() - - self.send_msg(query, request_id, process_result) + # acquire a new request id. Keep close() excluded through callback + # registration/write; otherwise close can drain requests between the + # health check and send_msg registering this stream. + request_id = None + send_state = {'phase': 'unknown'} + try: + with self.lock: + request_id = self.get_request_id() + self._send_msg_states[request_id] = send_state + try: + self.send_msg(query, request_id, process_result) + finally: + if self._send_msg_states.get(request_id) is send_state: + del self._send_msg_states[request_id] + except BaseException as exc: + if ( + request_id is not None and + send_state['phase'] in ( + 'before_registration', + 'registered')): + # No push was attempted, so no response can arrive for this + # stream and both pieces of ownership can be safely unwound. + with self.lock: + self._requests.pop(request_id, None) + if request_id not in self.request_ids: + self.request_ids.append(request_id) + elif request_id is not None: + # A custom send_msg implementation, or any failure once push() + # begins, is ambiguous: the peer may still answer. Closing the + # connection prevents that response from being delivered to a + # new request reusing this stream id. + self.defunct(exc) + raise @property def is_idle(self): @@ -1913,6 +2249,20 @@ def run(self): # TODO: move this, along with connection locks in pool, down into Connection with connection.lock: connection.in_flight -= 1 + on_connection_released = getattr( + f.owner, 'on_connection_released', None) + if on_connection_released is not None: + # The pool callback takes pool -> connection locks. + try: + on_connection_released(connection) + except Exception: + # Retirement bookkeeping must not reclassify a + # successful heartbeat as a transport failure; + # in_flight has already been balanced. + log.exception( + "Error releasing connection (%s) after " + "successful heartbeat", + id(connection)) connection.reset_idle() except Exception as e: log.warning("Heartbeat failed for connection (%s) to %s", diff --git a/cassandra/io/twistedreactor.py b/cassandra/io/twistedreactor.py index 446200bf63..a9ed645f5a 100644 --- a/cassandra/io/twistedreactor.py +++ b/cassandra/io/twistedreactor.py @@ -24,6 +24,7 @@ from twisted.internet import reactor, protocol from twisted.internet.endpoints import connectProtocol, TCP4ClientEndpoint, SSL4ClientEndpoint +from twisted.internet.error import ConnectionDone from twisted.internet.interfaces import IOpenSSLClientConnectionCreator from twisted.python.failure import Failure from zope.interface import implementer @@ -198,6 +199,9 @@ def create_timer(cls, timeout, callback): cls._loop.add_timer(timer) return timer + def _is_clean_close_error(self, exc): + return isinstance(exc, ConnectionDone) or Connection._is_clean_close_error(self, exc) + def __init__(self, *args, **kwargs): """ Initialization method. @@ -209,7 +213,10 @@ def __init__(self, *args, **kwargs): """ Connection.__init__(self, *args, **kwargs) - self.is_closed = True + # A scheduled or in-progress endpoint connection is still closeable. + # Keeping it logically open lets pool shutdown cancel it before + # connectionMade. + self.is_closed = False self.connector = None self.transport = None @@ -226,9 +233,12 @@ def _check_pyopenssl(self): def add_connection(self): """ - Convenience function to connect and store the resulting - connector. + Convenience function to connect and store the resulting Deferred. """ + with self.lock: + if self.is_closed: + return + host, port = self.endpoint.resolve() if self.ssl_context or self.ssl_options: # Can't use optionsForClientTLS here because it *forces* hostname verification. @@ -257,7 +267,25 @@ def add_connection(self): port, timeout=self.connect_timeout ) - connectProtocol(endpoint, TwistedConnectionProtocol(self)) + + # connectProtocol returns a cancellable Deferred. Keep the final + # shutdown check and assignment under the lock so close() either stops + # this attempt before it starts or observes the Deferred and cancels it. + with self.lock: + if self.is_closed: + return + self.connector = connectProtocol( + endpoint, TwistedConnectionProtocol(self)) + self.connector.addErrback(self._handle_connect_failure) + + def _handle_connect_failure(self, failure): + # Pool shutdown intentionally cancels the endpoint Deferred. Consume + # that failure instead of leaving an unhandled cancellation in Twisted. + with self.lock: + if self.is_closed: + return None + self.defunct(failure.value) + return None def client_connection_made(self, transport): """ @@ -265,8 +293,16 @@ def client_connection_made(self, transport): succeeded. """ with self.lock: - self.is_closed = False - self.transport = transport + if self.is_closed: + close_transport = True + else: + close_transport = False + self.transport = transport + + if close_transport: + reactor.callFromThread(transport.connector.disconnect) + return + self._send_options_message() def close(self): @@ -277,18 +313,29 @@ def close(self): if self.is_closed: return self.is_closed = True + connector = self.connector + transport = self.transport + + shutdown_error = None + if not self.is_defunct: + msg = "Connection to %s was closed" % self.endpoint + if self.last_error: + msg += ": %s" % (self.last_error,) + shutdown_error = ConnectionShutdown(msg) + if not self._startup_completed: + self.last_error = shutdown_error + # Wake Connection.factory before waiting for reactor cleanup. + self.connected_event.set() log.debug("Closing connection (%s) to %s", id(self), self.endpoint) - reactor.callFromThread(self.transport.connector.disconnect) + if transport is not None: + reactor.callFromThread(transport.connector.disconnect) + elif connector is not None: + reactor.callFromThread(connector.cancel) log.debug("Closed socket to %s", self.endpoint) - if not self.is_defunct: - msg = "Connection to %s was closed" % self.endpoint - if self.last_error: - msg += ": %s" % (self.last_error,) - self.error_all_requests(ConnectionShutdown(msg)) - # don't leave in-progress operations hanging - self.connected_event.set() + if shutdown_error is not None: + self.error_all_requests(shutdown_error) def handle_read(self): """ diff --git a/cassandra/pool.py b/cassandra/pool.py index 176751f60a..68ea510ac5 100644 --- a/cassandra/pool.py +++ b/cassandra/pool.py @@ -15,8 +15,11 @@ """ Connection pooling and host management. """ +from collections import deque from concurrent.futures import Future +from contextlib import ExitStack from functools import total_ordering +import inspect import logging import time import random @@ -28,12 +31,105 @@ except ImportError: from cassandra.util import WeakSet # NOQA -from cassandra import AuthenticationFailed -from cassandra.connection import ConnectionException, EndPoint, DefaultEndPoint +from cassandra import AuthenticationFailed, RequestValidationException +from cassandra.connection import (ConnectionException, EndPoint, DefaultEndPoint, + _ConnectionClosedDuringStartup, + _set_keyspace_blocking, + _startup_close_error) from cassandra.policies import HostDistance +from cassandra.protocol import ( + RequestValidationException as ProtocolRequestValidationException) log = logging.getLogger(__name__) +_REQUEST_VALIDATION_EXCEPTIONS = ( + RequestValidationException, + ProtocolRequestValidationException) + + +def _signal_connection_failure( + cluster, host, connection_exc, is_host_addition, + expect_host_to_be_down=False, force=False): + """Invoke the failure hook without breaking pre-``force`` overrides.""" + method = cluster.signal_connection_failure + try: + parameters = tuple(inspect.signature(method).parameters.values()) + except (TypeError, ValueError): + # This is the historical four-positional-argument API. In particular, + # do not guess that an uninspectable override accepts ``force``. + return method( + host, + connection_exc, + is_host_addition, + expect_host_to_be_down) + + positional_parameters = tuple( + parameter for parameter in parameters + if parameter.kind in ( + inspect.Parameter.POSITIONAL_ONLY, + inspect.Parameter.POSITIONAL_OR_KEYWORD)) + has_varargs = any( + parameter.kind == inspect.Parameter.VAR_POSITIONAL + for parameter in parameters) + has_varkw = any( + parameter.kind == inspect.Parameter.VAR_KEYWORD + for parameter in parameters) + parameters_by_name = { + parameter.name: parameter for parameter in parameters} + + args = [host, connection_exc] + kwargs = {} + addition_parameter = parameters_by_name.get('is_host_addition') + if ( + addition_parameter is not None and + addition_parameter.kind == inspect.Parameter.KEYWORD_ONLY): + kwargs['is_host_addition'] = is_host_addition + elif len(positional_parameters) >= 3 or has_varargs: + args.append(is_host_addition) + elif has_varkw: + kwargs['is_host_addition'] = is_host_addition + else: + # Preserve the historical positional call and let an actually + # incompatible override raise a useful TypeError. + args.append(is_host_addition) + + expect_parameter = parameters_by_name.get('expect_host_to_be_down') + if ( + expect_parameter is not None and + expect_parameter.kind in ( + inspect.Parameter.POSITIONAL_OR_KEYWORD, + inspect.Parameter.KEYWORD_ONLY)): + kwargs['expect_host_to_be_down'] = expect_host_to_be_down + elif len(positional_parameters) >= 4: + # Some historical overrides renamed this argument. Pass it by + # position so those overrides continue to work. + args.append(expect_host_to_be_down) + elif has_varkw: + kwargs['expect_host_to_be_down'] = expect_host_to_be_down + elif has_varargs: + args.append(expect_host_to_be_down) + + force_parameter = parameters_by_name.get('force') + if ( + force_parameter is not None and + force_parameter.kind == inspect.Parameter.POSITIONAL_ONLY): + # A positional-only force parameter necessarily follows the + # historical expectation slot. + if len(args) == 3 and 'expect_host_to_be_down' in kwargs: + args.append(expect_host_to_be_down) + kwargs.pop('expect_host_to_be_down', None) + args.append(force) + elif ( + has_varkw or + ( + force_parameter is not None and + force_parameter.kind in ( + inspect.Parameter.POSITIONAL_OR_KEYWORD, + inspect.Parameter.KEYWORD_ONLY))): + kwargs['force'] = force + + return method(*args, **kwargs) + class NoConnectionsAvailable(Exception): """ @@ -162,6 +258,18 @@ class Host(object): lock = None _currently_handling_node_up = False + _currently_handling_node_add = False + _recovery_epoch = 0 + _is_removed = False + _transition_lock = None + _transition_queue = None + _transition_running = False + _transition_notification_queue = None + _transition_notification_running = False + _transition_event_sequence = 0 + _latest_endpoint_change_sequence = 0 + _latest_preemptive_down_sequence = 0 + _pending_down_notification_epoch = None sharding_info = None @@ -178,6 +286,23 @@ def __init__(self, endpoint, conviction_policy_factory, datacenter=None, rack=No self.host_id = host_id self.set_location_info(datacenter, rack) self.lock = RLock() + # Cluster topology side effects are drained serially without holding + # this lock while callbacks execute. Reentrant/cross-host callbacks + # enqueue and return instead of nesting host locks. + self._transition_lock = RLock() + self._transition_queue = deque() + self._transition_running = False + self._transition_owner = None + self._transition_notification_queue = deque() + self._transition_notification_running = False + self._transition_event_sequence = 0 + self._latest_endpoint_change_sequence = 0 + self._latest_preemptive_down_sequence = 0 + self._pending_down_notification_epoch = None + self._recovery_epoch = 0 + self._currently_handling_node_up = False + self._currently_handling_node_add = False + self._is_removed = False @property def address(self): @@ -231,6 +356,14 @@ def get_and_set_reconnection_handler(self, new_handler): self._reconnection_handler = new_handler return old + def clear_reconnection_handler(self, expected_handler): + """Clear a reconnector only if it still owns the host slot.""" + with self.lock: + if self._reconnection_handler is expected_handler: + self._reconnection_handler = None + return expected_handler + return None + def __eq__(self, other): if isinstance(other, Host): return self.endpoint == other.endpoint @@ -279,6 +412,7 @@ def run(self): return conn = None + connection_transferred = False try: conn = self.try_reconnect() except Exception as exc: @@ -298,10 +432,10 @@ def run(self): self.scheduler.schedule(next_delay, self.run) else: if not self._cancelled: - self.on_reconnection(conn) + connection_transferred = self.on_reconnection(conn) is True self.callback(*(self.callback_args), **(self.callback_kwargs)) finally: - if conn: + if conn and not connection_transferred: conn.close() def cancel(self): @@ -351,7 +485,12 @@ def __init__(self, host, connection_factory, is_host_addition, on_add, on_up, *a self.connection_factory = connection_factory def try_reconnect(self): - return self.connection_factory() + connection = self.connection_factory() + if connection.is_closed: + error = _startup_close_error(connection, self.host.endpoint) + connection.close() + raise error + return connection def on_reconnection(self, connection): log.info("Successful reconnection to %s, marking node up if it isn't already", self.host) @@ -361,7 +500,13 @@ def on_reconnection(self, connection): self.on_up(self.host) def on_exception(self, exc, next_delay): - if isinstance(exc, AuthenticationFailed): + if isinstance(exc, _ConnectionClosedDuringStartup): + log.info( + "Connection to %s closed cleanly during startup; leaving the " + "host down until an UP event or verified control reconnect", + self.host) + return False + elif isinstance(exc, AuthenticationFailed): return False else: log.warning("Error attempting to reconnect to %s, scheduling retry in %s seconds: %s", @@ -391,17 +536,51 @@ class HostConnection(object): tablets_routing_v1 = False - def __init__(self, host, host_distance, session): + @staticmethod + def _session_keyspace_snapshot(session): + """ + Return the Session keyspace together with its change generation. + + Real Sessions publish both values under their lock. Lightweight + Session stand-ins used by applications and tests may not expose the + generation, in which case comparing the keyspace itself still + provides the historical best-effort behavior. + """ + lock = getattr(session, '_lock', None) + generation = getattr(session, '_keyspace_generation', None) + if ( + isinstance(generation, int) and + hasattr(lock, '__enter__')): + with lock: + return ( + session.keyspace, + getattr(session, '_keyspace_generation', generation)) + return session.keyspace, None + + def __init__( + self, host, host_distance, session, pool_generation=None, + host_recovery_epoch=None): self.host = host self.host_distance = host_distance self._session = weakref.proxy(session) + # Session supplies the generation for pools built asynchronously. + # This lets a shard worker hand recovery off even if it fails before + # the candidate pool is published. + self._pool_generation = pool_generation + self._host_recovery_epoch = host_recovery_epoch self._lock = Lock() # this is used in conjunction with the connection streams. Not using the connection lock because the connection can be replaced in the lifetime of the pool. self._stream_available_condition = Condition(Lock()) self._is_replacing = False + self._regular_replacement_future = None self._connecting = set() + # Maps a shard id to the unique token for the attempt which currently + # owns its `_connecting` marker. `_connecting` is retained for + # compatibility with code which introspects pool state. + self._shard_connection_attempts = {} self._connections = {} self._pending_connections = [] + self._shutdown_owned_pending_ids = set() # A pool of additional connections which are not used but affect how Scylla # assigns shards to them. Scylla tends to assign the shard which has # the lowest number of connections. If connections are not distributed @@ -417,6 +596,12 @@ def __init__(self, host, host_distance, session): self._trash = set() self._shard_connections_futures = [] self.advanced_shardaware_block_until = 0 + self._keyspace_generation = 0 + self._keyspace_update_queue = deque() + self._keyspace_update_in_progress = False + self._keyspace_update_runner_active = False + self._keyspace_update_run_requested = False + self._keyspace_update_current = None if host_distance == HostDistance.IGNORED: log.debug("Not opening connection to ignored host %s", self.host) @@ -426,28 +611,433 @@ def __init__(self, host, host_distance, session): return log.debug("Initializing connection for host %s", self.host) - first_connection = session.cluster.connection_factory(self.host.endpoint, on_orphaned_stream_released=self.on_orphaned_stream_released) - log.debug("First connection created to %s for shard_id=%i", self.host, first_connection.features.shard_id) - self._connections[first_connection.features.shard_id] = first_connection - self._keyspace = session.keyspace - - if self._keyspace: - first_connection.set_keyspace_blocking(self._keyspace) - if first_connection.features.sharding_info and not self._session.cluster.shard_aware_options.disable: - self.host.sharding_info = first_connection.features.sharding_info - self._open_connections_for_all_shards(first_connection.features.shard_id) - self.tablets_routing_v1 = first_connection.features.tablets_routing_v1 + first_connection = None + try: + first_connection = session.cluster.connection_factory( + self.host.endpoint, + host_conn=self, + on_orphaned_stream_released=self.on_orphaned_stream_released) + if not self._register_pending_connection(first_connection): + if not self._shutdown_owns_pending_connection( + first_connection): + first_connection.close() + raise ConnectionException( + "Pool for %s was shutdown during initialization" % + (self.host,), + self.host) + if first_connection.is_closed: + raise _startup_close_error(first_connection, self.host.endpoint) + log.debug("First connection created to %s for shard_id=%i", self.host, first_connection.features.shard_id) + self._keyspace, keyspace_generation = \ + self._session_keyspace_snapshot(session) + + while self._keyspace: + try: + _set_keyspace_blocking( + first_connection, + self._keyspace, + session.cluster.connect_timeout) + break + except _REQUEST_VALIDATION_EXCEPTIONS: + current_keyspace, current_generation = \ + self._session_keyspace_snapshot(session) + if ( + current_keyspace == self._keyspace and + ( + keyspace_generation is None or + current_generation == keyspace_generation)): + raise + # The failed USE targeted a keyspace generation that is no + # longer current. Reconcile on this still-owned socket + # before deciding that pool construction failed. + self._keyspace = current_keyspace + keyspace_generation = current_generation + + with self._lock: + if self.is_shutdown: + raise ConnectionException( + "Pool for %s was shutdown during initialization" % + (self.host,), + self.host) + if not self._remove_pending_connection_locked( + first_connection): + raise ConnectionException( + "Initial connection ownership was lost", + self.host.endpoint) + self._connections[ + first_connection.features.shard_id] = first_connection + + if first_connection.features.sharding_info and not self._session.cluster.shard_aware_options.disable: + self.host.sharding_info = first_connection.features.sharding_info + self._open_connections_for_all_shards(first_connection.features.shard_id) + self.tablets_routing_v1 = first_connection.features.tablets_routing_v1 + except BaseException: + # A constructor failure leaves no published pool for Session to + # shut down. Take ownership here, including any shard attempts + # which may already have been submitted. + self.shutdown() + raise 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 _is_shard_aware(self): + return bool( + self.host.sharding_info and + not self._session.cluster.shard_aware_options.disable) + + def _remove_pending_connection_locked(self, connection): + for i, pending in enumerate(self._pending_connections): + if pending is connection: + del self._pending_connections[i] + return True + return False + + def _register_pending_connection(self, connection): + """Publish a factory-owned socket atomically with shutdown.""" + with self._lock: + if self.is_shutdown: + return False + try: + if connection._owning_pool is None: + connection._owning_pool = self + except AttributeError: + # Preserve compatibility with third-party Connection + # stand-ins which do not expose the optional owner hook. + pass + if any( + pending is connection + for pending in self._pending_connections): + # Connection.factory keeps a successful candidate registered + # until its caller adopts it. Recognize that handoff without + # inserting the same socket twice. + return True + self._pending_connections.append(connection) + return True + + def _unregister_pending_connection(self, connection): + """Release factory ownership without racing adoption or shutdown.""" + with self._lock: + return self._remove_pending_connection_locked(connection) + + def _shutdown_owns_pending_connection(self, connection): + """Return whether shutdown already claimed this pending candidate.""" + with self._lock: + return id(connection) in self._shutdown_owned_pending_ids + + def _handoff_replacement_failure(self, exc): + """ + Hand an unusable replacement to host recovery. + + Validate in the cluster-wide Cluster-lock, Host-lock, then Session-lock + order. The user-extensible conviction policy runs without those locks; + the pool is revalidated before committing the forced DOWN transition. + """ + try: + cluster_lock = self._session.cluster._lock + host_lock = self.host.lock + session_lock = self._session._lock + pools = self._session._pools + except (AttributeError, ReferenceError): + cluster_lock = None + host_lock = None + session_lock = None + pools = None + if not hasattr(cluster_lock, '__enter__'): + cluster_lock = None + if not hasattr(host_lock, '__enter__'): + host_lock = None + if not hasattr(session_lock, '__enter__'): + session_lock = None + + # Some unit-test and third-party Session stand-ins do not expose a real + # pool mapping. Preserve the historical behavior for those objects. + has_current_pool_fence = ( + isinstance(pools, dict) and + session_lock is not None) + cluster = self._session.cluster + down_locked = getattr(type(cluster), '_on_down_locked', None) + uses_default_hooks = getattr( + cluster, '_uses_default_failure_hooks', None) + can_use_private_down = bool( + down_locked is not None and + callable(uses_default_hooks) and + uses_default_hooks() is True) + + if not has_current_pool_fence: + if self.is_shutdown: + return False + _signal_connection_failure( + cluster, + self.host, + exc, + is_host_addition=False, + expect_host_to_be_down=True, + force=True) + return True + + locks = [ + lock for lock in (cluster_lock, host_lock, session_lock) + if lock is not None and ( + lock is not session_lock or has_current_pool_fence)] + + def is_current_locked(): + if self.is_shutdown: + return False + if ( + self._host_recovery_epoch is not None and + getattr( + self.host, + '_recovery_epoch', + self._host_recovery_epoch) != + self._host_recovery_epoch): + return False + current_pool = pools.get(self.host) + if current_pool is None: + for pool_host, pool in tuple(pools.items()): + if pool_host is self.host: + current_pool = pool + break + generation_is_current = self._pool_generation is None + if self._pool_generation is not None: + get_generation = getattr( + self._session, '_get_pool_generation_locked', None) + if get_generation is not None: + generation_is_current = ( + get_generation(self.host) == + self._pool_generation) + candidate_is_owned = bool( + current_pool is self or + self._pool_generation is not None) + return not ( + self._session.is_shutdown or + not generation_is_current or + not candidate_is_owned) + + with ExitStack() as stack: + for lock in locks: + stack.enter_context(lock) + if not is_current_locked(): + return False + + if not can_use_private_down: + # Notify legacy/custom public hooks through their compatible + # signature first. If they cannot express the newer + # authoritative ``force`` transition, commit the inherited + # private DOWN path while this pool is still current. + recovery_epoch = getattr(self.host, '_recovery_epoch', None) + if not isinstance(recovery_epoch, int): + recovery_epoch = None + try: + _signal_connection_failure( + cluster, + self.host, + exc, + is_host_addition=False, + expect_host_to_be_down=True, + force=True) + finally: + if down_locked is not None: + with ExitStack() as stack: + for lock in locks: + stack.enter_context(lock) + recovery_was_started = bool( + recovery_epoch is not None and + getattr( + self.host, + '_recovery_epoch', + recovery_epoch) != recovery_epoch) + if ( + is_current_locked() and + not recovery_was_started): + down_locked( + cluster, + self.host, + is_host_addition=False, + expect_host_to_be_down=True, + force=True) + return True + + # Conviction policies may call back into Session/Cluster code. This + # recovery handoff is authoritative, so a broken user policy must not + # strand an empty pool without a reconnector. + try: + self.host.signal_connection_failure(exc) + except BaseException: + log.exception( + "Conviction policy failed while handing replacement failure " + "for host %s to recovery", + self.host) + + with ExitStack() as stack: + for lock in locks: + stack.enter_context(lock) + if not is_current_locked(): + return False + down_locked( + cluster, + self.host, + is_host_addition=False, + expect_host_to_be_down=True, + force=True) + return True + + def _retire_connection_locked(self, connection): + """ + Stop using a connection and atomically decide who closes it. + + The pool lock must be held by the caller. Taking the connection lock + while publishing it to `_trash` closes the race with its last return. + """ + with connection.lock: + connection._pool_retired = True + remaining = self._connection_active_uses_locked(connection) + if remaining <= 0: + return True, remaining + self._trash.add(connection) + return False, remaining + + @staticmethod + def _connection_active_uses_locked(connection): + """ + Count work which still requires an open retired connection. + + Continuous paging keeps its stream after the initial request has been + returned to the pool, so it is not represented by ``in_flight``. + """ + pending_requests = max( + connection.in_flight - + len(connection.orphaned_request_ids), + 0) + paging_sessions = len(getattr( + connection, '_continuous_paging_sessions', ())) + return pending_requests + paging_sessions + + @staticmethod + def _connection_needs_replacement(connection): + return bool( + connection.is_closed or + connection.is_defunct or + getattr(connection, '_pool_retired', False) is True or + connection.orphaned_threshold_reached) + + def _schedule_connection_to_missing_shard( + self, shard_id, is_replacement=False, track_future=False): + """ + Start at most one connection attempt for a shard. + + A replacement request promotes an already-running optimistic attempt, + so its eventual failure is handed to host recovery. + """ + with self._lock: + if self.is_shutdown: + return None + if is_replacement: + current = self._connections.get(shard_id) + if ( + current is not None and + not self._connection_needs_replacement(current)): + return None + + attempt = self._shard_connection_attempts.get(shard_id) + if attempt is not None: + if is_replacement: + attempt['is_replacement'] = True + return attempt.get('future') + + token = object() + attempt = { + 'token': token, + 'is_replacement': is_replacement, + 'future': None, + } + self._shard_connection_attempts[shard_id] = attempt + self._connecting.add(shard_id) + + # Session.submit may be implemented by an application and execute the + # work synchronously. Never invoke it while holding the pool lock. + try: + future = self._session.submit( + self._open_connection_to_missing_shard, + shard_id, + token) + except BaseException: + self._finish_shard_connection_attempt(shard_id, token) + raise + + with self._lock: + current_attempt = self._shard_connection_attempts.get(shard_id) + if ( + current_attempt is not None and + current_attempt['token'] is token): + if future is None: + del self._shard_connection_attempts[shard_id] + self._connecting.discard(shard_id) + return None + current_attempt['future'] = future + if track_future and isinstance(future, Future): + self._shard_connections_futures.append(future) + + if isinstance(future, Future): + def attempt_completed(completed_future): + if not completed_future.cancelled(): + return + should_recover = False + was_replacement = self._finish_shard_connection_attempt( + shard_id, token) + with self._lock: + should_recover = bool( + was_replacement and + not self.is_shutdown) + if should_recover: + try: + self._handoff_replacement_failure( + ConnectionException( + "Shard replacement task was cancelled", + self.host)) + except BaseException: + log.exception( + "Failed handing cancelled shard replacement for " + "host %s to recovery", + self.host) + + future.add_done_callback(attempt_completed) + return future + + def _finish_shard_connection_attempt(self, shard_id, token): + with self._lock: + attempt = self._shard_connection_attempts.get(shard_id) + if attempt is None or attempt['token'] is not token: + return False + is_replacement = attempt['is_replacement'] + del self._shard_connection_attempts[shard_id] + self._connecting.discard(shard_id) + return is_replacement + + def _claim_failed_shard_attempt(self, shard_id, token): + """ + Atomically classify a failed attempt against replacement promotion. + + Replacement attempts remain published until host recovery is + signaled. Opportunistic attempts are removed immediately so a later + request can retry. This makes promotion and failure linearizable. + """ + with self._lock: + attempt = self._shard_connection_attempts.get(shard_id) + if attempt is None or attempt['token'] is not token: + return False + if attempt['is_replacement']: + return True + del self._shard_connection_attempts[shard_id] + self._connecting.discard(shard_id) + return False + + def _get_connection_for_routing_key( + self, routing_key=None, keyspace=None, table=None, + allow_keyspace_mismatch=False): if self.is_shutdown: raise ConnectionException( "Pool for %s is shutdown" % (self.host,), self.host) - if not self._connections: - raise NoConnectionsAvailable() - shard_id = None if not self._session.cluster.shard_aware_options.disable and self.host.sharding_info and routing_key: t = self._session.cluster.metadata.token_map.token_class.from_key(routing_key) @@ -468,59 +1058,214 @@ def _get_connection_for_routing_key(self, routing_key=None, keyspace=None, table if shard_id is None: shard_id = self.host.sharding_info.shard_id_from_token(t.value) - conn = self._connections.get(shard_id) + with self._lock: + conn = self._connections.get(shard_id) + + def keyspace_is_usable(connection): + return bool( + allow_keyspace_mismatch or + getattr( + connection, + '_pool_keyspace_mismatch', + False) is not True) + + # Routed queries normally have a healthy connection for their target + # shard. Keep that path O(1); the borrow-side validation below closes + # races with replacement or shutdown after this snapshot. + if ( + shard_id is not None and + conn and + not self._connection_needs_replacement(conn) and + keyspace_is_usable(conn)): + return conn + + with self._lock: + connections = list(self._connections.values()) + + # A connection can close while idle, so there may be no request + # callback to return it to the pool and trigger replacement. Retire + # unhealthy mappings during selection, including unkeyed requests. + unhealthy_connections = [] + seen = set() + for connection in connections: + if ( + id(connection) not in seen and + (connection.is_closed or connection.is_defunct)): + seen.add(id(connection)) + unhealthy_connections.append(connection) + for connection in unhealthy_connections: + self.return_connection( + connection, stream_was_orphaned=True) + + if unhealthy_connections: + if self.is_shutdown: + raise ConnectionException( + "Pool for %s is shutdown" % (self.host,), + self.host) + with self._lock: + conn = self._connections.get(shard_id) + connections = list(self._connections.values()) + + shard_aware = self._is_shard_aware() + if shard_aware: + for connection in connections: + if ( + not (connection.is_closed or connection.is_defunct) and + connection.orphaned_threshold_reached): + try: + self._schedule_connection_to_missing_shard( + connection.features.shard_id, + is_replacement=True) + except BaseException as exc: + self._handoff_replacement_failure(exc) + if not isinstance(exc, Exception): + raise + else: + regular_replacement = next( + ( + connection for connection in connections + if ( + not ( + connection.is_closed or + connection.is_defunct) and + connection.orphaned_threshold_reached) + ), + None) + if regular_replacement is not None: + self._submit_regular_replacement(regular_replacement) # missing shard aware connection to shard_id, let's schedule an # optimistic try to connect to it if shard_id is not None: - if conn: - if conn.orphaned_threshold_reached and shard_id not in self._connecting: - # The connection has met its orphaned stream ID limit - # and needs to be replaced. Start opening a connection - # to the same shard and replace when it is opened. - self._connecting.add(shard_id) - self._session.submit(self._open_connection_to_missing_shard, shard_id) + needs_replacement = bool( + conn and self._connection_needs_replacement(conn)) + if conn is None or needs_replacement: + try: + self._schedule_connection_to_missing_shard( + shard_id, + is_replacement=needs_replacement) + except BaseException as exc: + if needs_replacement: + self._handoff_replacement_failure(exc) + else: + log.warning( + "Unable to open an optional connection to missing " + "shard %i on host %s; using another shard", + shard_id, + self.host, + exc_info=True) + if not isinstance(exc, Exception): + raise + if needs_replacement: log.debug( - "Connection to shard_id=%i reached orphaned stream limit, replacing on host %s (%s/%i)", + "Connection to shard_id=%i needs replacement on host %s (%s/%i)", shard_id, self.host, - len(self._connections), + len(connections), + self.host.sharding_info.shards_count + ) + else: + # rate controlled optimistic attempt to connect to a + # missing shard + log.debug( + "Trying to connect to missing shard_id=%i on host %s (%s/%i)", + shard_id, + self.host, + len(connections), self.host.sharding_info.shards_count ) - elif shard_id not in self._connecting: - # rate controlled optimistic attempt to connect to a missing shard - self._connecting.add(shard_id) - self._session.submit(self._open_connection_to_missing_shard, shard_id) - log.debug( - "Trying to connect to missing shard_id=%i on host %s (%s/%i)", - shard_id, - self.host, - len(self._connections), - self.host.sharding_info.shards_count - ) - if conn and not conn.is_closed: + # Session.submit is intentionally allowed to execute synchronously, + # and a fast executor can also install a replacement before scheduling + # returns. Select from the current mappings rather than the stale + # pre-scheduling snapshot. + with self._lock: + if self.is_shutdown: + raise ConnectionException( + "Pool for %s is shutdown" % (self.host,), + self.host) + conn = self._connections.get(shard_id) + connections = list(self._connections.values()) + + if ( + conn and + not self._connection_needs_replacement(conn) and + keyspace_is_usable(conn)): return conn - active_connections = [conn for conn in list(self._connections.values()) if not conn.is_closed] + active_connections = [ + connection for connection in connections + if ( + not self._connection_needs_replacement(connection) and + keyspace_is_usable(connection))] if active_connections: 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) + awaiting_replacement = [ + connection for connection in connections + if not ( + connection.is_closed or + connection.is_defunct or + getattr(connection, '_pool_retired', False) is True or + not keyspace_is_usable(connection))] + if awaiting_replacement: + return random.choice(awaiting_replacement) + raise NoConnectionsAvailable( + "No open connections to host %s" % (self.host,)) + + def borrow_connection( + self, timeout, routing_key=None, keyspace=None, table=None, + allow_keyspace_mismatch=False): + def select_connection(): + return self._get_connection_for_routing_key( + routing_key, + keyspace, + table, + allow_keyspace_mismatch=allow_keyspace_mismatch) + + conn = select_connection() 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) - 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 - # raise Exception - conn.in_flight += 1 - return conn, conn.get_request_id() + if ( + conn.is_closed or + conn.is_defunct or + getattr(conn, '_pool_retired', False) is True or + ( + not allow_keyspace_mismatch and + getattr( + conn, + '_pool_keyspace_mismatch', + False) is True)): + # The connection might have failed or been superseded in the + # meantime; select the current mapping and schedule repair. + conn = select_connection() + with self._lock: + if self.is_shutdown: + raise ConnectionException( + "Pool for %s is shutdown" % (self.host,), + self.host) + current = self._connections.get( + conn.features.shard_id) + with conn.lock: + if ( + current is conn and + not (conn.is_closed or conn.is_defunct) and + getattr( + conn, + '_pool_retired', + False) is not True and + ( + allow_keyspace_mismatch or + getattr( + conn, + '_pool_keyspace_mismatch', + False) is not True) and + conn.in_flight < conn.max_request_id): + conn.in_flight += 1 + return conn, conn.get_request_id() + if current is not conn: + conn = select_connection() + continue if timeout is not None: remaining = timeout - time.time() + start if remaining < 0: @@ -529,11 +1274,37 @@ def borrow_connection(self, timeout, routing_key=None, keyspace=None, table=None break last_retry = True continue + retry_selection = False with self._stream_available_condition: - if conn.orphaned_threshold_reached and conn.is_closed: - conn = self._get_connection() - else: + # Recheck after acquiring the condition so a stream return or + # shutdown between the first check and wait cannot be lost. + with self._lock: + if self.is_shutdown: + raise ConnectionException( + "Pool for %s is shutdown" % (self.host,), + self.host) + current = self._connections.get( + conn.features.shard_id) + with conn.lock: + retry_selection = bool( + current is not conn or + conn.is_closed or + conn.is_defunct or + getattr(conn, '_pool_retired', False) is True or + ( + not allow_keyspace_mismatch and + getattr( + conn, + '_pool_keyspace_mismatch', + False) is True) or + conn.in_flight < conn.max_request_id) + if not retry_selection: self._stream_available_condition.wait(remaining) + if retry_selection: + # Selection may synchronously install a replacement and + # notify this same non-reentrant Condition. + conn = select_connection() + continue raise NoConnectionsAvailable("All request IDs are currently in use") @@ -544,120 +1315,542 @@ def return_connection(self, connection, stream_was_orphaned=False): with self._stream_available_condition: self._stream_available_condition.notify() - if connection.is_defunct or connection.is_closed: - if connection.signaled_error and not self.shutdown_on_error: + connection_is_bad = False + connection_error = None + connection_was_retired = False + retired_connection_needs_close = False + should_signal_error = False + + # Snapshot shutdown and connection health atomically in the pool -> + # connection lock order used by retirement. This prevents a deliberate + # shutdown close from being mistaken for a host failure. + with self._lock: + if self.is_shutdown: return + with connection.lock: + connection_is_bad = bool( + connection.is_defunct or connection.is_closed) + connection_error = connection.last_error + connection_was_retired = ( + getattr(connection, '_pool_retired', False) is True) + retired_connection_needs_close = bool( + connection_was_retired and + not connection.is_closed and + self._connection_active_uses_locked(connection) <= 0) + should_signal_error = bool( + connection_is_bad and + not connection_was_retired and + not connection.signaled_error) + if should_signal_error: + connection.signaled_error = True + + if connection_is_bad and connection_was_retired: + # A replacement deliberately closed this connection after its + # last request decremented in_flight but before this health + # classification. It is not evidence that the host failed. + with self._lock: + self._trash.discard(connection) + if retired_connection_needs_close: + connection.close() + return + if connection_is_bad: is_down = False - if not connection.signaled_error: + if should_signal_error: log.debug("Defunct or closed connection (%s) returned to pool, potentially " "marking host %s as down", id(connection), self.host) - is_down = self.host.signal_connection_failure(connection.last_error) - connection.signaled_error = True + try: + is_down = self.host.signal_connection_failure( + connection_error) + except BaseException: + # Conviction policies are user-extensible. If one fails, + # release this connection's claim so a later return can + # retry instead of permanently suppressing host recovery. + with self._lock: + if not self.is_shutdown: + with connection.lock: + connection.signaled_error = False + raise if self.shutdown_on_error and not is_down: is_down = True if is_down: - self.shutdown() + # Conviction callbacks are user-extensible and run without + # pool locks. A healthy replacement may have been installed + # while the callback was blocked; atomically fence adoption + # and revalidate before committing pool shutdown. + if not self._shutdown_if_connection_failure_current( + connection): + return self._session.cluster.on_down(self.host, is_host_addition=False) - else: - connection.close() - with self._lock: - if self.is_shutdown: - return - self._connections.pop(connection.features.shard_id, None) - if self._is_replacing: - return - self._is_replacing = True - self._session.submit(self._replace, connection) - elif connection in self._trash: - with connection.lock: - no_pending_requests = connection.in_flight <= len(connection.orphaned_request_ids) - if no_pending_requests: - with self._lock: - close_connection = False - if connection in self._trash: - self._trash.remove(connection) - close_connection = True - if close_connection: - log.debug("Closing trashed connection (%s) to %s", id(connection), self.host) - connection.close() + return + + connection.close() + shard_id = connection.features.shard_id + schedule_shard_replacement = False + schedule_regular_replacement = False + with self._lock: + self._trash.discard(connection) + current = self._connections.get(shard_id) + if current is connection: + del self._connections[shard_id] + current = None + + if self.is_shutdown: + return + + # A returned retired connection must not evict or replace the + # healthy connection which superseded it. + if current is None: + if self._is_shard_aware(): + schedule_shard_replacement = True + elif not self._is_replacing: + self._is_replacing = True + schedule_regular_replacement = True + + if schedule_shard_replacement: + try: + self._schedule_connection_to_missing_shard( + shard_id, + is_replacement=True) + except BaseException as exc: + self._handoff_replacement_failure(exc) + if not isinstance(exc, Exception): + raise + elif schedule_regular_replacement: + self._submit_regular_replacement(connection, claimed=True) + return + + self.on_connection_released(connection) + + def on_connection_released(self, connection): + """ + Close a retired connection after its final active use is released. + + Some internal users, notably successful heartbeats and continuous + paging, release connection state without going through + :meth:`return_connection`; they notify the pool through this hook. + """ + close_connection = False + with self._lock: + if connection in self._trash: + with connection.lock: + no_active_uses = ( + self._connection_active_uses_locked(connection) <= 0) + if no_active_uses: + self._trash.remove(connection) + close_connection = True + + if close_connection: + log.debug( + "Closing trashed connection (%s) to %s", + id(connection), + self.host) + connection.close() def on_orphaned_stream_released(self): """ Called when a response for an orphaned stream (timed out on the client side) was received. """ + # The callback predates passing the releasing Connection. Scan only + # retired sockets; unlike the routed query path this is not per-query. + with self._lock: + retired_connections = list(self._trash) + for connection in retired_connections: + self.on_connection_released(connection) with self._stream_available_condition: self._stream_available_condition.notify() - def _replace(self, connection): - with self._lock: - if self.is_shutdown: - return + def _submit_regular_replacement(self, connection, claimed=False): + """ + Submit one replacement for a non-shard-aware connection. - log.debug("Replacing connection (%s) to %s", id(connection), self.host) + ``claimed`` means the caller already set ``_is_replacing`` while + removing a failed mapping. + """ + if not claimed: + shard_id = connection.features.shard_id + with self._lock: + if self.is_shutdown or self._is_replacing: + return None + current = self._connections.get(shard_id) + if current is not connection: + return None + self._is_replacing = True + + try: + future = self._session.submit(self._replace, connection) + except Exception as exc: + with self._lock: + self._is_replacing = False + self._handoff_replacement_failure(exc) + return None + except BaseException as exc: + with self._lock: + self._is_replacing = False try: - if connection.features.shard_id in self._connections: - del self._connections[connection.features.shard_id] - if self.host.sharding_info and not self._session.cluster.shard_aware_options.disable: - self._connecting.add(connection.features.shard_id) - self._session.submit(self._open_connection_to_missing_shard, connection.features.shard_id) - else: - connection = self._session.cluster.connection_factory(self.host.endpoint, - on_orphaned_stream_released=self.on_orphaned_stream_released) - if self._keyspace: - connection.set_keyspace_blocking(self._keyspace) - self._connections[connection.features.shard_id] = connection - except Exception: - log.warning("Failed reconnecting %s. Retrying." % (self.host.endpoint,)) - self._session.submit(self._replace, connection) - else: + self._handoff_replacement_failure(exc) + except BaseException: + log.exception( + "Additional failure handing replacement submission " + "cancellation to host recovery for %s", + self.host) + raise + + if future is None: + with self._lock: self._is_replacing = False - with self._stream_available_condition: - self._stream_available_condition.notify() + return None - def shutdown(self): - log.debug("Shutting down connections to %s", self.host) - with self._lock: - if self.is_shutdown: + if isinstance(future, Future): + with self._lock: + if self._is_replacing: + self._regular_replacement_future = future + + def replacement_completed(completed_future): + cancelled_replacement = False + with self._lock: + if self._regular_replacement_future is not \ + completed_future: + return + self._regular_replacement_future = None + if ( + completed_future.cancelled() and + self._is_replacing and + not self.is_shutdown): + self._is_replacing = False + cancelled_replacement = True + + if cancelled_replacement: + try: + self._handoff_replacement_failure( + ConnectionException( + "Regular replacement task was cancelled", + self.host)) + except BaseException: + log.exception( + "Failed handing cancelled replacement for host " + "%s to recovery", + self.host) + + future.add_done_callback(replacement_completed) + return future + + def _replace(self, connection): + replacement = None + replacement_is_pending = False + replacement_was_adopted = False + retired_connection = None + close_retired_connection = False + abandon_replacement = False + shard_id = connection.features.shard_id + + try: + with self._lock: + if self.is_shutdown: + return + + log.debug("Replacing connection (%s) to %s", id(connection), self.host) + current = self._connections.get(shard_id) + if current is not None and current is not connection and not ( + current.is_closed or + current.is_defunct or + current.orphaned_threshold_reached): + # This work item is stale; preserve the newer healthy + # mapping instead of deleting it by shard id. + self._is_replacing = False + return + shard_aware = self._is_shard_aware() + if shard_aware and current is not None: + del self._connections[shard_id] + retired_connection = current + ( + close_retired_connection, + _ + ) = self._retire_connection_locked(current) + if shard_aware: + # `_is_replacing` is retained only for non-shard pools. + self._is_replacing = False + + if close_retired_connection: + retired_connection.close() + retired_connection = None + close_retired_connection = False + + if shard_aware: + self._schedule_connection_to_missing_shard( + shard_id, + is_replacement=True) return - else: - self.is_shutdown = True + + replacement = self._session.cluster.connection_factory( + self.host.endpoint, + host_conn=self, + on_orphaned_stream_released=self.on_orphaned_stream_released) + if not self._register_pending_connection(replacement): + if not self._shutdown_owns_pending_connection(replacement): + replacement.close() + return + replacement_is_pending = True + if replacement.is_closed: + raise _startup_close_error(replacement, self.host.endpoint) + + while True: + with self._lock: + if self.is_shutdown: + # shutdown() drained the pending list and owns the + # replacement. + return + keyspace = self._keyspace + keyspace_generation = self._keyspace_generation + + if keyspace: + try: + _set_keyspace_blocking( + replacement, + keyspace, + self._session.cluster.connect_timeout) + except _REQUEST_VALIDATION_EXCEPTIONS: + # A dropped/invalid Session keyspace does not make the + # host unavailable. Keep ownership so a later USE can + # repair the Session, but quarantine this socket from + # ordinary requests until it selects a keyspace. + replacement._pool_keyspace_mismatch = True + log.warning( + "Unable to set keyspace %s on replacement " + "connection to %s; quarantining the open " + "connection until a successful USE", + keyspace, + self.host.endpoint) + else: + replacement._pool_keyspace_mismatch = False + if replacement.is_closed: + raise _startup_close_error( + replacement, + self.host.endpoint) + + with self._lock: + if self.is_shutdown: + # shutdown() drained the pending list and owns the + # replacement. + return + if keyspace_generation != self._keyspace_generation: + continue + current = self._connections.get(shard_id) + abandon_replacement = bool( + current is not None and + current is not connection and + not self._connection_needs_replacement(current)) + if not self._remove_pending_connection_locked(replacement): + raise ConnectionException( + "Replacement connection ownership was lost", + self.host.endpoint) + replacement_is_pending = False + + if abandon_replacement: + # Another worker installed a healthy connection while + # this factory/setup was running. Release the pending + # candidate without disturbing the newer mapping. + self._is_replacing = False + else: + if ( + current is not None and + current is not replacement): + retired_connection = current + ( + close_retired_connection, + _ + ) = self._retire_connection_locked(current) + if ( + replacement.features.shard_id != shard_id and + shard_id in self._connections and + self._connections.get(shard_id) is current): + del self._connections[shard_id] + self._connections[ + replacement.features.shard_id] = replacement + replacement_was_adopted = True + self._is_replacing = False + + if abandon_replacement: + replacement.close() + return + break + + if close_retired_connection: + retired_connection.close() + with self._stream_available_condition: self._stream_available_condition.notify_all() + except BaseException as exc: + shutdown_owns_replacement = False + with self._lock: + if replacement_is_pending: + if not self._remove_pending_connection_locked(replacement): + # shutdown() already drained and owns closing it. + shutdown_owns_replacement = self.is_shutdown + is_shutdown = self.is_shutdown + + if ( + replacement is not None and + not replacement_was_adopted and + not shutdown_owns_replacement): + replacement.close() + + if not isinstance(exc, Exception): + if not is_shutdown: + try: + self._handoff_replacement_failure(exc) + except BaseException: + log.exception( + "Additional failure handing replacement " + "cancellation to host recovery for %s", + self.host) + with self._lock: + self._is_replacing = False + raise + + if is_shutdown: + return + + # A replacement failure cannot leave this pool serviceable. Hand + # it to the host reconnector, which supplies the retry backoff. + try: + self._handoff_replacement_failure(exc) + except BaseException: + with self._lock: + self._is_replacing = False + raise + if isinstance(exc, _ConnectionClosedDuringStartup): + log.info( + "Replacement connection to %s closed cleanly during " + "startup; handing recovery to the host reconnector", + self.host.endpoint) + else: + log.warning( + "Failed reconnecting %s; handing recovery to the host " + "reconnector", + self.host.endpoint, + exc_info=True) + + def _begin_shutdown_locked(self): + if self.is_shutdown: + return None - for future in self._shard_connections_futures: - future.cancel() + futures_to_cancel = list(self._shard_connections_futures) + if isinstance(self._regular_replacement_future, Future): + futures_to_cancel.append(self._regular_replacement_future) + connections_to_close = list(self._connections.values()) + connections_to_close.extend(self._pending_connections) + self._shutdown_owned_pending_ids.update( + id(connection) + for connection in self._pending_connections) + connections_to_close.extend(self._excess_connections) + connections_to_close.extend(self._trash) + current_keyspace_update = self._keyspace_update_current + claimed_keyspace_update = False + if current_keyspace_update is not None: + shutdown_error = ConnectionException( + "Pool for %s is shutdown" % (self.host,), + self.host) + claimed_keyspace_update = self._claim_keyspace_update( + current_keyspace_update, shutdown_error) + + # Publish shutdown only after the active keyspace generation has been + # atomically classified. A response which already owned the update + # lock wins before shutdown; no response can win in a post-commit gap. + self.is_shutdown = True + self._connections.clear() + self._pending_connections.clear() + self._excess_connections.clear() + self._trash.clear() + self._shard_connections_futures = [] + self._regular_replacement_future = None + self._shard_connection_attempts.clear() + self._connecting.clear() + self._is_replacing = False + return ( + futures_to_cancel, + connections_to_close, + current_keyspace_update, + claimed_keyspace_update) + + def _finish_shutdown(self, shutdown_state): + ( + futures_to_cancel, + connections_to_close, + current_keyspace_update, + claimed_keyspace_update + ) = shutdown_state - connections_to_close = self._connections.copy() - pending_connections_to_close = self._pending_connections.copy() - self._connections.clear() - self._pending_connections.clear() + with self._stream_available_condition: + self._stream_available_condition.notify_all() # connection.close can call pool.return_connection, which will # obtain self._lock via self._stream_available_condition. # So, it never should be called within self._lock context - for connection in connections_to_close.values(): + seen = set() + for connection in connections_to_close: + if id(connection) in seen: + continue + seen.add(id(connection)) log.debug("Closing connection (%s) to %s", id(connection), self.host) connection.close() - for connection in pending_connections_to_close: - log.debug("Closing pending connection (%s) to %s", id(connection), self.host) - connection.close() + # Future cancellation may synchronously invoke callbacks which reenter + # or block. All sockets are already closed and waiters notified before + # any such callback can delay shutdown. + for future in futures_to_cancel: + future.cancel() + + if claimed_keyspace_update: + self._deliver_keyspace_update(current_keyspace_update) + elif current_keyspace_update is None: + # shutdown may win before the runner publishes its current + # generation. Requesting a run makes every queued callback + # complete in FIFO order with the shutdown error. + self._request_next_keyspace_update() + + def _shutdown_if_connection_failure_current(self, connection): + with self._lock: + if self.is_shutdown: + return False + current = self._connections.get( + connection.features.shard_id) + with connection.lock: + retired = ( + getattr(connection, '_pool_retired', False) is True) + if self._is_shard_aware(): + replacement_candidates = (current,) + else: + # A regular Scylla connection can land on a different random + # shard id even though this pool is operating in non-shard + # mode. Any healthy mapped socket supersedes the failed one. + replacement_candidates = tuple( + self._connections.values()) + healthy_replacement = any( + candidate is not None and + candidate is not connection and + not self._connection_needs_replacement(candidate) + for candidate in replacement_candidates) + if retired or healthy_replacement: + return False + shutdown_state = self._begin_shutdown_locked() - self._close_excess_connections() + log.debug( + "Shutting down connections to %s after connection failure", + self.host) + self._finish_shutdown(shutdown_state) + return True - trash_conns = None + def shutdown(self): + log.debug("Shutting down connections to %s", self.host) with self._lock: - if self._trash: - trash_conns = self._trash - self._trash = set() - - if trash_conns: - for conn in trash_conns: - conn.close() + shutdown_state = self._begin_shutdown_locked() + if shutdown_state is None: + return + self._finish_shutdown(shutdown_state) def _close_excess_connections(self): with self._lock: @@ -700,7 +1893,82 @@ def _get_shard_aware_endpoint(self): return endpoint - def _open_connection_to_missing_shard(self, shard_id): + def _open_connection_to_missing_shard(self, shard_id, attempt_token=None): + # Direct callers from older integrations did not pass a token. Give + # them ownership only when there is no real attempt already running. + if attempt_token is None: + with self._lock: + attempt = self._shard_connection_attempts.get(shard_id) + if attempt is None: + attempt_token = object() + self._shard_connection_attempts[shard_id] = { + 'token': attempt_token, + 'is_replacement': False, + 'future': None, + } + self._connecting.add(shard_id) + else: + # A direct/stale invocation must not clear the live + # attempt's marker when it finishes. + attempt_token = object() + + try: + result = self._open_connection_to_missing_shard_impl(shard_id) + except Exception as exc: + clean_startup_close = isinstance( + exc, + _ConnectionClosedDuringStartup) + if clean_startup_close: + # A clean close during startup is authoritative maintenance + # evidence, even if this began as an optimistic attempt. Keep + # the marker published until the handoff completes. + is_replacement = True + else: + is_replacement = self._claim_failed_shard_attempt( + shard_id, + attempt_token) + try: + if is_replacement: + self._handoff_replacement_failure(exc) + log.warning( + "Failed replacing connection to shard %i on host %s; " + "handing recovery to the host reconnector", + shard_id, + self.host, + exc_info=( + type(exc), + exc, + exc.__traceback__)) + return None + + raise + finally: + self._finish_shard_connection_attempt( + shard_id, + attempt_token) + except BaseException as exc: + is_replacement = self._claim_failed_shard_attempt( + shard_id, + attempt_token) + try: + if is_replacement: + try: + self._handoff_replacement_failure(exc) + except BaseException: + log.exception( + "Additional failure handing interrupted shard " + "replacement to host recovery for %s", + self.host) + finally: + self._finish_shard_connection_attempt( + shard_id, + attempt_token) + raise + else: + self._finish_shard_connection_attempt(shard_id, attempt_token) + return result + + def _open_connection_to_missing_shard_impl(self, shard_id): """ Creates a new connection, checks its shard_id and populates our shard aware connections if the current shard_id is missing a connection. @@ -721,137 +1989,192 @@ def _open_connection_to_missing_shard(self, shard_id): with self._lock: if self.is_shutdown: return + shard_aware_endpoint = self._get_shard_aware_endpoint() log.debug("shard_aware_endpoint=%r", shard_aware_endpoint) - if shard_aware_endpoint: - try: - conn = self._session.cluster.connection_factory(shard_aware_endpoint, host_conn=self, on_orphaned_stream_released=self.on_orphaned_stream_released, - shard_id=shard_id, - total_shards=self.host.sharding_info.shards_count) - conn.original_endpoint = self.host.endpoint - except Exception as exc: - log.error("Failed to open connection to %s, on shard_id=%i: %s", self.host, shard_id, exc) - raise - else: - conn = self._session.cluster.connection_factory(self.host.endpoint, host_conn=self, on_orphaned_stream_released=self.on_orphaned_stream_released) - - log.debug( - "Received a connection %s for shard_id=%i on host %s", - id(conn), - conn.features.shard_id if conn.features.shard_id is not None else -1, - self.host) - if self.is_shutdown: - log.debug("Pool for host %s is in shutdown, closing the new connection (%s)", self.host, id(conn)) - conn.close() - return + conn = None + pending = False + adopted = False + shutdown_owns_connection = False + try: + if shard_aware_endpoint: + conn = self._session.cluster.connection_factory( + shard_aware_endpoint, + host_conn=self, + on_orphaned_stream_released=self.on_orphaned_stream_released, + shard_id=shard_id, + total_shards=self.host.sharding_info.shards_count) + else: + conn = self._session.cluster.connection_factory( + self.host.endpoint, + host_conn=self, + on_orphaned_stream_released=self.on_orphaned_stream_released) + + if not self._register_pending_connection(conn): + if not self._shutdown_owns_pending_connection(conn): + conn.close() + return + pending = True + if conn.is_closed: + raise _startup_close_error(conn, self.host.endpoint) - if shard_aware_endpoint and shard_id != conn.features.shard_id: - # connection didn't land on expected shared - # assuming behind a NAT, disabling advanced shard aware for a while - self.disable_advanced_shard_aware(10 * 60) + if shard_aware_endpoint: + conn.original_endpoint = self.host.endpoint - old_conn = self._connections.get(conn.features.shard_id) - if old_conn is None or old_conn.orphaned_threshold_reached: + actual_shard_id = conn.features.shard_id log.debug( - "New connection (%s) created to shard_id=%i on host %s", + "Received a connection %s for shard_id=%i on host %s", id(conn), - conn.features.shard_id, - self.host - ) - old_conn = None - with self._lock: - is_shutdown = self.is_shutdown - if not is_shutdown: - if conn.features.shard_id in self._connections: - # Move the current connection to the trash and use the new one from now on - old_conn = self._connections[conn.features.shard_id] - log.debug( - "Replacing overloaded connection (%s) with (%s) for shard %i for host %s", - id(old_conn), - id(conn), - conn.features.shard_id, - self.host - ) - if self._keyspace: - conn.set_keyspace_blocking(self._keyspace) - self._connections[conn.features.shard_id] = conn + actual_shard_id if actual_shard_id is not None else -1, + self.host) - if is_shutdown: - conn.close() - return + if shard_aware_endpoint and shard_id != actual_shard_id: + # The connection did not land on the requested shard, which + # commonly indicates NAT in front of the shard-aware port. + self.disable_advanced_shard_aware(10 * 60) - if old_conn is not None: - remaining = old_conn.in_flight - len(old_conn.orphaned_request_ids) - if remaining == 0: + while True: + with self._lock: + if self.is_shutdown: + shutdown_owns_connection = not ( + self._remove_pending_connection_locked(conn)) + pending = False + close_connection = not shutdown_owns_connection + keyspace = None + keyspace_generation = None + else: + close_connection = False + keyspace = self._keyspace + keyspace_generation = self._keyspace_generation + + if close_connection: + conn.close() + return + if shutdown_owns_connection: + return + + if keyspace: + try: + _set_keyspace_blocking( + conn, + keyspace, + self._session.cluster.connect_timeout) + except _REQUEST_VALIDATION_EXCEPTIONS: + # An invalid/dropped Session keyspace is not a host + # failure. Keep ownership so a later USE can repair + # the Session, but quarantine this socket from + # ordinary requests until it selects a keyspace. + conn._pool_keyspace_mismatch = True + log.warning( + "Unable to set keyspace %s on connection to shard " + "%i of %s; quarantining it until a successful USE", + keyspace, + actual_shard_id, + self.host) + else: + conn._pool_keyspace_mismatch = False + if conn.is_closed: + raise _startup_close_error(conn, self.host.endpoint) + + old_connection = None + close_old_connection = False + old_remaining = 0 + excess_to_close = [] + close_connection = False + mapped_connection_adopted = False + with self._lock: + if self.is_shutdown: + shutdown_owns_connection = not ( + self._remove_pending_connection_locked(conn)) + pending = False + close_connection = not shutdown_owns_connection + elif keyspace_generation != self._keyspace_generation: + continue + elif not self._remove_pending_connection_locked(conn): + raise ConnectionException( + "Missing-shard connection ownership was lost", + self.host.endpoint) + else: + pending = False + old_connection = self._connections.get( + actual_shard_id) + if ( + old_connection is None or + self._connection_needs_replacement( + old_connection)): + self._connections[actual_shard_id] = conn + adopted = True + mapped_connection_adopted = True + if old_connection is not None: + ( + close_old_connection, + old_remaining + ) = self._retire_connection_locked( + old_connection) + + if self.num_missing_or_needing_replacement == 0: + excess_to_close = list( + self._excess_connections) + self._excess_connections.clear() + elif ( + len(self._connections) == + self.host.sharding_info.shards_count and + self.num_missing_or_needing_replacement == 0): + close_connection = True + else: + if ( + len(self._excess_connections) >= + self._excess_connection_limit): + excess_to_close = list( + self._excess_connections) + self._excess_connections.clear() + self._excess_connections.add(conn) + adopted = True + + if shutdown_owns_connection: + return + if close_connection: + conn.close() + if close_old_connection: log.debug( - "Immediately closing the old connection (%s) for shard %i on host %s", - id(old_conn), - old_conn.features.shard_id, - self.host - ) - old_conn.close() - else: + "Immediately closing retired connection (%s) for " + "shard %i on host %s", + id(old_connection), + actual_shard_id, + self.host) + old_connection.close() + elif old_connection is not None: log.debug( - "Moving the connection (%s) for shard %i to trash on host %s, %i requests remaining", - id(old_conn), - old_conn.features.shard_id, + "Moved connection (%s) for shard %i to trash on host " + "%s, %i requests remaining", + id(old_connection), + actual_shard_id, self.host, - remaining, - ) - with self._lock: - is_shutdown = self.is_shutdown - if not is_shutdown: - self._trash.add(old_conn) - if is_shutdown: - conn.close() - num_missing_or_needing_replacement = self.num_missing_or_needing_replacement - log.debug( - "Connected to %s/%i shards on host %s (%i missing or needs replacement)", - len(self._connections), - self.host.sharding_info.shards_count, - self.host, - num_missing_or_needing_replacement - ) - if num_missing_or_needing_replacement == 0: - log.debug( - "All shards of host %s have at least one connection, closing %i excess connections", - self.host, - len(self._excess_connections) - ) - self._close_excess_connections() - elif self.host.sharding_info.shards_count == len(self._connections) and self.num_missing_or_needing_replacement == 0: - log.debug( - "All shards are already covered, closing newly opened excess connection %s for host %s", - id(self), - self.host - ) - conn.close() - else: - if len(self._excess_connections) >= self._excess_connection_limit: - log.debug( - "After connection %s is created excess connection pool size limit (%i) reached for host %s, closing all %i of them", - id(conn), - self._excess_connection_limit, - self.host, - len(self._excess_connections) - ) - self._close_excess_connections() + old_remaining) + for excess_connection in excess_to_close: + excess_connection.close() - log.debug( - "Putting a connection %s to shard %i to the excess pool of host %s", - id(conn), - conn.features.shard_id, - self.host - ) - close_connection = False - with self._lock: - if self.is_shutdown: - close_connection = True - else: - self._excess_connections.add(conn) - if close_connection: - conn.close() - self._connecting.discard(shard_id) + if mapped_connection_adopted: + with self._stream_available_condition: + self._stream_available_condition.notify_all() + log.debug( + "Connected to %s/%i shards on host %s (%i missing or " + "needs replacement)", + len(self._connections), + self.host.sharding_info.shards_count, + self.host, + self.num_missing_or_needing_replacement) + return + except BaseException: + if conn is not None and not adopted: + with self._lock: + if pending: + if not self._remove_pending_connection_locked(conn): + shutdown_owns_connection = self.is_shutdown + pending = False + if not shutdown_owns_connection: + conn.close() + raise def _open_connections_for_all_shards(self, skip_shard_id=None): """ @@ -860,14 +2183,14 @@ def _open_connections_for_all_shards(self, skip_shard_id=None): with self._lock: if self.is_shutdown: return + shards_count = self.host.sharding_info.shards_count - for shard_id in range(self.host.sharding_info.shards_count): - if skip_shard_id is not None and skip_shard_id == shard_id: - continue - future = self._session.submit(self._open_connection_to_missing_shard, shard_id) - if isinstance(future, Future): - self._connecting.add(shard_id) - self._shard_connections_futures.append(future) + for shard_id in range(shards_count): + if skip_shard_id is not None and skip_shard_id == shard_id: + continue + self._schedule_connection_to_missing_shard( + shard_id, + track_future=True) trash_conns = None with self._lock: @@ -876,7 +2199,7 @@ def _open_connections_for_all_shards(self, skip_shard_id=None): self._trash = set() if trash_conns is not None: - for conn in self._trash: + for conn in trash_conns: conn.close() def _set_keyspace_for_all_conns(self, keyspace, callback): @@ -885,31 +2208,162 @@ def _set_keyspace_for_all_conns(self, keyspace, callback): connections have been set, `callback` will be called with two arguments: this pool, and a list of any errors that occurred. """ - remaining_callbacks = set(self._connections.values()) - remaining_callbacks_lock = Lock() - errors = [] + with self._lock: + self._keyspace = keyspace + self._keyspace_generation += 1 + self._keyspace_update_queue.append((keyspace, callback)) + if self._keyspace_update_in_progress: + return + self._keyspace_update_in_progress = True + + self._request_next_keyspace_update() + + def _request_next_keyspace_update(self): + with self._lock: + self._keyspace_update_run_requested = True + if self._keyspace_update_runner_active: + return + self._keyspace_update_runner_active = True - if not remaining_callbacks: - callback(self, errors) + while True: + with self._lock: + if not self._keyspace_update_run_requested: + self._keyspace_update_runner_active = False + return + self._keyspace_update_run_requested = False + try: + self._run_next_keyspace_update() + except BaseException: + with self._lock: + self._keyspace_update_runner_active = False + rerun = self._keyspace_update_run_requested + if rerun: + # A synchronous completion callback can enqueue/request the + # next generation and then raise a cancellation-style + # BaseException. Preserve propagation to its caller, but + # first transfer runner ownership so queued USE operations + # cannot remain permanently unresolved. + try: + self._request_next_keyspace_update() + except BaseException: + log.exception( + "Additional failure while resuming queued " + "keyspace updates for host %s", + self.host) + raise + + def _run_next_keyspace_update(self): + with self._lock: + if not self._keyspace_update_queue: + self._keyspace_update_in_progress = False + return + keyspace, callback = self._keyspace_update_queue.popleft() + connections = list(self._connections.values()) + shutdown_error = ConnectionException( + "Pool for %s is shutdown" % (self.host,), + self.host) if self.is_shutdown else None + + update = { + 'callback': callback, + 'errors': ( + [shutdown_error] + if shutdown_error is not None else []), + 'finished': False, + 'lock': Lock(), + 'remaining': set(connections), + } + self._keyspace_update_current = update + + if not update['remaining']: + self._finish_keyspace_update(update) return def connection_finished_setting_keyspace(conn, error): - self.return_connection(conn) - with remaining_callbacks_lock: - remaining_callbacks.remove(conn) - if error: - errors.append(error) + with update['lock']: + if conn not in update['remaining']: + return + update['remaining'].remove(conn) + conn._pool_keyspace_mismatch = error is not None + update_already_finished = update['finished'] + if error and not update_already_finished: + update['errors'].append(error) + claimed = bool( + not update_already_finished and + not update['remaining']) + if claimed: + # The last response owns completion before releasing this + # lock. In particular, shutdown cannot overtake a response + # here while return_connection balances in_flight. + update['finished'] = True - if not remaining_callbacks: - callback(self, errors) + try: + self.return_connection(conn) + except BaseException: + log.exception( + "Error returning connection after setting keyspace on " + "host %s", + self.host) + finally: + if claimed: + self._deliver_keyspace_update(update) + + for conn in connections: + with update['lock']: + if update['finished']: + break + try: + conn.set_keyspace_async( + keyspace, + connection_finished_setting_keyspace) + except BaseException as exc: + # Connection.set_keyspace_async maintains the in-flight + # invariant before issuing work, even for synchronous send + # failures, so finish it through the normal callback path. + connection_finished_setting_keyspace(conn, exc) + + def _finish_keyspace_update(self, update, error=None): + """ + Complete one serialized keyspace generation exactly once. - self._keyspace = keyspace - for conn in list(self._connections.values()): - conn.set_keyspace_async(keyspace, connection_finished_setting_keyspace) + ``shutdown`` uses this path to force the active generation to finish + before queued generations, even on reactors whose connection close + callbacks run later. + """ + if not self._claim_keyspace_update(update, error): + return False + self._deliver_keyspace_update(update) + return True + + @staticmethod + def _claim_keyspace_update(update, error=None): + with update['lock']: + if update['finished']: + return False + update['finished'] = True + if error is not None: + update['errors'].append(error) + return True + + def _deliver_keyspace_update(self, update): + try: + update['callback'](self, update['errors']) + except Exception: + log.exception( + "Error completing keyspace update on host %s", + self.host) + finally: + # Keep the completed generation published until its user callback + # returns. Otherwise a concurrent shutdown can advance the queue + # and invoke the next callback out of FIFO order. + with self._lock: + if self._keyspace_update_current is update: + self._keyspace_update_current = None + self._request_next_keyspace_update() def get_connections(self): - connections = self._connections - return list(connections.values()) if connections else [] + with self._lock: + connections = self._connections + return list(connections.values()) if connections else [] def get_state(self): in_flights = [c.in_flight for c in list(self._connections.values())] @@ -920,11 +2374,19 @@ def get_state(self): @property def num_missing_or_needing_replacement(self): return self.host.sharding_info.shards_count \ - - sum(1 for c in list(self._connections.values()) if not c.orphaned_threshold_reached) + - sum( + 1 for c in list(self._connections.values()) + if not self._connection_needs_replacement(c)) @property def open_count(self): - return sum([1 if c and not (c.is_closed or c.is_defunct) else 0 for c in list(self._connections.values())]) + return sum([ + 1 + if ( + c and + not (c.is_closed or c.is_defunct)) + else 0 + for c in list(self._connections.values())]) @property def _excess_connection_limit(self): diff --git a/tests/integration/standard/test_maintenance_mode_connection.py b/tests/integration/standard/test_maintenance_mode_connection.py new file mode 100644 index 0000000000..f08b989673 --- /dev/null +++ b/tests/integration/standard/test_maintenance_mode_connection.py @@ -0,0 +1,151 @@ +# Copyright DataStax, Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import socket +import unittest +from threading import Event, Thread + +import pytest + +from cassandra.cluster import NoHostAvailable +from cassandra.connection import ConnectionShutdown, DefaultEndPoint +from tests.integration import TestCluster + + +class MaintenanceModeCqlServer(object): + """ + Minimal CQL listener that accepts a startup attempt and then closes the + socket without replying, matching Scylla's maintenance-mode failure shape. + """ + + def __init__(self, max_connections=1): + self._closed = False + self._max_connections = max_connections + self._sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) + self._sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) + self._sock.bind(('127.0.0.1', 0)) + self._sock.listen(max_connections) + self._sock.settimeout(0.2) + self.port = self._sock.getsockname()[1] + self.frames = [] + self.ready = Event() + self.received_frame = Event() + self.error = None + self.thread = Thread(target=self._run) + self.thread.daemon = True + self.thread.start() + + def _run(self): + self.ready.set() + try: + while len(self.frames) < self._max_connections and not self._closed: + try: + client, _ = self._sock.accept() + except socket.timeout: + continue + + with client: + client.settimeout(2) + frame = b'' + while len(frame) < 9: + chunk = client.recv(9 - len(frame)) + if not chunk: + break + frame += chunk + self.frames.append(frame) + self.received_frame.set() + except Exception as exc: + if not self._closed: + self.error = exc + finally: + self.received_frame.set() + + def close(self): + self._closed = True + try: + self._sock.close() + except socket.error: + # Best-effort test cleanup; close can race with the listener thread. + pass + self.thread.join(2) + + +class MaintenanceModeConnectionTest(unittest.TestCase): + + def setUp(self): + self.cluster = TestCluster(contact_points=[], connect_timeout=2) + self.cluster.connection_class.initialize_reactor() + + def tearDown(self): + self.cluster.shutdown() + + def test_startup_close_is_observable_through_connection_factories(self): + """ + Exercise the real reactor/factory path against a socket that behaves + like a node rejecting regular CQL traffic while in maintenance mode. + """ + endpoint, server = self._new_endpoint_and_server() + try: + conn = self.cluster.connection_class.factory( + endpoint, + self.cluster.connect_timeout, + **self.cluster._make_connection_kwargs(endpoint, {})) + + assert conn.is_closed + self._assert_server_saw_options_frame(server) + finally: + server.close() + + endpoint, server = self._new_endpoint_and_server() + try: + conn = self.cluster.connection_factory(endpoint) + + assert conn.is_closed + self._assert_server_saw_options_frame(server) + finally: + server.close() + + def test_cluster_connect_reports_startup_close_as_unavailable_host(self): + endpoint, server = self._new_endpoint_and_server(max_connections=4) + cluster = TestCluster( + contact_points=[endpoint], + connect_timeout=2, + control_connection_timeout=2) + cluster.connection_class.initialize_reactor() + + try: + with pytest.raises(NoHostAvailable) as exc_info: + cluster.connect() + + errors = exc_info.value.errors + assert errors + assert any(isinstance(exc, ConnectionShutdown) for exc in errors.values()) + assert any("closed during the startup handshake" in str(exc) + for exc in errors.values()) + self._assert_server_saw_options_frame(server) + finally: + cluster.shutdown() + server.close() + + def _new_endpoint_and_server(self, max_connections=1): + server = MaintenanceModeCqlServer(max_connections=max_connections) + assert server.ready.wait(2) + return DefaultEndPoint('127.0.0.1', server.port), server + + def _assert_server_saw_options_frame(self, server): + assert server.received_frame.wait(2) + assert server.error is None + assert server.frames + assert len(server.frames[0]) >= 5 + assert server.frames[0][4] == 0x05 # OPTIONS diff --git a/tests/unit/io/test_twistedreactor.py b/tests/unit/io/test_twistedreactor.py index 02bac10d8e..f02d6cba61 100644 --- a/tests/unit/io/test_twistedreactor.py +++ b/tests/unit/io/test_twistedreactor.py @@ -19,13 +19,15 @@ try: from twisted.test import proto_helpers + from twisted.internet.error import ConnectionDone, ConnectionLost from cassandra.io import twistedreactor from cassandra.io.twistedreactor import TwistedConnection except ImportError: - twistedreactor = TwistedConnection = None # NOQA + twistedreactor = TwistedConnection = ConnectionDone = ConnectionLost = None # NOQA -from cassandra.connection import _Frame +from cassandra.connection import ConnectionShutdown, _Frame +from cassandra.protocol import ReadyMessage from tests.unit.io.utils import TimerTestMixin @@ -132,6 +134,95 @@ def test_client_connection_made(self): self.obj_ut.client_connection_made(Mock()) self.obj_ut._send_options_message.assert_called_with() + def test_clean_close_error(self): + assert self.obj_ut._is_clean_close_error(ConnectionDone()) + assert not self.obj_ut._is_clean_close_error(ConnectionLost()) + + self.obj_ut.is_defunct = False + assert self.obj_ut._is_clean_close_error(ConnectionShutdown("closed")) + + self.obj_ut.is_defunct = True + assert not self.obj_ut._is_clean_close_error(ConnectionShutdown("defunct")) + + def test_factory_returns_connection_done_before_startup(self): + class StartupConnectionDone(TwistedConnection): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.defunct(ConnectionDone()) + + conn = StartupConnectionDone.factory( + DefaultEndPoint('1.2.3.4'), timeout=1) + + assert conn.is_closed + assert not conn._startup_completed + assert isinstance(conn.last_error, ConnectionDone) + + def test_factory_raises_connection_done_after_ready(self): + class ReadyThenConnectionDone(TwistedConnection): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self._compressor = None + self._handle_startup_response(ReadyMessage()) + self.defunct(ConnectionDone()) + + with self.assertRaises(ConnectionDone): + ReadyThenConnectionDone.factory( + DefaultEndPoint('1.2.3.4'), timeout=1) + + @patch('cassandra.io.twistedreactor.connectProtocol') + def test_close_cancels_pending_connect(self, mock_connect_protocol): + connector = Mock() + mock_connect_protocol.return_value = connector + self.obj_ut.error_all_requests = Mock() + self.obj_ut.add_connection() + + self.obj_ut.close() + + assert self.obj_ut.is_closed + assert self.obj_ut.connected_event.is_set() + assert isinstance(self.obj_ut.last_error, ConnectionShutdown) + connector.addErrback.assert_called_once_with( + self.obj_ut._handle_connect_failure) + self.mock_reactor_cft.assert_called_with(connector.cancel) + self.obj_ut.error_all_requests.assert_called_once_with( + self.obj_ut.last_error) + + def test_connect_failure_defuncts_connection_and_wakes_startup(self): + connect_error = RuntimeError("endpoint connect failed") + failure = Mock(value=connect_error) + + result = self.obj_ut._handle_connect_failure(failure) + + assert result is None + assert self.obj_ut.is_defunct + assert self.obj_ut.is_closed + assert self.obj_ut.last_error is connect_error + assert self.obj_ut.connected_event.is_set() + + @patch('cassandra.io.twistedreactor.connectProtocol') + def test_close_before_scheduled_connect(self, mock_connect_protocol): + self.obj_ut.error_all_requests = Mock() + + self.obj_ut.close() + self.obj_ut.add_connection() + + assert self.obj_ut.is_closed + assert self.obj_ut.connected_event.is_set() + mock_connect_protocol.assert_not_called() + + def test_connection_made_after_close_does_not_reopen(self): + transport = Mock() + self.obj_ut._send_options_message = Mock() + self.obj_ut.close() + + self.obj_ut.client_connection_made(transport) + + assert self.obj_ut.is_closed + assert self.obj_ut.transport is None + self.obj_ut._send_options_message.assert_not_called() + self.mock_reactor_cft.assert_called_with( + transport.connector.disconnect) + @patch('twisted.internet.reactor.connectTCP') def test_close(self, mock_connectTCP): """ diff --git a/tests/unit/test_cluster.py b/tests/unit/test_cluster.py index 3d55bc1860..fc5ac4ee7c 100644 --- a/tests/unit/test_cluster.py +++ b/tests/unit/test_cluster.py @@ -16,6 +16,7 @@ from concurrent.futures import Future import logging import socket +from threading import Event, Lock, RLock, Thread from types import SimpleNamespace from unittest.mock import patch, Mock @@ -24,9 +25,14 @@ 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.pool import Host + ExecutionProfile, _ConfigMode, EXEC_PROFILE_DEFAULT, \ + _HostTransitionResult, _SESSION_LOCAL_POOL_FAILURE, \ + _STALE_POOL_ATTEMPT +from cassandra.connection import (ConnectionBusy, ConnectionException, + DefaultEndPoint, + _ConnectionClosedDuringStartup) +from cassandra.metadata import Metadata +from cassandra.pool import Host, _HostReconnectionHandler from cassandra.policies import HostDistance, RetryPolicy, RoundRobinPolicy, DowngradingConsistencyRetryPolicy, SimpleConvictionPolicy from cassandra.query import SimpleStatement, named_tuple_factory, tuple_factory from tests.unit.utils import mock_session_pools @@ -150,6 +156,31 @@ def test_backward_compat_positional(self): class ClusterTest(unittest.TestCase): + @staticmethod + def _new_transition_cluster(): + cluster = object.__new__(Cluster) + cluster.is_shutdown = False + cluster.allow_control_connection_query_fallback = ( + ControlConnectionQueryFallback.Disabled) + cluster._lock = RLock() + cluster._discount_down_events = False + cluster.sessions = [] + cluster._listener_lock = RLock() + cluster._listeners = set() + cluster.profile_manager = Mock() + cluster.profile_manager.distance.return_value = HostDistance.LOCAL + cluster.control_connection = Mock() + cluster._prepare_all_queries = Mock() + cluster._start_reconnector = Mock() + cluster.on_down_potentially_blocking = ( + lambda host, is_host_addition, recovery_epoch: + cluster._run_host_transition( + host, + lambda: cluster._on_down_potentially_blocking_serialized( + host, is_host_addition, recovery_epoch))) + return cluster + + def test_tuple_for_contact_points(self): cluster = Cluster(contact_points=[('localhost', 9045), ('127.0.0.2', 9046), '127.0.0.3'], port=9999) # Refactored for clarity @@ -233,6 +264,141 @@ def test_control_connection_query_fallback_fallback_tolerates_empty_initial_pool assert session._initial_connect_futures == {future} assert session._pools == {} + def test_session_returns_after_first_pool_without_waiting_for_pending_pool( + self): + cluster = Cluster(monitor_reporting_enabled=False) + hosts = [ + Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()), + Host( + "127.0.0.2", + SimpleConvictionPolicy, + host_id=uuid.uuid4()), + ] + created = Future() + created.set_result(True) + pending = Future() + futures = { + id(hosts[0]): created, + id(hosts[1]): pending, + } + + def add_pool(session, host, is_host_addition): + return futures[id(host)] + + with patch.object( + Session, + "add_or_renew_pool", + new=add_pool): + session = Session(cluster, hosts) + + assert session._initial_connect_futures == {created, pending} + pending.set_result(False) + session.shutdown() + + def test_session_constructor_failure_closes_published_pool(self): + cluster = Cluster( + column_encryption_policy=object(), + monitor_reporting_enabled=False) + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + pool = Mock() + + def publish_pool(session, pool_host, is_host_addition): + session._pools[pool_host] = pool + future = Future() + future.set_result(True) + return future + + with patch.object( + Session, + "add_or_renew_pool", + new=publish_pool): + with pytest.raises( + Exception, + match="column_encryption_policy is temporary disabled"): + Session(cluster, [host]) + + pool.shutdown.assert_called_once_with() + + def test_new_session_reconciles_down_host_before_registration(self): + cluster = object.__new__(Cluster) + cluster._lock = RLock() + cluster.is_shutdown = False + cluster.sessions = set() + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.set_up() + cluster.metadata = Mock() + cluster.metadata.all_hosts.return_value = [host] + + pool = Mock() + session = Mock() + session.hosts = [host] + session._lock = RLock() + session._initial_connect_host_futures = [] + session._pop_pools_locked.return_value = [pool] + cluster._session_register_user_types = Mock( + side_effect=lambda _: host.set_down()) + + with patch('cassandra.cluster.Session', return_value=session): + created = cluster._new_session(None) + + assert created is session + assert session in cluster.sessions + session._advance_pool_generation_locked.assert_called_once_with( + host) + pool.shutdown.assert_called_once_with() + session.update_created_pools.assert_called_once_with( + skip_host_ids=set()) + + def test_new_session_reconciles_each_initial_host_once(self): + cluster = object.__new__(Cluster) + cluster._lock = RLock() + cluster.is_shutdown = False + cluster.sessions = set() + hosts = [ + Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()), + Host( + "127.0.0.2", + SimpleConvictionPolicy, + host_id=uuid.uuid4()), + ] + cluster.metadata = Mock() + cluster.metadata.all_hosts.return_value = hosts + cluster._session_register_user_types = Mock() + + initial_futures = [Future(), Future()] + session = Mock() + session.hosts = hosts + session._lock = RLock() + session._initial_connect_host_futures = list(zip( + hosts, initial_futures)) + + with patch('cassandra.cluster.Session', return_value=session): + assert cluster._new_session(None) is session + + session.update_created_pools.assert_called_once_with( + skip_host_ids={id(host) for host in hosts}) + + initial_futures[0].set_result(False) + assert session.update_created_pools.call_args_list[-1].kwargs == { + 'only_host_ids': (id(hosts[0]),)} + + initial_futures[1].set_result(False) + assert session.update_created_pools.call_args_list[-1].kwargs == { + 'only_host_ids': (id(hosts[1]),)} + assert session.update_created_pools.call_count == 3 + 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: @@ -268,16 +434,2526 @@ def test_connection_factory_passes_compression_kwarg(self): ] for supported, configured, expected in scenarios: + connection = Mock(is_closed=False) with patch.dict('cassandra.cluster.locally_supported_compressions', supported, clear=True): - with patch.object(Cluster.connection_class, 'factory', autospec=True, return_value='connection') as factory: + with patch.object(Cluster.connection_class, 'factory', autospec=True, return_value=connection) as factory: cluster = Cluster(compression=configured) conn = cluster.connection_factory(endpoint) - assert conn == 'connection' + assert conn is connection assert factory.call_count == 1 assert factory.call_args.kwargs['compression'] == expected assert cluster.compression == expected + def test_reconnection_factory_returns_open_connection(self): + endpoint = Mock(address='127.0.0.1') + host = Mock(endpoint=endpoint) + connection = Mock(is_closed=False) + + with patch.object(Cluster.connection_class, 'factory', autospec=True, return_value=connection) as factory: + cluster = Cluster() + conn_factory = cluster._make_connection_factory(host) + conn = conn_factory() + + assert conn is connection + assert factory.call_count == 1 + + def test_connection_factory_preserves_positional_arguments(self): + endpoint = Mock(address='127.0.0.1') + positional_option = object() + connection = Mock(is_closed=False) + + with patch.object( + Cluster.connection_class, + 'factory', + autospec=True, + return_value=connection) as factory: + cluster = Cluster() + conn = cluster.connection_factory(endpoint, positional_option) + + assert conn is connection + assert factory.call_args.args[2] is positional_option + + def test_reconnection_factory_preserves_positional_arguments(self): + endpoint = Mock(address='127.0.0.1') + host = Mock(endpoint=endpoint) + positional_option = object() + connection = Mock(is_closed=False) + + with patch.object( + Cluster.connection_class, + 'factory', + autospec=True, + return_value=connection) as factory: + cluster = Cluster() + conn_factory = cluster._make_connection_factory( + host, positional_option) + conn = conn_factory() + + assert conn is connection + assert factory.call_args.args[2] is positional_option + + def test_connection_factory_returns_startup_close(self): + endpoint = Mock(address='127.0.0.1') + connection = Mock(is_closed=True) + connection.endpoint = endpoint + + with patch.object(Cluster.connection_class, 'factory', autospec=True, return_value=connection) as factory: + cluster = Cluster() + conn = cluster.connection_factory(endpoint) + + assert conn is connection + factory.assert_called_once() + + def test_reconnection_factory_returns_startup_close_result(self): + endpoint = Mock(address='127.0.0.1') + host = Mock(endpoint=endpoint) + connection = Mock(is_closed=True) + connection.endpoint = endpoint + + with patch.object(Cluster.connection_class, 'factory', autospec=True, return_value=connection): + cluster = Cluster() + conn_factory = cluster._make_connection_factory(host) + conn = conn_factory() + + assert conn is connection + + def test_host_reconnection_handler_parks_startup_close(self): + endpoint = Mock(address='127.0.0.1') + host = Mock(endpoint=endpoint) + connection = Mock(is_closed=True, endpoint=endpoint) + connection_factory = Mock(return_value=connection) + scheduler = Mock() + on_add = Mock() + on_up = Mock() + callback = Mock() + handler = _HostReconnectionHandler( + host, connection_factory, False, on_add, on_up, + scheduler, iter([1]), callback) + + handler.run() + + connection_factory.assert_called_once_with() + connection.close.assert_called_once_with() + on_add.assert_not_called() + on_up.assert_not_called() + callback.assert_not_called() + scheduler.schedule.assert_not_called() + + def test_prepare_all_queries_rejects_startup_close(self): + endpoint = Mock(address='127.0.0.1') + host = Mock(endpoint=endpoint) + connection = Mock(is_closed=True, endpoint=endpoint) + cluster = Cluster() + cluster._prepared_statements = {b'query-id': Mock()} + + with patch.object(cluster, 'connection_factory', return_value=connection), \ + patch.object(cluster, '_send_chunks') as send_chunks: + cluster._prepare_all_queries(host) + + connection.set_keyspace_blocking.assert_not_called() + send_chunks.assert_not_called() + connection.close.assert_called_once_with() + + def test_force_on_down_bypasses_open_pool_discount(self): + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.set_up() + session = Mock() + session.get_pool_state.return_value = { + host: {'open_count': 1}, + } + cluster = object.__new__(Cluster) + cluster.is_shutdown = False + cluster.allow_control_connection_query_fallback = ( + ControlConnectionQueryFallback.Disabled) + cluster._lock = RLock() + cluster._discount_down_events = True + cluster.profile_manager = Mock() + cluster.profile_manager.distance.return_value = HostDistance.LOCAL + cluster.sessions = [session] + cluster.on_down_potentially_blocking = Mock() + + cluster.on_down(host, is_host_addition=False) + + assert host.is_up + cluster.on_down_potentially_blocking.assert_not_called() + + cluster.on_down( + host, + is_host_addition=False, + expect_host_to_be_down=True, + force=True) + + assert not host.is_up + cluster.on_down_potentially_blocking.assert_called_once_with( + host, False, host._recovery_epoch) + + def test_endpoint_relocation_cleans_old_hash_before_readding_host(self): + cluster = self._new_transition_cluster() + host = Host( + DefaultEndPoint("127.0.0.1"), + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.set_up() + policy = RoundRobinPolicy() + policy.populate(cluster, [host]) + profile_manager = Mock() + profile_manager.distance.return_value = HostDistance.LOCAL + profile_manager.on_down.side_effect = policy.on_down + profile_manager.on_up.side_effect = policy.on_up + cluster.profile_manager = profile_manager + cluster.metadata = Mock() + new_endpoint = DefaultEndPoint("127.0.0.2") + + assert cluster._force_down_for_endpoint_change( + host, + new_endpoint) + + assert host.endpoint == new_endpoint + assert host.is_up is True + assert tuple(policy.make_query_plan()) == (host,) + assert len(policy._live_hosts) == 1 + cluster.metadata.update_host.assert_called_once_with( + host, + DefaultEndPoint("127.0.0.1")) + cluster._start_reconnector.assert_not_called() + + def test_endpoint_relocation_mutates_before_reentrant_up(self): + cluster = self._new_transition_cluster() + host = Host( + DefaultEndPoint("127.0.0.1"), + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.set_up() + policy = RoundRobinPolicy() + policy.populate(cluster, [host]) + reentered = [] + + def on_down(down_host): + policy.on_down(down_host) + if not reentered: + reentered.append(True) + cluster.on_up(down_host) + + cluster.profile_manager.on_down.side_effect = on_down + cluster.profile_manager.on_up.side_effect = policy.on_up + cluster.metadata = Mock() + new_endpoint = DefaultEndPoint("127.0.0.2") + + assert cluster._force_down_for_endpoint_change( + host, + new_endpoint) + + assert host.endpoint == new_endpoint + assert host.is_up is True + assert policy._live_hosts == frozenset((host,)) + + policy.on_down(host) + assert not policy._live_hosts + + def test_endpoint_relocation_updates_in_skip_pool_creation_mode(self): + cluster = self._new_transition_cluster() + cluster.allow_control_connection_query_fallback = ( + ControlConnectionQueryFallback.SkipPoolCreation) + cluster.metadata = Mock() + old_endpoint = DefaultEndPoint("127.0.0.1") + new_endpoint = DefaultEndPoint("127.0.0.2") + host = Host( + old_endpoint, + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.set_up() + policy = RoundRobinPolicy() + policy.populate(cluster, [host]) + cluster.profile_manager.on_down.side_effect = policy.on_down + cluster.profile_manager.on_up.side_effect = policy.on_up + + assert cluster._force_down_for_endpoint_change( + host, + new_endpoint) + + assert host.endpoint == new_endpoint + assert host.is_up is True + assert policy._live_hosts == frozenset((host,)) + cluster.profile_manager.on_down.assert_called_once_with(host) + cluster.profile_manager.on_up.assert_called_once_with(host) + cluster.metadata.update_host.assert_called_once_with( + host, + old_endpoint) + + def test_endpoint_relocation_racing_up_keeps_one_policy_entry(self): + cluster = self._new_transition_cluster() + old_endpoint = DefaultEndPoint("127.0.0.1") + new_endpoint = DefaultEndPoint("127.0.0.2") + host = Host( + old_endpoint, + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.set_up() + policy = RoundRobinPolicy() + policy.populate(cluster, [host]) + cluster.profile_manager.on_down.side_effect = policy.on_down + cluster.profile_manager.on_up.side_effect = policy.on_up + cluster.metadata = Mock() + lane_entered = Event() + release_lane = Event() + relocation_result = [] + + def hold_lane(): + lane_entered.set() + assert release_lane.wait(2) + + lane_thread = Thread( + target=lambda: cluster._run_host_transition( + host, hold_lane)) + lane_thread.start() + assert lane_entered.wait(2) + + relocation_thread = Thread( + target=lambda: relocation_result.append( + cluster._force_down_for_endpoint_change( + host, new_endpoint))) + relocation_thread.start() + cluster.on_up(host) + release_lane.set() + lane_thread.join(2) + relocation_thread.join(2) + + assert not lane_thread.is_alive() + assert not relocation_thread.is_alive() + assert relocation_result == [True] + assert host.endpoint == new_endpoint + assert host.is_up is True + assert policy._live_hosts == frozenset((host,)) + assert len(policy._live_hosts) == 1 + + def test_endpoint_relocation_does_not_resurrect_removed_host(self): + cluster = self._new_transition_cluster() + old_endpoint = DefaultEndPoint("127.0.0.1") + new_endpoint = DefaultEndPoint("127.0.0.2") + host = Host( + old_endpoint, + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.set_up() + cluster.metadata = Mock() + lane_entered = Event() + release_lane = Event() + relocation_result = [] + + def hold_lane(): + lane_entered.set() + assert release_lane.wait(2) + + lane_thread = Thread( + target=lambda: cluster._run_host_transition( + host, hold_lane)) + lane_thread.start() + assert lane_entered.wait(2) + + relocation_thread = Thread( + target=lambda: relocation_result.append( + cluster._force_down_for_endpoint_change( + host, new_endpoint))) + relocation_thread.start() + cluster.on_remove(host) + release_lane.set() + lane_thread.join(2) + relocation_thread.join(2) + + assert not lane_thread.is_alive() + assert not relocation_thread.is_alive() + assert relocation_result == [False] + assert host.endpoint == old_endpoint + assert host._is_removed + cluster.metadata.update_host.assert_not_called() + + def test_endpoint_relocation_honors_metadata_removal_gap(self): + cluster = self._new_transition_cluster() + old_endpoint = DefaultEndPoint("127.0.0.1") + new_endpoint = DefaultEndPoint("127.0.0.2") + host = Host( + old_endpoint, + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.set_up() + cluster.metadata = Metadata() + cluster.metadata.add_or_return_host(host) + lane_entered = Event() + release_lane = Event() + relocation_result = [] + + def hold_lane(): + lane_entered.set() + assert release_lane.wait(2) + + lane_thread = Thread( + target=lambda: cluster._run_host_transition( + host, hold_lane)) + lane_thread.start() + assert lane_entered.wait(2) + relocation_thread = Thread( + target=lambda: relocation_result.append( + cluster._force_down_for_endpoint_change( + host, new_endpoint))) + relocation_thread.start() + + # Model Metadata.remove_host() winning just before Cluster.on_remove() + # can mark the Host object itself. + assert cluster.metadata.remove_host(host) + assert not host._is_removed + release_lane.set() + lane_thread.join(2) + relocation_thread.join(2) + + assert not lane_thread.is_alive() + assert not relocation_thread.is_alive() + assert relocation_result == [False] + assert host.endpoint == old_endpoint + assert cluster.metadata.get_host_by_host_id(host.host_id) is None + + def test_reentrant_endpoint_relocation_keeps_newest_endpoint(self): + cluster = self._new_transition_cluster() + first_endpoint = DefaultEndPoint("127.0.0.1") + second_endpoint = DefaultEndPoint("127.0.0.2") + newest_endpoint = DefaultEndPoint("127.0.0.3") + host = Host( + first_endpoint, + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.set_up() + cluster.metadata = Mock() + reentered = [] + + def relocate_again(_): + if not reentered: + reentered.append(True) + assert cluster._force_down_for_endpoint_change( + host, newest_endpoint) + + cluster.profile_manager.on_down.side_effect = relocate_again + + assert not cluster._force_down_for_endpoint_change( + host, second_endpoint) + + assert host.endpoint == newest_endpoint + assert host.is_up is True + cluster.metadata.update_host.assert_called_once_with( + host, first_endpoint) + + def test_endpoint_relocation_publishes_down_before_up(self): + cluster = self._new_transition_cluster() + host = Host( + DefaultEndPoint("127.0.0.1"), + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.set_up() + cluster.metadata = Mock() + events = [] + listener = Mock() + listener.on_down.side_effect = lambda _: events.append("down") + listener.on_up.side_effect = lambda _: events.append("up") + cluster._listeners = {listener} + pool_created = Future() + pool_created.set_result(True) + session = Mock() + session.add_or_renew_pool.return_value = pool_created + cluster.sessions = [session] + + assert cluster._force_down_for_endpoint_change( + host, + DefaultEndPoint("127.0.0.2")) + + assert events == ["down", "up"] + + def test_later_down_wins_while_endpoint_relocation_cleans_up(self): + cluster = self._new_transition_cluster() + host = Host( + DefaultEndPoint("127.0.0.1"), + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.set_up() + cluster.metadata = Mock() + cleanup_entered = Event() + release_cleanup = Event() + relocation_result = [] + + def block_relocation_cleanup(_): + cleanup_entered.set() + assert release_cleanup.wait(2) + + cluster.profile_manager.on_down.side_effect = \ + block_relocation_cleanup + relocation_thread = Thread( + target=lambda: relocation_result.append( + cluster._force_down_for_endpoint_change( + host, + DefaultEndPoint("127.0.0.2")))) + relocation_thread.start() + assert cleanup_entered.wait(2) + + cluster.on_down( + host, + is_host_addition=False, + expect_host_to_be_down=True, + force=True) + release_cleanup.set() + relocation_thread.join(2) + + assert not relocation_thread.is_alive() + assert relocation_result == [True] + assert host.is_up is False + + def test_private_duplicate_down_preempts_relocation_recovery(self): + cluster = self._new_transition_cluster() + host = Host( + DefaultEndPoint("127.0.0.1"), + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.set_up() + cluster.metadata = Mock() + cleanup_entered = Event() + release_cleanup = Event() + + def block_relocation_cleanup(_): + cleanup_entered.set() + assert release_cleanup.wait(2) + + cluster.profile_manager.on_down.side_effect = \ + block_relocation_cleanup + relocation_thread = Thread( + target=lambda: cluster._force_down_for_endpoint_change( + host, + DefaultEndPoint("127.0.0.2"))) + relocation_thread.start() + assert cleanup_entered.wait(2) + + cluster._on_down_locked( + host, + is_host_addition=False) + release_cleanup.set() + relocation_thread.join(2) + + assert not relocation_thread.is_alive() + assert host.is_up is False + + def test_endpoint_relocation_preserves_legacy_on_down_override(self): + class LegacyCluster(Cluster): + + def on_down( + self, host, is_host_addition, + expect_host_to_be_down=False): + self.endpoint_events.append(( + "down", + host.endpoint, + is_host_addition, + expect_host_to_be_down)) + host.get_and_set_reconnection_handler( + self.override_reconnector) + + def on_up(self, host): + self.endpoint_events.append(("up", host.endpoint)) + return Cluster.on_up(self, host) + + cluster = object.__new__(LegacyCluster) + cluster.__dict__.update( + self._new_transition_cluster().__dict__) + cluster.endpoint_events = [] + cluster.override_reconnector = Mock() + cluster.metadata = Mock() + old_endpoint = DefaultEndPoint("127.0.0.1") + new_endpoint = DefaultEndPoint("127.0.0.2") + host = Host( + old_endpoint, + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.set_up() + + assert cluster._force_down_for_endpoint_change( + host, + new_endpoint) + + assert cluster.endpoint_events == [ + ("down", old_endpoint, False, True), + ("up", new_endpoint), + ] + assert host.is_up is True + cluster.override_reconnector.cancel.assert_called_once_with() + cluster.profile_manager.on_down.assert_called_once_with(host) + + def test_endpoint_relocation_with_delegating_down_override_recovers_up( + self): + class DelegatingCluster(Cluster): + + def on_down( + self, host, is_host_addition, + expect_host_to_be_down=False): + return super().on_down( + host, + is_host_addition, + expect_host_to_be_down) + + cluster = object.__new__(DelegatingCluster) + cluster.__dict__.update( + self._new_transition_cluster().__dict__) + cluster.metadata = Mock() + host = Host( + DefaultEndPoint("127.0.0.1"), + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.set_up() + + assert cluster._force_down_for_endpoint_change( + host, + DefaultEndPoint("127.0.0.2")) + + assert host.is_up is True + assert host.endpoint == DefaultEndPoint("127.0.0.2") + + def test_endpoint_relocation_aborts_when_down_cleanup_fails(self): + cluster = self._new_transition_cluster() + old_endpoint = DefaultEndPoint("127.0.0.1") + new_endpoint = DefaultEndPoint("127.0.0.2") + host = Host( + old_endpoint, + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.set_up() + old_hash = hash(host) + cluster.metadata = Mock() + cluster.profile_manager.on_down.side_effect = RuntimeError( + "policy cleanup failed") + + assert not cluster._force_down_for_endpoint_change( + host, + new_endpoint) + + assert host.endpoint == old_endpoint + assert hash(host) == old_hash + assert host.is_up is False + cluster.metadata.update_host.assert_not_called() + cluster.profile_manager.on_up.assert_not_called() + cluster._start_reconnector.assert_called_once_with( + host, + is_host_addition=False, + recovery_epoch=host._recovery_epoch) + + def test_duplicate_deferred_relocation_runs_one_down_up_cycle(self): + cluster = self._new_transition_cluster() + old_endpoint = DefaultEndPoint("127.0.0.1") + new_endpoint = DefaultEndPoint("127.0.0.2") + host = Host( + old_endpoint, + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.set_up() + cluster.metadata = Mock() + deferred_results = [] + + def enqueue_duplicate_relocations(): + deferred_results.append( + cluster._force_down_for_endpoint_change( + host, + new_endpoint)) + deferred_results.append( + cluster._force_down_for_endpoint_change( + host, + new_endpoint)) + + cluster._run_host_transition( + host, + enqueue_duplicate_relocations) + + assert len(deferred_results) == 2 + assert host.endpoint == new_endpoint + assert host.is_up is True + cluster.metadata.update_host.assert_called_once_with( + host, + old_endpoint) + cluster.profile_manager.on_down.assert_called_once_with(host) + cluster.profile_manager.on_up.assert_called_once_with(host) + cluster.control_connection.on_down.assert_called_once_with(host) + cluster.control_connection.on_up.assert_called_once_with(host) + + def test_older_blocked_delegating_relocation_cannot_overwrite_newer( + self): + class DelegatingCluster(Cluster): + + def on_down( + self, host, is_host_addition, + expect_host_to_be_down=False): + self.delegating_down_calls += 1 + if self.delegating_down_calls == 1: + self.first_down_entered.set() + assert self.release_first_down.wait(2) + return super().on_down( + host, + is_host_addition, + expect_host_to_be_down) + + cluster = object.__new__(DelegatingCluster) + cluster.__dict__.update( + self._new_transition_cluster().__dict__) + cluster.metadata = Mock() + cluster.delegating_down_calls = 0 + cluster.first_down_entered = Event() + cluster.release_first_down = Event() + second_reservation = Event() + reservation_count = [] + + def reserve_endpoint_change(host): + sequence = Cluster._reserve_endpoint_change_event(host) + reservation_count.append(sequence) + if len(reservation_count) == 2: + second_reservation.set() + return sequence + + cluster._reserve_endpoint_change_event = reserve_endpoint_change + old_endpoint = DefaultEndPoint("127.0.0.1") + intermediate_endpoint = DefaultEndPoint("127.0.0.2") + newest_endpoint = DefaultEndPoint("127.0.0.3") + host = Host( + old_endpoint, + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.set_up() + older_result = [] + newer_result = [] + + older_thread = Thread( + target=lambda: older_result.append( + cluster._force_down_for_endpoint_change( + host, + intermediate_endpoint))) + newer_thread = Thread( + target=lambda: newer_result.append( + cluster._force_down_for_endpoint_change( + host, + newest_endpoint))) + older_thread.start() + assert cluster.first_down_entered.wait(2) + newer_thread.start() + assert second_reservation.wait(2) + cluster.release_first_down.set() + older_thread.join(2) + newer_thread.join(2) + + assert not older_thread.is_alive() + assert not newer_thread.is_alive() + assert older_result == [False] + assert newer_result == [True] + assert host.endpoint == newest_endpoint + assert host.is_up is True + cluster.metadata.update_host.assert_called_once_with( + host, + old_endpoint) + cluster.profile_manager.on_down.assert_called_once_with(host) + cluster.profile_manager.on_up.assert_called_once_with(host) + + def test_force_connection_failure_bypasses_conviction_policy(self): + host = Mock() + host.signal_connection_failure.return_value = False + error = ConnectionException("closed during startup") + cluster = object.__new__(Cluster) + cluster.on_down = Mock() + cluster._on_down = Mock() + + is_down = cluster.signal_connection_failure( + host, + error, + is_host_addition=False, + expect_host_to_be_down=True, + force=True) + + assert not is_down + cluster._on_down.assert_called_once_with( + host, + False, + True, + force=True) + cluster.on_down.assert_not_called() + + def test_forced_connection_failure_survives_conviction_policy_error(self): + host = Mock() + host.signal_connection_failure.side_effect = RuntimeError( + "broken policy") + cluster = object.__new__(Cluster) + cluster._on_down = Mock() + + is_down = cluster.signal_connection_failure( + host, + ConnectionException("closed during startup"), + is_host_addition=False, + expect_host_to_be_down=True, + force=True) + + assert not is_down + cluster._on_down.assert_called_once_with( + host, + False, + True, + force=True) + + def test_forced_connection_failure_survives_conviction_base_exception( + self): + class PolicyCancelled(BaseException): + pass + + host = Mock() + cancellation = PolicyCancelled("cancelled") + host.signal_connection_failure.side_effect = cancellation + cluster = object.__new__(Cluster) + cluster._on_down = Mock() + + is_down = cluster.signal_connection_failure( + host, + ConnectionException("closed during startup"), + is_host_addition=False, + expect_host_to_be_down=True, + force=True) + + assert not is_down + cluster._on_down.assert_called_once_with( + host, + False, + True, + force=True) + + def test_down_queued_during_on_up_leaves_host_and_policies_down(self): + cluster = self._new_transition_cluster() + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.set_down() + entered_on_up = Event() + release_on_up = Event() + policy_events = [] + + def block_on_up(_): + policy_events.append("up") + entered_on_up.set() + assert release_on_up.wait(2) + + cluster.profile_manager.on_up.side_effect = block_on_up + cluster.profile_manager.on_down.side_effect = ( + lambda _: policy_events.append("down")) + + up_thread = Thread(target=cluster.on_up, args=(host,)) + up_thread.start() + assert entered_on_up.wait(2) + + cluster._on_down_locked( + host, + is_host_addition=False, + expect_host_to_be_down=True, + force=True) + release_on_up.set() + up_thread.join(2) + + assert not up_thread.is_alive() + assert host.is_up is False + assert policy_events == ["up", "down"] + cluster.control_connection.on_up.assert_not_called() + cluster.control_connection.on_down.assert_called_once_with(host) + + def test_unknown_host_failed_up_enters_down_recovery(self): + cluster = self._new_transition_cluster() + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + cluster._prepare_all_queries.side_effect = RuntimeError( + "query preparation failed") + + assert host.is_up is None + with pytest.raises(RuntimeError, match="query preparation failed"): + cluster.on_up(host) + + assert host.is_up is False + assert not host._currently_handling_node_up + cluster.profile_manager.on_down.assert_called_once_with(host) + cluster.control_connection.on_down.assert_called_once_with(host) + cluster._start_reconnector.assert_called_once_with( + host, + is_host_addition=False, + recovery_epoch=host._recovery_epoch) + + def test_up_during_add_preserves_add_callbacks(self): + cluster = self._new_transition_cluster() + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + listener = Mock() + cluster._listeners = {listener} + entered_add = Event() + release_add = Event() + errors = [] + + def block_add(_): + entered_add.set() + assert release_add.wait(2) + + cluster.profile_manager.on_add.side_effect = block_add + + def run_add(): + try: + cluster.on_add(host) + except BaseException as exc: + errors.append(exc) + + add_thread = Thread(target=run_add) + add_thread.start() + assert entered_add.wait(2) + + cluster.on_up(host) + release_add.set() + add_thread.join(2) + + assert not add_thread.is_alive() + assert not errors + assert host.is_up is True + cluster.control_connection.on_add.assert_called_once_with(host, True) + listener.on_add.assert_called_once_with(host) + cluster.control_connection.on_up.assert_not_called() + listener.on_up.assert_not_called() + + def test_down_during_add_preserves_add_before_down_callbacks(self): + cluster = self._new_transition_cluster() + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + listener = Mock() + cluster._listeners = {listener} + entered_add = Event() + release_add = Event() + event_lock = Lock() + events = [] + errors = [] + + def record(event): + with event_lock: + events.append(event) + + def block_add(_): + record("profile-add") + entered_add.set() + assert release_add.wait(2) + + cluster.profile_manager.on_add.side_effect = block_add + cluster.profile_manager.on_down.side_effect = \ + lambda _: record("profile-down") + cluster.control_connection.on_add.side_effect = \ + lambda *_: record("control-add") + cluster.control_connection.on_down.side_effect = \ + lambda _: record("control-down") + listener.on_add.side_effect = lambda _: record("listener-add") + listener.on_down.side_effect = lambda _: record("listener-down") + + def run_add(): + try: + cluster.on_add(host) + except BaseException as exc: + errors.append(exc) + + add_thread = Thread(target=run_add) + add_thread.start() + assert entered_add.wait(2) + + cluster.on_down( + host, + is_host_addition=True, + expect_host_to_be_down=True) + release_add.set() + add_thread.join(2) + + assert not add_thread.is_alive() + assert not errors + assert host.is_up is False + assert events.index("control-add") < events.index("control-down") + assert events.index("listener-add") < events.index("listener-down") + + def test_queued_add_then_down_finishes_down(self): + cluster = self._new_transition_cluster() + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + lane_entered = Event() + release_lane = Event() + + def hold_lane(): + lane_entered.set() + assert release_lane.wait(2) + + lane_thread = Thread( + target=lambda: cluster._run_host_transition( + host, hold_lane)) + lane_thread.start() + assert lane_entered.wait(2) + + cluster.on_add(host) + cluster.on_down( + host, + is_host_addition=True, + expect_host_to_be_down=True, + force=True) + release_lane.set() + lane_thread.join(2) + + assert not lane_thread.is_alive() + assert host.is_up is False + cluster.profile_manager.on_add.assert_called_once_with(host) + cluster.profile_manager.on_down.assert_called_once_with(host) + + def test_preemptive_private_down_fences_older_queued_add(self): + cluster = self._new_transition_cluster() + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + lane_entered = Event() + release_lane = Event() + + def hold_lane(): + lane_entered.set() + assert release_lane.wait(2) + + lane_thread = Thread( + target=lambda: cluster._run_host_transition( + host, hold_lane)) + lane_thread.start() + assert lane_entered.wait(2) + + cluster.on_add(host) + cluster._on_down_locked( + host, + is_host_addition=True, + expect_host_to_be_down=True, + force=True) + release_lane.set() + lane_thread.join(2) + + assert not lane_thread.is_alive() + assert host.is_up is False + cluster.profile_manager.on_add.assert_not_called() + cluster.profile_manager.on_down.assert_called_once_with(host) + + def test_down_during_pending_add_does_not_publish_pool_success(self): + cluster = self._new_transition_cluster() + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.conviction_policy = Mock() + listener = Mock() + cluster._listeners = {listener} + listener_events = [] + listener.on_add.side_effect = \ + lambda _: listener_events.append("add") + listener.on_down.side_effect = \ + lambda _: listener_events.append("down") + delayed_executor_submissions = [] + cluster.on_down_potentially_blocking = \ + lambda *args: delayed_executor_submissions.append(args) + pending_add = Future() + session = Mock() + session.add_or_renew_pool.return_value = pending_add + cluster.sessions = [session] + + cluster.on_add(host) + + assert host._currently_handling_node_add + assert host.is_up is None + + cluster.on_down( + host, + is_host_addition=True, + expect_host_to_be_down=True) + + assert host.is_up is False + assert not host._currently_handling_node_add + host.conviction_policy.reset.assert_not_called() + listener.on_add.assert_called_once_with(host) + listener.on_down.assert_not_called() + assert listener_events == ["add"] + assert len(delayed_executor_submissions) == 1 + + down_host, is_addition, down_epoch = \ + delayed_executor_submissions.pop() + cluster._run_host_transition( + down_host, + lambda: cluster._on_down_potentially_blocking_serialized( + down_host, + is_addition, + down_epoch)) + + listener.on_down.assert_called_once_with(host) + assert listener_events == ["add", "down"] + + # The superseded pool attempt may finish later, but must not publish + # the host UP or replace the newer DOWN transition. + pending_add.set_result(False) + + assert host.is_up is False + host.conviction_policy.reset.assert_not_called() + session.add_or_renew_pool.assert_called_once_with( + host, + is_host_addition=True, + recovery_epoch=1) + + def test_discounted_down_does_not_cancel_pending_add(self): + cluster = self._new_transition_cluster() + cluster._discount_down_events = True + cluster.on_down_potentially_blocking = Mock() + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + listener = Mock() + cluster._listeners = {listener} + pending_add = Future() + session = Mock() + session.add_or_renew_pool.return_value = pending_add + session.get_pool_state.return_value = { + host: {'open_count': 1}} + cluster.sessions = [session] + + cluster.on_add(host) + cluster.on_down( + host, + is_host_addition=True, + expect_host_to_be_down=True) + + assert host.is_up is None + assert host._currently_handling_node_add + assert host._recovery_epoch == 1 + cluster.on_down_potentially_blocking.assert_not_called() + listener.on_add.assert_not_called() + listener.on_down.assert_not_called() + + pending_add.set_result(True) + + assert host.is_up is True + assert not host._currently_handling_node_add + listener.on_add.assert_called_once_with(host) + listener.on_down.assert_not_called() + + def test_external_host_notifications_have_independent_fifo(self): + cluster = self._new_transition_cluster() + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + first_notification_entered = Event() + release_first_notification = Event() + second_core_done = Event() + events = [] + errors = [] + + def first_notification(): + first_notification_entered.set() + assert release_first_notification.wait(2) + events.append("add") + + def run_first(): + try: + cluster._run_host_transition( + host, + lambda: _HostTransitionResult( + notifications=(first_notification,))) + except BaseException as exc: + errors.append(exc) + + def second_notification(): + events.append("down") + raise RuntimeError("second listener failed") + + def run_second(): + try: + cluster._run_host_transition( + host, + lambda: _HostTransitionResult( + notifications=(second_notification,))) + except BaseException as exc: + errors.append(exc) + finally: + second_core_done.set() + + first_thread = Thread(target=run_first) + second_thread = Thread(target=run_second) + first_thread.start() + assert first_notification_entered.wait(2) + second_thread.start() + + assert second_core_done.wait(1) + assert events == [] + + release_first_notification.set() + first_thread.join(2) + second_thread.join(2) + + assert not first_thread.is_alive() + assert not second_thread.is_alive() + assert not errors + assert events == ["add", "down"] + + def test_cross_host_transition_wait_does_not_deadlock(self): + cluster = self._new_transition_cluster() + first_host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + second_host = Host( + "127.0.0.2", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + first_entered = Event() + second_entered = Event() + callbacks = [] + errors = [] + + def first_work(): + first_entered.set() + assert second_entered.wait(2) + cluster._run_host_transition_and_wait( + second_host, + lambda: callbacks.append("from-first")) + + def second_work(): + second_entered.set() + assert first_entered.wait(2) + cluster._run_host_transition_and_wait( + first_host, + lambda: callbacks.append("from-second")) + + def run(host, work): + try: + cluster._run_host_transition(host, work) + except BaseException as exc: + errors.append(exc) + + first_thread = Thread( + target=run, + args=(first_host, first_work), + daemon=True) + second_thread = Thread( + target=run, + args=(second_host, second_work), + daemon=True) + first_thread.start() + second_thread.start() + first_thread.join(2) + second_thread.join(2) + + assert not first_thread.is_alive() + assert not second_thread.is_alive() + assert not errors + assert sorted(callbacks) == ["from-first", "from-second"] + + def test_host_up_reconciles_sessions_before_external_listeners(self): + cluster = self._new_transition_cluster() + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.set_down() + events = [] + listener = Mock() + listener.on_up.side_effect = lambda _: events.append("listener") + cluster._listeners = {listener} + pool_created = Future() + pool_created.set_result(True) + session = Mock() + session.add_or_renew_pool.return_value = pool_created + session.update_created_pools.side_effect = \ + lambda: events.append("session") + cluster.sessions = [session] + + cluster.on_up(host) + + assert events == ["session", "listener"] + + def test_host_up_without_pool_futures_reconciles_and_notifies(self): + cluster = self._new_transition_cluster() + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.set_down() + listener = Mock() + cluster._listeners = {listener} + session = Mock() + session.add_or_renew_pool.return_value = None + cluster.sessions = [session] + + cluster.on_up(host) + + assert host.is_up is True + session.add_or_renew_pool.assert_called_once_with( + host, + is_host_addition=False, + recovery_epoch=host._recovery_epoch) + session.update_created_pools.assert_called_once_with() + listener.on_up.assert_called_once_with(host) + + def test_host_add_reconciles_sessions_before_external_listeners(self): + cluster = self._new_transition_cluster() + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + events = [] + listener = Mock() + listener.on_add.side_effect = lambda _: events.append("listener") + cluster._listeners = {listener} + pool_created = Future() + pool_created.set_result(True) + session = Mock() + session.add_or_renew_pool.return_value = pool_created + session.update_created_pools.side_effect = \ + lambda: events.append("session") + cluster.sessions = [session] + + cluster.on_add(host) + + assert events == ["session", "listener"] + + def test_blocking_up_listener_does_not_block_remove_cleanup(self): + cluster = self._new_transition_cluster() + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.set_down() + host._recovery_epoch = 4 + host._currently_handling_node_up = True + session = Mock() + cluster.sessions = [session] + listener = Mock() + cluster._listeners = {listener} + entered_listener = Event() + release_listener = Event() + remove_done = Event() + errors = [] + + def block_on_up(_): + entered_listener.set() + assert release_listener.wait(2) + + listener.on_up.side_effect = block_on_up + + def finish_up(): + try: + cluster._run_host_transition( + host, + lambda: cluster._finish_on_up(host, 4, [True])) + except BaseException as exc: + errors.append(exc) + + def remove(): + try: + cluster.on_remove(host) + except BaseException as exc: + errors.append(exc) + finally: + remove_done.set() + + up_thread = Thread(target=finish_up) + remove_thread = Thread(target=remove) + up_thread.start() + assert entered_listener.wait(2) + remove_thread.start() + cleanup_completed_while_listener_blocked = remove_done.wait(1) + release_listener.set() + up_thread.join(2) + remove_thread.join(2) + + assert cleanup_completed_while_listener_blocked + assert not up_thread.is_alive() + assert not remove_thread.is_alive() + assert not errors + assert host._is_removed + assert host.is_up is False + session.on_remove.assert_called_once_with(host) + cluster.control_connection.on_remove.assert_called_once_with(host) + + def test_up_queued_during_on_down_leaves_host_and_policies_up(self): + cluster = self._new_transition_cluster() + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.set_up() + entered_on_down = Event() + release_on_down = Event() + policy_events = [] + + def block_on_down(_): + policy_events.append("down") + entered_on_down.set() + assert release_on_down.wait(2) + + cluster.profile_manager.on_down.side_effect = block_on_down + cluster.profile_manager.on_up.side_effect = ( + lambda _: policy_events.append("up")) + + down_thread = Thread( + target=cluster._on_down_locked, + kwargs={ + 'host': host, + 'is_host_addition': False, + 'expect_host_to_be_down': True, + 'force': True, + }) + down_thread.start() + assert entered_on_down.wait(2) + + cluster.on_up(host) + release_on_down.set() + down_thread.join(2) + + assert not down_thread.is_alive() + assert host.is_up is True + assert policy_events == ["down", "up"] + cluster.control_connection.on_down.assert_called_once_with(host) + cluster.control_connection.on_up.assert_called_once_with(host) + + def test_failed_up_preserves_delayed_down_notification(self): + cluster = self._new_transition_cluster() + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.set_up() + listener = Mock() + cluster._listeners = {listener} + delayed_down = [] + cluster.on_down_potentially_blocking = ( + lambda *args: delayed_down.append(args)) + + cluster.on_down( + host, + is_host_addition=False, + expect_host_to_be_down=True, + force=True) + + pool_future = Future() + session = Mock() + session.add_or_renew_pool.return_value = pool_future + cluster.sessions = [session] + cluster.on_up(host) + + assert host._currently_handling_node_up + assert listener.on_down.call_count == 0 + down_args = delayed_down.pop() + cluster._run_host_transition( + host, + lambda: cluster._on_down_potentially_blocking_serialized( + *down_args)) + pool_future.set_result(False) + + assert host.is_up is False + listener.on_down.assert_called_once_with(host) + listener.on_up.assert_not_called() + + def test_failed_up_does_not_repeat_published_down_notification(self): + cluster = self._new_transition_cluster() + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.set_up() + listener = Mock() + cluster._listeners = {listener} + + cluster.on_down( + host, + is_host_addition=False, + expect_host_to_be_down=True, + force=True) + listener.on_down.assert_called_once_with(host) + + pool_future = Future() + session = Mock() + session.add_or_renew_pool.return_value = pool_future + cluster.sessions = [session] + cluster.on_up(host) + pool_future.set_result(False) + + assert host.is_up is False + listener.on_down.assert_called_once_with(host) + listener.on_up.assert_not_called() + + def test_successful_up_suppresses_delayed_down_notification(self): + cluster = self._new_transition_cluster() + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.set_up() + listener = Mock() + cluster._listeners = {listener} + delayed_down = [] + cluster.on_down_potentially_blocking = ( + lambda *args: delayed_down.append(args)) + + cluster.on_down( + host, + is_host_addition=False, + expect_host_to_be_down=True, + force=True) + + pool_future = Future() + session = Mock() + session.add_or_renew_pool.return_value = pool_future + cluster.sessions = [session] + cluster.on_up(host) + pool_future.set_result(True) + + down_args = delayed_down.pop() + cluster._run_host_transition( + host, + lambda: cluster._on_down_potentially_blocking_serialized( + *down_args)) + + assert host.is_up is True + listener.on_down.assert_not_called() + listener.on_up.assert_called_once_with(host) + + def test_reentrant_add_failure_transitions_to_recovery_and_propagates(self): + cluster = self._new_transition_cluster() + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + cluster.profile_manager.on_add.side_effect = RuntimeError( + "policy add failed") + + cluster._run_host_transition( + host, + lambda: cluster.on_add(host)) + + assert host.is_up is False + assert not host._currently_handling_node_add + cluster.profile_manager.on_down.assert_called_once_with(host) + cluster.control_connection.on_down.assert_called_once_with(host) + cluster._start_reconnector.assert_called_once_with( + host, + True, + recovery_epoch=host._recovery_epoch) + + def test_add_resets_conviction_without_holding_host_lock(self): + cluster = self._new_transition_cluster() + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + acquired = Event() + reset_thread = [] + + def reset_from_another_thread(): + def acquire_host_lock(): + with host.lock: + acquired.set() + + thread = Thread(target=acquire_host_lock) + reset_thread.append(thread) + thread.start() + assert acquired.wait(1) + thread.join(1) + + host.conviction_policy.reset = reset_from_another_thread + + cluster.on_add(host) + + assert acquired.is_set() + assert not reset_thread[0].is_alive() + assert host.is_up is True + assert not host._currently_handling_node_add + + def test_add_reset_failure_transitions_to_recovery(self): + class ResetCancelled(BaseException): + pass + + cluster = self._new_transition_cluster() + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.conviction_policy.reset = Mock( + side_effect=ResetCancelled("cancelled reset")) + + with pytest.raises(ResetCancelled): + cluster.on_add(host) + + assert host.is_up is False + assert not host._currently_handling_node_add + cluster._start_reconnector.assert_called_once_with( + host, + True, + recovery_epoch=host._recovery_epoch) + + def test_up_reset_failure_restores_reconnector(self): + class ResetCancelled(BaseException): + pass + + cluster = self._new_transition_cluster() + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.set_down() + host.conviction_policy.reset = Mock( + side_effect=ResetCancelled("cancelled reset")) + + with pytest.raises(ResetCancelled): + cluster.on_up(host) + + assert host.is_up is False + assert not host._currently_handling_node_up + cluster._start_reconnector.assert_called_once_with( + host, + is_host_addition=False, + recovery_epoch=host._recovery_epoch) + + def test_down_callback_failure_still_starts_reconnector(self): + cluster = self._new_transition_cluster() + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.set_up() + session = Mock() + listener = Mock() + cluster.sessions = [session] + cluster._listeners = {listener} + cluster.profile_manager.on_down.side_effect = RuntimeError( + "policy down failed") + + cluster._on_down_locked( + host, + is_host_addition=False, + expect_host_to_be_down=True, + force=True) + + assert host.is_up is False + cluster.control_connection.on_down.assert_called_once_with(host) + session.on_down.assert_called_once_with(host) + listener.on_down.assert_called_once_with(host) + cluster._start_reconnector.assert_called_once_with( + host, + False, + recovery_epoch=host._recovery_epoch) + + def test_rejected_down_executor_uses_fallback_cleanup(self): + cluster = self._new_transition_cluster() + del cluster.on_down_potentially_blocking + cluster.executor = Mock() + cluster.executor.submit.side_effect = RuntimeError( + "executor rejected cleanup") + host = Host( + DefaultEndPoint("127.0.0.1"), + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.set_up() + session = Mock() + listener = Mock() + cluster.sessions = [session] + cluster._listeners = {listener} + cleanup_completed = Event() + cluster._start_reconnector.side_effect = \ + lambda *args, **kwargs: cleanup_completed.set() + startup_error = _ConnectionClosedDuringStartup( + "closed during startup", + host.endpoint) + + cluster.signal_connection_failure( + host, + startup_error, + is_host_addition=False, + expect_host_to_be_down=True, + force=True) + + assert cleanup_completed.wait(2) + assert host.is_up is False + cluster.profile_manager.on_down.assert_called_once_with(host) + cluster.control_connection.on_down.assert_called_once_with(host) + session.on_down.assert_called_once_with(host) + listener.on_down.assert_called_once_with(host) + cluster._start_reconnector.assert_called_once_with( + host, + False, + recovery_epoch=host._recovery_epoch) + + def test_reconnector_cancel_error_does_not_strand_transitions(self): + for transition in ("up", "add", "remove"): + with self.subTest(transition=transition): + cluster = self._new_transition_cluster() + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + if transition == "up": + host.set_down() + elif transition == "remove": + host.set_up() + reconnector = Mock() + reconnector.cancel.side_effect = RuntimeError( + "cancel failed") + host.get_and_set_reconnection_handler(reconnector) + + if transition == "up": + cluster.on_up(host) + assert host.is_up is True + cluster.profile_manager.on_up.assert_called_once_with( + host) + cluster.control_connection.on_up.assert_called_once_with( + host) + elif transition == "add": + cluster.on_add(host) + assert host.is_up is True + cluster.profile_manager.on_add.assert_called_once_with( + host) + cluster.control_connection.on_add.assert_called_once_with( + host, True) + else: + cluster.on_remove(host) + assert host._is_removed + assert host.is_up is False + cluster.profile_manager.on_remove.assert_called_once_with( + host) + cluster.control_connection.on_remove.assert_called_once_with( + host) + + reconnector.cancel.assert_called_once_with() + + def test_session_invalidation_error_does_not_suppress_down_cleanup(self): + cluster = self._new_transition_cluster() + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.set_up() + failing_session = Mock() + failing_session._invalidate_pool_attempts.side_effect = RuntimeError( + "invalidation failed") + healthy_session = Mock() + cluster.sessions = [failing_session, healthy_session] + + cluster.on_down( + host, + is_host_addition=False, + expect_host_to_be_down=True, + force=True) + + assert host.is_up is False + failing_session._invalidate_pool_attempts.assert_called_once_with( + host) + healthy_session._invalidate_pool_attempts.assert_called_once_with( + host) + failing_session.on_down.assert_called_once_with(host) + healthy_session.on_down.assert_called_once_with(host) + cluster.profile_manager.on_down.assert_called_once_with(host) + cluster.control_connection.on_down.assert_called_once_with(host) + cluster._start_reconnector.assert_called_once_with( + host, + False, + recovery_epoch=host._recovery_epoch) + + def test_session_invalidation_error_does_not_suppress_remove_cleanup(self): + cluster = self._new_transition_cluster() + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.set_up() + failing_session = Mock() + failing_session._invalidate_pool_attempts.side_effect = RuntimeError( + "invalidation failed") + healthy_session = Mock() + cluster.sessions = [failing_session, healthy_session] + + cluster.on_remove(host) + + assert host._is_removed + assert host.is_up is False + failing_session._invalidate_pool_attempts.assert_called_once_with( + host) + healthy_session._invalidate_pool_attempts.assert_called_once_with( + host) + failing_session.on_remove.assert_called_once_with(host) + healthy_session.on_remove.assert_called_once_with(host) + cluster.profile_manager.on_remove.assert_called_once_with(host) + cluster.control_connection.on_remove.assert_called_once_with(host) + + def test_failed_up_cleanup_error_does_not_skip_pool_teardown(self): + cluster = self._new_transition_cluster() + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.set_down() + host._recovery_epoch = 3 + first_session = Mock() + second_session = Mock() + cluster.sessions = [first_session, second_session] + cluster.profile_manager.on_down.side_effect = RuntimeError( + "policy cleanup failed") + + with pytest.raises(RuntimeError, match="policy cleanup failed"): + cluster._cleanup_failed_on_up_handling( + host, + recovery_epoch=3) + + cluster.control_connection.on_down.assert_called_once_with(host) + first_session.remove_pool.assert_called_once_with(host) + second_session.remove_pool.assert_called_once_with(host) + cluster._start_reconnector.assert_called_once_with( + host, + is_host_addition=False, + recovery_epoch=3) + + def test_up_listener_error_does_not_skip_remaining_reconciliation(self): + cluster = self._new_transition_cluster() + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.set_down() + host._recovery_epoch = 4 + host._currently_handling_node_up = True + failing_listener = Mock() + healthy_listener = Mock() + failing_listener.on_up.side_effect = RuntimeError( + "listener failed") + cluster._listeners = {failing_listener, healthy_listener} + session = Mock() + cluster.sessions = [session] + + cluster._run_host_transition( + host, + lambda: cluster._finish_on_up(host, 4, [True])) + + assert host.is_up is True + healthy_listener.on_up.assert_called_once_with(host) + session.update_created_pools.assert_called_once_with() + + def test_remove_callback_failure_does_not_skip_pool_cleanup(self): + cluster = self._new_transition_cluster() + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.set_up() + session = Mock() + listener = Mock() + cluster.sessions = [session] + cluster._listeners = {listener} + cluster.profile_manager.on_remove.side_effect = RuntimeError( + "policy remove failed") + + cluster.on_remove(host) + + assert host._is_removed + assert host.is_up is False + session.on_remove.assert_called_once_with(host) + listener.on_remove.assert_called_once_with(host) + cluster.control_connection.on_remove.assert_called_once_with(host) + + def test_async_add_base_exception_transitions_to_recovery(self): + class PoolSetupCancelled(BaseException): + pass + + cluster = self._new_transition_cluster() + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + pool_future = Future() + session = Mock() + session.add_or_renew_pool.return_value = pool_future + cluster.sessions = [session] + + cluster.on_add(host) + assert host._currently_handling_node_add + + cancellation = PoolSetupCancelled("cancelled") + pool_future.set_exception(cancellation) + + assert pool_future.exception() is cancellation + assert host.is_up is False + assert not host._currently_handling_node_add + cluster.profile_manager.on_down.assert_called_once_with(host) + cluster.control_connection.on_down.assert_called_once_with(host) + session.on_down.assert_called_once_with(host) + cluster._start_reconnector.assert_called_once_with( + host, + True, + recovery_epoch=host._recovery_epoch) + + def test_async_add_false_result_transitions_to_recovery(self): + cluster = self._new_transition_cluster() + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + pool_future = Future() + session = Mock() + session.add_or_renew_pool.return_value = pool_future + cluster.sessions = [session] + + cluster.on_add(host) + assert host._currently_handling_node_add + + pool_future.set_result(False) + + assert host.is_up is False + assert not host._currently_handling_node_add + cluster.profile_manager.on_down.assert_called_once_with(host) + cluster.control_connection.on_down.assert_called_once_with(host) + session.on_down.assert_called_once_with(host) + cluster._start_reconnector.assert_called_once_with( + host, + True, + recovery_epoch=host._recovery_epoch) + + def test_stale_pool_result_is_neutral_for_add(self): + cluster = self._new_transition_cluster() + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host._recovery_epoch = 3 + host._currently_handling_node_add = True + + cluster._run_host_transition( + host, + lambda: cluster._finish_add( + host, + 3, + [True, _STALE_POOL_ATTEMPT])) + + assert host.is_up is True + cluster._start_reconnector.assert_not_called() + cluster.profile_manager.on_down.assert_not_called() + + def test_stale_pool_result_is_neutral_for_up(self): + cluster = self._new_transition_cluster() + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.set_down() + host._recovery_epoch = 3 + host._currently_handling_node_up = True + + cluster._run_host_transition( + host, + lambda: cluster._finish_on_up( + host, + 3, + [True, _STALE_POOL_ATTEMPT])) + + assert host.is_up is True + cluster._start_reconnector.assert_not_called() + cluster.profile_manager.on_down.assert_not_called() + + def test_async_up_base_exception_restores_reconnector(self): + class PoolSetupCancelled(BaseException): + pass + + cluster = self._new_transition_cluster() + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.set_down() + pool_future = Future() + session = Mock() + session.add_or_renew_pool.return_value = pool_future + cluster.sessions = [session] + + cluster.on_up(host) + recovery_epoch = host._recovery_epoch + assert host._currently_handling_node_up + + cancellation = PoolSetupCancelled("cancelled") + pool_future.set_exception(cancellation) + + assert pool_future.exception() is cancellation + assert host.is_up is False + assert not host._currently_handling_node_up + cluster.profile_manager.on_down.assert_called_once_with(host) + cluster.control_connection.on_down.assert_called_once_with(host) + assert session.remove_pool.call_count == 2 + cluster._start_reconnector.assert_called_once_with( + host, + is_host_addition=False, + recovery_epoch=recovery_epoch) + + def test_connection_failure_preserves_legacy_on_down_override(self): + class LegacyCluster(Cluster): + + def on_down(self, host, is_host_addition, + expect_host_to_be_down=False): + self.down_call = ( + host, is_host_addition, expect_host_to_be_down) + + host = Mock() + host.signal_connection_failure.return_value = True + cluster = object.__new__(LegacyCluster) + + cluster.signal_connection_failure( + host, + ConnectionException("boom"), + is_host_addition=False, + expect_host_to_be_down=True) + + assert cluster.down_call == (host, False, True) + + def test_prepare_all_queries_bounds_keyspace_setup(self): + endpoint = Mock(address='127.0.0.1') + host = Mock(endpoint=endpoint) + connection = Mock(is_closed=False) + statement = Mock(keyspace='ks', query_string='SELECT * FROM tbl') + cluster = Cluster(protocol_version=ProtocolVersion.V4) + cluster._prepared_statements = {b'query-id': statement} + + with patch.object( + cluster, 'connection_factory', return_value=connection), \ + patch.object(cluster, '_send_chunks'): + cluster._prepare_all_queries(host) + + connection.set_keyspace_blocking.assert_called_once_with( + 'ks', timeout=cluster.connect_timeout) + + +class SessionPoolLifecycleTest(unittest.TestCase): + + @staticmethod + def _new_session(): + cluster = Mock() + cluster._lock = RLock() + cluster.is_shutdown = False + cluster.allow_control_connection_query_fallback = \ + ControlConnectionQueryFallback.Disabled + cluster.connect_timeout = 1 + cluster.profile_manager.distance.return_value = HostDistance.LOCAL + + session = object.__new__(Session) + session.cluster = cluster + session.keyspace = None + session.is_shutdown = False + session._lock = RLock() + session._keyspace_dispatch_lock = RLock() + session._keyspace_completion_lock = Lock() + session._keyspace_completion_queue = [] + session._keyspace_completion_runner_active = False + session._keyspace_generation = 0 + session._pools = {} + session._pool_generations = {} + session._profile_manager = cluster.profile_manager + session._initial_connect_futures = set() + session._monitor_reporter = None + session._pool_repair_schedule = None + session._pool_repair_scheduled = False + return session + + @staticmethod + def _new_host(): + host = Host( + "127.0.0.1", + SimpleConvictionPolicy, + host_id=uuid.uuid4()) + host.set_up() + return host + + @staticmethod + def _capture_submissions(session): + tasks = [] + + def submit(fn, *args, **kwargs): + tasks.append(lambda: fn(*args, **kwargs)) + return Future() + + session.submit = submit + return tasks + + def test_initial_pool_clean_close_forces_host_recovery(self): + session = self._new_session() + host = self._new_host() + tasks = self._capture_submissions(session) + startup_error = _ConnectionClosedDuringStartup( + "closed during startup", host.endpoint) + + with patch( + 'cassandra.cluster.HostConnection', + side_effect=startup_error): + session.add_or_renew_pool(host, is_host_addition=False) + task_result = tasks.pop()() + assert not task_result + + session.cluster.signal_connection_failure.assert_called_once_with( + host, + startup_error, + False, + expect_host_to_be_down=True, + force=True) + + def test_clean_startup_close_recovers_when_conviction_policy_raises(self): + class PolicyCancelled(BaseException): + pass + + class ClusterStandin(object): + + def __init__(self): + self._lock = RLock() + self.is_shutdown = False + self.allow_control_connection_query_fallback = ( + ControlConnectionQueryFallback.Disabled) + self.connect_timeout = 1 + self.profile_manager = Mock() + self.profile_manager.distance.return_value = ( + HostDistance.LOCAL) + self.down_calls = [] + + def _uses_default_failure_hooks(self): + return True + + def _on_down_locked( + self, host, is_host_addition, + expect_host_to_be_down=False, force=False): + self.down_calls.append(( + host, + is_host_addition, + expect_host_to_be_down, + force)) + + session = self._new_session() + session.cluster = ClusterStandin() + session._profile_manager = session.cluster.profile_manager + host = self._new_host() + host.signal_connection_failure = Mock( + side_effect=PolicyCancelled("cancelled policy")) + tasks = self._capture_submissions(session) + startup_error = _ConnectionClosedDuringStartup( + "closed during startup", host.endpoint) + + with patch( + 'cassandra.cluster.HostConnection', + side_effect=startup_error): + session.add_or_renew_pool(host, is_host_addition=False) + task_result = tasks.pop()() + assert not task_result + + assert session.cluster.down_calls == [ + (host, False, True, True), + ] + + def test_clean_startup_delegated_force_is_not_committed_twice(self): + calls = [] + + class DelegatingCluster(object): + def __init__(self): + self._lock = RLock() + self.is_shutdown = False + self.allow_control_connection_query_fallback = ( + ControlConnectionQueryFallback.Disabled) + self.connect_timeout = 1 + self.profile_manager = Mock() + self.profile_manager.distance.return_value = ( + HostDistance.LOCAL) + + def _uses_default_failure_hooks(self): + return False + + def signal_connection_failure( + self, host, exc, is_host_addition, + expect_host_to_be_down=False, force=False): + calls.append(("public", force)) + self._on_down_locked( + host, + is_host_addition, + expect_host_to_be_down, + force) + return False + + def _on_down_locked( + self, host, is_host_addition, + expect_host_to_be_down=False, force=False): + calls.append(("private", force)) + host._recovery_epoch += 1 + + session = self._new_session() + session.cluster = DelegatingCluster() + session._profile_manager = session.cluster.profile_manager + host = self._new_host() + starting_epoch = host._recovery_epoch + tasks = self._capture_submissions(session) + startup_error = _ConnectionClosedDuringStartup( + "closed during startup", host.endpoint) + + with patch( + 'cassandra.cluster.HostConnection', + side_effect=startup_error): + session.add_or_renew_pool(host, is_host_addition=False) + task_result = tasks.pop()() + assert not task_result + + assert calls == [ + ("public", True), + ("private", True), + ] + assert host._recovery_epoch == starting_epoch + 1 + + def test_pool_failure_preserves_legacy_cluster_on_down_override(self): + class LegacyCluster(Cluster): + + def on_down( + self, host, is_host_addition, + expect_host_to_be_down=False): + self.down_calls.append(( + host, + is_host_addition, + expect_host_to_be_down)) + + session = self._new_session() + session.cluster = object.__new__(LegacyCluster) + session.cluster._lock = RLock() + session.cluster.is_shutdown = False + session.cluster.allow_control_connection_query_fallback = ( + ControlConnectionQueryFallback.Disabled) + session.cluster.connect_timeout = 1 + session.cluster.profile_manager = Mock() + session.cluster.profile_manager.distance.return_value = ( + HostDistance.LOCAL) + session.cluster.down_calls = [] + session.cluster.sessions = [session] + session.cluster._discount_down_events = False + session.cluster.on_down_potentially_blocking = Mock() + session._profile_manager = session.cluster.profile_manager + host = self._new_host() + tasks = self._capture_submissions(session) + startup_error = _ConnectionClosedDuringStartup( + "closed during startup", host.endpoint) + + with patch( + 'cassandra.cluster.HostConnection', + side_effect=startup_error): + session.add_or_renew_pool(host, is_host_addition=False) + task_result = tasks.pop()() + assert not task_result + + assert session.cluster.down_calls == [ + (host, False, True), + ] + assert host.is_up is False + session.cluster.on_down_potentially_blocking.assert_called_once_with( + host, + False, + host._recovery_epoch) + + def test_stale_startup_close_does_not_mark_host_down(self): + session = self._new_session() + host = self._new_host() + tasks = self._capture_submissions(session) + startup_error = _ConnectionClosedDuringStartup( + "closed during startup", host.endpoint) + + with patch( + 'cassandra.cluster.HostConnection', + side_effect=startup_error): + session.add_or_renew_pool(host, is_host_addition=False) + session.remove_pool(host) + task_result = tasks.pop()() + assert not task_result + + session.cluster.signal_connection_failure.assert_not_called() + + def test_invalid_initial_keyspace_does_not_mark_host_down(self): + session = self._new_session() + host = self._new_host() + tasks = self._capture_submissions(session) + keyspace_error = InvalidRequest("keyspace does not exist") + + with patch( + 'cassandra.cluster.HostConnection', + side_effect=keyspace_error): + session.add_or_renew_pool(host, is_host_addition=False) + task_result = tasks.pop()() + assert not task_result + + session.cluster.signal_connection_failure.assert_not_called() + + def test_invalid_changed_keyspace_does_not_mark_host_down(self): + session = self._new_session() + session.keyspace = 'new_ks' + host = self._new_host() + tasks = self._capture_submissions(session) + keyspace_error = InvalidRequest("keyspace does not exist") + new_pool = Mock(_keyspace='old_ks') + new_pool._set_keyspace_for_all_conns.side_effect = \ + lambda keyspace, callback: callback( + new_pool, [keyspace_error]) + + with patch( + 'cassandra.cluster.HostConnection', + return_value=new_pool): + session.add_or_renew_pool(host, is_host_addition=False) + task_result = tasks.pop()() + assert not task_result + + session.cluster.on_down.assert_not_called() + new_pool.shutdown.assert_called_once_with() + + def test_session_local_pool_failure_retries_until_success(self): + session = self._new_session() + host = self._new_host() + session.cluster.metadata.all_hosts.return_value = [host] + session.cluster.reconnection_policy.new_schedule.return_value = iter( + [0.1, 0.2]) + + first_failure = Future() + first_failure.set_result(_SESSION_LOCAL_POOL_FAILURE) + second_failure = Future() + second_failure.set_result(_SESSION_LOCAL_POOL_FAILURE) + success = Future() + session.add_or_renew_pool = Mock( + side_effect=[first_failure, second_failure, success]) + + assert session.update_created_pools() == {first_failure} + first_retry = session.cluster.scheduler.schedule.call_args_list[0] + assert first_retry.args[0] == 0.1 + first_retry.args[1]() + + assert session._pool_repair_scheduled + second_retry = session.cluster.scheduler.schedule.call_args_list[1] + assert second_retry.args[0] == 0.2 + second_retry.args[1]() + + assert not session._pool_repair_scheduled + assert session._pool_repair_schedule is not None + success.set_result(True) + + assert session.add_or_renew_pool.call_count == 3 + assert session.cluster.scheduler.schedule.call_count == 2 + assert not session._pool_repair_scheduled + assert session._pool_repair_schedule is None + + def test_post_connect_cancellation_closes_unpublished_pool(self): + class SetupCancelled(BaseException): + pass + + session = self._new_session() + session.keyspace = 'new_ks' + host = self._new_host() + tasks = self._capture_submissions(session) + new_pool = Mock(_keyspace='old_ks') + new_pool._set_keyspace_for_all_conns.side_effect = SetupCancelled( + "cancelled keyspace reconciliation") + + with patch( + 'cassandra.cluster.HostConnection', + return_value=new_pool): + session.add_or_renew_pool(host, is_host_addition=False) + with pytest.raises(SetupCancelled): + tasks.pop()() + + assert host not in session._pools + new_pool.shutdown.assert_called_once_with() + + def test_pool_finishing_after_shutdown_is_closed_not_published(self): + session = self._new_session() + host = self._new_host() + tasks = self._capture_submissions(session) + new_pool = Mock(_keyspace=None) + + with patch( + 'cassandra.cluster.HostConnection', + return_value=new_pool): + session.add_or_renew_pool(host, is_host_addition=False) + session.shutdown() + task_result = tasks.pop()() + assert task_result is _STALE_POOL_ATTEMPT + + assert host not in session._pools + new_pool.shutdown.assert_called_once_with() + + def test_pool_finishing_after_host_down_is_closed_not_published(self): + session = self._new_session() + host = self._new_host() + tasks = self._capture_submissions(session) + new_pool = Mock(_keyspace=None) + + with patch( + 'cassandra.cluster.HostConnection', + return_value=new_pool): + session.add_or_renew_pool(host, is_host_addition=False) + host.set_down() + session.on_down(host) + task_result = tasks.pop()() + assert task_result is _STALE_POOL_ATTEMPT + + assert host not in session._pools + new_pool.shutdown.assert_called_once_with() + + def test_unregistered_session_cannot_publish_pool_after_host_down(self): + session = self._new_session() + host = self._new_host() + tasks = self._capture_submissions(session) + new_pool = Mock(_keyspace=None) + + with patch( + 'cassandra.cluster.HostConnection', + return_value=new_pool): + session.add_or_renew_pool(host, is_host_addition=False) + # Session.__init__ is not yet registered in Cluster.sessions, so a + # concurrent DOWN cannot invalidate this attempt through + # Session.on_down. Host state must fence publication itself. + host.set_down() + task_result = tasks.pop()() + assert task_result is _STALE_POOL_ATTEMPT + + assert host not in session._pools + new_pool.shutdown.assert_called_once_with() + + def test_concurrent_renewals_publish_only_one_pool(self): + session = self._new_session() + host = self._new_host() + old_pool = Mock() + first_pool = Mock(_keyspace=None) + second_pool = Mock(_keyspace=None) + session._pools[host] = old_pool + tasks = self._capture_submissions(session) + + with patch( + 'cassandra.cluster.HostConnection', + side_effect=[first_pool, second_pool]): + session.add_or_renew_pool(host, is_host_addition=False) + session.add_or_renew_pool(host, is_host_addition=False) + task_result = tasks.pop(0)() + assert task_result is True + task_result = tasks.pop(0)() + assert task_result is _STALE_POOL_ATTEMPT + + assert session._pools[host] is first_pool + old_pool.shutdown.assert_called_once_with() + first_pool.shutdown.assert_not_called() + second_pool.shutdown.assert_called_once_with() + + def test_remove_pool_closes_synchronously_if_submit_loses_shutdown_race(self): + session = self._new_session() + host = self._new_host() + pool = Mock() + session._pools[host] = pool + session.submit = Mock(return_value=None) + + assert session.remove_pool(host) is None + + assert host not in session._pools + pool.shutdown.assert_called_once_with() + + def test_remove_pool_cancellation_retains_shutdown_ownership(self): + session = self._new_session() + host = self._new_host() + pool = Mock() + session._pools[host] = pool + shutdown_future = Future() + session.submit = Mock(return_value=shutdown_future) + + assert session.remove_pool(host) is shutdown_future + assert shutdown_future.cancel() + + assert host not in session._pools + pool.shutdown.assert_called_once_with() + + def test_shutdown_atomically_drains_pool_mapping(self): + session = self._new_session() + host = self._new_host() + pool = Mock() + session._pools[host] = pool + + session.shutdown() + + assert session._pools == {} + pool.shutdown.assert_called_once_with() + + def test_keyspace_completion_uses_snapshot_and_is_exactly_once(self): + session = self._new_session() + first_pool = Mock(host='host1') + second_pool = Mock(host='host2') + callbacks = {} + first_pool._set_keyspace_for_all_conns.side_effect = \ + lambda keyspace, callback: callbacks.setdefault( + first_pool, callback) + second_pool._set_keyspace_for_all_conns.side_effect = \ + lambda keyspace, callback: callbacks.setdefault( + second_pool, callback) + session._pools = { + first_pool.host: first_pool, + second_pool.host: second_pool, + } + completed = Mock() + + session._set_keyspace_for_all_pools('ks', completed) + session._pools = {} + callbacks[first_pool](first_pool, []) + callbacks[first_pool](first_pool, []) + callbacks[second_pool](second_pool, []) + + completed.assert_called_once_with({}) + + def test_keyspace_dispatch_base_exception_completes_generation(self): + class PoolUpdateCancelled(BaseException): + pass + + session = self._new_session() + pool = Mock(host='host1') + cancellation = PoolUpdateCancelled("cancelled") + pool._set_keyspace_for_all_conns.side_effect = cancellation + session._pools = {pool.host: pool} + completed = Mock() + + session._set_keyspace_for_all_pools('ks', completed) + + completed.assert_called_once_with({ + pool.host: [cancellation], + }) + assert session._keyspace_completion_queue == [] + + def test_keyspace_completion_callbacks_preserve_dispatch_order(self): + session = self._new_session() + pool = Mock(host='host1') + pool_callbacks = [] + pool._set_keyspace_for_all_conns.side_effect = ( + lambda keyspace, callback: pool_callbacks.append(callback)) + session._pools = {pool.host: pool} + completions = [] + + session._set_keyspace_for_all_pools( + 'first', + lambda errors: completions.append(('first', errors))) + session._pools = {} + session._set_keyspace_for_all_pools( + 'second', + lambda errors: completions.append(('second', errors))) + + # The second generation completed synchronously on an empty snapshot, + # but it must remain behind the first generation. + assert completions == [] + pool_callbacks[0](pool, []) + + assert completions == [ + ('first', {}), + ('second', {}), + ] + class SchedulerTest(unittest.TestCase): # TODO: this suite could be expanded; for now just adding a test covering a ticket diff --git a/tests/unit/test_connection.py b/tests/unit/test_connection.py index 1f9a3f682c..2a730f00f5 100644 --- a/tests/unit/test_connection.py +++ b/tests/unit/test_connection.py @@ -12,20 +12,27 @@ # See the License for the specific language governing permissions and # limitations under the License. import itertools +import socket import unittest from io import BytesIO import time -from threading import Lock +from threading import Event, Lock, RLock, Thread, get_ident from unittest.mock import Mock, ANY, call, patch -from cassandra import OperationTimedOut +from cassandra import InvalidRequest, OperationTimedOut, Unauthorized from cassandra.cluster import Cluster -from cassandra.connection import (Connection, HEADER_DIRECTION_TO_CLIENT, ProtocolError, +from cassandra.connection import (Connection, ConnectionBusy, HEADER_DIRECTION_TO_CLIENT, ProtocolError, locally_supported_compressions, ConnectionHeartbeat, HeartbeatFuture, _Frame, Timer, TimerManager, - ConnectionException, ConnectionShutdown, DefaultEndPoint, ShardAwarePortGenerator) + ConnectionException, ConnectionShutdown, DefaultEndPoint, ShardAwarePortGenerator, + _ConnectionClosedDuringStartup, + _set_keyspace_blocking, + _startup_close_error) from cassandra.marshal import uint8_pack, uint32_pack, int32_pack from cassandra.protocol import (write_stringmultimap, write_int, write_string, SupportedMessage, ProtocolHandler, ResultMessage, + ReadyMessage, AuthSuccessMessage, + InvalidRequestException, + UnauthorizedErrorMessage, RESULT_KIND_SET_KEYSPACE) from tests.util import wait_until, assertRegex @@ -272,6 +279,73 @@ def test_set_keyspace_blocking_escapes_quotes(self): assert query_msg.query == 'USE "my""ks"', ( "Double quotes in keyspace name must be escaped as double-double quotes") + def test_set_keyspace_blocking_passes_timeout(self): + c = self.make_connection() + c.wait_for_response = Mock( + return_value=ResultMessage(kind=RESULT_KIND_SET_KEYSPACE)) + + c.set_keyspace_blocking('ks', timeout=1.25) + + query_msg = c.wait_for_response.call_args[0][0] + assert query_msg.query == 'USE "ks"' + c.wait_for_response.assert_called_once_with( + query_msg, timeout=1.25) + + def test_keyspace_timeout_adapter_preserves_custom_signatures(self): + calls = [] + + class HistoricalConnection(object): + def set_keyspace_blocking(self, keyspace): + calls.append(("historical", keyspace)) + + class PositionalTimeoutConnection(object): + def set_keyspace_blocking(self, keyspace, timeout, /): + calls.append(("positional", keyspace, timeout)) + + class VarargsTimeoutConnection(object): + def set_keyspace_blocking(self, keyspace, *args): + calls.append(("varargs", keyspace, args)) + + _set_keyspace_blocking(HistoricalConnection(), "ks1", 1.5) + _set_keyspace_blocking(PositionalTimeoutConnection(), "ks2", 2.5) + _set_keyspace_blocking(VarargsTimeoutConnection(), "ks3", 3.5) + + assert calls == [ + ("historical", "ks1"), + ("positional", "ks2", 2.5), + ("varargs", "ks3", (3.5,)), + ] + + def test_validation_response_does_not_defunct_transport(self): + cases = ( + ( + InvalidRequestException( + code=0x2200, message="invalid", info=None), + InvalidRequest), + ( + UnauthorizedErrorMessage( + code=0x2100, message="unauthorized", info=None), + Unauthorized), + ) + for response, error_type in cases: + with self.subTest(error_type=error_type): + c = self.make_connection() + success = ResultMessage(kind=RESULT_KIND_SET_KEYSPACE) + responses = [response, success] + + def send_response(message, request_id, callback): + callback(responses.pop(0)) + + c.send_msg = send_response + + with pytest.raises(error_type): + c.wait_for_response(Mock()) + + assert not c.is_defunct + assert not c.is_closed + assert c.last_error is None + assert c.wait_for_response(Mock()) is success + def test_set_keyspace_async_escapes_quotes(self): """ Test that set_keyspace_async properly escapes double quotes in @@ -291,6 +365,85 @@ def test_set_keyspace_async_escapes_quotes(self): assert query_msg.query == 'USE "my""ks"', ( "Double quotes in keyspace name must be escaped as double-double quotes") + def test_set_keyspace_async_pre_push_failure_restores_request_id(self): + c = self.make_connection() + original_request_ids = tuple(c.request_ids) + request_id = original_request_ids[0] + c._socket_writable = False + + with pytest.raises(ConnectionBusy): + c.set_keyspace_async("ks", Mock()) + + assert c._requests == {} + assert tuple(c.request_ids).count(request_id) == 1 + assert set(c.request_ids) == set(original_request_ids) + # The pool caller owns balancing the documented unconditional + # increment even when dispatch raises synchronously. + assert c.in_flight == 1 + + def test_set_keyspace_async_ambiguous_send_failure_defuncts(self): + c = self.make_connection() + original_request_ids = tuple(c.request_ids) + request_id = original_request_ids[0] + queued_frames = [] + send_error = RuntimeError("wakeup failed after enqueue") + + def close(): + with c.lock: + if c.is_closed: + return + c.is_closed = True + + def fail_after_enqueue(frame): + queued_frames.append(frame) + raise send_error + + c.close = close + c.push = fail_after_enqueue + callback = Mock() + + with pytest.raises(RuntimeError, match="wakeup failed after enqueue"): + c.set_keyspace_async("ks", callback) + + assert len(queued_frames) == 1 + assert c.is_defunct + assert c.is_closed + assert c.last_error is send_error + assert c._requests == {} + assert request_id not in c.request_ids + assert set(c.request_ids) == set(original_request_ids[1:]) + callback.assert_called_once() + + def test_final_continuous_paging_release_notifies_owner_outside_lock(self): + owner = Mock() + c = self.make_connection() + c._owning_pool = owner + c.lock = Lock() + c._continuous_paging_sessions = { + 301: Mock(), + 302: Mock(), + } + + callback_had_lock = [] + + def on_connection_released(connection): + acquired = connection.lock.acquire(False) + callback_had_lock.append(not acquired) + if acquired: + connection.lock.release() + + owner.on_connection_released.side_effect = on_connection_released + + c.remove_continuous_paging_session(301) + owner.on_connection_released.assert_not_called() + + c.remove_continuous_paging_session(302) + + owner.on_connection_released.assert_called_once_with(c) + assert callback_had_lock == [False] + assert 301 in c.request_ids + assert 302 in c.request_ids + def test_send_msg_passes_negotiated_features_to_encoder(self): """ send_msg must hand the connection's negotiated ProtocolFeatures to the @@ -383,6 +536,506 @@ def test_wait_for_responses_shutdown_includes_last_error(self): assert "already closed" in error_message assert "Bad file descriptor" in error_message + def test_factory_returns_maintenance_mode_startup_close(self): + """ + Maintenance mode accepts regular CQL sockets and closes them during + startup. The low-level factory keeps that close observable while + still tracking pool-owned startup connections for shutdown cleanup. + """ + + class MaintenanceModeCqlServer(object): + def __init__(self): + self._sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) + self._sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) + self._sock.bind(('127.0.0.1', 0)) + self._sock.listen(1) + self._sock.settimeout(2) + self.port = self._sock.getsockname()[1] + self.first_frame = b'' + self.ready = Event() + self.received_frame = Event() + self.error = None + self.thread = Thread(target=self._run) + self.thread.daemon = True + self.thread.start() + + def _run(self): + self.ready.set() + try: + client, _ = self._sock.accept() + with client: + client.settimeout(2) + while len(self.first_frame) < 9: + chunk = client.recv(9 - len(self.first_frame)) + if not chunk: + break + self.first_frame += chunk + except Exception as exc: + self.error = exc + finally: + self.received_frame.set() + + def close(self): + self._sock.close() + self.thread.join(2) + + class MaintenanceModeConnection(Connection): + def __init__(self, *args, **kwargs): + super(MaintenanceModeConnection, self).__init__(*args, **kwargs) + self._reader = None + self._connect_socket() + self._send_options_message() + self._reader = Thread(target=self._read_until_server_closes) + self._reader.daemon = True + self._reader.start() + + def push(self, data): + self._socket.sendall(data) + + def close(self): + with self.lock: + if self.is_closed: + return + self.is_closed = True + + if self._socket: + self._socket.close() + + if not self.is_defunct: + shutdown_error = ConnectionShutdown("Connection to %s was closed" % self.endpoint) + self.error_all_requests(shutdown_error) + if not self.connected_event.is_set(): + self.last_error = shutdown_error + self.connected_event.set() + + def _read_until_server_closes(self): + try: + while True: + data = self._socket.recv(self.in_buffer_size) + if not data: + self.close() + return + self._iobuf.write(data) + self.process_io_buffer() + except socket.error as exc: + if not self.is_closed: + self.defunct(exc) + + class PendingConnections(list): + def __init__(self): + super(PendingConnections, self).__init__() + self.appended = [] + + def append(self, conn): + self.appended.append(conn) + super(PendingConnections, self).append(conn) + + server = MaintenanceModeCqlServer() + try: + assert server.ready.wait(2) + conn = MaintenanceModeConnection.factory( + DefaultEndPoint('127.0.0.1', server.port), timeout=2) + + assert conn.is_closed + assert server.received_frame.wait(2) + assert server.error is None + assert len(server.first_frame) >= 5 + assert server.first_frame[4] == 0x05 # OPTIONS + finally: + server.close() + + server = MaintenanceModeCqlServer() + try: + assert server.ready.wait(2) + host_conn = Mock() + host_conn.is_shutdown = False + host_conn._pending_connections = PendingConnections() + + conn = MaintenanceModeConnection.factory( + DefaultEndPoint('127.0.0.1', server.port), + timeout=2, + host_conn=host_conn) + + assert conn.is_closed + assert server.received_frame.wait(2) + assert server.error is None + assert len(server.first_frame) >= 5 + assert server.first_frame[4] == 0x05 # OPTIONS + assert host_conn._pending_connections == [] + assert len(host_conn._pending_connections.appended) == 1 + assert host_conn._pending_connections.appended[0] is conn + finally: + server.close() + + def test_factory_raises_clean_close_after_ready(self): + class ReadyThenCleanCloseConnection(Connection): + def __init__(self, *args, **kwargs): + super(ReadyThenCleanCloseConnection, self).__init__( + *args, **kwargs) + self._compressor = None + self._handle_startup_response(ReadyMessage()) + self.close() + + def close(self): + with self.lock: + if self.is_closed: + return + self.is_closed = True + self.connected_event.set() + + with pytest.raises(ConnectionShutdown) as exc_info: + ReadyThenCleanCloseConnection.factory( + DefaultEndPoint('127.0.0.1', 9042), timeout=2) + + assert "closed by server" in str(exc_info.value) + + def test_factory_marks_legacy_event_only_startup_complete(self): + class LegacyReadyConnection(Connection): + def __init__(self, *args, **kwargs): + super(LegacyReadyConnection, self).__init__(*args, **kwargs) + # Historical third-party reactors used only this event to + # publish successful startup. + self.connected_event.set() + + def close(self): + with self.lock: + self.is_closed = True + self.connected_event.set() + + connection = LegacyReadyConnection.factory( + DefaultEndPoint('127.0.0.1', 9042), timeout=2) + + assert connection._startup_completed is True + connection.close() + assert isinstance( + _startup_close_error(connection, connection.endpoint), + ConnectionShutdown) + assert not isinstance( + _startup_close_error(connection, connection.endpoint), + _ConnectionClosedDuringStartup) + + def test_factory_preserves_positional_constructor_arguments(self): + expected_option = object() + + class PositionalOptionConnection(Connection): + def __init__(self, endpoint, positional_option, *args, **kwargs): + self.positional_option = positional_option + super(PositionalOptionConnection, self).__init__( + endpoint, *args, **kwargs) + self.connected_event.set() + + def close(self): + with self.lock: + self.is_closed = True + self.connected_event.set() + + connection = PositionalOptionConnection.factory( + DefaultEndPoint('127.0.0.1', 9042), + 2, + expected_option) + + assert connection.positional_option is expected_option + connection.close() + + def test_factory_records_host_connection_as_owner(self): + class ImmediateConnection(Connection): + def __init__(self, *args, **kwargs): + super(ImmediateConnection, self).__init__(*args, **kwargs) + self.connected_event.set() + + def close(self): + with self.lock: + self.is_closed = True + self.connected_event.set() + + class Owner(object): + is_shutdown = False + + def __init__(self): + self._pending_connections = [] + + owner = Owner() + connection = ImmediateConnection.factory( + DefaultEndPoint('127.0.0.1', 9042), + timeout=2, + host_conn=owner) + + assert connection._owning_pool is owner + assert owner._pending_connections == [connection] + owner._pending_connections.remove(connection) + connection.close() + + def test_factory_preserves_legacy_ready_when_close_wins_snapshot(self): + class LegacyReadyThenClosedConnection(Connection): + instance = None + + def __init__(self, *args, **kwargs): + super(LegacyReadyThenClosedConnection, self).__init__( + *args, **kwargs) + LegacyReadyThenClosedConnection.instance = self + # Historical reactors published READY only through this event. + self.connected_event.set() + # Deterministically place the close between READY publication + # and the factory state snapshot. + self.close() + + def close(self): + with self.lock: + if self.is_closed: + return + self.is_closed = True + if not self._startup_completed: + self.last_error = ConnectionShutdown( + "legacy connection closed after READY", + self.endpoint) + self.connected_event.set() + + with pytest.raises(ConnectionShutdown) as exc_info: + LegacyReadyThenClosedConnection.factory( + DefaultEndPoint('127.0.0.1', 9042), timeout=2) + + connection = LegacyReadyThenClosedConnection.instance + assert connection._startup_completed is True + assert exc_info.value is connection.last_error + assert "closed after READY" in str(exc_info.value) + assert not isinstance( + _startup_close_error(connection, connection.endpoint), + _ConnectionClosedDuringStartup) + + def test_legacy_startup_event_publication_is_atomic_with_close(self): + c = self.make_connection() + acquire_attempted = Event() + + class ObservedLock(object): + def __init__(self): + self.inner = RLock() + + def __enter__(self): + acquire_attempted.set() + self.inner.acquire() + return self + + def __exit__(self, *args): + self.inner.release() + + observed_lock = ObservedLock() + c.lock = observed_lock + observed_lock.inner.acquire() + setter = Thread(target=c.connected_event.set) + setter.start() + try: + assert acquire_attempted.wait(1) + # Close wins while the event publisher is waiting for the same + # state lock; it must prevent publication of a legacy READY. + c.is_closed = True + finally: + observed_lock.inner.release() + setter.join(1) + + assert not setter.is_alive() + assert c.connected_event.is_set() + assert not c._startup_event_set_while_open + + def test_factory_raises_clean_close_after_auth_success(self): + class AuthSuccessThenCleanCloseConnection(Connection): + def __init__(self, *args, **kwargs): + super(AuthSuccessThenCleanCloseConnection, self).__init__( + *args, **kwargs) + self.authenticator = Mock() + self._compressor = None + self._handle_auth_response(AuthSuccessMessage(b'token')) + self.close() + + def close(self): + with self.lock: + if self.is_closed: + return + self.is_closed = True + self.connected_event.set() + + with pytest.raises(ConnectionShutdown) as exc_info: + AuthSuccessThenCleanCloseConnection.factory( + DefaultEndPoint('127.0.0.1', 9042), timeout=2) + + assert "closed by server" in str(exc_info.value) + + def test_factory_observes_defunct_error_atomically_after_ready(self): + expected_error = RuntimeError("post-READY transport failure") + setter_entered = Event() + factory_waiting_for_lock = Event() + release_setter = Event() + factory_thread_id = get_ident() + + class ObservedRLock(object): + def __init__(self): + self._lock = RLock() + + def __enter__(self): + if (get_ident() == factory_thread_id + and setter_entered.is_set()): + factory_waiting_for_lock.set() + self._lock.acquire() + return self + + def __exit__(self, *args): + self._lock.release() + + class AtomicDefunctConnection(Connection): + instance = None + + @property + def last_error(self): + return self._last_error + + @last_error.setter + def last_error(self, exc): + if self._block_error_publication: + setter_entered.set() + release_setter.wait(2) + self._last_error = exc + + def __init__(self, *args, **kwargs): + self._last_error = None + self._block_error_publication = False + super(AtomicDefunctConnection, self).__init__(*args, **kwargs) + self.lock = ObservedRLock() + self._compressor = None + self._handle_startup_response(ReadyMessage()) + self._block_error_publication = True + self.defunct_thread = Thread( + target=self.defunct, args=(expected_error,)) + self.defunct_thread.daemon = True + self.defunct_thread.start() + assert setter_entered.wait(2) + + def release_when_factory_takes_snapshot(): + factory_waiting_for_lock.wait(2) + release_setter.set() + + self.release_thread = Thread( + target=release_when_factory_takes_snapshot) + self.release_thread.daemon = True + self.release_thread.start() + AtomicDefunctConnection.instance = self + + def close(self): + with self.lock: + if self.is_closed: + return + self.is_closed = True + + try: + with pytest.raises(RuntimeError) as exc_info: + AtomicDefunctConnection.factory( + DefaultEndPoint('127.0.0.1', 9042), timeout=2) + assert exc_info.value is expected_error + assert factory_waiting_for_lock.is_set() + finally: + release_setter.set() + conn = AtomicDefunctConnection.instance + if conn is not None: + conn.defunct_thread.join(2) + conn.release_thread.join(2) + + def test_factory_returns_reactor_clean_startup_close(self): + class ReactorCleanClose(Exception): + pass + + class ReactorCleanCloseConnection(Connection): + def __init__(self, *args, **kwargs): + super(ReactorCleanCloseConnection, self).__init__(*args, **kwargs) + self.last_error = ReactorCleanClose("closed cleanly") + self.is_closed = True + self.is_defunct = True + self.connected_event.set() + + def close(self): + self.is_closed = True + + def _is_clean_close_error(self, exc): + return isinstance(exc, ReactorCleanClose) + + conn = ReactorCleanCloseConnection.factory( + DefaultEndPoint('127.0.0.1', 9042), timeout=2) + + assert conn.is_closed + assert conn.is_defunct + assert isinstance(conn.last_error, ReactorCleanClose) + + def test_factory_raises_reactor_unclean_startup_close(self): + class ReactorUncleanClose(Exception): + pass + + class ReactorUncleanCloseConnection(Connection): + def __init__(self, *args, **kwargs): + super(ReactorUncleanCloseConnection, self).__init__(*args, **kwargs) + self.last_error = ReactorUncleanClose("closed uncleanly") + self.is_closed = True + self.is_defunct = True + self.connected_event.set() + + def close(self): + self.is_closed = True + + with pytest.raises(ReactorUncleanClose): + ReactorUncleanCloseConnection.factory( + DefaultEndPoint('127.0.0.1', 9042), timeout=2) + + def test_factory_raises_defunct_connection_shutdown_startup_error(self): + class DefunctConnectionShutdownConnection(Connection): + def __init__(self, *args, **kwargs): + super(DefunctConnectionShutdownConnection, self).__init__(*args, **kwargs) + self.last_error = ConnectionShutdown("defunct startup failure") + self.is_closed = True + self.is_defunct = True + self.connected_event.set() + + def close(self): + self.is_closed = True + + with pytest.raises(ConnectionShutdown): + DefunctConnectionShutdownConnection.factory( + DefaultEndPoint('127.0.0.1', 9042), timeout=2) + + def test_factory_closes_socket_when_wait_is_cancelled(self): + class FactoryCancelled(BaseException): + pass + + class InterruptingEvent(object): + def wait(self, timeout): + raise FactoryCancelled() + + class InterruptedConnection(Connection): + instance = None + + def __init__(self, *args, **kwargs): + super(InterruptedConnection, self).__init__(*args, **kwargs) + self.connected_event = InterruptingEvent() + self.close_count = 0 + InterruptedConnection.instance = self + + def close(self): + self.close_count += 1 + self.is_closed = True + + with pytest.raises(FactoryCancelled): + InterruptedConnection.factory( + DefaultEndPoint('127.0.0.1', 9042), timeout=2) + + assert InterruptedConnection.instance.is_closed + assert InterruptedConnection.instance.close_count == 1 + + def test_owner_does_not_reclassify_post_startup_close(self): + connection = self.make_connection() + connection._startup_completed = True + connection.is_closed = True + + error = _startup_close_error(connection) + + assert isinstance(error, ConnectionShutdown) + assert not isinstance(error, _ConnectionClosedDuringStartup) + assert "after startup" in str(error) + @patch('cassandra.connection.ConnectionHeartbeat._raise_if_stopped') class ConnectionHeartbeatTest(unittest.TestCase): @@ -436,8 +1089,19 @@ def send_msg(msg, req_id, msg_callback): holder = get_holders.return_value[0] holder.get_connections.return_value.append(idle_connection) holder.get_connections.return_value.append(non_idle_connection) + callback_had_lock = [] - self.run_heartbeat(get_holders) + def on_connection_released(connection): + acquired = connection.lock.acquire(False) + callback_had_lock.append(not acquired) + if acquired: + connection.lock.release() + raise RuntimeError("retirement bookkeeping failed") + + holder.on_connection_released.side_effect = on_connection_released + + with patch('cassandra.connection.log.exception') as log_exception: + self.run_heartbeat(get_holders) holder.get_connections.assert_has_calls([call()] * get_holders.call_count) assert idle_connection.in_flight == 0 @@ -445,6 +1109,12 @@ def send_msg(msg, req_id, msg_callback): idle_connection.send_msg.assert_has_calls([call(ANY, request_id, ANY)] * get_holders.call_count) assert non_idle_connection.send_msg.call_count == 0 + holder.on_connection_released.assert_has_calls( + [call(idle_connection)] * get_holders.call_count) + assert callback_had_lock == [False] * get_holders.call_count + idle_connection.defunct.assert_not_called() + holder.return_connection.assert_not_called() + assert log_exception.call_count == get_holders.call_count def test_closed_defunct(self, *args): get_holders = self.make_get_holders(1) diff --git a/tests/unit/test_control_connection.py b/tests/unit/test_control_connection.py index fd62323f33..b205ce1b03 100644 --- a/tests/unit/test_control_connection.py +++ b/tests/unit/test_control_connection.py @@ -15,13 +15,19 @@ import unittest from concurrent.futures import ThreadPoolExecutor +from threading import Event, RLock, Thread 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 ControlConnection, _Scheduler, ProfileManager, EXEC_PROFILE_DEFAULT, ExecutionProfile +from cassandra.cluster import (Cluster, ControlConnection, _Scheduler, + ProfileManager, EXEC_PROFILE_DEFAULT, + ExecutionProfile, + _ControlReconnectionHandler, + _HOST_TRANSITION_DEFERRED) from cassandra.pool import Host -from cassandra.connection import EndPoint, DefaultEndPoint, DefaultEndPointFactory +from cassandra.connection import (ConnectionShutdown, EndPoint, DefaultEndPoint, + DefaultEndPointFactory) from cassandra.policies import (SimpleConvictionPolicy, RoundRobinPolicy, ConstantReconnectionPolicy, IdentityTranslator) @@ -206,6 +212,269 @@ def setUp(self): self.control_connection._connection = self.connection self.control_connection._time = self.time + def test_try_connect_rejects_connection_closed_during_startup(self): + endpoint = DefaultEndPoint("192.168.1.0") + closed_connection = Mock() + closed_connection.endpoint = endpoint + closed_connection.is_closed = True + self.cluster.connection_factory = Mock(return_value=closed_connection) + + with self.assertRaises(ConnectionShutdown) as exc_info: + self.control_connection._try_connect(endpoint) + + assert "closed during the startup handshake" in str(exc_info.exception) + assert self.control_connection._connection is self.connection + closed_connection.register_watchers.assert_not_called() + + def test_set_new_connection_resumes_down_host_after_reconnect(self): + host = self.cluster.metadata.get_host(DefaultEndPoint("192.168.1.0")) + host.is_up = False + self.connection.close = Mock() + new_connection = Mock() + new_connection.endpoint = host.endpoint + + self.control_connection._set_new_connection(new_connection) + + assert self.control_connection._connection is new_connection + self.connection.close.assert_called_once_with() + self.cluster.scheduler.schedule_unique.assert_called_once_with( + 0, self.cluster.on_up, host) + + def test_set_new_connection_survives_resume_scheduling_failure(self): + host = self.cluster.metadata.get_host( + DefaultEndPoint("192.168.1.0")) + host.is_up = False + self.cluster.on_up = Mock() + self.cluster.scheduler.schedule_unique.side_effect = RuntimeError( + "scheduler rejected recovery") + self.connection.close = Mock() + new_connection = Mock() + new_connection.endpoint = host.endpoint + + self.control_connection._set_new_connection(new_connection) + + assert self.control_connection._connection is new_connection + new_connection.close.assert_not_called() + self.connection.close.assert_called_once_with() + self.cluster.on_up.assert_called_once_with(host) + + def test_control_reconnection_handler_keeps_adopted_socket_open(self): + self.connection.close = Mock() + new_connection = Mock() + new_connection.endpoint = DefaultEndPoint("192.168.1.0") + self.control_connection._reconnect_internal = Mock( + return_value=new_connection) + completed = Mock() + handler = _ControlReconnectionHandler( + self.control_connection, + self.cluster.scheduler, + iter(()), + completed) + + handler.run() + + assert self.control_connection._connection is new_connection + self.connection.close.assert_called_once_with() + new_connection.close.assert_not_called() + completed.assert_called_once_with() + + def test_control_reconnection_rejects_candidate_after_shutdown(self): + new_connection = Mock() + new_connection.endpoint = DefaultEndPoint("192.168.1.0") + self.control_connection._reconnect_internal = Mock( + return_value=new_connection) + completed = Mock() + handler = _ControlReconnectionHandler( + self.control_connection, + self.cluster.scheduler, + iter(()), + completed) + self.control_connection._is_shutdown = True + + handler.run() + + assert self.control_connection._connection is self.connection + new_connection.close.assert_called_once_with() + completed.assert_called_once_with() + + def test_connect_returns_when_shutdown_rejects_initial_candidate(self): + candidate = Mock() + candidate.endpoint = DefaultEndPoint("192.168.1.0") + original_dbaas = object() + self.cluster.metadata.dbaas = original_dbaas + self.cluster.protocol_version = 4 + self.control_connection._connection = None + + def reconnect_during_shutdown(): + self.control_connection._is_shutdown = True + return candidate + + self.control_connection._reconnect_internal = Mock( + side_effect=reconnect_during_shutdown) + + self.control_connection.connect() + + assert self.control_connection._connection is None + candidate.close.assert_called_once_with() + assert self.cluster.metadata.dbaas is original_dbaas + + def test_try_connect_closes_candidate_on_setup_base_exception(self): + class SetupCancelled(BaseException): + pass + + endpoint = DefaultEndPoint("192.168.1.0") + candidate = Mock(is_closed=False) + candidate.features.sharding_info = None + candidate.features.tablets_routing_v1 = False + candidate.register_watchers.side_effect = SetupCancelled( + "cancelled watcher registration") + self.cluster.connection_factory = Mock(return_value=candidate) + self.cluster.metadata_request_timeout = 0 + self.cluster._client_routes_handler = None + + with self.assertRaises(SetupCancelled): + self.control_connection._try_connect(endpoint) + + candidate.close.assert_called_once_with() + + def test_set_new_connection_uses_verified_host_for_aliased_endpoint(self): + host = self.cluster.metadata.get_host(DefaultEndPoint("192.168.1.0")) + host.is_up = False + self.connection.close = Mock() + new_connection = Mock() + new_connection.endpoint = DefaultEndPoint("shared-proxy") + new_connection._control_connection_host = host + + self.control_connection._set_new_connection(new_connection) + + assert self.cluster.metadata.get_host(new_connection.endpoint) is None + self.cluster.scheduler.schedule_unique.assert_called_once_with( + 0, self.cluster.on_up, host) + + def test_set_new_connection_preserves_parked_addition_semantics(self): + host = self.cluster.metadata.get_host(DefaultEndPoint("192.168.1.0")) + host.is_up = False + parked_reconnector = Mock(is_host_addition=True) + host._reconnection_handler = parked_reconnector + self.cluster.on_add = Mock() + self.connection.close = Mock() + new_connection = Mock() + new_connection.endpoint = DefaultEndPoint("shared-proxy") + new_connection._control_connection_host = host + + self.control_connection._set_new_connection(new_connection) + + self.cluster.scheduler.schedule_unique.assert_called_once_with( + 0, self.cluster.on_add, host) + parked_reconnector.cancel.assert_called_once_with() + assert host._reconnection_handler is None + + def test_set_initial_connection_does_not_resume_down_host(self): + host = self.cluster.metadata.get_host(DefaultEndPoint("192.168.1.0")) + host.is_up = False + self.control_connection._connection = None + new_connection = Mock() + new_connection.endpoint = host.endpoint + + self.control_connection._set_new_connection(new_connection) + + assert self.control_connection._connection is new_connection + self.cluster.scheduler.schedule_unique.assert_not_called() + + def test_default_hooks_stale_control_error_cannot_down_new_connection(self): + host = self.cluster.metadata.get_host(DefaultEndPoint("192.168.1.0")) + self.cluster._lock = RLock() + self.cluster._on_down_locked = Mock() + self.cluster._uses_default_failure_hooks = lambda: True + policy_entered = Event() + release_policy = Event() + thread_errors = [] + + def blocking_conviction(_): + policy_entered.set() + assert release_policy.wait(2) + return True + + host.signal_connection_failure = blocking_conviction + + old_connection = Mock( + endpoint=host.endpoint, + is_defunct=True, + last_error=ConnectionShutdown("old connection failed")) + old_connection._control_connection_host = host + self.control_connection._connection = old_connection + + new_connection = Mock( + endpoint=host.endpoint, + is_defunct=False) + new_connection._control_connection_host = host + + def signal_error(): + try: + self.control_connection._signal_error() + except BaseException as exc: + thread_errors.append(exc) + + signal_thread = Thread(target=signal_error) + signal_thread.start() + assert policy_entered.wait(2) + + replace_thread = Thread( + target=self.control_connection._set_new_connection, + args=(new_connection,)) + replace_thread.start() + replace_thread.join(2) + assert not replace_thread.is_alive() + + release_policy.set() + signal_thread.join(2) + + assert not signal_thread.is_alive() + assert thread_errors == [] + assert self.control_connection._connection is new_connection + assert host.is_up + self.cluster._on_down_locked.assert_not_called() + + def test_control_error_preserves_legacy_cluster_on_down_override(self): + class LegacyCluster(Cluster): + + def on_down( + self, host, is_host_addition, + expect_host_to_be_down=False): + self.down_calls.append(( + host, + is_host_addition, + expect_host_to_be_down)) + + cluster = object.__new__(LegacyCluster) + cluster._lock = RLock() + cluster.metadata = MockMetadata() + cluster.down_calls = [] + control_connection = ControlConnection(cluster, 1, 0, 0, 0) + host = cluster.metadata.get_host(DefaultEndPoint("192.168.1.0")) + connection = Mock( + endpoint=host.endpoint, + is_defunct=True, + last_error=ConnectionShutdown("control failed")) + connection._control_connection_host = host + control_connection._connection = connection + + control_connection._signal_error() + + assert cluster.down_calls == [ + (host, False, False), + ] + + def test_topology_refresh_retains_system_local_host_identity(self): + self.connection.endpoint = DefaultEndPoint("shared-proxy") + + self.control_connection._refresh_node_list_and_token_map( + self.connection, + preloaded_results=self._matching_schema_preloaded_results) + + assert self.connection._control_connection_host is \ + self.cluster.metadata.get_host_by_host_id('uuid1') + def test_wait_for_schema_agreement(self): """ Basic test with all schema versions agreeing @@ -389,6 +658,34 @@ def test_change_ip(self): assert 3 == len(self.cluster.metadata.all_hosts()) + def test_deferred_ip_change_reschedules_before_using_old_endpoint(self): + host = self.cluster.metadata.get_host_by_host_id('uuid2') + old_endpoint = host.endpoint + new_endpoint = DefaultEndPoint("192.168.1.5") + self.cluster._force_down_for_endpoint_change = Mock( + return_value=_HOST_TRANSITION_DEFERRED) + del self.connection.peer_results[:] + self.connection.peer_results.extend([ + [ + "rpc_address", "peer", "schema_version", "data_center", + "rack", "tokens", "host_id"], + [[ + new_endpoint.address, "10.0.0.5", "a", "dc1", "rack1", + ["2", "102", "202"], 'uuid2']]]) + preloaded_results = _node_meta_results( + self.connection.local_results, + self.connection.peer_results) + + self.control_connection._refresh_node_list_and_token_map( + self.connection, + preloaded_results=preloaded_results) + + assert host.endpoint == old_endpoint + self.cluster.scheduler.schedule_unique.assert_called_once_with( + 0, + self.control_connection.refresh_node_list_and_token_map, + force_token_rebuild=True) + def test_refresh_nodes_and_tokens_uses_preloaded_results_if_given(self): """ diff --git a/tests/unit/test_host_connection_pool.py b/tests/unit/test_host_connection_pool.py index 8bb57d0dc0..ea9ba785bc 100644 --- a/tests/unit/test_host_connection_pool.py +++ b/tests/unit/test_host_connection_pool.py @@ -11,20 +11,22 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. -from concurrent.futures import ThreadPoolExecutor +from concurrent.futures import Future, ThreadPoolExecutor import logging import uuid +from cassandra import InvalidRequest from cassandra.protocol_features import ProtocolFeatures from cassandra.shard_info import _ShardingInfo import unittest -from threading import Thread, Event, Lock +from threading import Thread, Event, Lock, RLock from unittest.mock import Mock, NonCallableMagicMock, MagicMock from cassandra.cluster import Session, ShardAwareOptions -from cassandra.connection import Connection -from cassandra.pool import HostConnection +from cassandra.connection import ( + Connection, ConnectionException, ConnectionShutdown) +from cassandra.pool import HostConnection, _signal_connection_failure from cassandra.pool import Host, NoConnectionsAvailable from cassandra.policies import HostDistance, SimpleConvictionPolicy import pytest @@ -41,21 +43,29 @@ class _PoolTests(unittest.TestCase): def make_session(self): session = NonCallableMagicMock(spec=Session, keyspace='foobarkeyspace', _trash=[]) + session.cluster.connect_timeout = 5 return session def test_borrow_and_return(self): host = Mock(spec=Host, address='ip1') session = self.make_session() - conn = HashableMock(spec=Connection, in_flight=0, is_defunct=False, is_closed=False, max_request_id=100) + conn = HashableMock( + spec=Connection, in_flight=0, is_defunct=False, is_closed=False, + max_request_id=100, orphaned_threshold_reached=False) session.cluster.connection_factory.return_value = conn pool = self.PoolImpl(host, HostDistance.LOCAL, session) - session.cluster.connection_factory.assert_called_once_with(host.endpoint, on_orphaned_stream_released=pool.on_orphaned_stream_released) + session.cluster.connection_factory.assert_called_once_with( + host.endpoint, + host_conn=pool, + on_orphaned_stream_released=pool.on_orphaned_stream_released) c, request_id = pool.borrow_connection(timeout=0.01) assert c is conn assert 1 == conn.in_flight - conn.set_keyspace_blocking.assert_called_once_with('foobarkeyspace') + conn.set_keyspace_blocking.assert_called_once_with( + 'foobarkeyspace', + timeout=session.cluster.connect_timeout) pool.return_connection(conn) assert 0 == conn.in_flight @@ -65,11 +75,16 @@ def test_borrow_and_return(self): def test_failed_wait_for_connection(self): host = Mock(spec=Host, address='ip1') session = self.make_session() - conn = HashableMock(spec=Connection, in_flight=0, is_defunct=False, is_closed=False, max_request_id=100) + conn = HashableMock( + spec=Connection, in_flight=0, is_defunct=False, is_closed=False, + max_request_id=100, orphaned_threshold_reached=False) session.cluster.connection_factory.return_value = conn pool = self.PoolImpl(host, HostDistance.LOCAL, session) - session.cluster.connection_factory.assert_called_once_with(host.endpoint, on_orphaned_stream_released=pool.on_orphaned_stream_released) + session.cluster.connection_factory.assert_called_once_with( + host.endpoint, + host_conn=pool, + on_orphaned_stream_released=pool.on_orphaned_stream_released) pool.borrow_connection(timeout=0.01) assert 1 == conn.in_flight @@ -81,15 +96,47 @@ def test_failed_wait_for_connection(self): with pytest.raises(NoConnectionsAvailable): pool.borrow_connection(0) + def test_orphan_threshold_connection_remains_borrowable_during_repair( + self): + host = Mock(spec=Host, address='ip1') + session = self.make_session() + connection = HashableMock( + spec=Connection, + in_flight=0, + is_defunct=False, + is_closed=False, + max_request_id=100, + orphaned_threshold_reached=True, + orphaned_request_ids=set(), + lock=RLock()) + connection.features = ProtocolFeatures(shard_id=0) + session.cluster.connection_factory.return_value = connection + session.submit.return_value = Future() + pool = self.PoolImpl(host, HostDistance.LOCAL, session) + session.submit.reset_mock() + + borrowed, request_id = pool.borrow_connection(0.1) + + assert borrowed is connection + assert connection.in_flight == 1 + if self.uses_single_connection: + session.submit.assert_called_once_with( + pool._replace, connection) + def test_successful_wait_for_connection(self): host = Mock(spec=Host, address='ip1') session = self.make_session() - conn = HashableMock(spec=Connection, in_flight=0, is_defunct=False, is_closed=False, max_request_id=100, - lock=Lock()) + conn = HashableMock( + spec=Connection, in_flight=0, is_defunct=False, is_closed=False, + max_request_id=100, lock=Lock(), + orphaned_threshold_reached=False) session.cluster.connection_factory.return_value = conn pool = self.PoolImpl(host, HostDistance.LOCAL, session) - session.cluster.connection_factory.assert_called_once_with(host.endpoint, on_orphaned_stream_released=pool.on_orphaned_stream_released) + session.cluster.connection_factory.assert_called_once_with( + host.endpoint, + host_conn=pool, + on_orphaned_stream_released=pool.on_orphaned_stream_released) pool.borrow_connection(timeout=0.01) assert 1 == conn.in_flight @@ -109,12 +156,17 @@ def get_second_conn(): def test_spawn_when_at_max(self): host = Mock(spec=Host, address='ip1') session = self.make_session() - conn = HashableMock(spec=Connection, in_flight=0, is_defunct=False, is_closed=False, max_request_id=100) + conn = HashableMock( + spec=Connection, in_flight=0, is_defunct=False, is_closed=False, + max_request_id=100, orphaned_threshold_reached=False) conn.max_request_id = 100 session.cluster.connection_factory.return_value = conn pool = self.PoolImpl(host, HostDistance.LOCAL, session) - session.cluster.connection_factory.assert_called_once_with(host.endpoint, on_orphaned_stream_released=pool.on_orphaned_stream_released) + session.cluster.connection_factory.assert_called_once_with( + host.endpoint, + host_conn=pool, + on_orphaned_stream_released=pool.on_orphaned_stream_released) pool.borrow_connection(timeout=0.01) assert 1 == conn.in_flight @@ -133,12 +185,17 @@ def test_spawn_when_at_max(self): def test_return_defunct_connection(self): host = Mock(spec=Host, address='ip1') session = self.make_session() - conn = HashableMock(spec=Connection, in_flight=0, is_defunct=False, is_closed=False, - max_request_id=100, signaled_error=False) + conn = HashableMock( + spec=Connection, in_flight=0, is_defunct=False, is_closed=False, + max_request_id=100, signaled_error=False, + orphaned_threshold_reached=False) session.cluster.connection_factory.return_value = conn pool = self.PoolImpl(host, HostDistance.LOCAL, session) - session.cluster.connection_factory.assert_called_once_with(host.endpoint, on_orphaned_stream_released=pool.on_orphaned_stream_released) + session.cluster.connection_factory.assert_called_once_with( + host.endpoint, + host_conn=pool, + on_orphaned_stream_released=pool.on_orphaned_stream_released) pool.borrow_connection(timeout=0.01) conn.is_defunct = True @@ -160,7 +217,10 @@ def test_return_defunct_connection_on_down_host(self): session.cluster.shard_aware_options = ShardAwareOptions() pool = self.PoolImpl(host, HostDistance.LOCAL, session) - session.cluster.connection_factory.assert_called_once_with(host.endpoint, on_orphaned_stream_released=pool.on_orphaned_stream_released) + session.cluster.connection_factory.assert_called_once_with( + host.endpoint, + host_conn=pool, + on_orphaned_stream_released=pool.on_orphaned_stream_released) pool.borrow_connection(timeout=0.01) conn.is_defunct = True @@ -182,12 +242,15 @@ def test_return_defunct_connection_on_down_host(self): def test_return_closed_connection(self): host = Mock(spec=Host, address='ip1') session = self.make_session() - conn = HashableMock(spec=Connection, in_flight=0, is_defunct=False, is_closed=True, max_request_id=100, + conn = HashableMock(spec=Connection, in_flight=0, is_defunct=False, is_closed=False, max_request_id=100, signaled_error=False, orphaned_threshold_reached=False) session.cluster.connection_factory.return_value = conn pool = self.PoolImpl(host, HostDistance.LOCAL, session) - session.cluster.connection_factory.assert_called_once_with(host.endpoint, on_orphaned_stream_released=pool.on_orphaned_stream_released) + session.cluster.connection_factory.assert_called_once_with( + host.endpoint, + host_conn=pool, + on_orphaned_stream_released=pool.on_orphaned_stream_released) pool.borrow_connection(timeout=0.01) conn.is_closed = True @@ -230,6 +293,1924 @@ class HostConnectionTests(_PoolTests): PoolImpl = HostConnection uses_single_connection = True + def test_get_connections_takes_pool_lock(self): + pool = object.__new__(HostConnection) + pool._lock = Lock() + connection = object() + pool._connections = {0: connection} + started = Event() + finished = Event() + result = [] + + def get_connections(): + started.set() + result.extend(pool.get_connections()) + finished.set() + + pool._lock.acquire() + thread = Thread(target=get_connections) + thread.start() + try: + assert started.wait(2) + assert not finished.wait(0.05) + finally: + pool._lock.release() + thread.join(2) + + assert not thread.is_alive() + assert result == [connection] + + def make_connection( + self, shard_id=0, sharding_info=None, in_flight=0, + orphaned_threshold_reached=False): + connection = HashableMock( + spec=Connection, + in_flight=in_flight, + is_defunct=False, + is_closed=False, + max_request_id=100, + signaled_error=False, + orphaned_threshold_reached=orphaned_threshold_reached, + orphaned_request_ids=set(), + lock=RLock()) + connection.features = ProtocolFeatures( + shard_id=shard_id, + sharding_info=sharding_info) + return connection + + def make_pool(self, shard_aware=False, shards_count=2): + host = Mock(spec=Host, address='ip1') + host.lock = RLock() + host.sharding_info = None + session = self.make_session() + session.cluster.shard_aware_options = ShardAwareOptions(disable=True) + first_connection = self.make_connection() + session.cluster.connection_factory.return_value = first_connection + pool = self.PoolImpl(host, HostDistance.LOCAL, session) + session.cluster.connection_factory.reset_mock() + session.submit.reset_mock() + + if shard_aware: + host.sharding_info = _ShardingInfo( + shard_id=0, + shards_count=shards_count, + partitioner="", + sharding_algorithm="", + sharding_ignore_msb=0, + shard_aware_port="", + shard_aware_port_ssl="") + session.cluster.shard_aware_options = ShardAwareOptions( + disable=False, + disable_shardaware_port=True) + + return pool, host, session, first_connection + + @staticmethod + def factory_returning_registered(candidate): + """ + Model Connection.factory's successful ownership handoff. + + The factory leaves the candidate registered with HostConnection until + the pool either adopts it or shutdown drains it. + """ + def factory(endpoint, host_conn=None, **kwargs): + assert host_conn._register_pending_connection(candidate) + return candidate + + return factory + + @staticmethod + def factory_returning_registered_with_adoption_barrier(candidate): + """ + Pause the caller's adoption immediately after the factory returns. + + This exposes the exact handoff interval to a concurrent shutdown. + """ + factory_returned = Event() + adoption_started = Event() + release_adoption = Event() + captured_pools = [] + + def factory(endpoint, host_conn=None, **kwargs): + original_register = host_conn._register_pending_connection + assert original_register(candidate) + captured_pools.append(host_conn) + + def block_adoption(connection): + assert connection is candidate + assert factory_returned.is_set() + adoption_started.set() + assert release_adoption.wait(2) + return original_register(connection) + + host_conn._register_pending_connection = block_adoption + factory_returned.set() + return candidate + + return ( + factory, + adoption_started, + release_adoption, + captured_pools, + ) + + def test_initial_factory_registered_connection_is_adopted_once(self): + host = Mock(spec=Host, address='ip1') + host.lock = RLock() + host.sharding_info = None + session = self.make_session() + session.cluster.shard_aware_options = ShardAwareOptions(disable=True) + connection = self.make_connection() + session.cluster.connection_factory.side_effect = \ + self.factory_returning_registered(connection) + + pool = self.PoolImpl(host, HostDistance.LOCAL, session) + + assert pool._connections == {0: connection} + assert pool._pending_connections == [] + connection.close.assert_not_called() + + def test_replacement_factory_registered_connection_is_adopted_once(self): + pool, host, session, old_connection = self.make_pool() + replacement = self.make_connection() + session.cluster.connection_factory.side_effect = \ + self.factory_returning_registered(replacement) + pool._is_replacing = True + + pool._replace(old_connection) + + assert pool._connections == {0: replacement} + assert pool._pending_connections == [] + replacement.close.assert_not_called() + old_connection.close.assert_called_once_with() + + def test_missing_shard_factory_registered_connection_is_adopted_once(self): + pool, host, session, first_connection = self.make_pool( + shard_aware=True) + connection = self.make_connection(shard_id=1) + session.cluster.connection_factory.side_effect = \ + self.factory_returning_registered(connection) + + pool._open_connection_to_missing_shard(1) + + assert pool._connections == { + 0: first_connection, + 1: connection, + } + assert pool._pending_connections == [] + connection.close.assert_not_called() + + def test_shutdown_owns_registered_initial_connection_before_adoption( + self): + host = Mock(spec=Host, address='ip1') + host.lock = RLock() + host.sharding_info = None + session = self.make_session() + session.cluster.shard_aware_options = ShardAwareOptions(disable=True) + connection = self.make_connection() + connection_closed = Event() + + def close_connection(): + connection.is_closed = True + connection_closed.set() + + connection.close.side_effect = close_connection + ( + factory, + adoption_started, + release_adoption, + captured_pools, + ) = self.factory_returning_registered_with_adoption_barrier(connection) + session.cluster.connection_factory.side_effect = factory + constructed_pools = [] + construction_errors = [] + + def construct_pool(): + try: + constructed_pools.append( + self.PoolImpl(host, HostDistance.LOCAL, session)) + except BaseException as exc: + construction_errors.append(exc) + + constructor = Thread(target=construct_pool) + constructor.start() + try: + assert adoption_started.wait(2) + captured_pools[0].shutdown() + assert connection_closed.wait(2) + finally: + release_adoption.set() + constructor.join(2) + + assert not constructor.is_alive() + assert constructed_pools == [] + assert len(construction_errors) == 1 + assert isinstance(construction_errors[0], ConnectionException) + assert captured_pools[0]._connections == {} + assert captured_pools[0]._pending_connections == [] + connection.close.assert_called_once_with() + + def test_shutdown_owns_registered_replacement_before_adoption(self): + pool, host, session, old_connection = self.make_pool() + replacement = self.make_connection() + replacement_closed = Event() + + def close_replacement(): + replacement.is_closed = True + replacement_closed.set() + + replacement.close.side_effect = close_replacement + ( + factory, + adoption_started, + release_adoption, + captured_pools, + ) = self.factory_returning_registered_with_adoption_barrier( + replacement) + session.cluster.connection_factory.side_effect = factory + pool._is_replacing = True + replacement_errors = [] + + def replace(): + try: + pool._replace(old_connection) + except BaseException as exc: + replacement_errors.append(exc) + + replacement_thread = Thread(target=replace) + replacement_thread.start() + try: + assert adoption_started.wait(2) + assert captured_pools == [pool] + pool.shutdown() + assert replacement_closed.wait(2) + finally: + release_adoption.set() + replacement_thread.join(2) + + assert not replacement_thread.is_alive() + assert replacement_errors == [] + assert pool._connections == {} + assert pool._pending_connections == [] + replacement.close.assert_called_once_with() + + def test_shutdown_owns_registered_missing_shard_before_adoption(self): + pool, host, session, first_connection = self.make_pool( + shard_aware=True) + connection = self.make_connection(shard_id=1) + connection_closed = Event() + + def close_connection(): + connection.is_closed = True + connection_closed.set() + + connection.close.side_effect = close_connection + ( + factory, + adoption_started, + release_adoption, + captured_pools, + ) = self.factory_returning_registered_with_adoption_barrier(connection) + session.cluster.connection_factory.side_effect = factory + opening_errors = [] + + def open_missing_shard(): + try: + pool._open_connection_to_missing_shard(1) + except BaseException as exc: + opening_errors.append(exc) + + opening_thread = Thread(target=open_missing_shard) + opening_thread.start() + try: + assert adoption_started.wait(2) + assert captured_pools == [pool] + pool.shutdown() + assert connection_closed.wait(2) + finally: + release_adoption.set() + opening_thread.join(2) + + assert not opening_thread.is_alive() + assert opening_errors == [] + assert pool._connections == {} + assert pool._pending_connections == [] + connection.close.assert_called_once_with() + + def test_initial_closed_connection_is_rejected(self): + host = Mock(spec=Host, address='ip1') + session = self.make_session() + connection = HashableMock( + spec=Connection, + is_closed=True, + is_defunct=False, + in_flight=0, + max_request_id=100) + session.cluster.connection_factory.return_value = connection + + with pytest.raises(ConnectionShutdown) as exc_info: + self.PoolImpl(host, HostDistance.LOCAL, session) + + assert "closed during the startup handshake" in str(exc_info.value) + connection.set_keyspace_blocking.assert_not_called() + + def test_replace_tracks_pending_connection(self): + host = Mock(spec=Host, address='ip1') + host.sharding_info = None + session = self.make_session() + first_conn = HashableMock(spec=Connection, in_flight=0, is_defunct=False, + is_closed=False, max_request_id=100) + first_conn.features = ProtocolFeatures(shard_id=0) + replacement_conn = HashableMock(spec=Connection, in_flight=0, is_defunct=False, + is_closed=False, max_request_id=100) + replacement_conn.features = ProtocolFeatures(shard_id=0) + session.cluster.connection_factory.side_effect = [first_conn, replacement_conn] + + pool = self.PoolImpl(host, HostDistance.LOCAL, session) + session.cluster.connection_factory.reset_mock() + pool._is_replacing = True + + pool._replace(first_conn) + + session.cluster.connection_factory.assert_called_once_with( + host.endpoint, + host_conn=pool, + on_orphaned_stream_released=pool.on_orphaned_stream_released) + assert pool._pending_connections == [] + assert pool._connections[0] is replacement_conn + assert not pool._is_replacing + + def test_regular_replacement_keeps_old_mapping_until_candidate_is_ready( + self): + pool, host, session, old_connection = self.make_pool() + old_connection.orphaned_threshold_reached = True + replacement = self.make_connection() + keyspace_started = Event() + release_keyspace = Event() + replacement_errors = [] + + def set_keyspace(keyspace, timeout=None): + keyspace_started.set() + assert release_keyspace.wait(2) + + replacement.set_keyspace_blocking.side_effect = set_keyspace + session.cluster.connection_factory.return_value = replacement + pool._is_replacing = True + + def replace(): + try: + pool._replace(old_connection) + except BaseException as exc: + replacement_errors.append(exc) + + replace_thread = Thread(target=replace) + replace_thread.start() + try: + assert keyspace_started.wait(2) + assert pool._connections[0] is old_connection + + borrowed, _ = pool.borrow_connection(timeout=0) + assert borrowed is old_connection + pool.return_connection(borrowed) + finally: + release_keyspace.set() + replace_thread.join(2) + + assert not replace_thread.is_alive() + assert replacement_errors == [] + assert pool._connections[0] is replacement + old_connection.close.assert_called_once_with() + + def test_replace_closed_connection_hands_off_to_host_reconnector(self): + host = Mock(spec=Host, address='ip1') + host.sharding_info = None + session = self.make_session() + first_conn = HashableMock( + spec=Connection, + in_flight=0, + is_defunct=False, + is_closed=False, + max_request_id=100) + first_conn.features = ProtocolFeatures(shard_id=0) + replacement_conn = HashableMock( + spec=Connection, + in_flight=0, + is_defunct=False, + is_closed=True, + max_request_id=100) + session.cluster.connection_factory.side_effect = [ + first_conn, replacement_conn] + session.cluster.signal_connection_failure.return_value = True + + pool = self.PoolImpl(host, HostDistance.LOCAL, session) + session.cluster.connection_factory.reset_mock() + session.submit.reset_mock() + pool._is_replacing = True + + pool._replace(first_conn) + + replacement_conn.close.assert_called_once_with() + replacement_conn.set_keyspace_blocking.assert_not_called() + assert pool._pending_connections == [] + assert pool._connections == {0: first_conn} + first_conn.close.assert_not_called() + assert pool._is_replacing + session.submit.assert_not_called() + session.cluster.signal_connection_failure.assert_called_once() + failure_args = session.cluster.signal_connection_failure.call_args + assert failure_args.args[0] is host + assert isinstance(failure_args.args[1], ConnectionShutdown) + assert failure_args.args[2] is False + assert failure_args.kwargs == { + 'expect_host_to_be_down': True, + 'force': True, + } + + def test_replace_startup_close_during_shutdown_does_not_mark_host_down(self): + host = Mock(spec=Host, address='ip1') + host.sharding_info = None + session = self.make_session() + first_conn = HashableMock( + spec=Connection, + in_flight=0, + is_defunct=False, + is_closed=False, + max_request_id=100) + first_conn.features = ProtocolFeatures(shard_id=0) + replacement_conn = HashableMock( + spec=Connection, + in_flight=0, + is_defunct=False, + is_closed=True, + max_request_id=100) + session.cluster.connection_factory.return_value = first_conn + + pool = self.PoolImpl(host, HostDistance.LOCAL, session) + session.cluster.connection_factory.reset_mock() + pool._is_replacing = True + + def close_pool_during_factory(*args, **kwargs): + pool.is_shutdown = True + return replacement_conn + + session.cluster.connection_factory.side_effect = close_pool_during_factory + pool._replace(first_conn) + + replacement_conn.close.assert_called_once_with() + session.cluster.signal_connection_failure.assert_not_called() + session.cluster.scheduler.schedule.assert_not_called() + + def test_replace_allows_shutdown_to_close_pending_connection(self): + host = Mock(spec=Host, address='ip1') + host.sharding_info = None + session = self.make_session() + first_conn = HashableMock(spec=Connection, in_flight=0, is_defunct=False, + is_closed=False, max_request_id=100) + first_conn.features = ProtocolFeatures(shard_id=0) + session.cluster.connection_factory.return_value = first_conn + + pool = self.PoolImpl(host, HostDistance.LOCAL, session) + session.cluster.connection_factory.reset_mock() + pool._is_replacing = True + + factory_entered = Event() + release_factory = Event() + pending_closed = Event() + + pending_conn = HashableMock(spec=Connection, in_flight=0, is_defunct=False, + is_closed=False, max_request_id=100) + pending_conn.features = ProtocolFeatures(shard_id=0) + + def close_pending(): + pending_conn.is_closed = True + pending_closed.set() + + pending_conn.close.side_effect = close_pending + + replacement_conn = HashableMock(spec=Connection, in_flight=0, is_defunct=False, + is_closed=False, max_request_id=100) + replacement_conn.features = ProtocolFeatures(shard_id=0) + + def blocking_factory(endpoint, host_conn=None, **kwargs): + host_conn._pending_connections.append(pending_conn) + factory_entered.set() + try: + release_factory.wait(2) + return replacement_conn + finally: + try: + host_conn._pending_connections.remove(pending_conn) + except ValueError: + pass + + session.cluster.connection_factory.side_effect = blocking_factory + + replace_thread = Thread(target=pool._replace, args=(first_conn,)) + replace_thread.start() + assert factory_entered.wait(2) + + shutdown_thread = Thread(target=pool.shutdown) + shutdown_thread.start() + try: + assert pending_closed.wait(2) + finally: + release_factory.set() + replace_thread.join(2) + shutdown_thread.join(2) + + assert not replace_thread.is_alive() + assert not shutdown_thread.is_alive() + assert pool.is_shutdown + replacement_conn.close.assert_called_once() + + def test_replace_allows_shutdown_during_keyspace_setup(self): + host = Mock(spec=Host, address='ip1') + host.sharding_info = None + session = self.make_session() + first_conn = HashableMock(spec=Connection, in_flight=0, is_defunct=False, + is_closed=False, max_request_id=100) + first_conn.features = ProtocolFeatures(shard_id=0) + replacement_conn = HashableMock(spec=Connection, in_flight=0, is_defunct=False, + is_closed=False, max_request_id=100) + replacement_conn.features = ProtocolFeatures(shard_id=0) + session.cluster.connection_factory.side_effect = [first_conn, replacement_conn] + + pool = self.PoolImpl(host, HostDistance.LOCAL, session) + session.cluster.connection_factory.reset_mock() + pool._is_replacing = True + + keyspace_started = Event() + release_keyspace = Event() + replacement_closed = Event() + + def set_keyspace_blocking(keyspace, timeout=None): + keyspace_started.set() + release_keyspace.wait(2) + + def close_replacement(): + replacement_conn.is_closed = True + replacement_closed.set() + release_keyspace.set() + + replacement_conn.set_keyspace_blocking.side_effect = set_keyspace_blocking + replacement_conn.close.side_effect = close_replacement + + replace_thread = Thread(target=pool._replace, args=(first_conn,)) + replace_thread.start() + try: + assert keyspace_started.wait(2) + assert replacement_conn in pool._pending_connections + + pool.shutdown() + + assert replacement_closed.is_set() + finally: + release_keyspace.set() + replace_thread.join(2) + + assert not replace_thread.is_alive() + assert pool.is_shutdown + assert pool._pending_connections == [] + assert pool._connections == {} + replacement_conn.close.assert_called_once_with() + + def test_initial_keyspace_failure_closes_unpublished_connection(self): + host = Mock(spec=Host, address='ip1') + session = self.make_session() + connection = self.make_connection() + connection.set_keyspace_blocking.side_effect = RuntimeError( + "keyspace setup failed") + session.cluster.connection_factory.return_value = connection + + with pytest.raises(RuntimeError, match="keyspace setup failed"): + self.PoolImpl(host, HostDistance.LOCAL, session) + + connection.close.assert_called_once_with() + + def test_initial_keyspace_cancellation_closes_unpublished_connection(self): + class SetupCancelled(BaseException): + pass + + host = Mock(spec=Host, address='ip1') + session = self.make_session() + connection = self.make_connection() + connection.set_keyspace_blocking.side_effect = SetupCancelled() + session.cluster.connection_factory.return_value = connection + + with pytest.raises(SetupCancelled): + self.PoolImpl(host, HostDistance.LOCAL, session) + + connection.close.assert_called_once_with() + + def test_initial_stale_keyspace_validation_retries_current_generation(self): + host = Mock(spec=Host, address='ip1') + session = self.make_session() + session._lock = RLock() + session._keyspace_generation = 0 + session.keyspace = "old_keyspace" + connection = self.make_connection() + seen_keyspaces = [] + + def set_keyspace(keyspace, timeout=None): + seen_keyspaces.append((keyspace, timeout)) + if keyspace == "old_keyspace": + with session._lock: + session.keyspace = "new_keyspace" + session._keyspace_generation += 1 + raise InvalidRequest("old keyspace was dropped") + + connection.set_keyspace_blocking.side_effect = set_keyspace + session.cluster.connection_factory.return_value = connection + + pool = self.PoolImpl(host, HostDistance.LOCAL, session) + + assert seen_keyspaces == [ + ("old_keyspace", session.cluster.connect_timeout), + ("new_keyspace", session.cluster.connect_timeout), + ] + assert pool._keyspace == "new_keyspace" + assert pool._connections[0] is connection + connection.close.assert_not_called() + + def test_initial_current_keyspace_validation_still_fails(self): + host = Mock(spec=Host, address='ip1') + session = self.make_session() + session._lock = RLock() + session._keyspace_generation = 0 + connection = self.make_connection() + connection.set_keyspace_blocking.side_effect = InvalidRequest( + "current keyspace is invalid") + session.cluster.connection_factory.return_value = connection + + with pytest.raises( + InvalidRequest, match="current keyspace is invalid"): + self.PoolImpl(host, HostDistance.LOCAL, session) + + connection.close.assert_called_once_with() + + def test_failure_adapter_preserves_pre_force_override_signature(self): + calls = [] + + class LegacyCluster(object): + def signal_connection_failure( + self, host, exc, adding, expected_down=False, **kwargs): + calls.append( + (host, exc, adding, expected_down, kwargs)) + return "legacy-result" + + host = object() + error = ConnectionShutdown("startup closed") + + result = _signal_connection_failure( + LegacyCluster(), + host, + error, + is_host_addition=False, + expect_host_to_be_down=True, + force=True) + + assert result == "legacy-result" + assert calls == [( + host, + error, + False, + True, + {'force': True}, + )] + + def test_failure_adapter_supports_keyword_only_override(self): + calls = [] + + class KeywordOnlyCluster(object): + def signal_connection_failure( + self, host, exc, *, is_host_addition, + expect_host_to_be_down=False, force=False): + calls.append(( + host, + exc, + is_host_addition, + expect_host_to_be_down, + force)) + + host = object() + error = ConnectionShutdown("startup closed") + _signal_connection_failure( + KeywordOnlyCluster(), + host, + error, + is_host_addition=False, + expect_host_to_be_down=True, + force=True) + + assert calls == [(host, error, False, True, True)] + + def test_replacement_handoff_forces_recovery_after_legacy_false(self): + pool, host, session, connection = self.make_pool() + calls = [] + + class LegacyCluster(object): + def __init__(self): + self._lock = RLock() + + def _uses_default_failure_hooks(self): + return False + + def signal_connection_failure( + self, failed_host, exc, adding, expected_down=False): + calls.append( + ("public", failed_host, exc, adding, expected_down)) + return False + + def _on_down_locked( + self, failed_host, is_host_addition, + expect_host_to_be_down=False, force=False): + calls.append(( + "private", + failed_host, + is_host_addition, + expect_host_to_be_down, + force)) + + cluster = LegacyCluster() + session.cluster = cluster + session._lock = RLock() + session.is_shutdown = False + session._pools = {host: pool} + error = ConnectionShutdown("replacement closed") + + assert pool._handoff_replacement_failure(error) + + assert calls == [ + ("public", host, error, False, True), + ("private", host, False, True, True), + ] + + def test_replacement_handoff_does_not_duplicate_delegated_force(self): + pool, host, session, connection = self.make_pool() + calls = [] + host._recovery_epoch = 3 + + class DelegatingCluster(object): + def __init__(self): + self._lock = RLock() + + def _uses_default_failure_hooks(self): + return False + + def signal_connection_failure( + self, failed_host, exc, is_host_addition, + expect_host_to_be_down=False, force=False): + calls.append(("public", force)) + self._on_down_locked( + failed_host, + is_host_addition, + expect_host_to_be_down, + force) + return False + + def _on_down_locked( + self, failed_host, is_host_addition, + expect_host_to_be_down=False, force=False): + calls.append(("private", force)) + failed_host._recovery_epoch += 1 + + session.cluster = DelegatingCluster() + session._lock = RLock() + session.is_shutdown = False + session._pools = {host: pool} + + assert pool._handoff_replacement_failure( + ConnectionShutdown("replacement closed")) + + assert calls == [ + ("public", True), + ("private", True), + ] + assert host._recovery_epoch == 4 + + def test_replacement_handoff_survives_conviction_base_exception(self): + class PolicyCancelled(BaseException): + pass + + pool, host, session, connection = self.make_pool() + calls = [] + + class DefaultClusterStandin(object): + def __init__(self): + self._lock = RLock() + + def _uses_default_failure_hooks(self): + return True + + def _on_down_locked( + self, failed_host, is_host_addition, + expect_host_to_be_down=False, force=False): + calls.append(( + failed_host, + is_host_addition, + expect_host_to_be_down, + force)) + + session.cluster = DefaultClusterStandin() + session._lock = RLock() + session.is_shutdown = False + session._pools = {host: pool} + host.signal_connection_failure.side_effect = PolicyCancelled( + "cancelled policy") + + assert pool._handoff_replacement_failure( + ConnectionShutdown("replacement closed")) + + assert calls == [(host, False, True, True)] + + def test_replace_invalid_keyspace_quarantines_open_connection(self): + pool, host, session, first_connection = self.make_pool() + replacement = self.make_connection() + replacement.set_keyspace_blocking.side_effect = InvalidRequest( + "keyspace was dropped") + session.cluster.connection_factory.return_value = replacement + pool._is_replacing = True + + pool._replace(first_connection) + + assert pool._connections[0] is replacement + assert pool._pending_connections == [] + assert not pool._is_replacing + assert replacement._pool_keyspace_mismatch is True + replacement.close.assert_not_called() + session.cluster.signal_connection_failure.assert_not_called() + session.submit.reset_mock() + session.cluster.connection_factory.reset_mock() + for _ in range(2): + with pytest.raises(NoConnectionsAvailable): + pool.borrow_connection(timeout=0) + session.submit.assert_not_called() + session.cluster.connection_factory.assert_not_called() + borrowed, _ = pool.borrow_connection( + timeout=0, + allow_keyspace_mismatch=True) + assert borrowed is replacement + pool.return_connection(borrowed) + + def test_quarantined_connection_counts_as_open(self): + pool, host, session, connection = self.make_pool() + connection._pool_keyspace_mismatch = True + + assert pool.open_count == 1 + assert pool.get_state()['open_count'] == 1 + + def test_replace_revalidates_keyspace_before_adoption(self): + pool, host, session, first_connection = self.make_pool() + replacement = self.make_connection() + seen_keyspaces = [] + first_connection.set_keyspace_async.side_effect = ( + lambda keyspace, callback: + callback(first_connection, None)) + + def set_keyspace(keyspace, timeout=None): + seen_keyspaces.append((keyspace, timeout)) + if len(seen_keyspaces) == 1: + keyspace_set = Mock() + pool._set_keyspace_for_all_conns("new_keyspace", keyspace_set) + keyspace_set.assert_called_once_with(pool, []) + + replacement.set_keyspace_blocking.side_effect = set_keyspace + session.cluster.connection_factory.return_value = replacement + pool._is_replacing = True + + pool._replace(first_connection) + + assert seen_keyspaces == [ + ("foobarkeyspace", session.cluster.connect_timeout), + ("new_keyspace", session.cluster.connect_timeout), + ] + assert pool._connections[0] is replacement + assert pool._keyspace == "new_keyspace" + + def test_regular_replacement_failure_uses_host_recovery_backoff(self): + pool, host, session, first_connection = self.make_pool() + replacement = self.make_connection() + replacement.set_keyspace_blocking.side_effect = RuntimeError( + "socket failed") + session.cluster.connection_factory.return_value = replacement + pool._is_replacing = True + + pool._replace(first_connection) + + replacement.close.assert_called_once_with() + assert pool._connections == {0: first_connection} + first_connection.close.assert_not_called() + assert pool._pending_connections == [] + assert pool._is_replacing + session.submit.assert_not_called() + session.cluster.signal_connection_failure.assert_called_once() + + def test_regular_replacement_cancellation_cleans_socket_and_marker(self): + class SetupCancelled(BaseException): + pass + + pool, host, session, first_connection = self.make_pool() + replacement = self.make_connection() + replacement.set_keyspace_blocking.side_effect = SetupCancelled() + session.cluster.connection_factory.return_value = replacement + pool._is_replacing = True + + with pytest.raises(SetupCancelled): + pool._replace(first_connection) + + replacement.close.assert_called_once_with() + assert pool._pending_connections == [] + assert not pool._is_replacing + session.cluster.signal_connection_failure.assert_called_once() + + def test_regular_replacement_submit_cancellation_hands_off_recovery(self): + class SubmitCancelled(BaseException): + pass + + pool, host, session, connection = self.make_pool() + cancellation = SubmitCancelled("cancelled") + with pool._lock: + del pool._connections[connection.features.shard_id] + pool._is_replacing = True + session.submit.side_effect = cancellation + + with pytest.raises(SubmitCancelled): + pool._submit_regular_replacement(connection, claimed=True) + + assert not pool._is_replacing + session.cluster.signal_connection_failure.assert_called_once() + + def test_cancelled_regular_replacement_future_hands_off_recovery(self): + pool, host, session, connection = self.make_pool() + replacement_future = Future() + session.submit.return_value = replacement_future + + assert pool._submit_regular_replacement(connection) is \ + replacement_future + assert pool._is_replacing + + assert replacement_future.cancel() + + assert not pool._is_replacing + session.cluster.signal_connection_failure.assert_called_once() + + def test_clean_shard_replacement_hands_off_to_host_recovery(self): + pool, host, session, first_connection = self.make_pool( + shard_aware=True) + submitted = Future() + session.submit.return_value = submitted + pool._replace(first_connection) + attempt = pool._shard_connection_attempts[0] + + closed_replacement = self.make_connection() + closed_replacement.is_closed = True + session.cluster.connection_factory.return_value = closed_replacement + + pool._open_connection_to_missing_shard(0, attempt['token']) + + assert pool._connections == {} + assert 0 not in pool._connecting + assert 0 not in pool._shard_connection_attempts + closed_replacement.close.assert_called_once_with() + session.cluster.signal_connection_failure.assert_called_once() + failure = session.cluster.signal_connection_failure.call_args + assert failure.kwargs['force'] is True + assert failure.kwargs['expect_host_to_be_down'] is True + + def test_cancelled_shard_replacement_future_hands_off_recovery(self): + pool, host, session, first_connection = self.make_pool( + shard_aware=True) + replacement_future = Future() + session.submit.return_value = replacement_future + + pool._replace(first_connection) + assert 0 in pool._shard_connection_attempts + + assert replacement_future.cancel() + + assert 0 not in pool._shard_connection_attempts + assert 0 not in pool._connecting + session.cluster.signal_connection_failure.assert_called_once() + + def test_interrupted_shard_replacement_with_live_shard_hands_off_recovery( + self): + class SetupCancelled(BaseException): + pass + + pool, host, session, first_connection = self.make_pool( + shard_aware=True) + live_connection = self.make_connection(shard_id=1) + pool._connections[1] = live_connection + replacement_future = Future() + session.submit.return_value = replacement_future + pool._replace(first_connection) + attempt = pool._shard_connection_attempts[0] + candidate = self.make_connection(shard_id=0) + candidate.set_keyspace_blocking.side_effect = SetupCancelled( + "cancelled replacement setup") + session.cluster.connection_factory.return_value = candidate + + with pytest.raises(SetupCancelled): + pool._open_connection_to_missing_shard( + 0, + attempt['token']) + + assert pool._connections == {1: live_connection} + assert 0 not in pool._connecting + assert 0 not in pool._shard_connection_attempts + candidate.close.assert_called_once_with() + session.cluster.signal_connection_failure.assert_called_once() + + def test_missing_shard_setup_failure_closes_unadopted_socket(self): + pool, host, session, first_connection = self.make_pool( + shard_aware=True) + connection = self.make_connection(shard_id=1) + connection.set_keyspace_blocking.side_effect = RuntimeError( + "USE failed") + session.cluster.connection_factory.return_value = connection + pool._connections.pop(0) + + with pytest.raises(RuntimeError, match="USE failed"): + pool._open_connection_to_missing_shard(1) + + connection.close.assert_called_once_with() + assert pool._connections == {} + assert pool._pending_connections == [] + assert 1 not in pool._connecting + session.cluster.signal_connection_failure.assert_not_called() + + def test_missing_shard_setup_cancellation_cleans_socket_and_marker(self): + class SetupCancelled(BaseException): + pass + + pool, host, session, first_connection = self.make_pool( + shard_aware=True) + connection = self.make_connection(shard_id=1) + connection.set_keyspace_blocking.side_effect = SetupCancelled() + session.cluster.connection_factory.return_value = connection + + with pytest.raises(SetupCancelled): + pool._open_connection_to_missing_shard(1) + + connection.close.assert_called_once_with() + assert pool._pending_connections == [] + assert 1 not in pool._connecting + assert 1 not in pool._shard_connection_attempts + + def test_shard_submit_cancellation_clears_attempt_marker(self): + class SubmitCancelled(BaseException): + pass + + pool, host, session, first_connection = self.make_pool( + shard_aware=True) + session.submit.side_effect = SubmitCancelled() + + with pytest.raises(SubmitCancelled): + pool._schedule_connection_to_missing_shard(1) + + assert 1 not in pool._connecting + assert 1 not in pool._shard_connection_attempts + + def test_missing_shard_invalid_keyspace_quarantines_open_socket(self): + pool, host, session, first_connection = self.make_pool( + shard_aware=True) + connection = self.make_connection(shard_id=1) + connection.set_keyspace_blocking.side_effect = InvalidRequest( + "keyspace was dropped") + session.cluster.connection_factory.return_value = connection + + pool._open_connection_to_missing_shard(1) + + assert pool._connections[1] is connection + assert connection._pool_keyspace_mismatch is True + assert 1 not in pool._connecting + assert 1 not in pool._shard_connection_attempts + connection.close.assert_not_called() + session.cluster.signal_connection_failure.assert_not_called() + pool._connections.pop(0) + session.submit.reset_mock() + session.cluster.connection_factory.reset_mock() + for _ in range(2): + with pytest.raises(NoConnectionsAvailable): + pool.borrow_connection(timeout=0) + session.submit.assert_not_called() + session.cluster.connection_factory.assert_not_called() + borrowed, _ = pool.borrow_connection( + timeout=0, + allow_keyspace_mismatch=True) + assert borrowed is connection + pool.return_connection(borrowed) + pool.shutdown() + connection.close.assert_called_once_with() + + def test_shard_attempt_tokens_are_identity_safe(self): + pool, host, session, first_connection = self.make_pool( + shard_aware=True) + session.submit.side_effect = lambda *args, **kwargs: Future() + + pool._schedule_connection_to_missing_shard(1) + first_token = pool._shard_connection_attempts[1]['token'] + pool._finish_shard_connection_attempt(1, first_token) + pool._schedule_connection_to_missing_shard(1) + second_token = pool._shard_connection_attempts[1]['token'] + + assert second_token is not first_token + pool._finish_shard_connection_attempt(1, first_token) + assert pool._shard_connection_attempts[1]['token'] is second_token + assert 1 in pool._connecting + + pool._finish_shard_connection_attempt(1, second_token) + assert 1 not in pool._connecting + + def test_shard_failure_and_promotion_are_linearizable(self): + pool, host, session, first_connection = self.make_pool( + shard_aware=True) + session.submit.side_effect = lambda *args, **kwargs: Future() + + # Promotion wins the lock first: the failure remains published for + # forced host recovery. + pool._schedule_connection_to_missing_shard(1) + promoted_token = pool._shard_connection_attempts[1]['token'] + pool._schedule_connection_to_missing_shard( + 1, + is_replacement=True) + assert pool._claim_failed_shard_attempt(1, promoted_token) + assert pool._shard_connection_attempts[1]['token'] is promoted_token + pool._finish_shard_connection_attempt(1, promoted_token) + + # Failure wins first: it clears only its token, and the later + # replacement receives a distinct attempt. + pool._schedule_connection_to_missing_shard(1) + failed_token = pool._shard_connection_attempts[1]['token'] + assert not pool._claim_failed_shard_attempt(1, failed_token) + pool._schedule_connection_to_missing_shard( + 1, + is_replacement=True) + replacement_token = pool._shard_connection_attempts[1]['token'] + assert replacement_token is not failed_token + assert pool._shard_connection_attempts[1]['is_replacement'] + + def test_ordinary_shard_failure_clears_marker_for_later_retry(self): + pool, host, session, first_connection = self.make_pool( + shard_aware=True) + session.submit.side_effect = lambda *args, **kwargs: Future() + session.cluster.connection_factory.side_effect = RuntimeError( + "temporary failure") + pool._schedule_connection_to_missing_shard(1) + token = pool._shard_connection_attempts[1]['token'] + + with pytest.raises(RuntimeError, match="temporary failure"): + pool._open_connection_to_missing_shard(1, token) + + assert 1 not in pool._connecting + assert 1 not in pool._shard_connection_attempts + session.cluster.signal_connection_failure.assert_not_called() + + session.cluster.connection_factory.side_effect = None + session.cluster.connection_factory.return_value = self.make_connection( + shard_id=1) + pool._schedule_connection_to_missing_shard(1) + assert 1 in pool._connecting + + def test_failed_shards_schedule_independent_replacements(self): + pool, host, session, first_connection = self.make_pool( + shard_aware=True) + second_connection = self.make_connection(shard_id=1, in_flight=1) + pool._connections[1] = second_connection + first_connection.in_flight = 1 + session.submit.side_effect = lambda *args, **kwargs: Future() + host.signal_connection_failure.return_value = False + + first_connection.is_closed = True + second_connection.is_closed = True + pool.return_connection(first_connection) + pool.return_connection(second_connection) + + assert set(pool._shard_connection_attempts) == {0, 1} + assert pool._connecting == {0, 1} + assert session.submit.call_count == 2 + + def test_stale_mapping_removal_preserves_new_connection(self): + pool, host, session, old_connection = self.make_pool( + shard_aware=True) + new_connection = self.make_connection(shard_id=0) + pool._connections[0] = new_connection + + pool._replace(old_connection) + assert pool._connections[0] is new_connection + session.cluster.connection_factory.assert_not_called() + + old_connection.in_flight = 1 + old_connection.is_closed = True + old_connection.signaled_error = True + pool._trash.add(old_connection) + pool.return_connection(old_connection) + + assert pool._connections[0] is new_connection + assert old_connection not in pool._trash + session.submit.assert_not_called() + + def test_mapped_closed_shard_schedules_repair_and_uses_fallback(self): + pool, host, session, closed_connection = self.make_pool( + shard_aware=True) + fallback = self.make_connection(shard_id=1) + pool._connections[1] = fallback + closed_connection.is_closed = True + host.signal_connection_failure.return_value = False + session.submit.side_effect = lambda *args, **kwargs: Future() + token = Mock(value=123) + session.cluster.metadata.token_map.token_class.from_key.return_value = ( + token) + host.sharding_info.shard_id_from_token = Mock(return_value=0) + + selected = pool._get_connection_for_routing_key(b"routing-key") + + assert selected is fallback + attempt = pool._shard_connection_attempts[0] + assert attempt['is_replacement'] + assert pool.num_missing_or_needing_replacement == 1 + + def test_routed_healthy_shard_does_not_scan_all_connections(self): + pool, host, session, target = self.make_pool( + shard_aware=True, + shards_count=64) + + class GetOnlyConnections(dict): + + def values(self): + raise AssertionError( + "healthy routed selection scanned every connection") + + pool._connections = GetOnlyConnections({0: target}) + token = Mock(value=123) + session.cluster.metadata.token_map.token_class.from_key.return_value = ( + token) + host.sharding_info.shard_id_from_token = Mock(return_value=0) + + selected = pool._get_connection_for_routing_key(b"routing-key") + + assert selected is target + session.submit.assert_not_called() + + def test_synchronous_regular_replacement_is_selected_immediately(self): + pool, host, session, old_connection = self.make_pool() + old_connection.orphaned_threshold_reached = True + replacement = self.make_connection(shard_id=1) + session.cluster.connection_factory.return_value = replacement + session.submit.side_effect = ( + lambda function, *args, **kwargs: + function(*args, **kwargs)) + + selected = pool._get_connection_for_routing_key() + + assert selected is replacement + assert pool._connections == {1: replacement} + old_connection.close.assert_called_once_with() + + def test_synchronous_optional_shard_failure_uses_healthy_fallback(self): + pool, host, session, fallback = self.make_pool(shard_aware=True) + token = Mock(value=123) + session.cluster.metadata.token_map.token_class.from_key.return_value = ( + token) + host.sharding_info.shard_id_from_token = Mock(return_value=1) + session.cluster.connection_factory.side_effect = RuntimeError( + "optional shard failed") + session.submit.side_effect = ( + lambda function, *args, **kwargs: + function(*args, **kwargs)) + + selected = pool._get_connection_for_routing_key(b"routing-key") + + assert selected is fallback + assert 1 not in pool._connecting + session.cluster.signal_connection_failure.assert_not_called() + + def test_last_return_before_retirement_still_closes_old_connection(self): + pool, host, session, old_connection = self.make_pool( + shard_aware=True) + old_connection.in_flight = 1 + old_connection.orphaned_threshold_reached = True + + # The final request returns just before replacement adoption. At this + # point the old connection is not yet in trash. + pool.return_connection(old_connection) + assert old_connection.in_flight == 0 + assert old_connection not in pool._trash + + replacement = self.make_connection(shard_id=0) + session.cluster.connection_factory.return_value = replacement + pool._open_connection_to_missing_shard(0) + + assert pool._connections[0] is replacement + old_connection.close.assert_called_once_with() + assert old_connection not in pool._trash + + def test_retired_connection_closes_on_last_return_without_evicting_new(self): + pool, host, session, old_connection = self.make_pool( + shard_aware=True) + old_connection.in_flight = 1 + old_connection.orphaned_threshold_reached = True + replacement = self.make_connection(shard_id=0) + session.cluster.connection_factory.return_value = replacement + + pool._open_connection_to_missing_shard(0) + + assert pool._connections[0] is replacement + assert old_connection in pool._trash + old_connection.close.assert_not_called() + + pool.return_connection(old_connection) + + assert pool._connections[0] is replacement + assert old_connection not in pool._trash + old_connection.close.assert_called_once_with() + + def test_retirement_waits_for_active_continuous_paging_session(self): + pool, host, session, old_connection = self.make_pool( + shard_aware=True) + old_connection.orphaned_threshold_reached = True + old_connection._continuous_paging_sessions = {7: Mock()} + replacement = self.make_connection(shard_id=0) + session.cluster.connection_factory.return_value = replacement + + pool._open_connection_to_missing_shard(0) + + assert pool._connections[0] is replacement + assert old_connection in pool._trash + old_connection.close.assert_not_called() + + old_connection._continuous_paging_sessions.clear() + pool.on_connection_released(old_connection) + + assert old_connection not in pool._trash + old_connection.close.assert_called_once_with() + + def test_direct_release_closes_idle_retired_connection(self): + pool, host, session, old_connection = self.make_pool( + shard_aware=True) + old_connection.in_flight = 1 + old_connection.orphaned_threshold_reached = True + replacement = self.make_connection(shard_id=0) + session.cluster.connection_factory.return_value = replacement + + pool._open_connection_to_missing_shard(0) + + assert old_connection in pool._trash + old_connection.close.assert_not_called() + + with old_connection.lock: + old_connection.in_flight -= 1 + pool.on_connection_released(old_connection) + + assert old_connection not in pool._trash + old_connection.close.assert_called_once_with() + + def test_defunct_retired_connection_does_not_convict_healthy_replacement( + self): + pool, host, session, old_connection = self.make_pool( + shard_aware=True) + old_connection.in_flight = 1 + old_connection.orphaned_threshold_reached = True + replacement = self.make_connection(shard_id=0) + session.cluster.connection_factory.return_value = replacement + + pool._open_connection_to_missing_shard(0) + old_connection.is_defunct = True + old_connection.is_closed = False + pool.return_connection(old_connection) + + assert pool._connections[0] is replacement + assert old_connection not in pool._trash + old_connection.close.assert_called_once_with() + host.signal_connection_failure.assert_not_called() + session.cluster.on_down.assert_not_called() + + def test_return_racing_deliberate_retirement_does_not_convict_host(self): + pool, host, session, connection = self.make_pool() + connection.in_flight = 1 + decremented = Event() + release_return = Event() + thread_errors = [] + + class PauseAfterDecrementLock(object): + + def __init__(self): + self._lock = RLock() + self._exits = 0 + + def __enter__(self): + self._lock.acquire() + return self + + def __exit__(self, *args): + self._lock.release() + self._exits += 1 + if self._exits == 1: + decremented.set() + assert release_return.wait(2) + + connection.lock = PauseAfterDecrementLock() + + def return_connection(): + try: + pool.return_connection(connection) + except BaseException as exc: + thread_errors.append(exc) + + return_thread = Thread(target=return_connection) + return_thread.start() + assert decremented.wait(2) + + with pool._lock: + del pool._connections[connection.features.shard_id] + close_connection, remaining = ( + pool._retire_connection_locked(connection)) + assert close_connection + assert remaining == 0 + connection.is_closed = True + connection.close() + + release_return.set() + return_thread.join(2) + + assert not return_thread.is_alive() + assert thread_errors == [] + host.signal_connection_failure.assert_not_called() + session.cluster.on_down.assert_not_called() + session.cluster.signal_connection_failure.assert_not_called() + + def test_blocked_conviction_cannot_down_healthy_replacement(self): + pool, host, session, connection = self.make_pool() + connection.in_flight = 1 + connection.is_closed = True + conviction_entered = Event() + release_conviction = Event() + thread_errors = [] + + def blocking_conviction(_): + conviction_entered.set() + assert release_conviction.wait(2) + return True + + host.signal_connection_failure.side_effect = blocking_conviction + + def return_connection(): + try: + pool.return_connection(connection) + except BaseException as exc: + thread_errors.append(exc) + + return_thread = Thread(target=return_connection) + return_thread.start() + assert conviction_entered.wait(2) + + # In non-shard-aware mode Scylla may assign the replacement a + # different shard id. + replacement = self.make_connection(shard_id=1) + with pool._lock: + del pool._connections[0] + pool._connections[1] = replacement + + release_conviction.set() + return_thread.join(2) + + assert not return_thread.is_alive() + assert thread_errors == [] + assert not pool.is_shutdown + assert pool._connections[1] is replacement + replacement.close.assert_not_called() + session.cluster.on_down.assert_not_called() + + def test_conviction_base_exception_releases_signal_claim(self): + class PolicyCancelled(BaseException): + pass + + pool, host, session, connection = self.make_pool() + connection.in_flight = 1 + connection.is_closed = True + connection.last_error = ConnectionShutdown("closed") + host.signal_connection_failure.side_effect = PolicyCancelled( + "cancelled policy") + + with pytest.raises(PolicyCancelled): + pool.return_connection(connection) + + assert not connection.signaled_error + host.signal_connection_failure.side_effect = None + host.signal_connection_failure.return_value = False + pool.return_connection(connection, stream_was_orphaned=True) + assert host.signal_connection_failure.call_count == 2 + + def test_shutdown_closes_adopted_and_retired_connections_once(self): + pool, host, session, old_connection = self.make_pool( + shard_aware=True) + old_connection.in_flight = 1 + old_connection.orphaned_threshold_reached = True + replacement = self.make_connection(shard_id=0) + session.cluster.connection_factory.return_value = replacement + + pool._open_connection_to_missing_shard(0) + assert old_connection in pool._trash + + pool.shutdown() + + replacement.close.assert_called_once_with() + old_connection.close.assert_called_once_with() + assert pool._connections == {} + assert pool._trash == set() + + def test_missing_shard_keyspace_setup_does_not_block_shutdown(self): + pool, host, session, first_connection = self.make_pool( + shard_aware=True) + connection = self.make_connection(shard_id=1) + keyspace_started = Event() + release_keyspace = Event() + connection_closed = Event() + + def set_keyspace(keyspace, timeout=None): + keyspace_started.set() + release_keyspace.wait(2) + + def close_connection(): + connection.is_closed = True + connection_closed.set() + release_keyspace.set() + + connection.set_keyspace_blocking.side_effect = set_keyspace + connection.close.side_effect = close_connection + session.cluster.connection_factory.return_value = connection + + open_thread = Thread( + target=pool._open_connection_to_missing_shard, + args=(1,)) + open_thread.start() + assert keyspace_started.wait(2) + assert connection in pool._pending_connections + + shutdown_thread = Thread(target=pool.shutdown) + shutdown_thread.start() + try: + assert connection_closed.wait(2) + finally: + release_keyspace.set() + open_thread.join(2) + shutdown_thread.join(2) + + assert not open_thread.is_alive() + assert not shutdown_thread.is_alive() + connection.close.assert_called_once_with() + assert pool._pending_connections == [] + + def test_missing_shard_revalidates_keyspace_before_adoption(self): + pool, host, session, first_connection = self.make_pool( + shard_aware=True) + connection = self.make_connection(shard_id=1) + seen_keyspaces = [] + + def set_keyspace(keyspace, timeout=None): + seen_keyspaces.append((keyspace, timeout)) + if len(seen_keyspaces) == 1: + keyspace_set = Mock() + pool._set_keyspace_for_all_conns("new_keyspace", keyspace_set) + callbacks = [ + call.args[1] + for call in first_connection.set_keyspace_async.call_args_list + ] + assert len(callbacks) == 1 + first_connection.in_flight = 1 + callbacks[0](first_connection, None) + keyspace_set.assert_called_once_with(pool, []) + + connection.set_keyspace_blocking.side_effect = set_keyspace + session.cluster.connection_factory.return_value = connection + + pool._open_connection_to_missing_shard(1) + + assert seen_keyspaces == [ + ("foobarkeyspace", session.cluster.connect_timeout), + ("new_keyspace", session.cluster.connect_timeout), + ] + assert pool._connections[1] is connection + + def test_empty_pool_keyspace_update_changes_pool_state(self): + pool, host, session, first_connection = self.make_pool() + pool._connections.clear() + callback = Mock() + previous_generation = pool._keyspace_generation + + pool._set_keyspace_for_all_conns("new_keyspace", callback) + + assert pool._keyspace == "new_keyspace" + assert pool._keyspace_generation == previous_generation + 1 + callback.assert_called_once_with(pool, []) + + def test_keyspace_completion_callback_is_exactly_once(self): + pool, host, session, first_connection = self.make_pool() + second_connection = self.make_connection(shard_id=1) + pool._connections[1] = second_connection + callbacks = {} + + def capture_callback(connection): + def set_keyspace_async(keyspace, callback): + connection.in_flight += 1 + callbacks[connection] = callback + connection.set_keyspace_async.side_effect = set_keyspace_async + + capture_callback(first_connection) + capture_callback(second_connection) + finished = Mock() + pool._set_keyspace_for_all_conns("new_keyspace", finished) + + callbacks[first_connection](first_connection, None) + callbacks[first_connection](first_connection, RuntimeError("late")) + callbacks[second_connection](second_connection, None) + callbacks[second_connection](second_connection, None) + + finished.assert_called_once_with(pool, []) + assert first_connection.in_flight == 0 + assert second_connection.in_flight == 0 + assert first_connection._pool_keyspace_mismatch is False + + def test_late_keyspace_callback_cannot_clear_newer_quarantine(self): + pool, host, session, connection = self.make_pool() + callbacks = [] + + def capture_callback(keyspace, callback): + connection.in_flight += 1 + callbacks.append(callback) + + connection.set_keyspace_async.side_effect = capture_callback + first_finished = Mock() + second_finished = Mock() + + pool._set_keyspace_for_all_conns("first_keyspace", first_finished) + callbacks[0](connection, None) + pool._set_keyspace_for_all_conns("second_keyspace", second_finished) + validation_error = InvalidRequest("schema has not converged") + callbacks[1](connection, validation_error) + + assert connection._pool_keyspace_mismatch is True + callbacks[0](connection, None) + + assert connection._pool_keyspace_mismatch is True + first_finished.assert_called_once_with(pool, []) + second_finished.assert_called_once_with( + pool, + [validation_error]) + + def test_keyspace_completion_survives_sync_and_return_errors(self): + pool, host, session, connection = self.make_pool() + connection.in_flight = 1 + sync_error = RuntimeError("send failed") + connection.set_keyspace_async.side_effect = sync_error + pool.return_connection = Mock(side_effect=RuntimeError("return failed")) + finished = Mock() + + pool._set_keyspace_for_all_conns("new_keyspace", finished) + + finished.assert_called_once() + errors = finished.call_args.args[1] + assert errors == [sync_error] + assert connection._pool_keyspace_mismatch is True + pool.return_connection.assert_called_once_with(connection) + + def test_keyspace_completion_survives_sync_base_exception(self): + class SetupCancelled(BaseException): + pass + + pool, host, session, connection = self.make_pool() + connection.in_flight = 1 + cancellation = SetupCancelled("cancelled") + connection.set_keyspace_async.side_effect = cancellation + finished = Mock() + + pool._set_keyspace_for_all_conns("new_keyspace", finished) + + finished.assert_called_once_with(pool, [cancellation]) + assert connection.in_flight == 0 + assert pool._keyspace_update_current is None + assert not pool._keyspace_update_in_progress + + def test_keyspace_callback_base_exception_does_not_strand_next_update( + self): + class CompletionCancelled(BaseException): + pass + + pool, host, session, connection = self.make_pool() + pool._connections = {} + completions = [] + cancellation = CompletionCancelled("cancelled completion") + + def first_completed(_, errors): + pool._set_keyspace_for_all_conns( + "second", + lambda _, second_errors: completions.append( + ("second", second_errors))) + raise cancellation + + with pytest.raises(CompletionCancelled) as raised: + pool._set_keyspace_for_all_conns( + "first", + first_completed) + + assert raised.value is cancellation + assert completions == [("second", [])] + assert not pool._keyspace_update_queue + assert pool._keyspace_update_current is None + assert not pool._keyspace_update_in_progress + assert not pool._keyspace_update_runner_active + + def test_shutdown_completes_keyspace_updates_in_fifo_order(self): + pool, host, session, connection = self.make_pool() + callbacks = [] + completions = [] + + def delay_keyspace_update(keyspace, callback): + connection.in_flight += 1 + callbacks.append(callback) + + connection.set_keyspace_async.side_effect = delay_keyspace_update + pool._set_keyspace_for_all_conns( + "first", + lambda _, errors: completions.append(("first", errors))) + pool._set_keyspace_for_all_conns( + "second", + lambda _, errors: completions.append(("second", errors))) + + assert len(callbacks) == 1 + assert completions == [] + + pool.shutdown() + + assert [name for name, _ in completions] == ["first", "second"] + assert all(errors for _, errors in completions) + + # A reactor may deliver the close callback after shutdown returns. + # It must neither repeat the first completion nor disturb FIFO order. + callbacks[0](connection, ConnectionShutdown("closed")) + assert [name for name, _ in completions] == ["first", "second"] + + def test_shutdown_does_not_overtake_active_keyspace_callback(self): + pool, host, session, connection = self.make_pool() + connection_callbacks = [] + completions = [] + first_callback_entered = Event() + release_first_callback = Event() + + def delay_keyspace_update(keyspace, callback): + connection.in_flight += 1 + connection_callbacks.append(callback) + + def finish_first(_, errors): + first_callback_entered.set() + assert release_first_callback.wait(2) + completions.append(("first", errors)) + + connection.set_keyspace_async.side_effect = delay_keyspace_update + pool._set_keyspace_for_all_conns("first", finish_first) + pool._set_keyspace_for_all_conns( + "second", + lambda _, errors: completions.append(("second", errors))) + + callback_thread = Thread( + target=connection_callbacks[0], + args=(connection, None)) + callback_thread.start() + assert first_callback_entered.wait(2) + + pool.shutdown() + assert completions == [] + + release_first_callback.set() + callback_thread.join(2) + + assert not callback_thread.is_alive() + assert [name for name, _ in completions] == ["first", "second"] + assert completions[1][1] + + def test_last_keyspace_response_wins_before_return_during_shutdown(self): + pool, host, session, connection = self.make_pool() + connection_callbacks = [] + completions = [] + return_entered = Event() + release_return = Event() + + def capture_keyspace_callback(keyspace, callback): + connection.in_flight += 1 + connection_callbacks.append(callback) + + def block_return(conn): + return_entered.set() + assert release_return.wait(2) + conn.in_flight -= 1 + + connection.set_keyspace_async.side_effect = \ + capture_keyspace_callback + pool.return_connection = block_return + pool._set_keyspace_for_all_conns( + "new_keyspace", + lambda _, errors: completions.append(errors)) + + response_thread = Thread( + target=connection_callbacks[0], + args=(connection, None)) + response_thread.start() + assert return_entered.wait(2) + + # The final response removed the final pending connection while + # holding the update lock. Shutdown must not replace that success + # with its own error while in_flight is being balanced. + pool.shutdown() + assert completions == [] + + release_return.set() + response_thread.join(2) + + assert not response_thread.is_alive() + assert completions == [[]] + + def test_shutdown_stops_dispatch_to_remaining_connections(self): + pool, host, session, first_connection = self.make_pool() + second_connection = self.make_connection(shard_id=1) + pool._connections[1] = second_connection + first_dispatch_entered = Event() + release_first_dispatch = Event() + completed = Mock() + + def block_first_dispatch(keyspace, callback): + first_dispatch_entered.set() + assert release_first_dispatch.wait(2) + + first_connection.set_keyspace_async.side_effect = block_first_dispatch + + dispatch_thread = Thread( + target=pool._set_keyspace_for_all_conns, + args=("new_keyspace", completed)) + dispatch_thread.start() + assert first_dispatch_entered.wait(2) + + pool.shutdown() + release_first_dispatch.set() + dispatch_thread.join(2) + + assert not dispatch_thread.is_alive() + second_connection.set_keyspace_async.assert_not_called() + completed.assert_called_once() + assert completed.call_args.args[1] + + def test_forced_handoff_uses_global_lock_order_and_fences_stale_pool(self): + pool, host, session, first_connection = self.make_pool() + events = [] + + class RecordingLock(object): + def __init__(self, name): + self.name = name + + def __enter__(self): + events.append("enter_" + self.name) + + def __exit__(self, *args): + events.append("exit_" + self.name) + + session.cluster._lock = RecordingLock("cluster") + host.lock = RecordingLock("host") + session._lock = RecordingLock("session") + session._pools = {host: pool} + session.is_shutdown = False + session.cluster.signal_connection_failure.side_effect = ( + lambda *args, **kwargs: events.append("signal")) + + assert pool._handoff_replacement_failure(RuntimeError("failed")) + # User-extensible conviction callbacks run outside driver locks. + assert events == [ + "enter_cluster", + "enter_host", + "enter_session", + "exit_session", + "exit_host", + "exit_cluster", + "signal", + ] + + events[:] = [] + session.cluster.signal_connection_failure.reset_mock() + session._pools[host] = object() + assert not pool._handoff_replacement_failure( + RuntimeError("stale failure")) + session.cluster.signal_connection_failure.assert_not_called() + + def test_open_all_shards_closes_saved_trash_snapshot(self): + pool, host, session, first_connection = self.make_pool( + shard_aware=True) + trashed = self.make_connection(shard_id=1) + pool._trash.add(trashed) + pool._schedule_connection_to_missing_shard = Mock() + + pool._open_connections_for_all_shards(skip_shard_id=0) + + trashed.close.assert_called_once_with() + assert pool._trash == set() + def test_fast_shutdown(self): class MockSession(MagicMock): is_shutdown = False diff --git a/tests/unit/test_response_future.py b/tests/unit/test_response_future.py index cf1194a91f..594b453781 100644 --- a/tests/unit/test_response_future.py +++ b/tests/unit/test_response_future.py @@ -15,17 +15,20 @@ import unittest from collections import deque -from threading import RLock -from unittest.mock import Mock, MagicMock, ANY +from threading import Lock, RLock +from unittest.mock import Mock, MagicMock, ANY, patch -from cassandra import ConsistencyLevel, Unavailable, SchemaTargetType, SchemaChangeType, OperationTimedOut +from cassandra import (ConsistencyLevel, InvalidRequest, OperationTimedOut, + SchemaChangeType, SchemaTargetType, Unavailable, + UnsupportedOperation) from cassandra.cluster import Session, ResponseFuture, NoHostAvailable, ProtocolVersion, ControlConnectionQueryFallback from cassandra.connection import Connection, ConnectionException from cassandra.protocol import (ReadTimeoutErrorMessage, WriteTimeoutErrorMessage, UnavailableErrorMessage, ResultMessage, QueryMessage, ExecuteMessage, OverloadedErrorMessage, IsBootstrappingErrorMessage, - PreparedQueryNotFound, PrepareMessage, ServerError, + PreparedQueryNotFound, PrepareMessage, + BatchMessage, ServerError, RESULT_KIND_ROWS, RESULT_KIND_SET_KEYSPACE, RESULT_KIND_SCHEMA_CHANGE, RESULT_KIND_PREPARED, ProtocolHandler) @@ -40,7 +43,11 @@ class ResponseFutureTests(unittest.TestCase): def make_basic_session(self): s = Mock(spec=Session) + s._lock = RLock() + s.keyspace = None + s._protocol_version = ProtocolVersion.V5 s.row_factory = lambda col_names, rows: [(col_names, rows)] + s.cluster.protocol_version = ProtocolVersion.V5 s.cluster.control_connection._tablets_routing_v1 = False s.cluster.allow_control_connection_query_fallback = ControlConnectionQueryFallback.Disabled return s @@ -63,6 +70,8 @@ def make_control_connection(self): connection.orphaned_threshold = 75 connection.orphaned_threshold_reached = False connection.is_control_connection = True + connection.protocol_version = ProtocolVersion.V5 + connection.keyspace = None connection.get_request_id.return_value = 7 connection.send_msg.return_value = 128 return connection @@ -103,6 +112,135 @@ def test_result_message(self): result = rf.result()[0] assert result == expected_result + def test_continuous_paging_session_registered_before_pool_return(self): + session = self.make_basic_session() + session.cluster._default_load_balancing_policy.make_query_plan.return_value = [] + query = SimpleStatement("SELECT * FROM foo") + message = QueryMessage( + query=query, + consistency_level=ConsistencyLevel.ONE) + # QueryMessage's protocol-specific construction normally supplies + # this attribute through Session message creation. + message.continuous_paging_options = object() + rf = ResponseFuture(session, message, query, 1) + + connection = Connection('127.0.0.1') + connection.lock = Lock() + connection.in_flight = 1 + + response = ResultMessage(RESULT_KIND_ROWS) + response.stream_id = 301 + response.column_names = ['value'] + response.column_types = [object()] + response.parsed_rows = [(1,)] + response.paging_state = None + response.continuous_paging_last = False + + pool = Mock() + pool.is_shutdown = False + release_observations = [] + + def return_connection(conn): + lock_was_free = conn.lock.acquire(False) + try: + paging_registered = bool( + conn._continuous_paging_sessions) + release_observations.append( + (lock_was_free, paging_registered)) + conn.in_flight -= 1 + if not paging_registered: + # Model retirement of a connection whose last ordinary + # request was just returned. + conn.is_closed = True + finally: + if lock_was_free: + conn.lock.release() + + pool.return_connection.side_effect = return_connection + + rf._set_result(None, connection, pool, response) + + assert release_observations == [(True, True)] + assert not connection.is_closed + assert connection._continuous_paging_sessions[301] is \ + rf._continuous_paging_session + assert rf._event.is_set() + + def test_continuous_paging_registration_removed_when_handoff_fails(self): + session = self.make_basic_session() + session.cluster._default_load_balancing_policy.make_query_plan.return_value = [] + query = SimpleStatement("SELECT * FROM foo") + message = QueryMessage( + query=query, + consistency_level=ConsistencyLevel.ONE) + message.continuous_paging_options = object() + rf = ResponseFuture(session, message, query, 1) + + connection = Connection('127.0.0.1') + paging_session = Mock() + paging_session.on_message.side_effect = RuntimeError( + "failed to accept first page") + + def register_paging_session(stream_id, *args): + with connection.lock: + connection._continuous_paging_sessions[stream_id] = \ + paging_session + return paging_session + + connection.new_continuous_paging_session = Mock( + side_effect=register_paging_session) + + response = ResultMessage(RESULT_KIND_ROWS) + response.stream_id = 301 + response.column_names = ['value'] + response.column_types = [object()] + response.parsed_rows = [(1,)] + response.paging_state = None + response.continuous_paging_last = False + + pool = Mock() + pool.is_shutdown = False + + # Avoid exercising the host Python 3.12 logging crash while verifying + # the intentional handler-failure path. + with patch('cassandra.cluster.log.exception'): + rf._set_result(None, connection, pool, response) + + pool.return_connection.assert_called_once_with(connection) + pool.on_connection_released.assert_called_once_with(connection) + assert connection._continuous_paging_sessions == {} + assert rf._continuous_paging_session is None + with pytest.raises( + RuntimeError, match="failed to accept first page"): + rf.result() + + def test_use_with_legal_prefix_can_borrow_quarantined_connection(self): + for query_string in ( + "USE\tsystem", + " \nUSE\nsystem", + "-- app=foo\nUSE system", + "// app=foo\r\nUSE system", + "/* app=foo */ USE system"): + with self.subTest(query_string=query_string): + session = self.make_session() + pool = session._pools.get.return_value + connection = Mock(spec=Connection) + pool.borrow_connection.return_value = (connection, 1) + query = SimpleStatement(query_string) + message = QueryMessage( + query=query_string, + consistency_level=ConsistencyLevel.ONE) + rf = ResponseFuture(session, message, query, 1) + + rf.send_request() + + pool.borrow_connection.assert_called_once_with( + timeout=ANY, + routing_key=ANY, + keyspace=ANY, + table=ANY, + allow_keyspace_mismatch=True) + def test_unknown_result_class(self): session = self.make_session() pool = session._pools.get.return_value @@ -124,9 +262,65 @@ def test_set_keyspace_result(self): kind=RESULT_KIND_SET_KEYSPACE, results="keyspace1") rf._set_result(None, None, None, result) + session.submit.reset_mock() rf._set_keyspace_completed({}) + session.submit.assert_called_once_with( + session.update_created_pools) assert not rf.result() + def test_successful_use_clears_quarantine_before_pool_return(self): + session = self.make_session() + rf = self.make_response_future(session) + connection = Mock(_pool_keyspace_mismatch=True) + pool = Mock(is_shutdown=False) + returned_state = [] + + def return_connection(returned_connection): + returned_state.append(( + returned_connection, + returned_connection._pool_keyspace_mismatch, + returned_connection.keyspace)) + + pool.return_connection.side_effect = return_connection + result = Mock( + spec=ResultMessage, + kind=RESULT_KIND_SET_KEYSPACE, + new_keyspace='keyspace1') + + rf._set_result('host', connection, pool, result) + + pool.return_connection.assert_called_once_with(connection) + assert returned_state == [(connection, False, 'keyspace1')] + + def test_set_keyspace_error_does_not_repair_missing_pools(self): + session = self.make_session() + rf = self.make_response_future(session) + session.submit.reset_mock() + + rf._set_keyspace_completed({'host': [InvalidRequest('invalid')]}) + + session.submit.assert_not_called() + with pytest.raises(ConnectionException): + rf.result() + + def test_set_keyspace_success_survives_repair_submission_failure(self): + class SubmitCancelled(BaseException): + pass + + failures = ( + ValueError("executor rejected repair"), + SubmitCancelled("executor cancelled repair"), + ) + for failure in failures: + with self.subTest(failure=type(failure)): + session = self.make_session() + session.submit.side_effect = failure + rf = self.make_response_future(session) + + rf._set_keyspace_completed({}) + + assert rf.result().one() is None + def test_schema_change_result(self): session = self.make_session() rf = self.make_response_future(session) @@ -428,6 +622,7 @@ def test_control_connection_fallback_updates_connection_keyspace(self): session.cluster.allow_control_connection_query_fallback = ControlConnectionQueryFallback.Fallback session.cluster._default_load_balancing_policy.make_query_plan.return_value = ['ip1'] session._pools = {} + session.keyspace = 'oldks' def set_keyspace_for_all_pools(keyspace, callback): session.keyspace = keyspace @@ -451,6 +646,214 @@ def set_keyspace_for_all_pools(keyspace, callback): assert session.keyspace == 'newks' assert rf.result().current_rows == [] + def test_control_connection_fallback_scopes_session_keyspace(self): + for fallback_mode in ( + ControlConnectionQueryFallback.Fallback, + ControlConnectionQueryFallback.SkipPoolCreation): + with self.subTest(fallback_mode=fallback_mode): + session = self.make_basic_session() + session.keyspace = 'session_ks' + session.cluster.allow_control_connection_query_fallback = \ + fallback_mode + session.cluster._default_load_balancing_policy.\ + make_query_plan.return_value = ['ip1'] + session._pools = {} + connection = self.make_control_connection() + connection.keyspace = 'another_session_ks' + session.cluster.control_connection._connection = connection + + rf = self.make_response_future(session) + assert rf.message.keyspace is None + assert rf.send_request() + + sent_message = connection.send_msg.call_args.args[0] + assert sent_message is rf.message + assert sent_message.keyspace == 'session_ks' + + def test_control_connection_fallback_scopes_supported_messages(self): + message_factories = ( + ( + 'query', + lambda: QueryMessage( + query="SELECT * FROM foo", + consistency_level=ConsistencyLevel.ONE)), + ( + 'execute', + lambda: ExecuteMessage( + query_id=b'query-id', + query_params=(), + consistency_level=ConsistencyLevel.ONE)), + ( + 'batch', + lambda: BatchMessage( + batch_type=Mock(), + queries=(), + consistency_level=ConsistencyLevel.ONE)), + ( + 'prepare', + lambda: PrepareMessage(query="SELECT * FROM foo")), + ) + for message_name, message_factory in message_factories: + with self.subTest(message_name=message_name): + session = self.make_basic_session() + session.keyspace = 'session_ks' + session.cluster.allow_control_connection_query_fallback = \ + ControlConnectionQueryFallback.SkipPoolCreation + session.cluster._default_load_balancing_policy.\ + make_query_plan.return_value = [] + session._pools = {} + connection = self.make_control_connection() + session.cluster.control_connection._connection = connection + query = SimpleStatement("SELECT * FROM foo") + rf = ResponseFuture( + session, + message_factory(), + query, + 1) + + assert rf.send_request() + + assert connection.send_msg.call_args.args[0].keyspace == \ + 'session_ks' + + def test_control_connection_fallback_preserves_keyspace_snapshot(self): + session = self.make_basic_session() + session.keyspace = 'original_ks' + 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 + + rf = self.make_response_future(session) + session.keyspace = 'new_ks' + assert rf.send_request() + + assert connection.send_msg.call_args.args[0].keyspace == 'original_ks' + + def test_control_connection_fallback_explicit_keyspace_wins(self): + session = self.make_basic_session() + session.keyspace = 'session_ks' + 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 + query = SimpleStatement( + "SELECT * FROM foo", + keyspace='statement_ks') + message = QueryMessage( + query=query.query_string, + consistency_level=ConsistencyLevel.ONE, + keyspace='message_ks') + rf = ResponseFuture(session, message, query, 1) + + assert rf.send_request() + + assert connection.send_msg.call_args.args[0].keyspace == 'message_ks' + + def test_control_connection_fallback_keyspace_requires_v5(self): + for fallback_mode in ( + ControlConnectionQueryFallback.Fallback, + ControlConnectionQueryFallback.SkipPoolCreation): + with self.subTest(fallback_mode=fallback_mode): + session = self.make_basic_session() + session.keyspace = 'session_ks' + session.cluster.allow_control_connection_query_fallback = \ + fallback_mode + session.cluster._default_load_balancing_policy.\ + make_query_plan.return_value = ['ip1'] + session._pools = {} + connection = self.make_control_connection() + connection.protocol_version = ProtocolVersion.V4 + connection.keyspace = 'session_ks' + session.cluster.control_connection._connection = connection + + rf = self.make_response_future(session) + rf.send_request() + + connection.send_msg.assert_not_called() + with pytest.raises(NoHostAvailable) as exc_info: + rf.result() + assert any( + isinstance(error, UnsupportedOperation) + for error in exc_info.value.errors.values()) + + def test_control_connection_fallback_rejects_use(self): + session = self.make_basic_session() + session.cluster.allow_control_connection_query_fallback = \ + ControlConnectionQueryFallback.SkipPoolCreation + session.cluster._default_load_balancing_policy.\ + make_query_plan.return_value = [] + session._pools = {} + connection = self.make_control_connection() + session.cluster.control_connection._connection = connection + query = SimpleStatement("USE other_ks") + message = QueryMessage( + query=query.query_string, + consistency_level=ConsistencyLevel.ONE) + rf = ResponseFuture(session, message, query, 1) + + rf.send_request() + + connection.send_msg.assert_not_called() + assert connection.keyspace is None + with pytest.raises(NoHostAvailable) as exc_info: + rf.result() + assert any( + isinstance(error, UnsupportedOperation) + for error in exc_info.value.errors.values()) + + def test_control_connection_fallback_rejects_unscoped_keyed_socket(self): + session = self.make_basic_session() + session.cluster.allow_control_connection_query_fallback = \ + ControlConnectionQueryFallback.SkipPoolCreation + session.cluster._default_load_balancing_policy.\ + make_query_plan.return_value = [] + session._pools = {} + connection = self.make_control_connection() + connection.keyspace = 'another_session_ks' + session.cluster.control_connection._connection = connection + + rf = self.make_response_future(session) + rf.send_request() + + connection.send_msg.assert_not_called() + with pytest.raises(NoHostAvailable) as exc_info: + rf.result() + assert any( + isinstance(error, UnsupportedOperation) + for error in exc_info.value.errors.values()) + + def test_continuous_paging_never_uses_control_connection_fallback(self): + for fallback_mode in ( + ControlConnectionQueryFallback.Fallback, + ControlConnectionQueryFallback.SkipPoolCreation): + with self.subTest(fallback_mode=fallback_mode): + session = self.make_basic_session() + session.cluster.allow_control_connection_query_fallback = \ + fallback_mode + 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 + rf = self.make_response_future(session) + rf.message.continuous_paging_options = object() + + rf.send_request() + + connection.send_msg.assert_not_called() + with pytest.raises(NoHostAvailable) as exc_info: + rf.result() + assert any( + isinstance(error, UnsupportedOperation) + for error in exc_info.value.errors.values()) + def test_control_connection_fallback_when_no_usable_pools(self): session = self.make_basic_session() session.cluster.allow_control_connection_query_fallback = ControlConnectionQueryFallback.SkipPoolCreation @@ -517,6 +920,7 @@ def test_control_connection_fallback_retries_after_server_error(self): def test_control_connection_fallback_fetches_next_page(self): session = self.make_basic_session() + session.keyspace = 'original_ks' session.cluster.allow_control_connection_query_fallback = ControlConnectionQueryFallback.Fallback session.cluster._default_load_balancing_policy.make_query_plan.return_value = ['ip1'] session._pools = {} @@ -536,12 +940,14 @@ def test_control_connection_fallback_fetches_next_page(self): assert rf.result().current_rows == [(['col'], [(1,)])] assert rf.has_more_pages + session.keyspace = 'changed_ks' rf.start_fetching_next_page() assert connection.send_msg.call_count == 2 assert connection.send_msg.call_args_list[1][0][0] is rf.message assert connection.send_msg.call_args_list[1][0][1] == 8 assert rf.message.paging_state == b'next-page' + assert rf.message.keyspace == 'original_ks' second_response = self.make_mock_response(['col'], [(2,)]) connection.send_msg.call_args_list[1][1]['cb'](second_response) @@ -552,7 +958,7 @@ def test_control_connection_fallback_fetches_next_page(self): def test_control_connection_fallback_reprepares_prepared_statement(self): session = self.make_basic_session() session.cluster.allow_control_connection_query_fallback = ControlConnectionQueryFallback.Fallback - session.cluster.protocol_version = ProtocolVersion.V4 + session.cluster.protocol_version = ProtocolVersion.V5 session.cluster._default_load_balancing_policy.make_query_plan.return_value = ['ip1'] session._pools = {} session.submit.side_effect = lambda fn, *args, **kwargs: fn(*args, **kwargs) @@ -567,7 +973,7 @@ def test_control_connection_fallback_reprepares_prepared_statement(self): session.cluster._prepared_statements = {query_id: prepared_statement} connection = self.make_control_connection() - connection.keyspace = "FooKeyspace" + connection.keyspace = "AnotherKeyspace" connection.get_request_id.side_effect = [7, 8, 9] session.cluster.control_connection._connection = connection control_host = Mock(endpoint=connection.endpoint) @@ -575,7 +981,10 @@ def test_control_connection_fallback_reprepares_prepared_statement(self): rf = self.make_response_future(session) rf.prepared_statement = prepared_statement + rf._control_connection_keyspace = prepared_statement.keyspace assert rf.send_request() + assert rf.message.keyspace == "FooKeyspace" + session.keyspace = "ChangedKeyspace" missing = Mock(spec=PreparedQueryNotFound, info=query_id) connection.send_msg.call_args_list[0][1]['cb'](missing) @@ -584,6 +993,7 @@ def test_control_connection_fallback_reprepares_prepared_statement(self): prepare_message = connection.send_msg.call_args_list[1][0][0] assert isinstance(prepare_message, PrepareMessage) assert prepare_message.query == "SELECT * FROM foobar" + assert prepare_message.keyspace == "FooKeyspace" assert connection.send_msg.call_args_list[1][0][1] == 8 prepared_response = Mock( @@ -596,6 +1006,7 @@ def test_control_connection_fallback_reprepares_prepared_statement(self): assert connection.send_msg.call_count == 3 assert connection.send_msg.call_args_list[2][0][0] is rf.message + assert rf.message.keyspace == "FooKeyspace" assert connection.send_msg.call_args_list[2][0][1] == 9 expected_result = (['col'], [(1,)]) @@ -622,6 +1033,67 @@ def test_control_connection_fallback_not_used_when_pool_can_serve(self): with pytest.raises(NoHostAvailable): rf.result() + def test_control_connection_fallback_used_when_pool_is_quarantined(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'] + quarantined_connection = Mock( + is_closed=False, + is_defunct=False, + _pool_retired=False, + _pool_keyspace_mismatch=True) + pool = Mock(is_shutdown=False) + pool.get_connections.return_value = [quarantined_connection] + pool.borrow_connection.side_effect = NoConnectionsAvailable() + session._pools = {'ip1': pool} + control_connection = self.make_control_connection() + session.cluster.control_connection._connection = control_connection + control_host = Mock(endpoint=control_connection.endpoint) + session.cluster.get_control_connection_host.return_value = control_host + + rf = self.make_response_future(session) + assert rf.send_request() + + control_connection.send_msg.assert_called_once_with( + rf.message, 7, cb=ANY, + encoder=ProtocolHandler.encode_message, + decoder=ProtocolHandler.decode_message, + result_metadata=[]) + pool.borrow_connection.assert_called_once_with( + timeout=ANY, + routing_key=ANY, + keyspace=ANY, + table=ANY) + assert rf.attempted_hosts == [control_host] + + def test_usable_pool_snapshot_holds_session_lock(self): + session = self.make_basic_session() + session._lock = Lock() + session.cluster._default_load_balancing_policy.\ + make_query_plan.return_value = [] + connection = Mock( + is_closed=False, + is_defunct=False, + _pool_retired=False, + _pool_keyspace_mismatch=False) + pool = Mock(is_shutdown=False) + pool.get_connections.return_value = [connection] + + class LockCheckingPools(dict): + def values(mapping): + if session._lock.acquire(False): + session._lock.release() + raise AssertionError( + "pool mapping was read without Session lock") + return super(LockCheckingPools, mapping).values() + + session._pools = LockCheckingPools(ip1=pool) + + assert self.make_response_future( + session)._has_usable_node_pool() + def test_control_connection_fallback_orphans_stream_on_timeout(self): session = self.make_basic_session() session.cluster.allow_control_connection_query_fallback = ControlConnectionQueryFallback.Fallback diff --git a/tests/unit/test_shard_aware.py b/tests/unit/test_shard_aware.py index 902b48a276..1319709e32 100644 --- a/tests/unit/test_shard_aware.py +++ b/tests/unit/test_shard_aware.py @@ -39,6 +39,7 @@ def __init__(self, ssl_options=None, ssl_context=None, sharding_info=None, self.cluster.ssl_options = ssl_options self.cluster.ssl_context = ssl_context self.cluster.shard_aware_options = ShardAwareOptions() + self.cluster.connect_timeout = 5 self.cluster.executor = ThreadPoolExecutor(max_workers=2) self.cluster.signal_connection_failure = lambda *args, **kwargs: False self.cluster.connection_factory = self.mock_connection_factory @@ -73,6 +74,75 @@ def mock_connection_factory(self, *args, **kwargs): class TestShardAware(unittest.TestCase): + def _assert_closed_missing_shard_connection_is_discarded( + self, use_shard_aware_endpoint): + host = MagicMock() + host.endpoint = DefaultEndPoint("1.2.3.4") + session = MockSession() + session.cluster.shard_aware_options.disable_shardaware_port = ( + not use_shard_aware_endpoint) + pool = HostConnection( + host=host, host_distance=HostDistance.REMOTE, session=session) + + try: + for future in session.futures: + future.result() + + requested_shard = 2 + + class ClosedConnection(object): + def __init__(self): + self.is_closed = True + self.close = MagicMock() + self.set_keyspace_blocking = MagicMock() + self.features_access_count = 0 + self._features = ProtocolFeatures( + shard_id=requested_shard) + + @property + def features(self): + self.features_access_count += 1 + return self._features + + closed_connection = ClosedConnection() + session.cluster.connection_factory = MagicMock( + return_value=closed_connection) + pool._connections.clear() + pool._excess_connections.clear() + pool._connecting.add(requested_shard) + + pool._open_connection_to_missing_shard(requested_shard) + + assert pool._connections == {} + assert pool._excess_connections == set() + assert closed_connection.features_access_count == 0 + closed_connection.set_keyspace_blocking.assert_not_called() + closed_connection.close.assert_called_once_with() + assert requested_shard not in pool._connecting + + factory_args = session.cluster.connection_factory.call_args + expected_endpoint = ( + DefaultEndPoint("1.2.3.4", port=19042) + if use_shard_aware_endpoint else host.endpoint) + assert factory_args.args[0] == expected_endpoint + if use_shard_aware_endpoint: + assert factory_args.kwargs['shard_id'] == requested_shard + assert factory_args.kwargs['total_shards'] == 4 + else: + assert 'shard_id' not in factory_args.kwargs + assert 'total_shards' not in factory_args.kwargs + finally: + pool.shutdown() + session.cluster.executor.shutdown(wait=True) + + def test_closed_connection_to_missing_shard_is_discarded(self): + self._assert_closed_missing_shard_connection_is_discarded( + use_shard_aware_endpoint=True) + + def test_closed_connection_to_missing_shard_fallback_is_discarded(self): + self._assert_closed_missing_shard_connection_is_discarded( + use_shard_aware_endpoint=False) + def test_parsing_and_calculating_shard_id(self): """ Testing the parsing of the options command