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