Files
AI_engine/core/message_bus.py
2026-06-26 16:08:31 +05:30

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