rebase
This commit is contained in:
84
app/engine/sse.py
Normal file
84
app/engine/sse.py
Normal file
@@ -0,0 +1,84 @@
|
||||
"""SSE event emitter — wraps sse-starlette to emit typed events."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import uuid
|
||||
from collections.abc import AsyncIterator
|
||||
from typing import Any
|
||||
|
||||
from app.core.logging import get_logger
|
||||
|
||||
_logger = get_logger(__name__)
|
||||
|
||||
|
||||
class SseEmitter:
|
||||
"""Async queue-based SSE emitter.
|
||||
|
||||
Usage:
|
||||
emitter = SseEmitter()
|
||||
async with emitter.stream() as stream:
|
||||
async for event in stream:
|
||||
yield event
|
||||
|
||||
In a producer task:
|
||||
await emitter.emit("tool_call", {...})
|
||||
await emitter.done({"result": "ok"})
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._queue: asyncio.Queue[tuple[str, str, str] | None] = asyncio.Queue()
|
||||
# (event_type, data_json, event_id)
|
||||
self._event_counter = 0
|
||||
self._closed = False
|
||||
|
||||
async def emit(self, event_type: str, data: Any) -> None:
|
||||
if self._closed:
|
||||
return
|
||||
self._event_counter += 1
|
||||
event_id = f"evt_{self._event_counter}"
|
||||
try:
|
||||
data_str = json.dumps(data, ensure_ascii=False, default=str)
|
||||
except (TypeError, ValueError):
|
||||
data_str = json.dumps({"error": "serialization_failed"})
|
||||
await self._queue.put((event_type, data_str, event_id))
|
||||
|
||||
async def ping(self) -> None:
|
||||
await self.emit("ping", {"ts": _now_iso()})
|
||||
|
||||
async def done(self, result: Any = None) -> None:
|
||||
await self.emit("done", result if result is not None else {})
|
||||
await self._queue.put(None) # sentinel
|
||||
self._closed = True
|
||||
|
||||
async def error(self, code: str, message: str, details: Any = None) -> None:
|
||||
payload: dict[str, Any] = {"code": code, "message": message}
|
||||
if details is not None:
|
||||
payload["details"] = details
|
||||
await self.emit("error", payload)
|
||||
await self._queue.put(None)
|
||||
self._closed = True
|
||||
|
||||
async def stream(self) -> AsyncIterator[dict[str, str]]:
|
||||
"""Yield SSE-formatted dicts until the emitter is closed."""
|
||||
try:
|
||||
while True:
|
||||
item = await self._queue.get()
|
||||
if item is None:
|
||||
break
|
||||
event_type, data_str, event_id = item
|
||||
yield {
|
||||
"event": event_type,
|
||||
"data": data_str,
|
||||
"id": event_id,
|
||||
}
|
||||
except asyncio.CancelledError:
|
||||
_logger.info("sse_stream_cancelled")
|
||||
raise
|
||||
|
||||
|
||||
def _now_iso() -> str:
|
||||
from datetime import datetime, timezone
|
||||
|
||||
return datetime.now(timezone.utc).isoformat()
|
||||
Reference in New Issue
Block a user