Files
LLM-Zabbix-HABR/alert-processor/app/matrix_token_manager.py
T

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()