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;
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"
)
}
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");
}
}
}