Coverage for src / agent_contracts / runtime / streaming.py: 73%
113 statements
« prev ^ index » next coverage.py v7.13.1, created at 2026-01-09 00:42 +0900
« prev ^ index » next coverage.py v7.13.1, created at 2026-01-09 00:42 +0900
1"""Streaming execution support for agent runtime.
3Provides streaming capabilities for node-by-node execution,
4enabling SSE (Server-Sent Events) and progressive response patterns.
5"""
6from __future__ import annotations
8from dataclasses import dataclass, field
9from typing import Any, AsyncIterator, Callable, Literal
10from enum import Enum
11import logging
13from agent_contracts.runtime.context import RequestContext, ExecutionResult
14from agent_contracts.runtime.hooks import RuntimeHooks, DefaultHooks
15from agent_contracts.runtime.session import SessionStore
16from agent_contracts.runtime.state_ops import create_base_state, merge_session
17from agent_contracts.state_accessors import Internal
18from agent_contracts.state import apply_slice_updates
20logger = logging.getLogger(__name__)
23class StreamEventType(str, Enum):
24 """Types of streaming events."""
25 NODE_START = "node_start" # Node execution starting
26 NODE_END = "node_end" # Node execution completed
27 STATUS = "status" # Status update message
28 PROGRESS = "progress" # Progress indicator
29 DATA = "data" # Intermediate data
30 ERROR = "error" # Error occurred
31 DONE = "done" # Execution complete
34@dataclass
35class StreamEvent:
36 """Event emitted during streaming execution.
38 Attributes:
39 type: Type of event (node_start, node_end, status, etc.)
40 node_name: Name of node (for node events)
41 data: Event data payload
42 message: Human-readable message
43 state: State snapshot (optional)
45 Example:
46 >>> event = StreamEvent(
47 ... type=StreamEventType.NODE_END,
48 ... node_name="search",
49 ... data={"results_count": 10},
50 ... )
51 """
52 type: StreamEventType | str
53 node_name: str | None = None
54 data: dict[str, Any] | None = None
55 message: str | None = None
56 state: dict[str, Any] | None = None
58 def to_dict(self) -> dict[str, Any]:
59 """Convert to dictionary for JSON serialization."""
60 result: dict[str, Any] = {
61 "type": self.type.value if isinstance(self.type, StreamEventType) else self.type,
62 }
63 if self.node_name:
64 result["node_name"] = self.node_name
65 if self.data:
66 result["data"] = self.data
67 if self.message:
68 result["message"] = self.message
69 return result
71 def to_sse(self) -> str:
72 """Format as Server-Sent Event string."""
73 import json
74 event_type = self.type.value if isinstance(self.type, StreamEventType) else self.type
75 data = json.dumps(self.to_dict(), ensure_ascii=False)
76 return f"event: {event_type}\ndata: {data}\n\n"
79@dataclass
80class NodeExecutor:
81 """Wrapper for executing a single node.
83 Attributes:
84 name: Node name
85 func: Async function that takes state and returns updates
86 description: Optional description for status messages
87 """
88 name: str
89 func: Callable[[dict], Any] # async (state) -> updates
90 description: str | None = None
93class StreamingRuntime:
94 """Runtime with streaming execution support.
96 Enables node-by-node execution with events emitted for each step,
97 suitable for SSE streaming to clients.
99 Example:
100 >>> runtime = StreamingRuntime(nodes=[
101 ... NodeExecutor("search", search_node),
102 ... NodeExecutor("stylist", stylist_node),
103 ... ])
104 >>>
105 >>> async for event in runtime.stream(request):
106 ... yield event.to_sse()
107 """
109 def __init__(
110 self,
111 nodes: list[NodeExecutor] | None = None,
112 hooks: RuntimeHooks | None = None,
113 session_store: SessionStore | None = None,
114 slices_to_restore: list[str] | None = None,
115 ) -> None:
116 """Initialize the streaming runtime.
118 Args:
119 nodes: List of node executors to run in sequence
120 hooks: Custom runtime hooks
121 session_store: Session persistence store
122 slices_to_restore: Slice names to restore from session
123 """
124 self.nodes = nodes or []
125 self.hooks = hooks or DefaultHooks()
126 self.session_store = session_store
127 self.slices_to_restore = slices_to_restore or ["_internal", "interview", "shopping"]
129 def add_node(self, name: str, func: Callable, description: str | None = None) -> "StreamingRuntime":
130 """Add a node to the execution pipeline (fluent API).
132 Args:
133 name: Node name
134 func: Async function that takes state and returns updates
135 description: Optional description
137 Returns:
138 self for chaining
139 """
140 self.nodes.append(NodeExecutor(name=name, func=func, description=description))
141 return self
143 async def stream(
144 self,
145 request: RequestContext,
146 initial_state: dict[str, Any] | None = None,
147 ) -> AsyncIterator[StreamEvent]:
148 """Execute nodes and yield streaming events.
150 Args:
151 request: Execution request context
152 initial_state: Optional pre-built initial state
154 Yields:
155 StreamEvent for each execution step
156 """
157 try:
158 # 1. Build initial state
159 if initial_state is not None:
160 state = initial_state
161 else:
162 state = create_base_state(
163 session_id=request.session_id,
164 action=request.action,
165 params=request.params,
166 message=request.message,
167 image=request.image,
168 )
170 # 2. Restore session if resuming
171 if request.resume_session and self.session_store:
172 session_data = await self.session_store.load(request.session_id)
173 if session_data:
174 state = merge_session(state, session_data, self.slices_to_restore)
175 logger.debug(f"Restored session {request.session_id}")
177 # 3. Apply prepare_state hook
178 state = await self.hooks.prepare_state(state, request)
180 # 4. Execute nodes in sequence
181 for node_executor in self.nodes:
182 # Emit node start
183 yield StreamEvent(
184 type=StreamEventType.NODE_START,
185 node_name=node_executor.name,
186 message=node_executor.description or f"Executing {node_executor.name}...",
187 )
189 try:
190 # Execute node
191 updates = await node_executor.func(state)
193 # Apply updates
194 if updates:
195 state = apply_slice_updates(state, updates)
197 # Emit node end
198 yield StreamEvent(
199 type=StreamEventType.NODE_END,
200 node_name=node_executor.name,
201 data=updates,
202 state=state,
203 )
205 except Exception as e:
206 logger.error(f"Node {node_executor.name} failed: {e}", exc_info=True)
207 yield StreamEvent(
208 type=StreamEventType.ERROR,
209 node_name=node_executor.name,
210 message=str(e),
211 )
212 return
214 # 5. Build final result
215 result = ExecutionResult.from_state(state)
217 # 6. Apply after_execution hook
218 await self.hooks.after_execution(state, result)
220 # 7. Emit done
221 yield StreamEvent(
222 type=StreamEventType.DONE,
223 data=result.to_response_dict(),
224 state=state,
225 )
227 except Exception as e:
228 logger.error(f"Streaming execution failed: {e}", exc_info=True)
229 yield StreamEvent(
230 type=StreamEventType.ERROR,
231 message=str(e),
232 )
234 async def stream_with_graph(
235 self,
236 request: RequestContext,
237 graph: Any,
238 stream_mode: str = "updates",
239 ) -> AsyncIterator[StreamEvent]:
240 """Stream execution using LangGraph's native streaming.
242 Uses LangGraph's astream() method for true streaming support.
244 Args:
245 request: Execution request context
246 graph: Compiled LangGraph graph
247 stream_mode: LangGraph stream mode ("values", "updates", "debug")
249 Yields:
250 StreamEvent for each update from the graph
251 """
252 try:
253 # Build initial state
254 state = create_base_state(
255 session_id=request.session_id,
256 action=request.action,
257 params=request.params,
258 message=request.message,
259 image=request.image,
260 )
262 # Restore session
263 if request.resume_session and self.session_store:
264 session_data = await self.session_store.load(request.session_id)
265 if session_data:
266 state = merge_session(state, session_data, self.slices_to_restore)
268 # Apply hooks
269 state = await self.hooks.prepare_state(state, request)
271 # Stream from graph
272 final_state = state
273 async for chunk in graph.astream(state, stream_mode=stream_mode):
274 if stream_mode == "updates":
275 # chunk is dict of {node_name: update}
276 for node_name, update in chunk.items():
277 yield StreamEvent(
278 type=StreamEventType.NODE_END,
279 node_name=node_name,
280 data=update if isinstance(update, dict) else {"value": update},
281 )
282 if isinstance(update, dict):
283 final_state = apply_slice_updates(final_state, update)
284 else:
285 # chunk is the state
286 yield StreamEvent(
287 type=StreamEventType.DATA,
288 data=chunk if isinstance(chunk, dict) else {"value": chunk},
289 )
290 if isinstance(chunk, dict):
291 final_state = chunk
293 # Build result
294 result = ExecutionResult.from_state(final_state)
295 await self.hooks.after_execution(final_state, result)
297 yield StreamEvent(
298 type=StreamEventType.DONE,
299 data=result.to_response_dict(),
300 )
302 except Exception as e:
303 logger.error(f"Graph streaming failed: {e}", exc_info=True)
304 yield StreamEvent(
305 type=StreamEventType.ERROR,
306 message=str(e),
307 )
310def create_status_event(message: str) -> StreamEvent:
311 """Helper to create a status event."""
312 return StreamEvent(type=StreamEventType.STATUS, message=message)
315def create_progress_event(current: int, total: int, message: str | None = None) -> StreamEvent:
316 """Helper to create a progress event."""
317 return StreamEvent(
318 type=StreamEventType.PROGRESS,
319 data={"current": current, "total": total},
320 message=message,
321 )
324def create_data_event(data: dict[str, Any], message: str | None = None) -> StreamEvent:
325 """Helper to create a data event."""
326 return StreamEvent(type=StreamEventType.DATA, data=data, message=message)