from __future__ import annotations import asyncio import time from dataclasses import dataclass, field from typing import Any @dataclass class AuthHealth: status: str message: str = "" expires_at: float | None = None last_checked_at: float = field(default_factory=time.monotonic) metadata: dict[str, Any] = field(default_factory=dict) @property def is_healthy(self) -> bool: return self.status == "healthy" class AuthHealthMonitor: def __init__(self): self._health_states: dict[str, AuthHealth] = {} self._check_tasks: dict[str, asyncio.Task] = {} self._on_expiring_callbacks: list = [] self._on_expired_callbacks: list = [] self._on_revoked_callbacks: list = [] self._on_failure_callbacks: list = [] self._locks: dict[str, asyncio.Lock] = {} def _get_lock(self, channel_id: str) -> asyncio.Lock: if channel_id not in self._locks: self._locks[channel_id] = asyncio.Lock() return self._locks[channel_id] def update_health(self, channel_id: str, health: AuthHealth) -> None: old = self._health_states.get(channel_id) self._health_states[channel_id] = health if old and old.status == "healthy" and health.status != "healthy": try: asyncio.create_task(self._notify_failure(channel_id, health)) except RuntimeError: pass def get_health(self, channel_id: str) -> AuthHealth | None: return self._health_states.get(channel_id) def get_all_health(self) -> dict[str, AuthHealth]: return dict(self._health_states) def on_token_expiring(self, callback) -> None: self._on_expiring_callbacks.append(callback) def on_token_expired(self, callback) -> None: self._on_expired_callbacks.append(callback) def on_auth_revoked(self, callback) -> None: self._on_revoked_callbacks.append(callback) def on_auth_failure(self, callback) -> None: self._on_failure_callbacks.append(callback) async def _notify_expiring(self, channel_id: str) -> None: for cb in self._on_expiring_callbacks: try: await cb(channel_id) except Exception: pass async def _notify_expired(self, channel_id: str) -> None: for cb in self._on_expired_callbacks: try: await cb(channel_id) except Exception: pass async def _notify_revoked(self, channel_id: str) -> None: for cb in self._on_revoked_callbacks: try: await cb(channel_id) except Exception: pass async def _notify_failure(self, channel_id: str, health: AuthHealth) -> None: for cb in self._on_failure_callbacks: try: await cb(channel_id, health) except Exception: pass def start_periodic_check( self, channel_id: str, check_func, interval_seconds: float = 60.0, ) -> None: if channel_id in self._check_tasks and not self._check_tasks[channel_id].done(): return self._check_tasks[channel_id] = asyncio.create_task( self._periodic_check_loop(channel_id, check_func, interval_seconds) ) def stop_periodic_check(self, channel_id: str) -> None: task = self._check_tasks.pop(channel_id, None) if task and not task.done(): task.cancel() async def _periodic_check_loop( self, channel_id: str, check_func, interval_seconds: float, ) -> None: from yuxi.utils.logging_config import logger while True: try: await asyncio.sleep(interval_seconds) health = await check_func() self.update_health(channel_id, health) except asyncio.CancelledError: break except Exception as e: logger.warning(f"[AuthHealthMonitor] Health check failed for {channel_id}: {e}") self.update_health( channel_id, AuthHealth(status="unhealthy", message=str(e)), ) _health_monitor: AuthHealthMonitor | None = None def get_health_monitor() -> AuthHealthMonitor: global _health_monitor if _health_monitor is None: _health_monitor = AuthHealthMonitor() return _health_monitor