Skip to main content

rig_core/providers/mistral/
extension.rs

1//! Mistral's typed request options and reply extras
2//! (<https://docs.mistral.ai/api/endpoint/chat>).
3//!
4//! ```
5//! use rig_core::completion::CompletionRequest;
6//! use rig_core::providers::mistral::extension::{MistralOptions, PromptMode};
7//!
8//! let options = MistralOptions::new().prompt_mode(PromptMode::Reasoning).safe_prompt(true);
9//! let request = CompletionRequest::new("hi").provider_option(options);
10//! # let _ = request;
11//! ```
12
13use serde::{Deserialize, Serialize};
14use serde_json::Value;
15
16use crate::completion::provider_options::reply_field;
17use crate::completion::{ExtensionOptions, ProviderExtension, ReplyExtras};
18use crate::message::Api;
19use crate::providers::openai::extension::Prediction;
20
21/// Mistral's extension marker.
22#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
23pub struct MistralExt;
24
25impl ProviderExtension for MistralExt {
26    const PROVIDER: &'static str = super::PROVIDER_NAME;
27    type Options = MistralOptions;
28    type Extras = MistralExtras;
29}
30
31/// Mistral's request options.
32#[non_exhaustive]
33#[derive(Clone, Debug, Default, PartialEq, Serialize)]
34pub struct MistralOptions {
35    /// The fields every route takes.
36    #[serde(rename = "*")]
37    pub shared: MistralShared,
38}
39
40/// The fields Mistral takes.
41#[non_exhaustive]
42#[derive(Clone, Debug, Default, PartialEq, Serialize)]
43pub struct MistralShared {
44    /// The system prompt Mistral adds for a reasoning model.
45    #[serde(skip_serializing_if = "Option::is_none")]
46    pub prompt_mode: Option<PromptMode>,
47    /// Whether Mistral prepends its safety prompt.
48    #[serde(skip_serializing_if = "Option::is_none")]
49    pub safe_prompt: Option<bool>,
50    /// The prompt-cache routing key.
51    #[serde(skip_serializing_if = "Option::is_none")]
52    pub prompt_cache_key: Option<String>,
53    /// Penalizes tokens by how often they already appear.
54    #[serde(skip_serializing_if = "Option::is_none")]
55    pub frequency_penalty: Option<f64>,
56    /// Penalizes tokens that already appear.
57    #[serde(skip_serializing_if = "Option::is_none")]
58    pub presence_penalty: Option<f64>,
59    /// Predicted output.
60    #[serde(skip_serializing_if = "Option::is_none")]
61    pub prediction: Option<Prediction>,
62}
63
64/// The system prompt Mistral adds.
65#[non_exhaustive]
66#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize)]
67#[serde(rename_all = "lowercase")]
68pub enum PromptMode {
69    /// The reasoning system prompt.
70    Reasoning,
71}
72
73impl MistralOptions {
74    /// No option set.
75    pub fn new() -> Self {
76        Self::default()
77    }
78
79    /// Add Mistral's `mode` system prompt.
80    pub fn prompt_mode(mut self, mode: PromptMode) -> Self {
81        self.shared.prompt_mode = Some(mode);
82        self
83    }
84
85    /// Whether Mistral prepends its safety prompt.
86    pub fn safe_prompt(mut self, safe: bool) -> Self {
87        self.shared.safe_prompt = Some(safe);
88        self
89    }
90
91    /// Route the prompt cache by `key`.
92    pub fn prompt_cache_key(mut self, key: impl Into<String>) -> Self {
93        self.shared.prompt_cache_key = Some(key.into());
94        self
95    }
96
97    /// Set the frequency penalty.
98    pub fn frequency_penalty(mut self, penalty: f64) -> Self {
99        self.shared.frequency_penalty = Some(penalty);
100        self
101    }
102
103    /// Set the presence penalty.
104    pub fn presence_penalty(mut self, penalty: f64) -> Self {
105        self.shared.presence_penalty = Some(penalty);
106        self
107    }
108
109    /// Predict the output as `content`.
110    pub fn prediction(mut self, content: impl Into<String>) -> Self {
111        self.shared.prediction = Some(Prediction::content(content));
112        self
113    }
114}
115
116impl ExtensionOptions for MistralOptions {
117    type Ext = MistralExt;
118}
119
120/// Mistral's reply fields. Each is `None` when the reply lacks it.
121#[non_exhaustive]
122#[derive(Clone, Debug, Default, PartialEq)]
123pub struct MistralExtras {
124    /// The service tier that served the request.
125    pub service_tier: Option<String>,
126    /// Seconds of audio in the prompt.
127    pub prompt_audio_seconds: Option<u64>,
128    /// Prompt tokens read from the cache, as older replies report them.
129    pub num_cached_tokens: Option<u64>,
130    /// Prompt token details.
131    pub prompt_tokens_details: Option<MistralPromptTokens>,
132}
133
134/// Where a Mistral reply's prompt tokens went.
135#[non_exhaustive]
136#[derive(Clone, Debug, Default, PartialEq, Eq, Deserialize)]
137pub struct MistralPromptTokens {
138    /// Tokens read from the cache.
139    #[serde(default)]
140    pub cached_tokens: Option<u64>,
141    /// Audio tokens.
142    #[serde(default)]
143    pub audio_tokens: Option<u64>,
144}
145
146impl ReplyExtras for MistralExtras {
147    fn from_reply(_api: &Api, raw: &Value) -> Result<Self, serde_json::Error> {
148        Ok(Self {
149            service_tier: reply_field(raw, "/usage/service_tier")?,
150            prompt_audio_seconds: reply_field(raw, "/usage/prompt_audio_seconds")?,
151            num_cached_tokens: reply_field(raw, "/usage/num_cached_tokens")?,
152            prompt_tokens_details: reply_field(raw, "/usage/prompt_tokens_details")?,
153        })
154    }
155}
156
157#[cfg(test)]
158mod tests;