Coverage for agentos/api/streaming.py: 40%

103 statements  

« prev     ^ index     » next       coverage.py v7.14.3, created at 2026-07-06 21:19 +0800

1""" 

2Streaming SSE (Server-Sent Events) endpoint for agent interactions. 

3 

4Provides real-time streaming of agent outputs via HTTP SSE, enabling 

5browser-based chat UIs and real-time monitoring dashboards. 

6""" 

7 

8from __future__ import annotations 

9 

10import asyncio 

11import json 

12import time 

13from collections import defaultdict 

14from collections.abc import AsyncIterator 

15from dataclasses import dataclass, field 

16from typing import Any 

17 

18 

19@dataclass 

20class StreamEvent: 

21 """Single SSE event emitted by the stream.""" 

22 

23 event: str 

24 """Event type: 'chunk', 'tool_call', 'tool_result', 'done', 'error'.""" 

25 

26 data: dict[str, Any] 

27 """Event payload as JSON-serializable dict.""" 

28 

29 id: str | None = None 

30 """Optional event ID for resume support.""" 

31 

32 retry: int | None = None 

33 """Reconnection retry interval in milliseconds.""" 

34 

35 def to_sse(self) -> str: 

36 """Format as SSE wire format.""" 

37 lines: list[str] = [] 

38 if self.id: 

39 lines.append(f"id: {self.id}") 

40 if self.event: 

41 lines.append(f"event: {self.event}") 

42 lines.append(f"data: {json.dumps(self.data, ensure_ascii=False)}") 

43 if self.retry: 

44 lines.append(f"retry: {self.retry}") 

45 lines.append("") # blank line terminates event 

46 return "\n".join(lines) 

47 

48 

49@dataclass 

50class StreamSession: 

51 """Track an active streaming session.""" 

52 

53 session_id: str 

54 started_at: float = field(default_factory=time.time) 

55 events_emitted: int = 0 

56 last_event_at: float = 0.0 

57 metadata: dict[str, Any] = field(default_factory=dict) 

58 

59 

60class StreamingAgent: 

61 """ 

62 Agent that emits Server-Sent Events for real-time streaming. 

63 

64 Example (FastAPI integration):: 

65 

66 streaming = StreamingAgent(agent_loop) 

67 

68 @app.get("/agent/stream") 

69 async def stream(): 

70 return StreamingResponse( 

71 streaming.stream_chat("What is quantum computing?", "session-1"), 

72 media_type="text/event-stream" 

73 ) 

74 """ 

75 

76 def __init__( 

77 self, 

78 agent_loop: Any = None, 

79 heartbeat_interval: float = 15.0, 

80 ): 

81 """ 

82 Args: 

83 agent_loop: The underlying agent loop (sync or async). 

84 heartbeat_interval: Seconds between heartbeat keepalive events. 

85 """ 

86 self._loop = agent_loop 

87 self._heartbeat = heartbeat_interval 

88 self._sessions: dict[str, StreamSession] = defaultdict(StreamSession) 

89 

90 async def stream_chat( 

91 self, 

92 message: str, 

93 session_id: str = "default", 

94 ) -> AsyncIterator[str]: 

95 """ 

96 Stream a chat interaction as SSE events. 

97 

98 Yields: 

99 SSE-formatted strings suitable for HTTP response body. 

100 """ 

101 session = self._sessions[session_id] 

102 session.session_id = session_id 

103 t_start = time.time() 

104 

105 # Emit start event 

106 yield StreamEvent( 

107 event="start", 

108 data={"session_id": session_id, "message": message}, 

109 ).to_sse() 

110 session.events_emitted += 1 

111 

112 # Simulate streaming chunks (integrate with real agent loop) 

113 chunks = self._generate_chunks(message) 

114 heartbeat_task = asyncio.create_task(self._heartbeat_loop(session_id)) 

115 

116 try: 

117 async for chunk in chunks: 

