Coverage for agentos/api/sse.py: 47%
110 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-08 21:26 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-08 21:26 +0800
1"""
2SSE (Server-Sent Events) Streaming — production-grade async streaming endpoint.
4Provides an ASGI-compatible SSE stream with automatic reconnection,
5client heartbeat, backpressure control, and typed event dispatching.
6"""
8import asyncio
9import json
10import time
11from collections.abc import AsyncIterator
12from dataclasses import dataclass
13from enum import StrEnum
14from typing import Any
16DEFAULT_RETRY_MS = 3000
17DEFAULT_HEARTBEAT_S = 30
18MAX_QUEUE_SIZE = 256
21class SSEEventType(StrEnum):
22 """Standard SSE event types plus AgentOS extensions."""
24 MESSAGE = "message"
25 TOKEN = "token"
26 TOOL_CALL = "tool_call"
27 TOOL_RESULT = "tool_result"
28 ERROR = "error"
29 DONE = "done"
30 PING = "ping"
31 HEARTBEAT = "heartbeat"
32 METADATA = "metadata"
35@dataclass
36class SSEEvent:
37 """A single SSE event to be serialized to the wire."""
39 event: str = SSEEventType.MESSAGE
40 data: Any = ""
41 id: str = ""
42 retry: int = DEFAULT_RETRY_MS
44 def serialize(self) -> str:
45 """Serialize to raw SSE wire format."""
46 lines: list[str] = []
47 if self.event:
48 lines.append(f"event: {self.event.value}")
49 if self.id:
50 lines.append(f"id: {self.id}")
51 if self.retry != DEFAULT_RETRY_MS:
52 lines.append(f"retry: {self.retry}")
54 if isinstance(self.data, (dict, list)):
55 data_str = json.dumps(self.data, ensure_ascii=False)
56 else:
57 data_str = str(self.data)
59 # Multi-line data
60 for line in data_str.split("\n"):
61 lines.append(f"data: {line}")
62 return "\n".join(lines) + "\n\n"
64 @classmethod
65 def token(cls, text: str, seq: int = 0) -> "SSEEvent":
66 return cls(event=SSEEventType.TOKEN, data={"text": text, "seq": seq})
68 @classmethod
69 def tool_call(cls, name: str, args: dict) -> "SSEEvent":
70 return cls(
71 event=SSEEventType.TOOL_CALL,
72 data={"name": name, "arguments": args},
73 )
75 @classmethod
76 def tool_result(cls, name: str, result: Any) -> "SSEEvent":
77 return cls(
78 event=SSEEventType.TOOL_RESULT,
79 data={"name": name, "result": result},
80 )
82 @classmethod
83 def error(cls, message: str, code: str = "UNKNOWN") -> "SSEEvent":
84 return cls(
85 event=SSEEventType.ERROR,
86 data={"message": message, "code": code},
87 )
89 @classmethod
90 def done(cls, metadata: dict[str, Any] | None = None) -> "SSEEvent":
91 return cls(
92 event=SSEEventType.DONE,
93 data=metadata or {},
94 )
96 @classmethod
97 def metadata(cls, meta: dict[str, Any]) -> "SSEEvent":
98 return cls(event=SSEEventType.METADATA, data=meta)
101class SSEStream:
102 """SSE stream with heartbeats and backpressure handling.
104 Usage::
106 stream = SSEStream(retry_ms=3000)
107 # Producer
108 await stream.queue.put(SSEEvent.token("Hello"))
109 await stream.queue.put(SSEEvent.done())
110 await stream.close()
112 # Consumer (ASGI)
113 async for chunk in stream.iter_chunks():
114 yield chunk
115 """
117 def __init__(
118 self,
119 retry_ms: int = DEFAULT_RETRY_MS,
120 heartbeat_s: float = DEFAULT_HEARTBEAT_S,
121 max_queue: int = MAX_QUEUE_SIZE,
122 ):
123 self.retry_ms = retry_ms
124 self.heartbeat_s = heartbeat_s
125 self.queue: asyncio.Queue[SSEEvent | None] = asyncio.Queue(maxsize=max_queue)
126 self._closed = False
127 self._heartbeat_task: asyncio.Task | None = None
128 self._last_event_id = 0
130 async def start(self):
131 """Start the heartbeat background task."""
132 self._heartbeat_task = asyncio.create_task(self._heartbeat_loop())
134 async def _heartbeat_loop(self):
135 """Send periodic heartbeat pings."""
136 try:
137 while not self._closed:
138 await asyncio.sleep(self.heartbeat_s)
139 if not self._closed:
140 await self.queue.put(
141 SSEEvent(
142 event=SSEEventType.HEARTBEAT,
143 data={"ts": time.time()},
144 )
145 )
146 except asyncio.CancelledError:
147 pass
149 async def send(self, event: SSEEvent):
150 """Enqueue an event. Raises QueueFull if backpressure exceeded."""
151 if self._closed:
152 raise RuntimeError("Stream is closed")
153 self._last_event_id += 1
154 if not event.id:
155 event.id = str(self._last_event_id)
156 self.queue.put_nowait(event)
158 async def close(self):
159 """Signal end of stream."""
160 self._closed = True
161 await self.queue.put(None) # Sentinel
162 if self._heartbeat_task:
163 self._heartbeat_task.cancel()
164 try:
165 await self._heartbeat_task
166 except asyncio.CancelledError:
167 pass
169 async def iter_events(self) -> AsyncIterator[SSEEvent]:
170 """Async iterator over enqueued events."""
171 while True:
172 event = await self.queue.get()
173 if event is None:
174 break
175 yield event
177 async def iter_chunks(self) -> AsyncIterator[str]:
178 """Async iterator yielding raw SSE wire-format chunks."""
179 async for event in self.iter_events():
180 yield event.serialize()
183class SSEResponse:
184 """Factory for generating ASGI-compatible SSE HTTP responses.
186 Usage (Starlette / FastAPI)::
188 from starlette.responses import StreamingResponse
190 sse = SSEResponse(stream)
191 return StreamingResponse(
192 sse.body(),
193 media_type="text/event-stream",
194 headers=sse.headers(),
195 )
196 """
198 HEADERS = {
199 "Content-Type": "text/event-stream",
200 "Cache-Control": "no-cache",
201 "Connection": "keep-alive",
202 "X-Accel-Buffering": "no",
203 }
205 def __init__(self, stream: SSEStream):
206 self.stream = stream
208 def headers(self) -> dict[str, str]:
209 return dict(self.HEADERS)
211 async def body(self) -> AsyncIterator[str]:
212 """ASGI-compatible body iterator."""
213 async for chunk in self.stream.iter_chunks():
214 yield chunk