Coverage for agentos/protocols/a2a_streaming.py: 37%

127 statements  

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

1""" 

2A2A Streaming — real-time task status updates via SSE for A2A protocol. 

3 

4Provides push-based task lifecycle notifications so agents don't poll. 

5""" 

6 

7from __future__ import annotations 

8 

9import asyncio 

10import json 

11import time 

12from collections.abc import AsyncIterator, Callable 

13from dataclasses import dataclass, field 

14from enum import StrEnum 

15from typing import Any 

16 

17from agentos.protocols.a2a import A2ATask, TaskState 

18 

19 

20class A2AStreamEvent(StrEnum): 

21 """A2A-specific streaming event types.""" 

22 

23 TASK_CREATED = "task.created" 

24 TASK_STARTED = "task.started" 

25 TASK_PROGRESS = "task.progress" 

26 TASK_COMPLETED = "task.completed" 

27 TASK_FAILED = "task.failed" 

28 TASK_CANCELLED = "task.cancelled" 

29 ARTIFACT_ADDED = "artifact.added" 

30 HEARTBEAT = "heartbeat" 

31 

32 

33@dataclass 

34class TaskProgress: 

35 """Progress update within a running task.""" 

36 

37 percent: float = 0.0 

38 message: str = "" 

39 step: str = "" 

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

41 

42 

43class A2AStreamSession: 

44 """Manages a streaming connection for a single task. 

45 

46 Agents subscribe to receive push updates as the task progresses. 

47 """ 

48 

49 def __init__(self, task: A2ATask): 

50 self.task_id = task.task_id 

51 self._subscribers: list[asyncio.Queue[dict]] = [] 

52 self._closed = False 

53 self._heartbeat_task: asyncio.Task | None = None 

54 

55 async def start(self, heartbeat_s: float = 30.0): 

56 """Start heartbeat loop.""" 

57 

58 async def _pulse(): 

59 while not self._closed: 

60 await asyncio.sleep(heartbeat_s) 

61 if not self._closed: 

62 await self._broadcast( 

63 { 

64 "event": A2AStreamEvent.HEARTBEAT, 

65 "task_id": self.task_id, 

66 "timestamp": time.time(), 

67 } 

68 ) 

69 

70 self._heartbeat_task = asyncio.create_task(_pulse()) 

71 

72 def subscribe(self) -> asyncio.Queue[dict]: 

73 """Register a new subscriber. Returns a queue of SSE events.""" 

74 q: asyncio.Queue[dict] = asyncio.Queue(maxsize=64) 

75 self._subscribers.append(q) 

76 return q 

77 

78 def unsubscribe(self, sub: asyncio.Queue): 

79 """Remove a subscriber.""" 

80 try: 

81 self._subscribers.remove(sub) 

82 except ValueError: 

83 pass 

84 

85 async def emit(self, event: A2AStreamEvent, data: dict | None = None): 

86 """Push an event to all subscribers.""" 

87 payload = { 

88 "event": event.value, 

89 "task_id": self.task_id, 

90 "timestamp": time.time(), 

91 } 

92 if data: 

93 payload["data"] = data 

94 await self._broadcast(payload) 

95 

96 async def _broadcast(self, payload: dict): 

97 dead: list[asyncio.Queue] = [] 

98 for q in self._subscribers: 

99 try: 

100 q.put_nowait(payload) 

101 except asyncio.QueueFull: 

102 dead.append(q) 

103 for q in dead: 

104 self.unsubscribe(q) 

105 

106 async def close(self): 

107 """Shut down the stream.""" 

108 self._closed = True 

109 if self._heartbeat_task: 

110 self._heartbeat_task.cancel() 

111 # Close all subscriber queues 

112 for q in self._subscribers: 

113 try: 

114 q.put_nowait(None) # Sentinel 

115 except asyncio.QueueFull: 

116 pass 

117 self._subscribers.clear() 

118 

119 async def iter_events(self, subscriber: asyncio.Queue) -> AsyncIterator[dict]: 

120 """Async iterator yielding SSE-compatible event dicts.""" 

121 while True: 

