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