新增了完整的认证工具库模块,包含以下功能: 1. 指数退避重试组件 2. 敏感数据过滤与日志脱敏 3. 安全策略管理引擎 4. 认证健康监控模块 5. SSRF防护工具集 6. 多类型token管理系统 7. 密钥管理器加解密工具
184 lines
6.2 KiB
Python
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
|