1import asyncio
2import collections
3import logging
4import random
5import socket
6import threading
7import time
8import warnings
9import weakref
10from abc import ABC, abstractmethod
11from collections import defaultdict
12from contextlib import nullcontext
13from copy import copy
14from itertools import chain
15from types import MethodType
16from typing import (
17 TYPE_CHECKING,
18 Any,
19 Callable,
20 Coroutine,
21 Deque,
22 Dict,
23 Generator,
24 Iterable,
25 List,
26 Literal,
27 Mapping,
28 Optional,
29 Set,
30 Tuple,
31 Type,
32 TypeVar,
33 Union,
34)
35
36if TYPE_CHECKING:
37 from redis.asyncio.keyspace_notifications import (
38 AsyncClusterKeyspaceNotifications,
39 )
40
41from redis._defaults import (
42 DEFAULT_RETRY_BASE,
43 DEFAULT_RETRY_CAP,
44 DEFAULT_RETRY_COUNT,
45 DEFAULT_SOCKET_CONNECT_TIMEOUT,
46 DEFAULT_SOCKET_READ_SIZE,
47 DEFAULT_SOCKET_TIMEOUT,
48)
49from redis._parsers import AsyncCommandsParser, Encoder
50from redis._parsers.helpers import get_response_callbacks
51from redis.asyncio import _himport_exec
52from redis.asyncio.client import PubSub, ResponseCallbackT
53from redis.asyncio.connection import (
54 AbstractConnection,
55 Connection,
56 ConnectionPoolInterface,
57 SSLConnection,
58 parse_url,
59)
60from redis.asyncio.lock import Lock
61from redis.asyncio.maint_notifications import AsyncOSSMaintNotificationsHandler
62from redis.asyncio.observability.recorder import (
63 record_error_count,
64 record_operation_duration,
65)
66from redis.asyncio.retry import Retry
67from redis.auth.token import TokenInterface
68from redis.backoff import ExponentialWithJitterBackoff, NoBackoff
69from redis.client import EMPTY_RESPONSE, NEVER_DECODE, AbstractRedis
70from redis.cluster import (
71 _REPLICAS_ONLY_STRATEGIES,
72 PIPELINE_BLOCKED_COMMANDS,
73 PRIMARY,
74 REPLICA,
75 SLOT_ID,
76 AbstractRedisCluster,
77 LoadBalancer,
78 LoadBalancingStrategy,
79 block_pipeline_command,
80 get_node_name,
81 is_zero_key_eval_command,
82 parse_cluster_shards,
83 parse_cluster_shards_unified,
84 parse_cluster_shards_with_str_keys,
85 parse_cluster_slots,
86)
87from redis.commands import AsyncRedisClusterCommands
88from redis.commands.helpers import list_or_args, parse_pubsub_subscriptions
89from redis.commands.metadata import (
90 _DEFAULT_KEYED_METADATA,
91 _DEFAULT_KEYLESS_METADATA,
92 _METADATA_BY_REQUEST_POLICY,
93 AsyncMetadataResolver,
94 AsyncStaticMetadataResolver,
95 CommandMetadata,
96 CommandPolicies,
97 RequestPolicy,
98 ResponsePolicy,
99)
100from redis.commands.policies import AsyncPolicyResolver, AsyncStaticPolicyResolver
101from redis.crc import REDIS_CLUSTER_HASH_SLOTS, key_slot
102from redis.credentials import CredentialProvider
103from redis.driver_info import DriverInfo, resolve_driver_info
104from redis.event import (
105 AfterAsyncClusterInstantiationEvent,
106 AsyncAfterSlotsCacheRefreshEvent,
107 AsyncEventListenerInterface,
108 EventDispatcher,
109)
110from redis.exceptions import (
111 AskError,
112 AuthenticationError,
113 AuthorizationError,
114 BusyLoadingError,
115 ClusterDownError,
116 ClusterError,
117 ConnectionError,
118 CrossSlotTransactionError,
119 DataError,
120 ExecAbortError,
121 InvalidPipelineStack,
122 MaxConnectionsError,
123 MovedError,
124 RedisClusterException,
125 RedisClusterUnreachableError,
126 RedisError,
127 ResponseError,
128 SlotNotCoveredError,
129 TimeoutError,
130 TryAgainError,
131 WatchError,
132)
133from redis.himport import HImportRegistry, parse_himport_set_args
134from redis.maint_notifications import MaintNotificationsConfig
135from redis.typing import (
136 AnyKeyT,
137 ChannelT,
138 EncodableT,
139 FieldT,
140 KeyT,
141 PubSubHandler,
142 Subscription,
143)
144from redis.utils import (
145 SENTINEL,
146 SSL_AVAILABLE,
147 check_protocol_version,
148 deprecated_args,
149 deprecated_function,
150 experimental_method,
151 safe_str,
152 str_if_bytes,
153 truncate_text,
154)
155
156if SSL_AVAILABLE:
157 from ssl import TLSVersion, VerifyFlags, VerifyMode
158else:
159 TLSVersion = None
160 VerifyMode = None
161 VerifyFlags = None
162
163logger = logging.getLogger(__name__)
164
165TargetNodesT = TypeVar(
166 "TargetNodesT", str, "ClusterNode", List["ClusterNode"], Dict[Any, "ClusterNode"]
167)
168
169_T = TypeVar("_T")
170
171
172def _run_coroutine_in_thread(coro: Coroutine[Any, Any, _T]) -> _T:
173 """
174 Runs ``coro`` to completion on a private event loop in a worker thread.
175
176 The bridge for the places where a synchronous entry point has to drive an awaitable
177 body. The caller may already be running an event loop, which rules out both awaiting
178 the coroutine and running it here with ``asyncio.run``, so it gets a loop of its own
179 and the calling thread blocks until it finishes. Whatever the coroutine raised is
180 re-raised to the caller, so the entry point fails the way it would have had the body
181 run in place.
182
183 .. warning::
184
185 **Deadlock hazard. Pass only a coroutine that touches nothing bound to the
186 caller's loop.** The ``thread.join()`` below blocks the calling thread for the
187 whole run, so when that thread is the one running an event loop, the loop stops
188 making progress for the duration. A body that then awaits anything owned by that
189 loop - a future, a task, a lock, or I/O over a pooled connection - is waiting on a
190 loop that is itself blocked on this very ``join()``. The result is a permanent
191 stall, not an exception: nothing times out and nothing is raised. Loop-bound
192 futures may instead fail with a cross-loop ``RuntimeError``, which is the luckier
193 outcome because at least it is visible.
194
195 Pass only a body whose awaits are self-contained, and prefer restructuring the entry
196 point to be ``async`` over reaching for this.
197 """
198 result: Any = None
199 error: BaseException | None = None
200
201 def runner() -> None:
202 nonlocal result
203 nonlocal error
204 try:
205 result = asyncio.run(coro)
206 except BaseException as e:
207 error = e
208
209 thread = threading.Thread(target=runner)
210 thread.start()
211 # Unconditional and untimed by design: the caller expects a value, so there is no
212 # partial result to return early with. If the caller's thread is running an event
213 # loop, that loop is blocked here, which is why the coroutine must never await
214 # anything owned by it. See the deadlock hazard in the docstring above.
215 thread.join()
216
217 if error is not None:
218 raise error
219
220 return result
221
222
223class AsyncMaintNotificationsAbstractRedisCluster:
224 """
225 Mixin for async cluster maintenance notifications handling.
226
227 Intended to be used with multiple inheritance alongside RedisCluster.
228 All logic related to cluster-level maintenance notifications is encapsulated here.
229 """
230
231 def __init__(
232 self,
233 maint_notifications_config: MaintNotificationsConfig | None,
234 **kwargs,
235 ) -> None:
236 # The RESP3 requirement is validated in RedisCluster.__init__ before the
237 # NodesManager is constructed; this mixin is only ever run from there, so
238 # the config it receives has already been validated.
239 is_protocol_supported = check_protocol_version(kwargs.get("protocol"), 3)
240
241 if maint_notifications_config is None and is_protocol_supported:
242 maint_notifications_config = MaintNotificationsConfig()
243
244 self.maint_notifications_config = maint_notifications_config
245
246 if self.maint_notifications_config and self.maint_notifications_config.enabled:
247 self._oss_cluster_maint_notifications_handler = (
248 AsyncOSSMaintNotificationsHandler(self, self.maint_notifications_config)
249 )
250 self._update_connection_kwargs_for_maint_notifications(
251 self._oss_cluster_maint_notifications_handler
252 )
253 # Connections are created lazily via ClusterNode.acquire_connection()
254 # during nodes_manager.initialize() (which runs after __init__), so
255 # injecting into the shared connection_kwargs covers nodes discovered
256 # later. Startup nodes are the exception — they were built before this
257 # runs with their own kwargs snapshot — so the helper above also
258 # updates them directly.
259 else:
260 self._oss_cluster_maint_notifications_handler = None
261
262 def _update_connection_kwargs_for_maint_notifications(
263 self,
264 oss_cluster_maint_notifications_handler: AsyncOSSMaintNotificationsHandler,
265 ) -> None:
266 maint_kwargs = {
267 "oss_cluster_maint_notifications_handler": oss_cluster_maint_notifications_handler,
268 "maint_notifications_config": oss_cluster_maint_notifications_handler.config,
269 }
270 # Shared template used for every node created from now on (e.g. nodes
271 # discovered during nodes_manager.initialize()).
272 self.nodes_manager.connection_kwargs.update(maint_kwargs)
273 # Startup nodes were constructed before this mixin ran, so each one
274 # snapshotted connection_kwargs without the handler. Their connections
275 # are created lazily, so updating their per-node kwargs now is in time —
276 # otherwise initialize() opens the topology-discovery connection (CLUSTER
277 # SLOTS) on a startup node with no push handler wired and silently drops
278 # the maintenance notifications carried on that connection.
279 for node in self.nodes_manager.startup_nodes.values():
280 node.connection_kwargs.update(maint_kwargs)
281
282
283class RedisCluster(
284 AbstractRedis,
285 AbstractRedisCluster,
286 AsyncMaintNotificationsAbstractRedisCluster,
287 AsyncRedisClusterCommands,
288):
289 """
290 Create a new RedisCluster client.
291
292 Pass one of parameters:
293
294 - `host` & `port`
295 - `startup_nodes`
296
297 | Use ``await`` :meth:`initialize` to find cluster nodes & create connections.
298 | Use ``await`` :meth:`close` to disconnect connections & close client.
299
300 Many commands support the target_nodes kwarg. It can be one of the
301 :attr:`NODE_FLAGS`:
302
303 - :attr:`PRIMARIES`
304 - :attr:`REPLICAS`
305 - :attr:`ALL_NODES`
306 - :attr:`RANDOM`
307 - :attr:`DEFAULT_NODE`
308
309 Note: This client is not thread/process/fork safe.
310
311 :param host:
312 | Can be used to point to a startup node
313 :param port:
314 | Port used if **host** is provided
315 :param startup_nodes:
316 | :class:`~.ClusterNode` to used as a startup node
317 :param require_full_coverage:
318 | When set to ``False``: the client will not require a full coverage of
319 the slots. However, if not all slots are covered, and at least one node
320 has ``cluster-require-full-coverage`` set to ``yes``, the server will throw
321 a :class:`~.ClusterDownError` for some key-based commands.
322 | When set to ``True``: all slots must be covered to construct the cluster
323 client. If not all slots are covered, :class:`~.RedisClusterException` will be
324 thrown.
325 | See:
326 https://redis.io/docs/manual/scaling/#redis-cluster-configuration-parameters
327 :param read_from_replicas:
328 | @deprecated - please use load_balancing_strategy instead
329 | Enable read from replicas in READONLY mode.
330 When set to true, read commands will be assigned between the primary and
331 its replications in a Round-Robin manner.
332 The data read from replicas is eventually consistent with the data in primary nodes.
333 :param load_balancing_strategy:
334 | Enable read from replicas in READONLY mode and defines the load balancing
335 strategy that will be used for cluster node selection.
336 The data read from replicas is eventually consistent with the data in primary nodes.
337 :param metadata_resolver:
338 | Optional :class:`~.AsyncMetadataResolver` instance used to map command names
339 to replica-safe routing rules. If not provided, an AsyncStaticMetadataResolver
340 is used by default. The routing view of it is derived into the default
341 ``policy_resolver``; an explicit ``policy_resolver`` supersedes that view, but
342 replica safety keeps resolving through this argument either way.
343 :param dynamic_startup_nodes:
344 | Set the RedisCluster's startup nodes to all the discovered nodes.
345 If true (default value), the cluster's discovered nodes will be used to
346 determine the cluster nodes-slots mapping in the next topology refresh.
347 It will remove the initial passed startup nodes if their endpoints aren't
348 listed in the CLUSTER SLOTS output.
349 If you use dynamic DNS endpoints for startup nodes but CLUSTER SLOTS lists
350 specific IP addresses, it is best to set it to false.
351 :param reinitialize_steps:
352 | Specifies the number of MOVED errors that need to occur before reinitializing
353 the whole cluster topology. If a MOVED error occurs and the cluster does not
354 need to be reinitialized on this current error handling, only the MOVED slot
355 will be patched with the redirected node.
356 To reinitialize the cluster on every MOVED error, set reinitialize_steps to 1.
357 To avoid reinitializing the cluster on moved errors, set reinitialize_steps to
358 0.
359 :param cluster_error_retry_attempts:
360 | @deprecated - Please configure the 'retry' object instead
361 In case 'retry' object is set - this argument is ignored!
362
363 Number of times to retry before raising an error when :class:`~.TimeoutError`,
364 :class:`~.ConnectionError`, :class:`~.SlotNotCoveredError`
365 or :class:`~.ClusterDownError` are encountered
366 :param retry:
367 | A retry object that defines the retry strategy and the number of
368 retries for the cluster client.
369 In current implementation for the cluster client (starting form redis-py version 6.0.0)
370 the retry object is not yet fully utilized, instead it is used just to determine
371 the number of retries for the cluster client.
372 In the future releases the retry object will be used to handle the cluster client retries!
373 :param max_connections:
374 | Maximum number of connections per node. If there are no free connections & the
375 maximum number of connections are already created, a
376 :class:`~.MaxConnectionsError` is raised.
377 :param socket_keepalive:
378 | If ``True``, TCP keepalive is enabled for TCP socket connections.
379 :param socket_keepalive_options:
380 | Mapping of TCP keepalive socket option constants to values, for
381 example ``{socket.TCP_KEEPIDLE: 30}``. If left unspecified, redis-py
382 uses TCP keepalive defaults when ``socket_keepalive`` is enabled:
383 idle 30 seconds, interval 5 seconds, and 3 probes.
384 Platform-specific options that are not available are skipped.
385 Pass ``None`` or ``{}`` to avoid setting additional TCP keepalive
386 options.
387 :param address_remap:
388 | An optional callable which, when provided with an internal network
389 address of a node, e.g. a `(host, port)` tuple, will return the address
390 where the node is reachable. This can be used to map the addresses at
391 which the nodes _think_ they are, to addresses at which a client may
392 reach them, such as when they sit behind a proxy.
393
394 | Rest of the arguments will be passed to the
395 :class:`~redis.asyncio.connection.Connection` instances when created
396
397 :raises RedisClusterException:
398 if any arguments are invalid or unknown. Eg:
399
400 - `db` != 0 or None
401 - `path` argument for unix socket connection
402 - none of the `host`/`port` & `startup_nodes` were provided
403
404 """
405
406 @classmethod
407 def from_url(cls, url: str, **kwargs: Any) -> "RedisCluster":
408 """
409 Return a Redis client object configured from the given URL.
410
411 For example::
412
413 redis://[[username]:[password]]@localhost:6379/0
414 rediss://[[username]:[password]]@localhost:6379/0
415
416 Three URL schemes are supported:
417
418 - `redis://` creates a TCP socket connection. See more at:
419 <https://www.iana.org/assignments/uri-schemes/prov/redis>
420 - `rediss://` creates a SSL wrapped TCP socket connection. See more at:
421 <https://www.iana.org/assignments/uri-schemes/prov/rediss>
422
423 The username, password, hostname and path are passed through
424 ``urllib.parse.unquote`` in order to replace any percent-encoded values with
425 their corresponding characters. Querystring values are decoded by
426 ``urllib.parse.parse_qs`` and are not unquoted again.
427
428 All querystring options are cast to their appropriate Python types. Boolean
429 arguments can be specified with string values "True"/"False" or "Yes"/"No".
430 Values that cannot be properly cast cause a ``ValueError`` to be raised. Once
431 parsed, the querystring arguments and keyword arguments are passed to
432 :class:`~redis.asyncio.connection.Connection` when created.
433 In the case of conflicting arguments, querystring arguments are used.
434 """
435 kwargs.update(parse_url(url))
436 if kwargs.pop("connection_class", None) is SSLConnection:
437 kwargs["ssl"] = True
438 return cls(**kwargs)
439
440 # Type discrimination marker for @overload self-type pattern
441 _is_async_client: Literal[True] = True
442
443 __slots__ = (
444 "_initialize",
445 "_lock",
446 "maint_notifications_config",
447 "_oss_cluster_maint_notifications_handler",
448 "_himport_registry",
449 "retry",
450 "command_flags",
451 "commands_parser",
452 "connection_kwargs",
453 "encoder",
454 "node_flags",
455 "nodes_manager",
456 "read_from_replicas",
457 "reinitialize_counter",
458 "reinitialize_steps",
459 "response_callbacks",
460 "result_callbacks",
461 )
462
463 @deprecated_args(
464 args_to_warn=["read_from_replicas"],
465 reason="Please configure the 'load_balancing_strategy' instead",
466 version="5.3.0",
467 )
468 @deprecated_args(
469 args_to_warn=[
470 "cluster_error_retry_attempts",
471 ],
472 reason="Please configure the 'retry' object instead",
473 version="6.0.0",
474 )
475 @deprecated_args(
476 args_to_warn=["lib_name", "lib_version"],
477 reason="Use 'driver_info' parameter instead. "
478 "lib_name and lib_version will be removed in a future version.",
479 )
480 def __init__(
481 self,
482 host: str | None = None,
483 port: str | int = 6379,
484 # Cluster related kwargs
485 startup_nodes: List["ClusterNode"] | None = None,
486 require_full_coverage: bool = True,
487 read_from_replicas: bool = False,
488 load_balancing_strategy: LoadBalancingStrategy | None = None,
489 dynamic_startup_nodes: bool = True,
490 reinitialize_steps: int = 5,
491 cluster_error_retry_attempts: int = DEFAULT_RETRY_COUNT,
492 max_connections: int = 100,
493 retry: Retry | None = None,
494 retry_on_error: List[Type[Exception]] | None = None,
495 # Client related kwargs
496 db: str | int = 0,
497 path: str | None = None,
498 credential_provider: CredentialProvider | None = None,
499 username: str | None = None,
500 password: str | None = None,
501 client_name: str | None = None,
502 lib_name: str | object | None = SENTINEL,
503 lib_version: str | object | None = SENTINEL,
504 driver_info: DriverInfo | object | None = SENTINEL,
505 # Encoding related kwargs
506 encoding: str = "utf-8",
507 encoding_errors: str = "strict",
508 decode_responses: bool = False,
509 # Connection related kwargs
510 health_check_interval: float = 0,
511 socket_timeout: float | None = DEFAULT_SOCKET_TIMEOUT,
512 socket_connect_timeout: float | None = DEFAULT_SOCKET_CONNECT_TIMEOUT,
513 socket_read_size: int = DEFAULT_SOCKET_READ_SIZE,
514 socket_keepalive: bool = True,
515 socket_keepalive_options: Mapping[int, int | bytes] | object | None = SENTINEL,
516 # SSL related kwargs
517 ssl: bool = False,
518 ssl_ca_certs: str | None = None,
519 ssl_ca_data: str | None = None,
520 ssl_cert_reqs: "str | VerifyMode" = "required",
521 ssl_include_verify_flags: List["VerifyFlags"] | None = None,
522 ssl_exclude_verify_flags: List["VerifyFlags"] | None = None,
523 ssl_certfile: str | None = None,
524 ssl_check_hostname: bool = True,
525 ssl_keyfile: str | None = None,
526 ssl_min_version: "TLSVersion | None" = None,
527 ssl_ciphers: str | None = None,
528 protocol: int | None = None,
529 legacy_responses: bool = True,
530 address_remap: Callable[[Tuple[str, int]], Tuple[str, int]] | None = None,
531 event_dispatcher: EventDispatcher | None = None,
532 policy_resolver: AsyncPolicyResolver | None = None,
533 maint_notifications_config: MaintNotificationsConfig | None = None,
534 metadata_resolver: AsyncMetadataResolver | None = None,
535 ) -> None:
536 if db:
537 raise RedisClusterException(
538 "Argument 'db' must be 0 or None in cluster mode"
539 )
540
541 if path:
542 raise RedisClusterException(
543 "Unix domain socket is not supported in cluster mode"
544 )
545
546 port_is_provided = bool(port) or type(port) is int
547 if (not host or not port_is_provided) and not startup_nodes:
548 raise RedisClusterException(
549 "RedisCluster requires at least one node to discover the cluster.\n"
550 "Please provide one of the following or use RedisCluster.from_url:\n"
551 ' - host and port: RedisCluster(host="localhost", port=6379)\n'
552 " - startup_nodes: RedisCluster(startup_nodes=["
553 'ClusterNode("localhost", 6379), ClusterNode("localhost", 6380)])'
554 )
555
556 computed_driver_info = resolve_driver_info(driver_info, lib_name, lib_version)
557
558 kwargs: Dict[str, Any] = {
559 "max_connections": max_connections,
560 "connection_class": Connection,
561 # Client related kwargs
562 "credential_provider": credential_provider,
563 "username": username,
564 "password": password,
565 "client_name": client_name,
566 "driver_info": computed_driver_info,
567 # Encoding related kwargs
568 "encoding": encoding,
569 "encoding_errors": encoding_errors,
570 "decode_responses": decode_responses,
571 # Connection related kwargs
572 "health_check_interval": health_check_interval,
573 "socket_connect_timeout": socket_connect_timeout,
574 "socket_keepalive": socket_keepalive,
575 "socket_keepalive_options": socket_keepalive_options,
576 "socket_read_size": socket_read_size,
577 "socket_timeout": socket_timeout,
578 "protocol": protocol,
579 "legacy_responses": legacy_responses,
580 }
581
582 if ssl:
583 # SSL related kwargs
584 kwargs.update(
585 {
586 "connection_class": SSLConnection,
587 "ssl_ca_certs": ssl_ca_certs,
588 "ssl_ca_data": ssl_ca_data,
589 "ssl_cert_reqs": ssl_cert_reqs,
590 "ssl_include_verify_flags": ssl_include_verify_flags,
591 "ssl_exclude_verify_flags": ssl_exclude_verify_flags,
592 "ssl_certfile": ssl_certfile,
593 "ssl_check_hostname": ssl_check_hostname,
594 "ssl_keyfile": ssl_keyfile,
595 "ssl_min_version": ssl_min_version,
596 "ssl_ciphers": ssl_ciphers,
597 }
598 )
599
600 if read_from_replicas or load_balancing_strategy:
601 # Call our on_connect function to configure READONLY mode
602 kwargs["redis_connect_func"] = self.on_connect
603
604 if retry:
605 self.retry = retry
606 else:
607 self.retry = Retry(
608 backoff=ExponentialWithJitterBackoff(
609 base=DEFAULT_RETRY_BASE, cap=DEFAULT_RETRY_CAP
610 ),
611 retries=cluster_error_retry_attempts,
612 )
613 if retry_on_error:
614 self.retry.update_supported_errors(retry_on_error)
615
616 kwargs["response_callbacks"] = get_response_callbacks(
617 user_protocol=kwargs.get("protocol"),
618 legacy_responses=kwargs.get("legacy_responses", True),
619 )
620 if not kwargs.get("legacy_responses", True):
621 kwargs["response_callbacks"]["CLUSTER SHARDS"] = (
622 parse_cluster_shards_unified
623 )
624 elif kwargs.get("protocol") is None:
625 kwargs["response_callbacks"]["CLUSTER SHARDS"] = (
626 parse_cluster_shards_with_str_keys
627 )
628 else:
629 kwargs["response_callbacks"]["CLUSTER SHARDS"] = parse_cluster_shards
630
631 # Build the client-level HIMPORT registry once (always empty at construction)
632 # and share the same object with every node connection. It rides in
633 # connection_kwargs -> ClusterNode -> each node's Connection, so the registry is
634 # shared cluster-wide and runtime himport_prepare mutates one object. (Async has
635 # no per-node Redis client, so the object flows via connection_kwargs directly to
636 # the Connection, which is internal plumbing, not a public param.)
637 self._himport_registry = HImportRegistry()
638 kwargs["himport_registry"] = self._himport_registry
639
640 self.connection_kwargs = kwargs
641
642 # Validate maint_notifications_config before NodesManager is constructed
643 # so that a bad config doesn't leak an open NodesManager.
644 if (
645 maint_notifications_config
646 and maint_notifications_config.enabled
647 and not check_protocol_version(protocol, 3)
648 ):
649 raise RedisError(
650 "Maintenance notifications are only supported with RESP version 3"
651 )
652 if check_protocol_version(protocol, 3) and maint_notifications_config is None:
653 maint_notifications_config = MaintNotificationsConfig()
654 # Initialize to None so aclose() and any error-path code never sees an
655 # unset slot, even if __init__ raises before the mixin runs.
656 self._oss_cluster_maint_notifications_handler = None
657
658 if startup_nodes:
659 passed_nodes = []
660 for node in startup_nodes:
661 passed_nodes.append(
662 ClusterNode(node.host, node.port, **self.connection_kwargs)
663 )
664 startup_nodes = passed_nodes
665 else:
666 startup_nodes = []
667 if host and port_is_provided:
668 startup_nodes.append(ClusterNode(host, port, **self.connection_kwargs))
669
670 if event_dispatcher is None:
671 self._event_dispatcher = EventDispatcher()
672 else:
673 self._event_dispatcher = event_dispatcher
674
675 self.startup_nodes = startup_nodes
676 self.nodes_manager = NodesManager(
677 startup_nodes,
678 require_full_coverage,
679 kwargs,
680 dynamic_startup_nodes=dynamic_startup_nodes,
681 address_remap=address_remap,
682 event_dispatcher=self._event_dispatcher,
683 )
684 AsyncMaintNotificationsAbstractRedisCluster.__init__(
685 self,
686 maint_notifications_config=maint_notifications_config,
687 protocol=protocol,
688 )
689 self.encoder = Encoder(encoding, encoding_errors, decode_responses)
690 self.read_from_replicas = read_from_replicas
691 self.load_balancing_strategy = load_balancing_strategy
692 self.reinitialize_steps = reinitialize_steps
693 self.reinitialize_counter = 0
694
695 # For backward compatibility, mapping from existing policies to new one
696 self._command_flags_mapping: dict[str, Union[RequestPolicy, ResponsePolicy]] = {
697 self.__class__.RANDOM: RequestPolicy.DEFAULT_KEYLESS,
698 self.__class__.PRIMARIES: RequestPolicy.ALL_SHARDS,
699 self.__class__.ALL_NODES: RequestPolicy.ALL_NODES,
700 self.__class__.REPLICAS: RequestPolicy.ALL_REPLICAS,
701 self.__class__.DEFAULT_NODE: RequestPolicy.DEFAULT_NODE,
702 SLOT_ID: RequestPolicy.DEFAULT_KEYED,
703 }
704
705 self._policies_callback_mapping: dict[
706 Union[RequestPolicy, ResponsePolicy], Callable
707 ] = {
708 RequestPolicy.DEFAULT_KEYLESS: lambda command_name: (
709 self.get_keyless_target_node(command_name)
710 ),
711 RequestPolicy.DEFAULT_KEYED: self.get_nodes_from_slot,
712 RequestPolicy.DEFAULT_NODE: lambda: [self.get_default_node()],
713 RequestPolicy.ALL_SHARDS: self.get_primaries,
714 RequestPolicy.ALL_NODES: self.get_nodes,
715 RequestPolicy.ALL_REPLICAS: self.get_replicas,
716 RequestPolicy.SPECIAL: self.get_special_nodes,
717 ResponsePolicy.DEFAULT_KEYLESS: lambda res: res,
718 ResponsePolicy.DEFAULT_KEYED: lambda res: res,
719 }
720
721 # Built here rather than defaulted in the signature, so that each client owns its
722 # resolver and the memos it accumulates are released with the client. The async
723 # ClusterPipeline holds the client and reads this attribute, so a pipeline routes by
724 # whatever the client routes by without any propagation of its own.
725 if metadata_resolver is None:
726 self._metadata_resolver: AsyncMetadataResolver = (
727 AsyncStaticMetadataResolver()
728 )
729 else:
730 self._metadata_resolver = metadata_resolver
731
732 # ``policy_resolver`` is the routing view of a metadata resolver, so the two
733 # arguments overlap. Resolved by precedence rather than by rejecting the
734 # combination, because a user migrating from one to the other will legitimately pass
735 # both: an explicit ``policy_resolver`` - the extension point that shipped in 7.1.0
736 # - keeps deciding which nodes a command targets, and otherwise those policies are
737 # derived from the metadata resolver.
738 #
739 # The precedence covers the routing view only. Replica safety is read from
740 # ``_metadata_resolver`` either way, because the projection a policy resolver serves
741 # drops the flag it is decided from - a ``CommandPolicies`` record has no
742 # ``is_readonly``. Mirrors the sync stack, including the log below.
743 if policy_resolver is None:
744 self._policy_resolver: AsyncPolicyResolver = AsyncStaticPolicyResolver(
745 metadata_resolver=self._metadata_resolver
746 )
747 else:
748 self._policy_resolver = policy_resolver
749 if metadata_resolver is not None:
750 logger.debug(
751 "Both policy_resolver and metadata_resolver were given; the nodes a "
752 "command targets resolve through policy_resolver and ignore "
753 "metadata_resolver. Replica safety still resolves through "
754 "metadata_resolver."
755 )
756 self.commands_parser = AsyncCommandsParser()
757 self._aggregate_nodes = None
758 self.node_flags = self.__class__.NODE_FLAGS.copy()
759 self.command_flags = self.__class__.COMMAND_FLAGS.copy()
760 self.response_callbacks = kwargs["response_callbacks"]
761 self.result_callbacks = self.__class__.RESULT_CALLBACKS.copy()
762 self.result_callbacks["CLUSTER SLOTS"] = lambda cmd, res, **kwargs: (
763 parse_cluster_slots(list(res.values())[0], **kwargs)
764 )
765
766 self._initialize = True
767 self._lock: Optional[asyncio.Lock] = None
768
769 # When used as an async context manager, we need to increment and decrement
770 # a usage counter so that we can close the connection pool when no one is
771 # using the client.
772 self._usage_counter = 0
773 self._usage_lock = asyncio.Lock()
774
775 async def initialize(
776 self,
777 additional_startup_nodes_info: Optional[List[Tuple[str, int]]] = None,
778 last_failed_node_name: Optional[str] = None,
779 ) -> "RedisCluster":
780 """Get all nodes from startup nodes & creates connections if not initialized."""
781 if self._initialize:
782 if not self._lock:
783 self._lock = asyncio.Lock()
784 async with self._lock:
785 if self._initialize:
786 try:
787 await self.nodes_manager.initialize(
788 additional_startup_nodes_info=additional_startup_nodes_info,
789 last_failed_node_name=last_failed_node_name,
790 )
791 await self.commands_parser.initialize(
792 self.nodes_manager.default_node
793 )
794 self._initialize = False
795 except BaseException:
796 await self.nodes_manager.aclose()
797 await self.nodes_manager.aclose("startup_nodes")
798 raise
799 return self
800
801 async def aclose(self) -> None:
802 """Close all connections & client if initialized."""
803 if not self._initialize:
804 if not self._lock:
805 self._lock = asyncio.Lock()
806 async with self._lock:
807 if not self._initialize:
808 self._initialize = True
809 if self._oss_cluster_maint_notifications_handler:
810 tasks = list(
811 self._oss_cluster_maint_notifications_handler._background_tasks
812 )
813 for task in tasks:
814 task.cancel()
815 await asyncio.gather(*tasks, return_exceptions=True)
816 await self.nodes_manager.aclose()
817 await self.nodes_manager.aclose("startup_nodes")
818
819 @deprecated_function(version="5.0.0", reason="Use aclose() instead", name="close")
820 async def close(self) -> None:
821 """alias for aclose() for backwards compatibility"""
822 await self.aclose()
823
824 async def __aenter__(self) -> "RedisCluster":
825 """
826 Async context manager entry. Increments a usage counter so that the
827 connection pool is only closed (via aclose()) when no context is using
828 the client.
829 """
830 await self._increment_usage()
831 try:
832 # Initialize the client (i.e. establish connection, etc.)
833 return await self.initialize()
834 except Exception:
835 # If initialization fails, decrement the counter to keep it in sync
836 await self._decrement_usage()
837 raise
838
839 async def _increment_usage(self) -> int:
840 """
841 Helper coroutine to increment the usage counter while holding the lock.
842 Returns the new value of the usage counter.
843 """
844 async with self._usage_lock:
845 self._usage_counter += 1
846 return self._usage_counter
847
848 async def _decrement_usage(self) -> int:
849 """
850 Helper coroutine to decrement the usage counter while holding the lock.
851 Returns the new value of the usage counter.
852 """
853 async with self._usage_lock:
854 self._usage_counter -= 1
855 return self._usage_counter
856
857 async def __aexit__(self, exc_type, exc_value, traceback):
858 """
859 Async context manager exit. Decrements a usage counter. If this is the
860 last exit (counter becomes zero), the client closes its connection pool.
861 """
862 current_usage = await asyncio.shield(self._decrement_usage())
863 if current_usage == 0:
864 # This was the last active context, so disconnect the pool.
865 await asyncio.shield(self.aclose())
866
867 def __await__(self) -> Generator[Any, None, "RedisCluster"]:
868 return self.initialize().__await__()
869
870 _DEL_MESSAGE = "Unclosed RedisCluster client"
871
872 def __del__(
873 self,
874 _warn: Any = warnings.warn,
875 _grl: Any = asyncio.get_running_loop,
876 ) -> None:
877 if hasattr(self, "_initialize") and not self._initialize:
878 _warn(f"{self._DEL_MESSAGE} {self!r}", ResourceWarning, source=self)
879 try:
880 context = {"client": self, "message": self._DEL_MESSAGE}
881 _grl().call_exception_handler(context)
882 except RuntimeError:
883 pass
884
885 async def on_connect(self, connection: Connection) -> None:
886 await connection.on_connect()
887
888 # Sending READONLY command to server to configure connection as
889 # readonly. Since each cluster node may change its server type due
890 # to a failover, we should establish a READONLY connection
891 # regardless of the server type. If this is a primary connection,
892 # READONLY would not affect executing write commands.
893 await connection.send_command("READONLY")
894 if str_if_bytes(await connection.read_response()) != "OK":
895 raise ConnectionError("READONLY command failed")
896
897 def get_nodes(self) -> List["ClusterNode"]:
898 """Get all nodes of the cluster."""
899 return list(self.nodes_manager.nodes_cache.values())
900
901 def get_primaries(self) -> List["ClusterNode"]:
902 """Get the primary nodes of the cluster."""
903 return self.nodes_manager.get_nodes_by_server_type(PRIMARY)
904
905 def get_replicas(self) -> List["ClusterNode"]:
906 """Get the replica nodes of the cluster."""
907 return self.nodes_manager.get_nodes_by_server_type(REPLICA)
908
909 def get_random_node(self) -> "ClusterNode":
910 """Get a random node of the cluster."""
911 return random.choice(list(self.nodes_manager.nodes_cache.values()))
912
913 def get_default_node(self) -> "ClusterNode":
914 """Get the default node of the client."""
915 return self.nodes_manager.default_node
916
917 def set_default_node(self, node: "ClusterNode") -> None:
918 """
919 Set the default node of the client.
920
921 :raises DataError: if None is passed or node does not exist in cluster.
922 """
923 if not node or not self.get_node(node_name=node.name):
924 raise DataError("The requested node does not exist in the cluster.")
925
926 self.nodes_manager.default_node = node
927
928 def get_node(
929 self,
930 host: Optional[str] = None,
931 port: Optional[int] = None,
932 node_name: Optional[str] = None,
933 ) -> Optional["ClusterNode"]:
934 """Get node by (host, port) or node_name."""
935 return self.nodes_manager.get_node(host, port, node_name)
936
937 def get_node_from_key(
938 self, key: str, replica: bool = False
939 ) -> Optional["ClusterNode"]:
940 """
941 Get the cluster node corresponding to the provided key.
942
943 :param key:
944 :param replica:
945 | Indicates if a replica should be returned
946 |
947 None will returned if no replica holds this key
948
949 :raises SlotNotCoveredError: if the key is not covered by any slot.
950 """
951 slot = self.keyslot(key)
952 slot_cache = self.nodes_manager.slots_cache.get(slot)
953 if not slot_cache:
954 raise SlotNotCoveredError(f'Slot "{slot}" is not covered by the cluster.')
955
956 if replica:
957 if len(self.nodes_manager.slots_cache[slot]) < 2:
958 return None
959 node_idx = 1
960 else:
961 node_idx = 0
962
963 return slot_cache[node_idx]
964
965 async def get_keyless_target_node(self, command_name: str) -> "ClusterNode":
966 """
967 Returns the node a keyless command is routed to: a random node when replica reads
968 are enabled and the command is safe to serve from a replica, a random primary
969 otherwise.
970
971 A replicas-only ``load_balancing_strategy`` is honored by picking from the replicas
972 alone, so a strategy that asks for replicas cannot land on a primary here. The
973 strategy is not applied any further than that: the rest of it is an index into one
974 shard's node list and a round-robin counter kept per primary name, and a keyless
975 command has no shard to index - so the pick is uniform over the eligible nodes.
976
977 Falls back to the whole node set when the cluster has no replicas to pick from,
978 which is every primary. That is also the answer for the two strategies that
979 include the primary, and for ``read_from_replicas`` on its own, which is what this
980 method has returned for a replica-safe command since 7.1.0.
981 """
982 replica_safe = (
983 self.read_from_replicas or self.load_balancing_strategy is not None
984 ) and await self._is_replica_safe(command_name)
985 if replica_safe:
986 if self.load_balancing_strategy in _REPLICAS_ONLY_STRATEGIES:
987 replicas = self.get_replicas()
988 if replicas:
989 return random.choice(replicas)
990
991 return self.get_random_node()
992
993 return self.get_random_primary_node()
994
995 @deprecated_function(
996 version="8.2.0",
997 reason="Use get_keyless_target_node() instead.",
998 )
999 def get_random_primary_or_all_nodes(self, command_name: str) -> "ClusterNode":
1000 """
1001 Returns random primary or all nodes depends on READONLY mode.
1002
1003 Deprecated wrapper over :meth:`get_keyless_target_node`, kept so the name that
1004 has been public since 7.1.0 keeps working. Replica safety is resolved through the
1005 metadata resolver, which is awaitable, so the coroutine runs on its own event loop
1006 in a worker thread to keep this entry point synchronous.
1007
1008 .. warning::
1009
1010 **Do not extend this method, and do not call it from a running event loop.**
1011 Keeping the 7.1.0 signature synchronous costs a thread bridge
1012 (``_run_coroutine_in_thread``) that blocks the calling thread until the
1013 resolver answers. Two rules follow, and breaking either one stalls the caller
1014 permanently rather than raising:
1015
1016 1. **Override the right method.** A subclass that customizes keyless routing
1017 must override :meth:`get_keyless_target_node`, the coroutine the client
1018 actually awaits. Overriding this name changes nothing, because no code path
1019 inside the client calls it.
1020 2. **Keep a custom resolver synchronous in effect.** A caller-supplied
1021 ``AsyncMetadataResolver`` must implement ``is_replica_safe`` as a
1022 self-contained in-memory lookup. One that awaits work owned by the caller's
1023 loop - I/O over a pooled connection, say - waits on a loop this bridge has
1024 already blocked, and hangs.
1025
1026 That bridge is what confines this method to the deprecation window. It is sound
1027 for the resolvers the library ships, whose ``is_replica_safe`` suspends on nothing,
1028 and nothing inside the client reaches this path: keyless routing goes through
1029 :meth:`get_keyless_target_node`, which is awaited normally. So the hazard is
1030 confined to callers of this deprecated name.
1031 """
1032 return _run_coroutine_in_thread(self.get_keyless_target_node(command_name))
1033
1034 async def _is_replica_safe(self, command_name: str) -> bool:
1035 return await self._metadata_resolver.is_replica_safe(command_name)
1036
1037 def get_random_primary_node(self) -> "ClusterNode":
1038 """
1039 Returns a random primary node
1040 """
1041 return random.choice(self.get_primaries())
1042
1043 async def get_nodes_from_slot(self, command: str, *args):
1044 """
1045 Returns a list of nodes that hold the specified keys' slots.
1046 """
1047 # get the node that holds the key's slot
1048 replica_safe = (
1049 self.read_from_replicas or self.load_balancing_strategy is not None
1050 ) and await self._is_replica_safe(command)
1051 return [
1052 self.nodes_manager.get_node_from_slot(
1053 await self._determine_slot(command, *args),
1054 replica_safe,
1055 self.load_balancing_strategy if replica_safe else None,
1056 )
1057 ]
1058
1059 def get_special_nodes(self) -> Optional[list["ClusterNode"]]:
1060 """
1061 Returns a list of nodes for commands with a special policy.
1062 """
1063 if not self._aggregate_nodes:
1064 raise RedisClusterException(
1065 "Cannot execute FT.CURSOR commands without FT.AGGREGATE"
1066 )
1067
1068 return self._aggregate_nodes
1069
1070 def keyslot(self, key: EncodableT) -> int:
1071 """
1072 Find the keyslot for a given key.
1073
1074 See: https://redis.io/docs/manual/scaling/#redis-cluster-data-sharding
1075 """
1076 return key_slot(self.encoder.encode(key))
1077
1078 # HIMPORT orchestration (async mirror of redis.cluster.RedisCluster). The one
1079 # shared HImportRegistry is mutated once by PREPARE/DISCARD/DISCARDALL and applied
1080 # lazily per node; SET routes by key slot to the owning primary's ClusterNode.
1081 # See ``.agents/himport_client_support_spec.md``.
1082
1083 @property
1084 def himport_registry(self) -> HImportRegistry:
1085 """The cluster-wide HIMPORT fieldset registry (empty if none was declared).
1086
1087 Read-only: the registry is mutated only through the HIMPORT command methods.
1088 """
1089 return self._himport_registry
1090
1091 @experimental_method()
1092 async def himport_prepare(
1093 self, fieldset_name: str, fields: Iterable[FieldT]
1094 ) -> bool:
1095 """Declare an HIMPORT fieldset cluster-wide (shared registry, applied lazily)."""
1096 await self.initialize()
1097 self._himport_registry.prepare(fieldset_name, fields)
1098 return True
1099
1100 @experimental_method()
1101 async def himport_discard(self, fieldset_name: str) -> int:
1102 """Remove an HIMPORT fieldset cluster-wide (shared registry, applied lazily)."""
1103 await self.initialize()
1104 return 1 if self._himport_registry.discard(fieldset_name) else 0
1105
1106 @experimental_method()
1107 async def himport_discard_all(self) -> int:
1108 """Remove all HIMPORT fieldsets cluster-wide (shared registry, applied lazily)."""
1109 await self.initialize()
1110 return self._himport_registry.discard_all()
1111
1112 def get_encoder(self) -> Encoder:
1113 """Get the encoder object of the client."""
1114 return self.encoder
1115
1116 def get_connection_kwargs(self) -> Dict[str, Optional[Any]]:
1117 """Get the kwargs passed to :class:`~redis.asyncio.connection.Connection`."""
1118 return self.connection_kwargs
1119
1120 def set_retry(self, retry: Retry) -> None:
1121 self.retry = retry
1122
1123 def set_response_callback(self, command: str, callback: ResponseCallbackT) -> None:
1124 """Set a custom response callback."""
1125 self.response_callbacks[command] = callback
1126
1127 async def _resolve_command_policies(
1128 self, *args: Any, target_nodes_specified: bool = False
1129 ) -> Tuple[str, Union[CommandPolicies, CommandMetadata]]:
1130 """
1131 Resolves the policies a command routes and aggregates by.
1132
1133 Returns the name the policies were decided by along with the record, because the
1134 name a command is known by is not always ``args[0]``: a container command is
1135 keyed by both of its words, and the flag tables are keyed in upper case. Callers
1136 that look the command up elsewhere - result callbacks, observability - use the
1137 returned name so they agree with the routing decision.
1138
1139 The name is normalized before any branch, so one command answers with one name
1140 however it got here. The result callbacks are keyed in upper case, so a name that
1141 kept the caller's spelling on only some paths would fire them on only some paths -
1142 ``execute_command("dbsize")`` would be summed and ``execute_command("dbsize",
1143 target_nodes=...)`` would not.
1144
1145 First choice is the policy resolver. When it does not know the command, the
1146 fallbacks are, in order: nothing to route at all because the caller named its
1147 targets, the command's ``COMMAND_FLAGS`` entry, and finally whether the command
1148 carries a key.
1149 """
1150 command = args[0].upper()
1151 if len(args) >= 2 and f"{args[0]} {args[1]}".upper() in self.command_flags:
1152 command = f"{args[0]} {args[1]}".upper()
1153
1154 if target_nodes_specified:
1155 # The caller named its targets, so nothing is routed from here - and the
1156 # command's own aggregation must not apply either. A response policy resolved
1157 # for the whole cluster (``ONE_SUCCEEDED``, an ``AGG_*``) would short-circuit
1158 # the loop over the nodes the caller picked, or fold their replies into one.
1159 # Answer with the record that aggregates nothing, and skip the resolver: with
1160 # the targets given, neither of its answers is used.
1161 return command, _DEFAULT_KEYLESS_METADATA
1162
1163 policies = await self._policy_resolver.resolve(args[0].lower())
1164 if policies:
1165 return command, policies
1166
1167 command_flag = self.command_flags.get(command)
1168 if command_flag:
1169 if command_flag in self._command_flags_mapping:
1170 return command, _METADATA_BY_REQUEST_POLICY[
1171 self._command_flags_mapping[command_flag]
1172 ]
1173 return command, _DEFAULT_KEYLESS_METADATA
1174
1175 # Unflagged and unresolved, so the command routes by its key. Without a default
1176 # node the topology is not known yet and there is no slot to route by.
1177 if not self.get_default_node():
1178 return command, _DEFAULT_KEYLESS_METADATA
1179
1180 slot = await self._determine_slot(*args)
1181 if slot is None:
1182 return command, _DEFAULT_KEYLESS_METADATA
1183
1184 return command, _DEFAULT_KEYED_METADATA
1185
1186 async def _determine_nodes(
1187 self,
1188 command: str,
1189 *args: Any,
1190 request_policy: Optional[RequestPolicy] = None,
1191 node_flag: Optional[str] = None,
1192 ) -> List["ClusterNode"]:
1193 # Determine which nodes should be executed the command on.
1194 # Returns a list of target nodes.
1195 # The caller resolves the command's own policy - see
1196 # ``_resolve_command_policies`` - so the only decision left here is an explicit
1197 # node flag, which overrides it.
1198 if node_flag and self._is_node_flag(node_flag):
1199 if node_flag in self._command_flags_mapping:
1200 request_policy = self._command_flags_mapping[node_flag]
1201
1202 if request_policy is None:
1203 raise RedisClusterException(
1204 f"No targets were found to execute {command} command on"
1205 )
1206
1207 policy_callback = self._policies_callback_mapping[request_policy]
1208
1209 if request_policy == RequestPolicy.DEFAULT_KEYED:
1210 nodes = await policy_callback(command, *args)
1211 elif request_policy == RequestPolicy.DEFAULT_KEYLESS:
1212 nodes = [await policy_callback(command)]
1213 else:
1214 nodes = policy_callback()
1215
1216 if command.lower() == "ft.aggregate":
1217 self._aggregate_nodes = nodes
1218
1219 return nodes
1220
1221 async def _determine_slot(self, command: str, *args: Any) -> int:
1222 if self.command_flags.get(command.upper()) == SLOT_ID:
1223 # The command contains the slot ID
1224 return int(args[0])
1225
1226 # Get the keys in the command
1227
1228 # EVAL and EVALSHA are common enough that it's wasteful to go to the
1229 # redis server to parse the keys. Besides, there is a bug in redis<7.0
1230 # where `self._get_command_keys()` fails anyway. So, we special case
1231 # EVAL/EVALSHA.
1232 # - issue: https://github.com/redis/redis/issues/9493
1233 # - fix: https://github.com/redis/redis/pull/9733
1234 if command.upper() in ("EVAL", "EVALSHA"):
1235 # command syntax: EVAL "script body" num_keys ...
1236 if len(args) < 2:
1237 raise RedisClusterException(
1238 f"Invalid args in command: {command, *args}"
1239 )
1240 keys = args[2 : 2 + int(args[1])]
1241 # if there are 0 keys, that means the script can be run on any node
1242 # so we can just return a random slot
1243 if not keys:
1244 return random.randrange(0, REDIS_CLUSTER_HASH_SLOTS)
1245 else:
1246 keys = await self.commands_parser.get_keys(command, *args)
1247 if not keys:
1248 # FCALL can call a function with 0 keys, that means the function
1249 # can be run on any node so we can just return a random slot
1250 if command.upper() in ("FCALL", "FCALL_RO"):
1251 return random.randrange(0, REDIS_CLUSTER_HASH_SLOTS)
1252 raise RedisClusterException(
1253 "No way to dispatch this command to Redis Cluster. "
1254 "Missing key.\nYou can execute the command by specifying "
1255 f"target nodes.\nCommand: {args}"
1256 )
1257
1258 # single key command
1259 if len(keys) == 1:
1260 return self.keyslot(keys[0])
1261
1262 # multi-key command; we need to make sure all keys are mapped to
1263 # the same slot
1264 slots = {self.keyslot(key) for key in keys}
1265 if len(slots) != 1:
1266 raise RedisClusterException(
1267 f"{command} - all keys must map to the same key slot"
1268 )
1269
1270 return slots.pop()
1271
1272 def _is_node_flag(self, target_nodes: Any) -> bool:
1273 return isinstance(target_nodes, str) and target_nodes in self.node_flags
1274
1275 def _parse_target_nodes(self, target_nodes: Any) -> List["ClusterNode"]:
1276 if isinstance(target_nodes, list):
1277 nodes = target_nodes
1278 elif isinstance(target_nodes, ClusterNode):
1279 # Supports passing a single ClusterNode as a variable
1280 nodes = [target_nodes]
1281 elif isinstance(target_nodes, dict):
1282 # Supports dictionaries of the format {node_name: node}.
1283 # It enables to execute commands with multi nodes as follows:
1284 # rc.cluster_save_config(rc.get_primaries())
1285 nodes = list(target_nodes.values())
1286 else:
1287 raise TypeError(
1288 "target_nodes type can be one of the following: "
1289 "node_flag (PRIMARIES, REPLICAS, RANDOM, ALL_NODES),"
1290 "ClusterNode, list<ClusterNode>, or dict<any, ClusterNode>. "
1291 f"The passed type is {type(target_nodes)}"
1292 )
1293 return nodes
1294
1295 async def _record_error_metric(
1296 self,
1297 error: Exception,
1298 connection: Union[Connection, "ClusterNode"],
1299 is_internal: bool = True,
1300 retry_attempts: Optional[int] = None,
1301 ):
1302 """
1303 Records error count metric directly.
1304 Accepts either a Connection or ClusterNode object.
1305 """
1306 await record_error_count(
1307 server_address=connection.host,
1308 server_port=connection.port,
1309 network_peer_address=connection.host,
1310 network_peer_port=connection.port,
1311 error_type=error,
1312 retry_attempts=retry_attempts if retry_attempts is not None else 0,
1313 is_internal=is_internal,
1314 )
1315
1316 async def _record_command_metric(
1317 self,
1318 command_name: str,
1319 duration_seconds: float,
1320 connection: Union[Connection, "ClusterNode"],
1321 error: Optional[Exception] = None,
1322 ):
1323 """
1324 Records operation duration metric directly.
1325 Accepts either a Connection or ClusterNode object.
1326 """
1327 # Connection has db attribute, ClusterNode has connection_kwargs
1328 if hasattr(connection, "db"):
1329 db = connection.db
1330 else:
1331 db = connection.connection_kwargs.get("db", 0)
1332 await record_operation_duration(
1333 command_name=command_name,
1334 duration_seconds=duration_seconds,
1335 server_address=connection.host,
1336 server_port=connection.port,
1337 db_namespace=str(db) if db is not None else None,
1338 error=error,
1339 )
1340
1341 async def execute_command(self, *args: EncodableT, **kwargs: Any) -> Any:
1342 """
1343 Execute a raw command on the appropriate cluster node or target_nodes.
1344
1345 It will retry the command as specified by the retries property of
1346 the :attr:`retry` & then raise an exception.
1347
1348 :param args:
1349 | Raw command args
1350 :param kwargs:
1351
1352 - target_nodes: :attr:`NODE_FLAGS` or :class:`~.ClusterNode`
1353 or List[:class:`~.ClusterNode`] or Dict[Any, :class:`~.ClusterNode`]
1354 - Rest of the kwargs are passed to the Redis connection
1355
1356 :raises RedisClusterException: if target_nodes is not provided & the command
1357 can't be mapped to a slot
1358 """
1359 target_nodes = []
1360 target_nodes_specified = False
1361 retry_attempts = self.retry.get_retries()
1362
1363 passed_targets = kwargs.pop("target_nodes", None)
1364 if (
1365 passed_targets is not None
1366 and not self._is_node_flag(passed_targets)
1367 and not (
1368 isinstance(passed_targets, (list, dict, str)) and not passed_targets
1369 )
1370 ):
1371 target_nodes = self._parse_target_nodes(passed_targets)
1372 target_nodes_specified = True
1373 retry_attempts = 0
1374
1375 command, command_policies = await self._resolve_command_policies(
1376 *args, target_nodes_specified=target_nodes_specified
1377 )
1378
1379 # Add one for the first execution
1380 execute_attempts = 1 + retry_attempts
1381 failure_count = 0
1382
1383 # Start timing for observability
1384 start_time = time.monotonic()
1385 last_failed_node_name = None
1386
1387 for _ in range(execute_attempts):
1388 if self._initialize:
1389 await self.initialize(last_failed_node_name=last_failed_node_name)
1390 last_failed_node_name = None
1391 if (
1392 len(target_nodes) == 1
1393 and target_nodes[0] == self.get_default_node()
1394 ):
1395 # Replace the default cluster node
1396 self.replace_default_node()
1397 try:
1398 if not target_nodes_specified:
1399 # Determine the nodes to execute the command on
1400 target_nodes = await self._determine_nodes(
1401 *args,
1402 request_policy=command_policies.request_policy,
1403 node_flag=passed_targets,
1404 )
1405 if not target_nodes:
1406 raise RedisClusterException(
1407 f"No targets were found to execute {args} command on"
1408 )
1409
1410 if len(target_nodes) == 1:
1411 # Return the processed result
1412 ret = await self._execute_command(target_nodes[0], *args, **kwargs)
1413 if command in self.result_callbacks:
1414 ret = self.result_callbacks[command](
1415 command, {target_nodes[0].name: ret}, **kwargs
1416 )
1417 return self._policies_callback_mapping[
1418 command_policies.response_policy
1419 ](ret)
1420 else:
1421 keys = [node.name for node in target_nodes]
1422 values = await asyncio.gather(
1423 *(
1424 asyncio.create_task(
1425 self._execute_command(node, *args, **kwargs)
1426 )
1427 for node in target_nodes
1428 )
1429 )
1430 if command in self.result_callbacks:
1431 return self.result_callbacks[command](
1432 command, dict(zip(keys, values)), **kwargs
1433 )
1434 return self._policies_callback_mapping[
1435 command_policies.response_policy
1436 ](dict(zip(keys, values)))
1437 except Exception as e:
1438 if retry_attempts > 0 and type(e) in self.__class__.ERRORS_ALLOW_RETRY:
1439 # The nodes and slots cache were should be reinitialized.
1440 # Try again with the new cluster setup.
1441 retry_attempts -= 1
1442 failure_count += 1
1443 last_failed_node_name = getattr(e, "last_failed_node_name", None)
1444
1445 if hasattr(e, "connection"):
1446 # ``args[0]``, not the resolved ``command``: every other metric in
1447 # this class - including the per-command one ``_execute_command``
1448 # records on success - names the command the caller spelled. Using
1449 # the routed name only here would report one command under two
1450 # names depending on whether it was retried.
1451 await self._record_command_metric(
1452 command_name=args[0],
1453 duration_seconds=time.monotonic() - start_time,
1454 connection=e.connection,
1455 error=e,
1456 )
1457 await self._record_error_metric(
1458 error=e,
1459 connection=e.connection,
1460 retry_attempts=failure_count,
1461 )
1462 continue
1463 else:
1464 # raise the exception
1465 if hasattr(e, "connection"):
1466 await self._record_error_metric(
1467 error=e,
1468 connection=e.connection,
1469 retry_attempts=failure_count,
1470 is_internal=False,
1471 )
1472 raise e
1473
1474 async def _execute_command(
1475 self, target_node: "ClusterNode", *args: Union[KeyT, EncodableT], **kwargs: Any
1476 ) -> Any:
1477 asking = moved = False
1478 redirect_addr = None
1479 ttl = self.RedisClusterRequestTTL
1480 command = args[0]
1481 start_time = time.monotonic()
1482
1483 while ttl > 0:
1484 ttl -= 1
1485 ask_himport = False
1486 try:
1487 if asking:
1488 target_node = self.get_node(node_name=redirect_addr)
1489 if parse_himport_set_args(args) is not None:
1490 # ASKING must sit on the same connection as the SET,
1491 # immediately before it. HIMPORT SET's own executor folds
1492 # ASKING into the SET's packed write after the session setup,
1493 # so don't send it here as a separately pooled command.
1494 ask_himport = True
1495 else:
1496 await target_node.execute_command("ASKING")
1497 asking = False
1498 elif moved:
1499 # MOVED occurred and the slots cache was updated,
1500 # refresh the target node
1501 slot = await self._determine_slot(*args)
1502 replica_safe = (
1503 self.read_from_replicas
1504 or self.load_balancing_strategy is not None
1505 ) and await self._is_replica_safe(args[0])
1506 target_node = self.nodes_manager.get_node_from_slot(
1507 slot,
1508 replica_safe,
1509 self.load_balancing_strategy if replica_safe else None,
1510 )
1511 moved = False
1512
1513 response = await target_node.execute_command(
1514 *args, asking=ask_himport, **kwargs
1515 )
1516 await self._record_command_metric(
1517 command_name=command,
1518 duration_seconds=time.monotonic() - start_time,
1519 connection=target_node,
1520 )
1521 return response
1522 except BusyLoadingError as e:
1523 e.connection = target_node
1524 await self._record_command_metric(
1525 command_name=command,
1526 duration_seconds=time.monotonic() - start_time,
1527 connection=target_node,
1528 error=e,
1529 )
1530 raise
1531 except MaxConnectionsError as e:
1532 # MaxConnectionsError indicates client-side resource exhaustion
1533 # (too many connections in the pool), not a node failure.
1534 # Don't treat this as a node failure - just re-raise the error
1535 # without reinitializing the cluster.
1536 e.connection = target_node
1537 await self._record_command_metric(
1538 command_name=command,
1539 duration_seconds=time.monotonic() - start_time,
1540 connection=target_node,
1541 error=e,
1542 )
1543 raise
1544 except (ConnectionError, TimeoutError) as e:
1545 # Connection retries are being handled in the node's
1546 # Retry object.
1547 # Mark active connections for reconnect and disconnect free ones
1548 # This handles connection state (like READONLY) that may be stale
1549 target_node.update_active_connections_for_reconnect()
1550 await target_node.disconnect_free_connections()
1551
1552 # Move the failed node to the end of the cached nodes list
1553 # so it's tried last during reinitialization
1554 self.nodes_manager.move_node_to_end_of_cached_nodes(target_node.name)
1555 e.last_failed_node_name = target_node.name
1556
1557 # Signal that reinitialization is needed
1558 # The retry loop will handle initialize() AND replace_default_node()
1559 self._initialize = True
1560 e.connection = target_node
1561 await self._record_command_metric(
1562 command_name=command,
1563 duration_seconds=time.monotonic() - start_time,
1564 connection=target_node,
1565 error=e,
1566 )
1567 raise
1568 except (ClusterDownError, SlotNotCoveredError) as e:
1569 # ClusterDownError can occur during a failover and to get
1570 # self-healed, we will try to reinitialize the cluster layout
1571 # and retry executing the command
1572
1573 # SlotNotCoveredError can occur when the cluster is not fully
1574 # initialized or can be temporary issue.
1575 # We will try to reinitialize the cluster topology
1576 # and retry executing the command
1577
1578 await self.aclose()
1579 await asyncio.sleep(0.25)
1580 e.connection = target_node
1581 await self._record_command_metric(
1582 command_name=command,
1583 duration_seconds=time.monotonic() - start_time,
1584 connection=target_node,
1585 error=e,
1586 )
1587 raise
1588 except MovedError as e:
1589 # First, we will try to patch the slots/nodes cache with the
1590 # redirected node output and try again. If MovedError exceeds
1591 # 'reinitialize_steps' number of times, we will force
1592 # reinitializing the tables, and then try again.
1593 # 'reinitialize_steps' counter will increase faster when
1594 # the same client object is shared between multiple threads. To
1595 # reduce the frequency you can set this variable in the
1596 # RedisCluster constructor.
1597 self.reinitialize_counter += 1
1598 if (
1599 self.reinitialize_steps
1600 and self.reinitialize_counter % self.reinitialize_steps == 0
1601 ):
1602 await self.aclose()
1603 await self.initialize(
1604 additional_startup_nodes_info=[(e.host, e.port)]
1605 )
1606 # Reset the counter
1607 self.reinitialize_counter = 0
1608 else:
1609 await self.nodes_manager.move_slot(e)
1610 moved = True
1611 await self._record_command_metric(
1612 command_name=command,
1613 duration_seconds=time.monotonic() - start_time,
1614 connection=target_node,
1615 error=e,
1616 )
1617 await self._record_error_metric(
1618 error=e,
1619 connection=target_node,
1620 )
1621 except AskError as e:
1622 redirect_addr = get_node_name(host=e.host, port=e.port)
1623 asking = True
1624 await self._record_command_metric(
1625 command_name=command,
1626 duration_seconds=time.monotonic() - start_time,
1627 connection=target_node,
1628 error=e,
1629 )
1630 await self._record_error_metric(
1631 error=e,
1632 connection=target_node,
1633 )
1634 except TryAgainError as e:
1635 if ttl < self.RedisClusterRequestTTL / 2:
1636 await asyncio.sleep(0.05)
1637 await self._record_command_metric(
1638 command_name=command,
1639 duration_seconds=time.monotonic() - start_time,
1640 connection=target_node,
1641 error=e,
1642 )
1643 await self._record_error_metric(
1644 error=e,
1645 connection=target_node,
1646 )
1647 except ResponseError as e:
1648 e.connection = target_node
1649 await self._record_command_metric(
1650 command_name=command,
1651 duration_seconds=time.monotonic() - start_time,
1652 connection=target_node,
1653 error=e,
1654 )
1655 raise
1656 except Exception as e:
1657 e.connection = target_node
1658 await self._record_command_metric(
1659 command_name=command,
1660 duration_seconds=time.monotonic() - start_time,
1661 connection=target_node,
1662 error=e,
1663 )
1664 raise
1665
1666 e = ClusterError("TTL exhausted.")
1667 e.connection = target_node
1668 await self._record_command_metric(
1669 command_name=command,
1670 duration_seconds=time.monotonic() - start_time,
1671 connection=target_node,
1672 error=e,
1673 )
1674 raise e
1675
1676 def pipeline(
1677 self, transaction: Optional[Any] = None, shard_hint: Optional[Any] = None
1678 ) -> "ClusterPipeline":
1679 """
1680 Create & return a new :class:`~.ClusterPipeline` object.
1681
1682 Cluster implementation of pipeline does not support transaction or shard_hint.
1683
1684 :raises RedisClusterException: if transaction or shard_hint are truthy values
1685 """
1686 if shard_hint:
1687 raise RedisClusterException("shard_hint is deprecated in cluster mode")
1688
1689 return ClusterPipeline(self, transaction)
1690
1691 def pubsub(
1692 self,
1693 node: Optional["ClusterNode"] = None,
1694 host: Optional[str] = None,
1695 port: Optional[int] = None,
1696 **kwargs: Any,
1697 ) -> "ClusterPubSub":
1698 """
1699 Create and return a ClusterPubSub instance.
1700
1701 Allows passing a ClusterNode, or host&port, to get a pubsub instance
1702 connected to the specified node
1703
1704 :param node: ClusterNode to connect to
1705 :param host: Host of the node to connect to
1706 :param port: Port of the node to connect to
1707 :param kwargs: Additional keyword arguments
1708 :return: ClusterPubSub instance
1709 """
1710 return ClusterPubSub(self, node=node, host=host, port=port, **kwargs)
1711
1712 def keyspace_notifications(
1713 self,
1714 key_prefix: Union[str, bytes, None] = None,
1715 ignore_subscribe_messages: bool = True,
1716 ) -> "AsyncClusterKeyspaceNotifications":
1717 """
1718 Return an
1719 :class:`~redis.asyncio.keyspace_notifications.AsyncClusterKeyspaceNotifications`
1720 object for subscribing to keyspace and keyevent notifications across
1721 all primary nodes in the cluster.
1722
1723 Note: Keyspace notifications must be enabled on all Redis cluster nodes
1724 via the ``notify-keyspace-events`` configuration option.
1725
1726 Args:
1727 key_prefix: Optional prefix to filter and strip from keys in
1728 notifications.
1729 ignore_subscribe_messages: If True, subscribe/unsubscribe
1730 confirmations are not returned by
1731 get_message/listen.
1732 """
1733 from redis.asyncio.keyspace_notifications import (
1734 AsyncClusterKeyspaceNotifications,
1735 )
1736
1737 return AsyncClusterKeyspaceNotifications(
1738 self,
1739 key_prefix=key_prefix,
1740 ignore_subscribe_messages=ignore_subscribe_messages,
1741 )
1742
1743 def lock(
1744 self,
1745 name: KeyT,
1746 timeout: Optional[float] = None,
1747 sleep: float = 0.1,
1748 blocking: bool = True,
1749 blocking_timeout: Optional[float] = None,
1750 lock_class: Optional[Type[Lock]] = None,
1751 thread_local: bool = True,
1752 raise_on_release_error: bool = True,
1753 ) -> Lock:
1754 """
1755 Return a new Lock object using key ``name`` that mimics
1756 the behavior of threading.Lock.
1757
1758 If specified, ``timeout`` indicates a maximum life for the lock.
1759 By default, it will remain locked until release() is called.
1760
1761 ``sleep`` indicates the amount of time to sleep per loop iteration
1762 when the lock is in blocking mode and another client is currently
1763 holding the lock.
1764
1765 ``blocking`` indicates whether calling ``acquire`` should block until
1766 the lock has been acquired or to fail immediately, causing ``acquire``
1767 to return False and the lock not being acquired. Defaults to True.
1768 Note this value can be overridden by passing a ``blocking``
1769 argument to ``acquire``.
1770
1771 ``blocking_timeout`` indicates the maximum amount of time in seconds to
1772 spend trying to acquire the lock. A value of ``None`` indicates
1773 continue trying forever. ``blocking_timeout`` can be specified as a
1774 float or integer, both representing the number of seconds to wait.
1775
1776 ``lock_class`` forces the specified lock implementation. Note that as
1777 of redis-py 3.0, the only lock class we implement is ``Lock`` (which is
1778 a Lua-based lock). So, it's unlikely you'll need this parameter, unless
1779 you have created your own custom lock class.
1780
1781 ``thread_local`` indicates whether the lock token is placed in
1782 thread-local storage. By default, the token is placed in thread local
1783 storage so that a thread only sees its token, not a token set by
1784 another thread. Consider the following timeline:
1785
1786 time: 0, thread-1 acquires `my-lock`, with a timeout of 5 seconds.
1787 thread-1 sets the token to "abc"
1788 time: 1, thread-2 blocks trying to acquire `my-lock` using the
1789 Lock instance.
1790 time: 5, thread-1 has not yet completed. redis expires the lock
1791 key.
1792 time: 5, thread-2 acquired `my-lock` now that it's available.
1793 thread-2 sets the token to "xyz"
1794 time: 6, thread-1 finishes its work and calls release(). if the
1795 token is *not* stored in thread local storage, then
1796 thread-1 would see the token value as "xyz" and would be
1797 able to successfully release the thread-2's lock.
1798
1799 ``raise_on_release_error`` indicates whether to raise an exception when
1800 the lock is no longer owned when exiting the context manager. By default,
1801 this is True, meaning an exception will be raised. If False, the warning
1802 will be logged and the exception will be suppressed.
1803
1804 In some use cases it's necessary to disable thread local storage. For
1805 example, if you have code where one thread acquires a lock and passes
1806 that lock instance to a worker thread to release later. If thread
1807 local storage isn't disabled in this case, the worker thread won't see
1808 the token set by the thread that acquired the lock. Our assumption
1809 is that these cases aren't common and as such default to using
1810 thread local storage."""
1811 if lock_class is None:
1812 lock_class = Lock
1813 return lock_class(
1814 self,
1815 name,
1816 timeout=timeout,
1817 sleep=sleep,
1818 blocking=blocking,
1819 blocking_timeout=blocking_timeout,
1820 thread_local=thread_local,
1821 raise_on_release_error=raise_on_release_error,
1822 )
1823
1824 async def transaction(
1825 self, func: Coroutine[None, "ClusterPipeline", Any], *watches, **kwargs
1826 ):
1827 """
1828 Convenience method for executing the callable `func` as a transaction
1829 while watching all keys specified in `watches`. The 'func' callable
1830 should expect a single argument which is a Pipeline object.
1831 """
1832 shard_hint = kwargs.pop("shard_hint", None)
1833 value_from_callable = kwargs.pop("value_from_callable", False)
1834 watch_delay = kwargs.pop("watch_delay", None)
1835 async with self.pipeline(True, shard_hint) as pipe:
1836 while True:
1837 try:
1838 if watches:
1839 await pipe.watch(*watches)
1840 func_value = await func(pipe)
1841 exec_value = await pipe.execute()
1842 return func_value if value_from_callable else exec_value
1843 except WatchError:
1844 if watch_delay is not None and watch_delay > 0:
1845 time.sleep(watch_delay)
1846 continue
1847
1848
1849class ClusterNode:
1850 """
1851 Create a new ClusterNode.
1852
1853 Each ClusterNode manages multiple :class:`~redis.asyncio.connection.Connection`
1854 objects for the (host, port).
1855 """
1856
1857 __slots__ = (
1858 "_background_tasks",
1859 "_connections",
1860 "_free",
1861 "_lock",
1862 "_event_dispatcher",
1863 "connection_class",
1864 "connection_kwargs",
1865 "host",
1866 "max_connections",
1867 "name",
1868 "port",
1869 "response_callbacks",
1870 "server_type",
1871 )
1872
1873 def __init__(
1874 self,
1875 host: str,
1876 port: Union[str, int],
1877 server_type: Optional[str] = None,
1878 *,
1879 max_connections: int = 100,
1880 connection_class: Type[Connection] = Connection,
1881 **connection_kwargs: Any,
1882 ) -> None:
1883 if host == "localhost":
1884 host = socket.gethostbyname(host)
1885
1886 connection_kwargs["host"] = host
1887 connection_kwargs["port"] = port
1888 self.host = host
1889 self.port = port
1890 self.name = get_node_name(host, port)
1891 self.server_type = server_type
1892
1893 self.max_connections = max_connections
1894 self.connection_class = connection_class
1895 self.connection_kwargs = connection_kwargs
1896 self.response_callbacks = connection_kwargs.pop("response_callbacks", {})
1897
1898 self._connections: List[Connection] = []
1899 self._free: Deque[Connection] = collections.deque(maxlen=self.max_connections)
1900 self._background_tasks: Set[asyncio.Task] = set()
1901 self._event_dispatcher = self.connection_kwargs.get("event_dispatcher", None)
1902 if self._event_dispatcher is None:
1903 self._event_dispatcher = EventDispatcher()
1904
1905 def __repr__(self) -> str:
1906 return (
1907 f"[host={self.host}, port={self.port}, "
1908 f"name={self.name}, server_type={self.server_type}]"
1909 )
1910
1911 def __eq__(self, obj: Any) -> bool:
1912 return isinstance(obj, ClusterNode) and obj.name == self.name
1913
1914 def __hash__(self) -> int:
1915 return hash(self.name)
1916
1917 _DEL_MESSAGE = "Unclosed ClusterNode object"
1918
1919 def __del__(
1920 self,
1921 _warn: Any = warnings.warn,
1922 _grl: Any = asyncio.get_running_loop,
1923 ) -> None:
1924 for connection in self._connections:
1925 if connection.is_connected:
1926 _warn(f"{self._DEL_MESSAGE} {self!r}", ResourceWarning, source=self)
1927
1928 try:
1929 context = {"client": self, "message": self._DEL_MESSAGE}
1930 _grl().call_exception_handler(context)
1931 except RuntimeError:
1932 pass
1933 break
1934
1935 async def disconnect(self) -> None:
1936 ret = await asyncio.gather(
1937 *(
1938 asyncio.create_task(connection.disconnect())
1939 for connection in self._connections
1940 ),
1941 return_exceptions=True,
1942 )
1943 exc = next((res for res in ret if isinstance(res, Exception)), None)
1944 if exc:
1945 raise exc
1946
1947 def acquire_connection(self) -> Connection:
1948 try:
1949 return self._free.popleft()
1950 except IndexError:
1951 if len(self._connections) < self.max_connections:
1952 # We are configuring the connection pool not to retry
1953 # connections on lower level clients to avoid retrying
1954 # connections to nodes that are not reachable
1955 # and to avoid blocking the connection pool.
1956 # The only error that will have some handling in the lower
1957 # level clients is ConnectionError which will trigger disconnection
1958 # of the socket.
1959 # The retries will be handled on cluster client level
1960 # where we will have proper handling of the cluster topology
1961 retry = Retry(
1962 backoff=NoBackoff(),
1963 retries=0,
1964 supported_errors=(ConnectionError,),
1965 )
1966 connection_kwargs = self.connection_kwargs.copy()
1967 connection_kwargs["retry"] = retry
1968 connection = self.connection_class(**connection_kwargs)
1969 self._connections.append(connection)
1970 return connection
1971
1972 raise MaxConnectionsError()
1973
1974 async def disconnect_if_needed(self, connection: Connection) -> None:
1975 """
1976 Disconnect a connection if it's marked for reconnect.
1977 This implements lazy disconnection to avoid race conditions.
1978 The connection will auto-reconnect on next use.
1979 """
1980 if connection.should_reconnect():
1981 # Render the connection before disconnecting: extract_connection_details()
1982 # reads the local port and the in-flight read deadline off the transport, so
1983 # after disconnect() it can only report "not connected". This line is what
1984 # records the maintenance state and relaxed timeout at the moment they are
1985 # discarded, so a maintenance-driven recycle is attributable in the logs.
1986 if logger.isEnabledFor(logging.DEBUG):
1987 logger.debug(
1988 "Disconnecting acquired connection marked for reconnect: "
1989 f"{connection}, {connection.extract_connection_details()}"
1990 )
1991 await connection.disconnect()
1992
1993 def release(self, connection: Connection) -> None:
1994 """
1995 Release connection back to free queue.
1996 If a connected connection is marked for reconnect, disconnect it before
1997 returning it to the free queue. An already-closed connection can be
1998 returned immediately after clearing its reconnect flag.
1999 """
2000 if connection.should_reconnect():
2001 if connection.is_connected:
2002 # Logged here rather than in _disconnect_and_release: that runs as a
2003 # task after the fact, by which point extract_connection_details()
2004 # may already have lost the transport it reads from.
2005 if logger.isEnabledFor(logging.DEBUG):
2006 logger.debug(
2007 "Disconnecting released connection marked for reconnect: "
2008 f"{connection}, {connection.extract_connection_details()}"
2009 )
2010 task = asyncio.create_task(self._disconnect_and_release(connection))
2011 self._background_tasks.add(task)
2012 task.add_done_callback(self._background_tasks.discard)
2013 return
2014 # It may have been re-marked while its own disconnect was in progress.
2015 connection.reset_should_reconnect()
2016 self._free.append(connection)
2017
2018 async def _disconnect_and_release(self, connection: Connection) -> None:
2019 try:
2020 await connection.disconnect()
2021 except Exception as exc:
2022 logger.debug(
2023 "disconnecting released cluster connection failed: %r",
2024 exc,
2025 exc_info=True,
2026 )
2027 try:
2028 self._connections.remove(connection)
2029 except ValueError:
2030 pass
2031 return
2032
2033 self._free.append(connection)
2034
2035 def get_encoder(self) -> Encoder:
2036 """Return an :class:`Encoder` derived from this node's connection kwargs."""
2037 kwargs = self.connection_kwargs
2038 encoder_class = kwargs.get("encoder_class", Encoder)
2039 return encoder_class(
2040 encoding=kwargs.get("encoding", "utf-8"),
2041 encoding_errors=kwargs.get("encoding_errors", "strict"),
2042 decode_responses=kwargs.get("decode_responses", False),
2043 )
2044
2045 def update_active_connections_for_reconnect(self) -> None:
2046 """
2047 Mark all in-use (active) connections for reconnect.
2048 In-use connections are those in _connections but not currently in _free.
2049 They will be disconnected after their current operation completes.
2050 """
2051 free_set = set(self._free)
2052 for connection in self._connections:
2053 if connection not in free_set:
2054 connection.mark_for_reconnect()
2055
2056 async def disconnect_free_connections(self) -> None:
2057 """
2058 Disconnect all free/idle connections in the pool.
2059 This is useful after topology changes (e.g., failover) to clear
2060 stale connection state like READONLY mode.
2061 The connections remain in the pool and will reconnect on next use.
2062 """
2063 if self._free:
2064 # Take a snapshot to avoid issues if _free changes during await
2065 await asyncio.gather(
2066 *(connection.disconnect() for connection in tuple(self._free)),
2067 return_exceptions=True,
2068 )
2069
2070 async def parse_response(
2071 self, connection: Connection, command: str, **kwargs: Any
2072 ) -> Any:
2073 try:
2074 if NEVER_DECODE in kwargs:
2075 response = await connection.read_response(disable_decoding=True)
2076 kwargs.pop(NEVER_DECODE)
2077 else:
2078 response = await connection.read_response()
2079 except ResponseError:
2080 if EMPTY_RESPONSE in kwargs:
2081 return kwargs[EMPTY_RESPONSE]
2082 raise
2083
2084 if EMPTY_RESPONSE in kwargs:
2085 kwargs.pop(EMPTY_RESPONSE)
2086
2087 # Remove keys entry, it needs only for cache.
2088 kwargs.pop("keys", None)
2089
2090 # Return response
2091 if command in self.response_callbacks:
2092 return self.response_callbacks[command](response, **kwargs)
2093
2094 return response
2095
2096 async def execute_command(
2097 self, *args: Any, asking: bool = False, **kwargs: Any
2098 ) -> Any:
2099 # Acquire connection
2100 connection = self.acquire_connection()
2101 try:
2102 # Handle lazy disconnect for connections marked for reconnect
2103 await self.disconnect_if_needed(connection)
2104
2105 # HIMPORT SET is the one command whose wire form depends on
2106 # per-connection state: the fieldset must be PREPAREd on this
2107 # connection first, and any fieldset discarded since this connection
2108 # last reconciled must be dropped. Doing it here (rather than in
2109 # RedisCluster.himport_set) reuses the caller's full retry, MOVED/ASK
2110 # and disconnect-on-error handling for HIMPORT SET too.
2111 # This per-command branch in the hot dispatch path is deliberate and has
2112 # no cleaner alternative: this is the only seam where the concrete routed
2113 # connection is known, and connection-scoped session setup can only happen
2114 # once that connection is chosen. The overhead is one comparison per
2115 # command.
2116 # On an ASK redirect ``asking`` is passed here rather than sent as a
2117 # separate ASKING command so the allowance sits on this same connection,
2118 # folded into the SET's own write immediately before the SET.
2119 himport_set = parse_himport_set_args(args)
2120 if himport_set is not None:
2121 # HIMPORT SET in the joined ("HIMPORT SET", key, ...) or split
2122 # ("HIMPORT", "SET", key, ...) raw form; operands at the right
2123 # offsets. Too few operands returns None and falls through to the
2124 # normal send path below so the server returns its arity error
2125 # instead of a client-side IndexError.
2126 key, fieldset_name, values = himport_set
2127 return await self._himport_execute_set(
2128 connection, key, fieldset_name, values, asking=asking
2129 )
2130
2131 # Execute command
2132 await connection.send_packed_command(connection.pack_command(*args))
2133
2134 # Read response
2135 return await self.parse_response(connection, args[0], **kwargs)
2136 finally:
2137 try:
2138 await self.disconnect_if_needed(connection)
2139 finally:
2140 # Release connection
2141 self.release(connection)
2142
2143 async def _himport_reconcile_discards(self, conn: "Connection") -> None:
2144 """Delegate to the shared async HIMPORT executor."""
2145 return await _himport_exec.reconcile_discards(self, conn)
2146
2147 async def _himport_prepare_and_set(
2148 self,
2149 conn: "Connection",
2150 key: KeyT,
2151 fieldset_name: str,
2152 values: List,
2153 fieldset,
2154 asking: bool = False,
2155 ) -> Any:
2156 """Delegate to the shared async HIMPORT executor."""
2157 return await _himport_exec.prepare_and_set(
2158 self, conn, key, fieldset_name, values, fieldset, asking=asking
2159 )
2160
2161 async def _himport_execute_set(
2162 self,
2163 conn: "Connection",
2164 key: KeyT,
2165 fieldset_name: str,
2166 values: List,
2167 asking: bool = False,
2168 ) -> Any:
2169 """Delegate to the shared async HIMPORT executor."""
2170 return await _himport_exec.execute_set(
2171 self, conn, key, fieldset_name, values, asking=asking
2172 )
2173
2174 async def _himport_prepare_pipeline(
2175 self, conn: "Connection", commands: List["PipelineCommand"]
2176 ) -> None:
2177 """Delegate to the shared async HIMPORT executor."""
2178 await _himport_exec.prepare_pipeline(self, conn, [cmd.args for cmd in commands])
2179
2180 async def execute_pipeline(self, commands: List["PipelineCommand"]) -> bool:
2181 # Acquire connection
2182 connection = self.acquire_connection()
2183 try:
2184 # Handle lazy disconnect for connections marked for reconnect
2185 await self.disconnect_if_needed(connection)
2186
2187 # PREPARE fieldsets referenced by buffered HIMPORT SETs before the
2188 # batched write (it bypasses the per-command lazy prepare path).
2189 await self._himport_prepare_pipeline(connection, commands)
2190
2191 # Execute command
2192 await connection.send_packed_command(
2193 connection.pack_commands(cmd.args for cmd in commands)
2194 )
2195
2196 # Read responses
2197 ret = False
2198 for cmd in commands:
2199 try:
2200 cmd.result = await self.parse_response(
2201 connection, cmd.args[0], **cmd.kwargs
2202 )
2203 except Exception as e:
2204 cmd.result = e
2205 ret = True
2206
2207 return ret
2208 finally:
2209 try:
2210 await self.disconnect_if_needed(connection)
2211 finally:
2212 # Release connection
2213 self.release(connection)
2214
2215 async def re_auth_callback(self, token: TokenInterface):
2216 tmp_queue = collections.deque()
2217 while self._free:
2218 conn = self._free.popleft()
2219 await conn.retry.call_with_retry(
2220 lambda: conn.send_command(
2221 "AUTH", token.try_get("oid"), token.get_value()
2222 ),
2223 lambda error: self._mock(error),
2224 )
2225 await conn.retry.call_with_retry(
2226 lambda: conn.read_response(), lambda error: self._mock(error)
2227 )
2228 tmp_queue.append(conn)
2229
2230 while tmp_queue:
2231 conn = tmp_queue.popleft()
2232 self._free.append(conn)
2233
2234 async def _mock(self, error: RedisError):
2235 """
2236 Dummy functions, needs to be passed as error callback to retry object.
2237 :param error:
2238 :return:
2239 """
2240 pass
2241
2242
2243class NodesManager:
2244 __slots__ = (
2245 "_dynamic_startup_nodes",
2246 "_event_dispatcher",
2247 "_background_tasks",
2248 "connection_kwargs",
2249 "default_node",
2250 "nodes_cache",
2251 "_epoch",
2252 "read_load_balancer",
2253 "_initialize_lock",
2254 "require_full_coverage",
2255 "slots_cache",
2256 "startup_nodes",
2257 "address_remap",
2258 )
2259
2260 def __init__(
2261 self,
2262 startup_nodes: List["ClusterNode"],
2263 require_full_coverage: bool,
2264 connection_kwargs: Dict[str, Any],
2265 dynamic_startup_nodes: bool = True,
2266 address_remap: Optional[Callable[[Tuple[str, int]], Tuple[str, int]]] = None,
2267 event_dispatcher: Optional[EventDispatcher] = None,
2268 ) -> None:
2269 self.startup_nodes = {node.name: node for node in startup_nodes}
2270 self.require_full_coverage = require_full_coverage
2271 self.connection_kwargs = connection_kwargs
2272 self.address_remap = address_remap
2273
2274 self.default_node: "ClusterNode" = None
2275 self.nodes_cache: Dict[str, "ClusterNode"] = {}
2276 self.slots_cache: Dict[int, List["ClusterNode"]] = {}
2277 self._epoch: int = 0
2278 self.read_load_balancer = LoadBalancer()
2279 self._initialize_lock: asyncio.Lock = asyncio.Lock()
2280
2281 self._background_tasks: Set[asyncio.Task] = set()
2282 self._dynamic_startup_nodes: bool = dynamic_startup_nodes
2283 if event_dispatcher is None:
2284 self._event_dispatcher = EventDispatcher()
2285 else:
2286 self._event_dispatcher = event_dispatcher
2287
2288 def get_node(
2289 self,
2290 host: Optional[str] = None,
2291 port: Optional[int] = None,
2292 node_name: Optional[str] = None,
2293 ) -> Optional["ClusterNode"]:
2294 if host and port:
2295 # the user passed host and port
2296 if host == "localhost":
2297 host = socket.gethostbyname(host)
2298 return self.nodes_cache.get(get_node_name(host=host, port=port))
2299 elif node_name:
2300 return self.nodes_cache.get(node_name)
2301 else:
2302 raise DataError(
2303 "get_node requires one of the following: 1. node name 2. host and port"
2304 )
2305
2306 def set_nodes(
2307 self,
2308 old: Dict[str, "ClusterNode"],
2309 new: Dict[str, "ClusterNode"],
2310 remove_old: bool = False,
2311 ) -> None:
2312 if remove_old:
2313 for name in list(old.keys()):
2314 if name not in new:
2315 # Node is removed from cache before disconnect starts,
2316 # so it won't be found in lookups during disconnect
2317 # Mark active connections so in-flight commands can
2318 # finish, then disconnect them when their current
2319 # operation completes. Free connections can be
2320 # disconnected immediately.
2321 removed_node = old.pop(name)
2322 removed_node.update_active_connections_for_reconnect()
2323 task = asyncio.create_task(
2324 removed_node.disconnect_free_connections()
2325 )
2326 self._background_tasks.add(task)
2327 task.add_done_callback(self._background_tasks.discard)
2328
2329 for name, node in new.items():
2330 if name in old:
2331 # Preserve the existing node but mark ALL its connections for
2332 # reconnect on every topology refresh.
2333 #
2334 # Why recycle every preserved node's connections, not just the
2335 # ones whose slots/role changed?
2336 # set_nodes only sees the old vs new node dicts; it does not
2337 # track which specific nodes had slot-ownership or role changes
2338 # during this refresh. Rather than try to diff that (and risk
2339 # serving a connection whose cached routing/READONLY state is
2340 # now stale), we conservatively refresh every preserved node.
2341 # Reconnect is lazy and cheap, so the extra churn is acceptable
2342 # in exchange for never serving a stale connection after a
2343 # topology change.
2344 #
2345 # Why mark-for-reconnect instead of disconnecting here?
2346 # set_nodes is sync but disconnect_free_connections() is async,
2347 # so we cannot disconnect inline. Marking both in-use and free
2348 # connections for reconnect lets them be lazily disconnected on
2349 # next acquire via disconnect_if_needed(), which avoids races.
2350 #
2351 # TODO: Make this method async in the next major release to allow
2352 # immediate disconnection of free connections.
2353 existing_node = old[name]
2354 existing_node.server_type = node.server_type
2355 existing_node.update_active_connections_for_reconnect()
2356 for conn in existing_node._free:
2357 conn.mark_for_reconnect()
2358 continue
2359 # New node is detected and should be added to the pool
2360 old[name] = node
2361
2362 def move_node_to_end_of_cached_nodes(self, node_name: str) -> None:
2363 """
2364 Move a failing node to the end of startup_nodes and nodes_cache so it's
2365 tried last during reinitialization and when selecting the default node.
2366 If the node is not in the respective list, nothing is done.
2367 """
2368 # Move in startup_nodes
2369 if node_name in self.startup_nodes and len(self.startup_nodes) > 1:
2370 node = self.startup_nodes.pop(node_name)
2371 self.startup_nodes[node_name] = node # Re-insert at end
2372
2373 # Move in nodes_cache - this affects get_nodes_by_server_type ordering
2374 # which is used to select the default_node during initialize()
2375 if node_name in self.nodes_cache and len(self.nodes_cache) > 1:
2376 node = self.nodes_cache.pop(node_name)
2377 self.nodes_cache[node_name] = node # Re-insert at end
2378
2379 async def move_slot(self, e: AskError | MovedError):
2380 node_changed = False
2381 redirected_node = self.get_node(host=e.host, port=e.port)
2382 if redirected_node:
2383 # The node already exists
2384 if redirected_node.server_type != PRIMARY:
2385 # Update the node's server type
2386 redirected_node.server_type = PRIMARY
2387 else:
2388 # This is a new node, we will add it to the nodes cache
2389 redirected_node = ClusterNode(
2390 e.host, e.port, PRIMARY, **self.connection_kwargs
2391 )
2392 self.set_nodes(self.nodes_cache, {redirected_node.name: redirected_node})
2393 slot_nodes = self.slots_cache[e.slot_id]
2394 if redirected_node not in slot_nodes:
2395 # The new slot owner is a new server, or a server from a different
2396 # shard. We need to remove all current nodes from the slot's list
2397 # (including replications) and add just the new node.
2398 self.slots_cache[e.slot_id] = [redirected_node]
2399 node_changed = True
2400 elif redirected_node is not slot_nodes[0]:
2401 # The MOVED error resulted from a failover, and the new slot owner
2402 # had previously been a replica.
2403 old_primary = slot_nodes[0]
2404 # Update the old primary to be a replica and add it to the end of
2405 # the slot's node list
2406 old_primary.server_type = REPLICA
2407 slot_nodes.append(old_primary)
2408 # Remove the old replica, which is now a primary, from the slot's
2409 # node list
2410 slot_nodes.remove(redirected_node)
2411 # Override the old primary with the new one
2412 slot_nodes[0] = redirected_node
2413 if self.default_node == old_primary:
2414 # Update the default node with the new primary
2415 self.default_node = redirected_node
2416 node_changed = True
2417 # else: circular MOVED to current primary -> no-op
2418 # Dispatch so listeners can run shard-pubsub reconciliation; skipped on
2419 # the no-op branch to avoid needless walks under MOVED storms. A
2420 # listener must not break slots-cache refresh; log and continue so a
2421 # single buggy listener cannot starve the rest.
2422 if node_changed:
2423 try:
2424 await self._event_dispatcher.dispatch_async(
2425 AsyncAfterSlotsCacheRefreshEvent()
2426 )
2427 except Exception as exc:
2428 # Don't shadow the method parameter ``e``: ``except as`` binds
2429 # the listener exception in the function scope and ``del``s
2430 # the name on block exit (PEP 3134), which would also wipe
2431 # out the original AskError/MovedError parameter.
2432 logger.exception(
2433 "listener raised during slots-cache refresh: %s: %s",
2434 type(exc).__name__,
2435 exc,
2436 )
2437
2438 def get_node_from_slot(
2439 self,
2440 slot: int,
2441 read_from_replicas: bool = False,
2442 load_balancing_strategy=None,
2443 ) -> "ClusterNode":
2444 if read_from_replicas is True and load_balancing_strategy is None:
2445 load_balancing_strategy = LoadBalancingStrategy.ROUND_ROBIN
2446
2447 try:
2448 if len(self.slots_cache[slot]) > 1 and load_balancing_strategy:
2449 # get the server index using the strategy defined in load_balancing_strategy
2450 primary_name = self.slots_cache[slot][0].name
2451 node_idx = self.read_load_balancer.get_server_index(
2452 primary_name, len(self.slots_cache[slot]), load_balancing_strategy
2453 )
2454 return self.slots_cache[slot][node_idx]
2455 return self.slots_cache[slot][0]
2456 except (IndexError, KeyError, TypeError):
2457 raise SlotNotCoveredError(
2458 f'Slot "{slot}" not covered by the cluster. '
2459 f'"require_full_coverage={self.require_full_coverage}"'
2460 )
2461
2462 def get_nodes_by_server_type(self, server_type: str) -> List["ClusterNode"]:
2463 return [
2464 node
2465 for node in self.nodes_cache.values()
2466 if node.server_type == server_type
2467 ]
2468
2469 async def initialize(
2470 self,
2471 additional_startup_nodes_info: Optional[List[Tuple[str, int]]] = None,
2472 last_failed_node_name: Optional[str] = None,
2473 ) -> None:
2474 self.read_load_balancer.reset()
2475 tmp_nodes_cache: Dict[str, "ClusterNode"] = {}
2476 tmp_slots: Dict[int, List["ClusterNode"]] = {}
2477 disagreements = []
2478 startup_nodes_reachable = False
2479 fully_covered = False
2480 exception = None
2481 epoch = self._epoch
2482 if additional_startup_nodes_info is None:
2483 additional_startup_nodes_info = []
2484
2485 async with self._initialize_lock:
2486 if self._epoch != epoch:
2487 # another initialize call has already reinitialized the
2488 # nodes since we started waiting for the lock;
2489 # we don't need to do it again.
2490 return
2491
2492 # Copy to a list to prevent RuntimeError if self.startup_nodes
2493 # is modified during iteration, then shuffle the iteration order.
2494 startup_nodes = list(self.startup_nodes.values())
2495 deferred_failed_nodes = []
2496 if last_failed_node_name is not None:
2497 for index, node in enumerate(startup_nodes):
2498 if node.name == last_failed_node_name:
2499 deferred_failed_nodes.append(startup_nodes.pop(index))
2500 break
2501 if len(startup_nodes) > 1:
2502 # Vary which startup node is queried first so clients do not
2503 # all reinitialize through the same node.
2504 random.shuffle(startup_nodes)
2505 additional_startup_nodes = [
2506 ClusterNode(host, port, **self.connection_kwargs)
2507 for host, port in additional_startup_nodes_info
2508 ]
2509 if last_failed_node_name is not None:
2510 for index, node in enumerate(additional_startup_nodes):
2511 if node.name == last_failed_node_name:
2512 if not deferred_failed_nodes:
2513 deferred_failed_nodes.append(node)
2514 additional_startup_nodes.pop(index)
2515 break
2516 for startup_node in chain(
2517 startup_nodes,
2518 additional_startup_nodes,
2519 deferred_failed_nodes,
2520 ):
2521 try:
2522 # Make sure cluster mode is enabled on this node
2523 try:
2524 self._event_dispatcher.dispatch(
2525 AfterAsyncClusterInstantiationEvent(
2526 self.nodes_cache,
2527 self.connection_kwargs.get("credential_provider", None),
2528 )
2529 )
2530 if logger.isEnabledFor(logging.DEBUG):
2531 logger.debug(
2532 "Topology refresh: querying CLUSTER SLOTS on "
2533 f"{startup_node.name}"
2534 )
2535 cluster_slots = await startup_node.execute_command(
2536 "CLUSTER SLOTS"
2537 )
2538 except ResponseError:
2539 raise RedisClusterException(
2540 "Cluster mode is not enabled on this node"
2541 )
2542 startup_nodes_reachable = True
2543 except Exception as e:
2544 # Try the next startup node.
2545 # The exception is saved and raised only if we have no more nodes.
2546 if logger.isEnabledFor(logging.DEBUG):
2547 logger.debug(
2548 "Topology refresh: CLUSTER SLOTS failed on "
2549 f"{startup_node.name}: {type(e).__name__}: {e}"
2550 )
2551 exception = e
2552 continue
2553
2554 # CLUSTER SLOTS command results in the following output:
2555 # [[slot_section[from_slot,to_slot,master,replica1,...,replicaN]]]
2556 # where each node contains the following list: [IP, port, node_id]
2557 # Therefore, cluster_slots[0][2][0] will be the IP address of the
2558 # primary node of the first slot section.
2559 # If there's only one server in the cluster, its ``host`` is ''
2560 # Fix it to the host in startup_nodes
2561 if (
2562 len(cluster_slots) == 1
2563 and not cluster_slots[0][2][0]
2564 and len(self.startup_nodes) == 1
2565 ):
2566 cluster_slots[0][2][0] = startup_node.host
2567
2568 for slot in cluster_slots:
2569 for i in range(2, len(slot)):
2570 slot[i] = [str_if_bytes(val) for val in slot[i]]
2571 primary_node = slot[2]
2572 host = primary_node[0]
2573 if host == "":
2574 host = startup_node.host
2575 port = int(primary_node[1])
2576 host, port = self.remap_host_port(host, port)
2577
2578 nodes_for_slot = []
2579
2580 target_node = tmp_nodes_cache.get(get_node_name(host, port))
2581 if not target_node:
2582 target_node = ClusterNode(
2583 host, port, PRIMARY, **self.connection_kwargs
2584 )
2585 # add this node to the nodes cache
2586 tmp_nodes_cache[target_node.name] = target_node
2587 nodes_for_slot.append(target_node)
2588
2589 replica_nodes = slot[3:]
2590 for replica_node in replica_nodes:
2591 host = replica_node[0]
2592 port = replica_node[1]
2593 host, port = self.remap_host_port(host, port)
2594
2595 target_replica_node = tmp_nodes_cache.get(
2596 get_node_name(host, port)
2597 )
2598 if not target_replica_node:
2599 target_replica_node = ClusterNode(
2600 host, port, REPLICA, **self.connection_kwargs
2601 )
2602 # add this node to the nodes cache
2603 tmp_nodes_cache[target_replica_node.name] = target_replica_node
2604 nodes_for_slot.append(target_replica_node)
2605
2606 for i in range(int(slot[0]), int(slot[1]) + 1):
2607 if i not in tmp_slots:
2608 tmp_slots[i] = nodes_for_slot
2609 else:
2610 # Validate that 2 nodes want to use the same slot cache
2611 # setup
2612 tmp_slot = tmp_slots[i][0]
2613 if tmp_slot.name != target_node.name:
2614 disagreements.append(
2615 f"{tmp_slot.name} vs {target_node.name} on slot: {i}"
2616 )
2617
2618 if len(disagreements) > 5:
2619 raise RedisClusterException(
2620 f"startup_nodes could not agree on a valid "
2621 f"slots cache: {', '.join(disagreements)}"
2622 )
2623
2624 # Validate if all slots are covered or if we should try next startup node
2625 fully_covered = True
2626 for i in range(REDIS_CLUSTER_HASH_SLOTS):
2627 if i not in tmp_slots:
2628 fully_covered = False
2629 break
2630 if logger.isEnabledFor(logging.DEBUG):
2631 logger.debug(
2632 f"Topology refresh: CLUSTER SLOTS from {startup_node.name} "
2633 f"reported nodes {sorted(tmp_nodes_cache)}; "
2634 f"slots fully covered: {fully_covered}"
2635 )
2636 if fully_covered:
2637 break
2638
2639 if not startup_nodes_reachable:
2640 # The unreachable subtype is reserved for connectivity failures:
2641 # MultiDB registers it as retryable, so a deterministic
2642 # server/configuration error (e.g. cluster mode disabled or
2643 # invalid credentials - AuthenticationError and
2644 # AuthorizationError subclass ConnectionError but cannot be
2645 # repaired by a failover) must keep surfacing as a plain
2646 # RedisClusterException.
2647 if isinstance(
2648 exception, (ConnectionError, TimeoutError, OSError)
2649 ) and not isinstance(
2650 exception, (AuthenticationError, AuthorizationError)
2651 ):
2652 raise RedisClusterUnreachableError(
2653 f"Redis Cluster cannot be connected. Please provide at least "
2654 f"one reachable node: {str(exception)}"
2655 ) from exception
2656 raise RedisClusterException(
2657 f"Redis Cluster cannot be connected. Please provide at least "
2658 f"one reachable node: {str(exception)}"
2659 ) from exception
2660
2661 # Check if the slots are not fully covered
2662 if not fully_covered and self.require_full_coverage:
2663 # Despite the requirement that the slots be covered, there
2664 # isn't a full coverage
2665 raise RedisClusterException(
2666 f"All slots are not covered after query all startup_nodes. "
2667 f"{len(tmp_slots)} of {REDIS_CLUSTER_HASH_SLOTS} "
2668 f"covered..."
2669 )
2670
2671 # Set the tmp variables to the real variables
2672 self.set_nodes(self.nodes_cache, tmp_nodes_cache, remove_old=True)
2673 # tmp_slots was built from CLUSTER SLOTS responses and can contain
2674 # newly-created ClusterNode objects for nodes we already know about.
2675 # Rebuild the slots cache with the preserved nodes_cache instances
2676 # so existing per-node connection pools stay in use after refresh.
2677 # Keep the shared node-list-per-slot-range shape from tmp_slots to
2678 # avoid allocating a separate list for every slot.
2679 node_lists_by_id: Dict[int, List["ClusterNode"]] = {}
2680 new_slots_cache: Dict[int, List["ClusterNode"]] = {}
2681 for slot, nodes in tmp_slots.items():
2682 node_list_id = id(nodes)
2683 slot_nodes = node_lists_by_id.get(node_list_id)
2684 if slot_nodes is None:
2685 slot_nodes = [self.nodes_cache[node.name] for node in nodes]
2686 node_lists_by_id[node_list_id] = slot_nodes
2687 new_slots_cache[slot] = slot_nodes
2688 self.slots_cache = new_slots_cache
2689
2690 if self._dynamic_startup_nodes:
2691 # Populate the startup nodes with all discovered nodes
2692 self.set_nodes(self.startup_nodes, self.nodes_cache, remove_old=True)
2693
2694 # Set the default node
2695 self.default_node = self.get_nodes_by_server_type(PRIMARY)[0]
2696 self._epoch += 1
2697 # Dispatch so listeners (e.g. ClusterPubSub) can reconcile per-node
2698 # state after slot ownership may have changed. A listener must not
2699 # break slots-cache refresh; log and continue so a single buggy
2700 # listener cannot starve the rest.
2701 try:
2702 await self._event_dispatcher.dispatch_async(
2703 AsyncAfterSlotsCacheRefreshEvent()
2704 )
2705 except Exception as e:
2706 logger.exception(
2707 "listener raised during slots-cache refresh: %s: %s",
2708 type(e).__name__,
2709 e,
2710 )
2711
2712 async def aclose(self, attr: str = "nodes_cache") -> None:
2713 self.default_node = None
2714 await asyncio.gather(
2715 *(
2716 asyncio.create_task(node.disconnect())
2717 for node in getattr(self, attr).values()
2718 )
2719 )
2720
2721 def remap_host_port(self, host: str, port: int) -> Tuple[str, int]:
2722 """
2723 Remap the host and port returned from the cluster to a different
2724 internal value. Useful if the client is not connecting directly
2725 to the cluster.
2726 """
2727 if self.address_remap:
2728 return self.address_remap((host, port))
2729 return host, port
2730
2731
2732class ClusterPipeline(AbstractRedis, AbstractRedisCluster, AsyncRedisClusterCommands):
2733 """
2734 Create a new ClusterPipeline object.
2735
2736 Usage::
2737
2738 result = await (
2739 rc.pipeline()
2740 .set("A", 1)
2741 .get("A")
2742 .hset("K", "F", "V")
2743 .hgetall("K")
2744 .mset_nonatomic({"A": 2, "B": 3})
2745 .get("A")
2746 .get("B")
2747 .delete("A", "B", "K")
2748 .execute()
2749 )
2750 # result = [True, "1", 1, {"F": "V"}, True, True, "2", "3", 1, 1, 1]
2751
2752 Note: For commands `DELETE`, `EXISTS`, `TOUCH`, `UNLINK`, `mset_nonatomic`, which
2753 are split across multiple nodes, you'll get multiple results for them in the array.
2754
2755 Retryable errors:
2756 - :class:`~.ClusterDownError`
2757 - :class:`~.ConnectionError`
2758 - :class:`~.TimeoutError`
2759
2760 Redirection errors:
2761 - :class:`~.TryAgainError`
2762 - :class:`~.MovedError`
2763 - :class:`~.AskError`
2764
2765 :param client:
2766 | Existing :class:`~.RedisCluster` client
2767 """
2768
2769 __slots__ = (
2770 "cluster_client",
2771 "_transaction",
2772 "_execution_strategy",
2773 )
2774
2775 # Type discrimination marker for @overload self-type pattern
2776 _is_async_client: Literal[True] = True
2777
2778 def __init__(
2779 self, client: RedisCluster, transaction: Optional[bool] = None
2780 ) -> None:
2781 self.cluster_client = client
2782 self._transaction = transaction
2783 self._execution_strategy: ExecutionStrategy = (
2784 PipelineStrategy(self)
2785 if not self._transaction
2786 else TransactionStrategy(self)
2787 )
2788
2789 @property
2790 def nodes_manager(self) -> "NodesManager":
2791 """Get the nodes manager from the cluster client."""
2792 return self.cluster_client.nodes_manager
2793
2794 # HIMPORT lifecycle on a cluster pipeline delegates to the parent client, mutating
2795 # the one shared registry that every node pool references. A fieldset declared here
2796 # is therefore visible to the batched himport_set pre-flight, mirroring the sync
2797 # ClusterPipeline (which inherits these from RedisCluster over the shared registry).
2798
2799 @property
2800 def himport_registry(self) -> HImportRegistry:
2801 """The cluster-wide HIMPORT fieldset registry (empty if none was declared).
2802
2803 Read-only: the registry is mutated only through the HIMPORT command methods.
2804 """
2805 return self.cluster_client.himport_registry
2806
2807 async def himport_prepare(
2808 self, fieldset_name: str, fields: Iterable[FieldT]
2809 ) -> bool:
2810 """Declare an HIMPORT fieldset cluster-wide (shared registry, applied lazily)."""
2811 return await self.cluster_client.himport_prepare(fieldset_name, fields)
2812
2813 async def himport_discard(self, fieldset_name: str) -> int:
2814 """Remove an HIMPORT fieldset cluster-wide (shared registry, applied lazily)."""
2815 return await self.cluster_client.himport_discard(fieldset_name)
2816
2817 async def himport_discard_all(self) -> int:
2818 """Remove all HIMPORT fieldsets cluster-wide (shared registry, applied lazily)."""
2819 return await self.cluster_client.himport_discard_all()
2820
2821 def set_response_callback(self, command: str, callback: ResponseCallbackT) -> None:
2822 """Set a custom response callback on the cluster client."""
2823 self.cluster_client.set_response_callback(command, callback)
2824
2825 async def initialize(self) -> "ClusterPipeline":
2826 await self._execution_strategy.initialize()
2827 return self
2828
2829 async def __aenter__(self) -> "ClusterPipeline":
2830 return await self.initialize()
2831
2832 async def __aexit__(self, exc_type: None, exc_value: None, traceback: None) -> None:
2833 await self.reset()
2834
2835 def __await__(self) -> Generator[Any, None, "ClusterPipeline"]:
2836 return self.initialize().__await__()
2837
2838 def __bool__(self) -> bool:
2839 "Pipeline instances should always evaluate to True on Python 3+"
2840 return True
2841
2842 def __len__(self) -> int:
2843 return len(self._execution_strategy)
2844
2845 def execute_command(
2846 self, *args: Union[KeyT, EncodableT], **kwargs: Any
2847 ) -> "ClusterPipeline":
2848 """
2849 Append a raw command to the pipeline.
2850
2851 :param args:
2852 | Raw command args
2853 :param kwargs:
2854
2855 - target_nodes: :attr:`NODE_FLAGS` or :class:`~.ClusterNode`
2856 or List[:class:`~.ClusterNode`] or Dict[Any, :class:`~.ClusterNode`]
2857 - Rest of the kwargs are passed to the Redis connection
2858 """
2859 return self._execution_strategy.execute_command(*args, **kwargs)
2860
2861 async def execute(
2862 self, raise_on_error: bool = True, allow_redirections: bool = True
2863 ) -> List[Any]:
2864 """
2865 Execute the pipeline.
2866
2867 It will retry the commands as specified by retries specified in :attr:`retry`
2868 & then raise an exception.
2869
2870 :param raise_on_error:
2871 | Raise the first error if there are any errors
2872 :param allow_redirections:
2873 | Whether to retry each failed command individually in case of redirection
2874 errors
2875
2876 :raises RedisClusterException: if target_nodes is not provided & the command
2877 can't be mapped to a slot
2878 """
2879 try:
2880 return await self._execution_strategy.execute(
2881 raise_on_error, allow_redirections
2882 )
2883 finally:
2884 await self.reset()
2885
2886 def _split_command_across_slots(
2887 self, command: str, *keys: KeyT
2888 ) -> "ClusterPipeline":
2889 for slot_keys in self.cluster_client._partition_keys_by_slot(keys).values():
2890 self.execute_command(command, *slot_keys)
2891
2892 return self
2893
2894 async def reset(self):
2895 """
2896 Reset back to empty pipeline.
2897 """
2898 await self._execution_strategy.reset()
2899
2900 def multi(self):
2901 """
2902 Start a transactional block of the pipeline after WATCH commands
2903 are issued. End the transactional block with `execute`.
2904 """
2905 self._execution_strategy.multi()
2906
2907 async def discard(self):
2908 """ """
2909 await self._execution_strategy.discard()
2910
2911 async def watch(self, *names):
2912 """Watches the values at keys ``names``"""
2913 await self._execution_strategy.watch(*names)
2914
2915 async def unwatch(self):
2916 """Unwatches all previously specified keys"""
2917 await self._execution_strategy.unwatch()
2918
2919 async def unlink(self, *names):
2920 await self._execution_strategy.unlink(*names)
2921
2922 def mset_nonatomic(
2923 self, mapping: Mapping[AnyKeyT, EncodableT]
2924 ) -> "ClusterPipeline":
2925 return self._execution_strategy.mset_nonatomic(mapping)
2926
2927
2928for command in PIPELINE_BLOCKED_COMMANDS:
2929 command = command.replace(" ", "_").lower()
2930 if command == "mset_nonatomic":
2931 continue
2932
2933 setattr(ClusterPipeline, command, block_pipeline_command(command))
2934
2935# client_list_iter has no wire command of its own to add to
2936# PIPELINE_BLOCKED_COMMANDS - it sends CLIENT LIST, blocked above under its
2937# own name - so block it explicitly here too, or it would fall through to
2938# the inherited implementation and queue CLIENT LIST like a real pipelined
2939# command instead of raising.
2940setattr(ClusterPipeline, "client_list_iter", block_pipeline_command("client_list_iter"))
2941
2942
2943class PipelineCommand:
2944 def __init__(self, position: int, *args: Any, **kwargs: Any) -> None:
2945 self.args = args
2946 self.kwargs = kwargs
2947 self.position = position
2948 self.result: Union[Any, Exception] = None
2949 # Either record type: a policy resolver serves ``CommandPolicies``, while the
2950 # fallbacks below reuse the shared ``CommandMetadata`` defaults. Only the two routing
2951 # policies, which both carry, are ever read.
2952 self.command_policies: Optional[Union[CommandPolicies, CommandMetadata]] = None
2953
2954 def __repr__(self) -> str:
2955 return f"[{self.position}] {self.args} ({self.kwargs})"
2956
2957
2958class ExecutionStrategy(ABC):
2959 @abstractmethod
2960 async def initialize(self) -> "ClusterPipeline":
2961 """
2962 Initialize the execution strategy.
2963
2964 See ClusterPipeline.initialize()
2965 """
2966 pass
2967
2968 @abstractmethod
2969 def execute_command(
2970 self, *args: Union[KeyT, EncodableT], **kwargs: Any
2971 ) -> "ClusterPipeline":
2972 """
2973 Append a raw command to the pipeline.
2974
2975 See ClusterPipeline.execute_command()
2976 """
2977 pass
2978
2979 @abstractmethod
2980 async def execute(
2981 self, raise_on_error: bool = True, allow_redirections: bool = True
2982 ) -> List[Any]:
2983 """
2984 Execute the pipeline.
2985
2986 It will retry the commands as specified by retries specified in :attr:`retry`
2987 & then raise an exception.
2988
2989 See ClusterPipeline.execute()
2990 """
2991 pass
2992
2993 @abstractmethod
2994 def mset_nonatomic(
2995 self, mapping: Mapping[AnyKeyT, EncodableT]
2996 ) -> "ClusterPipeline":
2997 """
2998 Executes multiple MSET commands according to the provided slot/pairs mapping.
2999
3000 See ClusterPipeline.mset_nonatomic()
3001 """
3002 pass
3003
3004 @abstractmethod
3005 async def reset(self):
3006 """
3007 Resets current execution strategy.
3008
3009 See: ClusterPipeline.reset()
3010 """
3011 pass
3012
3013 @abstractmethod
3014 def multi(self):
3015 """
3016 Starts transactional context.
3017
3018 See: ClusterPipeline.multi()
3019 """
3020 pass
3021
3022 @abstractmethod
3023 async def watch(self, *names):
3024 """
3025 Watch given keys.
3026
3027 See: ClusterPipeline.watch()
3028 """
3029 pass
3030
3031 @abstractmethod
3032 async def unwatch(self):
3033 """
3034 Unwatches all previously specified keys
3035
3036 See: ClusterPipeline.unwatch()
3037 """
3038 pass
3039
3040 @abstractmethod
3041 async def discard(self):
3042 pass
3043
3044 @abstractmethod
3045 async def unlink(self, *names):
3046 """
3047 "Unlink a key specified by ``names``"
3048
3049 See: ClusterPipeline.unlink()
3050 """
3051 pass
3052
3053 @abstractmethod
3054 def __len__(self) -> int:
3055 pass
3056
3057
3058class AbstractStrategy(ExecutionStrategy):
3059 def __init__(self, pipe: ClusterPipeline) -> None:
3060 self._pipe: ClusterPipeline = pipe
3061 self._command_queue: List["PipelineCommand"] = []
3062
3063 async def initialize(self) -> "ClusterPipeline":
3064 if self._pipe.cluster_client._initialize:
3065 await self._pipe.cluster_client.initialize()
3066 self._command_queue = []
3067 return self._pipe
3068
3069 def execute_command(
3070 self, *args: Union[KeyT, EncodableT], **kwargs: Any
3071 ) -> "ClusterPipeline":
3072 self._command_queue.append(
3073 PipelineCommand(len(self._command_queue), *args, **kwargs)
3074 )
3075 return self._pipe
3076
3077 def _annotate_exception(self, exception, number, command):
3078 """
3079 Provides extra context to the exception prior to it being handled
3080 """
3081 cmd = " ".join(map(safe_str, command))
3082 msg = (
3083 f"Command # {number} ({truncate_text(cmd)}) of pipeline "
3084 f"caused error: {exception.args[0]}"
3085 )
3086 exception.args = (msg,) + exception.args[1:]
3087
3088 @abstractmethod
3089 def mset_nonatomic(
3090 self, mapping: Mapping[AnyKeyT, EncodableT]
3091 ) -> "ClusterPipeline":
3092 pass
3093
3094 @abstractmethod
3095 async def execute(
3096 self, raise_on_error: bool = True, allow_redirections: bool = True
3097 ) -> List[Any]:
3098 pass
3099
3100 @abstractmethod
3101 async def reset(self):
3102 pass
3103
3104 @abstractmethod
3105 def multi(self):
3106 pass
3107
3108 @abstractmethod
3109 async def watch(self, *names):
3110 pass
3111
3112 @abstractmethod
3113 async def unwatch(self):
3114 pass
3115
3116 @abstractmethod
3117 async def discard(self):
3118 pass
3119
3120 @abstractmethod
3121 async def unlink(self, *names):
3122 pass
3123
3124 def __len__(self) -> int:
3125 return len(self._command_queue)
3126
3127
3128class PipelineStrategy(AbstractStrategy):
3129 def __init__(self, pipe: ClusterPipeline) -> None:
3130 super().__init__(pipe)
3131
3132 def mset_nonatomic(
3133 self, mapping: Mapping[AnyKeyT, EncodableT]
3134 ) -> "ClusterPipeline":
3135 encoder = self._pipe.cluster_client.encoder
3136
3137 slots_pairs = {}
3138 for pair in mapping.items():
3139 slot = key_slot(encoder.encode(pair[0]))
3140 slots_pairs.setdefault(slot, []).extend(pair)
3141
3142 for pairs in slots_pairs.values():
3143 self.execute_command("MSET", *pairs)
3144
3145 return self._pipe
3146
3147 async def execute(
3148 self, raise_on_error: bool = True, allow_redirections: bool = True
3149 ) -> List[Any]:
3150 if not self._command_queue:
3151 return []
3152
3153 try:
3154 retry_attempts = self._pipe.cluster_client.retry.get_retries()
3155 while True:
3156 try:
3157 if self._pipe.cluster_client._initialize:
3158 await self._pipe.cluster_client.initialize()
3159 return await self._execute(
3160 self._pipe.cluster_client,
3161 self._command_queue,
3162 raise_on_error=raise_on_error,
3163 allow_redirections=allow_redirections,
3164 )
3165
3166 except RedisCluster.ERRORS_ALLOW_RETRY as e:
3167 if retry_attempts > 0:
3168 # Try again with the new cluster setup. All other errors
3169 # should be raised.
3170 retry_attempts -= 1
3171 await self._pipe.cluster_client.aclose()
3172 await asyncio.sleep(0.25)
3173 else:
3174 # All other errors should be raised.
3175 raise e
3176 finally:
3177 await self.reset()
3178
3179 async def _execute(
3180 self,
3181 client: "RedisCluster",
3182 stack: List["PipelineCommand"],
3183 raise_on_error: bool = True,
3184 allow_redirections: bool = True,
3185 ) -> List[Any]:
3186 todo = [
3187 cmd for cmd in stack if not cmd.result or isinstance(cmd.result, Exception)
3188 ]
3189
3190 nodes = {}
3191 for cmd in todo:
3192 passed_targets = cmd.kwargs.pop("target_nodes", None)
3193 target_nodes_specified = bool(passed_targets) and not client._is_node_flag(
3194 passed_targets
3195 )
3196 _, command_policies = await client._resolve_command_policies(
3197 *cmd.args, target_nodes_specified=target_nodes_specified
3198 )
3199
3200 if target_nodes_specified:
3201 target_nodes = client._parse_target_nodes(passed_targets)
3202 else:
3203 target_nodes = await client._determine_nodes(
3204 *cmd.args,
3205 request_policy=command_policies.request_policy,
3206 node_flag=passed_targets,
3207 )
3208 if not target_nodes:
3209 raise RedisClusterException(
3210 f"No targets were found to execute {cmd.args} command on"
3211 )
3212 cmd.command_policies = command_policies
3213 if len(target_nodes) > 1:
3214 raise RedisClusterException(f"Too many targets for command {cmd.args}")
3215 node = target_nodes[0]
3216 if node.name not in nodes:
3217 nodes[node.name] = (node, [])
3218 nodes[node.name][1].append(cmd)
3219
3220 # Start timing for observability
3221 start_time = time.monotonic()
3222
3223 errors = await asyncio.gather(
3224 *(
3225 asyncio.create_task(node[0].execute_pipeline(node[1]))
3226 for node in nodes.values()
3227 )
3228 )
3229
3230 # Record operation duration for each node
3231 for node_name, (node, commands) in nodes.items():
3232 # Find the first error in this node's commands, if any
3233 node_error = None
3234 for cmd in commands:
3235 if isinstance(cmd.result, Exception):
3236 node_error = cmd.result
3237 break
3238
3239 db = node.connection_kwargs.get("db", 0)
3240 await record_operation_duration(
3241 command_name="PIPELINE",
3242 duration_seconds=time.monotonic() - start_time,
3243 server_address=node.host,
3244 server_port=node.port,
3245 db_namespace=str(db) if db is not None else None,
3246 error=node_error,
3247 )
3248
3249 if any(errors):
3250 if allow_redirections:
3251 # send each errored command individually
3252 for cmd in todo:
3253 if isinstance(cmd.result, (TryAgainError, MovedError, AskError)):
3254 try:
3255 cmd.result = client._policies_callback_mapping[
3256 cmd.command_policies.response_policy
3257 ](await client.execute_command(*cmd.args, **cmd.kwargs))
3258 except Exception as e:
3259 cmd.result = e
3260
3261 if raise_on_error:
3262 for cmd in todo:
3263 result = cmd.result
3264 if isinstance(result, Exception):
3265 command = " ".join(map(safe_str, cmd.args))
3266 msg = (
3267 f"Command # {cmd.position + 1} "
3268 f"({truncate_text(command)}) "
3269 f"of pipeline caused error: {result.args}"
3270 )
3271 result.args = (msg,) + result.args[1:]
3272 raise result
3273
3274 default_cluster_node = client.get_default_node()
3275
3276 # Check whether the default node was used. In some cases,
3277 # 'client.get_default_node()' may return None. The check below
3278 # prevents a potential AttributeError.
3279 if default_cluster_node is not None:
3280 default_node = nodes.get(default_cluster_node.name)
3281 if default_node is not None:
3282 # This pipeline execution used the default node, check if we need
3283 # to replace it.
3284 # Note: when the error is raised we'll reset the default node in the
3285 # caller function.
3286 for cmd in default_node[1]:
3287 # Check if it has a command that failed with a relevant
3288 # exception
3289 if type(cmd.result) in RedisCluster.ERRORS_ALLOW_RETRY:
3290 client.replace_default_node()
3291 break
3292
3293 return [cmd.result for cmd in stack]
3294
3295 async def reset(self):
3296 """
3297 Reset back to empty pipeline.
3298 """
3299 self._command_queue = []
3300
3301 def multi(self):
3302 raise RedisClusterException(
3303 "method multi() is not supported outside of transactional context"
3304 )
3305
3306 async def watch(self, *names):
3307 raise RedisClusterException(
3308 "method watch() is not supported outside of transactional context"
3309 )
3310
3311 async def unwatch(self):
3312 raise RedisClusterException(
3313 "method unwatch() is not supported outside of transactional context"
3314 )
3315
3316 async def discard(self):
3317 raise RedisClusterException(
3318 "method discard() is not supported outside of transactional context"
3319 )
3320
3321 async def unlink(self, *names):
3322 if len(names) != 1:
3323 raise RedisClusterException(
3324 "unlinking multiple keys is not implemented in pipeline command"
3325 )
3326
3327 return self.execute_command("UNLINK", names[0])
3328
3329
3330class TransactionStrategy(AbstractStrategy):
3331 NO_SLOTS_COMMANDS = {"UNWATCH"}
3332 IMMEDIATE_EXECUTE_COMMANDS = {"WATCH", "UNWATCH"}
3333 UNWATCH_COMMANDS = {"DISCARD", "EXEC", "UNWATCH"}
3334 SLOT_REDIRECT_ERRORS = (AskError, MovedError)
3335 CONNECTION_ERRORS = (
3336 ConnectionError,
3337 OSError,
3338 ClusterDownError,
3339 SlotNotCoveredError,
3340 )
3341
3342 def __init__(self, pipe: ClusterPipeline) -> None:
3343 super().__init__(pipe)
3344 self._explicit_transaction = False
3345 self._watching = False
3346 self._pipeline_slots: Set[int] = set()
3347 # True once a keyed (non-slot-agnostic) command has fixed the slot
3348 self._transaction_has_keyed_slot = False
3349 self._transaction_node: Optional[ClusterNode] = None
3350 self._transaction_connection: Optional[Connection] = None
3351 self._executing = False
3352 self._retry = copy(self._pipe.cluster_client.retry)
3353 self._retry.update_supported_errors(
3354 RedisCluster.ERRORS_ALLOW_RETRY + self.SLOT_REDIRECT_ERRORS
3355 )
3356
3357 async def _resolve_transaction_slot(self, *args) -> Optional[int]:
3358 """
3359 Pick a slot for a transactional pipeline command.
3360
3361 Zero-key EVAL/EVALSHA can run on any primary. Reuse an existing
3362 transaction slot when present so multiple zero-key scripts (or a
3363 mix with keyed commands) stay single-slot.
3364 """
3365 if args[0] in self.NO_SLOTS_COMMANDS:
3366 return None
3367
3368 if is_zero_key_eval_command(*args):
3369 if self._pipeline_slots:
3370 return next(iter(self._pipeline_slots))
3371 return await self._pipe.cluster_client._determine_slot(*args)
3372
3373 slot_number = await self._pipe.cluster_client._determine_slot(*args)
3374 if (
3375 slot_number is not None
3376 and self._pipeline_slots
3377 and slot_number not in self._pipeline_slots
3378 and not self._transaction_has_keyed_slot
3379 ):
3380 # Prior slots came only from zero-key scripts; retarget.
3381 self._pipeline_slots.clear()
3382 if slot_number is not None:
3383 self._transaction_has_keyed_slot = True
3384 return slot_number
3385
3386 def _get_client_and_connection_for_transaction(
3387 self,
3388 ) -> Tuple[ClusterNode, Connection]:
3389 """
3390 Find a connection for a pipeline transaction.
3391
3392 For running an atomic transaction, watch keys ensure that contents have not been
3393 altered as long as the watch commands for those keys were sent over the same
3394 connection. So once we start watching a key, we fetch a connection to the
3395 node that owns that slot and reuse it.
3396 """
3397 if not self._pipeline_slots:
3398 raise RedisClusterException(
3399 "At least a command with a key is needed to identify a node"
3400 )
3401
3402 node: ClusterNode = self._pipe.cluster_client.nodes_manager.get_node_from_slot(
3403 list(self._pipeline_slots)[0], False
3404 )
3405 self._transaction_node = node
3406
3407 if not self._transaction_connection:
3408 connection: Connection = self._transaction_node.acquire_connection()
3409 self._transaction_connection = connection
3410
3411 return self._transaction_node, self._transaction_connection
3412
3413 def execute_command(self, *args: Union[KeyT, EncodableT], **kwargs: Any) -> "Any":
3414 # Given the limitation of ClusterPipeline sync API, we have to run it in thread.
3415 return _run_coroutine_in_thread(self._execute_command(*args, **kwargs))
3416
3417 async def _execute_command(
3418 self, *args: Union[KeyT, EncodableT], **kwargs: Any
3419 ) -> Any:
3420 if self._pipe.cluster_client._initialize:
3421 await self._pipe.cluster_client.initialize()
3422
3423 slot_number: Optional[int] = None
3424 if args[0] not in self.NO_SLOTS_COMMANDS:
3425 slot_number = await self._resolve_transaction_slot(*args)
3426
3427 if (
3428 self._watching or args[0] in self.IMMEDIATE_EXECUTE_COMMANDS
3429 ) and not self._explicit_transaction:
3430 if args[0] == "WATCH":
3431 self._validate_watch()
3432
3433 if slot_number is not None:
3434 if self._pipeline_slots and slot_number not in self._pipeline_slots:
3435 raise CrossSlotTransactionError(
3436 "Cannot watch or send commands on different slots"
3437 )
3438
3439 self._pipeline_slots.add(slot_number)
3440 elif args[0] not in self.NO_SLOTS_COMMANDS:
3441 raise RedisClusterException(
3442 f"Cannot identify slot number for command: {args[0]},"
3443 "it cannot be triggered in a transaction"
3444 )
3445
3446 return self._immediate_execute_command(*args, **kwargs)
3447 else:
3448 if slot_number is not None:
3449 self._pipeline_slots.add(slot_number)
3450
3451 return super().execute_command(*args, **kwargs)
3452
3453 def _validate_watch(self):
3454 if self._explicit_transaction:
3455 raise RedisError("Cannot issue a WATCH after a MULTI")
3456
3457 self._watching = True
3458
3459 async def _immediate_execute_command(self, *args, **options):
3460 return await self._retry.call_with_retry(
3461 lambda: self._get_connection_and_send_command(*args, **options),
3462 self._reinitialize_on_error,
3463 with_failure_count=True,
3464 )
3465
3466 async def _get_connection_and_send_command(self, *args, **options):
3467 redis_node, connection = self._get_client_and_connection_for_transaction()
3468 # Only disconnect if not watching - disconnecting would lose WATCH state
3469 if not self._watching:
3470 await redis_node.disconnect_if_needed(connection)
3471
3472 # Start timing for observability
3473 start_time = time.monotonic()
3474
3475 try:
3476 response = await self._send_command_parse_response(
3477 connection, redis_node, args[0], *args, **options
3478 )
3479
3480 await record_operation_duration(
3481 command_name=args[0],
3482 duration_seconds=time.monotonic() - start_time,
3483 server_address=connection.host,
3484 server_port=connection.port,
3485 db_namespace=str(connection.db),
3486 )
3487
3488 return response
3489 except Exception as e:
3490 e.connection = connection
3491 await record_operation_duration(
3492 command_name=args[0],
3493 duration_seconds=time.monotonic() - start_time,
3494 server_address=connection.host,
3495 server_port=connection.port,
3496 db_namespace=str(connection.db),
3497 error=e,
3498 )
3499 raise
3500
3501 async def _send_command_parse_response(
3502 self,
3503 connection: Connection,
3504 redis_node: ClusterNode,
3505 command_name,
3506 *args,
3507 **options,
3508 ):
3509 """
3510 Send a command and parse the response
3511 """
3512
3513 # HIMPORT SET's wire form depends on per-connection state: the fieldset
3514 # must be PREPAREd on this connection first, and any fieldset discarded
3515 # since this connection last reconciled must be dropped. The
3516 # immediate/watched path (commands issued after WATCH, before MULTI)
3517 # would otherwise send a bare HIMPORT SET and fail with "no such
3518 # fieldset". Route it through the node's HIMPORT executor, the same way
3519 # the normal cluster path, the batched MULTI/EXEC path, and standalone
3520 # watched pipelines all do.
3521 himport_set = parse_himport_set_args(args)
3522 if himport_set is not None:
3523 # HIMPORT SET in the joined or split raw form; operands at the right
3524 # offsets. Too few operands returns None and falls through to the bare
3525 # send so the server returns its arity error.
3526 key, fieldset_name, values = himport_set
3527 output = await redis_node._himport_execute_set(
3528 connection, key, fieldset_name, values
3529 )
3530 else:
3531 await connection.send_command(*args)
3532 output = await redis_node.parse_response(
3533 connection, command_name, **options
3534 )
3535
3536 if command_name in self.UNWATCH_COMMANDS:
3537 self._watching = False
3538 return output
3539
3540 async def _reinitialize_on_error(self, error, failure_count):
3541 if hasattr(error, "connection"):
3542 await record_error_count(
3543 server_address=error.connection.host,
3544 server_port=error.connection.port,
3545 network_peer_address=error.connection.host,
3546 network_peer_port=error.connection.port,
3547 error_type=error,
3548 retry_attempts=failure_count,
3549 is_internal=True,
3550 )
3551
3552 if self._watching:
3553 if type(error) in self.SLOT_REDIRECT_ERRORS and self._executing:
3554 raise WatchError("Slot rebalancing occurred while watching keys")
3555
3556 if (
3557 type(error) in self.SLOT_REDIRECT_ERRORS
3558 or type(error) in self.CONNECTION_ERRORS
3559 ):
3560 if self._transaction_connection and self._transaction_node:
3561 # Disconnect and release back to pool
3562 await self._transaction_connection.disconnect()
3563 self._transaction_node.release(self._transaction_connection)
3564 self._transaction_connection = None
3565
3566 self._pipe.cluster_client.reinitialize_counter += 1
3567 if (
3568 self._pipe.cluster_client.reinitialize_steps
3569 and self._pipe.cluster_client.reinitialize_counter
3570 % self._pipe.cluster_client.reinitialize_steps
3571 == 0
3572 ):
3573 await self._pipe.cluster_client.nodes_manager.initialize()
3574 self.reinitialize_counter = 0
3575 else:
3576 if isinstance(error, AskError):
3577 await self._pipe.cluster_client.nodes_manager.move_slot(error)
3578
3579 self._executing = False
3580
3581 async def _raise_first_error(self, responses, stack, start_time):
3582 """
3583 Raise the first exception on the stack
3584 """
3585 for r, cmd in zip(responses, stack):
3586 if isinstance(r, Exception):
3587 self._annotate_exception(r, cmd.position + 1, cmd.args)
3588
3589 await record_operation_duration(
3590 command_name="TRANSACTION",
3591 duration_seconds=time.monotonic() - start_time,
3592 server_address=self._transaction_connection.host,
3593 server_port=self._transaction_connection.port,
3594 db_namespace=str(self._transaction_connection.db),
3595 error=r,
3596 )
3597
3598 raise r
3599
3600 def mset_nonatomic(
3601 self, mapping: Mapping[AnyKeyT, EncodableT]
3602 ) -> "ClusterPipeline":
3603 raise NotImplementedError("Method is not supported in transactional context.")
3604
3605 async def execute(
3606 self, raise_on_error: bool = True, allow_redirections: bool = True
3607 ) -> List[Any]:
3608 stack = self._command_queue
3609 if not stack and (not self._watching or not self._pipeline_slots):
3610 return []
3611
3612 return await self._execute_transaction_with_retries(stack, raise_on_error)
3613
3614 async def _execute_transaction_with_retries(
3615 self, stack: List["PipelineCommand"], raise_on_error: bool
3616 ):
3617 return await self._retry.call_with_retry(
3618 lambda: self._execute_transaction(stack, raise_on_error),
3619 lambda error, failure_count: self._reinitialize_on_error(
3620 error, failure_count
3621 ),
3622 with_failure_count=True,
3623 )
3624
3625 async def _execute_transaction(
3626 self, stack: List["PipelineCommand"], raise_on_error: bool
3627 ):
3628 if len(self._pipeline_slots) > 1:
3629 raise CrossSlotTransactionError(
3630 "All keys involved in a cluster transaction must map to the same slot"
3631 )
3632
3633 self._executing = True
3634
3635 redis_node, connection = self._get_client_and_connection_for_transaction()
3636 # Only disconnect if not watching - disconnecting would lose WATCH state
3637 if not self._watching:
3638 await redis_node.disconnect_if_needed(connection)
3639
3640 # Ensure fieldsets referenced by buffered HIMPORT SETs are prepared on this
3641 # node's connection before the MULTI/EXEC block (session state, not
3642 # transactional). All keys share one slot here, so it is a single node.
3643 await redis_node._himport_prepare_pipeline(connection, stack)
3644
3645 stack = chain(
3646 [PipelineCommand(0, "MULTI")],
3647 stack,
3648 [PipelineCommand(0, "EXEC")],
3649 )
3650 commands = [c.args for c in stack if EMPTY_RESPONSE not in c.kwargs]
3651 packed_commands = connection.pack_commands(commands)
3652
3653 # Start timing for observability
3654 start_time = time.monotonic()
3655
3656 await connection.send_packed_command(packed_commands)
3657 errors = []
3658
3659 # parse off the response for MULTI
3660 # NOTE: we need to handle ResponseErrors here and continue
3661 # so that we read all the additional command messages from
3662 # the socket
3663 try:
3664 await redis_node.parse_response(connection, "MULTI")
3665 except ResponseError as e:
3666 self._annotate_exception(e, 0, "MULTI")
3667 errors.append(e)
3668 except self.CONNECTION_ERRORS as cluster_error:
3669 self._annotate_exception(cluster_error, 0, "MULTI")
3670 cluster_error.connection = connection
3671 raise
3672
3673 # and all the other commands
3674 for i, command in enumerate(self._command_queue):
3675 if EMPTY_RESPONSE in command.kwargs:
3676 errors.append((i, command.kwargs[EMPTY_RESPONSE]))
3677 else:
3678 try:
3679 _ = await redis_node.parse_response(connection, "_")
3680 except self.SLOT_REDIRECT_ERRORS as slot_error:
3681 self._annotate_exception(slot_error, i + 1, command.args)
3682 errors.append(slot_error)
3683 except self.CONNECTION_ERRORS as cluster_error:
3684 self._annotate_exception(cluster_error, i + 1, command.args)
3685 cluster_error.connection = connection
3686 raise
3687 except ResponseError as e:
3688 self._annotate_exception(e, i + 1, command.args)
3689 errors.append(e)
3690
3691 response = None
3692 # parse the EXEC.
3693 try:
3694 response = await redis_node.parse_response(connection, "EXEC")
3695 except ExecAbortError:
3696 if errors:
3697 raise errors[0]
3698 raise
3699
3700 self._executing = False
3701
3702 # EXEC clears any watched keys
3703 self._watching = False
3704
3705 if response is None:
3706 raise WatchError("Watched variable changed.")
3707
3708 # put any parse errors into the response
3709 for i, e in errors:
3710 response.insert(i, e)
3711
3712 if len(response) != len(self._command_queue):
3713 raise InvalidPipelineStack(
3714 "Unexpected response length for cluster pipeline EXEC."
3715 " Command stack was {} but response had length {}".format(
3716 [c.args[0] for c in self._command_queue], len(response)
3717 )
3718 )
3719
3720 # find any errors in the response and raise if necessary
3721 if raise_on_error or len(errors) > 0:
3722 await self._raise_first_error(
3723 response,
3724 self._command_queue,
3725 start_time,
3726 )
3727
3728 # We have to run response callbacks manually
3729 data = []
3730 for r, cmd in zip(response, self._command_queue):
3731 if not isinstance(r, Exception):
3732 command_name = cmd.args[0]
3733 if command_name in self._pipe.cluster_client.response_callbacks:
3734 r = self._pipe.cluster_client.response_callbacks[command_name](
3735 r, **cmd.kwargs
3736 )
3737 data.append(r)
3738
3739 await record_operation_duration(
3740 command_name="TRANSACTION",
3741 duration_seconds=time.monotonic() - start_time,
3742 server_address=connection.host,
3743 server_port=connection.port,
3744 db_namespace=str(connection.db),
3745 )
3746
3747 return data
3748
3749 async def reset(self):
3750 self._command_queue = []
3751
3752 try:
3753 # make sure to reset the connection state in the event that we
3754 # were watching something
3755 if self._transaction_connection:
3756 try:
3757 if self._watching:
3758 # call this manually since our unwatch or
3759 # immediate_execute_command methods can call reset()
3760 await self._transaction_connection.send_command("UNWATCH")
3761 await self._transaction_connection.read_response()
3762 except self.CONNECTION_ERRORS:
3763 # disconnect will also remove any previous WATCHes
3764 if self._transaction_connection:
3765 await self._transaction_connection.disconnect()
3766 except asyncio.CancelledError:
3767 # Disconnect so any unread UNWATCH reply does not get
3768 # served to the next caller that takes the connection.
3769 if self._transaction_connection:
3770 await self._transaction_connection.disconnect()
3771 raise
3772 else:
3773 # On the happy path, honor lazy reconnect before release.
3774 await self._transaction_node.disconnect_if_needed(
3775 self._transaction_connection
3776 )
3777 finally:
3778 # Always return the connection to the node's free queue, even on
3779 # cancellation, so cancelled resets do not leak pooled
3780 # connections. Detach the reference before releasing so the
3781 # strategy never holds a pointer to a returned connection.
3782 # ClusterNode.release is synchronous, so no shield is required.
3783 if self._transaction_connection and self._transaction_node:
3784 connection, self._transaction_connection = (
3785 self._transaction_connection,
3786 None,
3787 )
3788 self._transaction_node.release(connection)
3789 # clean up the other instance attributes
3790 self._transaction_connection = None
3791 self._transaction_node = None
3792 self._watching = False
3793 self._explicit_transaction = False
3794 self._pipeline_slots = set()
3795 self._transaction_has_keyed_slot = False
3796 self._executing = False
3797
3798 def multi(self):
3799 if self._explicit_transaction:
3800 raise RedisError("Cannot issue nested calls to MULTI")
3801 if self._command_queue:
3802 raise RedisError(
3803 "Commands without an initial WATCH have already been issued"
3804 )
3805 self._explicit_transaction = True
3806
3807 async def watch(self, *names):
3808 if self._explicit_transaction:
3809 raise RedisError("Cannot issue a WATCH after a MULTI")
3810
3811 return await self.execute_command("WATCH", *names)
3812
3813 async def unwatch(self):
3814 if self._watching:
3815 return await self.execute_command("UNWATCH")
3816
3817 return True
3818
3819 async def discard(self):
3820 await self.reset()
3821
3822 async def unlink(self, *names):
3823 return self.execute_command("UNLINK", *names)
3824
3825
3826class _ClusterNodePoolAdapter(ConnectionPoolInterface):
3827 """Thin adapter exposing the :class:`ConnectionPoolInterface` that
3828 :class:`PubSub` requires, backed by a :class:`ClusterNode`'s own
3829 connection pool.
3830
3831 Connections are acquired from the node via
3832 :meth:`ClusterNode.acquire_connection` and returned via
3833 :meth:`ClusterNode.release`. :meth:`PubSub.aclose` already
3834 disconnects the connection *before* calling :meth:`release`, so the
3835 connection is returned to the node's free-queue in a disconnected
3836 state — guaranteeing that a subscribed socket is never silently
3837 reused for regular commands.
3838
3839 Methods that do not apply to this adapter (the underlying node's
3840 lifecycle is managed by the cluster, not by individual PubSub
3841 instances) are implemented as no-ops so the adapter remains a valid
3842 :class:`ConnectionPoolInterface`.
3843 """
3844
3845 def __init__(self, node: "ClusterNode") -> None:
3846 self._node = node
3847 self.connection_kwargs = node.connection_kwargs
3848
3849 # -- methods used by PubSub ------------------------------------------------
3850
3851 def get_encoder(self) -> Encoder:
3852 return self._node.get_encoder()
3853
3854 async def get_connection(
3855 self, command_name: Optional[str] = None, *keys: Any, **options: Any
3856 ) -> AbstractConnection:
3857 connection = self._node.acquire_connection()
3858 try:
3859 await connection.connect()
3860 except BaseException:
3861 # connect() may fail mid-handshake (e.g. after the TCP socket
3862 # is established but before AUTH/HELLO completes) leaving the
3863 # connection in a partially-connected state. Disconnect before
3864 # returning it to the node's free queue so it is not reused.
3865 await connection.disconnect()
3866 self._node.release(connection)
3867 raise
3868 return connection
3869
3870 async def release(self, connection: AbstractConnection) -> None:
3871 # PubSub.aclose() disconnects the connection before calling
3872 # release(), so it is safe to put it back in the node's free
3873 # queue – it will reconnect lazily on next use.
3874 await self._node.disconnect_if_needed(connection)
3875 self._node.release(connection)
3876
3877 # -- no-op stubs for the rest of ConnectionPoolInterface -------------------
3878 # The node's connections are shared with regular cluster traffic and its
3879 # lifecycle is managed by RedisCluster / NodesManager, so the adapter must
3880 # not reset, disconnect, retry-configure or re-auth them on behalf of a
3881 # single PubSub instance.
3882
3883 def get_protocol(self):
3884 return self.connection_kwargs.get("protocol", None)
3885
3886 def reset(self) -> None:
3887 pass
3888
3889 async def disconnect(self, inuse_connections: bool = True) -> None:
3890 pass
3891
3892 async def aclose(self) -> None:
3893 pass
3894
3895 def set_retry(self, retry: "Retry") -> None:
3896 pass
3897
3898 async def re_auth_callback(self, token: TokenInterface) -> None:
3899 pass
3900
3901 def get_connection_count(self) -> List[Tuple[int, dict]]:
3902 return []
3903
3904
3905def _unregister_slots_cache_listener(
3906 dispatcher_ref: "weakref.ref[EventDispatcher]",
3907 listener: AsyncEventListenerInterface,
3908 event_type: Type[object],
3909) -> None:
3910 # Module-level finalizer callback. Kept free of strong references to the
3911 # owning ClusterPubSub so attaching it via weakref.finalize does not
3912 # extend the pubsub's lifetime.
3913 dispatcher = dispatcher_ref()
3914 if dispatcher is not None:
3915 dispatcher.unregister_listeners({event_type: [listener]})
3916
3917
3918class ClusterPubSubSlotsCacheListener(AsyncEventListenerInterface):
3919 """
3920 Async listener that forwards AsyncAfterSlotsCacheRefreshEvent to a
3921 ClusterPubSub.
3922
3923 Holds a weak reference to the pubsub so it does not keep the instance
3924 alive. Deterministic cleanup of the dispatcher's strong reference to this
3925 listener is performed by a ``weakref.finalize`` attached to the owning
3926 ClusterPubSub in ``ClusterPubSub.__init__``.
3927 """
3928
3929 def __init__(self, pubsub: "ClusterPubSub") -> None:
3930 self._pubsub_ref: "weakref.ref[ClusterPubSub]" = weakref.ref(pubsub)
3931
3932 async def listen(self, event: object) -> None:
3933 pubsub = self._pubsub_ref()
3934 if pubsub is None:
3935 # Race window between pubsub GC and the finalizer running; safe
3936 # no-op, finalizer will remove this listener shortly.
3937 return
3938 try:
3939 await pubsub.on_slots_changed()
3940 except Exception as e:
3941 # Listeners must not break slots-cache refresh; log and continue so
3942 # a single buggy pubsub cannot starve the rest.
3943 logger.exception(
3944 "pubsub %r raised during slots-cache change: %s: %s",
3945 pubsub,
3946 type(e).__name__,
3947 e,
3948 )
3949
3950
3951# How long a per-node sharded-pubsub connection is skipped by the round robin
3952# after a failed poll. PubSub._execute reconnects and retries through the
3953# connection's own Retry, so one poll on an unreachable node can cost its whole
3954# retry budget rather than the timeout the caller asked for; a cool-off keeps
3955# the single reader from spending every pass on that node while its healthy
3956# siblings hold undelivered messages.
3957SHARD_POLL_COOL_OFF_SECONDS = 1.0
3958
3959# How often a failed sharded-pubsub poll may trigger a slots-cache refresh.
3960# Reconciliation is otherwise purely event-driven, and a node that has left the
3961# deployment answers ECONNREFUSED rather than MOVED - so without this the reader
3962# would cool off against the departed node forever and the shard channels pinned
3963# to it would never move to their new owner. Throttled because the refresh costs
3964# a CLUSTER SLOTS round trip and a failing node fails every poll.
3965SHARD_TOPOLOGY_REPAIR_INTERVAL_SECONDS = 5.0
3966
3967
3968class ClusterPubSub(PubSub):
3969 """
3970 Async cluster implementation for pub/sub.
3971
3972 IMPORTANT: before using ClusterPubSub, read about the known limitations
3973 with pubsub in Cluster mode and learn how to workaround them:
3974 https://redis.readthedocs.io/en/stable/clustering.html#known-pubsub-limitations
3975 """
3976
3977 def __init__(
3978 self,
3979 redis_cluster: "RedisCluster",
3980 node: Optional["ClusterNode"] = None,
3981 host: Optional[str] = None,
3982 port: Optional[int] = None,
3983 push_handler_func: Optional[Callable] = None,
3984 event_dispatcher: Optional[EventDispatcher] = None,
3985 **kwargs: Any,
3986 ) -> None:
3987 """
3988 When a pubsub instance is created without specifying a node, a single
3989 node will be transparently chosen for the pubsub connection on the
3990 first command execution. The node will be determined by:
3991 1. Hashing the channel name in the request to find its keyslot
3992 2. Selecting a node that handles the keyslot: If read_from_replicas is
3993 set to true or load_balancing_strategy is set, a replica can be selected.
3994
3995 :param redis_cluster: RedisCluster instance
3996 :param node: ClusterNode to connect to
3997 :param host: Host of the node to connect to
3998 :param port: Port of the node to connect to
3999 :param push_handler_func: Optional push handler function
4000 :param event_dispatcher: Optional event dispatcher
4001 :param kwargs: Additional keyword arguments
4002 """
4003 self.node = None
4004 self.set_pubsub_node(redis_cluster, node, host, port)
4005
4006 # Borrow the node's own connection pool via an adapter rather than
4007 # creating a second, detached ConnectionPool for pubsub.
4008 if self.node is not None:
4009 connection_pool = _ClusterNodePoolAdapter(self.node)
4010 else:
4011 connection_pool = None
4012
4013 self.cluster = redis_cluster
4014 self.node_pubsub_mapping: Dict[str, PubSub] = {}
4015 # Reverse index: shard channel (normalized) -> owning node.name. Used to
4016 # route sunsubscribe calls and reconcile subscriptions after slot
4017 # migration / failover.
4018 self._shard_channel_to_node: Dict[Any, str] = {}
4019 # Per-node poll cool-off deadlines (monotonic). Weak-keyed so a
4020 # per-node pubsub dropped from node_pubsub_mapping takes its entry with
4021 # it instead of leaking one per migration.
4022 self._poll_cool_off: "weakref.WeakKeyDictionary[PubSub, float]" = (
4023 weakref.WeakKeyDictionary()
4024 )
4025 # Node names whose last poll failed to connect. Read by
4026 # _migrate_shard_channel to skip a wire SUNSUBSCRIBE that cannot
4027 # succeed, and cleared as soon as a poll on that node works again.
4028 self._unreachable_nodes: Set[str] = set()
4029 # Monotonic deadline before which a failed poll must not trigger
4030 # another slots-cache refresh. 0.0 means "never refreshed".
4031 self._next_topology_repair: float = 0.0
4032 # Dedicated lock for shard-subscription bookkeeping. Distinct from
4033 # PubSub.self._lock (which serializes wire I/O on the cluster-level
4034 # connection used by aclose / send_command / regular subscribe) so
4035 # that reconciliation cannot starve those unrelated coroutines
4036 # during long per-channel migrations.
4037 self._shard_state_lock: asyncio.Lock = asyncio.Lock()
4038 # Background tasks created by on_slots_changed; kept to prevent GC.
4039 self._reconcile_tasks: Set[asyncio.Task] = set()
4040 # Background NodesManager.initialize() tasks created by
4041 # _schedule_topology_repair; kept to prevent GC, and kept apart from
4042 # _reconcile_tasks because aclose() cancels that set and this work is
4043 # not ours to cancel - see _schedule_topology_repair.
4044 self._topology_repair_tasks: Set[asyncio.Task] = set()
4045 self._pubsubs_generator = self._pubsubs_generator()
4046 if event_dispatcher is None:
4047 self._event_dispatcher = EventDispatcher()
4048 else:
4049 self._event_dispatcher = event_dispatcher
4050 super().__init__(
4051 connection_pool=connection_pool,
4052 encoder=redis_cluster.encoder,
4053 push_handler_func=push_handler_func,
4054 event_dispatcher=self._event_dispatcher,
4055 **kwargs,
4056 )
4057 # Subscribe to slots-cache change notifications so shard subscriptions
4058 # can be reconciled automatically after topology refreshes.
4059 nm_dispatcher = redis_cluster.nodes_manager._event_dispatcher
4060 self._slots_cache_listener = ClusterPubSubSlotsCacheListener(self)
4061 nm_dispatcher.register_listeners(
4062 {AsyncAfterSlotsCacheRefreshEvent: [self._slots_cache_listener]}
4063 )
4064 # Deterministic GC-time cleanup so short-lived pubsubs do not leak
4065 # listeners in the dispatcher when no slots-refresh event ever fires.
4066 weakref.finalize(
4067 self,
4068 _unregister_slots_cache_listener,
4069 weakref.ref(nm_dispatcher),
4070 self._slots_cache_listener,
4071 AsyncAfterSlotsCacheRefreshEvent,
4072 )
4073
4074 async def _ensure_cluster_initialized(self) -> None:
4075 if self.cluster._initialize:
4076 await self.cluster.initialize()
4077
4078 def set_pubsub_node(
4079 self,
4080 cluster: "RedisCluster",
4081 node: Optional["ClusterNode"] = None,
4082 host: Optional[str] = None,
4083 port: Optional[int] = None,
4084 ) -> None:
4085 """
4086 The pubsub node will be set according to the passed node, host and port
4087 When none of the node, host, or port are specified - the node is set
4088 to None and will be determined by the keyslot of the channel in the
4089 first command to be executed.
4090 RedisClusterException will be thrown if the passed node does not exist
4091 in the cluster.
4092 If host is passed without port, or vice versa, a DataError will be
4093 thrown.
4094 """
4095 if node is not None:
4096 # node is passed by the user
4097 self._raise_on_invalid_node(cluster, node, node.host, node.port)
4098 pubsub_node = node
4099 elif host is not None and port is not None:
4100 # host and port passed by the user
4101 node = cluster.get_node(host=host, port=port)
4102 self._raise_on_invalid_node(cluster, node, host, port)
4103 pubsub_node = node
4104 elif host is not None or port is not None:
4105 # only one of host and port is specified
4106 raise DataError("Specify both host and port")
4107 else:
4108 # nothing specified by the user
4109 pubsub_node = None
4110 self.node = pubsub_node
4111
4112 def get_pubsub_node(self) -> Optional["ClusterNode"]:
4113 """
4114 Get the node that is being used as the pubsub connection.
4115
4116 :return: The ClusterNode being used for pubsub, or None if not yet determined
4117 """
4118 return self.node
4119
4120 async def _resubscribe_shard_channels(self) -> None:
4121 # A single node can own multiple slot ranges, so a batched
4122 # ``SSUBSCRIBE`` covering every tracked channel would be rejected by
4123 # Redis with a ``CROSSSLOT`` error. Group by hash slot and emit one
4124 # ``SSUBSCRIBE`` per slot.
4125 by_slot: defaultdict[int, dict] = defaultdict(dict)
4126 for k, v in self.shard_channels.items():
4127 by_slot[key_slot(self.encoder.encode(k))][k] = v
4128 for subscriptions in by_slot.values():
4129 await self._resubscribe(subscriptions, self.ssubscribe)
4130
4131 def _get_node_pubsub(self, node: "ClusterNode") -> PubSub:
4132 """Get or create a PubSub instance for the given node."""
4133 try:
4134 return self.node_pubsub_mapping[node.name]
4135 except KeyError:
4136 pubsub = PubSub(
4137 connection_pool=_ClusterNodePoolAdapter(node),
4138 encoder=self.cluster.encoder,
4139 push_handler_func=self.push_handler_func,
4140 event_dispatcher=self._event_dispatcher,
4141 )
4142 # Replay shard subscriptions on reconnect with slot-aware grouping
4143 # so that channels spanning multiple slots owned by this node do
4144 # not trigger a CROSSSLOT error.
4145 pubsub._resubscribe_shard_channels = MethodType(
4146 ClusterPubSub._resubscribe_shard_channels, pubsub
4147 )
4148 self._pubsub_io_lock(pubsub)
4149 self.node_pubsub_mapping[node.name] = pubsub
4150 return pubsub
4151
4152 def _find_node_name_for_pubsub(self, pubsub: PubSub) -> Optional[str]:
4153 for name, candidate in self.node_pubsub_mapping.items():
4154 if candidate is pubsub:
4155 return name
4156 return None
4157
4158 @staticmethod
4159 def _pubsub_io_lock(pubsub: PubSub) -> asyncio.Lock:
4160 """Return the per-node pubsub's wire I/O lock, creating it on first use.
4161
4162 A per-node ``PubSub`` is read by whichever task polls
4163 ``get_sharded_message`` and written by the reconciliation task
4164 (``_migrate_shard_channel``) and by any caller of ``ssubscribe`` /
4165 ``sunsubscribe``. ``PubSub`` guards writes with its own ``_lock``
4166 (``PubSub.execute_command``) but reads take no lock at all, so without
4167 this the reader can be awaiting ``read_response`` while another task's
4168 ``_execute`` disconnects and reconnects the same socket underneath it -
4169 which loses the reply to the handshake and surfaces as a read timeout
4170 followed by a broken connection.
4171
4172 Kept on the pubsub rather than in a dict keyed by node name so it
4173 travels with the object through ``node_pubsub_mapping`` and cannot go
4174 stale when a per-node pubsub is dropped and recreated.
4175 """
4176 lock = getattr(pubsub, "_shard_io_lock", None)
4177 if lock is None:
4178 lock = asyncio.Lock()
4179 pubsub._shard_io_lock = lock
4180 return lock
4181
4182 @staticmethod
4183 def _detach_shard_channel(pubsub: PubSub, channel: Any) -> None:
4184 """Forget a shard channel on a per-node pubsub without a wire round trip.
4185
4186 ``PubSub.sunsubscribe`` only records the intent in
4187 ``pending_unsubscribe_shard_channels``; the channel leaves
4188 ``shard_channels`` when the server confirmation is read. So if the
4189 ``SUNSUBSCRIBE`` never reaches the server, ``on_connect`` clears the
4190 pending set and replays ``SSUBSCRIBE`` for the channel - on the node it
4191 is being migrated away from, on every reconnect. Once the caller has
4192 decided the channel belongs to a different node, the local intent is
4193 the only truth left, so drop it here.
4194
4195 Unlike the sync counterpart there is no ``subscribed_event`` to clear
4196 once the last subscription is gone: the async ``PubSub`` has no such
4197 event, and its ``subscribed`` is a plain property derived from the
4198 subscription dicts - including the ``shard_channels`` entry just
4199 dropped here.
4200
4201 Nor is the per-node I/O lock taken, as in the sync counterpart - and
4202 here it is not even needed for correctness: this runs to completion
4203 without an ``await``, so the event loop already makes it atomic against
4204 ``handle_message``'s unsubscribe bookkeeping. Should an ``await`` ever
4205 land in this body, that bookkeeping still discards rather than removes
4206 (see ``PubSub.handle_message``), so a racing detach cannot make it raise
4207 ``KeyError`` into a poll no caller catches.
4208 """
4209 pubsub.shard_channels.pop(channel, None)
4210 pubsub.pending_unsubscribe_shard_channels.discard(channel)
4211
4212 async def _drop_node_pubsub(self, name: str, pubsub: PubSub) -> None:
4213 """Retire a per-node pubsub and drop it from ``node_pubsub_mapping``.
4214
4215 Callers hold ``_shard_state_lock``, the lock that every mutation of
4216 that mapping observes. ``aclose()`` runs under the per-node I/O lock so
4217 the socket is not torn down beneath a concurrent bounded poll parked in
4218 ``parse_response``, and its errors are swallowed: retiring one node's
4219 pubsub must not abort the caller's pass.
4220
4221 Every caller must leave nothing subscribed on ``pubsub`` (or have lost
4222 the node itself). An empty per-node pubsub left in the mapping keeps a
4223 live connection with nothing subscribed on it, so ``_poll_node_pubsub``
4224 passes its ``connection is None`` guard and blocks reading a socket no
4225 message can arrive on - for the whole timeout of every pass, and until
4226 the connection is torn down when the caller passed ``timeout=None``.
4227
4228 Popping it from the mapping is not enough to keep it out of a poll; the
4229 rebind below is what stops the round robin from handing it out.
4230 """
4231 try:
4232 async with self._pubsub_io_lock(pubsub):
4233 await pubsub.aclose()
4234 except Exception:
4235 pass
4236 self.node_pubsub_mapping.pop(name, None)
4237 self._unreachable_nodes.discard(name)
4238 # Same snapshot reason ``aclose()`` recreates this: ``_pubsubs_generator``
4239 # captures node_pubsub_mapping.values() into a local list inside
4240 # ``yield from``, which the pop above does not reach - so a generator
4241 # suspended mid-yield-from would still hand the object we just retired
4242 # to the next poll. That cannot stall the reader the way it can in the
4243 # sync stack (there is no subscription wait to park on, and
4244 # ``_poll_node_pubsub``'s ``connection is None`` guard skips it), but it
4245 # still spends one of the pass's ``range(len(node_pubsub_mapping))``
4246 # slots on a dead entry - so a healthy sibling loses its turn until the
4247 # captured snapshot drains. ``type(self)`` bypasses the instance-level
4248 # self-shadow established at __init__. Costs nothing: constructing a
4249 # generator runs no frame, so the per-node collection loop in
4250 # reinitialize_shard_subscriptions can rebind once per dropped node.
4251 self._pubsubs_generator = type(self)._pubsubs_generator( # type: ignore[method-assign]
4252 self
4253 )
4254
4255 async def _sharded_message_generator(
4256 self, timeout: float = 0.0
4257 ) -> Tuple[Optional[PubSub], Optional[Dict[str, Any]]]:
4258 """Generate messages from shard channels across all nodes."""
4259 first_error: Optional[BaseException] = None
4260 polled = 0
4261 failed = 0
4262 next_ready: Optional[float] = None
4263 for _ in range(len(self.node_pubsub_mapping)):
4264 pubsub = next(self._pubsubs_generator)
4265 if pubsub is None:
4266 # node_pubsub_mapping was emptied between the len() above and
4267 # here; nothing left to poll in this pass.
4268 break
4269 if self._poll_cool_off:
4270 deadline = self._poll_cool_off.get(pubsub, 0.0)
4271 if time.monotonic() < deadline:
4272 # In cool-off after a failed poll: skip it so the reader
4273 # spends this pass on the nodes that can still deliver.
4274 if next_ready is None or deadline < next_ready:
4275 next_ready = deadline
4276 continue
4277 polled += 1
4278 try:
4279 message = await self._poll_node_pubsub(pubsub, timeout)
4280 except MovedError as e:
4281 # Handled, not failed: _handle_moved_on_read re-routes the
4282 # offending channels and schedules reconciliation, so the next
4283 # pass recovers. Re-raising a MovedError out of a pubsub read
4284 # would only hand the caller an error it cannot act on. Still
4285 # cool off: if the slots cache cannot be corrected the repair
4286 # would otherwise re-run on every poll.
4287 self._poll_cool_off[pubsub] = (
4288 time.monotonic() + SHARD_POLL_COOL_OFF_SECONDS
4289 )
4290 await self._handle_moved_on_read(pubsub, e)
4291 continue
4292 except (ConnectionError, TimeoutError, OSError) as e:
4293 # One unhealthy node must not starve its healthy siblings. A
4294 # single reader serves every per-node pubsub, so aborting the
4295 # pass here stops delivery cluster-wide for as long as this one
4296 # node stays unreachable - even though the slots it no longer
4297 # serves are the only ones affected. Keep polling the rest and
4298 # surface an error only if nothing in the pass worked, the same
4299 # made-progress rule reinitialize_shard_subscriptions applies.
4300 failed += 1
4301 if first_error is None:
4302 first_error = e
4303 # Cool off before polling this one again. PubSub._execute
4304 # reconnects and then retries through the connection's own
4305 # Retry, so a single "bounded" poll on an unreachable node can
4306 # cost its whole retry budget - far longer than the timeout the
4307 # caller asked for. Without a cool-off the reader goes straight
4308 # back to that node on the next pass and pays it again, which
4309 # is what turns one sick node into a cluster-wide delivery
4310 # stall.
4311 self._poll_cool_off[pubsub] = (
4312 time.monotonic() + SHARD_POLL_COOL_OFF_SECONDS
4313 )
4314 # Mutated without _shard_state_lock, mirroring the sync stack:
4315 # this is an advisory hint for _migrate_shard_channel's fast
4316 # path, and both misread directions are handled there and
4317 # self-heal - a stale entry only skips a SUNSUBSCRIBE to a dead
4318 # node, a missing one only pays a reconnect before the same
4319 # local forget.
4320 node_name = self._find_node_name_for_pubsub(pubsub)
4321 if node_name is not None:
4322 self._unreachable_nodes.add(node_name)
4323 if logger.isEnabledFor(logging.DEBUG):
4324 logger.debug(
4325 "sharded pubsub poll failed on %s: %s: %s",
4326 node_name,
4327 type(e).__name__,
4328 e,
4329 )
4330 # A node that has left the deployment never answers MOVED, so
4331 # this branch is the only signal that its shard channels may
4332 # need a new owner. Ask for a slots-cache refresh; its dispatch
4333 # reaches on_slots_changed and reconciles.
4334 self._schedule_topology_repair()
4335 continue
4336 # Emptiness check first: this is the per-message hot path, and the
4337 # weakref lookup a WeakKeyDictionary pop needs is pure overhead
4338 # while no node is in cool-off, which is the normal case.
4339 if self._poll_cool_off:
4340 self._poll_cool_off.pop(pubsub, None)
4341 if self._unreachable_nodes:
4342 node_name = self._find_node_name_for_pubsub(pubsub)
4343 if node_name is not None:
4344 self._unreachable_nodes.discard(node_name)
4345 if message is not None:
4346 return pubsub, message
4347 if first_error is not None and failed == polled:
4348 raise first_error
4349 if polled == 0 and next_ready is not None:
4350 await self._wait_out_cool_off(next_ready, timeout)
4351 return None, None
4352
4353 @staticmethod
4354 async def _wait_out_cool_off(next_ready: float, timeout: Optional[float]) -> None:
4355 """Wait out a pass in which every node was skipped for cool-off.
4356
4357 Such a pass does no wire read at all, so returning straight away
4358 ignores the timeout the caller asked to block for - and a reader loop
4359 on ``get_sharded_message`` polls back immediately, spinning until the
4360 cool-off expires instead of blocking. Sleep instead: until the
4361 earliest cool-off is over, never longer than the caller's timeout, and
4362 not at all for a non-blocking poll.
4363 """
4364 if timeout is not None and timeout <= 0:
4365 return
4366 delay = next_ready - time.monotonic()
4367 if delay <= 0:
4368 return
4369 if timeout is not None:
4370 delay = min(delay, timeout)
4371 await asyncio.sleep(delay)
4372
4373 def _poll_io_lock(self, pubsub: PubSub, timeout: Optional[float]):
4374 """Guard a per-node poll against concurrent writers on the same socket.
4375
4376 ``timeout=None`` makes ``_poll_node_pubsub``'s read wait indefinitely,
4377 so holding the lock across it would block reconciliation for as long as
4378 no message arrives. Such a caller drives the pubsub itself and gets the
4379 pre-existing unguarded behavior; every bounded poll - which is what
4380 ``ClusterPubSub``'s own callers use - is serialized.
4381 """
4382 if timeout is None:
4383 return nullcontext()
4384 return self._pubsub_io_lock(pubsub)
4385
4386 async def _poll_node_pubsub(
4387 self, pubsub: PubSub, timeout: Optional[float]
4388 ) -> Optional[Dict[str, Any]]:
4389 """Read one message from a per-node pubsub, dispatching outside the lock.
4390
4391 Splits ``PubSub.get_message`` so the per-node I/O lock covers the wire
4392 read only. ``handle_message`` awaits a subscribed channel's user
4393 handler inline, and a handler is free to await ``ssubscribe`` /
4394 ``sunsubscribe`` on this ``ClusterPubSub`` - which takes
4395 ``_shard_state_lock`` and then the same I/O lock. ``asyncio.Lock`` is
4396 not reentrant, so holding the I/O lock across the handler hangs the
4397 reader task permanently and silently on the re-acquire; even a
4398 task-reentrant lock would still deadlock against the reconciliation
4399 task, which holds ``_shard_state_lock`` and awaits that I/O lock.
4400
4401 The two halves of ``handle_message`` are mutually exclusive:
4402 ``UNSUBSCRIBE_MESSAGE_TYPES`` does subscription bookkeeping and never
4403 reaches a handler, ``PUBLISH_MESSAGE_TYPES`` only dispatches. So
4404 bookkeeping stays inside the lock - it mutates the very
4405 ``shard_channels`` / ``pending_unsubscribe_shard_channels`` that
4406 ``ssubscribe`` / ``sunsubscribe`` mutate under this lock - and only the
4407 dispatch moves out. The cost is a narrow race: a reconciliation pass
4408 that detaches the channel between the read and the dispatch makes the
4409 handler lookup miss, so the message is returned to the caller instead
4410 of dispatched. That is the same in-flight-during-unsubscribe race
4411 ``PubSub`` itself has, and closing it would mean duplicating
4412 ``handle_message``'s dispatch here.
4413
4414 The async ``PubSub`` has no ``subscribed_event`` for the sync
4415 counterpart's subscription wait to mirror, but the connectionless state
4416 that wait guards against still has to be handled: a per-node pubsub
4417 enters ``node_pubsub_mapping`` before its first ``SSUBSCRIBE``
4418 (``_get_node_pubsub``) and is left connectionless by ``aclose()`` (the
4419 GC in ``reinitialize_shard_subscriptions``), while ``parse_response``
4420 raises ``RuntimeError`` on a ``None`` connection - which neither poll
4421 site catches. ``_pubsubs_generator`` yields from a snapshot of the
4422 mapping, so it can hand out a pubsub the GC has just dropped. Checking
4423 under the I/O lock rather than before it closes that window for every
4424 bounded poll, because the GC ``aclose()``s under the same lock.
4425 """
4426 async with self._poll_io_lock(pubsub, timeout):
4427 if pubsub.connection is None:
4428 # Not connected yet, or closed by a concurrent teardown.
4429 # Reading would raise RuntimeError; skip this node instead.
4430 return None
4431 response = await pubsub.parse_response(
4432 block=(timeout is None), timeout=timeout
4433 )
4434 # get_message's truthiness test, not "is None": a health check
4435 # reply filtered out by parse_response, or an empty bulk, is "no
4436 # message" rather than a message to parse.
4437 if not response:
4438 return None
4439 if not self._is_publish_response(response):
4440 # Don't pass ignore_subscribe_messages here - let
4441 # get_sharded_message handle the filtering after processing
4442 # subscription state changes
4443 return await pubsub.handle_message(
4444 response, ignore_subscribe_messages=False
4445 )
4446 return await pubsub.handle_message(response, ignore_subscribe_messages=False)
4447
4448 @staticmethod
4449 def _is_publish_response(response: Any) -> bool:
4450 """Whether a raw pubsub reply can make ``handle_message`` dispatch.
4451
4452 ``handle_message`` awaits a user handler only for
4453 ``PUBLISH_MESSAGE_TYPES``; every other reply either does subscription
4454 bookkeeping (``UNSUBSCRIBE_MESSAGE_TYPES``) or is a pong, and the two
4455 branches are mutually exclusive. A non-sequence reply is the bare-PING
4456 shape ``handle_message`` rewrites into a pong, so it cannot dispatch
4457 either.
4458 """
4459 if not isinstance(response, (list, tuple)):
4460 return False
4461 return str_if_bytes(response[0]) in PubSub.PUBLISH_MESSAGE_TYPES
4462
4463 def _schedule_topology_repair(self) -> None:
4464 """Ask for a slots-cache refresh after a poll could not reach a node.
4465
4466 ``reinitialize_shard_subscriptions`` only ever runs from a slots-cache
4467 change notification, and a node that has been rebooted or taken out of
4468 the deployment answers ``ECONNREFUSED`` rather than ``MOVED`` - so the
4469 read path itself has to ask, or the shard channels pinned to that node
4470 stay there for the lifetime of the pubsub.
4471
4472 ``NodesManager.initialize`` serializes concurrent callers, drops nodes
4473 that have left the topology and dispatches
4474 ``AsyncAfterSlotsCacheRefreshEvent``, which reaches
4475 ``on_slots_changed``; run it as a task so a bounded poll does not pay
4476 for a ``CLUSTER SLOTS`` round trip, and throttle it because a node that
4477 is down fails every poll.
4478 """
4479 if not self.shard_channels:
4480 return
4481 now = time.monotonic()
4482 if now < self._next_topology_repair:
4483 return
4484 self._next_topology_repair = now + SHARD_TOPOLOGY_REPAIR_INTERVAL_SECONDS
4485 # Not tracked in _reconcile_tasks: that set is cancelled by aclose(),
4486 # and ``initialize`` refreshes the slots cache of the whole cluster
4487 # client, which every other command on it reads - closing one pubsub
4488 # must not abort a refresh in flight, leaving the shared cache stale
4489 # until some later MOVED or failed poll asks again. The sync
4490 # counterpart has the same property for free: ``reset()`` retires the
4491 # reconciliation executor with ``cancel_futures=True``, which drops
4492 # queued work but lets a running ``initialize`` finish.
4493 task = asyncio.create_task(self.cluster.nodes_manager.initialize())
4494 self._topology_repair_tasks.add(task)
4495 task.add_done_callback(self._topology_repair_tasks.discard)
4496 task.add_done_callback(self._log_reconcile_task_exception)
4497
4498 async def _handle_moved_on_read(self, pubsub: PubSub, error: MovedError) -> None:
4499 """Re-route shard channels pinned to a node that lost their slot.
4500
4501 ``PubSub.on_connect`` replays ``SSUBSCRIBE`` to the node its connection
4502 is bound to, so after a slot migration that node answers ``MOVED``.
4503 ``MovedError`` is not in ``Retry.supported_errors`` and no other code on
4504 the read path refreshes the slots cache, so a shard channel left on a
4505 former owner could never recover. Drop the offending channels from this
4506 pubsub so the replay stops, forget their recorded owner so
4507 ``reinitialize_shard_subscriptions`` does not short-circuit on an
4508 already-advanced reverse index, then apply the redirect and reconcile.
4509 """
4510 node_name = self._find_node_name_for_pubsub(pubsub)
4511 logger.debug(
4512 "sharded pubsub: %s no longer owns slot %s; re-routing its shard channels",
4513 node_name,
4514 error.slot_id,
4515 )
4516 async with self._shard_state_lock:
4517 for channel in list(pubsub.shard_channels):
4518 if key_slot(self.encoder.encode(channel)) != error.slot_id:
4519 continue
4520 self._detach_shard_channel(pubsub, channel)
4521 if self._shard_channel_to_node.get(channel) == node_name:
4522 del self._shard_channel_to_node[channel]
4523 # The detach above can leave this pubsub with nothing subscribed -
4524 # a node that lost its only slot answers MOVED for every channel it
4525 # held. Retire it here rather than leave it in the mapping for a
4526 # collector elsewhere: no SUNSUBSCRIBE confirmation will arrive for
4527 # a channel forgotten locally, so get_sharded_message's collector
4528 # cannot reach it, and the reconciliation pass scheduled below only
4529 # GCs it once its task gets to run - a whole poll cool-off later,
4530 # at best, while a poll that reaches the empty pubsub first blocks
4531 # on a socket no message can arrive on (see _drop_node_pubsub).
4532 if node_name is not None and not pubsub.subscribed:
4533 await self._drop_node_pubsub(node_name, pubsub)
4534 # move_slot applies the redirect to the slots cache and dispatches
4535 # AsyncAfterSlotsCacheRefreshEvent, which reaches on_slots_changed. Call
4536 # on_slots_changed unconditionally too: move_slot skips the dispatch on
4537 # a circular MOVED, and a duplicate reconciliation pass is a no-op.
4538 # move_slot indexes slots_cache by the redirected slot, so an
4539 # as-yet-uncovered slot raises: log and still reconcile rather than let
4540 # a repair attempt break a pubsub read.
4541 try:
4542 await self.cluster.nodes_manager.move_slot(error)
4543 except Exception as exc:
4544 logger.debug(
4545 "sharded pubsub: could not apply the redirect for slot %s: %s: %s",
4546 error.slot_id,
4547 type(exc).__name__,
4548 exc,
4549 )
4550 await self.on_slots_changed()
4551
4552 def _pubsubs_generator(self) -> Generator[Optional[PubSub], None, None]:
4553 """Generator that yields PubSub instances in round-robin fashion.
4554
4555 Never returns: a generator that returns is exhausted for good and only
4556 ``reset`` recreates this one, so a momentarily empty
4557 ``node_pubsub_mapping`` - reconciliation drops a per-node pubsub before
4558 creating its replacement - would stop the round robin permanently.
4559 Yields ``None`` for an empty mapping instead, which lets the caller skip
4560 the slot without this loop spinning on an empty list.
4561 """
4562 while True:
4563 current_nodes = list(self.node_pubsub_mapping.values())
4564 if not current_nodes:
4565 yield None
4566 else:
4567 yield from current_nodes
4568
4569 async def get_sharded_message(
4570 self,
4571 ignore_subscribe_messages: bool = False,
4572 timeout: float = 0.0,
4573 target_node: Optional["ClusterNode"] = None,
4574 ) -> Optional[Dict[str, Any]]:
4575 """
4576 Get the next sharded pubsub message, or ``None`` if none is available.
4577
4578 Polls the per-node connections in round robin unless ``target_node`` is
4579 given, and keeps shard channels attached to the node that currently
4580 owns their slot: a failed poll cools that node off and asks for a
4581 slots-cache refresh, and a ``MOVED`` reply re-routes the affected
4582 channels to their new owner. Neither reaches the caller. A connection
4583 failure is surfaced only when every node polled in the pass failed, so
4584 one unreachable node does not stop delivery from its healthy siblings.
4585
4586 ``target_node`` opts out of that shielding: a caller that names a
4587 single node has no sibling to protect, so connection errors propagate.
4588 A ``MOVED`` reply is still handled rather than raised.
4589
4590 :param ignore_subscribe_messages: Whether to ignore subscribe messages
4591 :param timeout: Timeout for message retrieval
4592 :param target_node: Specific node to get message from
4593 :return: Message dictionary or None
4594 """
4595 pubsub: Optional[PubSub]
4596 if target_node:
4597 pubsub = self.node_pubsub_mapping.get(target_node.name)
4598 if pubsub:
4599 try:
4600 message = await self._poll_node_pubsub(pubsub, timeout)
4601 except MovedError as e:
4602 # Same handling as the round-robin path: the caller cannot
4603 # act on a MovedError raised out of a pubsub read, and the
4604 # channels this node no longer owns have to be re-routed or
4605 # they never recover. Cool off too, so a slots cache that
4606 # cannot be corrected does not re-run the repair on every
4607 # poll. Unlike that path, connectivity errors still
4608 # propagate: they are swallowed there only to keep one sick
4609 # node from starving its healthy siblings, and a caller that
4610 # named a single node has no sibling to protect.
4611 self._poll_cool_off[pubsub] = (
4612 time.monotonic() + SHARD_POLL_COOL_OFF_SECONDS
4613 )
4614 await self._handle_moved_on_read(pubsub, e)
4615 message = None
4616 else:
4617 message = None
4618 else:
4619 pubsub, message = await self._sharded_message_generator(timeout=timeout)
4620
4621 if message is None:
4622 return None
4623 # Only sunsubscribe mutates cluster-level shard state; bypassing the
4624 # lock on the data-message hot path keeps smessage delivery from
4625 # competing with the reconciliation task for _shard_state_lock.
4626 if str_if_bytes(message["type"]) == "sunsubscribe":
4627 # Serialize state mutation against reinitialize_shard_subscriptions
4628 # (background task). The blocking _poll_node_pubsub above
4629 # intentionally runs outside the lock so reconciliation is not
4630 # stalled by long polls.
4631 async with self._shard_state_lock:
4632 if message["channel"] in self.pending_unsubscribe_shard_channels:
4633 # User-initiated sunsubscribe: drop from cluster-level tracking.
4634 self.pending_unsubscribe_shard_channels.remove(message["channel"])
4635 self.shard_channels.pop(message["channel"], None)
4636 self._shard_channel_to_node.pop(message["channel"], None)
4637 # Drop the per-node pubsub that delivered the confirmation once
4638 # it no longer holds any shard subscriptions, regardless of
4639 # whether the sunsubscribe was user-initiated or driven by
4640 # slot-migration reconciliation (_migrate_shard_channel, which
4641 # intentionally does not add the channel to
4642 # pending_unsubscribe_shard_channels). This releases the
4643 # dedicated connection that would otherwise linger.
4644 # Identifying the receiving pubsub directly (rather than via
4645 # the cluster's current slot map) is required after slot
4646 # migration, where the channel's owner is no longer the node
4647 # that received our original SSUBSCRIBE.
4648 if pubsub is not None and not pubsub.subscribed:
4649 name = self._find_node_name_for_pubsub(pubsub)
4650 if name is not None:
4651 await self._drop_node_pubsub(name, pubsub)
4652
4653 # Only suppress subscribe/unsubscribe messages, not data messages (smessage)
4654 if str_if_bytes(message["type"]) in ("ssubscribe", "sunsubscribe"):
4655 if self.ignore_subscribe_messages or ignore_subscribe_messages:
4656 return None
4657 return message
4658
4659 async def ssubscribe(
4660 self, *args: ChannelT | Subscription, **kwargs: PubSubHandler
4661 ) -> None:
4662 """
4663 Subscribe to shard channels.
4664
4665 :param args: Channel names or ``Subscription`` objects
4666 :param kwargs: Channel names with handlers
4667 """
4668 s_channels = parse_pubsub_subscriptions(args, kwargs)
4669
4670 if not s_channels:
4671 return
4672 await self._ensure_cluster_initialized()
4673
4674 # Serialize against reinitialize_shard_subscriptions (background
4675 # task) so the reverse index, shard_channels, and node_pubsub_mapping
4676 # are not mutated concurrently. _migrate_shard_channel below does not
4677 # re-acquire this lock (asyncio.Lock is non-reentrant).
4678 async with self._shard_state_lock:
4679 for s_channel, handler in s_channels.items():
4680 node = self.cluster.get_node_from_key(s_channel)
4681 if not node:
4682 continue
4683 # Lazy re-route: if this channel is already tracked against a
4684 # different node (e.g. after a slot migration), migrate it now
4685 # so the caller's intent is applied on the current owner.
4686 normalized_key = next(iter(self._normalize_keys({s_channel: None})))
4687 old_name = self._shard_channel_to_node.get(normalized_key)
4688 if old_name and old_name != node.name:
4689 # Match PubSub.ssubscribe() dict.update() semantics: the
4690 # caller's newly supplied handler (including None) always
4691 # overrides any previously registered handler.
4692 await self._migrate_shard_channel(
4693 normalized_key,
4694 handler,
4695 old_name,
4696 node,
4697 )
4698 continue
4699 pubsub = self._get_node_pubsub(node)
4700 async with self._pubsub_io_lock(pubsub):
4701 if handler:
4702 await pubsub.ssubscribe(Subscription(s_channel, handler))
4703 else:
4704 await pubsub.ssubscribe(s_channel)
4705 self.shard_channels.update(pubsub.shard_channels)
4706 self._shard_channel_to_node[normalized_key] = node.name
4707 self.pending_unsubscribe_shard_channels.difference_update(
4708 self._normalize_keys({s_channel: None})
4709 )
4710
4711 async def sunsubscribe(self, *args: Any) -> None:
4712 """
4713 Unsubscribe from shard channels.
4714
4715 :param args: Channel names to unsubscribe from. If empty, unsubscribe from all.
4716 """
4717 if args:
4718 args = list_or_args(args[0], args[1:])
4719 else:
4720 args = list(self.shard_channels.keys())
4721
4722 if not self.node_pubsub_mapping:
4723 return
4724 # Keep initialization outside the shard-state lock: it can dispatch a
4725 # topology refresh whose reconciler must remain able to make progress.
4726 await self._ensure_cluster_initialized()
4727
4728 # Serialize against reinitialize_shard_subscriptions: the reverse
4729 # index and node_pubsub_mapping must not change between the lookup
4730 # and the per-node sunsubscribe call below.
4731 async with self._shard_state_lock:
4732 for s_channel in args:
4733 normalized_key = next(iter(self._normalize_keys({s_channel: None})))
4734 # Route via the reverse index so we unsubscribe on the node
4735 # that actually holds the subscription. After a slot migration
4736 # the cluster's current owner may no longer be that node.
4737 name = self._shard_channel_to_node.get(normalized_key)
4738 if name and name in self.node_pubsub_mapping:
4739 pubsub = self.node_pubsub_mapping[name]
4740 else:
4741 node = self.cluster.get_node_from_key(s_channel)
4742 if not node or node.name not in self.node_pubsub_mapping:
4743 continue
4744 pubsub = self.node_pubsub_mapping[node.name]
4745 async with self._pubsub_io_lock(pubsub):
4746 await pubsub.sunsubscribe(s_channel)
4747 self.pending_unsubscribe_shard_channels.update(
4748 pubsub.pending_unsubscribe_shard_channels
4749 )
4750
4751 async def reinitialize_shard_subscriptions(self) -> None:
4752 """
4753 Reconcile per-node shard subscriptions against the cluster's current
4754 slot ownership map. For each tracked shard channel whose owning node
4755 has changed (e.g. after CLUSTER SETSLOT / failover), sunsubscribe on
4756 the old node's pubsub and ssubscribe on the new owner's pubsub,
4757 preserving any registered handler.
4758 """
4759 uncovered: list = []
4760 made_progress = False
4761 first_migrate_error: Optional[BaseException] = None
4762 async with self._shard_state_lock:
4763 for channel, handler in list(self.shard_channels.items()):
4764 if channel in self.pending_unsubscribe_shard_channels:
4765 continue
4766 try:
4767 new_node = self.cluster.get_node_from_key(channel)
4768 except SlotNotCoveredError:
4769 # Slot is transiently uncovered (mid-migration / partial
4770 # topology refresh). Defer this channel so coverable
4771 # siblings still reconcile this pass; we surface the
4772 # error below so the caller (and logs) know not every
4773 # channel was reconciled. Retry happens on the next
4774 # slots-cache change notification.
4775 uncovered.append(channel)
4776 continue
4777 old_name = self._shard_channel_to_node.get(channel)
4778 if old_name == new_node.name:
4779 owner = self.node_pubsub_mapping.get(new_node.name)
4780 if owner is not None and channel in owner.shard_channels:
4781 continue
4782 # The reverse index names this node but the subscription is
4783 # not there. _migrate_shard_channel detaches from the old
4784 # owner before it advances the index, so a pass that failed
4785 # to attach leaves the channel subscribed nowhere - and once
4786 # ownership moves back, this short-circuit would skip it for
4787 # the lifetime of the pubsub. Re-attach instead of trusting
4788 # the index; there is nothing to sunsubscribe from.
4789 old_name = None
4790 try:
4791 await self._migrate_shard_channel(
4792 channel, handler, old_name, new_node
4793 )
4794 made_progress = True
4795 except (ConnectionError, TimeoutError, OSError) as e:
4796 # Transient connectivity error while subscribing on the
4797 # new owner (or unsubscribing on the old owner if its
4798 # handler chose to re-raise). Do not abort reconciliation
4799 # for sibling channels: _shard_channel_to_node was not
4800 # advanced for this channel, so the next slots-cache
4801 # change notification will retry it.
4802 logger.warning(
4803 "shard channel %r migration deferred: %s: %s",
4804 channel,
4805 type(e).__name__,
4806 e,
4807 )
4808 if first_migrate_error is None:
4809 first_migrate_error = e
4810 continue
4811 # Garbage-collect per-node pubsubs that no longer hold any
4812 # subscription so their connections are released.
4813 for name, pubsub in list(self.node_pubsub_mapping.items()):
4814 if not pubsub.subscribed:
4815 await self._drop_node_pubsub(name, pubsub)
4816 if uncovered:
4817 # Surface the uncovered channels so the caller (and observer
4818 # notification path) knows reconciliation was incomplete. All
4819 # coverable siblings have already been migrated above.
4820 raise SlotNotCoveredError(
4821 f"{len(uncovered)} shard channel(s) left unreconciled; "
4822 f"slot(s) not covered by the cluster: {uncovered!r}"
4823 )
4824 if first_migrate_error is not None and not made_progress:
4825 # Every migration attempted in this pass failed transiently and
4826 # nothing else made progress. Re-raise the first caught error
4827 # (typically the root cause; later failures are often downstream
4828 # symptoms of the same unreachable node) so the task's done-
4829 # callback surfaces a single representative failure through the
4830 # same logger channel used for SlotNotCoveredError. Per-channel
4831 # WARNINGs above preserve the full forensic detail.
4832 raise first_migrate_error
4833
4834 async def _forget_shard_channel_on_old_node(
4835 self, old_pubsub: PubSub, channel: Any, old_name: str
4836 ) -> None:
4837 """Drop a migrating shard channel from a node we could not tell about it.
4838
4839 Forget the channel locally: the caller advances the reverse index to the
4840 new owner, so reconciliation will never revisit this channel, while
4841 ``on_connect`` would keep replaying ``SSUBSCRIBE`` for it to this very
4842 node on every reconnect - the server would answer ``MOVED`` and the
4843 subscription would never work again.
4844 """
4845 self._detach_shard_channel(old_pubsub, channel)
4846 # Drop the per-node pubsub when either the old node has left the cluster
4847 # topology - no reconnect target, so the round-robin generator must stop
4848 # yielding a dead one, and any sibling subscription it still holds
4849 # recovers through ``PubSub._execute``'s reconnect and ``on_connect``
4850 # replay - or the detach above left it with nothing subscribed.
4851 #
4852 # The empty case cannot be deferred to a collector elsewhere, because
4853 # neither of the other two can reach it. ``get_sharded_message``'s
4854 # unsubscribe branch needs a ``SUNSUBSCRIBE`` confirmation, and none will
4855 # arrive for a channel this method forgot locally - that is the whole
4856 # reason it is forgotten. ``reinitialize_shard_subscriptions``'s
4857 # end-of-pass GC only runs for the reconciliation caller, while
4858 # ``ssubscribe``'s lazy re-route reaches here without it. An empty pubsub
4859 # left in the mapping keeps a live connection with nothing subscribed on
4860 # it, so ``_poll_node_pubsub`` passes its ``connection is None`` guard and
4861 # blocks reading a socket no message can arrive on: forever when the
4862 # caller passed ``timeout=None``, and for the whole timeout of every pass
4863 # otherwise, before a single healthy node is read.
4864 if (
4865 self.cluster.get_node(node_name=old_name) is None
4866 or not old_pubsub.subscribed
4867 ):
4868 await self._drop_node_pubsub(old_name, old_pubsub)
4869
4870 async def _migrate_shard_channel(
4871 self,
4872 channel: Any,
4873 handler: Optional[Callable],
4874 old_name: Optional[str],
4875 new_node: "ClusterNode",
4876 ) -> None:
4877 # Detach from the old per-node pubsub, best-effort: the old node may
4878 # already be unreachable during migration / failover.
4879 if old_name and old_name in self.node_pubsub_mapping:
4880 old_pubsub = self.node_pubsub_mapping[old_name]
4881 if old_name in self._unreachable_nodes:
4882 # The reader has just failed to reach this node, so a
4883 # ``SUNSUBSCRIBE`` cannot arrive. Skip it: the attempt would pay
4884 # a full reconnect (and the client's whole retry budget) behind
4885 # the reader on the same per-node io lock, all while this pass
4886 # holds ``_shard_state_lock`` - which is what turns one departed
4887 # node into a migration slow enough to look like a permanent
4888 # delivery stall.
4889 await self._forget_shard_channel_on_old_node(
4890 old_pubsub, channel, old_name
4891 )
4892 else:
4893 try:
4894 async with self._pubsub_io_lock(old_pubsub):
4895 await old_pubsub.sunsubscribe(channel)
4896 except (ConnectionError, TimeoutError, OSError):
4897 # redis-py's Connection has already called ``disconnect()``
4898 # before raising (see Connection.read_response /
4899 # send_packed_command with ``disconnect_on_error=True``), so
4900 # ``old_pubsub``'s dedicated socket is gone and the
4901 # ``SUNSUBSCRIBE`` never reached the server.
4902 await self._forget_shard_channel_on_old_node(
4903 old_pubsub, channel, old_name
4904 )
4905 # Attach to the new per-node pubsub, preserving the handler. Decode to
4906 # a text key only when we must pass it as a kwarg (handler present).
4907 new_pubsub = self._get_node_pubsub(new_node)
4908 async with self._pubsub_io_lock(new_pubsub):
4909 if handler:
4910 await new_pubsub.ssubscribe(Subscription(channel, handler))
4911 else:
4912 await new_pubsub.ssubscribe(channel)
4913 self.shard_channels.update(new_pubsub.shard_channels)
4914 normalized_key = next(iter(self._normalize_keys({channel: None})))
4915 self._shard_channel_to_node[normalized_key] = new_node.name
4916 self.pending_unsubscribe_shard_channels.difference_update(
4917 self._normalize_keys({channel: None})
4918 )
4919
4920 async def on_slots_changed(self) -> None:
4921 # Observer hook invoked by NodesManager after a slots-cache refresh.
4922 # Schedule reconciliation as a separate task so the caller's code
4923 # path (typically MovedError handling in _execute_command) is not
4924 # blocked on the network I/O performed by reinitialize_shard_
4925 # subscriptions. No-op when there are no shard subscriptions to
4926 # reconcile.
4927 if not self.shard_channels:
4928 return
4929 task = asyncio.create_task(self.reinitialize_shard_subscriptions())
4930 self._reconcile_tasks.add(task)
4931 task.add_done_callback(self._reconcile_tasks.discard)
4932 # Consume the task's exception (if any) so Python does not emit a
4933 # "Task exception was never retrieved" warning. reinitialize_shard_
4934 # subscriptions surfaces SlotNotCoveredError when a slot is still
4935 # transiently uncovered; route it through the same logger channel
4936 # as sync ClusterPubSubSlotsCacheListener for consistent observability.
4937 task.add_done_callback(self._log_reconcile_task_exception)
4938
4939 @staticmethod
4940 def _log_reconcile_task_exception(task: "asyncio.Task") -> None:
4941 if task.cancelled():
4942 return
4943 exc = task.exception()
4944 if exc is not None:
4945 logger.error(
4946 "shard subscription reconciliation failed: %r", exc, exc_info=exc
4947 )
4948
4949 def get_redis_connection(self) -> Optional["AbstractConnection"]:
4950 """
4951 Get the Redis connection of the pubsub connected node.
4952
4953 Returns the pubsub's dedicated connection (acquired from its own
4954 connection pool), not from the ClusterNode's connection pool.
4955 This avoids the connection pool resource leak that would occur
4956 if we called node.acquire_connection() without releasing.
4957 """
4958 # Return the pubsub's own dedicated connection, which is acquired
4959 # from self.connection_pool when executing pubsub commands.
4960 # This is safe because it's the connection dedicated to this pubsub
4961 # instance, not a shared pool connection from the ClusterNode.
4962 return self.connection
4963
4964 async def aclose(self) -> None:
4965 """
4966 Disconnect the pubsub connection.
4967 """
4968 # Cancel and gather in-flight reconciliation tasks BEFORE acquiring
4969 # _shard_state_lock. The tasks themselves take that lock inside
4970 # reinitialize_shard_subscriptions; since asyncio.Lock is non-
4971 # reentrant, gathering while holding it would deadlock. Awaiting
4972 # each task with suppressed CancelledError also avoids unhandled-
4973 # exception warnings if the task was created but not yet scheduled.
4974 # _topology_repair_tasks is deliberately left alone: it holds
4975 # NodesManager.initialize() calls that refresh the whole client's
4976 # slots cache, which is not this pubsub's to abort (see
4977 # _schedule_topology_repair). The set keeps them referenced until
4978 # they finish and discard themselves.
4979 if self._reconcile_tasks:
4980 tasks = list(self._reconcile_tasks)
4981 for task in tasks:
4982 task.cancel()
4983 await asyncio.gather(*tasks, return_exceptions=True)
4984 # Hold _shard_state_lock across the rest of the teardown so it
4985 # observes the same mutual-exclusion discipline as ssubscribe /
4986 # sunsubscribe / get_sharded_message / reinitialize_shard_
4987 # subscriptions, which all mutate shard_channels,
4988 # _shard_channel_to_node, and node_pubsub_mapping under this lock.
4989 # Without it, super().aclose() rebinds shard_channels and
4990 # pending_unsubscribe_shard_channels in parallel with a concurrent
4991 # user-coroutine mutation that resumes during one of the awaits
4992 # below, silently dropping subscription intent.
4993 async with self._shard_state_lock:
4994 self._reconcile_tasks.clear()
4995 # Close all shard pubsub instances first, under the per-node I/O
4996 # lock so the socket is not torn down beneath a concurrent bounded
4997 # poll parked in parse_response. A bounded poll holds the lock only
4998 # for its timeout, and an unbounded one holds nullcontext() (see
4999 # _poll_io_lock), so teardown never waits indefinitely here.
5000 for pubsub in self.node_pubsub_mapping.values():
5001 async with self._pubsub_io_lock(pubsub):
5002 await pubsub.aclose()
5003 # Drop the now-dead per-node pubsubs from the mapping so the
5004 # round-robin in _pubsubs_generator / _sharded_message_generator
5005 # cannot yield them between teardown and re-subscription.
5006 self.node_pubsub_mapping.clear()
5007 self._unreachable_nodes.clear()
5008 # Drop the throttle window too: a reused pubsub that keeps a
5009 # deadline armed before the teardown would skip the first repair
5010 # after resubscribing, delaying the move of its shard channels
5011 # off a node that is already gone.
5012 self._next_topology_repair = 0.0
5013 # _pubsubs_generator captures node_pubsub_mapping.values() into
5014 # a local list inside ``yield from``; clearing the mapping does
5015 # not reach references already held by that captured snapshot,
5016 # so a generator suspended mid-yield-from would still surface
5017 # the now-aclose()'d per-node pubsubs after re-subscription.
5018 # Recreate it to drop the captured list. type(self) bypasses
5019 # the instance-level self-shadow established at __init__
5020 # (self._pubsubs_generator = self._pubsubs_generator()).
5021 self._pubsubs_generator = type(self)._pubsubs_generator( # type: ignore[method-assign]
5022 self
5023 )
5024 # Let parent handle self.connection disconnect under the lock
5025 # (includes disconnect, release to pool, and clearing
5026 # self.connection)
5027 await super().aclose()
5028 # Clear the reverse index so a reused instance doesn't route
5029 # against stale mappings. super().aclose() has already cleared
5030 # shard_channels.
5031 self._shard_channel_to_node.clear()
5032
5033 def _raise_on_invalid_node(
5034 self,
5035 redis_cluster: "RedisCluster",
5036 node: Optional["ClusterNode"],
5037 host: Optional[str],
5038 port: Optional[int],
5039 ) -> None:
5040 """
5041 Raise a RedisClusterException if the node is None or doesn't exist in
5042 the cluster.
5043 """
5044 if node is None or redis_cluster.get_node(node_name=node.name) is None:
5045 raise RedisClusterException(
5046 f"Node {host}:{port} doesn't exist in the cluster"
5047 )
5048
5049 async def execute_command(self, *args: Any, **kwargs: Any) -> Any:
5050 """
5051 Execute a command on the appropriate cluster node.
5052
5053 Taken code from redis-py and tweaked to make it work within a cluster.
5054 """
5055 # NOTE: don't parse the response in this function -- it could pull a
5056 # legitimate message off the stack if the connection is already
5057 # subscribed to one or more channels
5058
5059 await self._ensure_cluster_initialized()
5060
5061 # For shard commands, route to appropriate node
5062 command = args[0].upper() if args else ""
5063 if command in ("SSUBSCRIBE", "SUNSUBSCRIBE", "SPUBLISH"):
5064 if len(args) > 1:
5065 # ssubscribe / sunsubscribe own both the per-node I/O lock and
5066 # the shard_channels / _shard_channel_to_node bookkeeping, so
5067 # delegate to them instead of dispatching raw. A raw dispatch
5068 # writes the socket unguarded against a concurrent poll and
5069 # records nothing, leaving the channel invisible to the reader
5070 # loop and to on_connect's replay.
5071 if command == "SSUBSCRIBE":
5072 return await self.ssubscribe(*args[1:])
5073 if command == "SUNSUBSCRIBE":
5074 return await self.sunsubscribe(*args[1:])
5075 channel = args[1]
5076 node = self.cluster.get_node_from_key(channel)
5077 if node:
5078 pubsub = self._get_node_pubsub(node)
5079 async with self._pubsub_io_lock(pubsub):
5080 return await pubsub.execute_command(*args, **kwargs)
5081
5082 # For other commands, use the set node or lazily discover one
5083 if self.connection is None:
5084 if self.connection_pool is None:
5085 if len(args) > 1:
5086 # Hash the first channel and get one of the nodes holding
5087 # this slot
5088 channel = args[1]
5089 slot = self.cluster.keyslot(channel)
5090 node = self.cluster.nodes_manager.get_node_from_slot(
5091 slot,
5092 self.cluster.read_from_replicas,
5093 self.cluster.load_balancing_strategy,
5094 )
5095 else:
5096 # Get a random node
5097 node = self.cluster.get_random_node()
5098 self.node = node
5099 self.connection_pool = _ClusterNodePoolAdapter(node)
5100
5101 # Now we have a connection_pool, use parent's execute_command
5102 return await super().execute_command(*args, **kwargs)