ForcePilot/backend/package/yuxi/channel/gateway/client.py
Kris ecd3c90e80 feat(channel/gateway): 新增完整网关通道模块
新增设备身份管理、认证限流、并发通道、Webhook路由、RBAC权限控制、SSE/轮询降级等全套网关通道功能,包含:
1. 设备身份生成与签名验证
2. 设备令牌认证与速率限制
3. 内存+数据库双重设备注册表
4. 并发通道限流管理
5. Webhook安全处理与路由
6. RBAC权限校验系统
7. OpenAI API兼容适配层
8. Tailscale认证支持
9. HTTP轮询降级机制
2026-05-21 10:26:33 +08:00

558 lines
20 KiB
Python

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")