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

73 lines
2.7 KiB
Python

"""In-process fan-out for low-latency Deep Research SSE updates.
The event repository remains the durable recovery log. This app-layer hub only
delivers best-effort live model deltas to connections served by this worker.
The persisted report checkpoints keep reconnect recovery correct.
"""
from __future__ import annotations
import asyncio
import uuid
from collections import defaultdict
from collections.abc import Mapping
from typing import Any
class DeepResearchLiveHub:
"""Best-effort per-job fan-out that never blocks model generation."""
def __init__(self, *, queue_size: int = 512) -> None:
self._queue_size = queue_size
self._subscribers: dict[str, set[asyncio.Queue[dict[str, Any]]]] = defaultdict(set)
def subscribe(self, job_id: str) -> asyncio.Queue[dict[str, Any]]:
queue: asyncio.Queue[dict[str, Any]] = asyncio.Queue(maxsize=self._queue_size)
self._subscribers[job_id].add(queue)
return queue
def has_subscribers(self, job_id: str) -> bool:
"""Whether an SSE connection for this job is attached to this worker.
Deployments run several uvicorn workers, and a job runs on whichever
worker won its durable lease while ``GET /jobs/{id}/stream`` is served
by whichever worker accepted that connection. When those differ, this
in-process hub cannot reach the viewer at all, and a live-only frame
(model thinking) would be lost with no durable row to replay. Callers
use this probe to mirror those frames into the event log instead.
"""
return bool(self._subscribers.get(job_id))
def unsubscribe(self, job_id: str, queue: asyncio.Queue[dict[str, Any]]) -> None:
subscribers = self._subscribers.get(job_id)
if not subscribers:
return
subscribers.discard(queue)
if not subscribers:
self._subscribers.pop(job_id, None)
async def publish(self, event: Mapping[str, Any]) -> None:
"""Publish a persisted or ephemeral envelope without waiting on a client."""
job_id = str(event.get("jobId") or "")
if not job_id:
return
subscribers = tuple(self._subscribers.get(job_id, ()))
if not subscribers:
return
payload = dict(event)
payload.setdefault("liveId", uuid.uuid4().hex)
for queue in subscribers:
if queue.full():
try:
queue.get_nowait()
except asyncio.QueueEmpty:
pass
try:
queue.put_nowait(payload)
except asyncio.QueueFull:
# A durable checkpoint will repair any skipped live delta.
continue
__all__ = ["DeepResearchLiveHub"]