Skip to content
Closed
43 changes: 39 additions & 4 deletions crates/apr-cli/src/commands/serve/chat.rs
Original file line number Diff line number Diff line change
Expand Up @@ -266,7 +266,7 @@ pub(crate) async fn safetensors_chat_completions_handler(
.get("temperature")
.and_then(|t| t.as_f64())
.unwrap_or(0.0) as f32;
let output_ids = {
let (output_ids, max_tokens) = {
// PMAT-189: Handle transformer lock poisoning gracefully
let t = match transformer.lock() {
Ok(guard) => guard,
Expand All @@ -280,8 +280,12 @@ pub(crate) async fn safetensors_chat_completions_handler(
.into_response();
}
};
match st_cpu_generate(&t, &input_ids, max_tokens, temperature) {
Ok(ids) => ids,
let budget = match st_context_budget(&t, input_ids.len(), max_tokens) {
Ok(budget) => budget,
Err(refusal) => return refusal,
};
match st_cpu_generate(&t, &input_ids, budget, temperature) {
Ok(ids) => (ids, budget),
Err(e) => {
return (
StatusCode::INTERNAL_SERVER_ERROR,
Expand Down Expand Up @@ -336,6 +340,7 @@ pub(crate) async fn safetensors_chat_completions_handler(
stream_mode,
input_ids.len(),
tokens_generated,
max_tokens,
elapsed,
tok_per_sec,
)
Expand Down Expand Up @@ -419,6 +424,7 @@ fn build_chat_response(
stream_mode: bool,
prompt_tokens: usize,
tokens_generated: usize,
max_tokens: usize,
elapsed: std::time::Duration,
tok_per_sec: f64,
) -> axum::response::Response {
Expand All @@ -427,7 +433,12 @@ fn build_chat_response(

let request_id = generate_request_id();
let has_tool_calls = tool_calls.is_some();
let finish_reason = if has_tool_calls { "tool_calls" } else { "stop" };
// #3718: a reply cut at `max_tokens` is "length", never "stop".
let finish_reason = if has_tool_calls {
"tool_calls"
} else {
super::handlers::finish_reason_for(tokens_generated, max_tokens)
};
let created = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
Expand Down Expand Up @@ -688,6 +699,26 @@ mod chat_helper_tests {

// ---- build_chat_response ------------------------------------------

/// #3718: a SafeTensors reply that used its whole `max_tokens` budget was cut,
/// and says "length"; the hardcoded "stop" read it as finished.
#[tokio::test]
async fn a_reply_cut_at_max_tokens_is_length_not_stop() {
use axum::body::to_bytes;
let resp = build_chat_response(
"cut".to_string(),
None,
false,
5,
16,
16,
std::time::Duration::from_millis(10),
300.0,
);
let bytes = to_bytes(resp.into_body(), 64 * 1024).await.expect("body");
let v: serde_json::Value = serde_json::from_slice(&bytes).expect("json");
assert_eq!(v["choices"][0]["finish_reason"], "length");
}

#[tokio::test]
async fn build_chat_response_non_streaming_json_body() {
use axum::body::to_bytes;
Expand All @@ -697,6 +728,7 @@ mod chat_helper_tests {
false,
5,
3,
16,
std::time::Duration::from_millis(10),
300.0,
);
Expand Down Expand Up @@ -727,6 +759,7 @@ mod chat_helper_tests {
false,
2,
0,
16,
std::time::Duration::from_millis(1),
0.0,
);
Expand All @@ -748,6 +781,7 @@ mod chat_helper_tests {
true,
1,
1,
16,
std::time::Duration::from_millis(1),
1.0,
);
Expand All @@ -773,6 +807,7 @@ mod chat_helper_tests {
true,
1,
1,
16,
std::time::Duration::from_millis(1),
1.0,
);
Expand Down
103 changes: 103 additions & 0 deletions crates/apr-cli/src/commands/serve/context_budget_3718.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,103 @@
// #3718 done_when 3: a prompt that does not fit the context window says so,
// never a silent cut. The wgpu handler capped `max_tokens` at 4096 and nothing
// else, so a prompt at or past the context window prefilled anyway and the
// decode loop grew the KV cache past what the model was trained on, reporting
// `finish_reason: "stop"`. These two pure functions are the rule it now follows,
// the same rule the CPU (`effective_max_tokens`) and Qwen3.5 (`Session`) paths
// already enforce. The SafeTensors handlers (chat.rs, simple.rs) apply it too,
// through `st_context_budget`. They are not gated on `wgpu`, so default-feature
// CI tests them.

/// Tokens the reply may use once the prompt is in a `context_length` window:
/// `min(requested, context_length - prompt_len)`.
///
/// # Errors
/// `(prompt_len, context_length)` when the prompt leaves no room for even one
/// generated token. That is decided by the request alone, so the caller refuses
/// it as a client error rather than truncating.
#[cfg_attr(not(any(feature = "wgpu", feature = "inference")), allow(dead_code))]
pub(super) fn context_token_budget(
prompt_len: usize,
requested: usize,
context_length: usize,
) -> std::result::Result<usize, (usize, usize)> {
if prompt_len >= context_length {
return Err((prompt_len, context_length));
}
Ok(requested.min(context_length - prompt_len))
}

/// The OpenAI-shaped 400 body for a prompt refused for length
/// (`error.code = "context_length_exceeded"`), so a client can tell it from a
/// server fault without parsing the message.
#[cfg_attr(not(any(feature = "wgpu", feature = "inference")), allow(dead_code))]
pub(super) fn context_length_exceeded_body(prompt_len: usize, context_length: usize) -> serde_json::Value {
serde_json::json!({
"error": {
"message": format!(
"prompt is {prompt_len} tokens; the model's context window is {context_length}, \
which leaves no room to generate. The prompt was refused whole, not truncated."
),
"type": "invalid_request_error",
"param": "messages",
"code": "context_length_exceeded",
"prompt_tokens": prompt_len,
"context_length": context_length,
}
})
}

/// OpenAI `finish_reason` for a reply of `generated` tokens under a `max_tokens`
/// budget: `"length"` when the budget ran out, else `"stop"`. Every serve path
/// (wgpu, CUDA, the CUDA->CPU fallback) reports through this one rule, so a cut
/// reply is never labelled a natural stop.
pub(super) fn finish_reason_for(generated: usize, max_tokens: usize) -> &'static str {
if generated >= max_tokens {
"length"
} else {
"stop"
}
}

#[cfg(test)]
mod tests_context_budget_3718 {
use super::{context_length_exceeded_body, context_token_budget, finish_reason_for};

#[test]
fn a_reply_that_used_the_whole_budget_is_length_not_stop() {
assert_eq!(finish_reason_for(64, 64), "length");
assert_eq!(finish_reason_for(65, 64), "length");
assert_eq!(finish_reason_for(63, 64), "stop");
assert_eq!(finish_reason_for(0, 64), "stop");
}

#[test]
fn budget_is_the_request_when_it_fits() {
assert_eq!(context_token_budget(10, 64, 2048), Ok(64));
}

#[test]
fn budget_is_clamped_to_the_room_left() {
// 2040 prompt tokens in a 2048 window leave 8, whatever was asked.
assert_eq!(context_token_budget(2040, 64, 2048), Ok(8));
assert_eq!(context_token_budget(2047, 4096, 2048), Ok(1));
}

#[test]
fn a_prompt_that_fills_the_window_is_refused_not_cut() {
assert_eq!(context_token_budget(2048, 64, 2048), Err((2048, 2048)));
assert_eq!(context_token_budget(9000, 1, 2048), Err((9000, 2048)));
}

#[test]
fn refusal_body_names_the_code_and_both_counts() {
let body = context_length_exceeded_body(9000, 2048);
let e = &body["error"];
assert_eq!(e["code"], "context_length_exceeded");
assert_eq!(e["type"], "invalid_request_error");
assert_eq!(e["prompt_tokens"], 9000);
assert_eq!(e["context_length"], 2048);
let msg = e["message"].as_str().expect("message is a string");
assert!(msg.contains("9000") && msg.contains("2048") && msg.contains("not truncated"));
}
}
14 changes: 10 additions & 4 deletions crates/apr-cli/src/commands/serve/handler_gpu_completion.rs
Original file line number Diff line number Diff line change
Expand Up @@ -138,7 +138,7 @@ async fn gpu_cpu_fallback(
.await;

match result {
Ok(Ok(out)) => build_cpu_fallback_response(&out, start),
Ok(Ok(out)) => build_cpu_fallback_response(&out, max_tokens, start),
Ok(Err(cpu_err)) => {
Json(serde_json::json!({
"error": format!("GPU failed: {gpu_err}; CPU fallback also failed: {cpu_err}")
Expand All @@ -158,7 +158,11 @@ async fn gpu_cpu_fallback(
#[cfg_attr(coverage_nightly, coverage(off))]
#[cfg(all(feature = "inference", feature = "cuda"))]
#[allow(clippy::disallowed_methods)]
fn build_cpu_fallback_response(out: &AprInferenceOutput, start: Instant) -> axum::response::Response {
fn build_cpu_fallback_response(
out: &AprInferenceOutput,
max_tokens: usize,
start: Instant,
) -> axum::response::Response {
use axum::{response::IntoResponse, Json};

let request_id = generate_request_id();
Expand All @@ -171,7 +175,8 @@ fn build_cpu_fallback_response(out: &AprInferenceOutput, start: Instant) -> axum
"object": "chat.completion",
"created": created,
"model": "apr-cpu-fallback",
"choices": [{"index": 0, "message": {"role": "assistant", "content": out.text}, "finish_reason": "stop"}],
// #3718: a reply cut at the budget is "length", never "stop".
"choices": [{"index": 0, "message": {"role": "assistant", "content": out.text}, "finish_reason": finish_reason_for(out.tokens_generated, max_tokens)}],
"usage": {
"prompt_tokens": out.input_token_count,
"completion_tokens": out.tokens_generated,
Expand Down Expand Up @@ -324,7 +329,8 @@ async fn handle_gpu_chat_completion(
"object": "chat.completion",
"created": created,
"model": &response_model,
"choices": [{"index": 0, "message": {"role": "assistant", "content": output_text}, "finish_reason": "stop"}],
// #3718: a reply cut at the budget is "length", never "stop".
"choices": [{"index": 0, "message": {"role": "assistant", "content": output_text}, "finish_reason": finish_reason_for(tokens_generated, max_tokens_clamped)}],
"usage": {
"prompt_tokens": input_tokens.len(),
"completion_tokens": tokens_generated,
Expand Down
37 changes: 30 additions & 7 deletions crates/apr-cli/src/commands/serve/handlers.rs
Original file line number Diff line number Diff line change
Expand Up @@ -37,6 +37,8 @@ struct WgpuInferenceState {
num_layers: usize,
vocab_size: usize,
hidden_dim: usize,
/// #3718: the model's context window; a prompt that fills it is refused.
context_length: usize,
}

/// PMAT-355: How one character of a GPT-2 byte-level BPE token maps to bytes.
Expand Down Expand Up @@ -230,12 +232,13 @@ fn wgpu_stream_done_chunk(
id: &str,
prompt_len: usize,
completion_tokens: u32,
finish_reason: &str,
elapsed: std::time::Duration,
) -> String {
let tok_s = wgpu_tokens_per_second(completion_tokens as f64, elapsed);
serde_json::json!({
"id": id, "object": "chat.completion.chunk", "model": "qwen-wgpu",
"choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}],
"choices": [{"index": 0, "delta": {}, "finish_reason": finish_reason}],
"usage": {"prompt_tokens": prompt_len, "completion_tokens": completion_tokens,
"total_tokens": prompt_len as u32 + completion_tokens},
"x_wgpu_tok_s": tok_s,
Expand Down Expand Up @@ -281,7 +284,15 @@ fn wgpu_stream_generate(
}
}

let done = wgpu_stream_done_chunk(id, prompt_ids.len(), completion_tokens, gen_start.elapsed());
// #3718: a reply cut at the budget is "length", never "stop".
let finish_reason = finish_reason_for(completion_tokens as usize, max_tokens);
let done = wgpu_stream_done_chunk(
id,
prompt_ids.len(),
completion_tokens,
finish_reason,
gen_start.elapsed(),
);
let _ = tx.blocking_send(done);
let _ = tx.blocking_send("[DONE]".to_string());
}
Expand Down Expand Up @@ -347,11 +358,7 @@ fn wgpu_chat_completion_blocking(
.map(|&tok| wgpu_detokenize_one(tok, &state.vocab))
.collect();
let tok_s = wgpu_tokens_per_second(output_ids.len() as f64, elapsed);
let finish_reason = if output_ids.len() >= max_tokens {
"length"
} else {
"stop"
};
let finish_reason = finish_reason_for(output_ids.len(), max_tokens);

axum::Json(serde_json::json!({
"id": id, "object": "chat.completion", "model": "qwen-wgpu",
Expand All @@ -377,6 +384,20 @@ async fn wgpu_chat_completion(
let prompt_ids = wgpu_prompt_ids(&state, &body);
let id = wgpu_completion_id();

// #3718: refuse a prompt that fills the context window, and clamp the budget
// to the room left, before any prefill runs.
let max_tokens = match context_token_budget(prompt_ids.len(), max_tokens, state.context_length)
{
Ok(budget) => budget,
Err((prompt_len, context_length)) => {
return (
axum::http::StatusCode::BAD_REQUEST,
axum::Json(context_length_exceeded_body(prompt_len, context_length)),
)
.into_response();
}
};

if stream {
// PMAT-355: Streaming SSE via spawn_blocking + channel
wgpu_chat_completion_streaming(state, prompt_ids, max_tokens, id)
Expand Down Expand Up @@ -854,6 +875,7 @@ fn serve_wgpu_backend(
num_layers,
vocab_size,
hidden_dim: dims.hidden_dim,
context_length: quantized.config().context_length,
});

run_wgpu_server(build_wgpu_router(wgpu_state), config)?;
Expand Down Expand Up @@ -1715,6 +1737,7 @@ pub fn build_demo_streaming_apr_cpu_router_for_test() -> axum::Router {
build_apr_cpu_router(state, super::auth::AuthGate::disabled())
}

include!("context_budget_3718.rs");
include!("handler_apr_cpu_completion.rs");
include!("handler_gpu_completion.rs");
include!("server.rs");
29 changes: 29 additions & 0 deletions crates/apr-cli/src/commands/serve/safetensors.rs
Original file line number Diff line number Diff line change
Expand Up @@ -417,6 +417,35 @@ fn st_cpu_generate(
.map_err(|e| e.to_string())
}

/// #3718: the SafeTensors handlers apply the context rule the wgpu handler does,
/// BEFORE generating. Without it `Session` clamped the budget to the room left in
/// the window on its own and the reply was reported `"stop"`, and a prompt that
/// filled the window came back as a 500. Returns the budget to generate with and
/// to judge `finish_reason` against, or the 400 `context_length_exceeded` reply.
#[cfg(feature = "inference")]
fn st_context_budget(
model: &realizar::apr_transformer::AprTransformer,
prompt_len: usize,
max_tokens: usize,
) -> std::result::Result<usize, axum::response::Response> {
use axum::response::IntoResponse;
use realizar::session::ArchForward;
// The same length `Session` checks against (StCpuForward::context_length).
let context_length = realizar::safetensors_infer::StCpuForward::new(model).context_length();
super::handlers::context_token_budget(prompt_len, max_tokens, context_length).map_err(
|(prompt_len, context_length)| {
(
axum::http::StatusCode::BAD_REQUEST,
axum::Json(super::handlers::context_length_exceeded_body(
prompt_len,
context_length,
)),
)
.into_response()
},
)
}

#[cfg(all(test, feature = "inference"))]
#[path = "tests_st_serve_session_4269.rs"]
mod tests_st_serve_session_4269;
Expand Down
Loading
Loading