Coverage for agentos/api/sse.py: 47%

110 statements  

« prev     ^ index     » next       coverage.py v7.14.3, created at 2026-07-08 10:59 +0800

1""" 

2SSE (Server-Sent Events) Streaming — production-grade async streaming endpoint. 

3 

4Provides an ASGI-compatible SSE stream with automatic reconnection, 

5client heartbeat, backpressure control, and typed event dispatching. 

6""" 

7 

8import asyncio 

9import json 

10import time 

11from collections.abc import AsyncIterator 

12from dataclasses import dataclass 

13from enum import StrEnum 

14from typing import Any 

15 

16DEFAULT_RETRY_MS = 3000 

17DEFAULT_HEARTBEAT_S = 30 

18MAX_QUEUE_SIZE = 256 

19 

20 

21class SSEEventType(StrEnum): 

22 """Standard SSE event types plus AgentOS extensions.""" 

23 

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" 

33 

34 

35@dataclass 

36class SSEEvent: 

37 """A single SSE event to be serialized to the wire.""" 

38 

39 event: str = SSEEventType.MESSAGE 

40 data: Any = "" 

41 id: str = "" 

42 retry: int = DEFAULT_RETRY_MS 

43 

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

53 

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) 

58 

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" 

63 

64 @classmethod 

65 def token(cls, text: str, seq: int = 0) -> "SSEEvent": 

66 return cls(event=SSEEventType.TOKEN, data={"text": text, "seq": seq}) 

67 

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 ) 

74 

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 ) 

81 

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 ) 

88 

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 ) 

95 

96 @classmethod 

97 def metadata(cls, meta: dict[str, Any]) -> "SSEEvent": 

98 return cls(event=SSEEventType.METADATA, data=meta) 

99 

100 

101class SSEStream: 

102 """SSE stream with heartbeats and backpressure handling. 

103 

104 Usage:: 

105 

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

111 

112 # Consumer (ASGI) 

113 async for chunk in stream.iter_chunks(): 

114 yield chunk 

115 """ 

116 

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 

129 

130 async def start(self): 

131 """Start the heartbeat background task.""" 

132 self._heartbeat_task = asyncio.create_task(self._heartbeat_loop()) 

133 

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 

148 

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) 

157 

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 

168 

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 

176 

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

181 

182 

183class SSEResponse: 

184 """Factory for generating ASGI-compatible SSE HTTP responses. 

185 

186 Usage (Starlette / FastAPI):: 

187 

188 from starlette.responses import StreamingResponse 

189 

190 sse = SSEResponse(stream) 

191 return StreamingResponse( 

192 sse.body(), 

193 media_type="text/event-stream", 

194 headers=sse.headers(), 

195 ) 

196 """ 

197 

198 HEADERS = { 

199 "Content-Type": "text/event-stream", 

200 "Cache-Control": "no-cache", 

201 "Connection": "keep-alive", 

202 "X-Accel-Buffering": "no", 

203 } 

204 

205 def __init__(self, stream: SSEStream): 

206 self.stream = stream 

207 

208 def headers(self) -> dict[str, str]: 

209 return dict(self.HEADERS) 

210 

211 async def body(self) -> AsyncIterator[str]: 

212 """ASGI-compatible body iterator.""" 

213 async for chunk in self.stream.iter_chunks(): 

214 yield chunk