ForcePilot/backend/package/yuxi/channels/infra/broadcast.py

64 lines
2.3 KiB
Python
Raw Normal View History

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