from __future__ import annotations import asyncio import json import logging import time from dataclasses import asdict, dataclass from pathlib import Path import httpx logger = logging.getLogger(__name__) @dataclass class MatrixTokenState: access_token: str refresh_token: str expires_at: float | None = None class MatrixTokenManager: def __init__( self, token_endpoint: str, client_id: str, client_secret: str | None, initial_access_token: str, initial_refresh_token: str, initial_expires_in_seconds: int | None, refresh_margin_seconds: int, state_file: str, timeout_seconds: float = 10, verify_tls: bool = True, ) -> None: self.token_endpoint = token_endpoint self.client_id = client_id self.client_secret = client_secret or None self.refresh_margin_seconds = refresh_margin_seconds self.state_path = Path(state_file) self._lock = asyncio.Lock() self._stop_event = asyncio.Event() self._refresh_task: asyncio.Task | None = None self._refresh_failures = 0 self.client = httpx.AsyncClient( timeout=timeout_seconds, verify=verify_tls, ) state = self._load_state() if state is None: expires_at = None if initial_expires_in_seconds and initial_expires_in_seconds > 0: expires_at = time.time() + initial_expires_in_seconds state = MatrixTokenState( access_token=initial_access_token, refresh_token=initial_refresh_token, expires_at=expires_at, ) self._save_state(state) self._state = state def _load_state(self) -> MatrixTokenState | None: if not self.state_path.exists(): return None raw = json.loads(self.state_path.read_text(encoding="utf-8")) return MatrixTokenState( access_token=raw["access_token"], refresh_token=raw["refresh_token"], expires_at=raw.get("expires_at"), ) def _save_state(self, state: MatrixTokenState) -> None: self.state_path.write_text( json.dumps(asdict(state), ensure_ascii=False, indent=2), encoding="utf-8", ) def _needs_refresh(self) -> bool: if not self._state.refresh_token: return False if self._state.expires_at is None: return False return time.time() >= (self._state.expires_at - self.refresh_margin_seconds) async def get_access_token(self) -> str: await self.refresh_if_needed() return self._state.access_token async def refresh_if_needed(self, force: bool = False) -> MatrixTokenState: async with self._lock: if not force and not self._needs_refresh(): return self._state data = { "grant_type": "refresh_token", "refresh_token": self._state.refresh_token, "client_id": self.client_id, } if self.client_secret: data["client_secret"] = self.client_secret response = await self.client.post( self.token_endpoint, data=data, headers={"Content-Type": "application/x-www-form-urlencoded"}, ) response.raise_for_status() payload = response.json() access_token = payload["access_token"] refresh_token = payload.get("refresh_token", self._state.refresh_token) expires_in = payload.get("expires_in") expires_at = None if expires_in is not None: expires_at = time.time() + int(expires_in) self._state = MatrixTokenState( access_token=access_token, refresh_token=refresh_token, expires_at=expires_at, ) self._save_state(self._state) self._refresh_failures = 0 logger.info("Matrix OAuth token refreshed successfully") return self._state def start_background_refresh(self) -> None: if self._refresh_task is None: self._refresh_task = asyncio.create_task(self._refresh_loop()) async def _refresh_loop(self) -> None: while not self._stop_event.is_set(): try: await self.refresh_if_needed() except Exception as exc: self._refresh_failures += 1 logger.exception("Matrix token background refresh failed: %s", exc) if self._refresh_failures > 0: sleep_for = min(300, 30 * self._refresh_failures) else: sleep_for = 30 if self._state.expires_at is not None: remaining = int(self._state.expires_at - time.time() - self.refresh_margin_seconds) sleep_for = max(5, min(60, remaining if remaining > 0 else 5)) try: await asyncio.wait_for(self._stop_event.wait(), timeout=sleep_for) except asyncio.TimeoutError: pass async def close(self) -> None: self._stop_event.set() if self._refresh_task is not None: self._refresh_task.cancel() try: await self._refresh_task except asyncio.CancelledError: pass await self.client.aclose()