use crate::{
client::{self, BearerAuth, DebugExt, Provider},
providers::mistral::MistralModelLister,
};
use serde::{Deserialize, Serialize};
use std::fmt::Debug;
const MISTRAL_API_BASE_URL: &str = "https://api.mistral.ai";
#[derive(Debug, Default, Clone, Copy)]
pub struct MistralExt;
#[derive(Debug, Default, Clone, Copy)]
pub struct MistralBuilder;
type MistralApiKey = BearerAuth;
pub type Client<H = reqwest::Client> = client::Client<MistralExt, H>;
pub type ClientBuilder<H = crate::markers::Missing> =
client::ClientBuilder<MistralBuilder, MistralApiKey, H>;
impl Provider for MistralExt {
type Builder = MistralBuilder;
const VERIFY_PATH: &'static str = "/v1/models";
}
impl crate::providers::openai::completion::OpenAICompatibleProvider for MistralExt {
const PROVIDER_NAME: &'static str = "mistral";
const REQUEST_ID_HEADER: Option<&'static str> = Some("mistral-correlation-id");
type StreamingUsage = Usage;
const EMITS_COMPLETE_SINGLE_CHUNK_TOOL_CALLS: bool = true;
const STREAM_INCLUDE_USAGE: bool = false;
type Response = super::CompletionResponse;
fn completion_path(&self, _model: &str) -> String {
"/v1/chat/completions".to_string()
}
fn finalize_request_body(
&self,
body: &mut serde_json::Value,
) -> Result<(), crate::completion::CompletionError> {
let Some(map) = body.as_object_mut() else {
return Ok(());
};
if let Some(tool_choice) = map.get_mut("tool_choice")
&& tool_choice.as_str() == Some("required")
{
*tool_choice = serde_json::Value::String("any".to_string());
}
let forces_a_tool_call = map
.get("tool_choice")
.is_some_and(|choice| !matches!(choice.as_str(), Some("auto" | "none")));
let has_tools = map
.get("tools")
.and_then(serde_json::Value::as_array)
.is_some_and(|tools| !tools.is_empty());
let has_structured_format = map
.get("response_format")
.and_then(|format| format.get("type"))
.and_then(serde_json::Value::as_str)
.is_some_and(|kind| matches!(kind, "json_schema" | "json_object"));
if forces_a_tool_call && has_tools && has_structured_format {
tracing::debug!(
"relaxing tool_choice to `auto`: Mistral rejects a forced tool choice \
alongside a response format"
);
map.insert(
"tool_choice".to_string(),
serde_json::Value::String("auto".to_string()),
);
}
if let Some(messages) = map
.get_mut("messages")
.and_then(serde_json::Value::as_array_mut)
{
for message in messages {
let Some(message) = message.as_object_mut() else {
continue;
};
let is_assistant =
message.get("role").and_then(serde_json::Value::as_str) == Some("assistant");
if let Some(content) = message.get_mut("content") {
super::completion::normalize_request_content(content)?;
}
if is_assistant {
if !message.contains_key("content") {
message.insert(
"content".to_string(),
serde_json::Value::String(String::new()),
);
}
message
.entry("prefix")
.or_insert(serde_json::Value::Bool(false));
message.remove("reasoning_content");
}
}
}
Ok(())
}
}
client::impl_capabilities!(
MistralExt,
completion = super::CompletionModel<H>,
embeddings = super::EmbeddingModel<H>,
transcription = super::TranscriptionModel<H>,
model_listing = MistralModelLister<H>,
);
impl DebugExt for MistralExt {}
client::impl_default_provider_builder!(
MistralBuilder => MistralExt,
api_key = MistralApiKey,
base_url = MISTRAL_API_BASE_URL,
);
client::impl_provider_client!(Client, input = String, api_key_env = "MISTRAL_API_KEY");
#[derive(Clone, Debug, Default, Deserialize, Serialize)]
pub struct PromptTokensDetails {
#[serde(default)]
pub cached_tokens: u64,
#[serde(default)]
pub audio_tokens: u64,
}
#[derive(Clone, Debug, Default, Deserialize, Serialize)]
pub struct Usage {
pub completion_tokens: usize,
pub prompt_tokens: usize,
pub total_tokens: usize,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub service_tier: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub prompt_audio_seconds: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub num_cached_tokens: Option<u64>,
#[serde(
default,
alias = "prompt_token_details",
skip_serializing_if = "Option::is_none"
)]
pub prompt_tokens_details: Option<PromptTokensDetails>,
}
impl Usage {
pub fn cached_tokens(&self) -> u64 {
self.prompt_tokens_details
.as_ref()
.map(|d| d.cached_tokens)
.or(self.num_cached_tokens)
.unwrap_or(0)
}
pub fn audio_tokens(&self) -> u64 {
self.prompt_tokens_details
.as_ref()
.map_or(0, |details| details.audio_tokens)
}
pub fn input_tokens(&self) -> u64 {
self.prompt_tokens as u64 + self.audio_tokens()
}
}
impl From<&Usage> for crate::completion::Usage {
fn from(usage: &Usage) -> Self {
crate::providers::internal::completion_usage(
usage.input_tokens(),
usage.completion_tokens as u64,
usage.total_tokens as u64,
usage.cached_tokens(),
)
}
}
impl From<Usage> for crate::completion::Usage {
fn from(usage: Usage) -> Self {
Self::from(&usage)
}
}
impl std::fmt::Display for Usage {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"Prompt tokens: {} Total tokens: {}",
self.prompt_tokens, self.total_tokens
)
}
}
#[cfg(test)]
mod tests {
use super::Usage;
#[test]
fn test_client_initialization() {
let _client =
crate::providers::mistral::Client::new("dummy-key").expect("Client::new() failed");
let builder: crate::providers::mistral::ClientBuilder =
crate::providers::mistral::Client::builder().api_key("dummy-key");
let _client_from_builder = builder.build().expect("Client::builder() failed");
}
#[test]
fn usage_retains_live_service_tier() {
let usage: Usage = serde_json::from_value(serde_json::json!({
"completion_tokens": 4,
"prompt_tokens": 20,
"total_tokens": 24,
"prompt_tokens_details": { "cached_tokens": 0 },
"service_tier": "standard"
}))
.expect("live Mistral usage should deserialize");
assert_eq!(usage.service_tier.as_deref(), Some("standard"));
assert_eq!(
serde_json::to_value(usage).expect("Mistral usage should serialize")["service_tier"],
"standard"
);
}
}