1import asyncio
2import functools
3import logging
4from typing import TYPE_CHECKING, Any, Awaitable, Callable
5
6from redis.asyncio.observability.recorder import (
7 record_connection_handoff,
8 record_connection_relaxed_timeout,
9 record_maint_notification_count,
10)
11from redis.maint_notifications import (
12 MaintenanceNotification,
13 MaintenanceState,
14 MaintNotificationsConfig,
15 NodeMovingNotification,
16 OSSNodeMigratedNotification,
17 OSSNodeMigratingNotification,
18 _get_maintenance_notification_name,
19 _get_maintenance_notification_type,
20 _should_skip_connection_timeout_update,
21)
22from redis.observability.attributes import get_pool_name
23
24if TYPE_CHECKING:
25 from redis.asyncio.cluster import RedisCluster
26 from redis.asyncio.connection import AsyncMaintNotificationsAbstractConnection
27
28logger = logging.getLogger(__name__)
29
30_ScheduledCallback = Callable[..., Awaitable[None]]
31
32
33def _log_task_exception(message: str, task: "asyncio.Task[Any]") -> None:
34 """Task done-callback that surfaces a failed task's exception in the logs.
35
36 Without it, an unhandled exception in a fire-and-forget task is only
37 reported by asyncio as a noisy "Task exception was never retrieved"
38 warning. Retrieving the exception here suppresses that warning and logs a
39 meaningful error instead. Cancellation is expected and left unlogged.
40
41 Bind ``message`` with ``functools.partial`` before passing to
42 ``add_done_callback`` (which supplies ``task``).
43 """
44 try:
45 exc = task.exception()
46 except asyncio.CancelledError:
47 return
48 if exc:
49 logger.error(message, exc_info=exc)
50
51
52def add_debug_log_for_notification(
53 connection: object,
54 notification: str | MaintenanceNotification,
55) -> None:
56 if not logger.isEnabledFor(logging.DEBUG):
57 return
58
59 details = "no connection"
60 try:
61 extract = getattr(connection, "extract_connection_details", None)
62 if callable(extract):
63 details = extract()
64 except (AttributeError, OSError, TypeError):
65 pass
66
67 logger.debug(
68 f"Handling maintenance notification: {notification}, "
69 f"with connection: {connection}, {details}",
70 )
71
72
73class AsyncMaintNotificationsPoolHandler:
74 def __init__(
75 self,
76 pool: Any,
77 config: MaintNotificationsConfig,
78 ) -> None:
79 self.pool = pool
80 self.config = config
81 self._processed_notifications: set[MaintenanceNotification] = set()
82 self._scheduled_tasks: set[asyncio.Task[None]] = set()
83 self._lock = asyncio.Lock()
84 self.connection: Any | None = None
85
86 def set_connection(
87 self, connection: "AsyncMaintNotificationsAbstractConnection"
88 ) -> None:
89 self.connection = connection
90
91 def get_handler_for_connection(self) -> "AsyncMaintNotificationsPoolHandler":
92 # Copy all data that should be shared between connections, while each
93 # connection gets its own handler instance and current connection state.
94 copy = AsyncMaintNotificationsPoolHandler(self.pool, self.config)
95 copy._processed_notifications = self._processed_notifications
96 copy._scheduled_tasks = self._scheduled_tasks
97 copy._lock = self._lock
98 copy.connection = None
99 return copy
100
101 async def remove_expired_notifications(self) -> None:
102 async with self._lock:
103 for notification in tuple(self._processed_notifications):
104 if notification.is_expired():
105 self._processed_notifications.remove(notification)
106
107 async def handle_notification(self, notification: MaintenanceNotification) -> None:
108 await self.remove_expired_notifications()
109
110 if isinstance(notification, NodeMovingNotification):
111 await self.handle_node_moving_notification(notification)
112 else:
113 logger.error(f"Unhandled notification type: {notification}")
114
115 async def handle_node_moving_notification(
116 self, notification: NodeMovingNotification
117 ) -> None:
118 if (
119 not self.config.proactive_reconnect
120 and not self.config.is_relaxed_timeouts_enabled()
121 ):
122 return
123
124 async with self._lock:
125 if notification in self._processed_notifications:
126 # nothing to do in the connection pool handling
127 # the notification has already been handled or is expired
128 # just return
129 return
130 if logger.isEnabledFor(logging.DEBUG):
131 logger.debug(
132 f"Handling node MOVING notification: {notification}, "
133 f"with connection: {self.connection}, connected to ip "
134 f"{self.connection.get_resolved_ip() if self.connection else None}"
135 )
136 # Get the current connected address - if any
137 # This is the address that is being moved
138 # and we need to handle only connections
139 # connected to the same address
140 moving_address_src = (
141 self.connection.getpeername() if self.connection else None
142 )
143
144 # The async pool owns the active/free connection collections and
145 # asyncio.Lock is not reentrant, so the whole MOVING pool mutation
146 # has to be one pool-owned atomic operation. The handler still owns
147 # the notification policy and passes the already-decided inputs.
148 await self.pool.apply_moving_notification(
149 notification=notification,
150 config=self.config,
151 moving_address_src=moving_address_src,
152 run_proactive_reconnect=(
153 self.config.proactive_reconnect
154 and notification.new_node_host is not None
155 ),
156 )
157
158 if self.config.proactive_reconnect and notification.new_node_host is None:
159 self._schedule(
160 notification.ttl / 2,
161 self.run_proactive_reconnect,
162 moving_address_src,
163 )
164
165 self._schedule(
166 notification.ttl,
167 self.handle_node_moved_notification,
168 notification,
169 )
170
171 await record_connection_handoff(
172 pool_name=get_pool_name(self.pool),
173 )
174
175 self._processed_notifications.add(notification)
176
177 async def run_proactive_reconnect(
178 self, moving_address_src: str | None = None
179 ) -> None:
180 """
181 Run proactive reconnect for the pool.
182 Active connections are marked for reconnect after they complete the current command.
183 Inactive connections are disconnected and will be connected on next use.
184 """
185 async with self._lock:
186 # This delayed reconnect must be atomic for the same reason as the
187 # initial MOVING mutation: a connection can move between active and
188 # free lists while another task is acquiring or releasing it.
189 await self.pool.run_proactive_reconnect(
190 moving_address_src=moving_address_src,
191 )
192
193 async def handle_node_moved_notification(
194 self, notification: NodeMovingNotification
195 ) -> None:
196 """
197 Handle the cleanup after a node moving notification expires.
198 """
199 notification_hash = hash(notification)
200
201 async with self._lock:
202 if logger.isEnabledFor(logging.DEBUG):
203 logger.debug(
204 f"Reverting temporary changes related to notification: {notification}, "
205 f"with connection: {self.connection}, connected to ip "
206 f"{self.connection.get_resolved_ip() if self.connection else None}"
207 )
208 reset_relaxed_timeout = self.config.is_relaxed_timeouts_enabled()
209 reset_host_address = self.config.proactive_reconnect
210
211 # Cleanup has to reset future connection kwargs and existing
212 # matching connections together under the pool lock. Splitting it
213 # lets an acquire/release interleave and leaves stale MOVING state.
214 await self.pool.cleanup_moving_notification(
215 notification_hash=notification_hash,
216 reset_relaxed_timeout=reset_relaxed_timeout,
217 reset_host_address=reset_host_address,
218 )
219
220 def _schedule(
221 self,
222 delay: float,
223 callback: _ScheduledCallback,
224 *args: Any,
225 ) -> None:
226 # Record the absolute deadline now so that any lag between create_task
227 # and the task's first execution slice does not push the fire time out.
228 deadline = asyncio.get_running_loop().time() + delay
229 task = asyncio.create_task(self._run_after(deadline, callback, *args))
230 self._scheduled_tasks.add(task)
231 task.add_done_callback(self._scheduled_tasks.discard)
232 task.add_done_callback(
233 functools.partial(
234 _log_task_exception,
235 "Error handling scheduled maintenance notification",
236 )
237 )
238
239 async def _run_after(
240 self,
241 deadline: float,
242 callback: _ScheduledCallback,
243 *args: Any,
244 ) -> None:
245 remaining = deadline - asyncio.get_running_loop().time()
246 if remaining > 0:
247 await asyncio.sleep(remaining)
248 await callback(*args)
249
250 async def cancel_scheduled_tasks(self) -> None:
251 if not self._scheduled_tasks:
252 return
253 tasks = tuple(self._scheduled_tasks)
254 for task in tasks:
255 task.cancel()
256 await asyncio.gather(*tasks, return_exceptions=True)
257
258
259class AsyncMaintNotificationsConnectionHandler:
260 def __init__(
261 self,
262 connection: "AsyncMaintNotificationsAbstractConnection",
263 config: MaintNotificationsConfig,
264 ) -> None:
265 self.connection = connection
266 self.config = config
267
268 def _get_pool_name(self) -> str:
269 """
270 Get the pool name from the connection's pool handler.
271 Falls back to connection representation if pool is not available.
272 """
273 pool_handler = getattr(
274 self.connection, "_maint_notifications_pool_handler", None
275 )
276 if pool_handler and getattr(pool_handler, "pool", None):
277 return get_pool_name(pool_handler.pool)
278 # Fallback for standalone connections without a pool
279 return repr(self.connection)
280
281 async def handle_notification(self, notification: MaintenanceNotification) -> None:
282 # 1 for start, 0 for end notification type, None for unknown.
283 notification_type = _get_maintenance_notification_type(notification)
284 maint_notification = _get_maintenance_notification_name(notification)
285
286 await record_maint_notification_count(
287 server_address=self.connection.host,
288 server_port=self.connection.port,
289 network_peer_address=self.connection.host,
290 network_peer_port=self.connection.port,
291 maint_notification=maint_notification,
292 )
293
294 if notification_type is None:
295 logger.error(f"Unhandled notification type: {notification}")
296 return
297
298 if notification_type:
299 await self.handle_maintenance_start_notification(
300 MaintenanceState.MAINTENANCE, notification
301 )
302 else:
303 await self.handle_maintenance_completed_notification(
304 notification=notification
305 )
306
307 async def handle_maintenance_start_notification(
308 self,
309 maintenance_state: MaintenanceState,
310 notification: MaintenanceNotification,
311 ) -> None:
312 add_debug_log_for_notification(self.connection, notification)
313
314 if _should_skip_connection_timeout_update(
315 self.connection.maintenance_state, self.config
316 ):
317 return
318
319 self.connection.maintenance_state = maintenance_state
320 self.connection.set_tmp_settings(
321 tmp_relaxed_timeout=self.config.relaxed_timeout
322 )
323 self.connection.update_current_socket_timeout(self.config.relaxed_timeout)
324 if isinstance(notification, OSSNodeMigratingNotification):
325 # add the notification id to the set of processed start maint notifications
326 # this is used to skip the unrelaxing of the timeouts if we have received more than
327 # one start notification before the final end notification
328 self.connection.add_maint_start_notification(notification.id)
329
330 maint_notification = _get_maintenance_notification_name(notification)
331 await record_connection_relaxed_timeout(
332 connection_name=self._get_pool_name(),
333 maint_notification=maint_notification,
334 relaxed=True,
335 )
336
337 async def handle_maintenance_completed_notification(self, **kwargs: Any) -> None:
338 # Only reset timeouts if state is not MOVING and relaxed timeouts are enabled
339 if _should_skip_connection_timeout_update(
340 self.connection.maintenance_state, self.config
341 ):
342 return
343
344 notification = None
345 if kwargs.get("notification"):
346 notification = kwargs["notification"]
347 add_debug_log_for_notification(
348 self.connection, notification if notification else "MAINTENANCE_COMPLETED"
349 )
350 self.connection.reset_tmp_settings(reset_relaxed_timeout=True)
351 # Maintenance completed - reset the connection
352 # timeouts by providing -1 as the relaxed timeout
353 self.connection.update_current_socket_timeout(-1)
354 self.connection.maintenance_state = MaintenanceState.NONE
355 # reset the sets that keep track of received start maint
356 # notifications and skipped end maint notifications
357 self.connection.reset_received_notifications()
358
359 if notification:
360 maint_notification = _get_maintenance_notification_name(notification)
361 await record_connection_relaxed_timeout(
362 connection_name=self._get_pool_name(),
363 maint_notification=maint_notification,
364 relaxed=False,
365 )
366
367
368class AsyncOSSMaintNotificationsHandler:
369 """
370 Cluster-wide handler for OSS (open-source) cluster maintenance notifications.
371
372 Reacts to SMIGRATED (slot migration completed) push notifications and
373 triggers topology re-initialization via nodes_manager.initialize().
374
375 Lock discipline: _lock is held across the whole pool mutation, including
376 await initialize(), so the topology refresh and the subsequent connection
377 marking/disconnect happen as one atomic operation — releasing the lock
378 mid-mutation would let other tasks observe partial state.
379
380 Re-entrancy is not a deadlock risk here even though asyncio.Lock is
381 non-reentrant: initialize() may dispatch commands whose responses carry new
382 push notifications, but handle_notification schedules the actual handling as
383 a separate background task rather than calling into it inline, so the
384 re-entrant arrival never tries to re-acquire the lock on this call stack.
385 The cheap _in_progress/_processed dedup that gates that scheduling runs
386 without the lock — those sets are only mutated from the single event loop.
387
388 For the same reason the sync handler's lock-ordering rule (acquire
389 NodesManager._initialization_lock before the handler's _lock) does not apply
390 here: because the handling never runs inline on the call stack that is inside
391 initialize(), no task ever holds NodesManager._initialize_lock while waiting
392 for this _lock, so the two can be acquired in this order safely. The
393 divergence from the sync handler is deliberate, not drift.
394 """
395
396 def __init__(
397 self,
398 cluster_client: "RedisCluster",
399 config: MaintNotificationsConfig,
400 ) -> None:
401 self.cluster_client = cluster_client
402 self.config = config
403 self._processed_notifications: set[MaintenanceNotification] = set()
404 self._in_progress: set[MaintenanceNotification] = set()
405 self._lock = asyncio.Lock()
406 self._background_tasks: set[asyncio.Task] = set()
407
408 async def remove_expired_notifications(self) -> None:
409 async with self._lock:
410 for n in tuple(self._processed_notifications):
411 if n.is_expired():
412 self._processed_notifications.remove(n)
413
414 async def handle_notification(self, notification: MaintenanceNotification) -> None:
415 # Synchronous pre-dedup BEFORE scheduling a task.
416 #
417 # The same SMIGRATED notification is delivered by the server on every
418 # connection, so under load this callback fires many times for the same
419 # notification. Without this guard we would schedule one background task
420 # per arrival - a flood of tasks that each acquire the lock and dedup,
421 # saturating the single event loop. Reserving the notification in
422 # _in_progress here (and skipping if already in-progress/processed) caps
423 # it at exactly ONE handling task per unique notification.
424 #
425 # These sets are only mutated from the single event loop, so this
426 # check-and-add is race-free without holding the lock.
427 if (
428 notification in self._in_progress
429 or notification in self._processed_notifications
430 ):
431 return
432 self._in_progress.add(notification)
433
434 # Schedule as a background task so the parser's read path is not blocked.
435 # This also breaks the inline call chain that would otherwise deadlock:
436 # initialize() dispatches commands whose responses may carry more push
437 # notifications; as a separate task it can run while this one awaits.
438 #
439 # If scheduling the task (or wiring its callbacks) fails, release the
440 # in-progress reservation made above. Otherwise the notification would be
441 # stuck in _in_progress forever - the dedup guard would skip it on every
442 # future arrival, and _do_handle_notification's finally (which normally
443 # clears it) never runs because the task was never started.
444 try:
445 task = asyncio.get_running_loop().create_task(
446 self._do_handle_notification(notification)
447 )
448 self._background_tasks.add(task)
449 task.add_done_callback(self._background_tasks.discard)
450 task.add_done_callback(
451 functools.partial(
452 _log_task_exception,
453 "Error handling maintenance notification background task",
454 )
455 )
456 except Exception:
457 self._in_progress.discard(notification)
458 raise
459
460 async def _do_handle_notification(
461 self, notification: MaintenanceNotification
462 ) -> None:
463 try:
464 if isinstance(notification, OSSNodeMigratedNotification):
465 await self.handle_oss_maintenance_completed_notification(notification)
466 else:
467 logger.error(f"Unhandled notification type: {notification}")
468 finally:
469 # Release the in-progress reservation. On success the notification is
470 # also in _processed_notifications (so it won't be re-handled); on
471 # failure it is not, allowing a later retry.
472 self._in_progress.discard(notification)
473
474 async def handle_oss_maintenance_completed_notification(
475 self, notification: OSSNodeMigratedNotification
476 ) -> None:
477 await self.remove_expired_notifications()
478
479 async with self._lock:
480 # handle_notification already reserved this notification in
481 # _in_progress and guaranteed uniqueness; the processed check here is
482 # defensive (e.g. across handler copies sharing the same sets).
483 if notification in self._processed_notifications:
484 return
485 if logger.isEnabledFor(logging.DEBUG):
486 logger.debug(f"Handling SMIGRATED notification: {notification}")
487
488 # Extract the information about the src and destination nodes that are
489 # affected by the maintenance. nodes_to_slots_mapping structure:
490 # {
491 # "src_host:port": [
492 # {"dest_host:port": "slot_range"},
493 # ...
494 # ],
495 # ...
496 # }
497 additional_startup_nodes_info = []
498 affected_nodes = set()
499 for (
500 src_address,
501 dest_mappings,
502 ) in notification.nodes_to_slots_mapping.items():
503 src_host, src_port = src_address.rsplit(":", 1)
504 src_node = self.cluster_client.nodes_manager.get_node(
505 host=src_host, port=int(src_port)
506 )
507 if src_node is not None:
508 affected_nodes.add(src_node)
509 for dest_mapping in dest_mappings:
510 for dest_address in dest_mapping.keys():
511 dest_host, dest_port = dest_address.rsplit(":", 1)
512 additional_startup_nodes_info.append(
513 (dest_host, int(dest_port))
514 )
515 # Updates the cluster slots cache with the new slots mapping
516 # This will also update the nodes cache with the new nodes mapping
517 await self.cluster_client.nodes_manager.initialize(
518 additional_startup_nodes_info=additional_startup_nodes_info,
519 )
520
521 all_nodes = set(affected_nodes)
522 all_nodes = all_nodes.union(
523 self.cluster_client.nodes_manager.nodes_cache.values()
524 )
525 for current_node in all_nodes:
526 handoff_recorded = False
527 if current_node in affected_nodes:
528 # mark for reconnect all in-use connections to the node — this
529 # forces them to disconnect after completing their current command
530 free_set = set(current_node._free)
531 for conn in current_node._connections:
532 if conn not in free_set:
533 add_debug_log_for_notification(
534 conn, "SMIGRATED - mark for reconnect"
535 )
536 conn.mark_for_reconnect()
537 await record_connection_handoff(
538 pool_name=f"{current_node.host}:{current_node.port}"
539 )
540 handoff_recorded = True
541 else:
542 if logger.isEnabledFor(logging.DEBUG):
543 logger.debug(
544 f"SMIGRATED: Node {current_node.name} not affected "
545 f"by maintenance, skipping mark for reconnect"
546 )
547 if (
548 current_node
549 not in self.cluster_client.nodes_manager.nodes_cache.values()
550 ):
551 task = asyncio.get_running_loop().create_task(
552 current_node.disconnect_free_connections()
553 )
554 self._background_tasks.add(task)
555 task.add_done_callback(self._background_tasks.discard)
556 task.add_done_callback(
557 functools.partial(
558 _log_task_exception,
559 "Error disconnecting free connections after "
560 "maintenance notification",
561 )
562 )
563 if not handoff_recorded:
564 await record_connection_handoff(
565 pool_name=f"{current_node.host}:{current_node.port}"
566 )
567
568 self._processed_notifications.add(notification)