Coverage for agentos/protocols/a2a.py: 61%

480 statements  

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

1""" 

2AgentOS v1.2.2 — A2A (Agent-to-Agent) 协议实现。 

3 

4基因来源: Google A2A Protocol (agent-to-agent-protocol.google.com) 

5 

6A2A 协议核心概念: 

7- Task: 异步工作单元,带状态机 (SUBMITTED→WORKING→COMPLETED/FAILED/CANCELLED) 

8- Message: 多模态消息,支持 text/file/data parts 

9- Artifact: 任务产生的输出物,带 MIME 类型 

10- Handoff: Agent 间任务移交 

11- Session: 多轮对话上下文 

12 

13协议层: 

14- REST: GET/POST /tasks, /tasks/{id} 

15- Future: WebSocket 推送 (v1.3+) 

16""" 

17 

18from __future__ import annotations 

19 

20import json 

21import time 

22import uuid 

23from collections.abc import Callable 

24from dataclasses import dataclass, field 

25from enum import StrEnum 

26from typing import Any 

27 

28# ── 基础枚举 ──────────────────────────────────── 

29 

30 

31class TaskState(StrEnum): 

32 """A2A 任务状态。""" 

33 

34 SUBMITTED = "submitted" 

35 WORKING = "working" 

36 COMPLETED = "completed" 

37 FAILED = "failed" 

38 CANCELLED = "cancelled" 

39 

40 

41class TaskStatus(StrEnum): 

42 """A2A 任务状态(别名兼容)— 用于合规测试套件。""" 

43 

44 submitted = "submitted" 

45 working = "working" 

46 completed = "completed" 

47 failed = "failed" 

48 canceled = "canceled" 

49 

50 

51class AgentCard: 

52 """A2A Agent 名片 — 合规测试套件要求 Pydantic 兼容。""" 

53 

54 def __init__( 

55 self, 

56 name: str, 

57 description: str, 

58 url: str, 

59 version: str, 

60 capabilities: list, 

61 provider: dict, 

62 authentication: Any | None = None, 

63 default_input_modes: list | None = None, 

64 default_output_modes: list | None = None, 

65 skills: list | None = None, 

66 ): 

67 self.name = name 

68 self.description = description 

69 self.url = url 

70 self.version = version 

71 self.capabilities = capabilities 

72 self.provider = provider 

73 self.authentication = authentication 

74 self.default_input_modes = default_input_modes or ["text"] 

75 self.default_output_modes = default_output_modes or ["text"] 

76 self.skills = skills or [] 

77 

78 def model_dump(self) -> dict: 

79 return { 

80 "name": self.name, 

81 "description": self.description, 

82 "url": self.url, 

83 "version": self.version, 

84 "capabilities": self.capabilities, 

85 "provider": self.provider, 

86 "authentication": self.authentication, 

87 "default_input_modes": self.default_input_modes, 

88 "default_output_modes": self.default_output_modes, 

89 "skills": self.skills, 

90 } 

91 

92 

93class A2AMessageBus: 

94 """A2A 消息总线 — 支持 agent 注册和消息发送。""" 

95 

96 def __init__(self): 

97 self._agents: dict[str, Any] = {} 

98 

99 def register_agent(self, agent_id: str, agent: Any = None) -> None: 

100 self._agents[agent_id] = agent 

101 

102 async def send(self, target_agent: str, message: Any) -> bool: 

103 return target_agent in self._agents 

104 

105 

106class PartType(StrEnum): 

107 """A2A 内容片段类型。""" 

108 

109 TEXT = "text" 

110 FILE = "file" 

111 DATA = "data" 

112 

113 

114class MessageRole(StrEnum): 

115 """A2A 消息角色。""" 

116 

117 USER = "user" 

118 AGENT = "agent" 

119 

120 

121# ── Message Parts ────────────────────────────── 

122 

123 

124@dataclass 

125class TextPart: 

126 """文本消息片段。""" 

127 

128 text: str 

129 meta: dict[str, str] = field(default_factory=dict) 

130 

131 def to_dict(self) -> dict: 