118 yield StreamEvent( 

119 event="chunk", 

120 data={"content": chunk, "session_id": session_id}, 

121 ).to_sse() 

122 session.events_emitted += 1 

123 session.last_event_at = time.time() 

124 finally: 

125 heartbeat_task.cancel() 

126 try: 

127 await heartbeat_task 

128 except asyncio.CancelledError: 

129 pass 

130 

131 # Emit done event 

132 total_ms = (time.time() - t_start) * 1000 

133 yield StreamEvent( 

134 event="done", 

135 data={ 

136 "session_id": session_id, 

137 "total_latency_ms": total_ms, 

138 "events_emitted": session.events_emitted, 

139 }, 

140 ).to_sse() 

141 

142 def stream_chat_sync(self, message: str, session_id: str = "default"): 

143 """Synchronous wrapper for stream_chat.""" 

144 loop = asyncio.get_event_loop() 

145 return _SyncSSEWrapper(loop.run_until_complete(self._collect_events(message, session_id))) 

146 

147 async def _collect_events(self, message: str, session_id: str) -> list[str]: 

148 events: list[str] = [] 

149 async for sse in self.stream_chat(message, session_id): 

150 events.append(sse) 

151 return events 

152 

153 async def _generate_chunks(self, message: str) -> AsyncIterator[str]: 

154 """Generate streaming text chunks. Override with real LLM integration.""" 

155 if self._loop and hasattr(self._loop, "run"): 

156 # Integrate with actual agent loop 

157 result = self._loop.run(message) 

158 text = str(result.output) if hasattr(result, "output") else str(result) 

159 words = text.split() 

160 for i, word in enumerate(words): 

161 chunk = word + (" " if i < len(words) - 1 else "") 

162 yield chunk 

163 await asyncio.sleep(0.02) # simulate streaming 

164 else: 

165 # Fallback: simulate streaming 

166 words = message.split() 

167 yield f"Processing: {message}\n" 

168 await asyncio.sleep(0.3) 

169 for i in range(3): 

170 yield f"Agent step {i + 1}: analyzing...\n" 

171 await asyncio.sleep(0.5) 

172 yield f"Complete. Response for: {message}" 

173 

174 async def _heartbeat_loop(self, session_id: str) -> None: 

175 """Send periodic heartbeat comments to keep connection alive.""" 

176 while True: 

177 await asyncio.sleep(self._heartbeat) 

178 

179 def emit_tool_call(self, session_id: str, tool_name: str, args: dict) -> str: 

180 """Emit a tool_call SSE event (non-streaming helper).""" 

181 return StreamEvent( 

182 event="tool_call", 

183 data={ 

184 "session_id": session_id, 

185 "tool": tool_name, 

186 "arguments": args, 

187 }, 

188 ).to_sse() 

189 

190 def emit_tool_result(self, session_id: str, tool_name: str, result: Any) -> str: 

191 """Emit a tool_result SSE event.""" 

192 return StreamEvent( 

193 event="tool_result", 

194 data={ 

195 "session_id": session_id, 

196 "tool": tool_name, 

197 "result": result, 

198 }, 

199 ).to_sse() 

200 

201 def emit_error(self, session_id: str, error: str) -> str: 

202 """Emit an error SSE event.""" 

203 return StreamEvent( 

204 event="error", 

205 data={"session_id": session_id, "error": error}, 

206 ).to_sse() 

207 

208 def get_session(self, session_id: str) -> StreamSession | None: 

209 return self._sessions.get(session_id) 

210 

211 def list_sessions(self) -> dict[str, StreamSession]: 

212 return dict(self._sessions) 

213 

214 

215class _SyncSSEWrapper: 

216 """Makes a list of SSE strings iterable for sync streaming.""" 

217 

218 def __init__(self, events: list[str]): 

219 self._events = events 

220 

221 def __iter__(self): 

222 return iter(self._events)