use std::collections::HashMap;
use std::sync::Arc;
use tracing::{debug, info, warn};
use crate::middleware::ToolMiddleware;
use crate::normalizer;
use crate::traits::{BaseTool, ToolContext, ToolResult};
use crate::validation;
use super::ToolRegistry;
use super::helpers::{camel_to_snake_name, edit_distance, make_dedup_key};
impl ToolRegistry {
fn suggest_tool_names(&self, name: &str) -> Vec<String> {
let tools = self.tools.read().expect("ToolRegistry lock poisoned");
let lower = name.to_lowercase();
let mut suggestions: Vec<String> = Vec::new();
for registered in tools.keys() {
let reg_lower = registered.to_lowercase();
if reg_lower.contains(&lower)
|| lower.contains(®_lower)
|| edit_distance(&lower, ®_lower) <= 3
{
suggestions.push(registered.clone());
}
}
suggestions.sort();
suggestions.truncate(5);
suggestions
}
fn resolve_tool(&self, name: &str) -> Option<(Arc<dyn BaseTool>, String)> {
let tools = self.tools.read().expect("ToolRegistry lock poisoned");
let name = name.strip_prefix("functions.").unwrap_or(name);
if let Some(t) = tools.get(name) {
return Some((Arc::clone(t), name.to_string()));
}
let lower = name.to_lowercase();
for (registered_name, tool) in tools.iter() {
if registered_name.to_lowercase() == lower {
info!(
requested = %name,
resolved = %registered_name,
"Fuzzy tool name match (case-insensitive)"
);
return Some((Arc::clone(tool), registered_name.clone()));
}
}
let snake = camel_to_snake_name(name);
if snake != name
&& let Some(t) = tools.get(&snake)
{
info!(
requested = %name,
resolved = %snake,
"Fuzzy tool name match (camelCase -> snake_case)"
);
return Some((Arc::clone(t), snake));
}
None
}
pub async fn execute(
&self,
tool_name: &str,
args: HashMap<String, serde_json::Value>,
ctx: &ToolContext,
) -> ToolResult {
let (tool, resolved_name) = match self.resolve_tool(tool_name) {
Some((t, name)) => (t, name),
None => {
let suggestions = self.suggest_tool_names(tool_name);
let hint = if suggestions.is_empty() {
String::new()
} else {
format!(". Did you mean: {}?", suggestions.join(", "))
};
warn!(tool = %tool_name, "Unknown tool");
return ToolResult::fail(format!("Unknown tool: {tool_name}{hint}"));
}
};
let tool_name = &resolved_name;
let working_dir = ctx.working_dir.to_string_lossy().to_string();
let normalized = normalizer::normalize_params(tool_name, args, Some(&working_dir));
debug!(tool = %tool_name, params = ?normalized, "Normalized tool params");
const NO_DEDUP: &[&str] = &["spawn_subagent"];
let skip_dedup = NO_DEDUP.contains(&tool_name.as_str());
let dedup_key = make_dedup_key(tool_name, &normalized);
if !skip_dedup
&& let Ok(cache) = self.dedup_cache.lock()
&& let Some(cached) = cache.get(&dedup_key)
{
info!(tool = %tool_name, "Returning cached result (dedup)");
return cached.clone();
}
let schema = tool.parameter_schema();
let validation_errors = validation::validate_args_detailed(&normalized, &schema);
if !validation_errors.is_empty() {
let error_msg = tool
.format_validation_error(&validation_errors)
.unwrap_or_else(|| {
let details: Vec<String> =
validation_errors.iter().map(|e| e.to_string()).collect();
format!(
"The {} tool was called with invalid arguments:\n - {}\nPlease fix the arguments and try again.",
tool_name,
details.join("\n - ")
)
});
warn!(tool = %tool_name, error = %error_msg, "Parameter validation failed");
return ToolResult::fail(error_msg);
}
let middleware: Vec<Arc<dyn ToolMiddleware>> = {
let mw = self.middleware.read().expect("ToolRegistry lock poisoned");
mw.clone()
};
for mw in &middleware {
if let Err(err) = mw.before_execute(tool_name, &normalized, ctx).await {
warn!(tool = %tool_name, error = %err, "Middleware rejected execution");
return ToolResult::fail(format!("Middleware error: {err}"));
}
}
let exec_ctx = {
let timeouts = self
.tool_timeouts
.read()
.expect("ToolRegistry lock poisoned");
if let Some(timeout_config) = timeouts.get(tool_name) {
let mut new_ctx = ctx.clone();
new_ctx.timeout_config = Some(timeout_config.clone());
new_ctx
} else {
ctx.clone()
}
};
let start = std::time::Instant::now();
let mut result = tool.execute(normalized, &exec_ctx).await;
result.duration_ms = Some(start.elapsed().as_millis() as u64);
let sanitized = self.sanitizer.sanitize_with_mcp_fallback(
tool_name,
result.success,
result.output.as_deref(),
result.error.as_deref(),
);
if sanitized.was_truncated {
result.output = sanitized.output;
result.error = sanitized.error;
}
for mw in &middleware {
if let Err(err) = mw.after_execute(tool_name, &result).await {
warn!(tool = %tool_name, error = %err, "Middleware after_execute error");
}
}
if !skip_dedup && let Ok(mut cache) = self.dedup_cache.lock() {
cache.insert(dedup_key, result.clone());
}
result
}
}