132 return {"type": PartType.TEXT.value, "text": self.text, "meta": self.meta} 

133 

134 @classmethod 

135 def from_dict(cls, d: dict) -> TextPart: 

136 return cls(text=d.get("text", ""), meta=d.get("meta", {})) 

137 

138 

139@dataclass 

140class FilePart: 

141 """文件引用消息片段。""" 

142 

143 url: str = "" 

144 filename: str = "" 

145 mime_type: str = "application/octet-stream" 

146 size: int = 0 

147 meta: dict[str, str] = field(default_factory=dict) 

148 

149 def to_dict(self) -> dict: 

150 return { 

151 "type": PartType.FILE.value, 

152 "url": self.url, 

153 "filename": self.filename, 

154 "mime_type": self.mime_type, 

155 "size": self.size, 

156 "meta": self.meta, 

157 } 

158 

159 @classmethod 

160 def from_dict(cls, d: dict) -> FilePart: 

161 return cls( 

162 url=d.get("url", ""), 

163 filename=d.get("filename", ""), 

164 mime_type=d.get("mime_type", "application/octet-stream"), 

165 size=d.get("size", 0), 

166 meta=d.get("meta", {}), 

167 ) 

168 

169 

170@dataclass 

171class DataPart: 

172 """结构化数据消息片段。""" 

173 

174 data: dict[str, Any] = field(default_factory=dict) 

175 schema_uri: str = "" 

176 meta: dict[str, str] = field(default_factory=dict) 

177 

178 def to_dict(self) -> dict: 

179 return { 

180 "type": PartType.DATA.value, 

181 "data": self.data, 

182 "schema_uri": self.schema_uri, 

183 "meta": self.meta, 

184 } 

185 

186 @classmethod 

187 def from_dict(cls, d: dict) -> DataPart: 

188 return cls( 

189 data=d.get("data", {}), 

190 schema_uri=d.get("schema_uri", ""), 

191 meta=d.get("meta", {}), 

192 ) 

193 

194 

195def part_from_dict(d: dict): 

196 """从字典反序列化任意 Part。""" 

197 ptype = d.get("type", "") 

198 if ptype == PartType.TEXT.value: 

199 return TextPart.from_dict(d) 

200 elif ptype == PartType.FILE.value: 

201 return FilePart.from_dict(d) 

202 elif ptype == PartType.DATA.value: 

203 return DataPart.from_dict(d) 

204 raise ValueError(f"Unknown part type: {ptype}") 

205 

206 

207# ── A2A Artifact ─────────────────────────────── 

208 

209 

210@dataclass 

211class A2AArtifact: 

212 """任务产出物。 

213 可以是内联数据 (blob) 或外部引用 (url)。 

214 """ 

215 

216 name: str 

217 mime_type: str = "application/octet-stream" 

218 blob: bytes | None = None 

219 url: str = "" 

220 size: int = 0 

221 description: str = "" 

222 meta: dict[str, str] = field(default_factory=dict) 

223 

224 def to_dict(self) -> dict: 

225 d = { 

226 "name": self.name, 

227 "mime_type": self.mime_type, 

228 "size": self.size, 

229 "description": self.description, 

230 "meta": self.meta, 

231 } 

232 if self.url: 

233 d["url"] = self.url 

234 if self.blob: 

235 import base64 

236 

237 d["blob_base64"] = base64.b64encode(self.blob).decode("ascii") 

238 return d 

239 

240 @classmethod 

241 def from_dict(cls, d: dict) -> A2AArtifact: 

242 artifact = cls( 

243 name=d.get("name", ""), 

244 mime_type=d.get("mime_type", "application/octet-stream"), 

245 url=d.get("url", ""), 

246 size=d.get("size", 0), 

247 description=d.get("description", ""), 

248 meta=d.get("meta", {}), 

249 ) 

250 if "blob_base64" in d: 

251 import base64 

252 

253 artifact.blob = base64.b64decode(d["blob_base64"]) 

254 return artifact 

255 

256 

257# ── A2A Message ──────────────────────────────── 

