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#[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 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}