1import asyncio
2import contextlib
3import copy
4import inspect
5import logging
6import math
7import socket
8import sys
9import time
10import warnings
11import weakref
12from abc import ABC, abstractmethod
13from itertools import chain
14from types import MappingProxyType
15from typing import (
16 Any,
17 AsyncIterator,
18 Callable,
19 Iterable,
20 List,
21 Literal,
22 Mapping,
23 Optional,
24 Protocol,
25 Set,
26 Tuple,
27 Type,
28 TypedDict,
29 TypeVar,
30 Union,
31)
32from urllib.parse import ParseResult, parse_qs, unquote, urlparse
33
34from ..observability.attributes import (
35 DB_CLIENT_CONNECTION_POOL_NAME,
36 DB_CLIENT_CONNECTION_STATE,
37 AttributeBuilder,
38 ConnectionState,
39 get_pool_name,
40)
41from ..utils import SSL_AVAILABLE, deprecated_function
42
43if SSL_AVAILABLE:
44 import ssl
45 from ssl import SSLContext, TLSVersion, VerifyFlags
46else:
47 ssl = None
48 TLSVersion = None
49 SSLContext = None
50 VerifyFlags = None
51
52from ..auth.token import TokenInterface
53from ..driver_info import DriverInfo, resolve_driver_info
54from ..event import AsyncAfterConnectionReleasedEvent, EventDispatcher
55from ..utils import deprecated_args, format_error_message
56
57# the functionality is available in 3.11.x but has a major issue before
58# 3.11.3. See https://github.com/redis/redis-py/issues/2633
59if sys.version_info >= (3, 11, 3):
60 from asyncio import timeout as async_timeout
61else:
62 from async_timeout import timeout as async_timeout
63
64from redis import exceptions as redis_exceptions
65from redis.asyncio.maint_notifications import (
66 AsyncMaintNotificationsConnectionHandler,
67 AsyncMaintNotificationsPoolHandler,
68 AsyncOSSMaintNotificationsHandler,
69)
70from redis.asyncio.observability.recorder import (
71 record_connection_closed,
72 record_connection_count,
73 record_connection_create_time,
74 record_connection_wait_time,
75 record_error_count,
76)
77from redis.asyncio.retry import Retry
78from redis.backoff import NoBackoff
79from redis.credentials import CredentialProvider, UsernamePasswordCredentialProvider
80from redis.exceptions import (
81 AuthenticationError,
82 AuthenticationWrongNumberOfArgsError,
83 ConnectionError,
84 DataError,
85 MaxConnectionsError,
86 RedisError,
87 ResponseError,
88 TimeoutError,
89)
90from redis.himport import HImportRegistry
91from redis.maint_notifications import (
92 MaintenanceState,
93 MaintNotificationsConfig,
94 NodeMovingNotification,
95 _build_moving_cleanup_connection_kwargs,
96 _build_moving_connection_kwargs,
97)
98from redis.observability.metrics import CloseReason
99from redis.typing import EncodableT
100from redis.utils import (
101 DEFAULT_RESP_VERSION,
102 HIREDIS_AVAILABLE,
103 SENTINEL,
104 check_protocol_version,
105 str_if_bytes,
106)
107
108from .._defaults import (
109 DEFAULT_SOCKET_CONNECT_TIMEOUT,
110 DEFAULT_SOCKET_READ_SIZE,
111 DEFAULT_SOCKET_TIMEOUT,
112 get_default_socket_keepalive_options,
113)
114from .._parsers import (
115 AsyncPushNotificationsParser,
116 BaseParser,
117 Encoder,
118 _AsyncHiredisParser,
119 _AsyncRESP2Parser,
120 _AsyncRESP3Parser,
121)
122
123SYM_STAR = b"*"
124SYM_DOLLAR = b"$"
125SYM_CRLF = b"\r\n"
126SYM_LF = b"\n"
127SYM_EMPTY = b""
128
129DefaultParser: Type[Union[_AsyncRESP2Parser, _AsyncRESP3Parser, _AsyncHiredisParser]]
130if HIREDIS_AVAILABLE:
131 DefaultParser = _AsyncHiredisParser
132else:
133 DefaultParser = _AsyncRESP3Parser
134
135logger = logging.getLogger(__name__)
136
137
138def add_debug_log_for_connection_failure(
139 connection: "AbstractConnection",
140 error: BaseException,
141 operation: str,
142) -> None:
143 """
144 Render the connection's live state on a failure that is about to close it.
145
146 Must be called *before* ``disconnect()``. ``extract_connection_details()``
147 reads the local port and the in-flight read deadline off the transport, so
148 once it is gone it can only report ``not connected`` - which hides exactly
149 the state needed to explain the failure. In particular it is what tells
150 apart a read that ran under the original timeout from one that ran under a
151 relaxed maintenance timeout.
152 """
153 if logger.isEnabledFor(logging.DEBUG):
154 logger.debug(
155 f"{type(error).__name__} while {operation}, "
156 f"with connection: {connection}, "
157 f"details: {connection.extract_connection_details()}, "
158 f"error: {error}",
159 )
160
161
162class ConnectCallbackProtocol(Protocol):
163 def __call__(self, connection: "AbstractConnection"): ...
164
165
166class AsyncConnectCallbackProtocol(Protocol):
167 async def __call__(self, connection: "AbstractConnection"): ...
168
169
170ConnectCallbackT = Union[ConnectCallbackProtocol, AsyncConnectCallbackProtocol]
171
172
173class AsyncMaintNotificationsAbstractConnection:
174 """
175 Internal mixin for async maintenance notification state and parser handlers.
176
177 The sync implementation uses the same mixin-style structure. The async
178 version keeps the notification state and parser handler installation close
179 to the connection without sending the server-side handshake; that is wired
180 in a later step.
181 """
182
183 __slots__ = ()
184
185 def __init__(
186 self,
187 maint_notifications_config: MaintNotificationsConfig | None,
188 maint_notifications_pool_handler: (
189 AsyncMaintNotificationsPoolHandler | None
190 ) = None,
191 maintenance_state: MaintenanceState = MaintenanceState.NONE,
192 maintenance_notification_hash: int | None = None,
193 orig_host_address: str | None = None,
194 orig_socket_timeout: float | None = None,
195 orig_socket_connect_timeout: float | None = None,
196 oss_cluster_maint_notifications_handler: (
197 AsyncOSSMaintNotificationsHandler | None
198 ) = None,
199 parser: BaseParser | None = None,
200 ) -> None:
201 self.maint_notifications_config = maint_notifications_config
202 self.maintenance_state = maintenance_state
203 self.maintenance_notification_hash = maintenance_notification_hash
204 self._processed_start_maint_notifications: set[int] = set()
205 self._skipped_end_maint_notifications: set[int] = set()
206 self._configure_maintenance_notifications(
207 maint_notifications_pool_handler,
208 orig_host_address,
209 orig_socket_timeout,
210 orig_socket_connect_timeout,
211 oss_cluster_maint_notifications_handler,
212 parser,
213 )
214
215 @abstractmethod
216 def _get_parser(self) -> BaseParser:
217 pass
218
219 def _get_push_notifications_parser(self) -> AsyncPushNotificationsParser:
220 parser = self._get_parser()
221 if not isinstance(parser, (_AsyncHiredisParser, _AsyncRESP3Parser)):
222 raise RedisError(
223 "Maintenance notifications are only supported with hiredis and RESP3 parsers!"
224 )
225 return parser
226
227 @abstractmethod
228 def get_protocol(self):
229 pass
230
231 @abstractmethod
232 async def send_command(self, *args: Any, **kwargs: Any) -> None:
233 pass
234
235 @abstractmethod
236 async def read_response(
237 self,
238 disable_decoding: bool = False,
239 timeout: float | None = None,
240 *,
241 disconnect_on_error: bool = True,
242 push_request: bool | None = False,
243 ) -> Any:
244 pass
245
246 @abstractmethod
247 def getpeername(self) -> str | None:
248 pass
249
250 @abstractmethod
251 def extract_connection_details(self) -> str:
252 pass
253
254 def _configure_maintenance_notifications(
255 self,
256 maint_notifications_pool_handler: (
257 AsyncMaintNotificationsPoolHandler | None
258 ) = None,
259 orig_host_address: str | None = None,
260 orig_socket_timeout: float | None = None,
261 orig_socket_connect_timeout: float | None = None,
262 oss_cluster_maint_notifications_handler: (
263 AsyncOSSMaintNotificationsHandler | None
264 ) = None,
265 parser: BaseParser | None = None,
266 ) -> None:
267 if (
268 not self.maint_notifications_config
269 or not self.maint_notifications_config.enabled
270 ):
271 self._maint_notifications_pool_handler = None
272 self._maint_notifications_connection_handler = None
273 self._oss_cluster_maint_notifications_handler = None
274 return
275
276 if not parser:
277 raise RedisError(
278 "To configure maintenance notifications, a parser must be provided!"
279 )
280
281 if not isinstance(parser, _AsyncHiredisParser) and not isinstance(
282 parser, _AsyncRESP3Parser
283 ):
284 raise RedisError(
285 "Maintenance notifications are only supported with hiredis and RESP3 parsers!"
286 )
287
288 if maint_notifications_pool_handler:
289 # Extract a reference to a new pool handler that copies all properties
290 # of the original one and has a different connection reference
291 # This is needed because when we attach the handler to the parser
292 # we need to make sure that the handler has a reference to the
293 # connection that the parser is attached to.
294 self._maint_notifications_pool_handler = (
295 maint_notifications_pool_handler.get_handler_for_connection()
296 )
297 self._maint_notifications_pool_handler.set_connection(self)
298 else:
299 self._maint_notifications_pool_handler = None
300
301 self._maint_notifications_connection_handler = (
302 AsyncMaintNotificationsConnectionHandler(
303 self, self.maint_notifications_config
304 )
305 )
306
307 if oss_cluster_maint_notifications_handler:
308 self._oss_cluster_maint_notifications_handler = (
309 oss_cluster_maint_notifications_handler
310 )
311 parser.set_oss_cluster_maint_push_handler(
312 oss_cluster_maint_notifications_handler.handle_notification
313 )
314 else:
315 self._oss_cluster_maint_notifications_handler = None
316
317 # Set up pool handler to parser if available
318 if self._maint_notifications_pool_handler:
319 parser.set_node_moving_push_handler(
320 self._maint_notifications_pool_handler.handle_notification
321 )
322
323 # Set up connection handler
324 parser.set_maintenance_push_handler(
325 self._maint_notifications_connection_handler.handle_notification
326 )
327
328 self.orig_host_address = orig_host_address if orig_host_address else self.host
329 self.orig_socket_timeout = (
330 orig_socket_timeout if orig_socket_timeout else self.socket_timeout
331 )
332 self.orig_socket_connect_timeout = (
333 orig_socket_connect_timeout
334 if orig_socket_connect_timeout
335 else self.socket_connect_timeout
336 )
337
338 def set_maint_notifications_pool_handler_for_connection(
339 self, maint_notifications_pool_handler: AsyncMaintNotificationsPoolHandler
340 ) -> None:
341 # Deep copy the pool handler to avoid sharing the same pool handler
342 # between multiple connections, because otherwise each connection will override
343 # the connection reference and the pool handler will only hold a reference
344 # to the last connection that was set.
345 maint_notifications_pool_handler_copy = (
346 maint_notifications_pool_handler.get_handler_for_connection()
347 )
348 maint_notifications_pool_handler_copy.set_connection(self)
349 parser = self._get_push_notifications_parser()
350 parser.set_node_moving_push_handler(
351 maint_notifications_pool_handler_copy.handle_notification
352 )
353 self._maint_notifications_pool_handler = maint_notifications_pool_handler_copy
354
355 # Update maintenance notification connection handler if it doesn't exist
356 if not self._maint_notifications_connection_handler:
357 self._maint_notifications_connection_handler = (
358 AsyncMaintNotificationsConnectionHandler(
359 self, maint_notifications_pool_handler.config
360 )
361 )
362 parser.set_maintenance_push_handler(
363 self._maint_notifications_connection_handler.handle_notification
364 )
365 else:
366 self._maint_notifications_connection_handler.config = (
367 maint_notifications_pool_handler.config
368 )
369
370 def set_maint_notifications_cluster_handler_for_connection(
371 self,
372 oss_cluster_maint_notifications_handler: AsyncOSSMaintNotificationsHandler,
373 ) -> None:
374 parser = self._get_push_notifications_parser()
375 parser.set_oss_cluster_maint_push_handler(
376 oss_cluster_maint_notifications_handler.handle_notification
377 )
378 # OSS cluster mode and pool-handler mode are mutually exclusive. Clear
379 # any node-moving/pool handler a default (RESP3 "auto") pool wired in
380 # __init__ so this existing connection is not configured with both.
381 parser.set_node_moving_push_handler(None)
382 self._maint_notifications_pool_handler = None
383
384 self._oss_cluster_maint_notifications_handler = (
385 oss_cluster_maint_notifications_handler
386 )
387
388 # Update maintenance notification connection handler if it doesn't exist
389 if not self._maint_notifications_connection_handler:
390 self._maint_notifications_connection_handler = (
391 AsyncMaintNotificationsConnectionHandler(
392 self, oss_cluster_maint_notifications_handler.config
393 )
394 )
395 parser.set_maintenance_push_handler(
396 self._maint_notifications_connection_handler.handle_notification
397 )
398 else:
399 self._maint_notifications_connection_handler.config = (
400 oss_cluster_maint_notifications_handler.config
401 )
402
403 async def activate_maint_notifications_handling_if_enabled(
404 self, check_health: bool = True
405 ) -> None:
406 # Send maintenance notifications handshake if RESP3 is active
407 # and maintenance notifications are enabled
408 # and we have a host to determine the endpoint type from
409 # When the maint_notifications_config enabled mode is "auto",
410 # we just log a warning if the handshake fails
411 # When the mode is enabled=True, we raise an exception in case of failure
412 host = getattr(self, "host", None)
413 if (
414 check_protocol_version(self.get_protocol(), 3)
415 and self.maint_notifications_config
416 and self.maint_notifications_config.enabled
417 and self._maint_notifications_connection_handler
418 and host is not None
419 ):
420 await self._enable_maintenance_notifications(
421 maint_notifications_config=self.maint_notifications_config,
422 check_health=check_health,
423 )
424
425 async def _enable_maintenance_notifications(
426 self,
427 maint_notifications_config: MaintNotificationsConfig,
428 check_health: bool = True,
429 ) -> None:
430 try:
431 host = getattr(self, "host", None)
432 if host is None:
433 raise ValueError(
434 "Cannot enable maintenance notifications for connection"
435 " object that doesn't have a host attribute."
436 )
437
438 endpoint_type = maint_notifications_config.get_endpoint_type(host, self)
439 await self.send_command(
440 "CLIENT",
441 "MAINT_NOTIFICATIONS",
442 "ON",
443 "moving-endpoint-type",
444 endpoint_type.value,
445 check_health=check_health,
446 )
447 response = await self.read_response()
448 if not response or str_if_bytes(response) != "OK":
449 raise ResponseError(
450 "The server doesn't support maintenance notifications"
451 )
452 except Exception as e:
453 if (
454 isinstance(e, ResponseError)
455 and maint_notifications_config.enabled == "auto"
456 ):
457 # Log warning but don't fail the connection
458 logger.debug(f"Failed to enable maintenance notifications: {e}")
459 else:
460 raise
461
462 def get_resolved_ip(self) -> str | None:
463 """
464 Extract the resolved IP address from an established connection or host.
465
466 First tries to get the actual peer IP from the async stream writer, then
467 falls back to DNS resolution if needed.
468
469 Returns:
470 The resolved IP address, or None if it cannot be determined.
471 """
472
473 # Method 1: Try to get the actual IP from the established stream.
474 # This is most accurate as it shows the exact IP being used.
475 try:
476 peer_addr = self.getpeername()
477 if peer_addr:
478 return peer_addr
479 except (AttributeError, OSError):
480 # Stream might not be connected or peer address lookup might fail.
481 pass
482
483 # Method 2: Fall back to the configured host (which may be an IP or an
484 # FQDN). Unlike the sync client we intentionally do NOT call
485 # socket.getaddrinfo() here: this method runs on the event loop, so a
486 # blocking DNS resolution would stall it. On the endpoint-type handshake
487 # path (get_endpoint_type) getpeername() above always succeeds because the
488 # writer was just connected, so this fallback is only reached by the
489 # debug-log call sites during reconnects — where returning the host is
490 # fine. A blocking getaddrinfo on an FQDN host there can freeze the loop
491 # for seconds and trip unrelated connect timeouts.
492 return getattr(self, "host", None)
493
494 @property
495 def maintenance_state(self) -> MaintenanceState:
496 return self._maintenance_state
497
498 @maintenance_state.setter
499 def maintenance_state(self, state: MaintenanceState) -> None:
500 self._maintenance_state = state
501
502 def add_maint_start_notification(self, id: int) -> None:
503 self._processed_start_maint_notifications.add(id)
504
505 def get_processed_start_notifications(self) -> set[int]:
506 return self._processed_start_maint_notifications
507
508 def add_skipped_end_notification(self, id: int) -> None:
509 self._skipped_end_maint_notifications.add(id)
510
511 def get_skipped_end_notifications(self) -> set[int]:
512 return self._skipped_end_maint_notifications
513
514 def reset_received_notifications(self) -> None:
515 self._processed_start_maint_notifications.clear()
516 self._skipped_end_maint_notifications.clear()
517
518 def update_current_socket_timeout(
519 self, relaxed_timeout: float | None = None
520 ) -> None:
521 timeout = relaxed_timeout if relaxed_timeout != -1 else self.socket_timeout
522 self._reschedule_active_read_timeout(timeout)
523
524 def _reschedule_active_read_timeout(self, timeout: float | None) -> None:
525 timeout_context = getattr(self, "_active_read_timeout", None)
526 if timeout_context is None:
527 # No read_response call is currently inside its socket timeout
528 # context, so there is no in-flight deadline to relax or restore.
529 return
530
531 if timeout is None:
532 # A None socket timeout means the active read should become blocking.
533 # Python 3.11's timeout context supports clearing the deadline.
534 if hasattr(timeout_context, "reschedule"):
535 timeout_context.reschedule(None)
536 # Older async-timeout contexts cannot clear a deadline, so reject the
537 # current timeout instead of leaving a stale relaxed deadline active.
538 elif hasattr(timeout_context, "reject"):
539 timeout_context.reject()
540 return
541
542 # Active read timeouts are stored as loop-time deadlines, not durations.
543 deadline = asyncio.get_running_loop().time() + timeout
544 if hasattr(timeout_context, "reschedule"):
545 # Python 3.11 asyncio.timeout exposes reschedule().
546 timeout_context.reschedule(deadline)
547 elif hasattr(timeout_context, "update"):
548 # async-timeout exposes update() for the same deadline adjustment.
549 timeout_context.update(deadline)
550
551 def set_tmp_settings(
552 self,
553 tmp_host_address: str | object | None = SENTINEL,
554 tmp_relaxed_timeout: float | None = -1,
555 ) -> None:
556 """
557 SENTINEL keeps the host unchanged. -1 keeps the relaxed timeout unchanged.
558 """
559 if tmp_host_address and tmp_host_address != SENTINEL:
560 self.host = str(tmp_host_address)
561 if tmp_relaxed_timeout != -1:
562 self.socket_timeout = tmp_relaxed_timeout
563 self.socket_connect_timeout = tmp_relaxed_timeout
564
565 def reset_tmp_settings(
566 self,
567 reset_host_address: bool = False,
568 reset_relaxed_timeout: bool = False,
569 ) -> None:
570 if reset_host_address:
571 self.host = self.orig_host_address
572 if reset_relaxed_timeout:
573 self.socket_timeout = self.orig_socket_timeout
574 self.socket_connect_timeout = self.orig_socket_connect_timeout
575
576
577class AbstractConnection(AsyncMaintNotificationsAbstractConnection):
578 """Manages communication to and from a Redis server"""
579
580 __slots__ = (
581 "db",
582 "username",
583 "client_name",
584 "lib_name",
585 "lib_version",
586 "credential_provider",
587 "password",
588 "socket_timeout",
589 "socket_connect_timeout",
590 "redis_connect_func",
591 "retry_on_timeout",
592 "retry_on_error",
593 "health_check_interval",
594 "next_health_check",
595 "last_active_at",
596 "encoder",
597 "ssl_context",
598 "protocol",
599 "_reader",
600 "_writer",
601 "_parser",
602 "_active_read_timeout",
603 "_connect_callbacks",
604 "_buffer_cutoff",
605 "_lock",
606 "_socket_read_size",
607 "__dict__",
608 )
609
610 @deprecated_args(
611 args_to_warn=["lib_name", "lib_version"],
612 reason="Use 'driver_info' parameter instead. "
613 "lib_name and lib_version will be removed in a future version.",
614 )
615 def __init__(
616 self,
617 *,
618 db: str | int = 0,
619 password: str | None = None,
620 socket_timeout: float | None = DEFAULT_SOCKET_TIMEOUT,
621 socket_connect_timeout: float | None = DEFAULT_SOCKET_CONNECT_TIMEOUT,
622 retry_on_timeout: bool = False,
623 retry_on_error: Iterable[Type[Exception]] | object = SENTINEL,
624 encoding: str = "utf-8",
625 encoding_errors: str = "strict",
626 decode_responses: bool = False,
627 parser_class: Type[BaseParser] = DefaultParser,
628 socket_read_size: int = DEFAULT_SOCKET_READ_SIZE,
629 health_check_interval: float = 0,
630 client_name: str | None = None,
631 lib_name: str | object | None = SENTINEL,
632 lib_version: str | object | None = SENTINEL,
633 driver_info: DriverInfo | object | None = SENTINEL,
634 username: str | None = None,
635 retry: Retry | None = None,
636 redis_connect_func: ConnectCallbackT | None = None,
637 encoder_class: Type[Encoder] = Encoder,
638 credential_provider: CredentialProvider | None = None,
639 protocol: int | None = None,
640 legacy_responses: bool = True,
641 event_dispatcher: EventDispatcher | None = None,
642 maint_notifications_config: MaintNotificationsConfig | None = None,
643 maint_notifications_pool_handler: (
644 AsyncMaintNotificationsPoolHandler | None
645 ) = None,
646 maintenance_state: MaintenanceState = MaintenanceState.NONE,
647 maintenance_notification_hash: int | None = None,
648 orig_host_address: str | None = None,
649 orig_socket_timeout: float | None = None,
650 orig_socket_connect_timeout: float | None = None,
651 oss_cluster_maint_notifications_handler: (
652 AsyncOSSMaintNotificationsHandler | None
653 ) = None,
654 himport_registry: HImportRegistry | None = None,
655 ):
656 """
657 Initialize a new async Connection.
658
659 Parameters
660 ----------
661 driver_info : DriverInfo, optional
662 Driver metadata for CLIENT SETINFO. If provided, lib_name and lib_version
663 are ignored. If not provided, a DriverInfo will be created from lib_name
664 and lib_version. Explicit None disables CLIENT SETINFO.
665 lib_name : str, optional
666 **Deprecated.** Use driver_info instead. Library name for CLIENT SETINFO.
667 lib_version : str, optional
668 **Deprecated.** Use driver_info instead. Library version for CLIENT SETINFO.
669 """
670 if (username or password) and credential_provider is not None:
671 raise DataError(
672 "'username' and 'password' cannot be passed along with 'credential_"
673 "provider'. Please provide only one of the following arguments: \n"
674 "1. 'password' and (optional) 'username'\n"
675 "2. 'credential_provider'"
676 )
677 if event_dispatcher is None:
678 self._event_dispatcher = EventDispatcher()
679 else:
680 self._event_dispatcher = event_dispatcher
681 self.db = db
682 self.client_name = client_name
683
684 # Handle driver_info: if provided, use it; otherwise create from lib_name/lib_version.
685 self.driver_info = resolve_driver_info(driver_info, lib_name, lib_version)
686
687 self.credential_provider = credential_provider
688 self.password = password
689 self.username = username
690 self.socket_timeout = socket_timeout
691 if socket_connect_timeout is None:
692 socket_connect_timeout = socket_timeout
693 self.socket_connect_timeout = socket_connect_timeout
694 self.retry_on_timeout = retry_on_timeout
695 if retry_on_error is SENTINEL:
696 retry_on_error = []
697 else:
698 # Copy so we never mutate the caller-supplied list (parity with sync).
699 retry_on_error = list(retry_on_error)
700 if retry_on_timeout:
701 retry_on_error.append(TimeoutError)
702 retry_on_error.append(socket.timeout)
703 retry_on_error.append(asyncio.TimeoutError)
704 self.retry_on_error = retry_on_error
705 if retry or retry_on_error:
706 if not retry:
707 self.retry = Retry(NoBackoff(), 1)
708 else:
709 # deep-copy the Retry object as it is mutable
710 self.retry = copy.deepcopy(retry)
711 # Update the retry's supported errors with the specified errors
712 self.retry.update_supported_errors(retry_on_error)
713 else:
714 self.retry = Retry(NoBackoff(), 0)
715 self.health_check_interval = health_check_interval
716 self.next_health_check: float = -1
717 self.encoder = encoder_class(encoding, encoding_errors, decode_responses)
718 self.redis_connect_func = redis_connect_func
719 self._reader: Optional[asyncio.StreamReader] = None
720 self._writer: Optional[asyncio.StreamWriter] = None
721 self._socket_read_size = socket_read_size
722 self._active_read_timeout = None
723 self._connect_callbacks: List[weakref.WeakMethod[ConnectCallbackT]] = []
724 self._buffer_cutoff = 6000
725 self._re_auth_token: Optional[TokenInterface] = None
726 self._should_reconnect = False
727
728 try:
729 p = int(protocol)
730 except TypeError:
731 p = DEFAULT_RESP_VERSION
732 except ValueError:
733 raise ConnectionError("protocol must be an integer")
734 else:
735 if p < 2 or p > 3:
736 raise ConnectionError("protocol must be either 2 or 3")
737 self.protocol = p
738 self.legacy_responses = legacy_responses
739 if parser_class != _AsyncHiredisParser:
740 # The Python parsers are protocol-specific; hiredis supports both.
741 if self.protocol == 3 and parser_class == _AsyncRESP2Parser:
742 parser_class = _AsyncRESP3Parser
743 elif self.protocol == 2 and parser_class == _AsyncRESP3Parser:
744 parser_class = _AsyncRESP2Parser
745 self.set_parser(parser_class)
746
747 # HIMPORT client-side state. `himport_registry` is the shared client-level
748 # registry (empty if unconfigured) and persists across reconnects.
749 self.himport_registry = himport_registry
750 self._reset_himport_state()
751
752 AsyncMaintNotificationsAbstractConnection.__init__(
753 self,
754 maint_notifications_config,
755 maint_notifications_pool_handler,
756 maintenance_state,
757 maintenance_notification_hash,
758 orig_host_address,
759 orig_socket_timeout,
760 orig_socket_connect_timeout,
761 oss_cluster_maint_notifications_handler,
762 self._parser,
763 )
764
765 def __del__(self, _warnings: Any = warnings):
766 # For some reason, the individual streams don't get properly garbage
767 # collected and therefore produce no resource warnings. We add one
768 # here, in the same style as those from the stdlib.
769 if getattr(self, "_writer", None):
770 _warnings.warn(
771 f"unclosed Connection {self!r}", ResourceWarning, source=self
772 )
773
774 try:
775 asyncio.get_running_loop()
776 self._close()
777 except RuntimeError:
778 # No actions been taken if pool already closed.
779 pass
780
781 def _close(self):
782 """
783 Internal method to silently close the connection without waiting
784 """
785 if self._writer:
786 self._writer.close()
787 self._writer = self._reader = None
788
789 def __repr__(self):
790 repr_args = ",".join((f"{k}={v}" for k, v in self.repr_pieces()))
791 return f"<{self.__class__.__module__}.{self.__class__.__name__}({repr_args})>"
792
793 @abstractmethod
794 def repr_pieces(self):
795 pass
796
797 @property
798 def is_connected(self):
799 return self._reader is not None and self._writer is not None
800
801 def register_connect_callback(self, callback):
802 """
803 Register a callback to be called when the connection is established either
804 initially or reconnected. This allows listeners to issue commands that
805 are ephemeral to the connection, for example pub/sub subscription or
806 key tracking. The callback must be a _method_ and will be kept as
807 a weak reference.
808 """
809 wm = weakref.WeakMethod(callback)
810 if wm not in self._connect_callbacks:
811 self._connect_callbacks.append(wm)
812
813 def deregister_connect_callback(self, callback):
814 """
815 De-register a previously registered callback. It will no-longer receive
816 notifications on connection events. Calling this is not required when the
817 listener goes away, since the callbacks are kept as weak methods.
818 """
819 try:
820 self._connect_callbacks.remove(weakref.WeakMethod(callback))
821 except ValueError:
822 pass
823
824 def set_parser(self, parser_class: Type[BaseParser]) -> None:
825 """
826 Creates a new instance of parser_class with socket size:
827 _socket_read_size and assigns it to the parser for the connection
828 :param parser_class: The required parser class
829 """
830 self._parser = parser_class(socket_read_size=self._socket_read_size)
831
832 def _get_parser(self) -> BaseParser:
833 return self._parser
834
835 def getpeername(self) -> str | None:
836 """
837 Returns the peer name of the connection.
838 """
839 writer = self._writer
840 if writer is None:
841 return None
842 peername = writer.get_extra_info("peername")
843 if isinstance(peername, tuple) and peername:
844 return str(peername[0])
845 return None
846
847 def extract_connection_details(self) -> str:
848 """
849 Render the connection's identity, maintenance state and effective timeouts.
850
851 This is what the debug logs use to explain a failed or timed out command:
852
853 - ``host`` vs ``orig host`` says whether the connection still points at the
854 node being moved away from, or has already been repointed at the new one.
855 - ``state`` says whether maintenance handling touched this connection at
856 all, so an unaffected node's connections are distinguishable.
857 - ``socket_timeout`` vs ``active read timeout`` says which timeout the read
858 actually ran under. ``active read timeout`` is the remaining deadline of
859 the in-flight ``read_response`` (``None`` when no read is in flight), so
860 the two diverge for a command that was already reading when the relaxed
861 timeout was applied.
862 """
863 writer = self._writer
864 if writer is None:
865 return "not connected"
866
867 socket_address = None
868 try:
869 socket_name = writer.get_extra_info("sockname")
870 # AF_UNIX sockets report a path string rather than a (host, port) tuple
871 if isinstance(socket_name, tuple) and len(socket_name) > 1:
872 socket_address = socket_name[1]
873 except (AttributeError, OSError):
874 pass
875
876 # Unlike the sync client there is no timeout armed on the socket; the
877 # deadline lives in the timeout context wrapping the in-flight read.
878 active_read_timeout = None
879 timeout_context = self._active_read_timeout
880 if timeout_context is not None:
881 try:
882 when = timeout_context.when()
883 if when is not None:
884 active_read_timeout = round(
885 when - asyncio.get_running_loop().time(), 3
886 )
887 except (AttributeError, RuntimeError):
888 pass
889
890 state = getattr(self.maintenance_state, "value", self.maintenance_state)
891 return (
892 f"connected to ip {self.get_resolved_ip()}, "
893 f"local socket port: {socket_address}, "
894 f"host: {self._host_error()} "
895 f"(orig: {getattr(self, 'orig_host_address', None)}), "
896 f"state: {state}, "
897 f"socket_timeout: {self.socket_timeout} "
898 f"(orig: {getattr(self, 'orig_socket_timeout', None)}), "
899 f"active read timeout: {active_read_timeout}, "
900 f"should_reconnect: {self.should_reconnect()}, "
901 f"notification_hash: {self.maintenance_notification_hash}"
902 )
903
904 async def connect(self):
905 """Connects to the Redis server if not already connected"""
906 # try once the socket connect with the handshake, retry the whole
907 # connect/handshake flow based on retry policy
908 await self.retry.call_with_retry(
909 lambda: self.connect_check_health(
910 check_health=True, retry_socket_connect=False
911 ),
912 lambda error, failure_count: self.disconnect(
913 error=error, failure_count=failure_count
914 ),
915 with_failure_count=True,
916 )
917
918 async def connect_check_health(
919 self, check_health: bool = True, retry_socket_connect: bool = True
920 ):
921 if self.is_connected:
922 return
923 # Track actual retry attempts for error reporting
924 actual_retry_attempts = 0
925
926 def failure_callback(error, failure_count):
927 nonlocal actual_retry_attempts
928 actual_retry_attempts = failure_count
929 return self.disconnect(error=error, failure_count=failure_count)
930
931 try:
932 if retry_socket_connect:
933 await self.retry.call_with_retry(
934 lambda: self._connect(),
935 failure_callback,
936 with_failure_count=True,
937 )
938 else:
939 await self._connect()
940 except asyncio.CancelledError:
941 raise # in 3.7 and earlier, this is an Exception, not BaseException
942 except (socket.timeout, asyncio.TimeoutError):
943 e = TimeoutError("Timeout connecting to server")
944 await record_error_count(
945 server_address=getattr(self, "host", None),
946 server_port=getattr(self, "port", None),
947 network_peer_address=getattr(self, "host", None),
948 network_peer_port=getattr(self, "port", None),
949 error_type=e,
950 retry_attempts=actual_retry_attempts,
951 is_internal=False,
952 )
953 raise e
954 except OSError as e:
955 e = ConnectionError(self._error_message(e))
956 await record_error_count(
957 server_address=getattr(self, "host", None),
958 server_port=getattr(self, "port", None),
959 network_peer_address=getattr(self, "host", None),
960 network_peer_port=getattr(self, "port", None),
961 error_type=e,
962 retry_attempts=actual_retry_attempts,
963 is_internal=False,
964 )
965 raise e
966 except Exception as exc:
967 raise ConnectionError(exc) from exc
968
969 try:
970 if not self.redis_connect_func:
971 # Use the default on_connect function
972 await self.on_connect_check_health(check_health=check_health)
973 else:
974 # Use the passed function redis_connect_func
975 (
976 await self.redis_connect_func(self)
977 if asyncio.iscoroutinefunction(self.redis_connect_func)
978 else self.redis_connect_func(self)
979 )
980 except RedisError:
981 # clean up after any error in on_connect
982 await self.disconnect()
983 raise
984
985 # run any user callbacks. right now the only internal callback
986 # is for pubsub channel/pattern resubscription
987 # first, remove any dead weakrefs
988 self._connect_callbacks = [ref for ref in self._connect_callbacks if ref()]
989 for ref in self._connect_callbacks:
990 callback = ref()
991 task = callback(self)
992 if task and inspect.isawaitable(task):
993 await task
994
995 def mark_for_reconnect(self):
996 self._should_reconnect = True
997
998 def should_reconnect(self):
999 return self._should_reconnect
1000
1001 def reset_should_reconnect(self):
1002 self._should_reconnect = False
1003
1004 @abstractmethod
1005 async def _connect(self):
1006 pass
1007
1008 @abstractmethod
1009 def _host_error(self) -> str:
1010 pass
1011
1012 def _error_message(self, exception: BaseException) -> str:
1013 return format_error_message(self._host_error(), exception)
1014
1015 def get_protocol(self):
1016 return self.protocol
1017
1018 def _reset_himport_state(self) -> None:
1019 # A fresh server session has no prepared HIMPORT fieldsets, so the next
1020 # himport_set must re-prepare on this connection. ``_himport_prepared`` maps
1021 # fieldset name -> the version prepared on the server; ``_himport_reconciled
1022 # _revision`` is the registry revision this connection last reconciled discards
1023 # against. Both are reset on connect/disconnect since the session is gone.
1024 self._himport_prepared: dict[str, int] = {}
1025 self._himport_reconciled_revision: int = 0
1026
1027 async def on_connect(self) -> None:
1028 """Initialize the connection, authenticate and select a database"""
1029 await self.on_connect_check_health(check_health=True)
1030
1031 async def on_connect_check_health(self, check_health: bool = True) -> None:
1032 # A fresh socket is a new server session: no prepared HIMPORT fieldsets.
1033 self._reset_himport_state()
1034 self._parser.on_connect(self)
1035 parser = self._parser
1036
1037 auth_args = None
1038 # if credential provider or username and/or password are set, authenticate
1039 if self.credential_provider or (self.username or self.password):
1040 cred_provider = (
1041 self.credential_provider
1042 or UsernamePasswordCredentialProvider(self.username, self.password)
1043 )
1044 auth_args = await cred_provider.get_credentials_async()
1045
1046 # if resp version is specified and we have auth args,
1047 # we need to send them via HELLO
1048 if auth_args and check_protocol_version(self.protocol, 3):
1049 if isinstance(self._parser, _AsyncRESP2Parser):
1050 self.set_parser(_AsyncRESP3Parser)
1051 # update cluster exception classes
1052 self._parser.EXCEPTION_CLASSES = parser.EXCEPTION_CLASSES
1053 self._parser.on_connect(self)
1054 if len(auth_args) == 1:
1055 auth_args = ["default", auth_args[0]]
1056 # avoid checking health here -- PING will fail if we try
1057 # to check the health prior to the AUTH
1058 await self.send_command(
1059 "HELLO", self.protocol, "AUTH", *auth_args, check_health=False
1060 )
1061 response = await self.read_response()
1062 if response.get(b"proto") != int(self.protocol) and response.get(
1063 "proto"
1064 ) != int(self.protocol):
1065 raise ConnectionError("Invalid RESP version")
1066 # avoid checking health here -- PING will fail if we try
1067 # to check the health prior to the AUTH
1068 elif auth_args:
1069 await self.send_command("AUTH", *auth_args, check_health=False)
1070
1071 try:
1072 auth_response = await self.read_response()
1073 except AuthenticationWrongNumberOfArgsError:
1074 # a username and password were specified but the Redis
1075 # server seems to be < 6.0.0 which expects a single password
1076 # arg. retry auth with just the password.
1077 # https://github.com/andymccurdy/redis-py/issues/1274
1078 await self.send_command("AUTH", auth_args[-1], check_health=False)
1079 auth_response = await self.read_response()
1080
1081 if str_if_bytes(auth_response) != "OK":
1082 raise AuthenticationError("Invalid Username or Password")
1083
1084 # if resp version is specified, switch to it
1085 elif check_protocol_version(self.protocol, 3):
1086 if isinstance(self._parser, _AsyncRESP2Parser):
1087 self.set_parser(_AsyncRESP3Parser)
1088 # update cluster exception classes
1089 self._parser.EXCEPTION_CLASSES = parser.EXCEPTION_CLASSES
1090 self._parser.on_connect(self)
1091 await self.send_command("HELLO", self.protocol, check_health=check_health)
1092 response = await self.read_response()
1093 # if response.get(b"proto") != self.protocol and response.get(
1094 # "proto"
1095 # ) != self.protocol:
1096 # raise ConnectionError("Invalid RESP version")
1097
1098 # Activate maintenance notifications for this connection
1099 # if enabled in the configuration
1100 # This is a no-op if maintenance notifications are not enabled
1101 await self.activate_maint_notifications_handling_if_enabled(
1102 check_health=check_health
1103 )
1104
1105 # if a client_name is given, set it
1106 if self.client_name:
1107 await self.send_command(
1108 "CLIENT",
1109 "SETNAME",
1110 self.client_name,
1111 check_health=check_health,
1112 )
1113 if str_if_bytes(await self.read_response()) != "OK":
1114 raise ConnectionError("Error setting client name")
1115
1116 # Set the library name and version from driver_info, pipeline for lower startup latency
1117 lib_name_sent = False
1118 lib_version_sent = False
1119
1120 if self.driver_info and self.driver_info.formatted_name:
1121 await self.send_command(
1122 "CLIENT",
1123 "SETINFO",
1124 "LIB-NAME",
1125 self.driver_info.formatted_name,
1126 check_health=check_health,
1127 )
1128 lib_name_sent = True
1129
1130 if self.driver_info and self.driver_info.lib_version:
1131 await self.send_command(
1132 "CLIENT",
1133 "SETINFO",
1134 "LIB-VER",
1135 self.driver_info.lib_version,
1136 check_health=check_health,
1137 )
1138 lib_version_sent = True
1139
1140 # if a database is specified, switch to it. Also pipeline this
1141 if self.db:
1142 await self.send_command("SELECT", self.db, check_health=check_health)
1143
1144 # read responses from pipeline
1145 for _ in range(sum([lib_name_sent, lib_version_sent])):
1146 try:
1147 await self.read_response()
1148 except ResponseError:
1149 pass
1150
1151 if self.db:
1152 if str_if_bytes(await self.read_response()) != "OK":
1153 raise ConnectionError("Invalid Database")
1154
1155 async def disconnect(
1156 self,
1157 nowait: bool = False,
1158 error: Optional[Exception] = None,
1159 failure_count: Optional[int] = None,
1160 health_check_failed: bool = False,
1161 ) -> None:
1162 """Disconnects from the Redis server"""
1163 # The server session is gone, so any HIMPORT fieldsets prepared on this
1164 # socket no longer exist; reset the tracking.
1165 self._reset_himport_state()
1166 # On Python 3.13+, asyncio.timeout() raises RuntimeError when called
1167 # outside a running Task (e.g. during GC finalization or event-loop
1168 # callbacks). In that context we fall back to a synchronous close.
1169 # See https://github.com/redis/redis-py/issues/3856
1170 if asyncio.current_task() is None:
1171 self._parser.on_disconnect()
1172 self.reset_should_reconnect()
1173 self._close()
1174 return
1175
1176 try:
1177 async with async_timeout(self.socket_connect_timeout):
1178 self._parser.on_disconnect()
1179 # Reset the reconnect flag
1180 self.reset_should_reconnect()
1181 if not self.is_connected:
1182 return
1183 try:
1184 self._writer.close() # type: ignore[union-attr]
1185 # wait for close to finish, except when handling errors and
1186 # forcefully disconnecting.
1187 if not nowait:
1188 await self._writer.wait_closed() # type: ignore[union-attr]
1189 except OSError:
1190 pass
1191 finally:
1192 self._reader = None
1193 self._writer = None
1194 except asyncio.TimeoutError:
1195 raise TimeoutError(
1196 f"Timed out closing connection after {self.socket_connect_timeout}"
1197 ) from None
1198
1199 if error:
1200 if health_check_failed:
1201 close_reason = CloseReason.HEALTHCHECK_FAILED
1202 else:
1203 close_reason = CloseReason.ERROR
1204
1205 if failure_count is not None and failure_count > self.retry.get_retries():
1206 await record_error_count(
1207 server_address=getattr(self, "host", None),
1208 server_port=getattr(self, "port", None),
1209 network_peer_address=getattr(self, "host", None),
1210 network_peer_port=getattr(self, "port", None),
1211 error_type=error,
1212 retry_attempts=failure_count,
1213 )
1214
1215 await record_connection_closed(
1216 close_reason=close_reason,
1217 error_type=error,
1218 )
1219 else:
1220 await record_connection_closed(
1221 close_reason=CloseReason.APPLICATION_CLOSE,
1222 )
1223
1224 if self.maintenance_state == MaintenanceState.MAINTENANCE:
1225 # MOVING state is owned by the pool-level TTL cleanup. Regular
1226 # maintenance timeout relaxation can be restored when this
1227 # connection closes, matching the sync lifecycle.
1228 self.reset_tmp_settings(reset_relaxed_timeout=True)
1229 self.maintenance_state = MaintenanceState.NONE
1230 # reset the sets that keep track of received start maint
1231 # notifications and skipped end maint notifications
1232 self.reset_received_notifications()
1233
1234 async def _send_ping(self):
1235 """Send PING, expect PONG in return"""
1236 await self.send_command("PING", check_health=False)
1237 if str_if_bytes(await self.read_response()) != "PONG":
1238 raise ConnectionError("Bad response from PING health check")
1239
1240 async def _ping_failed(self, error, failure_count):
1241 """Function to call when PING fails"""
1242 await self.disconnect(
1243 error=error, failure_count=failure_count, health_check_failed=True
1244 )
1245
1246 async def check_health(self):
1247 """Check the health of the connection with a PING/PONG"""
1248 if (
1249 self.health_check_interval
1250 and asyncio.get_running_loop().time() > self.next_health_check
1251 ):
1252 await self.retry.call_with_retry(
1253 self._send_ping, self._ping_failed, with_failure_count=True
1254 )
1255
1256 async def _send_packed_command(self, command: Iterable[bytes]) -> None:
1257 writer = self._writer
1258 if writer is None or writer.transport.is_closing():
1259 raise ConnectionError("Connection closed by the server before write")
1260 try:
1261 writer.writelines(command)
1262 await writer.drain()
1263 except (TypeError, AttributeError) as e:
1264 # CPython gh-136234 adds the missing connection-lost check in 3.13.10+
1265 # and 3.14.1+. Python 3.12, 3.13.0-3.13.9, and 3.14.0 can instead
1266 # leak TypeError or AttributeError from the transport (#4287).
1267 if writer.transport.is_closing():
1268 raise ConnectionError(
1269 "Connection closed by the server while writing"
1270 ) from e
1271 raise
1272
1273 async def send_packed_command(
1274 self, command: Union[bytes, str, Iterable[bytes]], check_health: bool = True
1275 ) -> None:
1276 if not self.is_connected:
1277 await self.connect_check_health(check_health=False)
1278 if check_health:
1279 await self.check_health()
1280
1281 try:
1282 if isinstance(command, str):
1283 command = command.encode()
1284 if isinstance(command, bytes):
1285 command = [command]
1286 if self.socket_timeout:
1287 await asyncio.wait_for(
1288 self._send_packed_command(command), self.socket_timeout
1289 )
1290 else:
1291 await self._send_packed_command(command)
1292 except asyncio.TimeoutError as e:
1293 add_debug_log_for_connection_failure(self, e, "writing command")
1294 await self.disconnect(nowait=True)
1295 raise TimeoutError("Timeout writing to socket") from None
1296 except OSError as e:
1297 add_debug_log_for_connection_failure(self, e, "writing command")
1298 await self.disconnect(nowait=True)
1299 if len(e.args) == 1:
1300 err_no, errmsg = "UNKNOWN", e.args[0]
1301 else:
1302 err_no = e.args[0]
1303 errmsg = e.args[1]
1304 raise ConnectionError(
1305 f"Error {err_no} while writing to socket. {errmsg}."
1306 ) from e
1307 except BaseException as e:
1308 # BaseExceptions can be raised when a socket send operation is not
1309 # finished, e.g. due to a timeout. Ideally, a caller could then re-try
1310 # to send un-sent data. However, the send_packed_command() API
1311 # does not support it so there is no point in keeping the connection open.
1312 add_debug_log_for_connection_failure(self, e, "writing command")
1313 await self.disconnect(nowait=True)
1314 raise
1315
1316 async def send_command(self, *args: Any, **kwargs: Any) -> None:
1317 """Pack and send a command to the Redis server"""
1318 await self.send_packed_command(
1319 self.pack_command(*args), check_health=kwargs.get("check_health", True)
1320 )
1321
1322 @deprecated_function(
1323 version="8.0.0", reason="Use can_read() instead", name="can_read_destructive"
1324 )
1325 async def can_read_destructive(self) -> bool:
1326 """Check the socket to see if there's data loaded in the buffer."""
1327 try:
1328 return await self._parser.can_read()
1329 except OSError as e:
1330 await self.disconnect(nowait=True)
1331 host_error = self._host_error()
1332 raise ConnectionError(f"Error while reading from {host_error}: {e.args}")
1333
1334 async def can_read(self) -> bool:
1335 """Check the socket to see if there's data loaded in the buffer."""
1336 # TODO: Rename this API; it detects pending data or dirty/closed
1337 # connection state, not only whether application data can be read.
1338 try:
1339 return await self._parser.can_read()
1340 except OSError as e:
1341 await self.disconnect(nowait=True)
1342 host_error = self._host_error()
1343 raise ConnectionError(f"Error while reading from {host_error}: {e.args}")
1344
1345 async def read_response(
1346 self,
1347 disable_decoding: bool = False,
1348 timeout: float | None = None,
1349 *,
1350 disconnect_on_error: bool = True,
1351 push_request: bool | None = False,
1352 ):
1353 """Read the response from a previously sent command.
1354
1355 ``timeout`` semantics:
1356 - ``None`` (default): fall back to ``self.socket_timeout``.
1357 - ``math.inf``: block indefinitely with no timeout. Used by PubSub
1358 blocking reads (``listen()`` / ``get_message(timeout=None)`` /
1359 ``parse_response(block=True)``) where the configured
1360 ``socket_timeout`` must not abort the read.
1361 - ``float``: apply that timeout in seconds for this single read.
1362
1363 TODO(next-major): replace the ``math.inf`` opt-in with a SENTINEL
1364 default for ``timeout``. After that change, ``timeout=None`` will
1365 mean "no timeout, block until a response arrives" (matching the
1366 long-standing PubSub docstring contract) and the SENTINEL default
1367 will be the value that falls back to ``self.socket_timeout``.
1368 That swap is a breaking change, so it must wait for a major
1369 release. Until then, callers that need an indefinitely blocking
1370 read pass ``math.inf`` explicitly.
1371 """
1372 # TODO(next-major): drop the math.inf branch. Use SENTINEL as the
1373 # default for ``timeout`` and treat ``timeout is None`` as the
1374 # "no timeout" signal (matching the PubSub docstring contract).
1375 # Match only positive infinity here. ``-math.inf`` is not a valid
1376 # "block forever" signal and historically behaved as an already-
1377 # expired timeout; preserve that.
1378 if timeout == math.inf:
1379 read_timeout = None
1380 else:
1381 read_timeout = timeout if timeout is not None else self.socket_timeout
1382 host_error = self._host_error()
1383 try:
1384 if read_timeout is not None:
1385 timeout_context = async_timeout(read_timeout)
1386 if timeout is None:
1387 async with timeout_context as active_timeout:
1388 self._active_read_timeout = active_timeout
1389 try:
1390 response = await self._read_response_from_parser(
1391 disable_decoding=disable_decoding,
1392 push_request=push_request,
1393 )
1394 finally:
1395 self._active_read_timeout = None
1396 else:
1397 async with timeout_context:
1398 response = await self._read_response_from_parser(
1399 disable_decoding=disable_decoding,
1400 push_request=push_request,
1401 )
1402 else:
1403 response = await self._read_response_from_parser(
1404 disable_decoding=disable_decoding,
1405 push_request=push_request,
1406 )
1407 except asyncio.TimeoutError as e:
1408 if timeout is not None:
1409 # user requested timeout, return None. Operation can be retried
1410 return None
1411 # it was a self.socket_timeout error.
1412 if disconnect_on_error:
1413 add_debug_log_for_connection_failure(self, e, "reading response")
1414 await self.disconnect(nowait=True)
1415 raise TimeoutError(f"Timeout reading from {host_error}")
1416 except OSError as e:
1417 if disconnect_on_error:
1418 add_debug_log_for_connection_failure(self, e, "reading response")
1419 await self.disconnect(nowait=True)
1420 raise ConnectionError(f"Error while reading from {host_error} : {e.args}")
1421 except BaseException as e:
1422 # Also by default close in case of BaseException. A lot of code
1423 # relies on this behaviour when doing Command/Response pairs.
1424 # See #1128.
1425 if disconnect_on_error:
1426 add_debug_log_for_connection_failure(self, e, "reading response")
1427 await self.disconnect(nowait=True)
1428 raise
1429
1430 if self.health_check_interval:
1431 next_time = asyncio.get_running_loop().time() + self.health_check_interval
1432 self.next_health_check = next_time
1433
1434 if isinstance(response, ResponseError):
1435 raise response from None
1436 return response
1437
1438 async def _read_response_from_parser(
1439 self, disable_decoding: bool = False, push_request: bool | None = False
1440 ):
1441 if check_protocol_version(self.protocol, 3):
1442 return await self._parser.read_response(
1443 disable_decoding=disable_decoding, push_request=push_request
1444 )
1445 return await self._parser.read_response(disable_decoding=disable_decoding)
1446
1447 def pack_command(self, *args: EncodableT) -> List[bytes]:
1448 """Pack a series of arguments into the Redis protocol"""
1449 output = []
1450 # the client might have included 1 or more literal arguments in
1451 # the command name, e.g., 'CONFIG GET'. The Redis server expects these
1452 # arguments to be sent separately, so split the first argument
1453 # manually. These arguments should be bytestrings so that they are
1454 # not encoded.
1455 assert not isinstance(args[0], float)
1456 if isinstance(args[0], str):
1457 args = tuple(args[0].encode().split()) + args[1:]
1458 elif b" " in args[0]:
1459 args = tuple(args[0].split()) + args[1:]
1460
1461 buff = SYM_EMPTY.join((SYM_STAR, str(len(args)).encode(), SYM_CRLF))
1462
1463 buffer_cutoff = self._buffer_cutoff
1464 for arg in map(self.encoder.encode, args):
1465 # to avoid large string mallocs, chunk the command into the
1466 # output list if we're sending large values or memoryviews
1467 arg_length = len(arg)
1468 if (
1469 len(buff) > buffer_cutoff
1470 or arg_length > buffer_cutoff
1471 or isinstance(arg, memoryview)
1472 ):
1473 buff = SYM_EMPTY.join(
1474 (buff, SYM_DOLLAR, str(arg_length).encode(), SYM_CRLF)
1475 )
1476 output.append(buff)
1477 output.append(arg)
1478 buff = SYM_CRLF
1479 else:
1480 buff = SYM_EMPTY.join(
1481 (
1482 buff,
1483 SYM_DOLLAR,
1484 str(arg_length).encode(),
1485 SYM_CRLF,
1486 arg,
1487 SYM_CRLF,
1488 )
1489 )
1490 output.append(buff)
1491 return output
1492
1493 def pack_commands(self, commands: Iterable[Iterable[EncodableT]]) -> List[bytes]:
1494 """Pack multiple commands into the Redis protocol"""
1495 output: List[bytes] = []
1496 pieces: List[bytes] = []
1497 buffer_length = 0
1498 buffer_cutoff = self._buffer_cutoff
1499
1500 for cmd in commands:
1501 for chunk in self.pack_command(*cmd):
1502 chunklen = len(chunk)
1503 if (
1504 buffer_length > buffer_cutoff
1505 or chunklen > buffer_cutoff
1506 or isinstance(chunk, memoryview)
1507 ):
1508 if pieces:
1509 output.append(SYM_EMPTY.join(pieces))
1510 buffer_length = 0
1511 pieces = []
1512
1513 if chunklen > buffer_cutoff or isinstance(chunk, memoryview):
1514 output.append(chunk)
1515 else:
1516 pieces.append(chunk)
1517 buffer_length += chunklen
1518
1519 if pieces:
1520 output.append(SYM_EMPTY.join(pieces))
1521 return output
1522
1523 def _socket_is_empty(self):
1524 """Check if the socket is empty"""
1525 return len(self._reader._buffer) == 0
1526
1527 async def process_invalidation_messages(self):
1528 while not self._socket_is_empty():
1529 await self.read_response(push_request=True)
1530
1531 def set_re_auth_token(self, token: TokenInterface):
1532 self._re_auth_token = token
1533
1534 async def re_auth(self):
1535 if self._re_auth_token is not None:
1536 await self.send_command(
1537 "AUTH",
1538 self._re_auth_token.try_get("oid"),
1539 self._re_auth_token.get_value(),
1540 )
1541 await self.read_response()
1542 self._re_auth_token = None
1543
1544
1545class Connection(AbstractConnection):
1546 "Manages TCP communication to and from a Redis server"
1547
1548 def __init__(
1549 self,
1550 *,
1551 host: str = "localhost",
1552 port: str | int = 6379,
1553 socket_keepalive: bool = True,
1554 socket_keepalive_options: Mapping[int, int | bytes] | object | None = SENTINEL,
1555 socket_type: int = 0,
1556 **kwargs,
1557 ):
1558 """
1559 Initialize a TCP connection.
1560
1561 Parameters
1562 ----------
1563 socket_keepalive : bool
1564 If `True`, TCP keepalive is enabled for TCP socket connections.
1565 socket_keepalive_options : Mapping[int, int | bytes] | object | None
1566 Mapping of TCP keepalive socket option constants to values, for
1567 example `{socket.TCP_KEEPIDLE: 30}`. If left unspecified, redis-py
1568 uses TCP keepalive defaults when `socket_keepalive` is enabled:
1569 idle 30 seconds, interval 5 seconds, and 3 probes. Platform-specific
1570 options that are not available are skipped. Pass `None` or `{}` to
1571 avoid setting additional TCP keepalive options.
1572 """
1573 self.host = host
1574 # bool subclasses int; port=True would become privileged port 1.
1575 # Numeric strings stay valid. Callers still pass "6379".
1576 if isinstance(port, bool):
1577 raise TypeError("port must be an integer, not bool")
1578 if isinstance(port, str):
1579 try:
1580 port = int(port)
1581 except ValueError:
1582 raise TypeError("port must be an integer, not str") from None
1583 elif not isinstance(port, int):
1584 raise TypeError(f"port must be an integer, not {type(port).__name__}")
1585 if not 0 <= port <= 65535:
1586 raise ValueError(f"port must be in 0..65535, got {port}")
1587 self.port = port
1588 self.socket_keepalive = socket_keepalive
1589 if socket_keepalive_options is SENTINEL:
1590 socket_keepalive_options = get_default_socket_keepalive_options()
1591 self.socket_keepalive_options = socket_keepalive_options or {}
1592 self.socket_type = socket_type
1593 super().__init__(**kwargs)
1594
1595 def repr_pieces(self):
1596 pieces = [("host", self.host), ("port", self.port), ("db", self.db)]
1597 if self.client_name:
1598 pieces.append(("client_name", self.client_name))
1599 return pieces
1600
1601 def _connection_arguments(self) -> Mapping:
1602 return {"host": self.host, "port": self.port}
1603
1604 async def _connect(self):
1605 """Create a TCP socket connection"""
1606 async with async_timeout(self.socket_connect_timeout):
1607 reader, writer = await asyncio.open_connection(
1608 **self._connection_arguments()
1609 )
1610 self._reader = reader
1611 self._writer = writer
1612 sock = writer.transport.get_extra_info("socket")
1613 if sock:
1614 sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1)
1615 try:
1616 # TCP_KEEPALIVE
1617 if self.socket_keepalive:
1618 sock.setsockopt(socket.SOL_SOCKET, socket.SO_KEEPALIVE, 1)
1619 for k, v in self.socket_keepalive_options.items():
1620 sock.setsockopt(socket.SOL_TCP, k, v)
1621
1622 except (OSError, TypeError):
1623 # `socket_keepalive_options` might contain invalid options
1624 # causing an error. Do not leave the connection open.
1625 writer.close()
1626 raise
1627
1628 def _host_error(self) -> str:
1629 return f"{self.host}:{self.port}"
1630
1631
1632class SSLConnection(Connection):
1633 """Manages SSL connections to and from the Redis server(s).
1634 This class extends the Connection class, adding SSL functionality, and making
1635 use of ssl.SSLContext (https://docs.python.org/3/library/ssl.html#ssl.SSLContext)
1636 """
1637
1638 def __init__(
1639 self,
1640 ssl_keyfile: Optional[str] = None,
1641 ssl_certfile: Optional[str] = None,
1642 ssl_cert_reqs: Union[str, ssl.VerifyMode] = "required",
1643 ssl_include_verify_flags: Optional[List["ssl.VerifyFlags"]] = None,
1644 ssl_exclude_verify_flags: Optional[List["ssl.VerifyFlags"]] = None,
1645 ssl_ca_certs: Optional[str] = None,
1646 ssl_ca_data: Optional[str] = None,
1647 ssl_ca_path: Optional[str] = None,
1648 ssl_check_hostname: bool = True,
1649 ssl_min_version: Optional[TLSVersion] = None,
1650 ssl_ciphers: Optional[str] = None,
1651 ssl_password: Optional[str] = None,
1652 **kwargs,
1653 ):
1654 if not SSL_AVAILABLE:
1655 raise RedisError("Python wasn't built with SSL support")
1656
1657 self.ssl_context: RedisSSLContext = RedisSSLContext(
1658 keyfile=ssl_keyfile,
1659 certfile=ssl_certfile,
1660 cert_reqs=ssl_cert_reqs,
1661 include_verify_flags=ssl_include_verify_flags,
1662 exclude_verify_flags=ssl_exclude_verify_flags,
1663 ca_certs=ssl_ca_certs,
1664 ca_data=ssl_ca_data,
1665 ca_path=ssl_ca_path,
1666 check_hostname=ssl_check_hostname,
1667 min_version=ssl_min_version,
1668 ciphers=ssl_ciphers,
1669 password=ssl_password,
1670 )
1671 super().__init__(**kwargs)
1672
1673 def _connection_arguments(self) -> Mapping:
1674 kwargs = super()._connection_arguments()
1675 kwargs["ssl"] = self.ssl_context.get()
1676 return kwargs
1677
1678 @property
1679 def keyfile(self):
1680 return self.ssl_context.keyfile
1681
1682 @property
1683 def certfile(self):
1684 return self.ssl_context.certfile
1685
1686 @property
1687 def cert_reqs(self):
1688 return self.ssl_context.cert_reqs
1689
1690 @property
1691 def include_verify_flags(self):
1692 return self.ssl_context.include_verify_flags
1693
1694 @property
1695 def exclude_verify_flags(self):
1696 return self.ssl_context.exclude_verify_flags
1697
1698 @property
1699 def ca_certs(self):
1700 return self.ssl_context.ca_certs
1701
1702 @property
1703 def ca_data(self):
1704 return self.ssl_context.ca_data
1705
1706 @property
1707 def check_hostname(self):
1708 return self.ssl_context.check_hostname
1709
1710 @property
1711 def min_version(self):
1712 return self.ssl_context.min_version
1713
1714
1715class RedisSSLContext:
1716 __slots__ = (
1717 "keyfile",
1718 "certfile",
1719 "cert_reqs",
1720 "include_verify_flags",
1721 "exclude_verify_flags",
1722 "ca_certs",
1723 "ca_data",
1724 "ca_path",
1725 "context",
1726 "check_hostname",
1727 "min_version",
1728 "ciphers",
1729 "password",
1730 )
1731
1732 def __init__(
1733 self,
1734 keyfile: Optional[str] = None,
1735 certfile: Optional[str] = None,
1736 cert_reqs: Optional[Union[str, ssl.VerifyMode]] = None,
1737 include_verify_flags: Optional[List["ssl.VerifyFlags"]] = None,
1738 exclude_verify_flags: Optional[List["ssl.VerifyFlags"]] = None,
1739 ca_certs: Optional[str] = None,
1740 ca_data: Optional[str] = None,
1741 ca_path: Optional[str] = None,
1742 check_hostname: bool = False,
1743 min_version: Optional[TLSVersion] = None,
1744 ciphers: Optional[str] = None,
1745 password: Optional[str] = None,
1746 ):
1747 if not SSL_AVAILABLE:
1748 raise RedisError("Python wasn't built with SSL support")
1749
1750 self.keyfile = keyfile
1751 self.certfile = certfile
1752 if cert_reqs is None:
1753 cert_reqs = ssl.CERT_NONE
1754 elif isinstance(cert_reqs, str):
1755 CERT_REQS = { # noqa: N806
1756 "none": ssl.CERT_NONE,
1757 "optional": ssl.CERT_OPTIONAL,
1758 "required": ssl.CERT_REQUIRED,
1759 }
1760 if cert_reqs not in CERT_REQS:
1761 raise RedisError(
1762 f"Invalid SSL Certificate Requirements Flag: {cert_reqs}"
1763 )
1764 cert_reqs = CERT_REQS[cert_reqs]
1765 self.cert_reqs = cert_reqs
1766 self.include_verify_flags = include_verify_flags
1767 self.exclude_verify_flags = exclude_verify_flags
1768 self.ca_certs = ca_certs
1769 self.ca_data = ca_data
1770 self.ca_path = ca_path
1771 self.check_hostname = (
1772 check_hostname if self.cert_reqs != ssl.CERT_NONE else False
1773 )
1774 self.min_version = min_version
1775 self.ciphers = ciphers
1776 self.password = password
1777 self.context: Optional[SSLContext] = None
1778
1779 def get(self) -> SSLContext:
1780 if not self.context:
1781 context = ssl.create_default_context()
1782 context.check_hostname = self.check_hostname
1783 context.verify_mode = self.cert_reqs
1784 if self.include_verify_flags:
1785 for flag in self.include_verify_flags:
1786 context.verify_flags |= flag
1787 if self.exclude_verify_flags:
1788 for flag in self.exclude_verify_flags:
1789 context.verify_flags &= ~flag
1790 if self.certfile or self.keyfile:
1791 context.load_cert_chain(
1792 certfile=self.certfile,
1793 keyfile=self.keyfile,
1794 password=self.password,
1795 )
1796 if self.ca_certs or self.ca_data or self.ca_path:
1797 context.load_verify_locations(
1798 cafile=self.ca_certs, capath=self.ca_path, cadata=self.ca_data
1799 )
1800 if self.min_version is not None:
1801 context.minimum_version = self.min_version
1802 if self.ciphers is not None:
1803 context.set_ciphers(self.ciphers)
1804 self.context = context
1805 return self.context
1806
1807
1808class UnixDomainSocketConnection(AbstractConnection):
1809 "Manages UDS communication to and from a Redis server"
1810
1811 def __init__(self, *, path: str = "", **kwargs):
1812 self.path = path
1813 super().__init__(**kwargs)
1814
1815 def repr_pieces(self) -> Iterable[Tuple[str, Union[str, int]]]:
1816 pieces = [("path", self.path), ("db", self.db)]
1817 if self.client_name:
1818 pieces.append(("client_name", self.client_name))
1819 return pieces
1820
1821 async def _connect(self):
1822 async with async_timeout(self.socket_connect_timeout):
1823 reader, writer = await asyncio.open_unix_connection(path=self.path)
1824 self._reader = reader
1825 self._writer = writer
1826 await self.on_connect()
1827
1828 def _host_error(self) -> str:
1829 return self.path
1830
1831
1832FALSE_STRINGS = ("0", "F", "FALSE", "N", "NO")
1833
1834
1835def to_bool(value) -> Optional[bool]:
1836 if value is None or value == "":
1837 return None
1838 if isinstance(value, str) and value.upper() in FALSE_STRINGS:
1839 return False
1840 return bool(value)
1841
1842
1843def parse_ssl_verify_flags(value):
1844 # flags are passed in as a string representation of a list,
1845 # e.g. VERIFY_X509_STRICT, VERIFY_X509_PARTIAL_CHAIN
1846 verify_flags_str = value.replace("[", "").replace("]", "")
1847
1848 verify_flags = []
1849 for flag in verify_flags_str.split(","):
1850 flag = flag.strip()
1851 if not hasattr(VerifyFlags, flag):
1852 raise ValueError(f"Invalid ssl verify flag: {flag}")
1853 verify_flags.append(getattr(VerifyFlags, flag))
1854 return verify_flags
1855
1856
1857def parse_retry_on_error(value):
1858 # exception class names are passed as a comma-separated list,
1859 # e.g. ConnectionError,TimeoutError
1860 retry_on_error = []
1861 for name in value.replace("[", "").replace("]", "").split(","):
1862 name = name.strip()
1863 if not name:
1864 raise ValueError("Empty retry_on_error entry")
1865 exc = getattr(redis_exceptions, name, None)
1866 if not (isinstance(exc, type) and issubclass(exc, Exception)):
1867 raise ValueError(f"Unknown redis exception {name!r}")
1868 retry_on_error.append(exc)
1869 return retry_on_error
1870
1871
1872URL_QUERY_ARGUMENT_PARSERS: Mapping[str, Callable[..., object]] = MappingProxyType(
1873 {
1874 "db": int,
1875 "socket_timeout": float,
1876 "socket_connect_timeout": float,
1877 "socket_read_size": int,
1878 "socket_keepalive": to_bool,
1879 "retry_on_timeout": to_bool,
1880 "retry_on_error": parse_retry_on_error,
1881 "max_connections": int,
1882 "health_check_interval": int,
1883 "ssl_check_hostname": to_bool,
1884 "ssl_include_verify_flags": parse_ssl_verify_flags,
1885 "ssl_exclude_verify_flags": parse_ssl_verify_flags,
1886 "ssl_min_version": int,
1887 "timeout": float,
1888 "protocol": int,
1889 "legacy_responses": to_bool,
1890 }
1891)
1892
1893
1894class ConnectKwargs(TypedDict, total=False):
1895 username: str
1896 password: str
1897 connection_class: Type[AbstractConnection]
1898 host: str
1899 port: int
1900 db: int
1901 path: str
1902
1903
1904def parse_url(url: str) -> ConnectKwargs:
1905 # Scheme names are case-insensitive (RFC 3986), so normalize before the
1906 # prefix check; the "://" is required so a URL like "redis:foo" (which
1907 # urlparse would still report as the "redis" scheme) is rejected.
1908 if not url.lower().startswith(("redis://", "rediss://", "unix://")):
1909 raise ValueError(
1910 "Redis URL must specify one of the following schemes "
1911 "(redis://, rediss://, unix://)"
1912 )
1913
1914 parsed: ParseResult = urlparse(url)
1915 kwargs: ConnectKwargs = {}
1916
1917 for name, value_list in parse_qs(parsed.query).items():
1918 if value_list and len(value_list) > 0:
1919 # parse_qs() already percent-decodes query values, so use the value
1920 # as-is; unquoting again here would double-decode (e.g. "%2520" ->
1921 # "%20" -> " "). See issue #4208.
1922 value = value_list[0]
1923 parser = URL_QUERY_ARGUMENT_PARSERS.get(name)
1924 if parser:
1925 try:
1926 kwargs[name] = parser(value)
1927 except (TypeError, ValueError):
1928 raise ValueError(f"Invalid value for '{name}' in connection URL.")
1929 else:
1930 kwargs[name] = value
1931
1932 if parsed.username:
1933 kwargs["username"] = unquote(parsed.username)
1934 if parsed.password:
1935 kwargs["password"] = unquote(parsed.password)
1936
1937 # We only support redis://, rediss:// and unix:// schemes.
1938 if parsed.scheme == "unix":
1939 if parsed.path:
1940 kwargs["path"] = unquote(parsed.path)
1941 kwargs["connection_class"] = UnixDomainSocketConnection
1942
1943 else: # implied: parsed.scheme in ("redis", "rediss")
1944 if parsed.hostname:
1945 kwargs["host"] = unquote(parsed.hostname)
1946 if parsed.port is not None:
1947 kwargs["port"] = int(parsed.port)
1948
1949 # If there's a path argument, use it as the db argument if a
1950 # querystring value wasn't specified
1951 if parsed.path and "db" not in kwargs:
1952 try:
1953 kwargs["db"] = int(unquote(parsed.path).replace("/", ""))
1954 except (AttributeError, ValueError):
1955 pass
1956
1957 if parsed.scheme == "rediss":
1958 kwargs["connection_class"] = SSLConnection
1959
1960 return kwargs
1961
1962
1963_CP = TypeVar("_CP", bound="ConnectionPool")
1964
1965
1966class ConnectionPoolInterface(ABC):
1967 @abstractmethod
1968 def get_protocol(self):
1969 pass
1970
1971 @abstractmethod
1972 def reset(self) -> None:
1973 pass
1974
1975 @abstractmethod
1976 @deprecated_args(
1977 args_to_warn=["*"],
1978 reason="Use get_connection() without args instead",
1979 version="5.3.0",
1980 )
1981 async def get_connection(
1982 self, command_name: Optional[str] = None, *keys: Any, **options: Any
1983 ) -> "AbstractConnection":
1984 pass
1985
1986 @abstractmethod
1987 def get_encoder(self) -> "Encoder":
1988 pass
1989
1990 @abstractmethod
1991 async def release(self, connection: "AbstractConnection") -> None:
1992 pass
1993
1994 @abstractmethod
1995 async def disconnect(self, inuse_connections: bool = True) -> None:
1996 pass
1997
1998 @abstractmethod
1999 async def aclose(self) -> None:
2000 pass
2001
2002 @abstractmethod
2003 def set_retry(self, retry: "Retry") -> None:
2004 pass
2005
2006 @abstractmethod
2007 async def re_auth_callback(self, token: TokenInterface) -> None:
2008 pass
2009
2010 @abstractmethod
2011 def get_connection_count(self) -> List[Tuple[int, dict]]:
2012 """
2013 Returns a connection count (both idle and in use).
2014 """
2015 pass
2016
2017
2018class AsyncMaintNotificationsAbstractConnectionPool:
2019 """
2020 Internal mixin for async maintenance notification pool wiring.
2021
2022 The handler owns notification policy.
2023 This mixin owns pool state mutation because `_available_connections`,
2024 `_in_use_connections`, `connection_kwargs`, and the non-reentrant `asyncio.Lock`
2025 all live on the pool.
2026 """
2027
2028 def __init__(
2029 self,
2030 maint_notifications_config: MaintNotificationsConfig | None = None,
2031 oss_cluster_maint_notifications_handler: (
2032 "AsyncOSSMaintNotificationsHandler | None"
2033 ) = None,
2034 **kwargs: Any,
2035 ) -> None:
2036 protocol = kwargs.get("protocol")
2037 is_protocol_supported = check_protocol_version(protocol, 3)
2038 is_connection_supported = self._maintenance_notifications_supported()
2039
2040 if (
2041 maint_notifications_config is None
2042 and is_protocol_supported
2043 and is_connection_supported
2044 ):
2045 maint_notifications_config = MaintNotificationsConfig()
2046
2047 if maint_notifications_config and maint_notifications_config.enabled:
2048 if not is_connection_supported:
2049 if maint_notifications_config.enabled is True:
2050 # Unix sockets do not have a host endpoint for CLIENT
2051 # MAINT_NOTIFICATIONS to describe.
2052 if "path" in self.connection_kwargs:
2053 raise RedisError(
2054 "Maintenance notifications are not supported for "
2055 "Unix domain socket connections"
2056 )
2057
2058 # Custom connection classes must inherit the async maintenance
2059 # mixin so handlers can update connection state safely.
2060 if not self._maintenance_notifications_connection_class_supported():
2061 connection_class = getattr(self, "connection_class", None)
2062 connection_class_name = getattr(
2063 connection_class, "__name__", connection_class
2064 )
2065 raise RedisError(
2066 "Maintenance notifications are not supported for "
2067 f"connection class {connection_class_name}"
2068 )
2069
2070 # TCP-like connections still need a host to identify the
2071 # endpoint that can move during maintenance.
2072 raise RedisError(
2073 "Maintenance notifications are not supported for connections "
2074 "without a host"
2075 )
2076 self._maint_notifications_pool_handler = None
2077 self._oss_cluster_maint_notifications_handler = None
2078 return
2079
2080 if not is_protocol_supported:
2081 raise RedisError(
2082 "Maintenance notifications handlers on connection are only supported with RESP version 3"
2083 )
2084
2085 if oss_cluster_maint_notifications_handler:
2086 self._oss_cluster_maint_notifications_handler = (
2087 oss_cluster_maint_notifications_handler
2088 )
2089 self._update_connection_kwargs_for_maint_notifications(
2090 oss_cluster_maint_notifications_handler=self._oss_cluster_maint_notifications_handler
2091 )
2092 self._maint_notifications_pool_handler = None
2093 else:
2094 self._oss_cluster_maint_notifications_handler = None
2095 self._maint_notifications_pool_handler = (
2096 AsyncMaintNotificationsPoolHandler(self, maint_notifications_config)
2097 )
2098 self._update_connection_kwargs_for_maint_notifications(
2099 maint_notifications_pool_handler=self._maint_notifications_pool_handler
2100 )
2101 else:
2102 self._maint_notifications_pool_handler = None
2103 self._oss_cluster_maint_notifications_handler = None
2104
2105 async def _on_close(self) -> None:
2106 """Hook invoked from the pool's ``aclose()`` before the pool is shut down."""
2107 if self._maint_notifications_pool_handler is not None:
2108 await self._maint_notifications_pool_handler.cancel_scheduled_tasks()
2109
2110 @property
2111 @abstractmethod
2112 def connection_kwargs(self) -> dict[str, Any]:
2113 pass
2114
2115 @connection_kwargs.setter
2116 @abstractmethod
2117 def connection_kwargs(self, value: dict[str, Any]) -> None:
2118 pass
2119
2120 @abstractmethod
2121 def _get_pool_lock(self) -> asyncio.Lock:
2122 pass
2123
2124 @abstractmethod
2125 def _get_free_connections(self) -> Iterable["AbstractConnection"]:
2126 pass
2127
2128 @abstractmethod
2129 def _get_in_use_connections(self) -> Iterable["AbstractConnection"]:
2130 pass
2131
2132 def _maintenance_notifications_supported(self) -> bool:
2133 if "path" in self.connection_kwargs:
2134 return False
2135 if not self._maintenance_notifications_connection_class_supported():
2136 return False
2137 return bool(self.connection_kwargs.get("host"))
2138
2139 def _maintenance_notifications_connection_class_supported(self) -> bool:
2140 connection_class = getattr(self, "connection_class", None)
2141 if connection_class is None:
2142 return False
2143 try:
2144 return issubclass(
2145 connection_class, AsyncMaintNotificationsAbstractConnection
2146 )
2147 except TypeError:
2148 return False
2149
2150 def maint_notifications_enabled(self):
2151 """
2152 Returns:
2153 True if the maintenance notifications are enabled, False otherwise.
2154 The maintenance notifications config is stored in the pool handler.
2155 If the pool handler is not set, the maintenance notifications are not enabled.
2156 """
2157 if self._oss_cluster_maint_notifications_handler:
2158 maint_notifications_config = (
2159 self._oss_cluster_maint_notifications_handler.config
2160 )
2161 else:
2162 maint_notifications_config = (
2163 self._maint_notifications_pool_handler.config
2164 if self._maint_notifications_pool_handler
2165 else None
2166 )
2167 return maint_notifications_config and maint_notifications_config.enabled
2168
2169 async def update_maint_notifications_config(
2170 self,
2171 maint_notifications_config: MaintNotificationsConfig,
2172 oss_cluster_maint_notifications_handler: (
2173 AsyncOSSMaintNotificationsHandler | None
2174 ) = None,
2175 ) -> None:
2176 """
2177 Updates the maintenance notifications configuration.
2178 This method should be called only if the pool was created
2179 without enabling the maintenance notifications and
2180 in a later point in time maintenance notifications
2181 are requested to be enabled.
2182 """
2183 if (
2184 self.maint_notifications_enabled()
2185 and not maint_notifications_config.enabled
2186 ):
2187 raise ValueError(
2188 "Cannot disable maintenance notifications after enabling them"
2189 )
2190
2191 if oss_cluster_maint_notifications_handler:
2192 self._oss_cluster_maint_notifications_handler = (
2193 oss_cluster_maint_notifications_handler
2194 )
2195 # OSS cluster mode and pool-handler mode are mutually exclusive
2196 # (see __init__). A pool created with the default RESP3 "auto"
2197 # config wires a pool handler before this method runs; clear it so
2198 # new and existing connections are not configured with both handlers.
2199 self._maint_notifications_pool_handler = None
2200 else:
2201 if (
2202 maint_notifications_config.enabled
2203 and not self._maintenance_notifications_supported()
2204 ):
2205 if maint_notifications_config.enabled is True:
2206 # Unix sockets do not have a host endpoint for CLIENT
2207 # MAINT_NOTIFICATIONS to describe.
2208 if "path" in self.connection_kwargs:
2209 raise RedisError(
2210 "Maintenance notifications are not supported for "
2211 "Unix domain socket connections"
2212 )
2213
2214 # Custom connection classes must inherit the async maintenance
2215 # mixin so handlers can update connection state safely.
2216 if not self._maintenance_notifications_connection_class_supported():
2217 connection_class = getattr(self, "connection_class", None)
2218 connection_class_name = getattr(
2219 connection_class, "__name__", connection_class
2220 )
2221 raise RedisError(
2222 "Maintenance notifications are not supported for "
2223 f"connection class {connection_class_name}"
2224 )
2225
2226 # TCP-like connections still need a host to identify the
2227 # endpoint that can move during maintenance.
2228 raise RedisError(
2229 "Maintenance notifications are not supported for connections "
2230 "without a host"
2231 )
2232 self._maint_notifications_pool_handler = None
2233 return
2234
2235 if self._oss_cluster_maint_notifications_handler:
2236 # Pool already in OSS cluster mode; update the OSS handler config
2237 # instead of creating a mutually-exclusive pool handler (which
2238 # would be silently ignored because the OSS handler wins priority
2239 # in both update helpers below).
2240 self._oss_cluster_maint_notifications_handler.config = (
2241 maint_notifications_config
2242 )
2243 elif not self._maint_notifications_pool_handler:
2244 self._maint_notifications_pool_handler = (
2245 AsyncMaintNotificationsPoolHandler(self, maint_notifications_config)
2246 )
2247 else:
2248 self._maint_notifications_pool_handler.config = (
2249 maint_notifications_config
2250 )
2251
2252 self._update_connection_kwargs_for_maint_notifications(
2253 maint_notifications_pool_handler=self._maint_notifications_pool_handler,
2254 oss_cluster_maint_notifications_handler=self._oss_cluster_maint_notifications_handler,
2255 )
2256 await self._update_maint_notifications_configs_for_connections(
2257 maint_notifications_pool_handler=self._maint_notifications_pool_handler,
2258 oss_cluster_maint_notifications_handler=self._oss_cluster_maint_notifications_handler,
2259 )
2260
2261 def _update_connection_kwargs_for_maint_notifications(
2262 self,
2263 maint_notifications_pool_handler: (
2264 AsyncMaintNotificationsPoolHandler | None
2265 ) = None,
2266 oss_cluster_maint_notifications_handler: (
2267 AsyncOSSMaintNotificationsHandler | None
2268 ) = None,
2269 ) -> None:
2270 """
2271 Update the connection kwargs for all future connections.
2272 """
2273 if not self.maint_notifications_enabled():
2274 return
2275
2276 if maint_notifications_pool_handler:
2277 self.connection_kwargs.update(
2278 {
2279 "maint_notifications_pool_handler": maint_notifications_pool_handler,
2280 "maint_notifications_config": maint_notifications_pool_handler.config,
2281 }
2282 )
2283 if oss_cluster_maint_notifications_handler:
2284 self.connection_kwargs.update(
2285 {
2286 "oss_cluster_maint_notifications_handler": oss_cluster_maint_notifications_handler,
2287 "maint_notifications_config": oss_cluster_maint_notifications_handler.config,
2288 }
2289 )
2290 # OSS cluster mode and pool-handler mode are mutually exclusive.
2291 # Drop any pool handler a default (RESP3 "auto") pool creation may
2292 # have wired so future connections are not configured with both.
2293 self.connection_kwargs.pop("maint_notifications_pool_handler", None)
2294
2295 # Store original connection parameters for maintenance notifications.
2296 if self.connection_kwargs.get("orig_host_address", None) is None:
2297 # If orig_host_address is None it means we haven't
2298 # configured the original values yet
2299 self.connection_kwargs.update(
2300 {
2301 "orig_host_address": self.connection_kwargs.get("host"),
2302 "orig_socket_timeout": self.connection_kwargs.get(
2303 "socket_timeout", DEFAULT_SOCKET_TIMEOUT
2304 ),
2305 "orig_socket_connect_timeout": self.connection_kwargs.get(
2306 "socket_connect_timeout", DEFAULT_SOCKET_CONNECT_TIMEOUT
2307 ),
2308 }
2309 )
2310
2311 async def _update_maint_notifications_configs_for_connections(
2312 self,
2313 maint_notifications_pool_handler: (
2314 AsyncMaintNotificationsPoolHandler | None
2315 ) = None,
2316 oss_cluster_maint_notifications_handler: (
2317 AsyncOSSMaintNotificationsHandler | None
2318 ) = None,
2319 ) -> None:
2320 """Update the maintenance notifications config for all connections in the pool."""
2321 async with self._get_pool_lock():
2322 for conn in list(self._get_free_connections()):
2323 if oss_cluster_maint_notifications_handler:
2324 conn.set_maint_notifications_cluster_handler_for_connection(
2325 oss_cluster_maint_notifications_handler
2326 )
2327 conn.maint_notifications_config = (
2328 oss_cluster_maint_notifications_handler.config
2329 )
2330 elif maint_notifications_pool_handler:
2331 conn.set_maint_notifications_pool_handler_for_connection(
2332 maint_notifications_pool_handler
2333 )
2334 conn.maint_notifications_config = (
2335 maint_notifications_pool_handler.config
2336 )
2337 else:
2338 raise ValueError(
2339 "Either maint_notifications_pool_handler or "
2340 "oss_cluster_maint_notifications_handler must be set"
2341 )
2342 await conn.disconnect()
2343
2344 for conn in list(self._get_in_use_connections()):
2345 if oss_cluster_maint_notifications_handler:
2346 # Use set_maint_notifications_cluster_handler_for_connection
2347 # (not _configure_maintenance_notifications) so the parser is
2348 # obtained from the connection itself. _configure_* requires a
2349 # parser argument and would raise here; it would also reset the
2350 # connection's orig_* settings, which is wrong for an in-use
2351 # (active) connection. This mirrors the idle-connection branch
2352 # above and the pool-handler branches.
2353 conn.set_maint_notifications_cluster_handler_for_connection(
2354 oss_cluster_maint_notifications_handler
2355 )
2356 conn.maint_notifications_config = (
2357 oss_cluster_maint_notifications_handler.config
2358 )
2359 elif maint_notifications_pool_handler:
2360 conn.set_maint_notifications_pool_handler_for_connection(
2361 maint_notifications_pool_handler
2362 )
2363 conn.maint_notifications_config = (
2364 maint_notifications_pool_handler.config
2365 )
2366 else:
2367 raise ValueError(
2368 "Either maint_notifications_pool_handler or "
2369 "oss_cluster_maint_notifications_handler must be set"
2370 )
2371 if logger.isEnabledFor(logging.DEBUG):
2372 logger.debug(
2373 "Marking active connection for reconnect after config update "
2374 f"config update: {conn}, {conn.extract_connection_details()}"
2375 )
2376 conn.mark_for_reconnect()
2377
2378 def _should_update_connection(
2379 self,
2380 conn: "AbstractConnection",
2381 matching_pattern: str = "connected_address",
2382 matching_address: str | None = None,
2383 matching_notification_hash: int | None = None,
2384 ) -> bool:
2385 """
2386 Check if the connection should be updated based on the matching criteria.
2387 """
2388 if matching_pattern == "connected_address":
2389 if matching_address and conn.getpeername() != matching_address:
2390 return False
2391 elif matching_pattern == "configured_address":
2392 if matching_address and conn.host != matching_address:
2393 return False
2394 elif matching_pattern == "notification_hash":
2395 if (
2396 matching_notification_hash is not None
2397 and conn.maintenance_notification_hash != matching_notification_hash
2398 ):
2399 return False
2400 return True
2401
2402 def update_connection_settings(
2403 self,
2404 conn: "AsyncMaintNotificationsAbstractConnection",
2405 state: MaintenanceState | None = None,
2406 maintenance_notification_hash: int | None = None,
2407 host_address: str | None = None,
2408 relaxed_timeout: float | None = None,
2409 update_notification_hash: bool = False,
2410 reset_host_address: bool = False,
2411 reset_relaxed_timeout: bool = False,
2412 ) -> None:
2413 """
2414 Update the settings for a single connection.
2415 """
2416 if state:
2417 conn.maintenance_state = state
2418
2419 if update_notification_hash:
2420 # update the notification hash only if requested
2421 conn.maintenance_notification_hash = maintenance_notification_hash
2422
2423 if host_address is not None:
2424 conn.set_tmp_settings(tmp_host_address=host_address)
2425
2426 if relaxed_timeout is not None:
2427 conn.set_tmp_settings(tmp_relaxed_timeout=relaxed_timeout)
2428
2429 if reset_relaxed_timeout or reset_host_address:
2430 conn.reset_tmp_settings(
2431 reset_host_address=reset_host_address,
2432 reset_relaxed_timeout=reset_relaxed_timeout,
2433 )
2434
2435 conn.update_current_socket_timeout(relaxed_timeout)
2436
2437 async def update_connections_settings(
2438 self,
2439 state: MaintenanceState | None = None,
2440 maintenance_notification_hash: int | None = None,
2441 host_address: str | None = None,
2442 relaxed_timeout: float | None = None,
2443 matching_address: str | None = None,
2444 matching_notification_hash: int | None = None,
2445 matching_pattern: Literal[
2446 "connected_address", "configured_address", "notification_hash"
2447 ] = "connected_address",
2448 update_notification_hash: bool = False,
2449 reset_host_address: bool = False,
2450 reset_relaxed_timeout: bool = False,
2451 include_free_connections: bool = True,
2452 ) -> None:
2453 """
2454 Update the settings for all matching connections in the pool.
2455
2456 This method does not create new connections.
2457 This method does not affect the connection kwargs.
2458
2459 :param state: The maintenance state to set for the connection.
2460 :param maintenance_notification_hash: The hash of the maintenance notification
2461 to set for the connection.
2462 :param host_address: The host address to set for the connection.
2463 :param relaxed_timeout: The relaxed timeout to set for the connection.
2464 :param matching_address: The address to match for the connection.
2465 :param matching_notification_hash: The notification hash to match for the connection.
2466 :param matching_pattern: The pattern to match for the connection.
2467 :param update_notification_hash: Whether to update the notification hash for the connection.
2468 :param reset_host_address: Whether to reset the host address to the original address.
2469 :param reset_relaxed_timeout: Whether to reset the relaxed timeout to the original timeout.
2470 :param include_free_connections: Whether to include free/available connections.
2471 """
2472 async with self._get_pool_lock():
2473 self._update_connections_settings_without_locking(
2474 state=state,
2475 maintenance_notification_hash=maintenance_notification_hash,
2476 host_address=host_address,
2477 relaxed_timeout=relaxed_timeout,
2478 matching_address=matching_address,
2479 matching_notification_hash=matching_notification_hash,
2480 matching_pattern=matching_pattern,
2481 update_notification_hash=update_notification_hash,
2482 reset_host_address=reset_host_address,
2483 reset_relaxed_timeout=reset_relaxed_timeout,
2484 include_free_connections=include_free_connections,
2485 )
2486
2487 def _update_connections_settings_without_locking(
2488 self,
2489 state: MaintenanceState | None = None,
2490 maintenance_notification_hash: int | None = None,
2491 host_address: str | None = None,
2492 relaxed_timeout: float | None = None,
2493 matching_address: str | None = None,
2494 matching_notification_hash: int | None = None,
2495 matching_pattern: Literal[
2496 "connected_address", "configured_address", "notification_hash"
2497 ] = "connected_address",
2498 update_notification_hash: bool = False,
2499 reset_host_address: bool = False,
2500 reset_relaxed_timeout: bool = False,
2501 include_free_connections: bool = True,
2502 ) -> None:
2503 """
2504 Update matching connections while the caller already holds the pool lock.
2505
2506 This helper intentionally does not acquire the pool lock so callers can
2507 compose several pool mutations inside one critical section without
2508 deadlocking the non-reentrant `asyncio.Lock`.
2509 """
2510 for conn in self._get_in_use_connections():
2511 if self._should_update_connection(
2512 conn,
2513 matching_pattern,
2514 matching_address,
2515 matching_notification_hash,
2516 ):
2517 self.update_connection_settings(
2518 conn,
2519 state=state,
2520 maintenance_notification_hash=maintenance_notification_hash,
2521 host_address=host_address,
2522 relaxed_timeout=relaxed_timeout,
2523 update_notification_hash=update_notification_hash,
2524 reset_host_address=reset_host_address,
2525 reset_relaxed_timeout=reset_relaxed_timeout,
2526 )
2527
2528 if include_free_connections:
2529 for conn in self._get_free_connections():
2530 if self._should_update_connection(
2531 conn,
2532 matching_pattern,
2533 matching_address,
2534 matching_notification_hash,
2535 ):
2536 self.update_connection_settings(
2537 conn,
2538 state=state,
2539 maintenance_notification_hash=maintenance_notification_hash,
2540 host_address=host_address,
2541 relaxed_timeout=relaxed_timeout,
2542 update_notification_hash=update_notification_hash,
2543 reset_host_address=reset_host_address,
2544 reset_relaxed_timeout=reset_relaxed_timeout,
2545 )
2546
2547 def update_connection_kwargs(self, **kwargs: Any) -> None:
2548 """
2549 Update the connection kwargs for all future connections.
2550
2551 This method updates the connection kwargs for all future connections created by the pool.
2552 Existing connections are not affected.
2553 """
2554 self.connection_kwargs.update(kwargs)
2555
2556 async def apply_moving_notification(
2557 self,
2558 notification: NodeMovingNotification,
2559 config: MaintNotificationsConfig,
2560 moving_address_src: str | None,
2561 run_proactive_reconnect: bool = False,
2562 ) -> None:
2563 """
2564 Apply the pool state transition for a MOVING notification atomically.
2565
2566 Async pools use a non-reentrant `asyncio.Lock`, so the handler cannot
2567 safely compose several separately locked calls. Existing connection
2568 updates, optional proactive reconnect, and future `connection_kwargs`
2569 changes must happen under one pool-owned lock; otherwise a connection
2570 can move between active/free lists and escape handling.
2571 """
2572 if logger.isEnabledFor(logging.DEBUG):
2573 logger.debug(
2574 f"Applying MOVING notification to pool: {notification}, "
2575 f"moving address src: {moving_address_src}, "
2576 f"proactive reconnect: {run_proactive_reconnect}"
2577 )
2578 async with self._get_pool_lock():
2579 # Opt BlockingConnectionPool into serializing its get/release
2580 # with this critical section. Other pools do not define
2581 # set_in_maintenance and this is a no-op for them.
2582 self._set_in_maintenance(True)
2583 try:
2584 self._update_connections_settings_without_locking(
2585 state=MaintenanceState.MOVING,
2586 maintenance_notification_hash=hash(notification),
2587 relaxed_timeout=config.relaxed_timeout,
2588 host_address=notification.new_node_host,
2589 matching_address=moving_address_src,
2590 matching_pattern="connected_address",
2591 update_notification_hash=True,
2592 include_free_connections=True,
2593 )
2594
2595 if run_proactive_reconnect:
2596 await self._run_proactive_reconnect_without_locking(
2597 moving_address_src
2598 )
2599
2600 self.update_connection_kwargs(
2601 **_build_moving_connection_kwargs(notification, config)
2602 )
2603 finally:
2604 self._set_in_maintenance(False)
2605
2606 async def run_proactive_reconnect(
2607 self,
2608 moving_address_src: str | None = None,
2609 ) -> None:
2610 """
2611 Mark active connections and disconnect free connections atomically.
2612
2613 This operation is pool-owned because the active/free lists can change
2614 while tasks acquire or release connections. Keeping the mark/disconnect
2615 pass under one lock avoids a connection moving between lists between
2616 separately locked calls.
2617 """
2618 async with self._get_pool_lock():
2619 await self._run_proactive_reconnect_without_locking(moving_address_src)
2620
2621 async def _run_proactive_reconnect_without_locking(
2622 self,
2623 moving_address_src: str | None = None,
2624 ) -> None:
2625 """
2626 Mark and disconnect matching connections while the caller holds the pool lock.
2627
2628 This helper intentionally does not acquire the pool lock so it can be
2629 reused by larger atomic operations that already hold the non-reentrant
2630 `asyncio.Lock`.
2631 """
2632 debug = logger.isEnabledFor(logging.DEBUG)
2633 for conn in self._get_in_use_connections():
2634 if self._should_update_connection(
2635 conn, "connected_address", moving_address_src
2636 ):
2637 if debug:
2638 logger.debug(
2639 f"Marking active connection for reconnect: {conn}, "
2640 f"{conn.extract_connection_details()}"
2641 )
2642 conn.mark_for_reconnect()
2643
2644 free_connections = [
2645 conn
2646 for conn in self._get_free_connections()
2647 if self._should_update_connection(
2648 conn, "connected_address", moving_address_src
2649 )
2650 ]
2651 if debug:
2652 for conn in free_connections:
2653 logger.debug(
2654 f"Disconnecting free connection: {conn}, "
2655 f"{conn.extract_connection_details()}"
2656 )
2657 await self._disconnect_connections(free_connections)
2658
2659 async def cleanup_moving_notification(
2660 self,
2661 notification_hash: int,
2662 reset_relaxed_timeout: bool,
2663 reset_host_address: bool,
2664 ) -> None:
2665 """
2666 Revert MOVING pool state atomically after the notification TTL.
2667
2668 Future connection kwargs and existing connection state must be cleaned
2669 up in the same critical section. Splitting the cleanup lets an
2670 acquire/release interleave, which can leave stale MOVING state or undo a
2671 newer overlapping MOVING notification.
2672 """
2673 if logger.isEnabledFor(logging.DEBUG):
2674 logger.debug(
2675 "Cleaning up MOVING pool state for notification hash "
2676 f"{notification_hash}, reset_relaxed_timeout="
2677 f"{reset_relaxed_timeout}, reset_host_address={reset_host_address}"
2678 )
2679 async with self._get_pool_lock():
2680 kwargs = _build_moving_cleanup_connection_kwargs(
2681 self.connection_kwargs, notification_hash
2682 )
2683 if kwargs is not None:
2684 self.update_connection_kwargs(**kwargs)
2685
2686 self._update_connections_settings_without_locking(
2687 relaxed_timeout=-1,
2688 state=MaintenanceState.NONE,
2689 maintenance_notification_hash=None,
2690 matching_notification_hash=notification_hash,
2691 matching_pattern="notification_hash",
2692 update_notification_hash=True,
2693 reset_relaxed_timeout=reset_relaxed_timeout,
2694 reset_host_address=reset_host_address,
2695 include_free_connections=True,
2696 )
2697
2698 async def _disconnect_connections(
2699 self, connections: Iterable["AbstractConnection"]
2700 ) -> None:
2701 connections = tuple(connections)
2702 if not connections:
2703 return
2704 results = await asyncio.gather(
2705 *(connection.disconnect() for connection in connections),
2706 return_exceptions=True,
2707 )
2708 exc = next(
2709 (result for result in results if isinstance(result, BaseException)), None
2710 )
2711 if exc:
2712 raise exc
2713
2714 def _set_in_maintenance(self, in_maintenance: bool) -> None:
2715 """Flip the pool's maintenance flag if it exposes one (BlockingConnectionPool)."""
2716 set_in_maintenance = getattr(self, "set_in_maintenance", None)
2717 if callable(set_in_maintenance):
2718 set_in_maintenance(in_maintenance)
2719
2720
2721class ConnectionPool(
2722 AsyncMaintNotificationsAbstractConnectionPool, ConnectionPoolInterface
2723):
2724 """
2725 Create a connection pool. ``If max_connections`` is set, then this
2726 object raises :py:class:`~redis.ConnectionError` when the pool's
2727 limit is reached.
2728
2729 By default, TCP connections are created unless ``connection_class``
2730 is specified. Use :py:class:`~redis.UnixDomainSocketConnection` for
2731 unix sockets.
2732 :py:class:`~redis.SSLConnection` can be used for SSL enabled connections.
2733
2734 Any additional keyword arguments are passed to the constructor of
2735 ``connection_class``.
2736 """
2737
2738 @classmethod
2739 def from_url(cls: Type[_CP], url: str, **kwargs) -> _CP:
2740 """
2741 Return a connection pool configured from the given URL.
2742
2743 For example::
2744
2745 redis://[[username]:[password]]@localhost:6379/0
2746 rediss://[[username]:[password]]@localhost:6379/0
2747 unix://[username@]/path/to/socket.sock?db=0[&password=password]
2748
2749 Three URL schemes are supported:
2750
2751 - `redis://` creates a TCP socket connection. See more at:
2752 <https://www.iana.org/assignments/uri-schemes/prov/redis>
2753 - `rediss://` creates a SSL wrapped TCP socket connection. See more at:
2754 <https://www.iana.org/assignments/uri-schemes/prov/rediss>
2755 - ``unix://``: creates a Unix Domain Socket connection.
2756
2757 The username, password, hostname and path are passed through
2758 urllib.parse.unquote in order to replace any percent-encoded values
2759 with their corresponding characters. Querystring values are decoded
2760 by urllib.parse.parse_qs and are not unquoted again.
2761
2762 There are several ways to specify a database number. The first value
2763 found will be used:
2764
2765 1. A ``db`` querystring option, e.g. redis://localhost?db=0
2766
2767 2. If using the redis:// or rediss:// schemes, the path argument
2768 of the url, e.g. redis://localhost/0
2769
2770 3. A ``db`` keyword argument to this function.
2771
2772 If none of these options are specified, the default db=0 is used.
2773
2774 All querystring options are cast to their appropriate Python types.
2775 Boolean arguments can be specified with string values "True"/"False"
2776 or "Yes"/"No". Values that cannot be properly cast cause a
2777 ``ValueError`` to be raised. Once parsed, the querystring arguments
2778 and keyword arguments are passed to the ``ConnectionPool``'s
2779 class initializer. In the case of conflicting arguments, querystring
2780 arguments always win.
2781 """
2782 url_options = parse_url(url)
2783 kwargs.update(url_options)
2784 return cls(**kwargs)
2785
2786 def __init__(
2787 self,
2788 connection_class: Type[AbstractConnection] = Connection,
2789 max_connections: Optional[int] = None,
2790 maint_notifications_config: MaintNotificationsConfig | None = None,
2791 **connection_kwargs,
2792 ):
2793 max_connections = max_connections or 100
2794 if not isinstance(max_connections, int) or max_connections < 0:
2795 raise ValueError('"max_connections" must be a positive integer')
2796
2797 self.connection_class = connection_class
2798 self._connection_kwargs = connection_kwargs
2799 self.max_connections = max_connections
2800
2801 # Resolve the HIMPORT registry. A pre-built ``himport_registry`` (shared, e.g.
2802 # from the cluster client) takes precedence; otherwise build a fresh empty one.
2803 # A registry always exists so runtime ``himport_prepare`` mutates a single object
2804 # every connection already shares. The object stays in ``connection_kwargs`` so
2805 # it reaches every connection. It is injected unconditionally (like other
2806 # auto-added pool kwargs), so a custom ``connection_class`` must accept
2807 # ``**kwargs`` (or a ``himport_registry`` parameter), as built-ins do.
2808 himport_registry = connection_kwargs.get("himport_registry")
2809 if himport_registry is None:
2810 himport_registry = HImportRegistry()
2811 connection_kwargs["himport_registry"] = himport_registry
2812 self.himport_registry = himport_registry
2813
2814 self._available_connections: List[AbstractConnection] = []
2815 self._in_use_connections: Set[AbstractConnection] = set()
2816 self.encoder_class = self.connection_kwargs.get("encoder_class", Encoder)
2817 self._lock = asyncio.Lock()
2818 self._event_dispatcher = self.connection_kwargs.get("event_dispatcher", None)
2819 if self._event_dispatcher is None:
2820 self._event_dispatcher = EventDispatcher()
2821
2822 AsyncMaintNotificationsAbstractConnectionPool.__init__(
2823 self,
2824 maint_notifications_config=maint_notifications_config,
2825 **connection_kwargs,
2826 )
2827
2828 # Keys that should be redacted in __repr__ to avoid exposing sensitive information
2829 SENSITIVE_REPR_KEYS = frozenset(
2830 {
2831 "password",
2832 "username",
2833 "ssl_password",
2834 "credential_provider",
2835 }
2836 )
2837
2838 # Internal plumbing kwargs omitted from __repr__ (not user-facing config).
2839 OMIT_REPR_KEYS = frozenset({"himport_registry"})
2840
2841 def __repr__(self):
2842 conn_kwargs = ",".join(
2843 [
2844 f"{k}={'<REDACTED>' if k in self.SENSITIVE_REPR_KEYS else v}"
2845 for k, v in self.connection_kwargs.items()
2846 if k not in self.OMIT_REPR_KEYS
2847 ]
2848 )
2849 return (
2850 f"<{self.__class__.__module__}.{self.__class__.__name__}"
2851 f"(<{self.connection_class.__module__}.{self.connection_class.__name__}"
2852 f"({conn_kwargs})>)>"
2853 )
2854
2855 @property
2856 def connection_kwargs(self) -> dict[str, Any]:
2857 return self._connection_kwargs
2858
2859 @connection_kwargs.setter
2860 def connection_kwargs(self, value: dict[str, Any]) -> None:
2861 self._connection_kwargs = value
2862
2863 def _get_pool_lock(self) -> asyncio.Lock:
2864 return self._lock
2865
2866 def _get_free_connections(self) -> Iterable[AbstractConnection]:
2867 return self._available_connections
2868
2869 def _get_in_use_connections(self) -> Iterable[AbstractConnection]:
2870 return self._in_use_connections
2871
2872 def get_protocol(self):
2873 """
2874 Returns:
2875 The RESP protocol version, or ``None`` if the protocol is not specified,
2876 in which case the server default will be used.
2877 """
2878 return self.connection_kwargs.get("protocol", None)
2879
2880 def reset(self):
2881 # Record metrics for connections being removed before clearing
2882 # (only if attributes exist - they won't during __init__)
2883 if hasattr(self, "_available_connections") and hasattr(
2884 self, "_in_use_connections"
2885 ):
2886 idle_count = len(self._available_connections)
2887 in_use_count = len(self._in_use_connections)
2888 if idle_count > 0 or in_use_count > 0:
2889 pool_name = get_pool_name(self)
2890 # Note: Using sync version since reset() is sync
2891 from redis.observability.recorder import (
2892 record_connection_count as sync_record_connection_count,
2893 )
2894
2895 if idle_count > 0:
2896 sync_record_connection_count(
2897 pool_name=pool_name,
2898 connection_state=ConnectionState.IDLE,
2899 counter=-idle_count,
2900 )
2901 if in_use_count > 0:
2902 sync_record_connection_count(
2903 pool_name=pool_name,
2904 connection_state=ConnectionState.USED,
2905 counter=-in_use_count,
2906 )
2907
2908 self._available_connections = []
2909 self._in_use_connections = weakref.WeakSet()
2910
2911 def __del__(self) -> None:
2912 """Clean up connection pool and record metrics when garbage collected."""
2913 try:
2914 if not hasattr(self, "_available_connections") or not hasattr(
2915 self, "_in_use_connections"
2916 ):
2917 return
2918 idle_count = len(self._available_connections)
2919 in_use_count = len(self._in_use_connections)
2920 if idle_count > 0 or in_use_count > 0:
2921 pool_name = get_pool_name(self)
2922 # Note: Using sync version since __del__ is sync
2923 from redis.observability.recorder import (
2924 record_connection_count as sync_record_connection_count,
2925 )
2926
2927 if idle_count > 0:
2928 sync_record_connection_count(
2929 pool_name=pool_name,
2930 connection_state=ConnectionState.IDLE,
2931 counter=-idle_count,
2932 )
2933 if in_use_count > 0:
2934 sync_record_connection_count(
2935 pool_name=pool_name,
2936 connection_state=ConnectionState.USED,
2937 counter=-in_use_count,
2938 )
2939 except Exception:
2940 pass
2941
2942 def can_get_connection(self) -> bool:
2943 """Return True if a connection can be retrieved from the pool."""
2944 return (
2945 self._available_connections
2946 or len(self._in_use_connections) < self.max_connections
2947 )
2948
2949 @deprecated_args(
2950 args_to_warn=["*"],
2951 reason="Use get_connection() without args instead",
2952 version="5.3.0",
2953 )
2954 async def get_connection(self, command_name=None, *keys, **options):
2955 """Get a connected connection from the pool"""
2956 # Track connection count before to detect if a new connection is created
2957 async with self._lock:
2958 connections_before = len(self._available_connections) + len(
2959 self._in_use_connections
2960 )
2961 start_time_created = time.monotonic()
2962 connection = self.get_available_connection()
2963 connections_after = len(self._available_connections) + len(
2964 self._in_use_connections
2965 )
2966 is_created = connections_after > connections_before
2967
2968 # Record state transition for observability
2969 # This ensures counters stay balanced if ensure_connection() fails and release() is called
2970 pool_name = get_pool_name(self)
2971 if is_created:
2972 # New connection created and acquired: just USED +1
2973 await record_connection_count(
2974 pool_name=pool_name,
2975 connection_state=ConnectionState.USED,
2976 counter=1,
2977 )
2978 else:
2979 # Existing connection acquired from pool: IDLE -> USED
2980 await record_connection_count(
2981 pool_name=pool_name,
2982 connection_state=ConnectionState.IDLE,
2983 counter=-1,
2984 )
2985 await record_connection_count(
2986 pool_name=pool_name,
2987 connection_state=ConnectionState.USED,
2988 counter=1,
2989 )
2990
2991 # We now perform the connection check outside of the lock.
2992 try:
2993 await self.ensure_connection(connection)
2994
2995 if is_created:
2996 await record_connection_create_time(
2997 connection_pool=self,
2998 duration_seconds=time.monotonic() - start_time_created,
2999 )
3000
3001 return connection
3002 except BaseException:
3003 await self.release(connection)
3004 raise
3005
3006 def get_available_connection(self):
3007 """Get a connection from the pool, without making sure it is connected"""
3008 try:
3009 connection = self._available_connections.pop()
3010 except IndexError:
3011 if len(self._in_use_connections) >= self.max_connections:
3012 raise MaxConnectionsError("Too many connections") from None
3013 connection = self.make_connection()
3014 self._in_use_connections.add(connection)
3015 return connection
3016
3017 def get_encoder(self):
3018 """Return an encoder based on encoding settings"""
3019 kwargs = self.connection_kwargs
3020 return self.encoder_class(
3021 encoding=kwargs.get("encoding", "utf-8"),
3022 encoding_errors=kwargs.get("encoding_errors", "strict"),
3023 decode_responses=kwargs.get("decode_responses", False),
3024 )
3025
3026 def make_connection(self):
3027 """Create a new connection. Can be overridden by child classes."""
3028 # Note: We don't record IDLE here because async uses a sync make_connection
3029 # but async record_connection_count. The recording is handled in get_connection.
3030 return self.connection_class(**self.connection_kwargs)
3031
3032 async def ensure_connection(self, connection: AbstractConnection):
3033 """Ensure that the connection object is connected and valid"""
3034 await connection.connect()
3035 # connections that the pool provides should be ready to send
3036 # a command. if not, the connection was either returned to the
3037 # pool before all data has been read or the socket has been
3038 # closed. either way, reconnect and verify everything is good.
3039 try:
3040 if await connection.can_read() and not self.maint_notifications_enabled():
3041 raise ConnectionError("Connection has data") from None
3042 except (ConnectionError, TimeoutError, OSError):
3043 await connection.disconnect()
3044 await connection.connect()
3045 if await connection.can_read() and not self.maint_notifications_enabled():
3046 raise ConnectionError("Connection not ready") from None
3047
3048 async def release(self, connection: AbstractConnection):
3049 """Releases the connection back to the pool"""
3050 # Connections should always be returned to the correct pool,
3051 # not doing so is an error that will cause an exception here.
3052 async with self._lock:
3053 self._in_use_connections.remove(connection)
3054
3055 if connection.should_reconnect():
3056 if logger.isEnabledFor(logging.DEBUG):
3057 logger.debug(
3058 "Disconnecting released connection marked for reconnect: "
3059 f"{connection}, {connection.extract_connection_details()}"
3060 )
3061 await connection.disconnect()
3062
3063 self._available_connections.append(connection)
3064
3065 await self._event_dispatcher.dispatch_async(
3066 AsyncAfterConnectionReleasedEvent(connection)
3067 )
3068
3069 # Record state transition: USED -> IDLE
3070 pool_name = get_pool_name(self)
3071 await record_connection_count(
3072 pool_name=pool_name,
3073 connection_state=ConnectionState.USED,
3074 counter=-1,
3075 )
3076 await record_connection_count(
3077 pool_name=pool_name,
3078 connection_state=ConnectionState.IDLE,
3079 counter=1,
3080 )
3081
3082 async def disconnect(self, inuse_connections: bool = True):
3083 """
3084 Disconnects connections in the pool
3085
3086 If ``inuse_connections`` is True, disconnect connections that are
3087 current in use, potentially by other tasks. Otherwise only disconnect
3088 connections that are idle in the pool.
3089 """
3090 if inuse_connections:
3091 connections: Iterable[AbstractConnection] = chain(
3092 self._available_connections, self._in_use_connections
3093 )
3094 else:
3095 connections = self._available_connections
3096 resp = await asyncio.gather(
3097 *(connection.disconnect() for connection in connections),
3098 return_exceptions=True,
3099 )
3100
3101 exc = next((r for r in resp if isinstance(r, BaseException)), None)
3102 if exc:
3103 raise exc
3104
3105 async def update_active_connections_for_reconnect(self):
3106 """
3107 Mark all active connections for reconnect.
3108 """
3109 debug = logger.isEnabledFor(logging.DEBUG)
3110 async with self._lock:
3111 for conn in self._in_use_connections:
3112 if debug:
3113 logger.debug(
3114 f"Marking active connection for reconnect: {conn}, "
3115 f"{conn.extract_connection_details()}"
3116 )
3117 conn.mark_for_reconnect()
3118
3119 async def aclose(self) -> None:
3120 """Close the pool, disconnecting all connections"""
3121 await self._on_close()
3122 await self.disconnect()
3123
3124 async def __aenter__(self: _CP) -> _CP:
3125 return self
3126
3127 async def __aexit__(self, exc_type, exc_value, traceback) -> None:
3128 await self.aclose()
3129
3130 def set_retry(self, retry: "Retry") -> None:
3131 for conn in self._available_connections:
3132 conn.retry = retry
3133 for conn in self._in_use_connections:
3134 conn.retry = retry
3135
3136 async def re_auth_callback(self, token: TokenInterface):
3137 async with self._lock:
3138 for conn in self._available_connections:
3139 await conn.retry.call_with_retry(
3140 lambda: conn.send_command(
3141 "AUTH", token.try_get("oid"), token.get_value()
3142 ),
3143 lambda error: self._mock(error),
3144 )
3145 await conn.retry.call_with_retry(
3146 lambda: conn.read_response(), lambda error: self._mock(error)
3147 )
3148 for conn in self._in_use_connections:
3149 conn.set_re_auth_token(token)
3150
3151 async def _mock(self, error: RedisError):
3152 """
3153 Dummy functions, needs to be passed as error callback to retry object.
3154 :param error:
3155 :return:
3156 """
3157 pass
3158
3159 def get_connection_count(self) -> List[tuple[int, dict]]:
3160 """
3161 Returns a connection count (both idle and in use).
3162 """
3163 attributes = AttributeBuilder.build_base_attributes()
3164 attributes[DB_CLIENT_CONNECTION_POOL_NAME] = get_pool_name(self)
3165 free_connections_attributes = attributes.copy()
3166 in_use_connections_attributes = attributes.copy()
3167
3168 free_connections_attributes[DB_CLIENT_CONNECTION_STATE] = (
3169 ConnectionState.IDLE.value
3170 )
3171 in_use_connections_attributes[DB_CLIENT_CONNECTION_STATE] = (
3172 ConnectionState.USED.value
3173 )
3174
3175 return [
3176 (len(self._available_connections), free_connections_attributes),
3177 (len(self._in_use_connections), in_use_connections_attributes),
3178 ]
3179
3180
3181class BlockingConnectionPool(ConnectionPool):
3182 """
3183 A blocking connection pool::
3184
3185 >>> from redis.asyncio import Redis, BlockingConnectionPool
3186 >>> client = Redis.from_pool(BlockingConnectionPool())
3187
3188 It performs the same function as the default
3189 :py:class:`~redis.asyncio.ConnectionPool` implementation, in that,
3190 it maintains a pool of reusable connections that can be shared by
3191 multiple async redis clients.
3192
3193 The difference is that, in the event that a client tries to get a
3194 connection from the pool when all of connections are in use, rather than
3195 raising a :py:class:`~redis.ConnectionError` (as the default
3196 :py:class:`~redis.asyncio.ConnectionPool` implementation does), it
3197 blocks the current `Task` for a specified number of seconds until
3198 a connection becomes available.
3199
3200 Use ``max_connections`` to increase / decrease the pool size::
3201
3202 >>> pool = BlockingConnectionPool(max_connections=10)
3203
3204 Use ``timeout`` to tell it either how many seconds to wait for a connection
3205 to become available, or to block forever:
3206
3207 >>> # Block forever.
3208 >>> pool = BlockingConnectionPool(timeout=None)
3209
3210 >>> # Raise a ``ConnectionError`` after five seconds if a connection is
3211 >>> # not available.
3212 >>> pool = BlockingConnectionPool(timeout=5)
3213 """
3214
3215 def __init__(
3216 self,
3217 max_connections: int = 50,
3218 timeout: Optional[float] = 20,
3219 connection_class: Type[AbstractConnection] = Connection,
3220 queue_class: Type[asyncio.Queue] = asyncio.LifoQueue, # deprecated
3221 **connection_kwargs,
3222 ):
3223 super().__init__(
3224 connection_class=connection_class,
3225 max_connections=max_connections,
3226 **connection_kwargs,
3227 )
3228 self._condition = asyncio.Condition()
3229 self.timeout = timeout
3230 self._in_maintenance = False
3231
3232 def set_in_maintenance(self, in_maintenance: bool) -> None:
3233 """
3234 Toggle the pool's maintenance mode.
3235
3236 While maintenance mode is on, ``get_connection`` and ``release``
3237 serialize their pool mutations through ``self._lock`` so they cannot
3238 interleave with a MOVING notification handler that is currently
3239 rewriting pool state under the same lock. Outside of maintenance the
3240 mutations skip the lock, since their critical sections are pure-Python
3241 and already atomic under asyncio's single-threaded scheduling.
3242 """
3243 self._in_maintenance = in_maintenance
3244
3245 @contextlib.asynccontextmanager
3246 async def _maybe_pool_lock(self) -> AsyncIterator[None]:
3247 if self._in_maintenance:
3248 async with self._lock:
3249 yield
3250 else:
3251 yield
3252
3253 @deprecated_args(
3254 args_to_warn=["*"],
3255 reason="Use get_connection() without args instead",
3256 version="5.3.0",
3257 )
3258 async def get_connection(self, command_name=None, *keys, **options):
3259 """Gets a connection from the pool, blocking until one is available"""
3260 # Start timing for wait time observability
3261 start_time_acquired = time.monotonic()
3262
3263 try:
3264 async with self._condition:
3265 async with async_timeout(self.timeout):
3266 await self._condition.wait_for(self.can_get_connection)
3267 async with self._maybe_pool_lock():
3268 # Track connection count before to detect if a new connection is created
3269 connections_before = len(self._available_connections) + len(
3270 self._in_use_connections
3271 )
3272 start_time_created = time.monotonic()
3273 connection = super().get_available_connection()
3274 connections_after = len(self._available_connections) + len(
3275 self._in_use_connections
3276 )
3277 is_created = connections_after > connections_before
3278 except asyncio.TimeoutError as err:
3279 raise ConnectionError("No connection available.") from err
3280
3281 # We now perform the connection check outside of the lock.
3282 try:
3283 await self.ensure_connection(connection)
3284
3285 if is_created:
3286 await record_connection_create_time(
3287 connection_pool=self,
3288 duration_seconds=time.monotonic() - start_time_created,
3289 )
3290
3291 await record_connection_wait_time(
3292 pool_name=get_pool_name(self),
3293 duration_seconds=time.monotonic() - start_time_acquired,
3294 )
3295
3296 return connection
3297 except BaseException:
3298 await self.release(connection)
3299 raise
3300
3301 async def release(self, connection: AbstractConnection):
3302 """Releases the connection back to the pool."""
3303 async with self._condition:
3304 await super().release(connection)
3305 self._condition.notify()