Skip to main content

rig_core/tool/
contextual.rs

1//! Contextual tool authoring and JSON-argument dispatch adapters.
2//! [`ToolContext`] carries typed inbound values and host-only result metadata;
3//! model-visible outputs retain their text, JSON, or multimodal representation.
4//!
5//! ```
6//! use rig_core::tool::{DynamicTool, ToolOutput};
7//!
8//! let tool = DynamicTool::new_with_context("echo", "Echo JSON", serde_json::json!({}),
9//!     |_context, args| Box::pin(async move { Ok(ToolOutput::json(args)) }));
10//! assert_eq!(tool.name(), "echo");
11//! ```
12
13use 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
28/// A typed LLM tool.
29///
30/// Tool authors provide metadata and exactly one execution method. Runtime
31/// context and host-only result metadata share the [`ToolContext`] path. Rig's
32/// object-safe dispatch boundary is private; use [`DynamicTool`] when the tool
33/// name or callback is only known at runtime.
34pub trait Tool: Sized + WasmCompatSend + WasmCompatSync {
35    /// Unique registration and provider-facing name.
36    const NAME: &'static str;
37    /// Typed JSON arguments.
38    type Args: for<'de> Deserialize<'de> + WasmCompatSend + WasmCompatSync;
39    /// Output convertible into Rig's canonical model presentation.
40    ///
41    /// Every owned serializable value implements [`IntoToolOutput`]
42    /// automatically. [`ToolResultContent`](crate::message::ToolResultContent)
43    /// and `Vec<ToolResultContent>` preserve rich content when returned
44    /// directly; use [`ToolOutput`] when constructing the presentation
45    /// explicitly.
46    type Output: IntoToolOutput;
47    /// Typed error returned by direct calls to this tool.
48    ///
49    /// Rig normalizes this error into [`ToolExecutionError`] only at the erased
50    /// dispatch boundary. This keeps ordinary `?` propagation and typed unit
51    /// tests available to tool authors without creating a second runtime error
52    /// representation.
53    type Error: std::error::Error + WasmCompatSend + WasmCompatSync + 'static;
54
55    /// Model-facing description.
56    fn description(&self) -> String;
57
58    /// JSON Schema for arguments.
59    fn parameters(&self) -> serde_json::Value;
60
61    /// Normalize a typed author-facing error for runtime policy and telemetry.
62    ///
63    /// The default preserves concrete sources for operators and exposes safe
64    /// kind-level model feedback. An existing [`ToolExecutionError`] retains its
65    /// classification and model output. Override to supply deliberate domain feedback.
66    fn map_error(&self, error: Self::Error) -> ToolExecutionError {
67        ToolExecutionError::from_error(error)
68    }
69
70    /// Execute the tool.
71    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
108/// A tool that can be stored in a vector store and reconstructed for RAG.
109pub trait ToolEmbedding: Tool {
110    /// Error returned while reconstructing the tool.
111    type InitError: std::error::Error + WasmCompatSend + WasmCompatSync + 'static;
112    /// Serializable static context.
113    type Context: for<'de> Deserialize<'de> + Serialize;
114    /// Runtime initialization state.
115    type State: WasmCompatSend;
116
117    /// Documents used to retrieve the tool.
118    fn embedding_docs(&self) -> Vec<String>;
119    /// Serializable tool context.
120    fn context(&self) -> Self::Context;
121    /// Reconstruct the tool.
122    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
163/// Parses JSON arguments and runs a contextual callback, returning parse,
164/// execution, or output-conversion failures as failed tool results.
165pub(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
193/// The object-safe form of [`Tool`]: raw JSON arguments in, a
194/// [`ToolResult`] out. This is the impl-side contract the bus's
195/// `ToolAdapter` calls; nothing stores it behind a vtable.
196pub trait ErasedTool: WasmCompatSend + WasmCompatSync {
197    /// The tool's name.
198    fn name(&self) -> String;
199    /// The tool's description.
200    fn description(&self) -> String;
201    /// The JSON schema of the tool's arguments.
202    fn parameters(&self) -> serde_json::Value;
203    /// Run the tool on raw arguments, shaping the answer into a result.
204    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/// Reports whether a tool's owner still serves it. Registries check lazily on
247/// reads or reconciliation; calls before retirement may fail with a transport
248/// error rather than `HandlerUnavailable`.
249#[cfg(not(target_family = "wasm"))]
250pub type LivenessFn = Arc<dyn Fn() -> bool + Send + Sync>;
251/// A liveness probe (browser wasm: no `Send + Sync`, no threads).
252#[cfg(target_family = "wasm")]
253pub type LivenessFn = Arc<dyn Fn() -> bool>;
254
255/// A tool defined at runtime: a name, a schema and a callback. The callback
256/// is the handler ([`ToolFn`]); this struct is its definition plus the
257/// erased handler a registry stages until a bus takes it. The optional
258/// liveness probe supports registry retirement; inline execution does not
259/// consult it.
260#[derive(Clone)]
261pub struct DynamicTool {
262    definition: ToolDefinition,
263    handler: ErasedHandler,
264    liveness: Option<LivenessFn>,
265}
266
267impl DynamicTool {
268    /// Define a tool from a context-free callback over owned arguments.
269    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    /// Define a tool from a callback over the dispatch-scoped context.
292    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    /// Attach a liveness probe.
321    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    /// The tool's name.
330    pub fn name(&self) -> &str {
331        &self.definition.name
332    }
333
334    /// The tool's definition.
335    pub fn definition(&self) -> ToolDefinition {
336        self.definition.clone()
337    }
338
339    /// The erased handler behind this definition.
340    pub fn handler(&self) -> &ErasedHandler {
341        &self.handler
342    }
343
344    /// Consumes the tool into its definition, handler, and optional liveness probe.
345    pub fn into_parts(self) -> (ToolDefinition, ErasedHandler, Option<LivenessFn>) {
346        (self.definition, self.handler, self.liveness)
347    }
348
349    /// Whether the tool's owner still serves it (`true` without a probe).
350    pub fn is_live(&self) -> bool {
351        self.liveness.as_ref().is_none_or(|probe| probe())
352    }
353
354    /// Run the tool inline with an empty context.
355    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    /// Run the tool inline with isolated inbound values. A completed call
364    /// replaces only `context`'s result metadata, including when the tool
365    /// returns an error. Dropping the execution future leaves the caller's
366    /// context unchanged; it does not publish partial mutations.
367    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
407/// A tool's [`ToolDefinition`] from a typed tool.
408pub 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;