from __future__ import annotations import asyncio import time from dataclasses import dataclass, field from enum import Enum from typing import Any @dataclass class TokenState: token: str expires_at: float | None = None token_type: str = "bearer" scopes: list[str] = field(default_factory=list) source: str = "config" metadata: dict[str, Any] = field(default_factory=dict) @property def is_expired(self) -> bool: if self.expires_at is None: return False return time.monotonic() > self.expires_at - 300 @property def remaining_seconds(self) -> float | None: if self.expires_at is None: return None return max(0, self.expires_at - time.monotonic()) class TokenProviderType(Enum): STATIC = "static" OAUTH2 = "oauth2" QR_CODE = "qr_code" CERTIFICATE = "certificate" class BaseTokenProvider: provider_type: TokenProviderType async def get_token(self) -> TokenState: ... async def refresh_token(self) -> TokenState: ... async def validate_token(self, token: TokenState) -> bool: ... async def revoke_token(self) -> None: ... class UnifiedTokenManager: REFRESH_MARGIN_SECONDS = 300 def __init__(self): self._providers: dict[str, BaseTokenProvider] = {} self._cache: dict[str, TokenState] = {} self._locks: dict[str, asyncio.Lock] = {} self._refresh_events: dict[str, asyncio.Event] = {} self._refresh_tasks: dict[str, asyncio.Task] = {} self._on_token_expiring_callbacks: list = [] def register_provider(self, channel_id: str, provider: BaseTokenProvider) -> None: self._providers[channel_id] = provider def unregister_provider(self, channel_id: str) -> None: self._providers.pop(channel_id, None) self._cache.pop(channel_id, None) self._locks.pop(channel_id, None) self._refresh_events.pop(channel_id, None) task = self._refresh_tasks.pop(channel_id, None) if task and not task.done(): task.cancel() 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] async def get_token(self, channel_id: str) -> TokenState: provider = self._providers.get(channel_id) if not provider: raise ValueError(f"No token provider registered for channel '{channel_id}'") lock = self._get_lock(channel_id) async with lock: cached = self._cache.get(channel_id) if cached and not cached.is_expired: return cached if channel_id in self._refresh_events: await self._refresh_events[channel_id].wait() cached = self._cache.get(channel_id) if cached: return cached event = asyncio.Event() self._refresh_events[channel_id] = event try: token_state = await provider.get_token() self._cache[channel_id] = token_state return token_state finally: event.set() self._refresh_events.pop(channel_id, None) async def refresh_token(self, channel_id: str) -> TokenState: provider = self._providers.get(channel_id) if not provider: raise ValueError(f"No token provider registered for channel '{channel_id}'") lock = self._get_lock(channel_id) async with lock: token_state = await provider.refresh_token() self._cache[channel_id] = token_state return token_state async def invalidate_token(self, channel_id: str) -> None: lock = self._get_lock(channel_id) async with lock: self._cache.pop(channel_id, None) provider = self._providers.get(channel_id) if provider: await provider.revoke_token() def on_token_expiring(self, callback) -> None: self._on_token_expiring_callbacks.append(callback) async def _notify_token_expiring(self, channel_id: str, token_state: TokenState) -> None: for callback in self._on_token_expiring_callbacks: try: await callback(channel_id, token_state) except Exception: pass def start_background_refresh(self, channel_id: str, interval_seconds: float = 60.0) -> None: if channel_id in self._refresh_tasks and not self._refresh_tasks[channel_id].done(): return self._refresh_tasks[channel_id] = asyncio.create_task( self._background_refresh_loop(channel_id, interval_seconds) ) def stop_background_refresh(self, channel_id: str) -> None: task = self._refresh_tasks.pop(channel_id, None) if task and not task.done(): task.cancel() async def _background_refresh_loop(self, channel_id: str, interval_seconds: float) -> None: from yuxi.utils.logging_config import logger while True: try: await asyncio.sleep(interval_seconds) cached = self._cache.get(channel_id) if cached and cached.expires_at: remaining = cached.expires_at - time.monotonic() - self.REFRESH_MARGIN_SECONDS if remaining > 0: sleep_time = max(30, min(interval_seconds, remaining)) await asyncio.sleep(sleep_time) await self.refresh_token(channel_id) except asyncio.CancelledError: break except Exception as e: logger.warning(f"[TokenManager] Background refresh failed for {channel_id}: {e}") async def get_all_token_status(self) -> dict[str, TokenState]: result = {} for channel_id in self._providers: try: result[channel_id] = await self.get_token(channel_id) except Exception: result[channel_id] = TokenState(token="***error***") return result _token_manager: UnifiedTokenManager | None = None def get_token_manager() -> UnifiedTokenManager: global _token_manager if _token_manager is None: _token_manager = UnifiedTokenManager() return _token_manager