ForcePilot/backend/package/yuxi/channels/auth/token_manager.py
Kris 655eab7230 feat(auth): add complete auth toolkit module
新增了完整的认证工具库模块,包含以下功能:
1. 指数退避重试组件
2. 敏感数据过滤与日志脱敏
3. 安全策略管理引擎
4. 认证健康监控模块
5. SSRF防护工具集
6. 多类型token管理系统
7. 密钥管理器加解密工具
2026-05-13 16:18:51 +08:00

184 lines
6.2 KiB
Python

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