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