1use std::{future::Future, sync::Arc};
15
16use serde::{Deserialize, Serialize};
17
18use crate::{
19 completion::{ToolDefinition, message::ToolName},
20 effect::{EffectKind, Outcome},
21 serve::{ErasedHandler, adapters::ToolCallback, adapters::ToolFn},
22 wasm_compat::{WasmBoxedFuture, WasmCompatSend, WasmCompatSync},
23};
24
25use super::{
26 IntoToolOutput, PublishedContext, ToolContext, ToolExecutionError, ToolOutput, ToolResult,
27};
28
29pub trait Tool: Sized + WasmCompatSend + WasmCompatSync {
36 const NAME: &'static str;
38 type Args: for<'de> Deserialize<'de> + WasmCompatSend + WasmCompatSync;
40 type Output: IntoToolOutput;
48 type Error: std::error::Error + WasmCompatSend + WasmCompatSync + 'static;
55
56 fn description(&self) -> String;
58
59 fn parameters(&self) -> serde_json::Value;
61
62 fn map_error(&self, error: Self::Error) -> ToolExecutionError {
68 ToolExecutionError::from_error(error)
69 }
70
71 fn call(
73 &self,
74 context: &mut ToolContext,
75 args: Self::Args,
76 ) -> impl Future<Output = Result<Self::Output, Self::Error>> + WasmCompatSend;
77}
78
79impl<T> Tool for T
80where
81 T: super::PortableTool,
82{
83 const NAME: &'static str = <T as super::PortableTool>::NAME;
84 type Args = <T as super::PortableTool>::Args;
85 type Output = <T as super::PortableTool>::Output;
86 type Error = <T as super::PortableTool>::Error;
87
88 fn description(&self) -> String {
89 super::PortableTool::description(self)
90 }
91
92 fn parameters(&self) -> serde_json::Value {
93 super::PortableTool::parameters(self)
94 }
95
96 fn map_error(&self, error: Self::Error) -> ToolExecutionError {
97 super::PortableTool::map_error(self, error)
98 }
99
100 async fn call(
101 &self,
102 _context: &mut ToolContext,
103 args: Self::Args,
104 ) -> Result<Self::Output, Self::Error> {
105 super::PortableTool::call(self, args).await
106 }
107}
108
109pub trait ToolEmbedding: Tool {
111 type InitError: std::error::Error + WasmCompatSend + WasmCompatSync + 'static;
113 type Context: for<'de> Deserialize<'de> + Serialize;
115 type State: WasmCompatSend;
117
118 fn embedding_docs(&self) -> Vec<String>;
120 fn context(&self) -> Self::Context;
122 fn init(state: Self::State, context: Self::Context) -> Result<Self, Self::InitError>;
124}
125
126impl<T> ToolEmbedding for T
127where
128 T: super::PortableToolEmbedding,
129{
130 type InitError = <T as super::PortableToolEmbedding>::InitError;
131 type Context = <T as super::PortableToolEmbedding>::Context;
132 type State = <T as super::PortableToolEmbedding>::State;
133
134 fn embedding_docs(&self) -> Vec<String> {
135 super::PortableToolEmbedding::embedding_docs(self)
136 }
137
138 fn context(&self) -> Self::Context {
139 super::PortableToolEmbedding::context(self)
140 }
141
142 fn init(state: Self::State, context: Self::Context) -> Result<Self, Self::InitError> {
143 super::PortableToolEmbedding::init(state, context)
144 }
145}
146
147fn parse_tool_args<A>(args: &str) -> Result<A, ToolExecutionError>
148where
149 A: serde::de::DeserializeOwned,
150{
151 match serde_json::from_str(args) {
152 Ok(parsed) => Ok(parsed),
153 Err(original) if args.trim() == "null" => serde_json::from_str("{}").map_err(|_| {
154 ToolExecutionError::invalid_args(format!("failed to parse tool arguments: {original}"))
155 .with_source(original)
156 }),
157 Err(error) => Err(ToolExecutionError::invalid_args(format!(
158 "failed to parse tool arguments: {error}"
159 ))
160 .with_source(error)),
161 }
162}
163
164pub(crate) async fn execute_callback<F>(
167 callback: &F,
168 args: String,
169 context: &mut ToolContext,
170) -> ToolResult
171where
172 F: for<'a> Fn(
173 &'a mut ToolContext,
174 serde_json::Value,
175 ) -> WasmBoxedFuture<'a, Result<ToolOutput, ToolExecutionError>>,
176{
177 let args = match parse_tool_args::<serde_json::Value>(&args) {
178 Ok(args) => args,
179 Err(error) => return ToolResult::failed(error),
180 };
181 tool_result_from(callback(context, args).await)
182}
183
184fn tool_result_from<O>(outcome: Result<O, ToolExecutionError>) -> ToolResult
185where
186 O: IntoToolOutput,
187{
188 match outcome.and_then(IntoToolOutput::into_tool_output) {
189 Ok(output) => ToolResult::success(output),
190 Err(error) => ToolResult::failed(error),
191 }
192}
193
194pub trait ErasedTool: WasmCompatSend + WasmCompatSync {
198 fn name(&self) -> String;
200 fn description(&self) -> String;
202 fn parameters(&self) -> serde_json::Value;
204 fn execute<'a>(
206 &'a self,
207 args: String,
208 context: &'a mut ToolContext,
209 ) -> WasmBoxedFuture<'a, ToolResult>;
210}
211
212impl<T> ErasedTool for T
213where
214 T: Tool,
215{
216 fn name(&self) -> String {
217 T::NAME.to_string()
218 }
219
220 fn description(&self) -> String {
221 Tool::description(self)
222 }
223
224 fn parameters(&self) -> serde_json::Value {
225 Tool::parameters(self)
226 }
227
228 fn execute<'a>(
229 &'a self,
230 args: String,
231 context: &'a mut ToolContext,
232 ) -> WasmBoxedFuture<'a, ToolResult> {
233 Box::pin(async move {
234 let args = match parse_tool_args::<T::Args>(&args) {
235 Ok(args) => args,
236 Err(error) => return ToolResult::failed(error),
237 };
238 tool_result_from(
239 Tool::call(self, context, args)
240 .await
241 .map_err(|error| Tool::map_error(self, error)),
242 )
243 })
244 }
245}
246
247#[cfg(not(target_family = "wasm"))]
251pub type LivenessFn = Arc<dyn Fn() -> bool + Send + Sync>;
252#[cfg(target_family = "wasm")]
254pub type LivenessFn = Arc<dyn Fn() -> bool>;
255
256#[derive(Clone)]
262pub struct DynamicTool {
263 definition: ToolDefinition,
264 handler: ErasedHandler,
265 liveness: Option<LivenessFn>,
266}
267
268impl DynamicTool {
269 pub fn new<F>(
271 name: ToolName,
272 description: impl Into<String>,
273 parameters: serde_json::Value,
274 callback: F,
275 ) -> Self
276 where
277 F: Fn(
278 serde_json::Value,
279 ) -> WasmBoxedFuture<'static, Result<ToolOutput, ToolExecutionError>>
280 + WasmCompatSend
281 + WasmCompatSync
282 + 'static,
283 {
284 Self::new_with_context(
285 name,
286 description,
287 parameters,
288 move |_context: &mut ToolContext, arguments| callback(arguments),
289 )
290 }
291
292 pub fn new_with_context<F>(
294 name: ToolName,
295 description: impl Into<String>,
296 parameters: serde_json::Value,
297 callback: F,
298 ) -> Self
299 where
300 F: ToolCallback + 'static,
301 {
302 let description = description.into();
303 let handler = ErasedHandler::new(ToolFn::new(
304 name.to_string(),
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) -> &ToolName {
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.to_string(),
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(report.into()),
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_name<T: Tool>() -> ToolName {
444 const { assert!(!T::NAME.is_empty(), "Tool::NAME cannot be empty") };
445 ToolName::new_unchecked(T::NAME)
446}
447
448pub fn tool_definition<T: Tool>(tool: &T) -> ToolDefinition {
450 ToolDefinition {
451 name: tool_name::<T>(),
452 description: tool.description(),
453 parameters: tool.parameters(),
454 }
455}