deerflow-code/offline-backend-20260512/backend/app/gateway/workflow_live_hub.py
2026-09-07 18:24:55 +08:00

60 lines
2.1 KiB
Python

"""In-process fan-out of workflow events to connected SSE clients.
The durable log in ``workflow_run_events`` is the source of truth; this hub only
shortens latency for the tail. Slow consumers lose the *oldest* queued frame
rather than blocking the executor, and the SSE endpoint reconciles any gap from
the database using the ``seq`` cursor.
"""
from __future__ import annotations
import asyncio
import logging
from collections import defaultdict
from collections.abc import Mapping
from typing import Any
logger = logging.getLogger(__name__)
class WorkflowLiveHub:
def __init__(self, *, queue_size: int = 512) -> None:
self._subscribers: dict[str, set[asyncio.Queue[dict[str, Any]]]] = defaultdict(set)
self._queue_size = queue_size
def subscribe(self, run_id: str) -> asyncio.Queue[dict[str, Any]]:
queue: asyncio.Queue[dict[str, Any]] = asyncio.Queue(maxsize=self._queue_size)
self._subscribers[run_id].add(queue)
return queue
def unsubscribe(self, run_id: str, queue: asyncio.Queue[dict[str, Any]]) -> None:
subscribers = self._subscribers.get(run_id)
if not subscribers:
return
subscribers.discard(queue)
if not subscribers:
self._subscribers.pop(run_id, None)
def subscriber_count(self, run_id: str) -> int:
return len(self._subscribers.get(run_id) or ())
async def publish(self, event: Mapping[str, Any]) -> None:
run_id = str(event.get("runId") or "")
subscribers = self._subscribers.get(run_id)
if not run_id or not subscribers:
return
payload = dict(event)
for queue in list(subscribers):
if queue.full():
try:
queue.get_nowait()
except asyncio.QueueEmpty: # pragma: no cover - race with a reader
pass
try:
queue.put_nowait(payload)
except asyncio.QueueFull: # pragma: no cover - race with a reader
logger.debug("workflow live queue still full for run %s", run_id)
__all__ = ["WorkflowLiveHub"]