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

1""" 

2AgentOS v1.14.0 — A2A Agent Registry(服务发现)。 

3 

4Agent 注册中心实现: 

5- 服务注册/注销/心跳 

6- 能力广播 (Agent Card) 

7- 服务发现 (按能力/角色查询) 

8- 健康检查自动摘除 

9""" 

10 

11from __future__ import annotations 

12 

13import asyncio 

14import time 

15from dataclasses import dataclass, field 

16from enum import Enum 

17from typing import Any, Callable, Dict, List, Optional, Set, Union 

18 

19 

20# ── Agent Info ─────────────────────────────── 

21 

22 

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) 

33 

34 

35# ── Agent Card ─────────────────────────────── 

36 

37 

38@dataclass 

39class DiscoveryCapability: 

40 """服务发现用 Agent 能力描述。 

41 

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) # 性能指标 

50 

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 } 

59 

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 ) 

69 

70 

71@dataclass 

72class DiscoveryCard: 

73 """Agent 发现名片 — 向注册中心声明的自身信息。 

74 

75 符合 Google A2A AgentCard 规范。 

76 不同于 protocols.agent_card.AgentCard(面向本地广播), 

77 本类面向远程注册中心的服务发现。 

78 """ 

79 

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) 

96 

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 } 

116 

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 ) 

140 

141 def has_capability(self, cap_name: str) -> bool: 

142 """检查是否具备某项能力。""" 

143 return any(c.name == cap_name for c in self.capabilities) 

144 

145 

146# ── Registry Entry ────────────────────────── 

147 

148 

149class AgentStatus(str, Enum): 

150 """Agent 运行状态。""" 

151 ONLINE = "online" 

152 BUSY = "busy" 

153 DEGRADED = "degraded" 

154 OFFLINE = "offline" 

155 

156 

157@dataclass 

158class RegistryEntry: 

159 """注册中心中的单条 Agent 记录。""" 

160 

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) 

171 

172 def is_healthy(self, heartbeat_timeout: float = 30.0) -> bool: 

173 """心跳是否正常。""" 

174 return (time.time() - self.last_heartbeat) < heartbeat_timeout 

175 

176 @property 

177 def uptime_seconds(self) -> float: 

178 """注册后的运行时长。""" 

179 return time.time() - self.registered_at 

180 

181 @property 

182 def endpoint(self) -> str: 

183 """Agent 服务端点(映射 card.url)。""" 

184 return self.card.url 

185 

186 

187# ── Agent Registry ────────────────────────── 

188 

189 

190class AgentRegistry: 

191 """A2A Agent 注册中心。 

192 

193 功能: 

194 - register: 注册 Agent(带心跳保活) 

195 - discover: 按能力/角色/标签发现 Agent 

196 - health_check: 自动摘除失联 Agent 

197 - subscribe: 订阅注册事件 

198 

199 Usage: 

200 registry = AgentRegistry(heartbeat_timeout=30) 

201 registry.register(agent_card) 

202 

203 # 发现能处理 code_review 的 Agent 

204 agents = registry.discover(capability="code_review") 

205 

206 # 按标签查找 

207 agents = registry.discover(tags=["production", "high-priority"]) 

208 """ 

209 

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 

223 

224 # Event subscribers 

225 self._subscribers: Dict[str, List[Callable]] = { 

226 "register": [], 

227 "deregister": [], 

228 "status_change": [], 

229 "heartbeat": [], 

230 } 

231 

232 # Health check task 

233 self._health_check_task: Optional[asyncio.Task] = None 

234 self._running = False 

235 

236 # ── Lifecycle ───────────────────────────── 

237 

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()) 

245 

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 

256 

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) 

265 

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) 

275 

276 for agent_id in to_remove: 

277 self._cleanup_indices(agent_id) 

278 

279 # ── Registration ───────────────────────── 

280 

281 def register( 

282 self, 

283 card_or_info: Union[DiscoveryCard, "AgentInfo"], 

284 tags: Optional[List[str]] = None, 

285 ) -> str: 

286 """注册一个 Agent。 

287 

288 Args: 

289 card_or_info: Agent 名片 (DiscoveryCard) 或轻量信息 (AgentInfo) 

290 tags: 自定义标签 

291 

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 

310 

311 agent_id = card.agent_id 

312 now = time.time() 

313 

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 

324 

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 

336 

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 

345 

346 def heartbeat(self, agent_id: str, load: Optional[float] = None) -> bool: 

