Coverage for agentos/protocols/registry.py: 39%
358 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-06 10:59 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-06 10:59 +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 dataclasses import dataclass, field
16from enum import Enum
17from typing import Any, Callable, Dict, List, Optional, Set, Union
20# ── Agent Info ───────────────────────────────
23@dataclass
24class AgentInfo:
25 """Lightweight agent info for registration and discovery."""
26 agent_id: str
27 endpoint: str = ""
28 capabilities: List[str] = field(default_factory=list)
29 version: str = "1.0.0"
30 transport: str = "grpc"
31 status: str = "active"
32 heartbeat_ts: float = field(default_factory=time.time)
35# ── Agent Card ───────────────────────────────
38@dataclass
39class DiscoveryCapability:
40 """服务发现用 Agent 能力描述。
42 不同于 protocols.contracts.AgentCapability(面向合约),
43 本类面向注册中心发现和匹配。
44 """
45 name: str # 能力名称 (e.g. "code_review", "pdf_parsing")
46 description: str = ""
47 version: str = "1.0.0"
48 input_schema: Optional[Dict[str, Any]] = None
49 performance: Dict[str, Any] = field(default_factory=dict) # 性能指标
51 def to_dict(self) -> dict:
52 return {
53 "name": self.name,
54 "description": self.description,
55 "version": self.version,
56 "input_schema": self.input_schema,
57 "performance": self.performance,
58 }
60 @classmethod
61 def from_dict(cls, d: dict) -> "DiscoveryCapability":
62 return cls(
63 name=d.get("name", ""),
64 description=d.get("description", ""),
65 version=d.get("version", "1.0.0"),
66 input_schema=d.get("input_schema"),
67 performance=d.get("performance", {}),
68 )
71@dataclass
72class DiscoveryCard:
73 """Agent 发现名片 — 向注册中心声明的自身信息。
75 符合 Google A2A AgentCard 规范。
76 不同于 protocols.agent_card.AgentCard(面向本地广播),
77 本类面向远程注册中心的服务发现。
78 """
80 agent_id: str
81 name: str
82 description: str = ""
83 url: str = "" # Agent 的服务端点
84 version: str = "1.0.0"
85 capabilities: List[DiscoveryCapability] = field(default_factory=list)
86 provider: Dict[str, Any] = field(default_factory=dict) # 组织/团队信息
87 default_input_modes: List[str] = field(default_factory=lambda: ["text"])
88 default_output_modes: List[str] = field(default_factory=lambda: ["text"])
89 skills: List[Dict[str, Any]] = field(default_factory=list)
90 supports_streaming: bool = False
91 supports_handoff: bool = True
92 max_context_length: int = 128000
93 preferred_model: str = ""
94 service_tier: str = "standard" # standard | premium | enterprise
95 metadata: Dict[str, Any] = field(default_factory=dict)
97 def to_dict(self) -> dict:
98 return {
99 "agent_id": self.agent_id,
100 "name": self.name,
101 "description": self.description,
102 "url": self.url,
103 "version": self.version,
104 "capabilities": [c.to_dict() for c in self.capabilities],
105 "provider": self.provider,
106 "defaultInputModes": self.default_input_modes,
107 "defaultOutputModes": self.default_output_modes,
108 "skills": self.skills,
109 "supportsStreaming": self.supports_streaming,
110 "supportsHandoff": self.supports_handoff,
111 "maxContextLength": self.max_context_length,
112 "preferredModel": self.preferred_model,
113 "serviceTier": self.service_tier,
114 "metadata": self.metadata,
115 }
117 @classmethod
118 def from_dict(cls, d: dict) -> "DiscoveryCard":
119 return cls(
120 agent_id=d.get("agent_id", ""),
121 name=d.get("name", ""),
122 description=d.get("description", ""),
123 url=d.get("url", ""),
124 version=d.get("version", "1.0.0"),
125 capabilities=[
126 DiscoveryCapability.from_dict(c)
127 for c in d.get("capabilities", [])
128 ],
129 provider=d.get("provider", {}),
130 default_input_modes=d.get("defaultInputModes", ["text"]),
131 default_output_modes=d.get("defaultOutputModes", ["text"]),
132 skills=d.get("skills", []),
133 supports_streaming=d.get("supportsStreaming", False),
134 supports_handoff=d.get("supportsHandoff", True),
135 max_context_length=d.get("maxContextLength", 128000),
136 preferred_model=d.get("preferredModel", ""),
137 service_tier=d.get("serviceTier", "standard"),
138 metadata=d.get("metadata", {}),
139 )
141 def has_capability(self, cap_name: str) -> bool:
142 """检查是否具备某项能力。"""
143 return any(c.name == cap_name for c in self.capabilities)
146# ── Registry Entry ──────────────────────────
149class AgentStatus(str, Enum):
150 """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: Optional[asyncio.Task] = 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(
271 self.heartbeat_timeout
272 ):
273 self._update_status(agent_id, AgentStatus.OFFLINE)
274 to_remove.append(agent_id)
276 for agent_id in to_remove:
277 self._cleanup_indices(agent_id)
279 # ── Registration ─────────────────────────
281 def register(
282 self,
283 card_or_info: Union[DiscoveryCard, "AgentInfo"],
284 tags: Optional[List[str]] = None,
285 ) -> str:
286 """注册一个 Agent。
288 Args:
289 card_or_info: Agent 名片 (DiscoveryCard) 或轻量信息 (AgentInfo)
290 tags: 自定义标签
292 Returns:
293 agent_id
294 """
295 # Accept both DiscoveryCard and AgentInfo
296 if isinstance(card_or_info, AgentInfo):
297 info = card_or_info
298 card = DiscoveryCard(
299 agent_id=info.agent_id,
300 name=info.agent_id,
301 description=f"Agent {info.agent_id} v{info.version}",
302 version=info.version,
303 capabilities=[
304 DiscoveryCapability(name=c) for c in info.capabilities
305 ],
306 url=info.endpoint,
307 )
308 else:
309 card = card_or_info
311 agent_id = card.agent_id
312 now = time.time()
314 if agent_id in self._entries:
315 # Re-registration: update card, refresh heartbeat
316 entry = self._entries[agent_id]
317 old_status = entry.status
318 entry.card = card
319 entry.last_heartbeat = now
320 entry.tags = tags or entry.tags
321 if old_status == AgentStatus.OFFLINE:
322 self._update_status(agent_id, AgentStatus.ONLINE)
323 return agent_id
325 entry = RegistryEntry(
326 card=card,
327 status=AgentStatus.ONLINE,
328 registered_at=now,
329 last_heartbeat=now,
330 tags=tags or [],
331 )
332 self._entries[agent_id] = entry
333 self._build_indices(agent_id, card, tags or [])
334 self._emit("register", {"agent_id": agent_id, "card": card.to_dict()})
335 return agent_id
337 def deregister(self, agent_id: str) -> bool:
338 """注销一个 Agent。"""
339 if agent_id not in self._entries:
340 return False
341 entry = self._entries.pop(agent_id)
342 self._cleanup_indices(agent_id)
343 self._emit("deregister", {"agent_id": agent_id, "card": entry.card.to_dict()})
344 return True
346 def heartbeat(self, agent_id: str, load: Optional[float] = None) -> bool:
347 """Agent 心跳上报。
349 Args:
350 agent_id: Agent ID
351 load: 当前负载 0.0~1.0
353 Returns:
354 是否成功
355 """
356 if agent_id not in self._entries:
357 return False
358 entry = self._entries[agent_id]
359 entry.last_heartbeat = time.time()
360 if load is not None:
361 entry.load = max(0.0, min(1.0, load))
362 if entry.status == AgentStatus.OFFLINE:
363 self._update_status(agent_id, AgentStatus.ONLINE)
364 self._emit("heartbeat", {"agent_id": agent_id, "load": entry.load})
365 return True
367 def update_stats(
368 self,
369 agent_id: str,
370 task_count: Optional[int] = None,
371 error_count: Optional[int] = None,
372 avg_response_ms: Optional[float] = None,
373 ) -> bool:
374 """更新 Agent 统计信息。"""
375 if agent_id not in self._entries:
376 return False
377 entry = self._entries[agent_id]
378 if task_count is not None:
379 entry.task_count = task_count
380 if error_count is not None:
381 entry.error_count = error_count
382 if avg_response_ms is not None:
383 entry.avg_response_ms = avg_response_ms
384 return True
386 # ── Discovery ────────────────────────────
388 def discover(
389 self,
390 capability: Optional[str] = None,
391 tags: Optional[List[str]] = None,
392 role: Optional[str] = None,
393 status: Optional[AgentStatus] = None,
394 service_tier: Optional[str] = None,
395 supports_streaming: Optional[bool] = None,
396 min_health: bool = True,
397 limit: int = 50,
398 ) -> List[RegistryEntry]:
399 """服务发现:按条件筛选 Agent。
401 Args:
402 capability: 按能力名称筛选
403 tags: 按标签筛选(AND 逻辑)
404 role: 按角色筛选
405 status: 按状态筛选
406 service_tier: 按服务等级筛选
407 supports_streaming: 是否支持流式
408 min_health: 仅返回心跳正常的 Agent
409 limit: 最大返回数
411 Returns:
412 匹配的 RegistryEntry 列表
413 """
414 candidates: Set[str] = set(self._entries.keys())
416 # Capability filter
417 if capability:
418 cap_agents = self._capability_index.get(capability, set())
419 candidates &= cap_agents
421 # Tag filter (AND)
422 if tags:
423 for tag in tags:
424 tag_agents = self._tag_index.get(tag, set())
425 candidates &= tag_agents
427 # Role filter
428 if role:
429 role_agents = self._role_index.get(role, set())
430 candidates &= role_agents
432 # Collect results
433 results = []
434 for agent_id in candidates:
435 entry = self._entries[agent_id]
437 # Status filter
438 if status and entry.status != status:
439 continue
441 # Health filter
442 if min_health and not entry.is_healthy(self.heartbeat_timeout):
443 continue
445 # Service tier filter
446 if service_tier and entry.card.service_tier != service_tier:
447 continue
449 # Streaming filter
450 if supports_streaming is not None and \
451 entry.card.supports_streaming != supports_streaming:
452 continue
454 results.append(entry)
456 # Sort: online first, then by load (least loaded first)
457 status_order = {
458 AgentStatus.ONLINE: 0,
459 AgentStatus.BUSY: 1,
460 AgentStatus.DEGRADED: 2,
461 AgentStatus.OFFLINE: 3,
462 }
463 results.sort(key=lambda e: (status_order.get(e.status, 9), e.load))
465 return results[:limit]
467 def discover_one(
468 self,
469 capability: Optional[str] = None,
470 load_balanced: bool = True,
471 **kwargs,
472 ) -> Optional[RegistryEntry]:
473 """发现单个最优 Agent。
475 Args:
476 capability: 按能力筛选
477 load_balanced: 是否负载均衡(选负载最低的)
478 """
479 results = self.discover(
480 capability=capability,
481 status=AgentStatus.ONLINE,
482 min_health=True,
483 limit=10,
484 **kwargs,
485 )
486 if not results:
487 return None
488 if load_balanced:
489 # Already sorted by load ascending
490 return results[0]
491 return results[0]
493 def get_agent(self, agent_id: str) -> Optional[RegistryEntry]:
494 """按 ID 获取 Agent。"""
495 return self._entries.get(agent_id)
497 def list_all(self, include_offline: bool = False) -> List[RegistryEntry]:
498 """列出所有 Agent。"""
499 entries = list(self._entries.values())
500 if not include_offline:
501 entries = [
502 e for e in entries
503 if e.status != AgentStatus.OFFLINE
504 or e.is_healthy(self.heartbeat_timeout)
505 ]
506 return entries
508 # ── Stats ─────────────────────────────────
510 @property
511 def total_agents(self) -> int:
512 """已注册的 Agent 总数。"""
513 return len(self._entries)
515 @property
516 def online_count(self) -> int:
517 """在线 Agent 数。"""
518 return sum(
519 1 for e in self._entries.values()
520 if e.status == AgentStatus.ONLINE and e.is_healthy(self.heartbeat_timeout)
521 )
523 def get_stats(self) -> Dict[str, Any]:
524 """获取注册中心统计。"""
525 online = 0
526 busy = 0
527 degraded = 0
528 offline = 0
529 for e in self._entries.values():
530 if not e.is_healthy(self.heartbeat_timeout):
531 offline += 1
532 elif e.status == AgentStatus.ONLINE:
533 online += 1
534 elif e.status == AgentStatus.BUSY:
535 busy += 1
536 elif e.status == AgentStatus.DEGRADED:
537 degraded += 1
539 return {
540 "total": self.total_agents,
541 "online": online,
542 "busy": busy,
543 "degraded": degraded,
544 "offline": offline,
545 "capabilities": list(self._capability_index.keys()),
546 "tags": list(self._tag_index.keys()),
547 "roles": list(self._role_index.keys()),
548 }
550 # ── Events ────────────────────────────────
552 def subscribe(
553 self,
554 event: str,
555 callback: Callable[[Dict[str, Any]], Any],
556 ) -> None:
557 """订阅注册事件。
559 Args:
560 event: 'register' | 'deregister' | 'status_change' | 'heartbeat'
561 callback: 回调函数,接收事件 dict
562 """
563 if event in self._subscribers:
564 self._subscribers[event].append(callback)
566 def _emit(self, event: str, data: Dict[str, Any]) -> None:
567 """触发事件。"""
568 for cb in self._subscribers.get(event, []):
569 try:
570 cb(data)
571 except Exception:
572 pass
574 # ── Internal Helpers ─────────────────────
576 def _update_status(self, agent_id: str, new_status: AgentStatus) -> None:
577 """更新 Agent 状态并触发事件。"""
578 if agent_id in self._entries:
579 old_status = self._entries[agent_id].status
580 self._entries[agent_id].status = new_status
581 if old_status != new_status:
582 self._emit("status_change", {
583 "agent_id": agent_id,
584 "old_status": old_status.value,
585 "new_status": new_status.value,
586 })
588 def _build_indices(
589 self,
590 agent_id: str,
591 card: DiscoveryCard,
592 tags: List[str],
593 ) -> None:
594 """构建反向索引。"""
595 # Capability index
596 for cap in card.capabilities:
597 if cap.name not in self._capability_index:
598 self._capability_index[cap.name] = set()
599 self._capability_index[cap.name].add(agent_id)
601 # Tag index
602 for tag in tags:
603 if tag not in self._tag_index:
604 self._tag_index[tag] = set()
605 self._tag_index[tag].add(agent_id)
607 # Role index (from provider)
608 role = card.provider.get("role", "")
609 if role:
610 if role not in self._role_index:
611 self._role_index[role] = set()
612 self._role_index[role].add(agent_id)
614 def _cleanup_indices(self, agent_id: str) -> None:
615 """从所有索引中移除 Agent。"""
616 for cap_set in self._capability_index.values():
617 cap_set.discard(agent_id)
618 for tag_set in self._tag_index.values():
619 tag_set.discard(agent_id)
620 for role_set in self._role_index.values():
621 role_set.discard(agent_id)
624# ── A2A Registry Bridge ─────────────────────
627class A2ARegistryBridge:
628 """A2A Registry + A2A Client 桥接。
630 自动从 Registry 发现 Agent 并创建 A2A Client。
632 Usage:
633 bridge = A2ARegistryBridge(registry)
634 client = await bridge.get_client(capability="code_review")
635 result = await client.send_and_wait_for_reply("Review this code...")
636 """
638 def __init__(self, registry: AgentRegistry):
639 self._registry = registry
640 self._clients: Dict[str, Any] = {} # agent_id -> A2AClient
642 async def get_client(
643 self,
644 capability: Optional[str] = None,
645 agent_id: Optional[str] = None,
646 **kwargs,
647 ) -> Any:
648 """获取已发现的 Agent 的 A2A Client。
650 Args:
651 capability: 按能力发现
652 agent_id: 直接指定 Agent ID
653 """
654 from agentos.protocols.a2a import A2AClient
656 if agent_id and agent_id in self._clients:
657 return self._clients[agent_id]
659 entry = None
660 if agent_id:
661 entry = self._registry.get_agent(agent_id)
662 else:
663 entry = self._registry.discover_one(capability=capability, **kwargs)
665 if not entry:
666 raise RuntimeError(
667 f"No available agent found for capability={capability}"
668 )
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) ──
696from dataclasses import dataclass, field
697import time
699@dataclass
700class AgentRecord:
701 agent_id: str
702 capabilities: list = field(default_factory=list)
703 endpoint: str = ""
704 load: float = 0.0
705 healthy: bool = True
706 name: str = ""
707 version: str = "1.0"
708 registered_at: float = field(default_factory=time.time)
710# ==========================================================================
711# Compat: AgentRecord + lightweight compat API for test_core.py
712# ==========================================================================
715# Re-use the AgentRecord from a2a module to avoid duplication
716try:
717 from agentos.protocols.a2a import AgentRecord
718except ImportError:
719 from dataclasses import dataclass, field
720 import time as _time
722 @dataclass
723 class AgentRecord:
724 agent_id: str
725 capabilities: list = field(default_factory=list)
726 endpoint: str = ""
727 load: float = 0.0
728 _last_heartbeat: float = field(default_factory=_time.time, repr=False)
730 @property
731 def healthy(self) -> bool:
732 return (_time.time() - self._last_heartbeat) < 60.0
735def _patch_agent_registry():
736 """Extend AgentRegistry with compat methods."""
737 import time as _time
739 def compat_register(self, record):
740 from agentos.protocols.a2a import AgentRecord as AR
741 if not hasattr(self, '_compat_records'):
742 self._compat_records = {}
743 if isinstance(record, AR):
744 record._last_heartbeat = _time.time()
745 self._compat_records[record.agent_id] = record
746 else:
747 self._orig_register(record)
749 def compat_get(self, agent_id):
750 if not hasattr(self, '_compat_records'):
751 self._compat_records = {}
752 return self._compat_records.get(agent_id)
754 def compat_find_by_capability(self, capability):
755 if not hasattr(self, '_compat_records'):
756 self._compat_records = {}
757 return [r for r in self._compat_records.values() if capability in r.capabilities]
759 def compat_heartbeat(self, agent_id):
760 if not hasattr(self, '_compat_records'):
761 self._compat_records = {}
762 r = self._compat_records.get(agent_id)
763 if r:
764 r._last_heartbeat = _time.time()
766 def compat_pick_least_loaded(self, capability):
767 candidates = compat_find_by_capability(self, capability)
768 if not candidates:
769 return None
770 return min(candidates, key=lambda r: r.load)
772 if not hasattr(AgentRegistry, '_orig_register'):
773 AgentRegistry._orig_register = AgentRegistry.register
774 AgentRegistry.register = compat_register
775 AgentRegistry.get = compat_get
776 AgentRegistry.find_by_capability = compat_find_by_capability
777 AgentRegistry.heartbeat = compat_heartbeat
778 AgentRegistry.pick_least_loaded = compat_pick_least_loaded
781_patch_agent_registry()