Files
AI_engine/core/message_bus.py

334 lines
13 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 _deliver_to_agent(self, agent_id: str, message: AgentMessage):
"""Hand a directed message to a locally-registered agent so it is
actually consumed (via its ``deliver``). Only if no such agent exists in
this process does it fall back to the pull queue — so a message is never
silently lost, but is delivered live whenever the recipient is present."""
agent = self._agents.get(agent_id)
if agent is not None and hasattr(agent, "deliver"):
try:
await agent.deliver(message)
return
except Exception as e:
logger.error(f"Delivery to {agent_id} failed: {e}")
return
async with self._lock:
self._queues[agent_id].append(QueuedMessage(message))
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, [])]
# ------------------------------------------------------------------ #
# Telemetry #
# ------------------------------------------------------------------ #
async def publish_telemetry(self, kind: str, payload: Dict[str, Any]):
"""
Fire-and-forget observability event on plain NATS (subject `telemetry.<kind>`).
Deliberately NOT JetStream: telemetry is ephemeral fan-out for dashboards
and must never accumulate in the persistent `logistics` stream. No-op
when NATS is not connected — telemetry must never affect agent behavior.
"""
if self._nc is None or not self._nc.is_connected:
return
try:
body = {"kind": kind, "ts": datetime.now().isoformat(), **payload}
await self._nc.publish(f"telemetry.{kind}", json.dumps(body).encode())
except Exception as e:
logger.debug(f"Telemetry publish failed [{kind}]: {e}")
# ------------------------------------------------------------------ #
# 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())
await self._deliver_to_agent(agent_id, 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":
await self._deliver_to_agent(message.recipient, 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()