import asyncio import logging import time from collections.abc import Awaitable, Callable import aiohttp from yuxi.channel.config.defaults import TIMEOUT from yuxi.channel.runtime.backoff import ( BackoffConfig, ErrorBackoff, ) from yuxi.channel.gateway.net_utils import is_secure_ws_url from yuxi.channel.gateway.protocol import ( GatewayErrorCode, RpcEvent, RpcRequest, RpcResponse, marshal_frame, unmarshal_frame, ) logger = logging.getLogger(__name__) DEFAULT_GATEWAY_URL = "ws://127.0.0.1:9001" DEFAULT_REQUEST_TIMEOUT = TIMEOUT.http.normal DEFAULT_RECONNECT_BASE_MS = 1000.0 DEFAULT_RECONNECT_MAX_MS = 30000.0 DEFAULT_RECONNECT_FACTOR = 2.0 DEFAULT_RECONNECT_JITTER = 0.1 OnConnectCallback = Callable[["GatewayClient"], Awaitable[None]] OnDisconnectCallback = Callable[["GatewayClient", int, str], Awaitable[None]] OnEventCallback = Callable[["GatewayClient", RpcEvent], Awaitable[None]] OnErrorCallback = Callable[["GatewayClient", Exception], Awaitable[None]] OnCloseCallback = Callable[["GatewayClient", int, str], Awaitable[None]] class GatewayClientError(Exception): def __init__(self, code: GatewayErrorCode | str, message: str, details: dict | None = None): self.code = code self.details = details or {} super().__init__(message) class _PendingRequest: __slots__ = ("future", "timeout_handle") future: asyncio.Future timeout_handle: asyncio.TimerHandle | None def __init__(self): self.future = asyncio.get_running_loop().create_future() self.timeout_handle = None class GatewayClient: """统一 Gateway WebSocket 客户端 SDK。 连接 → 认证 → 收发 RPC 请求/事件 → 自动重连。 对齐 OpenClaw GatewayClient 语义,适配项目现有 Gateway 协议。 用法:: client = GatewayClient( url="ws://127.0.0.1:9001", token="my-token", on_event=my_event_handler, ) await client.start() resp = await client.request("system.health") ... await client.stop() """ def __init__( self, *, url: str = DEFAULT_GATEWAY_URL, token: str | None = None, password: str | None = None, device_token: str | None = None, request_timeout: float = DEFAULT_REQUEST_TIMEOUT, reconnect_base_ms: float = DEFAULT_RECONNECT_BASE_MS, reconnect_max_ms: float = DEFAULT_RECONNECT_MAX_MS, reconnect_factor: float = DEFAULT_RECONNECT_FACTOR, reconnect_jitter: float = DEFAULT_RECONNECT_JITTER, on_connect: OnConnectCallback | None = None, on_disconnect: OnDisconnectCallback | None = None, on_event: OnEventCallback | None = None, on_error: OnErrorCallback | None = None, on_close: OnCloseCallback | None = None, ): self._url = url self._token = token self._password = password self._device_token = device_token self._request_timeout = max(1.0, request_timeout) self._reconnect_base_ms = max(100.0, reconnect_base_ms) self._reconnect_max_ms = max(reconnect_base_ms, reconnect_max_ms) self._reconnect_factor = max(1.0, reconnect_factor) self._reconnect_jitter = float(reconnect_jitter) self._on_connect = on_connect self._on_disconnect = on_disconnect self._on_event = on_event self._on_error = on_error self._on_close = on_close self._ws: aiohttp.ClientWebSocketResponse | None = None self._session: aiohttp.ClientSession | None = None self._pending: dict[str, _PendingRequest] = {} self._running = False self._close_event = asyncio.Event() self._reconnect_task: asyncio.Task | None = None self._device_id: str | None = None self._device_private_key_pem: str | None = None self._last_tick: float | None = None self._tick_interval_ms = 30_000 self._tick_watch_task: asyncio.Task | None = None self._reconnect_backoff = ErrorBackoff( config=BackoffConfig( base_delay=reconnect_base_ms / 1000.0, max_delay=reconnect_max_ms / 1000.0, exponent=reconnect_factor, jitter=True, jitter_factor=reconnect_jitter, max_retries=0, ), ) @property def connected(self) -> bool: return self._ws is not None and not self._ws.closed @property def url(self) -> str: return self._url async def start(self) -> None: """启动客户端,开始连接和自动重连循环。""" if self._running: return self._running = True self._close_event.clear() self._validate_url() self._reconnect_task = asyncio.create_task(self._reconnect_loop()) async def stop(self) -> None: """停止客户端,断开连接并等待清理完成。""" self._running = False self._close_event.set() self._stop_tick_watch() self._cancel_all_pending(RuntimeError("gateway client stopped")) await self._disconnect() if self._session and not self._session.closed: await self._session.close() self._session = None if self._reconnect_task: self._reconnect_task.cancel() try: await self._reconnect_task except asyncio.CancelledError: pass self._reconnect_task = None async def request(self, method: str, params: dict | None = None, timeout: float | None = None) -> dict: """发送 RPC 请求,返回响应结果。 Raises: GatewayClientError: 服务端返回错误 RuntimeError: 未连接 TimeoutError: 请求超时 """ if not self.connected: raise RuntimeError("gateway not connected") req = RpcRequest(method=method, params=params) frame = marshal_frame(req) request_id = req.id effective_timeout = timeout if timeout is not None else self._request_timeout pending = _PendingRequest() self._pending[request_id] = pending timeout_coro: asyncio.Task | None = None try: timeout_coro = asyncio.create_task(asyncio.sleep(effective_timeout)) send_task = asyncio.create_task(self._ws.send_str(frame)) done, _pending_set = await asyncio.wait( [pending.future, send_task, timeout_coro], return_when=asyncio.FIRST_COMPLETED, ) for task in _pending_set: task.cancel() try: await task except (asyncio.CancelledError, Exception): pass if timeout_coro in done and not pending.future.done(): raise TimeoutError(f"gateway request timeout for {method}") if send_task in done: send_exc = send_task.exception() if send_exc: raise send_exc if pending.future in done: exc = pending.future.exception() if exc: raise exc return pending.future.result() raise RuntimeError(f"gateway request failed for {method}: unexpected state") finally: self._pending.pop(request_id, None) if timeout_coro and not timeout_coro.done(): timeout_coro.cancel() try: await timeout_coro except asyncio.CancelledError: pass async def request_with_retry( self, method: str, params: dict | None = None, *, timeout: float | None = None, max_attempts: int = 3, ) -> dict: """发送 RPC 请求,失败时自动重试。 Raises: GatewayClientError: 服务端返回错误 RuntimeError: 未连接 TimeoutError: 请求超时 """ backoff = ErrorBackoff( config=BackoffConfig( base_delay=0.4, max_delay=5.0, exponent=2.0, jitter=True, jitter_factor=0.1, max_retries=max_attempts - 1, ), ) async def _do_request() -> dict: return await self.request(method, params=params, timeout=timeout) return await backoff.execute(f"gw-rpc:{method}", _do_request) async def init_device_identity(self) -> str: """生成并注册设备身份,返回 device_token。 首次调用生成 Ed25519 密钥对,注册到设备注册表, 并构建可用于认证的 device_token。 """ from yuxi.channel.gateway.device_auth import build_device_token as _build_token from yuxi.channel.gateway.device_identity import generate_device_identity from yuxi.channel.gateway.device_registry import register_device device_id, public_key_pem, private_key_pem = generate_device_identity() self._device_id = device_id self._device_private_key_pem = private_key_pem register_device(device_id, public_key_pem) token = self._build_device_token_internal() self._device_token = token logger.info("GatewayClient device identity initialized: device_id=%s", device_id) return token async def refresh_device_token(self) -> str | None: """刷新设备令牌(重新签名时间戳),返回新的 device_token。""" if not self._device_id or not self._device_private_key_pem: return None token = self._build_device_token_internal() self._device_token = token return token @property def device_id(self) -> str | None: return self._device_id # ---- internal ---- def _validate_url(self) -> None: if not is_secure_ws_url(self._url): raise ValueError( f"不安全连接: 不允许通过明文 ws:// 连接到非回环地址 " f"(CWE-319)。请使用 wss:// 或将 gateway.bind 设为 loopback。" f" 当前 URL: {self._url}" ) async def _ensure_session(self) -> None: if self._session is None or self._session.closed: self._session = aiohttp.ClientSession() async def _connect(self) -> bool: await self._ensure_session() headers = {} auth_value = self._resolve_auth_header() if auth_value: headers["Authorization"] = auth_value try: self._ws = await self._session.ws_connect(self._url, headers=headers, heartbeat=30.0) except aiohttp.ClientError as e: logger.warning("Gateway client connect failed: %s", e) await self._fire_error(e) return False except Exception as e: logger.warning("Gateway client connect error: %s", e) await self._fire_error(e) return False try: request_id = f"auth-check-{id(self):x}" await asyncio.wait_for( self._ws.send_str(marshal_frame(RpcRequest(id=request_id, method="system.health", params={}))), timeout=TIMEOUT.gateway.connect, ) auth_ok = False async for msg in self._ws: if msg.type == aiohttp.WSMsgType.TEXT: frame = unmarshal_frame(msg.data) if isinstance(frame, RpcResponse) and frame.id == request_id: auth_ok = frame.ok break elif msg.type in (aiohttp.WSMsgType.CLOSED, aiohttp.WSMsgType.ERROR): break if not auth_ok: logger.warning("Gateway client auth check failed") await self._ws.close(code=4001, message=b"auth failed") self._ws = None return False return True except Exception as e: logger.warning("Gateway client auth check error: %s", e) ws = self._ws self._ws = None if ws and not ws.closed: try: await ws.close() except Exception: pass return False def _resolve_auth_header(self) -> str | None: if self._device_token: return f"Bearer {self._device_token}" if self._password: return f"Bearer {self._password}" if self._token: return f"Bearer {self._token}" return None async def _disconnect(self) -> None: self._stop_tick_watch() self._cancel_all_pending(RuntimeError("gateway disconnected")) ws = self._ws self._ws = None if ws and not ws.closed: try: await ws.close(code=1000, message=b"client shutdown") except Exception: pass def _start_tick_watch(self) -> None: self._last_tick = None async def _watch(): try: while True: await asyncio.sleep(self._tick_interval_ms / 1000.0 * 2) if self._last_tick is None: continue gap_ms = (asyncio.get_event_loop().time() - self._last_tick) * 1000 if gap_ms > self._tick_interval_ms * 3: logger.warning("Gateway tick timeout: gap=%.0fms", gap_ms) await self._disconnect() break except asyncio.CancelledError: pass self._tick_watch_task = asyncio.create_task(_watch()) def _stop_tick_watch(self) -> None: if self._tick_watch_task and not self._tick_watch_task.done(): self._tick_watch_task.cancel() self._tick_watch_task = None self._last_tick = None async def _reconnect_loop(self) -> None: backoff_ms = self._reconnect_base_ms while self._running: connected = await self._connect() if not connected: delay = self._reconnect_backoff.config.compute_delay(0) try: await asyncio.wait_for(self._close_event.wait(), timeout=delay) return except TimeoutError: backoff_ms = min(backoff_ms * self._reconnect_factor, self._reconnect_max_ms) continue backoff_ms = self._reconnect_base_ms self._start_tick_watch() await self._fire_connect() try: await self._recv_loop() except aiohttp.ClientError as e: logger.warning("Gateway client recv error: %s", e) await self._fire_error(e) except asyncio.CancelledError: return except Exception as e: logger.exception("Gateway client recv unexpected error") await self._fire_error(e) close_code = 1006 close_reason = "connection lost" if self._ws is not None: close_code = self._ws.close_code or close_code close_reason = self._ws._close_reason or close_reason await self._disconnect() await self._fire_close(close_code, close_reason) if not self._running: return delay = self._reconnect_backoff.config.compute_delay(0) try: await asyncio.wait_for(self._close_event.wait(), timeout=delay) return except TimeoutError: backoff_ms = min(backoff_ms * self._reconnect_factor, self._reconnect_max_ms) async def _recv_loop(self) -> None: ws = self._ws if ws is None: return async for msg in ws: if msg.type == aiohttp.WSMsgType.TEXT: try: frame = unmarshal_frame(msg.data) except Exception: logger.warning("Gateway client: invalid frame received") continue if isinstance(frame, RpcResponse): self._handle_response(frame) elif isinstance(frame, RpcEvent): self._handle_event(frame) elif msg.type == aiohttp.WSMsgType.CLOSED: break elif msg.type == aiohttp.WSMsgType.ERROR: break def _handle_response(self, resp: RpcResponse) -> None: pending = self._pending.get(resp.id) if pending is None: return self._pending.pop(resp.id, None) if resp.ok: pending.future.set_result(resp.result or {}) else: code = resp.error_code or GatewayErrorCode.INTERNAL_ERROR msg = resp.error_message or "unknown error" err = GatewayClientError(code=code, message=msg) pending.future.set_exception(err) def _handle_event(self, evt: RpcEvent) -> None: if evt.event == "tick": self._last_tick = asyncio.get_event_loop().time() if self._on_event: asyncio.create_task(self._fire_event(evt)) def _cancel_all_pending(self, exc: Exception) -> None: for pending in list(self._pending.values()): if not pending.future.done(): pending.future.set_exception(exc) self._pending.clear() def _compute_reconnect_levels(self) -> list[float]: levels = [] current = self._reconnect_base_ms / 1000.0 for _ in range(10): if current >= self._reconnect_max_ms / 1000.0: break levels.append(current) current *= self._reconnect_factor if not levels or levels[-1] < self._reconnect_max_ms / 1000.0: levels.append(self._reconnect_max_ms / 1000.0) return levels def _build_device_token_internal(self) -> str: from yuxi.channel.gateway.device_auth import build_device_token as _build_token, build_challenge from yuxi.channel.gateway.device_identity import sign_challenge timestamp_ms = int(time.time() * 1000) challenge = build_challenge(self._device_id, timestamp_ms) signature = sign_challenge(self._device_private_key_pem, challenge) return _build_token(self._device_id, timestamp_ms, signature) async def _fire_connect(self) -> None: if self._on_connect: try: await self._on_connect(self) except Exception: logger.exception("on_connect callback failed") async def _fire_close(self, code: int, reason: str) -> None: await self._fire_disconnect(code, reason) if self._on_close: try: await self._on_close(self, code, reason) except Exception: logger.exception("on_close callback failed") async def _fire_disconnect(self, code: int, reason: str) -> None: if self._on_disconnect: try: await self._on_disconnect(self, code, reason) except Exception: logger.exception("on_disconnect callback failed") async def _fire_event(self, evt: RpcEvent) -> None: if self._on_event: try: await self._on_event(self, evt) except Exception: logger.exception("on_event callback failed") async def _fire_error(self, err: Exception) -> None: if self._on_error: try: await self._on_error(self, err) except Exception: logger.exception("on_error callback failed")