use std::path::Path;
use axum::extract::{Path as AxPath, State};
use axum::http::StatusCode;
use axum::response::{IntoResponse, Response};
use axum::Json;
use std::sync::Arc;
use super::engine::{self, Engine, SamplingParams};
use super::grammar;
use super::registry;
use super::schema::{
ApiError, ApiErrorBody, ChatCompletionChoice, ChatCompletionRequest, ChatCompletionResponse,
ChatMessage, ChoiceLogprobs, EmbeddingObject, EmbeddingPayload, EmbeddingRequest,
EmbeddingResponse, EmbeddingUsage, HealthResponse, MessageContent, ModelListResponse,
ModelObject, PromptTokensDetails, ReadyzResponse, ResponseFormat, TokenLogprob, UsageStats,
};
use super::state::AppState;
use crate::serve::auto_pipeline;
use crate::serve::multi_model::{EngineConfig, HotSwapError, LoadedEngine, PoolError};
use crate::serve::quant_select::QuantType;
pub(crate) const BOS_PROBE_FRAGMENTS: &[&str] =
&["<bos>", "<|begin_of_text|>", "<s>", "<|im_start|>"];
pub(crate) fn probe_bos_token_id(tokenizer: &tokenizers::Tokenizer) -> Option<u32> {
BOS_PROBE_FRAGMENTS
.iter()
.find_map(|t| tokenizer.token_to_id(t))
}
async fn resolve_engine_for_request(
state: &AppState,
requested_model: &str,
) -> std::result::Result<Arc<LoadedEngine<Engine>>, Response> {
let trimmed = requested_model.trim();
let model_arg: String = if !trimmed.is_empty() {
trimmed.to_string()
} else if let Some(d) = state.default_model.as_deref() {
d.to_string()
} else {
return Err(ApiError::model_not_loaded(requested_model).into_response());
};
if let Ok(pool_guard) = state.pool.read() {
for le in pool_guard.snapshot_engines() {
if le.engine.model_id() == model_arg {
drop(pool_guard);
return Ok(le);
}
}
drop(pool_guard);
}
let cache_arc = state.cache.clone();
let hardware = state.hardware.clone();
let no_integrity = state.no_integrity;
let model_arg_for_resolve = model_arg.clone();
let resolve_outcome = tokio::task::spawn_blocking(move || {
let mut cache_guard = cache_arc
.lock()
.map_err(|e| anyhow::anyhow!("cache mutex poisoned: {e}"))?;
auto_pipeline::resolve_or_prepare_model(
&model_arg_for_resolve,
&mut cache_guard,
hardware.as_ref(),
no_integrity,
)
})
.await;
let resolved = match resolve_outcome {
Ok(Ok(r)) => r,
Ok(Err(e)) => {
tracing::debug!(model = %model_arg, error = %e, "auto-pipeline rejected request model");
return Err(ApiError::model_not_loaded(&model_arg).into_response());
}
Err(join_err) => {
tracing::error!(error = %join_err, "auto-pipeline blocking task panicked");
return Err(ApiError::internal_error().into_response());
}
};
let pool_repo = resolved
.repo_id
.clone()
.unwrap_or_else(|| crate::serve::pool_key_for_path(&resolved.gguf_path));
let pool_quant = resolved.quant.unwrap_or(QuantType::Q4_K_M);
let engine_config = EngineConfig {
tokenizer_path: None,
config_path: None,
queue_capacity: state.engine_queue_capacity,
warmup_synchronously: true,
kv_metrics_sink: Some(Arc::clone(&state.kv_spill_counters)
as Arc<dyn crate::serve::kv_persist::metrics::KvCacheMetricsSink>),
dwq_overlay_path: None,
engine_mode: crate::serve::api::engine::EngineMode::SerialFifo,
};
let pool_arc = state.pool.clone();
let pool_repo_blocking = pool_repo.clone();
let gguf_path = resolved.gguf_path.clone();
let load_outcome: Result<
Result<Arc<LoadedEngine<Engine>>, HotSwapError>,
tokio::task::JoinError,
> = tokio::task::spawn_blocking(move || {
let mut w = pool_arc.write().map_err(|e| {
HotSwapError::LoaderFailed(anyhow::anyhow!("pool rwlock poisoned: {e}"))
})?;
w.load_or_get(&pool_repo_blocking, pool_quant, &gguf_path, &engine_config)
})
.await;
match load_outcome {
Ok(Ok(arc)) => Ok(arc),
Ok(Err(e)) => {
tracing::error!(model = %model_arg, error = %e, "hot-swap load failed");
Err(map_hotswap_error_to_response(e))
}
Err(join_err) => {
tracing::error!(error = %join_err, "hot-swap blocking task panicked");
Err(ApiError::internal_error().into_response())
}
}
}
fn map_hotswap_error_to_response(e: HotSwapError) -> Response {
match e {
HotSwapError::PoolRefused(ref pool_err) => {
let detail = match pool_err {
PoolError::OversizedHandle { ref repo_id, .. } => format!(
"model {repo_id} exceeds pool memory budget; \
try a smaller quant or raise the budget"
),
PoolError::ZeroCapacity => "multi-model pool has capacity_models = 0; \
operator must raise capacity"
.to_string(),
};
ApiError {
status: StatusCode::SERVICE_UNAVAILABLE,
retry_after_seconds: Some(5),
error: ApiErrorBody {
message: detail,
error_type: "server_error".to_string(),
code: Some("pool_refused".to_string()),
param: None,
},
}
.into_response()
}
HotSwapError::LoaderFailed(ref inner) => {
tracing::error!(error = %inner, "hot-swap loader failed");
ApiError::generation_error(format!("model load failed: {inner}")).into_response()
}
HotSwapError::FileSize {
ref path,
ref source,
} => {
tracing::error!(
path = %path.display(), error = %source,
"hot-swap file size read failed"
);
ApiError::generation_error(format!(
"cannot read GGUF size at {}: {source}",
path.display()
))
.into_response()
}
}
}
pub async fn chat_completions(
State(state): State<AppState>,
Json(req): Json<ChatCompletionRequest>,
) -> Response {
use std::sync::atomic::Ordering;
state.metrics.requests_total.fetch_add(1, Ordering::Relaxed);
let prepared = match prepare_chat_generation(&state, &req).await {
Ok(p) => p,
Err(resp) => {
state
.metrics
.requests_rejected_total
.fetch_add(1, Ordering::Relaxed);
return resp;
}
};
state
.metrics
.chat_completions_started
.fetch_add(1, Ordering::Relaxed);
if req.stream.unwrap_or(false) {
return chat_completions_stream(state.clone(), req, prepared).await;
}
chat_completions_with_prepared(state, req, prepared).await
}
async fn chat_completions_with_prepared(
state: AppState,
req: ChatCompletionRequest,
prepared: PreparedChatContext,
) -> Response {
use std::sync::atomic::Ordering;
let PreparedChatContext {
loaded_engine,
prompt_tokens,
params,
summarized_messages,
summary_tokens,
soft_tokens,
vit_forward_ms,
vit_images,
vit_soft_tokens_total,
deepstack_data,
positions_flat,
} = prepared;
let engine: &Engine = &loaded_engine.engine;
let tool_call_policy = params.tool_call_policy;
let pre_dispatches = mlx_native::dispatch_count();
let pre_syncs = mlx_native::sync_count();
let gen_started = std::time::Instant::now();
let gen_outcome = if soft_tokens.is_empty() {
engine.generate(prompt_tokens, params).await
} else if deepstack_data.is_some() || positions_flat.is_some() {
engine
.generate_with_soft_tokens_and_deepstack(
prompt_tokens,
soft_tokens,
params,
deepstack_data,
positions_flat,
)
.await
} else {
engine
.generate_with_soft_tokens(prompt_tokens, soft_tokens, params)
.await
};
let result = match gen_outcome {
Ok(r) => r,
Err(e) => {
let msg = format!("{e:#}");
if msg.contains("queue_full") {
state
.metrics
.chat_completions_queue_full
.fetch_add(1, Ordering::Relaxed);
return queue_full_with_rate_limit_headers(&state);
}
if msg.contains("slot_budget_exceeded") {
let (needed, budget) = parse_slot_budget_exceeded(&msg);
return ApiError::slot_budget_exceeded(needed, budget).into_response();
}
if msg.contains(engine::QWEN35_NOT_IMPLEMENTED_SENTINEL) {
tracing::info!(
error = %msg,
"chat_completion routed to Qwen3.5/3.6 SERVE arm (Wedge-3 pending) — returning 501"
);
return ApiError::not_implemented(
engine::QWEN35_NOT_IMPLEMENTED_MESSAGE.to_string(),
)
.into_response();
}
if msg.contains(
crate::inference::models::qwen3vl_text::forward::QWEN3VL_TEXT_FORWARD_PENDING_SENTINEL,
) {
tracing::info!(
error = %msg,
"chat_completion routed to Qwen3-VL text SERVE arm (iter-228b pending) — returning 501"
);
return ApiError::not_implemented(
crate::inference::models::qwen3vl_text::forward::QWEN3VL_TEXT_FORWARD_PENDING_MESSAGE
.to_string(),
)
.into_response();
}
tracing::error!(error = %msg, "chat_completion generation failed");
return ApiError::generation_error(msg).into_response();
}
};
let total_time = gen_started.elapsed();
state
.metrics
.chat_completions_completed
.fetch_add(1, Ordering::Relaxed);
state
.metrics
.prompt_tokens_total
.fetch_add(result.prompt_tokens as u64, Ordering::Relaxed);
state
.metrics
.decode_tokens_total
.fetch_add(result.completion_tokens as u64, Ordering::Relaxed);
let request_id = format!("chatcmpl-{}", uuid::Uuid::new_v4());
let created = chrono_seconds();
let system_fingerprint = state.config.system_fingerprint.clone();
let prefill_time_secs = result.prefill_duration.as_secs_f64();
let decode_time_secs = result.decode_duration.as_secs_f64();
let total_time_secs = total_time.as_secs_f64();
let prefill_tokens_per_sec = if prefill_time_secs > 0.0 {
result.prompt_tokens as f64 / prefill_time_secs
} else {
0.0
};
let decode_tokens_per_sec = if decode_time_secs > 0.0 {
result.completion_tokens as f64 / decode_time_secs
} else {
0.0
};
let ttft_ms = prefill_time_secs * 1000.0;
let post_dispatches = mlx_native::dispatch_count();
let post_syncs = mlx_native::sync_count();
let timing = super::schema::TimingInfo {
prefill_time_secs,
decode_time_secs,
total_time_secs,
time_to_first_token_ms: ttft_ms,
prefill_tokens_per_sec,
decode_tokens_per_sec,
gpu_sync_count: post_syncs.saturating_sub(pre_syncs),
gpu_dispatch_count: post_dispatches.saturating_sub(pre_dispatches),
};
let reasoning_tokens = result.reasoning_tokens;
let registration = engine.registration();
let extracted = extract_tool_calls_from_text(&result.text, registration, tool_call_policy);
if let Some(failed_body) = extracted.constrained_parse_failure {
tracing::error!(
body = %failed_body,
"tool_call_parse_failure: Constrained tool_choice but model emitted unparseable body"
);
return ApiError::generation_error("tool_call_parse_failure".to_string()).into_response();
}
if let Some(resp) = defensive_no_call_under_constrained(
tool_call_policy,
extracted.tool_calls.is_empty(),
result.finish_reason,
result.completion_tokens,
result.text.len(),
) {
return resp;
}
let (message_content, message_tool_calls, effective_finish_reason) =
if extracted.tool_calls.is_empty() {
(
Some(MessageContent::Text(result.text)),
None,
result.finish_reason.to_string(),
)
} else {
let content = tool_turn_message_content(extracted.content);
(
content,
Some(extracted.tool_calls),
"tool_calls".to_string(),
)
};
let resp = ChatCompletionResponse {
id: request_id,
object: "chat.completion",
created,
model: req.model.clone(),
system_fingerprint,
choices: vec![ChatCompletionChoice {
index: 0,
message: ChatMessage {
role: "assistant".into(),
content: message_content,
reasoning_content: result.reasoning_text,
tool_calls: message_tool_calls,
tool_call_id: None,
name: None,
},
finish_reason: effective_finish_reason,
logprobs: result.logprobs.as_ref().map(|lps| ChoiceLogprobs {
content: lps
.iter()
.map(|&lp| TokenLogprob {
token: String::new(),
logprob: lp,
bytes: None,
top_logprobs: Vec::new(),
})
.collect(),
}),
}],
usage: UsageStats {
prompt_tokens: result.prompt_tokens,
completion_tokens: result.completion_tokens,
total_tokens: result.prompt_tokens + result.completion_tokens,
prompt_tokens_details: Some(PromptTokensDetails {
cached_tokens: result.cached_tokens,
}),
completion_tokens_details: reasoning_tokens.map(|reasoning_tokens| {
super::schema::CompletionTokensDetails { reasoning_tokens }
}),
},
x_hf2q_timing: Some(timing),
};
let mut response = (StatusCode::OK, Json(resp)).into_response();
apply_transparency_headers(
&state,
&req,
&mut response,
summarized_messages,
summary_tokens,
);
apply_vit_transparency_headers(
&mut response,
vit_forward_ms,
vit_images,
vit_soft_tokens_total,
);
response
}
struct ExtractedToolCalls {
pub content: String,
pub tool_calls: Vec<super::schema::ToolCall>,
pub constrained_parse_failure: Option<String>,
}
fn tool_turn_message_content(content: String) -> Option<MessageContent> {
(!content.trim().is_empty()).then_some(MessageContent::Text(content))
}
fn extract_tool_calls_from_text(
text: &str,
registration: Option<®istry::ModelRegistration>,
policy: engine::ToolCallPolicy,
) -> ExtractedToolCalls {
let Some(reg) = registration else {
return ExtractedToolCalls {
content: text.to_string(),
tool_calls: Vec::new(),
constrained_parse_failure: None,
};
};
let Some(mut tool_splitter) = registry::ToolCallSplitter::from_registration(reg) else {
return ExtractedToolCalls {
content: text.to_string(),
tool_calls: Vec::new(),
constrained_parse_failure: None,
};
};
let mut content = String::new();
let mut tool_calls: Vec<super::schema::ToolCall> = Vec::new();
let mut tc_body = String::new();
let mut tc_index: usize = 0;
let mut parse_failure: Option<String> = None;
let mut process_events = |events: Vec<registry::ToolCallEvent>| {
for ev in events {
if parse_failure.is_some() {
break; }
match ev {
registry::ToolCallEvent::Content(t) => {
content.push_str(&t);
}
registry::ToolCallEvent::ToolCallOpen => {
tc_body.clear();
}
registry::ToolCallEvent::ToolCallText(t) => {
tc_body.push_str(&t);
}
registry::ToolCallEvent::ToolCallClose => {
let body = std::mem::take(&mut tc_body);
match registry::parse_tool_call_bodies(reg, &body) {
Some(parsed_calls) if !parsed_calls.is_empty() => {
for parsed in parsed_calls {
let id = format!(
"call_hf2q_{:016x}",
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_nanos() as u64)
.unwrap_or(0)
^ (tc_index as u64).wrapping_mul(0x9e3779b97f4a7c15)
);
tool_calls.push(super::schema::ToolCall {
id,
call_type: "function".to_string(),
function: super::schema::ToolCallFunction {
name: parsed.name,
arguments: parsed.arguments_json,
},
});
tc_index += 1;
}
}
_ => {
if policy.enforces_body_grammar() {
let policy_label = match policy {
engine::ToolCallPolicy::Constrained => "required/function",
engine::ToolCallPolicy::AutoLazyGrammar => {
"auto (lazy grammar active)"
}
engine::ToolCallPolicy::Auto => {
unreachable!("enforces_body_grammar gate")
}
};
tracing::error!(
body = %body,
policy = policy_label,
"non-stream tool-call parse failure under {} policy \
— per-model body grammar should have prevented this",
policy_label
);
parse_failure = Some(body);
} else {
let scrubbed = registry::scrub_special_tokens(&body);
tracing::warn!(
body = %body,
scrubbed = %scrubbed,
"non-stream tool-call body unparseable; emitting as \
content (tool_choice=auto with no active grammar). \
Special-token markers scrubbed before emit."
);
content.push_str(&scrubbed);
}
}
}
}
}
}
};
let events = tool_splitter.feed(text);
process_events(events);
if parse_failure.is_none() {
if let Some(tail_ev) = tool_splitter.finish() {
match tail_ev {
registry::ToolCallEvent::Content(t) => content.push_str(&t),
registry::ToolCallEvent::ToolCallText(t) => {
let prefix = reg.tool_open.unwrap_or("");
content.push_str(&format!("{prefix}{t}"));
}
registry::ToolCallEvent::ToolCallOpen | registry::ToolCallEvent::ToolCallClose => {}
}
}
}
ExtractedToolCalls {
content,
tool_calls,
constrained_parse_failure: parse_failure,
}
}
fn apply_transparency_headers(
state: &AppState,
req: &ChatCompletionRequest,
resp: &mut Response,
summarized_messages: Option<usize>,
summary_tokens: Option<usize>,
) {
use super::schema::OverflowPolicy;
use axum::http::{header::HeaderName, HeaderValue};
let policy = req
.hf2q_overflow_policy
.unwrap_or(state.config.default_overflow_policy);
let policy_str: &'static str = match policy {
OverflowPolicy::Reject => "reject",
OverflowPolicy::TruncateLeft => "truncate_left",
OverflowPolicy::Summarize => "summarize",
};
let headers = resp.headers_mut();
headers.insert(
HeaderName::from_static("x-hf2q-overflow-policy"),
HeaderValue::from_static(policy_str),
);
if let Some(n) = summarized_messages {
if let Ok(v) = HeaderValue::from_str(&n.to_string()) {
headers.insert(HeaderName::from_static("x-hf2q-summarized-messages"), v);
}
}
if let Some(n) = summary_tokens {
if let Ok(v) = HeaderValue::from_str(&n.to_string()) {
headers.insert(HeaderName::from_static("x-hf2q-summary-tokens"), v);
}
}
}
type ResolverBoxFuture<'s> = std::pin::Pin<
Box<
dyn std::future::Future<Output = std::result::Result<Arc<LoadedEngine<Engine>>, Response>>
+ Send
+ 's,
>,
>;
struct PreparedChatContext {
loaded_engine: Arc<LoadedEngine<Engine>>,
prompt_tokens: Vec<u32>,
params: SamplingParams,
summarized_messages: Option<usize>,
summary_tokens: Option<usize>,
soft_tokens: Vec<engine::SoftTokenData>,
vit_forward_ms: Option<u64>,
vit_images: Option<usize>,
vit_soft_tokens_total: Option<usize>,
deepstack_data: Option<engine::DeepstackData>,
positions_flat: Option<Vec<i32>>,
}
async fn prepare_chat_generation(
state: &AppState,
req: &ChatCompletionRequest,
) -> std::result::Result<PreparedChatContext, Response> {
prepare_chat_generation_core(state, req, |s, m| {
Box::pin(async move { resolve_engine_for_request(s, &m).await })
})
.await
}
async fn prepare_chat_generation_core<'s, F>(
state: &'s AppState,
req: &ChatCompletionRequest,
resolver: F,
) -> std::result::Result<PreparedChatContext, Response>
where
F: FnOnce(&'s AppState, String) -> ResolverBoxFuture<'s>,
{
if !state.is_ready_for_gen() {
return Err(ApiError::not_ready().into_response());
}
let loaded_engine = resolver(state, req.model.clone()).await?;
let engine: &Engine = &loaded_engine.engine;
if req.messages.is_empty() {
return Err(ApiError::invalid_request(
"messages must contain at least one entry",
Some("messages".into()),
)
.into_response());
}
let preprocessed_inputs = match process_multimodal_content(&req.messages, state.mmproj.as_ref())
{
Ok(imgs) => imgs,
Err(resp) => return Err(resp),
};
let (
messages_for_render,
vision_embeddings,
vit_forward_ms_v,
vit_images_v,
vision_family,
per_row_floats,
qwen3vl_image_grids,
) = if preprocessed_inputs.is_empty() {
(
req.messages.clone(),
Vec::new(),
None,
None,
crate::inference::vision::mmproj::VisionFamily::Gemma,
engine.hidden_size(),
Vec::<(u32, u32)>::new(),
)
} else {
let mmproj = state
.mmproj
.as_ref()
.expect("mmproj checked in process_multimodal_content");
let pipeline_out = match crate::inference::vision::pipeline::run_vit_forward(
&preprocessed_inputs,
mmproj,
engine.hidden_size(),
) {
Ok(out) => out,
Err(e) => {
return Err(
ApiError::generation_error(format!("ViT forward failed: {e:#}"))
.into_response(),
);
}
};
let n_images = pipeline_out.embeddings.len();
let rewritten =
rewrite_messages_for_vision_placeholders_family(&req.messages, pipeline_out.family);
(
rewritten,
pipeline_out.embeddings,
Some(pipeline_out.forward_ms),
Some(n_images),
pipeline_out.family,
pipeline_out.per_row_floats,
pipeline_out.qwen3vl_image_grids,
)
};
let response_grammar: Option<grammar::Grammar> = match req.response_format.as_ref() {
Some(rf) => match compile_response_format(rf) {
Ok(g) => g,
Err(resp) => return Err(resp),
},
None => None,
};
let policy = req
.hf2q_overflow_policy
.unwrap_or(state.config.default_overflow_policy);
let tool_choice = super::schema::ToolChoiceValue::parse(req.tool_choice.as_ref());
let req_tools: Option<&[super::schema::Tool]> = match tool_choice {
super::schema::ToolChoiceValue::None => None,
_ => req.tools.as_deref(),
};
let tool_grammar: Option<grammar::Grammar> = match compile_tool_grammar(req, &tool_choice) {
Ok(g) => g,
Err(resp) => return Err(resp),
};
let tool_grammar_kind = tool_grammar_kind_for(&tool_choice);
let (effective_grammar, effective_grammar_kind) =
select_effective_grammar(tool_grammar, tool_grammar_kind, response_grammar);
let enable_thinking = req.hf2q_enable_thinking.unwrap_or(false);
let template_kwargs = req.chat_template_kwargs.as_ref();
let (prompt_tokens, _prompt_len, summarized_messages, summary_tokens) =
match apply_overflow_policy(
engine,
&messages_for_render,
policy,
req_tools,
enable_thinking,
template_kwargs,
)
.await
{
Ok(r) => r,
Err(resp) => return Err(resp),
};
let reasoning_forced_open = match super::registry::find_for(&req.model) {
Some(reg) => {
let probe_msgs = [super::schema::ChatMessage {
role: "user".to_string(),
content: Some(super::schema::MessageContent::Text("x".to_string())),
reasoning_content: None,
tool_calls: None,
tool_call_id: None,
name: None,
}];
let rendered = render_chat_prompt_or_400(
engine.chat_template(),
&probe_msgs,
req_tools,
enable_thinking,
template_kwargs,
)?;
super::registry::prompt_seeds_reasoning_open(&rendered, ®)
}
None => false,
};
let max_tokens = req
.max_completion_tokens
.or(req.max_tokens)
.unwrap_or(SamplingParams::default().max_tokens);
let stop_strings = req.stop.clone().map(|s| s.into_vec()).unwrap_or_default();
let logit_bias: std::collections::HashMap<u32, f32> = req
.logit_bias
.as_ref()
.map(|m| {
m.iter()
.filter_map(|(k, v)| k.parse::<u32>().ok().map(|id| (id, *v)))
.collect()
})
.unwrap_or_default();
let grammar_token_bytes: Option<std::sync::Arc<Vec<Vec<u8>>>> = if effective_grammar.is_some() {
Some(engine.token_bytes_table())
} else {
None
};
let tc_policy = match &tool_choice {
super::schema::ToolChoiceValue::Required | super::schema::ToolChoiceValue::Function(_) => {
engine::ToolCallPolicy::Constrained
}
super::schema::ToolChoiceValue::Auto
if effective_grammar.is_some()
&& matches!(
effective_grammar_kind,
engine::GrammarKind::ToolCallBodyAuto
) =>
{
engine::ToolCallPolicy::AutoLazyGrammar
}
_ => engine::ToolCallPolicy::Auto,
};
let params = SamplingParams {
temperature: req.temperature.unwrap_or(0.0),
top_p: req.top_p.unwrap_or(1.0),
top_k: req.top_k.map(|v| v as usize).unwrap_or(0),
repetition_penalty: req.repetition_penalty.unwrap_or(1.0),
max_tokens,
stop_strings,
frequency_penalty: req.frequency_penalty.unwrap_or(0.0),
presence_penalty: req.presence_penalty.unwrap_or(0.0),
seed: req.seed,
min_p: req.min_p.unwrap_or(0.0),
logit_bias,
logprobs: req.logprobs.unwrap_or(false),
top_logprobs: req.top_logprobs.unwrap_or(0),
parallel_tool_calls: req.parallel_tool_calls.unwrap_or(true),
grammar: effective_grammar,
token_bytes: grammar_token_bytes,
grammar_kind: effective_grammar_kind,
tool_call_policy: tc_policy,
reasoning_forced_open,
};
let (final_prompt_tokens, soft_tokens, image_token_positions_per_image): (
Vec<u32>,
Vec<engine::SoftTokenData>,
Vec<Vec<u32>>,
) = if vision_embeddings.is_empty() {
(
prompt_tokens,
Vec::<engine::SoftTokenData>::new(),
Vec::new(),
)
} else {
match expand_image_placeholders_family(
engine,
&prompt_tokens,
&vision_embeddings,
vision_family,
per_row_floats,
) {
Ok(triple) => triple,
Err(resp) => return Err(resp),
}
};
if let Some(ctx_len) = engine.context_length() {
if final_prompt_tokens.len() >= ctx_len {
return Err(
ApiError::context_length_exceeded(ctx_len, final_prompt_tokens.len())
.into_response(),
);
}
}
let vit_soft_tokens_total: Option<usize> = if soft_tokens.is_empty() {
None
} else {
Some(soft_tokens.iter().map(|s| s.range.len()).sum())
};
let (deepstack_data, positions_flat) = if matches!(
vision_family,
crate::inference::vision::mmproj::VisionFamily::Qwen3Vl,
) && !vision_embeddings.is_empty()
{
match dispatch_qwen3vl_seam_split(
engine,
state.mmproj.as_ref().expect("mmproj checked above"),
&vision_embeddings,
&image_token_positions_per_image,
&qwen3vl_image_grids,
&final_prompt_tokens,
) {
Ok((ds, pos)) => (Some(ds), Some(pos)),
Err(resp) => return Err(resp),
}
} else {
(None, None)
};
Ok(PreparedChatContext {
loaded_engine,
prompt_tokens: final_prompt_tokens,
params,
summarized_messages,
summary_tokens,
soft_tokens,
vit_forward_ms: vit_forward_ms_v,
vit_images: vit_images_v,
vit_soft_tokens_total,
deepstack_data,
positions_flat,
})
}
fn dispatch_qwen3vl_seam_split(
engine: &engine::Engine,
mmproj: &super::state::LoadedMmproj,
vision_embeddings: &[Vec<f32>],
image_token_positions_per_image: &[Vec<u32>],
qwen3vl_image_grids: &[(u32, u32)],
final_prompt_tokens: &[u32],
) -> std::result::Result<(engine::DeepstackData, Vec<i32>), Response> {
use crate::serve::forward_prefill::{build_qwen3vl_positions, Qwen3VlImageGrid};
let hidden = engine.hidden_size();
let n_deepstack = mmproj
.config
.deepstack_indexes
.as_ref()
.map(|v| v.len())
.unwrap_or(0);
let per_row_floats = hidden.saturating_mul(1 + n_deepstack);
if hidden == 0 || per_row_floats == 0 {
return Err(ApiError::generation_error(format!(
"dispatch_qwen3vl_seam_split: degenerate hidden ({}) or \
n_deepstack ({})",
hidden, n_deepstack
))
.into_response());
}
if vision_embeddings.len() != image_token_positions_per_image.len() {
return Err(ApiError::generation_error(format!(
"dispatch_qwen3vl_seam_split: vision_embeddings.len()={} != \
image_token_positions_per_image.len()={}",
vision_embeddings.len(),
image_token_positions_per_image.len()
))
.into_response());
}
let stride = mmproj
.config
.patch_size
.saturating_mul(mmproj.config.spatial_merge_size.unwrap_or(1));
if stride == 0 {
return Err(ApiError::generation_error(format!(
"dispatch_qwen3vl_seam_split: patch_size ({}) * \
spatial_merge_size ({:?}) = 0",
mmproj.config.patch_size, mmproj.config.spatial_merge_size
))
.into_response());
}
let image_size = mmproj.config.image_size;
if image_size == 0 || image_size % stride != 0 {
return Err(ApiError::generation_error(format!(
"dispatch_qwen3vl_seam_split: image_size ({}) is not a \
positive multiple of stride ({})",
image_size, stride
))
.into_response());
}
let mut per_image_n_image_tokens: Vec<usize> = Vec::with_capacity(vision_embeddings.len());
for (i, e) in vision_embeddings.iter().enumerate() {
if e.len() % per_row_floats != 0 {
return Err(ApiError::generation_error(format!(
"dispatch_qwen3vl_seam_split: vision_embeddings[{i}].len()={} \
not divisible by per_row_floats={per_row_floats}",
e.len()
))
.into_response());
}
let observed = e.len() / per_row_floats;
if observed == 0 {
return Err(ApiError::generation_error(format!(
"dispatch_qwen3vl_seam_split: vision_embeddings[{i}] has 0 \
image tokens (augmented-embed empty)"
))
.into_response());
}
if image_token_positions_per_image[i].len() != observed {
return Err(ApiError::generation_error(format!(
"dispatch_qwen3vl_seam_split: image_token_positions_per_image[{i}].len()={} \
!= observed n_image_tokens={observed} (must match the \
augmented-embed row count after per_row_floats={per_row_floats} split)",
image_token_positions_per_image[i].len()
))
.into_response());
}
per_image_n_image_tokens.push(observed);
}
let total_n_image_tokens: usize = per_image_n_image_tokens.iter().sum();
if total_n_image_tokens == 0 || n_deepstack == 0 {
let flat = build_qwen3vl_positions(final_prompt_tokens.len(), &[]).map_err(|e| {
ApiError::generation_error(format!("build_qwen3vl_positions (no images): {e}"))
.into_response()
})?;
return Ok((
engine::DeepstackData {
image_token_positions: Vec::new(),
chunks: Vec::new(),
},
flat,
));
}
let mlx_dev = mlx_native::MlxDevice::new().map_err(|e| {
ApiError::generation_error(format!("dispatch_qwen3vl_seam_split: MlxDevice::new: {e}"))
.into_response()
})?;
let chunk_byte_len = total_n_image_tokens * hidden * std::mem::size_of::<f32>();
let mut chunks: Vec<mlx_native::MlxBuffer> = Vec::with_capacity(n_deepstack);
for j in 0..n_deepstack {
let buf = mlx_dev
.alloc_buffer(
chunk_byte_len,
mlx_native::DType::F32,
vec![total_n_image_tokens, hidden],
)
.map_err(|e| {
ApiError::generation_error(format!(
"dispatch_qwen3vl_seam_split: alloc deepstack chunk {j}: {e}"
))
.into_response()
})?;
chunks.push(buf);
}
for j in 0..n_deepstack {
let dst = chunks[j].as_mut_slice::<f32>().map_err(|e| {
ApiError::generation_error(format!(
"dispatch_qwen3vl_seam_split: as_mut_slice chunk {j}: {e}"
))
.into_response()
})?;
let mut row_offset: usize = 0;
for (i, src) in vision_embeddings.iter().enumerate() {
let n_img_i = per_image_n_image_tokens[i];
for r in 0..n_img_i {
let src_base = r * per_row_floats + (j + 1) * hidden;
let dst_base = (row_offset + r) * hidden;
dst[dst_base..dst_base + hidden].copy_from_slice(&src[src_base..src_base + hidden]);
}
row_offset += n_img_i;
}
}
let mut all_positions: Vec<u32> = Vec::with_capacity(total_n_image_tokens);
for img in image_token_positions_per_image.iter() {
all_positions.extend_from_slice(img);
}
debug_assert_eq!(all_positions.len(), total_n_image_tokens);
if !qwen3vl_image_grids.is_empty() && qwen3vl_image_grids.len() != vision_embeddings.len() {
return Err(ApiError::generation_error(format!(
"dispatch_qwen3vl_seam_split: qwen3vl_image_grids.len()={} != \
vision_embeddings.len()={} — one (n_x, n_y) entry required \
per image",
qwen3vl_image_grids.len(),
vision_embeddings.len(),
))
.into_response());
}
let canonical_per_side = image_size / stride;
let mut image_grids: Vec<(Qwen3VlImageGrid, u32)> =
Vec::with_capacity(image_token_positions_per_image.len());
for (i, img_positions) in image_token_positions_per_image.iter().enumerate() {
if img_positions.is_empty() {
continue; }
let seq_start = img_positions[0];
let n_tokens = per_image_n_image_tokens[i] as u32;
let (n_x, n_y) = if !qwen3vl_image_grids.is_empty() {
let (gx, gy) = qwen3vl_image_grids[i];
if gx == 0 || gy == 0 {
return Err(ApiError::generation_error(format!(
"dispatch_qwen3vl_seam_split: image[{i}] grid \
({gx},{gy}) has zero axis (preprocessor must report \
positive (n_x, n_y))"
))
.into_response());
}
if gx > canonical_per_side || gy > canonical_per_side {
return Err(ApiError::generation_error(format!(
"dispatch_qwen3vl_seam_split: image[{i}] grid \
({gx},{gy}) exceeds canonical canvas grid \
({canonical_per_side},{canonical_per_side}) — \
preprocessor canvas clamp violated"
))
.into_response());
}
if gx.saturating_mul(gy) != n_tokens {
return Err(ApiError::generation_error(format!(
"dispatch_qwen3vl_seam_split: image[{i}] grid \
({gx},{gy}) product {} != observed n_image_tokens \
{n_tokens} — producer/consumer mismatch between \
preprocess_qwen3vl and compute_vision_embeddings_gpu_qwen3vl",
gx.saturating_mul(gy)
))
.into_response());
}
(gx, gy)
} else {
let side = (n_tokens as f64).sqrt() as u32;
if side == 0 || side > canonical_per_side || side * side != n_tokens {
return Err(ApiError::generation_error(format!(
"dispatch_qwen3vl_seam_split: image[{i}] n_image_tokens={n_tokens} \
is not a square ≤ canonical_per_side² ({canonical_per_side}²) \
and no explicit per-image grid was supplied — caller must \
thread `qwen3vl_image_grids` for non-square Phase-2 inputs"
))
.into_response());
}
(side, side)
};
image_grids.push((Qwen3VlImageGrid { n_x, n_y }, seq_start));
}
let positions_flat =
build_qwen3vl_positions(final_prompt_tokens.len(), &image_grids).map_err(|e| {
ApiError::generation_error(format!(
"dispatch_qwen3vl_seam_split: build_qwen3vl_positions: {e}"
))
.into_response()
})?;
Ok((
engine::DeepstackData {
image_token_positions: all_positions,
chunks,
},
positions_flat,
))
}
async fn chat_completions_stream(
state: AppState,
req: ChatCompletionRequest,
prepared: PreparedChatContext,
) -> Response {
use super::sse::{generation_events_to_sse, SseStreamOptions};
let engine: Engine = prepared.loaded_engine.engine.clone();
let summarized_messages = prepared.summarized_messages;
let summary_tokens = prepared.summary_tokens;
let (events_tx, events_rx) = tokio::sync::mpsc::channel(64);
let cancellation_counter = Some(state.metrics.sse_cancellations_counter_arc());
if let Err(e) = engine.try_admit_budget(
prepared.prompt_tokens.len() as u32,
prepared.params.max_tokens as u32,
) {
match e {
engine::EngineAdmitError::SlotBudgetExceeded {
needed_bytes,
budget_bytes,
} => {
tracing::info!(
needed_bytes,
budget_bytes,
"chat_completions_stream: ADR-040 §3.5 A5b pre-stream slot_budget_exceeded"
);
return ApiError::slot_budget_exceeded(needed_bytes, budget_bytes).into_response();
}
}
}
if let Err(e) = engine
.generate_stream_with_deepstack(
prepared.prompt_tokens,
prepared.params,
events_tx,
cancellation_counter,
prepared.soft_tokens,
prepared.deepstack_data,
prepared.positions_flat,
)
.await
{
let msg = format!("{e}");
if msg.contains("queue_full") {
state
.metrics
.chat_completions_queue_full
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
return queue_full_with_rate_limit_headers(&state);
}
if msg.contains("slot_budget_exceeded") {
let (needed, budget) = parse_slot_budget_exceeded(&msg);
return ApiError::slot_budget_exceeded(needed, budget).into_response();
}
tracing::error!(error = %msg, "chat_completions_stream enqueue failed");
return ApiError::generation_error(msg).into_response();
}
let opts = SseStreamOptions {
include_usage: req
.stream_options
.as_ref()
.and_then(|s| s.include_usage)
.unwrap_or(false),
logprobs: req.logprobs.unwrap_or(false),
system_fingerprint: state.config.system_fingerprint.clone(),
};
let request_id = format!("chatcmpl-{}", uuid::Uuid::new_v4());
let created = chrono_seconds();
let sse = generation_events_to_sse(events_rx, request_id, req.model.clone(), created, opts);
let mut response = sse.into_response();
apply_transparency_headers(
&state,
&req,
&mut response,
summarized_messages,
summary_tokens,
);
response
}
#[derive(Debug, PartialEq)]
pub(crate) enum TruncateOutcome<E> {
Fits(Vec<super::schema::ChatMessage>, usize),
Truncated(Vec<super::schema::ChatMessage>, usize, usize),
CannotShrink { ctx_len: usize, actual: usize },
TokenizeErr(E),
}
pub(crate) fn truncate_left<F, E>(
messages: &[super::schema::ChatMessage],
ctx_len: usize,
mut tokenize: F,
) -> TruncateOutcome<E>
where
F: FnMut(&[super::schema::ChatMessage]) -> Result<usize, E>,
{
let initial_n = match tokenize(messages) {
Ok(n) => n,
Err(e) => return TruncateOutcome::TokenizeErr(e),
};
if initial_n < ctx_len {
return TruncateOutcome::Fits(messages.to_vec(), initial_n);
}
let mut msgs = messages.to_vec();
let mut iterations = 0usize;
loop {
let last_user_idx = msgs.iter().rposition(|m| m.role == "user");
let drop_idx = msgs
.iter()
.position(|m| m.role != "system")
.filter(|idx| Some(*idx) != last_user_idx);
let Some(drop_idx) = drop_idx else {
let final_n = tokenize(&msgs).unwrap_or(usize::MAX);
return TruncateOutcome::CannotShrink {
ctx_len,
actual: final_n,
};
};
if Some(drop_idx) == last_user_idx {
return TruncateOutcome::CannotShrink {
ctx_len,
actual: initial_n,
};
}
msgs.remove(drop_idx);
let n = match tokenize(&msgs) {
Ok(n) => n,
Err(e) => return TruncateOutcome::TokenizeErr(e),
};
iterations += 1;
if n < ctx_len {
return TruncateOutcome::Truncated(msgs, n, iterations);
}
}
}
const SUMMARIZE_KEEP_RECENT_MSGS: usize = 4;
const SUMMARIZE_MAX_TOKENS: usize = 160;
pub(crate) fn split_for_summarize(
messages: &[super::schema::ChatMessage],
keep_recent_count: usize,
) -> SummarySplit {
let prefix_end = messages
.iter()
.position(|m| m.role != "system")
.unwrap_or(messages.len());
let system_prefix: Vec<_> = messages[..prefix_end].to_vec();
let non_system: &[super::schema::ChatMessage] = &messages[prefix_end..];
let n_tail = non_system.len();
let recent_count = keep_recent_count.min(n_tail);
let summary_end = n_tail - recent_count;
SummarySplit {
system_prefix,
summary_window: non_system[..summary_end].to_vec(),
recent_window: non_system[summary_end..].to_vec(),
}
}
#[derive(Debug, Clone, PartialEq)]
pub(crate) struct SummarySplit {
pub system_prefix: Vec<super::schema::ChatMessage>,
pub summary_window: Vec<super::schema::ChatMessage>,
pub recent_window: Vec<super::schema::ChatMessage>,
}
pub(crate) fn build_summary_user_text(summary_window: &[super::schema::ChatMessage]) -> String {
use super::schema::MessageContent;
let mut buf = String::with_capacity(summary_window.len() * 80);
buf.push_str(
"Summarize the following conversation in 2-3 sentences. \
Be concise; preserve key facts and decisions:\n\n",
);
for m in summary_window {
let text = match &m.content {
None => "".to_string(),
Some(MessageContent::Text(s)) => s.clone(),
Some(MessageContent::Parts(parts)) => parts
.iter()
.filter_map(|p| match p {
super::schema::ContentPart::Text { text } => Some(text.as_str()),
_ => None,
})
.collect::<Vec<_>>()
.join(" "),
};
if text.trim().is_empty() {
continue;
}
buf.push_str(&m.role.to_uppercase());
buf.push_str(": ");
buf.push_str(&text);
buf.push('\n');
}
buf
}
pub(crate) fn build_synthetic_summary_message(summary_text: &str) -> super::schema::ChatMessage {
super::schema::ChatMessage {
role: "system".to_string(),
content: Some(super::schema::MessageContent::Text(format!(
"[Summary of prior conversation]: {}",
summary_text.trim()
))),
reasoning_content: None,
tool_calls: None,
tool_call_id: None,
name: None,
}
}
async fn apply_overflow_policy(
engine: &engine::Engine,
messages: &[super::schema::ChatMessage],
policy: super::schema::OverflowPolicy,
tools: Option<&[super::schema::Tool]>,
enable_thinking: bool,
template_kwargs: Option<&serde_json::Map<String, serde_json::Value>>,
) -> std::result::Result<(Vec<u32>, usize, Option<usize>, Option<usize>), Response> {
use super::schema::OverflowPolicy;
let (tokens, n) = render_and_tokenize_for_overflow(
engine,
messages,
tools,
enable_thinking,
template_kwargs,
)?;
let ctx_len = match engine.context_length() {
Some(c) => c,
None => return Ok((tokens, n, None, None)), };
if n < ctx_len {
return Ok((tokens, n, None, None));
}
match policy {
OverflowPolicy::Reject => {
Err(ApiError::context_length_exceeded(ctx_len, n).into_response())
}
OverflowPolicy::TruncateLeft => {
let (t, n) = apply_truncate_left(
engine,
messages,
ctx_len,
tokens,
n,
tools,
enable_thinking,
template_kwargs,
)?;
Ok((t, n, None, None))
}
OverflowPolicy::Summarize => {
apply_summarize(
engine,
messages,
ctx_len,
tools,
enable_thinking,
template_kwargs,
)
.await
}
}
}
fn render_and_tokenize_for_overflow(
engine: &engine::Engine,
msgs: &[super::schema::ChatMessage],
tools: Option<&[super::schema::Tool]>,
enable_thinking: bool,
template_kwargs: Option<&serde_json::Map<String, serde_json::Value>>,
) -> Result<(Vec<u32>, usize), Response> {
let rendered = render_chat_prompt_or_400(
engine.chat_template(),
msgs,
tools,
enable_thinking,
template_kwargs,
)?;
let encoding = match engine.tokenizer().encode(rendered.as_str(), false) {
Ok(e) => e,
Err(e) => {
tracing::warn!(error = %e, "tokenization failed");
return Err(ApiError::internal_error().into_response());
}
};
let tokens: Vec<u32> = encoding.get_ids().to_vec();
let n = tokens.len();
Ok((tokens, n))
}
fn render_chat_prompt_or_400(
template: &str,
msgs: &[super::schema::ChatMessage],
tools: Option<&[super::schema::Tool]>,
enable_thinking: bool,
template_kwargs: Option<&serde_json::Map<String, serde_json::Value>>,
) -> Result<String, Response> {
engine::render_chat_prompt_with_tools(template, msgs, tools, enable_thinking, template_kwargs)
.map_err(|e| {
tracing::warn!(error = %e, "chat template render failed");
ApiError::invalid_request(
format!("chat template render failed: {e:#}"),
Some("messages".into()),
)
.into_response()
})
}
fn apply_truncate_left(
engine: &engine::Engine,
messages: &[super::schema::ChatMessage],
ctx_len: usize,
initial_tokens: Vec<u32>,
initial_n: usize,
tools: Option<&[super::schema::Tool]>,
enable_thinking: bool,
template_kwargs: Option<&serde_json::Map<String, serde_json::Value>>,
) -> Result<(Vec<u32>, usize), Response> {
let outcome = truncate_left(messages, ctx_len, |msgs| {
render_and_tokenize_for_overflow(engine, msgs, tools, enable_thinking, template_kwargs)
.map(|(_, n)| n)
.map_err(|r| r)
});
match outcome {
TruncateOutcome::Fits(_, _) => Ok((initial_tokens, initial_n)),
TruncateOutcome::Truncated(shrunk, _, _) => render_and_tokenize_for_overflow(
engine,
&shrunk,
tools,
enable_thinking,
template_kwargs,
),
TruncateOutcome::CannotShrink { ctx_len, actual } => {
Err(ApiError::context_length_exceeded(ctx_len, actual).into_response())
}
TruncateOutcome::TokenizeErr(resp) => Err(resp),
}
}
async fn apply_summarize(
engine: &engine::Engine,
messages: &[super::schema::ChatMessage],
ctx_len: usize,
tools: Option<&[super::schema::Tool]>,
enable_thinking: bool,
template_kwargs: Option<&serde_json::Map<String, serde_json::Value>>,
) -> Result<(Vec<u32>, usize, Option<usize>, Option<usize>), Response> {
let split = split_for_summarize(messages, SUMMARIZE_KEEP_RECENT_MSGS);
if split.summary_window.is_empty() {
let (initial_tokens, initial_n) = render_and_tokenize_for_overflow(
engine,
messages,
tools,
enable_thinking,
template_kwargs,
)?;
let (t, n) = apply_truncate_left(
engine,
messages,
ctx_len,
initial_tokens,
initial_n,
tools,
enable_thinking,
template_kwargs,
)?;
return Ok((t, n, None, None));
}
let summary_user_text = build_summary_user_text(&split.summary_window);
let summary_request_msgs = vec![
super::schema::ChatMessage {
role: "system".to_string(),
content: Some(super::schema::MessageContent::Text(
"You are a concise summarization assistant. Produce 2-3 \
sentence summaries of conversations. Preserve key facts \
and decisions; drop pleasantries."
.to_string(),
)),
reasoning_content: None,
tool_calls: None,
tool_call_id: None,
name: None,
},
super::schema::ChatMessage {
role: "user".to_string(),
content: Some(super::schema::MessageContent::Text(summary_user_text)),
reasoning_content: None,
tool_calls: None,
tool_call_id: None,
name: None,
},
];
let (summary_prompt_tokens, summary_prompt_n) =
render_and_tokenize_for_overflow(engine, &summary_request_msgs, None, false, None)?;
if summary_prompt_n + SUMMARIZE_MAX_TOKENS >= ctx_len {
let (initial_tokens, initial_n) = render_and_tokenize_for_overflow(
engine,
messages,
tools,
enable_thinking,
template_kwargs,
)?;
let (t, n) = apply_truncate_left(
engine,
messages,
ctx_len,
initial_tokens,
initial_n,
tools,
enable_thinking,
template_kwargs,
)?;
return Ok((t, n, None, None));
}
let params = super::engine::SamplingParams {
temperature: 0.0,
max_tokens: SUMMARIZE_MAX_TOKENS,
..super::engine::SamplingParams::default()
};
let summary_text = match engine.generate(summary_prompt_tokens, params).await {
Ok(result) => {
let text = result.text.trim().to_string();
if text.is_empty() {
tracing::warn!("summarize produced empty text; falling back to truncate_left");
let (initial_tokens, initial_n) = render_and_tokenize_for_overflow(
engine,
messages,
tools,
enable_thinking,
template_kwargs,
)?;
let (t, n) = apply_truncate_left(
engine,
messages,
ctx_len,
initial_tokens,
initial_n,
tools,
enable_thinking,
template_kwargs,
)?;
return Ok((t, n, None, None));
}
(text, result.completion_tokens)
}
Err(e) => {
tracing::warn!(error = %e, "summarize engine.generate failed");
return Err(
ApiError::generation_error(format!("summarize forward pass failed: {e}"))
.into_response(),
);
}
};
let (summary_text, summary_completion_tokens) = summary_text;
let synthetic = build_synthetic_summary_message(&summary_text);
let mut new_messages = split.system_prefix.clone();
new_messages.push(synthetic);
new_messages.extend(split.recent_window.iter().cloned());
let summarized_count = split.summary_window.len();
let (new_tokens, new_n) = render_and_tokenize_for_overflow(
engine,
&new_messages,
tools,
enable_thinking,
template_kwargs,
)?;
if new_n < ctx_len {
return Ok((
new_tokens,
new_n,
Some(summarized_count),
Some(summary_completion_tokens),
));
}
let outcome = truncate_left(&new_messages, ctx_len, |msgs| {
render_and_tokenize_for_overflow(engine, msgs, tools, enable_thinking, template_kwargs)
.map(|(_, n)| n)
.map_err(|r| r)
});
match outcome {
TruncateOutcome::Fits(_, _) => Ok((
new_tokens,
new_n,
Some(summarized_count),
Some(summary_completion_tokens),
)),
TruncateOutcome::Truncated(shrunk, _, _) => {
let (t, n) = render_and_tokenize_for_overflow(
engine,
&shrunk,
tools,
enable_thinking,
template_kwargs,
)?;
Ok((
t,
n,
Some(summarized_count),
Some(summary_completion_tokens),
))
}
TruncateOutcome::CannotShrink { ctx_len, actual } => {
Err(ApiError::context_length_exceeded(ctx_len, actual).into_response())
}
TruncateOutcome::TokenizeErr(resp) => Err(resp),
}
}
#[cfg(test)]
mod truncate_tests {
use super::super::schema::{ChatMessage, MessageContent};
use super::*;
fn msg(role: &str, content: &str) -> ChatMessage {
ChatMessage {
role: role.into(),
content: Some(MessageContent::Text(content.into())),
reasoning_content: None,
tool_calls: None,
tool_call_id: None,
name: None,
}
}
fn toks_per_msg(msgs: &[ChatMessage], per_msg: usize) -> Result<usize, ()> {
Ok(msgs.len() * per_msg)
}
#[test]
fn truncate_left_fits_returns_initial_count() {
let msgs = vec![msg("system", "s"), msg("user", "u")];
let out = truncate_left::<_, ()>(&msgs, 100, |m| toks_per_msg(m, 10));
match out {
TruncateOutcome::Fits(m, n) => {
assert_eq!(m.len(), 2);
assert_eq!(n, 20);
}
other => panic!("expected Fits, got {:?}", other),
}
}
#[test]
fn truncate_left_drops_oldest_nonsystem_until_fits() {
let msgs = vec![
msg("system", "s"),
msg("user", "u1"),
msg("assistant", "a1"),
msg("user", "u2"),
msg("assistant", "a2"),
msg("user", "final"),
];
let out = truncate_left::<_, ()>(&msgs, 25, |m| toks_per_msg(m, 10));
match out {
TruncateOutcome::Truncated(retained, n, iters) => {
assert_eq!(retained.len(), 2);
assert_eq!(retained[0].role, "system");
assert_eq!(retained[1].role, "user");
assert_eq!(
retained[1].content.as_ref().map(|c| c.text()).unwrap(),
"final"
);
assert_eq!(n, 20);
assert_eq!(iters, 4); }
other => panic!("expected Truncated, got {:?}", other),
}
}
#[test]
fn truncate_left_cannot_shrink_when_system_plus_last_user_overflow() {
let msgs = vec![msg("system", "s"), msg("user", "u")];
let out = truncate_left::<_, ()>(&msgs, 15, |m| toks_per_msg(m, 10));
assert!(
matches!(out, TruncateOutcome::CannotShrink { .. }),
"got {:?}",
out
);
}
#[test]
fn truncate_left_no_system_keeps_last_user_only() {
let msgs = vec![
msg("user", "u1"),
msg("assistant", "a1"),
msg("user", "final"),
];
let out = truncate_left::<_, ()>(&msgs, 15, |m| toks_per_msg(m, 10));
match out {
TruncateOutcome::Truncated(retained, n, _) => {
assert_eq!(retained.len(), 1);
assert_eq!(retained[0].role, "user");
assert_eq!(
retained[0].content.as_ref().map(|c| c.text()).unwrap(),
"final"
);
assert_eq!(n, 10);
}
other => panic!("expected Truncated, got {:?}", other),
}
}
#[test]
fn truncate_left_multiple_system_messages_all_preserved() {
let msgs = vec![
msg("system", "s1"),
msg("system", "s2"),
msg("user", "u1"),
msg("assistant", "a1"),
msg("user", "final"),
];
let out = truncate_left::<_, ()>(&msgs, 35, |m| toks_per_msg(m, 10));
match out {
TruncateOutcome::Truncated(retained, _, _) => {
assert!(retained.iter().filter(|m| m.role == "system").count() == 2);
assert_eq!(retained.last().unwrap().role, "user");
}
other => panic!("expected Truncated, got {:?}", other),
}
}
#[test]
fn truncate_left_propagates_tokenizer_error() {
let msgs = vec![msg("user", "u")];
let out = truncate_left::<_, &'static str>(&msgs, 100, |_| Err("fake tokenize error"));
assert_eq!(out, TruncateOutcome::TokenizeErr("fake tokenize error"));
}
#[test]
fn truncate_left_is_deterministic() {
let msgs = vec![
msg("system", "s"),
msg("user", "u1"),
msg("assistant", "a1"),
msg("user", "final"),
];
let a = truncate_left::<_, ()>(&msgs, 25, |m| toks_per_msg(m, 10));
let b = truncate_left::<_, ()>(&msgs, 25, |m| toks_per_msg(m, 10));
assert_eq!(format!("{:?}", a), format!("{:?}", b));
}
#[test]
fn split_for_summarize_no_messages_yields_empty_split() {
let s = split_for_summarize(&[], 4);
assert!(s.system_prefix.is_empty());
assert!(s.summary_window.is_empty());
assert!(s.recent_window.is_empty());
}
#[test]
fn split_for_summarize_only_system_prefix() {
let msgs = vec![msg("system", "s1"), msg("system", "s2")];
let s = split_for_summarize(&msgs, 4);
assert_eq!(s.system_prefix.len(), 2);
assert!(s.summary_window.is_empty());
assert!(s.recent_window.is_empty());
}
#[test]
fn split_for_summarize_keep_count_exceeds_tail_means_no_summary() {
let msgs = vec![msg("system", "s"), msg("user", "u"), msg("assistant", "a")];
let s = split_for_summarize(&msgs, 4);
assert_eq!(s.system_prefix.len(), 1);
assert!(s.summary_window.is_empty());
assert_eq!(s.recent_window.len(), 2);
}
#[test]
fn split_for_summarize_5_messages_keep_2_recent() {
let msgs = vec![
msg("system", "s"),
msg("user", "u1"),
msg("assistant", "a1"),
msg("user", "u2"),
msg("assistant", "a2"),
msg("user", "u3"),
];
let s = split_for_summarize(&msgs, 2);
assert_eq!(s.system_prefix.len(), 1);
assert_eq!(s.summary_window.len(), 3);
assert_eq!(
s.summary_window[0].content,
Some(MessageContent::Text("u1".into()))
);
assert_eq!(
s.summary_window[2].content,
Some(MessageContent::Text("u2".into()))
);
assert_eq!(s.recent_window.len(), 2);
assert_eq!(
s.recent_window[1].content,
Some(MessageContent::Text("u3".into()))
);
}
#[test]
fn split_for_summarize_no_system_prefix() {
let msgs = vec![msg("user", "u1"), msg("assistant", "a1"), msg("user", "u2")];
let s = split_for_summarize(&msgs, 2);
assert!(s.system_prefix.is_empty());
assert_eq!(s.summary_window.len(), 1);
assert_eq!(
s.summary_window[0].content,
Some(MessageContent::Text("u1".into()))
);
assert_eq!(s.recent_window.len(), 2);
}
#[test]
fn build_summary_user_text_skips_empty_content() {
let window = vec![
msg("user", "hello"),
msg("assistant", ""),
msg("user", "world"),
];
let out = build_summary_user_text(&window);
assert!(out.contains("USER: hello"));
assert!(out.contains("USER: world"));
assert!(!out.contains("ASSISTANT:"));
}
#[test]
fn build_summary_user_text_uppercases_role() {
let window = vec![msg("user", "x"), msg("tool", "y")];
let out = build_summary_user_text(&window);
assert!(out.contains("USER: x"));
assert!(out.contains("TOOL: y"));
}
#[test]
fn build_summary_user_text_handles_multipart_content() {
use super::super::schema::{ContentPart, ImageUrl};
let m = ChatMessage {
role: "user".into(),
content: Some(MessageContent::Parts(vec![
ContentPart::Text {
text: "see this:".into(),
},
ContentPart::ImageUrl {
image_url: ImageUrl {
url: "data:img...".into(),
detail: None,
},
},
ContentPart::Text {
text: "and this".into(),
},
])),
reasoning_content: None,
tool_calls: None,
tool_call_id: None,
name: None,
};
let out = build_summary_user_text(&[m]);
assert!(out.contains("USER: see this: and this"));
assert!(!out.contains("data:img"));
}
#[test]
fn build_synthetic_summary_message_wraps_with_marker() {
let m = build_synthetic_summary_message("user discussed deployment");
assert_eq!(m.role, "system");
let content = match m.content {
Some(MessageContent::Text(s)) => s,
_ => panic!("expected text content"),
};
assert!(content.starts_with("[Summary of prior conversation]:"));
assert!(content.contains("user discussed deployment"));
}
#[test]
fn build_synthetic_summary_message_trims_whitespace() {
let m = build_synthetic_summary_message(" trimmed text \n");
let content = match m.content {
Some(MessageContent::Text(s)) => s,
_ => unreachable!(),
};
assert!(content.ends_with("trimmed text"));
assert!(!content.ends_with("trimmed text "));
}
}
fn process_multimodal_content(
messages: &[super::schema::ChatMessage],
mmproj: Option<&super::state::LoadedMmproj>,
) -> std::result::Result<Vec<crate::inference::vision::vit_gpu::VisionInput>, Response> {
use super::schema::{ContentPart, MessageContent};
use crate::inference::vision::mmproj::ArchProfile;
use crate::inference::vision::vit_gpu::{Gemma4vPreprocessedImage, VisionInput};
let mut image_refs: Vec<(usize, usize, &super::schema::ImageUrl)> = Vec::new();
for (mi, msg) in messages.iter().enumerate() {
if let Some(MessageContent::Parts(parts)) = msg.content.as_ref() {
for (pi, p) in parts.iter().enumerate() {
if let ContentPart::ImageUrl { image_url } = p {
image_refs.push((mi, pi, image_url));
}
}
}
}
if image_refs.is_empty() {
return Ok(Vec::new());
}
let mmproj = match mmproj {
Some(m) => m,
None => return Err(ApiError::no_mmproj_loaded().into_response()),
};
if !mmproj.arch.is_supported() {
return Err(ApiError::invalid_request(
format!(
"mmproj arch profile is '{}' — supported profiles: 'gemma4_siglip', \
'clip_classic'. Reload server with a compatible mmproj GGUF.",
mmproj.arch.as_str()
),
Some("messages".into()),
)
.into_response());
}
let mut out: Vec<VisionInput> = Vec::with_capacity(image_refs.len());
for (mi, pi, image_url) in image_refs {
let parsed = crate::inference::vision::parse_image_url(&image_url.url).map_err(|e| {
ApiError::invalid_request(
format!(
"messages[{}].content[{}].image_url parse failed: {}",
mi, pi, e
),
Some(format!("messages[{}].content[{}]", mi, pi)),
)
.into_response()
})?;
let source_label = match &parsed {
crate::inference::vision::ImageInput::DataUri { mime_type, .. } => mime_type.clone(),
crate::inference::vision::ImageInput::FilePath(p) => p
.file_name()
.and_then(|s| s.to_str())
.unwrap_or("file")
.to_string(),
crate::inference::vision::ImageInput::HttpUrl(u) => u.clone(),
};
let bytes = crate::inference::vision::load_image_bytes(&parsed).map_err(|e| {
ApiError::invalid_request(
format!(
"messages[{}].content[{}].image_url load failed: {}",
mi, pi, e
),
Some(format!("messages[{}].content[{}]", mi, pi)),
)
.into_response()
})?;
match mmproj.arch {
ArchProfile::ClipClassic => {
let preprocess_cfg = mmproj.config.preprocess_config();
let pixel_values =
crate::inference::vision::preprocess_rgb_chw(&bytes, &preprocess_cfg).map_err(
|e| {
ApiError::invalid_request(
format!(
"messages[{}].content[{}].image_url preprocess failed: {}",
mi, pi, e
),
Some(format!("messages[{}].content[{}]", mi, pi)),
)
.into_response()
},
)?;
out.push(VisionInput::Siglip49(
crate::inference::vision::PreprocessedImage {
pixel_values,
target_size: preprocess_cfg.target_size,
pixel_w: None,
pixel_h: None,
source_label,
},
));
}
ArchProfile::Gemma4Siglip => {
let cfg = &crate::inference::vision::preprocess::GEMMA4V_PREPROCESS_DEFAULT;
let preprocessed = crate::inference::vision::preprocess::preprocess_gemma4v(
&bytes, cfg,
)
.map_err(|e| {
ApiError::invalid_request(
format!(
"messages[{}].content[{}].image_url gemma4v preprocess failed: {}",
mi, pi, e
),
Some(format!("messages[{}].content[{}]", mi, pi)),
)
.into_response()
})?;
out.push(VisionInput::Gemma4v(Gemma4vPreprocessedImage {
patches: preprocessed.patches,
pos_x: preprocessed.pos_x,
pos_y: preprocessed.pos_y,
n_x: preprocessed.n_x,
n_y: preprocessed.n_y,
source_label,
}));
}
ArchProfile::Qwen3VlSiglip => {
let qcfg =
crate::inference::vision::preprocess::Qwen3VlPreprocessConfig::from_mmproj(
&mmproj.config,
)
.map_err(|e| {
ApiError::generation_error(format!(
"messages[{}].content[{}].image_url qwen3vl preprocess \
config: {e}",
mi, pi
))
.into_response()
})?;
let preprocessed = crate::inference::vision::preprocess::preprocess_qwen3vl(
&bytes,
&qcfg,
mmproj.config.image_size,
)
.map_err(|e| {
ApiError::invalid_request(
format!(
"messages[{}].content[{}].image_url qwen3vl preprocess \
failed: {e}",
mi, pi
),
Some(format!("messages[{}].content[{}]", mi, pi)),
)
.into_response()
})?;
let (target_w, target_h) = preprocessed.target_pixel_grid();
out.push(VisionInput::Siglip49(
crate::inference::vision::PreprocessedImage {
pixel_values: preprocessed.pixel_values,
target_size: preprocessed.target_size,
pixel_w: Some(target_w),
pixel_h: Some(target_h),
source_label,
},
));
}
ArchProfile::Unknown => {
unreachable!(
"process_multimodal_content: ArchProfile::Unknown should have been \
rejected by `is_supported` check above"
);
}
}
}
Ok(out)
}
fn apply_vit_transparency_headers(
resp: &mut Response,
forward_ms: Option<u64>,
n_images: Option<usize>,
soft_tokens_total: Option<usize>,
) {
use axum::http::{header::HeaderName, HeaderValue};
let headers = resp.headers_mut();
if let Some(ms) = forward_ms {
if let Ok(v) = HeaderValue::from_str(&ms.to_string()) {
headers.insert(HeaderName::from_static("x-hf2q-vit-forward-ms"), v);
}
}
if let Some(n) = n_images {
if let Ok(v) = HeaderValue::from_str(&n.to_string()) {
headers.insert(HeaderName::from_static("x-hf2q-vit-images"), v);
}
}
if let Some(n) = soft_tokens_total {
if let Ok(v) = HeaderValue::from_str(&n.to_string()) {
headers.insert(HeaderName::from_static("x-hf2q-soft-tokens-total"), v);
}
}
}
fn rewrite_messages_for_vision_placeholders(messages: &[ChatMessage]) -> Vec<ChatMessage> {
rewrite_messages_for_vision_placeholders_family(
messages,
crate::inference::vision::mmproj::VisionFamily::Gemma,
)
}
fn rewrite_messages_for_vision_placeholders_family(
messages: &[ChatMessage],
family: crate::inference::vision::mmproj::VisionFamily,
) -> Vec<ChatMessage> {
use super::schema::ContentPart;
use crate::inference::vision::mmproj::VisionFamily;
let placeholder_marker: String = match family {
VisionFamily::Gemma => "<|image|>".to_string(),
VisionFamily::Qwen3Vl => "<|vision_start|><|image_pad|><|vision_end|>".to_string(),
VisionFamily::Unknown => "".to_string(),
};
messages
.iter()
.map(|msg| {
let new_content = msg.content.as_ref().map(|c| match c {
MessageContent::Text(_) => c.clone(),
MessageContent::Parts(parts) => {
let any_image = parts
.iter()
.any(|p| matches!(p, ContentPart::ImageUrl { .. }));
if !any_image {
c.clone()
} else {
let mut buf = String::new();
for part in parts {
match part {
ContentPart::Text { text } => buf.push_str(text),
ContentPart::ImageUrl { .. } => buf.push_str(&placeholder_marker),
}
}
MessageContent::Text(buf)
}
}
});
ChatMessage {
role: msg.role.clone(),
content: new_content,
reasoning_content: msg.reasoning_content.clone(),
tool_calls: msg.tool_calls.clone(),
tool_call_id: msg.tool_call_id.clone(),
name: msg.name.clone(),
}
})
.collect()
}
#[derive(Debug, PartialEq)]
pub(crate) struct PlaceholderCountMismatch {
pub placeholder_positions_found: usize,
pub n_image_tokens_supplied: usize,
}
pub(crate) fn compute_soft_token_layout(
img_token_id: u32,
prompt_tokens: &[u32],
n_image_tokens_per_image: &[usize],
) -> std::result::Result<(Vec<u32>, Vec<std::ops::Range<usize>>), PlaceholderCountMismatch> {
let placeholder_positions: Vec<usize> = prompt_tokens
.iter()
.enumerate()
.filter_map(|(p, t)| if *t == img_token_id { Some(p) } else { None })
.collect();
if placeholder_positions.len() != n_image_tokens_per_image.len() {
return Err(PlaceholderCountMismatch {
placeholder_positions_found: placeholder_positions.len(),
n_image_tokens_supplied: n_image_tokens_per_image.len(),
});
}
let total_extra: usize = n_image_tokens_per_image
.iter()
.copied()
.sum::<usize>()
.saturating_sub(placeholder_positions.len()); let mut prompt_expanded: Vec<u32> = Vec::with_capacity(prompt_tokens.len() + total_extra);
let mut ranges: Vec<std::ops::Range<usize>> = Vec::with_capacity(placeholder_positions.len());
let mut last_pos = 0usize;
for (i, &pos) in placeholder_positions.iter().enumerate() {
prompt_expanded.extend_from_slice(&prompt_tokens[last_pos..pos]);
let n = n_image_tokens_per_image[i];
let start = prompt_expanded.len();
for _ in 0..n {
prompt_expanded.push(img_token_id);
}
let end = prompt_expanded.len();
ranges.push(start..end);
last_pos = pos + 1;
}
prompt_expanded.extend_from_slice(&prompt_tokens[last_pos..]);
Ok((prompt_expanded, ranges))
}
fn expand_image_placeholders(
engine: &engine::Engine,
prompt_tokens: &[u32],
embeddings: &[Vec<f32>],
) -> std::result::Result<(Vec<u32>, Vec<engine::SoftTokenData>), Response> {
expand_image_placeholders_family(
engine,
prompt_tokens,
embeddings,
crate::inference::vision::mmproj::VisionFamily::Gemma,
engine.hidden_size(),
)
.map(|(toks, soft, _positions)| (toks, soft))
}
fn expand_image_placeholders_family(
engine: &engine::Engine,
prompt_tokens: &[u32],
embeddings: &[Vec<f32>],
family: crate::inference::vision::mmproj::VisionFamily,
per_row_floats: usize,
) -> std::result::Result<(Vec<u32>, Vec<engine::SoftTokenData>, Vec<Vec<u32>>), Response> {
crate::inference::vision::pipeline::expand_image_placeholders(
engine.tokenizer(),
prompt_tokens,
embeddings,
family,
per_row_floats,
engine.hidden_size(),
)
.map_err(|e| ApiError::generation_error(format!("{e:#}")).into_response())
}
fn compile_response_format(
rf: &ResponseFormat,
) -> std::result::Result<Option<grammar::Grammar>, Response> {
let gbnf = match rf {
ResponseFormat::Text => return Ok(None),
ResponseFormat::JsonObject => {
static JSON_OBJECT_GRAMMAR: &str = r#"root ::= object
value ::= object | array | string | number | ("true" | "false" | "null") ws
object ::=
"{" ws (
string ":" ws value
("," ws string ":" ws value)*
)? "}" ws
array ::=
"[" ws (
value
("," ws value)*
)? "]" ws
string ::=
"\"" (
[^"\\\x7F\x00-\x1F] |
"\\" (["\\bfnrt] | "u" [0-9a-fA-F]{4})
)* "\"" ws
number ::= ("-"? ([0-9] | [1-9] [0-9]{0,15})) ("." [0-9]+)? ([eE] [-+]? [0-9] [1-9]{0,15})? ws
ws ::= | " " | "\n" [ \t]{0,20}
"#;
JSON_OBJECT_GRAMMAR.to_string()
}
ResponseFormat::JsonSchema { json_schema } => {
match grammar::json_schema::schema_to_gbnf(&json_schema.schema) {
Ok(g) => g,
Err(e) => {
return Err(ApiError::grammar_error(format!(
"json_schema → GBNF failed: {}",
e
))
.into_response())
}
}
}
};
match grammar::parser::parse(&gbnf) {
Ok(g) => Ok(Some(g)),
Err(e) => Err(ApiError::grammar_error(format!("GBNF parse failed: {}", e)).into_response()),
}
}
fn select_effective_grammar(
tool_grammar: Option<grammar::Grammar>,
tool_grammar_kind: engine::GrammarKind,
response_grammar: Option<grammar::Grammar>,
) -> (Option<grammar::Grammar>, engine::GrammarKind) {
match (tool_grammar, response_grammar) {
(Some(g), _) => (Some(g), tool_grammar_kind),
(None, Some(g)) => (Some(g), engine::GrammarKind::ResponseFormat),
(None, None) => (None, engine::GrammarKind::default()),
}
}
fn tool_grammar_kind_for(tool_choice: &super::schema::ToolChoiceValue) -> engine::GrammarKind {
use super::schema::ToolChoiceValue;
match tool_choice {
ToolChoiceValue::Auto => engine::GrammarKind::ToolCallBodyAuto,
ToolChoiceValue::Required | ToolChoiceValue::Function(_) => {
engine::GrammarKind::ToolCallBodyRequired
}
ToolChoiceValue::None => engine::GrammarKind::default(),
}
}
fn defensive_no_call_under_constrained(
tool_call_policy: engine::ToolCallPolicy,
tool_calls_is_empty: bool,
finish_reason: &str,
completion_tokens: usize,
text_len: usize,
) -> Option<Response> {
if matches!(tool_call_policy, engine::ToolCallPolicy::Constrained) && tool_calls_is_empty {
tracing::error!(
finish_reason = %finish_reason,
completion_tokens = completion_tokens,
text_len = text_len,
"tool_call_no_call_under_constrained: eager grammar should have prevented \
this — either max_tokens cut a call mid-emission or the grammar emitter \
has a bug"
);
return Some(
ApiError::generation_error("tool_call_no_call_under_constrained".to_string())
.into_response(),
);
}
None
}
fn compile_tool_grammar(
req: &super::schema::ChatCompletionRequest,
tool_choice: &super::schema::ToolChoiceValue,
) -> std::result::Result<Option<grammar::Grammar>, Response> {
use super::schema::ToolChoiceValue;
if matches!(tool_choice, ToolChoiceValue::None) {
return Ok(None);
}
let constrain = matches!(
tool_choice,
ToolChoiceValue::Required | ToolChoiceValue::Function(_)
);
let tools = match req.tools.as_deref() {
Some(t) if !t.is_empty() => t,
_ => {
if !constrain {
return Ok(None);
}
let label = match tool_choice {
ToolChoiceValue::Required => "required",
ToolChoiceValue::Function(_) => "function",
_ => unreachable!("constrain gated above"),
};
tracing::error!(
model = %req.model,
tool_choice = label,
"tool_choice_constrained_without_tools: tool_choice={} but \
request has no tools[] defined; rejecting with 400 (silent \
Ok(None) fallback would re-open the wave-2.6 \
'Required/Function enforcement before ToolCallOpen' divergence)",
label
);
return Err(ApiError::invalid_request(
format!(
"tool_choice={} but request has no tools[] defined; \
declare at least one tool or use tool_choice=auto/none",
label
),
Some("tool_choice".into()),
)
.into_response());
}
};
let reg = match registry::find_for(&req.model) {
Some(r) => r,
None => {
if !constrain {
return Ok(None);
}
let label = match tool_choice {
ToolChoiceValue::Required => "required",
ToolChoiceValue::Function(_) => "function",
_ => unreachable!("constrain gated above"),
};
tracing::error!(
model = %req.model,
tool_choice = label,
"tool_choice_constrained_unknown_model: tool_choice={} but \
model '{}' has no registered tool_call_gbnf emitter; rejecting \
with 400 (silent Ok(None) fallback would re-open the wave-2.6 \
'Required/Function enforcement before ToolCallOpen' divergence)",
label,
req.model
);
return Err(ApiError::invalid_request(
format!(
"tool_choice={} but model '{}' has no registered \
tool_call_gbnf emitter; use a registered model \
(e.g. gemma4-27b-it, qwen3.6-27b-dwq46) or \
tool_choice=auto/none",
label, req.model
),
Some("tool_choice".into()),
)
.into_response());
}
};
let matching_tools: Vec<&super::schema::Tool> = match tool_choice {
ToolChoiceValue::Function(name) => {
let found: Vec<_> = tools.iter().filter(|t| t.function.name == *name).collect();
if found.is_empty() {
return Err(ApiError::invalid_request(
format!(
"tool_choice.function.name '{}' not found in tools list",
name
),
Some("tool_choice".into()),
)
.into_response());
}
found
}
ToolChoiceValue::Required | ToolChoiceValue::Auto => tools.iter().collect(),
_ => unreachable!("None gated above"),
};
let parallel = req.parallel_tool_calls.unwrap_or(false);
let auto_lazy = matches!(tool_choice, ToolChoiceValue::Auto);
let deepseek_multi = reg.family == "deepseek4" && matching_tools.len() > 1;
let per_fn_shape = if deepseek_multi {
registry::GrammarShape::SingleBody
} else if matching_tools.len() == 1 {
if auto_lazy {
registry::GrammarShape::OneOrMoreCallsBodyOnly { parallel }
} else {
registry::GrammarShape::OneOrMoreCalls { parallel }
}
} else if auto_lazy {
registry::GrammarShape::OneOrMoreCallsBodyOnly { parallel: false }
} else {
registry::GrammarShape::OneOrMoreCalls { parallel: false }
};
let mut fn_gbnfs: Vec<String> = Vec::new();
for tool in &matching_tools {
let params_schema = tool
.function
.parameters
.as_ref()
.cloned()
.unwrap_or(serde_json::Value::Object(serde_json::Map::new()));
match reg.tool_call_gbnf(&tool.function.name, ¶ms_schema, per_fn_shape) {
Ok(gbnf) => fn_gbnfs.push(gbnf),
Err(e) => {
return Err(ApiError::grammar_error(format!(
"tool_call_gbnf for '{}': {}",
tool.function.name, e
))
.into_response());
}
}
}
let combined_gbnf = if deepseek_multi {
combine_deepseek4_function_grammars(fn_gbnfs, auto_lazy, parallel)
} else if fn_gbnfs.len() == 1 {
fn_gbnfs.into_iter().next().unwrap()
} else {
let separator = if auto_lazy {
reg.auto_lazy_multi_fn_inter_call()
} else {
reg.parallel_call_separator()
};
combine_function_grammars(fn_gbnfs, parallel, separator)
};
match grammar::parser::parse(&combined_gbnf) {
Ok(g) => Ok(Some(g)),
Err(e) => Err(
ApiError::grammar_error(format!("tool call GBNF parse failed: {}", e)).into_response(),
),
}
}
#[cfg(test)]
mod compile_tool_grammar_precondition_tests {
use super::super::schema::{
ChatCompletionRequest, ChatMessage, MessageContent, Tool, ToolChoiceValue, ToolFunction,
};
use super::*;
fn req_with(model: &str, tools: Option<Vec<Tool>>) -> ChatCompletionRequest {
ChatCompletionRequest {
model: model.to_string(),
messages: vec![ChatMessage {
role: "user".to_string(),
content: Some(MessageContent::Text("call a tool".to_string())),
reasoning_content: None,
tool_calls: None,
tool_call_id: None,
name: None,
}],
stream: None,
max_tokens: None,
max_completion_tokens: None,
temperature: None,
stop: None,
tools,
tool_choice: None,
response_format: None,
top_p: None,
seed: None,
frequency_penalty: None,
presence_penalty: None,
stream_options: None,
top_k: None,
repetition_penalty: None,
min_p: None,
logprobs: None,
top_logprobs: None,
logit_bias: None,
parallel_tool_calls: None,
hf2q_overflow_policy: None,
hf2q_enable_thinking: None,
chat_template_kwargs: None,
}
}
fn one_scalar_tool(name: &str) -> Vec<Tool> {
vec![Tool {
tool_type: "function".to_string(),
function: ToolFunction {
name: name.to_string(),
description: None,
parameters: Some(serde_json::json!({
"type": "object",
"properties": {"q": {"type": "string"}},
"required": ["q"]
})),
},
}]
}
#[test]
fn compile_tool_grammar_required_with_no_tools_returns_error() {
let req = req_with("gemma4-27b-it", None);
let res = compile_tool_grammar(&req, &ToolChoiceValue::Required);
let resp = res.expect_err("Required + no tools MUST hard-error (400)");
assert_eq!(
resp.status(),
StatusCode::BAD_REQUEST,
"must be 400 — silent Ok(None) re-opens the wave-2.6 \
'Required/Function enforcement before ToolCallOpen' divergence"
);
}
#[test]
fn compile_tool_grammar_required_with_empty_tools_returns_error() {
let req = req_with("gemma4-27b-it", Some(Vec::new()));
let res = compile_tool_grammar(&req, &ToolChoiceValue::Required);
let resp = res.expect_err("Required + empty tools[] MUST hard-error (400)");
assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
}
#[test]
fn compile_tool_grammar_function_with_no_tools_returns_error() {
let req = req_with("gemma4-27b-it", None);
let res = compile_tool_grammar(&req, &ToolChoiceValue::Function("get_weather".to_string()));
let resp = res.expect_err("Function + no tools MUST hard-error (400)");
assert_eq!(
resp.status(),
StatusCode::BAD_REQUEST,
"must be 400 — silent Ok(None) here would let the engine \
run unconstrained under tool_choice=function (no grammar guard)"
);
}
#[test]
fn compile_tool_grammar_function_with_empty_tools_returns_error() {
let req = req_with("gemma4-27b-it", Some(Vec::new()));
let res = compile_tool_grammar(&req, &ToolChoiceValue::Function("get_weather".to_string()));
let resp = res.expect_err("Function + empty tools[] MUST hard-error (400)");
assert_eq!(
resp.status(),
StatusCode::BAD_REQUEST,
"must be 400 — silent Ok(None) re-opens the wave-2.6 \
'Required/Function enforcement before ToolCallOpen' divergence"
);
}
#[test]
fn compile_tool_grammar_function_with_unknown_family_returns_error() {
assert!(
registry::find_for("unknown-fake-model-zzzz").is_none(),
"test fixture must use an unregistered model id; \
update if registry adds a 'unknown-fake-model-zzzz' family"
);
let req = req_with(
"unknown-fake-model-zzzz",
Some(one_scalar_tool("get_weather")),
);
let res = compile_tool_grammar(&req, &ToolChoiceValue::Function("get_weather".to_string()));
let resp = res.expect_err("Function + unknown model MUST hard-error (400)");
assert_eq!(
resp.status(),
StatusCode::BAD_REQUEST,
"must be 400 — silent Ok(None) here would let the engine \
run unconstrained under tool_choice=function (no grammar guard)"
);
}
#[test]
fn compile_tool_grammar_required_with_unknown_family_returns_error() {
assert!(registry::find_for("unknown-fake-model-zzzz").is_none());
let req = req_with("unknown-fake-model-zzzz", Some(one_scalar_tool("foo")));
let res = compile_tool_grammar(&req, &ToolChoiceValue::Required);
let resp = res.expect_err("Required + unknown model MUST hard-error (400)");
assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
}
#[test]
fn compile_tool_grammar_auto_with_no_tools_returns_ok_none() {
let req = req_with("gemma4-27b-it", None);
let res = compile_tool_grammar(&req, &ToolChoiceValue::Auto);
let g = res.expect("Auto + no tools MUST stay Ok(None)");
assert!(g.is_none(), "Auto policy is allowed to compile no grammar");
}
#[test]
fn compile_tool_grammar_auto_with_unknown_family_returns_ok_none() {
let req = req_with("unknown-fake-model-zzzz", Some(one_scalar_tool("foo")));
let res = compile_tool_grammar(&req, &ToolChoiceValue::Auto);
let g = res.expect("Auto + unknown model MUST stay Ok(None)");
assert!(g.is_none());
}
#[test]
fn compile_tool_grammar_none_choice_returns_ok_none() {
let req = req_with("gemma4-27b-it", Some(one_scalar_tool("foo")));
let res = compile_tool_grammar(&req, &ToolChoiceValue::None);
let g = res.expect("None choice MUST stay Ok(None)");
assert!(g.is_none());
}
#[test]
fn compile_tool_grammar_auto_default_parallel_is_false() {
if registry::find_for("gemma4-27b-it").is_none() {
return;
}
let mut req = req_with("gemma4-27b-it", Some(one_scalar_tool("get_weather")));
req.parallel_tool_calls = None;
let res = compile_tool_grammar(&req, &ToolChoiceValue::Auto);
let g = res
.expect("Auto + tools + registered family MUST compile a grammar")
.expect("expected Some(grammar) under Auto+lazy with registered family");
let serialized = format!("{g:?}");
assert!(
!serialized.contains("gemma4-call"),
"iter-218: with `parallel_tool_calls` unset, the default MUST be \
FALSE (single-call grammar, no `gemma4-call` recursion rule). \
Found `gemma4-call` in serialized grammar — default flip regressed. \
Serialized: {serialized}"
);
}
#[test]
fn compile_tool_grammar_auto_explicit_parallel_true_emits_recursion() {
if registry::find_for("gemma4-27b-it").is_none() {
return;
}
let mut req = req_with("gemma4-27b-it", Some(one_scalar_tool("get_weather")));
req.parallel_tool_calls = Some(true);
let res = compile_tool_grammar(&req, &ToolChoiceValue::Auto);
let g = res
.expect("Auto + tools MUST compile")
.expect("expected Some(grammar) under Auto+lazy with parallel=true");
let serialized = format!("{g:?}");
assert!(
serialized.contains("gemma4-call"),
"explicit parallel_tool_calls=true MUST emit `gemma4-call` recursion \
rule (operator opt-in to multi-call). Serialized: {serialized}"
);
}
#[test]
fn compile_tool_grammar_required_happy_path_compiles() {
if registry::find_for("gemma4-27b-it").is_none() {
return;
}
let req = req_with("gemma4-27b-it", Some(one_scalar_tool("get_weather")));
let res = compile_tool_grammar(&req, &ToolChoiceValue::Required);
let g = res.expect("Required + valid tools + registered model MUST compile");
assert!(
g.is_some(),
"happy path MUST return Some(grammar); got None"
);
}
#[test]
fn compile_tool_grammar_auto_with_tools_compiles_lazy_grammar() {
if registry::find_for("gemma4-27b-it").is_none() {
return;
}
let req = req_with("gemma4-27b-it", Some(one_scalar_tool("get_weather")));
let res = compile_tool_grammar(&req, &ToolChoiceValue::Auto);
let g = res.expect(
"Auto + tools[] + registered family MUST compile a grammar (W-B2 lazy wiring); \
previously short-circuited to Ok(None) before W-B2",
);
let g = g.expect(
"Auto + tools[] + registered family MUST return Some(grammar) — the Auto \
branch now compiles the same per-model body grammar as Required, \
tagged with GrammarKind::ToolCallBodyAuto by the caller",
);
assert!(
g.rule_id("root").is_some(),
"Auto-mode compiled grammar MUST have a root rule (combiner output)"
);
}
#[test]
fn compile_tool_grammar_auto_strips_first_open_marker_relative_to_required() {
if registry::find_for("gemma4-27b-it").is_none() {
return;
}
let req = req_with_parallel(
"gemma4-27b-it",
Some(one_scalar_tool("get_weather")),
true,
);
let g_auto = compile_tool_grammar(&req, &ToolChoiceValue::Auto)
.expect("Auto compile must succeed")
.expect("Auto must return Some(grammar) under W-B2");
let g_req = compile_tool_grammar(&req, &ToolChoiceValue::Required)
.expect("Required compile must succeed")
.expect("Required must return Some(grammar)");
let auto_gbnf = grammar::serialize::serialize(&g_auto);
let req_gbnf = grammar::serialize::serialize(&g_req);
assert_ne!(
auto_gbnf, req_gbnf,
"Wave 3.5 HIGH-1 invariant: Auto and Required produce DIFFERENT \
grammars; Auto's body-only root is the production-correct match \
for the awaiting_trigger swallow-then-trigger byte-stream order. \
A regression that re-equates them would re-open the audit's \
'W-B2 Auto lazy grammar production correctness' divergence."
);
let open_marker_encoded = "[<] [|] [t] [o] [o] [l] [_] [c] [a] [l] [l] [>]";
let body_opener_encoded = "[c] [a] [l] [l] [:]"; assert!(
req_gbnf.contains(open_marker_encoded),
"Required grammar MUST contain the encoded open marker; \
grammar:\n{req_gbnf}"
);
assert!(
auto_gbnf.contains(open_marker_encoded),
"Auto grammar with parallel=true MUST still contain the open \
marker for inter-call positions (calls 2+); only the FIRST \
open marker is stripped. grammar:\n{auto_gbnf}"
);
let auto_root_line = auto_gbnf
.lines()
.find(|l| l.starts_with("root ::= "))
.expect("Auto grammar must have a root rule");
assert!(
auto_root_line.contains(body_opener_encoded),
"Wave 3.5 HIGH-1 invariant: Auto's `root` rule MUST inline the \
body opener (encoded `call:`) directly (the body-only shape \
produces `body close space` for parallel=false or \
`body close call_rule* space` for parallel=true at the root \
level). A regression to the old marker-wrapped grammar \
would put `gemma4-call` at the root and hide the body inline. \
root line: {auto_root_line}"
);
let req_root_line = req_gbnf
.lines()
.find(|l| l.starts_with("root ::= "))
.expect("Required grammar must have a root rule");
assert!(
!req_root_line.contains(body_opener_encoded),
"Wave 3.5 HIGH-1 control: Required's `root` rule MUST NOT inline \
the body opener — it expands via `gemma4-call` which has the \
open marker at its start. A regression that put body bytes \
directly in Required's root would break eager enforcement. \
root line: {req_root_line}"
);
}
#[test]
fn auto_lazy_grammar_allows_preamble_content() {
if registry::find_for("gemma4-27b-it").is_none() {
return;
}
let req = req_with("gemma4-27b-it", Some(one_scalar_tool("get_weather")));
let g = compile_tool_grammar(&req, &ToolChoiceValue::Auto)
.expect("Auto compile")
.expect("Auto must return Some(grammar)");
let root_id = g.rule_id("root").expect("compiled grammar has root");
let mut rt = grammar::sampler::GrammarRuntime::new(g, root_id)
.expect("GrammarRuntime::new must succeed");
rt.set_awaiting_trigger(true);
assert!(
rt.is_awaiting_trigger(),
"lazy Auto runtime starts in awaiting_trigger=true"
);
let preamble: &[u8] = b"Sure! Let me think about whether this needs a tool call. ";
let alive = rt.accept_bytes(preamble);
assert!(
alive,
"lazy Auto runtime MUST accept (return true) for preamble bytes \
before the trigger fires — the awaiting_trigger gate makes \
accept_bytes a no-op returning true"
);
assert!(
rt.is_awaiting_trigger(),
"preamble feed MUST NOT advance the runtime nor flip the trigger; \
trigger only flips when the engine calls runtime.trigger() in \
response to ToolCallSplitter's ToolCallOpen event"
);
assert!(
!rt.is_dead(),
"suspended (awaiting_trigger) runtime MUST never report dead — \
this is the wave-2.5 audit-divergence-A1 fix at sampler.rs:611"
);
assert!(
!rt.is_accepted(),
"suspended runtime is also NOT in accepting state — the gate \
is symmetric on both ends (sampler.rs:595)"
);
}
#[test]
fn auto_lazy_grammar_constrains_body_after_marker() {
if registry::find_for("gemma4-27b-it").is_none() {
return;
}
let req = req_with("gemma4-27b-it", Some(one_scalar_tool("get_weather")));
let g = compile_tool_grammar(&req, &ToolChoiceValue::Auto)
.expect("Auto compile")
.expect("Auto must return Some(grammar)");
let root_id = g.rule_id("root").expect("compiled grammar has root");
let mut rt = grammar::sampler::GrammarRuntime::new(g, root_id)
.expect("GrammarRuntime::new must succeed");
rt.set_awaiting_trigger(true);
rt.trigger();
assert!(
!rt.is_awaiting_trigger(),
"post-trigger the gate MUST be flipped false so subsequent \
mask + accept calls enforce the body grammar"
);
let garbage: &[u8] = b"this is not a tool call body";
let alive = rt.accept_bytes(garbage);
assert!(
!alive,
"post-trigger lazy Auto runtime MUST drive its stacks dead on \
non-call bytes — the grammar is now enforcing the per-model \
body shape (gemma4: call:NAME{{...}}<tool_call|>)"
);
assert!(
rt.is_dead(),
"post-trigger lazy Auto runtime that consumed invalid body bytes \
MUST report dead so the engine's defensive 500 + streaming \
error event fires (engine.rs:3673 dead-check global)"
);
let g2 = compile_tool_grammar(&req, &ToolChoiceValue::Auto)
.expect("compile")
.expect("Some");
let rid2 = g2.rule_id("root").expect("root");
let mut rt2 = grammar::sampler::GrammarRuntime::new(g2, rid2).expect("rt");
rt2.set_awaiting_trigger(true);
rt2.trigger();
let valid_prefix = b"call:get_weather{q:";
let alive2 = rt2.accept_bytes(valid_prefix);
assert!(
alive2,
"post-trigger runtime MUST stay alive on valid Gemma4 body prefix \
(no leading open marker — splitter swallowed it); a regression \
that marks all post-trigger feeds dead would trip this control"
);
assert!(
!rt2.is_dead(),
"post-trigger runtime fed valid body prefix bytes MUST NOT report dead"
);
}
#[test]
fn auto_lazy_grammar_accepts_production_order_byte_stream() {
let reg = match registry::find_for("gemma4-27b-it") {
Some(r) => r,
None => return,
};
let open = reg.tool_open.expect("gemma4 has tool_open");
let close = reg.tool_close.expect("gemma4 has tool_close");
let req = req_with("gemma4-27b-it", Some(one_scalar_tool("get_weather")));
let g = compile_tool_grammar(&req, &ToolChoiceValue::Auto)
.expect("Auto compile must succeed")
.expect("Auto must return Some(grammar)");
let root_id = g.rule_id("root").expect("compiled grammar has root");
let mut rt = grammar::sampler::GrammarRuntime::new(g, root_id)
.expect("GrammarRuntime::new must succeed");
rt.set_awaiting_trigger(true);
let pre_trigger_alive = rt.accept_bytes(open.as_bytes());
assert!(
pre_trigger_alive,
"while awaiting_trigger=true, accept_bytes MUST be a no-op \
returning true regardless of input bytes — the lazy gate \
prevents the open marker from advancing the grammar; \
this is the apply+accept atomicity invariant from \
research-report.md Q2"
);
assert!(
rt.is_awaiting_trigger(),
"open-marker bytes MUST NOT flip the trigger; only the engine \
explicitly calling rt.trigger() in response to the splitter's \
ToolCallOpen event should flip the gate"
);
rt.trigger();
assert!(
!rt.is_awaiting_trigger(),
"post-trigger the gate MUST be flipped false"
);
let body = b"call:get_weather{q:<|\"|>SF<|\"|>}";
let body_alive = rt.accept_bytes(body);
assert!(
body_alive,
"production-order regression: body-only grammar MUST accept \
body bytes after trigger. A regression to the old marker- \
wrapped grammar (OneOrMoreCalls) would mark the runtime dead \
because root would still expect the (already-consumed) open \
marker. This is the exact divergence the Wave 3.5 audit \
caught at /tmp/cfa-cfa-20260427-adr005-wave3/codex-review-last.txt \
('W-B2 Auto lazy grammar production correctness')."
);
assert!(
!rt.is_dead(),
"production-order body bytes MUST NOT drive the runtime dead \
on a body-only Auto-lazy grammar"
);
let close_alive = rt.accept_bytes(close.as_bytes());
assert!(
close_alive,
"close marker bytes MUST be accepted by the body-only grammar — \
the close marker IS in the grammar (only the FIRST open \
marker is stripped)"
);
assert!(
rt.is_accepted(),
"after body + close marker the runtime MUST be in accepting \
state (single-call body-only grammar = `body close space`); \
the engine's is_accepted check at engine.rs:3678 then drives \
early termination"
);
assert!(!rt.is_dead(), "an accepted runtime MUST NOT also be dead");
}
#[test]
fn auto_lazy_grammar_accepts_qwen35_production_order_byte_stream() {
let reg = match registry::find_for("qwen3.5-72b-instruct") {
Some(r) => r,
None => return,
};
let open = reg.tool_open.expect("qwen35 has tool_open: <tool_call>");
let close = reg.tool_close.expect("qwen35 has tool_close: </tool_call>");
let req = req_with("qwen3.5-72b-instruct", Some(one_scalar_tool("get_weather")));
let g = compile_tool_grammar(&req, &ToolChoiceValue::Auto)
.expect("Auto compile must succeed for qwen35")
.expect("Auto must return Some(grammar) for registered qwen35");
let root_id = g.rule_id("root").expect("compiled grammar has root");
let mut rt = grammar::sampler::GrammarRuntime::new(g, root_id)
.expect("GrammarRuntime::new must succeed");
rt.set_awaiting_trigger(true);
let pre_trigger_alive = rt.accept_bytes(open.as_bytes());
assert!(
pre_trigger_alive,
"qwen35: while awaiting_trigger=true, accept_bytes MUST be a \
no-op returning true for the open marker `{open}`"
);
assert!(
rt.is_awaiting_trigger(),
"qwen35: open-marker bytes MUST NOT flip the trigger; only \
runtime.trigger() in response to ToolCallOpen should flip it"
);
rt.trigger();
assert!(
!rt.is_awaiting_trigger(),
"qwen35: post-trigger gate MUST be flipped false"
);
let body: &[u8] = b"<function=get_weather>\n<parameter=q>\nSF\n</parameter>\n</function>";
let body_alive = rt.accept_bytes(body);
assert!(
body_alive,
"qwen35 production-order regression: body-only grammar MUST accept \
body bytes (starting with `<function=`) after trigger with no \
leading `<tool_call>` marker. A regression to OneOrMoreCalls would \
mark the runtime dead because root would still expect the already- \
consumed open marker `<tool_call>`."
);
assert!(
!rt.is_dead(),
"qwen35: body bytes MUST NOT drive the runtime dead on the \
body-only Auto-lazy grammar"
);
let close_bytes = format!("\n{close}");
let close_alive = rt.accept_bytes(close_bytes.as_bytes());
assert!(
close_alive,
"qwen35: close marker bytes (`\\n{close}`) MUST be accepted by the \
body-only grammar — only the FIRST open marker is stripped"
);
assert!(
rt.is_accepted(),
"qwen35: after body + close marker the runtime MUST be in accepting \
state (single-call body-only grammar = `body \\n close space`)"
);
assert!(
!rt.is_dead(),
"qwen35: accepted runtime MUST NOT also be dead"
);
}
#[test]
fn auto_lazy_grammar_accepts_two_call_production_order() {
{
let reg = match registry::find_for("gemma4-27b-it") {
Some(r) => r,
None => {
eprintln!("gemma4 absent; skipping two-call parallel test for Gemma");
return;
}
};
let open = reg.tool_open.expect("gemma4 has tool_open");
let close = reg.tool_close.expect("gemma4 has tool_close");
let req = req_with_parallel(
"gemma4-27b-it",
Some(one_scalar_tool("get_weather")),
true,
);
let g = compile_tool_grammar(&req, &ToolChoiceValue::Auto)
.expect("Auto compile must succeed for gemma4 parallel")
.expect("Auto must return Some(grammar)");
let root_id = g
.rule_id("root")
.expect("gemma4 parallel Auto grammar has root");
let mut rt =
grammar::sampler::GrammarRuntime::new(g, root_id).expect("GrammarRuntime::new");
rt.set_awaiting_trigger(true);
let pre = rt.accept_bytes(open.as_bytes());
assert!(
pre,
"gemma4 parallel: pre-trigger open marker MUST be no-op"
);
assert!(rt.is_awaiting_trigger());
rt.trigger();
assert!(!rt.is_awaiting_trigger());
let body1 = b"call:get_weather{q:<|\"|>SF<|\"|>}";
let alive1 = rt.accept_bytes(body1);
assert!(
alive1,
"gemma4 parallel: body1 bytes MUST be accepted by body-only grammar"
);
assert!(!rt.is_dead());
let alive_c1 = rt.accept_bytes(close.as_bytes());
assert!(
alive_c1 || rt.is_accepted(),
"gemma4 parallel: close1 must be accepted"
);
let alive_inter = rt.accept_bytes(open.as_bytes());
assert!(
alive_inter,
"gemma4 parallel: inter-call open marker 2 (`{open}`) MUST be \
accepted — it is in the grammar's inter-call rule (only the \
FIRST open marker is stripped)"
);
assert!(!rt.is_dead());
let body2 = b"call:get_weather{q:<|\"|>NYC<|\"|>}";
let alive2 = rt.accept_bytes(body2);
assert!(
alive2,
"gemma4 parallel: body2 bytes MUST be accepted by inter-call body rule"
);
assert!(!rt.is_dead());
let alive_c2 = rt.accept_bytes(close.as_bytes());
assert!(
alive_c2 || rt.is_accepted(),
"gemma4 parallel: close2 must be accepted"
);
assert!(
rt.is_accepted(),
"gemma4 parallel: after two complete calls the runtime MUST be \
in accepting state (body-only parallel root = \
`body close (sep open body close)* space`)"
);
assert!(!rt.is_dead());
}
{
let reg = match registry::find_for("qwen3.5-72b-instruct") {
Some(r) => r,
None => {
eprintln!("qwen35 absent; skipping two-call parallel test for Qwen35");
return;
}
};
let open = reg.tool_open.expect("qwen35 has tool_open: <tool_call>");
let close = reg.tool_close.expect("qwen35 has tool_close: </tool_call>");
let separator = reg.parallel_call_separator();
let req = req_with_parallel(
"qwen3.5-72b-instruct",
Some(one_scalar_tool("get_weather")),
true,
);
let g = compile_tool_grammar(&req, &ToolChoiceValue::Auto)
.expect("Auto compile must succeed for qwen35 parallel")
.expect("Auto must return Some(grammar)");
let root_id = g
.rule_id("root")
.expect("qwen35 parallel Auto grammar has root");
let mut rt =
grammar::sampler::GrammarRuntime::new(g, root_id).expect("GrammarRuntime::new");
rt.set_awaiting_trigger(true);
let pre = rt.accept_bytes(open.as_bytes());
assert!(
pre,
"qwen35 parallel: pre-trigger open marker MUST be no-op"
);
assert!(rt.is_awaiting_trigger());
rt.trigger();
assert!(!rt.is_awaiting_trigger());
let body1 = b"<function=get_weather>\n<parameter=q>\nSF\n</parameter>\n</function>";
let alive1 = rt.accept_bytes(body1);
assert!(
alive1,
"qwen35 parallel: body1 bytes MUST be accepted by body-only grammar"
);
assert!(!rt.is_dead());
let close1 = format!("\n{close}");
let alive_c1 = rt.accept_bytes(close1.as_bytes());
assert!(
alive_c1 || rt.is_accepted(),
"qwen35 parallel: close1 (`\\n{close}`) must be accepted"
);
let inter_call = format!("{separator}{open}\n");
let alive_inter = rt.accept_bytes(inter_call.as_bytes());
assert!(
alive_inter,
"qwen35 parallel: inter-call bytes `{inter_call:?}` MUST be \
accepted — the second open marker IS in the grammar's inter-call \
rule (only the FIRST open marker is stripped from the root)"
);
assert!(!rt.is_dead());
let body2 = b"<function=get_weather>\n<parameter=q>\nNYC\n</parameter>\n</function>";
let alive2 = rt.accept_bytes(body2);
assert!(
alive2,
"qwen35 parallel: body2 bytes MUST be accepted in inter-call body rule"
);
assert!(!rt.is_dead());
let close2 = format!("\n{close}");
let alive_c2 = rt.accept_bytes(close2.as_bytes());
assert!(
alive_c2 || rt.is_accepted(),
"qwen35 parallel: close2 must be accepted"
);
assert!(
rt.is_accepted(),
"qwen35 parallel: after two complete calls the runtime MUST be \
in accepting state (body-only parallel root = \
`body \\n close (\\n open \\n body \\n close)* space`)"
);
assert!(!rt.is_dead());
}
}
fn req_with_parallel(
model: &str,
tools: Option<Vec<Tool>>,
parallel: bool,
) -> ChatCompletionRequest {
let mut r = req_with(model, tools);
r.parallel_tool_calls = Some(parallel);
r
}
#[test]
fn compile_tool_grammar_parallel_two_real_tools_alternates() {
if registry::find_for("gemma4-27b-it").is_none() {
return;
}
let tool_search = Tool {
tool_type: "function".to_string(),
function: ToolFunction {
name: "search".to_string(),
description: None,
parameters: Some(serde_json::json!({
"type": "object",
"properties": {"query": {"type": "string"}},
"required": ["query"]
})),
},
};
let tool_add = Tool {
tool_type: "function".to_string(),
function: ToolFunction {
name: "add_numbers".to_string(),
description: None,
parameters: Some(serde_json::json!({
"type": "object",
"properties": {
"x": {"type": "integer"},
"y": {"type": "integer"}
},
"required": ["x", "y"]
})),
},
};
let req = req_with_parallel(
"gemma4-27b-it",
Some(vec![tool_search, tool_add]),
true,
);
let g = compile_tool_grammar(&req, &ToolChoiceValue::Required)
.expect("compile_tool_grammar must succeed for registered gemma4")
.expect("must return Some(grammar) for Required + 2 tools + gemma4");
let gbnf = grammar::serialize::serialize(&g);
assert!(
g.rule_id("fn-0-root").is_some(),
"parallel combine must have fn-0-root; grammar:\n{gbnf}"
);
assert!(
g.rule_id("fn-1-root").is_some(),
"parallel combine must have fn-1-root; grammar:\n{gbnf}"
);
let root_id = g
.rule_id("root")
.expect("combined grammar must have root rule");
let mut rt = grammar::sampler::GrammarRuntime::new(g, root_id)
.expect("GrammarRuntime::new must succeed for valid combined grammar");
let call_0: &[u8] = b"<|tool_call>call:search{query:<|\"|>hello<|\"|>}<tool_call|>";
let call_1: &[u8] = b"<|tool_call>call:add_numbers{x:1,y:2}<tool_call|>";
let mut payload = Vec::new();
payload.extend_from_slice(call_0);
payload.extend_from_slice(call_1);
let alive = rt.accept_bytes(&payload);
assert!(
alive || rt.is_accepted(),
"parallel gemma4 two-tool runtime did not accept alternating calls.\n\
payload: {}\ngrammar:\n{gbnf}",
String::from_utf8_lossy(&payload)
);
assert!(
rt.is_accepted(),
"parallel gemma4 two-tool runtime must be in accepting state after \
two alternating calls.\npayload: {}\ngrammar:\n{gbnf}",
String::from_utf8_lossy(&payload)
);
}
#[test]
fn compile_tool_grammar_deepseek_parallel_tools_share_outer_dsml_block() {
let search = Tool {
tool_type: "function".into(),
function: ToolFunction {
name: "search".into(),
description: None,
parameters: Some(serde_json::json!({
"type": "object",
"properties": {"query": {"type": "string"}},
"required": ["query"]
})),
},
};
let read = Tool {
tool_type: "function".into(),
function: ToolFunction {
name: "read".into(),
description: None,
parameters: Some(serde_json::json!({
"type": "object",
"properties": {"path": {"type": "string"}},
"required": ["path"]
})),
},
};
let req = req_with_parallel("DeepSeek-V4-Flash-0731", Some(vec![search, read]), true);
let grammar = compile_tool_grammar(&req, &ToolChoiceValue::Required)
.expect("compile DeepSeek multi-tool grammar")
.expect("required tools must produce grammar");
let root = grammar.rule_id("root").expect("root rule");
let mut runtime =
grammar::sampler::GrammarRuntime::new(grammar, root).expect("DeepSeek grammar runtime");
let payload = "<|DSML|tool_calls>\n<|DSML|invoke name=\"search\">\n<|DSML|parameter name=\"query\" string=\"true\">kv cache</|DSML|parameter>\n</|DSML|invoke>\n<|DSML|invoke name=\"read\">\n<|DSML|parameter name=\"path\" string=\"true\">src/main.rs</|DSML|parameter>\n</|DSML|invoke>\n</|DSML|tool_calls>";
assert!(runtime.accept_bytes(payload.as_bytes()));
assert!(runtime.is_accepted());
}
#[test]
fn compile_tool_grammar_deepseek_auto_rejects_bare_dsml_atom_before_invoke() {
let bash = Tool {
tool_type: "function".into(),
function: ToolFunction {
name: "bash".into(),
description: None,
parameters: Some(serde_json::json!({
"type": "object",
"properties": {"command": {"type": "string"}},
"required": ["command"]
})),
},
};
let read = Tool {
tool_type: "function".into(),
function: ToolFunction {
name: "read".into(),
description: None,
parameters: Some(serde_json::json!({
"type": "object",
"properties": {"filePath": {"type": "string"}},
"required": ["filePath"]
})),
},
};
let req = req_with_parallel("DeepSeek-V4-Flash-0731", Some(vec![bash, read]), false);
let grammar = compile_tool_grammar(&req, &ToolChoiceValue::Auto)
.expect("compile DeepSeek auto multi-tool grammar")
.expect("auto tools must produce a lazy grammar");
let root = grammar.rule_id("root").expect("root rule");
let valid = "\n<|DSML|invoke name=\"bash\">\n<|DSML|parameter name=\"command\" string=\"true\">ls -la</|DSML|parameter>\n</|DSML|invoke>\n</|DSML|tool_calls>";
let mut valid_runtime = grammar::sampler::GrammarRuntime::new(grammar.clone(), root)
.expect("DeepSeek auto grammar runtime");
valid_runtime.set_awaiting_trigger(true);
valid_runtime.trigger();
assert!(valid_runtime.accept_bytes(valid.as_bytes()));
assert!(valid_runtime.is_accepted());
let leaked = "\n<|DSML|\n<|DSML|invoke name=\"bash\">\n<|DSML|parameter name=\"command\" string=\"true\">ls -la</|DSML|parameter>\n</|DSML|invoke>\n</|DSML|tool_calls>";
let mut leaked_runtime = grammar::sampler::GrammarRuntime::new(grammar, root)
.expect("DeepSeek auto grammar runtime");
leaked_runtime.set_awaiting_trigger(true);
leaked_runtime.trigger();
assert!(
!leaked_runtime.accept_bytes(leaked.as_bytes()),
"bare DSML atom before invoke must be rejected by the lazy grammar"
);
}
}
fn combine_function_grammars(
gbnfs: Vec<String>,
parallel: bool,
separator_literal: &str,
) -> String {
use grammar::parser::parse;
use grammar::serialize::{rename_rules, serialize};
let mut combined = String::with_capacity(gbnfs.iter().map(|g| g.len()).sum::<usize>() + 64);
if parallel {
combined.push_str("alt ::= ");
for i in 0..gbnfs.len() {
if i > 0 {
combined.push_str(" | ");
}
combined.push_str(&format!("fn-{}-root", i));
}
combined.push('\n');
if separator_literal.is_empty() {
combined.push_str("root ::= alt alt*\n");
} else {
let mut esc = String::with_capacity(separator_literal.len() + 2);
esc.push('"');
for c in separator_literal.chars() {
match c {
'\n' => esc.push_str("\\n"),
'\r' => esc.push_str("\\r"),
'"' => esc.push_str("\\\""),
'\\' => esc.push_str("\\\\"),
_ => esc.push(c),
}
}
esc.push('"');
combined.push_str(&format!("root ::= alt ( {} alt )*\n", esc));
}
} else {
combined.push_str("root ::= ");
for i in 0..gbnfs.len() {
if i > 0 {
combined.push_str(" | ");
}
combined.push_str(&format!("fn-{}-root", i));
}
combined.push('\n');
}
for (i, gbnf) in gbnfs.iter().enumerate() {
let parsed = parse(gbnf).unwrap_or_else(|e| {
panic!(
"combine_function_grammars: per-function GBNF #{} from emitter \
failed to parse — this is a registry.rs emitter bug: {}\n\
grammar:\n{}",
i, e, gbnf
)
});
let renamed = rename_rules(&parsed, |name| format!("fn-{}-{}", i, name));
combined.push_str(&serialize(&renamed));
}
if let Err(e) = parse(&combined) {
panic!(
"combine_function_grammars: combined output failed to re-parse — \
this is an AST combiner invariant violation: {}\n\
combined grammar was:\n{}",
e, combined
);
}
combined
}
fn combine_deepseek4_function_grammars(
gbnfs: Vec<String>,
auto_lazy: bool,
parallel: bool,
) -> String {
use grammar::parser::parse;
use grammar::serialize::{rename_rules, serialize};
debug_assert!(gbnfs.len() > 1);
let alternatives = (0..gbnfs.len())
.map(|index| format!("fn-{index}-root"))
.collect::<Vec<_>>()
.join(" | ");
let mut combined = format!("dsml-alt ::= {alternatives}\n");
let open = if auto_lazy {
String::new()
} else {
r#""<|DSML|tool_calls>" "#.to_string()
};
let invokes = if parallel {
r#"dsml-alt ( "\n" dsml-alt )*"#
} else {
"dsml-alt"
};
combined.push_str(&format!(
"root ::= {open}\"\\n\" {invokes} \"\\n\" \"</|DSML|tool_calls>\"\n"
));
for (index, gbnf) in gbnfs.iter().enumerate() {
let parsed = parse(gbnf).unwrap_or_else(|error| {
panic!(
"combine_deepseek4_function_grammars: per-function grammar #{index} \
failed to parse: {error}\ngrammar:\n{gbnf}"
)
});
let renamed = rename_rules(&parsed, |name| format!("fn-{index}-{name}"));
combined.push_str(&serialize(&renamed));
}
if let Err(error) = parse(&combined) {
panic!("combine_deepseek4_function_grammars produced invalid grammar: {error}\n{combined}");
}
combined
}
#[cfg(test)]
mod combine_function_grammars_tests {
use super::grammar::parser::parse;
use super::grammar::sampler::GrammarRuntime;
use super::*;
fn runtime_from_gbnf(gbnf: &str) -> GrammarRuntime {
let g = parse(gbnf).expect("parse combined grammar");
let rid = g.rule_id("root").expect("root rule exists");
GrammarRuntime::new(g, rid).expect("runtime init")
}
fn gbnf_for_tool_a() -> String {
concat!(
"root ::= \"call:tool_a\" \"{\" space tool_a-q-kv space \"}\"\n",
"tool_a-q-kv ::= \"\\\"q\\\"\" \":\" space string\n",
"string ::= \"\\\"\" ( [^\"\\\\] )* \"\\\"\"\n",
"space ::= \" \"?\n",
)
.to_string()
}
fn gbnf_for_tool_b() -> String {
concat!(
"root ::= \"call:tool_b\" \"{\" space tool_b-q-kv space \"}\"\n",
"tool_b-q-kv ::= \"\\\"q\\\"\" \":\" space integer\n",
"integer ::= \"-\"? [0-9]+\n",
"space ::= \" \"?\n",
)
.to_string()
}
#[test]
fn multi_tool_shared_param_name_different_types() {
let combined =
combine_function_grammars(vec![gbnf_for_tool_a(), gbnf_for_tool_b()], false, "");
let g = parse(&combined).expect("combined grammar must re-parse");
assert!(
g.rule_id("fn-0-root").is_some(),
"combined grammar missing fn-0-root: {combined}"
);
assert!(
g.rule_id("fn-1-root").is_some(),
"combined grammar missing fn-1-root: {combined}"
);
let mut rt_a = runtime_from_gbnf(&combined);
let alive = rt_a.accept_bytes(b"call:tool_a{\"q\":\"hi\"}");
assert!(
alive,
"tool_a string call rejected by combined grammar: {combined}"
);
assert!(
rt_a.is_accepted(),
"tool_a string call not in accepting state"
);
let mut rt_b = runtime_from_gbnf(&combined);
let alive = rt_b.accept_bytes(b"call:tool_b{\"q\":42}");
assert!(
alive,
"tool_b integer call rejected by combined grammar: {combined}"
);
assert!(
rt_b.is_accepted(),
"tool_b integer call not in accepting state"
);
}
#[test]
fn single_tool_combine_roundtrip() {
let combined = combine_function_grammars(vec![gbnf_for_tool_a()], false, "");
let g = parse(&combined).expect("combined re-parses");
assert!(
g.rule_id("fn-0-root").is_some(),
"single-tool combine missing fn-0-root: {combined}"
);
let mut rt = runtime_from_gbnf(&combined);
let alive = rt.accept_bytes(b"call:tool_a{\"q\":\"hi\"}");
assert!(
alive,
"single-tool combine rejected valid tool_a call: {combined}"
);
assert!(
rt.is_accepted(),
"single-tool combine not accepting valid call"
);
}
#[test]
fn effective_grammar_precedence_tool_wins_over_response_format() {
let tool_grammar: Option<u32> = Some(1); let response_grammar: Option<u32> = Some(2); let effective = tool_grammar.or(response_grammar);
assert_eq!(
effective,
Some(1),
"tool grammar must win over response_format grammar (.or() semantics)"
);
let no_tool: Option<u32> = None;
let effective_resp_only = no_tool.or(Some(2));
assert_eq!(
effective_resp_only,
Some(2),
"response_format grammar must apply when tool_choice is absent"
);
let none_a: Option<u32> = None;
let none_b: Option<u32> = None;
assert!(
none_a.or(none_b).is_none(),
"no tool_grammar + no response_grammar must produce no constraint"
);
}
#[test]
fn combine_two_gemma_grammars_preserves_negated_char_class() {
let gbnf_a = "root ::= \"A\" [^<\\\\]+\n".to_string();
let gbnf_b = "root ::= \"B\" [^<\\\\]+\n".to_string();
let combined = combine_function_grammars(vec![gbnf_a.clone(), gbnf_b.clone()], false, "");
parse(&combined).expect("combined grammar must re-parse");
let mut rt = runtime_from_gbnf(&combined);
let alive = rt.accept_bytes(b"A<");
assert!(
!alive,
"negated char class `[^<\\\\]` was corrupted by combine: \
a `<` byte was ACCEPTED through the fn-0 (gemma_a) path, but \
the source grammar should REJECT it.\nCombined grammar:\n{combined}"
);
let mut rt = runtime_from_gbnf(&combined);
let alive = rt.accept_bytes(b"B<");
assert!(
!alive,
"negated char class corrupted on fn-1 (gemma_b) path: \
`<` ACCEPTED. Combined grammar:\n{combined}"
);
let mut rt = runtime_from_gbnf(&combined);
let alive = rt.accept_bytes(b"Ahello");
assert!(
alive,
"valid string `Ahello` rejected through combined grammar (fn-0): \
combined output corrupted in a non-negation way.\nCombined:\n{combined}"
);
assert!(
rt.is_accepted(),
"valid string `Ahello` consumed but runtime not in accepting state \
(fn-0). Combined:\n{combined}"
);
let mut rt = runtime_from_gbnf(&combined);
let alive = rt.accept_bytes(b"Bworld");
assert!(alive, "valid string `Bworld` rejected (fn-1): {combined}");
assert!(
rt.is_accepted(),
"valid string `Bworld` not in accepting state (fn-1): {combined}"
);
}
#[test]
fn combine_two_qwen_grammars_round_trips() {
let gbnf_a = concat!(
"root ::= \"<function=fa>\\n<parameter=q>\\n\" qstr \"\\n</parameter>\\n</function>\"\n",
"qstr ::= qchar*\n",
"qchar ::= [^<\\\\] | [\\\\] [^\\x00-\\x1F]\n",
)
.to_string();
let gbnf_b = concat!(
"root ::= \"<function=fb>\\n<parameter=q>\\n\" qstr \"\\n</parameter>\\n</function>\"\n",
"qstr ::= qchar*\n",
"qchar ::= [^<\\\\] | [\\\\] [^\\x00-\\x1F]\n",
)
.to_string();
let combined = combine_function_grammars(vec![gbnf_a, gbnf_b], false, "");
let g = parse(&combined).expect("combined re-parses");
assert!(
g.rule_id("fn-0-root").is_some(),
"fn-0-root missing: {combined}"
);
assert!(
g.rule_id("fn-1-root").is_some(),
"fn-1-root missing: {combined}"
);
let mut rt = runtime_from_gbnf(&combined);
let alive =
rt.accept_bytes(b"<function=fa>\n<parameter=q>\nhello\n</parameter>\n</function>");
assert!(alive, "canonical Qwen call (fa) rejected: {combined}");
assert!(
rt.is_accepted(),
"canonical Qwen call (fa) not accepting: {combined}"
);
let mut rt = runtime_from_gbnf(&combined);
let alive =
rt.accept_bytes(b"<function=fb>\n<parameter=q>\nworld\n</parameter>\n</function>");
assert!(alive, "canonical Qwen call (fb) rejected: {combined}");
assert!(
rt.is_accepted(),
"canonical Qwen call (fb) not accepting: {combined}"
);
}
#[test]
fn combine_namespacing_isolates_shared_param_names() {
let gbnf_a = concat!(
"root ::= \"A\" qval\n",
"qval ::= \"\\\"\" qchar* \"\\\"\"\n",
"qchar ::= [^\"\\\\]\n",
)
.to_string();
let gbnf_b = concat!("root ::= \"B\" qval\n", "qval ::= [0-9]+\n",).to_string();
let combined = combine_function_grammars(vec![gbnf_a, gbnf_b], false, "");
parse(&combined).expect("combined re-parses");
let mut rt = runtime_from_gbnf(&combined);
let alive = rt.accept_bytes(b"A\"hello\"");
assert!(alive, "fn-0 string call rejected: {combined}");
assert!(
rt.is_accepted(),
"fn-0 string call not accepting: {combined}"
);
let mut rt = runtime_from_gbnf(&combined);
let alive = rt.accept_bytes(b"B42");
assert!(alive, "fn-1 number call rejected: {combined}");
assert!(
rt.is_accepted(),
"fn-1 number call not accepting: {combined}"
);
let mut rt = runtime_from_gbnf(&combined);
let alive = rt.accept_bytes(b"A42");
assert!(
!alive || !rt.is_accepted(),
"namespace isolation broken: fn-0 (A) accepted a numeric `q` value, \
but fn-0's `qval` rule is string-only. This means fn-1's number \
rule leaked into fn-0's namespace.\nCombined:\n{combined}"
);
let mut rt = runtime_from_gbnf(&combined);
let alive = rt.accept_bytes(b"B\"hello\"");
assert!(
!alive || !rt.is_accepted(),
"namespace isolation broken: fn-1 (B) accepted a string `q` value, \
but fn-1's `qval` rule is numeric-only.\nCombined:\n{combined}"
);
}
#[test]
fn combine_single_function_yields_only_root_alternative() {
let src = "root ::= \"X\" [a-z]+\n".to_string();
let combined = combine_function_grammars(vec![src.clone()], false, "");
let g = parse(&combined).expect("combined re-parses");
let root_id = g.rule_id("root").expect("root exists");
let root_rule = &g.rules[root_id as usize];
let alt_count = root_rule
.iter()
.filter(|e| e.ty == grammar::parser::GretType::Alt)
.count();
assert_eq!(
alt_count, 0,
"single-function combine root should have no Alt elements (no `|`); \
got {alt_count}.\nCombined:\n{combined}"
);
let mut rt = runtime_from_gbnf(&combined);
let alive = rt.accept_bytes(b"Xhello");
assert!(
alive,
"single-function combine rejected `Xhello`: {combined}"
);
assert!(
rt.is_accepted(),
"single-function combine not accepting `Xhello`: {combined}"
);
let mut src_rt = runtime_from_gbnf(&src);
let alive = src_rt.accept_bytes(b"Xhello");
assert!(alive, "source grammar disagrees on `Xhello`");
assert!(src_rt.is_accepted(), "source grammar disagrees on `Xhello`");
}
#[test]
fn parallel_multi_tool_runtime_accepts_alternating_calls_gemma_separator() {
let combined = combine_function_grammars(
vec![gbnf_for_tool_a(), gbnf_for_tool_b()],
true,
"",
);
let g = parse(&combined).expect("parallel multi-tool combined re-parses");
assert!(
g.rule_id("fn-0-root").is_some(),
"parallel combine missing fn-0-root: {combined}"
);
assert!(
g.rule_id("fn-1-root").is_some(),
"parallel combine missing fn-1-root: {combined}"
);
assert!(
g.rule_id("alt").is_some(),
"parallel combine missing synthetic alt rule: {combined}"
);
let mut rt = runtime_from_gbnf(&combined);
let payload = b"call:tool_a{\"q\":\"hi\"}call:tool_b{\"q\":42}call:tool_a{\"q\":\"bye\"}";
let alive = rt.accept_bytes(payload);
assert!(
alive,
"Gemma-shape parallel runtime rejected three alternating calls.\n\
payload: {}\n\
grammar:\n{combined}",
String::from_utf8_lossy(payload)
);
assert!(
rt.is_accepted(),
"Gemma-shape parallel runtime not in accepting state after three calls"
);
let mut rt_single = runtime_from_gbnf(&combined);
let alive = rt_single.accept_bytes(b"call:tool_b{\"q\":7}");
assert!(
alive,
"Gemma-shape parallel runtime rejected single tool_b call"
);
assert!(rt_single.is_accepted());
let rt_empty = runtime_from_gbnf(&combined);
assert!(
!rt_empty.is_accepted(),
"parallel grammar must NOT accept empty (min_calls=1)"
);
}
#[test]
fn parallel_multi_tool_runtime_accepts_alternating_calls_qwen_separator() {
let combined = combine_function_grammars(
vec![gbnf_for_tool_a(), gbnf_for_tool_b()],
true,
"\n",
);
let g = parse(&combined).expect("Qwen-separator combined re-parses");
assert!(g.rule_id("fn-0-root").is_some());
assert!(g.rule_id("fn-1-root").is_some());
assert!(g.rule_id("alt").is_some());
let mut rt = runtime_from_gbnf(&combined);
let payload =
b"call:tool_a{\"q\":\"hi\"}\ncall:tool_b{\"q\":42}\ncall:tool_a{\"q\":\"bye\"}";
let alive = rt.accept_bytes(payload);
assert!(
alive,
"Qwen-shape parallel runtime rejected three alternating newline-separated \
calls.\npayload: {}\ngrammar:\n{combined}",
String::from_utf8_lossy(payload)
);
assert!(
rt.is_accepted(),
"Qwen-shape parallel runtime not in accepting state after three calls"
);
let mut rt_no_sep = runtime_from_gbnf(&combined);
let no_sep_payload = b"call:tool_a{\"q\":\"hi\"}call:tool_b{\"q\":42}";
let alive = rt_no_sep.accept_bytes(no_sep_payload);
let rejected = !alive || !rt_no_sep.is_accepted();
assert!(
rejected,
"Qwen-shape parallel runtime accepted multi-call WITHOUT '\\n' \
separator — the family separator must be load-bearing.\n\
grammar:\n{combined}"
);
}
}
#[cfg(test)]
mod grammar_kind_selection_tests {
use super::engine::GrammarKind;
#[test]
fn default_grammar_kind_is_response_format() {
let k = GrammarKind::default();
assert_eq!(
k,
GrammarKind::ResponseFormat,
"Default grammar_kind MUST be ResponseFormat — any caller that \
omits the field must get unconditional grammar enforcement, \
not the trigger-gated tool-call-body path. This preserves \
pre-wave-2.5 response_format semantics on all paths."
);
}
#[test]
fn tool_choice_required_yields_tool_call_body_kind() {
fn select(tool: Option<u32>, resp: Option<u32>) -> (Option<u32>, GrammarKind) {
match (tool, resp) {
(Some(g), _) => (Some(g), GrammarKind::ToolCallBodyRequired),
(None, Some(g)) => (Some(g), GrammarKind::ResponseFormat),
(None, None) => (None, GrammarKind::default()),
}
}
let (g, k) = select(Some(1), Some(2));
assert_eq!(g, Some(1), "tool grammar must win precedence");
assert_eq!(
k,
GrammarKind::ToolCallBodyRequired,
"tool_choice=required/function MUST yield GrammarKind::ToolCallBodyRequired \
so the runtime is EAGER from token 0 — the grammar root already \
wraps the body in open/close markers, mirroring llama.cpp \
grammar_lazy=false at common/chat.cpp:898-913, 1177-1200."
);
let (g, k) = select(None, Some(2));
assert_eq!(g, Some(2), "response_format grammar must apply");
assert_eq!(
k,
GrammarKind::ResponseFormat,
"response_format=json_object/json_schema MUST yield \
GrammarKind::ResponseFormat. This is the wave-2.5 audit fix: \
without this, the runtime would sit in awaiting_trigger=true \
and never enforce the JSON grammar on registered Gemma/Qwen \
models because the tool open marker never fires."
);
let (g, k) = select(None, None);
assert!(g.is_none(), "no grammar means no constraint");
assert_eq!(
k,
GrammarKind::default(),
"kind defaults when grammar is None"
);
}
#[test]
fn response_format_json_schema_yields_response_format_kind_through_production_helper() {
use super::super::schema::{JsonSchemaSpec, ResponseFormat};
let rf = ResponseFormat::JsonSchema {
json_schema: JsonSchemaSpec {
name: "weather".to_string(),
description: None,
strict: None,
schema: serde_json::json!({
"type": "object",
"properties": {
"city": {"type": "string"},
"temp": {"type": "number"}
},
"required": ["city", "temp"]
}),
},
};
let response_grammar =
super::compile_response_format(&rf).expect("real compile must succeed");
assert!(
response_grammar.is_some(),
"compile_response_format(json_schema) MUST produce a grammar"
);
let (effective, kind) = super::select_effective_grammar(
None,
GrammarKind::ToolCallBodyAuto,
response_grammar,
);
assert!(
effective.is_some(),
"select_effective_grammar must propagate response grammar when tool=None"
);
assert_eq!(
kind,
GrammarKind::ResponseFormat,
"response_format=json_schema MUST yield GrammarKind::ResponseFormat \
through the PRODUCTION helper (not a stand-in match). This proves \
the wave-2.5 audit fix is wired through the real selection path."
);
}
#[test]
fn tool_grammar_wins_over_response_format_through_production_helper() {
use super::super::schema::{JsonSchemaSpec, ResponseFormat};
let rf_resp = ResponseFormat::JsonObject;
let resp_grammar = super::compile_response_format(&rf_resp).expect("compile json_object");
assert!(resp_grammar.is_some());
let rf_tool = ResponseFormat::JsonSchema {
json_schema: JsonSchemaSpec {
name: "tool_proxy".to_string(),
description: None,
strict: None,
schema: serde_json::json!({
"type": "object",
"properties": {"name": {"type": "string"}},
"required": ["name"]
}),
},
};
let tool_grammar = super::compile_response_format(&rf_tool).expect("compile schema");
assert!(tool_grammar.is_some());
let tool_gbnf = super::grammar::serialize::serialize(tool_grammar.as_ref().unwrap());
let resp_gbnf = super::grammar::serialize::serialize(resp_grammar.as_ref().unwrap());
assert_ne!(
tool_gbnf, resp_gbnf,
"test fixture broken: tool and response grammars must be distinct"
);
let (effective, kind) = super::select_effective_grammar(
tool_grammar,
GrammarKind::ToolCallBodyRequired,
resp_grammar,
);
assert!(
effective.is_some(),
"helper must propagate the tool grammar"
);
assert_eq!(
kind,
GrammarKind::ToolCallBodyRequired,
"tool grammar MUST win and yield the caller-provided kind \
(here ToolCallBodyRequired, mirroring tool_choice=required) \
through the PRODUCTION helper (not a stand-in match)."
);
let chosen = effective.expect("Some");
let chosen_gbnf = super::grammar::serialize::serialize(&chosen);
assert_eq!(
chosen_gbnf, tool_gbnf,
"select_effective_grammar MUST return the tool grammar, not the \
response grammar. chosen_gbnf differs from tool_gbnf.\n\
chosen:\n{chosen_gbnf}\ntool:\n{tool_gbnf}\nresp:\n{resp_gbnf}"
);
assert_ne!(
chosen_gbnf, resp_gbnf,
"select_effective_grammar returned the response grammar \
(chosen == resp_gbnf); tool grammar must win."
);
}
#[test]
fn no_grammar_yields_default_kind_through_production_helper() {
let (effective, kind) =
super::select_effective_grammar(None, GrammarKind::ToolCallBodyAuto, None);
assert!(effective.is_none(), "no grammar means no constraint");
assert_eq!(
kind,
GrammarKind::default(),
"kind defaults to ResponseFormat when both grammars are None \
(caller-provided tool kind is dropped because no tool grammar exists)"
);
}
#[test]
fn tool_grammar_kind_for_auto_is_lazy() {
use super::super::schema::ToolChoiceValue;
assert_eq!(
super::tool_grammar_kind_for(&ToolChoiceValue::Auto),
GrammarKind::ToolCallBodyAuto,
"Auto must produce the LAZY tool-call body kind \
(engine arms awaiting_trigger=true on construction)"
);
}
#[test]
fn tool_grammar_kind_for_required_is_eager() {
use super::super::schema::ToolChoiceValue;
assert_eq!(
super::tool_grammar_kind_for(&ToolChoiceValue::Required),
GrammarKind::ToolCallBodyRequired,
"Required must produce the EAGER tool-call body kind \
(engine leaves awaiting_trigger=false on construction)"
);
}
#[test]
fn tool_grammar_kind_for_function_is_eager() {
use super::super::schema::ToolChoiceValue;
assert_eq!(
super::tool_grammar_kind_for(&ToolChoiceValue::Function("get_weather".into())),
GrammarKind::ToolCallBodyRequired,
"Function(name) must produce the EAGER tool-call body kind"
);
}
#[test]
fn tool_grammar_kind_for_none_is_default() {
use super::super::schema::ToolChoiceValue;
assert_eq!(
super::tool_grammar_kind_for(&ToolChoiceValue::None),
GrammarKind::default(),
"None must produce the Default kind (ResponseFormat); \
this branch is never read because tool_grammar is also None"
);
}
}
fn parse_slot_budget_exceeded(msg: &str) -> (u64, u64) {
fn extract_u64(msg: &str, key: &str) -> u64 {
let needle = format!("{key}=");
let Some(start) = msg.find(&needle) else {
return 0;
};
let after = &msg[start + needle.len()..];
let end = after
.find(|c: char| !c.is_ascii_digit())
.unwrap_or(after.len());
after[..end].parse::<u64>().unwrap_or(0)
}
(
extract_u64(msg, "needed_bytes"),
extract_u64(msg, "budget_bytes"),
)
}
fn queue_full_with_rate_limit_headers(state: &AppState) -> Response {
use axum::http::{header::HeaderName, HeaderValue};
let err = ApiError::queue_full();
let mut resp = err.into_response();
let cap = state.config.queue_capacity as u64;
let headers = resp.headers_mut();
if let Ok(v) = HeaderValue::from_str(&cap.to_string()) {
headers.insert(HeaderName::from_static("x-ratelimit-limit"), v);
}
headers.insert(
HeaderName::from_static("x-ratelimit-remaining"),
HeaderValue::from_static("0"),
);
headers.insert(
HeaderName::from_static("x-ratelimit-reset"),
HeaderValue::from_static("1"),
);
resp
}
fn chrono_seconds() -> i64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs() as i64)
.unwrap_or(0)
}
const BERT_MIN_SEQ_LEN: usize = 32;
pub async fn embeddings(
State(state): State<AppState>,
Json(req): Json<EmbeddingRequest>,
) -> Response {
if state.embedding_config.is_none() {
let loaded = match resolve_engine_for_request(&state, &req.model).await {
Ok(arc) => arc,
Err(resp) => return resp,
};
return chat_model_embeddings(loaded.engine.clone(), req).await;
}
let em = match state.embedding_config.as_ref() {
Some(em) => em.clone(),
None => {
return ApiError::model_not_loaded(&req.model).into_response();
}
};
if req.model != em.model_id {
return ApiError::model_not_loaded(&req.model).into_response();
}
let arch = match em.arch.as_ref() {
Some(a) => a.clone(),
None => {
return ApiError::generation_error(
"embedding model has no loaded weights (server startup did not eagerly load)"
.to_string(),
)
.into_response();
}
};
let hidden_size_native = arch.hidden_size();
let max_pos = arch.max_position_embeddings();
let want_base64 = match req.encoding_format.as_deref() {
None | Some("float") => false,
Some("base64") => true,
Some(other) => {
return ApiError::invalid_request(
format!("encoding_format='{other}' not supported (only 'float' or 'base64')"),
Some("encoding_format".into()),
)
.into_response();
}
};
if let Some(d) = req.dimensions {
if d != hidden_size_native {
return ApiError::invalid_request(
format!(
"model '{}' does not support custom output dimensions (native dim is {}; \
only the text-embedding-3 family supports `dimensions`)",
em.model_id, hidden_size_native
),
Some("dimensions".into()),
)
.into_response();
}
}
let inputs = req.input.into_vec();
if inputs.is_empty() {
return ApiError::invalid_request(
"input must be a non-empty string or array of strings".to_string(),
Some("input".into()),
)
.into_response();
}
let model_id = em.model_id.clone();
let tokenizer = em.tokenizer.clone();
let shared_registry = state.embedding_registry.clone();
let result = tokio::task::spawn_blocking(move || -> anyhow::Result<EmbeddingResponse> {
use crate::inference::models::bert::bert_gpu::apply_bert_full_forward_gpu;
use crate::inference::models::nomic_bert::apply_nomic_bert_full_forward_gpu;
use crate::serve::api::state::EmbeddingArch;
use mlx_native::{GraphExecutor, MlxDevice};
let device = MlxDevice::new()
.map_err(|e| anyhow::anyhow!("create MlxDevice for embedding forward: {e}"))?;
let executor = GraphExecutor::new(device);
let registry_arc = shared_registry.ok_or_else(|| {
anyhow::anyhow!("embedding registry not pre-warmed (server boot did not initialize it)")
})?;
let mut data: Vec<EmbeddingObject> = Vec::with_capacity(inputs.len());
let mut total_tokens: usize = 0;
for (i, input) in inputs.into_iter().enumerate() {
let raw_ids: Vec<u32> = tokenizer.encode(input.as_str(), true);
total_tokens += raw_ids.len();
let mut ids: Vec<u32> = raw_ids.to_vec();
if ids.len() < BERT_MIN_SEQ_LEN {
ids.resize(BERT_MIN_SEQ_LEN, 0u32);
}
if ids.len() > max_pos {
ids.truncate(max_pos);
}
let seq_len = ids.len() as u32;
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let ids_buf = device
.alloc_buffer(ids.len() * 4, mlx_native::DType::U32, vec![ids.len()])
.map_err(|e| anyhow::anyhow!("alloc ids buf: {e}"))?;
let s: &mut [u32] = unsafe {
std::slice::from_raw_parts_mut(ids_buf.contents_ptr() as *mut u32, ids.len())
};
s.copy_from_slice(&ids);
let mut session = executor
.begin()
.map_err(|e| anyhow::anyhow!("begin session: {e}"))?;
let mut registry_guard = registry_arc
.lock()
.map_err(|e| anyhow::anyhow!("embedding registry mutex poisoned: {e}"))?;
let valid_token_count = raw_ids.len().min(ids.len()) as u32;
let out = match &arch {
EmbeddingArch::Bert { config, weights } => apply_bert_full_forward_gpu(
session.encoder_mut(),
&mut registry_guard,
device,
&ids_buf,
None, weights,
config,
seq_len,
valid_token_count,
)?,
EmbeddingArch::NomicBert { config, weights } => {
apply_nomic_bert_full_forward_gpu(
session.encoder_mut(),
&mut registry_guard,
device,
&ids_buf,
None, weights,
config,
seq_len,
valid_token_count,
)?
}
};
session
.finish()
.map_err(|e| anyhow::anyhow!("session finish: {e}"))?;
drop(registry_guard);
let slice = out
.as_slice::<f32>()
.map_err(|e| anyhow::anyhow!("readback: {e}"))?;
let payload = if want_base64 {
let mut bytes = Vec::with_capacity(slice.len() * 4);
for v in slice {
bytes.extend_from_slice(&v.to_le_bytes());
}
use base64::Engine;
EmbeddingPayload::Base64(base64::engine::general_purpose::STANDARD.encode(&bytes))
} else {
EmbeddingPayload::Float(slice.to_vec())
};
data.push(EmbeddingObject {
object: "embedding",
embedding: payload,
index: i,
});
}
Ok(EmbeddingResponse {
object: "list",
data,
model: model_id,
usage: EmbeddingUsage {
prompt_tokens: total_tokens,
total_tokens,
},
})
})
.await;
match result {
Ok(Ok(resp)) => (StatusCode::OK, Json(resp)).into_response(),
Ok(Err(e)) => {
ApiError::generation_error(format!("embedding forward: {e:#}")).into_response()
}
Err(join_err) => {
ApiError::generation_error(format!("embedding worker panicked: {join_err}"))
.into_response()
}
}
}
async fn chat_model_embeddings(engine: super::engine::Engine, req: EmbeddingRequest) -> Response {
let want_base64 = match req.encoding_format.as_deref() {
None | Some("float") => false,
Some("base64") => true,
Some(other) => {
return ApiError::invalid_request(
format!("encoding_format='{other}' not supported (only 'float' or 'base64')"),
Some("encoding_format".into()),
)
.into_response();
}
};
let hidden_size = engine.hidden_size();
if let Some(d) = req.dimensions {
if d != hidden_size {
return ApiError::invalid_request(
format!(
"model '{}' does not support custom output dimensions (native dim is {}; \
only the text-embedding-3 family supports `dimensions`)",
engine.model_id(),
hidden_size
),
Some("dimensions".into()),
)
.into_response();
}
}
let inputs = req.input.into_vec();
if inputs.is_empty() {
return ApiError::invalid_request(
"input must be a non-empty string or array of strings".to_string(),
Some("input".into()),
)
.into_response();
}
let model_id = engine.model_id().to_string();
let mut data: Vec<EmbeddingObject> = Vec::with_capacity(inputs.len());
let mut total_tokens: usize = 0;
let bos_id: Option<u32> = probe_bos_token_id(engine.tokenizer());
for (i, input) in inputs.into_iter().enumerate() {
let encoded = match engine.tokenizer().encode(input.as_str(), false) {
Ok(e) => e,
Err(e) => {
return ApiError::invalid_request(
format!("input[{i}] tokenization failed: {e}"),
Some("input".into()),
)
.into_response();
}
};
let mut prompt_tokens: Vec<u32> = encoded.get_ids().to_vec();
if let Some(b) = bos_id {
prompt_tokens.insert(0, b);
}
if prompt_tokens.is_empty() {
return ApiError::invalid_request(
format!("input[{i}] tokenized to zero tokens (empty after preprocessing)"),
Some("input".into()),
)
.into_response();
}
total_tokens += prompt_tokens.len();
if let Err(e) = engine.try_admit_budget(prompt_tokens.len() as u32, 0) {
match e {
engine::EngineAdmitError::SlotBudgetExceeded {
needed_bytes,
budget_bytes,
} => {
tracing::info!(
needed_bytes,
budget_bytes,
"embeddings: ADR-040 §3.5 A5b pre-dispatch slot_budget_exceeded"
);
return ApiError::slot_budget_exceeded(needed_bytes, budget_bytes)
.into_response();
}
}
}
let embedding: Vec<f32> = match engine.embed(prompt_tokens).await {
Ok(v) => v,
Err(e) => {
let msg = format!("{e}");
if msg.contains("queue_full") {
return (
StatusCode::TOO_MANY_REQUESTS,
[(axum::http::header::RETRY_AFTER, "1")],
Json(serde_json::json!({
"error": {
"message": "engine queue full; retry shortly",
"type": "rate_limit_exceeded",
"param": null,
"code": "queue_full"
}
})),
)
.into_response();
}
if msg.contains("slot_budget_exceeded") {
let (needed, budget) = parse_slot_budget_exceeded(&msg);
return ApiError::slot_budget_exceeded(needed, budget).into_response();
}
if msg.contains(engine::QWEN35_NOT_IMPLEMENTED_SENTINEL) {
return ApiError::not_implemented(
engine::QWEN35_NOT_IMPLEMENTED_MESSAGE.to_string(),
)
.into_response();
}
if msg.contains(
crate::inference::models::qwen3vl_text::forward::QWEN3VL_TEXT_FORWARD_PENDING_SENTINEL,
) {
return ApiError::not_implemented(
crate::inference::models::qwen3vl_text::forward::QWEN3VL_TEXT_FORWARD_PENDING_MESSAGE
.to_string(),
)
.into_response();
}
return ApiError::generation_error(format!("chat-model embedding: {e:#}"))
.into_response();
}
};
if embedding.len() != hidden_size {
return ApiError::generation_error(format!(
"chat-model embedding length {} != hidden_size {}",
embedding.len(),
hidden_size
))
.into_response();
}
let payload = if want_base64 {
let mut bytes = Vec::with_capacity(embedding.len() * 4);
for v in &embedding {
bytes.extend_from_slice(&v.to_le_bytes());
}
use base64::Engine;
EmbeddingPayload::Base64(base64::engine::general_purpose::STANDARD.encode(&bytes))
} else {
EmbeddingPayload::Float(embedding)
};
data.push(EmbeddingObject {
object: "embedding",
embedding: payload,
index: i,
});
}
let resp = EmbeddingResponse {
object: "list",
data,
model: model_id,
usage: EmbeddingUsage {
prompt_tokens: total_tokens,
total_tokens,
},
};
(StatusCode::OK, Json(resp)).into_response()
}
fn build_kv_counter_block(
metric_name: &str,
help_text: &str,
rows: &[((String, String), [u64; 4])],
) -> String {
use crate::serve::api::state::KV_SPILL_OUTCOMES;
let mut s = String::new();
s.push_str("# HELP ");
s.push_str(metric_name);
s.push(' ');
s.push_str(help_text);
s.push('\n');
s.push_str("# TYPE ");
s.push_str(metric_name);
s.push_str(" counter\n");
for ((repo, quant), counts) in rows {
for (i, outcome_label) in KV_SPILL_OUTCOMES.iter().enumerate() {
s.push_str(metric_name);
s.push_str("{repo=\"");
s.push_str(repo);
s.push_str("\",quant=\"");
s.push_str(quant);
s.push_str("\",outcome=\"");
s.push_str(outcome_label);
s.push_str("\"} ");
s.push_str(&counts[i].to_string());
s.push('\n');
}
}
s
}
pub async fn metrics(State(state): State<AppState>) -> Response {
use std::sync::atomic::Ordering;
let m = &state.metrics;
let kv_spill_rows = state.kv_spill_counters.snapshot_spills();
let kv_restore_rows = state.kv_spill_counters.snapshot_restores();
let ready = if state.is_ready_for_gen() { 1 } else { 0 };
let pool_stats_for_metrics = state.pool.read().ok().map(|m| m.pool_stats());
let (model_loaded, pool_loaded_models, pool_resident_bytes, pool_memory_budget_bytes) =
match pool_stats_for_metrics {
Some(stats) => (
if stats.loaded_count > 0 { 1 } else { 0 },
stats.loaded_count as u64,
stats.total_resident_bytes,
stats.memory_budget_bytes,
),
None => (0, 0, 0, 0),
};
let kv_spills_block = build_kv_counter_block(
"hf2q_pool_kv_spills_total",
"Count of KV-cache spill operations on engine eviction.",
&kv_spill_rows,
);
let kv_restores_block = build_kv_counter_block(
"hf2q_pool_kv_restores_total",
"Count of KV-cache restore operations on engine admission.",
&kv_restore_rows,
);
let kv_quarantine_counts = state.kv_spill_counters.snapshot_quarantines();
let kv_quarantined_block = {
use crate::serve::api::state::KV_QUARANTINE_REASONS;
let mut s = String::with_capacity(256);
s.push_str("# HELP hf2q_kv_quarantined_total Count of KV blocks moved to kv-quarantine/, by reason.\n");
s.push_str("# TYPE hf2q_kv_quarantined_total counter\n");
for (i, label) in KV_QUARANTINE_REASONS.iter().enumerate() {
s.push_str("hf2q_kv_quarantined_total{reason=\"");
s.push_str(label);
s.push_str("\"} ");
s.push_str(&kv_quarantine_counts[i].to_string());
s.push('\n');
}
s
};
let kv_eviction_counts = state.kv_spill_counters.snapshot_evictions();
let kv_evictions_block = {
use crate::serve::api::state::KV_EVICTION_TRIGGERS;
let mut s = String::with_capacity(192);
s.push_str("# HELP hf2q_kv_cache_evictions_total Count of KV blocks evicted from the on-disk cache, by trigger.\n");
s.push_str("# TYPE hf2q_kv_cache_evictions_total counter\n");
for (i, label) in KV_EVICTION_TRIGGERS.iter().enumerate() {
s.push_str("hf2q_kv_cache_evictions_total{trigger=\"");
s.push_str(label);
s.push_str("\"} ");
s.push_str(&kv_eviction_counts[i].to_string());
s.push('\n');
}
s
};
let (kv_lcp_lookups, kv_lcp_detected) = state.kv_spill_counters.snapshot_lcp();
let kv_lcp_block = {
let mut s = String::with_capacity(384);
s.push_str("# HELP hf2q_kv_lcp_lookups_total Total LCP probes after PromptCache full-equality miss (ADR-017 Phase E.a iter-2).\n");
s.push_str("# TYPE hf2q_kv_lcp_lookups_total counter\n");
s.push_str("hf2q_kv_lcp_lookups_total ");
s.push_str(&kv_lcp_lookups.to_string());
s.push('\n');
s.push_str("# HELP hf2q_kv_lcp_detected_total Probes that found a non-trivial partial-prefix opportunity (0 < K < N). Iter-2 reports only — partial-prefill resume path stays OFF.\n");
s.push_str("# TYPE hf2q_kv_lcp_detected_total counter\n");
s.push_str("hf2q_kv_lcp_detected_total ");
s.push_str(&kv_lcp_detected.to_string());
s.push('\n');
s
};
let (kv_cache_bytes_on_disk, kv_cache_blocks_total) =
if let Some(store) = state.kv_disk_store.as_ref() {
(
store.index().total_bytes_on_disk(),
store.index().block_count() as u64,
)
} else {
(0u64, 0u64)
};
let body = format!(
"\
# HELP hf2q_uptime_seconds Process uptime in seconds since bind.\n\
# TYPE hf2q_uptime_seconds gauge\n\
hf2q_uptime_seconds {uptime}\n\
# HELP hf2q_ready 1 if generation endpoints are ready, 0 during warmup.\n\
# TYPE hf2q_ready gauge\n\
hf2q_ready {ready}\n\
# HELP hf2q_model_loaded 1 if a model is loaded, 0 if HTTP-only backbone.\n\
# TYPE hf2q_model_loaded gauge\n\
hf2q_model_loaded {model}\n\
# HELP hf2q_pool_loaded_models Number of models currently resident in HotSwapManager.\n\
# TYPE hf2q_pool_loaded_models gauge\n\
hf2q_pool_loaded_models {pool_loaded}\n\
# HELP hf2q_pool_resident_bytes Total bytes of GGUF data resident in HotSwapManager.\n\
# TYPE hf2q_pool_resident_bytes gauge\n\
hf2q_pool_resident_bytes {pool_resident}\n\
# HELP hf2q_pool_memory_budget_bytes Memory budget for HotSwapManager (80% of unified RAM by default).\n\
# TYPE hf2q_pool_memory_budget_bytes gauge\n\
hf2q_pool_memory_budget_bytes {pool_budget}\n\
{kv_spills_block}\
{kv_restores_block}\
{kv_quarantined_block}\
# HELP hf2q_kv_cache_bytes_on_disk Sum of envelope-file bytes currently indexed in the on-disk KV cache.\n\
# TYPE hf2q_kv_cache_bytes_on_disk gauge\n\
hf2q_kv_cache_bytes_on_disk {kv_cache_bytes}\n\
# HELP hf2q_kv_cache_blocks_total Number of KV envelopes currently indexed in the on-disk cache.\n\
# TYPE hf2q_kv_cache_blocks_total gauge\n\
hf2q_kv_cache_blocks_total {kv_cache_blocks}\n\
{kv_evictions_block}\
{kv_lcp_block}\
# HELP hf2q_requests_total Total HTTP requests reaching a handler (post-auth).\n\
# TYPE hf2q_requests_total counter\n\
hf2q_requests_total {req_total}\n\
# HELP hf2q_requests_rejected_total HTTP requests rejected at handler (auth/malformed).\n\
# TYPE hf2q_requests_rejected_total counter\n\
hf2q_requests_rejected_total {req_rej}\n\
# HELP hf2q_chat_completions_started Chat completion generations started.\n\
# TYPE hf2q_chat_completions_started counter\n\
hf2q_chat_completions_started {chat_start}\n\
# HELP hf2q_chat_completions_completed Chat completion generations completed successfully.\n\
# TYPE hf2q_chat_completions_completed counter\n\
hf2q_chat_completions_completed {chat_done}\n\
# HELP hf2q_chat_completions_queue_full Chat completions rejected with 429 queue_full.\n\
# TYPE hf2q_chat_completions_queue_full counter\n\
hf2q_chat_completions_queue_full {chat_429}\n\
# HELP hf2q_sse_cancellations SSE streams cancelled by client drop mid-generation.\n\
# TYPE hf2q_sse_cancellations counter\n\
hf2q_sse_cancellations {sse_cancel}\n\
# HELP hf2q_decode_tokens_total Tokens decoded across all completions.\n\
# TYPE hf2q_decode_tokens_total counter\n\
hf2q_decode_tokens_total {decode_tok}\n\
# HELP hf2q_prompt_tokens_total Prompt tokens ingested across all completions.\n\
# TYPE hf2q_prompt_tokens_total counter\n\
hf2q_prompt_tokens_total {prompt_tok}\n\
",
uptime = state.uptime_seconds(),
ready = ready,
model = model_loaded,
pool_loaded = pool_loaded_models,
pool_resident = pool_resident_bytes,
pool_budget = pool_memory_budget_bytes,
kv_spills_block = kv_spills_block,
kv_restores_block = kv_restores_block,
kv_quarantined_block = kv_quarantined_block,
kv_evictions_block = kv_evictions_block,
kv_lcp_block = kv_lcp_block,
kv_cache_bytes = kv_cache_bytes_on_disk,
kv_cache_blocks = kv_cache_blocks_total,
req_total = m.requests_total.load(Ordering::Relaxed),
req_rej = m.requests_rejected_total.load(Ordering::Relaxed),
chat_start = m.chat_completions_started.load(Ordering::Relaxed),
chat_done = m.chat_completions_completed.load(Ordering::Relaxed),
chat_429 = m.chat_completions_queue_full.load(Ordering::Relaxed),
sse_cancel = m.sse_cancellations.load(Ordering::Relaxed),
decode_tok = m.decode_tokens_total.load(Ordering::Relaxed),
prompt_tok = m.prompt_tokens_total.load(Ordering::Relaxed),
);
let mut resp = (StatusCode::OK, body).into_response();
resp.headers_mut().insert(
axum::http::header::CONTENT_TYPE,
axum::http::HeaderValue::from_static("text/plain; version=0.0.4; charset=utf-8"),
);
resp
}
pub async fn health(State(state): State<AppState>) -> impl IntoResponse {
state
.metrics
.requests_total
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let snapshot: Vec<_> = state
.pool
.read()
.ok()
.map(|m| m.snapshot_engines())
.unwrap_or_default();
let (model, context_length) = match snapshot.last() {
Some(le) => (
Some(le.engine.model_id().to_string()),
le.engine.context_length(),
),
None => (None, None),
};
let resp = HealthResponse {
status: "ok".to_string(),
model,
backend: "mlx-native",
context_length,
uptime_seconds: state.uptime_seconds(),
};
(StatusCode::OK, Json(resp))
}
pub async fn readyz(State(state): State<AppState>) -> impl IntoResponse {
state
.metrics
.requests_total
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
if state.is_ready_for_gen() {
(
StatusCode::OK,
Json(ReadyzResponse {
ready: true,
detail: "ready",
}),
)
.into_response()
} else {
let mut resp = (
StatusCode::SERVICE_UNAVAILABLE,
Json(ReadyzResponse {
ready: false,
detail: "warming up",
}),
)
.into_response();
resp.headers_mut().insert(
axum::http::header::RETRY_AFTER,
axum::http::HeaderValue::from_static("1"),
);
resp
}
}
pub async fn shutdown(State(state): State<AppState>) -> Response {
state
.metrics
.requests_total
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let pre_shutdown_queue_depth = state
.kv_spiller
.as_ref()
.map(|sp| sp.pending_writer_queue_depth())
.unwrap_or(0);
let kv_persist_enabled = state.kv_spiller.is_some();
let pid = std::process::id();
tracing::info!(
pid = pid,
kv_persist_enabled = kv_persist_enabled,
pre_shutdown_queue_depth = pre_shutdown_queue_depth,
"ADR-017: POST /shutdown received; raising SIGTERM"
);
let raise_rc = unsafe { libc::raise(libc::SIGTERM) };
let raise_ok = raise_rc == 0;
if !raise_ok {
let errno = std::io::Error::last_os_error();
tracing::warn!(
errno = %errno,
"ADR-017: libc::raise(SIGTERM) returned non-zero; \
graceful-shutdown drain may not run"
);
}
let body = serde_json::json!({
"status": "accepted",
"pid": pid,
"kv_persist_enabled": kv_persist_enabled,
"pre_shutdown_queue_depth": pre_shutdown_queue_depth,
"raise_sigterm_rc": raise_rc,
"note": "graceful drain runs after this response; poll for process exit",
});
(StatusCode::ACCEPTED, Json(body)).into_response()
}
pub async fn list_models(State(state): State<AppState>) -> Response {
state
.metrics
.requests_total
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let cache_dir = state.config.cache_dir.clone();
let mut models =
match tokio::task::spawn_blocking(move || scan_cache_dir(cache_dir.as_deref())).await {
Ok(Ok(models)) => models,
Ok(Err(e)) => {
tracing::warn!(error = %e, "model cache scan failed");
return ApiError::internal_error().into_response();
}
Err(e) => {
tracing::error!(error = %e, "spawn_blocking panicked in list_models");
return ApiError::internal_error().into_response();
}
};
let pool_snapshot: Vec<_> = state
.pool
.read()
.ok()
.map(|m| m.snapshot_engines())
.unwrap_or_default();
for le in pool_snapshot.iter() {
let loaded_id = le.engine.model_id().to_string();
let info = le.engine.info();
match models.iter_mut().find(|m| m.id == loaded_id) {
Some(m) => {
m.loaded = true;
enrich_model_object_from_load_info(m, info);
}
None => {
models.insert(0, model_object_from_load_info(loaded_id, info, true));
}
}
}
if let Some(em) = state.embedding_config.as_ref() {
if !models.iter().any(|m| m.id == em.model_id) {
models.insert(
0,
ModelObject {
id: em.model_id.clone(),
object: "model",
created: chrono_seconds(),
owned_by: "hf2q",
context_length: em.arch.as_ref().map(|a| a.max_position_embeddings()),
quant_type: None,
backend: Some("mlx-native"),
loaded: true,
arch: None,
max_context_length: None,
provenance: None,
moe_experts: None,
moe_experts_per_tok: None,
sliding_window: None,
kv_spill_active: None,
quant_bpw: None,
},
);
}
}
if let Some(m) = state.mmproj.as_ref() {
if !models.iter().any(|existing| existing.id == m.model_id) {
models.insert(
0,
ModelObject {
id: m.model_id.clone(),
object: "model",
created: chrono_seconds(),
owned_by: "hf2q",
context_length: None,
quant_type: None,
backend: Some("mlx-native"),
loaded: true,
arch: None,
max_context_length: None,
provenance: None,
moe_experts: None,
moe_experts_per_tok: None,
sliding_window: None,
kv_spill_active: None,
quant_bpw: None,
},
);
}
}
let resp = ModelListResponse {
object: "list",
data: models,
};
(StatusCode::OK, Json(resp)).into_response()
}
pub async fn get_model(
State(state): State<AppState>,
AxPath(model_id): AxPath<String>,
) -> Response {
let cache_dir = state.config.cache_dir.clone();
let all = match tokio::task::spawn_blocking(move || scan_cache_dir(cache_dir.as_deref())).await
{
Ok(Ok(m)) => m,
Ok(Err(_)) | Err(_) => return ApiError::internal_error().into_response(),
};
match all.into_iter().find(|m| m.id == model_id) {
Some(m) => (StatusCode::OK, Json(m)).into_response(),
None => ApiError::model_not_found(&model_id).into_response(),
}
}
pub(crate) fn scan_cache_dir(cache_dir: Option<&Path>) -> std::io::Result<Vec<ModelObject>> {
let Some(dir) = cache_dir else {
return Ok(Vec::new());
};
if !dir.is_dir() {
return Ok(Vec::new());
}
let mut out = Vec::new();
visit_dir(dir, &mut out, 0, 6)?;
out.sort_by(|a, b| a.id.cmp(&b.id));
Ok(out)
}
fn visit_dir(
dir: &Path,
out: &mut Vec<ModelObject>,
depth: usize,
max_depth: usize,
) -> std::io::Result<()> {
if depth > max_depth {
return Ok(());
}
for entry in std::fs::read_dir(dir)? {
let entry = entry?;
let path = entry.path();
let ft = entry.file_type()?;
if ft.is_dir() {
if !ft.is_symlink() {
visit_dir(&path, out, depth + 1, max_depth)?;
}
} else if ft.is_file() {
if path.extension().and_then(|s| s.to_str()) == Some("gguf") {
if let Some(obj) = inspect_gguf(&path) {
out.push(obj);
}
}
}
}
Ok(())
}
fn inspect_gguf(path: &Path) -> Option<ModelObject> {
use mlx_native::gguf::GgufFile;
let gguf = match GgufFile::open(path) {
Ok(g) => g,
Err(e) => {
tracing::warn!(
path = %path.display(),
error = %e,
"skipping malformed GGUF in cache scan"
);
return None;
}
};
let stem = path.file_stem()?.to_string_lossy().into_owned();
let context_length = context_length_for_arch(&gguf);
let quant_type = infer_quant_type(&gguf);
let created = std::fs::metadata(path)
.ok()
.and_then(|m| m.modified().ok())
.and_then(|t| t.duration_since(std::time::UNIX_EPOCH).ok())
.map(|d| d.as_secs() as i64)
.unwrap_or(0);
Some(ModelObject {
id: stem,
object: "model",
created,
owned_by: "hf2q",
context_length,
quant_type,
backend: Some("mlx-native"),
loaded: false,
arch: None,
max_context_length: None,
provenance: None,
moe_experts: None,
moe_experts_per_tok: None,
sliding_window: None,
kv_spill_active: None,
quant_bpw: None,
})
}
fn model_object_from_load_info(
id: String,
info: &crate::serve::load_info::LoadInfo,
loaded: bool,
) -> ModelObject {
ModelObject {
id,
object: "model",
created: chrono_seconds(),
owned_by: "hf2q",
context_length: info.max_context_length.map(|v| v as usize),
quant_type: info.quant_label.clone(),
backend: Some("mlx-native"),
loaded,
arch: Some(info.arch_str.clone()),
max_context_length: info.max_context_length.map(u64::from),
provenance: Some(provenance_label(&info.provenance)),
moe_experts: info.moe.map(|m| m.n_experts),
moe_experts_per_tok: info.moe.map(|m| m.n_experts_per_tok),
sliding_window: info.sliding_window,
kv_spill_active: Some(info.kv_spill_active),
quant_bpw: info.quant_bpw,
}
}
fn enrich_model_object_from_load_info(
m: &mut ModelObject,
info: &crate::serve::load_info::LoadInfo,
) {
m.arch = Some(info.arch_str.clone());
m.max_context_length = info.max_context_length.map(u64::from);
m.provenance = Some(provenance_label(&info.provenance));
m.moe_experts = info.moe.map(|moe| moe.n_experts);
m.moe_experts_per_tok = info.moe.map(|moe| moe.n_experts_per_tok);
m.sliding_window = info.sliding_window;
m.kv_spill_active = Some(info.kv_spill_active);
m.quant_bpw = info.quant_bpw;
if m.quant_type.is_none() {
m.quant_type = info.quant_label.clone();
}
}
fn provenance_label(p: &crate::core::provenance::Provenance) -> &'static str {
use crate::core::provenance::Provenance;
match p {
Provenance::Hf2q { .. } => "hf2q",
Provenance::External => "external",
}
}
fn context_length_for_arch(gguf: &mlx_native::gguf::GgufFile) -> Option<usize> {
let arch = gguf.metadata_string("general.architecture")?;
let key = format!("{arch}.context_length");
gguf.metadata_u32(&key).map(|v| v as usize)
}
fn infer_quant_type(gguf: &mlx_native::gguf::GgufFile) -> Option<String> {
crate::serve::load_info::infer_quant_label(gguf)
}
#[cfg(test)]
pub(crate) fn test_inspect_gguf(path: &Path) -> Option<ModelObject> {
inspect_gguf(path)
}
#[cfg(test)]
pub(crate) fn test_scan(dir: &Path) -> std::io::Result<Vec<ModelObject>> {
scan_cache_dir(Some(dir))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn scan_missing_dir_returns_empty() {
let tmp = std::env::temp_dir().join("hf2q-test-does-not-exist-xyz");
let result = scan_cache_dir(Some(&tmp)).unwrap();
assert!(result.is_empty());
}
#[test]
fn scan_none_cache_dir_returns_empty() {
let result = scan_cache_dir(None).unwrap();
assert!(result.is_empty());
}
#[test]
fn scan_empty_dir_returns_empty() {
let tmp = tempdir_for("hf2q-scan-empty");
let result = scan_cache_dir(Some(&tmp)).unwrap();
assert!(result.is_empty());
std::fs::remove_dir_all(&tmp).ok();
}
#[test]
fn scan_skips_non_gguf_files() {
let tmp = tempdir_for("hf2q-scan-skip-nongguf");
std::fs::write(tmp.join("readme.txt"), "hello").unwrap();
std::fs::write(tmp.join("data.bin"), [0u8, 1, 2, 3]).unwrap();
let result = scan_cache_dir(Some(&tmp)).unwrap();
assert!(result.is_empty());
std::fs::remove_dir_all(&tmp).ok();
}
#[test]
fn scan_skips_malformed_gguf_but_succeeds() {
let tmp = tempdir_for("hf2q-scan-malformed");
std::fs::write(tmp.join("fake.gguf"), b"not a real gguf file").unwrap();
let result = scan_cache_dir(Some(&tmp)).unwrap();
assert!(result.is_empty(), "malformed GGUF should be skipped");
std::fs::remove_dir_all(&tmp).ok();
}
#[test]
fn scan_is_deterministic_ordering() {
let tmp = tempdir_for("hf2q-scan-determ");
std::fs::create_dir_all(tmp.join("a")).unwrap();
std::fs::create_dir_all(tmp.join("b")).unwrap();
let result = scan_cache_dir(Some(&tmp)).unwrap();
assert!(result.is_empty());
std::fs::remove_dir_all(&tmp).ok();
}
fn tempdir_for(tag: &str) -> std::path::PathBuf {
let pid = std::process::id();
let nanos = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.subsec_nanos();
let p = std::env::temp_dir().join(format!("{tag}-{pid}-{nanos}"));
std::fs::create_dir_all(&p).unwrap();
p
}
fn populated_qwen35_load_info() -> crate::serve::load_info::LoadInfo {
use crate::core::provenance::Provenance;
use crate::serve::load_info::{
ArchFamily, ChatTemplateSource, LoadInfo, MoeShape, TokenizerSource,
};
use std::path::PathBuf;
use std::time::Duration;
LoadInfo {
model_id: "Qwen3.6-27B-A3B-DWQ46-MoE".to_string(),
arch_str: "qwen35moe".to_string(),
arch_family: ArchFamily::Qwen35,
model_path: PathBuf::from("/cache/qwen35-27b-moe.gguf"),
on_disk_bytes: 16_000_000_000,
backend_chip: "Apple M5 Max".to_string(),
backend: "mlx-native",
n_layers: 64,
hidden_size: 4096,
vocab_size: 151_936,
n_attention_heads: 16,
n_key_value_heads: 4,
head_dim: 128,
sliding_window: None,
full_attention_interval: Some(4),
max_context_length: Some(262_144),
moe: Some(MoeShape {
n_experts: 128,
n_experts_per_tok: 8,
}),
quant_label: Some("Q4_K".to_string()),
quant_bpw: Some(4.55),
tokenizer_source: TokenizerSource::GgufEmbedded,
eos_token_ids: vec![151_645],
bos_token_id: Some(151_643),
chat_template_source: ChatTemplateSource::GgufEmbedded,
provenance: Provenance::Hf2q {
producer_version: "hf2q 0.1.0".to_string(),
source_sha256: "abcd".repeat(16),
mmproj_sha256: None,
},
vision_projector: None,
load_wall_clock: Duration::from_secs_f64(6.84),
resident_weight_bytes: Some(14_000_000_000),
kv_cache_budget_bytes: Some(4 * 1024 * 1024 * 1024),
kv_spill_active: false,
tq_kv_active: false,
kv_bytes_per_token_override: None,
}
}
fn populated_gemma4_load_info() -> crate::serve::load_info::LoadInfo {
use crate::core::provenance::Provenance;
use crate::serve::load_info::{ArchFamily, ChatTemplateSource, LoadInfo, TokenizerSource};
use std::path::PathBuf;
use std::time::Duration;
LoadInfo {
model_id: "gemma-4-27b-it-Q4_K_M".to_string(),
arch_str: "gemma4".to_string(),
arch_family: ArchFamily::Gemma4,
model_path: PathBuf::from("/cache/gemma-4-27b-it-Q4_K_M.gguf"),
on_disk_bytes: 18_000_000_000,
backend_chip: "Apple M5 Max".to_string(),
backend: "mlx-native",
n_layers: 62,
hidden_size: 5376,
vocab_size: 262_144,
n_attention_heads: 32,
n_key_value_heads: 16,
head_dim: 128,
sliding_window: Some(4096),
full_attention_interval: None,
max_context_length: Some(131_072),
moe: None,
quant_label: Some("Q4_K".to_string()),
quant_bpw: Some(4.83),
tokenizer_source: TokenizerSource::HfTokenizerJson {
path: PathBuf::from("/cache/tokenizer.json"),
},
eos_token_ids: vec![1, 106],
bos_token_id: Some(2),
chat_template_source: ChatTemplateSource::GgufEmbedded,
provenance: Provenance::External,
vision_projector: None,
load_wall_clock: Duration::from_secs_f64(2.41),
resident_weight_bytes: Some(17_000_000_000),
kv_cache_budget_bytes: None,
kv_spill_active: true,
tq_kv_active: false,
kv_bytes_per_token_override: None,
}
}
#[test]
fn provenance_label_maps_external_and_hf2q() {
use crate::core::provenance::Provenance;
assert_eq!(provenance_label(&Provenance::External), "external");
assert_eq!(
provenance_label(&Provenance::Hf2q {
producer_version: "hf2q 0.1.0".into(),
source_sha256: "0".repeat(64),
mmproj_sha256: None,
}),
"hf2q"
);
}
#[test]
fn model_object_from_load_info_populates_qwen35_moe_fields() {
let info = populated_qwen35_load_info();
let m = model_object_from_load_info("Qwen3.6-27B-A3B-DWQ46-MoE".into(), &info, true);
assert_eq!(m.id, "Qwen3.6-27B-A3B-DWQ46-MoE");
assert_eq!(m.object, "model");
assert_eq!(m.owned_by, "hf2q");
assert_eq!(m.backend, Some("mlx-native"));
assert!(m.loaded);
assert_eq!(m.context_length, Some(262_144));
assert_eq!(m.quant_type.as_deref(), Some("Q4_K"));
assert_eq!(m.arch.as_deref(), Some("qwen35moe"));
assert_eq!(m.max_context_length, Some(262_144));
assert_eq!(m.provenance, Some("hf2q"));
assert_eq!(m.moe_experts, Some(128));
assert_eq!(m.moe_experts_per_tok, Some(8));
assert_eq!(m.sliding_window, None);
assert_eq!(m.kv_spill_active, Some(false));
assert!(
m.quant_bpw
.map(|v| (v - 4.55).abs() < 1e-3)
.unwrap_or(false),
"expected quant_bpw ≈ 4.55, got {:?}",
m.quant_bpw
);
}
#[test]
fn model_object_from_load_info_populates_gemma4_dense_fields() {
let info = populated_gemma4_load_info();
let m = model_object_from_load_info("gemma-4-27b-it-Q4_K_M".into(), &info, true);
assert_eq!(m.id, "gemma-4-27b-it-Q4_K_M");
assert_eq!(m.arch.as_deref(), Some("gemma4"));
assert_eq!(m.provenance, Some("external"));
assert_eq!(m.moe_experts, None);
assert_eq!(m.moe_experts_per_tok, None);
assert_eq!(m.sliding_window, Some(4096));
assert_eq!(m.kv_spill_active, Some(true));
assert!(
m.quant_bpw
.map(|v| (v - 4.83).abs() < 1e-3)
.unwrap_or(false),
"expected quant_bpw ≈ 4.83, got {:?}",
m.quant_bpw
);
}
#[test]
fn enrich_model_object_from_load_info_overlays_new_fields_only() {
let preexisting_created = 1_700_000_000_i64;
let mut m = ModelObject {
id: "Qwen3.6-27B-A3B-DWQ46-MoE".into(),
object: "model",
created: preexisting_created,
owned_by: "hf2q",
context_length: Some(131_072),
quant_type: Some("Q4_K".into()),
backend: Some("mlx-native"),
loaded: true,
arch: None,
max_context_length: None,
provenance: None,
moe_experts: None,
moe_experts_per_tok: None,
sliding_window: None,
kv_spill_active: None,
quant_bpw: None,
};
let info = populated_qwen35_load_info();
enrich_model_object_from_load_info(&mut m, &info);
assert_eq!(m.id, "Qwen3.6-27B-A3B-DWQ46-MoE");
assert_eq!(m.created, preexisting_created);
assert_eq!(m.context_length, Some(131_072));
assert!(m.loaded);
assert_eq!(m.quant_type.as_deref(), Some("Q4_K"));
assert_eq!(m.arch.as_deref(), Some("qwen35moe"));
assert_eq!(m.max_context_length, Some(262_144));
assert_eq!(m.provenance, Some("hf2q"));
assert_eq!(m.moe_experts, Some(128));
assert_eq!(m.moe_experts_per_tok, Some(8));
assert_eq!(m.sliding_window, None);
assert_eq!(m.kv_spill_active, Some(false));
}
#[test]
fn enrich_model_object_fills_quant_type_when_cache_scan_missed_it() {
let mut m = ModelObject {
id: "weird-model".into(),
object: "model",
created: 0,
owned_by: "hf2q",
context_length: None,
quant_type: None, backend: Some("mlx-native"),
loaded: true,
arch: None,
max_context_length: None,
provenance: None,
moe_experts: None,
moe_experts_per_tok: None,
sliding_window: None,
kv_spill_active: None,
quant_bpw: None,
};
let info = populated_qwen35_load_info();
enrich_model_object_from_load_info(&mut m, &info);
assert_eq!(m.quant_type.as_deref(), Some("Q4_K"));
}
#[test]
fn model_object_from_load_info_serializes_with_new_fields() {
let info = populated_gemma4_load_info();
let m = model_object_from_load_info("gemma-4-27b-it-Q4_K_M".into(), &info, true);
let v = serde_json::to_value(&m).expect("serialize ModelObject");
assert_eq!(v["id"], "gemma-4-27b-it-Q4_K_M");
assert_eq!(v["arch"], "gemma4");
assert_eq!(v["max_context_length"], 131_072);
assert_eq!(v["provenance"], "external");
assert!(
v.get("moe_experts").is_none(),
"dense → moe_experts skipped"
);
assert!(
v.get("moe_experts_per_tok").is_none(),
"dense → moe_experts_per_tok skipped"
);
assert_eq!(v["sliding_window"], 4096);
assert_eq!(v["kv_spill_active"], true);
let bpw = v["quant_bpw"].as_f64().expect("quant_bpw f64");
assert!((bpw - 4.83_f64).abs() < 1e-3);
}
}
#[cfg(test)]
mod multimodal_tests {
use super::*;
use crate::inference::vision::mmproj::{MmprojConfig, ProjectorType};
use crate::serve::api::schema::{ChatMessage, ContentPart, ImageUrl, MessageContent};
use crate::serve::api::state::LoadedMmproj;
fn synthetic_png_data_uri() -> String {
use base64::Engine;
use image::{ImageBuffer, ImageFormat, Rgb, RgbImage};
use std::io::Cursor;
let img: RgbImage = ImageBuffer::from_fn(4, 4, |_x, _y| Rgb([200u8, 100, 50]));
let mut buf: Vec<u8> = Vec::new();
img.write_to(&mut Cursor::new(&mut buf), ImageFormat::Png)
.expect("encode png");
let b64 = base64::engine::general_purpose::STANDARD.encode(&buf);
format!("data:image/png;base64,{b64}")
}
fn synthetic_mmproj() -> LoadedMmproj {
use std::sync::Arc;
let cfg = MmprojConfig {
image_size: 8,
patch_size: 1,
num_patches_side: 8,
hidden_size: 1152,
intermediate_size: 4304,
num_attention_heads: 16,
num_hidden_layers: 27,
layer_norm_eps: 1e-6,
projector: ProjectorType::Mlp,
image_mean: [0.5, 0.5, 0.5],
image_std: [0.5, 0.5, 0.5],
spatial_merge_size: None,
projection_dim: None,
deepstack_indexes: None,
};
let device = mlx_native::MlxDevice::new().expect("create device");
let weights = crate::inference::vision::mmproj_weights::LoadedMmprojWeights::empty(device);
LoadedMmproj {
gguf_path: "/tmp/synthetic-mmproj.gguf".into(),
config: cfg,
arch: crate::inference::vision::mmproj::ArchProfile::ClipClassic,
weights: Arc::new(weights),
model_id: "synthetic-mmproj".into(),
}
}
fn user_text(text: &str) -> ChatMessage {
ChatMessage {
role: "user".into(),
content: Some(MessageContent::Text(text.into())),
name: None,
tool_calls: None,
tool_call_id: None,
reasoning_content: None,
}
}
fn user_with_image(text: &str, image_url: &str) -> ChatMessage {
ChatMessage {
role: "user".into(),
content: Some(MessageContent::Parts(vec![
ContentPart::Text { text: text.into() },
ContentPart::ImageUrl {
image_url: ImageUrl {
url: image_url.into(),
detail: None,
},
},
])),
name: None,
tool_calls: None,
tool_call_id: None,
reasoning_content: None,
}
}
#[test]
fn text_only_returns_empty_and_does_not_need_mmproj() {
let msgs = vec![user_text("hi")];
let got = process_multimodal_content(&msgs, None).expect("ok");
assert!(got.is_empty());
}
#[test]
fn images_without_mmproj_return_400_no_mmproj_loaded() {
let uri = synthetic_png_data_uri();
let msgs = vec![user_with_image("describe this", &uri)];
let resp =
process_multimodal_content(&msgs, None).expect_err("image without mmproj should 400");
assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
}
#[test]
fn single_image_preprocesses_to_chw_f32_tensor() {
use crate::inference::vision::vit_gpu::VisionInput;
let uri = synthetic_png_data_uri();
let msgs = vec![user_with_image("describe", &uri)];
let mmproj = synthetic_mmproj();
let got = process_multimodal_content(&msgs, Some(&mmproj)).expect("ok");
assert_eq!(got.len(), 1);
let img = match &got[0] {
VisionInput::Siglip49(img) => img,
VisionInput::Gemma4v(_) => panic!("expected Siglip49 variant for ClipClassic mmproj"),
};
assert_eq!(img.target_size, 8);
assert_eq!(img.pixel_values.len(), 3 * 8 * 8);
assert_eq!(img.source_label, "image/png");
let r = img.pixel_values[0];
let g = img.pixel_values[64];
let b = img.pixel_values[128];
assert!((r - 0.569).abs() < 0.05, "R approx 0.569, got {r}");
assert!((g - (-0.216)).abs() < 0.05, "G approx -0.216, got {g}");
assert!((b - (-0.608)).abs() < 0.05, "B approx -0.608, got {b}");
}
#[test]
fn multiple_images_preserve_message_order() {
use crate::inference::vision::vit_gpu::VisionInput;
let uri = synthetic_png_data_uri();
let msgs = vec![
user_with_image("first", &uri),
user_text("middle"),
user_with_image("second", &uri),
];
let mmproj = synthetic_mmproj();
let got = process_multimodal_content(&msgs, Some(&mmproj)).expect("ok");
assert_eq!(got.len(), 2);
for input in &got {
match input {
VisionInput::Siglip49(img) => assert_eq!(img.source_label, "image/png"),
VisionInput::Gemma4v(_) => panic!("expected Siglip49 variant"),
}
}
}
#[test]
fn malformed_url_returns_400_with_location() {
let msgs = vec![user_with_image("x", "not-a-url")];
let mmproj = synthetic_mmproj();
let resp =
process_multimodal_content(&msgs, Some(&mmproj)).expect_err("bad URL should 400");
assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
}
fn synthetic_mmproj_gemma4v() -> LoadedMmproj {
use std::sync::Arc;
let cfg = MmprojConfig {
image_size: 896,
patch_size: 16,
num_patches_side: 56,
hidden_size: 1152,
intermediate_size: 4304,
num_attention_heads: 16,
num_hidden_layers: 27,
layer_norm_eps: 1e-6,
projector: ProjectorType::Mlp,
image_mean: [0.5, 0.5, 0.5],
image_std: [0.5, 0.5, 0.5],
spatial_merge_size: None,
projection_dim: None,
deepstack_indexes: None,
};
let device = mlx_native::MlxDevice::new().expect("create device");
let weights = crate::inference::vision::mmproj_weights::LoadedMmprojWeights::empty(device);
LoadedMmproj {
gguf_path: "/tmp/synthetic-mmproj-gemma4v.gguf".into(),
config: cfg,
arch: crate::inference::vision::mmproj::ArchProfile::Gemma4Siglip,
weights: Arc::new(weights),
model_id: "synthetic-mmproj-gemma4v".into(),
}
}
fn synthetic_gemma4v_png_data_uri() -> String {
use base64::Engine;
use image::{ImageBuffer, ImageFormat, Rgb, RgbImage};
use std::io::Cursor;
let img: RgbImage = ImageBuffer::from_fn(256, 256, |_x, _y| Rgb([200u8, 100, 50]));
let mut buf: Vec<u8> = Vec::new();
img.write_to(&mut Cursor::new(&mut buf), ImageFormat::Png)
.expect("encode png");
let b64 = base64::engine::general_purpose::STANDARD.encode(&buf);
format!("data:image/png;base64,{b64}")
}
#[test]
fn gemma4v_arch_routes_through_variable_resolution_preprocess() {
use crate::inference::vision::vit_gpu::VisionInput;
let uri = synthetic_gemma4v_png_data_uri();
let msgs = vec![user_with_image("describe", &uri)];
let mmproj = synthetic_mmproj_gemma4v();
let got = process_multimodal_content(&msgs, Some(&mmproj)).expect("ok");
assert_eq!(got.len(), 1);
match &got[0] {
VisionInput::Gemma4v(g) => {
assert_eq!(g.n_x, 48);
assert_eq!(g.n_y, 48);
let n = (g.n_x as usize) * (g.n_y as usize);
assert_eq!(g.patches.len(), n * 16 * 16 * 3);
assert_eq!(g.pos_x.len(), n);
assert_eq!(g.pos_y.len(), n);
assert_eq!(g.pos_x[0], 0);
assert_eq!(g.pos_y[0], 0);
assert_eq!(g.pos_x[n - 1], g.n_x - 1);
assert_eq!(g.pos_y[n - 1], g.n_y - 1);
assert_eq!(g.source_label, "image/png");
}
VisionInput::Siglip49(_) => {
panic!("expected Gemma4v variant for Gemma4Siglip mmproj");
}
}
}
#[test]
fn unsupported_mime_type_returns_400() {
let msgs = vec![user_with_image("x", "data:image/gif;base64,R0lGODlh")];
let mmproj = synthetic_mmproj();
let resp = process_multimodal_content(&msgs, Some(&mmproj)).expect_err("gif should 400");
assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
}
#[test]
fn malformed_png_bytes_returns_400() {
use base64::Engine;
let payload = base64::engine::general_purpose::STANDARD.encode(b"this is not a png");
let uri = format!("data:image/png;base64,{payload}");
let msgs = vec![user_with_image("x", &uri)];
let mmproj = synthetic_mmproj();
let resp = process_multimodal_content(&msgs, Some(&mmproj))
.expect_err("bad PNG payload should 400");
assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
}
#[test]
fn rewrite_messages_for_vision_placeholders_passthrough_text() {
let msgs = vec![ChatMessage {
role: "user".into(),
content: Some(MessageContent::Text("hello".into())),
reasoning_content: None,
tool_calls: None,
tool_call_id: None,
name: None,
}];
let out = rewrite_messages_for_vision_placeholders(&msgs);
assert_eq!(out.len(), 1);
assert_eq!(out[0].content, msgs[0].content);
}
#[test]
fn rewrite_messages_for_vision_placeholders_pure_text_parts() {
let msgs = vec![ChatMessage {
role: "user".into(),
content: Some(MessageContent::Parts(vec![
ContentPart::Text {
text: "hello".into(),
},
ContentPart::Text {
text: " world".into(),
},
])),
reasoning_content: None,
tool_calls: None,
tool_call_id: None,
name: None,
}];
let out = rewrite_messages_for_vision_placeholders(&msgs);
match out[0].content.as_ref().expect("content") {
MessageContent::Parts(parts) => assert_eq!(parts.len(), 2),
other => panic!("expected Parts, got {:?}", other),
}
}
#[test]
fn rewrite_messages_for_vision_placeholders_one_image() {
let msgs = vec![ChatMessage {
role: "user".into(),
content: Some(MessageContent::Parts(vec![
ContentPart::Text {
text: "see this:".into(),
},
ContentPart::ImageUrl {
image_url: ImageUrl {
url: "data:image/png;base64,XXX".into(),
detail: None,
},
},
])),
reasoning_content: None,
tool_calls: None,
tool_call_id: None,
name: None,
}];
let out = rewrite_messages_for_vision_placeholders(&msgs);
match out[0].content.as_ref().expect("content") {
MessageContent::Text(t) => {
assert!(t.contains("<|image|>"), "got: {t}");
assert!(t.starts_with("see this:"));
}
other => panic!("expected Text, got {:?}", other),
}
}
#[test]
fn rewrite_messages_for_vision_placeholders_two_images_two_placeholders() {
let msgs = vec![ChatMessage {
role: "user".into(),
content: Some(MessageContent::Parts(vec![
ContentPart::Text { text: "a:".into() },
ContentPart::ImageUrl {
image_url: ImageUrl {
url: "data:image/png;base64,A".into(),
detail: None,
},
},
ContentPart::Text { text: " b:".into() },
ContentPart::ImageUrl {
image_url: ImageUrl {
url: "data:image/png;base64,B".into(),
detail: None,
},
},
])),
reasoning_content: None,
tool_calls: None,
tool_call_id: None,
name: None,
}];
let out = rewrite_messages_for_vision_placeholders(&msgs);
match out[0].content.as_ref().expect("content") {
MessageContent::Text(t) => {
let n_marks = t.matches("<|image|>").count();
assert_eq!(
n_marks, 2,
"expected 2 placeholders, got {n_marks} in {t:?}"
);
assert!(t.starts_with("a:"));
}
other => panic!("expected Text, got {:?}", other),
}
}
#[test]
fn compute_soft_token_layout_empty_prompt_zero_images_returns_empty() {
let (out, ranges) = compute_soft_token_layout(258880, &[], &[]).expect("ok");
assert!(out.is_empty());
assert!(ranges.is_empty());
}
#[test]
fn compute_soft_token_layout_no_placeholders_passes_through() {
let prompt = vec![1u32, 2, 3, 4, 5];
let (out, ranges) = compute_soft_token_layout(258880, &prompt, &[]).expect("ok");
assert_eq!(out, prompt);
assert!(ranges.is_empty());
}
#[test]
fn compute_soft_token_layout_single_placeholder_expands_to_n_copies() {
const IMG: u32 = 258880;
let prompt = vec![10u32, 20, IMG, 30, 40];
let (out, ranges) = compute_soft_token_layout(IMG, &prompt, &[4]).expect("ok");
assert_eq!(out, vec![10, 20, IMG, IMG, IMG, IMG, 30, 40]);
assert_eq!(ranges.len(), 1);
assert_eq!(ranges[0], 2..6);
for slot in &out[ranges[0].clone()] {
assert_eq!(*slot, IMG);
}
}
#[test]
fn compute_soft_token_layout_placeholder_at_start() {
const IMG: u32 = 258880;
let prompt = vec![IMG, 1, 2, 3];
let (out, ranges) = compute_soft_token_layout(IMG, &prompt, &[3]).expect("ok");
assert_eq!(out, vec![IMG, IMG, IMG, 1, 2, 3]);
assert_eq!(ranges[0], 0..3);
}
#[test]
fn compute_soft_token_layout_placeholder_at_end() {
const IMG: u32 = 258880;
let prompt = vec![1, 2, 3, IMG];
let (out, ranges) = compute_soft_token_layout(IMG, &prompt, &[2]).expect("ok");
assert_eq!(out, vec![1, 2, 3, IMG, IMG]);
assert_eq!(ranges[0], 3..5);
}
#[test]
fn compute_soft_token_layout_two_placeholders_independent_ranges() {
const IMG: u32 = 258880;
let prompt = vec![100u32, IMG, 200, IMG, 300];
let (out, ranges) = compute_soft_token_layout(IMG, &prompt, &[2, 4]).expect("ok");
assert_eq!(out.len(), 5 - 2 + 2 + 4);
assert_eq!(ranges[0], 1..3);
assert_eq!(ranges[1], 4..8);
assert_eq!(out[0], 100);
assert_eq!(out[3], 200);
assert_eq!(out[8], 300);
}
#[test]
fn compute_soft_token_layout_mismatch_reports_both_counts() {
const IMG: u32 = 258880;
let prompt = vec![1, IMG, 2, IMG, 3];
let err = compute_soft_token_layout(IMG, &prompt, &[5]).expect_err("must reject");
assert_eq!(
err,
PlaceholderCountMismatch {
placeholder_positions_found: 2,
n_image_tokens_supplied: 1,
}
);
let err = compute_soft_token_layout(IMG, &[1, 2, 3], &[5, 5]).expect_err("must reject");
assert_eq!(
err,
PlaceholderCountMismatch {
placeholder_positions_found: 0,
n_image_tokens_supplied: 2,
}
);
}
#[test]
fn compute_soft_token_layout_zero_image_tokens_drops_placeholder() {
const IMG: u32 = 258880;
let prompt = vec![10u32, IMG, 20];
let (out, ranges) = compute_soft_token_layout(IMG, &prompt, &[0]).expect("ok");
assert_eq!(out, vec![10, 20]);
assert_eq!(ranges[0], 1..1);
}
#[test]
fn compute_soft_token_layout_ranges_match_n_image_tokens() {
const IMG: u32 = 258880;
let prompt = vec![1u32, IMG, 2, 3, IMG, 4, IMG];
let n_per = vec![3usize, 7, 1];
let (_, ranges) = compute_soft_token_layout(IMG, &prompt, &n_per).expect("ok");
for (i, range) in ranges.iter().enumerate() {
assert_eq!(range.len(), n_per[i], "range {i} len");
}
}
#[test]
fn rewrite_messages_for_vision_placeholders_only_touches_image_messages() {
let msgs = vec![
ChatMessage {
role: "system".into(),
content: Some(MessageContent::Text("be helpful".into())),
reasoning_content: None,
tool_calls: None,
tool_call_id: None,
name: None,
},
ChatMessage {
role: "user".into(),
content: Some(MessageContent::Parts(vec![ContentPart::ImageUrl {
image_url: ImageUrl {
url: "data:image/png;base64,X".into(),
detail: None,
},
}])),
reasoning_content: None,
tool_calls: None,
tool_call_id: None,
name: None,
},
ChatMessage {
role: "assistant".into(),
content: Some(MessageContent::Text("ack".into())),
reasoning_content: None,
tool_calls: None,
tool_call_id: None,
name: None,
},
];
let out = rewrite_messages_for_vision_placeholders(&msgs);
assert_eq!(out.len(), 3);
assert_eq!(out[0].content, msgs[0].content);
assert_eq!(out[2].content, msgs[2].content);
match out[1].content.as_ref().expect("content") {
MessageContent::Text(t) => assert_eq!(t, "<|image|>"),
other => panic!("expected Text, got {:?}", other),
}
}
#[test]
fn rewrite_messages_family_qwen3vl_emits_vision_start_image_pad_vision_end() {
use crate::inference::vision::mmproj::VisionFamily;
let msgs = vec![ChatMessage {
role: "user".into(),
content: Some(MessageContent::Parts(vec![
ContentPart::Text {
text: "describe:".into(),
},
ContentPart::ImageUrl {
image_url: ImageUrl {
url: "data:image/png;base64,XXX".into(),
detail: None,
},
},
])),
reasoning_content: None,
tool_calls: None,
tool_call_id: None,
name: None,
}];
let out = rewrite_messages_for_vision_placeholders_family(&msgs, VisionFamily::Qwen3Vl);
match out[0].content.as_ref().expect("content") {
MessageContent::Text(t) => {
assert!(
t.contains("<|vision_start|><|image_pad|><|vision_end|>"),
"expected qwen3vl triplet, got: {t}"
);
assert!(t.starts_with("describe:"));
assert_eq!(t.matches("<|image_pad|>").count(), 1);
}
other => panic!("expected Text, got {:?}", other),
}
}
#[test]
fn rewrite_messages_family_qwen3vl_two_images_two_triplets() {
use crate::inference::vision::mmproj::VisionFamily;
let msgs = vec![ChatMessage {
role: "user".into(),
content: Some(MessageContent::Parts(vec![
ContentPart::Text { text: "a".into() },
ContentPart::ImageUrl {
image_url: ImageUrl {
url: "data:image/png;base64,A".into(),
detail: None,
},
},
ContentPart::Text { text: "b".into() },
ContentPart::ImageUrl {
image_url: ImageUrl {
url: "data:image/png;base64,B".into(),
detail: None,
},
},
])),
reasoning_content: None,
tool_calls: None,
tool_call_id: None,
name: None,
}];
let out = rewrite_messages_for_vision_placeholders_family(&msgs, VisionFamily::Qwen3Vl);
match out[0].content.as_ref().expect("content") {
MessageContent::Text(t) => {
assert_eq!(t.matches("<|image_pad|>").count(), 2);
assert_eq!(t.matches("<|vision_start|>").count(), 2);
assert_eq!(t.matches("<|vision_end|>").count(), 2);
}
other => panic!("expected Text, got {:?}", other),
}
}
#[test]
fn rewrite_messages_family_gemma_byte_identical_to_legacy() {
use crate::inference::vision::mmproj::VisionFamily;
let msgs = vec![ChatMessage {
role: "user".into(),
content: Some(MessageContent::Parts(vec![
ContentPart::Text { text: "x".into() },
ContentPart::ImageUrl {
image_url: ImageUrl {
url: "data:image/png;base64,A".into(),
detail: None,
},
},
ContentPart::Text { text: "y".into() },
])),
reasoning_content: None,
tool_calls: None,
tool_call_id: None,
name: None,
}];
let legacy = rewrite_messages_for_vision_placeholders(&msgs);
let family = rewrite_messages_for_vision_placeholders_family(&msgs, VisionFamily::Gemma);
assert_eq!(legacy.len(), family.len());
assert_eq!(legacy[0].content, family[0].content);
}
fn cpu_split_qwen3vl_augmented(
vision_embeddings: &[Vec<f32>],
per_image_n_image_tokens: usize,
hidden: usize,
n_deepstack: usize,
) -> Vec<Vec<f32>> {
let per_row_floats = hidden * (1 + n_deepstack);
let total = per_image_n_image_tokens * vision_embeddings.len();
let mut chunks: Vec<Vec<f32>> = (0..n_deepstack)
.map(|_| vec![0f32; total * hidden])
.collect();
for j in 0..n_deepstack {
let mut row_offset = 0usize;
for src in vision_embeddings {
for r in 0..per_image_n_image_tokens {
let src_base = r * per_row_floats + (j + 1) * hidden;
let dst_base = (row_offset + r) * hidden;
chunks[j][dst_base..dst_base + hidden]
.copy_from_slice(&src[src_base..src_base + hidden]);
}
row_offset += per_image_n_image_tokens;
}
}
chunks
}
#[test]
fn qwen3vl_seam_split_round_trip_identity() {
let hidden = 4;
let n_deepstack = 3;
let n_image_tokens = 6;
let per_row_floats = hidden * (1 + n_deepstack);
let mut augmented = vec![0f32; n_image_tokens * per_row_floats];
for r in 0..n_image_tokens {
for j in 0..(1 + n_deepstack) {
for h in 0..hidden {
augmented[r * per_row_floats + j * hidden + h] =
(r * 1000 + j * 100 + h) as f32;
}
}
}
let vision_embeddings = vec![augmented.clone()];
let chunks =
cpu_split_qwen3vl_augmented(&vision_embeddings, n_image_tokens, hidden, n_deepstack);
assert_eq!(chunks.len(), n_deepstack);
for r in 0..n_image_tokens {
let mut reconstructed = Vec::with_capacity(per_row_floats);
reconstructed
.extend_from_slice(&augmented[r * per_row_floats..r * per_row_floats + hidden]);
for j in 0..n_deepstack {
reconstructed.extend_from_slice(&chunks[j][r * hidden..(r + 1) * hidden]);
}
let original_row = &augmented[r * per_row_floats..(r + 1) * per_row_floats];
assert_eq!(reconstructed, original_row, "row {r} round-trip mismatch");
}
}
#[test]
fn qwen3vl_seam_split_handles_per_image_variable_n_image_tokens_iter225() {
let hidden = 2;
let n_deepstack = 1;
let per_row_floats = hidden * (1 + n_deepstack);
let n_img0 = 4usize;
let n_img1 = 3usize;
let mut img0 = vec![0f32; n_img0 * per_row_floats];
for r in 0..n_img0 {
for h in 0..hidden {
img0[r * per_row_floats + 1 * hidden + h] = (r * 10) as f32;
}
}
let mut img1 = vec![0f32; n_img1 * per_row_floats];
for r in 0..n_img1 {
for h in 0..hidden {
img1[r * per_row_floats + 1 * hidden + h] = (100 + r) as f32;
}
}
let total_n = n_img0 + n_img1;
let per_image_n_image_tokens = vec![n_img0, n_img1];
let vision_embeddings = vec![img0, img1];
let mut chunks: Vec<Vec<f32>> = (0..n_deepstack)
.map(|_| vec![0f32; total_n * hidden])
.collect();
for j in 0..n_deepstack {
let mut row_offset = 0usize;
for (i, src) in vision_embeddings.iter().enumerate() {
let n_img_i = per_image_n_image_tokens[i];
for r in 0..n_img_i {
let src_base = r * per_row_floats + (j + 1) * hidden;
let dst_base = (row_offset + r) * hidden;
chunks[j][dst_base..dst_base + hidden]
.copy_from_slice(&src[src_base..src_base + hidden]);
}
row_offset += n_img_i;
}
}
let row = |r: usize, h: usize| r * hidden + h;
for h in 0..hidden {
assert_eq!(chunks[0][row(0, h)], 0.0);
assert_eq!(chunks[0][row(1, h)], 10.0);
assert_eq!(chunks[0][row(2, h)], 20.0);
assert_eq!(chunks[0][row(3, h)], 30.0);
assert_eq!(chunks[0][row(4, h)], 100.0);
assert_eq!(chunks[0][row(5, h)], 101.0);
assert_eq!(chunks[0][row(6, h)], 102.0);
}
}
#[test]
fn qwen3vl_seam_split_concatenates_multi_image_in_order() {
let hidden = 2;
let n_deepstack = 1;
let n_image_tokens = 3;
let per_row_floats = hidden * (1 + n_deepstack);
let mut img0 = vec![0f32; n_image_tokens * per_row_floats];
let mut img1 = vec![0f32; n_image_tokens * per_row_floats];
for r in 0..n_image_tokens {
for h in 0..hidden {
img0[r * per_row_floats + 1 * hidden + h] = r as f32;
img1[r * per_row_floats + 1 * hidden + h] = (100 + r) as f32;
}
}
let vision_embeddings = vec![img0, img1];
let chunks =
cpu_split_qwen3vl_augmented(&vision_embeddings, n_image_tokens, hidden, n_deepstack);
let row = |r: usize, h: usize| r * hidden + h;
for h in 0..hidden {
assert_eq!(chunks[0][row(0, h)], 0.0);
assert_eq!(chunks[0][row(3, h)], 100.0);
assert_eq!(chunks[0][row(4, h)], 101.0);
assert_eq!(chunks[0][row(5, h)], 102.0);
}
}
}
#[cfg(test)]
mod readiness_guard_tests {
use super::super::state::{AppState, ServerConfig};
use super::*;
use axum::http::header::RETRY_AFTER;
#[test]
fn not_ready_error_is_503_with_retry_after_1() {
let resp = ApiError::not_ready().into_response();
assert_eq!(resp.status(), StatusCode::SERVICE_UNAVAILABLE);
let ra = resp
.headers()
.get(RETRY_AFTER)
.expect("Retry-After must be set on not_ready");
assert_eq!(ra, "1");
}
#[test]
fn app_state_mark_not_ready_disables_gen() {
let state = AppState::new(ServerConfig::default());
assert!(state.is_ready_for_gen(), "should start ready");
state.mark_not_ready();
assert!(!state.is_ready_for_gen(), "should be not-ready after mark");
let resp = ApiError::not_ready().into_response();
assert_eq!(resp.status(), StatusCode::SERVICE_UNAVAILABLE);
assert!(resp.headers().contains_key(RETRY_AFTER));
}
#[tokio::test]
async fn chat_completions_returns_503_before_resolve_when_not_ready() {
let state = AppState::new(ServerConfig::default());
state.mark_not_ready();
assert!(!state.is_ready_for_gen());
let req = super::super::schema::ChatCompletionRequest {
model: "any-model".to_string(),
messages: vec![super::super::schema::ChatMessage {
role: "user".to_string(),
content: Some(super::super::schema::MessageContent::Text(
"hello".to_string(),
)),
reasoning_content: None,
tool_calls: None,
tool_call_id: None,
name: None,
}],
stream: None,
max_tokens: None,
max_completion_tokens: None,
temperature: None,
stop: None,
tools: None,
tool_choice: None,
response_format: None,
top_p: None,
seed: None,
frequency_penalty: None,
presence_penalty: None,
stream_options: None,
top_k: None,
repetition_penalty: None,
min_p: None,
logprobs: None,
top_logprobs: None,
logit_bias: None,
parallel_tool_calls: None,
hf2q_overflow_policy: None,
hf2q_enable_thinking: None,
chat_template_kwargs: None,
};
let resolver = |_state: &AppState, _model: String| -> ResolverBoxFuture<'_> {
unreachable!(
"resolver must not be called when is_ready_for_gen() == false; \
the readiness guard at prepare_chat_generation_core should have \
returned early with 503 before reaching this call-site"
)
};
let resp = match prepare_chat_generation_core(&state, &req, resolver).await {
Err(r) => r,
Ok(_) => panic!("not-ready state must return Err(503 response), got Ok"),
};
assert_eq!(
resp.status(),
StatusCode::SERVICE_UNAVAILABLE,
"not-ready must yield 503 SERVICE_UNAVAILABLE"
);
let ra = resp
.headers()
.get(RETRY_AFTER)
.expect("Retry-After header must be present on not-ready 503");
assert_eq!(ra, "1", "not-ready Retry-After must be 1 second");
}
}
#[cfg(test)]
mod pool_error_tests {
use super::*;
use axum::http::header::RETRY_AFTER;
#[test]
fn oversized_handle_maps_to_503_with_retry_after() {
let err = HotSwapError::PoolRefused(PoolError::OversizedHandle {
repo_id: "acme/big-model".to_string(),
handle_bytes: 40_000_000_000,
budget_bytes: 20_000_000_000,
});
let resp = map_hotswap_error_to_response(err);
assert_eq!(resp.status(), StatusCode::SERVICE_UNAVAILABLE);
let ra = resp
.headers()
.get(RETRY_AFTER)
.expect("Retry-After header must be set");
assert_eq!(ra, "5", "Retry-After should be 5 seconds for pool refusal");
}
#[test]
fn zero_capacity_maps_to_503_with_retry_after() {
let err = HotSwapError::PoolRefused(PoolError::ZeroCapacity);
let resp = map_hotswap_error_to_response(err);
assert_eq!(resp.status(), StatusCode::SERVICE_UNAVAILABLE);
let ra = resp
.headers()
.get(RETRY_AFTER)
.expect("Retry-After header must be set");
assert_eq!(ra, "5");
}
#[test]
fn loader_failed_maps_to_500() {
let err = HotSwapError::LoaderFailed(anyhow::anyhow!("tokenizer missing"));
let resp = map_hotswap_error_to_response(err);
assert_eq!(resp.status(), StatusCode::INTERNAL_SERVER_ERROR);
assert!(resp.headers().get(RETRY_AFTER).is_none());
}
#[tokio::test]
async fn pool_refused_through_resolver_seam_yields_503_with_retry_after_5() {
use super::super::state::{AppState, ServerConfig};
let state = AppState::new(ServerConfig::default());
assert!(state.is_ready_for_gen(), "AppState::new should start ready");
let req = super::super::schema::ChatCompletionRequest {
model: "pool-refused-model".to_string(),
messages: vec![super::super::schema::ChatMessage {
role: "user".to_string(),
content: Some(super::super::schema::MessageContent::Text(
"test".to_string(),
)),
reasoning_content: None,
tool_calls: None,
tool_call_id: None,
name: None,
}],
stream: None,
max_tokens: None,
max_completion_tokens: None,
temperature: None,
stop: None,
tools: None,
tool_choice: None,
response_format: None,
top_p: None,
seed: None,
frequency_penalty: None,
presence_penalty: None,
stream_options: None,
top_k: None,
repetition_penalty: None,
min_p: None,
logprobs: None,
top_logprobs: None,
logit_bias: None,
parallel_tool_calls: None,
hf2q_overflow_policy: None,
hf2q_enable_thinking: None,
chat_template_kwargs: None,
};
let resolver = |_state: &AppState, _model: String| -> ResolverBoxFuture<'_> {
let err = HotSwapError::PoolRefused(PoolError::ZeroCapacity);
let response = map_hotswap_error_to_response(err);
Box::pin(async move { Err(response) })
};
let resp = match prepare_chat_generation_core(&state, &req, resolver).await {
Err(r) => r,
Ok(_) => panic!("PoolRefused resolver must yield Err(503 response), got Ok"),
};
assert_eq!(
resp.status(),
StatusCode::SERVICE_UNAVAILABLE,
"PoolRefused must yield 503 SERVICE_UNAVAILABLE at the handler boundary"
);
let ra = resp
.headers()
.get(RETRY_AFTER)
.expect("Retry-After header must be present on pool-refused 503");
assert_eq!(
ra, "5",
"pool-refused Retry-After must be 5 seconds (not 1, which is for not-ready)"
);
}
}
#[cfg(test)]
mod iter215_qwen35_chat_501_tests {
use super::*;
use axum::body::to_bytes;
use axum::http::StatusCode;
use std::sync::Arc;
use std::time::SystemTime;
use super::super::engine;
use super::super::schema::{ChatCompletionRequest, ChatMessage, MessageContent};
use super::super::state::{AppState, ServerConfig};
fn empty_request_for_model(model: &str) -> ChatCompletionRequest {
ChatCompletionRequest {
model: model.to_string(),
messages: vec![ChatMessage {
role: "user".to_string(),
content: Some(MessageContent::Text("hi".to_string())),
reasoning_content: None,
tool_calls: None,
tool_call_id: None,
name: None,
}],
stream: None,
max_tokens: None,
max_completion_tokens: None,
temperature: None,
stop: None,
tools: None,
tool_choice: None,
response_format: None,
top_p: None,
seed: None,
frequency_penalty: None,
presence_penalty: None,
stream_options: None,
top_k: None,
repetition_penalty: None,
min_p: None,
logprobs: None,
top_logprobs: None,
logit_bias: None,
parallel_tool_calls: None,
hf2q_overflow_policy: None,
hf2q_enable_thinking: None,
chat_template_kwargs: None,
}
}
#[tokio::test]
async fn qwen35_chat_completion_returns_501_with_actionable_message() {
let state = AppState::new(ServerConfig::default());
assert!(state.is_ready_for_gen());
let req = empty_request_for_model("iter-215-test-model");
let engine = engine::make_synthetic_engine_for_test(engine::LoadedArch::Qwen35);
let loaded_engine = Arc::new(LoadedEngine {
engine,
repo: "iter-215-test-model".to_string(),
quant: QuantType::Q4_K_M,
bytes_resident: 1024,
loaded_at: SystemTime::now(),
});
let resolver_le = loaded_engine.clone();
let resolver = move |_state: &AppState, _model: String| -> ResolverBoxFuture<'_> {
let le = resolver_le.clone();
Box::pin(async move { Ok(le) })
};
let prepared = match prepare_chat_generation_core(&state, &req, resolver).await {
Ok(p) => p,
Err(_resp) => panic!("prepare must succeed before the 501 short-circuit fires"),
};
assert_eq!(
prepared.loaded_engine.engine.arch(),
engine::LoadedArch::Qwen35,
"synthetic engine must report Qwen35 arch"
);
let resp = ApiError::not_implemented(engine::QWEN35_NOT_IMPLEMENTED_MESSAGE.to_string())
.into_response();
assert_eq!(
resp.status(),
StatusCode::NOT_IMPLEMENTED,
"Qwen3.5/3.6 chat must return HTTP 501"
);
let body_bytes = to_bytes(resp.into_body(), 1 << 20)
.await
.expect("read response body");
let body = String::from_utf8_lossy(&body_bytes);
assert!(
body.contains("hf2q generate"),
"501 body must name `hf2q generate`; got: {body}"
);
assert!(
body.contains("cmd_generate_qwen35"),
"501 body must name `cmd_generate_qwen35`; got: {body}"
);
}
}
#[cfg(test)]
mod test_a3_tool_call_extraction {
use super::{engine, extract_tool_calls_from_text, registry, tool_turn_message_content};
#[test]
fn pure_tool_turn_discards_whitespace_only_content() {
assert!(tool_turn_message_content("\n\n".to_string()).is_none());
assert!(tool_turn_message_content(" explanation".to_string()).is_some());
}
fn gemma4_reg() -> registry::ModelRegistration {
registry::find_for("gemma4-27b-it").expect("gemma4 registration must exist")
}
#[test]
fn extract_no_tool_calls_returns_text_unchanged() {
let reg = gemma4_reg();
let result =
extract_tool_calls_from_text("Hello, world!", Some(®), engine::ToolCallPolicy::Auto);
assert!(result.tool_calls.is_empty(), "no tool calls in plain text");
assert_eq!(result.content, "Hello, world!", "content must be preserved");
assert!(result.constrained_parse_failure.is_none());
}
#[test]
fn extract_tool_call_with_gemma4_markers() {
let reg = gemma4_reg();
let (open, close) = match (reg.tool_open, reg.tool_close) {
(Some(o), Some(c)) => (o, c),
_ => {
eprintln!("gemma4 registration has no tool markers, skipping");
return;
}
};
let body = "call:get_weather{location:<|\"|>Paris<|\"|>}";
let full_text = format!("Here is the weather: {open}{body}{close}");
let result =
extract_tool_calls_from_text(&full_text, Some(®), engine::ToolCallPolicy::Auto);
assert_eq!(
result.tool_calls.len(),
1,
"exactly one tool call must be extracted"
);
assert_eq!(result.tool_calls[0].function.name, "get_weather");
assert_eq!(result.tool_calls[0].call_type, "function");
assert!(
result.tool_calls[0].id.starts_with("call_hf2q_"),
"tool call id must start with call_hf2q_"
);
assert!(
result.content.contains("Here is the weather:"),
"pre-marker content must be preserved"
);
}
#[test]
fn extract_no_registration_returns_text_unchanged() {
let result = extract_tool_calls_from_text("plain text", None, engine::ToolCallPolicy::Auto);
assert!(result.tool_calls.is_empty());
assert_eq!(result.content, "plain text");
}
#[test]
fn no_marker_under_constrained_yields_empty_tool_calls() {
let reg = gemma4_reg();
let result = extract_tool_calls_from_text(
"I will call the weather tool but I forgot how.",
Some(®),
engine::ToolCallPolicy::Constrained,
);
assert!(
result.tool_calls.is_empty(),
"no marker span ⇒ no tool calls, even under Constrained policy"
);
assert!(
result.constrained_parse_failure.is_none(),
"constrained_parse_failure only set when a body was parsed and failed"
);
}
#[test]
fn malformed_body_under_auto_lazy_grammar_yields_parse_failure() {
let reg = gemma4_reg();
let (open, close) = match (reg.tool_open, reg.tool_close) {
(Some(o), Some(c)) => (o, c),
_ => return,
};
let bad_body = "totally malformed body bytes — not a call";
let full_text = format!("preamble {open}{bad_body}{close}");
let result = extract_tool_calls_from_text(
&full_text,
Some(®),
engine::ToolCallPolicy::AutoLazyGrammar,
);
assert!(
result.tool_calls.is_empty(),
"malformed body MUST NOT produce a parsed tool call"
);
assert!(
result.constrained_parse_failure.is_some(),
"AutoLazyGrammar parse failure MUST set constrained_parse_failure \
(loud-error promotion identical to Constrained); pre-W-B2 this \
would have demoted to Content"
);
let failed = result.constrained_parse_failure.unwrap();
assert_eq!(
failed, bad_body,
"constrained_parse_failure MUST carry the unparseable body bytes \
verbatim for the operator to inspect; got: {failed:?}"
);
}
#[test]
fn malformed_body_under_auto_no_grammar_demotes_to_content() {
let reg = gemma4_reg();
let (open, close) = match (reg.tool_open, reg.tool_close) {
(Some(o), Some(c)) => (o, c),
_ => return,
};
let bad_body = "totally malformed body bytes — not a call";
let full_text = format!("{open}{bad_body}{close}");
let result =
extract_tool_calls_from_text(&full_text, Some(®), engine::ToolCallPolicy::Auto);
assert!(
result.tool_calls.is_empty(),
"malformed body MUST NOT produce a parsed tool call"
);
assert!(
result.constrained_parse_failure.is_none(),
"Auto (no grammar) MUST keep the content-fallback semantics — \
parse failure does NOT promote to constrained_parse_failure \
(W-B2 only narrowed the registered-family path; vanilla Auto \
remains unchanged)"
);
assert!(
result.content.contains(bad_body),
"Auto (no grammar) MUST re-emit the malformed body as Content; \
got content: {:?}",
result.content
);
}
#[test]
fn iter219b_nonstream_auto_fallback_scrubs_special_tokens() {
let reg = gemma4_reg();
let (open, close) = match (reg.tool_open, reg.tool_close) {
(Some(o), Some(c)) => (o, c),
_ => return,
};
let bad_body =
"call:get_current<|tool_response>call:get_current_weather{location:<|\"|>Paris<|\"|>}";
let full_text = format!("{open}{bad_body}{close}");
let result =
extract_tool_calls_from_text(&full_text, Some(®), engine::ToolCallPolicy::Auto);
assert!(
result.tool_calls.is_empty(),
"polluted body must not produce a parsed tool call"
);
for marker in &[
"<|channel>",
"<channel|>",
"<|tool_call>",
"<tool_call|>",
"<|tool_response>",
"<tool_response|>",
"<|turn>",
"<turn|>",
] {
assert!(
!result.content.contains(marker),
"iter-219b non-stream content must not contain registered \
special-token marker {marker:?}; got content: {:?}",
result.content
);
}
assert!(
result.content.contains("call:get_current_weather"),
"scrubbed content should retain the legitimate body bytes; \
got: {:?}",
result.content
);
}
}
#[cfg(test)]
mod defensive_500_wire_shape_tests {
use super::*;
use axum::body::to_bytes;
#[tokio::test]
async fn defensive_500_under_constrained_returns_structured_envelope() {
let resp = defensive_no_call_under_constrained(
engine::ToolCallPolicy::Constrained,
true,
"length",
64,
0,
)
.expect("defensive helper MUST emit a Response under Constrained + empty");
assert_eq!(
resp.status(),
StatusCode::INTERNAL_SERVER_ERROR,
"defensive 500 MUST return HTTP 500 (server_error class)"
);
let bytes = to_bytes(resp.into_body(), 1 << 16).await.expect("body");
let body_str = String::from_utf8_lossy(&bytes);
let parsed: serde_json::Value =
serde_json::from_str(&body_str).expect("body must be valid JSON");
let err = parsed.get("error").expect("envelope must have `error` key");
assert_eq!(
err.get("type").and_then(|v| v.as_str()),
Some("server_error"),
"error.type MUST be `server_error` for the 500 class"
);
assert_eq!(
err.get("code").and_then(|v| v.as_str()),
Some("generation_error"),
"error.code MUST be `generation_error` (the OpenAI envelope code; \
the structured tool_call_no_call_under_constrained discriminant \
is in the message)"
);
let message = err
.get("message")
.and_then(|v| v.as_str())
.expect("error.message must be a string");
assert!(
message.contains("tool_call_no_call_under_constrained"),
"error.message MUST contain the structured discriminant \
`tool_call_no_call_under_constrained` so log/metrics matchers \
can target it; got: {message}"
);
}
#[test]
fn defensive_500_when_calls_present_returns_none() {
let resp = defensive_no_call_under_constrained(
engine::ToolCallPolicy::Constrained,
false,
"stop",
42,
128,
);
assert!(
resp.is_none(),
"Constrained + non-empty tool_calls MUST proceed to success path"
);
}
#[test]
fn defensive_500_under_auto_no_call_returns_none() {
let resp = defensive_no_call_under_constrained(
engine::ToolCallPolicy::Auto,
true,
"stop",
8,
32,
);
assert!(
resp.is_none(),
"Auto policy + zero tool_calls MUST proceed to success path \
(Auto allows no-call). Regression-preserve."
);
}
#[test]
fn defensive_500_under_auto_with_calls_returns_none() {
let resp = defensive_no_call_under_constrained(
engine::ToolCallPolicy::Auto,
false,
"tool_calls",
42,
128,
);
assert!(resp.is_none(), "Auto + non-empty must always be None");
}
}
#[cfg(test)]
mod test_a4_tool_call_policy {
use super::{engine, extract_tool_calls_from_text, registry};
fn gemma4_reg() -> registry::ModelRegistration {
registry::find_for("gemma4-27b-it").expect("gemma4 registration must exist")
}
#[test]
fn auto_parse_failure_fallback_to_content() {
let reg = gemma4_reg();
let (open, close) = match (reg.tool_open, reg.tool_close) {
(Some(o), Some(c)) => (o, c),
_ => return, };
let bad_body = "garbage{}";
let full_text = format!("{open}{bad_body}{close}");
let result =
extract_tool_calls_from_text(&full_text, Some(®), engine::ToolCallPolicy::Auto);
assert!(
result.constrained_parse_failure.is_none(),
"Auto mode must NOT set constrained_parse_failure"
);
assert!(
result.tool_calls.is_empty(),
"Auto mode parse failure must produce no tool calls"
);
assert!(
result.content.contains(bad_body),
"Auto mode must re-emit bad body as content, got: {:?}",
result.content
);
}
#[test]
fn constrained_parse_failure_signals_error() {
let reg = gemma4_reg();
let (open, close) = match (reg.tool_open, reg.tool_close) {
(Some(o), Some(c)) => (o, c),
_ => return, };
let bad_body = "garbage{}";
let full_text = format!("{open}{bad_body}{close}");
let result = extract_tool_calls_from_text(
&full_text,
Some(®),
engine::ToolCallPolicy::Constrained,
);
assert!(
result.constrained_parse_failure.is_some(),
"Constrained mode must set constrained_parse_failure on parse failure"
);
assert!(
result.tool_calls.is_empty(),
"Constrained mode parse failure must produce no tool calls"
);
}
}
#[cfg(test)]
mod bos_probe_tests {
use super::{probe_bos_token_id, BOS_PROBE_FRAGMENTS};
use tokenizers::{models::bpe::BPE, AddedToken, Tokenizer};
fn tokenizer_with_specials(fragments: &[&str]) -> Tokenizer {
let mut tok = Tokenizer::new(BPE::default());
let added: Vec<AddedToken> = fragments
.iter()
.map(|f| AddedToken::from((*f).to_string(), true))
.collect();
tok.add_special_tokens(&added);
tok
}
#[test]
fn probe_returns_bos_for_gemma_family_tokenizer() {
let tok = tokenizer_with_specials(&["<bos>"]);
let expected_id = tok.token_to_id("<bos>");
assert!(
expected_id.is_some(),
"fixture invariant: <bos> must register"
);
assert_eq!(probe_bos_token_id(&tok), expected_id);
}
#[test]
fn probe_returns_begin_of_text_for_llama3_family_tokenizer() {
let tok = tokenizer_with_specials(&["<|begin_of_text|>"]);
let expected_id = tok.token_to_id("<|begin_of_text|>");
assert!(
expected_id.is_some(),
"fixture invariant: <|begin_of_text|> must register"
);
assert_eq!(probe_bos_token_id(&tok), expected_id);
}
#[test]
fn probe_returns_s_for_llama2_mistral_family_tokenizer() {
let tok = tokenizer_with_specials(&["<s>"]);
let expected_id = tok.token_to_id("<s>");
assert!(
expected_id.is_some(),
"fixture invariant: <s> must register"
);
assert_eq!(probe_bos_token_id(&tok), expected_id);
}
#[test]
fn probe_returns_im_start_for_qwen_family_tokenizer() {
let tok = tokenizer_with_specials(&["<|im_start|>"]);
let expected_id = tok.token_to_id("<|im_start|>");
assert!(
expected_id.is_some(),
"fixture invariant: <|im_start|> must register"
);
assert_eq!(probe_bos_token_id(&tok), expected_id);
}
#[test]
fn probe_returns_none_when_no_fragment_matches() {
let tok = tokenizer_with_specials(&["<unk>", "<pad>"]);
assert_eq!(probe_bos_token_id(&tok), None);
}
#[test]
fn probe_first_match_wins_when_multiple_present() {
let tok = tokenizer_with_specials(&["<s>", "<bos>"]);
let bos_id = tok.token_to_id("<bos>");
let s_id = tok.token_to_id("<s>");
assert!(bos_id.is_some() && s_id.is_some());
assert_ne!(bos_id, s_id, "fixture invariant: distinct ids");
assert_eq!(
probe_bos_token_id(&tok),
bos_id,
"<bos> must win over <s> per BOS_PROBE_FRAGMENTS array order"
);
}
#[test]
fn probe_fragments_array_documented_order_invariant() {
assert_eq!(BOS_PROBE_FRAGMENTS.len(), 4);
assert_eq!(BOS_PROBE_FRAGMENTS[0], "<bos>", "Gemma slot");
assert_eq!(BOS_PROBE_FRAGMENTS[1], "<|begin_of_text|>", "Llama 3 slot");
assert_eq!(BOS_PROBE_FRAGMENTS[2], "<s>", "Llama 1/2 + Mistral slot");
assert_eq!(BOS_PROBE_FRAGMENTS[3], "<|im_start|>", "Qwen slot");
}
#[test]
fn probe_fragments_array_contains_no_duplicates() {
let mut seen = std::collections::HashSet::new();
for f in BOS_PROBE_FRAGMENTS {
assert!(
seen.insert(*f),
"duplicate fragment in BOS_PROBE_FRAGMENTS: {f}"
);
}
}
#[test]
fn a5b_parse_slot_budget_exceeded_extracts_both_numbers() {
let msg = "slot_budget_exceeded: ADR-040 §3.5 A5b — per-slot KV \
budget exceeded (needed_bytes=12345678, budget_bytes=4096000). \
Reduce max_tokens or use a shorter prompt.";
assert_eq!(
super::parse_slot_budget_exceeded(msg),
(12_345_678, 4_096_000)
);
}
#[test]
fn a5b_parse_slot_budget_exceeded_no_match_returns_zeros() {
let msg = "completely unrelated error string";
assert_eq!(super::parse_slot_budget_exceeded(msg), (0, 0));
}
#[test]
fn a5b_parse_slot_budget_exceeded_partial_match_returns_zero_for_missing() {
let msg = "slot_budget_exceeded: needed_bytes=999 (no budget here)";
assert_eq!(super::parse_slot_budget_exceeded(msg), (999, 0));
let msg2 = "slot_budget_exceeded: budget_bytes=4096 (no needed)";
assert_eq!(super::parse_slot_budget_exceeded(msg2), (0, 4096));
}
#[test]
fn a5b_parse_slot_budget_exceeded_handles_u64_max() {
let big = u64::MAX.to_string();
let msg = format!("slot_budget_exceeded: needed_bytes={big}, budget_bytes={big}");
assert_eq!(
super::parse_slot_budget_exceeded(&msg),
(u64::MAX, u64::MAX)
);
}
#[test]
fn a5b_parse_slot_budget_exceeded_extracts_from_streaming_error_format() {
let msg = "slot_budget_exceeded: ADR-040 §3.5 A5b — per-slot KV \
budget exceeded for GenerateStream (needed_bytes=200, \
budget_bytes=100). Reduce max_tokens or use a shorter prompt.";
assert_eq!(super::parse_slot_budget_exceeded(msg), (200, 100));
}
}
#[cfg(test)]
mod a5d_handler_429_tests {
use super::super::engine::{
make_synthetic_engine_over_budget, make_synthetic_engine_with_slot_budget_exceeded_worker,
LoadedArch,
};
use super::super::state::{AppState, ServerConfig};
use super::*;
use crate::serve::multi_model::LoadedEngine;
use crate::serve::quant_select::QuantType;
use axum::body::to_bytes;
use axum::http::{header, StatusCode};
use std::sync::Arc;
use std::time::SystemTime;
fn build_prepared_context(
engine: Engine,
prompt_tokens: Vec<u32>,
max_tokens: usize,
) -> PreparedChatContext {
let loaded_engine = Arc::new(LoadedEngine {
engine,
repo: "a5d-handler-test".to_string(),
quant: QuantType::Q4_K_M,
bytes_resident: 0,
loaded_at: SystemTime::now(),
});
let mut params = SamplingParams::default();
params.max_tokens = max_tokens;
PreparedChatContext {
loaded_engine,
prompt_tokens,
params,
summarized_messages: None,
summary_tokens: None,
soft_tokens: Vec::new(),
vit_forward_ms: None,
vit_images: None,
vit_soft_tokens_total: None,
deepstack_data: None,
positions_flat: None,
}
}
fn minimal_request(stream: bool, max_tokens: usize) -> ChatCompletionRequest {
ChatCompletionRequest {
model: "a5d-handler-test".to_string(),
messages: vec![ChatMessage {
role: "user".to_string(),
content: Some(MessageContent::Text("hello".to_string())),
reasoning_content: None,
tool_calls: None,
tool_call_id: None,
name: None,
}],
stream: Some(stream),
max_tokens: Some(max_tokens),
max_completion_tokens: None,
temperature: None,
stop: None,
tools: None,
tool_choice: None,
response_format: None,
top_p: None,
seed: None,
frequency_penalty: None,
presence_penalty: None,
stream_options: None,
top_k: None,
repetition_penalty: None,
min_p: None,
logprobs: None,
top_logprobs: None,
logit_bias: None,
parallel_tool_calls: None,
hf2q_overflow_policy: None,
hf2q_enable_thinking: None,
chat_template_kwargs: None,
}
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn a5d_chat_completions_stream_handler_returns_429_application_json_not_sse_when_kv_budget_exceeded(
) {
let per_slot_kv_budget_bytes: u64 = 1024 * 1024;
let kv_bytes_per_token_cached: u64 = 1024;
let engine = make_synthetic_engine_over_budget(
LoadedArch::Gemma,
per_slot_kv_budget_bytes,
kv_bytes_per_token_cached,
);
let state = AppState::new(ServerConfig::default());
let req = minimal_request( true, 1000);
let prepared = build_prepared_context(engine, vec![0u32; 1000], 1000);
let response = chat_completions_stream(state, req, prepared).await;
assert_eq!(
response.status(),
StatusCode::TOO_MANY_REQUESTS,
"iter-A5d Critical #2: streaming handler MUST return 429 \
when over-budget; got {}",
response.status()
);
let retry_after = response
.headers()
.get(header::RETRY_AFTER)
.and_then(|v| v.to_str().ok());
assert_eq!(
retry_after,
Some("1"),
"iter-A5d Critical #2: streaming handler MUST set Retry-After: 1"
);
let ct = response
.headers()
.get(header::CONTENT_TYPE)
.and_then(|v| v.to_str().ok())
.unwrap_or("");
assert!(
ct.contains("application/json"),
"iter-A5d Critical #2 (load-bearing): streaming handler \
over-budget Response MUST have Content-Type \
application/json (proving handler short-circuited BEFORE \
SSE body construction); got Content-Type: {ct:?}"
);
assert!(
!ct.contains("text/event-stream"),
"iter-A5d Critical #2 (load-bearing): streaming handler \
over-budget Response MUST NOT have Content-Type \
text/event-stream (would prove handler entered SSE body \
construction — the regression codex flagged in iter-A5); \
got Content-Type: {ct:?}"
);
let body_bytes = to_bytes(response.into_body(), 1 << 20)
.await
.expect("collect body bytes");
let body_str = String::from_utf8_lossy(&body_bytes).into_owned();
assert!(
body_str.contains("slot_budget_exceeded"),
"iter-A5d Critical #2: streaming handler body MUST contain \
`slot_budget_exceeded` code; got: {body_str}"
);
assert!(
body_str.contains("2048000"),
"iter-A5d Critical #2: streaming handler body MUST embed \
needed_bytes=2048000 verbatim (parse_slot_budget_exceeded \
round-trip); got: {body_str}"
);
assert!(
body_str.contains("1048576"),
"iter-A5d Critical #2: streaming handler body MUST embed \
budget_bytes=1048576 verbatim; got: {body_str}"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn a5d_chat_completions_non_streaming_handler_returns_429_when_worker_signals_slot_budget_exceeded(
) {
let needed_bytes: u64 = 2_048_000;
let budget_bytes: u64 = 1_048_576;
let engine = make_synthetic_engine_with_slot_budget_exceeded_worker(
LoadedArch::Gemma,
needed_bytes,
budget_bytes,
);
let state = AppState::new(ServerConfig::default());
let req = minimal_request( false, 1000);
let prepared = build_prepared_context(engine, vec![0u32; 1000], 1000);
let response = chat_completions_with_prepared(state, req, prepared).await;
assert_eq!(
response.status(),
StatusCode::TOO_MANY_REQUESTS,
"iter-A5d Critical #2: non-streaming handler MUST return \
429 when worker signals slot_budget_exceeded; got {}",
response.status()
);
let retry_after = response
.headers()
.get(header::RETRY_AFTER)
.and_then(|v| v.to_str().ok());
assert_eq!(
retry_after,
Some("1"),
"iter-A5d Critical #2: non-streaming handler MUST set \
Retry-After: 1"
);
let ct = response
.headers()
.get(header::CONTENT_TYPE)
.and_then(|v| v.to_str().ok())
.unwrap_or("");
assert!(
ct.contains("application/json"),
"iter-A5d Critical #2: non-streaming handler body \
Content-Type MUST be application/json; got {ct:?}"
);
let body_bytes = to_bytes(response.into_body(), 1 << 20)
.await
.expect("collect body bytes");
let body_str = String::from_utf8_lossy(&body_bytes).into_owned();
assert!(
body_str.contains("slot_budget_exceeded"),
"iter-A5d Critical #2: non-streaming handler body MUST \
contain `slot_budget_exceeded` code; got: {body_str}"
);
assert!(
body_str.contains(&needed_bytes.to_string()),
"iter-A5d Critical #2: non-streaming handler body MUST \
embed needed_bytes={needed_bytes} verbatim; got: {body_str}"
);
assert!(
body_str.contains(&budget_bytes.to_string()),
"iter-A5d Critical #2: non-streaming handler body MUST \
embed budget_bytes={budget_bytes} verbatim; got: {body_str}"
);
}
}
#[cfg(test)]
mod iter229_ac5_http_tests {
use super::super::schema::{ChatMessage, MessageContent};
use super::render_chat_prompt_or_400;
use axum::body::to_bytes;
use axum::http::StatusCode;
fn msg(role: &str, content: &str) -> ChatMessage {
ChatMessage {
role: role.into(),
content: Some(MessageContent::Text(content.into())),
reasoning_content: None,
tool_calls: None,
tool_call_id: None,
name: None,
}
}
async fn assert_400_with(msgs: Vec<ChatMessage>, expect: &str) {
let resp = render_chat_prompt_or_400(
crate::core::chat_templates::QWEN3_CHATML,
&msgs,
None,
false,
None,
)
.expect_err("raise-site transcript must map to a Response");
assert_eq!(resp.status(), StatusCode::BAD_REQUEST, "expect={expect}");
let body = to_bytes(resp.into_body(), usize::MAX).await.unwrap();
let body_str = String::from_utf8_lossy(&body);
assert!(
body_str.contains(expect),
"400 body must carry the template's raise_exception message \
{expect:?}; got: {body_str}"
);
}
#[tokio::test]
async fn ac5_no_messages_400() {
assert_400_with(vec![], "No messages provided").await;
}
#[tokio::test]
async fn ac5_no_user_query_400() {
assert_400_with(vec![msg("system", "s")], "No user query found").await;
}
#[tokio::test]
async fn ac5_system_not_first_400() {
assert_400_with(
vec![msg("user", "u"), msg("system", "late")],
"System message must be at the beginning",
)
.await;
}
#[tokio::test]
async fn ac5_unexpected_role_400() {
assert_400_with(
vec![msg("user", "u"), msg("narrator", "x")],
"Unexpected message role",
)
.await;
}
#[tokio::test]
async fn ac5_reserved_kwarg_400_names_key() {
let mut kwargs = serde_json::Map::new();
kwargs.insert("bos_token".into(), serde_json::Value::Bool(true));
let resp = render_chat_prompt_or_400(
crate::core::chat_templates::QWEN3_CHATML,
&[msg("user", "hi")],
None,
false,
Some(&kwargs),
)
.expect_err("reserved kwarg must map to a Response");
assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
let body = to_bytes(resp.into_body(), usize::MAX).await.unwrap();
let body_str = String::from_utf8_lossy(&body);
assert!(
body_str.contains("reserved chat_template_kwargs key")
&& body_str.contains("bos_token"),
"got: {body_str}"
);
}
}
#[cfg(test)]
mod iter230_b_probe_tests {
use super::super::registry;
use super::super::schema::{ChatMessage, MessageContent};
fn msg(role: &str, content: &str) -> ChatMessage {
ChatMessage {
role: role.into(),
content: Some(MessageContent::Text(content.into())),
reasoning_content: None,
tool_calls: None,
tool_call_id: None,
name: None,
}
}
#[test]
fn probe_render_tail_matches_real_render_tail() {
let qwen_reg = registry::find_for("qwen3.6-anything").expect("qwen registration");
let tools = vec![super::super::schema::Tool {
tool_type: "function".into(),
function: super::super::schema::ToolFunction {
name: "read_file".into(),
description: Some("Read a file".into()),
parameters: Some(serde_json::json!(
{"type": "object", "properties": {"path": {"type": "string"}}}
)),
},
}];
let probe = [msg("user", "x")];
let real = [
msg("system", "You are an agent."),
msg("user", "Fix the bug in foo.rs"),
msg("assistant", "Done."),
msg("user", "Now add a test"),
];
let gemma_tmpl = include_str!("test_fixtures/gemma4-apex-embedded-chat-template.jinja");
for (tmpl, tmpl_name) in [
(crate::core::chat_templates::QWEN3_CHATML, "qwen3-chatml"),
(gemma_tmpl, "gemma4-apex-embedded"),
] {
for enable_thinking in [true, false] {
let p = super::engine::render_chat_prompt_with_tools(
tmpl,
&probe,
Some(&tools),
enable_thinking,
None,
)
.unwrap();
let r = super::engine::render_chat_prompt_with_tools(
tmpl,
&real,
Some(&tools),
enable_thinking,
None,
)
.unwrap();
let (pb, rb) = (p.as_bytes(), r.as_bytes());
let lcs = pb
.iter()
.rev()
.zip(rb.iter().rev())
.take_while(|(a, b)| a == b)
.count();
assert!(
lcs >= 16,
"{tmpl_name} enable_thinking={enable_thinking}: probe \
and real renders share only a {lcs}-byte suffix -- \
the generation-prompt tail is not message-independent"
);
assert_eq!(
registry::prompt_seeds_reasoning_open(&p, &qwen_reg),
registry::prompt_seeds_reasoning_open(&r, &qwen_reg),
"{tmpl_name} enable_thinking={enable_thinking}: \
detector must agree between probe and real renders"
);
if tmpl_name == "qwen3-chatml" {
assert_eq!(
registry::prompt_seeds_reasoning_open(&r, &qwen_reg),
enable_thinking,
"qwen seed must be exactly enable_thinking"
);
}
}
}
}
}