64 lines
2.3 KiB
Python
64 lines
2.3 KiB
Python
|
|
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())
|