rig_core/embeddings/
embedding.rs1use crate::{
10 completion::Usage,
11 http_client, provider_response,
12 wasm_compat::{WasmCompatSend, WasmCompatSync},
13};
14use serde::{Deserialize, Serialize};
15
16#[derive(Debug, thiserror::Error)]
21#[non_exhaustive]
22pub enum EmbeddingError {
23 #[error("HttpError: {0}")]
25 HttpError(#[from] http_client::Error),
26
27 #[error("JsonError: {0}")]
29 JsonError(#[from] serde_json::Error),
30
31 #[error("UrlError: {0}")]
33 UrlError(#[from] url::ParseError),
34
35 #[cfg(not(target_family = "wasm"))]
36 #[error("DocumentError: {0}")]
38 DocumentError(Box<dyn std::error::Error + Send + Sync + 'static>),
39
40 #[cfg(target_family = "wasm")]
41 #[error("DocumentError: {0}")]
43 DocumentError(Box<dyn std::error::Error + 'static>),
44
45 #[error("ResponseError: {0}")]
47 ResponseError(String),
48
49 #[error("{provider} embeddings do not support the `{parameter}` parameter")]
51 UnsupportedParameter {
52 provider: &'static str,
54 parameter: &'static str,
56 },
57
58 #[error("{provider} embeddings require `{parameter}` {requirement}")]
61 InvalidParameterValue {
62 provider: &'static str,
64 parameter: &'static str,
66 requirement: &'static str,
68 },
69
70 #[error("Rig cannot decode {provider} embedding responses encoded as `{encoding_format}`")]
72 UnsupportedResponseEncoding {
73 provider: &'static str,
75 encoding_format: &'static str,
77 },
78
79 #[error("{provider} embedding response omitted required usage")]
81 MissingUsage {
82 provider: &'static str,
84 },
85
86 #[error("ProviderError: {0}")]
88 ProviderError(String),
89
90 #[error("ProviderResponseError: {0}")]
92 ProviderResponse(provider_response::ProviderResponseError),
93}
94
95crate::provider_response::impl_provider_response_helpers!(EmbeddingError);
96
97pub trait EmbeddingModel: WasmCompatSend + WasmCompatSync {
99 const MAX_DOCUMENTS: usize;
101
102 type Client;
104
105 fn make(client: &Self::Client, model: impl Into<String>, dims: Option<usize>) -> Self;
107
108 fn ndims(&self) -> usize;
110
111 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 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 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 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#[derive(Debug, Clone)]
175pub struct EmbeddingResponse {
176 pub embeddings: Vec<Embedding>,
178 pub usage: Usage,
180}
181
182pub trait ImageEmbeddingModel: Clone + WasmCompatSend + WasmCompatSync {
184 const MAX_DOCUMENTS: usize;
186
187 fn ndims(&self) -> usize;
189
190 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 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#[derive(Clone, Default, Deserialize, Serialize, Debug)]
216pub struct Embedding {
217 pub document: String,
219 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}