Coverage for agentos/protocols/a2a.py: 38%
480 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-08 12:20 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-08 12:20 +0800
1"""
2AgentOS v1.2.2 — A2A (Agent-to-Agent) 协议实现。
4基因来源: Google A2A Protocol (agent-to-agent-protocol.google.com)
6A2A 协议核心概念:
7- Task: 异步工作单元,带状态机 (SUBMITTED→WORKING→COMPLETED/FAILED/CANCELLED)
8- Message: 多模态消息,支持 text/file/data parts
9- Artifact: 任务产生的输出物,带 MIME 类型
10- Handoff: Agent 间任务移交
11- Session: 多轮对话上下文
13协议层:
14- REST: GET/POST /tasks, /tasks/{id}
15- Future: WebSocket 推送 (v1.3+)
16"""
18from __future__ import annotations
20import json
21import time
22import uuid
23from collections.abc import Callable
24from dataclasses import dataclass, field
25from enum import StrEnum
26from typing import Any
28# ── 基础枚举 ────────────────────────────────────
31class TaskState(StrEnum):
32 """A2A 任务状态。"""
34 SUBMITTED = "submitted"
35 WORKING = "working"
36 COMPLETED = "completed"
37 FAILED = "failed"
38 CANCELLED = "cancelled"
41class TaskStatus(StrEnum):
42 """A2A 任务状态(别名兼容)— 用于合规测试套件。"""
44 submitted = "submitted"
45 working = "working"
46 completed = "completed"
47 failed = "failed"
48 canceled = "canceled"
51class AgentCard:
52 """A2A Agent 名片 — 合规测试套件要求 Pydantic 兼容。"""
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 []
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 }
93class A2AMessageBus:
94 """A2A 消息总线 — 支持 agent 注册和消息发送。"""
96 def __init__(self):
97 self._agents: dict[str, Any] = {}
99 def register_agent(self, agent_id: str, agent: Any = None) -> None:
100 self._agents[agent_id] = agent
102 async def send(self, target_agent: str, message: Any) -> bool:
103 return target_agent in self._agents
106class PartType(StrEnum):
107 """A2A 内容片段类型。"""
109 TEXT = "text"
110 FILE = "file"
111 DATA = "data"
114class MessageRole(StrEnum):
115 """A2A 消息角色。"""
117 USER = "user"
118 AGENT = "agent"
121# ── Message Parts ──────────────────────────────
124@dataclass
125class TextPart:
126 """文本消息片段。"""
128 text: str
129 meta: dict[str, str] = field(default_factory=dict)
131 def to_dict(self) -> dict:
132 return {"type": PartType.TEXT.value, "text": self.text, "meta": self.meta}
134 @classmethod
135 def from_dict(cls, d: dict) -> TextPart:
136 return cls(text=d.get("text", ""), meta=d.get("meta", {}))
139@dataclass
140class FilePart:
141 """文件引用消息片段。"""
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)
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 }
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 )
170@dataclass
171class DataPart:
172 """结构化数据消息片段。"""
174 data: dict[str, Any] = field(default_factory=dict)
175 schema_uri: str = ""
176 meta: dict[str, str] = field(default_factory=dict)
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 }
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 )
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}")
207# ── A2A Artifact ───────────────────────────────
210@dataclass
211class A2AArtifact:
212 """任务产出物。
213 可以是内联数据 (blob) 或外部引用 (url)。
214 """
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)
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
237 d["blob_base64"] = base64.b64encode(self.blob).decode("ascii")
238 return d
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
253 artifact.blob = base64.b64decode(d["blob_base64"])
254 return artifact
257# ── A2A Message ────────────────────────────────
260@dataclass
261class A2AMessage:
262 """多模态消息。"""
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)
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 }
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 )
291 @classmethod
292 def user_text(cls, text: str) -> A2AMessage:
293 return cls(role=MessageRole.USER, parts=[TextPart(text=text)])
295 @classmethod
296 def agent_text(cls, text: str) -> A2AMessage:
297 return cls(role=MessageRole.AGENT, parts=[TextPart(text=text)])
299 def get_text(self) -> str:
300 """提取所有 text parts 拼接。"""
301 return " ".join(p.text for p in self.parts if isinstance(p, TextPart))
304# ── A2A Task ───────────────────────────────────
307@dataclass
308class A2ATask:
309 """A2A 异步任务。
311 状态机: SUBMITTED → WORKING → COMPLETED / FAILED / CANCELLED
312 """
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)]
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)
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)
339 def fail(self, error: str) -> None:
340 """任何状态 → FAILED"""
341 self.error = error
342 self._transition(TaskState.FAILED)
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)
350 def add_artifact(self, artifact: A2AArtifact) -> None:
351 self.artifacts.append(artifact)
353 def is_terminal(self) -> bool:
354 return self.state in (TaskState.COMPLETED, TaskState.FAILED, TaskState.CANCELLED)
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()
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 }
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
391 def to_json(self) -> str:
392 return json.dumps(self.to_dict(), ensure_ascii=False, indent=2)
394 @classmethod
395 def from_json(cls, json_str: str) -> A2ATask:
396 return cls.from_dict(json.loads(json_str))
399# ── A2A Handoff ────────────────────────────────
402@dataclass
403class A2AHandoff:
404 """Agent 间任务移交请求。"""
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)
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 }
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 )
440 def to_json(self) -> str:
441 return json.dumps(self.to_dict(), ensure_ascii=False, indent=2)
443 @classmethod
444 def from_json(cls, json_str: str) -> A2AHandoff:
445 return cls.from_dict(json.loads(json_str))
448# ── A2A Session ────────────────────────────────
451@dataclass
452class A2ASession:
453 """A2A 会话上下文。"""
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)
461 def add_message(self, msg: A2AMessage) -> None:
462 self.history.append(msg)
464 def add_task(self, task: A2ATask) -> None:
465 self.tasks.append(task)
467 def get_last_n_messages(self, n: int = 10) -> list[A2AMessage]:
468 return self.history[-n:]
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 }
480# ── A2A Client ─────────────────────────────────
483class A2AClient:
484 """A2A 协议客户端。
486 向远程 Agent 发送任务,查询状态,获取结果。
488 v1.3.13: 重试 + 认证头 + 流式订阅 + 持久化连接池。
489 """
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
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
514 async def _get_client(self) -> Any:
515 if self._client is None:
516 import httpx
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
525 async def close(self) -> None:
526 if self._client:
527 await self._client.aclose()
528 self._client = None
530 async def _retry(self, coro, *args, **kwargs):
531 import asyncio
533 import httpx
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
545 async def send_task(self, task: A2ATask) -> A2ATask:
546 """POST /tasks — 提交任务,返回带有 server 分配的 task_id 的任务。"""
547 client = await self._get_client()
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())
554 return await self._retry(_do)
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
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
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
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
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")
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 ""
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)
645# ── A2A Server ─────────────────────────────────
648class A2AServer:
649 """A2A 协议服务端。
651 接收并处理 Agent 间任务请求。
653 使用方式:
654 server = A2AServer()
655 server.register_handler("my-agent", my_handler)
656 # 集成到 FastAPI:
657 app = FastAPI()
658 server.mount_routes(app)
659 """
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
675 def _ensure_store(self):
676 if self._task_store is None:
677 from agentos.protocols.a2a_store import InMemoryTaskStore
679 self._task_store = InMemoryTaskStore()
680 self._default_store_created = True
682 @property
683 def task_store(self):
684 self._ensure_store()
685 return self._task_store
687 def register_handler(
688 self,
689 agent_name: str,
690 handler: Callable,
691 ) -> None:
692 """注册 Agent 处理函数。
694 handler 签名: async def handler(task: A2ATask) -> A2AMessage
695 """
696 self._handlers[agent_name] = handler
698 async def process_task(self, body: dict, auth_token: str = "") -> dict:
699 """处理传入任务:解析、执行 handler、返回。"""
700 self._ensure_store()
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()
708 task = A2ATask.from_dict(body)
709 old_state = task.state
710 self._task_store.save_task(task)
712 target = body.get("meta", {}).get("target_agent", "")
713 handler = self._handlers.get(target)
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
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))
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()
740 def get_task(self, task_id: str) -> A2ATask | None:
741 return self.task_store.get_task(task_id)
743 def list_tasks(self, state: TaskState | None = None) -> list[A2ATask]:
744 return self.task_store.list_tasks(state=state)
746 def cleanup_old(self, max_age_seconds: float = 3600.0) -> int:
747 return self.task_store.cleanup_terminal(max_age_seconds)
749 # ── FastAPI 路由构建器 ─────────────────────
751 def mount_routes(self, app, prefix: str = "") -> None:
752 """将 A2A 标准路由挂载到 FastAPI/Starlette app 上。
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()")
768 server = self
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)
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]
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()
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}
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")
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
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 )
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}
840# ── 便捷函数 ───────────────────────────────────
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 )
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 )
866# ==========================================================================
867# Compat: lightweight AgentRegistry for test_core.py
868# ==========================================================================
869import time # noqa: E402
870from dataclasses import dataclass, field # noqa: E402
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)
881 @property
882 def healthy(self) -> bool:
883 return (time.time() - self._last_heartbeat) < 60.0
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
892 def register(self, record: AgentRecord):
893 record._last_heartbeat = time.time()
894 self._records[record.agent_id] = record
896 def get(self, agent_id: str) -> AgentRecord | None:
897 return self._records.get(agent_id)
899 def find_by_capability(self, capability: str) -> list[AgentRecord]:
900 return [r for r in self._records.values() if capability in r.capabilities]
902 def heartbeat(self, agent_id: str):
903 r = self._records.get(agent_id)
904 if r:
905 r._last_heartbeat = time.time()
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)