pareidolia / server /sse.py
AndresCarreon's picture
PAREIDOLIA v0 — everything secretly has a face
772074b verified
Raw
History Blame Contribute Delete
6.94 kB
"""Asyncio broadcast hub for /api/stream (`awakened` events for the live wall).
Copied from godseed/server/sse.py — the proven June-12 SSEHub — with only the
logger renamed. Credit where due: design and hard-won constraints are godseed's.
One hub, many subscribers; every subscriber receives every published event.
Design constraints, in order of importance:
* The producer (a publish landing on the wall) must NEVER block or fail
because a browser tab is slow: each subscriber gets a bounded queue and the
OLDEST event is dropped on overflow. A reconnecting client re-syncs via
GET /api/menagerie.
* Idle connections receive ``{"type": "heartbeat"}`` every 15s so proxies
(and the HF Space router) keep the stream open.
* Disconnects clean up after themselves: the wire generator unsubscribes in a
``finally`` block, which Starlette triggers on client disconnect.
* Fan-out is bounded (security review #3): each subscriber costs a 256-slot
queue plus a generator, so ``subscribe`` refuses past a global ceiling
(and a smaller per-key/IP ceiling) instead of growing without bound.
"""
from __future__ import annotations
import asyncio
import json
import logging
from typing import Any, AsyncIterator, Iterable, Optional
log = logging.getLogger("pareidolia.sse")
HEARTBEAT_SECONDS = 15.0
DEFAULT_QUEUE_SIZE = 256
# Security review #3: ceilings on concurrent SSE subscribers. 256 streams is
# far beyond any plausible hackathon wall audience; 8 per IP allows a NAT'd
# classroom while stopping one host from holding every slot.
DEFAULT_MAX_SUBSCRIBERS = 256
DEFAULT_MAX_PER_KEY = 8
def format_sse(event: dict[str, Any]) -> str:
"""Wire-encode one event. Events are plain JSON on the default `message`
channel; clients dispatch on ``data.type``."""
return f"data: {json.dumps(event, ensure_ascii=False, separators=(',', ':'))}\n\n"
class SSEHub:
"""Fan-out broadcast hub. ``publish`` is synchronous and non-blocking;
``emit`` is the async alias for emit-callback call sites."""
def __init__(
self,
heartbeat: float = HEARTBEAT_SECONDS,
max_queue: int = DEFAULT_QUEUE_SIZE,
max_subscribers: int = DEFAULT_MAX_SUBSCRIBERS,
max_per_key: int = DEFAULT_MAX_PER_KEY,
) -> None:
self.heartbeat = heartbeat
self.max_queue = max_queue
self.max_subscribers = max_subscribers
self.max_per_key = max_per_key
# queue -> subscriber key (None for keyless/internal subscriptions)
self._subscribers: dict[asyncio.Queue[dict[str, Any]], Optional[str]] = {}
self._per_key: dict[str, int] = {}
@property
def subscriber_count(self) -> int:
return len(self._subscribers)
# ------------------------------------------------------------------ producer side
def publish(self, event: dict[str, Any]) -> None:
"""Deliver ``event`` to every subscriber. Never blocks: when a
subscriber's queue is full, its oldest pending event is dropped."""
for queue in tuple(self._subscribers):
try:
queue.put_nowait(event)
except asyncio.QueueFull:
try:
queue.get_nowait() # drop the oldest event for this slow client
except asyncio.QueueEmpty: # pragma: no cover — only under a race
pass
try:
queue.put_nowait(event)
except asyncio.QueueFull: # pragma: no cover — only under a race
log.debug("subscriber queue still full; event dropped")
async def emit(self, event: dict[str, Any]) -> None:
"""Async alias of :meth:`publish`."""
self.publish(event)
# ---------------------------------------------------------------- subscriber side
def subscribe(
self, key: Optional[str] = None
) -> Optional[asyncio.Queue[dict[str, Any]]]:
"""Reserve a subscriber slot; returns None when the global ceiling —
or, when ``key`` is given, that key's ceiling — is already full
(security review #3). The HTTP layer turns None into a 503 BEFORE the
stream starts, so the refusal is a real status code, not a dead pipe."""
if len(self._subscribers) >= self.max_subscribers:
log.warning("SSE subscriber ceiling (%d) reached", self.max_subscribers)
return None
if key is not None and self._per_key.get(key, 0) >= self.max_per_key:
log.warning("SSE per-key ceiling (%d) reached for %s", self.max_per_key, key)
return None
queue: asyncio.Queue[dict[str, Any]] = asyncio.Queue(maxsize=self.max_queue)
self._subscribers[queue] = key
if key is not None:
self._per_key[key] = self._per_key.get(key, 0) + 1
return queue
def unsubscribe(self, queue: asyncio.Queue[dict[str, Any]]) -> None:
key = self._subscribers.pop(queue, None)
if key is not None:
count = self._per_key.get(key, 0) - 1
if count > 0:
self._per_key[key] = count
else:
self._per_key.pop(key, None)
async def event_stream(
self,
initial_events: Iterable[dict[str, Any]] = (),
limit: Optional[int] = None,
queue: Optional[asyncio.Queue[dict[str, Any]]] = None,
) -> AsyncIterator[str]:
"""Async generator yielding wire-format SSE for one client.
``initial_events`` (hello snapshot) are sent first, then every broadcast
event; a heartbeat is sent whenever ``self.heartbeat`` seconds pass
without traffic. ``queue`` is the slot reserved by :meth:`subscribe` —
the HTTP handler subscribes BEFORE building the response so capacity
refusals can be a 503; passing none subscribes here (keyless), and the
stream ends immediately if the hub is at its ceiling.
``limit`` (optional) closes the stream after that many events — a
debug/test aid (`curl '/api/stream?limit=20'`); browsers stream unbounded.
"""
if queue is None:
queue = self.subscribe()
if queue is None:
return
sent = 0
try:
yield "retry: 3000\n\n" # client auto-reconnect hint
for event in initial_events:
yield format_sse(event)
sent += 1
if limit is not None and sent >= limit:
return
while True:
try:
event = await asyncio.wait_for(queue.get(), timeout=self.heartbeat)
except asyncio.TimeoutError:
event = {"type": "heartbeat"}
yield format_sse(event)
sent += 1
if limit is not None and sent >= limit:
return
finally:
self.unsubscribe(queue)