ForcePilot/backend/package/yuxi/channels/auth/token_providers.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

183 lines
5.5 KiB
Python

from __future__ import annotations
import asyncio
import time
from typing import Any, Protocol
from yuxi.channels.auth.token_manager import (
BaseTokenProvider,
TokenProviderType,
TokenState,
)
class TokenRefreshFunc(Protocol):
async def __call__(self) -> tuple[str, float | None]: ...
class StaticTokenProvider(BaseTokenProvider):
provider_type = TokenProviderType.STATIC
def __init__(self, token: str, scopes: list[str] | None = None):
self._token = token
self._scopes = scopes or []
async def get_token(self) -> TokenState:
return TokenState(
token=self._token,
expires_at=None,
token_type="bearer",
scopes=self._scopes,
source="config",
)
async def refresh_token(self) -> TokenState:
return await self.get_token()
async def validate_token(self, token: TokenState) -> bool:
return bool(token.token)
async def revoke_token(self) -> None:
pass
class OAuth2TokenProvider(BaseTokenProvider):
provider_type = TokenProviderType.OAUTH2
def __init__(self, refresh_func: TokenRefreshFunc, scopes: list[str] | None = None):
self._refresh_func = refresh_func
self._scopes = scopes or []
self._access_token: str | None = None
self._expires_at: float | None = None
self._lock = asyncio.Lock()
async def get_token(self) -> TokenState:
async with self._lock:
if not self._is_expired() and self._access_token:
return TokenState(
token=self._access_token,
expires_at=self._expires_at,
token_type="bearer",
scopes=self._scopes,
source="oauth2",
)
token, expires_at = await self._refresh_func()
self._access_token = token
self._expires_at = expires_at
return TokenState(
token=token,
expires_at=expires_at,
token_type="bearer",
scopes=self._scopes,
source="oauth2",
)
async def refresh_token(self) -> TokenState:
async with self._lock:
token, expires_at = await self._refresh_func()
self._access_token = token
self._expires_at = expires_at
return TokenState(
token=token,
expires_at=expires_at,
token_type="bearer",
scopes=self._scopes,
source="oauth2",
)
def _is_expired(self) -> bool:
if self._access_token is None or self._expires_at is None:
return True
return time.monotonic() > self._expires_at - 300
async def validate_token(self, token: TokenState) -> bool:
return bool(token.token) and not token.is_expired
async def revoke_token(self) -> None:
self._access_token = None
self._expires_at = None
class QRTokenProvider(BaseTokenProvider):
provider_type = TokenProviderType.QR_CODE
def __init__(self, get_credential_func):
self._get_credential = get_credential_func
async def get_token(self) -> TokenState:
credential = await self._get_credential()
return TokenState(
token=str(credential) if credential else "",
expires_at=None,
token_type="qr_credential",
source="qr_code",
)
async def refresh_token(self) -> TokenState:
return await self.get_token()
async def validate_token(self, token: TokenState) -> bool:
return bool(token.token)
async def revoke_token(self) -> None:
pass
class CertificateTokenProvider(BaseTokenProvider):
provider_type = TokenProviderType.CERTIFICATE
def __init__(self, certificate_data: dict[str, Any]):
self._cert = certificate_data
async def get_token(self) -> TokenState:
return TokenState(
token="certificate",
expires_at=None,
token_type="certificate",
source="certificate",
metadata=self._cert,
)
async def refresh_token(self) -> TokenState:
return await self.get_token()
async def validate_token(self, token: TokenState) -> bool:
return bool(token.metadata)
async def revoke_token(self) -> None:
pass
class CompatTokenProvider(BaseTokenProvider):
provider_type = TokenProviderType.OAUTH2
def __init__(self, compat_manager: Any):
self._compat = compat_manager
async def get_token(self) -> TokenState:
token = await self._compat.get_token()
expires_at = getattr(self._compat, "_token_expires_at", None) or getattr(self._compat, "_expires_at", None)
return TokenState(
token=token,
expires_at=expires_at,
token_type="bearer",
source="compat",
)
async def refresh_token(self) -> TokenState:
token = await self._compat.refresh_token()
expires_at = getattr(self._compat, "_token_expires_at", None) or getattr(self._compat, "_expires_at", None)
return TokenState(
token=token,
expires_at=expires_at,
token_type="bearer",
source="compat",
)
async def validate_token(self, token: TokenState) -> bool:
return bool(token.token) and not token.is_expired
async def revoke_token(self) -> None:
if hasattr(self._compat, "invalidate"):
self._compat.invalidate()