Загрузить файлы в «alert-processor/app»
This commit is contained in:
@@ -0,0 +1,168 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user