Skip to main content

rig_core/providers/cohere/
completion.rs

1//! Cohere chat request conversion and typed responses, usage, citations, and tools.
2//!
3//! ```
4//! use rig_core::providers::cohere::completion::FinishReason;
5//! let reason: FinishReason = serde_json::from_str("\"COMPLETE\"")?;
6//! assert_eq!(reason, FinishReason::Complete);
7//! # Ok::<(), serde_json::Error>(())
8//! ```
9
10use crate::error::EncodeError;
11use crate::error::ProviderError;
12use crate::{
13    completion, json_utils,
14    message::{self, ToolChoice},
15};
16use std::collections::HashMap;
17
18use crate::completion::CompletionRequest;
19use serde::{Deserialize, Serialize};
20
21/// Stable descriptor name recorded on normalized responses, streams, and
22/// telemetry spans for this provider.
23pub(crate) const PROVIDER_NAME: &str = "cohere";
24
25#[derive(Debug, Deserialize, Serialize)]
26pub struct CompletionResponse {
27    pub id: String,
28    pub finish_reason: FinishReason,
29    message: Message,
30    #[serde(default)]
31    pub usage: Option<Usage>,
32}
33
34impl CompletionResponse {
35    /// Clone assistant content, citations, and tool calls. Returns a response
36    /// error when the message is not an assistant message.
37    pub fn message(
38        &self,
39    ) -> Result<(Vec<AssistantContent>, Vec<Citation>, Vec<ToolCall>), ProviderError> {
40        let Message::Assistant {
41            content,
42            citations,
43            tool_calls,
44            ..
45        } = self.message.clone()
46        else {
47            return Err(ProviderError::Response(
48                "completion response did not contain an assistant message".into(),
49            ));
50        };
51
52        Ok((content, citations, tool_calls))
53    }
54}
55
56#[derive(Debug, Deserialize, PartialEq, Eq, Clone, Serialize)]
57#[serde(rename_all = "SCREAMING_SNAKE_CASE")]
58pub enum FinishReason {
59    MaxTokens,
60    StopSequence,
61    Complete,
62    Error,
63    ToolCall,
64    /// A reason outside the set Cohere documents today, kept verbatim in
65    /// Cohere's own spelling rather than failing deserialization.
66    #[serde(untagged)]
67    Other(String),
68}
69
70/// Normalize the terminal reason, preserving `ERROR` and unknown values as
71/// [`completion::FinishReason::Other`] rather than treating them as natural stops.
72pub(crate) fn map_finish_reason(reason: &FinishReason) -> completion::FinishReason {
73    match reason {
74        FinishReason::Complete | FinishReason::StopSequence => completion::FinishReason::Stop,
75        FinishReason::MaxTokens => completion::FinishReason::Length,
76        FinishReason::ToolCall => completion::FinishReason::ToolCalls,
77        FinishReason::Error => completion::FinishReason::Other("ERROR".to_owned()),
78        FinishReason::Other(other) => completion::FinishReason::Other(other.clone()),
79    }
80}
81
82#[derive(Copy, Debug, Deserialize, Clone, Serialize)]
83pub struct Usage {
84    #[serde(default)]
85    pub billed_units: Option<BilledUnits>,
86    #[serde(default)]
87    pub tokens: Option<Tokens>,
88    /// Subset of `tokens.input_tokens`; excluded from `billed_units.input_tokens`.
89    #[serde(default)]
90    pub cached_tokens: Option<f64>,
91}
92
93/// Normalize total token counters, not billed units, which exclude cached input
94/// and system overhead. A total requires both input and output counts.
95impl From<&Usage> for crate::completion::Usage {
96    fn from(usage: &Usage) -> crate::completion::Usage {
97        let tokens = usage.tokens.as_ref();
98        let input_tokens = tokens.and_then(|t| t.input_tokens).map(|n| n as u64);
99        let output_tokens = tokens.and_then(|t| t.output_tokens).map(|n| n as u64);
100        crate::completion::Usage {
101            input_tokens,
102            output_tokens,
103            total_tokens: input_tokens
104                .zip(output_tokens)
105                .map(|(input, output)| input + output),
106            // `cached_input_tokens` is a subset of `input_tokens`, so it's only
107            // reported when Cohere also reports `tokens`.
108            cached_input_tokens: tokens.and(usage.cached_tokens).map(|n| n as u64),
109            ..Default::default()
110        }
111    }
112}
113
114impl From<Usage> for crate::completion::Usage {
115    fn from(usage: Usage) -> crate::completion::Usage {
116        crate::completion::Usage::from(&usage)
117    }
118}
119
120#[derive(Copy, Debug, Deserialize, Clone, Serialize)]
121pub struct BilledUnits {
122    #[serde(default)]
123    pub output_tokens: Option<f64>,
124    #[serde(default)]
125    pub classifications: Option<f64>,
126    #[serde(default)]
127    pub search_units: Option<f64>,
128    #[serde(default)]
129    pub input_tokens: Option<f64>,
130}
131
132#[derive(Copy, Debug, Deserialize, Clone, Serialize)]
133pub struct Tokens {
134    #[serde(default)]
135    pub input_tokens: Option<f64>,
136    #[serde(default)]
137    pub output_tokens: Option<f64>,
138}
139
140#[derive(Clone, Debug, Deserialize, Serialize, PartialEq, Eq)]
141pub struct Document {
142    pub id: String,
143    /// Document text and metadata, serialized in sorted key order to keep
144    /// prompt-cache prefixes stable across map instances.
145    #[serde(serialize_with = "crate::json_utils::serialize_map_sorted")]
146    pub data: HashMap<String, serde_json::Value>,
147}
148
149impl From<completion::Document> for Document {
150    fn from(document: completion::Document) -> Self {
151        let mut data: HashMap<String, serde_json::Value> = HashMap::new();
152
153        document
154            .additional_props
155            .into_iter()
156            .for_each(|(key, value)| {
157                data.insert(key, value.into());
158            });
159
160        data.insert("text".to_string(), document.text.into());
161
162        Self {
163            id: document.id,
164            data,
165        }
166    }
167}
168
169#[derive(Clone, Debug, Deserialize, Serialize, PartialEq, Eq)]
170pub struct ToolCall {
171    #[serde(default)]
172    pub id: Option<String>,
173    #[serde(default)]
174    pub r#type: Option<ToolType>,
175    #[serde(default)]
176    pub function: Option<ToolCallFunction>,
177}
178
179#[derive(Clone, Debug, Deserialize, Serialize, PartialEq, Eq)]
180pub struct ToolCallFunction {
181    pub name: String,
182    #[serde(with = "json_utils::stringified_json")]
183    pub arguments: serde_json::Value,
184}
185
186#[derive(Clone, Default, Debug, Deserialize, Serialize, PartialEq, Eq)]
187#[serde(rename_all = "lowercase")]
188pub enum ToolType {
189    #[default]
190    Function,
191}
192
193#[derive(Clone, Debug, Deserialize, Serialize, PartialEq, Eq)]
194pub struct Tool {
195    pub r#type: ToolType,
196    pub function: Function,
197}
198
199#[derive(Clone, Debug, Deserialize, Serialize, PartialEq, Eq)]
200pub struct Function {
201    pub name: String,
202    #[serde(default)]
203    pub description: Option<String>,
204    pub parameters: serde_json::Value,
205}
206
207impl From<completion::ToolDefinition> for Tool {
208    fn from(tool: completion::ToolDefinition) -> Self {
209        Self {
210            r#type: ToolType::default(),
211            function: Function {
212                name: tool.name,
213                description: Some(tool.description),
214                parameters: tool.parameters,
215            },
216        }
217    }
218}
219
220#[derive(Debug, Clone, Deserialize, Serialize, PartialEq, Eq)]
221#[serde(tag = "role", rename_all = "lowercase")]
222pub enum Message {
223    User {
224        content: Vec<UserContent>,
225    },
226
227    Assistant {
228        #[serde(default)]
229        content: Vec<AssistantContent>,
230        #[serde(default)]
231        citations: Vec<Citation>,
232        #[serde(default)]
233        tool_calls: Vec<ToolCall>,
234        #[serde(default)]
235        tool_plan: Option<String>,
236    },
237
238    Tool {
239        content: Vec<ToolResultContent>,
240        tool_call_id: String,
241    },
242
243    System {
244        content: String,
245    },
246}
247
248#[derive(Debug, Clone, Deserialize, Serialize, PartialEq, Eq)]
249#[serde(tag = "type", rename_all = "lowercase")]
250pub enum UserContent {
251    Text { text: String },
252    ImageUrl { image_url: ImageUrl },
253}
254
255#[derive(Debug, Clone, Deserialize, Serialize, PartialEq, Eq)]
256#[serde(tag = "type", rename_all = "lowercase")]
257pub enum AssistantContent {
258    Text { text: String },
259    Thinking { thinking: String },
260}
261
262#[derive(Debug, Clone, Deserialize, Serialize, PartialEq, Eq)]
263pub struct ImageUrl {
264    pub url: String,
265}
266
267#[derive(Debug, Clone, Deserialize, Serialize, PartialEq, Eq)]
268#[serde(tag = "type", rename_all = "lowercase")]
269pub enum ToolResultContent {
270    Text { text: String },
271    Document { document: Document },
272}
273
274#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
275pub struct Citation {
276    #[serde(default)]
277    pub start: Option<u32>,
278    #[serde(default)]
279    pub end: Option<u32>,
280    #[serde(default)]
281    pub text: Option<String>,
282    #[serde(rename = "type")]
283    pub citation_type: Option<CitationType>,
284    #[serde(default)]
285    pub sources: Vec<Source>,
286}
287
288#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
289#[serde(tag = "type", rename_all = "lowercase")]
290pub enum Source {
291    Document {
292        id: Option<String>,
293        document: Option<serde_json::Map<String, serde_json::Value>>,
294    },
295    Tool {
296        id: Option<String>,
297        tool_output: Option<serde_json::Map<String, serde_json::Value>>,
298    },
299}
300
301#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
302#[serde(rename_all = "SCREAMING_SNAKE_CASE")]
303pub enum CitationType {
304    TextContent,
305    Plan,
306}
307
308impl TryFrom<message::Message> for Vec<Message> {
309    type Error = message::MessageError;
310
311    fn try_from(message: message::Message) -> Result<Self, Self::Error> {
312        Ok(match message {
313            message::Message::User { content } => content
314                .into_iter()
315                .map(|content| match content {
316                    message::UserContent::Text(message::Text { text, .. }) => Ok(Message::User {
317                        content: vec![UserContent::Text { text }],
318                    }),
319                    message::UserContent::ToolResult(tool_result) => Ok(Message::Tool {
320                        tool_call_id: tool_result.call.wire().into_owned(),
321                        content: tool_result
322                            .content
323                            .into_iter()
324                            .map(|content| match content {
325                                message::ToolResultContent::Text(text) => {
326                                    Ok(ToolResultContent::Text { text: text.text })
327                                }
328                                message::ToolResultContent::Json { value } => {
329                                    Ok(ToolResultContent::Text {
330                                        text: value.to_string(),
331                                    })
332                                }
333                                message::ToolResultContent::Image(_) => {
334                                    Err(message::MessageError::ConversionError(
335                                        "Only text tool result content is supported by Cohere"
336                                            .to_owned(),
337                                    ))
338                                }
339                            })
340                            .collect::<Result<Vec<_>, _>>()?,
341                    }),
342                    _ => Err(message::MessageError::ConversionError(
343                        "Only text content is supported by Cohere".to_owned(),
344                    )),
345                })
346                .collect::<Result<Vec<_>, _>>()?,
347            message::Message::System { content } => {
348                vec![Message::System { content }]
349            }
350            message::Message::Assistant { content, .. } => {
351                let mut text_content = vec![];
352                let mut tool_calls = vec![];
353
354                for content in content.into_iter() {
355                    match content {
356                        message::AssistantContent::Text(message::Text { text, .. }) => {
357                            text_content.push(AssistantContent::Text { text });
358                        }
359                        message::AssistantContent::ToolCall(message::ToolCall {
360                            id,
361                            function:
362                                message::ToolFunction {
363                                    name, arguments, ..
364                                },
365                            ..
366                        }) => {
367                            tool_calls.push(ToolCall {
368                                id: Some(id.wire().into_owned()),
369                                r#type: Some(ToolType::Function),
370                                function: Some(ToolCallFunction {
371                                    name: name.into(),
372                                    arguments: serde_json::to_value(arguments).unwrap_or_default(),
373                                }),
374                            });
375                        }
376                        message::AssistantContent::Reasoning(reasoning) => {
377                            // Reasoning another service issued is not replayed.
378                            if let Some(reasoning) = reasoning.open(&super::wire::ISSUER) {
379                                let thinking = reasoning.display_text();
380                                text_content.push(AssistantContent::Thinking { thinking });
381                            }
382                        }
383                        message::AssistantContent::Image(_) => {
384                            return Err(message::MessageError::ConversionError(
385                                "Cohere currently doesn't support images.".to_owned(),
386                            ));
387                        }
388                    }
389                }
390
391                vec![Message::Assistant {
392                    content: text_content,
393                    citations: vec![],
394                    tool_calls,
395                    tool_plan: None,
396                }]
397            }
398        })
399    }
400}
401
402/// Cohere's `tool_choice` is a bare string; only `REQUIRED`/`NONE` are valid.
403/// `Auto` errors below rather than silently mapping to the omitted-field
404/// behavior that would actually let the model decide.
405#[derive(Clone, Copy, Debug, Deserialize, Serialize, PartialEq, Eq)]
406#[serde(rename_all = "SCREAMING_SNAKE_CASE")]
407pub enum CohereToolChoice {
408    Required,
409    None,
410}
411
412impl TryFrom<ToolChoice> for CohereToolChoice {
413    type Error = EncodeError;
414
415    fn try_from(tool_choice: ToolChoice) -> Result<Self, Self::Error> {
416        match tool_choice {
417            ToolChoice::Required => Ok(Self::Required),
418            ToolChoice::None => Ok(Self::None),
419            ToolChoice::Auto => Err(EncodeError::request(
420                "\"auto\" is not an allowed tool_choice value in the Cohere API; \
421                 omit tool_choice to let the model decide",
422            )),
423            ToolChoice::Specific { .. } => Err(EncodeError::request(
424                "the Cohere API cannot be forced to call specific tools by name; \
425                 use ToolChoice::Required and restrict the tools you pass instead",
426            )),
427        }
428    }
429}
430
431#[derive(Debug, Serialize, Deserialize)]
432pub(super) struct CohereCompletionRequest {
433    pub(super) model: String,
434    pub messages: Vec<Message>,
435    documents: Vec<Document>,
436    #[serde(skip_serializing_if = "Option::is_none")]
437    temperature: Option<f64>,
438    #[serde(skip_serializing_if = "Option::is_none")]
439    max_tokens: Option<u64>,
440    #[serde(skip_serializing_if = "Vec::is_empty")]
441    tools: Vec<Tool>,
442    #[serde(skip_serializing_if = "Option::is_none")]
443    tool_choice: Option<CohereToolChoice>,
444    #[serde(flatten, skip_serializing_if = "Option::is_none")]
445    pub additional_params: Option<serde_json::Value>,
446}
447
448impl TryFrom<(&str, CompletionRequest)> for CohereCompletionRequest {
449    type Error = EncodeError;
450
451    fn try_from((model, req): (&str, CompletionRequest)) -> Result<Self, Self::Error> {
452        let documents = req
453            .documents
454            .iter()
455            .cloned()
456            .map(Document::from)
457            .collect::<Vec<_>>();
458        if req.output_schema.is_some() {
459            tracing::warn!("Structured outputs currently not supported for Cohere");
460        }
461
462        let model = req.model.clone().unwrap_or_else(|| model.to_string());
463        let mut partial_history = vec![];
464        partial_history.extend(req.chat_history);
465
466        let mut full_history: Vec<Message> = Vec::new();
467
468        let tool_ids = crate::providers::internal::wire_ids::WireIds::new(&partial_history);
469        for (position, message) in partial_history.into_iter().enumerate() {
470            let mut messages = Vec::<Message>::try_from(message)?;
471            let slots: Vec<&mut String> = messages
472                .iter_mut()
473                .flat_map(|message| match message {
474                    Message::Assistant { tool_calls, .. } => tool_calls
475                        .iter_mut()
476                        .filter_map(|call| call.id.as_mut())
477                        .collect(),
478                    Message::Tool { tool_call_id, .. } => vec![tool_call_id],
479                    _ => Vec::new(),
480                })
481                .collect();
482            tool_ids
483                .apply(position, slots)
484                .map_err(EncodeError::request)?;
485            full_history.extend(messages);
486        }
487
488        let tool_choice = req
489            .tool_choice
490            .map(CohereToolChoice::try_from)
491            .transpose()?;
492
493        // Count tools supplied through the provider escape hatch as well as
494        // typed tools so REQUIRED remains usable with Cohere-specific schemas.
495        let has_tools = !req.tools.is_empty()
496            || req
497                .additional_params
498                .as_ref()
499                .and_then(|params| params.get("tools"))
500                .and_then(serde_json::Value::as_array)
501                .is_some_and(|tools| !tools.is_empty());
502        if matches!(tool_choice, Some(CohereToolChoice::Required)) && !has_tools {
503            return Err(EncodeError::request(
504                "Cohere requires at least one tool when tool_choice is REQUIRED",
505            ));
506        }
507
508        Ok(Self {
509            model,
510            messages: full_history,
511            documents,
512            temperature: req.temperature,
513            max_tokens: req.max_tokens,
514            tools: req.tools.into_iter().map(Tool::from).collect::<Vec<_>>(),
515            tool_choice,
516            additional_params: req.additional_params,
517        })
518    }
519}
520
521#[cfg(test)]
522mod tests;