from __future__ import annotations import asyncio import logging from datetime import UTC, datetime, timedelta import httpx logger = logging.getLogger(__name__) PDD_TOKEN_URL = "https://open-api.pinduoduo.com/oauth/token" TOKEN_REFRESH_MARGIN = 600 REFRESH_CHECK_INTERVAL = 3600 class PddTokenManager: def __init__( self, client_id: str, client_secret: str, mall_id: str, refresh_token: str | None = None, access_token: str | None = None, expires_at: datetime | None = None, ): self._client_id = client_id self._client_secret = client_secret self._mall_id = mall_id self._access_token: str | None = access_token self._refresh_token: str | None = refresh_token self._expires_at: datetime | None = expires_at self._lock = asyncio.Lock() self._http_client: httpx.AsyncClient | None = None self._refresh_task: asyncio.Task | None = None 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._access_token is None: raise RuntimeError(f"Failed to obtain access_token for mall_id={self._mall_id}") return self._access_token async def invalidate_token(self) -> None: async with self._lock: self._access_token = None self._expires_at = None def _is_expired(self) -> bool: if self._access_token is None: return True if self._expires_at is None: return True return datetime.now(UTC) >= self._expires_at async def _refresh(self) -> None: if not self._refresh_token: raise RuntimeError(f"No refresh_token available for mall_id={self._mall_id}") client = self._get_client() try: resp = await client.post( PDD_TOKEN_URL, json={ "client_id": self._client_id, "client_secret": self._client_secret, "grant_type": "refresh_token", "refresh_token": self._refresh_token, }, ) resp.raise_for_status() data = resp.json() self._access_token = data.get("access_token") self._refresh_token = data.get("refresh_token") expires_in = data.get("expires_in", 86400) self._expires_at = datetime.now(UTC) + timedelta(seconds=expires_in - TOKEN_REFRESH_MARGIN) logger.info( "Token refreshed for mall_id=%s, expires_at=%s", self._mall_id, self._expires_at, ) except httpx.HTTPStatusError as e: logger.error("Token refresh HTTP error for mall_id=%s: %s", self._mall_id, e) raise except httpx.RequestError as e: logger.error("Token refresh network error for mall_id=%s: %s", self._mall_id, e) raise async def exchange_code(self, code: str, redirect_uri: str) -> dict: client = self._get_client() try: resp = await client.post( PDD_TOKEN_URL, json={ "client_id": self._client_id, "client_secret": self._client_secret, "grant_type": "authorization_code", "code": code, "redirect_uri": redirect_uri, }, ) resp.raise_for_status() data = resp.json() self._access_token = data.get("access_token") self._refresh_token = data.get("refresh_token") expires_in = data.get("expires_in", 86400) self._expires_at = datetime.now(UTC) + timedelta(seconds=expires_in - TOKEN_REFRESH_MARGIN) return data except httpx.HTTPStatusError as e: logger.error("Code exchange HTTP error: %s", e) raise 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"pdd-token-refresh-{self._mall_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: import random while True: offset = random.uniform(0, 300) await asyncio.sleep(REFRESH_CHECK_INTERVAL + offset) try: async with self._lock: if self._is_expired(): await self._refresh() except Exception: logger.exception( "Background token refresh failed for mall_id=%s", self._mall_id, ) async def close(self) -> None: await self.stop_background_refresh() if self._http_client: await self._http_client.aclose() self._http_client = None @property def mall_id(self) -> str: return self._mall_id @property def client_id(self) -> str: return self._client_id @property def client_secret(self) -> str: return self._client_secret