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;