v2.0
This commit is contained in:
224
app/services/remnawave_client.py
Normal file
224
app/services/remnawave_client.py
Normal file
@@ -0,0 +1,224 @@
|
||||
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,
|
||||
)
|
||||
Reference in New Issue
Block a user