use serde::{Deserialize, Serialize};
use serde_json::Value;
use crate::completion::provider_options::reply_field;
use crate::completion::{ExtensionOptions, ProviderExtension, ReplyExtras};
use crate::message::Api;
use crate::providers::openai::extension::Prediction;
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct MistralExt;
impl ProviderExtension for MistralExt {
const PROVIDER: &'static str = super::PROVIDER_NAME;
type Options = MistralOptions;
type Extras = MistralExtras;
}
#[non_exhaustive]
#[derive(Clone, Debug, Default, PartialEq, Serialize)]
pub struct MistralOptions {
#[serde(rename = "*")]
pub shared: MistralShared,
}
#[non_exhaustive]
#[derive(Clone, Debug, Default, PartialEq, Serialize)]
pub struct MistralShared {
#[serde(skip_serializing_if = "Option::is_none")]
pub prompt_mode: Option<PromptMode>,
#[serde(skip_serializing_if = "Option::is_none")]
pub safe_prompt: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub prompt_cache_key: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub frequency_penalty: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub presence_penalty: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub prediction: Option<Prediction>,
}
#[non_exhaustive]
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize)]
#[serde(rename_all = "lowercase")]
pub enum PromptMode {
Reasoning,
}
impl MistralOptions {
pub fn new() -> Self {
Self::default()
}
pub fn prompt_mode(mut self, mode: PromptMode) -> Self {
self.shared.prompt_mode = Some(mode);
self
}
pub fn safe_prompt(mut self, safe: bool) -> Self {
self.shared.safe_prompt = Some(safe);
self
}
pub fn prompt_cache_key(mut self, key: impl Into<String>) -> Self {
self.shared.prompt_cache_key = Some(key.into());
self
}
pub fn frequency_penalty(mut self, penalty: f64) -> Self {
self.shared.frequency_penalty = Some(penalty);
self
}
pub fn presence_penalty(mut self, penalty: f64) -> Self {
self.shared.presence_penalty = Some(penalty);
self
}
pub fn prediction(mut self, content: impl Into<String>) -> Self {
self.shared.prediction = Some(Prediction::content(content));
self
}
}
impl ExtensionOptions for MistralOptions {
type Ext = MistralExt;
}
#[non_exhaustive]
#[derive(Clone, Debug, Default, PartialEq)]
pub struct MistralExtras {
pub service_tier: Option<String>,
pub prompt_audio_seconds: Option<u64>,
pub num_cached_tokens: Option<u64>,
pub prompt_tokens_details: Option<MistralPromptTokens>,
}
#[non_exhaustive]
#[derive(Clone, Debug, Default, PartialEq, Eq, Deserialize)]
pub struct MistralPromptTokens {
#[serde(default)]
pub cached_tokens: Option<u64>,
#[serde(default)]
pub audio_tokens: Option<u64>,
}
impl ReplyExtras for MistralExtras {
fn from_reply(_api: &Api, raw: &Value) -> Result<Self, serde_json::Error> {
Ok(Self {
service_tier: reply_field(raw, "/usage/service_tier")?,
prompt_audio_seconds: reply_field(raw, "/usage/prompt_audio_seconds")?,
num_cached_tokens: reply_field(raw, "/usage/num_cached_tokens")?,
prompt_tokens_details: reply_field(raw, "/usage/prompt_tokens_details")?,
})
}
}
#[cfg(test)]
mod tests;