Files
RemnaWaveBOT/app/services/remnawave_client.py
2026-07-09 16:45:17 +03:00

261 lines
8.9 KiB
Python

from __future__ import annotations
import re
from collections.abc import Mapping
from typing import Any
import httpx
from app.schemas.remnawave import HwidDevicesResponse, 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 get_user_hwid_devices(self, user_uuid: str) -> HwidDevicesResponse:
payload = await self._request_json("GET", f"hwid/devices/{user_uuid}")
return HwidDevicesResponse.model_validate(payload)
async def delete_user_hwid_device(self, *, user_uuid: str, hwid: str) -> HwidDevicesResponse:
payload = await self._request_json(
"POST",
"hwid/devices/delete",
json={"userUuid": user_uuid, "hwid": hwid},
)
return HwidDevicesResponse.model_validate(payload)
async def delete_all_user_hwid_devices(self, user_uuid: str) -> HwidDevicesResponse:
payload = await self._request_json(
"POST",
"hwid/devices/delete-all",
json={"userUuid": user_uuid},
)
return HwidDevicesResponse.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:
try:
response = await self._client.request(
method=method,
url=path.lstrip("/"),
params=params,
json=json,
)
except httpx.TimeoutException as exc:
raise RemnawaveApiError(
status_code=0,
message=f"Remnawave API timeout: {exc}",
) from exc
except httpx.NetworkError as exc:
raise RemnawaveApiError(
status_code=0,
message=f"Remnawave API network error: {exc}",
) from exc
except httpx.HTTPError as exc:
raise RemnawaveApiError(
status_code=0,
message=f"Remnawave API HTTP error: {exc}",
) from exc
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,
)