meerkat-comms 0.5.2

Inter-agent communication for Meerkat
Documentation
//! CommsToolDispatcher - Implements AgentToolDispatcher for comms tools.

use crate::mcp::tools::{ToolContext, handle_tools_call, tools_list};
use crate::runtime::CommsRuntime;
use crate::{Router, TrustedPeers};
use async_trait::async_trait;
use meerkat_core::AgentToolDispatcher;
use meerkat_core::ToolDispatchOutcome;
use meerkat_core::error::ToolError;
use meerkat_core::types::{ToolCallView, ToolDef, ToolProvenance, ToolResult, ToolSourceKind};
use parking_lot::RwLock;
use serde_json::Value;
use std::sync::Arc;

/// Tool dispatcher that provides comms tools.
pub struct CommsToolDispatcher<T: AgentToolDispatcher = NoOpDispatcher> {
    tool_context: ToolContext,
    inner: Option<Arc<T>>,
    tool_defs: Arc<[Arc<ToolDef>]>,
}

impl CommsToolDispatcher<NoOpDispatcher> {
    pub fn new(router: Arc<Router>, trusted_peers: Arc<RwLock<TrustedPeers>>) -> Self {
        let tool_context = ToolContext {
            router,
            trusted_peers,
            runtime: None,
        };
        let tool_defs: Arc<[Arc<ToolDef>]> = comms_tool_defs().into();
        Self {
            tool_context,
            inner: None,
            tool_defs,
        }
    }

    pub fn new_with_runtime(
        router: Arc<Router>,
        trusted_peers: Arc<RwLock<TrustedPeers>>,
        runtime: Arc<dyn meerkat_core::agent::CommsRuntime>,
    ) -> Self {
        let tool_context = ToolContext {
            router,
            trusted_peers,
            runtime: Some(runtime),
        };
        let tool_defs: Arc<[Arc<ToolDef>]> = comms_tool_defs().into();
        Self {
            tool_context,
            inner: None,
            tool_defs,
        }
    }
}

impl<T: AgentToolDispatcher> CommsToolDispatcher<T> {
    pub fn with_inner(
        router: Arc<Router>,
        trusted_peers: Arc<RwLock<TrustedPeers>>,
        inner: Arc<T>,
    ) -> Self {
        let tool_context = ToolContext {
            router,
            trusted_peers,
            runtime: None,
        };
        let mut tools = comms_tool_defs();
        tools.extend(inner.tools().iter().map(Arc::clone));
        let tool_defs: Arc<[Arc<ToolDef>]> = tools.into();
        Self {
            tool_context,
            inner: Some(inner),
            tool_defs,
        }
    }
}

pub struct NoOpDispatcher;

#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
impl AgentToolDispatcher for NoOpDispatcher {
    fn tools(&self) -> Arc<[Arc<ToolDef>]> {
        Arc::from([])
    }
    async fn dispatch(&self, call: ToolCallView<'_>) -> Result<ToolDispatchOutcome, ToolError> {
        Err(ToolError::NotFound {
            name: call.name.to_string(),
        })
    }
}

fn is_comms_tool(name: &str) -> bool {
    matches!(
        name,
        "send" | "send_message" | "send_request" | "send_response" | "peers"
    )
}

/// Canonical JSON-to-ToolDef conversion for comms tools.
pub fn comms_tool_defs() -> Vec<Arc<ToolDef>> {
    tools_list()
        .into_iter()
        .map(|t| {
            Arc::new(ToolDef {
                name: t["name"].as_str().unwrap_or_default().to_string(),
                description: t["description"].as_str().unwrap_or_default().to_string(),
                input_schema: t["inputSchema"].clone(),
                provenance: Some(ToolProvenance {
                    kind: ToolSourceKind::Comms,
                    source_id: "comms".into(),
                }),
            })
        })
        .collect()
}

