from __future__ import annotations import re from collections.abc import Mapping from typing import Any import httpx from app.schemas.remnawave import PaginatedUsers, RemnawaveUser, ResolvedUser, SubscriptionRequestHistory UUID_RE = re.compile( r"^[0-9a-fA-F]{8}-" r"[0-9a-fA-F]{4}-" r"[0-9a-fA-F]{4}-" r"[0-9a-fA-F]{4}-" r"[0-9a-fA-F]{12}$" ) class RemnawaveApiError(RuntimeError): def __init__(self, *, status_code: int, message: str, payload: Any | None = None) -> None: self.status_code = status_code self.message = message self.payload = payload super().__init__(message) class RemnawaveApiClient: def __init__( self, *, base_url: str, api_token: str, caddy_api_key: str = "", timeout_seconds: float = 20.0, client: httpx.AsyncClient | None = None, ) -> None: self._managed_client = client is None self._client = client or httpx.AsyncClient( base_url=self.normalize_base_url(base_url), timeout=timeout_seconds, headers=self._build_headers( base_url=base_url, api_token=api_token, caddy_api_key=caddy_api_key, ), ) @staticmethod def normalize_base_url(base_url: str) -> str: normalized = base_url.strip().rstrip("/") if not normalized.endswith("/api"): normalized = f"{normalized}/api" return f"{normalized}/" @staticmethod def unwrap_payload(payload: Any) -> Any: if isinstance(payload, Mapping) and "response" in payload: return payload["response"] return payload @staticmethod def _build_headers( *, base_url: str, api_token: str, caddy_api_key: str, ) -> dict[str, str]: headers: dict[str, str] = {"Accept": "application/json"} if api_token: headers["Authorization"] = ( api_token if api_token.startswith("Bearer ") else f"Bearer {api_token}" ) if caddy_api_key: headers["X-Api-Key"] = caddy_api_key if base_url.startswith("http://"): headers["x-forwarded-proto"] = "https" headers["x-forwarded-for"] = "127.0.0.1" return headers async def close(self) -> None: if self._managed_client: await self._client.aclose() async def get_users_by_telegram_id(self, telegram_id: int) -> list[RemnawaveUser]: payload = await self._request_json("GET", f"users/by-telegram-id/{telegram_id}") return [RemnawaveUser.model_validate(item) for item in payload] async def get_user_by_uuid(self, user_uuid: str) -> RemnawaveUser: payload = await self._request_json("GET", f"users/{user_uuid}") return RemnawaveUser.model_validate(payload) async def get_user_by_short_uuid(self, short_uuid: str) -> RemnawaveUser: payload = await self._request_json("GET", f"users/by-short-uuid/{short_uuid}") return RemnawaveUser.model_validate(payload) async def get_user_by_username(self, username: str) -> RemnawaveUser: payload = await self._request_json("GET", f"users/by-username/{username}") return RemnawaveUser.model_validate(payload) async def create_user(self, body: dict[str, Any]) -> RemnawaveUser: payload = await self._request_json("POST", "users", json=body) return RemnawaveUser.model_validate(payload) async def update_user(self, body: dict[str, Any]) -> RemnawaveUser: payload = await self._request_json("PATCH", "users", json=body) return RemnawaveUser.model_validate(payload) async def get_all_users(self, *, start: int = 0, size: int = 100) -> PaginatedUsers: payload = await self._request_json( "GET", "users", params={"start": start, "size": size}, ) return PaginatedUsers.model_validate(payload) async def get_user_subscription_history(self, user_uuid: str) -> SubscriptionRequestHistory: payload = await self._request_json( "GET", f"users/{user_uuid}/subscription-request-history", ) return SubscriptionRequestHistory.model_validate(payload) async def resolve_user(self, identifier: str) -> ResolvedUser: body: dict[str, Any] if UUID_RE.match(identifier): body = {"uuid": identifier} elif identifier.isdigit(): body = {"id": int(identifier)} else: try: payload = await self._request_json( "POST", "users/resolve", json={"shortUuid": identifier}, ) return ResolvedUser.model_validate(payload) except RemnawaveApiError as exc: if exc.status_code != 404: raise body = {"username": identifier} payload = await self._request_json("POST", "users/resolve", json=body) return ResolvedUser.model_validate(payload) async def _request_json( self, method: str, path: str, *, params: dict[str, Any] | None = None, json: dict[str, Any] | None = None, ) -> Any: response = await self._client.request( method=method, url=path.lstrip("/"), params=params, json=json, ) if response.is_error: raise self._build_error(response) try: payload = response.json() except ValueError: return response.text return self.unwrap_payload(payload) def _build_error(self, response: httpx.Response) -> RemnawaveApiError: try: payload = response.json() except ValueError: payload = None message = "Unknown Remnawave API error" if isinstance(payload, Mapping): for key in ("message", "error", "code"): candidate = payload.get(key) if candidate: message = str(candidate) break error_code = payload.get("errorCode") if error_code: message = f"{message} [{error_code}]" error_items = payload.get("errors") if isinstance(error_items, list): details: list[str] = [] for item in error_items: if not isinstance(item, Mapping): continue path_value = item.get("path") if isinstance(path_value, list): path_text = ".".join(str(part) for part in path_value if part is not None) elif path_value is not None: path_text = str(path_value) else: path_text = "" detail_message = str(item.get("message") or item.get("code") or "").strip() if not detail_message: continue if path_text: details.append(f"{path_text}: {detail_message}") else: details.append(detail_message) if details: message = f"{message} | {'; '.join(details)}" elif response.text: message = response.text return RemnawaveApiError( status_code=response.status_code, message=message, payload=payload, )