Skip to main content

rig_core/providers/cohere/
wire.rs

1//! Cohere configuration and chat, text-embedding, and image-embedding wires.
2//!
3//! ```no_run
4//! use rig_core::providers::cohere::{Cohere, EMBED_V4};
5//! let wire = Cohere::from_env()?.embedding(EMBED_V4, None);
6//! # Ok::<(), Box<dyn std::error::Error>>(())
7//! ```
8
9use crate::client::env::{self, EnvError};
10use crate::completion::CompletionRequest;
11use crate::embeddings::Embedding as Vector;
12use crate::error::EncodeError;
13use crate::error::ProviderError;
14use crate::json_utils;
15use crate::operation::{Completion, Embedding, ImageEmbedding};
16use crate::wire::Flow;
17use crate::wire::{
18    Body, Capabilities, Decoder, Descriptor, Encoded, Framing, Mode, Out, Secret, Wire, WireEvent,
19    WireFrame,
20};
21use serde::{Deserialize, Serialize};
22
23use super::completion::{CohereCompletionRequest, PROVIDER_NAME};
24use crate::message::Issuer;
25
26/// The issuer of Cohere's reasoning, which is the only reasoning it replays.
27pub(crate) const ISSUER: Issuer = Issuer::from_static(PROVIDER_NAME);
28use super::embeddings::{
29    EmbeddingResponse as CohereEmbeddingResponse, ErrorEnvelope as CohereErrorEnvelope,
30    ImageEmbeddingResponse as CohereImageEmbeddingResponse, image_data_url, validate_image,
31};
32use super::streaming::ChatDecoder;
33
34/// Cohere's API root.
35const BASE_URL: &str = "https://api.cohere.ai";
36
37/// The environment variable carrying the API key.
38const API_KEY_ENV: &str = "COHERE_API_KEY";
39
40/// The settings of a Cohere provider: serializable, and the credential is
41/// never serialized. [`connect`](Self::connect) puts it on a transport as a
42/// [`Cohere`](super::Cohere) client.
43#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
44pub struct CohereConfig {
45    /// The API key, sent as `Authorization: Bearer`.
46    pub api_key: Secret,
47    /// The API root, without a trailing slash.
48    pub base_url: String,
49}
50
51impl CohereConfig {
52    /// Cohere with default settings.
53    pub fn new(api_key: impl Into<Secret>) -> Self {
54        Self {
55            api_key: api_key.into(),
56            base_url: BASE_URL.to_owned(),
57        }
58    }
59
60    /// Cohere from `COHERE_API_KEY`.
61    pub fn from_env() -> Result<Self, EnvError> {
62        Ok(Self::new(env::required(API_KEY_ENV)?))
63    }
64
65    /// Point the wires at another API root.
66    pub fn with_base_url(mut self, base_url: impl AsRef<str>) -> Self {
67        self.base_url = base_url.as_ref().trim_end_matches('/').to_owned();
68        self
69    }
70
71    /// The chat wire for `model`.
72    pub(crate) fn completion(&self, model: impl Into<String>) -> Chat {
73        Chat {
74            provider: self.clone(),
75            model: model.into(),
76        }
77    }
78
79    /// Build a text-embedding wire reporting the supplied or known model width,
80    /// or zero if unknown. This width is metadata, not a request parameter.
81    pub(crate) fn embedding(&self, model: impl Into<String>, ndims: Option<usize>) -> Embeddings {
82        let model = model.into();
83        let ndims = ndims
84            .or_else(|| super::model_dimensions_from_identifier(&model))
85            .unwrap_or_default();
86        Embeddings {
87            provider: self.clone(),
88            model,
89            ndims,
90            input_type: DEFAULT_INPUT_TYPE.to_owned(),
91        }
92    }
93
94    /// The image-embedding wire.
95    ///
96    /// Cohere Embed v3 embeds images with one fixed model, so this wire
97    /// names no model.
98    pub(crate) fn image_embedding(&self) -> ImageEmbeddings {
99        ImageEmbeddings {
100            provider: self.clone(),
101        }
102    }
103
104    /// One request, authenticated and typed as JSON.
105    fn post(&self, path: &str) -> http::request::Builder {
106        http::Request::post(format!("{}{path}", self.base_url))
107            .header(http::header::CONTENT_TYPE, "application/json")
108            .header(
109                http::header::AUTHORIZATION,
110                format!("Bearer {}", self.api_key.expose()),
111            )
112    }
113}
114
115/// The chat wire: `POST /v2/chat`, SSE when streamed.
116#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
117pub struct Chat {
118    /// The provider this wire speaks to.
119    pub provider: CohereConfig,
120    /// The model to address.
121    pub model: String,
122}
123
124impl Wire for Chat {
125    type Op = Completion;
126    type Payload = crate::wire::Encoded;
127    type Frame = crate::wire::WireFrame;
128    type Decoder<'id> = ChatDecoder<'id>;
129
130    fn describe(&self) -> Descriptor<'_> {
131        Descriptor::new(PROVIDER_NAME).model(self.model.as_str())
132    }
133
134    fn encode(&self, request: CompletionRequest, mode: Mode) -> Result<Encoded, EncodeError> {
135        let request = request.replayable_to(&[ISSUER])?;
136        let mut body = CohereCompletionRequest::try_from((self.model.as_str(), request))?;
137        if mode == Mode::Streaming {
138            body.additional_params = Some(json_utils::merge(
139                body.additional_params
140                    .take()
141                    .unwrap_or_else(|| serde_json::json!({})),
142                serde_json::json!({ "stream": true }),
143            ));
144        }
145        crate::providers::internal::trace_json(
146            crate::providers::internal::LogTarget::Completions,
147            "Cohere completion request",
148            &body,
149        );
150        let request = self
151            .provider
152            .post("/v2/chat")
153            .body(Body::Bytes(serde_json::to_vec(&body)?))?;
154        // Cohere reports no request-id response header: its
155        // `x-debug-trace-id` is a debug trace handle, not a documented
156        // request id, so the normalized id stays unset by design.
157        Ok(Encoded::new(
158            request,
159            match mode {
160                Mode::Unary => Framing::Whole,
161                Mode::Streaming => Framing::Sse,
162            },
163        ))
164    }
165
166    fn decoder<'id>(&self) -> Self::Decoder<'id> {
167        ChatDecoder::default()
168    }
169}
170
171/// Default retrieval role for embeddings of stored document chunks.
172const DEFAULT_INPUT_TYPE: &str = "search_document";
173
174/// The most texts Cohere embeds in one `/v1/embed` call.
175const MAX_DOCUMENTS: usize = 96;
176
177/// The width Cohere's image embeddings come back at.
178const IMAGE_NDIMS: usize = 1_024;
179
180/// Recognized embedding and error-envelope markers, including errors on HTTP 200.
181const EMBED_REPLY_MARKERS: &[&str] = &["embeddings", "message"];
182
183/// The key that recognizes the error envelope on its own.
184const EMBED_ERROR_MARKERS: &[&str] = &["message"];
185
186/// One `/v1/embed` reply: the answer, or the error envelope Cohere can
187/// answer a **200** with instead.
188pub enum EmbedReply<T> {
189    /// The vectors Cohere returned.
190    Reply(T),
191    /// The provider's error envelope, verbatim.
192    Failure(String),
193}
194
195/// Decode embeddings, then an error envelope on failure. Retain the embedding
196/// diagnostic if neither shape decodes.
197fn classify_embed_reply<T>(data: &str) -> WireEvent<EmbedReply<T>>
198where
199    T: serde::de::DeserializeOwned,
200{
201    crate::providers::internal::wire::classify_or(
202        data,
203        |data| {
204            crate::providers::internal::wire::classify_marker_keyed_frame::<T>(
205                data,
206                EMBED_REPLY_MARKERS,
207            )
208            .map(EmbedReply::Reply)
209        },
210        |data| {
211            crate::providers::internal::wire::classify_marker_keyed_frame::<CohereErrorEnvelope>(
212                data,
213                EMBED_ERROR_MARKERS,
214            )
215            .map(|_| EmbedReply::Failure(data.to_owned()))
216        },
217    )
218}
219
220/// The text-embedding wire: `POST /v1/embed`.
221#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
222pub struct Embeddings {
223    /// The provider this wire speaks to.
224    pub provider: CohereConfig,
225    /// The model to address.
226    pub model: String,
227    /// The width this wire reports, from the caller or the model's published
228    /// dimensions. `0` means neither named one.
229    pub ndims: usize,
230    /// Cohere's retrieval prompt: `search_document` for stored chunks,
231    /// `search_query` for queries, `classification`, `clustering`.
232    pub input_type: String,
233}
234
235impl Embeddings {
236    /// Embed for a different purpose than storing chunks.
237    pub fn with_input_type(mut self, input_type: impl Into<String>) -> Self {
238        self.input_type = input_type.into();
239        self
240    }
241}
242
243impl Wire for Embeddings {
244    type Op = Embedding;
245    type Payload = crate::wire::Encoded;
246    type Frame = crate::wire::WireFrame;
247    type Decoder<'id> = EmbeddingsDecoder;
248
249    fn describe(&self) -> Descriptor<'_> {
250        Descriptor::new(PROVIDER_NAME)
251            .model(self.model.as_str())
252            .capabilities(Capabilities::embedding(MAX_DOCUMENTS, self.ndims))
253    }
254
255    fn encode(&self, texts: Vec<String>, _mode: Mode) -> Result<Encoded, EncodeError> {
256        let body = serde_json::json!({
257            "model": self.model,
258            "texts": texts,
259            "input_type": self.input_type,
260        });
261        let request = self
262            .provider
263            .post("/v1/embed")
264            .body(Body::Bytes(serde_json::to_vec(&body)?))?;
265        Ok(Encoded::new(request, Framing::Whole))
266    }
267
268    fn decoder<'id>(&self) -> Self::Decoder<'id> {
269        EmbeddingsDecoder
270    }
271}
272
273/// Decodes one `/v1/embed` reply for texts.
274pub struct EmbeddingsDecoder;
275
276impl<'id> Decoder<'id, Embedding> for EmbeddingsDecoder {
277    type Event = EmbedReply<CohereEmbeddingResponse>;
278
279    fn classify(&self, frame: WireFrame) -> WireEvent<Self::Event> {
280        classify_embed_reply(&frame.as_str())
281    }
282
283    fn decode(
284        &mut self,
285        reply: Self::Event,
286        out: Out<'id, Embedding>,
287    ) -> Result<Flow, ProviderError> {
288        let reply = match reply {
289            EmbedReply::Reply(reply) => reply,
290            // Preserve the error body so the driver can attach its HTTP status.
291            EmbedReply::Failure(body) => {
292                return Err(ProviderError::from_provider_body(body));
293            }
294        };
295        let usage = reply
296            .meta
297            .as_ref()
298            .map(|meta| meta.billed_units.to_usage())
299            .unwrap_or_default();
300        // The vectors only; the operation's fold pairs them with the texts
301        // that were sent, which no `/v1/embed` reply is trusted to echo
302        // back in order.
303        let vectors = reply
304            .embeddings
305            .into_iter()
306            .map(|vector| Vector {
307                document: String::new(),
308                vec: vector.into_iter().filter_map(|n| n.as_f64()).collect(),
309            })
310            .collect();
311        // Cohere's `/v1/embed` reply names no model.
312        Ok(out.end(crate::embeddings::EmbeddingResponse {
313            response_id: Some(reply.id),
314            usage,
315            ..crate::embeddings::EmbeddingResponse::new(vectors)
316        }))
317    }
318}
319
320/// The image-embedding wire: `POST /v1/embed`, one image per request.
321///
322/// Cohere Embed v3 accepts a single image per call, so the wire's batch limit
323/// is one image.
324#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
325pub struct ImageEmbeddings {
326    /// The provider this wire speaks to.
327    pub provider: CohereConfig,
328}
329
330impl Wire for ImageEmbeddings {
331    type Op = ImageEmbedding;
332    type Payload = crate::wire::Encoded;
333    type Frame = crate::wire::WireFrame;
334    type Decoder<'id> = ImageEmbeddingsDecoder;
335
336    fn describe(&self) -> Descriptor<'_> {
337        Descriptor::new(PROVIDER_NAME)
338            .model(super::EMBED_ENGLISH_V3)
339            .capabilities(Capabilities::embedding(1, IMAGE_NDIMS))
340    }
341
342    fn encode(&self, images: Vec<Vec<u8>>, _mode: Mode) -> Result<Encoded, EncodeError> {
343        // The wire's batch limit is one image: a caller splits larger
344        // batches into one call per image.
345        let [image] = images.as_slice() else {
346            return Err(EncodeError::request(format!(
347                "Cohere embeds one image per request, not {}",
348                images.len()
349            )));
350        };
351        let media_type = validate_image(image)?;
352        let body = serde_json::json!({
353            "model": super::EMBED_ENGLISH_V3,
354            "images": [image_data_url(image, media_type)],
355            "input_type": "image",
356            "embedding_types": ["float"],
357        });
358        let request = self
359            .provider
360            .post("/v1/embed")
361            .body(Body::Bytes(serde_json::to_vec(&body)?))?;
362        Ok(Encoded::new(request, Framing::Whole))
363    }
364
365    fn decoder<'id>(&self) -> Self::Decoder<'id> {
366        ImageEmbeddingsDecoder
367    }
368}
369
370/// Decodes one `/v1/embed` reply for a single image.
371pub struct ImageEmbeddingsDecoder;
372
373impl<'id> Decoder<'id, ImageEmbedding> for ImageEmbeddingsDecoder {
374    type Event = EmbedReply<CohereImageEmbeddingResponse>;
375
376    fn classify(&self, frame: WireFrame) -> WireEvent<Self::Event> {
377        classify_embed_reply(&frame.as_str())
378    }
379
380    fn decode(
381        &mut self,
382        reply: Self::Event,
383        out: Out<'id, ImageEmbedding>,
384    ) -> Result<Flow, ProviderError> {
385        let reply = match reply {
386            EmbedReply::Reply(reply) => reply,
387            // Same 200-with-an-envelope reply as the text route: the body
388            // verbatim, with the driver stamping the status.
389            EmbedReply::Failure(body) => {
390                return Err(ProviderError::from_provider_body(body));
391            }
392        };
393        // Each request carries one image, so any other vector count is invalid.
394        let [vector] = reply.embeddings.values.as_slice() else {
395            return Err(ProviderError::Response(format!(
396                "Expected 1 image embedding, got {}",
397                reply.embeddings.values.len()
398            )));
399        };
400        let usage = reply
401            .meta
402            .as_ref()
403            .map(|meta| meta.billed_units.to_usage())
404            .unwrap_or_default();
405        let vector = Vector {
406            // The fold names the input: an image has no text, and its bytes
407            // must never travel back in a response.
408            document: String::new(),
409            vec: vector.iter().filter_map(|n| n.as_f64()).collect(),
410        };
411        Ok(out.end(crate::embeddings::ImageEmbeddingResponse {
412            usage,
413            response_id: reply.id,
414            ..crate::embeddings::ImageEmbeddingResponse::new(vec![vector])
415        }))
416    }
417}
418
419#[cfg(test)]
420mod tests;