258 

259 

260@dataclass 

261class A2AMessage: 

262 """多模态消息。""" 

263 

264 role: MessageRole = MessageRole.USER 

265 parts: list = field(default_factory=list) # List[TextPart|FilePart|DataPart] 

266 message_id: str = field(default_factory=lambda: f"msg-{uuid.uuid4().hex[:8]}") 

267 timestamp: float = field(default_factory=time.time) 

268 meta: dict[str, str] = field(default_factory=dict) 

269 

270 def to_dict(self) -> dict: 

271 return { 

272 "message_id": self.message_id, 

273 "role": self.role.value, 

274 "parts": [p.to_dict() for p in self.parts], 

275 "timestamp": self.timestamp, 

276 "meta": self.meta, 

277 } 

278 

279 @classmethod 

280 def from_dict(cls, d: dict) -> A2AMessage: 

281 role = MessageRole(d.get("role", "user")) 

282 parts = [part_from_dict(p) for p in d.get("parts", [])] 

283 return cls( 

284 message_id=d.get("message_id", f"msg-{uuid.uuid4().hex[:8]}"), 

285 role=role, 

286 parts=parts, 

287 timestamp=d.get("timestamp", time.time()), 

288 meta=d.get("meta", {}), 

289 ) 

290 

291 @classmethod 

292 def user_text(cls, text: str) -> A2AMessage: 

293 return cls(role=MessageRole.USER, parts=[TextPart(text=text)]) 

294 

295 @classmethod 

296 def agent_text(cls, text: str) -> A2AMessage: 

297 return cls(role=MessageRole.AGENT, parts=[TextPart(text=text)]) 

298 

299 def get_text(self) -> str: 

300 """提取所有 text parts 拼接。""" 

301 return " ".join(p.text for p in self.parts if isinstance(p, TextPart)) 

302 

303 

304# ── A2A Task ─────────────────────────────────── 

305 

306 

307@dataclass 

308class A2ATask: 

309 """A2A 异步任务。 

310 

311 状态机: SUBMITTED → WORKING → COMPLETED / FAILED / CANCELLED 

312 """ 

313 

314 task_id: str = field(default_factory=lambda: f"task-{uuid.uuid4().hex[:8]}") 

315 state: TaskState = TaskState.SUBMITTED 

316 input: A2AMessage | None = None 

317 output: A2AMessage | None = None 

318 artifacts: list[A2AArtifact] = field(default_factory=list) 

319 error: str | None = None 

320 meta: dict[str, Any] = field(default_factory=dict) 

321 _created: float = field(default_factory=time.time) 

322 _updated: float = field(default_factory=time.time) 

323 _state_history: list[tuple] = field(default_factory=list) # [(state, timestamp)] 

324 

325 def start_working(self) -> None: 

326 """SUBMITTED → WORKING""" 

327 if self.state != TaskState.SUBMITTED: 

328 raise ValueError(f"Cannot start from state {self.state}") 

329 self._transition(TaskState.WORKING) 

330 

331 def complete(self, output: A2AMessage | None = None) -> None: 

332 """WORKING → COMPLETED""" 

333 if self.state != TaskState.WORKING: 

334 raise ValueError(f"Cannot complete from state {self.state}") 

335 self.output = output 

336 self.error = None 

337 self._transition(TaskState.COMPLETED) 

338 

339 def fail(self, error: str) -> None: 

340 """任何状态 → FAILED""" 

341 self.error = error 

342 self._transition(TaskState.FAILED) 

343 

344 def cancel(self) -> None: 

345 """SUBMITTED/WORKING → CANCELLED""" 

346 if self.state not in (TaskState.SUBMITTED, TaskState.WORKING): 

347 raise ValueError(f"Cannot cancel from state {self.state}") 

348 self._transition(TaskState.CANCELLED) 

349 

350 def add_artifact(self, artifact: A2AArtifact) -> None: 

351 self.artifacts.append(artifact) 

352 

353 def is_terminal(self) -> bool: 

354 return self.state in (TaskState.COMPLETED, TaskState.FAILED, TaskState.CANCELLED) 

