99 lines
3.5 KiB
Python
99 lines
3.5 KiB
Python
|
|
from __future__ import annotations
|
||
|
|
|
||
|
|
import asyncio
|
||
|
|
import logging
|
||
|
|
import random
|
||
|
|
import time
|
||
|
|
|
||
|
|
import httpx
|
||
|
|
|
||
|
|
from yuxi.channel.extensions.qqbot.errors import QQBotError, QQBotErrorCode
|
||
|
|
from yuxi.channel.extensions.qqbot.api_routes import API_BASE_URL, TOKEN_URL
|
||
|
|
|
||
|
|
logger = logging.getLogger(__name__)
|
||
|
|
|
||
|
|
TOKEN_REFRESH_MARGIN = 300
|
||
|
|
BACKGROUND_REFRESH_INTERVAL = 3600
|
||
|
|
RANDOM_OFFSET_MAX = 600
|
||
|
|
|
||
|
|
|
||
|
|
class TokenManager:
|
||
|
|
def __init__(self, app_id: str, client_secret: str):
|
||
|
|
self._app_id = app_id
|
||
|
|
self._client_secret = client_secret
|
||
|
|
self._token: str | None = None
|
||
|
|
self._expires_at: float = 0.0
|
||
|
|
self._lock = asyncio.Lock()
|
||
|
|
self._refresh_task: asyncio.Task | None = None
|
||
|
|
self._http_client: httpx.AsyncClient | None = None
|
||
|
|
|
||
|
|
@property
|
||
|
|
def app_id(self) -> str:
|
||
|
|
return self._app_id
|
||
|
|
|
||
|
|
def _get_client(self) -> httpx.AsyncClient:
|
||
|
|
if self._http_client is None:
|
||
|
|
self._http_client = httpx.AsyncClient(timeout=httpx.Timeout(15.0))
|
||
|
|
return self._http_client
|
||
|
|
|
||
|
|
async def get_token(self) -> str:
|
||
|
|
async with self._lock:
|
||
|
|
if self._is_expired():
|
||
|
|
await self._refresh()
|
||
|
|
if self._token is None:
|
||
|
|
raise QQBotError(QQBotErrorCode.AUTH_FAILED, "Failed to obtain access token")
|
||
|
|
return self._token
|
||
|
|
|
||
|
|
async def _refresh(self) -> None:
|
||
|
|
client = self._get_client()
|
||
|
|
try:
|
||
|
|
resp = await client.post(
|
||
|
|
TOKEN_URL,
|
||
|
|
json={"appId": self._app_id, "clientSecret": self._client_secret},
|
||
|
|
)
|
||
|
|
resp.raise_for_status()
|
||
|
|
data = resp.json()
|
||
|
|
self._token = data.get("access_token")
|
||
|
|
expires_in = data.get("expires_in", 7200)
|
||
|
|
self._expires_at = time.time() + expires_in - TOKEN_REFRESH_MARGIN
|
||
|
|
logger.debug("Token refreshed for app_id=%s, expires_in=%d", self._app_id, expires_in)
|
||
|
|
except httpx.HTTPError as e:
|
||
|
|
raise QQBotError(QQBotErrorCode.AUTH_FAILED, f"Token refresh failed: {e}", retryable=True) from e
|
||
|
|
|
||
|
|
def _is_expired(self) -> bool:
|
||
|
|
return self._token is None or time.time() >= self._expires_at
|
||
|
|
|
||
|
|
def invalidate(self) -> None:
|
||
|
|
self._token = None
|
||
|
|
self._expires_at = 0.0
|
||
|
|
logger.debug("Token invalidated for app_id=%s", self._app_id)
|
||
|
|
|
||
|
|
async def start_background_refresh(self) -> None:
|
||
|
|
if self._refresh_task and not self._refresh_task.done():
|
||
|
|
return
|
||
|
|
self._refresh_task = asyncio.create_task(self._background_refresh_loop(), name=f"token-refresh-{self._app_id}")
|
||
|
|
|
||
|
|
async def stop_background_refresh(self) -> None:
|
||
|
|
if self._refresh_task and not self._refresh_task.done():
|
||
|
|
self._refresh_task.cancel()
|
||
|
|
try:
|
||
|
|
await self._refresh_task
|
||
|
|
except asyncio.CancelledError:
|
||
|
|
pass
|
||
|
|
self._refresh_task = None
|
||
|
|
|
||
|
|
async def _background_refresh_loop(self) -> None:
|
||
|
|
while True:
|
||
|
|
offset = random.uniform(0, RANDOM_OFFSET_MAX)
|
||
|
|
await asyncio.sleep(BACKGROUND_REFRESH_INTERVAL + offset)
|
||
|
|
try:
|
||
|
|
async with self._lock:
|
||
|
|
await self._refresh()
|
||
|
|
except Exception:
|
||
|
|
logger.exception("Background token refresh failed for app_id=%s", self._app_id)
|
||
|
|
|
||
|
|
async def close(self) -> None:
|
||
|
|
await self.stop_background_refresh()
|
||
|
|
if self._http_client:
|
||
|
|
await self._http_client.aclose()
|
||
|
|
self._http_client = None
|