ForcePilot/backend/package/yuxi/channels/auth/token_providers.py

183 lines
5.5 KiB
Python
Raw Normal View History

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