1use 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
21pub(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 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 #[serde(untagged)]
67 Other(String),
68}
69
70pub(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 #[serde(default)]
90 pub cached_tokens: Option<f64>,
91}
92
93impl 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: 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 #[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 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#[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 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;