Skip to main content

rig_core/embeddings/
tool.rs

1//! Embeddable tool descriptions and serialized tool context.
2//!
3//! ```
4//! use rig_core::embeddings::ToolSchema;
5//!
6//! let schema = ToolSchema {
7//!     name: "search".into(),
8//!     embedding_docs: vec!["Search documents".into()],
9//!     ..Default::default()
10//! };
11//! assert_eq!(schema.embedding_docs.len(), 1);
12//! ```
13
14use crate::{Embed, tool::PortableToolEmbedding};
15use serde::Serialize;
16
17use super::embed::EmbedError;
18
19/// Tool name, serialized context, and text descriptions for embedding-based retrieval.
20#[derive(Clone, Serialize, Default, Eq, PartialEq)]
21pub struct ToolSchema {
22    pub name: String,
23    pub context: serde_json::Value,
24    pub embedding_docs: Vec<String>,
25}
26
27impl Embed for ToolSchema {
28    fn embed(&self, embedder: &mut super::embed::TextEmbedder) -> Result<(), EmbedError> {
29        for doc in &self.embedding_docs {
30            embedder.embed(doc.clone());
31        }
32        Ok(())
33    }
34}
35
36impl ToolSchema {
37    /// Captures a tool's name, context, and embedding descriptions.
38    /// Returns an error if context serialization fails.
39    ///
40    /// ```
41    /// use rig_core::{embeddings::{ToolSchema, EmbedError}, tool::PortableToolEmbedding};
42    ///
43    /// fn schema(tool: &impl PortableToolEmbedding) -> Result<ToolSchema, EmbedError> {
44    ///     ToolSchema::try_from(tool)
45    /// }
46    /// ```
47    pub fn try_from<T>(tool: &T) -> Result<Self, EmbedError>
48    where
49        T: PortableToolEmbedding,
50    {
51        Ok(ToolSchema {
52            name: T::NAME.to_string(),
53            context: serde_json::to_value(tool.context()).map_err(EmbedError::new)?,
54            embedding_docs: tool.embedding_docs(),
55        })
56    }
57}
58
59#[cfg(test)]
60mod tests;