300 lines
11 KiB
Python
300 lines
11 KiB
Python
"""Agent Communication Bus - backed by NATS JetStream with local fallback."""
|
|
import asyncio
|
|
import json
|
|
import uuid
|
|
from datetime import datetime
|
|
from typing import Dict, List, Callable, Optional, Any
|
|
from collections import defaultdict
|
|
from dataclasses import dataclass
|
|
from enum import Enum
|
|
|
|
import nats
|
|
import nats.js.errors
|
|
|
|
from core.types import AgentMessage, MessageType
|
|
from core.logger import logger
|
|
from config.system_config import NATS_HOST, NATS_PORT, NATS_USER, NATS_PASSWORD
|
|
|
|
|
|
class QueuePriority(str, Enum):
|
|
HIGH = "high"
|
|
NORMAL = "normal"
|
|
LOW = "low"
|
|
|
|
|
|
@dataclass
|
|
class QueuedMessage:
|
|
message: AgentMessage
|
|
priority: QueuePriority = QueuePriority.NORMAL
|
|
retry_count: int = 0
|
|
max_retries: int = 3
|
|
|
|
|
|
class MessageBus:
|
|
"""
|
|
Central message bus backed by NATS JetStream.
|
|
|
|
Call `await message_bus.connect()` once at startup before starting agents.
|
|
Falls back to in-process local dispatch when not connected.
|
|
"""
|
|
|
|
def __init__(self):
|
|
self._nc = None # NATS client
|
|
self._js = None # JetStream context
|
|
self._nats_subs: list = [] # hold sub refs to prevent GC
|
|
|
|
self._subscribers: Dict[MessageType, List[Callable]] = defaultdict(list)
|
|
self._queues: Dict[str, List[QueuedMessage]] = defaultdict(list)
|
|
self._agents: Dict[str, Any] = {}
|
|
self._history: List[AgentMessage] = []
|
|
self._hooks: Dict[str, List[Callable]] = defaultdict(list)
|
|
self._lock = asyncio.Lock()
|
|
|
|
# ------------------------------------------------------------------ #
|
|
# Connection #
|
|
# ------------------------------------------------------------------ #
|
|
|
|
async def connect(self):
|
|
"""Connect to NATS and create the logistics JetStream stream."""
|
|
self._nc = await nats.connect(
|
|
servers=[f"nats://{NATS_HOST}:{NATS_PORT}"],
|
|
user=NATS_USER,
|
|
password=NATS_PASSWORD,
|
|
error_cb=self._on_nats_error,
|
|
reconnected_cb=self._on_reconnect,
|
|
max_reconnect_attempts=10,
|
|
)
|
|
self._js = self._nc.jetstream()
|
|
|
|
try:
|
|
await self._js.add_stream(name="logistics", subjects=["logistics.>"])
|
|
logger.info("NATS stream 'logistics' created")
|
|
except nats.js.errors.BadRequestError:
|
|
pass # stream already exists
|
|
except Exception as e:
|
|
logger.warning(f"Could not create NATS stream: {e}")
|
|
|
|
# Set up push subscriptions for all pre-registered message types
|
|
for msg_type in list(self._subscribers.keys()):
|
|
await self._setup_type_sub(msg_type)
|
|
|
|
# Set up direct-message subscriptions for all pre-registered agents
|
|
for agent_id in list(self._agents.keys()):
|
|
await self._setup_agent_sub(agent_id)
|
|
|
|
logger.info("MessageBus connected to NATS JetStream")
|
|
|
|
async def _on_nats_error(self, err):
|
|
logger.warning(f"NATS error: {err}")
|
|
|
|
async def _on_reconnect(self):
|
|
logger.info("NATS reconnected")
|
|
|
|
async def disconnect(self):
|
|
"""Drain and close NATS connection."""
|
|
if self._nc:
|
|
await self._nc.drain()
|
|
|
|
# ------------------------------------------------------------------ #
|
|
# Agent registry #
|
|
# ------------------------------------------------------------------ #
|
|
|
|
def register_agent(self, agent):
|
|
self._agents[agent.agent_id] = agent
|
|
logger.debug(f"Agent registered: {agent.agent_id} ({agent.agent_type})")
|
|
if self._js is not None:
|
|
asyncio.create_task(self._setup_agent_sub(agent.agent_id))
|
|
|
|
def unregister_agent(self, agent_id: str):
|
|
self._agents.pop(agent_id, None)
|
|
logger.debug(f"Agent unregistered: {agent_id}")
|
|
|
|
# ------------------------------------------------------------------ #
|
|
# Pub / Sub #
|
|
# ------------------------------------------------------------------ #
|
|
|
|
def subscribe(self, message_type: MessageType, callback: Callable):
|
|
"""Subscribe to a broadcast message type."""
|
|
self._subscribers[message_type].append(callback)
|
|
if self._js is not None:
|
|
asyncio.create_task(self._setup_type_sub(message_type))
|
|
|
|
def unsubscribe(self, message_type: MessageType, callback: Callable):
|
|
if message_type in self._subscribers:
|
|
try:
|
|
self._subscribers[message_type].remove(callback)
|
|
except ValueError:
|
|
pass
|
|
|
|
async def publish(self, message: AgentMessage):
|
|
"""Publish a message — via NATS if connected, otherwise local dispatch."""
|
|
async with self._lock:
|
|
self._history.append(message)
|
|
if len(self._history) > 1000:
|
|
self._history = self._history[-500:]
|
|
|
|
if self._js is not None:
|
|
await self._nats_publish(message)
|
|
else:
|
|
await self._local_dispatch(message)
|
|
|
|
await self._trigger_hook(f"on_{message.message_type.value}", message)
|
|
logger.debug(f"[{message.sender}] -> [{message.recipient}]: {message.message_type.value}")
|
|
|
|
async def send_to_agent(
|
|
self,
|
|
sender: str,
|
|
recipient: str,
|
|
message_type: MessageType,
|
|
payload: Dict[str, Any],
|
|
correlation_id: Optional[str] = None,
|
|
) -> str:
|
|
message_id = str(uuid.uuid4())
|
|
message = AgentMessage(
|
|
message_id=message_id,
|
|
sender=sender,
|
|
recipient=recipient,
|
|
message_type=message_type,
|
|
payload=payload,
|
|
timestamp=datetime.now(),
|
|
correlation_id=correlation_id or message_id,
|
|
)
|
|
await self.publish(message)
|
|
return message_id
|
|
|
|
async def broadcast(
|
|
self,
|
|
sender: str,
|
|
message_type: MessageType,
|
|
payload: Dict[str, Any],
|
|
correlation_id: Optional[str] = None,
|
|
) -> str:
|
|
return await self.send_to_agent(sender, "ALL", message_type, payload, correlation_id)
|
|
|
|
async def get_messages(self, agent_id: str) -> List[AgentMessage]:
|
|
async with self._lock:
|
|
queued = self._queues.pop(agent_id, [])
|
|
return [q.message for q in queued]
|
|
|
|
async def peek_messages(self, agent_id: str) -> List[AgentMessage]:
|
|
return [q.message for q in self._queues.get(agent_id, [])]
|
|
|
|
# ------------------------------------------------------------------ #
|
|
# Hooks #
|
|
# ------------------------------------------------------------------ #
|
|
|
|
def add_hook(self, event: str, callback: Callable):
|
|
self._hooks[event].append(callback)
|
|
|
|
async def _trigger_hook(self, event: str, message: AgentMessage):
|
|
for callback in self._hooks.get(event, []):
|
|
try:
|
|
if asyncio.iscoroutinefunction(callback):
|
|
await callback(message)
|
|
else:
|
|
callback(message)
|
|
except Exception as e:
|
|
logger.error(f"Hook error [{event}]: {e}")
|
|
|
|
# ------------------------------------------------------------------ #
|
|
# Agent / history helpers #
|
|
# ------------------------------------------------------------------ #
|
|
|
|
def get_agent(self, agent_id: str):
|
|
return self._agents.get(agent_id)
|
|
|
|
def get_all_agents(self) -> Dict[str, Any]:
|
|
return self._agents.copy()
|
|
|
|
def get_messages_by_type(self, message_type: MessageType) -> List[AgentMessage]:
|
|
return [m for m in self._history if m.message_type == message_type]
|
|
|
|
def get_messages_by_sender(self, sender: str) -> List[AgentMessage]:
|
|
return [m for m in self._history if m.sender == sender]
|
|
|
|
def clear_history(self):
|
|
self._history = []
|
|
|
|
# ------------------------------------------------------------------ #
|
|
# Internal: NATS helpers #
|
|
# ------------------------------------------------------------------ #
|
|
|
|
async def _nats_publish(self, message: AgentMessage):
|
|
if message.recipient != "ALL":
|
|
subject = f"logistics.direct.{message.recipient}"
|
|
else:
|
|
subject = f"logistics.{message.message_type.value}"
|
|
try:
|
|
await self._js.publish(subject, message.to_json().encode())
|
|
except Exception as e:
|
|
logger.warning(f"NATS publish error: {e} — falling back to local dispatch")
|
|
await self._local_dispatch(message)
|
|
|
|
async def _setup_type_sub(self, message_type: MessageType):
|
|
"""Create a JetStream push subscription for a broadcast message type."""
|
|
subject = f"logistics.{message_type.value}"
|
|
durable = f"logiflow-{message_type.value}"
|
|
callbacks = self._subscribers[message_type]
|
|
|
|
async def handler(msg):
|
|
try:
|
|
agent_msg = AgentMessage.from_json(msg.data.decode())
|
|
for cb in list(callbacks):
|
|
try:
|
|
if asyncio.iscoroutinefunction(cb):
|
|
await cb(agent_msg)
|
|
else:
|
|
cb(agent_msg)
|
|
except Exception as e:
|
|
logger.error(f"Subscriber callback error: {e}")
|
|
except Exception as e:
|
|
logger.error(f"NATS type-sub decode error [{message_type.value}]: {e}")
|
|
finally:
|
|
await msg.ack()
|
|
|
|
try:
|
|
sub = await self._js.subscribe(subject, durable=durable, cb=handler)
|
|
self._nats_subs.append(sub)
|
|
except Exception as e:
|
|
logger.warning(f"Could not subscribe to NATS subject {subject}: {e}")
|
|
|
|
async def _setup_agent_sub(self, agent_id: str):
|
|
"""Create a JetStream push subscription for directed messages to an agent."""
|
|
subject = f"logistics.direct.{agent_id}"
|
|
durable = f"logiflow-direct-{agent_id}"
|
|
|
|
async def handler(msg):
|
|
try:
|
|
agent_msg = AgentMessage.from_json(msg.data.decode())
|
|
async with self._lock:
|
|
self._queues[agent_id].append(QueuedMessage(agent_msg))
|
|
except Exception as e:
|
|
logger.error(f"NATS agent-sub decode error [{agent_id}]: {e}")
|
|
finally:
|
|
await msg.ack()
|
|
|
|
try:
|
|
sub = await self._js.subscribe(subject, durable=durable, cb=handler)
|
|
self._nats_subs.append(sub)
|
|
except Exception as e:
|
|
logger.warning(f"Could not subscribe to NATS subject {subject}: {e}")
|
|
|
|
async def _local_dispatch(self, message: AgentMessage):
|
|
"""In-process dispatch used when NATS is not connected."""
|
|
if message.recipient != "ALL":
|
|
async with self._lock:
|
|
self._queues[message.recipient].append(QueuedMessage(message))
|
|
else:
|
|
for callback in list(self._subscribers.get(message.message_type, [])):
|
|
try:
|
|
if asyncio.iscoroutinefunction(callback):
|
|
await callback(message)
|
|
else:
|
|
callback(message)
|
|
except Exception as e:
|
|
logger.error(f"Local dispatch callback error: {e}")
|
|
|
|
|
|
# Global message bus instance
|
|
message_bus = MessageBus()
|