from __future__ import annotations import asyncio from collections import defaultdict from collections.abc import Callable from typing import Any class EventBroadcaster: def __init__(self): self._subscribers: dict[str, list[asyncio.Queue]] = defaultdict(list) self._callbacks: dict[str, list[Callable]] = defaultdict(list) async def broadcast(self, event: str, payload: Any = None) -> None: for queue in self._subscribers.get(event, []): await queue.put({"event": event, "payload": payload}) for queue in self._subscribers.get("*", []): await queue.put({"event": event, "payload": payload}) for callback in self._callbacks.get(event, []): try: result = callback(event, payload) if asyncio.iscoroutine(result): await result except Exception: pass for callback in self._callbacks.get("*", []): try: result = callback(event, payload) if asyncio.iscoroutine(result): await result except Exception: pass async def node_send_to_session(self, session_key: str, event: str, payload: Any = None) -> None: target = self._subscribers.get(f"session:{session_key}", []) for queue in target: await queue.put({"event": event, "payload": payload}) def subscribe(self, event: str) -> asyncio.Queue: queue: asyncio.Queue = asyncio.Queue() self._subscribers[event].append(queue) return queue def subscribe_callback(self, event: str, callback: Callable) -> None: self._callbacks[event].append(callback) def unsubscribe_callback(self, event: str, callback: Callable) -> None: try: self._callbacks[event].remove(callback) except ValueError: pass def unsubscribe(self, event: str, queue: asyncio.Queue) -> None: try: self._subscribers[event].remove(queue) except ValueError: pass def subscriber_count(self, event: str | None = None) -> int: if event is not None: return len(self._subscribers.get(event, [])) + len(self._callbacks.get(event, [])) return sum(len(qs) for qs in self._subscribers.values()) + sum(len(cbs) for cbs in self._callbacks.values())