use std::fmt;
use std::str::FromStr;
use serde::{Deserialize, Deserializer, Serialize, Serializer};
use crate::providers::{ProviderConfig, ProviderKind};
pub const MAX: f64 = 2.0;
#[derive(Debug, Clone, Copy, PartialEq)]
pub enum Temperature {
Default,
Value(f64),
}
impl fmt::Display for Temperature {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Temperature::Default => f.write_str("model default"),
Temperature::Value(v) => write!(f, "{v}"),
}
}
}
impl FromStr for Temperature {
type Err = String;
fn from_str(s: &str) -> Result<Self, String> {
let s = s.trim();
if s.eq_ignore_ascii_case("default") || s.eq_ignore_ascii_case("model default") {
return Ok(Temperature::Default);
}
match s.parse::<f64>() {
Ok(v) if (0.0..=MAX).contains(&v) => Ok(Temperature::Value(v)),
_ => Err(format!("temperature must be a number from 0 to {MAX}, or \"default\"; got {s:?}")),
}
}
}
impl Serialize for Temperature {
fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
match self {
Temperature::Default => serializer.serialize_str("default"),
Temperature::Value(v) => serializer.serialize_f64(*v),
}
}
}
impl<'de> Deserialize<'de> for Temperature {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
#[derive(Deserialize)]
#[serde(untagged)]
enum Raw {
Number(f64),
Int(i64),
Text(String),
}
match Raw::deserialize(deserializer)? {
Raw::Number(v) if v.is_finite() => Ok(Temperature::Value(v)),
Raw::Number(v) => Err(serde::de::Error::custom(format!("temperature must be finite, not {v}"))),
Raw::Int(v) => Ok(Temperature::Value(v as f64)),
Raw::Text(s) if s.trim().eq_ignore_ascii_case("default") => Ok(Temperature::Default),
Raw::Text(s) => {
Err(serde::de::Error::custom(format!("temperature must be a number or \"default\", not {s:?}")))
}
}
}
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
#[serde(default)]
pub struct ModelSettings {
#[serde(skip_serializing_if = "Option::is_none")]
pub temperature: Option<Temperature>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Source {
Model,
Provider,
Global,
ExtraBody,
Required,
}
impl Source {
pub fn label(self) -> &'static str {
match self {
Source::Model => "set for this model",
Source::Provider => "set for this provider",
Source::Global => "global setting",
Source::ExtraBody => "from the provider's extra_body",
Source::Required => "required: the model accepts no temperature",
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct Resolved {
pub effective: Temperature,
pub source: Source,
pub fixed: Option<String>,
pub warning: Option<String>,
}
impl Resolved {
pub fn value(&self) -> Option<f64> {
match self.effective {
Temperature::Default => None,
Temperature::Value(v) => Some(v),
}
}
pub fn describe(&self) -> String {
format!("{} ({})", self.effective, self.source.label())
}
}
pub fn fixed_reason(kind: Option<ProviderKind>, provider: &ProviderConfig, model: &str) -> Option<String> {
if provider.drop_params.as_ref().is_some_and(|d| d.iter().any(|p| p == "temperature")) {
return Some("the provider drops temperature (drop_params)".to_string());
}
if kind == Some(ProviderKind::GithubCopilot) && crate::providers::github_copilot::is_reasoning_model(model) {
return Some(format!("{model} is a reasoning model and rejects a custom temperature"));
}
None
}
pub fn resolve(global: Temperature, kind: Option<ProviderKind>, provider: &ProviderConfig, model: &str) -> Resolved {
let (chosen, mut source) = if let Some(t) = provider.models.get(model).and_then(|m| m.temperature) {
(t, Source::Model)
} else if let Some(t) = provider.temperature {
(t, Source::Provider)
} else {
(global, Source::Global)
};
let mut effective = chosen;
let mut extra_body_warning = None;
if let Some(v) = provider.extra_body.as_ref().and_then(|b| b.get("temperature")).and_then(|v| match v {
toml::Value::Float(f) => Some(*f),
toml::Value::Integer(i) => Some(*i as f64),
_ => None,
}) {
if v.is_finite() {
effective = Temperature::Value(v);
source = Source::ExtraBody;
} else {
extra_body_warning =
Some(format!("extra_body temperature {v} is not finite and is ignored; using {chosen}"));
}
}
if let Some(reason) = fixed_reason(kind, provider, model) {
let warning = match (effective, source) {
(Temperature::Value(v), Source::Model | Source::Provider | Source::ExtraBody) => {
Some(format!("temperature {v} ({}) is ignored: {reason}; using the model default", source.label()))
}
_ => extra_body_warning,
};
return Resolved { effective: Temperature::Default, source: Source::Required, fixed: Some(reason), warning };
}
let mut warning = extra_body_warning;
let anthropic_messages = kind == Some(ProviderKind::Anthropic)
|| (kind == Some(ProviderKind::GithubCopilot)
&& crate::providers::github_copilot::uses_anthropic_messages(model));
if anthropic_messages && let Temperature::Value(v) = effective {
if v > 1.0 {
warning = Some(format!("temperature {v} is above Anthropic's maximum of 1; sending 1"));
effective = Temperature::Value(1.0);
} else if v < 0.0 {
warning = Some(format!("temperature {v} is below Anthropic's minimum of 0; sending 0"));
effective = Temperature::Value(0.0);
}
}
Resolved { effective, source, fixed: None, warning }
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::Config;
fn provider(toml_text: &str) -> ProviderConfig {
toml::from_str(toml_text).unwrap()
}
#[test]
fn parses_numbers_and_default() {
let config: Config = toml::from_str("temperature = 1").unwrap();
assert_eq!(config.temperature, Temperature::Value(1.0));
let config: Config = toml::from_str("temperature = 0.2").unwrap();
assert_eq!(config.temperature, Temperature::Value(0.2));
let config: Config = toml::from_str("temperature = \"default\"").unwrap();
assert_eq!(config.temperature, Temperature::Default);
assert!(toml::from_str::<Config>("temperature = \"warm\"").is_err());
assert_eq!("default".parse::<Temperature>(), Ok(Temperature::Default));
assert_eq!(" 0.5 ".parse::<Temperature>(), Ok(Temperature::Value(0.5)));
assert!("3".parse::<Temperature>().is_err());
assert!("hot".parse::<Temperature>().is_err());
let text = toml::to_string(&Config { temperature: Temperature::Default, ..Default::default() }).unwrap();
assert!(text.contains("temperature = \"default\""), "{text}");
}
#[test]
fn rejects_non_finite_values() {
for text in ["temperature = nan", "temperature = inf", "temperature = -inf"] {
let err = toml::from_str::<Config>(text).unwrap_err();
assert!(err.message().contains("temperature must be finite"), "{text}: {err}");
}
}
#[test]
fn most_specific_setting_wins() {
let global = Temperature::Value(0.7);
let p = provider("temperature = 0.4\n[models.\"big\"]\ntemperature = \"default\"\n");
let r = resolve(global, Some(ProviderKind::Openai), &p, "big");
assert_eq!((r.effective, r.source, r.value()), (Temperature::Default, Source::Model, None));
let r = resolve(global, Some(ProviderKind::Openai), &p, "small");
assert_eq!((r.value(), r.source), (Some(0.4), Source::Provider));
let r = resolve(global, Some(ProviderKind::Openai), &ProviderConfig::default(), "small");
assert_eq!((r.value(), r.source), (Some(0.7), Source::Global));
assert_eq!(r.describe(), "0.7 (global setting)");
assert_eq!(r.warning, None);
let p = provider("temperature = 0.4\nextra_body = { temperature = 0.1 }\n");
let r = resolve(global, Some(ProviderKind::Openai), &p, "m");
assert_eq!((r.value(), r.source), (Some(0.1), Source::ExtraBody));
for text in ["nan", "inf", "-inf"] {
let p = provider(&format!("temperature = 0.4\nextra_body = {{ temperature = {text} }}\n"));
let r = resolve(global, Some(ProviderKind::Openai), &p, "m");
assert_eq!((r.value(), r.source), (Some(0.4), Source::Provider), "{text}");
assert!(r.warning.as_deref().unwrap().contains("not finite"), "{text}: {:?}", r.warning);
}
}
#[test]
fn models_without_temperature_only_use_the_default() {
let global = Temperature::Value(0.7);
let r = resolve(global, Some(ProviderKind::GithubCopilot), &ProviderConfig::default(), "gpt-5");
assert_eq!((r.value(), r.source), (None, Source::Required));
assert!(r.fixed.as_deref().unwrap().contains("reasoning model"));
assert_eq!(r.warning, None);
let r = resolve(global, Some(ProviderKind::GithubCopilot), &ProviderConfig::default(), "gpt-4.1");
assert_eq!(r.value(), Some(0.7));
let p = provider("drop_params = [\"temperature\"]\n[models.\"k3\"]\ntemperature = 0.3\n");
let r = resolve(global, Some(ProviderKind::Openai), &p, "k3");
assert_eq!(r.value(), None);
let warning = r.warning.unwrap();
assert!(warning.contains("0.3") && warning.contains("drop_params"), "{warning}");
let p = provider("drop_params = [\"temperature\"]\ntemperature = \"default\"\n");
assert_eq!(resolve(global, Some(ProviderKind::Openai), &p, "k3").warning, None);
let p = provider("drop_params = [\"temperature\"]\nextra_body = { temperature = nan }\n");
let r = resolve(global, Some(ProviderKind::Openai), &p, "k3");
assert_eq!((r.value(), r.source), (None, Source::Required));
assert!(r.warning.as_deref().unwrap().contains("not finite"), "{:?}", r.warning);
}
#[test]
fn anthropic_range_is_checked() {
let p = provider("temperature = 1.5\n");
let r = resolve(Temperature::Value(0.7), Some(ProviderKind::Anthropic), &p, "claude");
assert_eq!(r.value(), Some(1.0));
assert!(r.warning.unwrap().contains("Anthropic's maximum"));
let p = provider("temperature = -0.5\n");
let r = resolve(Temperature::Value(0.7), Some(ProviderKind::Anthropic), &p, "claude");
assert_eq!(r.value(), Some(0.0));
assert!(r.warning.unwrap().contains("Anthropic's minimum"));
let r = resolve(Temperature::Value(0.7), Some(ProviderKind::Anthropic), &ProviderConfig::default(), "claude");
assert_eq!((r.value(), r.warning), (Some(0.7), None));
}
#[test]
fn copilot_claude_uses_the_anthropic_range() {
let p = provider("temperature = 1.5\n");
let r = resolve(Temperature::Value(0.7), Some(ProviderKind::GithubCopilot), &p, "claude-sonnet-4.5");
assert_eq!(r.value(), Some(1.0));
assert!(r.warning.unwrap().contains("Anthropic's maximum"));
let p = provider("temperature = -0.5\n");
let r = resolve(Temperature::Value(0.7), Some(ProviderKind::GithubCopilot), &p, "claude-opus-5");
assert_eq!(r.value(), Some(0.0));
assert!(r.warning.unwrap().contains("Anthropic's minimum"));
let r = resolve(
Temperature::Value(0.7),
Some(ProviderKind::GithubCopilot),
&ProviderConfig::default(),
"claude-haiku-4.5",
);
assert_eq!((r.value(), r.warning), (Some(0.7), None));
let p = provider("temperature = 1.5\n");
let r = resolve(Temperature::Value(0.7), Some(ProviderKind::GithubCopilot), &p, "gpt-4.1");
assert_eq!((r.value(), r.warning), (Some(1.5), None));
let r = resolve(Temperature::Value(0.7), Some(ProviderKind::GithubCopilot), &p, "claude-sonnet-3.5");
assert_eq!((r.value(), r.warning), (Some(1.5), None));
let r = resolve(Temperature::Value(0.7), None, &p, "claude-sonnet-4.5");
assert_eq!((r.value(), r.warning), (Some(1.5), None));
}
#[test]
fn user_model_settings_merge_over_presets() {
let preset = provider("temperature = 0.2\n[models.\"a\"]\ntemperature = 0.1\n");
let user = provider("[models.\"b\"]\ntemperature = \"default\"\n");
let merged = preset.merged_with(&user);
assert_eq!(merged.temperature, Some(Temperature::Value(0.2)));
assert_eq!(merged.models["a"].temperature, Some(Temperature::Value(0.1)));
assert_eq!(merged.models["b"].temperature, Some(Temperature::Default));
}
}