355 

356 def _transition(self, new_state: TaskState) -> None: 

357 self._state_history.append((self.state, self._updated)) 

358 self.state = new_state 

359 self._updated = time.time() 

360 

361 def to_dict(self) -> dict: 

362 return { 

363 "task_id": self.task_id, 

364 "state": self.state.value, 

365 "input": self.input.to_dict() if self.input else None, 

366 "output": self.output.to_dict() if self.output else None, 

367 "artifacts": [a.to_dict() for a in self.artifacts], 

368 "error": self.error, 

369 "meta": self.meta, 

370 "created": self._created, 

371 "updated": self._updated, 

372 } 

373 

374 @classmethod 

375 def from_dict(cls, d: dict) -> A2ATask: 

376 task = cls( 

377 task_id=d.get("task_id", f"task-{uuid.uuid4().hex[:8]}"), 

378 state=TaskState(d.get("state", "submitted")), 

379 error=d.get("error"), 

380 meta=d.get("meta", {}), 

381 _created=d.get("created", time.time()), 

382 _updated=d.get("updated", time.time()), 

383 ) 

384 if d.get("input"): 

385 task.input = A2AMessage.from_dict(d["input"]) 

386 if d.get("output"): 

387 task.output = A2AMessage.from_dict(d["output"]) 

388 task.artifacts = [A2AArtifact.from_dict(a) for a in d.get("artifacts", [])] 

389 return task 

390 

391 def to_json(self) -> str: 

392 return json.dumps(self.to_dict(), ensure_ascii=False, indent=2) 

393 

394 @classmethod 

395 def from_json(cls, json_str: str) -> A2ATask: 

396 return cls.from_dict(json.loads(json_str)) 

397 

398 

399# ── A2A Handoff ──────────────────────────────── 

400 

401 

402@dataclass 

403class A2AHandoff: 

404 """Agent 间任务移交请求。""" 

405 

406 handoff_id: str = field(default_factory=lambda: f"hoff-{uuid.uuid4().hex[:8]}") 

407 source_agent: str = "" 

408 target_agent: str = "" 

409 task: A2ATask | None = None 

410 reason: str = "" 

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

412 timestamp: float = field(default_factory=time.time) 

413 

414 def to_dict(self) -> dict: 

415 return { 

416 "handoff_id": self.handoff_id, 

417 "source_agent": self.source_agent, 

418 "target_agent": self.target_agent, 

419 "task": self.task.to_dict() if self.task else None, 

420 "reason": self.reason, 

421 "metadata": self.metadata, 

422 "timestamp": self.timestamp, 

423 } 

424 

425 @classmethod 

426 def from_dict(cls, d: dict) -> A2AHandoff: 

427 task = None 

428 if d.get("task"): 

429 task = A2ATask.from_dict(d["task"]) 

430 return cls( 

431 handoff_id=d.get("handoff_id", f"hoff-{uuid.uuid4().hex[:8]}"), 

432 source_agent=d.get("source_agent", ""), 

433 target_agent=d.get("target_agent", ""), 

434 task=task, 

435 reason=d.get("reason", ""), 

436 metadata=d.get("metadata", {}), 

437 timestamp=d.get("timestamp", time.time()), 

438 ) 

439 

440 def to_json(self) -> str: 

441 return json.dumps(self.to_dict(), ensure_ascii=False, indent=2) 

442 

443 @classmethod 

444 def from_json(cls, json_str: str) -> A2AHandoff: 

445 return cls.from_dict(json.loads(json_str)) 

446 

447 

448# ── A2A Session ──────────────────────────────── 

449 

450 

451@dataclass 

452class A2ASession: 

453 """A2A 会话上下文。""" 

454 

455 session_id: str = field(default_factory=lambda: f"sess-{uuid.uuid4().hex[:8]}") 

456 history: list[A2AMessage] = field(default_factory=list) 

457 tasks: list[A2ATask] = field(default_factory=list) 

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

459 created: float = field(default_factory=time.time) 

460 

