use std::{future::Future, sync::Arc};
use serde::{Deserialize, Serialize};
use crate::{
completion::{ToolDefinition, message::ToolName},
effect::{EffectKind, Outcome},
serve::{ErasedHandler, adapters::ToolCallback, adapters::ToolFn},
wasm_compat::{WasmBoxedFuture, WasmCompatSend, WasmCompatSync},
};
use super::{
IntoToolOutput, PublishedContext, ToolContext, ToolExecutionError, ToolOutput, ToolResult,
};
pub trait Tool: Sized + WasmCompatSend + WasmCompatSync {
const NAME: &'static str;
type Args: for<'de> Deserialize<'de> + WasmCompatSend + WasmCompatSync;
type Output: IntoToolOutput;
type Error: std::error::Error + WasmCompatSend + WasmCompatSync + 'static;
fn description(&self) -> String;
fn parameters(&self) -> serde_json::Value;
fn map_error(&self, error: Self::Error) -> ToolExecutionError {
ToolExecutionError::from_error(error)
}
fn call(
&self,
context: &mut ToolContext,
args: Self::Args,
) -> impl Future<Output = Result<Self::Output, Self::Error>> + WasmCompatSend;
}
impl<T> Tool for T
where
T: super::PortableTool,
{
const NAME: &'static str = <T as super::PortableTool>::NAME;
type Args = <T as super::PortableTool>::Args;
type Output = <T as super::PortableTool>::Output;
type Error = <T as super::PortableTool>::Error;
fn description(&self) -> String {
super::PortableTool::description(self)
}
fn parameters(&self) -> serde_json::Value {
super::PortableTool::parameters(self)
}
fn map_error(&self, error: Self::Error) -> ToolExecutionError {
super::PortableTool::map_error(self, error)
}
async fn call(
&self,
_context: &mut ToolContext,
args: Self::Args,
) -> Result<Self::Output, Self::Error> {
super::PortableTool::call(self, args).await
}
}
pub trait ToolEmbedding: Tool {
type InitError: std::error::Error + WasmCompatSend + WasmCompatSync + 'static;
type Context: for<'de> Deserialize<'de> + Serialize;
type State: WasmCompatSend;
fn embedding_docs(&self) -> Vec<String>;
fn context(&self) -> Self::Context;
fn init(state: Self::State, context: Self::Context) -> Result<Self, Self::InitError>;
}
impl<T> ToolEmbedding for T
where
T: super::PortableToolEmbedding,
{
type InitError = <T as super::PortableToolEmbedding>::InitError;
type Context = <T as super::PortableToolEmbedding>::Context;
type State = <T as super::PortableToolEmbedding>::State;
fn embedding_docs(&self) -> Vec<String> {
super::PortableToolEmbedding::embedding_docs(self)
}
fn context(&self) -> Self::Context {
super::PortableToolEmbedding::context(self)
}
fn init(state: Self::State, context: Self::Context) -> Result<Self, Self::InitError> {
super::PortableToolEmbedding::init(state, context)
}
}
fn parse_tool_args<A>(args: &str) -> Result<A, ToolExecutionError>
where
A: serde::de::DeserializeOwned,
{
match serde_json::from_str(args) {
Ok(parsed) => Ok(parsed),
Err(original) if args.trim() == "null" => serde_json::from_str("{}").map_err(|_| {
ToolExecutionError::invalid_args(format!("failed to parse tool arguments: {original}"))
.with_source(original)
}),
Err(error) => Err(ToolExecutionError::invalid_args(format!(
"failed to parse tool arguments: {error}"
))
.with_source(error)),
}
}
pub(crate) async fn execute_callback<F>(
callback: &F,
args: String,
context: &mut ToolContext,
) -> ToolResult
where
F: for<'a> Fn(
&'a mut ToolContext,
serde_json::Value,
) -> WasmBoxedFuture<'a, Result<ToolOutput, ToolExecutionError>>,
{
let args = match parse_tool_args::<serde_json::Value>(&args) {
Ok(args) => args,
Err(error) => return ToolResult::failed(error),
};
tool_result_from(callback(context, args).await)
}
fn tool_result_from<O>(outcome: Result<O, ToolExecutionError>) -> ToolResult
where
O: IntoToolOutput,
{
match outcome.and_then(IntoToolOutput::into_tool_output) {
Ok(output) => ToolResult::success(output),
Err(error) => ToolResult::failed(error),
}
}
pub trait ErasedTool: WasmCompatSend + WasmCompatSync {
fn name(&self) -> String;
fn description(&self) -> String;
fn parameters(&self) -> serde_json::Value;
fn execute<'a>(
&'a self,
args: String,
context: &'a mut ToolContext,
) -> WasmBoxedFuture<'a, ToolResult>;
}
impl<T> ErasedTool for T
where
T: Tool,
{
fn name(&self) -> String {
T::NAME.to_string()
}
fn description(&self) -> String {
Tool::description(self)
}
fn parameters(&self) -> serde_json::Value {
Tool::parameters(self)
}
fn execute<'a>(
&'a self,
args: String,
context: &'a mut ToolContext,
) -> WasmBoxedFuture<'a, ToolResult> {
Box::pin(async move {
let args = match parse_tool_args::<T::Args>(&args) {
Ok(args) => args,
Err(error) => return ToolResult::failed(error),
};
tool_result_from(
Tool::call(self, context, args)
.await
.map_err(|error| Tool::map_error(self, error)),
)
})
}
}
#[cfg(not(target_family = "wasm"))]
pub type LivenessFn = Arc<dyn Fn() -> bool + Send + Sync>;
#[cfg(target_family = "wasm")]
pub type LivenessFn = Arc<dyn Fn() -> bool>;
#[derive(Clone)]
pub struct DynamicTool {
definition: ToolDefinition,
handler: ErasedHandler,
liveness: Option<LivenessFn>,
}
impl DynamicTool {
pub fn new<F>(
name: ToolName,
description: impl Into<String>,
parameters: serde_json::Value,
callback: F,
) -> Self
where
F: Fn(
serde_json::Value,
) -> WasmBoxedFuture<'static, Result<ToolOutput, ToolExecutionError>>
+ WasmCompatSend
+ WasmCompatSync
+ 'static,
{
Self::new_with_context(
name,
description,
parameters,
move |_context: &mut ToolContext, arguments| callback(arguments),
)
}
pub fn new_with_context<F>(
name: ToolName,
description: impl Into<String>,
parameters: serde_json::Value,
callback: F,
) -> Self
where
F: ToolCallback + 'static,
{
let description = description.into();
let handler = ErasedHandler::new(ToolFn::new(
name.to_string(),
description.clone(),
parameters.clone(),
callback,
));
Self {
definition: ToolDefinition {
name,
description,
parameters,
},
handler,
liveness: None,
}
}
pub fn with_liveness<F>(mut self, is_live: F) -> Self
where
F: Fn() -> bool + WasmCompatSend + WasmCompatSync + 'static,
{
self.liveness = Some(Arc::new(is_live));
self
}
pub fn name(&self) -> &ToolName {
&self.definition.name
}
pub fn definition(&self) -> ToolDefinition {
self.definition.clone()
}
pub fn handler(&self) -> &ErasedHandler {
&self.handler
}
pub fn into_parts(self) -> (ToolDefinition, ErasedHandler, Option<LivenessFn>) {
(self.definition, self.handler, self.liveness)
}
pub fn is_live(&self) -> bool {
self.liveness.as_ref().is_none_or(|probe| probe())
}
pub async fn execute(
&self,
arguments: serde_json::Value,
) -> Result<ToolOutput, ToolExecutionError> {
let mut context = ToolContext::new();
self.execute_with(&mut context, arguments).await
}
pub async fn execute_with(
&self,
context: &mut ToolContext,
arguments: serde_json::Value,
) -> Result<ToolOutput, ToolExecutionError> {
let published = PublishedContext::new();
let outcome = crate::serve::serve_inline_with(
&self.handler,
EffectKind::ToolCall {
name: self.definition.name.to_string(),
args: arguments.to_string(),
},
vec![
Arc::new(context.for_dispatch()),
published.clone() as Arc<dyn std::any::Any + Send + Sync>,
],
)
.await;
match outcome {
Ok(Outcome::ToolResult { result }) => {
context.accept_dispatch_result(published.take().unwrap_or_default());
result.into_result()
}
Ok(other) => Err(ToolExecutionError::other(format!(
"tool handler answered with a {} outcome",
other.family()
))),
Err(report) => Err(report.into()),
}
}
}
impl std::fmt::Debug for DynamicTool {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("DynamicTool")
.field("name", &self.definition.name)
.finish_non_exhaustive()
}
}
pub fn tool_name<T: Tool>() -> ToolName {
const { assert!(!T::NAME.is_empty(), "Tool::NAME cannot be empty") };
ToolName::new_unchecked(T::NAME)
}
pub fn tool_definition<T: Tool>(tool: &T) -> ToolDefinition {
ToolDefinition {
name: tool_name::<T>(),
description: tool.description(),
parameters: tool.parameters(),
}
}