use crate::frontend::EmbeddedOpenAiRequestDefaults;
use crate::frontend::EmbeddedReasoningBudget;
use crate::frontend::EmbeddedReasoningEnabled;
use crate::frontend::EmbeddedReasoningFormat;
use base64::Engine;
use openai_frontend::ChatCompletionRequest;
use openai_frontend::ChatMessage;
use openai_frontend::CompletionRequest;
use openai_frontend::MessageContent;
use openai_frontend::MessageContentPart;
use openai_frontend::OpenAiError;
use openai_frontend::OpenAiResult;
use serde_json::Value;
use skippy_package_format::{
GenerationProfile, GenerationReasoningBudget, GenerationReasoningBudgetLevel,
GenerationReasoningEnabled, GenerationReasoningFormat,
};
use skippy_protocol::binary::MAX_STAGE_DRY_SEQUENCE_BREAKERS;
use skippy_protocol::binary::MAX_STAGE_LOGIT_BIAS;
use skippy_protocol::binary::MAX_STAGE_SAMPLERS;
use skippy_protocol::binary::StageLogitBias as WireLogitBias;
use skippy_protocol::binary::StageSamplingConfig as WireSamplingConfig;
use skippy_protocol::binary::sampling_flags;
use skippy_runtime::ChatReasoningFormat;
use skippy_runtime::ChatTemplateOptions;
use skippy_runtime::DEFAULT_PENALTY_LAST_N;
use skippy_runtime::DrySamplingConfig;
use skippy_runtime::LogitBias as RuntimeLogitBias;
use skippy_runtime::MAX_DRY_SEQUENCE_BREAKER_BYTES;
use skippy_runtime::MAX_LOGIT_BIAS;
use skippy_runtime::MediaInput;
use skippy_runtime::ReasoningBudget;
use skippy_runtime::SamplingConfig;
use skippy_runtime::XtcSamplingConfig;
use skippy_runtime::penalty_window;
use std::collections::BTreeMap;
const MAX_NATIVE_PARSER_INPUT_BYTES: usize = 1024 * 1024;
#[derive(Clone, Debug, PartialEq, Eq)]
pub(super) struct RequestDefaultsDiagnostics {
pub(super) selected_package_profile: Option<String>,
pub(super) max_tokens_source: &'static str,
pub(super) reasoning_budget_source: &'static str,
pub(super) field_sources: BTreeMap<&'static str, &'static str>,
}
pub(super) fn resolve_chat_request_defaults(
request: &ChatCompletionRequest,
configured: &EmbeddedOpenAiRequestDefaults,
) -> OpenAiResult<(EmbeddedOpenAiRequestDefaults, RequestDefaultsDiagnostics)> {
let template_reasoning = openai_frontend::normalize_reasoning_template_options(
request.reasoning.as_ref(),
request.reasoning_effort,
&request.extra,
)?
.enable_thinking;
let explicit_budget = request_reasoning_budget(request)?;
let resolved_reasoning = explicit_budget
.map(reasoning_budget_enables_thinking)
.or(template_reasoning)
.or_else(|| operator_reasoning_mode(configured));
let selected = configured
.package_request_defaults
.as_ref()
.and_then(|package| {
let profile_name = match resolved_reasoning {
Some(true) => package.selection.reasoning_enabled.as_ref(),
Some(false) => package.selection.reasoning_disabled.as_ref(),
None => Some(&package.selection.default),
}?;
package
.profiles
.get(profile_name)
.map(|profile| (profile_name.clone(), profile))
});
let mut resolved = configured.clone();
let selected_profile = selected.as_ref().map(|(name, _)| name.clone());
let field_sources = chat_field_sources(
request,
configured,
selected.as_ref().map(|(_, profile)| *profile),
explicit_budget,
template_reasoning,
);
if let Some((_, profile)) = selected {
apply_package_profile(&mut resolved, profile);
}
let max_tokens_source = if request.effective_max_tokens().is_some() {
"request"
} else if configured.max_tokens.is_some() {
"deployment"
} else if resolved.max_tokens.is_some() {
"package"
} else {
"fallback"
};
let reasoning_budget_source = if explicit_budget.is_some() {
"request"
} else if configured.reasoning_budget.is_some() {
"deployment"
} else if resolved.reasoning_budget.is_some() {
"package"
} else {
"fallback"
};
Ok((
resolved,
RequestDefaultsDiagnostics {
selected_package_profile: selected_profile,
max_tokens_source,
reasoning_budget_source,
field_sources,
},
))
}
pub(super) fn resolve_completion_request_defaults(
request: &CompletionRequest,
configured: &EmbeddedOpenAiRequestDefaults,
) -> (EmbeddedOpenAiRequestDefaults, RequestDefaultsDiagnostics) {
let selected = configured
.package_request_defaults
.as_ref()
.and_then(|package| {
let profile_name = package
.selection
.reasoning_disabled
.as_ref()
.unwrap_or(&package.selection.default);
package
.profiles
.get(profile_name)
.map(|profile| (profile_name.clone(), profile))
});
let mut resolved = configured.clone();
let selected_package_profile = selected.as_ref().map(|(name, _)| name.clone());
let field_sources = completion_field_sources(
request,
configured,
selected.as_ref().map(|(_, profile)| *profile),
);
if let Some((_, profile)) = selected {
apply_package_profile(&mut resolved, profile);
}
let max_tokens_source = if request.max_tokens.is_some() {
"request"
} else if configured.max_tokens.is_some() {
"deployment"
} else if resolved.max_tokens.is_some() {
"package"
} else {
"fallback"
};
(
resolved,
RequestDefaultsDiagnostics {
selected_package_profile,
max_tokens_source,
reasoning_budget_source: "not_applicable",
field_sources,
},
)
}
fn resolved_field_source(request: bool, deployment: bool, package: bool) -> &'static str {
if request {
"request"
} else if deployment {
"deployment"
} else if package {
"package"
} else {
"fallback"
}
}
fn request_extra_is_set(extra: &BTreeMap<String, Value>, field: &str) -> bool {
!extra_value_is_omitted(extra, field)
}
fn chat_field_sources(
request: &ChatCompletionRequest,
configured: &EmbeddedOpenAiRequestDefaults,
package: Option<&GenerationProfile>,
explicit_budget: Option<ReasoningBudget>,
template_reasoning: Option<bool>,
) -> BTreeMap<&'static str, &'static str> {
let mut sources = BTreeMap::new();
macro_rules! field {
($name:literal, $request:expr, $configured:ident, $package:ident) => {
sources.insert(
$name,
resolved_field_source(
$request,
configured.$configured.is_some(),
package.is_some_and(|profile| profile.$package.is_some()),
),
);
};
}
field!(
"max_tokens",
request.effective_max_tokens().is_some(),
max_tokens,
max_tokens
);
field!("stop", request.stop.is_some(), stop, stop);
field!(
"temperature",
request.temperature.is_some(),
temperature,
temperature
);
field!("top_p", request.top_p.is_some(), top_p, top_p);
field!(
"presence_penalty",
request.presence_penalty.is_some(),
presence_penalty,
presence_penalty
);
field!(
"frequency_penalty",
request.frequency_penalty.is_some(),
frequency_penalty,
frequency_penalty
);
field!("seed", request.seed.is_some(), seed, seed);
field!(
"logit_bias",
request.logit_bias.is_some(),
logit_bias,
logit_bias
);
for (name, requested, deployed, packaged) in [
(
"top_k",
request_extra_is_set(&request.extra, "top_k"),
configured.top_k.is_some(),
package.is_some_and(|profile| profile.top_k.is_some()),
),
(
"min_p",
request_extra_is_set(&request.extra, "min_p"),
configured.min_p.is_some(),
package.is_some_and(|profile| profile.min_p.is_some()),
),
(
"typical_p",
request_extra_is_set(&request.extra, "typical_p"),
configured.typical_p.is_some(),
package.is_some_and(|profile| profile.typical_p.is_some()),
),
(
"top_nsigma",
request_extra_is_set(&request.extra, "top_nsigma"),
configured.top_nsigma.is_some(),
package.is_some_and(|profile| profile.top_nsigma.is_some()),
),
(
"repeat_penalty",
request_extra_is_set(&request.extra, "repeat_penalty")
|| request_extra_is_set(&request.extra, "repetition_penalty"),
configured.repeat_penalty.is_some(),
package.is_some_and(|profile| profile.repeat_penalty.is_some()),
),
(
"repeat_last_n",
request_extra_is_set(&request.extra, "repeat_last_n"),
configured.repeat_last_n.is_some(),
package.is_some_and(|profile| profile.repeat_last_n.is_some()),
),
(
"dynatemp_range",
request_extra_is_set(&request.extra, "dynatemp_range"),
configured.dynatemp_range.is_some(),
package.is_some_and(|profile| profile.dynatemp_range.is_some()),
),
(
"dynatemp_exponent",
request_extra_is_set(&request.extra, "dynatemp_exponent"),
configured.dynatemp_exponent.is_some(),
package.is_some_and(|profile| profile.dynatemp_exponent.is_some()),
),
(
"dry",
request_extra_is_set(&request.extra, "dry"),
configured.dry.is_some(),
package.is_some_and(|profile| profile.dry.is_some()),
),
(
"xtc",
request_extra_is_set(&request.extra, "xtc"),
configured.xtc.is_some(),
package.is_some_and(|profile| profile.xtc.is_some()),
),
(
"mirostat_mode",
request_extra_is_set(&request.extra, "mirostat_mode"),
configured.mirostat_mode.is_some(),
package.is_some_and(|profile| profile.mirostat_mode.is_some()),
),
(
"mirostat_entropy",
request_extra_is_set(&request.extra, "mirostat_entropy"),
configured.mirostat_entropy.is_some(),
package.is_some_and(|profile| profile.mirostat_entropy.is_some()),
),
(
"mirostat_learning_rate",
request_extra_is_set(&request.extra, "mirostat_learning_rate"),
configured.mirostat_learning_rate.is_some(),
package.is_some_and(|profile| profile.mirostat_learning_rate.is_some()),
),
(
"samplers",
request_extra_is_set(&request.extra, "samplers"),
configured.samplers.is_some(),
package.is_some_and(|profile| profile.samplers.is_some()),
),
(
"sampler_sequence",
request_extra_is_set(&request.extra, "sampler_sequence"),
configured.sampler_sequence.is_some(),
package.is_some_and(|profile| profile.sampler_sequence.is_some()),
),
(
"ignore_eos",
request_extra_is_set(&request.extra, "ignore_eos"),
configured.ignore_eos.is_some(),
package.is_some_and(|profile| profile.ignore_eos.is_some()),
),
] {
sources.insert(name, resolved_field_source(requested, deployed, packaged));
}
let package_reasoning = package.and_then(|profile| profile.reasoning.as_ref());
sources.insert(
"reasoning_enabled",
resolved_field_source(
template_reasoning.is_some() || explicit_budget.is_some(),
configured.reasoning_enabled.is_some(),
package_reasoning.is_some_and(|reasoning| reasoning.enabled.is_some()),
),
);
sources.insert(
"reasoning_format",
resolved_field_source(
request_extra_is_set(&request.extra, "reasoning_format"),
configured.reasoning_format.is_some(),
package_reasoning.is_some_and(|reasoning| reasoning.format.is_some()),
),
);
sources.insert(
"reasoning_budget",
resolved_field_source(
explicit_budget.is_some(),
configured.reasoning_budget.is_some(),
package_reasoning.is_some_and(|reasoning| reasoning.budget.is_some()),
),
);
sources
}
fn completion_field_sources(
request: &CompletionRequest,
configured: &EmbeddedOpenAiRequestDefaults,
package: Option<&GenerationProfile>,
) -> BTreeMap<&'static str, &'static str> {
let mut sources = BTreeMap::new();
macro_rules! field {
($name:literal, $request:expr, $configured:ident, $package:ident) => {
sources.insert(
$name,
resolved_field_source(
$request,
configured.$configured.is_some(),
package.is_some_and(|profile| profile.$package.is_some()),
),
);
};
}
field!(
"max_tokens",
request.max_tokens.is_some(),
max_tokens,
max_tokens
);
field!("stop", request.stop.is_some(), stop, stop);
field!(
"temperature",
request.temperature.is_some(),
temperature,
temperature
);
field!("top_p", request.top_p.is_some(), top_p, top_p);
field!(
"presence_penalty",
request.presence_penalty.is_some(),
presence_penalty,
presence_penalty
);
field!(
"frequency_penalty",
request.frequency_penalty.is_some(),
frequency_penalty,
frequency_penalty
);
field!("seed", request.seed.is_some(), seed, seed);
field!(
"logit_bias",
request.logit_bias.is_some(),
logit_bias,
logit_bias
);
for (name, requested, deployed, packaged) in [
(
"top_k",
request_extra_is_set(&request.extra, "top_k"),
configured.top_k.is_some(),
package.is_some_and(|profile| profile.top_k.is_some()),
),
(
"min_p",
request_extra_is_set(&request.extra, "min_p"),
configured.min_p.is_some(),
package.is_some_and(|profile| profile.min_p.is_some()),
),
(
"typical_p",
request_extra_is_set(&request.extra, "typical_p"),
configured.typical_p.is_some(),
package.is_some_and(|profile| profile.typical_p.is_some()),
),
(
"top_nsigma",
request_extra_is_set(&request.extra, "top_nsigma"),
configured.top_nsigma.is_some(),
package.is_some_and(|profile| profile.top_nsigma.is_some()),
),
(
"repeat_penalty",
request_extra_is_set(&request.extra, "repeat_penalty")
|| request_extra_is_set(&request.extra, "repetition_penalty"),
configured.repeat_penalty.is_some(),
package.is_some_and(|profile| profile.repeat_penalty.is_some()),
),
(
"repeat_last_n",
request_extra_is_set(&request.extra, "repeat_last_n"),
configured.repeat_last_n.is_some(),
package.is_some_and(|profile| profile.repeat_last_n.is_some()),
),
(
"dynatemp_range",
request_extra_is_set(&request.extra, "dynatemp_range"),
configured.dynatemp_range.is_some(),
package.is_some_and(|profile| profile.dynatemp_range.is_some()),
),
(
"dynatemp_exponent",
request_extra_is_set(&request.extra, "dynatemp_exponent"),
configured.dynatemp_exponent.is_some(),
package.is_some_and(|profile| profile.dynatemp_exponent.is_some()),
),
(
"dry",
request_extra_is_set(&request.extra, "dry"),
configured.dry.is_some(),
package.is_some_and(|profile| profile.dry.is_some()),
),
(
"xtc",
request_extra_is_set(&request.extra, "xtc"),
configured.xtc.is_some(),
package.is_some_and(|profile| profile.xtc.is_some()),
),
(
"mirostat_mode",
request_extra_is_set(&request.extra, "mirostat_mode"),
configured.mirostat_mode.is_some(),
package.is_some_and(|profile| profile.mirostat_mode.is_some()),
),
(
"mirostat_entropy",
request_extra_is_set(&request.extra, "mirostat_entropy"),
configured.mirostat_entropy.is_some(),
package.is_some_and(|profile| profile.mirostat_entropy.is_some()),
),
(
"mirostat_learning_rate",
request_extra_is_set(&request.extra, "mirostat_learning_rate"),
configured.mirostat_learning_rate.is_some(),
package.is_some_and(|profile| profile.mirostat_learning_rate.is_some()),
),
(
"samplers",
request_extra_is_set(&request.extra, "samplers"),
configured.samplers.is_some(),
package.is_some_and(|profile| profile.samplers.is_some()),
),
(
"sampler_sequence",
request_extra_is_set(&request.extra, "sampler_sequence"),
configured.sampler_sequence.is_some(),
package.is_some_and(|profile| profile.sampler_sequence.is_some()),
),
(
"ignore_eos",
request_extra_is_set(&request.extra, "ignore_eos"),
configured.ignore_eos.is_some(),
package.is_some_and(|profile| profile.ignore_eos.is_some()),
),
] {
sources.insert(name, resolved_field_source(requested, deployed, packaged));
}
sources
}
fn reasoning_budget_enables_thinking(budget: ReasoningBudget) -> bool {
!matches!(
budget,
ReasoningBudget::Explicit(0) | ReasoningBudget::Resolved(0)
)
}
fn operator_reasoning_mode(defaults: &EmbeddedOpenAiRequestDefaults) -> Option<bool> {
match defaults.reasoning_enabled {
Some(EmbeddedReasoningEnabled::Enabled) => return Some(true),
Some(EmbeddedReasoningEnabled::Disabled) => return Some(false),
Some(EmbeddedReasoningEnabled::Auto) | None => {}
}
defaults.reasoning_budget.and_then(|budget| match budget {
EmbeddedReasoningBudget::Tokens(0) => Some(false),
EmbeddedReasoningBudget::Auto => None,
EmbeddedReasoningBudget::Unrestricted | EmbeddedReasoningBudget::Tokens(_) => Some(true),
EmbeddedReasoningBudget::Effort(openai_frontend::ReasoningEffort::None) => Some(false),
EmbeddedReasoningBudget::Effort(_) => Some(true),
})
}
fn apply_package_profile(
resolved: &mut EmbeddedOpenAiRequestDefaults,
profile: &GenerationProfile,
) {
macro_rules! fill {
($field:ident) => {
if resolved.$field.is_none() {
resolved.$field = profile.$field.map(|value| value as _);
}
};
}
fill!(max_tokens);
if resolved.stop.is_none() {
resolved.stop.clone_from(&profile.stop);
}
fill!(temperature);
fill!(top_p);
fill!(top_k);
fill!(min_p);
fill!(typical_p);
fill!(top_nsigma);
fill!(presence_penalty);
fill!(frequency_penalty);
fill!(seed);
if resolved.logit_bias.is_none() {
resolved.logit_bias = profile.logit_bias.as_ref().map(|biases| {
biases
.iter()
.map(|(token, bias)| (token.clone(), Value::from(*bias)))
.collect()
});
}
fill!(repeat_penalty);
fill!(repeat_last_n);
fill!(dynatemp_range);
fill!(dynatemp_exponent);
fill!(mirostat_mode);
fill!(mirostat_entropy);
fill!(mirostat_learning_rate);
if resolved.samplers.is_none() {
resolved.samplers.clone_from(&profile.samplers);
}
if resolved.sampler_sequence.is_none() {
resolved
.sampler_sequence
.clone_from(&profile.sampler_sequence);
}
if resolved.ignore_eos.is_none() {
resolved.ignore_eos = profile.ignore_eos;
}
if resolved.dry.is_none() {
resolved.dry = profile.dry.as_ref().map(|dry| DrySamplingConfig {
multiplier: dry.multiplier.unwrap_or(0.0) as f32,
base: dry.base.unwrap_or(1.75) as f32,
allowed_length: dry.allowed_length.unwrap_or(2),
penalty_last_n: dry.penalty_last_n.unwrap_or(DEFAULT_PENALTY_LAST_N),
sequence_breakers: dry
.sequence_breakers
.clone()
.unwrap_or_else(|| vec!["\n".into(), ":".into(), "\"".into(), "*".into()]),
});
}
if resolved.xtc.is_none() {
resolved.xtc = profile.xtc.as_ref().map(|xtc| XtcSamplingConfig {
probability: xtc.probability.unwrap_or(0.0) as f32,
threshold: xtc.threshold.unwrap_or(0.1) as f32,
});
}
if let Some(reasoning) = &profile.reasoning {
if resolved.reasoning_enabled.is_none() {
resolved.reasoning_enabled = reasoning.enabled.map(|enabled| match enabled {
GenerationReasoningEnabled::Auto => EmbeddedReasoningEnabled::Auto,
GenerationReasoningEnabled::Off => EmbeddedReasoningEnabled::Disabled,
GenerationReasoningEnabled::On => EmbeddedReasoningEnabled::Enabled,
});
}
if resolved.reasoning_format.is_none() {
resolved.reasoning_format = reasoning.format.map(|format| match format {
GenerationReasoningFormat::Auto => EmbeddedReasoningFormat::Auto,
GenerationReasoningFormat::None => EmbeddedReasoningFormat::None,
GenerationReasoningFormat::Deepseek => EmbeddedReasoningFormat::Deepseek,
GenerationReasoningFormat::DeepseekLegacy => {
EmbeddedReasoningFormat::DeepseekLegacy
}
GenerationReasoningFormat::Hidden => EmbeddedReasoningFormat::Hidden,
});
}
if resolved.reasoning_budget.is_none() {
resolved.reasoning_budget = reasoning.budget.as_ref().map(|budget| match budget {
GenerationReasoningBudget::Tokens(tokens) => {
EmbeddedReasoningBudget::Tokens(*tokens)
}
GenerationReasoningBudget::Level(level) => match level {
GenerationReasoningBudgetLevel::Auto => EmbeddedReasoningBudget::Auto,
GenerationReasoningBudgetLevel::Low => {
EmbeddedReasoningBudget::Effort(openai_frontend::ReasoningEffort::Low)
}
GenerationReasoningBudgetLevel::Medium => {
EmbeddedReasoningBudget::Effort(openai_frontend::ReasoningEffort::Medium)
}
GenerationReasoningBudgetLevel::High => {
EmbeddedReasoningBudget::Effort(openai_frontend::ReasoningEffort::High)
}
GenerationReasoningBudgetLevel::Unrestricted => {
EmbeddedReasoningBudget::Unrestricted
}
},
});
}
}
}
struct SharedRequestFields<'a> {
presence_penalty: &'a mut Option<f32>,
frequency_penalty: &'a mut Option<f32>,
seed: &'a mut Option<u64>,
logit_bias: &'a mut Option<std::collections::BTreeMap<String, serde_json::Value>>,
temperature: &'a mut Option<f32>,
top_p: &'a mut Option<f32>,
stop: &'a mut Option<openai_frontend::StopSequence>,
extra: &'a mut std::collections::BTreeMap<String, serde_json::Value>,
}
pub(super) fn apply_chat_request_defaults(
request: &mut ChatCompletionRequest,
defaults: &EmbeddedOpenAiRequestDefaults,
) -> OpenAiResult<()> {
if request.max_tokens.is_none() && request.max_completion_tokens.is_none() {
request.max_tokens = defaults.max_tokens;
}
apply_shared_request_defaults(
SharedRequestFields {
presence_penalty: &mut request.presence_penalty,
frequency_penalty: &mut request.frequency_penalty,
seed: &mut request.seed,
logit_bias: &mut request.logit_bias,
temperature: &mut request.temperature,
top_p: &mut request.top_p,
stop: &mut request.stop,
extra: &mut request.extra,
},
defaults,
);
apply_chat_only_request_defaults(request, defaults)
}
fn apply_chat_only_request_defaults(
request: &mut ChatCompletionRequest,
defaults: &EmbeddedOpenAiRequestDefaults,
) -> OpenAiResult<()> {
for (name, value) in [
(
"chat_template",
defaults.chat_template.clone().map(Value::from),
),
("jinja", defaults.jinja.map(Value::from)),
(
"chat_template_kwargs",
defaults.chat_template_kwargs.clone(),
),
(
"skip_chat_parsing",
defaults.skip_chat_parsing.map(Value::from),
),
("prefill_assistant", defaults.prefill_assistant.clone()),
(
"system_prompt",
defaults.system_prompt.clone().map(Value::from),
),
(
"reasoning_format",
defaults.reasoning_format.map(|value| {
Value::from(match value {
EmbeddedReasoningFormat::Auto => "auto",
EmbeddedReasoningFormat::None => "none",
EmbeddedReasoningFormat::Deepseek => "deepseek",
EmbeddedReasoningFormat::DeepseekLegacy => "deepseek-legacy",
EmbeddedReasoningFormat::Hidden => "hidden",
})
}),
),
] {
if extra_value_is_omitted(&request.extra, name)
&& let Some(value) = value
{
request.extra.insert(name.to_string(), value);
}
}
if extra_value_is_omitted(&request.extra, "grammar")
&& extra_value_is_omitted(&request.extra, "json_schema")
&& let Some((name, value)) = [
("grammar", defaults.grammar.clone()),
("json_schema", defaults.json_schema.clone()),
]
.into_iter()
.find_map(|(name, value)| value.map(|value| (name, value)))
{
request.extra.insert(name.to_string(), value);
}
if let Some(system_prompt) = optional_string_extra(&request.extra, "system_prompt")?
&& !request
.messages
.iter()
.any(|message| message.role == "system")
{
request.messages.insert(
0,
ChatMessage {
role: "system".to_string(),
content: Some(openai_frontend::MessageContent::Text(system_prompt)),
extra: std::collections::BTreeMap::new(),
},
);
}
if let Some(prefill) = request
.extra
.get("prefill_assistant")
.filter(|value| !value.is_null())
{
request.messages.push(prefill_assistant_message(prefill)?);
}
Ok(())
}
fn prefill_assistant_message(value: &Value) -> OpenAiResult<ChatMessage> {
if let Some(content) = value.as_str() {
return Ok(ChatMessage {
role: "assistant".to_string(),
content: Some(openai_frontend::MessageContent::Text(content.to_string())),
extra: std::collections::BTreeMap::new(),
});
}
let message = serde_json::from_value::<ChatMessage>(value.clone()).map_err(|_| {
OpenAiError::invalid_request("prefill_assistant must be a string or chat message object")
})?;
if message.role != "assistant" {
return Err(OpenAiError::invalid_request(
"prefill_assistant message role must be assistant",
));
}
Ok(message)
}
pub(super) fn apply_completion_request_defaults(
request: &mut CompletionRequest,
defaults: &EmbeddedOpenAiRequestDefaults,
) {
if request.max_tokens.is_none() {
request.max_tokens = defaults.max_tokens;
}
apply_shared_request_defaults(
SharedRequestFields {
presence_penalty: &mut request.presence_penalty,
frequency_penalty: &mut request.frequency_penalty,
seed: &mut request.seed,
logit_bias: &mut request.logit_bias,
temperature: &mut request.temperature,
top_p: &mut request.top_p,
stop: &mut request.stop,
extra: &mut request.extra,
},
defaults,
);
}
pub(super) fn message_content_to_generation_text(
content: &MessageContent,
marker: &str,
media: &mut Vec<MediaInput>,
) -> OpenAiResult<String> {
match content {
MessageContent::Text(text) => Ok(text.clone()),
MessageContent::Parts(parts) => {
let mut chunks = Vec::new();
for part in parts {
if part.content_type == "text" {
if let Some(text) = part.text.as_deref() {
chunks.push(text.to_string());
}
continue;
}
if let Some(bytes) = media_bytes_from_part(part)? {
media.push(MediaInput { bytes });
chunks.push(marker.to_string());
}
}
Ok(chunks.join("\n"))
}
MessageContent::Other(_) => Ok(String::new()),
}
}
pub(super) fn media_bytes_from_part(part: &MessageContentPart) -> OpenAiResult<Option<Vec<u8>>> {
let is_media = matches!(
part.content_type.as_str(),
"image_url" | "input_image" | "image" | "input_audio" | "audio" | "audio_url"
);
if !is_media {
return Ok(None);
}
if let Some(url) = media_url(part) {
return decode_media_url(&url).map(Some);
}
if let Some(data) = media_data(part) {
return decode_base64_payload(&data).map(Some);
}
Err(OpenAiError::invalid_request(format!(
"media content block '{}' is missing url or data",
part.content_type
)))
}
pub(super) fn media_url(part: &MessageContentPart) -> Option<String> {
part.media_url()
}
pub(super) fn media_data(part: &MessageContentPart) -> Option<String> {
part.media_data()
}
pub(super) fn decode_media_url(url: &str) -> OpenAiResult<Vec<u8>> {
match url.split_once(',') {
Some((prefix, payload)) if prefix.starts_with("data:") && prefix.contains(";base64") => {
return decode_base64_payload(payload);
}
_ => {}
}
if url.starts_with("http://") || url.starts_with("https://") {
return Err(OpenAiError::unsupported(
"remote multimodal URLs must be fetched by mesh before reaching skippy",
));
}
decode_base64_payload(url)
}
pub(super) fn decode_base64_payload(payload: &str) -> OpenAiResult<Vec<u8>> {
base64::engine::general_purpose::STANDARD
.decode(payload.as_bytes())
.or_else(|_| base64::engine::general_purpose::URL_SAFE_NO_PAD.decode(payload.as_bytes()))
.map_err(|error| OpenAiError::invalid_request(format!("invalid media base64: {error}")))
}
fn apply_shared_request_defaults(
fields: SharedRequestFields<'_>,
defaults: &EmbeddedOpenAiRequestDefaults,
) {
let SharedRequestFields {
presence_penalty,
frequency_penalty,
seed,
logit_bias,
temperature,
top_p,
stop,
extra,
} = fields;
if presence_penalty.is_none() {
*presence_penalty = defaults.presence_penalty;
}
if frequency_penalty.is_none() {
*frequency_penalty = defaults.frequency_penalty;
}
if seed.is_none() {
*seed = defaults.seed;
}
if logit_bias.is_none() {
*logit_bias = defaults.logit_bias.clone();
}
if temperature.is_none() {
*temperature = defaults.temperature;
}
if top_p.is_none() {
*top_p = defaults.top_p;
}
if stop.is_none() {
*stop = defaults
.stop
.as_ref()
.map(|values| openai_frontend::StopSequence::from_values(values.clone()));
}
if let (true, Some(value)) = (extra_value_is_omitted(extra, "top_k"), defaults.top_k) {
extra.insert("top_k".to_string(), serde_json::json!(value));
}
if let (true, Some(value)) = (extra_value_is_omitted(extra, "min_p"), defaults.min_p) {
extra.insert("min_p".to_string(), serde_json::json!(value));
}
if let (true, Some(value)) = (
extra_value_is_omitted(extra, "repeat_penalty")
&& extra_value_is_omitted(extra, "repetition_penalty"),
defaults.repeat_penalty,
) {
extra.insert("repeat_penalty".to_string(), serde_json::json!(value));
}
if let (true, Some(value)) = (
extra_value_is_omitted(extra, "repeat_last_n"),
defaults.repeat_last_n,
) {
extra.insert("repeat_last_n".to_string(), serde_json::json!(value));
}
for (name, value) in [
("typical_p", defaults.typical_p.map(Value::from)),
("top_nsigma", defaults.top_nsigma.map(Value::from)),
("dynatemp_range", defaults.dynatemp_range.map(Value::from)),
(
"dynatemp_exponent",
defaults.dynatemp_exponent.map(Value::from),
),
("mirostat_mode", defaults.mirostat_mode.map(Value::from)),
(
"mirostat_entropy",
defaults.mirostat_entropy.map(Value::from),
),
(
"mirostat_learning_rate",
defaults.mirostat_learning_rate.map(Value::from),
),
("ignore_eos", defaults.ignore_eos.map(Value::from)),
] {
if extra_value_is_omitted(extra, name)
&& let Some(value) = value
{
extra.insert(name.to_string(), value);
}
}
if extra_value_is_omitted(extra, "dry")
&& let Some(dry) = defaults.dry.as_ref()
{
extra.insert(
"dry".to_string(),
serde_json::json!({
"multiplier": dry.multiplier,
"base": dry.base,
"allowed_length": dry.allowed_length,
"penalty_last_n": dry.penalty_last_n,
"sequence_breakers": dry.sequence_breakers,
}),
);
}
if extra_value_is_omitted(extra, "xtc")
&& let Some(xtc) = defaults.xtc.as_ref()
{
extra.insert(
"xtc".to_string(),
serde_json::json!({"probability": xtc.probability, "threshold": xtc.threshold}),
);
}
if extra_value_is_omitted(extra, "samplers")
&& let Some(samplers) = defaults.samplers.as_ref()
{
extra.insert("samplers".to_string(), serde_json::json!(samplers));
}
if extra_value_is_omitted(extra, "sampler_sequence")
&& let Some(sequence) = defaults.sampler_sequence.as_ref()
{
extra.insert(
"sampler_sequence".to_string(),
Value::from(sequence.clone()),
);
}
}
fn extra_value_is_omitted(
extra: &std::collections::BTreeMap<String, serde_json::Value>,
field: &str,
) -> bool {
extra.get(field).is_none_or(Value::is_null)
}
pub(super) fn chat_sampling_config(
request: &ChatCompletionRequest,
defaults: &EmbeddedOpenAiRequestDefaults,
) -> OpenAiResult<SamplingConfig> {
let mut sampling = sampling_config(
request.temperature,
request.top_p,
request.presence_penalty,
request.frequency_penalty,
request.seed,
request.logit_bias.as_ref(),
&request.extra,
)?;
sampling.reasoning_budget =
request_reasoning_budget(request)?.unwrap_or_else(|| embedded_reasoning_budget(defaults));
Ok(sampling)
}
pub(super) fn completion_sampling_config(
request: &CompletionRequest,
) -> OpenAiResult<SamplingConfig> {
sampling_config(
request.temperature,
request.top_p,
request.presence_penalty,
request.frequency_penalty,
request.seed,
request.logit_bias.as_ref(),
&request.extra,
)
}
pub(super) fn chat_template_options(
request: &ChatCompletionRequest,
defaults: &EmbeddedOpenAiRequestDefaults,
) -> OpenAiResult<ChatTemplateOptions> {
let reasoning = openai_frontend::normalize_reasoning_template_options(
request.reasoning.as_ref(),
request.reasoning_effort,
&request.extra,
)?;
Ok(ChatTemplateOptions {
add_assistant: request
.extra
.get("prefill_assistant")
.is_none_or(Value::is_null),
reasoning_format: Some(request_reasoning_format(request, defaults)?),
enable_thinking: reasoning
.enable_thinking
.or_else(|| default_reasoning_enabled(defaults.reasoning_enabled))
.or_else(|| default_reasoning_budget_enabled(defaults.reasoning_budget)),
chat_template_kwargs: merged_chat_template_kwargs(
defaults,
&reasoning.chat_template_kwargs,
)?
.map(|kwargs| serialize_bounded_native_parser_json("chat_template_kwargs", &kwargs))
.transpose()
.map_err(|error| OpenAiError::invalid_request(error.to_string()))?,
chat_template: bounded_optional_string_extra(&request.extra, "chat_template")?,
use_jinja: optional_bool_extra(&request.extra, "jinja")?.unwrap_or(true),
grammar: structured_output_string(request, "grammar")?,
json_schema: structured_output_json(request, "json_schema")?,
skip_chat_parsing: optional_bool_extra(&request.extra, "skip_chat_parsing")?
.unwrap_or(false),
})
}
fn merged_chat_template_kwargs(
defaults: &EmbeddedOpenAiRequestDefaults,
request: &std::collections::BTreeMap<String, Value>,
) -> OpenAiResult<Option<std::collections::BTreeMap<String, Value>>> {
let mut merged = defaults
.chat_template_kwargs
.as_ref()
.map(|value| {
value.as_object().cloned().ok_or_else(|| {
OpenAiError::invalid_request("chat_template_kwargs must be an object")
})
})
.transpose()?
.unwrap_or_default()
.into_iter()
.collect::<std::collections::BTreeMap<_, _>>();
if let Some(budget) = defaults.reasoning_budget {
match budget {
EmbeddedReasoningBudget::Tokens(tokens) => {
merged.insert("reasoning_budget".to_string(), Value::from(tokens));
merged.insert("thinking_budget".to_string(), Value::from(tokens));
}
EmbeddedReasoningBudget::Effort(effort) => {
merged.insert(
"reasoning_effort".to_string(),
Value::from(match effort {
openai_frontend::ReasoningEffort::None => "none",
openai_frontend::ReasoningEffort::Minimal => "minimal",
openai_frontend::ReasoningEffort::Low => "low",
openai_frontend::ReasoningEffort::Medium => "medium",
openai_frontend::ReasoningEffort::High => "high",
openai_frontend::ReasoningEffort::Xhigh => "xhigh",
openai_frontend::ReasoningEffort::Max => "max",
}),
);
}
EmbeddedReasoningBudget::Auto | EmbeddedReasoningBudget::Unrestricted => {}
}
}
merged.extend(request.clone());
Ok((!merged.is_empty()).then_some(merged))
}
fn request_reasoning_format(
request: &ChatCompletionRequest,
defaults: &EmbeddedOpenAiRequestDefaults,
) -> OpenAiResult<ChatReasoningFormat> {
let value = optional_string_extra(&request.extra, "reasoning_format")?;
match value.as_deref() {
Some("auto") => Ok(ChatReasoningFormat::Auto),
None => Ok(chat_reasoning_format(defaults.reasoning_format)),
Some("none") => Ok(ChatReasoningFormat::None),
Some("deepseek") => Ok(ChatReasoningFormat::Deepseek),
Some("deepseek-legacy") => Ok(ChatReasoningFormat::DeepseekLegacy),
Some("hidden") => Ok(ChatReasoningFormat::Hidden),
Some(_) => Err(OpenAiError::invalid_request(
"reasoning_format must be auto, none, deepseek, deepseek-legacy, or hidden",
)),
}
}
fn request_reasoning_budget(
request: &ChatCompletionRequest,
) -> OpenAiResult<Option<ReasoningBudget>> {
let normalized = openai_frontend::normalize_reasoning_template_options(
request.reasoning.as_ref(),
request.reasoning_effort,
&request.extra,
)?;
if normalized.enable_thinking == Some(false) {
return Ok(Some(ReasoningBudget::Explicit(0)));
}
for field in [
"reasoning_budget_tokens",
"thinking_budget_tokens",
"reasoning_budget",
"thinking_budget",
] {
if let Some(tokens) = optional_i32_extra(&request.extra, field)? {
return match tokens {
-1 => Ok(Some(ReasoningBudget::Unrestricted)),
0.. => Ok(Some(ReasoningBudget::Explicit(tokens as u32))),
_ => Err(OpenAiError::invalid_request(format!(
"{field} must be -1 or greater"
))),
};
}
}
if let Some(max_tokens) = request
.reasoning
.as_ref()
.and_then(|reasoning| reasoning.max_tokens)
{
return Ok(Some(ReasoningBudget::Explicit(max_tokens)));
}
let effort = request.reasoning_effort.or_else(|| {
request
.reasoning
.as_ref()
.and_then(|reasoning| reasoning.effort)
});
Ok(effort.map(reasoning_effort_budget))
}
fn reasoning_effort_budget(effort: openai_frontend::ReasoningEffort) -> ReasoningBudget {
use openai_frontend::ReasoningEffort;
match effort {
ReasoningEffort::None => ReasoningBudget::Explicit(0),
ReasoningEffort::Minimal | ReasoningEffort::Low => ReasoningBudget::Capped(1_024),
ReasoningEffort::Medium => ReasoningBudget::Capped(4_096),
ReasoningEffort::High | ReasoningEffort::Xhigh | ReasoningEffort::Max => {
ReasoningBudget::Capped(8_192)
}
}
}
fn embedded_reasoning_budget(defaults: &EmbeddedOpenAiRequestDefaults) -> ReasoningBudget {
match defaults.reasoning_budget {
Some(EmbeddedReasoningBudget::Unrestricted) => ReasoningBudget::Unrestricted,
Some(EmbeddedReasoningBudget::Tokens(tokens)) => ReasoningBudget::Explicit(tokens),
Some(EmbeddedReasoningBudget::Effort(effort)) => reasoning_effort_budget(effort),
Some(EmbeddedReasoningBudget::Auto) | None => match defaults.reasoning_enabled {
Some(EmbeddedReasoningEnabled::Disabled) => ReasoningBudget::Explicit(0),
_ => ReasoningBudget::Capped(4_096),
},
}
}
fn optional_string_extra(
extra: &std::collections::BTreeMap<String, Value>,
name: &str,
) -> OpenAiResult<Option<String>> {
extra
.get(name)
.filter(|value| !value.is_null())
.map(|value| {
value
.as_str()
.map(str::to_string)
.ok_or_else(|| OpenAiError::invalid_request(format!("{name} must be a string")))
})
.transpose()
}
fn bounded_optional_string_extra(
extra: &std::collections::BTreeMap<String, Value>,
name: &str,
) -> OpenAiResult<Option<String>> {
let value = optional_string_extra(extra, name)?;
if value
.as_ref()
.is_some_and(|value| value.len() > MAX_NATIVE_PARSER_INPUT_BYTES)
{
return Err(OpenAiError::invalid_request(format!(
"{name} exceeds the {MAX_NATIVE_PARSER_INPUT_BYTES}-byte limit"
)));
}
Ok(value)
}
fn optional_bool_extra(
extra: &std::collections::BTreeMap<String, Value>,
name: &str,
) -> OpenAiResult<Option<bool>> {
extra
.get(name)
.filter(|value| !value.is_null())
.map(|value| {
value
.as_bool()
.ok_or_else(|| OpenAiError::invalid_request(format!("{name} must be a boolean")))
})
.transpose()
}
fn structured_output_string(
request: &ChatCompletionRequest,
name: &str,
) -> OpenAiResult<Option<String>> {
let value = bounded_optional_string_extra(&request.extra, name)?;
if value.is_some()
&& request
.extra
.get("json_schema")
.is_some_and(|value| !value.is_null())
{
return Err(OpenAiError::invalid_request(
"grammar and json_schema cannot both be set",
));
}
Ok(value)
}
fn structured_output_json(
request: &ChatCompletionRequest,
name: &str,
) -> OpenAiResult<Option<String>> {
request
.extra
.get(name)
.filter(|value| !value.is_null())
.map(|value| {
if !value.is_object() {
return Err(OpenAiError::invalid_request(
"json_schema must be an object",
));
}
serialize_bounded_native_parser_json("json_schema", value)
})
.transpose()
}
fn serialize_bounded_native_parser_json(
name: &str,
value: &impl serde::Serialize,
) -> OpenAiResult<String> {
let serialized = serde_json::to_string(value)
.map_err(|error| OpenAiError::invalid_request(format!("serialize {name}: {error}")))?;
if serialized.len() > MAX_NATIVE_PARSER_INPUT_BYTES {
return Err(OpenAiError::invalid_request(format!(
"{name} exceeds the {MAX_NATIVE_PARSER_INPUT_BYTES}-byte limit"
)));
}
Ok(serialized)
}
fn default_reasoning_enabled(value: Option<EmbeddedReasoningEnabled>) -> Option<bool> {
match value {
Some(EmbeddedReasoningEnabled::Disabled) => Some(false),
Some(EmbeddedReasoningEnabled::Enabled) => Some(true),
Some(EmbeddedReasoningEnabled::Auto) | None => None,
}
}
fn default_reasoning_budget_enabled(value: Option<EmbeddedReasoningBudget>) -> Option<bool> {
match value {
Some(EmbeddedReasoningBudget::Tokens(0)) => Some(false),
Some(EmbeddedReasoningBudget::Tokens(_)) => Some(true),
Some(EmbeddedReasoningBudget::Effort(openai_frontend::ReasoningEffort::None)) => {
Some(false)
}
Some(EmbeddedReasoningBudget::Effort(_)) => Some(true),
Some(EmbeddedReasoningBudget::Unrestricted) => Some(true),
Some(EmbeddedReasoningBudget::Auto) | None => None,
}
}
fn chat_reasoning_format(value: Option<EmbeddedReasoningFormat>) -> ChatReasoningFormat {
match value.unwrap_or(EmbeddedReasoningFormat::Auto) {
EmbeddedReasoningFormat::Auto => ChatReasoningFormat::Auto,
EmbeddedReasoningFormat::None => ChatReasoningFormat::None,
EmbeddedReasoningFormat::Deepseek => ChatReasoningFormat::Deepseek,
EmbeddedReasoningFormat::DeepseekLegacy => ChatReasoningFormat::DeepseekLegacy,
EmbeddedReasoningFormat::Hidden => ChatReasoningFormat::Hidden,
}
}
pub(super) fn ensure_chat_runtime_features_supported(
request: &ChatCompletionRequest,
) -> OpenAiResult<()> {
if request.logprobs.unwrap_or(false) || request.top_logprobs.is_some() {
return Err(OpenAiError::unsupported(
"chat logprobs are parsed by openai-frontend but not yet implemented by skippy runtime",
));
}
Ok(())
}
pub(super) fn ensure_completion_runtime_features_supported(
request: &CompletionRequest,
) -> OpenAiResult<()> {
if request.logprobs.is_some() {
return Err(OpenAiError::unsupported(
"completion logprobs are parsed by openai-frontend but not yet implemented by skippy runtime",
));
}
Ok(())
}
pub(super) fn has_requested_tools(value: &Value) -> bool {
!matches!(value, Value::Array(items) if items.is_empty())
}
pub(super) fn ensure_extra_generation_fields_absent(
extra: &std::collections::BTreeMap<String, serde_json::Value>,
) -> OpenAiResult<()> {
const UNSUPPORTED_FIELDS: &[&str] = &["adaptive", "backend_sampling"];
for field in UNSUPPORTED_FIELDS {
if extra.get(*field).is_some_and(|value| !value.is_null()) {
return Err(OpenAiError::unsupported(format!(
"{field} is parsed but not yet implemented"
)));
}
}
Ok(())
}
pub(super) fn sampling_config(
temperature: Option<f32>,
top_p: Option<f32>,
presence_penalty: Option<f32>,
frequency_penalty: Option<f32>,
seed: Option<u64>,
logit_bias: Option<&std::collections::BTreeMap<String, serde_json::Value>>,
extra: &std::collections::BTreeMap<String, serde_json::Value>,
) -> OpenAiResult<SamplingConfig> {
ensure_extra_generation_fields_absent(extra)?;
let temperature = temperature.unwrap_or(0.8);
let top_p = top_p.unwrap_or(0.95);
let presence_penalty = presence_penalty.unwrap_or(0.0);
let frequency_penalty = frequency_penalty.unwrap_or(0.0);
let top_k = optional_i32_extra(extra, "top_k")?.unwrap_or(40);
let min_p = optional_f32_extra(extra, "min_p")?.unwrap_or(0.05);
let ignore_eos = optional_bool_extra(extra, "ignore_eos")?.unwrap_or(false);
let repeat_penalty = optional_f32_extra(extra, "repeat_penalty")?
.or(optional_f32_extra(extra, "repetition_penalty")?)
.unwrap_or(1.0);
let penalty_last_n =
optional_i32_extra(extra, "repeat_last_n")?.unwrap_or(DEFAULT_PENALTY_LAST_N);
let typical_p = optional_f32_extra(extra, "typical_p")?.unwrap_or(1.0);
let top_nsigma = optional_f32_extra(extra, "top_nsigma")?.unwrap_or(-1.0);
let dynatemp_range = optional_f32_extra(extra, "dynatemp_range")?.unwrap_or(0.0);
let dynatemp_exponent = optional_f32_extra(extra, "dynatemp_exponent")?.unwrap_or(1.0);
let dry = parse_dry_sampling(extra.get("dry"))?;
let xtc = parse_xtc_sampling(extra.get("xtc"))?;
let mirostat_mode = optional_i32_extra(extra, "mirostat_mode")?.unwrap_or(0);
let mirostat_entropy = optional_f32_extra(extra, "mirostat_entropy")?.unwrap_or(5.0);
let mirostat_learning_rate =
optional_f32_extra(extra, "mirostat_learning_rate")?.unwrap_or(0.1);
let samplers = parse_sampler_order(extra)?;
validate_sampling_range("temperature", temperature, 0.0..=100.0)?;
validate_sampling_range("top_p", top_p, 0.0..=1.0)?;
validate_sampling_range("presence_penalty", presence_penalty, -2.0..=2.0)?;
validate_sampling_range("frequency_penalty", frequency_penalty, -2.0..=2.0)?;
validate_sampling_range("min_p", min_p, 0.0..=1.0)?;
validate_sampling_range("repeat_penalty", repeat_penalty, 0.0..=100.0)?;
validate_sampling_range("typical_p", typical_p, 0.0..=1.0)?;
validate_sampling_range("top_nsigma", top_nsigma, -1.0..=f32::MAX)?;
validate_sampling_range("dynatemp_range", dynatemp_range, 0.0..=f32::MAX)?;
validate_sampling_range("dynatemp_exponent", dynatemp_exponent, 0.0..=f32::MAX)?;
validate_dry_sampling(&dry)?;
validate_sampling_range("xtc.probability", xtc.probability, 0.0..=1.0)?;
validate_sampling_range("xtc.threshold", xtc.threshold, 0.0..=1.0)?;
validate_mirostat_sampling(mirostat_mode, mirostat_entropy, mirostat_learning_rate)?;
if top_k < 0 {
return Err(OpenAiError::invalid_request(
"top_k must be greater than or equal to zero",
));
}
if penalty_last_n < -1 {
return Err(OpenAiError::invalid_request(
"repeat_last_n must be greater than or equal to -1",
));
}
let penalty_last_n = penalty_window(penalty_last_n);
let dry = DrySamplingConfig {
penalty_last_n: penalty_window(dry.penalty_last_n),
..dry
};
let seed = match seed {
Some(seed) => u32::try_from(seed)
.map_err(|_| OpenAiError::invalid_request("seed exceeds u32 range"))?,
None => 0,
};
let logit_bias = parse_logit_bias(logit_bias)?;
let defaults = SamplingConfig::default();
let enabled = seed != 0
|| temperature <= 0.0
|| (temperature - 1.0).abs() > f32::EPSILON
|| (top_p - 1.0).abs() > f32::EPSILON
|| top_k > 0
|| min_p > 0.0
|| presence_penalty.abs() > f32::EPSILON
|| frequency_penalty.abs() > f32::EPSILON
|| (repeat_penalty - 1.0).abs() > f32::EPSILON
|| penalty_last_n != defaults.penalty_last_n
|| !logit_bias.is_empty()
|| (typical_p - defaults.typical_p).abs() > f32::EPSILON
|| (top_nsigma - defaults.top_nsigma).abs() > f32::EPSILON
|| dynatemp_range.abs() > f32::EPSILON
|| (dynatemp_exponent - defaults.dynatemp_exponent).abs() > f32::EPSILON
|| dry != defaults.dry
|| xtc != defaults.xtc
|| mirostat_mode != defaults.mirostat_mode
|| (mirostat_entropy - defaults.mirostat_entropy).abs() > f32::EPSILON
|| (mirostat_learning_rate - defaults.mirostat_learning_rate).abs() > f32::EPSILON
|| samplers != defaults.samplers
|| ignore_eos;
Ok(SamplingConfig {
enabled,
ignore_eos,
seed,
temperature,
top_p,
top_k,
min_p,
presence_penalty,
frequency_penalty,
repeat_penalty,
penalty_last_n,
logit_bias,
typical_p,
top_nsigma,
dynatemp_range,
dynatemp_exponent,
dry,
xtc,
mirostat_mode,
mirostat_entropy,
mirostat_learning_rate,
samplers,
reasoning_budget: ReasoningBudget::Unrestricted,
})
}
fn validate_dry_sampling(dry: &DrySamplingConfig) -> OpenAiResult<()> {
validate_sampling_range("dry.multiplier", dry.multiplier, 0.0..=f32::MAX)?;
validate_positive_sampling_value("dry.base", dry.base)?;
if dry.allowed_length < 0 {
return Err(OpenAiError::invalid_request(
"dry.allowed_length must be greater than or equal to zero",
));
}
if dry.penalty_last_n < -1 {
return Err(OpenAiError::invalid_request(
"dry.penalty_last_n must be greater than or equal to -1",
));
}
Ok(())
}
fn validate_mirostat_sampling(mode: i32, entropy: f32, learning_rate: f32) -> OpenAiResult<()> {
if !matches!(mode, 0..=2) {
return Err(OpenAiError::invalid_request(
"mirostat_mode must be one of: 0 (disabled), 1, 2",
));
}
validate_positive_sampling_value("mirostat_entropy", entropy)?;
validate_positive_sampling_value("mirostat_learning_rate", learning_rate)
}
fn validate_positive_sampling_value(name: &str, value: f32) -> OpenAiResult<()> {
if !value.is_finite() || value <= 0.0 {
return Err(OpenAiError::invalid_request(format!(
"{name} is outside the supported range"
)));
}
Ok(())
}
fn parse_dry_sampling(value: Option<&Value>) -> OpenAiResult<DrySamplingConfig> {
let Some(value) = value.filter(|value| !value.is_null()) else {
return Ok(SamplingConfig::default().dry);
};
let object = value
.as_object()
.ok_or_else(|| OpenAiError::invalid_request("dry must be an object"))?;
ensure_allowed_object_keys(
object,
"dry",
&[
"multiplier",
"base",
"allowed_length",
"penalty_last_n",
"sequence_breakers",
],
)?;
let defaults = SamplingConfig::default().dry;
let sequence_breakers = match object.get("sequence_breakers") {
None | Some(Value::Null) => defaults.sequence_breakers,
Some(Value::Array(values)) if values.len() <= MAX_STAGE_DRY_SEQUENCE_BREAKERS => values
.iter()
.map(|value| {
let value = value.as_str().ok_or_else(|| {
OpenAiError::invalid_request("dry.sequence_breakers must contain strings")
})?;
if value.len() >= MAX_DRY_SEQUENCE_BREAKER_BYTES {
return Err(OpenAiError::invalid_request(
"dry.sequence_breakers entry exceeds maximum length",
));
}
Ok(value.to_string())
})
.collect::<OpenAiResult<Vec<_>>>()?,
Some(Value::Array(_)) => {
return Err(OpenAiError::invalid_request(
"dry.sequence_breakers contains too many entries",
));
}
Some(_) => {
return Err(OpenAiError::invalid_request(
"dry.sequence_breakers must be an array",
));
}
};
Ok(DrySamplingConfig {
multiplier: optional_object_f32(object, "dry.multiplier", "multiplier")?
.unwrap_or(defaults.multiplier),
base: optional_object_f32(object, "dry.base", "base")?.unwrap_or(defaults.base),
allowed_length: optional_object_i32(object, "dry.allowed_length", "allowed_length")?
.unwrap_or(defaults.allowed_length),
penalty_last_n: optional_object_i32(object, "dry.penalty_last_n", "penalty_last_n")?
.unwrap_or(defaults.penalty_last_n),
sequence_breakers,
})
}
fn parse_xtc_sampling(value: Option<&Value>) -> OpenAiResult<XtcSamplingConfig> {
let Some(value) = value.filter(|value| !value.is_null()) else {
return Ok(SamplingConfig::default().xtc);
};
let object = value
.as_object()
.ok_or_else(|| OpenAiError::invalid_request("xtc must be an object"))?;
ensure_allowed_object_keys(object, "xtc", &["probability", "threshold"])?;
let defaults = SamplingConfig::default().xtc;
Ok(XtcSamplingConfig {
probability: optional_object_f32(object, "xtc.probability", "probability")?
.unwrap_or(defaults.probability),
threshold: optional_object_f32(object, "xtc.threshold", "threshold")?
.unwrap_or(defaults.threshold),
})
}
fn optional_object_f32(
object: &serde_json::Map<String, Value>,
name: &str,
key: &str,
) -> OpenAiResult<Option<f32>> {
object
.get(key)
.filter(|value| !value.is_null())
.map(|value| {
serde_json::from_value(value.clone())
.map_err(|_| OpenAiError::invalid_request(format!("{name} must be a number")))
})
.transpose()
}
fn optional_object_i32(
object: &serde_json::Map<String, Value>,
name: &str,
key: &str,
) -> OpenAiResult<Option<i32>> {
object
.get(key)
.filter(|value| !value.is_null())
.map(|value| {
serde_json::from_value(value.clone())
.map_err(|_| OpenAiError::invalid_request(format!("{name} must be an integer")))
})
.transpose()
}
fn ensure_allowed_object_keys(
object: &serde_json::Map<String, Value>,
name: &str,
allowed: &[&str],
) -> OpenAiResult<()> {
if let Some(key) = object.keys().find(|key| !allowed.contains(&key.as_str())) {
return Err(OpenAiError::invalid_request(format!(
"{name} contains unknown field {key}"
)));
}
Ok(())
}
fn parse_sampler_order(
extra: &std::collections::BTreeMap<String, Value>,
) -> OpenAiResult<Vec<String>> {
if let Some(value) = extra.get("samplers").filter(|value| !value.is_null()) {
let values = value
.as_array()
.ok_or_else(|| OpenAiError::invalid_request("samplers must be an array"))?;
if values.len() > MAX_STAGE_SAMPLERS {
return Err(OpenAiError::invalid_request(
"samplers contains too many entries",
));
}
return values
.iter()
.map(|value| {
let sampler = value
.as_str()
.ok_or_else(|| OpenAiError::invalid_request("samplers must contain strings"))?;
canonical_sampler_name(sampler).map(str::to_string)
})
.collect();
}
if let Some(value) = extra
.get("sampler_sequence")
.filter(|value| !value.is_null())
{
let sequence = value
.as_str()
.ok_or_else(|| OpenAiError::invalid_request("sampler_sequence must be a string"))?;
let samplers = sequence
.chars()
.filter(|value| !value.is_whitespace())
.map(|value| match value {
'e' => Ok("penalties"),
'd' => Ok("dry"),
's' => Ok("top_n_sigma"),
'k' => Ok("top_k"),
'y' => Ok("typical_p"),
'p' => Ok("top_p"),
'm' => Ok("min_p"),
'x' => Ok("xtc"),
't' => Ok("temperature"),
_ => Err(OpenAiError::invalid_request(
"sampler_sequence contains an unsupported sampler code",
)),
})
.map(|result| result.map(str::to_string))
.collect::<OpenAiResult<Vec<_>>>()?;
if samplers.len() > MAX_STAGE_SAMPLERS {
return Err(OpenAiError::invalid_request(
"sampler_sequence contains too many entries",
));
}
return Ok(samplers);
}
Ok(SamplingConfig::default().samplers)
}
fn canonical_sampler_name(value: &str) -> OpenAiResult<&'static str> {
match value {
"penalties" => Ok("penalties"),
"dry" => Ok("dry"),
"top_n_sigma" => Ok("top_n_sigma"),
"top_k" => Ok("top_k"),
"typical_p" | "typ_p" => Ok("typical_p"),
"top_p" => Ok("top_p"),
"min_p" => Ok("min_p"),
"xtc" => Ok("xtc"),
"temperature" | "temp" => Ok("temperature"),
_ => Err(OpenAiError::invalid_request(
"samplers contains an unsupported sampler name",
)),
}
}
pub(super) fn parse_logit_bias(
logit_bias: Option<&std::collections::BTreeMap<String, serde_json::Value>>,
) -> OpenAiResult<Vec<RuntimeLogitBias>> {
let Some(logit_bias) = logit_bias else {
return Ok(Vec::new());
};
if logit_bias.len() > MAX_LOGIT_BIAS {
return Err(OpenAiError::invalid_request(format!(
"logit_bias supports at most {MAX_LOGIT_BIAS} entries"
)));
}
let mut parsed = Vec::with_capacity(logit_bias.len());
for (token_id, bias) in logit_bias {
let token_id = token_id
.parse::<i32>()
.map_err(|_| OpenAiError::invalid_request("logit_bias token IDs must be integers"))?;
if token_id < 0 {
return Err(OpenAiError::invalid_request(
"logit_bias token IDs must be greater than or equal to zero",
));
}
let bias = serde_json::from_value::<f32>(bias.clone())
.map_err(|_| OpenAiError::invalid_request("logit_bias values must be numbers"))?;
validate_sampling_range("logit_bias", bias, -100.0..=100.0)?;
parsed.push(RuntimeLogitBias { token_id, bias });
}
Ok(parsed)
}
pub(super) fn validate_sampling_range(
name: &str,
value: f32,
range: std::ops::RangeInclusive<f32>,
) -> OpenAiResult<()> {
if !value.is_finite() || !range.contains(&value) {
return Err(OpenAiError::invalid_request(format!(
"{name} is outside the supported range"
)));
}
Ok(())
}
pub(super) fn optional_f32_extra(
extra: &std::collections::BTreeMap<String, serde_json::Value>,
field: &str,
) -> OpenAiResult<Option<f32>> {
extra
.get(field)
.filter(|value| !value.is_null())
.map(|value| {
serde_json::from_value::<f32>(value.clone())
.map_err(|_| OpenAiError::invalid_request(format!("{field} must be a number")))
})
.transpose()
}
pub(super) fn optional_i32_extra(
extra: &std::collections::BTreeMap<String, serde_json::Value>,
field: &str,
) -> OpenAiResult<Option<i32>> {
extra
.get(field)
.filter(|value| !value.is_null())
.map(|value| {
serde_json::from_value::<i32>(value.clone())
.map_err(|_| OpenAiError::invalid_request(format!("{field} must be an integer")))
})
.transpose()
}
pub(super) fn wire_sampling_config(sampling: &SamplingConfig) -> Option<WireSamplingConfig> {
if !sampling.enabled {
return None;
}
let mut wire = WireSamplingConfig {
flags: (if sampling.enabled {
sampling_flags::ENABLED
} else {
0
}) | (if sampling.ignore_eos {
sampling_flags::IGNORE_EOS
} else {
0
}),
seed: sampling.seed,
temperature: sampling.temperature,
top_p: sampling.top_p,
top_k: sampling.top_k,
min_p: sampling.min_p,
presence_penalty: sampling.presence_penalty,
frequency_penalty: sampling.frequency_penalty,
repeat_penalty: sampling.repeat_penalty,
penalty_last_n: sampling.penalty_last_n,
typical_p: sampling.typical_p,
top_nsigma: sampling.top_nsigma,
dynatemp_range: sampling.dynatemp_range,
dynatemp_exponent: sampling.dynatemp_exponent,
dry_multiplier: sampling.dry.multiplier,
dry_base: sampling.dry.base,
dry_allowed_length: sampling.dry.allowed_length,
dry_penalty_last_n: sampling.dry.penalty_last_n,
dry_sequence_breakers: sampling.dry.sequence_breakers.clone(),
xtc_probability: sampling.xtc.probability,
xtc_threshold: sampling.xtc.threshold,
mirostat_mode: sampling.mirostat_mode,
mirostat_entropy: sampling.mirostat_entropy,
mirostat_learning_rate: sampling.mirostat_learning_rate,
samplers: sampling.samplers.clone(),
reasoning_budget_tokens: match sampling.reasoning_budget {
ReasoningBudget::Unrestricted => -1,
ReasoningBudget::Explicit(tokens) => i32::try_from(tokens).unwrap_or(i32::MAX),
ReasoningBudget::Capped(_) => return None,
ReasoningBudget::Resolved(tokens) => tokens,
},
ignore_eos: sampling.ignore_eos,
..WireSamplingConfig::default()
};
wire.logit_bias = sampling
.logit_bias
.iter()
.take(MAX_STAGE_LOGIT_BIAS)
.map(|source| WireLogitBias {
token_id: source.token_id,
bias: source.bias,
})
.collect();
Some(wire)
}