261 lines
8.9 KiB
Python
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,
|
|
)
|