Initial commit
This commit is contained in:
273
core/agent.py
Normal file
273
core/agent.py
Normal file
@@ -0,0 +1,273 @@
|
||||
"""Base Agent class - All agents inherit from this."""
|
||||
import asyncio
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
from typing import Dict, List, Optional, Any
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
from core.types import AgentMessage, MessageType, AgentState, AgentTask, Priority
|
||||
from core.message_bus import message_bus
|
||||
from core.logger import logger
|
||||
|
||||
|
||||
class Agent(ABC):
|
||||
"""
|
||||
Base class for all AI agents in the logistics system.
|
||||
Each agent has its own task queue and can communicate via the message bus.
|
||||
"""
|
||||
|
||||
def __init__(self, agent_id: str, agent_type: str, description: str = ""):
|
||||
self.agent_id = agent_id
|
||||
self.agent_type = agent_type
|
||||
self.description = description
|
||||
self.state = AgentState(
|
||||
agent_id=agent_id,
|
||||
agent_type=agent_type,
|
||||
status="idle"
|
||||
)
|
||||
self._task_queue: asyncio.Queue = asyncio.Queue()
|
||||
self._running = False
|
||||
self._task_handlers: Dict[str, callable] = {}
|
||||
|
||||
# Register with message bus
|
||||
message_bus.register_agent(self)
|
||||
|
||||
async def start(self):
|
||||
"""Start the agent's processing loop."""
|
||||
self._running = True
|
||||
self.state.status = "idle"
|
||||
logger.info(f"Agent started: {self.agent_id}")
|
||||
|
||||
while self._running:
|
||||
try:
|
||||
try:
|
||||
task = await asyncio.wait_for(
|
||||
self._task_queue.get(),
|
||||
timeout=1.0
|
||||
)
|
||||
await self._process_task(task)
|
||||
except asyncio.TimeoutError:
|
||||
await self._heartbeat()
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error in agent {self.agent_id}: {e}")
|
||||
self.state.status = "error"
|
||||
await asyncio.sleep(5)
|
||||
|
||||
async def stop(self):
|
||||
"""Stop the agent."""
|
||||
self._running = False
|
||||
self.state.status = "stopped"
|
||||
logger.info(f"Agent stopped: {self.agent_id}")
|
||||
|
||||
async def _process_task(self, task: AgentTask):
|
||||
"""Process a task from the queue."""
|
||||
self.state.status = "working"
|
||||
self.state.current_task = task.task_id
|
||||
|
||||
try:
|
||||
if task.task_type in self._task_handlers:
|
||||
handler = self._task_handlers[task.task_type]
|
||||
result = await handler(task)
|
||||
task.status = "completed"
|
||||
task.result = result
|
||||
else:
|
||||
result = await self.handle_task(task)
|
||||
task.status = "completed"
|
||||
task.result = result
|
||||
|
||||
self.state.tasks_completed += 1
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Task {task.task_id} failed in {self.agent_id}: {e}")
|
||||
task.status = "failed"
|
||||
task.error = str(e)
|
||||
self.state.tasks_failed += 1
|
||||
|
||||
finally:
|
||||
self.state.current_task = None
|
||||
self.state.last_active = datetime.now()
|
||||
self.state.status = "idle"
|
||||
|
||||
@abstractmethod
|
||||
async def handle_task(self, task: AgentTask) -> Dict[str, Any]:
|
||||
"""Override this method to handle tasks. Return a dict with results."""
|
||||
pass
|
||||
|
||||
async def _heartbeat(self):
|
||||
"""Called periodically when idle. Override for custom behavior."""
|
||||
pass
|
||||
|
||||
def register_task_handler(self, task_type: str, handler: callable):
|
||||
self._task_handlers[task_type] = handler
|
||||
|
||||
async def submit_task(self, task: AgentTask):
|
||||
await self._task_queue.put(task)
|
||||
|
||||
async def send_message(self, recipient: str, message_type: MessageType, payload: Dict[str, Any], correlation_id: Optional[str] = None) -> str:
|
||||
return await message_bus.send_to_agent(
|
||||
sender=self.agent_id,
|
||||
recipient=recipient,
|
||||
message_type=message_type,
|
||||
payload=payload,
|
||||
correlation_id=correlation_id
|
||||
)
|
||||
|
||||
async def broadcast(self, message_type: MessageType, payload: Dict[str, Any], correlation_id: Optional[str] = None) -> str:
|
||||
return await message_bus.broadcast(
|
||||
sender=self.agent_id,
|
||||
message_type=message_type,
|
||||
payload=payload,
|
||||
correlation_id=correlation_id
|
||||
)
|
||||
|
||||
async def receive_messages(self) -> List[AgentMessage]:
|
||||
return await message_bus.get_messages(self.agent_id)
|
||||
|
||||
def subscribe_to(self, message_type: MessageType, callback: callable):
|
||||
message_bus.subscribe(message_type, callback)
|
||||
|
||||
async def think(self, context: str, options: List[str] = None) -> str:
|
||||
return context
|
||||
|
||||
|
||||
class MasterAgent(Agent):
|
||||
"""
|
||||
Master Agent (JARVIS) - Orchestrates all other agents.
|
||||
Makes high-level decisions and assigns work.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__(
|
||||
agent_id="JARVIS",
|
||||
agent_type="master",
|
||||
description="Central orchestrator for all logistics operations"
|
||||
)
|
||||
self._sub_agents: Dict[str, Agent] = {}
|
||||
self._active_orders: Dict[str, Dict] = {}
|
||||
self._decision_log: List[Dict] = []
|
||||
|
||||
def register_sub_agent(self, agent: Agent):
|
||||
self._sub_agents[agent.agent_id] = agent
|
||||
logger.info(f"JARVIS registered sub-agent: {agent.agent_id}")
|
||||
|
||||
async def handle_task(self, task: AgentTask) -> Dict[str, Any]:
|
||||
if task.task_type == "orchestrate_order":
|
||||
return await self._orchestrate_order(task)
|
||||
elif task.task_type == "monitor_operations":
|
||||
return await self._monitor_operations(task)
|
||||
elif task.task_type == "handle_exception":
|
||||
return await self._handle_exception(task)
|
||||
elif task.task_type == "generate_report":
|
||||
return await self._generate_report(task)
|
||||
else:
|
||||
return {"status": "unknown_task", "task_type": task.task_type}
|
||||
|
||||
async def _orchestrate_order(self, task: AgentTask) -> Dict[str, Any]:
|
||||
order_data = task.data.get("order", {})
|
||||
order_id = order_data.get("order_id", "unknown")
|
||||
|
||||
logger.info(f"JARVIS: Orchestrating order {order_id}")
|
||||
|
||||
order_agent = self._sub_agents.get("ORDER_AGENT")
|
||||
if order_agent:
|
||||
await order_agent.submit_task(AgentTask(
|
||||
task_id=f"{order_id}_validate",
|
||||
agent_type="order",
|
||||
task_type="validate_order",
|
||||
data={"order": order_data}
|
||||
))
|
||||
|
||||
dispatch_agent = self._sub_agents.get("DISPATCH_AGENT")
|
||||
if dispatch_agent:
|
||||
await dispatch_agent.submit_task(AgentTask(
|
||||
task_id=f"{order_id}_dispatch",
|
||||
agent_type="dispatch",
|
||||
task_type="analyze_and_assign",
|
||||
data={"order": order_data}
|
||||
))
|
||||
|
||||
customer_agent = self._sub_agents.get("CUSTOMER_AGENT")
|
||||
if customer_agent:
|
||||
await customer_agent.submit_task(AgentTask(
|
||||
task_id=f"{order_id}_notify",
|
||||
agent_type="customer",
|
||||
task_type="send_confirmation",
|
||||
data={"order": order_data}
|
||||
))
|
||||
|
||||
self._decision_log.append({
|
||||
"timestamp": datetime.now(),
|
||||
"action": "orchestrate_order",
|
||||
"order_id": order_id,
|
||||
"agents_involved": list(self._sub_agents.keys())
|
||||
})
|
||||
|
||||
return {
|
||||
"status": "orchestrated",
|
||||
"order_id": order_id,
|
||||
"assigned_agents": list(self._sub_agents.keys())
|
||||
}
|
||||
|
||||
async def _monitor_operations(self, task: AgentTask) -> Dict[str, Any]:
|
||||
status_report = {
|
||||
"timestamp": datetime.now().isoformat(),
|
||||
"total_agents": len(self._sub_agents),
|
||||
"agent_statuses": {},
|
||||
"active_orders": len(self._active_orders),
|
||||
"recent_decisions": self._decision_log[-10:]
|
||||
}
|
||||
|
||||
for agent_id, agent in self._sub_agents.items():
|
||||
status_report["agent_statuses"][agent_id] = {
|
||||
"status": agent.state.status,
|
||||
"tasks_completed": agent.state.tasks_completed,
|
||||
"tasks_failed": agent.state.tasks_failed,
|
||||
"current_task": agent.state.current_task
|
||||
}
|
||||
|
||||
return status_report
|
||||
|
||||
async def _handle_exception(self, task: AgentTask) -> Dict[str, Any]:
|
||||
exception_type = task.data.get("type", "unknown")
|
||||
logger.warning(f"JARVIS: Handling exception - {exception_type}")
|
||||
|
||||
exception_agent = self._sub_agents.get("EXCEPTION_AGENT")
|
||||
if exception_agent:
|
||||
await exception_agent.submit_task(task)
|
||||
return {"status": "delegated", "exception_type": exception_type}
|
||||
|
||||
return {"status": "no_exception_agent"}
|
||||
|
||||
async def _generate_report(self, task: AgentTask) -> Dict[str, Any]:
|
||||
return {
|
||||
"report_type": task.data.get("type", "summary"),
|
||||
"generated_at": datetime.now().isoformat(),
|
||||
"agents": [a.state for a in self._sub_agents.values()],
|
||||
"decisions": self._decision_log[-50:]
|
||||
}
|
||||
|
||||
async def _heartbeat(self):
|
||||
for agent_id, agent in self._sub_agents.items():
|
||||
if agent.state.status == "error":
|
||||
logger.warning(f"Agent {agent_id} in error state")
|
||||
|
||||
|
||||
class SpecializedAgent(Agent):
|
||||
"""Base class for specialized domain agents."""
|
||||
|
||||
def __init__(self, agent_id: str, domain: str, description: str):
|
||||
super().__init__(agent_id, domain, description)
|
||||
self._domain_knowledge: Dict[str, Any] = {}
|
||||
self._learning_history: List[Dict] = []
|
||||
|
||||
async def think(self, context: str, options: List[str] = None) -> str:
|
||||
return f"[{self.agent_type} reasoning]: {context}"
|
||||
|
||||
def learn_from(self, experience: Dict[str, Any]):
|
||||
self._learning_history.append({
|
||||
"timestamp": datetime.now(),
|
||||
"experience": experience
|
||||
})
|
||||
if len(self._learning_history) > 100:
|
||||
self._learning_history = self._learning_history[-100:]
|
||||
Reference in New Issue
Block a user