Skip to main content

rig_core/embeddings/
embedding.rs

1//! The module defines the [EmbeddingModel] and [ImageEmbeddingModel] traits, which represent
2//! embedding models that can generate embeddings for text documents and images.
3//!
4//! The module also defines the [Embedding] struct, which represents a single document embedding.
5//!
6//! Finally, the module defines the [EmbeddingError] enum, which represents various errors that
7//! can occur during embedding generation or processing.
8
9use crate::{
10    completion::Usage,
11    wasm_compat::{WasmCompatSend, WasmCompatSync},
12};
13use serde::{Deserialize, Serialize};
14
15crate::provider_response::provider_error_enum!(
16    EmbeddingError, "embedding" {
17    /// URL construction or parsing failed while preparing a provider request.
18    #[error("UrlError: {0}")]
19    UrlError(#[from] url::ParseError),
20
21    #[cfg(not(target_family = "wasm"))]
22    /// Error processing the document for embedding
23    #[error("DocumentError: {0}")]
24    DocumentError(Box<dyn std::error::Error + Send + Sync + 'static>),
25
26    #[cfg(target_family = "wasm")]
27    /// Error processing the document for embedding
28    #[error("DocumentError: {0}")]
29    DocumentError(Box<dyn std::error::Error + 'static>),
30    } {
31    /// The provider does not support an embedding request parameter configured on the model.
32    #[error("{provider} embeddings do not support the `{parameter}` parameter")]
33    UnsupportedParameter {
34        /// Provider whose embedding API rejected the parameter.
35        provider: &'static str,
36        /// Unsupported request parameter.
37        parameter: &'static str,
38    },
39
40    /// A provider request parameter was configured with a value outside the
41    /// provider's supported range.
42    #[error("{provider} embeddings require `{parameter}` {requirement}")]
43    InvalidParameterValue {
44        /// Provider whose embedding API constrains the parameter.
45        provider: &'static str,
46        /// Request parameter with the invalid value.
47        parameter: &'static str,
48        /// Concise description of the accepted values.
49        requirement: &'static str,
50    },
51
52    /// Rig cannot decode the requested provider response encoding.
53    #[error("Rig cannot decode {provider} embedding responses encoded as `{encoding_format}`")]
54    UnsupportedResponseEncoding {
55        /// Provider whose response encoding was requested.
56        provider: &'static str,
57        /// Response encoding that Rig cannot decode.
58        encoding_format: &'static str,
59    },
60
61    /// A provider that guarantees embedding usage omitted it from the response.
62    #[error("{provider} embedding response omitted required usage")]
63    MissingUsage {
64        /// Provider whose response omitted usage.
65        provider: &'static str,
66    },
67    }
68);
69
70/// Trait for embedding models that can generate embeddings for documents.
71pub trait EmbeddingModel: WasmCompatSend + WasmCompatSync {
72    /// The maximum number of documents that can be embedded in a single request.
73    const MAX_DOCUMENTS: usize;
74
75    /// Provider client type used to construct this embedding model.
76    type Client;
77
78    /// Construct a model handle from a provider client, model identifier, and optional dimensions.
79    fn make(client: &Self::Client, model: impl Into<String>, dims: Option<usize>) -> Self;
80
81    /// The number of dimensions in the embedding vector.
82    fn ndims(&self) -> usize;
83
84    /// Embed multiple text documents in a single request
85    fn embed_texts(
86        &self,
87        texts: impl IntoIterator<Item = String> + WasmCompatSend,
88    ) -> impl std::future::Future<Output = Result<Vec<Embedding>, EmbeddingError>> + WasmCompatSend;
89
90    /// Embed a single text document.
91    fn embed_text(
92        &self,
93        text: &str,
94    ) -> impl std::future::Future<Output = Result<Embedding, EmbeddingError>> + WasmCompatSend {
95        async {
96            let mut embeddings = self.embed_texts(vec![text.to_string()]).await?;
97            embeddings.pop().ok_or_else(|| {
98                EmbeddingError::ResponseError(
99                    "embedding provider returned an empty response for embed_text".to_string(),
100                )
101            })
102        }
103    }
104
105    /// Embed multiple text documents in a single request and return token usage.
106    ///
107    /// The default implementation delegates to [`EmbeddingModel::embed_texts`] and returns
108    /// zero-valued usage. Providers that expose usage information from their embedding API
109    /// should override this method.
110    fn embed_texts_with_usage(
111        &self,
112        texts: impl IntoIterator<Item = String> + WasmCompatSend,
113    ) -> impl std::future::Future<Output = Result<EmbeddingResponse, EmbeddingError>> + WasmCompatSend
114    {
115        async {
116            let embeddings = self.embed_texts(texts).await?;
117            Ok(EmbeddingResponse {
118                embeddings,
119                usage: Usage::default(),
120            })
121        }
122    }
123
124    /// Embed a single text document and return token usage.
125    ///
126    /// The default implementation delegates to
127    /// [`EmbeddingModel::embed_texts_with_usage`].
128    fn embed_text_with_usage(
129        &self,
130        text: &str,
131    ) -> impl std::future::Future<Output = Result<EmbeddingResponse, EmbeddingError>> + WasmCompatSend
132    {
133        async {
134            let response = self.embed_texts_with_usage(vec![text.to_string()]).await?;
135            if response.embeddings.is_empty() {
136                return Err(EmbeddingError::ResponseError(
137                    "embedding provider returned an empty response for embed_text_with_usage"
138                        .to_string(),
139                ));
140            }
141            Ok(response)
142        }
143    }
144}
145
146/// Response from an embedding request containing the embeddings and token usage.
147#[derive(Debug, Clone)]
148pub struct EmbeddingResponse {
149    /// The embeddings returned by the provider, one per input text.
150    pub embeddings: Vec<Embedding>,
151    /// Token usage for this embedding request.
152    pub usage: Usage,
153}
154
155/// Trait for embedding models that can generate embeddings for images.
156pub trait ImageEmbeddingModel: Clone + WasmCompatSend + WasmCompatSync {
157    /// The maximum number of images the provider accepts in one request.
158    const MAX_DOCUMENTS: usize;
159
160    /// The number of dimensions in the embedding vector.
161    fn ndims(&self) -> usize;
162
163    /// Embed a batch of images from their encoded file bytes.
164    ///
165    /// Implementations must preserve input order in the returned embeddings.
166    /// The returned [`Embedding::document`] should identify the input without
167    /// retaining the raw image or a reversible encoding of it.
168    fn embed_images(
169        &self,
170        images: impl IntoIterator<Item = Vec<u8>> + WasmCompatSend,
171    ) -> impl std::future::Future<Output = Result<Vec<Embedding>, EmbeddingError>> + WasmCompatSend;
172
173    /// Embed a single image from its encoded file bytes.
174    fn embed_image<'a>(
175        &'a self,
176        bytes: &'a [u8],
177    ) -> impl std::future::Future<Output = Result<Embedding, EmbeddingError>> + WasmCompatSend {
178        async move {
179            let mut embeddings = self.embed_images(vec![bytes.to_owned()]).await?;
180            embeddings.pop().ok_or_else(|| {
181                EmbeddingError::ResponseError(
182                    "embedding provider returned an empty response for embed_image".to_string(),
183                )
184            })
185        }
186    }
187}
188
189/// Struct that holds a single document and its embedding.
190#[derive(Clone, Default, Deserialize, Serialize, Debug)]
191pub struct Embedding {
192    /// The text that was embedded, or a non-sensitive input identifier for
193    /// non-text embeddings. Used for debugging and equality.
194    pub document: String,
195    /// The embedding vector
196    pub vec: Vec<f64>,
197}
198
199impl PartialEq for Embedding {
200    fn eq(&self, other: &Self) -> bool {
201        self.document == other.document
202    }
203}
204
205impl Eq for Embedding {}
206
207#[cfg(test)]
208mod provider_response_tests {
209    use super::*;
210    use crate::{http_client, provider_response};
211    use http::StatusCode;
212
213    #[test]
214    fn embedding_error_provider_response_helpers_with_preserved_json_body() {
215        let body = r#"{"error":{"message":"rate limited"}}"#;
216        let error = EmbeddingError::ProviderResponse(
217            provider_response::ProviderResponseError::without_status(body.to_string()),
218        );
219
220        assert_eq!(error.provider_response_body(), Some(body));
221        assert_eq!(error.provider_response_status(), None);
222        assert_eq!(
223            error.provider_response_json().expect("valid JSON"),
224            Some(serde_json::json!({ "error": { "message": "rate limited" } }))
225        );
226    }
227
228    #[test]
229    fn embedding_error_provider_error_is_not_a_provider_response() {
230        let error = EmbeddingError::ProviderError("internal diagnostic".to_string());
231
232        assert_eq!(error.provider_response_body(), None);
233        assert_eq!(error.provider_response_status(), None);
234        assert_eq!(error.provider_response_json().expect("no body"), None);
235    }
236
237    #[test]
238    fn embedding_error_provider_response_helpers_with_http_non_success() {
239        let body = r#"{"error":{"message":"bad request"}}"#;
240        let error = EmbeddingError::HttpError(http_client::Error::InvalidStatusCodeWithMessage(
241            StatusCode::BAD_REQUEST,
242            body.to_string(),
243        ));
244
245        assert_eq!(error.provider_response_body(), Some(body));
246        assert_eq!(
247            error.provider_response_status(),
248            Some(StatusCode::BAD_REQUEST)
249        );
250        assert_eq!(
251            error.provider_response_json().expect("valid JSON"),
252            Some(serde_json::json!({ "error": { "message": "bad request" } }))
253        );
254    }
255
256    #[test]
257    fn embedding_error_provider_response_helpers_with_preserved_plain_text_body() {
258        let error = EmbeddingError::ProviderResponse(
259            provider_response::ProviderResponseError::without_status("not json".to_string()),
260        );
261
262        assert_eq!(error.provider_response_body(), Some("not json"));
263        assert!(error.provider_response_json().is_err());
264    }
265
266    #[test]
267    fn embedding_error_provider_response_helpers_with_unrelated_variant() {
268        let error = EmbeddingError::ResponseError("parse failed".to_string());
269
270        assert_eq!(error.provider_response_body(), None);
271        assert_eq!(error.provider_response_status(), None);
272        assert_eq!(error.provider_response_json().expect("no body"), None);
273    }
274}