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