Skip to main content

rig_core/providers/cohere/
embeddings.rs

1use super::{client::ApiResponse, client::Client};
2use crate::{
3    embeddings::{self, EmbeddingError},
4    http_client::HttpClientExt,
5    wasm_compat::*,
6};
7use base64::{
8    Engine as _,
9    engine::general_purpose::{STANDARD, URL_SAFE_NO_PAD},
10};
11use serde::Deserialize;
12use serde_json::json;
13use sha2::{Digest, Sha256};
14
15const MAX_IMAGE_BYTES: usize = 5_000_000;
16
17#[derive(Deserialize)]
18pub struct EmbeddingResponse {
19    #[serde(default)]
20    pub response_type: Option<String>,
21    pub id: String,
22    pub embeddings: Vec<Vec<serde_json::Number>>,
23    pub texts: Vec<String>,
24    #[serde(default)]
25    pub meta: Option<Meta>,
26}
27
28#[derive(Deserialize)]
29pub struct Meta {
30    pub api_version: ApiVersion,
31    pub billed_units: BilledUnits,
32    #[serde(default)]
33    pub warnings: Vec<String>,
34}
35
36#[derive(Deserialize)]
37pub struct ApiVersion {
38    pub version: String,
39    #[serde(default)]
40    pub is_deprecated: Option<bool>,
41    #[serde(default)]
42    pub is_experimental: Option<bool>,
43}
44
45#[derive(Deserialize, Debug)]
46pub struct BilledUnits {
47    #[serde(default)]
48    pub input_tokens: u32,
49    #[serde(default)]
50    pub output_tokens: u32,
51    #[serde(default)]
52    pub search_units: u32,
53    #[serde(default)]
54    pub classifications: u32,
55    #[serde(default)]
56    pub images: u32,
57}
58
59impl std::fmt::Display for BilledUnits {
60    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
61        write!(
62            f,
63            "Input tokens: {}\nOutput tokens: {}\nSearch units: {}\nClassifications: {}",
64            self.input_tokens, self.output_tokens, self.search_units, self.classifications
65        )?;
66        if self.images > 0 {
67            write!(f, "\nImages: {}", self.images)?;
68        }
69        Ok(())
70    }
71}
72
73#[derive(Deserialize)]
74struct ImageEmbeddingResponse {
75    embeddings: FloatEmbeddings,
76    #[serde(default)]
77    meta: Option<Meta>,
78}
79
80#[derive(Deserialize)]
81struct FloatEmbeddings {
82    #[serde(rename = "float")]
83    values: Vec<Vec<serde_json::Number>>,
84}
85
86#[derive(Debug, thiserror::Error)]
87enum ImageInputError {
88    #[error("Cohere image embeddings support PNG, JPEG, WebP, or GIF file bytes")]
89    UnsupportedFormat,
90    #[error("Cohere image embeddings accept at most 5 MB per image; received {actual_bytes} bytes")]
91    TooLarge { actual_bytes: usize },
92}
93
94fn image_media_type(bytes: &[u8]) -> Option<&'static str> {
95    if bytes.starts_with(b"\x89PNG\r\n\x1a\n") {
96        Some("image/png")
97    } else if bytes.starts_with(b"\xff\xd8\xff") {
98        Some("image/jpeg")
99    } else if bytes.starts_with(b"GIF87a") || bytes.starts_with(b"GIF89a") {
100        Some("image/gif")
101    } else if bytes.starts_with(b"RIFF") && bytes.get(8..12) == Some(b"WEBP".as_slice()) {
102        Some("image/webp")
103    } else {
104        None
105    }
106}
107
108fn validate_image(bytes: &[u8]) -> Result<&'static str, EmbeddingError> {
109    if bytes.len() > MAX_IMAGE_BYTES {
110        return Err(EmbeddingError::DocumentError(Box::new(
111            ImageInputError::TooLarge {
112                actual_bytes: bytes.len(),
113            },
114        )));
115    }
116
117    image_media_type(bytes)
118        .ok_or_else(|| EmbeddingError::DocumentError(Box::new(ImageInputError::UnsupportedFormat)))
119}
120
121fn image_data_url(bytes: &[u8], media_type: &str) -> String {
122    format!("data:{media_type};base64,{}", STANDARD.encode(bytes))
123}
124
125fn image_document(bytes: &[u8], media_type: &str) -> String {
126    let digest = Sha256::digest(bytes);
127    format!("{media_type};sha256={}", URL_SAFE_NO_PAD.encode(digest))
128}
129
130#[derive(Clone)]
131pub struct EmbeddingModel<T = reqwest::Client> {
132    client: Client<T>,
133    pub model: String,
134    pub input_type: String,
135    ndims: usize,
136}
137
138/// Cohere `embed-english-v3.0` image embedding model.
139///
140/// Cohere Embed v3 accepts one image per request, so batch calls are sent as
141/// ordered individual requests.
142#[derive(Clone)]
143pub struct ImageEmbeddingModel<T = reqwest::Client> {
144    client: Client<T>,
145}
146
147impl<T> embeddings::EmbeddingModel for EmbeddingModel<T>
148where
149    T: HttpClientExt + Clone + WasmCompatSend + WasmCompatSync + 'static,
150{
151    const MAX_DOCUMENTS: usize = 96;
152    type Client = Client<T>;
153
154    fn make(client: &Self::Client, model: impl Into<String>, dims: Option<usize>) -> Self {
155        let model = model.into();
156        let dims = dims
157            .or(super::model_dimensions_from_identifier(&model))
158            .unwrap_or_default();
159
160        Self::new(client.clone(), model, "search_document", dims)
161    }
162
163    fn ndims(&self) -> usize {
164        self.ndims
165    }
166
167    async fn embed_texts(
168        &self,
169        documents: impl IntoIterator<Item = String>,
170    ) -> Result<Vec<embeddings::Embedding>, EmbeddingError> {
171        let documents = documents.into_iter().collect::<Vec<_>>();
172
173        let body = json!({
174            "model": self.model.to_string(),
175            "texts": documents,
176            "input_type": self.input_type
177        });
178
179        let body = serde_json::to_vec(&body)?;
180
181        let req = self
182            .client
183            .post("/v1/embed")?
184            .body(body)
185            .map_err(|e| EmbeddingError::HttpError(e.into()))?;
186
187        let response = self
188            .client
189            .send::<_, Vec<u8>>(req)
190            .await
191            .map_err(EmbeddingError::HttpError)?;
192
193        let status = response.status();
194        let raw_body = response.into_body().await?;
195
196        if status.is_success() {
197            let body: ApiResponse<EmbeddingResponse> = serde_json::from_slice(raw_body.as_slice())?;
198
199            match body {
200                ApiResponse::Ok(response) => {
201                    match response.meta {
202                        Some(meta) => tracing::info!(target: "rig",
203                            "Cohere embeddings billed units: {}",
204                            meta.billed_units,
205                        ),
206                        None => tracing::info!(target: "rig",
207                            "Cohere embeddings billed units: n/a",
208                        ),
209                    };
210
211                    if response.embeddings.len() != documents.len() {
212                        return Err(EmbeddingError::DocumentError(
213                            format!(
214                                "Expected {} embeddings, got {}",
215                                documents.len(),
216                                response.embeddings.len()
217                            )
218                            .into(),
219                        ));
220                    }
221
222                    Ok(response
223                        .embeddings
224                        .into_iter()
225                        .zip(documents.into_iter())
226                        .map(|(embedding, document)| embeddings::Embedding {
227                            document,
228                            vec: embedding.into_iter().filter_map(|n| n.as_f64()).collect(),
229                        })
230                        .collect())
231                }
232                ApiResponse::Err(error) => {
233                    tracing::warn!(
234                        message = %error.message,
235                        "Cohere returned an error response"
236                    );
237                    Err(EmbeddingError::from_http_response(
238                        status,
239                        String::from_utf8_lossy(&raw_body),
240                    ))
241                }
242            }
243        } else {
244            Err(EmbeddingError::from_http_response(
245                status,
246                String::from_utf8_lossy(&raw_body),
247            ))
248        }
249    }
250}
251
252impl<T> embeddings::ImageEmbeddingModel for ImageEmbeddingModel<T>
253where
254    T: HttpClientExt + Clone + WasmCompatSend + WasmCompatSync + 'static,
255{
256    const MAX_DOCUMENTS: usize = 1;
257
258    fn ndims(&self) -> usize {
259        1_024
260    }
261
262    async fn embed_images(
263        &self,
264        images: impl IntoIterator<Item = Vec<u8>> + WasmCompatSend,
265    ) -> Result<Vec<embeddings::Embedding>, EmbeddingError> {
266        let images = images
267            .into_iter()
268            .map(|bytes| {
269                let media_type = validate_image(&bytes)?;
270                let document = image_document(&bytes, media_type);
271                Ok((bytes, media_type, document))
272            })
273            .collect::<Result<Vec<_>, EmbeddingError>>()?;
274        let mut embeddings = Vec::with_capacity(images.len());
275
276        for (image, media_type, document) in images {
277            let data_url = image_data_url(&image, media_type);
278            embeddings.push(self.embed_image_data_url(data_url, document).await?);
279        }
280
281        Ok(embeddings)
282    }
283}
284
285impl<T> EmbeddingModel<T> {
286    pub fn new(
287        client: Client<T>,
288        model: impl Into<String>,
289        input_type: &str,
290        ndims: usize,
291    ) -> Self {
292        Self {
293            client,
294            model: model.into(),
295            input_type: input_type.to_string(),
296            ndims,
297        }
298    }
299
300    pub fn with_model(client: Client<T>, model: &str, input_type: &str, ndims: usize) -> Self {
301        Self {
302            client,
303            model: model.into(),
304            input_type: input_type.into(),
305            ndims,
306        }
307    }
308}
309
310impl<T> ImageEmbeddingModel<T> {
311    pub(crate) fn new(client: Client<T>) -> Self {
312        Self { client }
313    }
314}
315
316impl<T> ImageEmbeddingModel<T>
317where
318    T: HttpClientExt + Clone + WasmCompatSend + WasmCompatSync + 'static,
319{
320    async fn embed_image_data_url(
321        &self,
322        data_url: String,
323        document: String,
324    ) -> Result<embeddings::Embedding, EmbeddingError> {
325        let body = json!({
326            "model": super::EMBED_ENGLISH_V3,
327            "images": [&data_url],
328            "input_type": "image",
329            "embedding_types": ["float"],
330        });
331        let body = serde_json::to_vec(&body)?;
332
333        let request = self
334            .client
335            .post("/v1/embed")?
336            .body(body)
337            .map_err(|error| EmbeddingError::HttpError(error.into()))?;
338        let response = self
339            .client
340            .send::<_, Vec<u8>>(request)
341            .await
342            .map_err(EmbeddingError::HttpError)?;
343        let status = response.status();
344        let raw_body = response.into_body().await?;
345
346        if !status.is_success() {
347            return Err(EmbeddingError::from_http_response(
348                status,
349                String::from_utf8_lossy(&raw_body),
350            ));
351        }
352
353        let body: ApiResponse<ImageEmbeddingResponse> =
354            serde_json::from_slice(raw_body.as_slice())?;
355        let response = match body {
356            ApiResponse::Ok(response) => response,
357            ApiResponse::Err(error) => {
358                tracing::warn!(
359                    message = %error.message,
360                    "Cohere returned an error response"
361                );
362                return Err(EmbeddingError::from_http_response(
363                    status,
364                    String::from_utf8_lossy(&raw_body),
365                ));
366            }
367        };
368
369        match response.meta {
370            Some(meta) => tracing::info!(target: "rig",
371                "Cohere embeddings billed units: {}",
372                meta.billed_units,
373            ),
374            None => tracing::info!(target: "rig", "Cohere embeddings billed units: n/a"),
375        }
376
377        if response.embeddings.values.len() != 1 {
378            return Err(EmbeddingError::DocumentError(
379                format!(
380                    "Expected 1 image embedding, got {}",
381                    response.embeddings.values.len()
382                )
383                .into(),
384            ));
385        }
386
387        let vector = response
388            .embeddings
389            .values
390            .into_iter()
391            .next()
392            .ok_or_else(|| {
393                EmbeddingError::ResponseError(
394                    "Cohere returned an empty image embedding response".to_string(),
395                )
396            })?;
397
398        Ok(embeddings::Embedding {
399            document,
400            vec: vector
401                .into_iter()
402                .filter_map(|number| number.as_f64())
403                .collect(),
404        })
405    }
406}
407
408#[cfg(test)]
409mod tests {
410    use super::*;
411
412    #[tokio::test]
413    async fn embeddings_non_success_preserves_status_and_body() {
414        use crate::embeddings::EmbeddingModel as _;
415        use crate::test_utils::RecordingHttpClient;
416
417        let body = r#"{"error":{"message":"boom"}}"#;
418        let http_client =
419            RecordingHttpClient::with_error_response(http::StatusCode::SERVICE_UNAVAILABLE, body);
420        let client = crate::providers::cohere::Client::builder()
421            .api_key("test-key")
422            .http_client(http_client)
423            .build()
424            .expect("build client");
425        let model = client.embedding_model(
426            crate::providers::cohere::EMBED_ENGLISH_V3,
427            "search_document",
428        );
429
430        let error = model
431            .embed_texts(["hello".to_string()])
432            .await
433            .expect_err("should fail with non-success status");
434
435        assert!(matches!(error, EmbeddingError::HttpError(_)));
436        assert_eq!(
437            error.provider_response_status(),
438            Some(http::StatusCode::SERVICE_UNAVAILABLE)
439        );
440        assert_eq!(error.provider_response_body(), Some(body));
441    }
442
443    #[tokio::test]
444    async fn embeddings_2xx_error_envelope_preserves_status_and_body() {
445        use crate::embeddings::EmbeddingModel as _;
446        use crate::test_utils::RecordingHttpClient;
447
448        // Deserializes to `ApiResponse::Err(ApiErrorResponse { message })` on a 200 OK.
449        let body = r#"{"message":"boom"}"#;
450        let http_client = RecordingHttpClient::new(body);
451        let client = crate::providers::cohere::Client::builder()
452            .api_key("test-key")
453            .http_client(http_client)
454            .build()
455            .expect("build client");
456        let model = client.embedding_model(
457            crate::providers::cohere::EMBED_ENGLISH_V3,
458            "search_document",
459        );
460
461        let error = model
462            .embed_texts(["hello".to_string()])
463            .await
464            .expect_err("should fail with provider error envelope");
465
466        match &error {
467            EmbeddingError::ProviderResponse(stored) => {
468                assert_eq!(stored.body, body);
469                assert_eq!(stored.status, Some(http::StatusCode::OK));
470            }
471            other => panic!("expected ProviderResponse, got {other:?}"),
472        }
473    }
474
475    #[test]
476    fn image_data_urls_detect_every_cohere_image_format() {
477        let cases: &[(&[u8], &str)] = &[
478            (b"\x89PNG\r\n\x1a\n", "image/png"),
479            (b"\xff\xd8\xff", "image/jpeg"),
480            (b"GIF89a", "image/gif"),
481            (b"RIFF\0\0\0\0WEBP", "image/webp"),
482        ];
483
484        for &(bytes, expected_media_type) in cases {
485            let result = validate_image(bytes);
486            assert!(
487                matches!(result, Ok(media_type) if media_type == expected_media_type),
488                "expected {expected_media_type}"
489            );
490            assert!(
491                image_data_url(bytes, expected_media_type)
492                    .starts_with(&format!("data:{expected_media_type};base64,"))
493            );
494        }
495    }
496
497    #[test]
498    fn image_documents_are_stable_without_retaining_image_bytes() {
499        let first = image_document(b"\x89PNG\r\n\x1a\nfirst", "image/png");
500        let second = image_document(b"\x89PNG\r\n\x1a\nother", "image/png");
501
502        assert_eq!(
503            first,
504            image_document(b"\x89PNG\r\n\x1a\nfirst", "image/png")
505        );
506        assert_ne!(first, second);
507        assert!(first.starts_with("image/png;sha256="));
508        assert!(!first.contains("first"));
509    }
510
511    #[test]
512    fn image_data_url_rejects_unsupported_and_oversized_inputs() {
513        assert!(matches!(
514            validate_image(b"not an image"),
515            Err(EmbeddingError::DocumentError(_))
516        ));
517        assert!(matches!(
518            validate_image(&vec![0; MAX_IMAGE_BYTES + 1]),
519            Err(EmbeddingError::DocumentError(_))
520        ));
521    }
522
523    #[tokio::test]
524    async fn image_batches_are_fully_validated_before_any_request() {
525        use crate::embeddings::ImageEmbeddingModel as _;
526        use crate::test_utils::RecordingHttpClient;
527
528        let http_client = RecordingHttpClient::default();
529        let client = crate::providers::cohere::Client::builder()
530            .api_key("test-key")
531            .http_client(http_client.clone())
532            .build()
533            .expect("build client");
534
535        let error = client
536            .image_embedding_model()
537            .embed_images([b"\x89PNG\r\n\x1a\n".to_vec(), b"not an image".to_vec()])
538            .await
539            .expect_err("invalid batch should fail before transport");
540
541        assert!(matches!(error, EmbeddingError::DocumentError(_)));
542        assert!(http_client.requests().is_empty());
543    }
544
545    #[tokio::test]
546    async fn image_embeddings_non_success_preserves_status_and_body() {
547        use crate::embeddings::ImageEmbeddingModel as _;
548        use crate::test_utils::RecordingHttpClient;
549
550        let body = r#"{"error":{"message":"boom"}}"#;
551        let http_client =
552            RecordingHttpClient::with_error_response(http::StatusCode::SERVICE_UNAVAILABLE, body);
553        let client = crate::providers::cohere::Client::builder()
554            .api_key("test-key")
555            .http_client(http_client)
556            .build()
557            .expect("build client");
558
559        let error = client
560            .image_embedding_model()
561            .embed_image(b"\x89PNG\r\n\x1a\n")
562            .await
563            .expect_err("should fail with non-success status");
564
565        assert!(matches!(error, EmbeddingError::HttpError(_)));
566        assert_eq!(
567            error.provider_response_status(),
568            Some(http::StatusCode::SERVICE_UNAVAILABLE)
569        );
570        assert_eq!(error.provider_response_body(), Some(body));
571    }
572
573    #[tokio::test]
574    async fn image_embeddings_2xx_error_envelope_preserves_status_and_body() {
575        use crate::embeddings::ImageEmbeddingModel as _;
576        use crate::test_utils::RecordingHttpClient;
577
578        let body = r#"{"message":"boom"}"#;
579        let http_client = RecordingHttpClient::new(body);
580        let client = crate::providers::cohere::Client::builder()
581            .api_key("test-key")
582            .http_client(http_client)
583            .build()
584            .expect("build client");
585
586        let error = client
587            .image_embedding_model()
588            .embed_image(b"\x89PNG\r\n\x1a\n")
589            .await
590            .expect_err("should fail with provider error envelope");
591
592        match &error {
593            EmbeddingError::ProviderResponse(stored) => {
594                assert_eq!(stored.body, body);
595                assert_eq!(stored.status, Some(http::StatusCode::OK));
596            }
597            other => panic!("expected ProviderResponse, got {other:?}"),
598        }
599    }
600}