use std::collections::HashMap;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use schemars::JsonSchema;
use serde::{Deserialize, Serialize};
use crate::tools::error::ToolError;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ToolDefinition {
pub name: String,
pub description: String,
pub parameters: serde_json::Value,
}
pub trait Tool: Clone + Send + Sync + 'static {
const NAME: &'static str;
type Args: for<'de> Deserialize<'de> + Send + JsonSchema;
type Output: Serialize;
type Error: std::error::Error + Send + Sync + 'static;
fn name(&self) -> &str {
Self::NAME
}
fn definition(&self, prompt: String) -> impl Future<Output = ToolDefinition> + Send;
fn call(
&self,
args: Self::Args,
) -> impl Future<Output = Result<Self::Output, Self::Error>> + Send;
}
pub trait ToolDyn: Send + Sync {
fn name(&self) -> &str;
fn definition(
&self,
prompt: String,
) -> Pin<Box<dyn Future<Output = ToolDefinition> + Send + '_>>;
fn call_json(
&self,
args: &str,
) -> Pin<Box<dyn Future<Output = Result<serde_json::Value, ToolError>> + Send + '_>>;
}
struct ToolWrapper<T: Tool> {
tool: T,
}
impl<T: Tool<Error = ToolError>> ToolDyn for ToolWrapper<T>
where
T::Output: 'static,
{
fn name(&self) -> &str {
self.tool.name()
}
fn definition(
&self,
prompt: String,
) -> Pin<Box<dyn Future<Output = ToolDefinition> + Send + '_>> {
Box::pin(async move { self.tool.definition(prompt).await })
}
fn call_json(
&self,
args: &str,
) -> Pin<Box<dyn Future<Output = Result<serde_json::Value, ToolError>> + Send + '_>> {
let args_str = args.to_string();
Box::pin(async move {
let parsed_args: T::Args = serde_json::from_str(&args_str).map_err(ToolError::Json)?;
let result = self
.tool
.call(parsed_args)
.await
.map_err(|e| ToolError::Validation(e.to_string()))?;
serde_json::to_value(result).map_err(ToolError::Json)
})
}
}
pub struct ToolRegistry {
tools: HashMap<String, Arc<dyn ToolDyn>>,
}
impl Default for ToolRegistry {
fn default() -> Self {
Self::new()
}
}
impl ToolRegistry {
pub fn new() -> Self {
Self {
tools: HashMap::new(),
}
}
pub fn register<T: Tool<Error = ToolError>>(&mut self, tool: T)
where
T::Output: 'static,
{
let name = tool.name().to_string();
self.tools.insert(name, Arc::new(ToolWrapper { tool }));
}
pub fn get(&self, name: &str) -> Option<Arc<dyn ToolDyn>> {
self.tools.get(name).cloned()
}
pub fn list(&self) -> Vec<String> {
self.tools.keys().cloned().collect()
}
pub fn exists(&self, name: &str) -> bool {
self.tools.contains_key(name)
}
pub async fn definition(&self, name: &str, prompt: &str) -> Option<ToolDefinition> {
if let Some(tool) = self.tools.get(name) {
Some(tool.definition(prompt.to_string()).await)
} else {
None
}
}
pub async fn all_definitions(&self, prompt: &str) -> Vec<ToolDefinition> {
let mut definitions = Vec::new();
for tool in self.tools.values() {
definitions.push(tool.definition(prompt.to_string()).await);
}
definitions
}
pub async fn execute(
&self,
name: &str,
args_json: &str,
) -> Result<serde_json::Value, ToolError> {
let tool = self
.tools
.get(name)
.ok_or_else(|| ToolError::Validation(format!("Tool not found: {}", name)))?;
tool.call_json(args_json).await
}
}