1use 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
31const CHAT_PATH: &str = "/v2/chat";
33
34pub(crate) const API: &str = "cohere.chat";
36
37#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
40pub struct NativeChat {
41 pub provider: CohereConfig,
43 pub model: String,
45 pub strict_tools: bool,
48}
49
50impl NativeChat {
51 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 pub fn with_strict_tools(mut self) -> Self {
62 self.strict_tools = true;
63 self
64 }
65
66 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 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 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 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 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
287pub(crate) const PLAN: &str = "tool_plan";
289
290fn 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
303fn kind(block: &AssistantContent, target: &NativeChat) -> Option<String> {
305 replay_kind(&block.replay(target, &WireIds::default()))
306}
307
308fn 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
322fn 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
340fn 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
366fn unsendable(what: &str) -> EncodeError {
369 EncodeError::request(format!("Cohere chat cannot carry {what} in this form"))
370}
371
372fn 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
380fn 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
434const REQUEST_ID_HEADER: &str = "x-debug-trace-id";
436
437impl crate::completion::ReplayTarget for NativeChat {
438 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 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 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 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 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 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;