Skip to main content

rig_core/providers/mistral/
completion.rs

1//! Mistral chat model identifiers and typed responses with service-tier and audio usage.
2//!
3//! ```no_run
4//! use rig_core::providers::mistral;
5//! let model = mistral::from_env()?.chat(mistral::MISTRAL_SMALL);
6//! # Ok::<(), Box<dyn std::error::Error>>(())
7//! ```
8
9use serde::{Deserialize, Deserializer, Serialize};
10
11use crate::json_utils;
12use crate::providers::openai;
13
14/// The latest version of the `codestral` Mistral model
15pub const CODESTRAL: &str = "codestral-latest";
16/// The latest version of the `mistral-large` Mistral model
17pub const MISTRAL_LARGE: &str = "mistral-large-latest";
18/// The latest version of the `mistral-3b` Mistral completions model
19pub const MINISTRAL_3B: &str = "ministral-3b-latest";
20/// The latest version of the `mistral-8b` Mistral completions model
21pub const MINISTRAL_8B: &str = "ministral-8b-latest";
22
23/// The latest version of the `mistral-small` Mistral completions model
24pub const MISTRAL_SMALL: &str = "mistral-small-latest";
25
26fn mistral_content_value_to_text(value: serde_json::Value) -> String {
27    match value {
28        serde_json::Value::String(text) => text,
29        serde_json::Value::Array(parts) => openai::completion::joined_text_parts(&parts),
30        _ => String::new(),
31    }
32}
33
34fn deserialize_mistral_content_string<'de, D>(deserializer: D) -> Result<String, D::Error>
35where
36    D: Deserializer<'de>,
37{
38    Ok(Option::<serde_json::Value>::deserialize(deserializer)?
39        .map(mistral_content_value_to_text)
40        .unwrap_or_default())
41}
42
43#[derive(Debug, Serialize, Deserialize, Clone)]
44pub struct Choice {
45    pub index: usize,
46    pub message: Message,
47    pub logprobs: Option<serde_json::Value>,
48    pub finish_reason: String,
49}
50
51/// Mistral's provider-native message shape, as it appears in responses.
52#[derive(Debug, Serialize, Deserialize, PartialEq, Clone)]
53#[serde(tag = "role", rename_all = "lowercase")]
54pub enum Message {
55    User {
56        content: String,
57    },
58    Assistant {
59        #[serde(default, deserialize_with = "deserialize_mistral_content_string")]
60        content: String,
61        #[serde(
62            default,
63            deserialize_with = "json_utils::null_or_default",
64            skip_serializing_if = "Vec::is_empty"
65        )]
66        tool_calls: Vec<ToolCall>,
67        #[serde(default)]
68        prefix: bool,
69    },
70    System {
71        content: String,
72    },
73    Tool {
74        /// The name of the tool that was called
75        #[serde(skip_serializing_if = "String::is_empty")]
76        name: String,
77        /// The content of the tool call
78        content: String,
79        /// The id of the tool call
80        tool_call_id: String,
81    },
82}
83
84#[derive(Debug, Serialize, Deserialize, PartialEq, Clone)]
85pub struct ToolCall {
86    pub id: String,
87    #[serde(default)]
88    pub r#type: ToolType,
89    pub function: Function,
90}
91
92#[derive(Debug, Serialize, Deserialize, PartialEq, Clone)]
93pub struct Function {
94    pub name: String,
95    #[serde(with = "json_utils::stringified_json")]
96    pub arguments: serde_json::Value,
97}
98
99#[derive(Default, Debug, Serialize, Deserialize, PartialEq, Clone)]
100#[serde(rename_all = "lowercase")]
101pub enum ToolType {
102    #[default]
103    Function,
104}
105
106#[derive(Debug, Deserialize, Clone, Serialize)]
107pub struct CompletionResponse {
108    pub id: String,
109    pub object: String,
110    pub created: u64,
111    pub model: String,
112    pub system_fingerprint: Option<String>,
113    #[serde(
114        deserialize_with = "crate::providers::internal::openai_chat_completions_compatible::deserialize_choices_dropping_incomplete_tool_calls"
115    )]
116    pub choices: Vec<Choice>,
117    pub usage: Option<Usage>,
118}
119
120/// In-depth details on prompt tokens.
121///
122/// Mirrors Mistral's `PromptTokensDetails` schema.
123#[derive(Clone, Debug, Default, Deserialize, Serialize)]
124pub struct PromptTokensDetails {
125    /// Number of tokens served from the prompt cache.
126    #[serde(default)]
127    pub cached_tokens: u64,
128    /// Input audio tokens, separate from `prompt_tokens`; both contribute to total usage.
129    #[serde(default)]
130    pub audio_tokens: u64,
131}
132
133/// Token usage returned by Mistral's chat completions and embeddings endpoints.
134///
135/// See <https://docs.mistral.ai/api/> (`UsageInfo` schema). The three counts are
136/// always present; the remaining fields are populated by Mistral on a best-effort
137/// basis (e.g. cached-token information appears once a prompt is large enough to
138/// be cached).
139#[derive(Clone, Debug, Default, Deserialize, Serialize)]
140pub struct Usage {
141    pub completion_tokens: usize,
142    pub prompt_tokens: usize,
143    pub total_tokens: usize,
144    /// Capacity tier that served the request, when reported.
145    #[serde(default, skip_serializing_if = "Option::is_none")]
146    pub service_tier: Option<String>,
147    /// Duration in seconds of audio tokens in the prompt (audio-input models only).
148    #[serde(default, skip_serializing_if = "Option::is_none")]
149    pub prompt_audio_seconds: Option<u64>,
150    /// Total cached prompt tokens reported at the top level. Some Mistral
151    /// responses populate this in addition to (or instead of)
152    /// `prompt_tokens_details.cached_tokens`.
153    #[serde(default, skip_serializing_if = "Option::is_none")]
154    pub num_cached_tokens: Option<u64>,
155    /// Prompt token breakdown. The singular `prompt_token_details` key is not an
156    /// alias because responses may contain both spellings.
157    #[serde(default, skip_serializing_if = "Option::is_none")]
158    pub prompt_tokens_details: Option<PromptTokensDetails>,
159}
160
161#[cfg(test)]
162mod tests;