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

1"""Streaming execution support for agent runtime. 

2 

3Provides streaming capabilities for node-by-node execution, 

4enabling SSE (Server-Sent Events) and progressive response patterns. 

5""" 

6from __future__ import annotations 

7 

8from dataclasses import dataclass, field 

9from typing import Any, AsyncIterator, Callable, Literal 

10from enum import Enum 

11import logging 

12 

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 

19 

20logger = logging.getLogger(__name__) 

21 

22 

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 

32 

33 

34@dataclass 

35class StreamEvent: 

36 """Event emitted during streaming execution. 

37  

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) 

44  

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 

57 

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 

70 

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" 

77 

78 

79@dataclass 

80class NodeExecutor: 

81 """Wrapper for executing a single node. 

82  

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 

91 

92 

93class StreamingRuntime: 

94 """Runtime with streaming execution support. 

95  

96 Enables node-by-node execution with events emitted for each step, 

97 suitable for SSE streaming to clients. 

98  

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

108 

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. 

117  

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"] 

128 

129 def add_node(self, name: str, func: Callable, description: str | None = None) -> "StreamingRuntime": 

130 """Add a node to the execution pipeline (fluent API). 

131  

132 Args: 

133 name: Node name 

134 func: Async function that takes state and returns updates 

135 description: Optional description 

136  

137 Returns: 

138 self for chaining 

139 """ 

140 self.nodes.append(NodeExecutor(name=name, func=func, description=description)) 

141 return self 

142 

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. 

149  

150 Args: 

151 request: Execution request context 

152 initial_state: Optional pre-built initial state 

153  

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 ) 

169 

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

176 

177 # 3. Apply prepare_state hook 

178 state = await self.hooks.prepare_state(state, request) 

179 

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 ) 

188 

189 try: 

190 # Execute node 

191 updates = await node_executor.func(state) 

192 

193 # Apply updates 

194 if updates: 

195 state = apply_slice_updates(state, updates) 

196 

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 ) 

204 

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 

213 

214 # 5. Build final result 

215 result = ExecutionResult.from_state(state) 

216 

217 # 6. Apply after_execution hook 

218 await self.hooks.after_execution(state, result) 

219 

220 # 7. Emit done 

221 yield StreamEvent( 

222 type=StreamEventType.DONE, 

223 data=result.to_response_dict(), 

224 state=state, 

225 ) 

226 

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 ) 

233 

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. 

241  

242 Uses LangGraph's astream() method for true streaming support. 

243  

244 Args: 

245 request: Execution request context 

246 graph: Compiled LangGraph graph 

247 stream_mode: LangGraph stream mode ("values", "updates", "debug") 

248  

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 ) 

261 

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) 

267 

268 # Apply hooks 

269 state = await self.hooks.prepare_state(state, request) 

270 

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 

292 

293 # Build result 

294 result = ExecutionResult.from_state(final_state) 

295 await self.hooks.after_execution(final_state, result) 

296 

297 yield StreamEvent( 

298 type=StreamEventType.DONE, 

299 data=result.to_response_dict(), 

300 ) 

301 

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 ) 

308 

309 

310def create_status_event(message: str) -> StreamEvent: 

311 """Helper to create a status event.""" 

312 return StreamEvent(type=StreamEventType.STATUS, message=message) 

313 

314 

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 ) 

322 

323 

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)