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