use crate::traits::{AgentError, Result, ToolDefinition};
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::RwLock;
pub type ToolCallback = Arc<
dyn Fn(serde_json::Value) -> futures::future::BoxFuture<'static, Result<String>> + Send + Sync,
>;
pub struct RegisteredTool {
pub definition: ToolDefinition,
pub callback: ToolCallback,
pub source: ToolSource,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum ToolSource {
BuiltIn,
Mcp { server_name: String },
Custom { source: String },
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct McpServerConfig {
pub name: String,
pub connection: McpConnection,
pub env: Option<HashMap<String, String>>,
pub timeout_secs: Option<u64>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "type")]
pub enum McpConnection {
#[serde(rename = "http")]
Http { url: String },
#[serde(rename = "websocket")]
WebSocket { url: String },
#[serde(rename = "process")]
Process {
command: String,
args: Option<Vec<String>>,
},
}
pub struct ToolRegistry {
tools: RwLock<HashMap<String, RegisteredTool>>,
mcp_servers: RwLock<HashMap<String, McpServerConfig>>,
}
impl ToolRegistry {
pub fn new() -> Self {
Self {
tools: RwLock::new(HashMap::new()),
mcp_servers: RwLock::new(HashMap::new()),
}
}
pub fn with_defaults() -> Self {
let registry = Self::new();
registry
}
pub async fn register(
&self,
definition: ToolDefinition,
callback: ToolCallback,
source: ToolSource,
) {
let name = definition.name.clone();
let tool = RegisteredTool {
definition,
callback,
source,
};
self.tools.write().await.insert(name, tool);
}
pub async fn register_mcp_server(&self, config: McpServerConfig) -> Result<Vec<String>> {
let server_name = config.name.clone();
self.mcp_servers
.write()
.await
.insert(server_name.clone(), config.clone());
let tool_names = self.discover_mcp_tools(&config).await?;
Ok(tool_names)
}
async fn discover_mcp_tools(&self, config: &McpServerConfig) -> Result<Vec<String>> {
use hanzo_mcp::mcp_methods;
let tools = match &config.connection {
McpConnection::Http { url } => mcp_methods::list_tools_via_http(url, None)
.await
.map_err(|e| AgentError::McpError(e.message))?,
McpConnection::Process { command, args } => {
let cmd_str = if let Some(args) = args {
format!("{} {}", command, args.join(" "))
} else {
command.clone()
};
mcp_methods::list_tools_via_command(&cmd_str, config.env.clone())
.await
.map_err(|e| AgentError::McpError(e.message))?
}
McpConnection::WebSocket { url: _ } => {
return Err(AgentError::ConfigError(
"WebSocket MCP not yet implemented".to_string(),
));
}
};
let mut registered_names = Vec::new();
let server_name = config.name.clone();
for tool in tools {
let prefixed_name = format!("{}:{}", server_name, tool.name);
let definition = ToolDefinition {
name: prefixed_name.clone(),
description: tool.description.unwrap_or_default().to_string(),
parameters: serde_json::Value::Object((*tool.input_schema).clone()),
requires_confirmation: false,
};
let config_clone = config.clone();
let tool_name = tool.name.clone();
let callback: ToolCallback = Arc::new(move |args: serde_json::Value| {
let config = config_clone.clone();
let name = tool_name.clone();
Box::pin(async move {
let params = args.as_object().cloned().unwrap_or_default();
execute_mcp_tool(&config, &name, params).await
})
});
self.register(
definition,
callback,
ToolSource::Mcp {
server_name: server_name.clone(),
},
)
.await;
registered_names.push(prefixed_name);
}
Ok(registered_names)
}
pub async fn get(&self, name: &str) -> Option<ToolDefinition> {
self.tools
.read()
.await
.get(name)
.map(|t| t.definition.clone())
}
pub async fn execute(&self, name: &str, args: serde_json::Value) -> Result<String> {
let tools = self.tools.read().await;
let tool = tools.get(name).ok_or_else(|| AgentError::ToolError {
tool_name: name.to_string(),
message: "Tool not found".to_string(),
})?;
(tool.callback)(args).await
}
pub async fn list(&self) -> Vec<ToolDefinition> {
self.tools
.read()
.await
.values()
.map(|t| t.definition.clone())
.collect()
}
pub async fn list_by_source(&self, source_filter: &str) -> Vec<ToolDefinition> {
self.tools
.read()
.await
.values()
.filter(|t| match &t.source {
ToolSource::BuiltIn => source_filter == "builtin",
ToolSource::Mcp { server_name } => server_name == source_filter,
ToolSource::Custom { source } => source == source_filter,
})
.map(|t| t.definition.clone())
.collect()
}
}
impl Default for ToolRegistry {
fn default() -> Self {
Self::new()
}
}
async fn execute_mcp_tool(
config: &McpServerConfig,
tool_name: &str,
params: serde_json::Map<String, serde_json::Value>,
) -> Result<String> {
use hanzo_mcp::mcp_methods;
let result = match &config.connection {
McpConnection::Http { url } => {
mcp_methods::run_tool_via_http(url.clone(), tool_name.to_string(), params)
.await
.map_err(|e| AgentError::McpError(e.message))?
}
McpConnection::Process { command, args } => {
let cmd_str = if let Some(args) = args {
format!("{} {}", command, args.join(" "))
} else {
command.clone()
};
mcp_methods::run_tool_via_command(
cmd_str,
tool_name.to_string(),
config.env.clone().unwrap_or_default(),
params,
)
.await
.map_err(|e| AgentError::McpError(e.message))?
}
McpConnection::WebSocket { url: _ } => {
return Err(AgentError::ConfigError(
"WebSocket MCP not yet implemented".to_string(),
));
}
};
let text = result
.content
.into_iter()
.filter_map(|c| c.as_text().map(|t| t.text.clone()))
.collect::<Vec<_>>()
.join("\n");
Ok(text)
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_registry_creation() {
let registry = ToolRegistry::new();
let tools = registry.list().await;
assert!(tools.is_empty());
}
#[tokio::test]
async fn test_custom_tool_registration() {
let registry = ToolRegistry::new();
let definition = ToolDefinition::new("test_tool", "A test tool");
let callback: ToolCallback =
Arc::new(|_args| Box::pin(async { Ok("test result".to_string()) }));
registry
.register(
definition,
callback,
ToolSource::Custom {
source: "test".to_string(),
},
)
.await;
let tools = registry.list().await;
assert_eq!(tools.len(), 1);
assert_eq!(tools[0].name, "test_tool");
let result = registry
.execute("test_tool", serde_json::json!({}))
.await
.unwrap();
assert_eq!(result, "test result");
}
}