rig_core/providers/mistral/
extension.rs1use 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#[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#[non_exhaustive]
33#[derive(Clone, Debug, Default, PartialEq, Serialize)]
34pub struct MistralOptions {
35 #[serde(rename = "*")]
37 pub shared: MistralShared,
38}
39
40#[non_exhaustive]
42#[derive(Clone, Debug, Default, PartialEq, Serialize)]
43pub struct MistralShared {
44 #[serde(skip_serializing_if = "Option::is_none")]
46 pub prompt_mode: Option<PromptMode>,
47 #[serde(skip_serializing_if = "Option::is_none")]
49 pub safe_prompt: Option<bool>,
50 #[serde(skip_serializing_if = "Option::is_none")]
52 pub prompt_cache_key: Option<String>,
53 #[serde(skip_serializing_if = "Option::is_none")]
55 pub frequency_penalty: Option<f64>,
56 #[serde(skip_serializing_if = "Option::is_none")]
58 pub presence_penalty: Option<f64>,
59 #[serde(skip_serializing_if = "Option::is_none")]
61 pub prediction: Option<Prediction>,
62}
63
64#[non_exhaustive]
66#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize)]
67#[serde(rename_all = "lowercase")]
68pub enum PromptMode {
69 Reasoning,
71}
72
73impl MistralOptions {
74 pub fn new() -> Self {
76 Self::default()
77 }
78
79 pub fn prompt_mode(mut self, mode: PromptMode) -> Self {
81 self.shared.prompt_mode = Some(mode);
82 self
83 }
84
85 pub fn safe_prompt(mut self, safe: bool) -> Self {
87 self.shared.safe_prompt = Some(safe);
88 self
89 }
90
91 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 pub fn frequency_penalty(mut self, penalty: f64) -> Self {
99 self.shared.frequency_penalty = Some(penalty);
100 self
101 }
102
103 pub fn presence_penalty(mut self, penalty: f64) -> Self {
105 self.shared.presence_penalty = Some(penalty);
106 self
107 }
108
109 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#[non_exhaustive]
122#[derive(Clone, Debug, Default, PartialEq)]
123pub struct MistralExtras {
124 pub service_tier: Option<String>,
126 pub prompt_audio_seconds: Option<u64>,
128 pub num_cached_tokens: Option<u64>,
130 pub prompt_tokens_details: Option<MistralPromptTokens>,
132}
133
134#[non_exhaustive]
136#[derive(Clone, Debug, Default, PartialEq, Eq, Deserialize)]
137pub struct MistralPromptTokens {
138 #[serde(default)]
140 pub cached_tokens: Option<u64>,
141 #[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;