Coverage for agentos/protocols/registry.py: 39%
359 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-06 21:19 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-06 21:19 +0800
1"""
2AgentOS v1.14.0 — A2A Agent Registry(服务发现)。
4Agent 注册中心实现:
5- 服务注册/注销/心跳
6- 能力广播 (Agent Card)
7- 服务发现 (按能力/角色查询)
8- 健康检查自动摘除
9"""
11from __future__ import annotations
13import asyncio
14import time
15from collections.abc import Callable
16from dataclasses import dataclass, field
17from enum import StrEnum
18from typing import Any
20# ── Agent Info ───────────────────────────────
23@dataclass
24class AgentInfo:
25 """Lightweight agent info for registration and discovery."""
27 agent_id: str
28 endpoint: str = ""
29 capabilities: list[str] = field(default_factory=list)
30 version: str = "1.0.0"
31 transport: str = "grpc"
32 status: str = "active"
33 heartbeat_ts: float = field(default_factory=time.time)
36# ── Agent Card ───────────────────────────────
39@dataclass
40class DiscoveryCapability:
41 """服务发现用 Agent 能力描述。
43 不同于 protocols.contracts.AgentCapability(面向合约),
44 本类面向注册中心发现和匹配。
45 """
47 name: str # 能力名称 (e.g. "code_review", "pdf_parsing")
48 description: str = ""
49 version: str = "1.0.0"
50 input_schema: dict[str, Any] | None = None
51 performance: dict[str, Any] = field(default_factory=dict) # 性能指标
53 def to_dict(self) -> dict:
54 return {
55 "name": self.name,
56 "description": self.description,
57 "version": self.version,
58 "input_schema": self.input_schema,
59 "performance": self.performance,
60 }
62 @classmethod
63 def from_dict(cls, d: dict) -> DiscoveryCapability:
64 return cls(
65 name=d.get("name", ""),
66 description=d.get("description", ""),
67 version=d.get("version", "1.0.0"),
68 input_schema=d.get("input_schema"),
69 performance=d.get("performance", {}),
70 )
73@dataclass
74class DiscoveryCard:
75 """Agent 发现名片 — 向注册中心声明的自身信息。
77 符合 Google A2A AgentCard 规范。
78 不同于 protocols.agent_card.AgentCard(面向本地广播),
79 本类面向远程注册中心的服务发现。
80 """
82 agent_id: str
83 name: str
84 description: str = ""
85 url: str = "" # Agent 的服务端点
86 version: str = "1.0.0"
87 capabilities: list[DiscoveryCapability] = field(default_factory=list)
88 provider: dict[str, Any] = field(default_factory=dict) # 组织/团队信息
89 default_input_modes: list[str] = field(default_factory=lambda: ["text"])
90 default_output_modes: list[str] = field(default_factory=lambda: ["text"])
91 skills: list[dict[str, Any]] = field(default_factory=list)
92 supports_streaming: bool = False
93 supports_handoff: bool = True
94 max_context_length: int = 128000
95 preferred_model: str = ""
96 service_tier: str = "standard" # standard | premium | enterprise
97 metadata: dict[str, Any] = field(default_factory=dict)
99 def to_dict(self) -> dict:
100 return {
101 "agent_id": self.agent_id,
102 "name": self.name,
103 "description": self.description,
104 "url": self.url,
105 "version": self.version,
106 "capabilities": [c.to_dict() for c in self.capabilities],
107 "provider": self.provider,
108 "defaultInputModes": self.default_input_modes,
109 "defaultOutputModes": self.default_output_modes,
110 "skills": self.skills,
111 "supportsStreaming": self.supports_streaming,
112 "supportsHandoff": self.supports_handoff,
113 "maxContextLength": self.max_context_length,
114 "preferredModel": self.preferred_model,
115 "serviceTier": self.service_tier,
116 "metadata": self.metadata,
117 }
119 @classmethod
120 def from_dict(cls, d: dict) -> DiscoveryCard:
121 return cls(
122 agent_id=d.get("agent_id", ""),
123 name=d.get("name", ""),
124 description=d.get("description", ""),
125 url=d.get("url", ""),
126 version=d.get("version", "1.0.0"),
127 capabilities=[DiscoveryCapability.from_dict(c) for c in d.get("capabilities", [])],
128 provider=d.get("provider", {}),
129 default_input_modes=d.get("defaultInputModes", ["text"]),
130 default_output_modes=d.get("defaultOutputModes", ["text"]),
131 skills=d.get("skills", []),
132 supports_streaming=d.get("supportsStreaming", False),
133 supports_handoff=d.get("supportsHandoff", True),
134 max_context_length=d.get("maxContextLength", 128000),
135 preferred_model=d.get("preferredModel", ""),
136 service_tier=d.get("serviceTier", "standard"),
137 metadata=d.get("metadata", {}),
138 )
140 def has_capability(self, cap_name: str) -> bool:
141 """检查是否具备某项能力。"""
142 return any(c.name == cap_name for c in self.capabilities)
145# ── Registry Entry ──────────────────────────
148class AgentStatus(StrEnum):
149 """Agent 运行状态。"""
151 ONLINE = "online"
152 BUSY = "busy"
153 DEGRADED = "degraded"
154 OFFLINE = "offline"
157@dataclass
158class RegistryEntry:
159 """注册中心中的单条 Agent 记录。"""
161 card: DiscoveryCard
162 status: AgentStatus = AgentStatus.ONLINE
163 registered_at: float = field(default_factory=time.time)
164 last_heartbeat: float = field(default_factory=time.time)
165 load: float = 0.0 # 0.0~1.0 负载
166 task_count: int = 0 # 已完成任务数
167 error_count: int = 0
168 avg_response_ms: float = 0.0
169 tags: list[str] = field(default_factory=list)
170 endpoint_health: dict[str, Any] = field(default_factory=dict)
172 def is_healthy(self, heartbeat_timeout: float = 30.0) -> bool:
173 """心跳是否正常。"""
174 return (time.time() - self.last_heartbeat) < heartbeat_timeout
176 @property
177 def uptime_seconds(self) -> float:
178 """注册后的运行时长。"""
179 return time.time() - self.registered_at
181 @property
182 def endpoint(self) -> str:
183 """Agent 服务端点(映射 card.url)。"""
184 return self.card.url
187# ── Agent Registry ──────────────────────────
190class AgentRegistry:
191 """A2A Agent 注册中心。
193 功能:
194 - register: 注册 Agent(带心跳保活)
195 - discover: 按能力/角色/标签发现 Agent
196 - health_check: 自动摘除失联 Agent
197 - subscribe: 订阅注册事件
199 Usage:
200 registry = AgentRegistry(heartbeat_timeout=30)
201 registry.register(agent_card)
203 # 发现能处理 code_review 的 Agent
204 agents = registry.discover(capability="code_review")
206 # 按标签查找
207 agents = registry.discover(tags=["production", "high-priority"])
208 """
210 def __init__(
211 self,
212 heartbeat_timeout: float = 30.0,
213 health_check_interval: float = 10.0,
214 auto_cleanup: bool = True,
215 ):
216 self._entries: dict[str, RegistryEntry] = {}
217 self._capability_index: dict[str, set[str]] = {}
218 self._tag_index: dict[str, set[str]] = {}
219 self._role_index: dict[str, set[str]] = {}
220 self.heartbeat_timeout = heartbeat_timeout
221 self.health_check_interval = health_check_interval
222 self.auto_cleanup = auto_cleanup
224 # Event subscribers
225 self._subscribers: dict[str, list[Callable]] = {
226 "register": [],
227 "deregister": [],
228 "status_change": [],
229 "heartbeat": [],
230 }
232 # Health check task
233 self._health_check_task: asyncio.Task | None = None
234 self._running = False
236 # ── Lifecycle ─────────────────────────────
238 async def start(self) -> None:
239 """启动注册中心(开始健康检查循环)。"""
240 if self._running:
241 return
242 self._running = True
243 if self.auto_cleanup:
244 self._health_check_task = asyncio.create_task(self._health_check_loop())
246 async def stop(self) -> None:
247 """停止注册中心。"""
248 self._running = False
249 if self._health_check_task:
250 self._health_check_task.cancel()
251 try:
252 await self._health_check_task
253 except asyncio.CancelledError:
254 pass
255 self._health_check_task = None
257 async def _health_check_loop(self) -> None:
258 """后台健康检查循环。"""
259 while self._running:
260 try:
261 self._check_all_health()
262 except Exception:
263 pass
264 await asyncio.sleep(self.health_check_interval)
266 def _check_all_health(self) -> None:
267 """健康检查:摘除失联 Agent。"""
268 to_remove = []
269 for agent_id, entry in list(self._entries.items()):
270 if entry.status != AgentStatus.OFFLINE and not entry.is_healthy(self.heartbeat_timeout):
271 self._update_status(agent_id, AgentStatus.OFFLINE)
272 to_remove.append(agent_id)
274 for agent_id in to_remove:
275 self._cleanup_indices(agent_id)
277 # ── Registration ─────────────────────────
279 def register(
280 self,
281 card_or_info: DiscoveryCard | AgentInfo,
282 tags: list[str] | None = None,
283 ) -> str:
284 """注册一个 Agent。
286 Args:
287 card_or_info: Agent 名片 (DiscoveryCard) 或轻量信息 (AgentInfo)
288 tags: 自定义标签
290 Returns:
291 agent_id
292 """
293 # Accept both DiscoveryCard and AgentInfo
294 if isinstance(card_or_info, AgentInfo):
295 info = card_or_info
296 card = DiscoveryCard(
297 agent_id=info.agent_id,
298 name=info.agent_id,
299 description=f"Agent {info.agent_id} v{info.version}",
300 version=info.version,
301 capabilities=[DiscoveryCapability(name=c) for c in info.capabilities],
302 url=info.endpoint,
303 )
304 else:
305 card = card_or_info
307 agent_id = card.agent_id
308 now = time.time()
310 if agent_id in self._entries:
311 # Re-registration: update card, refresh heartbeat
312 entry = self._entries[agent_id]
313 old_status = entry.status
314 entry.card = card
315 entry.last_heartbeat = now
316 entry.tags = tags or entry.tags
317 if old_status == AgentStatus.OFFLINE:
318 self._update_status(agent_id, AgentStatus.ONLINE)
319 return agent_id
321 entry = RegistryEntry(
322 card=card,
323 status=AgentStatus.ONLINE,
324 registered_at=now,
325 last_heartbeat=now,
326 tags=tags or [],
327 )
328 self._entries[agent_id] = entry
329 self._build_indices(agent_id, card, tags or [])
330 self._emit("register", {"agent_id": agent_id, "card": card.to_dict()})
331 return agent_id
333 def deregister(self, agent_id: str) -> bool:
334 """注销一个 Agent。"""
335 if agent_id not in self._entries:
336 return False
337 entry = self._entries.pop(agent_id)
338 self._cleanup_indices(agent_id)
339 self._emit("deregister", {"agent_id": agent_id, "card": entry.card.to_dict()})
340 return True
342 def heartbeat(self, agent_id: str, load: float | None = None) -> bool:
343 """Agent 心跳上报。
345 Args:
346 agent_id: Agent ID
347 load: 当前负载 0.0~1.0
349 Returns:
350 是否成功
351 """
352 if agent_id not in self._entries:
353 return False
354 entry = self._entries[agent_id]
355 entry.last_heartbeat = time.time()
356 if load is not None:
357 entry.load = max(0.0, min(1.0, load))
358 if entry.status == AgentStatus.OFFLINE:
359 self._update_status(agent_id, AgentStatus.ONLINE)
360 self._emit("heartbeat", {"agent_id": agent_id, "load": entry.load})
361 return True
363 def update_stats(
364 self,
365 agent_id: str,
366 task_count: int | None = None,
367 error_count: int | None = None,
368 avg_response_ms: float | None = None,
369 ) -> bool:
370 """更新 Agent 统计信息。"""
371 if agent_id not in self._entries:
372 return False
373 entry = self._entries[agent_id]
374 if task_count is not None:
375 entry.task_count = task_count
376 if error_count is not None:
377 entry.error_count = error_count
378 if avg_response_ms is not None:
379 entry.avg_response_ms = avg_response_ms
380 return True
382 # ── Discovery ────────────────────────────
384 def discover(
385 self,
386 capability: str | None = None,
387 tags: list[str] | None = None,
388 role: str | None = None,
389 status: AgentStatus | None = None,
390 service_tier: str | None = None,
391 supports_streaming: bool | None = None,
392 min_health: bool = True,
393 limit: int = 50,
394 ) -> list[RegistryEntry]:
395 """服务发现:按条件筛选 Agent。
397 Args:
398 capability: 按能力名称筛选
399 tags: 按标签筛选(AND 逻辑)
400 role: 按角色筛选
401 status: 按状态筛选
402 service_tier: 按服务等级筛选
403 supports_streaming: 是否支持流式
404 min_health: 仅返回心跳正常的 Agent
405 limit: 最大返回数
407 Returns:
408 匹配的 RegistryEntry 列表
409 """
410 candidates: set[str] = set(self._entries.keys())
412 # Capability filter
413 if capability:
414 cap_agents = self._capability_index.get(capability, set())
415 candidates &= cap_agents
417 # Tag filter (AND)
418 if tags:
419 for tag in tags:
420 tag_agents = self._tag_index.get(tag, set())
421 candidates &= tag_agents
423 # Role filter
424 if role:
425 role_agents = self._role_index.get(role, set())
426 candidates &= role_agents
428 # Collect results
429 results = []
430 for agent_id in candidates:
431 entry = self._entries[agent_id]
433 # Status filter
434 if status and entry.status != status:
435 continue
437 # Health filter
438 if min_health and not entry.is_healthy(self.heartbeat_timeout):
439 continue
441 # Service tier filter
442 if service_tier and entry.card.service_tier != service_tier:
443 continue
445 # Streaming filter
446 if (
447 supports_streaming is not None
448 and entry.card.supports_streaming != supports_streaming
449 ):
450 continue
452 results.append(entry)
454 # Sort: online first, then by load (least loaded first)
455 status_order = {
456 AgentStatus.ONLINE: 0,
457 AgentStatus.BUSY: 1,
458 AgentStatus.DEGRADED: 2,
459 AgentStatus.OFFLINE: 3,
460 }
461 results.sort(key=lambda e: (status_order.get(e.status, 9), e.load))
463 return results[:limit]
465 def discover_one(
466 self,
467 capability: str | None = None,
468 load_balanced: bool = True,
469 **kwargs,
470 ) -> RegistryEntry | None:
471 """发现单个最优 Agent。
473 Args:
474 capability: 按能力筛选
475 load_balanced: 是否负载均衡(选负载最低的)
476 """
477 results = self.discover(
478 capability=capability,
479 status=AgentStatus.ONLINE,
480 min_health=True,
481 limit=10,
482 **kwargs,
483 )
484 if not results:
485 return None
486 if load_balanced:
487 # Already sorted by load ascending
488 return results[0]
489 return results[0]
491 def get_agent(self, agent_id: str) -> RegistryEntry | None:
492 """按 ID 获取 Agent。"""
493 return self._entries.get(agent_id)
495 def list_all(self, include_offline: bool = False) -> list[RegistryEntry]:
496 """列出所有 Agent。"""
497 entries = list(self._entries.values())
498 if not include_offline:
499 entries = [
500 e
501 for e in entries
502 if e.status != AgentStatus.OFFLINE or e.is_healthy(self.heartbeat_timeout)
503 ]
504 return entries
506 # ── Stats ─────────────────────────────────
508 @property
509 def total_agents(self) -> int:
510 """已注册的 Agent 总数。"""
511 return len(self._entries)
513 @property
514 def online_count(self) -> int:
515 """在线 Agent 数。"""
516 return sum(
517 1
518 for e in self._entries.values()
519 if e.status == AgentStatus.ONLINE and e.is_healthy(self.heartbeat_timeout)
520 )
522 def get_stats(self) -> dict[str, Any]:
523 """获取注册中心统计。"""
524 online = 0
525 busy = 0
526 degraded = 0
527 offline = 0
528 for e in self._entries.values():
529 if not e.is_healthy(self.heartbeat_timeout):
530 offline += 1
531 elif e.status == AgentStatus.ONLINE:
532 online += 1
533 elif e.status == AgentStatus.BUSY:
534 busy += 1
535 elif e.status == AgentStatus.DEGRADED:
536 degraded += 1
538 return {
539 "total": self.total_agents,
540 "online": online,
541 "busy": busy,
542 "degraded": degraded,
543 "offline": offline,
544 "capabilities": list(self._capability_index.keys()),
545 "tags": list(self._tag_index.keys()),
546 "roles": list(self._role_index.keys()),
547 }
549 # ── Events ────────────────────────────────
551 def subscribe(
552 self,
553 event: str,
554 callback: Callable[[dict[str, Any]], Any],
555 ) -> None:
556 """订阅注册事件。
558 Args:
559 event: 'register' | 'deregister' | 'status_change' | 'heartbeat'
560 callback: 回调函数,接收事件 dict
561 """
562 if event in self._subscribers:
563 self._subscribers[event].append(callback)
565 def _emit(self, event: str, data: dict[str, Any]) -> None:
566 """触发事件。"""
567 for cb in self._subscribers.get(event, []):
568 try:
569 cb(data)
570 except Exception:
571 pass
573 # ── Internal Helpers ─────────────────────
575 def _update_status(self, agent_id: str, new_status: AgentStatus) -> None:
576 """更新 Agent 状态并触发事件。"""
577 if agent_id in self._entries:
578 old_status = self._entries[agent_id].status
579 self._entries[agent_id].status = new_status
580 if old_status != new_status:
581 self._emit(
582 "status_change",
583 {
584 "agent_id": agent_id,
585 "old_status": old_status.value,
586 "new_status": new_status.value,
587 },
588 )
590 def _build_indices(
591 self,
592 agent_id: str,
593 card: DiscoveryCard,
594 tags: list[str],
595 ) -> None:
596 """构建反向索引。"""
597 # Capability index
598 for cap in card.capabilities:
599 if cap.name not in self._capability_index:
600 self._capability_index[cap.name] = set()
601 self._capability_index[cap.name].add(agent_id)
603 # Tag index
604 for tag in tags:
605 if tag not in self._tag_index:
606 self._tag_index[tag] = set()
607 self._tag_index[tag].add(agent_id)
609 # Role index (from provider)
610 role = card.provider.get("role", "")
611 if role:
612 if role not in self._role_index:
613 self._role_index[role] = set()
614 self._role_index[role].add(agent_id)
616 def _cleanup_indices(self, agent_id: str) -> None:
617 """从所有索引中移除 Agent。"""
618 for cap_set in self._capability_index.values():
619 cap_set.discard(agent_id)
620 for tag_set in self._tag_index.values():
621 tag_set.discard(agent_id)
622 for role_set in self._role_index.values():
623 role_set.discard(agent_id)
626# ── A2A Registry Bridge ─────────────────────
629class A2ARegistryBridge:
630 """A2A Registry + A2A Client 桥接。
632 自动从 Registry 发现 Agent 并创建 A2A Client。
634 Usage:
635 bridge = A2ARegistryBridge(registry)
636 client = await bridge.get_client(capability="code_review")
637 result = await client.send_and_wait_for_reply("Review this code...")
638 """
640 def __init__(self, registry: AgentRegistry):
641 self._registry = registry
642 self._clients: dict[str, Any] = {} # agent_id -> A2AClient
644 async def get_client(
645 self,
646 capability: str | None = None,
647 agent_id: str | None = None,
648 **kwargs,
649 ) -> Any:
650 """获取已发现的 Agent 的 A2A Client。
652 Args:
653 capability: 按能力发现
654 agent_id: 直接指定 Agent ID
655 """
656 from agentos.protocols.a2a import A2AClient
658 if agent_id and agent_id in self._clients:
659 return self._clients[agent_id]
661 entry = None
662 if agent_id:
663 entry = self._registry.get_agent(agent_id)
664 else:
665 entry = self._registry.discover_one(capability=capability, **kwargs)
667 if not entry:
668 raise RuntimeError(f"No available agent found for capability={capability}")
670 client = A2AClient(
671 base_url=entry.card.url,
672 agent_name=entry.card.name,
673 )
674 self._clients[entry.card.agent_id] = client
675 return client
677 async def close_all(self) -> None:
678 """关闭所有 Client 连接。"""
679 for client in self._clients.values():
680 try:
681 await client.close()
682 except Exception:
683 pass
684 self._clients.clear()
686 def invalidate(self, agent_id: str) -> None:
687 """使某个 Agent 的缓存 Client 失效。"""
688 self._clients.pop(agent_id, None)
691# ── 全局单例 ─────────────────────────────────
693default_registry = AgentRegistry()
695# ── AgentRecord (test compatibility) ──
696import time # noqa: E402
697from dataclasses import dataclass, field # noqa: E402
700@dataclass
701class AgentRecord:
702 agent_id: str
703 capabilities: list = field(default_factory=list)
704 endpoint: str = ""
705 load: float = 0.0
706 healthy: bool = True
707 name: str = ""
708 version: str = "1.0"
709 registered_at: float = field(default_factory=time.time)
712# ==========================================================================
713# Compat: AgentRecord + lightweight compat API for test_core.py
714# ==========================================================================
717# Re-use the AgentRecord from a2a module to avoid duplication
718try:
719 from agentos.protocols.a2a import AgentRecord
720except ImportError:
721 import time as _time
722 from dataclasses import dataclass, field
724 @dataclass
725 class AgentRecord:
726 agent_id: str
727 capabilities: list = field(default_factory=list)
728 endpoint: str = ""
729 load: float = 0.0
730 _last_heartbeat: float = field(default_factory=_time.time, repr=False)
732 @property
733 def healthy(self) -> bool:
734 return (_time.time() - self._last_heartbeat) < 60.0
737def _patch_agent_registry():
738 """Extend AgentRegistry with compat methods."""
739 import time as _time
741 def compat_register(self, record):
742 from agentos.protocols.a2a import AgentRecord
744 if not hasattr(self, "_compat_records"):
745 self._compat_records = {}
746 if isinstance(record, AgentRecord):
747 record._last_heartbeat = _time.time()
748 self._compat_records[record.agent_id] = record
749 else:
750 self._orig_register(record)
752 def compat_get(self, agent_id):
753 if not hasattr(self, "_compat_records"):
754 self._compat_records = {}
755 return self._compat_records.get(agent_id)
757 def compat_find_by_capability(self, capability):
758 if not hasattr(self, "_compat_records"):
759 self._compat_records = {}
760 return [r for r in self._compat_records.values() if capability in r.capabilities]
762 def compat_heartbeat(self, agent_id):
763 if not hasattr(self, "_compat_records"):
764 self._compat_records = {}
765 r = self._compat_records.get(agent_id)
766 if r:
767 r._last_heartbeat = _time.time()
769 def compat_pick_least_loaded(self, capability):
770 candidates = compat_find_by_capability(self, capability)
771 if not candidates:
772 return None
773 return min(candidates, key=lambda r: r.load)
775 if not hasattr(AgentRegistry, "_orig_register"):
776 AgentRegistry._orig_register = AgentRegistry.register
777 AgentRegistry.register = compat_register
778 AgentRegistry.get = compat_get
779 AgentRegistry.find_by_capability = compat_find_by_capability
780 AgentRegistry.heartbeat = compat_heartbeat
781 AgentRegistry.pick_least_loaded = compat_pick_least_loaded
784_patch_agent_registry()