rig_core/providers/cohere/
embeddings.rs1use 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
23pub const EMBED_V4: &str = "embed-v4.0";
25pub const EMBED_ENGLISH_V3: &str = "embed-english-v3.0";
27pub const EMBED_ENGLISH_LIGHT_V3: &str = "embed-english-light-v3.0";
29pub const EMBED_MULTILINGUAL_V3: &str = "embed-multilingual-v3.0";
31pub 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 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 pub(crate) fn image_embedding(&self) -> ImageEmbeddings {
64 ImageEmbeddings {
65 provider: self.clone(),
66 }
67 }
68
69 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
80const DEFAULT_INPUT_TYPE: &str = "search_document";
82
83const MAX_DOCUMENTS: usize = 96;
85
86const IMAGE_NDIMS: usize = 1_024;
88
89#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
91pub struct Embeddings {
92 pub provider: CohereConfig,
94 pub model: String,
96 pub ndims: usize,
99 pub input_type: String,
102}
103
104impl Embeddings {
105 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
143pub struct EmbeddingsDecoder;
145
146impl<'id> Decoder<'id, Embedding> for EmbeddingsDecoder {
147 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 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#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
183pub struct ImageEmbeddings {
184 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 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
229pub struct ImageEmbeddingsDecoder;
231
232impl<'id> Decoder<'id, ImageEmbedding> for ImageEmbeddingsDecoder {
233 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 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 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#[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#[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 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
378pub(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;