1import enum
2import ipaddress
3import logging
4import re
5import threading
6import time
7from abc import ABC, abstractmethod
8from typing import (
9 TYPE_CHECKING,
10 Any,
11 Dict,
12 List,
13 Literal,
14 Mapping,
15 Optional,
16 Union,
17)
18
19from redis.observability.attributes import get_pool_name
20from redis.observability.recorder import (
21 record_connection_handoff,
22 record_connection_relaxed_timeout,
23 record_maint_notification_count,
24)
25from redis.typing import Number
26
27if TYPE_CHECKING:
28 from redis.cluster import MaintNotificationsAbstractRedisCluster
29
30logger = logging.getLogger(__name__)
31
32
33class MaintenanceState(enum.Enum):
34 NONE = "none"
35 MOVING = "moving"
36 MAINTENANCE = "maintenance"
37
38
39class EndpointType(enum.Enum):
40 """Valid endpoint types used in CLIENT MAINT_NOTIFICATIONS command."""
41
42 INTERNAL_IP = "internal-ip"
43 INTERNAL_FQDN = "internal-fqdn"
44 EXTERNAL_IP = "external-ip"
45 EXTERNAL_FQDN = "external-fqdn"
46 NONE = "none"
47
48 def __str__(self):
49 """Return the string value of the enum."""
50 return self.value
51
52
53if TYPE_CHECKING:
54 from redis.asyncio.connection import AsyncMaintNotificationsAbstractConnection
55 from redis.connection import (
56 MaintNotificationsAbstractConnection,
57 MaintNotificationsAbstractConnectionPool,
58 )
59
60
61class MaintenanceNotification(ABC):
62 """
63 Base class for maintenance notifications sent through push messages by Redis server.
64
65 This class provides common functionality for all maintenance notifications including
66 unique identification and TTL (Time-To-Live) functionality.
67
68 Attributes:
69 id (int): Unique identifier for this notification
70 ttl (int): Time-to-live in seconds for this notification
71 creation_time (float): Timestamp when the notification was created/read
72 """
73
74 def __init__(self, id: int, ttl: int):
75 """
76 Initialize a new MaintenanceNotification with unique ID and TTL functionality.
77
78 Args:
79 id (int): Unique identifier for this notification
80 ttl (int): Time-to-live in seconds for this notification
81 """
82 self.id = id
83 self.ttl = ttl
84 self.creation_time = time.monotonic()
85 self.expire_at = self.creation_time + self.ttl
86
87 def is_expired(self) -> bool:
88 """
89 Check if this notification has expired based on its TTL
90 and creation time.
91
92 Returns:
93 bool: True if the notification has expired, False otherwise
94 """
95 return time.monotonic() > (self.creation_time + self.ttl)
96
97 @abstractmethod
98 def __repr__(self) -> str:
99 """
100 Return a string representation of the maintenance notification.
101
102 This method must be implemented by all concrete subclasses.
103
104 Returns:
105 str: String representation of the notification
106 """
107 pass
108
109 @abstractmethod
110 def __eq__(self, other) -> bool:
111 """
112 Compare two maintenance notifications for equality.
113
114 This method must be implemented by all concrete subclasses.
115 Notifications are typically considered equal if they have the same id
116 and are of the same type.
117
118 Args:
119 other: The other object to compare with
120
121 Returns:
122 bool: True if the notifications are equal, False otherwise
123 """
124 pass
125
126 @abstractmethod
127 def __hash__(self) -> int:
128 """
129 Return a hash value for the maintenance notification.
130
131 This method must be implemented by all concrete subclasses to allow
132 instances to be used in sets and as dictionary keys.
133
134 Returns:
135 int: Hash value for the notification
136 """
137 pass
138
139
140class NodeMovingNotification(MaintenanceNotification):
141 """
142 This notification is received when a node is replaced with a new node
143 during cluster rebalancing or maintenance operations.
144 """
145
146 def __init__(
147 self,
148 id: int,
149 new_node_host: Optional[str],
150 new_node_port: Optional[int],
151 ttl: int,
152 ):
153 """
154 Initialize a new NodeMovingNotification.
155
156 Args:
157 id (int): Unique identifier for this notification
158 new_node_host (str): Hostname or IP address of the new replacement node
159 new_node_port (int): Port number of the new replacement node
160 ttl (int): Time-to-live in seconds for this notification
161 """
162 super().__init__(id, ttl)
163 self.new_node_host = new_node_host
164 self.new_node_port = new_node_port
165
166 def __repr__(self) -> str:
167 expiry_time = self.expire_at
168 remaining = max(0, expiry_time - time.monotonic())
169
170 return (
171 f"{self.__class__.__name__}("
172 f"id={self.id}, "
173 f"new_node_host='{self.new_node_host}', "
174 f"new_node_port={self.new_node_port}, "
175 f"ttl={self.ttl}, "
176 f"creation_time={self.creation_time}, "
177 f"expires_at={expiry_time}, "
178 f"remaining={remaining:.1f}s, "
179 f"expired={self.is_expired()}"
180 f")"
181 )
182
183 def __eq__(self, other) -> bool:
184 """
185 Two NodeMovingNotification notifications are considered equal if they have the same
186 id, new_node_host, and new_node_port.
187 """
188 if not isinstance(other, NodeMovingNotification):
189 return False
190 return (
191 self.id == other.id
192 and self.new_node_host == other.new_node_host
193 and self.new_node_port == other.new_node_port
194 )
195
196 def __hash__(self) -> int:
197 """
198 Return a hash value for the notification to allow
199 instances to be used in sets and as dictionary keys.
200
201 Returns:
202 int: Hash value based on notification type class name, id,
203 new_node_host and new_node_port
204 """
205 try:
206 node_port = int(self.new_node_port) if self.new_node_port else None
207 except ValueError:
208 node_port = 0
209
210 return hash(
211 (
212 self.__class__.__name__,
213 int(self.id),
214 str(self.new_node_host),
215 node_port,
216 )
217 )
218
219
220class NodeMigratingNotification(MaintenanceNotification):
221 """
222 Notification for when a Redis cluster node is in the process of migrating slots.
223
224 This notification is received when a node starts migrating its slots to another node
225 during cluster rebalancing or maintenance operations.
226
227 Args:
228 id (int): Unique identifier for this notification
229 ttl (int): Time-to-live in seconds for this notification
230 """
231
232 def __init__(self, id: int, ttl: int):
233 super().__init__(id, ttl)
234
235 def __repr__(self) -> str:
236 expiry_time = self.creation_time + self.ttl
237 remaining = max(0, expiry_time - time.monotonic())
238 return (
239 f"{self.__class__.__name__}("
240 f"id={self.id}, "
241 f"ttl={self.ttl}, "
242 f"creation_time={self.creation_time}, "
243 f"expires_at={expiry_time}, "
244 f"remaining={remaining:.1f}s, "
245 f"expired={self.is_expired()}"
246 f")"
247 )
248
249 def __eq__(self, other) -> bool:
250 """
251 Two NodeMigratingNotification notifications are considered equal if they have the same
252 id and are of the same type.
253 """
254 if not isinstance(other, NodeMigratingNotification):
255 return False
256 return self.id == other.id and type(self) is type(other)
257
258 def __hash__(self) -> int:
259 """
260 Return a hash value for the notification to allow
261 instances to be used in sets and as dictionary keys.
262
263 Returns:
264 int: Hash value based on notification type and id
265 """
266 return hash((self.__class__.__name__, int(self.id)))
267
268
269class NodeMigratedNotification(MaintenanceNotification):
270 """
271 Notification for when a Redis cluster node has completed migrating slots.
272
273 This notification is received when a node has finished migrating all its slots
274 to other nodes during cluster rebalancing or maintenance operations.
275
276 Args:
277 id (int): Unique identifier for this notification
278 """
279
280 DEFAULT_TTL = 5
281
282 def __init__(self, id: int):
283 super().__init__(id, NodeMigratedNotification.DEFAULT_TTL)
284
285 def __repr__(self) -> str:
286 expiry_time = self.creation_time + self.ttl
287 remaining = max(0, expiry_time - time.monotonic())
288 return (
289 f"{self.__class__.__name__}("
290 f"id={self.id}, "
291 f"ttl={self.ttl}, "
292 f"creation_time={self.creation_time}, "
293 f"expires_at={expiry_time}, "
294 f"remaining={remaining:.1f}s, "
295 f"expired={self.is_expired()}"
296 f")"
297 )
298
299 def __eq__(self, other) -> bool:
300 """
301 Two NodeMigratedNotification notifications are considered equal if they have the same
302 id and are of the same type.
303 """
304 if not isinstance(other, NodeMigratedNotification):
305 return False
306 return self.id == other.id and type(self) is type(other)
307
308 def __hash__(self) -> int:
309 """
310 Return a hash value for the notification to allow
311 instances to be used in sets and as dictionary keys.
312
313 Returns:
314 int: Hash value based on notification type and id
315 """
316 return hash((self.__class__.__name__, int(self.id)))
317
318
319class NodeFailingOverNotification(MaintenanceNotification):
320 """
321 Notification for when a Redis cluster node is in the process of failing over.
322
323 This notification is received when a node starts a failover process during
324 cluster maintenance operations or when handling node failures.
325
326 Args:
327 id (int): Unique identifier for this notification
328 ttl (int): Time-to-live in seconds for this notification
329 """
330
331 def __init__(self, id: int, ttl: int):
332 super().__init__(id, ttl)
333
334 def __repr__(self) -> str:
335 expiry_time = self.creation_time + self.ttl
336 remaining = max(0, expiry_time - time.monotonic())
337 return (
338 f"{self.__class__.__name__}("
339 f"id={self.id}, "
340 f"ttl={self.ttl}, "
341 f"creation_time={self.creation_time}, "
342 f"expires_at={expiry_time}, "
343 f"remaining={remaining:.1f}s, "
344 f"expired={self.is_expired()}"
345 f")"
346 )
347
348 def __eq__(self, other) -> bool:
349 """
350 Two NodeFailingOverNotification notifications are considered equal if they have the same
351 id and are of the same type.
352 """
353 if not isinstance(other, NodeFailingOverNotification):
354 return False
355 return self.id == other.id and type(self) is type(other)
356
357 def __hash__(self) -> int:
358 """
359 Return a hash value for the notification to allow
360 instances to be used in sets and as dictionary keys.
361
362 Returns:
363 int: Hash value based on notification type and id
364 """
365 return hash((self.__class__.__name__, int(self.id)))
366
367
368class NodeFailedOverNotification(MaintenanceNotification):
369 """
370 Notification for when a Redis cluster node has completed a failover.
371
372 This notification is received when a node has finished the failover process
373 during cluster maintenance operations or after handling node failures.
374
375 Args:
376 id (int): Unique identifier for this notification
377 """
378
379 DEFAULT_TTL = 5
380
381 def __init__(self, id: int):
382 super().__init__(id, NodeFailedOverNotification.DEFAULT_TTL)
383
384 def __repr__(self) -> str:
385 expiry_time = self.creation_time + self.ttl
386 remaining = max(0, expiry_time - time.monotonic())
387 return (
388 f"{self.__class__.__name__}("
389 f"id={self.id}, "
390 f"ttl={self.ttl}, "
391 f"creation_time={self.creation_time}, "
392 f"expires_at={expiry_time}, "
393 f"remaining={remaining:.1f}s, "
394 f"expired={self.is_expired()}"
395 f")"
396 )
397
398 def __eq__(self, other) -> bool:
399 """
400 Two NodeFailedOverNotification notifications are considered equal if they have the same
401 id and are of the same type.
402 """
403 if not isinstance(other, NodeFailedOverNotification):
404 return False
405 return self.id == other.id and type(self) is type(other)
406
407 def __hash__(self) -> int:
408 """
409 Return a hash value for the notification to allow
410 instances to be used in sets and as dictionary keys.
411
412 Returns:
413 int: Hash value based on notification type and id
414 """
415 return hash((self.__class__.__name__, int(self.id)))
416
417
418class OSSNodeMigratingNotification(MaintenanceNotification):
419 """
420 Notification for when a Redis OSS API client is used and a node is in the process of migrating slots.
421
422 This notification is received when a node starts migrating its slots to another node
423 during cluster rebalancing or maintenance operations.
424
425 Args:
426 id (int): Unique identifier for this notification
427 slots (Optional[List[int]]): List of slots being migrated
428 """
429
430 DEFAULT_TTL = 30
431
432 def __init__(
433 self,
434 id: int,
435 slots: Optional[str] = None,
436 ):
437 super().__init__(id, OSSNodeMigratingNotification.DEFAULT_TTL)
438 self.slots = slots
439
440 def __repr__(self) -> str:
441 expiry_time = self.creation_time + self.ttl
442 remaining = max(0, expiry_time - time.monotonic())
443 return (
444 f"{self.__class__.__name__}("
445 f"id={self.id}, "
446 f"slots={self.slots}, "
447 f"ttl={self.ttl}, "
448 f"creation_time={self.creation_time}, "
449 f"expires_at={expiry_time}, "
450 f"remaining={remaining:.1f}s, "
451 f"expired={self.is_expired()}"
452 f")"
453 )
454
455 def __eq__(self, other) -> bool:
456 """
457 Two OSSNodeMigratingNotification notifications are considered equal if they have the same
458 id and are of the same type.
459 """
460 if not isinstance(other, OSSNodeMigratingNotification):
461 return False
462 return self.id == other.id and type(self) is type(other)
463
464 def __hash__(self) -> int:
465 """
466 Return a hash value for the notification to allow
467 instances to be used in sets and as dictionary keys.
468
469 Returns:
470 int: Hash value based on notification type and id
471 """
472 return hash((self.__class__.__name__, int(self.id)))
473
474
475class OSSNodeMigratedNotification(MaintenanceNotification):
476 """
477 Notification for when a Redis OSS API client is used and a node has completed migrating slots.
478
479 This notification is received when a node has finished migrating all its slots
480 to other nodes during cluster rebalancing or maintenance operations.
481
482 Args:
483 id (int): Unique identifier for this notification
484 nodes_to_slots_mapping (Dict[str, List[Dict[str, str]]]): Map of source node address
485 to list of destination mappings. Each destination mapping is a dict with
486 the destination node address as key and the slot range as value.
487
488 Structure example:
489 {
490 "127.0.0.1:6379": [
491 {"127.0.0.1:6380": "1-100"},
492 {"127.0.0.1:6381": "101-200"}
493 ],
494 "127.0.0.1:6382": [
495 {"127.0.0.1:6383": "201-300"}
496 ]
497 }
498
499 Where:
500 - Key (str): Source node address in "host:port" format
501 - Value (List[Dict[str, str]]): List of destination mappings where each dict
502 contains destination node address as key and slot range as value
503 """
504
505 DEFAULT_TTL = 120
506
507 def __init__(
508 self,
509 id: int,
510 nodes_to_slots_mapping: Dict[str, List[Dict[str, str]]],
511 ):
512 super().__init__(id, OSSNodeMigratedNotification.DEFAULT_TTL)
513 self.nodes_to_slots_mapping = nodes_to_slots_mapping
514
515 def __repr__(self) -> str:
516 expiry_time = self.creation_time + self.ttl
517 remaining = max(0, expiry_time - time.monotonic())
518 return (
519 f"{self.__class__.__name__}("
520 f"id={self.id}, "
521 f"nodes_to_slots_mapping={self.nodes_to_slots_mapping}, "
522 f"ttl={self.ttl}, "
523 f"creation_time={self.creation_time}, "
524 f"expires_at={expiry_time}, "
525 f"remaining={remaining:.1f}s, "
526 f"expired={self.is_expired()}"
527 f")"
528 )
529
530 def __eq__(self, other) -> bool:
531 """
532 Two OSSNodeMigratedNotification notifications are considered equal if they have the same
533 id and are of the same type.
534 """
535 if not isinstance(other, OSSNodeMigratedNotification):
536 return False
537 return self.id == other.id and type(self) is type(other)
538
539 def __hash__(self) -> int:
540 """
541 Return a hash value for the notification to allow
542 instances to be used in sets and as dictionary keys.
543
544 Returns:
545 int: Hash value based on notification type and id
546 """
547 return hash((self.__class__.__name__, int(self.id)))
548
549
550def _is_private_fqdn(host: str) -> bool:
551 """
552 Determine if an FQDN is likely to be internal/private.
553
554 This uses heuristics based on RFC 952 and RFC 1123 standards:
555 - .local domains (RFC 6762 - Multicast DNS)
556 - .internal domains (common internal convention)
557 - Single-label hostnames (no dots)
558 - Common internal TLDs
559
560 Args:
561 host (str): The FQDN to check
562
563 Returns:
564 bool: True if the FQDN appears to be internal/private
565 """
566 host_lower = host.lower().rstrip(".")
567
568 # Single-label hostnames (no dots) are typically internal
569 if "." not in host_lower:
570 return True
571
572 # Common internal/private domain patterns
573 internal_patterns = [
574 r"\.local$", # mDNS/Bonjour domains
575 r"\.internal$", # Common internal convention
576 r"\.corp$", # Corporate domains
577 r"\.lan$", # Local area network
578 r"\.intranet$", # Intranet domains
579 r"\.private$", # Private domains
580 ]
581
582 for pattern in internal_patterns:
583 if re.search(pattern, host_lower):
584 return True
585
586 # If none of the internal patterns match, assume it's external
587 return False
588
589
590notification_types_mapping: dict[type[MaintenanceNotification], str] = {
591 NodeMovingNotification: "MOVING",
592 NodeMigratingNotification: "MIGRATING",
593 NodeMigratedNotification: "MIGRATED",
594 NodeFailingOverNotification: "FAILING_OVER",
595 NodeFailedOverNotification: "FAILED_OVER",
596 OSSNodeMigratingNotification: "SMIGRATING",
597 OSSNodeMigratedNotification: "SMIGRATED",
598}
599
600
601def add_debug_log_for_notification(
602 connection: "MaintNotificationsAbstractConnection",
603 notification: Union[str, MaintenanceNotification],
604):
605 if logger.isEnabledFor(logging.DEBUG):
606 logger.debug(
607 f"Handling maintenance notification: {notification}, "
608 f"with connection: {connection}, "
609 f"{connection.extract_connection_details() if connection else 'no connection'}",
610 )
611
612
613class MaintNotificationsConfig:
614 """
615 Configuration class for maintenance notifications handling behaviour. Notifications are received through
616 push notifications.
617
618 This class defines how the Redis client should react to different push notifications
619 such as node moving, migrations, etc. in a Redis cluster.
620
621 """
622
623 def __init__(
624 self,
625 enabled: Union[bool, Literal["auto"]] = "auto",
626 proactive_reconnect: bool = True,
627 relaxed_timeout: Optional[Number] = 10,
628 endpoint_type: Optional[EndpointType] = None,
629 ):
630 """
631 Initialize a new MaintNotificationsConfig.
632
633 Args:
634 enabled (bool | "auto"): Controls maintenance notifications handling behavior.
635 - True: The CLIENT MAINT_NOTIFICATIONS command must succeed during connection setup,
636 otherwise a ResponseError is raised.
637 - "auto": The CLIENT MAINT_NOTIFICATIONS command is attempted but failures are
638 gracefully handled - a warning is logged and normal operation continues.
639 - False: Maintenance notifications are completely disabled.
640 Defaults to "auto".
641 proactive_reconnect (bool): Whether to proactively reconnect when a node is replaced.
642 Defaults to True.
643 relaxed_timeout (Number): The relaxed timeout to use for the connection during maintenance.
644 If -1 is provided - the relaxed timeout is disabled. If None is provided - the
645 affected operations become blocking. Defaults to 10.
646 endpoint_type (Optional[EndpointType]): Override for the endpoint type to use in CLIENT MAINT_NOTIFICATIONS.
647 If None, the endpoint type will be automatically determined based on the host and TLS configuration.
648 Defaults to None.
649
650 Raises:
651 ValueError: If endpoint_type is provided but is not a valid endpoint type.
652 """
653 self.enabled = enabled
654 self.relaxed_timeout = relaxed_timeout
655 self.proactive_reconnect = proactive_reconnect
656 self.endpoint_type = endpoint_type
657
658 def __repr__(self) -> str:
659 return (
660 f"{self.__class__.__name__}("
661 f"enabled={self.enabled}, "
662 f"proactive_reconnect={self.proactive_reconnect}, "
663 f"relaxed_timeout={self.relaxed_timeout}, "
664 f"endpoint_type={self.endpoint_type!r}"
665 f")"
666 )
667
668 def is_relaxed_timeouts_enabled(self) -> bool:
669 """
670 Check if the relaxed_timeout is enabled. The '-1' value is used to disable the relaxed_timeout.
671 If relaxed_timeout is set to None, it will make the operation blocking
672 and waiting until any response is received.
673
674 Returns:
675 True if the relaxed_timeout is enabled, False otherwise.
676 """
677 return self.relaxed_timeout != -1
678
679 def get_endpoint_type(
680 self,
681 host: str,
682 connection: "MaintNotificationsAbstractConnection | AsyncMaintNotificationsAbstractConnection",
683 ) -> EndpointType:
684 """
685 Determine the appropriate endpoint type for CLIENT MAINT_NOTIFICATIONS command.
686
687 Logic:
688 1. If endpoint_type is explicitly set, use it
689 2. Otherwise, check the original host from connection.host:
690 - If host is an IP address, use it directly to determine internal-ip vs external-ip
691 - If host is an FQDN, get the resolved IP to determine internal-fqdn vs external-fqdn
692
693 Args:
694 host: User provided hostname to analyze
695 connection: The connection object to analyze for endpoint type determination
696
697 Returns:
698 """
699
700 # If endpoint_type is explicitly set, use it
701 if self.endpoint_type is not None:
702 return self.endpoint_type
703
704 # Check if the host is an IP address
705 try:
706 ip_addr = ipaddress.ip_address(host)
707 # Host is an IP address - use it directly
708 is_private = ip_addr.is_private
709 return EndpointType.INTERNAL_IP if is_private else EndpointType.EXTERNAL_IP
710 except ValueError:
711 # Host is an FQDN - need to check resolved IP to determine internal vs external
712 pass
713
714 # Host is an FQDN, get the resolved IP to determine if it's internal or external
715 resolved_ip = connection.get_resolved_ip()
716
717 if resolved_ip:
718 try:
719 ip_addr = ipaddress.ip_address(resolved_ip)
720 is_private = ip_addr.is_private
721 # Use FQDN types since the original host was an FQDN
722 return (
723 EndpointType.INTERNAL_FQDN
724 if is_private
725 else EndpointType.EXTERNAL_FQDN
726 )
727 except ValueError:
728 # This shouldn't happen since we got the IP from the socket, but fallback
729 pass
730
731 # Final fallback: use heuristics on the FQDN itself
732 is_private = _is_private_fqdn(host)
733 return EndpointType.INTERNAL_FQDN if is_private else EndpointType.EXTERNAL_FQDN
734
735
736_MAINTENANCE_START_NOTIFICATION_TYPES = (
737 NodeMigratingNotification,
738 NodeFailingOverNotification,
739 OSSNodeMigratingNotification,
740)
741_MAINTENANCE_COMPLETED_NOTIFICATION_TYPES = (
742 NodeMigratedNotification,
743 NodeFailedOverNotification,
744 OSSNodeMigratedNotification,
745)
746
747
748def _get_maintenance_notification_type(
749 notification: MaintenanceNotification,
750) -> Optional[int]:
751 if notification.__class__ in _MAINTENANCE_START_NOTIFICATION_TYPES:
752 return 1
753 if notification.__class__ in _MAINTENANCE_COMPLETED_NOTIFICATION_TYPES:
754 return 0
755 return None
756
757
758def _get_maintenance_notification_name(
759 notification: MaintenanceNotification,
760) -> str:
761 return notification_types_mapping.get(notification.__class__, "")
762
763
764def _should_skip_connection_timeout_update(
765 maintenance_state: MaintenanceState,
766 config: MaintNotificationsConfig,
767) -> bool:
768 return (
769 maintenance_state == MaintenanceState.MOVING
770 or not config.is_relaxed_timeouts_enabled()
771 )
772
773
774def _build_moving_connection_kwargs(
775 notification: NodeMovingNotification,
776 config: MaintNotificationsConfig,
777) -> dict[str, Any]:
778 kwargs: dict[str, Any] = {
779 "maintenance_state": MaintenanceState.MOVING,
780 "maintenance_notification_hash": hash(notification),
781 }
782 if notification.new_node_host is not None:
783 # the host is not updated if the new node host is None
784 # this happens when the MOVING push notification does not contain
785 # the new node host - in this case we only update the timeouts
786 kwargs["host"] = notification.new_node_host
787 if config.is_relaxed_timeouts_enabled():
788 kwargs.update(
789 {
790 "socket_timeout": config.relaxed_timeout,
791 "socket_connect_timeout": config.relaxed_timeout,
792 }
793 )
794 return kwargs
795
796
797def _build_moving_cleanup_connection_kwargs(
798 connection_kwargs: Mapping[str, Any],
799 notification_hash: int,
800) -> Optional[dict[str, Any]]:
801 # if the current maintenance_notification_hash in kwargs is not matching the notification
802 # it means there has been a new moving notification after this one
803 # and we don't need to revert the kwargs yet
804 if connection_kwargs.get("maintenance_notification_hash") != notification_hash:
805 return None
806
807 return {
808 "maintenance_state": MaintenanceState.NONE,
809 "maintenance_notification_hash": None,
810 "host": connection_kwargs.get("orig_host_address"),
811 "socket_timeout": connection_kwargs.get("orig_socket_timeout"),
812 "socket_connect_timeout": connection_kwargs.get("orig_socket_connect_timeout"),
813 }
814
815
816class MaintNotificationsPoolHandler:
817 def __init__(
818 self,
819 pool: "MaintNotificationsAbstractConnectionPool",
820 config: MaintNotificationsConfig,
821 ) -> None:
822 self.pool = pool
823 self.config = config
824 self._processed_notifications = set()
825 self._lock = threading.RLock()
826 self.connection = None
827
828 def set_connection(self, connection: "MaintNotificationsAbstractConnection"):
829 self.connection = connection
830
831 def get_handler_for_connection(self):
832 # Copy all data that should be shared between connections
833 # but each connection should have its own pool handler
834 # since each connection can be in a different state
835 copy = MaintNotificationsPoolHandler(self.pool, self.config)
836 copy._processed_notifications = self._processed_notifications
837 copy._lock = self._lock
838 copy.connection = None
839 return copy
840
841 def remove_expired_notifications(self):
842 with self._lock:
843 for notification in tuple(self._processed_notifications):
844 if notification.is_expired():
845 self._processed_notifications.remove(notification)
846
847 def handle_notification(self, notification: MaintenanceNotification):
848 self.remove_expired_notifications()
849
850 if isinstance(notification, NodeMovingNotification):
851 return self.handle_node_moving_notification(notification)
852 else:
853 logger.error(f"Unhandled notification type: {notification}")
854
855 def handle_node_moving_notification(self, notification: NodeMovingNotification):
856 if (
857 not self.config.proactive_reconnect
858 and not self.config.is_relaxed_timeouts_enabled()
859 ):
860 return
861
862 with self._lock:
863 if notification in self._processed_notifications:
864 # nothing to do in the connection pool handling
865 # the notification has already been handled or is expired
866 # just return
867 return
868
869 with self.pool._lock:
870 if logger.isEnabledFor(logging.DEBUG):
871 logger.debug(
872 f"Handling node MOVING notification: {notification}, "
873 f"with connection: {self.connection}, connected to ip "
874 f"{self.connection.get_resolved_ip() if self.connection else None}"
875 )
876 # Get the current connected address - if any
877 # This is the address that is being moved
878 # and we need to handle only connections
879 # connected to the same address
880 moving_address_src = (
881 self.connection.getpeername() if self.connection else None
882 )
883
884 if getattr(self.pool, "set_in_maintenance", False):
885 # Set pool in maintenance mode - executed only if
886 # BlockingConnectionPool is used
887 self.pool.set_in_maintenance(True)
888
889 # Update maintenance state, timeout and optionally host address
890 # connection settings for matching connections
891 self.pool.update_connections_settings(
892 state=MaintenanceState.MOVING,
893 maintenance_notification_hash=hash(notification),
894 relaxed_timeout=self.config.relaxed_timeout,
895 host_address=notification.new_node_host,
896 matching_address=moving_address_src,
897 matching_pattern="connected_address",
898 update_notification_hash=True,
899 include_free_connections=True,
900 )
901
902 if self.config.proactive_reconnect:
903 if notification.new_node_host is not None:
904 self.run_proactive_reconnect(moving_address_src)
905 else:
906 threading.Timer(
907 notification.ttl / 2,
908 self.run_proactive_reconnect,
909 args=(moving_address_src,),
910 ).start()
911
912 # Update config for new connections:
913 # Set state to MOVING
914 # update host
915 # if relax timeouts are enabled - update timeouts
916 self.pool.update_connection_kwargs(
917 **_build_moving_connection_kwargs(notification, self.config)
918 )
919
920 if getattr(self.pool, "set_in_maintenance", False):
921 self.pool.set_in_maintenance(False)
922
923 threading.Timer(
924 notification.ttl,
925 self.handle_node_moved_notification,
926 args=(notification,),
927 ).start()
928
929 record_connection_handoff(
930 pool_name=get_pool_name(self.pool),
931 )
932
933 self._processed_notifications.add(notification)
934
935 def run_proactive_reconnect(self, moving_address_src: Optional[str] = None):
936 """
937 Run proactive reconnect for the pool.
938 Active connections are marked for reconnect after they complete the current command.
939 Inactive connections are disconnected and will be connected on next use.
940 """
941 with self._lock:
942 with self.pool._lock:
943 # take care for the active connections in the pool
944 # mark them for reconnect after they complete the current command
945 self.pool.update_active_connections_for_reconnect(
946 moving_address_src=moving_address_src,
947 )
948 # take care for the inactive connections in the pool
949 # delete them and create new ones
950 self.pool.disconnect_free_connections(
951 moving_address_src=moving_address_src,
952 )
953
954 def handle_node_moved_notification(self, notification: NodeMovingNotification):
955 """
956 Handle the cleanup after a node moving notification expires.
957 """
958 notification_hash = hash(notification)
959
960 with self._lock:
961 if logger.isEnabledFor(logging.DEBUG):
962 logger.debug(
963 f"Reverting temporary changes related to notification: {notification}, "
964 f"with connection: {self.connection}, connected to ip "
965 f"{self.connection.get_resolved_ip() if self.connection else None}"
966 )
967 kwargs = _build_moving_cleanup_connection_kwargs(
968 self.pool.connection_kwargs, notification_hash
969 )
970 if kwargs is not None:
971 self.pool.update_connection_kwargs(**kwargs)
972
973 with self.pool._lock:
974 reset_relaxed_timeout = self.config.is_relaxed_timeouts_enabled()
975 reset_host_address = self.config.proactive_reconnect
976
977 self.pool.update_connections_settings(
978 relaxed_timeout=-1,
979 state=MaintenanceState.NONE,
980 maintenance_notification_hash=None,
981 matching_notification_hash=notification_hash,
982 matching_pattern="notification_hash",
983 update_notification_hash=True,
984 reset_relaxed_timeout=reset_relaxed_timeout,
985 reset_host_address=reset_host_address,
986 include_free_connections=True,
987 )
988
989
990class MaintNotificationsConnectionHandler:
991 # 1 = "starting maintenance" notifications, 0 = "completed maintenance" notifications
992 _NOTIFICATION_TYPES: dict[type["MaintenanceNotification"], int] = {
993 NodeMigratingNotification: 1,
994 NodeFailingOverNotification: 1,
995 OSSNodeMigratingNotification: 1,
996 NodeMigratedNotification: 0,
997 NodeFailedOverNotification: 0,
998 OSSNodeMigratedNotification: 0,
999 }
1000
1001 def __init__(
1002 self,
1003 connection: "MaintNotificationsAbstractConnection",
1004 config: MaintNotificationsConfig,
1005 ) -> None:
1006 self.connection = connection
1007 self.config = config
1008
1009 def _get_pool_name(self) -> str:
1010 """
1011 Get the pool name from the connection's pool handler.
1012 Falls back to connection representation if pool is not available.
1013 """
1014 pool_handler = getattr(
1015 self.connection, "_maint_notifications_pool_handler", None
1016 )
1017 if pool_handler and getattr(pool_handler, "pool", None):
1018 return get_pool_name(pool_handler.pool)
1019 # Fallback for standalone connections without a pool
1020 return repr(self.connection)
1021
1022 def handle_notification(self, notification: MaintenanceNotification):
1023 # get the notification type by checking its class
1024 # 1 for start, 0 for end notification type, None for unknown
1025 notification_type = _get_maintenance_notification_type(notification)
1026 maint_notification = _get_maintenance_notification_name(notification)
1027
1028 record_maint_notification_count(
1029 server_address=self.connection.host,
1030 server_port=self.connection.port,
1031 network_peer_address=self.connection.host,
1032 network_peer_port=self.connection.port,
1033 maint_notification=maint_notification,
1034 )
1035
1036 if notification_type is None:
1037 logger.error(f"Unhandled notification type: {notification}")
1038 return
1039
1040 if notification_type:
1041 self.handle_maintenance_start_notification(
1042 MaintenanceState.MAINTENANCE, notification
1043 )
1044 else:
1045 self.handle_maintenance_completed_notification(notification=notification)
1046
1047 def handle_maintenance_start_notification(
1048 self, maintenance_state: MaintenanceState, notification: MaintenanceNotification
1049 ):
1050 add_debug_log_for_notification(self.connection, notification)
1051
1052 if _should_skip_connection_timeout_update(
1053 self.connection.maintenance_state, self.config
1054 ):
1055 return
1056
1057 self.connection.maintenance_state = maintenance_state
1058 self.connection.set_tmp_settings(
1059 tmp_relaxed_timeout=self.config.relaxed_timeout
1060 )
1061 # extend the timeout for all created connections
1062 self.connection.update_current_socket_timeout(self.config.relaxed_timeout)
1063 if isinstance(notification, OSSNodeMigratingNotification):
1064 # add the notification id to the set of processed start maint notifications
1065 # this is used to skip the unrelaxing of the timeouts if we have received more than
1066 # one start notification before the final end notification
1067 self.connection.add_maint_start_notification(notification.id)
1068
1069 maint_notification = _get_maintenance_notification_name(notification)
1070 record_connection_relaxed_timeout(
1071 connection_name=self._get_pool_name(),
1072 maint_notification=maint_notification,
1073 relaxed=True,
1074 )
1075
1076 def handle_maintenance_completed_notification(self, **kwargs: Any) -> None:
1077 # Only reset timeouts if state is not MOVING and relaxed timeouts are enabled
1078 if _should_skip_connection_timeout_update(
1079 self.connection.maintenance_state, self.config
1080 ):
1081 return
1082
1083 notification = None
1084 if kwargs.get("notification"):
1085 notification = kwargs["notification"]
1086 add_debug_log_for_notification(
1087 self.connection, notification if notification else "MAINTENANCE_COMPLETED"
1088 )
1089 self.connection.reset_tmp_settings(reset_relaxed_timeout=True)
1090 # Maintenance completed - reset the connection
1091 # timeouts by providing -1 as the relaxed timeout
1092 self.connection.update_current_socket_timeout(-1)
1093 self.connection.maintenance_state = MaintenanceState.NONE
1094 # reset the sets that keep track of received start maint
1095 # notifications and skipped end maint notifications
1096 self.connection.reset_received_notifications()
1097
1098 if notification:
1099 maint_notification = _get_maintenance_notification_name(notification)
1100 record_connection_relaxed_timeout(
1101 connection_name=self._get_pool_name(),
1102 maint_notification=maint_notification,
1103 relaxed=False,
1104 )
1105
1106
1107class OSSMaintNotificationsHandler:
1108 def __init__(
1109 self,
1110 cluster_client: "MaintNotificationsAbstractRedisCluster",
1111 config: MaintNotificationsConfig,
1112 ) -> None:
1113 self.cluster_client = cluster_client
1114 self.config = config
1115 self._processed_notifications = set()
1116 self._in_progress = set()
1117 self._lock = threading.RLock()
1118
1119 def remove_expired_notifications(self):
1120 with self._lock:
1121 for notification in tuple(self._processed_notifications):
1122 if notification.is_expired():
1123 self._processed_notifications.remove(notification)
1124
1125 def handle_notification(self, notification: MaintenanceNotification):
1126 if isinstance(notification, OSSNodeMigratedNotification):
1127 self.handle_oss_maintenance_completed_notification(notification)
1128 else:
1129 logger.error(f"Unhandled notification type: {notification}")
1130
1131 def handle_oss_maintenance_completed_notification(
1132 self, notification: OSSNodeMigratedNotification
1133 ):
1134 self.remove_expired_notifications()
1135
1136 with self._lock:
1137 if (
1138 notification in self._in_progress
1139 or notification in self._processed_notifications
1140 ):
1141 # we are already handling this notification or it has already been processed
1142 # we should skip in_progress notification since when we reinitialize the cluster
1143 # we execute a CLUSTER SLOTS command that can use a different connection
1144 # that has also has the notification and we don't want to
1145 # process the same notification twice
1146 #
1147 # This cheap pre-check runs under _lock alone, before the
1148 # topology refresh lock below. The server delivers the same
1149 # SMIGRATED on every connection, so the vast majority of
1150 # arrivals land here; making them wait for an in-flight CLUSTER
1151 # SLOTS round trip just to discard a notification would stall
1152 # arbitrary command threads inside read_response.
1153 return
1154
1155 # Lock order: _initialization_lock BEFORE _lock. initialize() below runs a
1156 # CLUSTER SLOTS round trip while holding _initialization_lock, and the
1157 # response can carry another SMIGRATED push that is handled inline on that
1158 # thread and needs _lock - so a thread holding _lock must never wait for
1159 # _initialization_lock. See the _initialization_lock comment in
1160 # NodesManager.__init__ for the full ordering rule.
1161 with self.cluster_client.nodes_manager._initialization_lock, self._lock:
1162 if (
1163 notification in self._in_progress
1164 or notification in self._processed_notifications
1165 ):
1166 # Re-check now that both locks are held: another thread may have
1167 # handled the notification while we waited for the refresh lock.
1168 return
1169
1170 if logger.isEnabledFor(logging.DEBUG):
1171 logger.debug(f"Handling SMIGRATED notification: {notification}")
1172 self._in_progress.add(notification)
1173
1174 try:
1175 # Extract the information about the src and destination nodes that are affected
1176 # by the maintenance. nodes_to_slots_mapping structure:
1177 # {
1178 # "src_host:port": [
1179 # {"dest_host:port": "slot_range"},
1180 # ...
1181 # ],
1182 # ...
1183 # }
1184 additional_startup_nodes_info = []
1185 affected_nodes = set()
1186 for (
1187 src_address,
1188 dest_mappings,
1189 ) in notification.nodes_to_slots_mapping.items():
1190 src_host, src_port = src_address.rsplit(":", 1)
1191 src_node = self.cluster_client.nodes_manager.get_node(
1192 host=src_host, port=int(src_port)
1193 )
1194 if src_node is not None:
1195 affected_nodes.add(src_node)
1196
1197 for dest_mapping in dest_mappings:
1198 for dest_address in dest_mapping.keys():
1199 dest_host, dest_port = dest_address.rsplit(":", 1)
1200 additional_startup_nodes_info.append(
1201 (dest_host, int(dest_port))
1202 )
1203
1204 # Updates the cluster slots cache with the new slots mapping
1205 # This will also update the nodes cache with the new nodes mapping
1206 self.cluster_client.nodes_manager.initialize(
1207 disconnect_startup_nodes_pools=False,
1208 additional_startup_nodes_info=additional_startup_nodes_info,
1209 )
1210
1211 all_nodes = set(affected_nodes)
1212 all_nodes = all_nodes.union(
1213 self.cluster_client.nodes_manager.nodes_cache.values()
1214 )
1215
1216 for current_node in all_nodes:
1217 if current_node.redis_connection is None:
1218 continue
1219 with current_node.redis_connection.connection_pool._lock:
1220 handoff_recorded = False
1221 if current_node in affected_nodes:
1222 # mark for reconnect all in use connections to the node - this will force them to
1223 # disconnect after they complete their current commands
1224 # Some of them might be used by sub sub and we don't know which ones - so we disconnect
1225 # all in flight connections after they are done with current command execution
1226 for conn in current_node.redis_connection.connection_pool._get_in_use_connections():
1227 add_debug_log_for_notification(
1228 conn, "SMIGRATED - mark for reconnect"
1229 )
1230 conn.mark_for_reconnect()
1231
1232 record_connection_handoff(
1233 pool_name=get_pool_name(
1234 current_node.redis_connection.connection_pool
1235 )
1236 )
1237 handoff_recorded = True
1238 else:
1239 if logger.isEnabledFor(logging.DEBUG):
1240 logger.debug(
1241 f"SMIGRATED: Node {current_node.name} not affected by maintenance, "
1242 f"skipping mark for reconnect"
1243 )
1244
1245 if (
1246 current_node
1247 not in self.cluster_client.nodes_manager.nodes_cache.values()
1248 ):
1249 # disconnect all free connections to the node - this node will be dropped
1250 # from the cluster, so we don't need to revert the timeouts
1251 for conn in current_node.redis_connection.connection_pool._get_free_connections():
1252 conn.disconnect()
1253
1254 # Only record handoff if not already recorded for this node
1255 if not handoff_recorded:
1256 record_connection_handoff(
1257 pool_name=get_pool_name(
1258 current_node.redis_connection.connection_pool
1259 )
1260 )
1261
1262 # mark the notification as processed
1263 self._processed_notifications.add(notification)
1264 finally:
1265 # Release the in-progress reservation. On success the notification
1266 # is also in _processed_notifications (so it won't be re-handled);
1267 # on failure it is not, allowing a later redelivery to retry.
1268 self._in_progress.discard(notification)