168 lines
5.4 KiB
Python
168 lines
5.4 KiB
Python
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() |