Skip to main content

rig_core/providers/cohere/
embeddings.rs

1//! Cohere's text-embedding and image-embedding wires (`POST /v1/embed`),
2//! their reply types, and image-input validation.
3//!
4//! ```
5//! use rig_core::providers::cohere::embeddings::FloatEmbeddings;
6//! let vectors: FloatEmbeddings = serde_json::from_str(r#"{"float": [[0.5]]}"#)?;
7//! assert_eq!(vectors.values.len(), 1);
8//! # Ok::<(), serde_json::Error>(())
9//! ```
10
11use crate::error::{EncodeError, ProviderError};
12use crate::operation::{Embedding, ImageEmbedding};
13use crate::providers::internal::wire::classify_reply_or_message_envelope;
14use crate::wire::{
15    Body, Capabilities, Decoder, Descriptor, Encoded, Flow, Framing, Mode, Out, Wire, WireEvent,
16    WireFrame,
17};
18use base64::{Engine as _, engine::general_purpose::STANDARD};
19use serde::{Deserialize, Serialize};
20
21use super::{CohereConfig, PROVIDER_NAME};
22
23/// `embed-v4.0` embedding model
24pub const EMBED_V4: &str = "embed-v4.0";
25/// `embed-english-v3.0` embedding model
26pub const EMBED_ENGLISH_V3: &str = "embed-english-v3.0";
27/// `embed-english-light-v3.0` embedding model
28pub const EMBED_ENGLISH_LIGHT_V3: &str = "embed-english-light-v3.0";
29/// `embed-multilingual-v3.0` embedding model
30pub const EMBED_MULTILINGUAL_V3: &str = "embed-multilingual-v3.0";
31/// `embed-multilingual-light-v3.0` embedding model
32pub const EMBED_MULTILINGUAL_LIGHT_V3: &str = "embed-multilingual-light-v3.0";
33
34pub(crate) fn model_dimensions_from_identifier(identifier: &str) -> Option<usize> {
35    match identifier {
36        EMBED_V4 => Some(1_536),
37        EMBED_ENGLISH_V3 | EMBED_MULTILINGUAL_V3 => Some(1_024),
38        EMBED_ENGLISH_LIGHT_V3 | EMBED_MULTILINGUAL_LIGHT_V3 => Some(384),
39        _ => None,
40    }
41}
42
43impl CohereConfig {
44    /// Build a text-embedding wire reporting the supplied or known model width,
45    /// or zero if unknown. This width is metadata, not a request parameter.
46    pub(crate) fn embedding(&self, model: impl Into<String>, ndims: Option<usize>) -> Embeddings {
47        let model = model.into();
48        let ndims = ndims
49            .or_else(|| model_dimensions_from_identifier(&model))
50            .unwrap_or_default();
51        Embeddings {
52            provider: self.clone(),
53            model,
54            ndims,
55            input_type: DEFAULT_INPUT_TYPE.to_owned(),
56        }
57    }
58
59    /// The image-embedding wire.
60    ///
61    /// Cohere Embed v3 embeds images with one fixed model, so this wire
62    /// names no model.
63    pub(crate) fn image_embedding(&self) -> ImageEmbeddings {
64        ImageEmbeddings {
65            provider: self.clone(),
66        }
67    }
68
69    /// One request, authenticated and typed as JSON.
70    pub(super) fn post(&self, path: &str) -> http::request::Builder {
71        http::Request::post(format!("{}{path}", self.base_url))
72            .header(http::header::CONTENT_TYPE, "application/json")
73            .header(
74                http::header::AUTHORIZATION,
75                format!("Bearer {}", self.api_key.expose()),
76            )
77    }
78}
79
80/// Default retrieval role for embeddings of stored document chunks.
81const DEFAULT_INPUT_TYPE: &str = "search_document";
82
83/// The most texts Cohere embeds in one `/v1/embed` call.
84const MAX_DOCUMENTS: usize = 96;
85
86/// The width Cohere's image embeddings come back at.
87const IMAGE_NDIMS: usize = 1_024;
88
89/// The text-embedding wire: `POST /v1/embed`.
90#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
91pub struct Embeddings {
92    /// The provider this wire speaks to.
93    pub provider: CohereConfig,
94    /// The model to address.
95    pub model: String,
96    /// The width this wire reports, from the caller or the model's published
97    /// dimensions. `0` means neither named one.
98    pub ndims: usize,
99    /// Cohere's retrieval prompt: `search_document` for stored chunks,
100    /// `search_query` for queries, `classification`, `clustering`.
101    pub input_type: String,
102}
103
104impl Embeddings {
105    /// Embed for a different purpose than storing chunks.
106    pub fn with_input_type(mut self, input_type: impl Into<String>) -> Self {
107        self.input_type = input_type.into();
108        self
109    }
110}
111
112impl Wire for Embeddings {
113    type Op = Embedding;
114    type Payload = crate::wire::Encoded;
115    type Frame = crate::wire::WireFrame;
116    type Decoder<'id> = EmbeddingsDecoder;
117    type Reassembler = crate::wire::document::Unreassembled;
118
119    fn describe(&self) -> Descriptor<'_> {
120        Descriptor::new(PROVIDER_NAME)
121            .model(self.model.as_str())
122            .capabilities(Capabilities::embedding(MAX_DOCUMENTS, self.ndims))
123    }
124
125    fn encode(&self, texts: Vec<String>, _mode: Mode) -> Result<Encoded, EncodeError> {
126        let body = serde_json::json!({
127            "model": self.model,
128            "texts": texts,
129            "input_type": self.input_type,
130        });
131        let request = self
132            .provider
133            .post("/v1/embed")
134            .body(Body::Bytes(serde_json::to_vec(&body)?))?;
135        Ok(Encoded::new(request, Framing::Whole))
136    }
137
138    fn decoder<'id>(&self) -> Self::Decoder<'id> {
139        EmbeddingsDecoder
140    }
141}
142
143/// Decodes one `/v1/embed` reply for texts.
144pub struct EmbeddingsDecoder;
145
146impl<'id> Decoder<'id, Embedding> for EmbeddingsDecoder {
147    /// The vectors, or the error envelope Cohere can answer a **200** with.
148    type Event = Result<EmbeddingResponse, String>;
149
150    fn classify(&self, frame: WireFrame) -> WireEvent<Self::Event> {
151        classify_reply_or_message_envelope(&frame.as_str(), "embeddings")
152    }
153
154    fn decode(
155        &mut self,
156        reply: Self::Event,
157        out: Out<'id, Embedding>,
158    ) -> Result<Flow, ProviderError> {
159        let reply = reply.map_err(ProviderError::from_provider_body)?;
160        let usage = reply
161            .meta
162            .as_ref()
163            .map(|meta| meta.billed_units.to_usage())
164            .unwrap_or_default();
165        let vectors = reply
166            .embeddings
167            .into_iter()
168            .map(|vector| vector.into_iter().filter_map(|n| n.as_f64()).collect());
169        // Cohere's `/v1/embed` reply names no model.
170        Ok(out.end(crate::embeddings::EmbeddingResponse {
171            response_id: Some(reply.id),
172            usage,
173            ..crate::embeddings::EmbeddingResponse::from_vectors(vectors)
174        }))
175    }
176}
177
178/// The image-embedding wire: `POST /v1/embed`, one image per request.
179///
180/// Cohere Embed v3 accepts a single image per call, so the wire's batch limit
181/// is one image.
182#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
183pub struct ImageEmbeddings {
184    /// The provider this wire speaks to.
185    pub provider: CohereConfig,
186}
187
188impl Wire for ImageEmbeddings {
189    type Op = ImageEmbedding;
190    type Payload = crate::wire::Encoded;
191    type Frame = crate::wire::WireFrame;
192    type Decoder<'id> = ImageEmbeddingsDecoder;
193    type Reassembler = crate::wire::document::Unreassembled;
194
195    fn describe(&self) -> Descriptor<'_> {
196        Descriptor::new(PROVIDER_NAME)
197            .model(EMBED_ENGLISH_V3)
198            .capabilities(Capabilities::embedding(1, IMAGE_NDIMS))
199    }
200
201    fn encode(&self, images: Vec<Vec<u8>>, _mode: Mode) -> Result<Encoded, EncodeError> {
202        // The wire's batch limit is one image: a caller splits larger
203        // batches into one call per image.
204        let [image] = images.as_slice() else {
205            return Err(EncodeError::request(format!(
206                "Cohere embeds one image per request, not {}",
207                images.len()
208            )));
209        };
210        let media_type = validate_image(image)?;
211        let body = serde_json::json!({
212            "model": EMBED_ENGLISH_V3,
213            "images": [image_data_url(image, media_type)],
214            "input_type": "image",
215            "embedding_types": ["float"],
216        });
217        let request = self
218            .provider
219            .post("/v1/embed")
220            .body(Body::Bytes(serde_json::to_vec(&body)?))?;
221        Ok(Encoded::new(request, Framing::Whole))
222    }
223
224    fn decoder<'id>(&self) -> Self::Decoder<'id> {
225        ImageEmbeddingsDecoder
226    }
227}
228
229/// Decodes one `/v1/embed` reply for a single image.
230pub struct ImageEmbeddingsDecoder;
231
232impl<'id> Decoder<'id, ImageEmbedding> for ImageEmbeddingsDecoder {
233    /// The vector, or the error envelope Cohere can answer a **200** with.
234    type Event = Result<ImageEmbeddingResponse, String>;
235
236    fn classify(&self, frame: WireFrame) -> WireEvent<Self::Event> {
237        classify_reply_or_message_envelope(&frame.as_str(), "embeddings")
238    }
239
240    fn decode(
241        &mut self,
242        reply: Self::Event,
243        out: Out<'id, ImageEmbedding>,
244    ) -> Result<Flow, ProviderError> {
245        let reply = reply.map_err(ProviderError::from_provider_body)?;
246        // Each request carries one image, so any other vector count is invalid.
247        let [vector] = reply.embeddings.values.as_slice() else {
248            return Err(ProviderError::Response(format!(
249                "Expected 1 image embedding, got {}",
250                reply.embeddings.values.len()
251            )));
252        };
253        let usage = reply
254            .meta
255            .as_ref()
256            .map(|meta| meta.billed_units.to_usage())
257            .unwrap_or_default();
258        // The fold names the input: an image has no text, and its bytes
259        // must never travel back in a response.
260        let vector = vector.iter().filter_map(|n| n.as_f64()).collect();
261        Ok(out.end(crate::embeddings::EmbeddingResponse {
262            usage,
263            response_id: reply.id,
264            ..crate::embeddings::EmbeddingResponse::from_vectors([vector])
265        }))
266    }
267}
268
269const MAX_IMAGE_BYTES: usize = 5_000_000;
270
271#[derive(Debug, Clone, Serialize, Deserialize)]
272pub struct EmbeddingResponse {
273    #[serde(default)]
274    pub response_type: Option<String>,
275    pub id: String,
276    pub embeddings: Vec<Vec<serde_json::Number>>,
277    pub texts: Vec<String>,
278    #[serde(default)]
279    pub meta: Option<Meta>,
280}
281
282#[derive(Debug, Clone, Serialize, Deserialize)]
283pub struct Meta {
284    pub api_version: ApiVersion,
285    pub billed_units: BilledUnits,
286    #[serde(default)]
287    pub warnings: Vec<String>,
288}
289
290#[derive(Debug, Clone, Serialize, Deserialize)]
291pub struct ApiVersion {
292    pub version: String,
293    #[serde(default)]
294    pub is_deprecated: Option<bool>,
295    #[serde(default)]
296    pub is_experimental: Option<bool>,
297}
298
299/// Cohere's `meta.billed_units`. Token counters are absent when the request
300/// was not billed in tokens (image embeds bill `images`), so they are
301/// `Option`; the non-token counters default to zero.
302#[derive(Debug, Clone, Serialize, Deserialize)]
303pub struct BilledUnits {
304    #[serde(skip_serializing_if = "Option::is_none")]
305    pub input_tokens: Option<u32>,
306    #[serde(skip_serializing_if = "Option::is_none")]
307    pub output_tokens: Option<u32>,
308    #[serde(default)]
309    pub search_units: u32,
310    #[serde(default)]
311    pub classifications: u32,
312    #[serde(default)]
313    pub images: u32,
314}
315
316impl std::fmt::Display for BilledUnits {
317    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
318        write!(
319            f,
320            "Input tokens: {}\nOutput tokens: {}\nSearch units: {}\nClassifications: {}",
321            self.input_tokens.unwrap_or(0),
322            self.output_tokens.unwrap_or(0),
323            self.search_units,
324            self.classifications
325        )?;
326        if self.images > 0 {
327            write!(f, "\nImages: {}", self.images)?;
328        }
329        Ok(())
330    }
331}
332
333/// One Cohere `/v1/embed` answer for a single image, one per input image:
334/// Cohere Embed v3 accepts a single image per call.
335#[derive(Debug, Clone, Serialize, Deserialize)]
336pub struct ImageEmbeddingResponse {
337    #[serde(default)]
338    pub id: Option<String>,
339    pub embeddings: FloatEmbeddings,
340    #[serde(default)]
341    pub meta: Option<Meta>,
342}
343
344#[derive(Debug, Clone, Serialize, Deserialize)]
345pub struct FloatEmbeddings {
346    #[serde(rename = "float")]
347    pub values: Vec<Vec<serde_json::Number>>,
348}
349
350impl BilledUnits {
351    /// Maps the billed token counters straight through; `total_tokens` is
352    /// the sum of whichever counters Cohere sent (an embed bills input only,
353    /// so it reports `input_tokens` and `total_tokens`, no `output_tokens`).
354    pub(super) fn to_usage(&self) -> crate::completion::Usage {
355        let input_tokens = self.input_tokens.map(u64::from);
356        let output_tokens = self.output_tokens.map(u64::from);
357        let total_tokens = match (input_tokens, output_tokens) {
358            (None, None) => None,
359            (input, output) => Some(input.unwrap_or(0) + output.unwrap_or(0)),
360        };
361        crate::completion::Usage {
362            input_tokens,
363            output_tokens,
364            total_tokens,
365            ..Default::default()
366        }
367    }
368}
369
370#[derive(Debug, thiserror::Error)]
371pub(super) enum ImageInputError {
372    #[error("Cohere image embeddings support PNG, JPEG, WebP, or GIF file bytes")]
373    UnsupportedFormat,
374    #[error("Cohere image embeddings accept at most 5 MB per image; received {actual_bytes} bytes")]
375    TooLarge { actual_bytes: usize },
376}
377
378/// Detect an accepted image media type. Returns a document error for unsupported
379/// formats or inputs exceeding 5,000,000 bytes.
380pub(super) fn validate_image(bytes: &[u8]) -> Result<&'static str, EncodeError> {
381    if bytes.len() > MAX_IMAGE_BYTES {
382        return Err(EncodeError::request(ImageInputError::TooLarge {
383            actual_bytes: bytes.len(),
384        }));
385    }
386
387    crate::embeddings::image_media_type(bytes)
388        .ok_or_else(|| EncodeError::request(ImageInputError::UnsupportedFormat))
389}
390
391pub(super) fn image_data_url(bytes: &[u8], media_type: &str) -> String {
392    format!("data:{media_type};base64,{}", STANDARD.encode(bytes))
393}
394
395#[cfg(test)]
396mod tests;