ForcePilot/backend/package/yuxi/channel/gateway/server.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

383 lines
14 KiB
Python

import asyncio
import logging
import os
import time as _time
from collections.abc import Callable, Collection
from fastapi import WebSocket, WebSocketDisconnect
from yuxi.channel.gateway.auth import GatewayAuthResult, authenticate_gateway_connect
from yuxi.channel.gateway.broadcaster import BroadcastFilter, BroadcastResult, gateway_broadcaster
from yuxi.channel.gateway.net_utils import (
_ENV_ALLOW_INSECURE_PRIVATE_WS,
is_loopback_address,
is_secure_ws_url,
)
from yuxi.channel.gateway.protocol import (
GatewayErrorCode,
GatewayRpcMethod,
HelloOk,
RpcEvent,
RpcRequest,
RpcResponse,
marshal_frame,
unmarshal_frame,
)
from yuxi.channel.config.defaults import TIMEOUT
from yuxi.channel.gateway.rbac import GatewayRole
from yuxi.channel.gateway.rpc_dispatcher import rpc_dispatcher
logger = logging.getLogger(__name__)
OnConnectCallback = Callable[[str, GatewayAuthResult], None]
OnDisconnectCallback = Callable[[str, GatewayAuthResult], None]
HEARTBEAT_INTERVAL = TIMEOUT.gateway.heartbeat_interval
HEARTBEAT_TIMEOUT = TIMEOUT.gateway.heartbeat
SEND_TIMEOUT = TIMEOUT.gateway.send
_SEND_SLOW_WARN_THRESHOLD = 5.0
CONNECT_NEGOTIATION_TIMEOUT = TIMEOUT.gateway.negotiation
TICK_INTERVAL_MS = 30_000
async def _send_with_timeout(ws: WebSocket, text: str, timeout: float = SEND_TIMEOUT) -> None:
try:
await asyncio.wait_for(ws.send_text(text), timeout=timeout)
except TimeoutError:
logger.warning(
"Gateway WS send timed out after %.0fs, closing connection",
timeout,
)
raise
class GatewayWsServer:
def __init__(self, allow_private_ws: bool | None = None):
self._active_connections: dict[str, WebSocket] = {}
self._connection_auth: dict[str, GatewayAuthResult] = {}
self._heartbeat_tasks: dict[str, asyncio.Task] = {}
self._tick_tasks: dict[str, asyncio.Task] = {}
self._on_connect: list[OnConnectCallback] = []
self._on_disconnect: list[OnDisconnectCallback] = []
self._shutting_down = False
self._allow_private_ws = (
allow_private_ws
if allow_private_ws is not None
else os.environ.get(_ENV_ALLOW_INSECURE_PRIVATE_WS) == "1"
)
@property
def active_count(self) -> int:
return len(self._active_connections)
@property
def shutting_down(self) -> bool:
return self._shutting_down
@property
def allow_private_ws(self) -> bool:
return self._allow_private_ws
def set_allow_private_ws(self, value: bool) -> None:
self._allow_private_ws = value
async def handle_connection(self, ws: WebSocket, token: str | None = None):
client_ip = _resolve_ws_client_ip(ws)
is_loopback_client = is_loopback_address(client_ip)
if not is_loopback_client and not is_secure_ws_url(str(ws.url), allow_private_ws=self._allow_private_ws):
display_host = _format_ws_display_host(str(ws.url))
logger.warning(
"Gateway WS: rejected insecure ws:// connection from %s to %s",
client_ip or "unknown",
display_host,
)
await ws.close(code=4400, reason="insecure_ws_connection")
return
await ws.accept()
auth_header = ws.headers.get("authorization")
auth_result = await authenticate_gateway_connect(
auth_header, token, client_ip=client_ip, headers=dict(ws.headers), remote_addr=client_ip
)
if not auth_result.authenticated:
await _send_with_timeout(
ws,
marshal_frame(
RpcResponse(
id="auth",
ok=False,
error_code=GatewayErrorCode.AUTH_ERROR,
error_message=auth_result.error or "认证失败",
)
),
)
await ws.close(code=4001, reason="unauthorized")
return
conn_id = f"{auth_result.user_id}-{id(ws):x}"
self._active_connections[conn_id] = ws
self._connection_auth[conn_id] = auth_result
async def _send_to_ws(text: str) -> None:
await ws.send_text(text)
gateway_broadcaster.register_connection(conn_id, auth_result, _send_to_ws)
heartbeat_task = asyncio.create_task(
self._heartbeat_loop(conn_id, ws),
name=f"gateway_heartbeat:{conn_id}",
)
self._heartbeat_tasks[conn_id] = heartbeat_task
tick_task = asyncio.create_task(
self._tick_loop(conn_id, ws),
name=f"gateway_tick:{conn_id}",
)
self._tick_tasks[conn_id] = tick_task
try:
first_frame = await self._negotiate_connect(conn_id, ws, auth_result)
except Exception:
logger.warning("Gateway connect negotiation failed: %s", conn_id)
first_frame = None
try:
for cb in self._on_connect:
try:
cb(conn_id, auth_result)
except Exception:
logger.exception("on_connect callback failed: %s", conn_id)
logger.info(
"Gateway WS connected: %s (user=%s, roles=%s)",
conn_id,
auth_result.user_id,
[r.value for r in auth_result.roles],
)
if first_frame is not None:
await self._dispatch_frame(conn_id, ws, first_frame)
while True:
raw = await ws.receive_text()
try:
frame = unmarshal_frame(raw)
except Exception:
logger.warning("Gateway WS: 无效帧 from %s", conn_id)
await _send_with_timeout(
ws,
marshal_frame(
RpcResponse(
id="",
ok=False,
error_code=GatewayErrorCode.INVALID_REQUEST,
error_message="无效的 RPC 帧格式",
)
),
)
continue
await self._dispatch_frame(conn_id, ws, frame)
except WebSocketDisconnect:
logger.info("Gateway WS disconnected: %s", conn_id)
except Exception:
logger.exception("Gateway WS error: %s", conn_id)
finally:
gateway_broadcaster.unregister_connection(conn_id)
for task_key, task_dict in [("heartbeat", self._heartbeat_tasks), ("tick", self._tick_tasks)]:
task = task_dict.get(conn_id)
if task and not task.done():
task.cancel()
try:
await task
except asyncio.CancelledError:
pass
task_dict.pop(conn_id, None)
self._active_connections.pop(conn_id, None)
auth = self._connection_auth.pop(conn_id, None)
for cb in self._on_disconnect:
try:
cb(conn_id, auth or GatewayAuthResult(authenticated=False))
except Exception:
logger.exception("on_disconnect callback failed: %s", conn_id)
async def _heartbeat_loop(self, conn_id: str, ws: WebSocket) -> None:
try:
while True:
await asyncio.sleep(HEARTBEAT_INTERVAL)
try:
await ws.send_text(marshal_frame(RpcEvent(event="ping")))
except Exception:
logger.warning("Gateway WS heartbeat failed: %s", conn_id)
break
except asyncio.CancelledError:
pass
async def _negotiate_connect(self, conn_id: str, ws: WebSocket, auth_result: GatewayAuthResult):
try:
raw = await asyncio.wait_for(ws.receive_text(), timeout=CONNECT_NEGOTIATION_TIMEOUT)
frame = unmarshal_frame(raw)
if isinstance(frame, RpcRequest) and frame.method == GatewayRpcMethod.CONNECT:
hello_ok = self._build_hello_ok(conn_id, auth_result)
await _send_with_timeout(ws, marshal_frame(hello_ok))
return None
return frame
except TimeoutError:
return None
def _build_hello_ok(self, conn_id: str, auth_result: GatewayAuthResult) -> HelloOk:
hello = HelloOk()
hello.result["server"]["connId"] = conn_id
hello.result["auth"] = {
"role": auth_result.roles[0].value if auth_result.roles else "viewer",
"scopes": [r.value for r in auth_result.roles],
}
return hello
async def _dispatch_frame(self, conn_id: str, ws: WebSocket, frame):
auth_result = self._connection_auth.get(conn_id)
if auth_result is None:
return
if isinstance(frame, RpcRequest):
frame.caller_user_id = auth_result.user_id
if frame.method == "event.subscribe":
resp = await self._handle_event_subscribe(conn_id, frame)
await _send_with_timeout(ws, marshal_frame(resp))
elif frame.method == "event.unsubscribe":
resp = await self._handle_event_unsubscribe(conn_id, frame)
await _send_with_timeout(ws, marshal_frame(resp))
elif rpc_dispatcher.is_stream(frame.method):
async for item in rpc_dispatcher.dispatch_stream(frame, auth_result.roles):
await _send_with_timeout(ws, marshal_frame(item), timeout=TIMEOUT.gateway.stream_rpc)
else:
resp = await rpc_dispatcher.dispatch(frame, auth_result.roles)
await _send_with_timeout(ws, marshal_frame(resp))
async def _tick_loop(self, conn_id: str, ws: WebSocket) -> None:
try:
while True:
await asyncio.sleep(TICK_INTERVAL_MS / 1000.0)
tick_event = RpcEvent(
event="tick",
data={"ts": int(_time.time() * 1000)},
)
await _send_with_timeout(ws, marshal_frame(tick_event))
except asyncio.CancelledError:
pass
except Exception:
logger.warning("Gateway tick failed for %s", conn_id)
async def _handle_event_subscribe(self, conn_id: str, req: RpcRequest) -> RpcResponse:
events = req.params.get("events", []) if req.params else []
if not isinstance(events, list) or not all(isinstance(e, str) for e in events):
return RpcResponse(
id=req.id,
ok=False,
error_code=GatewayErrorCode.INVALID_PARAMS,
error_message="events 必须是字符串数组",
)
gateway_broadcaster.subscribe(conn_id, events)
subscribed = gateway_broadcaster.get_subscriptions(conn_id)
return RpcResponse(id=req.id, ok=True, result={"subscribed": list(subscribed)})
async def _handle_event_unsubscribe(self, conn_id: str, req: RpcRequest) -> RpcResponse:
events = req.params.get("events", []) if req.params else []
if not isinstance(events, list) or not all(isinstance(e, str) for e in events):
return RpcResponse(
id=req.id,
ok=False,
error_code=GatewayErrorCode.INVALID_PARAMS,
error_message="events 必须是字符串数组",
)
gateway_broadcaster.unsubscribe(conn_id, events)
subscribed = gateway_broadcaster.get_subscriptions(conn_id)
return RpcResponse(id=req.id, ok=True, result={"subscribed": list(subscribed)})
async def send_event(self, conn_id: str, event: RpcEvent) -> bool:
return await gateway_broadcaster.send_event(conn_id, event.event, event.data)
async def broadcast_event(self, event: RpcEvent) -> int:
result = await gateway_broadcaster.broadcast(event.event, event.data)
return result.sent
async def broadcast_event_filtered(
self,
event: str,
data: dict | None = None,
*,
conn_ids: Collection[str] | None = None,
user_ids: Collection[str] | None = None,
roles: Collection[GatewayRole] | None = None,
drop_if_slow: bool = True,
) -> BroadcastResult:
filter_ = BroadcastFilter(
conn_ids=conn_ids,
user_ids=user_ids,
roles=roles,
drop_if_slow=drop_if_slow,
)
return await gateway_broadcaster.broadcast(event, data, filter_)
def register_on_connect(self, cb: OnConnectCallback) -> None:
self._on_connect.append(cb)
def register_on_disconnect(self, cb: OnDisconnectCallback) -> None:
self._on_disconnect.append(cb)
async def shutdown(self, timeout: float = 10.0) -> None:
self._shutting_down = True
logger.info("Gateway WS server shutting down (%d connections)", len(self._active_connections))
await gateway_broadcaster.broadcast("server.shutdown", {"reason": "server_shutdown"})
for task_dict in [self._heartbeat_tasks, self._tick_tasks]:
for conn_id, task in list(task_dict.items()):
task.cancel()
try:
await task
except asyncio.CancelledError:
pass
gateway_broadcaster.unregister_connection(conn_id)
async def _close_conn(conn_id: str, ws: WebSocket) -> None:
try:
await asyncio.wait_for(ws.close(code=1001, reason="server_shutdown"), timeout=timeout)
except Exception:
pass
tasks = [_close_conn(cid, ws) for cid, ws in self._active_connections.items()]
await asyncio.gather(*tasks, return_exceptions=True)
self._active_connections.clear()
self._connection_auth.clear()
self._heartbeat_tasks.clear()
self._tick_tasks.clear()
logger.info("Gateway WS server shut down complete")
gateway_ws_server = GatewayWsServer()
def _resolve_ws_client_ip(ws: WebSocket) -> str | None:
client = getattr(ws, "client", None)
if client and hasattr(client, "host"):
return client.host
return None
def _format_ws_display_host(url_str: str) -> str:
from urllib.parse import urlparse
try:
parsed = urlparse(url_str)
host = parsed.hostname or parsed.netloc or "unknown"
port = parsed.port
if port and port not in (80, 443):
return f"{host}:{port}"
return host
except Exception:
return url_str