Coverage for /pythoncovmergedfiles/medio/medio/usr/local/lib/python3.11/site-packages/redis/asyncio/client.py: 22%
Shortcuts on this page
r m x toggle line displays
j k next/prev highlighted chunk
0 (zero) top of page
1 (one) first highlighted chunk
Shortcuts on this page
r m x toggle line displays
j k next/prev highlighted chunk
0 (zero) top of page
1 (one) first highlighted chunk
1import asyncio
2import copy
3import inspect
4import logging
5import math
6import re
7import time
8import warnings
9from typing import (
10 TYPE_CHECKING,
11 Any,
12 AsyncIterator,
13 Awaitable,
14 Callable,
15 Dict,
16 Iterable,
17 List,
18 Literal,
19 Mapping,
20 MutableMapping,
21 Optional,
22 Protocol,
23 Sequence,
24 Set,
25 Tuple,
26 Type,
27 TypedDict,
28 TypeVar,
29 Union,
30 cast,
31)
33from redis._defaults import (
34 DEFAULT_RETRY_BASE,
35 DEFAULT_RETRY_CAP,
36 DEFAULT_RETRY_COUNT,
37 DEFAULT_SOCKET_CONNECT_TIMEOUT,
38 DEFAULT_SOCKET_READ_SIZE,
39 DEFAULT_SOCKET_TIMEOUT,
40)
41from redis._parsers.helpers import bool_ok, get_response_callbacks
42from redis.asyncio import _himport_exec
43from redis.asyncio.connection import (
44 AbstractConnection,
45 Connection,
46 ConnectionPool,
47 SSLConnection,
48 UnixDomainSocketConnection,
49)
50from redis.asyncio.lock import Lock
51from redis.asyncio.observability.recorder import (
52 record_error_count,
53 record_operation_duration,
54 record_pubsub_message,
55)
56from redis.asyncio.retry import Retry, _to_async_retry
57from redis.backoff import ExponentialWithJitterBackoff
58from redis.client import (
59 EMPTY_RESPONSE,
60 NEVER_DECODE,
61 AbstractRedis,
62 CaseInsensitiveDict,
63)
64from redis.commands import (
65 AsyncCoreCommands,
66 AsyncRedisModuleCommands,
67 AsyncSentinelCommands,
68 list_or_args,
69)
70from redis.commands.helpers import parse_pubsub_subscriptions, pubsub_subscription_args
71from redis.credentials import CredentialProvider
72from redis.driver_info import DriverInfo, resolve_driver_info
73from redis.event import (
74 AfterPooledConnectionsInstantiationEvent,
75 AfterPubSubConnectionInstantiationEvent,
76 AfterSingleConnectionInstantiationEvent,
77 ClientType,
78 EventDispatcher,
79)
80from redis.exceptions import (
81 ConnectionError,
82 ExecAbortError,
83 PubSubError,
84 RedisError,
85 ResponseError,
86 WatchError,
87)
88from redis.himport import HImportRegistry, parse_himport_set_args
89from redis.maint_notifications import MaintNotificationsConfig
90from redis.observability.attributes import PubSubDirection
91from redis.typing import (
92 ChannelT,
93 EncodableT,
94 FieldT,
95 KeyT,
96 PubSubHandler,
97 Subscription,
98)
99from redis.utils import (
100 SENTINEL,
101 SSL_AVAILABLE,
102 _set_info_logger,
103 check_protocol_version,
104 deprecated_args,
105 deprecated_function,
106 experimental_method,
107 safe_str,
108 str_if_bytes,
109 truncate_text,
110)
112if TYPE_CHECKING and SSL_AVAILABLE:
113 from ssl import TLSVersion, VerifyFlags, VerifyMode
114else:
115 TLSVersion = None
116 VerifyMode = None
117 VerifyFlags = None
119_KeyT = TypeVar("_KeyT", bound=KeyT)
120_ArgT = TypeVar("_ArgT", KeyT, EncodableT)
121_RedisT = TypeVar("_RedisT", bound="Redis")
122_NormalizeKeysT = TypeVar("_NormalizeKeysT", bound=Mapping[ChannelT, object])
123if TYPE_CHECKING:
124 from redis.asyncio.keyspace_notifications import AsyncKeyspaceNotifications
125 from redis.commands.core import Script
128logger = logging.getLogger(__name__)
131def is_debug_log_enabled():
132 return logger.isEnabledFor(logging.DEBUG)
135def add_debug_log_for_operation_failure(
136 connection: AbstractConnection,
137 error: BaseException | None = None,
138 args: Sequence[Any] | None = None,
139):
140 details = connection.extract_connection_details() if connection else "no connection"
141 prefix = (
142 f"{type(error).__name__} received" if error is not None else "Operation failed"
143 )
144 # Log only the command name - argument values can carry secrets
145 # (AUTH, CONFIG SET requirepass, ACL SETUSER) or user data.
146 command = f" for command {safe_str(args[0])}" if args else ""
147 suffix = f", error: {error}" if error is not None else ""
148 logger.debug(
149 f"{prefix}{command}, with connection: {connection}, details: {details}{suffix}",
150 )
153class ResponseCallbackProtocol(Protocol):
154 def __call__(self, response: Any, **kwargs): ...
157class AsyncResponseCallbackProtocol(Protocol):
158 async def __call__(self, response: Any, **kwargs): ...
161ResponseCallbackT = Union[ResponseCallbackProtocol, AsyncResponseCallbackProtocol]
164class Redis(
165 AbstractRedis, AsyncRedisModuleCommands, AsyncCoreCommands, AsyncSentinelCommands
166):
167 """
168 Implementation of the Redis protocol.
170 This abstract class provides a Python interface to all Redis commands
171 and an implementation of the Redis protocol.
173 Pipelines derive from this, implementing how
174 the commands are sent and received to the Redis server. Based on
175 configuration, an instance will either use a ConnectionPool, or
176 Connection object to talk to redis.
177 """
179 # Type discrimination marker for @overload self-type pattern
180 _is_async_client: Literal[True] = True
182 response_callbacks: MutableMapping[Union[str, bytes], ResponseCallbackT]
184 @classmethod
185 def from_url(
186 cls: Type["Redis"],
187 url: str,
188 single_connection_client: bool = False,
189 auto_close_connection_pool: Optional[bool] = None,
190 **kwargs,
191 ) -> "Redis":
192 """
193 Return a Redis client object configured from the given URL
195 For example::
197 redis://[[username]:[password]]@localhost:6379/0
198 rediss://[[username]:[password]]@localhost:6379/0
199 unix://[username@]/path/to/socket.sock?db=0[&password=password]
201 Three URL schemes are supported:
203 - `redis://` creates a TCP socket connection. See more at:
204 <https://www.iana.org/assignments/uri-schemes/prov/redis>
205 - `rediss://` creates a SSL wrapped TCP socket connection. See more at:
206 <https://www.iana.org/assignments/uri-schemes/prov/rediss>
207 - ``unix://``: creates a Unix Domain Socket connection.
209 The username, password, hostname and path are passed through
210 urllib.parse.unquote in order to replace any percent-encoded values
211 with their corresponding characters. Querystring values are decoded
212 by urllib.parse.parse_qs and are not unquoted again.
214 There are several ways to specify a database number. The first value
215 found will be used:
217 1. A ``db`` querystring option, e.g. redis://localhost?db=0
219 2. If using the redis:// or rediss:// schemes, the path argument
220 of the url, e.g. redis://localhost/0
222 3. A ``db`` keyword argument to this function.
224 If none of these options are specified, the default db=0 is used.
226 All querystring options are cast to their appropriate Python types.
227 Boolean arguments can be specified with string values "True"/"False"
228 or "Yes"/"No". Values that cannot be properly cast cause a
229 ``ValueError`` to be raised. Once parsed, the querystring arguments
230 and keyword arguments are passed to the ``ConnectionPool``'s
231 class initializer. In the case of conflicting arguments, querystring
232 arguments always win.
234 """
235 connection_pool = ConnectionPool.from_url(url, **kwargs)
236 client = cls(
237 connection_pool=connection_pool,
238 single_connection_client=single_connection_client,
239 )
240 if auto_close_connection_pool is not None:
241 warnings.warn(
242 DeprecationWarning(
243 '"auto_close_connection_pool" is deprecated '
244 "since version 5.0.1. "
245 "Please create a ConnectionPool explicitly and "
246 "provide to the Redis() constructor instead."
247 )
248 )
249 else:
250 auto_close_connection_pool = True
251 client.auto_close_connection_pool = auto_close_connection_pool
252 return client
254 @classmethod
255 def from_pool(
256 cls: Type["Redis"],
257 connection_pool: ConnectionPool,
258 ) -> "Redis":
259 """
260 Return a Redis client from the given connection pool.
261 The Redis client will take ownership of the connection pool and
262 close it when the Redis client is closed.
264 Because the client closes (disconnects all connections in) the pool
265 when it is closed or garbage-collected, the pool must not be shared
266 with other clients. Constructing multiple clients from the same pool
267 via ``from_pool`` -- for example one per request across tasks -- is
268 not safe: when one client is closed it will disconnect connections
269 still in use by the others.
271 To share a single pool across clients, construct the pool explicitly
272 and manage its lifecycle instead. Unlike ``from_pool``, the plain
273 ``Redis(connection_pool=pool)`` constructor does not take ownership of
274 the pool and will not close it, so a pool created this way can be
275 safely shared across clients. ``ConnectionPool`` supports the async
276 context manager protocol for this::
278 async with ConnectionPool.from_url(url) as pool:
279 r = Redis(connection_pool=pool)
280 """
281 client = cls(
282 connection_pool=connection_pool,
283 )
284 client.auto_close_connection_pool = True
285 return client
287 @deprecated_args(
288 args_to_warn=["retry_on_timeout"],
289 reason="TimeoutError is included by default.",
290 version="6.0.0",
291 )
292 @deprecated_args(
293 args_to_warn=["lib_name", "lib_version"],
294 reason="Use 'driver_info' parameter instead. "
295 "lib_name and lib_version will be removed in a future version.",
296 )
297 def __init__(
298 self,
299 *,
300 host: str = "localhost",
301 port: int = 6379,
302 db: str | int = 0,
303 password: str | None = None,
304 socket_timeout: float | None = DEFAULT_SOCKET_TIMEOUT,
305 socket_connect_timeout: float | None = DEFAULT_SOCKET_CONNECT_TIMEOUT,
306 socket_read_size: int = DEFAULT_SOCKET_READ_SIZE,
307 socket_keepalive: bool | None = True,
308 socket_keepalive_options: Mapping[int, int | bytes] | object | None = SENTINEL,
309 connection_pool: ConnectionPool | None = None,
310 unix_socket_path: str | None = None,
311 encoding: str = "utf-8",
312 encoding_errors: str = "strict",
313 decode_responses: bool = False,
314 retry_on_timeout: bool = False,
315 retry: Retry = Retry(
316 backoff=ExponentialWithJitterBackoff(
317 base=DEFAULT_RETRY_BASE, cap=DEFAULT_RETRY_CAP
318 ),
319 retries=DEFAULT_RETRY_COUNT,
320 ),
321 retry_on_error: list | None = None,
322 ssl: bool = False,
323 ssl_keyfile: str | None = None,
324 ssl_certfile: str | None = None,
325 ssl_cert_reqs: "str | VerifyMode" = "required",
326 ssl_include_verify_flags: List["VerifyFlags"] | None = None,
327 ssl_exclude_verify_flags: List["VerifyFlags"] | None = None,
328 ssl_ca_certs: str | None = None,
329 ssl_ca_data: str | None = None,
330 ssl_ca_path: str | None = None,
331 ssl_check_hostname: bool = True,
332 ssl_min_version: "TLSVersion | None" = None,
333 ssl_ciphers: str | None = None,
334 ssl_password: str | None = None,
335 max_connections: int | None = None,
336 single_connection_client: bool = False,
337 health_check_interval: int = 0,
338 client_name: str | None = None,
339 lib_name: str | object | None = SENTINEL,
340 lib_version: str | object | None = SENTINEL,
341 driver_info: DriverInfo | object | None = SENTINEL,
342 username: str | None = None,
343 auto_close_connection_pool: bool | None = None,
344 redis_connect_func=None,
345 credential_provider: CredentialProvider | None = None,
346 protocol: int | None = None,
347 legacy_responses: bool = True,
348 event_dispatcher: EventDispatcher | None = None,
349 maint_notifications_config: MaintNotificationsConfig | None = None,
350 ):
351 """
352 Initialize a new Redis client.
354 To specify a retry policy for specific errors, you have two options:
356 1. Set the `retry_on_error` to a list of the error/s to retry on, and
357 you can also set `retry` to a valid `Retry` object(in case the default
358 one is not appropriate) - with this approach the retries will be triggered
359 on the default errors specified in the Retry object enriched with the
360 errors specified in `retry_on_error`.
362 2. Define a `Retry` object with configured 'supported_errors' and set
363 it to the `retry` parameter - with this approach you completely redefine
364 the errors on which retries will happen.
366 `retry_on_timeout` is deprecated - please include the TimeoutError
367 either in the Retry object or in the `retry_on_error` list.
369 When 'connection_pool' is provided - the retry configuration of the
370 provided pool will be used.
372 Args:
374 socket_keepalive:
375 if `True`, TCP keepalive is enabled for TCP socket connections.
376 Argument is ignored when connection_pool is provided.
377 socket_keepalive_options:
378 mapping of TCP keepalive socket option constants to values, for
379 example `{socket.TCP_KEEPIDLE: 30}`. If left unspecified, redis-py
380 uses TCP keepalive defaults when `socket_keepalive` is enabled:
381 idle 30 seconds, interval 5 seconds, and 3 probes. Platform-specific
382 options that are not available are skipped. Pass `None` or `{}` to
383 avoid setting additional TCP keepalive options. Argument is ignored
384 when connection_pool is provided.
385 maint_notifications_config:
386 configures the pool to support maintenance notifications - see
387 `redis.maint_notifications.MaintNotificationsConfig` for details.
388 Only supported with RESP3
389 If not provided and protocol is RESP3, the maintenance notifications
390 will be enabled by default (logic is included in the connection pool
391 initialization).
392 Argument is ignored when connection_pool is provided.
393 """
394 kwargs: Dict[str, Any]
395 if event_dispatcher is None:
396 self._event_dispatcher = EventDispatcher()
397 else:
398 self._event_dispatcher = event_dispatcher
399 # auto_close_connection_pool only has an effect if connection_pool is
400 # None. It is assumed that if connection_pool is not None, the user
401 # wants to manage the connection pool themselves.
402 if auto_close_connection_pool is not None:
403 warnings.warn(
404 DeprecationWarning(
405 '"auto_close_connection_pool" is deprecated '
406 "since version 5.0.1. "
407 "Please create a ConnectionPool explicitly and "
408 "provide to the Redis() constructor instead."
409 )
410 )
411 else:
412 auto_close_connection_pool = True
414 if not connection_pool:
415 # Create internal connection pool, expected to be closed by Redis instance
416 if not retry_on_error:
417 retry_on_error = []
419 # Handle driver_info: if provided, use it; otherwise create from lib_name/lib_version.
420 computed_driver_info = resolve_driver_info(
421 driver_info, lib_name, lib_version
422 )
424 kwargs = {
425 "db": db,
426 "username": username,
427 "password": password,
428 "credential_provider": credential_provider,
429 "socket_timeout": socket_timeout,
430 "socket_read_size": socket_read_size,
431 "encoding": encoding,
432 "encoding_errors": encoding_errors,
433 "decode_responses": decode_responses,
434 "retry_on_error": retry_on_error,
435 "retry": copy.deepcopy(retry),
436 "max_connections": max_connections,
437 "health_check_interval": health_check_interval,
438 "client_name": client_name,
439 "driver_info": computed_driver_info,
440 "redis_connect_func": redis_connect_func,
441 "protocol": protocol,
442 "legacy_responses": legacy_responses,
443 }
444 # based on input, setup appropriate connection args
445 if unix_socket_path is not None:
446 if (
447 maint_notifications_config
448 and maint_notifications_config.enabled is True
449 ):
450 raise RedisError(
451 "Maintenance notifications are not supported with Unix "
452 "domain socket connections"
453 )
454 kwargs.update(
455 {
456 "path": unix_socket_path,
457 "connection_class": UnixDomainSocketConnection,
458 "maint_notifications_config": MaintNotificationsConfig(
459 enabled=False
460 ),
461 }
462 )
463 else:
464 # TCP specific options
465 kwargs.update(
466 {
467 "host": host,
468 "port": port,
469 "socket_connect_timeout": socket_connect_timeout,
470 "socket_keepalive": socket_keepalive,
471 "socket_keepalive_options": socket_keepalive_options,
472 }
473 )
475 if ssl:
476 kwargs.update(
477 {
478 "connection_class": SSLConnection,
479 "ssl_keyfile": ssl_keyfile,
480 "ssl_certfile": ssl_certfile,
481 "ssl_cert_reqs": ssl_cert_reqs,
482 "ssl_include_verify_flags": ssl_include_verify_flags,
483 "ssl_exclude_verify_flags": ssl_exclude_verify_flags,
484 "ssl_ca_certs": ssl_ca_certs,
485 "ssl_ca_data": ssl_ca_data,
486 "ssl_ca_path": ssl_ca_path,
487 "ssl_check_hostname": ssl_check_hostname,
488 "ssl_min_version": ssl_min_version,
489 "ssl_ciphers": ssl_ciphers,
490 "ssl_password": ssl_password,
491 }
492 )
493 maint_notifications_enabled = (
494 maint_notifications_config and maint_notifications_config.enabled
495 )
496 if maint_notifications_enabled and not check_protocol_version(protocol, 3):
497 raise RedisError(
498 "Maintenance notifications handlers on connection are only supported with RESP version 3"
499 )
500 if maint_notifications_config:
501 kwargs.update(
502 {
503 "maint_notifications_config": maint_notifications_config,
504 }
505 )
506 # This arg only used if no pool is passed in
507 self.auto_close_connection_pool = auto_close_connection_pool
508 connection_pool = ConnectionPool(**kwargs)
509 self._event_dispatcher.dispatch(
510 AfterPooledConnectionsInstantiationEvent(
511 [connection_pool], ClientType.ASYNC, credential_provider
512 )
513 )
514 else:
515 # If a pool is passed in, do not close it
516 self.auto_close_connection_pool = False
517 self._event_dispatcher.dispatch(
518 AfterPooledConnectionsInstantiationEvent(
519 [connection_pool], ClientType.ASYNC, credential_provider
520 )
521 )
523 self.connection_pool = connection_pool
524 self.single_connection_client = single_connection_client
525 self.connection: Optional[Connection] = None
527 connection_kwargs = self.connection_pool.connection_kwargs
528 self.response_callbacks = CaseInsensitiveDict(
529 get_response_callbacks(
530 user_protocol=connection_kwargs.get("protocol"),
531 legacy_responses=connection_kwargs.get("legacy_responses", True),
532 )
533 )
535 # If using a single connection client, we need to lock creation-of and use-of
536 # the client in order to avoid race conditions such as using asyncio.gather
537 # on a set of redis commands
538 self._single_conn_lock = asyncio.Lock()
540 # When used as an async context manager, we need to increment and decrement
541 # a usage counter so that we can close the connection pool when no one is
542 # using the client.
543 self._usage_counter = 0
544 self._usage_lock = asyncio.Lock()
546 def __repr__(self):
547 return (
548 f"<{self.__class__.__module__}.{self.__class__.__name__}"
549 f"({self.connection_pool!r})>"
550 )
552 def __await__(self):
553 return self.initialize().__await__()
555 async def initialize(self: _RedisT) -> _RedisT:
556 if self.single_connection_client:
557 async with self._single_conn_lock:
558 if self.connection is None:
559 self.connection = await self.connection_pool.get_connection()
561 self._event_dispatcher.dispatch(
562 AfterSingleConnectionInstantiationEvent(
563 self.connection, ClientType.ASYNC, self._single_conn_lock
564 )
565 )
566 return self
568 def set_response_callback(self, command: str, callback: ResponseCallbackT):
569 """Set a custom Response Callback"""
570 self.response_callbacks[command] = callback
572 def get_encoder(self):
573 """Get the connection pool's encoder"""
574 return self.connection_pool.get_encoder()
576 def get_connection_kwargs(self):
577 """Get the connection's key-word arguments"""
578 return self.connection_pool.connection_kwargs
580 @property
581 def himport_registry(self) -> HImportRegistry:
582 """The client's HIMPORT fieldset registry (empty if none was declared).
584 Read-only: the registry is mutated only through the HIMPORT command methods.
585 """
586 return self.connection_pool.himport_registry
588 def get_retry(self) -> Optional[Retry]:
589 return self.get_connection_kwargs().get("retry")
591 def set_retry(self, retry: Retry) -> None:
592 retry = _to_async_retry(retry)
593 self.get_connection_kwargs().update({"retry": retry})
594 self.connection_pool.set_retry(retry)
596 def load_external_module(self, funcname, func):
597 """
598 This function can be used to add externally defined redis modules,
599 and their namespaces to the redis client.
601 funcname - A string containing the name of the function to create
602 func - The function, being added to this class.
604 ex: Assume that one has a custom redis module named foomod that
605 creates command named 'foo.dothing' and 'foo.anotherthing' in redis.
606 To load function functions into this namespace:
608 from redis import Redis
609 from foomodule import F
610 r = Redis()
611 r.load_external_module("foo", F)
612 r.foo().dothing('your', 'arguments')
614 For a concrete example see the reimport of the redisjson module in
615 tests/test_connection.py::test_loading_external_modules
616 """
617 setattr(self, funcname, func)
619 def pipeline(
620 self, transaction: bool = True, shard_hint: Optional[str] = None
621 ) -> "Pipeline":
622 """
623 Return a new pipeline object that can queue multiple commands for
624 later execution. ``transaction`` indicates whether all commands
625 should be executed atomically. Apart from making a group of operations
626 atomic, pipelines are useful for reducing the back-and-forth overhead
627 between the client and server.
628 """
629 return Pipeline(
630 self.connection_pool, self.response_callbacks, transaction, shard_hint
631 )
633 async def transaction(
634 self,
635 func: Callable[["Pipeline"], Union[Any, Awaitable[Any]]],
636 *watches: KeyT,
637 shard_hint: Optional[str] = None,
638 value_from_callable: bool = False,
639 watch_delay: Optional[float] = None,
640 ):
641 """
642 Convenience method for executing the callable `func` as a transaction
643 while watching all keys specified in `watches`. The 'func' callable
644 should expect a single argument which is a Pipeline object.
645 """
646 pipe: Pipeline
647 async with self.pipeline(True, shard_hint) as pipe:
648 while True:
649 try:
650 if watches:
651 await pipe.watch(*watches)
652 func_value = func(pipe)
653 if inspect.isawaitable(func_value):
654 func_value = await func_value
655 exec_value = await pipe.execute()
656 return func_value if value_from_callable else exec_value
657 except WatchError:
658 if watch_delay is not None and watch_delay > 0:
659 await asyncio.sleep(watch_delay)
660 continue
662 def lock(
663 self,
664 name: KeyT,
665 timeout: Optional[float] = None,
666 sleep: float = 0.1,
667 blocking: bool = True,
668 blocking_timeout: Optional[float] = None,
669 lock_class: Optional[Type[Lock]] = None,
670 thread_local: bool = True,
671 raise_on_release_error: bool = True,
672 ) -> Lock:
673 """
674 Return a new Lock object using key ``name`` that mimics
675 the behavior of threading.Lock.
677 If specified, ``timeout`` indicates a maximum life for the lock.
678 By default, it will remain locked until release() is called.
680 ``sleep`` indicates the amount of time to sleep per loop iteration
681 when the lock is in blocking mode and another client is currently
682 holding the lock.
684 ``blocking`` indicates whether calling ``acquire`` should block until
685 the lock has been acquired or to fail immediately, causing ``acquire``
686 to return False and the lock not being acquired. Defaults to True.
687 Note this value can be overridden by passing a ``blocking``
688 argument to ``acquire``.
690 ``blocking_timeout`` indicates the maximum amount of time in seconds to
691 spend trying to acquire the lock. A value of ``None`` indicates
692 continue trying forever. ``blocking_timeout`` can be specified as a
693 float or integer, both representing the number of seconds to wait.
695 ``lock_class`` forces the specified lock implementation. Note that as
696 of redis-py 3.0, the only lock class we implement is ``Lock`` (which is
697 a Lua-based lock). So, it's unlikely you'll need this parameter, unless
698 you have created your own custom lock class.
700 ``thread_local`` indicates whether the lock token is placed in
701 thread-local storage. By default, the token is placed in thread local
702 storage so that a thread only sees its token, not a token set by
703 another thread. Consider the following timeline:
705 time: 0, thread-1 acquires `my-lock`, with a timeout of 5 seconds.
706 thread-1 sets the token to "abc"
707 time: 1, thread-2 blocks trying to acquire `my-lock` using the
708 Lock instance.
709 time: 5, thread-1 has not yet completed. redis expires the lock
710 key.
711 time: 5, thread-2 acquired `my-lock` now that it's available.
712 thread-2 sets the token to "xyz"
713 time: 6, thread-1 finishes its work and calls release(). if the
714 token is *not* stored in thread local storage, then
715 thread-1 would see the token value as "xyz" and would be
716 able to successfully release the thread-2's lock.
718 ``raise_on_release_error`` indicates whether to raise an exception when
719 the lock is no longer owned when exiting the context manager. By default,
720 this is True, meaning an exception will be raised. If False, the warning
721 will be logged and the exception will be suppressed.
723 In some use cases it's necessary to disable thread local storage. For
724 example, if you have code where one thread acquires a lock and passes
725 that lock instance to a worker thread to release later. If thread
726 local storage isn't disabled in this case, the worker thread won't see
727 the token set by the thread that acquired the lock. Our assumption
728 is that these cases aren't common and as such default to using
729 thread local storage."""
730 if lock_class is None:
731 lock_class = Lock
732 return lock_class(
733 self,
734 name,
735 timeout=timeout,
736 sleep=sleep,
737 blocking=blocking,
738 blocking_timeout=blocking_timeout,
739 thread_local=thread_local,
740 raise_on_release_error=raise_on_release_error,
741 )
743 def pubsub(self, **kwargs) -> "PubSub":
744 """
745 Return a Publish/Subscribe object. With this object, you can
746 subscribe to channels and listen for messages that get published to
747 them.
748 """
749 return PubSub(
750 self.connection_pool, event_dispatcher=self._event_dispatcher, **kwargs
751 )
753 def keyspace_notifications(
754 self,
755 key_prefix: Union[str, bytes, None] = None,
756 ignore_subscribe_messages: bool = True,
757 ) -> "AsyncKeyspaceNotifications":
758 """
759 Return an :class:`~redis.asyncio.keyspace_notifications.AsyncKeyspaceNotifications`
760 object for subscribing to keyspace and keyevent notifications.
762 Note: Keyspace notifications must be enabled on the Redis server via
763 the ``notify-keyspace-events`` configuration option.
765 Args:
766 key_prefix: Optional prefix to filter and strip from keys in
767 notifications.
768 ignore_subscribe_messages: If True, subscribe/unsubscribe
769 confirmations are not returned by
770 get_message/listen.
771 """
772 from redis.asyncio.keyspace_notifications import AsyncKeyspaceNotifications
774 return AsyncKeyspaceNotifications(
775 self,
776 key_prefix=key_prefix,
777 ignore_subscribe_messages=ignore_subscribe_messages,
778 )
780 def monitor(self) -> "Monitor":
781 return Monitor(self.connection_pool)
783 def client(self) -> "Redis":
784 return self.__class__(
785 connection_pool=self.connection_pool, single_connection_client=True
786 )
788 async def __aenter__(self: _RedisT) -> _RedisT:
789 """
790 Async context manager entry. Increments a usage counter so that the
791 connection pool is only closed (via aclose()) when no context is using
792 the client.
793 """
794 await self._increment_usage()
795 try:
796 # Initialize the client (i.e. establish connection, etc.)
797 return await self.initialize()
798 except Exception:
799 # If initialization fails, decrement the counter to keep it in sync
800 await self._decrement_usage()
801 raise
803 async def _increment_usage(self) -> int:
804 """
805 Helper coroutine to increment the usage counter while holding the lock.
806 Returns the new value of the usage counter.
807 """
808 async with self._usage_lock:
809 self._usage_counter += 1
810 return self._usage_counter
812 async def _decrement_usage(self) -> int:
813 """
814 Helper coroutine to decrement the usage counter while holding the lock.
815 Returns the new value of the usage counter.
816 """
817 async with self._usage_lock:
818 self._usage_counter -= 1
819 return self._usage_counter
821 async def __aexit__(self, exc_type, exc_value, traceback):
822 """
823 Async context manager exit. Decrements a usage counter. If this is the
824 last exit (counter becomes zero), the client closes its connection pool.
825 """
826 current_usage = await asyncio.shield(self._decrement_usage())
827 if current_usage == 0:
828 # This was the last active context, so disconnect the pool.
829 await asyncio.shield(self.aclose())
831 _DEL_MESSAGE = "Unclosed Redis client"
833 # passing _warnings and _grl as argument default since they may be gone
834 # by the time __del__ is called at shutdown
835 def __del__(
836 self,
837 _warn: Any = warnings.warn,
838 _grl: Any = asyncio.get_running_loop,
839 ) -> None:
840 if hasattr(self, "connection") and (self.connection is not None):
841 _warn(f"Unclosed client session {self!r}", ResourceWarning, source=self)
842 try:
843 context = {"client": self, "message": self._DEL_MESSAGE}
844 _grl().call_exception_handler(context)
845 except RuntimeError:
846 pass
847 self.connection._close()
849 async def aclose(self, close_connection_pool: Optional[bool] = None) -> None:
850 """
851 Closes Redis client connection
853 Args:
854 close_connection_pool:
855 decides whether to close the connection pool used by this Redis client,
856 overriding Redis.auto_close_connection_pool.
857 By default, let Redis.auto_close_connection_pool decide
858 whether to close the connection pool.
859 """
860 conn = self.connection
861 if conn:
862 self.connection = None
863 await self.connection_pool.release(conn)
864 if close_connection_pool or (
865 close_connection_pool is None and self.auto_close_connection_pool
866 ):
867 await self.connection_pool.aclose()
869 @deprecated_function(version="5.0.1", reason="Use aclose() instead", name="close")
870 async def close(self, close_connection_pool: Optional[bool] = None) -> None:
871 """
872 Alias for aclose(), for backwards compatibility
873 """
874 await self.aclose(close_connection_pool)
876 async def _send_command_parse_response(self, conn, command_name, *args, **options):
877 """
878 Send a command and parse the response
879 """
880 # HIMPORT SET is the one command whose wire form depends on per-connection
881 # state: the fieldset must be PREPAREd on this connection first, and any
882 # fieldset discarded since this connection last reconciled must be dropped.
883 # Handling it here (rather than in himport_set) lets himport_set reuse the
884 # full execute_command machinery — retry, disconnect-on-error, pooling — so
885 # a failed HIMPORT SET disconnects the connection like any other command.
886 # This per-command branch in the hot dispatch path is deliberate and has no
887 # cleaner alternative: this is the only seam where the concrete borrowed
888 # connection is known, and connection-scoped session setup can only happen
889 # once that connection is chosen. The overhead is one string compare per
890 # command.
891 himport_set = parse_himport_set_args(args)
892 if himport_set is not None:
893 # ``args`` is an HIMPORT SET in either the joined ("HIMPORT SET", key,
894 # ...) or split ("HIMPORT", "SET", key, ...) raw form; the operands come
895 # back at the right offsets for the form. A command with too few operands
896 # returns None and falls through to the normal send path so the server
897 # returns its arity error instead of a client-side IndexError here.
898 key, fieldset_name, values = himport_set
899 return await self._himport_execute_set(conn, key, fieldset_name, values)
900 await conn.send_command(*args)
901 return await self.parse_response(conn, command_name, **options)
903 async def _himport_reconcile_discards(self, conn):
904 """Delegate to the shared async HIMPORT executor."""
905 return await _himport_exec.reconcile_discards(self, conn)
907 async def _himport_prepare_and_set(
908 self, conn, key, fieldset_name, values, fieldset
909 ):
910 """Delegate to the shared async HIMPORT executor."""
911 return await _himport_exec.prepare_and_set(
912 self, conn, key, fieldset_name, values, fieldset
913 )
915 async def _himport_execute_set(self, conn, key, fieldset_name, values):
916 """Delegate to the shared async HIMPORT executor."""
917 return await _himport_exec.execute_set(self, conn, key, fieldset_name, values)
919 async def _close_connection(
920 self,
921 conn: Connection,
922 error: Optional[BaseException] = None,
923 failure_count: Optional[int] = None,
924 start_time: Optional[float] = None,
925 command_name: Optional[str] = None,
926 ):
927 """
928 Close the connection before retrying.
930 The supported exceptions are already checked in the
931 retry object so we don't need to do it here.
933 After we disconnect the connection, it will try to reconnect and
934 do a health check as part of the send_command logic(on connection level).
935 """
936 if (
937 error
938 and failure_count is not None
939 and failure_count <= conn.retry.get_retries()
940 ):
941 await record_operation_duration(
942 command_name=command_name,
943 duration_seconds=time.monotonic() - start_time,
944 server_address=getattr(conn, "host", None),
945 server_port=getattr(conn, "port", None),
946 db_namespace=str(conn.db),
947 error=error,
948 retry_attempts=failure_count,
949 )
951 await conn.disconnect(error=error, failure_count=failure_count)
953 # COMMAND EXECUTION AND PROTOCOL PARSING
954 async def execute_command(self, *args, **options):
955 """Execute a command and return a parsed response"""
956 await self.initialize()
957 pool = self.connection_pool
958 command_name = args[0]
959 conn = self.connection or await pool.get_connection()
961 # Start timing for observability
962 start_time = time.monotonic()
963 # Track actual retry attempts for error reporting
964 actual_retry_attempts = 0
966 def failure_callback(error, failure_count):
967 if is_debug_log_enabled():
968 add_debug_log_for_operation_failure(conn, error, args)
969 nonlocal actual_retry_attempts
970 actual_retry_attempts = failure_count
971 return self._close_connection(
972 conn, error, failure_count, start_time, command_name
973 )
975 if self.single_connection_client:
976 await self._single_conn_lock.acquire()
977 try:
978 result = await conn.retry.call_with_retry(
979 lambda: self._send_command_parse_response(
980 conn, command_name, *args, **options
981 ),
982 failure_callback,
983 with_failure_count=True,
984 )
986 await record_operation_duration(
987 command_name=command_name,
988 duration_seconds=time.monotonic() - start_time,
989 server_address=getattr(conn, "host", None),
990 server_port=getattr(conn, "port", None),
991 db_namespace=str(conn.db),
992 )
993 return result
994 except Exception as e:
995 await record_error_count(
996 server_address=getattr(conn, "host", None),
997 server_port=getattr(conn, "port", None),
998 network_peer_address=getattr(conn, "host", None),
999 network_peer_port=getattr(conn, "port", None),
1000 error_type=e,
1001 retry_attempts=actual_retry_attempts,
1002 is_internal=False,
1003 )
1004 raise
1005 finally:
1006 try:
1007 if self.single_connection_client and conn and conn.should_reconnect():
1008 await self._close_connection(conn)
1009 await conn.connect()
1010 finally:
1011 if self.single_connection_client:
1012 self._single_conn_lock.release()
1013 if not self.connection:
1014 await pool.release(conn)
1016 async def parse_response(
1017 self, connection: Connection, command_name: Union[str, bytes], **options
1018 ):
1019 """Parses a response from the Redis server"""
1020 try:
1021 if NEVER_DECODE in options:
1022 response = await connection.read_response(disable_decoding=True)
1023 options.pop(NEVER_DECODE)
1024 else:
1025 response = await connection.read_response()
1026 except ResponseError:
1027 if EMPTY_RESPONSE in options:
1028 return options[EMPTY_RESPONSE]
1029 raise
1031 if EMPTY_RESPONSE in options:
1032 options.pop(EMPTY_RESPONSE)
1034 # Remove keys entry, it needs only for cache.
1035 options.pop("keys", None)
1037 if command_name in self.response_callbacks:
1038 # Mypy bug: https://github.com/python/mypy/issues/10977
1039 command_name = cast(str, command_name)
1040 retval = self.response_callbacks[command_name](response, **options)
1041 return await retval if inspect.isawaitable(retval) else retval
1042 return response
1044 # HIMPORT orchestration (async mirror of redis.client.Redis). See
1045 # ``.agents/himport_client_support_spec.md``.
1047 @experimental_method()
1048 async def himport_prepare(
1049 self, fieldset_name: str, fields: Iterable[FieldT]
1050 ) -> bool:
1051 """Declare an HIMPORT fieldset for use by :meth:`himport_set`."""
1052 await self.initialize()
1053 fieldset = self.himport_registry.prepare(fieldset_name, fields)
1054 conn = self.connection
1055 if self.single_connection_client and conn is not None and conn.is_connected:
1056 await self.himport_prepare_internal(fieldset_name, fieldset.fields)
1057 conn._himport_prepared[fieldset_name] = fieldset.version
1058 return True
1060 @experimental_method()
1061 async def himport_discard(self, fieldset_name: str) -> int:
1062 """Remove a fieldset from the registry."""
1063 await self.initialize()
1064 removed = self.himport_registry.discard(fieldset_name)
1065 conn = self.connection
1066 if self.single_connection_client and conn is not None and conn.is_connected:
1067 if removed:
1068 await self.himport_discard_internal(fieldset_name)
1069 conn._himport_prepared.pop(fieldset_name, None)
1070 conn._himport_reconciled_revision = self.himport_registry.revision
1071 return 1 if removed else 0
1073 @experimental_method()
1074 async def himport_discard_all(self) -> int:
1075 """Remove all fieldsets from the registry."""
1076 await self.initialize()
1077 count = self.himport_registry.discard_all()
1078 conn = self.connection
1079 if self.single_connection_client and conn is not None and conn.is_connected:
1080 if count:
1081 await self.himport_discard_all_internal()
1082 conn._himport_prepared.clear()
1083 conn._himport_reconciled_revision = self.himport_registry.revision
1084 return count
1087StrictRedis = Redis
1090class MonitorCommandInfo(TypedDict):
1091 time: float
1092 db: int
1093 client_address: str
1094 client_port: str
1095 client_type: str
1096 command: str
1099class Monitor:
1100 """
1101 Monitor is useful for handling the MONITOR command to the redis server.
1102 next_command() method returns one command from monitor
1103 listen() method yields commands from monitor.
1104 """
1106 monitor_re = re.compile(r"\[(\d+) (.*?)\] (.*)")
1107 command_re = re.compile(r'"(.*?)(?<!\\)"')
1109 def __init__(self, connection_pool: ConnectionPool):
1110 self.connection_pool = connection_pool
1111 self.connection: Optional[Connection] = None
1113 async def connect(self):
1114 if self.connection is None:
1115 self.connection = await self.connection_pool.get_connection()
1117 async def __aenter__(self):
1118 await self.connect()
1119 await self.connection.send_command("MONITOR")
1120 # check that monitor returns 'OK', but don't return it to user
1121 response = await self.connection.read_response()
1122 if not bool_ok(response):
1123 raise RedisError(f"MONITOR failed: {response}")
1124 return self
1126 async def __aexit__(self, *args):
1127 await self.connection.disconnect()
1128 await self.connection_pool.release(self.connection)
1130 async def next_command(self) -> MonitorCommandInfo:
1131 """Parse the response from a monitor command"""
1132 await self.connect()
1133 response = await self.connection.read_response()
1134 if isinstance(response, bytes):
1135 response = self.connection.encoder.decode(response, force=True)
1136 command_time, command_data = response.split(" ", 1)
1137 m = self.monitor_re.match(command_data)
1138 db_id, client_info, command = m.groups()
1139 command = " ".join(self.command_re.findall(command))
1140 # Redis escapes double quotes because each piece of the command
1141 # string is surrounded by double quotes. We don't have that
1142 # requirement so remove the escaping and leave the quote.
1143 command = command.replace('\\"', '"')
1145 if client_info == "lua":
1146 client_address = "lua"
1147 client_port = ""
1148 client_type = "lua"
1149 elif client_info.startswith("unix"):
1150 client_address = "unix"
1151 client_port = client_info[5:]
1152 client_type = "unix"
1153 else:
1154 # use rsplit as ipv6 addresses contain colons
1155 client_address, client_port = client_info.rsplit(":", 1)
1156 client_type = "tcp"
1157 return {
1158 "time": float(command_time),
1159 "db": int(db_id),
1160 "client_address": client_address,
1161 "client_port": client_port,
1162 "client_type": client_type,
1163 "command": command,
1164 }
1166 async def listen(self) -> AsyncIterator[MonitorCommandInfo]:
1167 """Listen for commands coming to the server."""
1168 while True:
1169 yield await self.next_command()
1172class PubSub:
1173 """
1174 PubSub provides publish, subscribe and listen support to Redis channels.
1176 After subscribing to one or more channels, the listen() method will block
1177 until a message arrives on one of the subscribed channels. That message
1178 will be returned and it's safe to start listening again.
1179 """
1181 PUBLISH_MESSAGE_TYPES = ("message", "pmessage", "smessage")
1182 UNSUBSCRIBE_MESSAGE_TYPES = ("unsubscribe", "punsubscribe", "sunsubscribe")
1183 HEALTH_CHECK_MESSAGE = "redis-py-health-check"
1185 def __init__(
1186 self,
1187 connection_pool: ConnectionPool,
1188 shard_hint: Optional[str] = None,
1189 ignore_subscribe_messages: bool = False,
1190 encoder=None,
1191 push_handler_func: Optional[Callable] = None,
1192 event_dispatcher: Optional["EventDispatcher"] = None,
1193 ):
1194 if event_dispatcher is None:
1195 self._event_dispatcher = EventDispatcher()
1196 else:
1197 self._event_dispatcher = event_dispatcher
1198 self.connection_pool = connection_pool
1199 self.shard_hint = shard_hint
1200 self.ignore_subscribe_messages = ignore_subscribe_messages
1201 self.connection = None
1202 # we need to know the encoding options for this connection in order
1203 # to lookup channel and pattern names for callback handlers.
1204 self.encoder = encoder
1205 self.push_handler_func = push_handler_func
1206 if self.encoder is None:
1207 self.encoder = self.connection_pool.get_encoder()
1208 if self.encoder.decode_responses:
1209 self.health_check_response = [
1210 ["pong", self.HEALTH_CHECK_MESSAGE],
1211 self.HEALTH_CHECK_MESSAGE,
1212 ]
1213 else:
1214 self.health_check_response = [
1215 [b"pong", self.encoder.encode(self.HEALTH_CHECK_MESSAGE)],
1216 self.encoder.encode(self.HEALTH_CHECK_MESSAGE),
1217 ]
1218 if self.push_handler_func is None:
1219 _set_info_logger()
1220 self.channels = {}
1221 self.pending_unsubscribe_channels = set()
1222 self.patterns = {}
1223 self.pending_unsubscribe_patterns = set()
1224 self.shard_channels = {}
1225 self.pending_unsubscribe_shard_channels = set()
1226 self._lock = asyncio.Lock()
1228 async def __aenter__(self):
1229 return self
1231 async def __aexit__(self, exc_type, exc_value, traceback):
1232 await self.aclose()
1234 def __del__(self):
1235 if self.connection:
1236 self.connection.deregister_connect_callback(self.on_connect)
1238 async def aclose(self):
1239 # In case a connection property does not yet exist
1240 # (due to a crash earlier in the Redis() constructor), return
1241 # immediately as there is nothing to clean-up.
1242 if not hasattr(self, "connection"):
1243 return
1244 async with self._lock:
1245 if self.connection:
1246 # Use nowait=True to avoid awaiting StreamWriter.wait_closed(),
1247 # which can deadlock when a concurrent reader task (e.g. one
1248 # running pubsub.run() or get_message(block=True)) still holds
1249 # the transport. See https://github.com/redis/redis-py/issues/3941
1250 await self.connection.disconnect(nowait=True)
1251 self.connection.deregister_connect_callback(self.on_connect)
1252 await self.connection_pool.release(self.connection)
1253 self.connection = None
1254 self.channels = {}
1255 self.pending_unsubscribe_channels = set()
1256 self.patterns = {}
1257 self.pending_unsubscribe_patterns = set()
1258 self.shard_channels = {}
1259 self.pending_unsubscribe_shard_channels = set()
1261 @deprecated_function(version="5.0.1", reason="Use aclose() instead", name="close")
1262 async def close(self) -> None:
1263 """Alias for aclose(), for backwards compatibility"""
1264 await self.aclose()
1266 @deprecated_function(version="5.0.1", reason="Use aclose() instead", name="reset")
1267 async def reset(self) -> None:
1268 """Alias for aclose(), for backwards compatibility"""
1269 await self.aclose()
1271 async def _resubscribe(self, subscribed, subscribe_fn) -> None:
1272 # Replay handler-backed subscriptions as positional Subscription objects
1273 # so binary names never need to be decoded into keyword argument keys.
1274 subscriptions = pubsub_subscription_args(subscribed)
1275 if subscriptions:
1276 await subscribe_fn(*subscriptions)
1278 async def _resubscribe_shard_channels(self) -> None:
1279 await self._resubscribe(self.shard_channels, self.ssubscribe)
1281 async def on_connect(self, connection: Connection):
1282 """Re-subscribe to any channels and patterns previously subscribed to"""
1283 self.pending_unsubscribe_channels.clear()
1284 self.pending_unsubscribe_patterns.clear()
1285 self.pending_unsubscribe_shard_channels.clear()
1286 if self.channels:
1287 await self._resubscribe(self.channels, self.subscribe)
1288 if self.patterns:
1289 await self._resubscribe(self.patterns, self.psubscribe)
1290 if self.shard_channels:
1291 await self._resubscribe_shard_channels()
1293 @property
1294 def subscribed(self):
1295 """Indicates if there are subscriptions to any channels or patterns"""
1296 return bool(self.channels or self.patterns or self.shard_channels)
1298 async def execute_command(self, *args: EncodableT):
1299 """Execute a publish/subscribe command"""
1301 # NOTE: don't parse the response in this function -- it could pull a
1302 # legitimate message off the stack if the connection is already
1303 # subscribed to one or more channels
1305 await self.connect()
1306 connection = self.connection
1307 kwargs = {"check_health": not self.subscribed}
1308 await self._execute(connection, connection.send_command, *args, **kwargs)
1310 async def connect(self):
1311 """
1312 Ensure that the PubSub is connected
1313 """
1314 if self.connection is None:
1315 self.connection = await self.connection_pool.get_connection()
1316 # register a callback that re-subscribes to any channels we
1317 # were listening to when we were disconnected
1318 self.connection.register_connect_callback(self.on_connect)
1319 else:
1320 await self.connection.connect()
1321 if self.push_handler_func is not None:
1322 self.connection._parser.set_pubsub_push_handler(self.push_handler_func)
1324 self._event_dispatcher.dispatch(
1325 AfterPubSubConnectionInstantiationEvent(
1326 self.connection, self.connection_pool, ClientType.ASYNC, self._lock
1327 )
1328 )
1330 async def _reconnect(
1331 self,
1332 conn,
1333 error: Optional[BaseException] = None,
1334 failure_count: Optional[int] = None,
1335 start_time: Optional[float] = None,
1336 command_name: Optional[str] = None,
1337 ):
1338 """
1339 The supported exceptions are already checked in the
1340 retry object so we don't need to do it here.
1342 In this error handler we are trying to reconnect to the server.
1343 """
1344 if (
1345 error
1346 and failure_count is not None
1347 and failure_count <= conn.retry.get_retries()
1348 ):
1349 if command_name:
1350 await record_operation_duration(
1351 command_name=command_name,
1352 duration_seconds=time.monotonic() - start_time,
1353 server_address=getattr(conn, "host", None),
1354 server_port=getattr(conn, "port", None),
1355 db_namespace=str(conn.db),
1356 error=error,
1357 retry_attempts=failure_count,
1358 )
1359 await conn.disconnect(error=error, failure_count=failure_count)
1360 await conn.connect()
1362 async def _execute(self, conn, command, *args, **kwargs):
1363 """
1364 Connect manually upon disconnection. If the Redis server is down,
1365 this will fail and raise a ConnectionError as desired.
1366 After reconnection, the ``on_connect`` callback should have been
1367 called by the # connection to resubscribe us to any channels and
1368 patterns we were previously listening to
1369 """
1370 if not len(args) == 0:
1371 command_name = args[0]
1372 else:
1373 command_name = None
1375 # Start timing for observability
1376 start_time = time.monotonic()
1377 # Track actual retry attempts for error reporting
1378 actual_retry_attempts = 0
1380 def failure_callback(error, failure_count):
1381 if is_debug_log_enabled():
1382 add_debug_log_for_operation_failure(conn, error, args)
1383 nonlocal actual_retry_attempts
1384 actual_retry_attempts = failure_count
1385 return self._reconnect(conn, error, failure_count, start_time, command_name)
1387 try:
1388 response = await conn.retry.call_with_retry(
1389 lambda: command(*args, **kwargs),
1390 failure_callback,
1391 with_failure_count=True,
1392 )
1394 if command_name:
1395 await record_operation_duration(
1396 command_name=command_name,
1397 duration_seconds=time.monotonic() - start_time,
1398 server_address=getattr(conn, "host", None),
1399 server_port=getattr(conn, "port", None),
1400 db_namespace=str(conn.db),
1401 )
1403 return response
1404 except Exception as e:
1405 await record_error_count(
1406 server_address=getattr(conn, "host", None),
1407 server_port=getattr(conn, "port", None),
1408 network_peer_address=getattr(conn, "host", None),
1409 network_peer_port=getattr(conn, "port", None),
1410 error_type=e,
1411 retry_attempts=actual_retry_attempts,
1412 is_internal=False,
1413 )
1414 raise
1416 async def parse_response(self, block: bool = True, timeout: float = 0):
1417 """
1418 Parse the response from a publish/subscribe command.
1420 Args:
1421 block: If True, block indefinitely until a message is available.
1422 If False, return immediately if no message is available.
1423 Default: True
1424 timeout: The timeout in seconds for reading a response when block=False.
1425 This parameter is ignored when block=True.
1426 Default: 0 (return immediately if no data available)
1428 Returns:
1429 The parsed response from the server, or None if no message is available
1430 within the timeout period (when block=False).
1432 Important:
1433 The block and timeout parameters work together:
1434 - When block=True: timeout is IGNORED, method blocks indefinitely
1435 - When block=False: timeout is USED, method returns after timeout expires
1437 Typically, you should use get_message(timeout=X) instead of calling
1438 parse_response() directly. The get_message() method automatically sets
1439 block=False when a timeout is provided, and block=True when timeout=None.
1441 Example:
1442 # Block indefinitely (timeout is ignored)
1443 response = await pubsub.parse_response(block=True, timeout=0.1)
1445 # Non-blocking with 0.1 second timeout
1446 response = await pubsub.parse_response(block=False, timeout=0.1)
1448 # Non-blocking, return immediately
1449 response = await pubsub.parse_response(block=False, timeout=0)
1451 # Recommended: use get_message() instead
1452 msg = await pubsub.get_message(timeout=0.1) # automatically sets block=False
1453 msg = await pubsub.get_message(timeout=None) # automatically sets block=True
1454 """
1455 conn = self.connection
1456 if conn is None:
1457 raise RuntimeError(
1458 "pubsub connection not set: "
1459 "did you forget to call subscribe() or psubscribe()?"
1460 )
1462 await self.check_health()
1464 if not conn.is_connected:
1465 await conn.connect()
1467 # Block=True: signal "no timeout" to conn.read_response via
1468 # math.inf. The connection treats math.inf as the per-read
1469 # opt-in for blocking indefinitely without falling back to
1470 # self.socket_timeout. Reconnect/AUTH/HELLO/resubscribe
1471 # operations performed by the retry layer continue to honor
1472 # self.socket_timeout because they do not pass math.inf.
1473 #
1474 # TODO(next-major): when the async Connection.read_response
1475 # default for ``timeout`` is changed to SENTINEL, passing
1476 # ``timeout=None`` from this method will become the natural
1477 # "no timeout" signal and the math.inf hand-off can be
1478 # removed. That swap is a breaking change to the
1479 # Connection.read_response signature so it must wait for a
1480 # major release.
1481 read_timeout = math.inf if block else timeout
1482 response = await self._execute(
1483 conn,
1484 conn.read_response,
1485 timeout=read_timeout,
1486 disconnect_on_error=False,
1487 push_request=True,
1488 )
1490 if conn.health_check_interval and response in self.health_check_response:
1491 # ignore the health check message as user might not expect it
1492 return None
1493 return response
1495 async def check_health(self):
1496 conn = self.connection
1497 if conn is None:
1498 raise RuntimeError(
1499 "pubsub connection not set: "
1500 "did you forget to call subscribe() or psubscribe()?"
1501 )
1503 if (
1504 conn.health_check_interval
1505 and asyncio.get_running_loop().time() > conn.next_health_check
1506 ):
1507 await conn.send_command(
1508 "PING", self.HEALTH_CHECK_MESSAGE, check_health=False
1509 )
1511 def _normalize_keys(self, data: _NormalizeKeysT) -> _NormalizeKeysT:
1512 """
1513 normalize channel/pattern names to be either bytes or strings
1514 based on whether responses are automatically decoded. this saves us
1515 from coercing the value for each message coming in.
1516 """
1517 encode = self.encoder.encode
1518 decode = self.encoder.decode
1519 return {decode(encode(k)): v for k, v in data.items()} # type: ignore[return-value] # noqa: E501
1521 async def psubscribe(
1522 self, *args: ChannelT | Subscription, **kwargs: PubSubHandler
1523 ) -> None:
1524 """
1525 Subscribe to channel patterns.
1526 Patterns supplied as keyword arguments expect a pattern name as the
1527 key and a callable as the value.
1528 ``Subscription`` objects can also be supplied positionally with an
1529 optional handler.
1530 A pattern's callable will be invoked automatically
1531 when a message is received on that pattern rather than producing a
1532 message via ``listen()``.
1533 """
1534 new_patterns = parse_pubsub_subscriptions(args, kwargs)
1535 ret_val = await self.execute_command("PSUBSCRIBE", *new_patterns.keys())
1536 # update the patterns dict AFTER we send the command. we don't want to
1537 # subscribe twice to these patterns, once for the command and again
1538 # for the reconnection.
1539 new_patterns = self._normalize_keys(new_patterns)
1540 self.patterns.update(new_patterns)
1541 self.pending_unsubscribe_patterns.difference_update(new_patterns)
1542 return ret_val
1544 def punsubscribe(self, *args: ChannelT) -> Awaitable:
1545 """
1546 Unsubscribe from the supplied patterns. If empty, unsubscribe from
1547 all patterns.
1548 """
1549 patterns: Iterable[ChannelT]
1550 if args:
1551 parsed_args = list_or_args((args[0],), args[1:])
1552 patterns = self._normalize_keys(dict.fromkeys(parsed_args)).keys()
1553 else:
1554 parsed_args = []
1555 patterns = self.patterns
1556 self.pending_unsubscribe_patterns.update(patterns)
1557 return self.execute_command("PUNSUBSCRIBE", *parsed_args)
1559 async def subscribe(
1560 self, *args: ChannelT | Subscription, **kwargs: PubSubHandler
1561 ) -> None:
1562 """
1563 Subscribe to channels.
1564 Channels supplied as keyword arguments expect
1565 a channel name as the key and a callable as the value.
1566 ``Subscription`` objects can also be supplied positionally with an
1567 optional handler.
1568 A channel's callable will be invoked automatically
1569 when a message is received on that channel rather than producing a
1570 message via ``listen()`` or ``get_message()``.
1571 """
1572 new_channels = parse_pubsub_subscriptions(args, kwargs)
1573 ret_val = await self.execute_command("SUBSCRIBE", *new_channels.keys())
1574 # update the channels dict AFTER we send the command. we don't want to
1575 # subscribe twice to these channels, once for the command and again
1576 # for the reconnection.
1577 new_channels = self._normalize_keys(new_channels)
1578 self.channels.update(new_channels)
1579 self.pending_unsubscribe_channels.difference_update(new_channels)
1580 return ret_val
1582 def unsubscribe(self, *args) -> Awaitable:
1583 """
1584 Unsubscribe from the supplied channels. If empty, unsubscribe from
1585 all channels
1586 """
1587 if args:
1588 parsed_args = list_or_args(args[0], args[1:])
1589 channels = self._normalize_keys(dict.fromkeys(parsed_args))
1590 else:
1591 parsed_args = []
1592 channels = self.channels
1593 self.pending_unsubscribe_channels.update(channels)
1594 return self.execute_command("UNSUBSCRIBE", *parsed_args)
1596 async def ssubscribe(
1597 self,
1598 *args: ChannelT | Subscription,
1599 target_node: Any = None,
1600 **kwargs: PubSubHandler,
1601 ) -> None:
1602 """
1603 Subscribes the client to the specified shard channels.
1604 Channels supplied as keyword arguments expect a channel name as the key
1605 and a callable as the value.
1606 ``Subscription`` objects can also be supplied positionally
1607 with an optional handler.
1608 A channel's callable will be invoked automatically when a message
1609 is received on that channel rather than producing a message
1610 via ``listen()`` or ``get_sharded_message()``.
1611 """
1612 new_s_channels = parse_pubsub_subscriptions(args, kwargs)
1613 ret_val = await self.execute_command("SSUBSCRIBE", *new_s_channels.keys())
1614 # update the s_channels dict AFTER we send the command. we don't want to
1615 # subscribe twice to these channels, once for the command and again
1616 # for the reconnection.
1617 new_s_channels = self._normalize_keys(new_s_channels)
1618 self.shard_channels.update(new_s_channels)
1619 self.pending_unsubscribe_shard_channels.difference_update(new_s_channels)
1620 return ret_val
1622 def sunsubscribe(self, *args, target_node=None) -> Awaitable:
1623 """
1624 Unsubscribe from the supplied shard_channels. If empty, unsubscribe from
1625 all shard_channels
1626 """
1627 if args:
1628 args = list_or_args(args[0], args[1:])
1629 s_channels = self._normalize_keys(dict.fromkeys(args))
1630 else:
1631 s_channels = self.shard_channels
1632 self.pending_unsubscribe_shard_channels.update(s_channels)
1633 return self.execute_command("SUNSUBSCRIBE", *args)
1635 async def listen(self) -> AsyncIterator:
1636 """Listen for messages on channels this client has been subscribed to.
1638 Iteration ends once every channel and pattern has been unsubscribed
1639 from. If nothing is subscribed when iteration begins it ends
1640 immediately rather than waiting, so subscribe first: a listener
1641 started before any subscription finishes without yielding anything.
1642 """
1643 while self.subscribed:
1644 response = await self.handle_message(await self.parse_response(block=True))
1645 if response is not None:
1646 yield response
1648 async def get_message(
1649 self, ignore_subscribe_messages: bool = False, timeout: Optional[float] = 0.0
1650 ):
1651 """
1652 Get the next message if one is available, otherwise None.
1654 If timeout is specified, the system will wait for `timeout` seconds
1655 before returning. Timeout should be specified as a floating point
1656 number or None to wait indefinitely.
1657 """
1658 response = await self.parse_response(block=(timeout is None), timeout=timeout)
1659 if response:
1660 return await self.handle_message(response, ignore_subscribe_messages)
1661 return None
1663 get_sharded_message = get_message
1665 def ping(self, message=None) -> Awaitable[bool]:
1666 """
1667 Ping the Redis server to test connectivity.
1669 Sends a PING command to the Redis server and returns True if the server
1670 responds with "PONG".
1671 """
1672 args = ["PING", message] if message is not None else ["PING"]
1673 return self.execute_command(*args)
1675 async def handle_message(self, response, ignore_subscribe_messages=False):
1676 """
1677 Parses a pub/sub message. If the channel or pattern was subscribed to
1678 with a message handler, the handler is invoked instead of a parsed
1679 message being returned.
1680 """
1681 if response is None:
1682 return None
1683 if isinstance(response, bytes):
1684 response = [b"pong", response] if response != b"PONG" else [b"pong", b""]
1685 message_type = str_if_bytes(response[0])
1686 if message_type == "pmessage":
1687 message = {
1688 "type": message_type,
1689 "pattern": response[1],
1690 "channel": response[2],
1691 "data": response[3],
1692 }
1693 elif message_type == "pong":
1694 message = {
1695 "type": message_type,
1696 "pattern": None,
1697 "channel": None,
1698 "data": response[1],
1699 }
1700 else:
1701 message = {
1702 "type": message_type,
1703 "pattern": None,
1704 "channel": response[1],
1705 "data": response[2],
1706 }
1708 if message_type in ["message", "pmessage"]:
1709 channel = str_if_bytes(message["channel"])
1710 await record_pubsub_message(
1711 direction=PubSubDirection.RECEIVE,
1712 channel=channel,
1713 )
1714 elif message_type == "smessage":
1715 channel = str_if_bytes(message["channel"])
1716 await record_pubsub_message(
1717 direction=PubSubDirection.RECEIVE,
1718 channel=channel,
1719 sharded=True,
1720 )
1722 # if this is an unsubscribe message, remove it from memory.
1723 # ``discard`` rather than ``remove``: the guard above already makes the
1724 # removal conditional, so the two are equivalent for a single caller -
1725 # but another writer can drop the same entry between the check and the
1726 # removal, and ``remove`` would then raise ``KeyError`` out of a pubsub
1727 # read that no caller catches. ``ClusterPubSub._detach_shard_channel``
1728 # is such a writer: it forgets a migrating shard channel locally,
1729 # deliberately without the per-node I/O lock this bookkeeping runs
1730 # under, because waiting for that lock stalls reconciliation behind a
1731 # poll's whole retry budget on the node being migrated away from.
1732 if message_type in self.UNSUBSCRIBE_MESSAGE_TYPES:
1733 if message_type == "punsubscribe":
1734 pattern = response[1]
1735 if pattern in self.pending_unsubscribe_patterns:
1736 self.pending_unsubscribe_patterns.discard(pattern)
1737 self.patterns.pop(pattern, None)
1738 elif message_type == "sunsubscribe":
1739 s_channel = response[1]
1740 if s_channel in self.pending_unsubscribe_shard_channels:
1741 self.pending_unsubscribe_shard_channels.discard(s_channel)
1742 self.shard_channels.pop(s_channel, None)
1743 else:
1744 channel = response[1]
1745 if channel in self.pending_unsubscribe_channels:
1746 self.pending_unsubscribe_channels.discard(channel)
1747 self.channels.pop(channel, None)
1749 if message_type in self.PUBLISH_MESSAGE_TYPES:
1750 # if there's a message handler, invoke it
1751 if message_type == "pmessage":
1752 handler = self.patterns.get(message["pattern"], None)
1753 elif message_type == "smessage":
1754 handler = self.shard_channels.get(message["channel"], None)
1755 else:
1756 handler = self.channels.get(message["channel"], None)
1757 if handler:
1758 if inspect.iscoroutinefunction(handler):
1759 await handler(message)
1760 else:
1761 handler(message)
1762 return None
1763 elif message_type != "pong":
1764 # this is a subscribe/unsubscribe message. ignore if we don't
1765 # want them
1766 if ignore_subscribe_messages or self.ignore_subscribe_messages:
1767 return None
1769 return message
1771 async def run(
1772 self,
1773 *,
1774 exception_handler: Optional["PSWorkerThreadExcHandlerT"] = None,
1775 poll_timeout: float = 1.0,
1776 pubsub=None,
1777 ) -> None:
1778 """Process pub/sub messages using registered callbacks.
1780 This is the equivalent of :py:meth:`redis.PubSub.run_in_thread` in
1781 redis-py, but it is a coroutine. To launch it as a separate task, use
1782 ``asyncio.create_task``:
1784 >>> task = asyncio.create_task(pubsub.run())
1786 To shut it down, use asyncio cancellation:
1788 >>> task.cancel()
1789 >>> await task
1790 """
1791 for channel, handler in self.channels.items():
1792 if handler is None:
1793 raise PubSubError(f"Channel: '{channel}' has no handler registered")
1794 for pattern, handler in self.patterns.items():
1795 if handler is None:
1796 raise PubSubError(f"Pattern: '{pattern}' has no handler registered")
1798 await self.connect()
1799 while True:
1800 try:
1801 if pubsub is None:
1802 await self.get_message(
1803 ignore_subscribe_messages=True, timeout=poll_timeout
1804 )
1805 else:
1806 await pubsub.get_message(
1807 ignore_subscribe_messages=True, timeout=poll_timeout
1808 )
1809 except asyncio.CancelledError:
1810 raise
1811 except BaseException as e:
1812 if exception_handler is None:
1813 raise
1814 res = exception_handler(e, self)
1815 if inspect.isawaitable(res):
1816 await res
1817 # Ensure that other tasks on the event loop get a chance to run
1818 # if we didn't have to block for I/O anywhere.
1819 await asyncio.sleep(0)
1822class PubsubWorkerExceptionHandler(Protocol):
1823 def __call__(self, e: BaseException, pubsub: PubSub): ...
1826class AsyncPubsubWorkerExceptionHandler(Protocol):
1827 async def __call__(self, e: BaseException, pubsub: PubSub): ...
1830PSWorkerThreadExcHandlerT = Union[
1831 PubsubWorkerExceptionHandler, AsyncPubsubWorkerExceptionHandler
1832]
1835CommandT = Tuple[Tuple[Union[str, bytes], ...], Mapping[str, Any]]
1836CommandStackT = List[CommandT]
1839class Pipeline(Redis): # lgtm [py/init-calls-subclass]
1840 """
1841 Pipelines provide a way to transmit multiple commands to the Redis server
1842 in one transmission. This is convenient for batch processing, such as
1843 saving all the values in a list to Redis.
1845 All commands executed within a pipeline(when running in transactional mode,
1846 which is the default behavior) are wrapped with MULTI and EXEC
1847 calls. This guarantees all commands executed in the pipeline will be
1848 executed atomically.
1850 Any command raising an exception does *not* halt the execution of
1851 subsequent commands in the pipeline. Instead, the exception is caught
1852 and its instance is placed into the response list returned by execute().
1853 Code iterating over the response list should be able to deal with an
1854 instance of an exception as a potential value. In general, these will be
1855 ResponseError exceptions, such as those raised when issuing a command
1856 on a key of a different datatype.
1857 """
1859 UNWATCH_COMMANDS = {"DISCARD", "EXEC", "UNWATCH"}
1861 def __init__(
1862 self,
1863 connection_pool: ConnectionPool,
1864 response_callbacks: MutableMapping[Union[str, bytes], ResponseCallbackT],
1865 transaction: bool,
1866 shard_hint: Optional[str],
1867 ):
1868 self.connection_pool = connection_pool
1869 self.connection = None
1870 self.response_callbacks = response_callbacks
1871 self.is_transaction = transaction
1872 self.shard_hint = shard_hint
1873 self.watching = False
1874 self.command_stack: CommandStackT = []
1875 self.scripts: Set[Script] = set()
1876 self.explicit_transaction = False
1878 async def __aenter__(self: _RedisT) -> _RedisT:
1879 return self
1881 async def __aexit__(self, exc_type, exc_value, traceback):
1882 await self.reset()
1884 def __await__(self):
1885 return self._async_self().__await__()
1887 _DEL_MESSAGE = "Unclosed Pipeline client"
1889 def __len__(self):
1890 return len(self.command_stack)
1892 def __bool__(self):
1893 """Pipeline instances should always evaluate to True"""
1894 return True
1896 async def _async_self(self):
1897 return self
1899 async def reset(self):
1900 self.command_stack = []
1901 self.scripts = set()
1902 try:
1903 # make sure to reset the connection state in the event that we were
1904 # watching something
1905 if self.watching and self.connection:
1906 try:
1907 # call this manually since our unwatch or
1908 # immediate_execute_command methods can call reset()
1909 await self.connection.send_command("UNWATCH")
1910 await self.connection.read_response()
1911 except ConnectionError:
1912 # disconnect will also remove any previous WATCHes
1913 if self.connection:
1914 await self.connection.disconnect()
1915 except asyncio.CancelledError:
1916 # Disconnect so any unread UNWATCH reply does not get
1917 # served to the next caller that takes the connection.
1918 if self.connection:
1919 await self.connection.disconnect()
1920 raise
1921 finally:
1922 self.watching = False
1923 self.explicit_transaction = False
1924 # We can safely return the connection to the pool here since we're
1925 # sure we're no longer WATCHing anything. Detach self.connection
1926 # before awaiting release: if a second cancel aborts the await,
1927 # the pipeline must not be left holding a reference to a
1928 # connection that is being returned to the pool. Shield the
1929 # release itself so a second cancel cannot split the pool's
1930 # internal in-use/available bookkeeping mid-update.
1931 if self.connection:
1932 connection, self.connection = self.connection, None
1933 await asyncio.shield(self.connection_pool.release(connection))
1935 async def aclose(self) -> None:
1936 """Alias for reset(), a standard method name for cleanup"""
1937 await self.reset()
1939 def multi(self):
1940 """
1941 Start a transactional block of the pipeline after WATCH commands
1942 are issued. End the transactional block with `execute`.
1943 """
1944 if self.explicit_transaction:
1945 raise RedisError("Cannot issue nested calls to MULTI")
1946 if self.command_stack:
1947 raise RedisError(
1948 "Commands without an initial WATCH have already been issued"
1949 )
1950 self.explicit_transaction = True
1952 def execute_command(
1953 self, *args, **kwargs
1954 ) -> Union["Pipeline", Awaitable["Pipeline"]]:
1955 if (self.watching or args[0] == "WATCH") and not self.explicit_transaction:
1956 return self.immediate_execute_command(*args, **kwargs)
1957 return self.pipeline_execute_command(*args, **kwargs)
1959 async def _disconnect_reset_raise_on_watching(
1960 self,
1961 conn: Connection,
1962 error: Exception,
1963 failure_count: Optional[int] = None,
1964 start_time: Optional[float] = None,
1965 command_name: Optional[str] = None,
1966 ) -> None:
1967 """
1968 Close the connection reset watching state and
1969 raise an exception if we were watching.
1971 The supported exceptions are already checked in the
1972 retry object so we don't need to do it here.
1974 After we disconnect the connection, it will try to reconnect and
1975 do a health check as part of the send_command logic(on connection level).
1976 """
1977 if (
1978 error
1979 and failure_count is not None
1980 and failure_count <= conn.retry.get_retries()
1981 ):
1982 await record_operation_duration(
1983 command_name=command_name,
1984 duration_seconds=time.monotonic() - start_time,
1985 server_address=getattr(conn, "host", None),
1986 server_port=getattr(conn, "port", None),
1987 db_namespace=str(conn.db),
1988 error=error,
1989 retry_attempts=failure_count,
1990 )
1991 await conn.disconnect(error=error, failure_count=failure_count)
1992 # if we were already watching a variable, the watch is no longer
1993 # valid since this connection has died. raise a WatchError, which
1994 # indicates the user should retry this transaction.
1995 if self.watching:
1996 await self.reset()
1997 raise WatchError(
1998 f"A {type(error).__name__} occurred while watching one or more keys"
1999 )
2001 async def immediate_execute_command(self, *args, **options):
2002 """
2003 Execute a command immediately, but don't auto-retry on the supported
2004 errors for retry if we're already WATCHing a variable.
2005 Used when issuing WATCH or subsequent commands retrieving their values but before
2006 MULTI is called.
2007 """
2008 command_name = args[0]
2009 conn = self.connection
2010 # if this is the first call, we need a connection
2011 if not conn:
2012 conn = await self.connection_pool.get_connection()
2013 self.connection = conn
2015 # Start timing for observability
2016 start_time = time.monotonic()
2017 # Track actual retry attempts for error reporting
2018 actual_retry_attempts = 0
2020 def failure_callback(error, failure_count):
2021 if is_debug_log_enabled():
2022 add_debug_log_for_operation_failure(conn, error, args)
2023 nonlocal actual_retry_attempts
2024 actual_retry_attempts = failure_count
2025 return self._disconnect_reset_raise_on_watching(
2026 conn, error, failure_count, start_time, command_name
2027 )
2029 try:
2030 response = await conn.retry.call_with_retry(
2031 lambda: self._send_command_parse_response(
2032 conn, command_name, *args, **options
2033 ),
2034 failure_callback,
2035 with_failure_count=True,
2036 )
2038 await record_operation_duration(
2039 command_name=command_name,
2040 duration_seconds=time.monotonic() - start_time,
2041 server_address=getattr(conn, "host", None),
2042 server_port=getattr(conn, "port", None),
2043 db_namespace=str(conn.db),
2044 )
2046 return response
2047 except Exception as e:
2048 await record_error_count(
2049 server_address=getattr(conn, "host", None),
2050 server_port=getattr(conn, "port", None),
2051 network_peer_address=getattr(conn, "host", None),
2052 network_peer_port=getattr(conn, "port", None),
2053 error_type=e,
2054 retry_attempts=actual_retry_attempts,
2055 is_internal=False,
2056 )
2057 raise
2059 def pipeline_execute_command(self, *args, **options):
2060 """
2061 Stage a command to be executed when execute() is next called
2063 Returns the current Pipeline object back so commands can be
2064 chained together, such as:
2066 pipe = pipe.set('foo', 'bar').incr('baz').decr('bang')
2068 At some other point, you can then run: pipe.execute(),
2069 which will execute all commands queued in the pipe.
2070 """
2071 self.command_stack.append((args, options))
2072 return self
2074 async def _himport_prepare_pipeline(self, conn, commands):
2075 """Delegate to the shared async HIMPORT executor."""
2076 await _himport_exec.prepare_pipeline(self, conn, [args for args, _ in commands])
2078 async def _execute_transaction( # noqa: C901
2079 self, connection: Connection, commands: CommandStackT, raise_on_error
2080 ):
2081 # Ensure fieldsets referenced by buffered HIMPORT SETs are prepared on this
2082 # connection before the MULTI/EXEC block (session state, not transactional).
2083 await self._himport_prepare_pipeline(connection, commands)
2084 pre: CommandT = (("MULTI",), {})
2085 post: CommandT = (("EXEC",), {})
2086 cmds = (pre, *commands, post)
2087 all_cmds = connection.pack_commands(
2088 args for args, options in cmds if EMPTY_RESPONSE not in options
2089 )
2090 await connection.send_packed_command(all_cmds)
2091 errors = []
2093 # parse off the response for MULTI
2094 # NOTE: we need to handle ResponseErrors here and continue
2095 # so that we read all the additional command messages from
2096 # the socket
2097 try:
2098 await self.parse_response(connection, "_")
2099 except ResponseError as err:
2100 errors.append((0, err))
2102 # and all the other commands
2103 for i, command in enumerate(commands):
2104 if EMPTY_RESPONSE in command[1]:
2105 errors.append((i, command[1][EMPTY_RESPONSE]))
2106 else:
2107 try:
2108 await self.parse_response(connection, "_")
2109 except ResponseError as err:
2110 self.annotate_exception(err, i + 1, command[0])
2111 errors.append((i, err))
2113 # parse the EXEC.
2114 try:
2115 response = await self.parse_response(connection, "_")
2116 except ExecAbortError as err:
2117 if errors:
2118 raise errors[0][1] from err
2119 raise
2121 # EXEC clears any watched keys
2122 self.watching = False
2124 if response is None:
2125 raise WatchError("Watched variable changed.") from None
2127 # put any parse errors into the response
2128 for i, e in errors:
2129 response.insert(i, e)
2131 if len(response) != len(commands):
2132 if self.connection:
2133 await self.connection.disconnect()
2134 raise ResponseError(
2135 "Wrong number of response items from pipeline execution"
2136 ) from None
2138 # find any errors in the response and raise if necessary
2139 if raise_on_error:
2140 self.raise_first_error(commands, response)
2142 # We have to run response callbacks manually
2143 data = []
2144 for r, cmd in zip(response, commands):
2145 if not isinstance(r, Exception):
2146 args, options = cmd
2147 command_name = args[0]
2149 # Remove keys entry, it needs only for cache.
2150 options.pop("keys", None)
2152 if command_name in self.response_callbacks:
2153 r = self.response_callbacks[command_name](r, **options)
2154 if inspect.isawaitable(r):
2155 r = await r
2156 data.append(r)
2157 return data
2159 async def _execute_pipeline(
2160 self, connection: Connection, commands: CommandStackT, raise_on_error: bool
2161 ):
2162 # Fold any first-use HIMPORT PREPAREs for referenced fieldsets into the same
2163 # packed write as the queued commands, so a pipeline that lands on a fresh or
2164 # reconnected connection stays a single round trip (the batched write bypasses
2165 # the per-command lazy PREPARE path). Deferred-discard reconciliation happens
2166 # inside pipeline_prepares and only touches the socket when discards are
2167 # actually pending.
2168 fieldsets = await _himport_exec.pipeline_prepares(
2169 self, connection, [args for args, _ in commands]
2170 )
2171 preflight = _himport_exec.prepare_wire_commands(fieldsets)
2172 # build up all commands into a single request to increase network perf
2173 all_cmds = connection.pack_commands(preflight + [args for args, _ in commands])
2174 await connection.send_packed_command(all_cmds)
2176 # Drain the leading PREPARE replies (bookkeeping + capture the first error)
2177 # before the queued replies. Everything on the wire is read before raising so
2178 # the pooled socket never desyncs.
2179 prep_error = await _himport_exec.drain_pipeline_prepares(
2180 self, connection, fieldsets
2181 )
2183 response = []
2184 for args, options in commands:
2185 try:
2186 response.append(
2187 await self.parse_response(connection, args[0], **options)
2188 )
2189 except ResponseError as e:
2190 response.append(e)
2192 # A PREPARE failure (rare: an invalid fieldset definition) is a hard error,
2193 # raised regardless of raise_on_error as it was before folding -- only now
2194 # every reply has already been drained.
2195 if prep_error is not None:
2196 raise prep_error
2197 if raise_on_error:
2198 self.raise_first_error(commands, response)
2199 return response
2201 def raise_first_error(self, commands: CommandStackT, response: Iterable[Any]):
2202 for i, r in enumerate(response):
2203 if isinstance(r, ResponseError):
2204 self.annotate_exception(r, i + 1, commands[i][0])
2205 raise r
2207 def annotate_exception(
2208 self, exception: Exception, number: int, command: Iterable[object]
2209 ) -> None:
2210 cmd = " ".join(map(safe_str, command))
2211 msg = (
2212 f"Command # {number} ({truncate_text(cmd)}) "
2213 f"of pipeline caused error: {exception.args}"
2214 )
2215 exception.args = (msg,) + exception.args[1:]
2217 async def parse_response(
2218 self, connection: Connection, command_name: Union[str, bytes], **options
2219 ):
2220 result = await super().parse_response(connection, command_name, **options)
2221 if command_name in self.UNWATCH_COMMANDS:
2222 self.watching = False
2223 elif command_name == "WATCH":
2224 self.watching = True
2225 return result
2227 async def load_scripts(self):
2228 # make sure all scripts that are about to be run on this pipeline exist
2229 scripts = list(self.scripts)
2230 immediate = self.immediate_execute_command
2231 shas = [s.sha for s in scripts]
2232 # we can't use the normal script_* methods because they would just
2233 # get buffered in the pipeline.
2234 exists = await immediate("SCRIPT EXISTS", *shas)
2235 if not all(exists):
2236 for s, exist in zip(scripts, exists):
2237 if not exist:
2238 s.sha = await immediate("SCRIPT LOAD", s.script)
2240 async def _disconnect_raise_on_watching(
2241 self,
2242 conn: Connection,
2243 error: Exception,
2244 failure_count: Optional[int] = None,
2245 start_time: Optional[float] = None,
2246 command_name: Optional[str] = None,
2247 ):
2248 """
2249 Close the connection, raise an exception if we were watching.
2251 The supported exceptions are already checked in the
2252 retry object so we don't need to do it here.
2254 After we disconnect the connection, it will try to reconnect and
2255 do a health check as part of the send_command logic(on connection level).
2256 """
2257 if (
2258 error
2259 and failure_count is not None
2260 and failure_count <= conn.retry.get_retries()
2261 ):
2262 await record_operation_duration(
2263 command_name=command_name,
2264 duration_seconds=time.monotonic() - start_time,
2265 server_address=getattr(conn, "host", None),
2266 server_port=getattr(conn, "port", None),
2267 db_namespace=str(conn.db),
2268 error=error,
2269 retry_attempts=failure_count,
2270 )
2271 await conn.disconnect(error=error, failure_count=failure_count)
2272 # if we were watching a variable, the watch is no longer valid
2273 # since this connection has died. raise a WatchError, which
2274 # indicates the user should retry this transaction.
2275 if self.watching:
2276 raise WatchError(
2277 f"A {type(error).__name__} occurred while watching one or more keys"
2278 )
2280 async def execute(self, raise_on_error: bool = True) -> List[Any]:
2281 """Execute all the commands in the current pipeline"""
2282 stack = self.command_stack
2283 if not stack and not self.watching:
2284 return []
2285 if self.scripts:
2286 await self.load_scripts()
2287 if self.is_transaction or self.explicit_transaction:
2288 execute = self._execute_transaction
2289 operation_name = "MULTI"
2290 else:
2291 execute = self._execute_pipeline
2292 operation_name = "PIPELINE"
2294 conn = self.connection
2295 if not conn:
2296 conn = await self.connection_pool.get_connection()
2297 # assign to self.connection so reset() releases the connection
2298 # back to the pool after we're done
2299 self.connection = conn
2300 conn = cast(Connection, conn)
2302 # Start timing for observability
2303 start_time = time.monotonic()
2304 # Track actual retry attempts for error reporting
2305 actual_retry_attempts = 0
2307 def failure_callback(error, failure_count):
2308 if is_debug_log_enabled():
2309 add_debug_log_for_operation_failure(conn, error, (operation_name,))
2310 nonlocal actual_retry_attempts
2311 actual_retry_attempts = failure_count
2312 return self._disconnect_raise_on_watching(
2313 conn, error, failure_count, start_time, operation_name
2314 )
2316 try:
2317 response = await conn.retry.call_with_retry(
2318 lambda: execute(conn, stack, raise_on_error),
2319 failure_callback,
2320 with_failure_count=True,
2321 )
2323 await record_operation_duration(
2324 command_name=operation_name,
2325 duration_seconds=time.monotonic() - start_time,
2326 server_address=getattr(conn, "host", None),
2327 server_port=getattr(conn, "port", None),
2328 db_namespace=str(conn.db),
2329 )
2330 return response
2331 except Exception as e:
2332 await record_error_count(
2333 server_address=getattr(conn, "host", None),
2334 server_port=getattr(conn, "port", None),
2335 network_peer_address=getattr(conn, "host", None),
2336 network_peer_port=getattr(conn, "port", None),
2337 error_type=e,
2338 retry_attempts=actual_retry_attempts,
2339 is_internal=False,
2340 )
2341 raise
2342 finally:
2343 await self.reset()
2345 async def discard(self):
2346 """Flushes all previously queued commands
2347 See: https://redis.io/commands/DISCARD
2348 """
2349 await self.execute_command("DISCARD")
2351 async def watch(self, *names: KeyT):
2352 """Watches the values at keys ``names``"""
2353 if self.explicit_transaction:
2354 raise RedisError("Cannot issue a WATCH after a MULTI")
2355 return await self.execute_command("WATCH", *names)
2357 async def unwatch(self):
2358 """Unwatches all previously specified keys"""
2359 return self.watching and await self.execute_command("UNWATCH") or True