Skip to main content

rig_core/embeddings/
embedding.rs

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