use serde::{Deserialize, Serialize};
use crate::shared::config::CloudProvider;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum ReasoningEffort {
None,
Minimal,
Low,
Medium,
High,
XHigh,
}
impl ReasoningEffort {
pub fn as_wire(self) -> &'static str {
match self {
ReasoningEffort::None => "none",
ReasoningEffort::Minimal => "minimal",
ReasoningEffort::Low => "low",
ReasoningEffort::Medium => "medium",
ReasoningEffort::High => "high",
ReasoningEffort::XHigh => "xhigh",
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum Verbosity {
Low,
Medium,
High,
}
impl Verbosity {
pub fn as_wire(self) -> &'static str {
match self {
Verbosity::Low => "low",
Verbosity::Medium => "medium",
Verbosity::High => "high",
}
}
}
#[derive(Debug, Clone, PartialEq, Default, Serialize, Deserialize)]
#[serde(default)]
pub struct SamplingConfig {
pub temperature: Option<f32>,
pub dynatemp_range: Option<f32>,
pub dynatemp_exponent: Option<f32>,
pub top_k: Option<i64>,
pub top_p: Option<f32>,
pub min_p: Option<f32>,
pub top_n_sigma: Option<f32>,
pub typical_p: Option<f32>,
pub adaptive_target: Option<f32>,
pub adaptive_decay: Option<f32>,
pub frequency_penalty: Option<f32>,
pub presence_penalty: Option<f32>,
pub repeat_penalty: Option<f32>,
pub repeat_last_n: Option<i64>,
pub dry_multiplier: Option<f32>,
pub dry_base: Option<f32>,
pub dry_allowed_length: Option<i64>,
pub dry_penalty_last_n: Option<i64>,
pub dry_sequence_breakers: Option<Vec<String>>,
pub xtc_probability: Option<f32>,
pub xtc_threshold: Option<f32>,
pub mirostat: Option<i64>,
pub mirostat_tau: Option<f32>,
pub mirostat_eta: Option<f32>,
pub max_tokens: Option<usize>,
pub seed: Option<i64>,
pub samplers: Option<Vec<String>>,
pub thinking: Option<bool>,
pub reasoning_effort: Option<ReasoningEffort>,
pub reasoning_budget: Option<i64>,
pub verbosity: Option<Verbosity>,
}
impl SamplingConfig {
pub fn retain_supported(
&self,
provider: Option<CloudProvider>,
endpoint: Option<&[String]>,
) -> SamplingConfig {
let supported = available_sampling_fields(provider, endpoint);
let Ok(serde_json::Value::Object(mut map)) = serde_json::to_value(self) else {
return self.clone();
};
map.retain(|k, _| supported.contains(&k.as_str()));
serde_json::from_value(serde_json::Value::Object(map)).unwrap_or_else(|_| self.clone())
}
}
pub fn resolve(
chat_override: Option<&SamplingConfig>,
profile_default: Option<&SamplingConfig>,
global: &SamplingConfig,
) -> SamplingConfig {
chat_override
.or(profile_default)
.cloned()
.unwrap_or_else(|| global.clone())
}
pub const SETTABLE_SAMPLING_FIELDS: &[&str] = &[
"temperature",
"dynatemp_range",
"dynatemp_exponent",
"top_k",
"top_p",
"min_p",
"top_n_sigma",
"typical_p",
"adaptive_target",
"adaptive_decay",
"frequency_penalty",
"presence_penalty",
"repeat_penalty",
"repeat_last_n",
"dry_multiplier",
"dry_base",
"dry_allowed_length",
"dry_penalty_last_n",
"dry_sequence_breakers",
"xtc_probability",
"xtc_threshold",
"mirostat",
"mirostat_tau",
"mirostat_eta",
"max_tokens",
"seed",
"samplers",
"thinking",
"reasoning_effort",
"verbosity",
];
fn catalogue_aliases(field: &str) -> &'static [&'static str] {
match field {
"temperature" => &["temperature"],
"top_k" => &["top_k"],
"top_p" => &["top_p"],
"min_p" => &["min_p"],
"max_tokens" => &["max_tokens"],
"seed" => &["seed"],
"frequency_penalty" => &["frequency_penalty"],
"presence_penalty" => &["presence_penalty"],
"verbosity" => &["verbosity"],
"repeat_penalty" => &["repetition_penalty"],
"thinking" | "reasoning_effort" => &["reasoning", "include_reasoning"],
_ => &[],
}
}
pub fn available_sampling_fields(
provider: Option<CloudProvider>,
endpoint: Option<&[String]>,
) -> Vec<&'static str> {
let base = supported_sampling_fields(provider);
let Some(published) = endpoint
.filter(|_| provider.is_none())
.filter(|p| !p.is_empty())
else {
return base.to_vec();
};
base.iter()
.copied()
.filter(|field| {
catalogue_aliases(field)
.iter()
.any(|alias| published.iter().any(|p| p == alias))
})
.collect()
}
pub fn supported_sampling_fields(provider: Option<CloudProvider>) -> &'static [&'static str] {
match provider {
None => SETTABLE_SAMPLING_FIELDS,
Some(CloudProvider::OpenAi) => &["max_tokens", "thinking", "reasoning_effort", "verbosity"],
Some(CloudProvider::Gemini) => &[
"temperature",
"top_p",
"top_k",
"max_tokens",
"seed",
"frequency_penalty",
"presence_penalty",
"thinking",
"reasoning_effort",
],
Some(CloudProvider::Claude) => &["max_tokens", "thinking", "reasoning_effort"],
Some(CloudProvider::Grok) => &[
"temperature",
"top_p",
"max_tokens",
"seed",
"thinking",
"reasoning_effort",
],
}
}
#[cfg(test)]
mod tests {
use super::*;
fn r1_catalogue() -> Vec<String> {
[
"frequency_penalty",
"include_reasoning",
"max_tokens",
"presence_penalty",
"reasoning",
"repetition_penalty",
"response_format",
"seed",
"stop",
"temperature",
"tool_choice",
"tools",
"top_k",
"top_p",
]
.iter()
.map(|s| s.to_string())
.collect()
}
#[test]
fn a_published_catalogue_narrows_the_offer_to_what_it_lists() {
let fields = available_sampling_fields(None, Some(&r1_catalogue()));
for kept in [
"temperature",
"top_p",
"top_k",
"max_tokens",
"seed",
"frequency_penalty",
"presence_penalty",
"repeat_penalty",
"thinking",
"reasoning_effort",
] {
assert!(
fields.contains(&kept),
"{kept} is published, under some name"
);
}
for dropped in [
"min_p",
"typical_p",
"top_n_sigma",
"dynatemp_range",
"adaptive_target",
"mirostat",
"dry_multiplier",
"xtc_probability",
"samplers",
"repeat_last_n",
] {
assert!(
!fields.contains(&dropped),
"{dropped} is not in the catalogue and must stop being offered"
);
}
}
#[test]
fn silence_and_the_clouds_are_left_alone() {
assert_eq!(
available_sampling_fields(None, None),
SETTABLE_SAMPLING_FIELDS.to_vec(),
"no list — every field, as before"
);
assert_eq!(
available_sampling_fields(None, Some(&[])),
SETTABLE_SAMPLING_FIELDS.to_vec(),
"an empty list is silence, not a claim that nothing is taken"
);
let openai = available_sampling_fields(Some(CloudProvider::OpenAi), Some(&r1_catalogue()));
assert_eq!(
openai,
supported_sampling_fields(Some(CloudProvider::OpenAi)).to_vec(),
"a cloud's table is not a gateway catalogue's business"
);
}
#[test]
fn default_is_all_none() {
let s = SamplingConfig::default();
assert!(s.temperature.is_none());
assert!(s.max_tokens.is_none());
assert!(s.reasoning_effort.is_none());
}
#[test]
fn reasoning_effort_wire_strings() {
assert_eq!(ReasoningEffort::Medium.as_wire(), "medium");
assert_eq!(ReasoningEffort::None.as_wire(), "none");
}
#[test]
fn serde_roundtrip() {
let s = SamplingConfig {
temperature: Some(0.7),
top_k: Some(40),
thinking: Some(true),
reasoning_effort: Some(ReasoningEffort::High),
..Default::default()
};
let json = serde_json::to_string(&s).unwrap();
let back: SamplingConfig = serde_json::from_str(&json).unwrap();
assert_eq!(s, back);
}
#[test]
fn supported_fields_mirror_wire_dialect() {
assert_eq!(supported_sampling_fields(None), SETTABLE_SAMPLING_FIELDS);
let openai = supported_sampling_fields(Some(CloudProvider::OpenAi));
assert!(openai.contains(&"max_tokens"));
assert!(openai.contains(&"thinking"));
assert!(openai.contains(&"reasoning_effort"));
assert!(openai.contains(&"verbosity"));
assert!(!openai.contains(&"seed"));
assert!(!openai.contains(&"temperature"));
assert!(!openai.contains(&"top_p"));
assert!(!openai.contains(&"top_k"));
let gemini = supported_sampling_fields(Some(CloudProvider::Gemini));
assert!(gemini.contains(&"temperature"));
assert!(gemini.contains(&"top_p"));
assert!(gemini.contains(&"top_k"));
assert!(gemini.contains(&"seed"));
assert!(gemini.contains(&"thinking"));
assert!(gemini.contains(&"reasoning_effort"));
assert!(!gemini.contains(&"verbosity"));
let claude = supported_sampling_fields(Some(CloudProvider::Claude));
assert!(claude.contains(&"max_tokens"));
assert!(claude.contains(&"thinking"));
assert!(claude.contains(&"reasoning_effort"));
assert!(!claude.contains(&"top_k"));
let grok = supported_sampling_fields(Some(CloudProvider::Grok));
assert!(grok.contains(&"temperature"));
assert!(grok.contains(&"top_p"));
assert!(grok.contains(&"seed"));
assert!(grok.contains(&"max_tokens"));
assert!(grok.contains(&"reasoning_effort"));
assert!(!grok.contains(&"presence_penalty"));
assert!(!grok.contains(&"frequency_penalty"));
assert!(!grok.contains(&"top_k"));
assert!(!grok.contains(&"min_p"));
assert!(!grok.contains(&"verbosity"));
for f in openai.iter().chain(gemini).chain(claude).chain(grok) {
assert!(SETTABLE_SAMPLING_FIELDS.contains(f));
}
}
#[test]
fn retain_supported_drops_fields_by_mode() {
let s = SamplingConfig {
temperature: Some(0.7),
top_k: Some(40),
min_p: Some(0.05),
max_tokens: Some(256),
thinking: Some(true),
..Default::default()
};
let local = s.retain_supported(None, None);
assert_eq!(local.temperature, Some(0.7));
assert_eq!(local.top_k, Some(40));
assert_eq!(local.min_p, Some(0.05));
assert_eq!(local.thinking, Some(true));
let openai = s.retain_supported(Some(CloudProvider::OpenAi), None);
assert_eq!(openai.max_tokens, Some(256));
assert_eq!(openai.thinking, Some(true));
assert_eq!(openai.temperature, None);
assert_eq!(openai.top_k, None);
assert_eq!(openai.min_p, None);
let gemini = s.retain_supported(Some(CloudProvider::Gemini), None);
assert_eq!(gemini.temperature, Some(0.7));
assert_eq!(gemini.max_tokens, Some(256));
assert_eq!(gemini.top_k, Some(40));
assert_eq!(gemini.thinking, Some(true));
assert_eq!(gemini.min_p, None);
let claude = s.retain_supported(Some(CloudProvider::Claude), None);
assert_eq!(claude.max_tokens, Some(256));
assert_eq!(claude.thinking, Some(true));
assert_eq!(claude.temperature, None);
assert_eq!(claude.top_k, None);
let grok = s.retain_supported(Some(CloudProvider::Grok), None);
assert_eq!(grok.temperature, Some(0.7));
assert_eq!(grok.max_tokens, Some(256));
assert_eq!(grok.thinking, Some(true));
assert_eq!(grok.top_k, None);
assert_eq!(grok.min_p, None);
}
#[test]
fn retain_supported_drops_penalties_for_grok() {
let s = SamplingConfig {
presence_penalty: Some(0.5),
frequency_penalty: Some(0.5),
temperature: Some(0.7),
..Default::default()
};
let grok = s.retain_supported(Some(CloudProvider::Grok), None);
assert_eq!(grok.presence_penalty, None);
assert_eq!(grok.frequency_penalty, None);
assert_eq!(grok.temperature, Some(0.7));
}
#[test]
fn resolve_follows_priority() {
let global = SamplingConfig {
temperature: Some(0.1),
..Default::default()
};
let profile = SamplingConfig {
temperature: Some(0.5),
..Default::default()
};
let chat = SamplingConfig {
temperature: Some(0.9),
..Default::default()
};
assert_eq!(resolve(Some(&chat), Some(&profile), &global), chat);
assert_eq!(resolve(None, Some(&profile), &global), profile);
assert_eq!(resolve(None, None, &global), global);
assert_eq!(resolve(Some(&chat), None, &global), chat);
}
}