"""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()