Skip to main content

rig_core/embeddings/
embed.rs

1//! Text extraction for embedding builders, including implementations for common types.
2//!
3//! ```
4//! use rig_core::embeddings::to_texts;
5//!
6//! assert_eq!(to_texts("document")?, vec!["document"]);
7//! # Ok::<(), rig_core::embeddings::EmbedError>(())
8//! ```
9
10/// Failure extracting text for embedding.
11#[derive(Debug, thiserror::Error)]
12#[error("{0}")]
13pub struct EmbedError(#[from] Box<dyn std::error::Error + Send + Sync>);
14
15impl EmbedError {
16    pub fn new<E: std::error::Error + Send + Sync + 'static>(error: E) -> Self {
17        EmbedError(Box::new(error))
18    }
19}
20
21/// Extracts text fragments for vector embedding in append order.
22/// Implementations report extraction failures through [`EmbedError`].
23///
24/// ```
25/// use rig_core::{Embed, embeddings::{EmbedError, TextEmbedder, to_texts}};
26///
27/// struct Definitions(String);
28/// impl Embed for Definitions {
29///     fn embed(&self, out: &mut TextEmbedder) -> Result<(), EmbedError> {
30///         for definition in self.0.split(',') {
31///             out.embed(definition.trim().to_owned());
32///         }
33///         Ok(())
34///     }
35/// }
36/// assert_eq!(to_texts(Definitions("fruit, company".into()))?, vec!["fruit", "company"]);
37/// # Ok::<(), EmbedError>(())
38/// ```
39pub trait Embed {
40    /// Append all text fragments that should be embedded for this value.
41    fn embed(&self, embedder: &mut TextEmbedder) -> Result<(), EmbedError>;
42}
43
44/// Accumulates string values that need to be embedded.
45/// Used by the [Embed] trait.
46#[derive(Default)]
47pub struct TextEmbedder {
48    pub(crate) texts: Vec<String>,
49}
50
51impl TextEmbedder {
52    /// Appends a text fragment without filtering or deduplication.
53    pub fn embed(&mut self, text: String) {
54        self.texts.push(text);
55    }
56}
57
58/// Extracts text fragments in append order, propagating extraction errors.
59pub fn to_texts(item: impl Embed) -> Result<Vec<String>, EmbedError> {
60    let mut embedder = TextEmbedder::default();
61    item.embed(&mut embedder)?;
62    Ok(embedder.texts)
63}
64
65impl Embed for String {
66    fn embed(&self, embedder: &mut TextEmbedder) -> Result<(), EmbedError> {
67        embedder.embed(self.clone());
68        Ok(())
69    }
70}
71
72impl Embed for &str {
73    fn embed(&self, embedder: &mut TextEmbedder) -> Result<(), EmbedError> {
74        embedder.embed(self.to_string());
75        Ok(())
76    }
77}
78
79impl Embed for i8 {
80    fn embed(&self, embedder: &mut TextEmbedder) -> Result<(), EmbedError> {
81        embedder.embed(self.to_string());
82        Ok(())
83    }
84}
85
86impl Embed for i16 {
87    fn embed(&self, embedder: &mut TextEmbedder) -> Result<(), EmbedError> {
88        embedder.embed(self.to_string());
89        Ok(())
90    }
91}
92
93impl Embed for i32 {
94    fn embed(&self, embedder: &mut TextEmbedder) -> Result<(), EmbedError> {
95        embedder.embed(self.to_string());
96        Ok(())
97    }
98}
99
100impl Embed for i64 {
101    fn embed(&self, embedder: &mut TextEmbedder) -> Result<(), EmbedError> {
102        embedder.embed(self.to_string());
103        Ok(())
104    }
105}
106
107impl Embed for i128 {
108    fn embed(&self, embedder: &mut TextEmbedder) -> Result<(), EmbedError> {
109        embedder.embed(self.to_string());
110        Ok(())
111    }
112}
113
114impl Embed for f32 {
115    fn embed(&self, embedder: &mut TextEmbedder) -> Result<(), EmbedError> {
116        embedder.embed(self.to_string());
117        Ok(())
118    }
119}
120
121impl Embed for f64 {
122    fn embed(&self, embedder: &mut TextEmbedder) -> Result<(), EmbedError> {
123        embedder.embed(self.to_string());
124        Ok(())
125    }
126}
127
128impl Embed for bool {
129    fn embed(&self, embedder: &mut TextEmbedder) -> Result<(), EmbedError> {
130        embedder.embed(self.to_string());
131        Ok(())
132    }
133}
134
135impl Embed for char {
136    fn embed(&self, embedder: &mut TextEmbedder) -> Result<(), EmbedError> {
137        embedder.embed(self.to_string());
138        Ok(())
139    }
140}
141
142impl Embed for serde_json::Value {
143    fn embed(&self, embedder: &mut TextEmbedder) -> Result<(), EmbedError> {
144        embedder.embed(serde_json::to_string(self).map_err(EmbedError::new)?);
145        Ok(())
146    }
147}
148
149impl<T: Embed> Embed for &T {
150    fn embed(&self, embedder: &mut TextEmbedder) -> Result<(), EmbedError> {
151        (*self).embed(embedder)
152    }
153}
154
155impl<T: Embed> Embed for Vec<T> {
156    fn embed(&self, embedder: &mut TextEmbedder) -> Result<(), EmbedError> {
157        for item in self {
158            item.embed(embedder).map_err(EmbedError::new)?;
159        }
160        Ok(())
161    }
162}