1use std::{future::Future, sync::Arc};
14
15use serde::{Deserialize, Serialize};
16
17use crate::{
18 completion::ToolDefinition,
19 effect::{EffectKind, Outcome},
20 serve::{ErasedHandler, adapters::ToolCallback, adapters::ToolFn},
21 wasm_compat::{WasmBoxedFuture, WasmCompatSend, WasmCompatSync},
22};
23
24use super::{
25 IntoToolOutput, PublishedContext, ToolContext, ToolExecutionError, ToolOutput, ToolResult,
26};
27
28pub trait Tool: Sized + WasmCompatSend + WasmCompatSync {
35 const NAME: &'static str;
37 type Args: for<'de> Deserialize<'de> + WasmCompatSend + WasmCompatSync;
39 type Output: IntoToolOutput;
47 type Error: std::error::Error + WasmCompatSend + WasmCompatSync + 'static;
54
55 fn description(&self) -> String;
57
58 fn parameters(&self) -> serde_json::Value;
60
61 fn map_error(&self, error: Self::Error) -> ToolExecutionError {
67 ToolExecutionError::from_error(error)
68 }
69
70 fn call(
72 &self,
73 context: &mut ToolContext,
74 args: Self::Args,
75 ) -> impl Future<Output = Result<Self::Output, Self::Error>> + WasmCompatSend;
76}
77
78impl<T> Tool for T
79where
80 T: super::PortableTool,
81{
82 const NAME: &'static str = <T as super::PortableTool>::NAME;
83 type Args = <T as super::PortableTool>::Args;
84 type Output = <T as super::PortableTool>::Output;
85 type Error = <T as super::PortableTool>::Error;
86
87 fn description(&self) -> String {
88 super::PortableTool::description(self)
89 }
90
91 fn parameters(&self) -> serde_json::Value {
92 super::PortableTool::parameters(self)
93 }
94
95 fn map_error(&self, error: Self::Error) -> ToolExecutionError {
96 super::PortableTool::map_error(self, error)
97 }
98
99 async fn call(
100 &self,
101 _context: &mut ToolContext,
102 args: Self::Args,
103 ) -> Result<Self::Output, Self::Error> {
104 super::PortableTool::call(self, args).await
105 }
106}
107
108pub trait ToolEmbedding: Tool {
110 type InitError: std::error::Error + WasmCompatSend + WasmCompatSync + 'static;
112 type Context: for<'de> Deserialize<'de> + Serialize;
114 type State: WasmCompatSend;
116
117 fn embedding_docs(&self) -> Vec<String>;
119 fn context(&self) -> Self::Context;
121 fn init(state: Self::State, context: Self::Context) -> Result<Self, Self::InitError>;
123}
124
125impl<T> ToolEmbedding for T
126where
127 T: super::PortableToolEmbedding,
128{
129 type InitError = <T as super::PortableToolEmbedding>::InitError;
130 type Context = <T as super::PortableToolEmbedding>::Context;
131 type State = <T as super::PortableToolEmbedding>::State;
132
133 fn embedding_docs(&self) -> Vec<String> {
134 super::PortableToolEmbedding::embedding_docs(self)
135 }
136
137 fn context(&self) -> Self::Context {
138 super::PortableToolEmbedding::context(self)
139 }
140
141 fn init(state: Self::State, context: Self::Context) -> Result<Self, Self::InitError> {
142 super::PortableToolEmbedding::init(state, context)
143 }
144}
145
146fn parse_tool_args<A>(args: &str) -> Result<A, ToolExecutionError>
147where
148 A: serde::de::DeserializeOwned,
149{
150 match serde_json::from_str(args) {
151 Ok(parsed) => Ok(parsed),
152 Err(original) if args.trim() == "null" => serde_json::from_str("{}").map_err(|_| {
153 ToolExecutionError::invalid_args(format!("failed to parse tool arguments: {original}"))
154 .with_source(original)
155 }),
156 Err(error) => Err(ToolExecutionError::invalid_args(format!(
157 "failed to parse tool arguments: {error}"
158 ))
159 .with_source(error)),
160 }
161}
162
163pub(crate) async fn execute_callback<F>(
166 callback: &F,
167 args: String,
168 context: &mut ToolContext,
169) -> ToolResult
170where
171 F: for<'a> Fn(
172 &'a mut ToolContext,
173 serde_json::Value,
174 ) -> WasmBoxedFuture<'a, Result<ToolOutput, ToolExecutionError>>,
175{
176 let args = match parse_tool_args::<serde_json::Value>(&args) {
177 Ok(args) => args,
178 Err(error) => return ToolResult::failed(error),
179 };
180 tool_result_from(callback(context, args).await)
181}
182
183fn tool_result_from<O>(outcome: Result<O, ToolExecutionError>) -> ToolResult
184where
185 O: IntoToolOutput,
186{
187 match outcome.and_then(IntoToolOutput::into_tool_output) {
188 Ok(output) => ToolResult::success(output),
189 Err(error) => ToolResult::failed(error),
190 }
191}
192
193pub trait ErasedTool: WasmCompatSend + WasmCompatSync {
197 fn name(&self) -> String;
199 fn description(&self) -> String;
201 fn parameters(&self) -> serde_json::Value;
203 fn execute<'a>(
205 &'a self,
206 args: String,
207 context: &'a mut ToolContext,
208 ) -> WasmBoxedFuture<'a, ToolResult>;
209}
210
211impl<T> ErasedTool for T
212where
213 T: Tool,
214{
215 fn name(&self) -> String {
216 T::NAME.to_string()
217 }
218
219 fn description(&self) -> String {
220 Tool::description(self)
221 }
222
223 fn parameters(&self) -> serde_json::Value {
224 Tool::parameters(self)
225 }
226
227 fn execute<'a>(
228 &'a self,
229 args: String,
230 context: &'a mut ToolContext,
231 ) -> WasmBoxedFuture<'a, ToolResult> {
232 Box::pin(async move {
233 let args = match parse_tool_args::<T::Args>(&args) {
234 Ok(args) => args,
235 Err(error) => return ToolResult::failed(error),
236 };
237 tool_result_from(
238 Tool::call(self, context, args)
239 .await
240 .map_err(|error| Tool::map_error(self, error)),
241 )
242 })
243 }
244}
245
246#[cfg(not(target_family = "wasm"))]
250pub type LivenessFn = Arc<dyn Fn() -> bool + Send + Sync>;
251#[cfg(target_family = "wasm")]
253pub type LivenessFn = Arc<dyn Fn() -> bool>;
254
255#[derive(Clone)]
261pub struct DynamicTool {
262 definition: ToolDefinition,
263 handler: ErasedHandler,
264 liveness: Option<LivenessFn>,
265}
266
267impl DynamicTool {
268 pub fn new<F>(
270 name: impl Into<String>,
271 description: impl Into<String>,
272 parameters: serde_json::Value,
273 callback: F,
274 ) -> Self
275 where
276 F: Fn(
277 serde_json::Value,
278 ) -> WasmBoxedFuture<'static, Result<ToolOutput, ToolExecutionError>>
279 + WasmCompatSend
280 + WasmCompatSync
281 + 'static,
282 {
283 Self::new_with_context(
284 name,
285 description,
286 parameters,
287 move |_context: &mut ToolContext, arguments| callback(arguments),
288 )
289 }
290
291 pub fn new_with_context<F>(
293 name: impl Into<String>,
294 description: impl Into<String>,
295 parameters: serde_json::Value,
296 callback: F,
297 ) -> Self
298 where
299 F: ToolCallback + 'static,
300 {
301 let name = name.into();
302 let description = description.into();
303 let handler = ErasedHandler::new(ToolFn::new(
304 name.clone(),
305 description.clone(),
306 parameters.clone(),
307 callback,
308 ));
309 Self {
310 definition: ToolDefinition {
311 name,
312 description,
313 parameters,
314 },
315 handler,
316 liveness: None,
317 }
318 }
319
320 pub fn with_liveness<F>(mut self, is_live: F) -> Self
322 where
323 F: Fn() -> bool + WasmCompatSend + WasmCompatSync + 'static,
324 {
325 self.liveness = Some(Arc::new(is_live));
326 self
327 }
328
329 pub fn name(&self) -> &str {
331 &self.definition.name
332 }
333
334 pub fn definition(&self) -> ToolDefinition {
336 self.definition.clone()
337 }
338
339 pub fn handler(&self) -> &ErasedHandler {
341 &self.handler
342 }
343
344 pub fn into_parts(self) -> (ToolDefinition, ErasedHandler, Option<LivenessFn>) {
346 (self.definition, self.handler, self.liveness)
347 }
348
349 pub fn is_live(&self) -> bool {
351 self.liveness.as_ref().is_none_or(|probe| probe())
352 }
353
354 pub async fn execute(
356 &self,
357 arguments: serde_json::Value,
358 ) -> Result<ToolOutput, ToolExecutionError> {
359 let mut context = ToolContext::new();
360 self.execute_with(&mut context, arguments).await
361 }
362
363 pub async fn execute_with(
368 &self,
369 context: &mut ToolContext,
370 arguments: serde_json::Value,
371 ) -> Result<ToolOutput, ToolExecutionError> {
372 let published = PublishedContext::new();
373 let outcome = crate::serve::serve_inline_with(
374 &self.handler,
375 EffectKind::ToolCall {
376 name: self.definition.name.clone(),
377 args: arguments.to_string(),
378 },
379 vec![
380 Arc::new(context.for_dispatch()),
381 published.clone() as Arc<dyn std::any::Any + Send + Sync>,
382 ],
383 )
384 .await;
385 match outcome {
386 Ok(Outcome::ToolResult { result }) => {
387 context.accept_dispatch_result(published.take().unwrap_or_default());
388 result.into_result()
389 }
390 Ok(other) => Err(ToolExecutionError::other(format!(
391 "tool handler answered with a {} outcome",
392 other.family()
393 ))),
394 Err(report) => Err(ToolExecutionError::other(report.message)),
395 }
396 }
397}
398
399impl std::fmt::Debug for DynamicTool {
400 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
401 f.debug_struct("DynamicTool")
402 .field("name", &self.definition.name)
403 .finish_non_exhaustive()
404 }
405}
406
407pub fn tool_definition<T: Tool>(tool: &T) -> ToolDefinition {
409 ToolDefinition {
410 name: T::NAME.to_string(),
411 description: tool.description(),
412 parameters: tool.parameters(),
413 }
414}
415
416#[cfg(test)]
417mod tests;