461 def add_message(self, msg: A2AMessage) -> None: 

462 self.history.append(msg) 

463 

464 def add_task(self, task: A2ATask) -> None: 

465 self.tasks.append(task) 

466 

467 def get_last_n_messages(self, n: int = 10) -> list[A2AMessage]: 

468 return self.history[-n:] 

469 

470 def to_dict(self) -> dict: 

471 return { 

472 "session_id": self.session_id, 

473 "history": [m.to_dict() for m in self.history], 

474 "tasks": [t.to_dict() for t in self.tasks], 

475 "metadata": self.metadata, 

476 "created": self.created, 

477 } 

478 

479 

480# ── A2A Client ───────────────────────────────── 

481 

482 

483class A2AClient: 

484 """A2A 协议客户端。 

485 

486 向远程 Agent 发送任务,查询状态,获取结果。 

487 

488 v1.3.13: 重试 + 认证头 + 流式订阅 + 持久化连接池。 

489 """ 

490 

491 def __init__( 

492 self, 

493 base_url: str, 

494 timeout: float = 30.0, 

495 max_retries: int = 3, 

496 retry_backoff: float = 1.0, 

497 auth_token: str = "", 

498 agent_name: str = "", 

499 ): 

500 self.base_url = base_url.rstrip("/") 

501 self.timeout = timeout 

502 self.max_retries = max_retries 

503 self.retry_backoff = retry_backoff 

504 self.auth_token = auth_token 

505 self.agent_name = agent_name 

506 self._client: Any = None 

507 

508 def _headers(self) -> dict[str, str]: 

509 h = {"User-Agent": f"AgentOS-A2A/{self.agent_name}" if self.agent_name else "AgentOS-A2A"} 

510 if self.auth_token: 

511 h["Authorization"] = f"Bearer {self.auth_token}" 

512 return h 

513 

514 async def _get_client(self) -> Any: 

515 if self._client is None: 

516 import httpx 

517 

518 self._client = httpx.AsyncClient( 

519 timeout=self.timeout, 

520 headers=self._headers(), 

521 limits=httpx.Limits(max_keepalive_connections=10, max_connections=50), 

522 ) 

523 return self._client 

524 

525 async def close(self) -> None: 

526 if self._client: 

527 await self._client.aclose() 

528 self._client = None 

529 

530 async def _retry(self, coro, *args, **kwargs): 

531 import asyncio 

532 

533 import httpx 

534 

535 last_exc = None 

536 for attempt in range(self.max_retries): 

537 try: 

538 return await coro(*args, **kwargs) 

539 except (httpx.ConnectError, httpx.TimeoutException) as e: 

540 last_exc = e 

541 if attempt < self.max_retries - 1: 

542 await asyncio.sleep(self.retry_backoff * (2**attempt)) 

543 raise last_exc # type: ignore 

544 

545 async def send_task(self, task: A2ATask) -> A2ATask: 

546 """POST /tasks — 提交任务,返回带有 server 分配的 task_id 的任务。""" 

547 client = await self._get_client() 

548 

549 async def _do(): 

550 resp = await client.post(f"{self.base_url}/tasks", json=task.to_dict()) 

551 resp.raise_for_status() 

552 return A2ATask.from_dict(resp.json()) 

553 

554 return await self._retry(_do) 

555 

556 async def get_task(self, task_id: str) -> A2ATask | None: 

557 """GET /tasks/{id} — 查询任务状态和结果。""" 

558 client = await self._get_client() 

559 try: 

560 resp = await client.get(f"{self.base_url}/tasks/{task_id}") 

561 resp.raise_for_status() 

562 return A2ATask.from_dict(resp.json()) 

563 except Exception: 

564 return None 

565 

566 async def cancel_task(self, task_id: str) -> bool: 

567 """DELETE /tasks/{id} — 取消任务。""" 

568 client = await self._get_client() 

569 try: 

570 resp = await client.delete(f"{self.base_url}/tasks/{task_id}") 

571 return resp.status_code < 400 

572 except Exception: 

573 return False 

574 

