#[cfg(feature = "cuda")]
fn try_safetensors_cuda_backend(
state: &AppState,
request: &ChatCompletionRequest,
request_id: &str,
start: Instant,
cancel: &CancelToken,
) -> Option<Response> {
let model_lock = state.safetensors_cuda_model()?;
if let Some(r) = super::openai_handlers::reject_unsupported_ignore_eos(
state,
request,
"SafeTensors CUDA",
) {
return Some(r);
}
let tokenizer = match require_tokenizer(state) {
Ok(t) => t,
Err(r) => return Some(r),
};
let prompt = match crate::api::realize_handlers::format_chat_messages_for_state_thinking_tools(
state,
&request.messages,
Some(&request.model),
request.thinking(),
request.tools.as_deref(),
) {
Ok(p) => p,
Err(e) => return Some(fail_response(state, StatusCode::BAD_REQUEST, e.to_string())),
};
let input_ids = tokenizer.encode(&prompt);
let max_tokens = request.max_tokens.unwrap_or(256).min(4096) as usize;
let eos_id = 151645u32;
let mut model = match model_lock.lock() {
Ok(m) => m,
Err(e) => {
return Some(
(
axum::http::StatusCode::INTERNAL_SERVER_ERROR,
axum::Json(serde_json::json!({"error": format!("Model lock failed: {e}")})),
)
.into_response(),
);
}
};
let output_ids = match model.generate(&input_ids, max_tokens, eos_id) {
Ok(ids) => ids,
Err(e) => {
let msg = format!("SafeTensors CUDA generation failed: {e}");
return Some(
(
crate::api::generation_error_status(&e),
axum::Json(serde_json::json!({"error": msg})),
)
.into_response(),
);
}
};
let output_text = tokenizer.decode(&output_ids).unwrap_or_else(|_| String::from("[decode error]"));
let completion_tokens = output_ids.len();
let prompt_tokens = input_ids.len();
let (output_text, finish_reason) =
finalize_chat_text(output_text, request.stop.as_deref(), completion_tokens, max_tokens);
let body = format!(
r#"{{"id":"{}","object":"chat.completion","model":"{}","choices":[{{"index":0,"message":{{"role":"assistant","content":{}}},"finish_reason":"{}"}}],"usage":{{"prompt_tokens":{},"completion_tokens":{},"total_tokens":{}}}}}"#,
request_id,
request.model,
serde_json::to_string(&output_text).unwrap_or_default(),
finish_reason,
prompt_tokens,
completion_tokens,
prompt_tokens + completion_tokens,
);
Some((
[(axum::http::header::CONTENT_TYPE, "application/json")],
body,
).into_response())
}
#[cfg(feature = "cuda")]
async fn try_cuda_backend(
state: &AppState,
request: &ChatCompletionRequest,
request_id: &str,
trace_level: Option<&str>,
start: Instant,
cancel: &CancelToken,
) -> Option<Response> {
let ttft_trace = std::env::var("TTFT_TRACE").is_ok();
let t0 = if ttft_trace { Some(std::time::Instant::now()) } else { None };
let cuda_model_lock = state.cuda_model()?;
let tokenizer = match require_tokenizer(state) {
Ok(t) => t,
Err(r) => return Some(r),
};
let arch_hint = state.model_architecture();
let tokenized =
tokenize_chat_prompt(&tokenizer, &request.messages, arch_hint.as_deref(), request.thinking(), request.tools.as_deref(), state);
let prompt_ids =
match fit_serving_context(state, tokenized) {
Ok(ids) => ids,
Err(r) => return Some(r),
};
if let Some(t) = t0 {
eprintln!("[TTFT] {:>20}: {:>7.2}ms ({}tok)", "tokenize", t.elapsed().as_secs_f64() * 1000.0, prompt_ids.len());
}
let prompt_tokens = prompt_ids.len();
let q_config = chat_quantized_config(
request,
&tokenizer,
state.model_eos_token_id(),
state.should_trace(trace_level),
cancel,
);
let max_tokens = q_config.max_tokens;
if request.stream {
let (tx, rx) = tokio::sync::mpsc::channel::<Result<u32, String>>(16);
let (timing_tx, timing_rx) =
tokio::sync::oneshot::channel::<crate::api::PhaseTimings>();
if let Err(r) =
dispatch_cuda_stream(state, cuda_model_lock, prompt_ids, q_config, tx, timing_tx)
{
return Some(r);
}
return Some(true_streaming_sse_response(
rx,
tokenizer,
request_id.to_string(),
request.model.clone(),
state.metrics.clone(),
start,
max_tokens,
prompt_tokens,
Some(timing_rx),
request.stop.as_deref(),
));
}
let (timing_tx, timing_rx) = tokio::sync::oneshot::channel::<crate::api::PhaseTimings>();
let generated = if let Some(batch_tx) = state.cuda_batch_tx() {
cuda_batch_collect(state, batch_tx, prompt_ids, q_config, timing_tx, &tokenizer).await
} else {
cuda_direct_collect(state, cuda_model_lock, &prompt_ids, &q_config, timing_tx, &tokenizer, prompt_tokens)
};
let (token_ids, completion_tokens, response_text) = match generated {
Ok(g) => g,
Err(r) => return Some(r),
};
let latency = start.elapsed();
state.metrics.record_success(completion_tokens, latency);
let timings = timing_rx
.await
.ok()
.and_then(|phases| phases.to_timings(prompt_tokens, completion_tokens));
Some(build_chat_response(
request_id.to_string(),
request.model.clone(),
response_text,
prompt_tokens,
completion_tokens,
max_tokens,
request.stop.as_deref(),
trace_level,
latency,
request.tools.as_deref(),
request_tool_choice(request),
timings,
None,
))
}
#[cfg(feature = "cuda")]
#[allow(clippy::result_large_err)]
fn dispatch_cuda_stream(
state: &AppState,
cuda_model_lock: &Arc<std::sync::RwLock<crate::gguf::OwnedQuantizedModelCuda>>,
prompt_ids: Vec<u32>,
q_config: crate::gguf::QuantizedGenerateConfig,
tx: tokio::sync::mpsc::Sender<Result<u32, String>>,
timing_tx: tokio::sync::oneshot::Sender<crate::api::PhaseTimings>,
) -> Result<(), Response> {
if let Some(batch_tx) = state.cuda_batch_tx() {
let batch_req = super::cuda_batch_scheduler::CudaBatchRequest {
prompt_ids,
config: q_config,
token_tx: tx,
non_streaming: false,
enqueue_time: std::time::Instant::now(),
timing_tx: Some(timing_tx),
};
if let Err(e) = batch_tx.try_send(batch_req) {
state.record_admission_rejected();
return Err(fail_response(
state,
StatusCode::SERVICE_UNAVAILABLE,
format!("Batch queue full: {e}"),
));
}
} else {
let cuda_model_clone = cuda_model_lock.clone();
let sink_metrics = state.metrics.clone();
tokio::task::spawn_blocking(move || {
let mut cuda_model = cuda_model_clone.write().expect("operation failed");
let generate_start = std::time::Instant::now();
let sink =
crate::api::openai_handlers::streaming_token_sink(tx.clone(), sink_metrics);
let result = dense_cuda_turn(&mut cuda_model, &prompt_ids, &q_config, sink);
let _ = timing_tx.send(phase_split(&mut cuda_model, generate_start));
if let Err(e) = result {
let _ = tx.blocking_send(Err(e.to_string()));
}
});
}
Ok(())
}
#[cfg(feature = "cuda")]
async fn cuda_batch_collect(
state: &AppState,
batch_tx: &tokio::sync::mpsc::Sender<super::cuda_batch_scheduler::CudaBatchRequest>,
prompt_ids: Vec<u32>,
q_config: crate::gguf::QuantizedGenerateConfig,
timing_tx: tokio::sync::oneshot::Sender<crate::api::PhaseTimings>,
tokenizer: &BPETokenizer,
) -> Result<(Vec<u32>, usize, String), Response> {
let (tx, mut rx) = tokio::sync::mpsc::channel::<Result<u32, String>>(512);
let batch_req = super::cuda_batch_scheduler::CudaBatchRequest {
prompt_ids,
config: q_config,
token_tx: tx,
non_streaming: true, enqueue_time: std::time::Instant::now(),
timing_tx: Some(timing_tx),
};
if let Err(e) = batch_tx.try_send(batch_req) {
state.record_admission_rejected();
return Err(fail_response(
state,
StatusCode::SERVICE_UNAVAILABLE,
format!("Batch queue full: {e}"),
));
}
let mut tokens = Vec::new();
while let Some(result) = rx.recv().await {
match result {
Ok(token_id) => tokens.push(token_id),
Err(e) => return Err(fail_response(state, StatusCode::INTERNAL_SERVER_ERROR, e)),
}
}
let n = tokens.len();
let text = tokenizer.decode(&tokens).unwrap_or_else(|_| String::new());
Ok((tokens, n, clean_chat_output(&text)))
}
#[cfg(feature = "cuda")]
#[allow(clippy::result_large_err)]
fn cuda_direct_collect(
state: &AppState,
cuda_model_lock: &Arc<std::sync::RwLock<crate::gguf::OwnedQuantizedModelCuda>>,
prompt_ids: &[u32],
q_config: &crate::gguf::QuantizedGenerateConfig,
timing_tx: tokio::sync::oneshot::Sender<crate::api::PhaseTimings>,
tokenizer: &BPETokenizer,
prompt_tokens: usize,
) -> Result<(Vec<u32>, usize, String), Response> {
let mut cuda_model = cuda_model_lock.write().expect("operation failed");
let generate_start = std::time::Instant::now();
let generated = match dense_cuda_turn(&mut cuda_model, prompt_ids, q_config, |_| true) {
Ok(g) => g,
Err(e) => return Err(fail_response(state, crate::api::generation_error_status(&e), e)),
};
let _ = timing_tx.send(phase_split(&mut cuda_model, generate_start));
let tokens: Vec<u32> = generated.iter().skip(prompt_tokens).copied().collect();
let n = tokens.len();
let text = tokenizer.decode(&tokens).unwrap_or_else(|_| String::new());
Ok((tokens, n, clean_chat_output(&text)))
}
#[cfg(feature = "cuda")]
fn dense_cuda_turn(
cuda_model: &mut crate::gguf::OwnedQuantizedModelCuda,
prompt: &[u32],
config: &crate::gguf::QuantizedGenerateConfig,
mut on_token: impl FnMut(u32) -> bool,
) -> crate::error::Result<Vec<u32>> {
if config.trace {
return cuda_model.generate_gpu_resident_streaming(prompt, config, on_token);
}
let _ = cuda_model.take_phase_timings();
let mut session = crate::session::Session::new(
crate::gguf::dense_session_borrowed::BorrowedCudaForward::new(cuda_model),
);
crate::gguf::dense_session::dense_stream(&mut session, prompt, config, &mut on_token)
.map(|(tokens, _)| tokens)
}
#[cfg(feature = "cuda")]
fn phase_split(
cuda_model: &mut crate::gguf::OwnedQuantizedModelCuda,
generate_start: std::time::Instant,
) -> crate::api::PhaseTimings {
let mut phases = cuda_model.take_phase_timings();
if let Some(prefill_ms) = phases.prefill_ms {
let total_ms = generate_start.elapsed().as_secs_f64() * 1000.0;
phases.decode_ms = Some((total_ms - prefill_ms).max(0.0));
}
phases
}
fn dense_cpu_session(
model: &Arc<crate::gguf::OwnedQuantizedModel>,
) -> crate::gguf::dense_session::DenseSession {
crate::gguf::dense_session::DenseSession::new(
crate::gguf::dense_session::DenseForward::cpu(Arc::clone(model)),
)
}
fn try_quantized_backend(
state: &AppState,
request: &ChatCompletionRequest,
request_id: &str,
trace_level: Option<&str>,
start: Instant,
cancel: &CancelToken,
) -> Option<Response> {
let quantized_model = state.quantized_model()?;
let tokenizer = match require_tokenizer(state) {
Ok(t) => t,
Err(r) => return Some(r),
};
let arch_hint = state.model_architecture();
let prompt_ids =
match tokenize_chat_prompt(&tokenizer, &request.messages, arch_hint.as_deref(), request.thinking(), request.tools.as_deref(), state) {
Ok(ids) => ids,
Err(r) => return Some(r),
};
let prompt_tokens = prompt_ids.len();
let q_config = chat_quantized_config(
request,
&tokenizer,
state.model_eos_token_id(),
state.should_trace(trace_level),
cancel,
);
let max_tokens = match quantized_model.effective_max_tokens(prompt_tokens, q_config.max_tokens)
{
Ok(budget) => budget,
Err(e) => return Some(fail_response(state, crate::api::generation_error_status(&e), e)),
};
if request.stream {
let (tx, rx) = tokio::sync::mpsc::channel::<Result<u32, String>>(16);
let quantized_model_clone = quantized_model.clone();
let prompt_ids_clone = prompt_ids.clone();
let q_config_clone = q_config.clone();
let sink_metrics = state.metrics.clone();
tokio::task::spawn_blocking(move || {
let mut sink =
crate::api::openai_handlers::streaming_token_sink(tx.clone(), sink_metrics);
let result = if q_config_clone.trace {
quantized_model_clone
.generate_with_cache_streaming(&prompt_ids_clone, &q_config_clone, sink)
.map(drop)
} else {
crate::gguf::dense_session::dense_stream(
&mut dense_cpu_session(&quantized_model_clone),
&prompt_ids_clone,
&q_config_clone,
&mut sink,
)
.map(drop)
};
if let Err(e) = result {
let _ = tx.blocking_send(Err(e.to_string()));
}
});
return Some(true_streaming_sse_response(
rx,
tokenizer,
request_id.to_string(),
request.model.clone(),
state.metrics.clone(),
start,
max_tokens,
prompt_tokens,
None,
request.stop.as_deref(),
));
}
let generated = if q_config.trace {
quantized_model.generate_with_cache(&prompt_ids, &q_config)
} else {
crate::gguf::dense_session::dense_turn(
&mut dense_cpu_session(quantized_model),
&prompt_ids,
&q_config,
)
.map(|(tokens, _)| tokens)
};
let generated = match generated {
Ok(g) => g,
Err(e) => return Some(fail_response(state, crate::api::generation_error_status(&e), e)),
};
let token_ids: Vec<u32> = generated.iter().skip(prompt_tokens).copied().collect();
let completion_tokens = token_ids.len();
let text = match tokenizer.decode(&token_ids) {
Ok(t) => clean_chat_output(&t),
Err(e) => return Some(fail_response(state, StatusCode::INTERNAL_SERVER_ERROR, e)),
};
let latency = start.elapsed();
state.metrics.record_success(completion_tokens, latency);
Some(build_chat_response(
request_id.to_string(),
request.model.clone(),
text,
prompt_tokens,
completion_tokens,
max_tokens,
request.stop.as_deref(),
trace_level,
latency,
request.tools.as_deref(),
request_tool_choice(request),
None,
None,
))
}
fn try_apr_transformer_backend(
state: &AppState,
request: &ChatCompletionRequest,
request_id: &str,
trace_level: Option<&str>,
start: Instant,
cancel: &CancelToken,
) -> Option<Response> {
use crate::apr_transformer::GenerateConfig;
let apr_transformer = state.apr_transformer()?;
if let Some(r) = super::openai_handlers::reject_unsupported_ignore_eos(
state,
request,
"APR transformer (f32)",
) {
return Some(r);
}
let tokenizer = match require_tokenizer(state) {
Ok(t) => t,
Err(r) => return Some(r),
};
let arch_hint = state.model_architecture();
let prompt_ids =
match tokenize_chat_prompt(&tokenizer, &request.messages, arch_hint.as_deref(), request.thinking(), request.tools.as_deref(), state) {
Ok(ids) => ids,
Err(r) => return Some(r),
};
let prompt_tokens = prompt_ids.len();
let max_tokens = request.max_tokens.unwrap_or(256);
let gen_config = GenerateConfig {
max_tokens,
temperature: request.temperature.unwrap_or(0.7),
cancel: cancel.clone(),
..Default::default()
};
let generated = match apr_transformer.generate_with_cache(&prompt_ids, &gen_config) {
Ok(g) => g,
Err(e) => {
return Some(fail_response(
state,
super::generation_error_status(&e),
format!("APR generation failed: {e}"),
))
},
};
let token_ids: Vec<u32> = generated.iter().skip(prompt_tokens).copied().collect();
let completion_tokens = token_ids.len();
if request.stream {
state
.metrics
.record_success(completion_tokens, start.elapsed());
return Some(pregenerated_sse_response(
token_ids,
tokenizer,
request_id.to_string(),
request.model.clone(),
request.stop.as_deref(),
max_tokens,
prompt_tokens,
));
}
let text = match tokenizer.decode(&token_ids) {
Ok(t) => clean_chat_output(&t),
Err(e) => return Some(fail_response(state, StatusCode::INTERNAL_SERVER_ERROR, e)),
};
let latency = start.elapsed();
state.metrics.record_success(completion_tokens, latency);
Some(build_chat_response(
request_id.to_string(),
request.model.clone(),
text,
prompt_tokens,
completion_tokens,
max_tokens,
request.stop.as_deref(),
trace_level,
latency,
request.tools.as_deref(),
request_tool_choice(request),
None,
None,
))
}
fn convert_token_ids(ids: &[usize]) -> Result<Vec<u32>, String> {
ids.iter()
.map(|&id| u32::try_from(id).map_err(|_| format!("Token ID {id} exceeds u32 range")))
.collect()
}
fn build_gen_config(request: &ChatCompletionRequest) -> GenerationConfig {
crate::api::realize_handlers::resolve_dense_generation_config(
request.temperature.unwrap_or(0.7),
request.top_p,
request.max_tokens.unwrap_or(256),
)
}
#[allow(clippy::result_large_err)]
fn registry_prompt_ids(
state: &AppState,
request: &ChatCompletionRequest,
tokenizer: &BPETokenizer,
) -> Result<Vec<u32>, Response> {
let prompt_text = match crate::api::realize_handlers::format_chat_messages_for_state_thinking_tools(
state,
&request.messages,
Some(&request.model),
request.thinking(),
request.tools.as_deref(),
) {
Ok(p) => p,
Err(e) => return Err(fail_response(state, StatusCode::BAD_REQUEST, e.to_string())),
};
let prompt_ids = tokenizer.encode(&prompt_text);
if prompt_ids.is_empty() {
return Err(fail_response(state, StatusCode::BAD_REQUEST, "Messages cannot be empty"));
}
Ok(prompt_ids)
}
#[allow(clippy::result_large_err)]
fn registry_token_ids<E: std::fmt::Display>(
state: &AppState,
generated: Result<Vec<usize>, E>,
) -> Result<Vec<u32>, Response> {
let generated = match generated {
Ok(g) => g,
Err(e) => return Err(fail_response(state, StatusCode::INTERNAL_SERVER_ERROR, e)),
};
convert_token_ids(&generated).map_err(|e| fail_response(state, StatusCode::BAD_REQUEST, e))
}
fn registry_fallback(
state: &AppState,
request: &ChatCompletionRequest,
request_id: &str,
start: Instant,
cancel: &CancelToken,
) -> Response {
if let Some(r) = super::openai_handlers::reject_unsupported_ignore_eos(
state,
request,
"dense registry",
) {
return r;
}
let model_id = if request.model == "default" || request.model.is_empty() {
None
} else {
Some(request.model.as_str())
};
let (model, tokenizer) = match state.get_model(model_id) {
Ok((m, t)) => (m, t),
Err(e) => return fail_response(state, super::model_resolution_status(&e), e),
};
let prompt_ids = match registry_prompt_ids(state, request, &tokenizer) {
Ok(ids) => ids,
Err(r) => return r,
};
let prompt_tokens = prompt_ids.len();
let prompt: Vec<usize> = prompt_ids.iter().map(|&id| id as usize).collect();
let config = build_gen_config(request).with_cancel(cancel.clone());
let token_ids: Vec<u32> = match registry_token_ids(state, model.generate(&prompt, &config)) {
Ok(ids) => ids,
Err(r) => return r,
};
let generated_ids: Vec<u32> = token_ids[prompt.len()..].to_vec();
let completion_tokens = generated_ids.len();
if request.stream {
state
.metrics
.record_success(completion_tokens, start.elapsed());
return pregenerated_sse_response(
generated_ids,
tokenizer,
request_id.to_string(),
request.model.clone(),
request.stop.as_deref(),
request.max_tokens.unwrap_or(256),
prompt_tokens,
);
}
let response_text = match tokenizer.decode(&generated_ids) {
Ok(t) => t,
Err(e) => return fail_response(state, StatusCode::INTERNAL_SERVER_ERROR, e),
};
let duration = start.elapsed();
state.metrics.record_success(completion_tokens, duration);
let max_tokens = request.max_tokens.unwrap_or(256);
build_chat_response(
request_id.to_string(),
request.model.clone(),
response_text,
prompt_tokens,
completion_tokens,
max_tokens,
request.stop.as_deref(),
None,
duration,
request.tools.as_deref(),
request_tool_choice(request),
None,
None,
)
}
fn model_loaded_at_unix_secs(state: &AppState) -> i64 {
state.clock().started_unix_secs()
}
pub async fn openai_models_handler(State(state): State<AppState>) -> Json<OpenAIModelsResponse> {
let created = model_loaded_at_unix_secs(&state);
let models = if let Some(registry) = &state.registry {
registry
.list()
.into_iter()
.map(|m| OpenAIModel {
id: m.id,
object: "model".to_string(),
created,
owned_by: "realizar".to_string(),
})
.collect()
} else {
vec![OpenAIModel {
id: "default".to_string(),
object: "model".to_string(),
created,
owned_by: "realizar".to_string(),
}]
};
Json(OpenAIModelsResponse {
object: "list".to_string(),
data: models,
})
}
#[cfg(feature = "cuda")]
async fn try_apr_q4k_chat_backend(
state: &AppState,
request: &ChatCompletionRequest,
request_id: &str,
trace_level: Option<&str>,
start: Instant,
cancel: &crate::generate::CancelToken,
) -> Option<Response> {
use crate::api::apr_q4k_scheduler::AprQ4kRequest;
let q4k_tx = state.apr_q4k_tx()?;
let tokenizer = match require_tokenizer(state) {
Ok(t) => t,
Err(r) => return Some(r),
};
let arch_hint = state.model_architecture();
let prompt_ids =
match tokenize_chat_prompt(&tokenizer, &request.messages, arch_hint.as_deref(), request.thinking(), request.tools.as_deref(), state) {
Ok(ids) => ids,
Err(r) => return Some(r),
};
let prompt_tokens = prompt_ids.len();
let (max_tokens, temperature, _eos_single) =
chat_gen_params(request, &tokenizer, state.model_eos_token_id());
let eos_ids = if request.ignore_eos.unwrap_or(false) {
Vec::new()
} else {
state.model_eos_ids()
};
let (response_tx, response_rx) = tokio::sync::oneshot::channel();
if q4k_tx
.send(AprQ4kRequest {
prompt_ids,
max_tokens,
temperature,
seed: request.seed.unwrap_or(crate::sampling::DEFAULT_SEED),
eos_ids,
cancel: cancel.clone(),
response_tx,
})
.await
.is_err()
{
return Some(fail_response(
state,
StatusCode::INTERNAL_SERVER_ERROR,
"Q4K thread unavailable",
));
}
let result = match response_rx.await {
Ok(r) => r,
Err(_) => {
return Some(fail_response(
state,
StatusCode::INTERNAL_SERVER_ERROR,
"Q4K thread dropped response",
))
}
};
let resp = match result {
Ok(r) => r,
Err(e) => {
return Some(fail_response(
state,
StatusCode::INTERNAL_SERVER_ERROR,
format!("Q4K generation failed: {e}"),
))
}
};
let text = match tokenizer.decode(&resp.output_tokens) {
Ok(t) => clean_chat_output(&t),
Err(e) => {
return Some(fail_response(
state,
StatusCode::INTERNAL_SERVER_ERROR,
e,
))
}
};
let completion_tokens = resp.tokens_generated;
state
.metrics
.record_success(completion_tokens, start.elapsed());
Some(build_chat_response(
request_id.to_string(),
request.model.clone(),
text,
prompt_tokens,
completion_tokens,
max_tokens,
request.stop.as_deref(),
trace_level,
start.elapsed(),
request.tools.as_deref(),
request_tool_choice(request),
None,
None,
))
}
pub async fn openai_chat_completions_handler(
State(state): State<AppState>,
headers: HeaderMap,
Extension(cancel): Extension<CancelToken>,
Json(request): Json<ChatCompletionRequest>,
) -> Response {
let start = Instant::now();
if state.is_verbose() {
let msg_count = request.messages.len();
let last_msg = request
.messages
.last()
.map(|m| m.content.chars().take(50).collect::<String>())
.unwrap_or_default();
eprintln!(
"[VERBOSE] POST /v1/chat/completions model={} messages={} last={:?}",
request.model, msg_count, last_msg
);
}
let trace_level = headers
.get("X-Trace-Level")
.and_then(|v| v.to_str().ok())
.map(str::to_lowercase);
let request_id = format!(
"chatcmpl-q4k-{}",
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_millis()
);
if let Some(reason) = request.thinking_conflict() {
return fail_response(&state, StatusCode::BAD_REQUEST, reason);
}
if let Some(r) = try_qwen35_backend(&state, &request, &request_id, start, &cancel).await {
return r;
}
if let Some(r) = try_qwen3_moe_backend(&state, &request, &request_id, start, &cancel) {
return r;
}
#[cfg(feature = "gpu")]
if let Some(r) = try_gpu_backend(
&state,
&request,
&request_id,
trace_level.as_deref(),
start,
&cancel,
) {
return r;
}
#[cfg(feature = "gpu")]
if let Some(r) = try_cached_backend(
&state,
&request,
&request_id,
trace_level.as_deref(),
start,
&cancel,
) {
return r;
}
#[cfg(feature = "cuda")]
if let Some(r) =
try_cuda_backend(&state, &request, &request_id, trace_level.as_deref(), start, &cancel).await
{
return r;
}
#[cfg(feature = "cuda")]
if let Some(r) = try_apr_q4k_chat_backend(&state, &request, &request_id, trace_level.as_deref(), start, &cancel).await {
return r;
}
#[cfg(feature = "cuda")]
if let Some(r) = try_safetensors_cuda_backend(&state, &request, &request_id, start, &cancel) {
return r;
}
cpu_chat_backends(
&state,
&request,
&request_id,
trace_level.as_deref(),
start,
&cancel,
)
}
fn cpu_chat_backends(
state: &AppState,
request: &ChatCompletionRequest,
request_id: &str,
trace_level: Option<&str>,
start: Instant,
cancel: &CancelToken,
) -> Response {
if let Some(r) =
try_quantized_backend(state, request, request_id, trace_level, start, cancel)
{
return r;
}
if let Some(r) =
try_apr_transformer_backend(state, request, request_id, trace_level, start, cancel)
{
return r;
}
registry_fallback(state, request, request_id, start, cancel)
}
fn stop_tokens_unless_ignore_eos(
request: &ChatCompletionRequest,
eos: impl IntoIterator<Item = u32>,
) -> Vec<u32> {
if request.ignore_eos.unwrap_or(false) {
Vec::new()
} else {
eos.into_iter().collect()
}
}
#[allow(clippy::result_large_err)]
fn moe_models(
state: &AppState,
raw_arch: &str,
) -> Result<(Arc<crate::gguf::MappedGGUFModel>, Arc<crate::gguf::OwnedQuantizedModel>), Response> {
let mapped = match state.mapped_gguf_model() {
Some(m) => m,
None => {
eprintln!(
"[WARN] aprender#1789: qwen3_moe arch detected at \
/v1/chat/completions (raw_arch={raw_arch}, canonical=qwen3_moe) \
but AppState has no retained MappedGGUFModel. This means the \
CLI server-command load path didn't call \
.with_mapped_gguf_model(). Returning NOT_IMPLEMENTED. \
See contracts/qwen3-moe-serve-dispatch-v1.yaml + \
https://github.com/paiml/aprender/issues/1789"
);
return Err(fail_response(
state,
StatusCode::NOT_IMPLEMENTED,
"qwen3_moe arch detected but mapped GGUF not retained in AppState. \
See aprender#1789 + contracts/qwen3-moe-serve-dispatch-v1.yaml.",
));
}
};
let quantized = match state.quantized_model() {
Some(q) => q.clone(),
None => {
return Err(fail_response(
state,
StatusCode::NOT_IMPLEMENTED,
"qwen3_moe arch detected but no OwnedQuantizedModel in AppState. \
See aprender#1789.",
));
}
};
Ok((mapped, quantized))
}
fn moe_gen_config(
state: &AppState,
request: &ChatCompletionRequest,
tokenizer: &BPETokenizer,
max_tokens: usize,
cancel: &CancelToken,
) -> crate::gguf::QuantizedGenerateConfig {
use crate::gguf::QuantizedGenerateConfig;
let defaults = QuantizedGenerateConfig::default();
let eos_id = state.model_eos_token_id().or_else(|| {
tokenizer
.get_token_id("<|im_end|>")
.or_else(|| tokenizer.get_token_id("<|endoftext|>"))
});
let stop_tokens: Vec<u32> = stop_tokens_unless_ignore_eos(request, eos_id);
QuantizedGenerateConfig {
max_tokens,
temperature: request.temperature.unwrap_or(defaults.temperature),
top_k: request.top_k.unwrap_or(defaults.top_k),
top_p: request.top_p.unwrap_or(defaults.top_p),
repeat_penalty: request.repeat_penalty.unwrap_or(defaults.repeat_penalty),
repeat_last_n: request.repeat_last_n.unwrap_or(defaults.repeat_last_n),
seed: request.seed.unwrap_or(defaults.seed),
stop_tokens,
cancel: cancel.clone(),
..defaults
}
}
#[allow(clippy::too_many_arguments)]
fn moe_stream_cpu(
state: &AppState,
request: &ChatCompletionRequest,
request_id: &str,
start: Instant,
mapped: &Arc<crate::gguf::MappedGGUFModel>,
quantized: &Arc<crate::gguf::OwnedQuantizedModel>,
input_ids: &[u32],
gen_config: &crate::gguf::QuantizedGenerateConfig,
tokenizer: Arc<BPETokenizer>,
max_tokens: usize,
prompt_token_count: usize,
) -> Response {
let (tx, rx) = tokio::sync::mpsc::channel::<Result<u32, String>>(64);
let mapped_clone = mapped.clone();
let quantized_clone = quantized.clone();
let input_ids_clone = input_ids.to_vec();
let gen_config_clone = gen_config.clone();
let sink_metrics = state.metrics.clone();
tokio::task::spawn_blocking(move || {
let result = crate::infer::qwen3_moe_generate::run_qwen3_moe_generate_streaming(
&mapped_clone,
&quantized_clone,
&input_ids_clone,
&gen_config_clone,
crate::api::openai_handlers::streaming_token_sink(tx.clone(), sink_metrics),
);
if let Err(e) = result {
let _ = tx.blocking_send(Err(e.to_string()));
}
});
crate::api::openai_handlers::true_streaming_sse_response(
rx,
tokenizer,
request_id.to_string(),
request.model.clone(),
state.metrics.clone(),
start,
max_tokens,
prompt_token_count,
None,
request.stop.as_deref(),
)
}
fn try_qwen3_moe_backend(
state: &AppState,
request: &ChatCompletionRequest,
request_id: &str,
start: Instant,
cancel: &CancelToken,
) -> Option<Response> {
use crate::gguf::QuantizedGenerateConfig;
let raw_arch = state.model_architecture()?;
if !is_qwen3_moe_arch(&raw_arch) {
return None;
}
let (mapped, quantized) = match moe_models(state, &raw_arch) {
Ok(m) => m,
Err(r) => return Some(r),
};
let tokenizer = match require_tokenizer(state) {
Ok(t) => t,
Err(r) => return Some(r),
};
let input_ids = match tokenize_chat_prompt(
&tokenizer,
&request.messages,
Some(&request.model),
request.thinking(),
request.tools.as_deref(),
state,
) {
Ok(ids) => ids,
Err(r) => return Some(r),
};
let prompt_token_count = input_ids.len();
let max_tokens = request.max_tokens.unwrap_or(256).min(4096) as usize;
let gen_config = moe_gen_config(state, request, &tokenizer, max_tokens, cancel);
if request.stream && state.moe_no_gpu() {
return Some(moe_stream_cpu(
state, request, request_id, start, &mapped, &quantized, &input_ids, &gen_config, tokenizer,
max_tokens, prompt_token_count,
));
}
let (tokens, used_gpu) = match crate::infer::qwen3_moe_dispatch::run_qwen3_moe_generate_dispatch(
&mapped,
&quantized,
&input_ids,
&gen_config,
state.moe_no_gpu(),
) {
Ok(t) => t,
Err(e) => {
state.metrics.record_failure();
return Some(fail_response(
state,
StatusCode::INTERNAL_SERVER_ERROR,
format!("qwen3_moe generation failed: {e}"),
));
}
};
let generated_ids: Vec<u32> = tokens[input_ids.len()..].to_vec();
let completion_tokens = generated_ids.len();
if request.stream {
state.metrics.record_success(completion_tokens, start.elapsed());
return Some(pregenerated_sse_response(
generated_ids,
tokenizer,
request_id.to_string(),
request.model.clone(),
request.stop.as_deref(),
max_tokens,
prompt_token_count,
));
}
let response_text = match tokenizer.decode(&generated_ids) {
Ok(t) => clean_chat_output(&t),
Err(e) => {
state.metrics.record_failure();
return Some(fail_response(state, StatusCode::INTERNAL_SERVER_ERROR, e));
}
};
let duration = start.elapsed();
state.metrics.record_success(completion_tokens, duration);
Some(build_chat_response(
request_id.to_string(),
request.model.clone(),
response_text,
prompt_token_count,
completion_tokens,
max_tokens,
request.stop.as_deref(),
None,
duration,
request.tools.as_deref(),
request_tool_choice(request),
None,
Some(used_gpu),
))
}
fn is_qwen3_moe_arch(raw_arch: &str) -> bool {
crate::tensor_names::normalize_architecture(raw_arch) == "qwen3_moe"
}
#[cfg(test)]
mod qwen3_moe_dispatch_guard_tests {
use super::is_qwen3_moe_arch;
#[test]
fn canonical_qwen3_moe_matches() {
assert!(is_qwen3_moe_arch("qwen3_moe"));
}
#[test]
fn huggingface_class_names_canonicalize() {
assert!(is_qwen3_moe_arch("Qwen3MoeForCausalLM"));
assert!(is_qwen3_moe_arch("Qwen3MoEForCausalLM"));
assert!(is_qwen3_moe_arch("Qwen3CoderForCausalLM"));
assert!(is_qwen3_moe_arch("Qwen3_5MoeForCausalLM"));
assert!(is_qwen3_moe_arch("Qwen3_5MoeForConditionalGeneration"));
}
#[test]
fn lowercase_underscore_variants_match() {
assert!(is_qwen3_moe_arch("qwen3moe"));
}
#[test]
fn dense_archs_do_not_match() {
assert!(!is_qwen3_moe_arch("qwen2"));
assert!(!is_qwen3_moe_arch("qwen3"));
assert!(!is_qwen3_moe_arch("llama"));
assert!(!is_qwen3_moe_arch("mistral"));
assert!(!is_qwen3_moe_arch("phi"));
assert!(!is_qwen3_moe_arch("gemma"));
}
#[test]
fn unknown_arch_does_not_match() {
assert!(!is_qwen3_moe_arch("some-future-arch-3000"));
assert!(!is_qwen3_moe_arch(""));
}
}
#[cfg(all(test, feature = "gpu"))]
#[path = "tests/dense_session_4268.rs"]
mod dense_session_4268;