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()