use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use axum::body::Body;
use axum::http::{Request, StatusCode};
use tower::util::ServiceExt;
use crate::api::{create_router, request_cancel_token, AppState};
use crate::generate::CancelToken;
use crate::layers::{Model, ModelConfig};
fn tiny_dense_model() -> Model {
Model::new(ModelConfig {
vocab_size: 16,
hidden_dim: 8,
num_heads: 2,
num_layers: 1,
intermediate_dim: 16,
eps: 1e-5,
})
.expect("build tiny dense model")
}
#[cfg(feature = "gpu")]
fn quantized_state() -> AppState {
use crate::api::test_helpers::create_test_quantized_model;
use crate::gguf::{ArchConstraints, GGUFConfig};
let config = GGUFConfig {
architecture: "llama".to_string(),
constraints: ArchConstraints::from_architecture("llama"),
hidden_dim: 64,
intermediate_dim: 128,
num_layers: 2,
num_heads: 4,
num_kv_heads: 4,
vocab_size: 256,
context_length: 512,
rope_theta: 10000.0,
eps: 1e-5,
rope_type: 0,
explicit_head_dim: None,
query_pre_attn_scalar: None,
bos_token_id: None,
eos_token_id: None,
};
AppState::with_quantized_model(create_test_quantized_model(&config))
.expect("build quantized AppState")
}
#[test]
fn dense_generation_stops_at_the_cancel_point_not_max_tokens() {
use crate::generate::GenerationConfig;
let model = tiny_dense_model();
let prompt = [1_usize, 2];
const MAX_TOKENS: usize = 64;
const BUDGET: usize = 8;
let uncancelled = model
.generate(
&prompt,
&GenerationConfig::greedy().with_max_tokens(MAX_TOKENS),
)
.expect("uncancelled generation");
assert_eq!(
uncancelled.len(),
prompt.len() + MAX_TOKENS,
"control: with no cancellation the loop must run the full {MAX_TOKENS}-token budget, \
otherwise the cancelled case below is not measuring cancellation"
);
let token = CancelToken::with_budget(BUDGET);
let cancelled = model
.generate(
&prompt,
&GenerationConfig::greedy()
.with_max_tokens(MAX_TOKENS)
.with_cancel(token.clone()),
)
.expect("cancelled generation still returns what it produced");
let produced = cancelled.len() - prompt.len();
assert_eq!(
produced, BUDGET,
"generation must stop at the cancel point ({BUDGET} tokens), not run to \
max_tokens ({MAX_TOKENS}); it produced {produced}"
);
assert_eq!(
token.polls(),
BUDGET + 1,
"the loop must poll exactly once per decode step (BUDGET polls that returned \
false, plus the one that returned true and broke the loop)"
);
assert_eq!(
cancelled,
uncancelled[..cancelled.len()].to_vec(),
"a cancelled run must be a strict prefix of the uncancelled run"
);
}
#[test]
fn dense_generation_cancelled_before_start_produces_no_tokens() {
use crate::generate::GenerationConfig;
let model = tiny_dense_model();
let prompt = [1_usize, 2];
let token = CancelToken::new();
token.cancel();
let out = model
.generate(
&prompt,
&GenerationConfig::greedy()
.with_max_tokens(64)
.with_cancel(token),
)
.expect("an already-cancelled request is not an error");
assert_eq!(
out.len(),
prompt.len(),
"an already-cancelled request must do no decode work at all; it returned \
{} tokens beyond the prompt",
out.len() - prompt.len()
);
}
#[test]
#[cfg(feature = "gpu")]
fn quantized_generation_stops_at_the_cancel_point_not_max_tokens() {
use crate::api::test_helpers::create_test_quantized_model;
use crate::gguf::{ArchConstraints, GGUFConfig, QuantizedGenerateConfig};
let gguf_config = GGUFConfig {
architecture: "llama".to_string(),
constraints: ArchConstraints::from_architecture("llama"),
hidden_dim: 64,
intermediate_dim: 128,
num_layers: 2,
num_heads: 4,
num_kv_heads: 4,
vocab_size: 256,
context_length: 512,
rope_theta: 10000.0,
eps: 1e-5,
rope_type: 0,
explicit_head_dim: None,
query_pre_attn_scalar: None,
bos_token_id: None,
eos_token_id: None,
};
let model = create_test_quantized_model(&gguf_config);
let prompt = [3_u32, 4, 5];
const MAX_TOKENS: usize = 48;
const BUDGET: usize = 6;
let base = QuantizedGenerateConfig::deterministic(MAX_TOKENS);
let uncancelled = model
.generate_with_cache(&prompt, &base)
.expect("uncancelled quantized generation");
assert_eq!(
uncancelled.len(),
prompt.len() + MAX_TOKENS,
"control: the quantized loop must run its full {MAX_TOKENS}-token budget when \
nothing cancels it"
);
let token = CancelToken::with_budget(BUDGET);
let cancelled = model
.generate_with_cache(&prompt, &base.clone().with_cancel(token.clone()))
.expect("cancelled quantized generation");
let produced = cancelled.len() - prompt.len();
assert_eq!(
produced, BUDGET,
"the quantized loop must stop at the cancel point ({BUDGET}), not run to \
max_tokens ({MAX_TOKENS}); it produced {produced}"
);
assert_eq!(
token.polls(),
BUDGET + 1,
"the quantized loop must poll exactly once per decode step"
);
}
#[tokio::test]
async fn dropping_the_request_future_stops_a_running_generation() {
use axum::routing::post;
use axum::Router;
const BUDGET_CEILING: usize = 500_000;
let steps = Arc::new(AtomicUsize::new(0));
let (started_tx, started_rx) = std::sync::mpsc::channel::<()>();
let (finished_tx, finished_rx) = std::sync::mpsc::channel::<usize>();
let steps_for_handler = Arc::clone(&steps);
let started = Arc::new(std::sync::Mutex::new(Some(started_tx)));
let finished = Arc::new(std::sync::Mutex::new(Some(finished_tx)));
let handler = move |request: Request<Body>| {
let steps = Arc::clone(&steps_for_handler);
let started = Arc::clone(&started);
let finished = Arc::clone(&finished);
async move {
let cancel = request_cancel_token(&request);
tokio::task::spawn_blocking(move || {
for i in 0..BUDGET_CEILING {
if cancel.is_cancelled() {
break;
}
steps.store(i + 1, Ordering::SeqCst);
if i == 64 {
if let Some(tx) = started.lock().ok().and_then(|mut g| g.take()) {
let _ = tx.send(());
}
}
std::hint::spin_loop();
std::thread::yield_now();
}
if let Some(tx) = finished.lock().ok().and_then(|mut g| g.take()) {
let _ = tx.send(steps.load(Ordering::SeqCst));
}
});
std::future::pending::<StatusCode>().await
}
};
let app = Router::new()
.route("/decode", post(handler))
.layer(axum::middleware::from_fn(crate::api::cancel_on_disconnect));
let request = Request::builder()
.method("POST")
.uri("/decode")
.body(Body::empty())
.expect("build request");
let mut response_future = Box::pin(app.oneshot(request));
let started_rx = tokio::task::spawn_blocking(move || {
started_rx
.recv_timeout(std::time::Duration::from_secs(10))
.map(|()| ())
});
tokio::select! {
_ = &mut response_future => panic!("the handler must not complete on its own"),
r = started_rx => r.expect("join").expect("the decode loop must start"),
}
drop(response_future);
let observed = tokio::task::spawn_blocking(move || {
finished_rx.recv_timeout(std::time::Duration::from_secs(10))
})
.await
.expect("join")
.expect("the decode loop must terminate after the request future is dropped");
assert!(
observed < BUDGET_CEILING,
"dropping the request future must stop the decode loop; it ran {observed} of \
{BUDGET_CEILING} steps, i.e. it never observed cancellation"
);
}
#[tokio::test]
#[cfg(feature = "gpu")]
async fn the_real_router_hands_every_request_a_live_cancel_token() {
let observed = Arc::new(std::sync::Mutex::new(None::<CancelToken>));
let sink = Arc::clone(&observed);
let app = axum::Router::new()
.route(
"/probe",
axum::routing::get(move |request: Request<Body>| {
let sink = Arc::clone(&sink);
async move {
let token = request_cancel_token(&request);
if let Ok(mut g) = sink.lock() {
*g = Some(token);
}
StatusCode::OK
}
}),
)
.layer(axum::middleware::from_fn(crate::api::cancel_on_disconnect));
let response = app
.oneshot(
Request::builder()
.uri("/probe")
.body(Body::empty())
.expect("build request"),
)
.await
.expect("probe request");
assert_eq!(response.status(), StatusCode::OK);
let token = observed
.lock()
.ok()
.and_then(|g| g.clone())
.expect("handler must have seen a token");
assert!(
!token.peek_cancelled(),
"the layer cancelled a request that ran to completion; that stops the \
background decode loop behind a streaming response before it emits its \
first token"
);
token.cancel();
assert!(
token.peek_cancelled(),
"the layer must hand every request a LIVE token; a `never` token silently \
ignores cancel() and leaves the decode loop with nothing to poll"
);
}
#[tokio::test]
#[cfg(feature = "gpu")]
async fn a_completed_generate_request_is_unchanged_by_the_cancellation_layer() {
use axum::extract::State;
use axum::{Extension, Json};
const BODY: &str = r#"{"prompt":"token5","max_tokens":4}"#;
let response = create_router(quantized_state())
.oneshot(
Request::builder()
.method("POST")
.uri("/generate")
.header("content-type", "application/json")
.body(Body::from(BODY))
.expect("build request"),
)
.await
.expect("generate request");
assert_eq!(
response.status(),
StatusCode::OK,
"a normal request must still succeed"
);
let via_router = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.expect("read body");
let via_router: serde_json::Value = serde_json::from_slice(&via_router).expect("json body");
let request = serde_json::from_str(BODY).expect("parse request");
let baseline = crate::api::generate_handler(
State(quantized_state()),
Extension(CancelToken::never()),
Json(request),
)
.await;
let baseline = match baseline {
Ok(Json(resp)) => serde_json::to_value(resp).expect("serialize baseline"),
Err((status, Json(err))) => {
panic!(
"baseline generation must succeed, got {status}: {}",
err.error
)
},
};
assert_eq!(
via_router, baseline,
"the cancellation layer must not change a completed response; routing it \
through the layer produced {via_router} but the direct call produced {baseline}"
);
}