use std::collections::HashMap;
use std::sync::Arc;
use serde_json::Value;
use thiserror::Error;
use super::AgentTool;
#[derive(Debug, Error)]
pub enum ToolError {
#[error("tool not found: {0}")]
ToolNotFound(String),
#[error("invalid input for tool: {0}")]
InvalidInput(String),
#[error("tool execution error: {0}")]
ExecutionError(String),
}
#[derive(Debug, Clone, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub struct ToolContributionSource(String);
impl ToolContributionSource {
pub fn new(source: impl Into<String>) -> Self {
Self(source.into())
}
#[must_use]
pub fn as_str(&self) -> &str {
&self.0
}
}
impl std::fmt::Display for ToolContributionSource {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.write_str(&self.0)
}
}
#[derive(Clone)]
pub struct ToolContribution {
source: ToolContributionSource,
tool: Arc<dyn AgentTool>,
}
impl ToolContribution {
pub fn new(source: ToolContributionSource, tool: Arc<dyn AgentTool>) -> Self {
Self { source, tool }
}
#[must_use]
pub fn source(&self) -> &ToolContributionSource {
&self.source
}
#[must_use]
pub fn name(&self) -> &str {
self.tool.name()
}
#[must_use]
pub fn tool(&self) -> &Arc<dyn AgentTool> {
&self.tool
}
#[must_use]
pub fn map_tool(mut self, wrap: impl FnOnce(Arc<dyn AgentTool>) -> Arc<dyn AgentTool>) -> Self {
self.tool = wrap(self.tool);
self
}
}
#[derive(Debug, Error, Clone, PartialEq, Eq)]
#[error(
"duplicate tool registration '{tool_name}': existing source '{existing_source}', incoming source '{incoming_source}'"
)]
pub struct ToolRegistrationError {
pub tool_name: String,
pub existing_source: ToolContributionSource,
pub incoming_source: ToolContributionSource,
}
const LEGACY_TOOL_REGISTRATION_SOURCE: &str = "legacy:unchecked";
#[derive(Default)]
pub struct ToolRegistry {
tools: HashMap<String, Arc<dyn AgentTool>>,
contribution_sources: HashMap<String, ToolContributionSource>,
}
impl ToolRegistry {
pub fn new() -> Self {
Self::default()
}
pub fn register(&mut self, tool: Arc<dyn AgentTool>) {
let name = tool.name().to_owned();
self.contribution_sources.remove(&name);
self.tools.insert(name, tool);
}
pub fn register_contribution(
&mut self,
contribution: ToolContribution,
) -> Result<(), ToolRegistrationError> {
let ToolContribution { source, tool } = contribution;
let tool_name = tool.name().to_owned();
if let Some(existing_source) = self.registered_source(&tool_name) {
return Err(ToolRegistrationError {
tool_name,
existing_source,
incoming_source: source,
});
}
self.contribution_sources.insert(tool_name.clone(), source);
self.tools.insert(tool_name, tool);
Ok(())
}
pub fn register_contributions(
&mut self,
contributions: impl IntoIterator<Item = ToolContribution>,
) -> Result<(), ToolRegistrationError> {
let contributions = contributions.into_iter().collect::<Vec<_>>();
let mut pending_sources = HashMap::<String, ToolContributionSource>::new();
for contribution in &contributions {
let tool_name = contribution.name().to_owned();
if let Some(existing_source) = self.registered_source(&tool_name) {
return Err(ToolRegistrationError {
tool_name,
existing_source,
incoming_source: contribution.source().clone(),
});
}
if let Some(existing_source) = pending_sources.get(&tool_name) {
return Err(ToolRegistrationError {
tool_name,
existing_source: existing_source.clone(),
incoming_source: contribution.source().clone(),
});
}
pending_sources.insert(tool_name, contribution.source().clone());
}
for ToolContribution { source, tool } in contributions {
let tool_name = tool.name().to_owned();
self.contribution_sources.insert(tool_name.clone(), source);
self.tools.insert(tool_name, tool);
}
Ok(())
}
fn registered_source(&self, tool_name: &str) -> Option<ToolContributionSource> {
self.tools.contains_key(tool_name).then(|| {
self.contribution_sources
.get(tool_name)
.cloned()
.unwrap_or_else(|| ToolContributionSource::new(LEGACY_TOOL_REGISTRATION_SOURCE))
})
}
pub fn get(&self, name: &str) -> Option<&dyn AgentTool> {
self.tools.get(name).map(|t| t.as_ref())
}
pub fn list(&self) -> Vec<&dyn AgentTool> {
self.tools.values().map(|t| t.as_ref()).collect()
}
pub fn validate_input(&self, name: &str, input: &Value) -> Result<(), ToolError> {
let tool = self
.get(name)
.ok_or_else(|| ToolError::ToolNotFound(name.to_owned()))?;
let params = tool.parameters();
if !input.is_object() {
return Err(ToolError::InvalidInput(format!(
"expected object for tool '{name}', got {}",
input_type_name(input)
)));
}
if let Some(schema_obj) = params.as_object()
&& let Some(Value::Array(required)) = schema_obj.get("required")
&& let Some(input_obj) = input.as_object()
{
for req in required {
if let Some(req_key) = req.as_str()
&& !input_obj.contains_key(req_key)
{
return Err(ToolError::InvalidInput(format!(
"missing required field '{req_key}' for tool '{name}'"
)));
}
}
}
Ok(())
}
}
fn input_type_name(value: &Value) -> &'static str {
match value {
Value::Null => "null",
Value::Bool(_) => "boolean",
Value::Number(_) => "number",
Value::String(_) => "string",
Value::Array(_) => "array",
Value::Object(_) => "object",
}
}