575 async def handoff(self, handoff: A2AHandoff) -> bool: 

576 """POST /handoff — 移交任务到另一个 Agent。""" 

577 client = await self._get_client() 

578 try: 

579 resp = await client.post(f"{self.base_url}/handoff", json=handoff.to_dict()) 

580 return resp.status_code < 400 

581 except Exception: 

582 return False 

583 

584 async def wait_for_completion( 

585 self, 

586 task_id: str, 

587 poll_interval: float = 1.0, 

588 max_wait: float = 60.0, 

589 ) -> A2ATask: 

590 """轮询等待任务完成。""" 

591 import asyncio 

592 

593 elapsed = 0.0 

594 while elapsed < max_wait: 

595 task = await self.get_task(task_id) 

596 if task is None: 

597 raise RuntimeError(f"Task {task_id} not found") 

598 if task.is_terminal(): 

599 return task 

600 await asyncio.sleep(poll_interval) 

601 elapsed += poll_interval 

602 raise TimeoutError(f"Task {task_id} did not complete within {max_wait}s") 

603 

604 async def send_and_wait_for_reply( 

605 self, 

606 text: str, 

607 target_agent: str = "", 

608 poll_interval: float = 1.0, 

609 max_wait: float = 60.0, 

610 ) -> str: 

611 """便捷方法:发送文本任务并等待回复文本。""" 

612 task = new_task(text, target_agent=target_agent) 

613 task = await self.send_task(task) 

614 result = await self.wait_for_completion(task.task_id, poll_interval, max_wait) 

615 if result.output: 

616 return result.output.get_text() 

617 if result.error: 

618 return f"[Error] {result.error}" 

619 return "" 

620 

621 async def subscribe_task_stream( 

622 self, 

623 task_id: str, 

624 on_event: Callable[[dict], Any] | None = None, 

625 ) -> None: 

626 """SSE streaming subscribe: 连接到服务端 SSE 端点监听任务事件。""" 

627 client = await self._get_client() 

628 async with client.stream("GET", f"{self.base_url}/tasks/{task_id}/stream") as resp: 

629 resp.raise_for_status() 

630 buffer = "" 

631 async for chunk in resp.aiter_text(): 

632 buffer += chunk 

633 while "\n\n" in buffer: 

634 msg, buffer = buffer.split("\n\n", 1) 

635 event_data: dict[str, str] = {} 

636 for line in msg.split("\n"): 

637 if line.startswith("event: "): 

638 event_data["event"] = line[7:] 

639 elif line.startswith("data: "): 

640 event_data["data"] = line[6:] 

641 if on_event: 

642 on_event(event_data) 

643 

644 

645# ── A2A Server ───────────────────────────────── 

646 

647 

648class A2AServer: 

649 """A2A 协议服务端。 

650 

651 接收并处理 Agent 间任务请求。 

652 

653 使用方式: 

654 server = A2AServer() 

655 server.register_handler("my-agent", my_handler) 

656 # 集成到 FastAPI: 

657 app = FastAPI() 

658 server.mount_routes(app) 

659 """ 

660 

661 def __init__( 

662 self, 

663 task_store=None, 

664 stream_manager=None, 

665 require_auth: bool = False, 

666 auth_tokens: list[str] | None = None, 

667 ): 

668 self._handlers: dict[str, Callable] = {} 

669 self._task_store = task_store 

670 self._stream_manager = stream_manager 

671 self.require_auth = require_auth 

672 self.auth_tokens: set[str] = set(auth_tokens or []) 

673 self._default_store_created = False 

674 

675 def _ensure_store(self): 

676 if self._task_store is None: 

677 from agentos.protocols.a2a_store import InMemoryTaskStore 

678 

679 self._task_store = InMemoryTaskStore() 

680 self._default_store_created = True 

681 

682 @property 

683 def task_store(self): 

684 self._ensure_store() 

685 return self._task_store 

686 

687 def register_handler( 

688 self, 

689 agent_name: str, 

690 handler: Callable, 

691 ) -> None: 

