Skip to main content

rig_core/providers/cohere/
chat.rs

1//! Cohere's native chat API, `POST /v2/chat`: the request built as JSON,
2//! with the request's documents sent as Cohere `documents`. Its replies
3//! are read by [`ChatDecoder`].
4//!
5//! ```
6//! use rig_core::providers::cohere::{COMMAND_A_03_2025, CohereConfig, NativeChat};
7//!
8//! let wire = NativeChat::new(CohereConfig::new("key"), COMMAND_A_03_2025);
9//! assert_eq!(wire.model, COMMAND_A_03_2025);
10//! ```
11
12use std::collections::BTreeMap;
13
14use serde::{Deserialize, Serialize};
15use serde_json::{Map, Value, json};
16
17use super::CohereConfig;
18use super::streaming::ChatDecoder;
19use crate::completion::options::{BaseInput, FinalBody, RawAt, request_params};
20use crate::completion::{CompletionRequest, Document, ProviderCapabilities, Replay};
21use crate::error::EncodeError;
22use crate::json_utils::Lenient;
23use crate::message::{
24    AssistantContent, AssistantMessage, DocumentMediaType, DocumentSourceKind as Source, Message,
25    MimeType, ToolChoice, ToolResult, ToolResultContent, UserContent,
26};
27use crate::operation::Completion;
28use crate::providers::internal::wire_ids::WireIds;
29use crate::wire::{Capabilities, Descriptor, Encoded, Framing, Mode, Wire};
30
31/// Where the native chat endpoint sits under the API root.
32const CHAT_PATH: &str = "/v2/chat";
33
34/// The replay API of turns made on the native chat endpoint.
35pub(crate) const API: &str = "cohere.chat";
36
37/// The native chat wire: a provider configuration, a model, and the
38/// per-turn options the endpoint takes.
39#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
40pub struct NativeChat {
41    /// Which provider, and how to reach it.
42    pub provider: CohereConfig,
43    /// The model this wire addresses.
44    pub model: String,
45    /// Whether requests ask Cohere to hold tool calls to their schemas
46    /// (`strict_tools`).
47    pub strict_tools: bool,
48}
49
50impl NativeChat {
51    /// The wire for `model` on `provider`, with every option off.
52    pub fn new(provider: CohereConfig, model: impl Into<String>) -> Self {
53        Self {
54            provider,
55            model: model.into(),
56            strict_tools: false,
57        }
58    }
59
60    /// Ask Cohere to hold every tool call to its tool's schema.
61    pub fn with_strict_tools(mut self) -> Self {
62        self.strict_tools = true;
63        self
64    }
65
66    /// The request body: the wire's encoding of `request`, then the mapped
67    /// options, then `additional_params`, merged key by key.
68    fn body(&self, request: &CompletionRequest, mode: Mode) -> Result<FinalBody, EncodeError> {
69        request_params(
70            self,
71            request,
72            |input| self.base(request, mode, input),
73            RawAt::Top,
74            &[],
75        )
76    }
77
78    /// The wire's own encoding of `request`.
79    fn base(
80        &self,
81        request: &CompletionRequest,
82        mode: Mode,
83        input: &mut BaseInput<'_>,
84    ) -> Result<Map<String, Value>, EncodeError> {
85        let model = request.model.clone().unwrap_or_else(|| self.model.clone());
86        let messages = self.messages(&request.chat_history, &model)?;
87        let mut tools: Vec<Value> = request
88            .tools
89            .iter()
90            .filter(|tool| match &request.tool_choice {
91                // Cohere cannot name the tool it must call: only the named
92                // tools are offered, and one of them is required.
93                Some(ToolChoice::Specific { function_names }) => {
94                    function_names.contains(&tool.name)
95                }
96                _ => true,
97            })
98            .map(|tool| {
99                json!({"type": "function", "function": {
100                    "name": tool.name,
101                    "description": tool.description,
102                    "parameters": tool.parameters,
103                }})
104            })
105            .collect();
106        tools.extend(input.raw_tools()?);
107        let tool_choice = match &request.tool_choice {
108            None | Some(ToolChoice::Auto) => None,
109            Some(ToolChoice::None) => Some("NONE"),
110            Some(ToolChoice::Required | ToolChoice::Specific { .. }) => Some("REQUIRED"),
111        }
112        .filter(|_| !tools.is_empty());
113        let documents: Vec<Value> = request
114            .documents
115            .iter()
116            .enumerate()
117            .map(|(position, document)| document_value(position, document))
118            .collect();
119        let response_format = request
120            .output_schema
121            .clone()
122            .map(|schema| json!({"type": "json_object", "schema": schema.to_value()}));
123        let fields = [
124            ("model", Some(Value::String(model))),
125            ("messages", Some(Value::Array(messages))),
126            (
127                "documents",
128                (!documents.is_empty()).then_some(Value::Array(documents)),
129            ),
130            ("tools", (!tools.is_empty()).then_some(Value::Array(tools))),
131            ("tool_choice", tool_choice.map(Value::from)),
132            (
133                "strict_tools",
134                self.strict_tools.then_some(Value::Bool(true)),
135            ),
136            ("response_format", response_format),
137            ("temperature", request.temperature.map(Value::from)),
138            ("max_tokens", request.max_tokens.map(Value::from)),
139            (
140                "stream",
141                (mode == Mode::Streaming).then_some(Value::Bool(true)),
142            ),
143        ];
144        Ok(fields
145            .into_iter()
146            .filter_map(|(key, value)| Some((key.to_owned(), value?)))
147            .collect())
148    }
149
150    /// The history as Cohere messages, each call and result spelled by one
151    /// [`WireIds`].
152    fn messages(&self, history: &[Message], model: &str) -> Result<Vec<Value>, EncodeError> {
153        let ids = WireIds::for_target(history, self, model);
154        let mut messages = Vec::new();
155        for message in history {
156            match message {
157                Message::System { content } => {
158                    messages.push(json!({"role": "system", "content": content}));
159                }
160                Message::User { content } => {
161                    let mut parts = Vec::new();
162                    for part in content {
163                        if let UserContent::ToolResult(result) = part {
164                            user_message(&mut messages, &mut parts);
165                            messages.push(tool_message(result, &ids));
166                        } else {
167                            parts.push(user_part(part)?);
168                        }
169                    }
170                    user_message(&mut messages, &mut parts);
171                }
172                Message::Assistant(turn) => messages.extend(self.assistant(turn, &ids)),
173            }
174        }
175        if messages.is_empty() {
176            return Err(EncodeError::request(
177                "Cohere chat request has no messages after conversion",
178            ));
179        }
180        Ok(messages)
181    }
182
183    /// One assistant turn: text, thinking and unknown parts as content
184    /// parts, a tool plan as `tool_plan`, each call under `tool_calls`, and
185    /// the citations each block's item holds, pointed at the part they
186    /// cite. A current item goes back as it came, less its citations. A
187    /// turn with nothing to send is `None`.
188    fn assistant(&self, turn: &AssistantMessage, ids: &WireIds) -> Option<Value> {
189        let (mut content, mut plan, mut calls, mut citations) =
190            (Vec::new(), String::new(), Vec::new(), Vec::new());
191        for block in &turn.content {
192            let replay = block.replay(self, ids);
193            let kind = replay_kind(&replay);
194            let mut item = match replay {
195                Replay::Item(item) => match item.into_owned() {
196                    Value::Object(item) => item,
197                    _ => Map::new(),
198                },
199                Replay::Identity(_) | Replay::Rebuild => Map::new(),
200            };
201            let cited = match item.shift_remove("citations") {
202                Some(Value::Array(cited)) => cited,
203                _ => Vec::new(),
204            };
205            let part = match block {
206                AssistantContent::Text(text) if !text.text.is_empty() => {
207                    json!({"type": "text", "text": text.text})
208                }
209                AssistantContent::Reasoning(reasoning) if kind.as_deref() == Some(PLAN) => {
210                    citations.extend(pointed(cited, None));
211                    plan.push_str(&reasoning.text);
212                    continue;
213                }
214                AssistantContent::Reasoning(reasoning) if !reasoning.text.is_empty() => {
215                    json!({"type": "thinking", "thinking": reasoning.text})
216                }
217                AssistantContent::Opaque(opaque)
218                    if opaque.replay && opaque.item.get("type").is_some() =>
219                {
220                    content.push(opaque.item.clone());
221                    continue;
222                }
223                AssistantContent::ToolCall(call) => {
224                    item.entry("type")
225                        .or_insert_with(|| Value::from("function"));
226                    item.insert("id".to_owned(), Value::String(ids.spell(&call.id)));
227                    let function = item
228                        .entry("function")
229                        .or_insert_with(|| Value::Object(Map::new()));
230                    if !function.is_object() {
231                        *function = Value::Object(Map::new());
232                    }
233                    if let Value::Object(function) = function {
234                        function
235                            .insert("name".to_owned(), Value::from(call.function.name.as_str()));
236                        function.insert(
237                            "arguments".to_owned(),
238                            Value::String(call.function.arguments_value().to_string()),
239                        );
240                    }
241                    calls.push(Value::Object(item));
242                    continue;
243                }
244                AssistantContent::Text(_)
245                | AssistantContent::Reasoning(_)
246                | AssistantContent::Image(_)
247                | AssistantContent::Opaque(_) => continue,
248            };
249            citations.extend(pointed(cited, Some(content.len())));
250            content.push(if item.is_empty() {
251                part
252            } else {
253                Value::Object(item)
254            });
255        }
256        if content.is_empty() && calls.is_empty() {
257            return None;
258        }
259        let fields = [
260            ("role", Some(Value::from("assistant"))),
261            (
262                "content",
263                (!content.is_empty()).then_some(Value::Array(content)),
264            ),
265            (
266                "tool_plan",
267                (!plan.is_empty()).then_some(Value::String(plan)),
268            ),
269            (
270                "tool_calls",
271                (!calls.is_empty()).then_some(Value::Array(calls)),
272            ),
273            (
274                "citations",
275                (!citations.is_empty()).then_some(Value::Array(citations)),
276            ),
277        ];
278        Some(Value::Object(
279            fields
280                .into_iter()
281                .filter_map(|(key, value)| Some((key.to_owned(), value?)))
282                .collect(),
283        ))
284    }
285}
286
287/// The item kind a tool plan's reasoning block holds.
288pub(crate) const PLAN: &str = "tool_plan";
289
290/// The kind of item replay hands the encoder for a block: the current
291/// item's `type`, or the one an edited block kept.
292fn replay_kind(replay: &Replay<'_>) -> Option<String> {
293    match replay {
294        Replay::Item(item) => item.str("type").map(str::to_owned),
295        Replay::Identity(identity) => identity
296            .get("type")
297            .and_then(Value::as_str)
298            .map(str::to_owned),
299        Replay::Rebuild => None,
300    }
301}
302
303/// The kind of item `block` replays as on `target`.
304fn kind(block: &AssistantContent, target: &NativeChat) -> Option<String> {
305    replay_kind(&block.replay(target, &WireIds::default()))
306}
307
308/// `citations` as a request sends them back: each at `content_index`, the
309/// position of the part it cites, or with none for a tool plan's.
310fn pointed(citations: Vec<Value>, content_index: Option<usize>) -> Vec<Value> {
311    citations
312        .into_iter()
313        .map(|mut citation| {
314            if let (Some(citation), Some(index)) = (citation.as_object_mut(), content_index) {
315                citation.insert("content_index".to_owned(), Value::from(index));
316            }
317            citation
318        })
319        .collect()
320}
321
322/// `document` as a Cohere document: its metadata and `text` under `data`,
323/// keys sorted so the same document always sends the same bytes, under its
324/// id, or `doc_<position>` when it has none.
325fn document_value(position: usize, document: &Document) -> Value {
326    let mut data: BTreeMap<&str, &str> = document
327        .additional_props
328        .iter()
329        .map(|(key, value)| (key.as_str(), value.as_str()))
330        .collect();
331    data.insert("text", &document.text);
332    let id = if document.id.is_empty() {
333        format!("doc_{position}")
334    } else {
335        document.id.clone()
336    };
337    json!({"id": id, "data": data})
338}
339
340/// One user content part as Cohere carries it: text, a string document's
341/// text, or an image by URL. Replay leaves no other part.
342fn user_part(part: &UserContent) -> Result<Value, EncodeError> {
343    Ok(match part {
344        UserContent::Text(text) => json!({"type": "text", "text": text.text}),
345        UserContent::Image(image) => {
346            let mime = image.media_type.as_ref().map(MimeType::to_mime_type);
347            let url = match (&image.data, mime) {
348                (Source::Url(url), _) => url.clone(),
349                (Source::Base64(data), Some(mime)) => format!("data:{mime};base64,{data}"),
350                _ => return Err(unsendable("an image")),
351            };
352            json!({"type": "image_url", "image_url": {"url": url}})
353        }
354        UserContent::Document(document) => match &document.data {
355            Source::String(text) if document.media_type != Some(DocumentMediaType::PDF) => {
356                json!({"type": "text", "text": text})
357            }
358            _ => return Err(unsendable("a document")),
359        },
360        UserContent::Audio(_) => return Err(unsendable("audio")),
361        UserContent::Video(_) => return Err(unsendable("a video")),
362        UserContent::ToolResult(_) => return Err(unsendable("a tool result as a content part")),
363    })
364}
365
366/// The error for content Cohere chat cannot carry, which replay replaces
367/// before a request is encoded.
368fn unsendable(what: &str) -> EncodeError {
369    EncodeError::request(format!("Cohere chat cannot carry {what} in this form"))
370}
371
372/// Push the user parts gathered so far as one message.
373fn user_message(messages: &mut Vec<Value>, parts: &mut Vec<Value>) {
374    if parts.is_empty() {
375        return;
376    }
377    messages.push(json!({"role": "user", "content": std::mem::take(parts)}));
378}
379
380/// A result as the `tool` message that answers its call, its parts as text.
381fn tool_message(result: &ToolResult, ids: &WireIds) -> Value {
382    let content: Vec<Value> = result
383        .content
384        .iter()
385        .filter_map(|part| match part {
386            ToolResultContent::Text(text) => Some(text.text.clone()),
387            ToolResultContent::Json { value } => Some(value.to_string()),
388            ToolResultContent::Image(_) => None,
389        })
390        .map(|text| json!({"type": "text", "text": text}))
391        .collect();
392    json!({"role": "tool", "tool_call_id": ids.spell(&result.call), "content": content})
393}
394
395impl Wire for NativeChat {
396    type Op = Completion;
397    type Payload = Encoded;
398    type Frame = crate::wire::WireFrame;
399    type Decoder<'id> = ChatDecoder;
400    type Reassembler = super::streaming::document::ChatResponse;
401
402    fn describe(&self) -> Descriptor<'_> {
403        Descriptor::new(super::PROVIDER_NAME)
404            .model(self.model.as_str())
405            .capabilities(Capabilities::completion(
406                ProviderCapabilities::default().with_native_output_tool_composition(true),
407            ))
408            .replay(self)
409    }
410
411    fn encode(&self, request: CompletionRequest, mode: Mode) -> Result<Encoded, EncodeError> {
412        let body = self.body(&request, mode)?;
413        crate::providers::internal::trace_json(
414            crate::providers::internal::LogTarget::Completions,
415            "Cohere chat request",
416            &body,
417        );
418        let request = self.provider.post(CHAT_PATH).body(body.into_body())?;
419        let framing = match mode {
420            Mode::Streaming => Framing::Sse,
421            Mode::Unary => Framing::Whole,
422        };
423        Ok(Encoded::new(request, framing)
424            .with_request_id_header(Some(REQUEST_ID_HEADER))
425            .with_projection(ChatDecoder::project)
426            .with_route(Some(CHAT_PATH)))
427    }
428
429    fn decoder<'id>(&self) -> Self::Decoder<'id> {
430        ChatDecoder::default()
431    }
432}
433
434/// The reply header carrying Cohere's request id.
435const REQUEST_ID_HEADER: &str = "x-debug-trace-id";
436
437impl crate::completion::ReplayTarget for NativeChat {
438    /// Section 6.5 of the typed-options design, for the native API.
439    fn map_options(
440        &self,
441        request: &CompletionRequest,
442        fields: crate::completion::options::OptionFields<'_>,
443    ) -> crate::completion::options::OptionMap {
444        use crate::completion::options::{Mapping, OptionFields, OptionMap};
445        use crate::completion::{CacheRetention, Effort, Reasoning};
446        let OptionFields {
447            reasoning,
448            cache,
449            service_tier,
450            verbosity,
451            parallel_tool_calls,
452            top_p,
453            seed,
454            stop,
455        } = fields;
456        let model = request.model.as_deref().unwrap_or(&self.model);
457        // A model that thinks does so by default. An id the catalog does
458        // not list thinks when its name says `reasoning`.
459        let thinks = super::thinks(model);
460        let reasons = thinks.unwrap_or_else(|| model.contains("reasoning"));
461        const NO_FIELD: &str = "Cohere's chat API has no such field";
462        OptionMap {
463            reasoning: Mapping::of(reasoning, |reasoning| match reasoning {
464                Reasoning::Off if reasons => {
465                    Mapping::Send(json!({"thinking": {"type": "disabled"}}))
466                }
467                Reasoning::Off => Mapping::Omit("the model does not think"),
468                Reasoning::Effort(_) | Reasoning::Budget { .. } if thinks == Some(false) => {
469                    Mapping::unsupported("the model does not think")
470                }
471                Reasoning::Effort(Effort::High) => {
472                    Mapping::Send(json!({"thinking": {"type": "enabled"}}))
473                }
474                Reasoning::Effort(effort) => Mapping::unsupported(format!(
475                    "Cohere takes thinking on or a token budget, not `{}`",
476                    effort.as_str()
477                )),
478                Reasoning::Budget { tokens } => Mapping::Send(json!({
479                    "thinking": {"type": "enabled", "token_budget": tokens},
480                })),
481            }),
482            cache: Mapping::of(cache, |cache| match cache {
483                CacheRetention::None => Mapping::Omit("Cohere does not cache prompts"),
484                CacheRetention::Short | CacheRetention::Long => {
485                    Mapping::unsupported("Cohere has no prompt cache")
486                }
487            }),
488            service_tier: Mapping::of(service_tier, |_| {
489                Mapping::unsupported("Cohere's `priority` is a queue position, not a tier")
490            }),
491            verbosity: Mapping::of(verbosity, |_| Mapping::unsupported(NO_FIELD)),
492            parallel_tool_calls: Mapping::of(parallel_tool_calls, |_| {
493                Mapping::unsupported(NO_FIELD)
494            }),
495            top_p: Mapping::of(top_p, |top_p| {
496                if (0.01..=0.99).contains(&top_p) {
497                    Mapping::Send(json!({ "p": top_p }))
498                } else {
499                    Mapping::unsupported("Cohere takes `p` from 0.01 to 0.99")
500                }
501            }),
502            seed: Mapping::of(seed, |seed| Mapping::Send(json!({ "seed": seed }))),
503            stop: Mapping::of_stop(stop, |stop| match stop.len() {
504                0..=5 => Mapping::Send(json!({ "stop_sequences": stop })),
505                _ => Mapping::unsupported("Cohere takes at most 5 stop sequences"),
506            }),
507        }
508    }
509
510    fn api(&self) -> crate::message::Api {
511        crate::message::Api::from_static(API)
512    }
513
514    fn provider(&self) -> &str {
515        super::PROVIDER_NAME
516    }
517
518    fn model(&self) -> &str {
519        &self.model
520    }
521
522    /// Cohere's vision models read user images; no model reads images in
523    /// assistant turns or tool results.
524    fn accepts(&self, model: &str) -> crate::completion::Accepts {
525        crate::completion::Accepts {
526            user_images: crate::catalog::reads_images_or(
527                super::PROVIDER_NAME,
528                model,
529                super::reads_images,
530            ),
531            assistant_images: false,
532            tool_result_images: false,
533            tools: true,
534        }
535    }
536
537    /// The encoder carries a user image by URL or typed data, and a text
538    /// document as its text. It carries no other media.
539    fn encodes(&self, _model: &str, media: crate::completion::Media<'_>) -> bool {
540        use crate::completion::{Media, Place};
541        match media {
542            Media::Image(image, Place::User) => match &image.data {
543                Source::Url(_) => true,
544                Source::Base64(_) => image.media_type.is_some(),
545                _ => false,
546            },
547            Media::Document(document) => {
548                matches!(document.data, Source::String(_))
549                    && document.media_type != Some(DocumentMediaType::PDF)
550            }
551            Media::Image(..) | Media::Audio(_) | Media::Video(_) => false,
552        }
553    }
554
555    /// An edited block keeps its kind: a tool plan stays a plan.
556    fn identity(&self, item: &Value) -> Map<String, Value> {
557        item.str("type")
558            .filter(|kind| *kind == PLAN)
559            .map(|kind| Map::from_iter([("type".to_owned(), Value::from(kind))]))
560            .unwrap_or_default()
561    }
562
563    fn call_id_slot(&self) -> Option<&'static str> {
564        Some("/id")
565    }
566
567    fn takes_documents(&self) -> bool {
568        true
569    }
570
571    /// Text and thinking with text are content parts, an opaque item with
572    /// a `type` is a part, and a call is a call. A tool plan rides with its
573    /// calls, and images are never sent.
574    fn sends_alone(&self, block: &AssistantContent) -> bool {
575        match block {
576            AssistantContent::Text(text) => !text.text.is_empty(),
577            AssistantContent::Reasoning(reasoning) => {
578                !reasoning.text.is_empty() && kind(block, self).as_deref() != Some(PLAN)
579            }
580            AssistantContent::Opaque(opaque) => opaque.item.get("type").is_some(),
581            AssistantContent::ToolCall(_) => true,
582            AssistantContent::Image(_) => false,
583        }
584    }
585}
586
587#[cfg(test)]
588mod tests;