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
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 self.get_connection_kwargs().update({"retry": retry})
593 self.connection_pool.set_retry(retry)
595 def load_external_module(self, funcname, func):
596 """
597 This function can be used to add externally defined redis modules,
598 and their namespaces to the redis client.
600 funcname - A string containing the name of the function to create
601 func - The function, being added to this class.
603 ex: Assume that one has a custom redis module named foomod that
604 creates command named 'foo.dothing' and 'foo.anotherthing' in redis.
605 To load function functions into this namespace:
607 from redis import Redis
608 from foomodule import F
609 r = Redis()
610 r.load_external_module("foo", F)
611 r.foo().dothing('your', 'arguments')
613 For a concrete example see the reimport of the redisjson module in
614 tests/test_connection.py::test_loading_external_modules
615 """
616 setattr(self, funcname, func)
618 def pipeline(
619 self, transaction: bool = True, shard_hint: Optional[str] = None
620 ) -> "Pipeline":
621 """
622 Return a new pipeline object that can queue multiple commands for
623 later execution. ``transaction`` indicates whether all commands
624 should be executed atomically. Apart from making a group of operations
625 atomic, pipelines are useful for reducing the back-and-forth overhead
626 between the client and server.
627 """
628 return Pipeline(
629 self.connection_pool, self.response_callbacks, transaction, shard_hint
630 )
632 async def transaction(
633 self,
634 func: Callable[["Pipeline"], Union[Any, Awaitable[Any]]],
635 *watches: KeyT,
636 shard_hint: Optional[str] = None,
637 value_from_callable: bool = False,
638 watch_delay: Optional[float] = None,
639 ):
640 """
641 Convenience method for executing the callable `func` as a transaction
642 while watching all keys specified in `watches`. The 'func' callable
643 should expect a single argument which is a Pipeline object.
644 """
645 pipe: Pipeline
646 async with self.pipeline(True, shard_hint) as pipe:
647 while True:
648 try:
649 if watches:
650 await pipe.watch(*watches)
651 func_value = func(pipe)
652 if inspect.isawaitable(func_value):
653 func_value = await func_value
654 exec_value = await pipe.execute()
655 return func_value if value_from_callable else exec_value
656 except WatchError:
657 if watch_delay is not None and watch_delay > 0:
658 await asyncio.sleep(watch_delay)
659 continue
661 def lock(
662 self,
663 name: KeyT,
664 timeout: Optional[float] = None,
665 sleep: float = 0.1,
666 blocking: bool = True,
667 blocking_timeout: Optional[float] = None,
668 lock_class: Optional[Type[Lock]] = None,
669 thread_local: bool = True,
670 raise_on_release_error: bool = True,
671 ) -> Lock:
672 """
673 Return a new Lock object using key ``name`` that mimics
674 the behavior of threading.Lock.
676 If specified, ``timeout`` indicates a maximum life for the lock.
677 By default, it will remain locked until release() is called.
679 ``sleep`` indicates the amount of time to sleep per loop iteration
680 when the lock is in blocking mode and another client is currently
681 holding the lock.
683 ``blocking`` indicates whether calling ``acquire`` should block until
684 the lock has been acquired or to fail immediately, causing ``acquire``
685 to return False and the lock not being acquired. Defaults to True.
686 Note this value can be overridden by passing a ``blocking``
687 argument to ``acquire``.
689 ``blocking_timeout`` indicates the maximum amount of time in seconds to
690 spend trying to acquire the lock. A value of ``None`` indicates
691 continue trying forever. ``blocking_timeout`` can be specified as a
692 float or integer, both representing the number of seconds to wait.
694 ``lock_class`` forces the specified lock implementation. Note that as
695 of redis-py 3.0, the only lock class we implement is ``Lock`` (which is
696 a Lua-based lock). So, it's unlikely you'll need this parameter, unless
697 you have created your own custom lock class.
699 ``thread_local`` indicates whether the lock token is placed in
700 thread-local storage. By default, the token is placed in thread local
701 storage so that a thread only sees its token, not a token set by
702 another thread. Consider the following timeline:
704 time: 0, thread-1 acquires `my-lock`, with a timeout of 5 seconds.
705 thread-1 sets the token to "abc"
706 time: 1, thread-2 blocks trying to acquire `my-lock` using the
707 Lock instance.
708 time: 5, thread-1 has not yet completed. redis expires the lock
709 key.
710 time: 5, thread-2 acquired `my-lock` now that it's available.
711 thread-2 sets the token to "xyz"
712 time: 6, thread-1 finishes its work and calls release(). if the
713 token is *not* stored in thread local storage, then
714 thread-1 would see the token value as "xyz" and would be
715 able to successfully release the thread-2's lock.
717 ``raise_on_release_error`` indicates whether to raise an exception when
718 the lock is no longer owned when exiting the context manager. By default,
719 this is True, meaning an exception will be raised. If False, the warning
720 will be logged and the exception will be suppressed.
722 In some use cases it's necessary to disable thread local storage. For
723 example, if you have code where one thread acquires a lock and passes
724 that lock instance to a worker thread to release later. If thread
725 local storage isn't disabled in this case, the worker thread won't see
726 the token set by the thread that acquired the lock. Our assumption
727 is that these cases aren't common and as such default to using
728 thread local storage."""
729 if lock_class is None:
730 lock_class = Lock
731 return lock_class(
732 self,
733 name,
734 timeout=timeout,
735 sleep=sleep,
736 blocking=blocking,
737 blocking_timeout=blocking_timeout,
738 thread_local=thread_local,
739 raise_on_release_error=raise_on_release_error,
740 )
742 def pubsub(self, **kwargs) -> "PubSub":
743 """
744 Return a Publish/Subscribe object. With this object, you can
745 subscribe to channels and listen for messages that get published to
746 them.
747 """
748 return PubSub(
749 self.connection_pool, event_dispatcher=self._event_dispatcher, **kwargs
750 )
752 def keyspace_notifications(
753 self,
754 key_prefix: Union[str, bytes, None] = None,
755 ignore_subscribe_messages: bool = True,
756 ) -> "AsyncKeyspaceNotifications":
757 """
758 Return an :class:`~redis.asyncio.keyspace_notifications.AsyncKeyspaceNotifications`
759 object for subscribing to keyspace and keyevent notifications.
761 Note: Keyspace notifications must be enabled on the Redis server via
762 the ``notify-keyspace-events`` configuration option.
764 Args:
765 key_prefix: Optional prefix to filter and strip from keys in
766 notifications.
767 ignore_subscribe_messages: If True, subscribe/unsubscribe
768 confirmations are not returned by
769 get_message/listen.
770 """
771 from redis.asyncio.keyspace_notifications import AsyncKeyspaceNotifications
773 return AsyncKeyspaceNotifications(
774 self,
775 key_prefix=key_prefix,
776 ignore_subscribe_messages=ignore_subscribe_messages,
777 )
779 def monitor(self) -> "Monitor":
780 return Monitor(self.connection_pool)
782 def client(self) -> "Redis":
783 return self.__class__(
784 connection_pool=self.connection_pool, single_connection_client=True
785 )
787 async def __aenter__(self: _RedisT) -> _RedisT:
788 """
789 Async context manager entry. Increments a usage counter so that the
790 connection pool is only closed (via aclose()) when no context is using
791 the client.
792 """
793 await self._increment_usage()
794 try:
795 # Initialize the client (i.e. establish connection, etc.)
796 return await self.initialize()
797 except Exception:
798 # If initialization fails, decrement the counter to keep it in sync
799 await self._decrement_usage()
800 raise
802 async def _increment_usage(self) -> int:
803 """
804 Helper coroutine to increment the usage counter while holding the lock.
805 Returns the new value of the usage counter.
806 """
807 async with self._usage_lock:
808 self._usage_counter += 1
809 return self._usage_counter
811 async def _decrement_usage(self) -> int:
812 """
813 Helper coroutine to decrement the usage counter while holding the lock.
814 Returns the new value of the usage counter.
815 """
816 async with self._usage_lock:
817 self._usage_counter -= 1
818 return self._usage_counter
820 async def __aexit__(self, exc_type, exc_value, traceback):
821 """
822 Async context manager exit. Decrements a usage counter. If this is the
823 last exit (counter becomes zero), the client closes its connection pool.
824 """
825 current_usage = await asyncio.shield(self._decrement_usage())
826 if current_usage == 0:
827 # This was the last active context, so disconnect the pool.
828 await asyncio.shield(self.aclose())
830 _DEL_MESSAGE = "Unclosed Redis client"
832 # passing _warnings and _grl as argument default since they may be gone
833 # by the time __del__ is called at shutdown
834 def __del__(
835 self,
836 _warn: Any = warnings.warn,
837 _grl: Any = asyncio.get_running_loop,
838 ) -> None:
839 if hasattr(self, "connection") and (self.connection is not None):
840 _warn(f"Unclosed client session {self!r}", ResourceWarning, source=self)
841 try:
842 context = {"client": self, "message": self._DEL_MESSAGE}
843 _grl().call_exception_handler(context)
844 except RuntimeError:
845 pass
846 self.connection._close()
848 async def aclose(self, close_connection_pool: Optional[bool] = None) -> None:
849 """
850 Closes Redis client connection
852 Args:
853 close_connection_pool:
854 decides whether to close the connection pool used by this Redis client,
855 overriding Redis.auto_close_connection_pool.
856 By default, let Redis.auto_close_connection_pool decide
857 whether to close the connection pool.
858 """
859 conn = self.connection
860 if conn:
861 self.connection = None
862 await self.connection_pool.release(conn)
863 if close_connection_pool or (
864 close_connection_pool is None and self.auto_close_connection_pool
865 ):
866 await self.connection_pool.aclose()
868 @deprecated_function(version="5.0.1", reason="Use aclose() instead", name="close")
869 async def close(self, close_connection_pool: Optional[bool] = None) -> None:
870 """
871 Alias for aclose(), for backwards compatibility
872 """
873 await self.aclose(close_connection_pool)
875 async def _send_command_parse_response(self, conn, command_name, *args, **options):
876 """
877 Send a command and parse the response
878 """
879 # HIMPORT SET is the one command whose wire form depends on per-connection
880 # state: the fieldset must be PREPAREd on this connection first, and any
881 # fieldset discarded since this connection last reconciled must be dropped.
882 # Handling it here (rather than in himport_set) lets himport_set reuse the
883 # full execute_command machinery — retry, disconnect-on-error, pooling — so
884 # a failed HIMPORT SET disconnects the connection like any other command.
885 # This per-command branch in the hot dispatch path is deliberate and has no
886 # cleaner alternative: this is the only seam where the concrete borrowed
887 # connection is known, and connection-scoped session setup can only happen
888 # once that connection is chosen. The overhead is one string compare per
889 # command.
890 himport_set = parse_himport_set_args(args)
891 if himport_set is not None:
892 # ``args`` is an HIMPORT SET in either the joined ("HIMPORT SET", key,
893 # ...) or split ("HIMPORT", "SET", key, ...) raw form; the operands come
894 # back at the right offsets for the form. A command with too few operands
895 # returns None and falls through to the normal send path so the server
896 # returns its arity error instead of a client-side IndexError here.
897 key, fieldset_name, values = himport_set
898 return await self._himport_execute_set(conn, key, fieldset_name, values)
899 await conn.send_command(*args)
900 return await self.parse_response(conn, command_name, **options)
902 async def _himport_reconcile_discards(self, conn):
903 """Delegate to the shared async HIMPORT executor."""
904 return await _himport_exec.reconcile_discards(self, conn)
906 async def _himport_prepare_and_set(
907 self, conn, key, fieldset_name, values, fieldset
908 ):
909 """Delegate to the shared async HIMPORT executor."""
910 return await _himport_exec.prepare_and_set(
911 self, conn, key, fieldset_name, values, fieldset
912 )
914 async def _himport_execute_set(self, conn, key, fieldset_name, values):
915 """Delegate to the shared async HIMPORT executor."""
916 return await _himport_exec.execute_set(self, conn, key, fieldset_name, values)
918 async def _close_connection(
919 self,
920 conn: Connection,
921 error: Optional[BaseException] = None,
922 failure_count: Optional[int] = None,
923 start_time: Optional[float] = None,
924 command_name: Optional[str] = None,
925 ):
926 """
927 Close the connection before retrying.
929 The supported exceptions are already checked in the
930 retry object so we don't need to do it here.
932 After we disconnect the connection, it will try to reconnect and
933 do a health check as part of the send_command logic(on connection level).
934 """
935 if (
936 error
937 and failure_count is not None
938 and failure_count <= conn.retry.get_retries()
939 ):
940 await record_operation_duration(
941 command_name=command_name,
942 duration_seconds=time.monotonic() - start_time,
943 server_address=getattr(conn, "host", None),
944 server_port=getattr(conn, "port", None),
945 db_namespace=str(conn.db),
946 error=error,
947 retry_attempts=failure_count,
948 )
950 await conn.disconnect(error=error, failure_count=failure_count)
952 # COMMAND EXECUTION AND PROTOCOL PARSING
953 async def execute_command(self, *args, **options):
954 """Execute a command and return a parsed response"""
955 await self.initialize()
956 pool = self.connection_pool
957 command_name = args[0]
958 conn = self.connection or await pool.get_connection()
960 # Start timing for observability
961 start_time = time.monotonic()
962 # Track actual retry attempts for error reporting
963 actual_retry_attempts = 0
965 def failure_callback(error, failure_count):
966 if is_debug_log_enabled():
967 add_debug_log_for_operation_failure(conn, error, args)
968 nonlocal actual_retry_attempts
969 actual_retry_attempts = failure_count
970 return self._close_connection(
971 conn, error, failure_count, start_time, command_name
972 )
974 if self.single_connection_client:
975 await self._single_conn_lock.acquire()
976 try:
977 result = await conn.retry.call_with_retry(
978 lambda: self._send_command_parse_response(
979 conn, command_name, *args, **options
980 ),
981 failure_callback,
982 with_failure_count=True,
983 )
985 await record_operation_duration(
986 command_name=command_name,
987 duration_seconds=time.monotonic() - start_time,
988 server_address=getattr(conn, "host", None),
989 server_port=getattr(conn, "port", None),
990 db_namespace=str(conn.db),
991 )
992 return result
993 except Exception as e:
994 await record_error_count(
995 server_address=getattr(conn, "host", None),
996 server_port=getattr(conn, "port", None),
997 network_peer_address=getattr(conn, "host", None),
998 network_peer_port=getattr(conn, "port", None),
999 error_type=e,
1000 retry_attempts=actual_retry_attempts,
1001 is_internal=False,
1002 )
1003 raise
1004 finally:
1005 try:
1006 if self.single_connection_client and conn and conn.should_reconnect():
1007 await self._close_connection(conn)
1008 await conn.connect()
1009 finally:
1010 if self.single_connection_client:
1011 self._single_conn_lock.release()
1012 if not self.connection:
1013 await pool.release(conn)
1015 async def parse_response(
1016 self, connection: Connection, command_name: Union[str, bytes], **options
1017 ):
1018 """Parses a response from the Redis server"""
1019 try:
1020 if NEVER_DECODE in options:
1021 response = await connection.read_response(disable_decoding=True)
1022 options.pop(NEVER_DECODE)
1023 else:
1024 response = await connection.read_response()
1025 except ResponseError:
1026 if EMPTY_RESPONSE in options:
1027 return options[EMPTY_RESPONSE]
1028 raise
1030 if EMPTY_RESPONSE in options:
1031 options.pop(EMPTY_RESPONSE)
1033 # Remove keys entry, it needs only for cache.
1034 options.pop("keys", None)
1036 if command_name in self.response_callbacks:
1037 # Mypy bug: https://github.com/python/mypy/issues/10977
1038 command_name = cast(str, command_name)
1039 retval = self.response_callbacks[command_name](response, **options)
1040 return await retval if inspect.isawaitable(retval) else retval
1041 return response
1043 # HIMPORT orchestration (async mirror of redis.client.Redis). See
1044 # ``.agents/himport_client_support_spec.md``.
1046 @experimental_method()
1047 async def himport_prepare(
1048 self, fieldset_name: str, fields: Iterable[FieldT]
1049 ) -> bool:
1050 """Declare an HIMPORT fieldset for use by :meth:`himport_set`."""
1051 await self.initialize()
1052 fieldset = self.himport_registry.prepare(fieldset_name, fields)
1053 conn = self.connection
1054 if self.single_connection_client and conn is not None and conn.is_connected:
1055 await self.himport_prepare_internal(fieldset_name, fieldset.fields)
1056 conn._himport_prepared[fieldset_name] = fieldset.version
1057 return True
1059 @experimental_method()
1060 async def himport_discard(self, fieldset_name: str) -> int:
1061 """Remove a fieldset from the registry."""
1062 await self.initialize()
1063 removed = self.himport_registry.discard(fieldset_name)
1064 conn = self.connection
1065 if self.single_connection_client and conn is not None and conn.is_connected:
1066 if removed:
1067 await self.himport_discard_internal(fieldset_name)
1068 conn._himport_prepared.pop(fieldset_name, None)
1069 conn._himport_reconciled_revision = self.himport_registry.revision
1070 return 1 if removed else 0
1072 @experimental_method()
1073 async def himport_discard_all(self) -> int:
1074 """Remove all fieldsets from the registry."""
1075 await self.initialize()
1076 count = self.himport_registry.discard_all()
1077 conn = self.connection
1078 if self.single_connection_client and conn is not None and conn.is_connected:
1079 if count:
1080 await self.himport_discard_all_internal()
1081 conn._himport_prepared.clear()
1082 conn._himport_reconciled_revision = self.himport_registry.revision
1083 return count
1086StrictRedis = Redis
1089class MonitorCommandInfo(TypedDict):
1090 time: float
1091 db: int
1092 client_address: str
1093 client_port: str
1094 client_type: str
1095 command: str
1098class Monitor:
1099 """
1100 Monitor is useful for handling the MONITOR command to the redis server.
1101 next_command() method returns one command from monitor
1102 listen() method yields commands from monitor.
1103 """
1105 monitor_re = re.compile(r"\[(\d+) (.*?)\] (.*)")
1106 command_re = re.compile(r'"(.*?)(?<!\\)"')
1108 def __init__(self, connection_pool: ConnectionPool):
1109 self.connection_pool = connection_pool
1110 self.connection: Optional[Connection] = None
1112 async def connect(self):
1113 if self.connection is None:
1114 self.connection = await self.connection_pool.get_connection()
1116 async def __aenter__(self):
1117 await self.connect()
1118 await self.connection.send_command("MONITOR")
1119 # check that monitor returns 'OK', but don't return it to user
1120 response = await self.connection.read_response()
1121 if not bool_ok(response):
1122 raise RedisError(f"MONITOR failed: {response}")
1123 return self
1125 async def __aexit__(self, *args):
1126 await self.connection.disconnect()
1127 await self.connection_pool.release(self.connection)
1129 async def next_command(self) -> MonitorCommandInfo:
1130 """Parse the response from a monitor command"""
1131 await self.connect()
1132 response = await self.connection.read_response()
1133 if isinstance(response, bytes):
1134 response = self.connection.encoder.decode(response, force=True)
1135 command_time, command_data = response.split(" ", 1)
1136 m = self.monitor_re.match(command_data)
1137 db_id, client_info, command = m.groups()
1138 command = " ".join(self.command_re.findall(command))
1139 # Redis escapes double quotes because each piece of the command
1140 # string is surrounded by double quotes. We don't have that
1141 # requirement so remove the escaping and leave the quote.
1142 command = command.replace('\\"', '"')
1144 if client_info == "lua":
1145 client_address = "lua"
1146 client_port = ""
1147 client_type = "lua"
1148 elif client_info.startswith("unix"):
1149 client_address = "unix"
1150 client_port = client_info[5:]
1151 client_type = "unix"
1152 else:
1153 # use rsplit as ipv6 addresses contain colons
1154 client_address, client_port = client_info.rsplit(":", 1)
1155 client_type = "tcp"
1156 return {
1157 "time": float(command_time),
1158 "db": int(db_id),
1159 "client_address": client_address,
1160 "client_port": client_port,
1161 "client_type": client_type,
1162 "command": command,
1163 }
1165 async def listen(self) -> AsyncIterator[MonitorCommandInfo]:
1166 """Listen for commands coming to the server."""
1167 while True:
1168 yield await self.next_command()
1171class PubSub:
1172 """
1173 PubSub provides publish, subscribe and listen support to Redis channels.
1175 After subscribing to one or more channels, the listen() method will block
1176 until a message arrives on one of the subscribed channels. That message
1177 will be returned and it's safe to start listening again.
1178 """
1180 PUBLISH_MESSAGE_TYPES = ("message", "pmessage", "smessage")
1181 UNSUBSCRIBE_MESSAGE_TYPES = ("unsubscribe", "punsubscribe", "sunsubscribe")
1182 HEALTH_CHECK_MESSAGE = "redis-py-health-check"
1184 def __init__(
1185 self,
1186 connection_pool: ConnectionPool,
1187 shard_hint: Optional[str] = None,
1188 ignore_subscribe_messages: bool = False,
1189 encoder=None,
1190 push_handler_func: Optional[Callable] = None,
1191 event_dispatcher: Optional["EventDispatcher"] = None,
1192 ):
1193 if event_dispatcher is None:
1194 self._event_dispatcher = EventDispatcher()
1195 else:
1196 self._event_dispatcher = event_dispatcher
1197 self.connection_pool = connection_pool
1198 self.shard_hint = shard_hint
1199 self.ignore_subscribe_messages = ignore_subscribe_messages
1200 self.connection = None
1201 # we need to know the encoding options for this connection in order
1202 # to lookup channel and pattern names for callback handlers.
1203 self.encoder = encoder
1204 self.push_handler_func = push_handler_func
1205 if self.encoder is None:
1206 self.encoder = self.connection_pool.get_encoder()
1207 if self.encoder.decode_responses:
1208 self.health_check_response = [
1209 ["pong", self.HEALTH_CHECK_MESSAGE],
1210 self.HEALTH_CHECK_MESSAGE,
1211 ]
1212 else:
1213 self.health_check_response = [
1214 [b"pong", self.encoder.encode(self.HEALTH_CHECK_MESSAGE)],
1215 self.encoder.encode(self.HEALTH_CHECK_MESSAGE),
1216 ]
1217 if self.push_handler_func is None:
1218 _set_info_logger()
1219 self.channels = {}
1220 self.pending_unsubscribe_channels = set()
1221 self.patterns = {}
1222 self.pending_unsubscribe_patterns = set()
1223 self.shard_channels = {}
1224 self.pending_unsubscribe_shard_channels = set()
1225 self._lock = asyncio.Lock()
1227 async def __aenter__(self):
1228 return self
1230 async def __aexit__(self, exc_type, exc_value, traceback):
1231 await self.aclose()
1233 def __del__(self):
1234 if self.connection:
1235 self.connection.deregister_connect_callback(self.on_connect)
1237 async def aclose(self):
1238 # In case a connection property does not yet exist
1239 # (due to a crash earlier in the Redis() constructor), return
1240 # immediately as there is nothing to clean-up.
1241 if not hasattr(self, "connection"):
1242 return
1243 async with self._lock:
1244 if self.connection:
1245 # Use nowait=True to avoid awaiting StreamWriter.wait_closed(),
1246 # which can deadlock when a concurrent reader task (e.g. one
1247 # running pubsub.run() or get_message(block=True)) still holds
1248 # the transport. See https://github.com/redis/redis-py/issues/3941
1249 await self.connection.disconnect(nowait=True)
1250 self.connection.deregister_connect_callback(self.on_connect)
1251 await self.connection_pool.release(self.connection)
1252 self.connection = None
1253 self.channels = {}
1254 self.pending_unsubscribe_channels = set()
1255 self.patterns = {}
1256 self.pending_unsubscribe_patterns = set()
1257 self.shard_channels = {}
1258 self.pending_unsubscribe_shard_channels = set()
1260 @deprecated_function(version="5.0.1", reason="Use aclose() instead", name="close")
1261 async def close(self) -> None:
1262 """Alias for aclose(), for backwards compatibility"""
1263 await self.aclose()
1265 @deprecated_function(version="5.0.1", reason="Use aclose() instead", name="reset")
1266 async def reset(self) -> None:
1267 """Alias for aclose(), for backwards compatibility"""
1268 await self.aclose()
1270 async def _resubscribe(self, subscribed, subscribe_fn) -> None:
1271 # Replay handler-backed subscriptions as positional Subscription objects
1272 # so binary names never need to be decoded into keyword argument keys.
1273 subscriptions = pubsub_subscription_args(subscribed)
1274 if subscriptions:
1275 await subscribe_fn(*subscriptions)
1277 async def _resubscribe_shard_channels(self) -> None:
1278 await self._resubscribe(self.shard_channels, self.ssubscribe)
1280 async def on_connect(self, connection: Connection):
1281 """Re-subscribe to any channels and patterns previously subscribed to"""
1282 self.pending_unsubscribe_channels.clear()
1283 self.pending_unsubscribe_patterns.clear()
1284 self.pending_unsubscribe_shard_channels.clear()
1285 if self.channels:
1286 await self._resubscribe(self.channels, self.subscribe)
1287 if self.patterns:
1288 await self._resubscribe(self.patterns, self.psubscribe)
1289 if self.shard_channels:
1290 await self._resubscribe_shard_channels()
1292 @property
1293 def subscribed(self):
1294 """Indicates if there are subscriptions to any channels or patterns"""
1295 return bool(self.channels or self.patterns or self.shard_channels)
1297 async def execute_command(self, *args: EncodableT):
1298 """Execute a publish/subscribe command"""
1300 # NOTE: don't parse the response in this function -- it could pull a
1301 # legitimate message off the stack if the connection is already
1302 # subscribed to one or more channels
1304 await self.connect()
1305 connection = self.connection
1306 kwargs = {"check_health": not self.subscribed}
1307 await self._execute(connection, connection.send_command, *args, **kwargs)
1309 async def connect(self):
1310 """
1311 Ensure that the PubSub is connected
1312 """
1313 if self.connection is None:
1314 self.connection = await self.connection_pool.get_connection()
1315 # register a callback that re-subscribes to any channels we
1316 # were listening to when we were disconnected
1317 self.connection.register_connect_callback(self.on_connect)
1318 else:
1319 await self.connection.connect()
1320 if self.push_handler_func is not None:
1321 self.connection._parser.set_pubsub_push_handler(self.push_handler_func)
1323 self._event_dispatcher.dispatch(
1324 AfterPubSubConnectionInstantiationEvent(
1325 self.connection, self.connection_pool, ClientType.ASYNC, self._lock
1326 )
1327 )
1329 async def _reconnect(
1330 self,
1331 conn,
1332 error: Optional[BaseException] = None,
1333 failure_count: Optional[int] = None,
1334 start_time: Optional[float] = None,
1335 command_name: Optional[str] = None,
1336 ):
1337 """
1338 The supported exceptions are already checked in the
1339 retry object so we don't need to do it here.
1341 In this error handler we are trying to reconnect to the server.
1342 """
1343 if (
1344 error
1345 and failure_count is not None
1346 and failure_count <= conn.retry.get_retries()
1347 ):
1348 if command_name:
1349 await record_operation_duration(
1350 command_name=command_name,
1351 duration_seconds=time.monotonic() - start_time,
1352 server_address=getattr(conn, "host", None),
1353 server_port=getattr(conn, "port", None),
1354 db_namespace=str(conn.db),
1355 error=error,
1356 retry_attempts=failure_count,
1357 )
1358 await conn.disconnect(error=error, failure_count=failure_count)
1359 await conn.connect()
1361 async def _execute(self, conn, command, *args, **kwargs):
1362 """
1363 Connect manually upon disconnection. If the Redis server is down,
1364 this will fail and raise a ConnectionError as desired.
1365 After reconnection, the ``on_connect`` callback should have been
1366 called by the # connection to resubscribe us to any channels and
1367 patterns we were previously listening to
1368 """
1369 if not len(args) == 0:
1370 command_name = args[0]
1371 else:
1372 command_name = None
1374 # Start timing for observability
1375 start_time = time.monotonic()
1376 # Track actual retry attempts for error reporting
1377 actual_retry_attempts = 0
1379 def failure_callback(error, failure_count):
1380 if is_debug_log_enabled():
1381 add_debug_log_for_operation_failure(conn, error, args)
1382 nonlocal actual_retry_attempts
1383 actual_retry_attempts = failure_count
1384 return self._reconnect(conn, error, failure_count, start_time, command_name)
1386 try:
1387 response = await conn.retry.call_with_retry(
1388 lambda: command(*args, **kwargs),
1389 failure_callback,
1390 with_failure_count=True,
1391 )
1393 if command_name:
1394 await record_operation_duration(
1395 command_name=command_name,
1396 duration_seconds=time.monotonic() - start_time,
1397 server_address=getattr(conn, "host", None),
1398 server_port=getattr(conn, "port", None),
1399 db_namespace=str(conn.db),
1400 )
1402 return response
1403 except Exception as e:
1404 await record_error_count(
1405 server_address=getattr(conn, "host", None),
1406 server_port=getattr(conn, "port", None),
1407 network_peer_address=getattr(conn, "host", None),
1408 network_peer_port=getattr(conn, "port", None),
1409 error_type=e,
1410 retry_attempts=actual_retry_attempts,
1411 is_internal=False,
1412 )
1413 raise
1415 async def parse_response(self, block: bool = True, timeout: float = 0):
1416 """
1417 Parse the response from a publish/subscribe command.
1419 Args:
1420 block: If True, block indefinitely until a message is available.
1421 If False, return immediately if no message is available.
1422 Default: True
1423 timeout: The timeout in seconds for reading a response when block=False.
1424 This parameter is ignored when block=True.
1425 Default: 0 (return immediately if no data available)
1427 Returns:
1428 The parsed response from the server, or None if no message is available
1429 within the timeout period (when block=False).
1431 Important:
1432 The block and timeout parameters work together:
1433 - When block=True: timeout is IGNORED, method blocks indefinitely
1434 - When block=False: timeout is USED, method returns after timeout expires
1436 Typically, you should use get_message(timeout=X) instead of calling
1437 parse_response() directly. The get_message() method automatically sets
1438 block=False when a timeout is provided, and block=True when timeout=None.
1440 Example:
1441 # Block indefinitely (timeout is ignored)
1442 response = await pubsub.parse_response(block=True, timeout=0.1)
1444 # Non-blocking with 0.1 second timeout
1445 response = await pubsub.parse_response(block=False, timeout=0.1)
1447 # Non-blocking, return immediately
1448 response = await pubsub.parse_response(block=False, timeout=0)
1450 # Recommended: use get_message() instead
1451 msg = await pubsub.get_message(timeout=0.1) # automatically sets block=False
1452 msg = await pubsub.get_message(timeout=None) # automatically sets block=True
1453 """
1454 conn = self.connection
1455 if conn is None:
1456 raise RuntimeError(
1457 "pubsub connection not set: "
1458 "did you forget to call subscribe() or psubscribe()?"
1459 )
1461 await self.check_health()
1463 if not conn.is_connected:
1464 await conn.connect()
1466 # Block=True: signal "no timeout" to conn.read_response via
1467 # math.inf. The connection treats math.inf as the per-read
1468 # opt-in for blocking indefinitely without falling back to
1469 # self.socket_timeout. Reconnect/AUTH/HELLO/resubscribe
1470 # operations performed by the retry layer continue to honor
1471 # self.socket_timeout because they do not pass math.inf.
1472 #
1473 # TODO(next-major): when the async Connection.read_response
1474 # default for ``timeout`` is changed to SENTINEL, passing
1475 # ``timeout=None`` from this method will become the natural
1476 # "no timeout" signal and the math.inf hand-off can be
1477 # removed. That swap is a breaking change to the
1478 # Connection.read_response signature so it must wait for a
1479 # major release.
1480 read_timeout = math.inf if block else timeout
1481 response = await self._execute(
1482 conn,
1483 conn.read_response,
1484 timeout=read_timeout,
1485 disconnect_on_error=False,
1486 push_request=True,
1487 )
1489 if conn.health_check_interval and response in self.health_check_response:
1490 # ignore the health check message as user might not expect it
1491 return None
1492 return response
1494 async def check_health(self):
1495 conn = self.connection
1496 if conn is None:
1497 raise RuntimeError(
1498 "pubsub connection not set: "
1499 "did you forget to call subscribe() or psubscribe()?"
1500 )
1502 if (
1503 conn.health_check_interval
1504 and asyncio.get_running_loop().time() > conn.next_health_check
1505 ):
1506 await conn.send_command(
1507 "PING", self.HEALTH_CHECK_MESSAGE, check_health=False
1508 )
1510 def _normalize_keys(self, data: _NormalizeKeysT) -> _NormalizeKeysT:
1511 """
1512 normalize channel/pattern names to be either bytes or strings
1513 based on whether responses are automatically decoded. this saves us
1514 from coercing the value for each message coming in.
1515 """
1516 encode = self.encoder.encode
1517 decode = self.encoder.decode
1518 return {decode(encode(k)): v for k, v in data.items()} # type: ignore[return-value] # noqa: E501
1520 async def psubscribe(
1521 self, *args: ChannelT | Subscription, **kwargs: PubSubHandler
1522 ) -> None:
1523 """
1524 Subscribe to channel patterns.
1525 Patterns supplied as keyword arguments expect a pattern name as the
1526 key and a callable as the value.
1527 ``Subscription`` objects can also be supplied positionally with an
1528 optional handler.
1529 A pattern's callable will be invoked automatically
1530 when a message is received on that pattern rather than producing a
1531 message via ``listen()``.
1532 """
1533 new_patterns = parse_pubsub_subscriptions(args, kwargs)
1534 ret_val = await self.execute_command("PSUBSCRIBE", *new_patterns.keys())
1535 # update the patterns dict AFTER we send the command. we don't want to
1536 # subscribe twice to these patterns, once for the command and again
1537 # for the reconnection.
1538 new_patterns = self._normalize_keys(new_patterns)
1539 self.patterns.update(new_patterns)
1540 self.pending_unsubscribe_patterns.difference_update(new_patterns)
1541 return ret_val
1543 def punsubscribe(self, *args: ChannelT) -> Awaitable:
1544 """
1545 Unsubscribe from the supplied patterns. If empty, unsubscribe from
1546 all patterns.
1547 """
1548 patterns: Iterable[ChannelT]
1549 if args:
1550 parsed_args = list_or_args((args[0],), args[1:])
1551 patterns = self._normalize_keys(dict.fromkeys(parsed_args)).keys()
1552 else:
1553 parsed_args = []
1554 patterns = self.patterns
1555 self.pending_unsubscribe_patterns.update(patterns)
1556 return self.execute_command("PUNSUBSCRIBE", *parsed_args)
1558 async def subscribe(
1559 self, *args: ChannelT | Subscription, **kwargs: PubSubHandler
1560 ) -> None:
1561 """
1562 Subscribe to channels.
1563 Channels supplied as keyword arguments expect
1564 a channel name as the key and a callable as the value.
1565 ``Subscription`` objects can also be supplied positionally with an
1566 optional handler.
1567 A channel's callable will be invoked automatically
1568 when a message is received on that channel rather than producing a
1569 message via ``listen()`` or ``get_message()``.
1570 """
1571 new_channels = parse_pubsub_subscriptions(args, kwargs)
1572 ret_val = await self.execute_command("SUBSCRIBE", *new_channels.keys())
1573 # update the channels dict AFTER we send the command. we don't want to
1574 # subscribe twice to these channels, once for the command and again
1575 # for the reconnection.
1576 new_channels = self._normalize_keys(new_channels)
1577 self.channels.update(new_channels)
1578 self.pending_unsubscribe_channels.difference_update(new_channels)
1579 return ret_val
1581 def unsubscribe(self, *args) -> Awaitable:
1582 """
1583 Unsubscribe from the supplied channels. If empty, unsubscribe from
1584 all channels
1585 """
1586 if args:
1587 parsed_args = list_or_args(args[0], args[1:])
1588 channels = self._normalize_keys(dict.fromkeys(parsed_args))
1589 else:
1590 parsed_args = []
1591 channels = self.channels
1592 self.pending_unsubscribe_channels.update(channels)
1593 return self.execute_command("UNSUBSCRIBE", *parsed_args)
1595 async def ssubscribe(
1596 self,
1597 *args: ChannelT | Subscription,
1598 target_node: Any = None,
1599 **kwargs: PubSubHandler,
1600 ) -> None:
1601 """
1602 Subscribes the client to the specified shard channels.
1603 Channels supplied as keyword arguments expect a channel name as the key
1604 and a callable as the value.
1605 ``Subscription`` objects can also be supplied positionally
1606 with an optional handler.
1607 A channel's callable will be invoked automatically when a message
1608 is received on that channel rather than producing a message
1609 via ``listen()`` or ``get_sharded_message()``.
1610 """
1611 new_s_channels = parse_pubsub_subscriptions(args, kwargs)
1612 ret_val = await self.execute_command("SSUBSCRIBE", *new_s_channels.keys())
1613 # update the s_channels dict AFTER we send the command. we don't want to
1614 # subscribe twice to these channels, once for the command and again
1615 # for the reconnection.
1616 new_s_channels = self._normalize_keys(new_s_channels)
1617 self.shard_channels.update(new_s_channels)
1618 self.pending_unsubscribe_shard_channels.difference_update(new_s_channels)
1619 return ret_val
1621 def sunsubscribe(self, *args, target_node=None) -> Awaitable:
1622 """
1623 Unsubscribe from the supplied shard_channels. If empty, unsubscribe from
1624 all shard_channels
1625 """
1626 if args:
1627 args = list_or_args(args[0], args[1:])
1628 s_channels = self._normalize_keys(dict.fromkeys(args))
1629 else:
1630 s_channels = self.shard_channels
1631 self.pending_unsubscribe_shard_channels.update(s_channels)
1632 return self.execute_command("SUNSUBSCRIBE", *args)
1634 async def listen(self) -> AsyncIterator:
1635 """Listen for messages on channels this client has been subscribed to.
1637 Iteration ends once every channel and pattern has been unsubscribed
1638 from. If nothing is subscribed when iteration begins it ends
1639 immediately rather than waiting, so subscribe first: a listener
1640 started before any subscription finishes without yielding anything.
1641 """
1642 while self.subscribed:
1643 response = await self.handle_message(await self.parse_response(block=True))
1644 if response is not None:
1645 yield response
1647 async def get_message(
1648 self, ignore_subscribe_messages: bool = False, timeout: Optional[float] = 0.0
1649 ):
1650 """
1651 Get the next message if one is available, otherwise None.
1653 If timeout is specified, the system will wait for `timeout` seconds
1654 before returning. Timeout should be specified as a floating point
1655 number or None to wait indefinitely.
1656 """
1657 response = await self.parse_response(block=(timeout is None), timeout=timeout)
1658 if response:
1659 return await self.handle_message(response, ignore_subscribe_messages)
1660 return None
1662 def ping(self, message=None) -> Awaitable[bool]:
1663 """
1664 Ping the Redis server to test connectivity.
1666 Sends a PING command to the Redis server and returns True if the server
1667 responds with "PONG".
1668 """
1669 args = ["PING", message] if message is not None else ["PING"]
1670 return self.execute_command(*args)
1672 async def handle_message(self, response, ignore_subscribe_messages=False):
1673 """
1674 Parses a pub/sub message. If the channel or pattern was subscribed to
1675 with a message handler, the handler is invoked instead of a parsed
1676 message being returned.
1677 """
1678 if response is None:
1679 return None
1680 if isinstance(response, bytes):
1681 response = [b"pong", response] if response != b"PONG" else [b"pong", b""]
1682 message_type = str_if_bytes(response[0])
1683 if message_type == "pmessage":
1684 message = {
1685 "type": message_type,
1686 "pattern": response[1],
1687 "channel": response[2],
1688 "data": response[3],
1689 }
1690 elif message_type == "pong":
1691 message = {
1692 "type": message_type,
1693 "pattern": None,
1694 "channel": None,
1695 "data": response[1],
1696 }
1697 else:
1698 message = {
1699 "type": message_type,
1700 "pattern": None,
1701 "channel": response[1],
1702 "data": response[2],
1703 }
1705 if message_type in ["message", "pmessage"]:
1706 channel = str_if_bytes(message["channel"])
1707 await record_pubsub_message(
1708 direction=PubSubDirection.RECEIVE,
1709 channel=channel,
1710 )
1711 elif message_type == "smessage":
1712 channel = str_if_bytes(message["channel"])
1713 await record_pubsub_message(
1714 direction=PubSubDirection.RECEIVE,
1715 channel=channel,
1716 sharded=True,
1717 )
1719 # if this is an unsubscribe message, remove it from memory.
1720 # ``discard`` rather than ``remove``: the guard above already makes the
1721 # removal conditional, so the two are equivalent for a single caller -
1722 # but another writer can drop the same entry between the check and the
1723 # removal, and ``remove`` would then raise ``KeyError`` out of a pubsub
1724 # read that no caller catches. ``ClusterPubSub._detach_shard_channel``
1725 # is such a writer: it forgets a migrating shard channel locally,
1726 # deliberately without the per-node I/O lock this bookkeeping runs
1727 # under, because waiting for that lock stalls reconciliation behind a
1728 # poll's whole retry budget on the node being migrated away from.
1729 if message_type in self.UNSUBSCRIBE_MESSAGE_TYPES:
1730 if message_type == "punsubscribe":
1731 pattern = response[1]
1732 if pattern in self.pending_unsubscribe_patterns:
1733 self.pending_unsubscribe_patterns.discard(pattern)
1734 self.patterns.pop(pattern, None)
1735 elif message_type == "sunsubscribe":
1736 s_channel = response[1]
1737 if s_channel in self.pending_unsubscribe_shard_channels:
1738 self.pending_unsubscribe_shard_channels.discard(s_channel)
1739 self.shard_channels.pop(s_channel, None)
1740 else:
1741 channel = response[1]
1742 if channel in self.pending_unsubscribe_channels:
1743 self.pending_unsubscribe_channels.discard(channel)
1744 self.channels.pop(channel, None)
1746 if message_type in self.PUBLISH_MESSAGE_TYPES:
1747 # if there's a message handler, invoke it
1748 if message_type == "pmessage":
1749 handler = self.patterns.get(message["pattern"], None)
1750 elif message_type == "smessage":
1751 handler = self.shard_channels.get(message["channel"], None)
1752 else:
1753 handler = self.channels.get(message["channel"], None)
1754 if handler:
1755 if inspect.iscoroutinefunction(handler):
1756 await handler(message)
1757 else:
1758 handler(message)
1759 return None
1760 elif message_type != "pong":
1761 # this is a subscribe/unsubscribe message. ignore if we don't
1762 # want them
1763 if ignore_subscribe_messages or self.ignore_subscribe_messages:
1764 return None
1766 return message
1768 async def run(
1769 self,
1770 *,
1771 exception_handler: Optional["PSWorkerThreadExcHandlerT"] = None,
1772 poll_timeout: float = 1.0,
1773 pubsub=None,
1774 ) -> None:
1775 """Process pub/sub messages using registered callbacks.
1777 This is the equivalent of :py:meth:`redis.PubSub.run_in_thread` in
1778 redis-py, but it is a coroutine. To launch it as a separate task, use
1779 ``asyncio.create_task``:
1781 >>> task = asyncio.create_task(pubsub.run())
1783 To shut it down, use asyncio cancellation:
1785 >>> task.cancel()
1786 >>> await task
1787 """
1788 for channel, handler in self.channels.items():
1789 if handler is None:
1790 raise PubSubError(f"Channel: '{channel}' has no handler registered")
1791 for pattern, handler in self.patterns.items():
1792 if handler is None:
1793 raise PubSubError(f"Pattern: '{pattern}' has no handler registered")
1795 await self.connect()
1796 while True:
1797 try:
1798 if pubsub is None:
1799 await self.get_message(
1800 ignore_subscribe_messages=True, timeout=poll_timeout
1801 )
1802 else:
1803 await pubsub.get_message(
1804 ignore_subscribe_messages=True, timeout=poll_timeout
1805 )
1806 except asyncio.CancelledError:
1807 raise
1808 except BaseException as e:
1809 if exception_handler is None:
1810 raise
1811 res = exception_handler(e, self)
1812 if inspect.isawaitable(res):
1813 await res
1814 # Ensure that other tasks on the event loop get a chance to run
1815 # if we didn't have to block for I/O anywhere.
1816 await asyncio.sleep(0)
1819class PubsubWorkerExceptionHandler(Protocol):
1820 def __call__(self, e: BaseException, pubsub: PubSub): ...
1823class AsyncPubsubWorkerExceptionHandler(Protocol):
1824 async def __call__(self, e: BaseException, pubsub: PubSub): ...
1827PSWorkerThreadExcHandlerT = Union[
1828 PubsubWorkerExceptionHandler, AsyncPubsubWorkerExceptionHandler
1829]
1832CommandT = Tuple[Tuple[Union[str, bytes], ...], Mapping[str, Any]]
1833CommandStackT = List[CommandT]
1836class Pipeline(Redis): # lgtm [py/init-calls-subclass]
1837 """
1838 Pipelines provide a way to transmit multiple commands to the Redis server
1839 in one transmission. This is convenient for batch processing, such as
1840 saving all the values in a list to Redis.
1842 All commands executed within a pipeline(when running in transactional mode,
1843 which is the default behavior) are wrapped with MULTI and EXEC
1844 calls. This guarantees all commands executed in the pipeline will be
1845 executed atomically.
1847 Any command raising an exception does *not* halt the execution of
1848 subsequent commands in the pipeline. Instead, the exception is caught
1849 and its instance is placed into the response list returned by execute().
1850 Code iterating over the response list should be able to deal with an
1851 instance of an exception as a potential value. In general, these will be
1852 ResponseError exceptions, such as those raised when issuing a command
1853 on a key of a different datatype.
1854 """
1856 UNWATCH_COMMANDS = {"DISCARD", "EXEC", "UNWATCH"}
1858 def __init__(
1859 self,
1860 connection_pool: ConnectionPool,
1861 response_callbacks: MutableMapping[Union[str, bytes], ResponseCallbackT],
1862 transaction: bool,
1863 shard_hint: Optional[str],
1864 ):
1865 self.connection_pool = connection_pool
1866 self.connection = None
1867 self.response_callbacks = response_callbacks
1868 self.is_transaction = transaction
1869 self.shard_hint = shard_hint
1870 self.watching = False
1871 self.command_stack: CommandStackT = []
1872 self.scripts: Set[Script] = set()
1873 self.explicit_transaction = False
1875 async def __aenter__(self: _RedisT) -> _RedisT:
1876 return self
1878 async def __aexit__(self, exc_type, exc_value, traceback):
1879 await self.reset()
1881 def __await__(self):
1882 return self._async_self().__await__()
1884 _DEL_MESSAGE = "Unclosed Pipeline client"
1886 def __len__(self):
1887 return len(self.command_stack)
1889 def __bool__(self):
1890 """Pipeline instances should always evaluate to True"""
1891 return True
1893 async def _async_self(self):
1894 return self
1896 async def reset(self):
1897 self.command_stack = []
1898 self.scripts = set()
1899 try:
1900 # make sure to reset the connection state in the event that we were
1901 # watching something
1902 if self.watching and self.connection:
1903 try:
1904 # call this manually since our unwatch or
1905 # immediate_execute_command methods can call reset()
1906 await self.connection.send_command("UNWATCH")
1907 await self.connection.read_response()
1908 except ConnectionError:
1909 # disconnect will also remove any previous WATCHes
1910 if self.connection:
1911 await self.connection.disconnect()
1912 except asyncio.CancelledError:
1913 # Disconnect so any unread UNWATCH reply does not get
1914 # served to the next caller that takes the connection.
1915 if self.connection:
1916 await self.connection.disconnect()
1917 raise
1918 finally:
1919 self.watching = False
1920 self.explicit_transaction = False
1921 # We can safely return the connection to the pool here since we're
1922 # sure we're no longer WATCHing anything. Detach self.connection
1923 # before awaiting release: if a second cancel aborts the await,
1924 # the pipeline must not be left holding a reference to a
1925 # connection that is being returned to the pool. Shield the
1926 # release itself so a second cancel cannot split the pool's
1927 # internal in-use/available bookkeeping mid-update.
1928 if self.connection:
1929 connection, self.connection = self.connection, None
1930 await asyncio.shield(self.connection_pool.release(connection))
1932 async def aclose(self) -> None:
1933 """Alias for reset(), a standard method name for cleanup"""
1934 await self.reset()
1936 def multi(self):
1937 """
1938 Start a transactional block of the pipeline after WATCH commands
1939 are issued. End the transactional block with `execute`.
1940 """
1941 if self.explicit_transaction:
1942 raise RedisError("Cannot issue nested calls to MULTI")
1943 if self.command_stack:
1944 raise RedisError(
1945 "Commands without an initial WATCH have already been issued"
1946 )
1947 self.explicit_transaction = True
1949 def execute_command(
1950 self, *args, **kwargs
1951 ) -> Union["Pipeline", Awaitable["Pipeline"]]:
1952 if (self.watching or args[0] == "WATCH") and not self.explicit_transaction:
1953 return self.immediate_execute_command(*args, **kwargs)
1954 return self.pipeline_execute_command(*args, **kwargs)
1956 async def _disconnect_reset_raise_on_watching(
1957 self,
1958 conn: Connection,
1959 error: Exception,
1960 failure_count: Optional[int] = None,
1961 start_time: Optional[float] = None,
1962 command_name: Optional[str] = None,
1963 ) -> None:
1964 """
1965 Close the connection reset watching state and
1966 raise an exception if we were watching.
1968 The supported exceptions are already checked in the
1969 retry object so we don't need to do it here.
1971 After we disconnect the connection, it will try to reconnect and
1972 do a health check as part of the send_command logic(on connection level).
1973 """
1974 if (
1975 error
1976 and failure_count is not None
1977 and failure_count <= conn.retry.get_retries()
1978 ):
1979 await record_operation_duration(
1980 command_name=command_name,
1981 duration_seconds=time.monotonic() - start_time,
1982 server_address=getattr(conn, "host", None),
1983 server_port=getattr(conn, "port", None),
1984 db_namespace=str(conn.db),
1985 error=error,
1986 retry_attempts=failure_count,
1987 )
1988 await conn.disconnect(error=error, failure_count=failure_count)
1989 # if we were already watching a variable, the watch is no longer
1990 # valid since this connection has died. raise a WatchError, which
1991 # indicates the user should retry this transaction.
1992 if self.watching:
1993 await self.reset()
1994 raise WatchError(
1995 f"A {type(error).__name__} occurred while watching one or more keys"
1996 )
1998 async def immediate_execute_command(self, *args, **options):
1999 """
2000 Execute a command immediately, but don't auto-retry on the supported
2001 errors for retry if we're already WATCHing a variable.
2002 Used when issuing WATCH or subsequent commands retrieving their values but before
2003 MULTI is called.
2004 """
2005 command_name = args[0]
2006 conn = self.connection
2007 # if this is the first call, we need a connection
2008 if not conn:
2009 conn = await self.connection_pool.get_connection()
2010 self.connection = conn
2012 # Start timing for observability
2013 start_time = time.monotonic()
2014 # Track actual retry attempts for error reporting
2015 actual_retry_attempts = 0
2017 def failure_callback(error, failure_count):
2018 if is_debug_log_enabled():
2019 add_debug_log_for_operation_failure(conn, error, args)
2020 nonlocal actual_retry_attempts
2021 actual_retry_attempts = failure_count
2022 return self._disconnect_reset_raise_on_watching(
2023 conn, error, failure_count, start_time, command_name
2024 )
2026 try:
2027 response = await conn.retry.call_with_retry(
2028 lambda: self._send_command_parse_response(
2029 conn, command_name, *args, **options
2030 ),
2031 failure_callback,
2032 with_failure_count=True,
2033 )
2035 await record_operation_duration(
2036 command_name=command_name,
2037 duration_seconds=time.monotonic() - start_time,
2038 server_address=getattr(conn, "host", None),
2039 server_port=getattr(conn, "port", None),
2040 db_namespace=str(conn.db),
2041 )
2043 return response
2044 except Exception as e:
2045 await record_error_count(
2046 server_address=getattr(conn, "host", None),
2047 server_port=getattr(conn, "port", None),
2048 network_peer_address=getattr(conn, "host", None),
2049 network_peer_port=getattr(conn, "port", None),
2050 error_type=e,
2051 retry_attempts=actual_retry_attempts,
2052 is_internal=False,
2053 )
2054 raise
2056 def pipeline_execute_command(self, *args, **options):
2057 """
2058 Stage a command to be executed when execute() is next called
2060 Returns the current Pipeline object back so commands can be
2061 chained together, such as:
2063 pipe = pipe.set('foo', 'bar').incr('baz').decr('bang')
2065 At some other point, you can then run: pipe.execute(),
2066 which will execute all commands queued in the pipe.
2067 """
2068 self.command_stack.append((args, options))
2069 return self
2071 async def _himport_prepare_pipeline(self, conn, commands):
2072 """Delegate to the shared async HIMPORT executor."""
2073 await _himport_exec.prepare_pipeline(self, conn, [args for args, _ in commands])
2075 async def _execute_transaction( # noqa: C901
2076 self, connection: Connection, commands: CommandStackT, raise_on_error
2077 ):
2078 # Ensure fieldsets referenced by buffered HIMPORT SETs are prepared on this
2079 # connection before the MULTI/EXEC block (session state, not transactional).
2080 await self._himport_prepare_pipeline(connection, commands)
2081 pre: CommandT = (("MULTI",), {})
2082 post: CommandT = (("EXEC",), {})
2083 cmds = (pre, *commands, post)
2084 all_cmds = connection.pack_commands(
2085 args for args, options in cmds if EMPTY_RESPONSE not in options
2086 )
2087 await connection.send_packed_command(all_cmds)
2088 errors = []
2090 # parse off the response for MULTI
2091 # NOTE: we need to handle ResponseErrors here and continue
2092 # so that we read all the additional command messages from
2093 # the socket
2094 try:
2095 await self.parse_response(connection, "_")
2096 except ResponseError as err:
2097 errors.append((0, err))
2099 # and all the other commands
2100 for i, command in enumerate(commands):
2101 if EMPTY_RESPONSE in command[1]:
2102 errors.append((i, command[1][EMPTY_RESPONSE]))
2103 else:
2104 try:
2105 await self.parse_response(connection, "_")
2106 except ResponseError as err:
2107 self.annotate_exception(err, i + 1, command[0])
2108 errors.append((i, err))
2110 # parse the EXEC.
2111 try:
2112 response = await self.parse_response(connection, "_")
2113 except ExecAbortError as err:
2114 if errors:
2115 raise errors[0][1] from err
2116 raise
2118 # EXEC clears any watched keys
2119 self.watching = False
2121 if response is None:
2122 raise WatchError("Watched variable changed.") from None
2124 # put any parse errors into the response
2125 for i, e in errors:
2126 response.insert(i, e)
2128 if len(response) != len(commands):
2129 if self.connection:
2130 await self.connection.disconnect()
2131 raise ResponseError(
2132 "Wrong number of response items from pipeline execution"
2133 ) from None
2135 # find any errors in the response and raise if necessary
2136 if raise_on_error:
2137 self.raise_first_error(commands, response)
2139 # We have to run response callbacks manually
2140 data = []
2141 for r, cmd in zip(response, commands):
2142 if not isinstance(r, Exception):
2143 args, options = cmd
2144 command_name = args[0]
2146 # Remove keys entry, it needs only for cache.
2147 options.pop("keys", None)
2149 if command_name in self.response_callbacks:
2150 r = self.response_callbacks[command_name](r, **options)
2151 if inspect.isawaitable(r):
2152 r = await r
2153 data.append(r)
2154 return data
2156 async def _execute_pipeline(
2157 self, connection: Connection, commands: CommandStackT, raise_on_error: bool
2158 ):
2159 # Fold any first-use HIMPORT PREPAREs for referenced fieldsets into the same
2160 # packed write as the queued commands, so a pipeline that lands on a fresh or
2161 # reconnected connection stays a single round trip (the batched write bypasses
2162 # the per-command lazy PREPARE path). Deferred-discard reconciliation happens
2163 # inside pipeline_prepares and only touches the socket when discards are
2164 # actually pending.
2165 fieldsets = await _himport_exec.pipeline_prepares(
2166 self, connection, [args for args, _ in commands]
2167 )
2168 preflight = _himport_exec.prepare_wire_commands(fieldsets)
2169 # build up all commands into a single request to increase network perf
2170 all_cmds = connection.pack_commands(preflight + [args for args, _ in commands])
2171 await connection.send_packed_command(all_cmds)
2173 # Drain the leading PREPARE replies (bookkeeping + capture the first error)
2174 # before the queued replies. Everything on the wire is read before raising so
2175 # the pooled socket never desyncs.
2176 prep_error = await _himport_exec.drain_pipeline_prepares(
2177 self, connection, fieldsets
2178 )
2180 response = []
2181 for args, options in commands:
2182 try:
2183 response.append(
2184 await self.parse_response(connection, args[0], **options)
2185 )
2186 except ResponseError as e:
2187 response.append(e)
2189 # A PREPARE failure (rare: an invalid fieldset definition) is a hard error,
2190 # raised regardless of raise_on_error as it was before folding -- only now
2191 # every reply has already been drained.
2192 if prep_error is not None:
2193 raise prep_error
2194 if raise_on_error:
2195 self.raise_first_error(commands, response)
2196 return response
2198 def raise_first_error(self, commands: CommandStackT, response: Iterable[Any]):
2199 for i, r in enumerate(response):
2200 if isinstance(r, ResponseError):
2201 self.annotate_exception(r, i + 1, commands[i][0])
2202 raise r
2204 def annotate_exception(
2205 self, exception: Exception, number: int, command: Iterable[object]
2206 ) -> None:
2207 cmd = " ".join(map(safe_str, command))
2208 msg = (
2209 f"Command # {number} ({truncate_text(cmd)}) "
2210 f"of pipeline caused error: {exception.args}"
2211 )
2212 exception.args = (msg,) + exception.args[1:]
2214 async def parse_response(
2215 self, connection: Connection, command_name: Union[str, bytes], **options
2216 ):
2217 result = await super().parse_response(connection, command_name, **options)
2218 if command_name in self.UNWATCH_COMMANDS:
2219 self.watching = False
2220 elif command_name == "WATCH":
2221 self.watching = True
2222 return result
2224 async def load_scripts(self):
2225 # make sure all scripts that are about to be run on this pipeline exist
2226 scripts = list(self.scripts)
2227 immediate = self.immediate_execute_command
2228 shas = [s.sha for s in scripts]
2229 # we can't use the normal script_* methods because they would just
2230 # get buffered in the pipeline.
2231 exists = await immediate("SCRIPT EXISTS", *shas)
2232 if not all(exists):
2233 for s, exist in zip(scripts, exists):
2234 if not exist:
2235 s.sha = await immediate("SCRIPT LOAD", s.script)
2237 async def _disconnect_raise_on_watching(
2238 self,
2239 conn: Connection,
2240 error: Exception,
2241 failure_count: Optional[int] = None,
2242 start_time: Optional[float] = None,
2243 command_name: Optional[str] = None,
2244 ):
2245 """
2246 Close the connection, raise an exception if we were watching.
2248 The supported exceptions are already checked in the
2249 retry object so we don't need to do it here.
2251 After we disconnect the connection, it will try to reconnect and
2252 do a health check as part of the send_command logic(on connection level).
2253 """
2254 if (
2255 error
2256 and failure_count is not None
2257 and failure_count <= conn.retry.get_retries()
2258 ):
2259 await record_operation_duration(
2260 command_name=command_name,
2261 duration_seconds=time.monotonic() - start_time,
2262 server_address=getattr(conn, "host", None),
2263 server_port=getattr(conn, "port", None),
2264 db_namespace=str(conn.db),
2265 error=error,
2266 retry_attempts=failure_count,
2267 )
2268 await conn.disconnect(error=error, failure_count=failure_count)
2269 # if we were watching a variable, the watch is no longer valid
2270 # since this connection has died. raise a WatchError, which
2271 # indicates the user should retry this transaction.
2272 if self.watching:
2273 raise WatchError(
2274 f"A {type(error).__name__} occurred while watching one or more keys"
2275 )
2277 async def execute(self, raise_on_error: bool = True) -> List[Any]:
2278 """Execute all the commands in the current pipeline"""
2279 stack = self.command_stack
2280 if not stack and not self.watching:
2281 return []
2282 if self.scripts:
2283 await self.load_scripts()
2284 if self.is_transaction or self.explicit_transaction:
2285 execute = self._execute_transaction
2286 operation_name = "MULTI"
2287 else:
2288 execute = self._execute_pipeline
2289 operation_name = "PIPELINE"
2291 conn = self.connection
2292 if not conn:
2293 conn = await self.connection_pool.get_connection()
2294 # assign to self.connection so reset() releases the connection
2295 # back to the pool after we're done
2296 self.connection = conn
2297 conn = cast(Connection, conn)
2299 # Start timing for observability
2300 start_time = time.monotonic()
2301 # Track actual retry attempts for error reporting
2302 actual_retry_attempts = 0
2304 def failure_callback(error, failure_count):
2305 if is_debug_log_enabled():
2306 add_debug_log_for_operation_failure(conn, error, (operation_name,))
2307 nonlocal actual_retry_attempts
2308 actual_retry_attempts = failure_count
2309 return self._disconnect_raise_on_watching(
2310 conn, error, failure_count, start_time, operation_name
2311 )
2313 try:
2314 response = await conn.retry.call_with_retry(
2315 lambda: execute(conn, stack, raise_on_error),
2316 failure_callback,
2317 with_failure_count=True,
2318 )
2320 await record_operation_duration(
2321 command_name=operation_name,
2322 duration_seconds=time.monotonic() - start_time,
2323 server_address=getattr(conn, "host", None),
2324 server_port=getattr(conn, "port", None),
2325 db_namespace=str(conn.db),
2326 )
2327 return response
2328 except Exception as e:
2329 await record_error_count(
2330 server_address=getattr(conn, "host", None),
2331 server_port=getattr(conn, "port", None),
2332 network_peer_address=getattr(conn, "host", None),
2333 network_peer_port=getattr(conn, "port", None),
2334 error_type=e,
2335 retry_attempts=actual_retry_attempts,
2336 is_internal=False,
2337 )
2338 raise
2339 finally:
2340 await self.reset()
2342 async def discard(self):
2343 """Flushes all previously queued commands
2344 See: https://redis.io/commands/DISCARD
2345 """
2346 await self.execute_command("DISCARD")
2348 async def watch(self, *names: KeyT):
2349 """Watches the values at keys ``names``"""
2350 if self.explicit_transaction:
2351 raise RedisError("Cannot issue a WATCH after a MULTI")
2352 return await self.execute_command("WATCH", *names)
2354 async def unwatch(self):
2355 """Unwatches all previously specified keys"""
2356 return self.watching and await self.execute_command("UNWATCH") or True