新增设备身份管理、认证限流、并发通道、Webhook路由、RBAC权限控制、SSE/轮询降级等全套网关通道功能,包含: 1. 设备身份生成与签名验证 2. 设备令牌认证与速率限制 3. 内存+数据库双重设备注册表 4. 并发通道限流管理 5. Webhook安全处理与路由 6. RBAC权限校验系统 7. OpenAI API兼容适配层 8. Tailscale认证支持 9. HTTP轮询降级机制
558 lines
20 KiB
Python
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")
|