rig_core/embeddings/
embedding.rs1use 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 #[error("UrlError: {0}")]
19 UrlError(#[from] url::ParseError),
20
21 #[cfg(not(target_family = "wasm"))]
22 #[error("DocumentError: {0}")]
24 DocumentError(Box<dyn std::error::Error + Send + Sync + 'static>),
25
26 #[cfg(target_family = "wasm")]
27 #[error("DocumentError: {0}")]
29 DocumentError(Box<dyn std::error::Error + 'static>),
30 } {
31 #[error("{provider} embeddings do not support the `{parameter}` parameter")]
33 UnsupportedParameter {
34 provider: &'static str,
36 parameter: &'static str,
38 },
39
40 #[error("{provider} embeddings require `{parameter}` {requirement}")]
43 InvalidParameterValue {
44 provider: &'static str,
46 parameter: &'static str,
48 requirement: &'static str,
50 },
51
52 #[error("Rig cannot decode {provider} embedding responses encoded as `{encoding_format}`")]
54 UnsupportedResponseEncoding {
55 provider: &'static str,
57 encoding_format: &'static str,
59 },
60
61 #[error("{provider} embedding response omitted required usage")]
63 MissingUsage {
64 provider: &'static str,
66 },
67 }
68);
69
70pub trait EmbeddingModel: WasmCompatSend + WasmCompatSync {
72 const MAX_DOCUMENTS: usize;
74
75 type Client;
77
78 fn make(client: &Self::Client, model: impl Into<String>, dims: Option<usize>) -> Self;
80
81 fn ndims(&self) -> usize;
83
84 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 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 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 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#[derive(Debug, Clone)]
148pub struct EmbeddingResponse {
149 pub embeddings: Vec<Embedding>,
151 pub usage: Usage,
153}
154
155pub trait ImageEmbeddingModel: Clone + WasmCompatSend + WasmCompatSync {
157 const MAX_DOCUMENTS: usize;
159
160 fn ndims(&self) -> usize;
162
163 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 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#[derive(Clone, Default, Deserialize, Serialize, Debug)]
191pub struct Embedding {
192 pub document: String,
195 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}