Skip to main content

rig_core/providers/cohere/
completion.rs

1use crate::{
2    OneOrMany,
3    completion::{self, CompletionError, GetTokenUsage},
4    http_client::{self, HttpClientExt},
5    json_utils,
6    message::{self, Reasoning, ToolChoice},
7    telemetry::{CompletionOperation, CompletionSpanBuilder, SpanCombinator},
8};
9use std::collections::HashMap;
10
11use super::client::Client;
12use crate::completion::CompletionRequest;
13use crate::providers::cohere::streaming::StreamingCompletionResponse;
14use serde::{Deserialize, Serialize};
15use tracing::{Instrument, Level, enabled};
16
17#[derive(Debug, Deserialize, Serialize)]
18pub struct CompletionResponse {
19    pub id: String,
20    pub finish_reason: FinishReason,
21    message: Message,
22    #[serde(default)]
23    pub usage: Option<Usage>,
24}
25
26type AssistantMessageParts = (Vec<AssistantContent>, Vec<Citation>, Vec<ToolCall>);
27
28impl CompletionResponse {
29    /// Return that parts of the response for assistant messages w/o dealing with the other variants
30    pub fn message(&self) -> Result<AssistantMessageParts, CompletionError> {
31        let Message::Assistant {
32            content,
33            citations,
34            tool_calls,
35            ..
36        } = self.message.clone()
37        else {
38            return Err(CompletionError::ResponseError(
39                "completion response did not contain an assistant message".into(),
40            ));
41        };
42
43        Ok((content, citations, tool_calls))
44    }
45}
46
47impl crate::telemetry::ProviderResponseExt for CompletionResponse {
48    type OutputMessage = Message;
49    type Usage = Usage;
50
51    fn get_response_id(&self) -> Option<String> {
52        Some(self.id.clone())
53    }
54
55    fn get_response_model_name(&self) -> Option<String> {
56        None
57    }
58
59    fn get_output_messages(&self) -> Vec<Self::OutputMessage> {
60        vec![self.message.clone()]
61    }
62
63    fn get_text_response(&self) -> Option<String> {
64        let Message::Assistant { ref content, .. } = self.message else {
65            return None;
66        };
67
68        let res = content
69            .iter()
70            .filter_map(|x| {
71                if let AssistantContent::Text { text } = x {
72                    Some(text.to_string())
73                } else {
74                    None
75                }
76            })
77            .collect::<Vec<String>>()
78            .join("\n");
79
80        if res.is_empty() { None } else { Some(res) }
81    }
82
83    fn get_usage(&self) -> Option<Self::Usage> {
84        self.usage.clone()
85    }
86}
87
88#[derive(Debug, Deserialize, PartialEq, Eq, Clone, Serialize)]
89#[serde(rename_all = "SCREAMING_SNAKE_CASE")]
90pub enum FinishReason {
91    MaxTokens,
92    StopSequence,
93    Complete,
94    Error,
95    ToolCall,
96}
97
98#[derive(Debug, Deserialize, Clone, Serialize)]
99pub struct Usage {
100    #[serde(default)]
101    pub billed_units: Option<BilledUnits>,
102    #[serde(default)]
103    pub tokens: Option<Tokens>,
104}
105
106impl GetTokenUsage for Usage {
107    fn token_usage(&self) -> crate::completion::Usage {
108        let mut usage = crate::completion::Usage::new();
109
110        if let Some(ref billed_units) = self.billed_units {
111            usage.input_tokens = billed_units.input_tokens.unwrap_or_default() as u64;
112            usage.output_tokens = billed_units.output_tokens.unwrap_or_default() as u64;
113            usage.total_tokens = usage.input_tokens + usage.output_tokens;
114        }
115
116        usage
117    }
118}
119
120#[derive(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(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
140impl TryFrom<CompletionResponse> for completion::CompletionResponse<CompletionResponse> {
141    type Error = CompletionError;
142
143    fn try_from(response: CompletionResponse) -> Result<Self, Self::Error> {
144        let (content, _, tool_calls) = response.message()?;
145
146        let model_response = if !tool_calls.is_empty() {
147            OneOrMany::many(
148                tool_calls
149                    .into_iter()
150                    .filter_map(|tool_call| {
151                        let ToolCallFunction { name, arguments } = tool_call.function?;
152                        let id = tool_call.id.unwrap_or_else(|| name.clone());
153
154                        Some(completion::AssistantContent::tool_call(id, name, arguments))
155                    })
156                    .collect::<Vec<_>>(),
157            )
158            .map_err(|_| {
159                CompletionError::ResponseError(
160                    "response contained tool call metadata without any callable tool content"
161                        .to_owned(),
162                )
163            })?
164        } else {
165            OneOrMany::many(content.into_iter().map(|content| match content {
166                AssistantContent::Text { text } => completion::AssistantContent::text(text),
167                AssistantContent::Thinking { thinking } => {
168                    completion::AssistantContent::Reasoning(Reasoning::new(&thinking))
169                }
170            }))
171            .map_err(|_| {
172                CompletionError::ResponseError(
173                    "Response contained no message or tool call (empty)".to_owned(),
174                )
175            })?
176        };
177
178        let usage = response
179            .usage
180            .as_ref()
181            .and_then(|usage| usage.tokens.as_ref())
182            .map(|tokens| {
183                let input_tokens = tokens.input_tokens.unwrap_or(0.0);
184                let output_tokens = tokens.output_tokens.unwrap_or(0.0);
185
186                completion::Usage {
187                    input_tokens: input_tokens as u64,
188                    output_tokens: output_tokens as u64,
189                    total_tokens: (input_tokens + output_tokens) as u64,
190                    cached_input_tokens: 0,
191                    cache_creation_input_tokens: 0,
192                    tool_use_prompt_tokens: 0,
193                    reasoning_tokens: 0,
194                }
195            })
196            .unwrap_or_default();
197
198        Ok(completion::CompletionResponse {
199            choice: model_response,
200            usage,
201            raw_response: response,
202            message_id: None,
203        })
204    }
205}
206
207#[derive(Clone, Debug, Deserialize, Serialize, PartialEq, Eq)]
208pub struct Document {
209    pub id: String,
210    pub data: HashMap<String, serde_json::Value>,
211}
212
213impl From<completion::Document> for Document {
214    fn from(document: completion::Document) -> Self {
215        let mut data: HashMap<String, serde_json::Value> = HashMap::new();
216
217        // We use `.into()` here explicitly since the `document.additional_props` type will likely
218        //  evolve into `serde_json::Value` in the future.
219        document
220            .additional_props
221            .into_iter()
222            .for_each(|(key, value)| {
223                data.insert(key, value.into());
224            });
225
226        data.insert("text".to_string(), document.text.into());
227
228        Self {
229            id: document.id,
230            data,
231        }
232    }
233}
234
235#[derive(Clone, Debug, Deserialize, Serialize, PartialEq, Eq)]
236pub struct ToolCall {
237    #[serde(default)]
238    pub id: Option<String>,
239    #[serde(default)]
240    pub r#type: Option<ToolType>,
241    #[serde(default)]
242    pub function: Option<ToolCallFunction>,
243}
244
245#[derive(Clone, Debug, Deserialize, Serialize, PartialEq, Eq)]
246pub struct ToolCallFunction {
247    pub name: String,
248    #[serde(with = "json_utils::stringified_json")]
249    pub arguments: serde_json::Value,
250}
251
252#[derive(Clone, Default, Debug, Deserialize, Serialize, PartialEq, Eq)]
253#[serde(rename_all = "lowercase")]
254pub enum ToolType {
255    #[default]
256    Function,
257}
258
259#[derive(Clone, Debug, Deserialize, Serialize, PartialEq, Eq)]
260pub struct Tool {
261    pub r#type: ToolType,
262    pub function: Function,
263}
264
265#[derive(Clone, Debug, Deserialize, Serialize, PartialEq, Eq)]
266pub struct Function {
267    pub name: String,
268    #[serde(default)]
269    pub description: Option<String>,
270    pub parameters: serde_json::Value,
271}
272
273impl From<completion::ToolDefinition> for Tool {
274    fn from(tool: completion::ToolDefinition) -> Self {
275        Self {
276            r#type: ToolType::default(),
277            function: Function {
278                name: tool.name,
279                description: Some(tool.description),
280                parameters: tool.parameters,
281            },
282        }
283    }
284}
285
286#[derive(Debug, Clone, Deserialize, Serialize, PartialEq, Eq)]
287#[serde(tag = "role", rename_all = "lowercase")]
288pub enum Message {
289    User {
290        content: OneOrMany<UserContent>,
291    },
292
293    Assistant {
294        #[serde(default)]
295        content: Vec<AssistantContent>,
296        #[serde(default)]
297        citations: Vec<Citation>,
298        #[serde(default)]
299        tool_calls: Vec<ToolCall>,
300        #[serde(default)]
301        tool_plan: Option<String>,
302    },
303
304    Tool {
305        content: OneOrMany<ToolResultContent>,
306        tool_call_id: String,
307    },
308
309    System {
310        content: String,
311    },
312}
313
314#[derive(Debug, Clone, Deserialize, Serialize, PartialEq, Eq)]
315#[serde(tag = "type", rename_all = "lowercase")]
316pub enum UserContent {
317    Text { text: String },
318    ImageUrl { image_url: ImageUrl },
319}
320
321#[derive(Debug, Clone, Deserialize, Serialize, PartialEq, Eq)]
322#[serde(tag = "type", rename_all = "lowercase")]
323pub enum AssistantContent {
324    Text { text: String },
325    Thinking { thinking: String },
326}
327
328#[derive(Debug, Clone, Deserialize, Serialize, PartialEq, Eq)]
329pub struct ImageUrl {
330    pub url: String,
331}
332
333#[derive(Debug, Clone, Deserialize, Serialize, PartialEq, Eq)]
334pub enum ToolResultContent {
335    Text { text: String },
336    Document { document: Document },
337}
338
339#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
340pub struct Citation {
341    #[serde(default)]
342    pub start: Option<u32>,
343    #[serde(default)]
344    pub end: Option<u32>,
345    #[serde(default)]
346    pub text: Option<String>,
347    #[serde(rename = "type")]
348    pub citation_type: Option<CitationType>,
349    #[serde(default)]
350    pub sources: Vec<Source>,
351}
352
353#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
354#[serde(tag = "type", rename_all = "lowercase")]
355pub enum Source {
356    Document {
357        id: Option<String>,
358        document: Option<serde_json::Map<String, serde_json::Value>>,
359    },
360    Tool {
361        id: Option<String>,
362        tool_output: Option<serde_json::Map<String, serde_json::Value>>,
363    },
364}
365
366#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
367#[serde(rename_all = "SCREAMING_SNAKE_CASE")]
368pub enum CitationType {
369    TextContent,
370    Plan,
371}
372
373impl TryFrom<message::Message> for Vec<Message> {
374    type Error = message::MessageError;
375
376    fn try_from(message: message::Message) -> Result<Self, Self::Error> {
377        Ok(match message {
378            message::Message::User { content } => content
379                .into_iter()
380                .map(|content| match content {
381                    message::UserContent::Text(message::Text { text, .. }) => Ok(Message::User {
382                        content: OneOrMany::one(UserContent::Text { text }),
383                    }),
384                    message::UserContent::ToolResult(message::ToolResult {
385                        id, content, ..
386                    }) => Ok(Message::Tool {
387                        tool_call_id: id,
388                        content: content.try_map(|content| match content {
389                            message::ToolResultContent::Text(text) => {
390                                Ok(ToolResultContent::Text { text: text.text })
391                            }
392                            message::ToolResultContent::Json { value } => {
393                                Ok(ToolResultContent::Text {
394                                    text: value.to_string(),
395                                })
396                            }
397                            message::ToolResultContent::Image(_) => {
398                                Err(message::MessageError::ConversionError(
399                                    "Only text tool result content is supported by Cohere"
400                                        .to_owned(),
401                                ))
402                            }
403                        })?,
404                    }),
405                    _ => Err(message::MessageError::ConversionError(
406                        "Only text content is supported by Cohere".to_owned(),
407                    )),
408                })
409                .collect::<Result<Vec<_>, _>>()?,
410            message::Message::System { content } => {
411                vec![Message::System { content }]
412            }
413            message::Message::Assistant { content, .. } => {
414                let mut text_content = vec![];
415                let mut tool_calls = vec![];
416
417                for content in content.into_iter() {
418                    match content {
419                        message::AssistantContent::Text(message::Text { text, .. }) => {
420                            text_content.push(AssistantContent::Text { text });
421                        }
422                        message::AssistantContent::ToolCall(message::ToolCall {
423                            id,
424                            function:
425                                message::ToolFunction {
426                                    name, arguments, ..
427                                },
428                            ..
429                        }) => {
430                            tool_calls.push(ToolCall {
431                                id: Some(id),
432                                r#type: Some(ToolType::Function),
433                                function: Some(ToolCallFunction {
434                                    name,
435                                    arguments: serde_json::to_value(arguments).unwrap_or_default(),
436                                }),
437                            });
438                        }
439                        message::AssistantContent::Reasoning(reasoning) => {
440                            let thinking = reasoning.display_text();
441                            text_content.push(AssistantContent::Thinking { thinking });
442                        }
443                        message::AssistantContent::Image(_) => {
444                            return Err(message::MessageError::ConversionError(
445                                "Cohere currently doesn't support images.".to_owned(),
446                            ));
447                        }
448                    }
449                }
450
451                vec![Message::Assistant {
452                    content: text_content,
453                    citations: vec![],
454                    tool_calls,
455                    tool_plan: None,
456                }]
457            }
458        })
459    }
460}
461
462impl TryFrom<Message> for message::Message {
463    type Error = message::MessageError;
464
465    fn try_from(message: Message) -> Result<Self, Self::Error> {
466        match message {
467            Message::User { content } => Ok(message::Message::User {
468                content: content.map(|content| match content {
469                    UserContent::Text { text } => {
470                        message::UserContent::Text(message::Text::new(text))
471                    }
472                    UserContent::ImageUrl { image_url } => {
473                        message::UserContent::image_url(image_url.url, None, None)
474                    }
475                }),
476            }),
477            Message::Assistant {
478                content,
479                tool_calls,
480                ..
481            } => {
482                let mut content = content
483                    .into_iter()
484                    .map(|content| match content {
485                        AssistantContent::Text { text } => message::AssistantContent::text(text),
486                        AssistantContent::Thinking { thinking } => {
487                            message::AssistantContent::Reasoning(Reasoning::new(&thinking))
488                        }
489                    })
490                    .collect::<Vec<_>>();
491
492                content.extend(tool_calls.into_iter().filter_map(|tool_call| {
493                    let ToolCallFunction { name, arguments } = tool_call.function?;
494
495                    Some(message::AssistantContent::tool_call(
496                        tool_call.id.unwrap_or_else(|| name.clone()),
497                        name,
498                        arguments,
499                    ))
500                }));
501
502                let content = OneOrMany::many(content).map_err(|_| {
503                    message::MessageError::ConversionError(
504                        "Expected either text content or tool calls".to_string(),
505                    )
506                })?;
507
508                Ok(message::Message::Assistant { id: None, content })
509            }
510            Message::Tool {
511                content,
512                tool_call_id,
513            } => {
514                let content = content.try_map(|content| {
515                    Ok(match content {
516                        ToolResultContent::Text { text } => message::ToolResultContent::text(text),
517                        ToolResultContent::Document { document } => {
518                            message::ToolResultContent::json(
519                                serde_json::to_value(document.data).map_err(|e| {
520                                    message::MessageError::ConversionError(
521                                        format!("Failed to convert tool result document content into JSON: {e}"),
522                                    )
523                                })?,
524                            )
525                        }
526                    })
527                })?;
528
529                Ok(message::Message::User {
530                    content: OneOrMany::one(message::UserContent::tool_result(
531                        tool_call_id,
532                        content,
533                    )),
534                })
535            }
536            Message::System { content } => Ok(message::Message::user(content)),
537        }
538    }
539}
540
541#[derive(Clone)]
542pub struct CompletionModel<T = reqwest::Client> {
543    pub(crate) client: Client<T>,
544    pub model: String,
545}
546
547#[derive(Debug, Serialize, Deserialize)]
548pub(super) struct CohereCompletionRequest {
549    pub(super) model: String,
550    pub messages: Vec<Message>,
551    documents: Vec<crate::completion::Document>,
552    #[serde(skip_serializing_if = "Option::is_none")]
553    temperature: Option<f64>,
554    #[serde(skip_serializing_if = "Vec::is_empty")]
555    tools: Vec<Tool>,
556    #[serde(skip_serializing_if = "Option::is_none")]
557    tool_choice: Option<ToolChoice>,
558    #[serde(flatten, skip_serializing_if = "Option::is_none")]
559    pub additional_params: Option<serde_json::Value>,
560}
561
562impl TryFrom<(&str, CompletionRequest)> for CohereCompletionRequest {
563    type Error = CompletionError;
564
565    fn try_from((model, req): (&str, CompletionRequest)) -> Result<Self, Self::Error> {
566        let documents = req.documents.clone();
567        if req.output_schema.is_some() {
568            tracing::warn!("Structured outputs currently not supported for Cohere");
569        }
570
571        let model = req.model.clone().unwrap_or_else(|| model.to_string());
572        let mut partial_history = vec![];
573        partial_history.extend(req.chat_history);
574
575        let mut full_history: Vec<Message> = req.preamble.map_or_else(Vec::new, |preamble| {
576            vec![Message::System { content: preamble }]
577        });
578
579        full_history.extend(
580            partial_history
581                .into_iter()
582                .map(message::Message::try_into)
583                .collect::<Result<Vec<Vec<Message>>, _>>()?
584                .into_iter()
585                .flatten()
586                .collect::<Vec<_>>(),
587        );
588
589        let tool_choice = if let Some(tool_choice) = req.tool_choice {
590            if !matches!(tool_choice, ToolChoice::Auto) {
591                Some(tool_choice)
592            } else {
593                return Err(CompletionError::RequestError(
594                    "\"auto\" is not an allowed tool_choice value in the Cohere API".into(),
595                ));
596            }
597        } else {
598            None
599        };
600
601        Ok(Self {
602            model: model.to_string(),
603            messages: full_history,
604            documents,
605            temperature: req.temperature,
606            tools: req.tools.into_iter().map(Tool::from).collect::<Vec<_>>(),
607            tool_choice,
608            additional_params: req.additional_params,
609        })
610    }
611}
612
613impl<T> CompletionModel<T>
614where
615    T: HttpClientExt,
616{
617    pub fn new(client: Client<T>, model: impl Into<String>) -> Self {
618        Self {
619            client,
620            model: model.into(),
621        }
622    }
623}
624
625impl<T> completion::CompletionModel for CompletionModel<T>
626where
627    T: HttpClientExt + Clone + 'static,
628{
629    type Response = CompletionResponse;
630    type StreamingResponse = StreamingCompletionResponse;
631    type Client = Client<T>;
632
633    fn make(client: &Self::Client, model: impl Into<String>) -> Self {
634        Self::new(client.clone(), model.into())
635    }
636
637    async fn completion(
638        &self,
639        completion_request: completion::CompletionRequest,
640    ) -> Result<completion::CompletionResponse<CompletionResponse>, CompletionError> {
641        let system_instructions = completion_request.preamble.clone();
642        let record_telemetry_content = completion_request.record_telemetry_content;
643        let request = CohereCompletionRequest::try_from((self.model.as_ref(), completion_request))?;
644
645        let llm_span =
646            CompletionSpanBuilder::new("cohere", &request.model, CompletionOperation::Chat)
647                .system_instructions(system_instructions.as_deref(), record_telemetry_content)
648                .build();
649
650        if enabled!(Level::TRACE) {
651            tracing::trace!(
652                "Cohere completion request: {}",
653                serde_json::to_string_pretty(&request)?
654            );
655        }
656
657        let req_body = serde_json::to_vec(&request)?;
658
659        let req = self
660            .client
661            .post("/v2/chat")?
662            .body(req_body)
663            .map_err(|e| CompletionError::HttpError(e.into()))?;
664
665        async {
666            let response = self
667                .client
668                .send::<_, bytes::Bytes>(req)
669                .await
670                .map_err(|e| http_client::Error::Instance(e.into()))?;
671
672            let status = response.status();
673            let body = response.into_body().into_future().await?.to_owned();
674
675            if status.is_success() {
676                let json_response: CompletionResponse = serde_json::from_slice(&body)?;
677                let span = tracing::Span::current();
678                span.record_token_usage(&json_response.usage);
679                span.record_response_metadata(&json_response);
680
681                if enabled!(Level::TRACE) {
682                    tracing::trace!(
683                        target: "rig::completions",
684                        "Cohere completion response: {}",
685                        serde_json::to_string_pretty(&json_response)?
686                    );
687                }
688
689                let completion: completion::CompletionResponse<CompletionResponse> =
690                    json_response.try_into()?;
691                Ok(completion)
692            } else {
693                Err(CompletionError::from_http_response(
694                    status,
695                    String::from_utf8_lossy(&body),
696                ))
697            }
698        }
699        .instrument(llm_span)
700        .await
701    }
702
703    async fn stream(
704        &self,
705        request: CompletionRequest,
706    ) -> Result<
707        crate::streaming::StreamingCompletionResponse<Self::StreamingResponse>,
708        CompletionError,
709    > {
710        CompletionModel::stream(self, request).await
711    }
712}
713#[cfg(test)]
714mod tests {
715    use super::*;
716    use serde_path_to_error::deserialize;
717
718    #[test]
719    fn test_deserialize_completion_response() {
720        let json_data = r#"
721        {
722            "id": "abc123",
723            "message": {
724                "role": "assistant",
725                "tool_plan": "I will use the subtract tool to find the difference between 2 and 5.",
726                "tool_calls": [
727                        {
728                            "id": "subtract_sm6ps6fb6y9f",
729                            "type": "function",
730                            "function": {
731                                "name": "subtract",
732                                "arguments": "{\"x\":5,\"y\":2}"
733                            }
734                        }
735                    ]
736                },
737                "finish_reason": "TOOL_CALL",
738                "usage": {
739                "billed_units": {
740                    "input_tokens": 78,
741                    "output_tokens": 27
742                },
743                "tokens": {
744                    "input_tokens": 1028,
745                    "output_tokens": 63
746                }
747            }
748        }
749        "#;
750
751        let mut deserializer = serde_json::Deserializer::from_str(json_data);
752        let result: Result<CompletionResponse, _> = deserialize(&mut deserializer);
753
754        let response = result.unwrap();
755        let (_, citations, tool_calls) = response.message().expect("assistant message");
756        let CompletionResponse {
757            id,
758            finish_reason,
759            usage,
760            ..
761        } = response;
762
763        assert_eq!(id, "abc123");
764        assert_eq!(finish_reason, FinishReason::ToolCall);
765
766        let Usage {
767            billed_units,
768            tokens,
769        } = usage.unwrap();
770        let BilledUnits {
771            input_tokens: billed_input_tokens,
772            output_tokens: billed_output_tokens,
773            ..
774        } = billed_units.unwrap();
775        let Tokens {
776            input_tokens,
777            output_tokens,
778        } = tokens.unwrap();
779
780        assert_eq!(billed_input_tokens.unwrap(), 78.0);
781        assert_eq!(billed_output_tokens.unwrap(), 27.0);
782        assert_eq!(input_tokens.unwrap(), 1028.0);
783        assert_eq!(output_tokens.unwrap(), 63.0);
784
785        assert!(citations.is_empty());
786        assert_eq!(tool_calls.len(), 1);
787
788        let ToolCallFunction { name, arguments } = tool_calls[0].function.clone().unwrap();
789
790        assert_eq!(name, "subtract");
791        assert_eq!(arguments, serde_json::json!({"x": 5, "y": 2}));
792    }
793
794    #[test]
795    fn test_convert_completion_message_to_message_and_back() {
796        let completion_message = completion::Message::User {
797            content: OneOrMany::one(completion::message::UserContent::Text(
798                completion::message::Text::new("Hello, world!".to_string()),
799            )),
800        };
801
802        let messages: Vec<Message> = completion_message.clone().try_into().unwrap();
803        let _converted_back: Vec<completion::Message> = messages
804            .into_iter()
805            .map(|msg| msg.try_into().unwrap())
806            .collect::<Vec<_>>();
807    }
808
809    #[test]
810    fn test_convert_message_to_completion_message_and_back() {
811        let message = Message::User {
812            content: OneOrMany::one(UserContent::Text {
813                text: "Hello, world!".to_string(),
814            }),
815        };
816
817        let completion_message: completion::Message = message.clone().try_into().unwrap();
818        let _converted_back: Vec<Message> = completion_message.try_into().unwrap();
819    }
820
821    #[test]
822    fn cohere_builder_request_preserves_native_documents() {
823        let request = crate::completion::CompletionRequestBuilder::new(
824            crate::test_utils::MockCompletionModel::default(),
825            "What is glarb-glarb?",
826        )
827        .document(crate::completion::request::Document {
828            id: "doc_1".to_string(),
829            text: "Definition of glarb-glarb: an ancient tool.".to_string(),
830            additional_props: Default::default(),
831        })
832        .build();
833
834        let request = CohereCompletionRequest::try_from(("command-r", request))
835            .expect("request conversion should succeed");
836
837        assert_eq!(request.documents.len(), 1);
838        assert_eq!(request.documents[0].id, "doc_1");
839    }
840
841    #[tokio::test]
842    async fn completion_non_success_preserves_status_and_body() {
843        use crate::client::CompletionClient;
844        use crate::completion::CompletionModel as _;
845        use crate::test_utils::RecordingHttpClient;
846
847        let body = r#"{"error":{"message":"boom"}}"#;
848        let http_client =
849            RecordingHttpClient::with_error_response(http::StatusCode::SERVICE_UNAVAILABLE, body);
850        let client = crate::providers::cohere::Client::builder()
851            .api_key("test-key")
852            .http_client(http_client)
853            .build()
854            .expect("build client");
855        let model = client.completion_model(crate::providers::cohere::COMMAND_R);
856        let request = model.completion_request("hello").build();
857
858        let error = model
859            .completion(request)
860            .await
861            .expect_err("should fail with non-success status");
862
863        assert!(matches!(error, CompletionError::HttpError(_)));
864        assert_eq!(
865            error.provider_response_status(),
866            Some(http::StatusCode::SERVICE_UNAVAILABLE)
867        );
868        assert_eq!(error.provider_response_body(), Some(body));
869    }
870}