122 event = await subscriber.get() 

123 if event is None: 

124 break 

125 yield event 

126 

127 def to_sse(self, event: dict) -> str: 

128 """Format a single event dict into SSE wire format.""" 

129 lines: list[str] = [f"event: {event['event']}"] 

130 for key in ("task_id", "timestamp"): 

131 if key in event: 

132 lines.append(f"id: {key}={event[key]}") 

133 data_str = json.dumps(event.get("data", {}), ensure_ascii=False) 

134 for line in data_str.split("\n"): 

135 lines.append(f"data: {line}") 

136 return "\n".join(lines) + "\n\n" 

137 

138 

139class StreamingAggregator: 

140 """流式结果聚合器 — 合规测试套件要求。""" 

141 

142 def __init__(self): 

143 self._chunks: list[str] = [] 

144 

145 def collect(self, chunk: str) -> None: 

146 self._chunks.append(chunk) 

147 

148 def aggregated(self) -> str: 

149 return "".join(self._chunks) 

150 

151 

152class A2AStreamManager: 

153 """Global manager for A2A task streaming sessions. 

154 

155 Tracks all active task streams and dispatches events on state transitions. 

156 """ 

157 

158 def __init__(self): 

159 self._sessions: dict[str, A2AStreamSession] = {} 

160 self._on_state_change: Callable | None = None 

161 

162 def on_state_change(self, callback: Callable[[A2ATask, TaskState, TaskState], Any]): 

163 """Register a hook called on every state transition (old_state, new_state).""" 

164 self._on_state_change = callback 

165 

166 def create_session(self, task: A2ATask) -> A2AStreamSession: 

167 """Create a streaming session for a new task.""" 

168 session = A2AStreamSession(task) 

169 self._sessions[task.task_id] = session 

170 return session 

171 

172 def get_session(self, task_id: str) -> A2AStreamSession | None: 

173 return self._sessions.get(task_id) 

174 

175 async def notify_state_change(self, task: A2ATask, old_state: TaskState): 

176 """Called when a task transitions state.""" 

177 session = self._sessions.get(task.task_id) 

178 if not session: 

179 return 

180 

181 event_map = { 

182 TaskState.SUBMITTED: A2AStreamEvent.TASK_CREATED, 

183 TaskState.WORKING: A2AStreamEvent.TASK_STARTED, 

184 TaskState.COMPLETED: A2AStreamEvent.TASK_COMPLETED, 

185 TaskState.FAILED: A2AStreamEvent.TASK_FAILED, 

186 TaskState.CANCELLED: A2AStreamEvent.TASK_CANCELLED, 

187 } 

188 event = event_map.get(task.state, A2AStreamEvent.TASK_PROGRESS) 

189 await session.emit( 

190 event, 

191 { 

192 "previous_state": old_state.value, 

193 "current_state": task.state.value, 

194 "error": task.error, 

195 }, 

196 ) 

197 

198 if task.is_terminal(): 

199 await session.close() 

200 del self._sessions[task.task_id] 

201 

202 async def notify_artifact(self, task_id: str, artifact_name: str): 

203 """Called when an artifact is added to a task.""" 

204 session = self._sessions.get(task_id) 

205 if session: 

206 await session.emit( 

207 A2AStreamEvent.ARTIFACT_ADDED, 

208 { 

209 "artifact_name": artifact_name, 

210 }, 

211 ) 

212 

213 async def notify_progress( 

214 self, 

215 task_id: str, 

216 progress: TaskProgress, 

217 ): 

218 """Push a progress update to subscribers.""" 

219 session = self._sessions.get(task_id) 

220 if session: 

221 await session.emit( 

222 A2AStreamEvent.TASK_PROGRESS, 

223 { 

224 "percent": progress.percent, 

225 "message": progress.message, 

226 "step": progress.step, 

227 "metadata": progress.metadata, 

228 }, 

229 ) 

230 

231 async def shutdown(self): 

232 """Gracefully close all sessions.""" 

233 for sid in list(self._sessions.keys()): 

234 session = self._sessions[sid] 

235 await session.close() 

236 self._sessions.clear()