#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
impl<T: AgentToolDispatcher + 'static> AgentToolDispatcher for CommsToolDispatcher<T> {
    fn tools(&self) -> Arc<[Arc<ToolDef>]> {
        Arc::clone(&self.tool_defs)
    }

    async fn dispatch(&self, call: ToolCallView<'_>) -> Result<ToolDispatchOutcome, ToolError> {
        let args: Value = serde_json::from_str(call.args.get())
            .unwrap_or_else(|_| Value::String(call.args.get().to_string()));
        if is_comms_tool(call.name) {
            let result = handle_tools_call(&self.tool_context, call.name, &args)
                .await
                .map_err(|e| ToolError::ExecutionFailed { message: e })?;
            Ok(ToolResult::new(call.id.to_string(), result.to_string(), false).into())
        } else if let Some(inner) = &self.inner {
            inner.dispatch(call).await
        } else {
            Err(ToolError::NotFound {
                name: call.name.to_string(),
            })
        }
    }
}

pub struct DynCommsToolDispatcher {
    tool_context: ToolContext,
    inner: Arc<dyn AgentToolDispatcher>,
    tool_defs: Arc<[Arc<ToolDef>]>,
}

impl DynCommsToolDispatcher {
    pub fn new(
        router: Arc<Router>,
        trusted_peers: Arc<RwLock<TrustedPeers>>,
        inner: Arc<dyn AgentToolDispatcher>,
    ) -> Self {
        let tool_context = ToolContext {
            router,
            trusted_peers,
            runtime: None,
        };
        let mut tools = comms_tool_defs();
        tools.extend(inner.tools().iter().map(Arc::clone));
        let tool_defs: Arc<[Arc<ToolDef>]> = tools.into();
        Self {
            tool_context,
            inner,
            tool_defs,
        }
    }

    pub fn new_with_runtime(
        router: Arc<Router>,
        trusted_peers: Arc<RwLock<TrustedPeers>>,
        runtime: Arc<dyn meerkat_core::agent::CommsRuntime>,
        inner: Arc<dyn AgentToolDispatcher>,
    ) -> Self {
        let tool_context = ToolContext {
            router,
            trusted_peers,
            runtime: Some(runtime),
        };
        let mut tools = comms_tool_defs();
        tools.extend(inner.tools().iter().map(Arc::clone));
        let tool_defs: Arc<[Arc<ToolDef>]> = tools.into();
        Self {
            tool_context,
            inner,
            tool_defs,
        }
    }
}

#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
impl AgentToolDispatcher for DynCommsToolDispatcher {
    fn tools(&self) -> Arc<[Arc<ToolDef>]> {
        Arc::clone(&self.tool_defs)
    }

    async fn dispatch(&self, call: ToolCallView<'_>) -> Result<ToolDispatchOutcome, ToolError> {
        let args: Value = serde_json::from_str(call.args.get())
            .unwrap_or_else(|_| Value::String(call.args.get().to_string()));
        if is_comms_tool(call.name) {
            let result = handle_tools_call(&self.tool_context, call.name, &args)
                .await
                .map_err(|e| ToolError::ExecutionFailed { message: e })?;
            Ok(ToolResult::new(call.id.to_string(), result.to_string(), false).into())
        } else {
            self.inner.dispatch(call).await
        }
    }
}

pub fn wrap_with_comms(
    tools: Arc<dyn AgentToolDispatcher>,
    runtime: Arc<CommsRuntime>,
) -> Arc<dyn AgentToolDispatcher> {
    let router = runtime.router_arc();
    let trusted_peers = runtime.trusted_peers_shared();
    Arc::new(DynCommsToolDispatcher::new_with_runtime(
        router,
        trusted_peers,
        runtime as Arc<dyn meerkat_core::agent::CommsRuntime>,
        tools,
    ))
}

#[cfg(test)]
mod tests {
    use super::*;

    #[test]
    fn comms_tool_defs_have_comms_provenance() {
        let defs = comms_tool_defs();
        assert!(!defs.is_empty(), "comms should expose at least one tool");
        for def in &defs {
            let prov = def
                .provenance
                .as_ref()
                .unwrap_or_else(|| panic!("comms tool '{}' is missing provenance", def.name));
            assert_eq!(
                prov.kind,
                ToolSourceKind::Comms,
                "comms tool '{}' should have Comms provenance",
                def.name
            );
            assert_eq!(prov.source_id, "comms");
        }
    }
}