347 """Agent 心跳上报。 

348 

349 Args: 

350 agent_id: Agent ID 

351 load: 当前负载 0.0~1.0 

352 

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 

366 

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 

385 

386 # ── Discovery ──────────────────────────── 

387 

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。 

400 

401 Args: 

402 capability: 按能力名称筛选 

403 tags: 按标签筛选(AND 逻辑) 

404 role: 按角色筛选 

405 status: 按状态筛选 

406 service_tier: 按服务等级筛选 

407 supports_streaming: 是否支持流式 

408 min_health: 仅返回心跳正常的 Agent 

409 limit: 最大返回数 

410 

411 Returns: 

412 匹配的 RegistryEntry 列表 

413 """ 

414 candidates: Set[str] = set(self._entries.keys()) 

415 

416 # Capability filter 

417 if capability: 

418 cap_agents = self._capability_index.get(capability, set()) 

419 candidates &= cap_agents 

420 

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 

426 

427 # Role filter 

428 if role: 

429 role_agents = self._role_index.get(role, set()) 

430 candidates &= role_agents 

431 

432 # Collect results 

433 results = [] 

434 for agent_id in candidates: 

435 entry = self._entries[agent_id] 

436 

437 # Status filter 

438 if status and entry.status != status: 

439 continue 

440 

441 # Health filter 

442 if min_health and not entry.is_healthy(self.heartbeat_timeout): 

443 continue 

444 

445 # Service tier filter 

446 if service_tier and entry.card.service_tier != service_tier: 

447 continue 

448 

449 # Streaming filter 

450 if supports_streaming is not None and \ 

451 entry.card.supports_streaming != supports_streaming: 

452 continue 

453 

454 results.append(entry) 

455 

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)) 

464 

465 return results[:limit] 

466 

467 def discover_one( 

468 self, 

469 capability: Optional[str] = None, 

470 load_balanced: bool = True, 

471 **kwargs, 

472 ) -> Optional[RegistryEntry]: 

473 """发现单个最优 Agent。 

474 

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] 

492 

493 def get_agent(self, agent_id: str) -> Optional[RegistryEntry]: 

494 """按 ID 获取 Agent。""" 

495 return self._entries.get(agent_id) 

496 

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 

507 

508 # ── Stats ───────────────────────────────── 

509 

510 @property 

511 def total_agents(self) -> int: 

512 """已注册的 Agent 总数。""" 

513 return len(self._entries) 

514 

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 ) 

522 

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 

538 

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 } 

549 

550 # ── Events ──────────────────────────────── 

551 

552 def subscribe( 

553 self, 

554 event: str, 

555 callback: Callable[[Dict[str, Any]], Any], 

556 ) -> None: 

557 """订阅注册事件。 

558 

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) 

565 

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 

573 

574 # ── Internal Helpers ───────────────────── 

575 

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 }) 

587 

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) 

600 

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) 

606 

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) 

613 

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) 

622 

623 

624# ── A2A Registry Bridge ───────────────────── 

625 

626 

627class A2ARegistryBridge: 

628 """A2A Registry + A2A Client 桥接。 

629 

630 自动从 Registry 发现 Agent 并创建 A2A Client。 

631 

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 """ 

637 

638 def __init__(self, registry: AgentRegistry): 

639 self._registry = registry 

640 self._clients: Dict[str, Any] = {} # agent_id -> A2AClient 

641 

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。 

649 

650 Args: 

651 capability: 按能力发现 

652 agent_id: 直接指定 Agent ID 

653 """ 

654 from agentos.protocols.a2a import A2AClient 

655 

656 if agent_id and agent_id in self._clients: 

657 return self._clients[agent_id] 

658 

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) 

664 

665 if not entry: 

666 raise RuntimeError( 

667 f"No available agent found for capability={capability}" 

668 ) 

669 

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 

676 

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() 

685 

686 def invalidate(self, agent_id: str) -> None: 

687 """使某个 Agent 的缓存 Client 失效。""" 

688 self._clients.pop(agent_id, None) 

689 

690 

691# ── 全局单例 ───────────────────────────────── 

692 

693default_registry = AgentRegistry() 

694 

695# ── AgentRecord (test compatibility) ── 

696from dataclasses import dataclass, field 

697import time 

698 

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) 

709 

710# ========================================================================== 

711# Compat: AgentRecord + lightweight compat API for test_core.py 

712# ========================================================================== 

713 

714 

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 

721 

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) 

729 

730 @property 

731 def healthy(self) -> bool: 

732 return (_time.time() - self._last_heartbeat) < 60.0 

733 

734 

735def _patch_agent_registry(): 

736 """Extend AgentRegistry with compat methods.""" 

737 import time as _time 

738 

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) 

748 

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) 

753 

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] 

758 

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() 

765 

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) 

771 

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 

779 

780 

781_patch_agent_registry()