Coverage for agentos/api/streaming.py: 40%
103 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-09 07:12 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-09 07:12 +0800
1"""
2Streaming SSE (Server-Sent Events) endpoint for agent interactions.
4Provides real-time streaming of agent outputs via HTTP SSE, enabling
5browser-based chat UIs and real-time monitoring dashboards.
6"""
8from __future__ import annotations
10import asyncio
11import json
12import time
13from collections import defaultdict
14from collections.abc import AsyncIterator
15from dataclasses import dataclass, field
16from typing import Any
19@dataclass
20class StreamEvent:
21 """Single SSE event emitted by the stream."""
23 event: str
24 """Event type: 'chunk', 'tool_call', 'tool_result', 'done', 'error'."""
26 data: dict[str, Any]
27 """Event payload as JSON-serializable dict."""
29 id: str | None = None
30 """Optional event ID for resume support."""
32 retry: int | None = None
33 """Reconnection retry interval in milliseconds."""
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)
49@dataclass
50class StreamSession:
51 """Track an active streaming session."""
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)
60class StreamingAgent:
61 """
62 Agent that emits Server-Sent Events for real-time streaming.
64 Example (FastAPI integration)::
66 streaming = StreamingAgent(agent_loop)
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 """
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)
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.
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()
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
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))
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
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()
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)))
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
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}"
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)
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()
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()
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()
208 def get_session(self, session_id: str) -> StreamSession | None:
209 return self._sessions.get(session_id)
211 def list_sessions(self) -> dict[str, StreamSession]:
212 return dict(self._sessions)
215class _SyncSSEWrapper:
216 """Makes a list of SSE strings iterable for sync streaming."""
218 def __init__(self, events: list[str]):
219 self._events = events
221 def __iter__(self):
222 return iter(self._events)