use std::collections::{btree_map::Entry, BTreeMap};
use async_trait::async_trait;
use openai_client_base::models::{ChatCompletionMessageToolCallsInner, ChatCompletionTool};
use serde::{de::DeserializeOwned, Serialize};
use serde_json::Value;
use crate::{builders::chat::tool_function, responses::chat::ToolCallExt, Error, Result};
#[async_trait]
pub trait FunctionTool: Send + Sync {
type Input: DeserializeOwned + Send;
type Output: Serialize + Send;
fn name(&self) -> &str;
fn description(&self) -> &str;
fn parameters_schema(&self) -> Value;
async fn execute(&self, input: Self::Input) -> Result<Self::Output>;
}
#[derive(Debug, thiserror::Error)]
pub enum ToolError {
#[error("Tool {name:?} is already registered")]
Duplicate {
name: String,
},
#[error("Invalid definition for tool {name:?}: {reason}")]
InvalidDefinition {
name: String,
reason: String,
},
#[error("Unknown tool {name:?}")]
Unknown {
name: String,
},
#[error("Tool {name:?} argument decoding failed: {source}")]
Arguments {
name: String,
#[source]
source: serde_json::Error,
},
#[error("Tool {name:?} output encoding failed: {source}")]
Output {
name: String,
#[source]
source: serde_json::Error,
},
#[error("Tool {name:?} arguments must be a JSON object")]
InvalidArguments {
name: String,
},
#[error("Tool {name:?} execution failed: {source}")]
Execution {
name: String,
#[source]
source: Error,
},
#[error("Invalid tool call: {reason}")]
InvalidCall {
reason: String,
},
#[error("Tool call {call_id:?} failed: {source}")]
Call {
call_id: String,
#[source]
source: Box<Self>,
},
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ToolOutput {
pub call_id: String,
pub content: String,
}
#[async_trait]
trait ErasedTool: Send + Sync {
async fn execute(&self, name: &str, args: &str) -> std::result::Result<Value, ToolError>;
}
#[async_trait]
impl<T: FunctionTool> ErasedTool for T {
async fn execute(&self, name: &str, args: &str) -> std::result::Result<Value, ToolError> {
let argument_error = |source| ToolError::Arguments {
name: name.into(),
source,
};
let value: Value = serde_json::from_str(args).map_err(argument_error)?;
if !value.is_object() {
return Err(ToolError::InvalidArguments { name: name.into() });
}
let input = serde_json::from_str(args).map_err(argument_error)?;
let output =
FunctionTool::execute(self, input)
.await
.map_err(|source| ToolError::Execution {
name: name.into(),
source,
})?;
serde_json::to_value(output).map_err(|source| ToolError::Output {
name: name.into(),
source,
})
}
}
struct RegisteredTool {
definition: ChatCompletionTool,
handler: Box<dyn ErasedTool>,
}
#[derive(Default)]
pub struct ToolRegistry {
tools: BTreeMap<String, RegisteredTool>,
}
impl std::fmt::Debug for ToolRegistry {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ToolRegistry")
.field("names", &self.tools.keys())
.finish_non_exhaustive()
}
}
impl ToolRegistry {
#[must_use]
pub fn new() -> Self {
Self::default()
}
pub fn register<T: FunctionTool + 'static>(
&mut self,
tool: T,
) -> std::result::Result<(), ToolError> {
let name = tool.name().to_owned();
let invalid = |reason: &str| ToolError::InvalidDefinition {
name: name.clone(),
reason: reason.into(),
};
if name.is_empty()
|| name.len() > 64
|| !name
.bytes()
.all(|b| b.is_ascii_alphanumeric() || b == b'_' || b == b'-')
{
return Err(invalid(
"name must contain 1–64 ASCII letters, digits, underscores or hyphens",
));
}
let description = tool.description();
if description.trim().is_empty() {
return Err(invalid("description must not be blank"));
}
let schema = tool.parameters_schema();
if !schema.is_object() || schema.get("type").and_then(Value::as_str) != Some("object") {
return Err(invalid("parameters schema must have root type object"));
}
let definition = tool_function(&name, description, schema);
match self.tools.entry(name) {
Entry::Occupied(entry) => Err(ToolError::Duplicate {
name: entry.key().clone(),
}),
Entry::Vacant(entry) => {
entry.insert(RegisteredTool {
definition,
handler: Box::new(tool),
});
Ok(())
}
}
}
#[must_use]
pub fn tool_definitions(&self) -> Vec<ChatCompletionTool> {
self.tools
.values()
.map(|tool| tool.definition.clone())
.collect()
}
pub async fn execute(
&self,
name: &str,
arguments: &str,
) -> std::result::Result<Value, ToolError> {
let tool = self
.tools
.get(name)
.ok_or_else(|| ToolError::Unknown { name: name.into() })?;
tool.handler.execute(name, arguments).await
}
pub async fn execute_call(
&self,
call: &ChatCompletionMessageToolCallsInner,
) -> std::result::Result<ToolOutput, ToolError> {
let result = async {
if call.id().is_empty() {
return Err(ToolError::InvalidCall {
reason: "call identifier must not be empty".into(),
});
}
let ChatCompletionMessageToolCallsInner::ChatCompletionMessageToolCall(function_call) =
call
else {
return Err(ToolError::InvalidCall {
reason: "only function tool calls are supported".into(),
});
};
let output = self
.execute(
&function_call.function.name,
&function_call.function.arguments,
)
.await?;
Ok(ToolOutput {
call_id: call.id().into(),
content: output.to_string(),
})
}
.await;
result.map_err(|source| ToolError::Call {
call_id: call.id().into(),
source: Box::new(source),
})
}
}