692 """注册 Agent 处理函数。 

693 

694 handler 签名: async def handler(task: A2ATask) -> A2AMessage 

695 """ 

696 self._handlers[agent_name] = handler 

697 

698 async def process_task(self, body: dict, auth_token: str = "") -> dict: 

699 """处理传入任务:解析、执行 handler、返回。""" 

700 self._ensure_store() 

701 

702 if self.require_auth and auth_token not in self.auth_tokens: 

703 task = A2ATask.from_dict(body) 

704 task.fail("Unauthorized: invalid or missing A2A auth token") 

705 self._task_store.save_task(task) 

706 return task.to_dict() 

707 

708 task = A2ATask.from_dict(body) 

709 old_state = task.state 

710 self._task_store.save_task(task) 

711 

712 target = body.get("meta", {}).get("target_agent", "") 

713 handler = self._handlers.get(target) 

714 

715 if not handler and target: 

716 task.fail(f"No handler for agent '{target}'") 

717 elif not handler: 

718 task.fail("No target agent specified in meta") 

719 else: 

720 try: 

721 task.start_working() 

722 if self._stream_manager: 

723 await self._stream_manager.notify_state_change(task, old_state) 

724 result = handler(task) 

725 import inspect 

726 

727 if inspect.isawaitable(result): 

728 output = await result 

729 else: 

730 output = result 

731 task.complete(output) 

732 except Exception as e: 

733 task.fail(str(e)) 

734 

735 if self._stream_manager: 

736 await self._stream_manager.notify_state_change(task, old_state) 

737 self._task_store.save_task(task) 

738 return task.to_dict() 

739 

740 def get_task(self, task_id: str) -> A2ATask | None: 

741 return self.task_store.get_task(task_id) 

742 

743 def list_tasks(self, state: TaskState | None = None) -> list[A2ATask]: 

744 return self.task_store.list_tasks(state=state) 

745 

746 def cleanup_old(self, max_age_seconds: float = 3600.0) -> int: 

747 return self.task_store.cleanup_terminal(max_age_seconds) 

748 

749 # ── FastAPI 路由构建器 ───────────────────── 

750 

751 def mount_routes(self, app, prefix: str = "") -> None: 

752 """将 A2A 标准路由挂载到 FastAPI/Starlette app 上。 

753 

754 路由: 

755 POST {prefix}/tasks — 创建任务 

756 GET {prefix}/tasks — 列出任务 

757 GET {prefix}/tasks/{id} — 获取任务 

758 DELETE {prefix}/tasks/{id} — 取消任务 

759 GET {prefix}/tasks/{id}/stream — SSE 事件流 

760 POST {prefix}/handoff — 任务移交 

761 """ 

762 try: 

763 from fastapi import HTTPException, Request 

764 from starlette.responses import StreamingResponse 

765 except ImportError: 

766 raise ImportError("FastAPI and Starlette are required for mount_routes()") 

767 

768 server = self 

769 

770 @app.post(f"{prefix}/tasks") 

771 async def create_task(request: Request): 

772 body = await request.json() 

773 token = request.headers.get("Authorization", "").removeprefix("Bearer ") 

774 return await server.process_task(body, auth_token=token) 

775 

776 @app.get(f"{prefix}/tasks") 

777 async def list_tasks_endpoint(state: str = ""): 

778 task_state = TaskState(state) if state else None 

779 tasks = server.list_tasks(state=task_state) 

780 return [t.to_dict() for t in tasks] 

781 

782 @app.get(f"{prefix}/tasks/{{task_id}}") 

783 async def get_task_endpoint(task_id: str): 

784 task = server.get_task(task_id) 

785 if task is None: 

786 raise HTTPException(status_code=404, detail="Task not found") 

787 return task.to_dict() 

788 

789 @app.delete(f"{prefix}/tasks/{{task_id}}") 

790 async def cancel_task_endpoint(task_id: str): 

791 task = server.get_task(task_id) 

792 if task is None: 

793 raise HTTPException(status_code=404, detail="Task not found") 

794 if task.is_terminal(): 

795 raise HTTPException(status_code=400, detail="Task already terminal") 

