from __future__ import annotations import time from abc import ABC, abstractmethod from typing import Any from yuxi.utils.logging_config import logger from .probe import refresh_access_token, validate_token class AuthProvider(ABC): @abstractmethod async def get_access_token(self) -> str | None: ... @abstractmethod async def get_client_id(self) -> str: ... @abstractmethod async def validate(self) -> bool: ... class StaticAuthProvider(AuthProvider): def __init__(self, client_id: str, access_token: str): self._client_id = client_id self._access_token = access_token async def get_access_token(self) -> str | None: return self._access_token or None async def get_client_id(self) -> str: return self._client_id async def validate(self) -> bool: if not self._client_id or not self._access_token: return False user_info = await validate_token(self._client_id, self._access_token) return user_info is not None class RefreshingAuthProvider(AuthProvider): def __init__( self, client_id: str, client_secret: str, access_token: str = "", refresh_token: str = "", ): self._client_id = client_id self._client_secret = client_secret self._access_token = access_token self._refresh_token = refresh_token self._token_expires_at: float | None = None self._token_obtained_at: float | None = None self._token_expires_in: int = 0 async def get_access_token(self) -> str | None: if self._is_expired(): refreshed = await self._refresh() if not refreshed: logger.warning("[TwitchAuth] Token expired and refresh failed") return None return self._access_token or None async def get_client_id(self) -> str: return self._client_id async def validate(self) -> bool: if not self._client_id or not self._access_token: return False user_info = await validate_token(self._client_id, self._access_token) return user_info is not None async def _refresh(self) -> bool: if not self._client_id or not self._client_secret or not self._refresh_token: return False try: new_tokens = await refresh_access_token(self._client_id, self._client_secret, self._refresh_token) if new_tokens is None: return False new_access = new_tokens.get("access_token", "") new_refresh = new_tokens.get("refresh_token", "") expires_in = new_tokens.get("expires_in", 0) if new_access: self._access_token = new_access if new_refresh: self._refresh_token = new_refresh if expires_in: self._token_obtained_at = time.time() self._token_expires_in = expires_in self._token_expires_at = self._token_obtained_at + expires_in - 300 logger.info(f"[TwitchAuth] Token refreshed, expires in {expires_in}s") return True except Exception as e: logger.error(f"[TwitchAuth] Token refresh error: {e}") return False def _is_expired(self) -> bool: if self._token_expires_at is None: return False return time.time() >= self._token_expires_at @property def expires_at(self) -> float | None: return self._token_expires_at @property def expires_in(self) -> int: return self._token_expires_in def create_auth_provider(config: dict[str, Any]) -> AuthProvider: client_id = config.get("client_id", "") client_secret = config.get("client_secret", "") access_token = config.get("access_token", "") refresh_token = config.get("refresh_token", "") if client_secret and refresh_token: logger.info("[TwitchAuth] Using RefreshingAuthProvider (auto-refresh enabled)") return RefreshingAuthProvider( client_id=client_id, client_secret=client_secret, access_token=access_token, refresh_token=refresh_token, ) logger.info("[TwitchAuth] Using StaticAuthProvider (manual token rotation)") return StaticAuthProvider(client_id=client_id, access_token=access_token)