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
« 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.
4Provides push-based task lifecycle notifications so agents don't poll.
5"""
7from __future__ import annotations
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
17from agentos.protocols.a2a import A2ATask, TaskState
20class A2AStreamEvent(StrEnum):
21 """A2A-specific streaming event types."""
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"
33@dataclass
34class TaskProgress:
35 """Progress update within a running task."""
37 percent: float = 0.0
38 message: str = ""
39 step: str = ""
40 metadata: dict[str, Any] = field(default_factory=dict)
43class A2AStreamSession:
44 """Manages a streaming connection for a single task.
46 Agents subscribe to receive push updates as the task progresses.
47 """
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
55 async def start(self, heartbeat_s: float = 30.0):
56 """Start heartbeat loop."""
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 )
70 self._heartbeat_task = asyncio.create_task(_pulse())
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
78 def unsubscribe(self, sub: asyncio.Queue):
79 """Remove a subscriber."""
80 try:
81 self._subscribers.remove(sub)
82 except ValueError:
83 pass
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)
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)
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()
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
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"
139class StreamingAggregator:
140 """流式结果聚合器 — 合规测试套件要求。"""
142 def __init__(self):
143 self._chunks: list[str] = []
145 def collect(self, chunk: str) -> None:
146 self._chunks.append(chunk)
148 def aggregated(self) -> str:
149 return "".join(self._chunks)
152class A2AStreamManager:
153 """Global manager for A2A task streaming sessions.
155 Tracks all active task streams and dispatches events on state transitions.
156 """
158 def __init__(self):
159 self._sessions: dict[str, A2AStreamSession] = {}
160 self._on_state_change: Callable | None = None
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
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
172 def get_session(self, task_id: str) -> A2AStreamSession | None:
173 return self._sessions.get(task_id)
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
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 )
198 if task.is_terminal():
199 await session.close()
200 del self._sessions[task.task_id]
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 )
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 )
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()