796 task.cancel() 

797 server._task_store.save_task(task) 

798 return {"status": "cancelled", "task_id": task_id} 

799 

800 @app.get(f"{prefix}/tasks/{{task_id}}/stream") 

801 async def stream_task(task_id: str): 

802 if server._stream_manager is None: 

803 raise HTTPException(status_code=501, detail="Streaming not enabled") 

804 

805 async def event_generator(): 

806 stream = server._stream_manager 

807 session = stream.get_session(task_id) 

808 if session is None: 

809 yield 'event: error\ndata: {"error": "Session not found"}\n\n' 

810 return 

811 sub = session.subscribe() 

812 try: 

813 async for evt in session.iter_events(sub): 

814 yield session.to_sse(evt) 

815 except Exception: 

816 pass 

817 

818 return StreamingResponse( 

819 event_generator(), 

820 media_type="text/event-stream", 

821 headers={ 

822 "Cache-Control": "no-cache", 

823 "Connection": "keep-alive", 

824 "X-Accel-Buffering": "no", 

825 }, 

826 ) 

827 

828 @app.post(f"{prefix}/handoff") 

829 async def handoff_endpoint(request: Request): 

830 body = await request.json() 

831 token = request.headers.get("Authorization", "").removeprefix("Bearer ") 

832 if server.require_auth and token not in server.auth_tokens: 

833 raise HTTPException(status_code=401, detail="Unauthorized") 

834 handoff = A2AHandoff.from_dict(body) 

835 if handoff.task: 

836 server._task_store.save_task(handoff.task) 

837 return {"status": "received", "handoff_id": handoff.handoff_id} 

838 

839 

840# ── 便捷函数 ─────────────────────────────────── 

841 

842 

843def new_task(text: str, target_agent: str = "", **meta) -> A2ATask: 

844 """快速创建一个文本任务。""" 

845 return A2ATask( 

846 input=A2AMessage.user_text(text), 

847 meta={"target_agent": target_agent, **meta}, 

848 ) 

849 

850 

851def new_handoff( 

852 task: A2ATask, 

853 source: str, 

854 target: str, 

855 reason: str = "", 

856) -> A2AHandoff: 

857 """快速创建 Handoff。""" 

858 return A2AHandoff( 

859 source_agent=source, 

860 target_agent=target, 

861 task=task, 

862 reason=reason, 

863 ) 

864 

865 

866# ========================================================================== 

867# Compat: lightweight AgentRegistry for test_core.py 

868# ========================================================================== 

869import time # noqa: E402 

870from dataclasses import dataclass, field # noqa: E402 

871 

872 

873@dataclass 

874class AgentRecord: 

875 agent_id: str 

876 capabilities: list[str] = field(default_factory=list) 

877 endpoint: str = "" 

878 load: float = 0.0 

879 _last_heartbeat: float = field(default_factory=time.time, repr=False) 

880 

881 @property 

882 def healthy(self) -> bool: 

883 return (time.time() - self._last_heartbeat) < 60.0 

884 

885 

886class AgentRegistry: 

887 def __init__(self, name: str = "default", default_ttl: float = 60.0): 

888 self._records: dict[str, AgentRecord] = {} 

889 self.name = name 

890 self.default_ttl = default_ttl 

891 

892 def register(self, record: AgentRecord): 

893 record._last_heartbeat = time.time() 

894 self._records[record.agent_id] = record 

895 

896 def get(self, agent_id: str) -> AgentRecord | None: 

897 return self._records.get(agent_id) 

898 

899 def find_by_capability(self, capability: str) -> list[AgentRecord]: 

900 return [r for r in self._records.values() if capability in r.capabilities] 

901 

902 def heartbeat(self, agent_id: str): 

903 r = self._records.get(agent_id) 

904 if r: 

905 r._last_heartbeat = time.time() 

906 

907 def pick_least_loaded(self, capability: str) -> AgentRecord | None: 

908 candidates = self.find_by_capability(capability) 

909 if not candidates: 

910 return None 

911 return min(candidates, key=lambda r: r.load)