from __future__ import annotations import asyncio import ipaddress import logging import random from urllib.parse import urlparse import aiohttp logger = logging.getLogger(__name__) LOOPBACK_HOSTS = {"127.0.0.1", "localhost", "::1", "0.0.0.0"} PRIVATE_NETWORKS = [ ipaddress.IPv4Network("10.0.0.0/8"), ipaddress.IPv4Network("172.16.0.0/12"), ipaddress.IPv4Network("192.168.0.0/16"), ] class RpcError(Exception): def __init__(self, message: str, code: int = -1): super().__init__(message) self.code = code class RateLimitError(RpcError): def __init__(self, message: str, retry_after_seconds: int | None = None, token: str | None = None): super().__init__(message, code=-32603) self.retry_after_seconds = retry_after_seconds self.token = token def _detect_rate_limit(message: str, error_data: dict) -> RateLimitError | None: msg_lower = message.lower() if "rate limit" in msg_lower or "ratelimitexception" in msg_lower: retry_after = error_data.get("retry_after_seconds") token = error_data.get("token") return RateLimitError(message, retry_after_seconds=retry_after, token=token) return None RATE_LIMIT_CODE = -32603 class UnauthorizedError(RpcError): def __init__(self, message: str = "Unauthorized"): super().__init__(message, code=-32001) class RpcClient: DEFAULT_RETRY = { "attempts": 3, "min_delay_ms": 400, "max_delay_ms": 30000, "jitter": 0.1, } def __init__( self, base_url: str, retry_config: dict | None = None, timeout_ms: int = 30000, allow_remote_daemon: bool = False, ): self._base_url = base_url.rstrip("/") self._session: aiohttp.ClientSession | None = None self._request_id = 0 self._retry_config = {**self.DEFAULT_RETRY, **(retry_config or {})} self._timeout_ms = timeout_ms self._allow_remote_daemon = allow_remote_daemon async def connect(self) -> None: if not self._allow_remote_daemon: _validate_daemon_url(self._base_url) self._session = aiohttp.ClientSession( base_url=self._base_url, timeout=aiohttp.ClientTimeout(total=self._timeout_ms / 1000.0), ) async def disconnect(self) -> None: if self._session: await self._session.close() self._session = None async def call(self, method: str, params: dict | None = None) -> dict: self._request_id += 1 payload = { "jsonrpc": "2.0", "method": method, "params": params or {}, "id": str(self._request_id), } if not self._session: raise RuntimeError("RpcClient not connected") attempts = self._retry_config["attempts"] max_delay = self._retry_config["max_delay_ms"] / 1000.0 jitter = self._retry_config["jitter"] last_error: Exception | None = None for attempt in range(attempts): try: return await self._do_call(payload) except (aiohttp.ClientError, TimeoutError) as e: last_error = e if attempt < attempts - 1: delay = min( self._retry_config["min_delay_ms"] / 1000.0 * (2**attempt), max_delay, ) delay += delay * jitter * random.random() await asyncio.sleep(delay) except RpcError: raise except UnauthorizedError: raise raise last_error # type: ignore[misc] async def _do_call(self, payload: dict) -> dict: async with self._session.post("/api/v1/rpc", json=payload) as resp: if resp.status == 401: raise UnauthorizedError("HTTP 401 Unauthorized") result = await resp.json() if "error" in result: err = result["error"] message = err.get("message", "Unknown RPC error") code = err.get("code", -1) rate_limit = _detect_rate_limit(message, err) if rate_limit: logger.warning(f"Rate limit detected: {rate_limit}") raise rate_limit raise RpcError(f"RPC error [{code}]: {message}", code=code) return result.get("result", {}) async def multipart_upload(self, endpoint: str, data: aiohttp.FormData) -> dict: if not self._session: raise RuntimeError("RpcClient not connected") async with self._session.post(endpoint, data=data) as resp: return await resp.json() @property def base_url(self) -> str: return self._base_url def _validate_daemon_url(base_url: str) -> None: parsed = urlparse(base_url) hostname = parsed.hostname if not hostname: raise ValueError(f"Invalid daemon URL: cannot determine host from {base_url}") if hostname in LOOPBACK_HOSTS: return try: addr = ipaddress.IPv4Address(hostname) except ValueError: try: addr = ipaddress.IPv6Address(hostname) except ValueError: if hostname not in LOOPBACK_HOSTS: raise ValueError( f"Remote daemon URL rejected (SSRF guard): {base_url}. " "Set allow_remote_daemon=True to permit non-local connections." ) return if addr.is_loopback: return for network in PRIVATE_NETWORKS: if isinstance(addr, ipaddress.IPv4Address) and addr in network: return raise ValueError( f"Remote daemon URL rejected (SSRF guard): {base_url}. " "Set allow_remote_daemon=True to permit non-local connections." )