Files
Agentswarm/orchestrator/agent_registry.py
T
2026-06-08 17:32:34 +08:00

150 lines
4.5 KiB
Python

"""Agent registry with Redis-backed state management."""
import json
import time
import logging
from typing import Dict, List, Optional
from enum import Enum
from pydantic import BaseModel
from .redis_client import redis_client
logger = logging.getLogger(__name__)
class AgentStatus(str, Enum):
"""Agent status enumeration."""
IDLE = "idle"
BUSY = "busy"
HANDOFF_PENDING = "handoff-pending"
FAILED = "failed"
class AgentMetadata(BaseModel):
"""Agent metadata model."""
agent_id: str
status: AgentStatus
last_heartbeat: float
capabilities: List[str]
current_task_id: Optional[str] = None
class AgentRegistry:
"""Manages agent registration and heartbeat tracking."""
HEARTBEAT_TIMEOUT = 30 # seconds
AGENT_KEY_PREFIX = "agent:"
def __init__(self):
pass
async def register_agent(
self, agent_id: str, capabilities: List[str]
) -> AgentMetadata:
"""Register a new agent."""
metadata = AgentMetadata(
agent_id=agent_id,
status=AgentStatus.IDLE,
last_heartbeat=time.time(),
capabilities=capabilities,
)
key = f"{self.AGENT_KEY_PREFIX}{agent_id}"
await redis_client.set(key, metadata.model_dump_json())
logger.info(f"Registered agent {agent_id} with capabilities: {capabilities}")
return metadata
async def deregister_agent(self, agent_id: str):
"""Deregister an agent."""
key = f"{self.AGENT_KEY_PREFIX}{agent_id}"
await redis_client.delete(key)
logger.info(f"Deregistered agent {agent_id}")
async def update_heartbeat(self, agent_id: str) -> bool:
"""Update agent heartbeat timestamp."""
key = f"{self.AGENT_KEY_PREFIX}{agent_id}"
data = await redis_client.get(key)
if not data:
logger.warning(f"Agent {agent_id} not found for heartbeat update")
return False
metadata = AgentMetadata.model_validate_json(data)
metadata.last_heartbeat = time.time()
await redis_client.set(key, metadata.model_dump_json())
return True
async def update_status(
self, agent_id: str, status: AgentStatus, task_id: Optional[str] = None
) -> bool:
"""Update agent status."""
key = f"{self.AGENT_KEY_PREFIX}{agent_id}"
data = await redis_client.get(key)
if not data:
logger.warning(f"Agent {agent_id} not found for status update")
return False
metadata = AgentMetadata.model_validate_json(data)
metadata.status = status
metadata.current_task_id = task_id
await redis_client.set(key, metadata.model_dump_json())
logger.info(f"Updated agent {agent_id} status to {status}")
return True
async def get_agent(self, agent_id: str) -> Optional[AgentMetadata]:
"""Get agent metadata."""
key = f"{self.AGENT_KEY_PREFIX}{agent_id}"
data = await redis_client.get(key)
if not data:
return None
return AgentMetadata.model_validate_json(data)
async def get_all_agents(self) -> List[AgentMetadata]:
"""Get all registered agents."""
pattern = f"{self.AGENT_KEY_PREFIX}*"
keys = await redis_client.keys(pattern)
agents = []
for key in keys:
data = await redis_client.get(key)
if data:
agents.append(AgentMetadata.model_validate_json(data))
return agents
async def get_idle_agents(self) -> List[AgentMetadata]:
"""Get all idle agents."""
all_agents = await self.get_all_agents()
return [agent for agent in all_agents if agent.status == AgentStatus.IDLE]
async def check_failed_agents(self) -> List[str]:
"""Check for agents with expired heartbeats and mark as failed."""
current_time = time.time()
failed_agents = []
all_agents = await self.get_all_agents()
for agent in all_agents:
if agent.status == AgentStatus.FAILED:
continue
time_since_heartbeat = current_time - agent.last_heartbeat
if time_since_heartbeat > self.HEARTBEAT_TIMEOUT:
await self.update_status(agent.agent_id, AgentStatus.FAILED)
failed_agents.append(agent.agent_id)
logger.warning(
f"Agent {agent.agent_id} marked as failed "
f"(no heartbeat for {time_since_heartbeat:.1f}s)"
)
return failed_agents
# Global agent registry instance
agent_registry = AgentRegistry()