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

360 lines
10 KiB
Python

import logging
import time as _time
from collections.abc import Collection
from dataclasses import dataclass, field
from typing import Any
from yuxi.channel.gateway.auth import GatewayAuthResult
from yuxi.channel.gateway.protocol import RpcEvent, marshal_frame
from yuxi.channel.gateway.rbac import ROLE_HIERARCHY, GatewayRole
logger = logging.getLogger(__name__)
DEFAULT_BROADCAST_RATE = 100
DEFAULT_BROADCAST_BURST = 200
DEFAULT_MAX_BUFFERED_FRAMES = 256
DEFAULT_SLOW_CONSUMER_THRESHOLD = 50
PUBLIC_EVENTS: set[str] = {
"heartbeat",
"ping",
"server.shutdown",
"server.tick",
}
READ_EVENTS: set[str] = PUBLIC_EVENTS | {
"channel.created",
"channel.updated",
"channel.deleted",
"channel.state_change",
"channel.message",
"session.created",
"session.updated",
"session.deleted",
"session.message",
"session.tool",
"agent.status",
"agent.event",
"cron.status",
}
WRITE_EVENTS: set[str] = READ_EVENTS | {
"config.changed",
}
ADMIN_EVENTS: set[str] = WRITE_EVENTS | {
"plugin.installed",
"plugin.removed",
"device.pair.requested",
"device.pair.resolved",
}
EVENT_MIN_ROLE: dict[str, GatewayRole] = {}
for _event in PUBLIC_EVENTS:
EVENT_MIN_ROLE[_event] = GatewayRole.VIEWER
for _event in READ_EVENTS - PUBLIC_EVENTS:
EVENT_MIN_ROLE[_event] = GatewayRole.VIEWER
for _event in WRITE_EVENTS - READ_EVENTS:
EVENT_MIN_ROLE[_event] = GatewayRole.OPERATOR
for _event in ADMIN_EVENTS - WRITE_EVENTS:
EVENT_MIN_ROLE[_event] = GatewayRole.ADMIN
@dataclass
class _ConnInfo:
conn_id: str
user_id: str | None = None
roles: list[GatewayRole] = field(default_factory=list)
subscribed_events: set[str] = field(default_factory=set)
pending_frames: int = 0
@dataclass
class BroadcastFilter:
conn_ids: Collection[str] | None = None
user_ids: Collection[str] | None = None
roles: Collection[GatewayRole] | None = None
drop_if_slow: bool = True
@dataclass
class BroadcastResult:
sent: int = 0
dropped: int = 0
filtered: int = 0
class TokenBucket:
def __init__(self, rate: float, burst: float):
self.rate = rate
self.burst = burst
self.tokens = burst
self.last_refill = _time.monotonic()
def consume(self, n: float = 1.0) -> bool:
self._refill()
if self.tokens >= n:
self.tokens -= n
return True
return False
def _refill(self):
now = _time.monotonic()
elapsed = now - self.last_refill
self.tokens = min(self.burst, self.tokens + elapsed * self.rate)
self.last_refill = now
class GatewayBroadcaster:
def __init__(
self,
broadcast_rate: float = DEFAULT_BROADCAST_RATE,
broadcast_burst: float = DEFAULT_BROADCAST_BURST,
max_buffered_frames: int = DEFAULT_MAX_BUFFERED_FRAMES,
slow_consumer_threshold: int = DEFAULT_SLOW_CONSUMER_THRESHOLD,
):
self._connections: dict[str, _ConnInfo] = {}
self._rate_limiter = TokenBucket(broadcast_rate, broadcast_burst)
self._max_buffered_frames = max_buffered_frames
self._slow_consumer_threshold = slow_consumer_threshold
self._send_func: dict[str, Any] = {}
self._global_seq = 0
@property
def global_seq(self) -> int:
return self._global_seq
@property
def connection_count(self) -> int:
return len(self._connections)
def register_connection(
self,
conn_id: str,
auth: GatewayAuthResult,
send_fn: Any,
) -> None:
info = _ConnInfo(
conn_id=conn_id,
user_id=auth.user_id,
roles=list(auth.roles) if auth.roles else [],
)
self._connections[conn_id] = info
self._send_func[conn_id] = send_fn
def unregister_connection(self, conn_id: str) -> None:
self._connections.pop(conn_id, None)
self._send_func.pop(conn_id, None)
def subscribe(self, conn_id: str, events: Collection[str]) -> None:
info = self._connections.get(conn_id)
if info is None:
return
info.subscribed_events.update(events)
def unsubscribe(self, conn_id: str, events: Collection[str]) -> None:
info = self._connections.get(conn_id)
if info is None:
return
info.subscribed_events.difference_update(events)
if not events or "*" in events:
info.subscribed_events.clear()
def clear_subscriptions(self, conn_id: str) -> None:
info = self._connections.get(conn_id)
if info is None:
return
info.subscribed_events.clear()
def get_subscriptions(self, conn_id: str) -> set[str]:
info = self._connections.get(conn_id)
if info is None:
return set()
return set(info.subscribed_events)
async def broadcast(
self,
event: str,
data: dict | None = None,
filter_: BroadcastFilter | None = None,
) -> BroadcastResult:
return await self._broadcast_internal(event, data, filter_)
async def broadcast_to_conn_ids(
self,
event: str,
data: dict | None,
conn_ids: Collection[str],
) -> BroadcastResult:
return await self._broadcast_internal(event, data, BroadcastFilter(conn_ids=conn_ids))
async def broadcast_to_user(
self,
event: str,
data: dict | None,
user_id: str,
) -> BroadcastResult:
return await self._broadcast_internal(event, data, BroadcastFilter(user_ids={user_id}))
async def broadcast_to_role(
self,
event: str,
data: dict | None,
role: GatewayRole,
) -> BroadcastResult:
return await self._broadcast_internal(event, data, BroadcastFilter(roles={role}))
async def send_event(
self,
conn_id: str,
event: str,
data: dict | None = None,
) -> bool:
send_fn = self._send_func.get(conn_id)
if send_fn is None:
return False
self._global_seq += 1
rpc_event = RpcEvent(
event=event,
data=data,
seq=self._global_seq,
state_version=self._resolve_state_version(event),
)
try:
await send_fn(marshal_frame(rpc_event))
return True
except Exception:
return False
async def _broadcast_internal(
self,
event: str,
data: dict | None,
filter_: BroadcastFilter | None,
) -> BroadcastResult:
result = BroadcastResult()
if not self._connections:
return result
candidates = [
conn_id
for conn_id, info in self._connections.items()
if self._passes_filter(info, event, filter_) and self._has_event_scope(info, event)
]
candidate_count = len(candidates)
result.filtered = len(self._connections) - candidate_count
if not candidates:
return result
if not self._rate_limiter.consume(candidate_count):
logger.warning("Broadcast rate limited: event=%s candidates=%d", event, candidate_count)
return result
self._global_seq += 1
rpc_event = RpcEvent(
event=event,
data=data,
seq=self._global_seq,
state_version=self._resolve_state_version(event),
)
frame = marshal_frame(rpc_event)
for conn_id in candidates:
info = self._connections.get(conn_id)
if info is None:
continue
info.pending_frames += 1
is_slow = info.pending_frames > self._slow_consumer_threshold
drop = filter_ is not None and filter_.drop_if_slow
if is_slow and drop:
info.pending_frames = max(0, info.pending_frames - 1)
result.dropped += 1
continue
if info.pending_frames > self._max_buffered_frames:
logger.warning(
"Slow consumer exceeded max buffer: %s (event=%s pending=%d)",
conn_id,
event,
info.pending_frames,
)
info.pending_frames = max(0, info.pending_frames - 1)
result.dropped += 1
continue
send_fn = self._send_func.get(conn_id)
if send_fn is None:
continue
try:
await send_fn(frame)
info.pending_frames = max(0, info.pending_frames - 1)
result.sent += 1
except Exception:
logger.debug("Failed to send event %s to %s", event, conn_id)
return result
def _passes_filter(
self,
info: _ConnInfo,
event: str,
filter_: BroadcastFilter | None,
) -> bool:
if info.subscribed_events:
if event not in info.subscribed_events and "*" not in info.subscribed_events:
return False
if filter_ is None:
return True
if filter_.conn_ids is not None and info.conn_id not in filter_.conn_ids:
return False
if filter_.user_ids is not None and info.user_id not in filter_.user_ids:
return False
if filter_.roles is not None:
user_has_role = any(r in filter_.roles for r in info.roles)
if not user_has_role:
return False
return True
@staticmethod
def _has_event_scope(info: _ConnInfo, event: str) -> bool:
min_role = EVENT_MIN_ROLE.get(event)
if min_role is None:
return True
if not info.roles:
return min_role == GatewayRole.VIEWER
max_level = max(ROLE_HIERARCHY.get(r, 0) for r in info.roles)
required_level = ROLE_HIERARCHY.get(min_role, 0)
return max_level >= required_level
def pending_frame_count(self, conn_id: str) -> int:
info = self._connections.get(conn_id)
if info is None:
return 0
return info.pending_frames
def _resolve_state_version(self, event: str) -> int | None:
versioned_events = {
"channel.state_change",
"config.changed",
"plugin.installed",
"plugin.removed",
}
if event in versioned_events:
return self._global_seq
return None
gateway_broadcaster = GatewayBroadcaster()