add the rust server ingress tests, guard, and submit modules (#32874)

This commit is contained in:
Rain Jiang
2026-07-30 12:46:18 -07:00
committed by GitHub
parent e2c65af229
commit 4facc0e18a
3 changed files with 900 additions and 0 deletions
+130
View File
@@ -0,0 +1,130 @@
//! Abort-on-disconnect guard for in-flight requests. Handlers arm a guard per
//! submitted rid; axum dropping the handler/SSE stream (client disconnected)
//! drops the guard, which aborts whatever wasn't disarmed (mirrors Python's
//! `is_disconnected` abort).
use std::collections::HashSet;
use crate::ids::Rid;
use crate::tokenizer_manager::{AbortSource, Senders};
/// Aborts still-in-flight rids on drop. Each rid is disarmed on natural finish;
/// whatever remains at drop is aborted.
pub(super) struct AbortGuard {
senders: Senders,
/// Rids still in flight. `Rid` carries its own partition key, so there is no
/// separate routing value to keep alongside it.
///
/// A set, not a `Vec`: `disarm` runs once per request that finishes, and over a
/// batch a linear scan makes the guard quadratic in the batch size — measured
/// 13.3 ms for a 4096-item batch, more than all of that batch's real transform
/// work combined. `Rid`'s identity is its id string, so set membership is the
/// same relation `retain` was testing. The cost is two hashes of a ~40-byte
/// string on the single-request path (~80 ns against a ~40 µs request), which
/// is why the trade is worth making rather than threading slot indices in from
/// the batch call sites.
rids: HashSet<Rid>,
}
impl AbortGuard {
pub(super) fn new(senders: Senders, rid: Rid) -> Self {
Self {
senders,
rids: HashSet::from([rid]),
}
}
/// Guard covering no rids yet — a batch arms each as it's submitted so a
/// mid-fan-out disconnect aborts every request already handed to the scheduler.
pub(super) fn new_empty(senders: Senders) -> Self {
Self {
senders,
rids: HashSet::new(),
}
}
/// Track a request for abort-on-drop.
pub(super) fn arm(&mut self, rid: Rid) {
self.rids.insert(rid);
}
/// Request finished naturally — don't abort it on drop.
pub(super) fn disarm(&mut self, rid: &Rid) {
self.rids.remove(rid);
}
}
impl Drop for AbortGuard {
fn drop(&mut self) {
// Report the abort and nothing more. There is no in-flight rid registry to
// release from: `Rid::from_client` makes each client rid internally unique,
// so a resubmit of the "same" rid is a different `Rid` and cannot be caught
// up in this abort. That removes the ordering hazard split ownership created.
//
// The lane is unbounded, so this send only fails at shutdown, when the loop
// is gone and nothing is generating anyway.
for rid in self.rids.drain() {
let _ = self.senders.abort.send(AbortSource::Guard(rid));
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn senders_with_abort(abort: flume::Sender<AbortSource>) -> Senders {
Senders {
tm: flume::unbounded().0,
abort,
tok: flume::unbounded().0,
detok: vec![],
}
}
/// A batch guard aborts exactly the rids still armed at drop — the ones whose
/// requests never reached a terminal — and leaves the finished ones alone.
#[test]
fn guard_aborts_only_the_rids_still_armed() {
let (abort_tx, abort_rx) = flume::unbounded();
let done: Rid = "done".into();
let mut guard = AbortGuard::new(senders_with_abort(abort_tx), done.clone());
guard.arm("aborted".into());
guard.disarm(&done); // finished naturally
drop(guard);
assert!(
matches!(abort_rx.try_recv().unwrap(), AbortSource::Guard(r) if r.as_str() == "aborted")
);
assert!(
abort_rx.try_recv().is_err(),
"a disarmed rid must not be aborted"
);
}
/// An armed guard aborts its rid on drop — exactly the cleanup a busy-skipped
/// `/health_generate` probe relies on. It never sees a terminal frame here, so
/// dropping the guard is the only path that deregisters its detok sink (via the
/// ingress `on_abort`). Regression for the detok-entry leak per health probe.
#[test]
fn armed_guard_aborts_on_drop() {
let (tm_tx, tm_rx) = flume::unbounded();
drop(AbortGuard::new(senders_with_abort(tm_tx), "r7".into()));
assert!(
matches!(tm_rx.try_recv(), Ok(AbortSource::Guard(rid)) if rid.as_str() == "r7"),
"armed guard must abort its rid on drop",
);
assert!(tm_rx.try_recv().is_err(), "exactly one abort");
}
/// A disarmed rid (finished naturally) is not aborted on drop.
#[test]
fn disarmed_guard_does_not_abort() {
let (tm_tx, tm_rx) = flume::unbounded();
let id = Rid::from("r9");
let mut guard = AbortGuard::new(senders_with_abort(tm_tx), "r9".into());
guard.disarm(&id);
drop(guard);
assert!(tm_rx.try_recv().is_err(), "disarmed rid must not abort");
}
}
+126
View File
@@ -0,0 +1,126 @@
//! Request submission into the ingress pipeline, shared by every endpoint
//! module: mint the client-visible rid (uuid hex, Python-parity), build the
//! `Request`, and hand it to the TM with an egress receiver for the response.
use std::convert::Infallible;
use axum::{
Json,
http::StatusCode,
response::{
IntoResponse, Response,
sse::{Event, Sse},
},
};
use tokio::sync::mpsc;
use super::AppState;
use super::frame::error_value;
use crate::fsm::RequestState;
use crate::ids::Rid;
use crate::message::{EgressItem, EgressSink, Request, RequestKind};
use crate::tokenizer_manager::TmEvent;
/// Submit one request; returns the rid, its hashed routing key, and the egress
/// receiver. Every request arrives with its final rid — a generate request from
/// `into_requests` (or the `HEALTH_CHECK_<uuid>` the health probe sets), a
/// control request from its constructor — so this only echoes it back.
pub(super) async fn submit(
state: &AppState,
kind: RequestKind,
// `stream`: the client is reading an SSE stream, so it expects 200 plus an
// error frame rather than a 4xx — same rule `pre_submit_error` applies
// everywhere else.
stream: bool,
) -> Result<(Rid, mpsc::Receiver<EgressItem>), Response> {
let rid = match &kind {
// Generate rids are already final: `GenerateBody::into_requests` normalized the
// client's, or minted one. Control requests have no client-facing rid.
RequestKind::Generate(g) => g.rid.clone(),
RequestKind::Control(c) => c.rid().into(),
};
// Two in-flight requests can name the same client rid, but they cannot share a
// `Rid`: `into_requests` built each through `Rid::from_client`, which appends a
// uniquifier. So nothing here needs to check for a collision — the detok table
// key is unique by construction, and `client_facing` restores what the client
// sent for `meta_info.id`.
// Async-aware send so a full TM inbox yields (backpressure) instead of parking
// a thread; Err only when the inbox is closed (shutdown).
let (tx, rx) = mpsc::channel::<EgressItem>(state.egress_buf);
let request = Request {
rid: rid.clone(),
state: RequestState::Received,
sink: EgressSink::Local(tx),
kind,
};
match state.senders.tm.send_async(TmEvent::Ingress(request)).await {
Ok(()) => Ok((rid, rx)),
// `SendError` has a single meaning — the channel is disconnected.
Err(_) => {
tracing::error!(%rid, "tm inbox closed; request rejected");
// Return 503 so the client can retry.
Err(pre_submit_error(
StatusCode::SERVICE_UNAVAILABLE,
"service unavailable",
stream,
))
}
}
}
/// Shape an error that occurs *before* (or instead of) a successful submit into a
/// client response. Two parity points with Python's `generate_request`. The body
/// is the same `{"error": {...}}` object every other path emits — not bare text,
/// which a client parsing JSON chokes on. And a streaming request gets 200 plus
/// one SSE error frame and `[DONE]`, not a 4xx: the client has already committed
/// to reading a stream, and Python answers it inside `stream_results()`.
pub(super) fn pre_submit_error(code: StatusCode, message: &str, stream: bool) -> Response {
let body = error_value(code.as_u16(), message);
if !stream {
return (code, Json(body)).into_response();
}
let frames = [body.to_string(), "[DONE]".to_string()];
Sse::new(futures::stream::iter(
frames.map(|data| Ok::<_, Infallible>(Event::default().data(data))),
))
.into_response()
}
#[cfg(test)]
mod tests {
use super::*;
/// Unary pre-submit errors are a 4xx/5xx with a JSON `{"error":...}` body;
/// streaming ones are 200 + an SSE error frame + `[DONE]`, because Python
/// answers from inside `stream_results()` once the stream is committed.
#[tokio::test]
async fn pre_submit_errors_match_python_shape() {
let unary = pre_submit_error(StatusCode::BAD_REQUEST, "bad input", false);
assert_eq!(unary.status(), StatusCode::BAD_REQUEST);
let body = axum::body::to_bytes(unary.into_body(), 64 * 1024)
.await
.unwrap();
let v: serde_json::Value = serde_json::from_slice(&body).expect("JSON body");
assert_eq!(v["error"]["message"], "bad input");
assert_eq!(v["error"]["code"], 400);
let streamed = pre_submit_error(StatusCode::BAD_REQUEST, "bad input", true);
assert_eq!(
streamed.status(),
StatusCode::OK,
"the stream itself is 200"
);
let body = axum::body::to_bytes(streamed.into_body(), 64 * 1024)
.await
.unwrap();
let text = String::from_utf8(body.to_vec()).unwrap();
assert!(
text.contains(r#""code":400"#),
"carries the status in-band: {text}"
);
assert!(
text.trim_end().ends_with("data: [DONE]"),
"terminated: {text}"
);
}
}
@@ -564,3 +564,647 @@ fn check_total_tokens(g: &mut GenerateRequest, limits: &Limits) -> Result<(), Er
g.sampling_params.max_new_tokens = Some(clamped);
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::fsm::RequestState;
use crate::message::{EgressSink, GenerateRequest, SamplingParams};
use crate::ring::{IngressConsumer, ingress_ring};
use tokio::sync::mpsc;
/// An `Ingress` plus its detok-shard receiver, ring consumer (keep alive —
/// dropping it closes the ring → false QueueFull), and tm inbox sender.
fn make_ingress() -> (
Ingress,
flume::Receiver<DetokMsg>,
IngressConsumer,
flume::Sender<TmEvent>,
) {
make_ingress_with(test_limits())
}
fn make_ingress_with_abort(
abort_rx: flume::Receiver<AbortSource>,
) -> (
Ingress,
flume::Receiver<DetokMsg>,
IngressConsumer,
flume::Sender<TmEvent>,
) {
make_ingress_inner(test_limits(), abort_rx)
}
fn make_ingress_with(
limits: Limits,
) -> (
Ingress,
flume::Receiver<DetokMsg>,
IngressConsumer,
flume::Sender<TmEvent>,
) {
let (abort_tx, abort_rx) = flume::unbounded::<AbortSource>();
std::mem::forget(abort_tx); // keep the lane open; tests end by dropping tm_tx
make_ingress_inner(limits, abort_rx)
}
fn make_ingress_inner(
limits: Limits,
abort_rx: flume::Receiver<AbortSource>,
) -> (
Ingress,
flume::Receiver<DetokMsg>,
IngressConsumer,
flume::Sender<TmEvent>,
) {
let (tok_tx, _tok_rx) = flume::unbounded();
let (detok_tx, detok_rx) = flume::unbounded();
let senders = Senders {
tm: flume::unbounded().0,
abort: flume::unbounded().0,
tok: tok_tx,
detok: vec![detok_tx],
};
let (ingress_producer, consumer) = ingress_ring(16);
let (tm_tx, tm_rx) = flume::unbounded();
// Keep the shutdown sender alive (leak) so its branch never fires — tests
// end `run` by dropping `tm_tx`, not by shutdown.
let (sd_tx, sd_rx) = flume::unbounded::<()>();
std::mem::forget(sd_tx);
let ingress = Ingress::new(tm_rx, abort_rx, senders, ingress_producer, limits, sd_rx);
(ingress, detok_rx, consumer, tm_tx)
}
/// Both abort sources do the same two things: drop the detok entry so no
/// further chunk can be delivered, and tell the scheduler to stop generating.
///
/// Neither releases anything, and nothing needs them to. Release ordering used
/// to be the delicate part here — `AbortGuard::drop` releasing a rid right
/// after enqueuing the abort ordered the SEND, not the EFFECT, so a retry of
/// the same rid could `Register` ahead of the stale abort and be torn down by
/// it. `Rid::from_client` removes the premise: a retry carries a different
/// `Rid`, so no abort in flight can name it.
#[test]
fn every_abort_source_deregisters_and_stops_the_scheduler() {
for source in [
AbortSource::Guard("x".into()),
AbortSource::Detok("x".into()),
] {
let (detok_tx, detok_rx) = flume::unbounded::<DetokMsg>();
let (ingress_producer, consumer) = ingress_ring(16);
let (sd_tx, sd_rx) = flume::unbounded::<()>();
std::mem::forget(sd_tx);
let ingress = Ingress::new(
flume::unbounded().1,
flume::unbounded().1,
Senders {
tm: flume::unbounded().0,
abort: flume::unbounded().0,
tok: flume::unbounded().0,
detok: vec![detok_tx],
},
ingress_producer,
test_limits(),
sd_rx,
);
ingress.on_abort(source.clone());
assert!(
matches!(detok_rx.try_recv(), Ok(DetokMsg::Deregister { rid }) if rid.as_str() == "x"),
"{source:?} must drop the detok entry",
);
assert_eq!(
consumer.drain(8).headers.len(),
1,
"{source:?} must push an AbortReq so the scheduler stops",
);
}
}
/// A context ceiling high enough that only a test which sets one on purpose
/// can reach it. `context_len` is mandatory now, so "no ceiling" has to be a
/// large number rather than `None`; kept well below `u64::MAX` so the
/// `as i64` in the auto-truncate clamp cannot go negative if a future test
/// does reach this path.
const NO_CONTEXT_CEILING: u64 = 1 << 40;
/// The default test limits: a real tokenizer, vocab 1000, no context ceiling.
/// Spelled out rather than `..Default::default()` — `Limits` deliberately has
/// no `Default`, because a zero `vocab_size`/`context_len` would reject every
/// request instead of behaving like "unset".
fn test_limits() -> Limits {
Limits {
skip_tokenizer_init: false,
vocab_size: 1000,
context_len: NO_CONTEXT_CEILING,
num_reserved_tokens: 0,
allow_auto_truncate: false,
enable_return_hidden_states: false,
}
}
fn generate_req(id: u64, sampling_params: SamplingParams) -> Request {
let (tx, _rx) = mpsc::channel(8);
Request {
rid: id.to_string().into(),
state: RequestState::Received,
sink: EgressSink::Local(tx),
kind: RequestKind::Generate(Box::new(GenerateRequest {
rid: id.to_string().into(),
input_ids: Some(vec![1, 2, 3]),
sampling_params,
..Default::default()
})),
}
}
/// `input + max_new_tokens` past the context window is an actionable 400, not a
/// silently truncated 200 (Python `TokenizerManager._validate_one_request`).
/// The message names both halves so the client can fix the right one.
#[test]
fn total_tokens_over_context_is_rejected() {
let limits = Limits {
context_len: 10,
..test_limits()
};
let mut g = GenerateRequest {
input_ids: Some(vec![1, 2, 3]),
sampling_params: SamplingParams {
max_new_tokens: Some(100),
..Default::default()
},
..Default::default()
};
let err = check_total_tokens(&mut g, &limits).unwrap_err();
let msg = err.to_string();
assert_eq!(err.http_status(), 400);
assert!(msg.contains("total of 103 tokens"), "{msg}");
assert!(msg.contains("3 tokens from the input"), "{msg}");
assert!(msg.contains("100 tokens for the completion"), "{msg}");
// Exactly filling the window is allowed (Python compares with `>`).
g.sampling_params.max_new_tokens = Some(7);
assert!(check_total_tokens(&mut g, &limits).is_ok());
assert_eq!(g.sampling_params.max_new_tokens, Some(7), "left alone");
}
/// The reserved slots (eagle draft tokens) count as input, so a request can be
/// rejected for them even when the prompt alone would fit.
#[test]
fn reserved_tokens_count_toward_the_limit() {
let limits = Limits {
context_len: 10,
num_reserved_tokens: 5,
..test_limits()
};
let mut g = GenerateRequest {
input_ids: Some(vec![1, 2, 3]),
sampling_params: SamplingParams {
max_new_tokens: Some(3), // 3 + 3 fits, but 3 + 5 + 3 does not
..Default::default()
},
..Default::default()
};
let msg = check_total_tokens(&mut g, &limits).unwrap_err().to_string();
assert!(msg.contains("8 tokens from the input"), "{msg}");
}
/// `--allow-auto-truncate` opts into clamping instead of rejecting; with no
/// context length, or no `max_new_tokens` cap, there is nothing to check.
#[test]
fn auto_truncate_clamps_and_unknowns_skip() {
let sp = |max_new_tokens| SamplingParams {
max_new_tokens,
..Default::default()
};
let mut g = GenerateRequest {
input_ids: Some(vec![1, 2, 3]),
sampling_params: sp(Some(100)),
..Default::default()
};
let truncating = Limits {
context_len: 10,
allow_auto_truncate: true,
..test_limits()
};
assert!(check_total_tokens(&mut g, &truncating).is_ok());
assert_eq!(g.sampling_params.max_new_tokens, Some(7), "clamped to fit");
// Unknown context length → no ceiling to enforce.
g.sampling_params = sp(Some(100));
assert!(check_total_tokens(&mut g, &test_limits()).is_ok());
assert_eq!(g.sampling_params.max_new_tokens, Some(100), "untouched");
// No cap requested → nothing to add to the input length, but the input
// itself is still checked (see `input_length_is_checked_unconditionally`).
g.sampling_params = sp(None);
let roomy = Limits {
context_len: 100,
..test_limits()
};
assert!(check_total_tokens(&mut g, &roomy).is_ok());
}
/// `max_new_tokens: null` means "no cap", NOT "skip the checks" — the input
/// alone must still fit. Gating the whole function on `max_new_tokens` let an
/// over-long prompt through to the scheduler with no ingress error at all.
/// Python compares with `>=`: a prompt that exactly fills the window leaves no
/// room to generate.
#[test]
fn input_length_is_checked_unconditionally() {
let limits = Limits {
context_len: 3,
..test_limits()
};
let req = |max_new_tokens| GenerateRequest {
input_ids: Some(vec![1, 2, 3]), // exactly fills a 3-token window
sampling_params: SamplingParams {
max_new_tokens,
..Default::default()
},
..Default::default()
};
for max_new_tokens in [None, Some(1)] {
let err = check_total_tokens(&mut req(max_new_tokens), &limits)
.expect_err("input == context_len must be rejected (Python uses >=)");
assert_eq!(err.http_status(), 400);
assert!(err.to_string().contains("longer than the model's context"));
}
// One token shorter fits, with or without a cap.
let mut g = GenerateRequest {
input_ids: Some(vec![1, 2]),
..Default::default()
};
g.sampling_params.max_new_tokens = None;
assert!(check_total_tokens(&mut g, &limits).is_ok());
// Under auto-truncate the input is cut to fit instead of rejected.
let truncating = Limits {
allow_auto_truncate: true,
..limits.clone()
};
let mut g = req(None);
assert!(check_total_tokens(&mut g, &truncating).is_ok());
assert_eq!(
g.input_ids.as_deref(),
Some(&[1, 2, 3][..]),
"fits at the cap"
);
}
/// The clamp runs AFTER `verify` (which happens in `Normalizing`), so lowering
/// `max_new_tokens` can leave `min_new_tokens > max_new_tokens`. Nothing
/// downstream re-checks — `is_normalized: true` makes the scheduler's own
/// verify early-return — so the clamp has to re-assert it here.
#[test]
fn auto_truncate_cannot_invert_min_and_max_new_tokens() {
let limits = Limits {
context_len: 10,
allow_auto_truncate: true,
..test_limits()
};
let mut g = GenerateRequest {
input_ids: Some(vec![1, 2, 3]), // clamps max_new_tokens to 7
sampling_params: SamplingParams {
max_new_tokens: Some(100),
min_new_tokens: 50, // …which is below min_new_tokens
..Default::default()
},
..Default::default()
};
let err = check_total_tokens(&mut g, &limits)
.expect_err("a clamp that inverts min/max must 400, not ride the wire");
assert_eq!(err.http_status(), 400);
assert!(err.to_string().contains("min_new_tokens"), "{err}");
// A clamp that keeps the invariant still clamps.
g.sampling_params.min_new_tokens = 2;
g.sampling_params.max_new_tokens = Some(100);
assert!(check_total_tokens(&mut g, &limits).is_ok());
assert_eq!(g.sampling_params.max_new_tokens, Some(7));
}
/// `return_hidden_states` on a server not launched for it is a 400: the
/// scheduler never computes them, so the request would otherwise 200 with
/// `meta_info.hidden_states` silently missing.
#[test]
fn hidden_states_gated_on_server_support() {
let req = |want| {
let mut r = generate_req(31, SamplingParams::default());
if let RequestKind::Generate(g) = &mut r.kind {
g.return_hidden_states = want;
}
r
};
let disabled = test_limits();
let err = validate(&mut req(true), &disabled).unwrap_err();
assert_eq!(err.http_status(), 400);
assert!(
err.to_string().contains("--enable-return-hidden-states"),
"message must name the flag: {err}"
);
// Not asking for them (the client sent `false`, or sent nothing and
// `into_requests` resolved the default), or asking on a server that
// supports them, is fine.
assert!(validate(&mut req(false), &disabled).is_ok());
let enabled = Limits {
enable_return_hidden_states: true,
..test_limits()
};
assert!(validate(&mut req(true), &enabled).is_ok());
}
/// End-to-end through `drive`: an over-context request is rejected on the way
/// to the ring, after registration — so it must be deregistered, not leaked.
#[test]
fn over_context_request_deregisters_and_never_reaches_the_ring() {
let (ingress, detok_rx, consumer, _tm_tx) = make_ingress_with(Limits {
context_len: 4,
..test_limits()
});
ingress.drive(generate_req(
33,
SamplingParams {
max_new_tokens: Some(64),
..Default::default()
},
));
assert!(
matches!(detok_rx.try_recv(), Ok(DetokMsg::Register { rid, .. }) if rid.as_str() == "33"),
"registered before the check",
);
assert!(
matches!(detok_rx.try_recv(), Ok(DetokMsg::Deregister { rid }) if rid.as_str() == "33"),
"must deregister on reject",
);
assert!(
consumer.drain(16).headers.is_empty(),
"must not reach the scheduler"
);
}
/// A dropped ring push is survivable, and this pins WHY. The ring is bounded,
/// so under load the scheduler never learns to stop and keeps generating; its
/// chunks then arrive for a rid the detok table no longer holds and are
/// dropped. That wastes GPU work but cannot MISDELIVER, because
/// `Rid::from_client` guarantees no later request ever answers to that rid.
/// The detok entry is dropped either way — that is the half that must not
/// depend on the ring.
///
/// Ring capacity 1: the first abort pushes, the second finds it full.
#[test]
fn abort_deregisters_even_when_the_ring_push_is_dropped() {
let (tok_tx, _tok_rx) = flume::unbounded();
let (detok_tx, detok_rx) = flume::unbounded();
let (abort_tx, abort_rx) = flume::unbounded::<AbortSource>();
let senders = Senders {
tm: flume::unbounded().0,
abort: abort_tx,
tok: tok_tx,
detok: vec![detok_tx],
};
let (producer, _consumer) = ingress_ring(1);
let (_tm_tx, tm_rx) = flume::unbounded();
let (sd_tx, sd_rx) = flume::unbounded::<()>();
std::mem::forget(sd_tx);
let ingress = Ingress::new(tm_rx, abort_rx, senders, producer, test_limits(), sd_rx);
ingress.on_abort(AbortSource::Guard("pushed".into()));
ingress.on_abort(AbortSource::Guard("dropped".into()));
// Both deregisters land regardless of whether the ring accepted the push.
for expected in ["pushed", "dropped"] {
assert!(
matches!(detok_rx.try_recv(), Ok(DetokMsg::Deregister { rid }) if rid.as_str() == expected),
"{expected}: the detok entry must be dropped even when the ring is full",
);
}
}
/// The rid keys the detok table and rides on every chunk of every decode step,
/// so an unbounded client-supplied one is a recurring cost, not a one-off.
#[test]
fn oversized_rid_is_rejected() {
let mut req = generate_req(51, SamplingParams::default());
req.rid = "x".repeat(MAX_RID_LEN + 1).into();
let err = validate(&mut req, &test_limits()).expect_err("must be rejected");
assert_eq!(err.http_status(), 400);
assert!(err.to_string().contains("over the"), "{err}");
// A uuid-sized rid — what Python mints — is nowhere near the cap.
let mut req = generate_req(52, SamplingParams::default());
req.rid = "0123456789abcdef0123456789abcdef".into();
assert!(validate(&mut req, &test_limits()).is_ok());
}
/// A request rejected BEFORE `register_detok` must not send `Deregister`: the
/// handler is a bare `table.remove(&rid)`, so it would evict whatever entry
/// holds that key — a concurrent request's sink — leaving that client hung with
/// no terminal frame. Python validates before it inserts, so it cannot hit this.
#[test]
fn pre_registration_failure_does_not_deregister() {
// Rejected inside `validate` (out-of-vocab id), which runs before registration.
let (ingress, detok_rx, _consumer, _tm_tx) = make_ingress();
let mut req = generate_req(41, SamplingParams::default());
if let RequestKind::Generate(g) = &mut req.kind {
g.input_ids = Some(vec![2_000_000_000]);
}
ingress.drive(req);
assert!(
detok_rx.try_recv().is_err(),
"a pre-registration reject must send NOTHING to the shard — a Deregister \
here removes a live request's sink"
);
// A post-registration reject still deregisters (the leak fix stays fixed).
let (ingress, detok_rx, _consumer, _tm_tx) = make_ingress();
ingress.drive(generate_req(
42,
SamplingParams {
top_p: 2.0, // rejected by `normalize`, after registration
..Default::default()
},
));
assert!(matches!(detok_rx.try_recv(), Ok(DetokMsg::Register { .. })));
assert!(matches!(
detok_rx.try_recv(),
Ok(DetokMsg::Deregister { .. })
));
}
/// A request rejected at normalization (post-register) must not leak: the shard
/// sees `Register` then `Deregister`. Regression for RSS growth on bad input.
#[test]
fn rejected_request_deregisters_from_shard() {
let (ingress, detok_rx, _consumer, _tm_tx) = make_ingress();
// top_p = 2.0 is outside (0, 1], so `SamplingParams::normalize` rejects it.
let bad = SamplingParams {
top_p: 2.0,
..Default::default()
};
ingress.drive(generate_req(7, bad));
assert!(
matches!(detok_rx.try_recv(), Ok(DetokMsg::Register { rid, .. }) if rid.as_str() == "7"),
"expected Register for rid 7",
);
assert!(
matches!(detok_rx.try_recv(), Ok(DetokMsg::Deregister { rid }) if rid.as_str() == "7"),
"expected Deregister for rid 7 (leak fix)",
);
assert!(
detok_rx.try_recv().is_err(),
"no further shard messages — registration fully cleaned up",
);
}
/// Regression: an out-of-vocabulary client token id must be rejected at
/// ingress with a 400 — passed through, it reaches the embedding lookup
/// and kills the scheduler process (`make_ingress` bounds vocab at 1000).
#[test]
fn out_of_vocab_input_ids_rejected() {
let (ingress, detok_rx, _consumer, _tm_tx) = make_ingress();
let mut req = generate_req(21, SamplingParams::default());
if let RequestKind::Generate(g) = &mut req.kind {
g.input_ids = Some(vec![1, 2_000_000_000]);
}
ingress.drive(req);
// Rejected before registration: the only shard message is nothing at
// all, or a Deregister if registration happened first — never a push.
match detok_rx.try_recv() {
Err(_) => {}
Ok(DetokMsg::Deregister { .. }) => {}
Ok(_) => panic!("out-of-vocab request must not be admitted"),
}
}
/// Same guard for negative ids and for `token_ids_logprob` entries.
#[test]
fn negative_and_logprob_token_ids_rejected() {
let (ingress, detok_rx, _consumer, _tm_tx) = make_ingress();
let mut req = generate_req(22, SamplingParams::default());
if let RequestKind::Generate(g) = &mut req.kind {
g.input_ids = Some(vec![-1]);
}
ingress.drive(req);
match detok_rx.try_recv() {
Err(_) | Ok(DetokMsg::Deregister { .. }) => {}
Ok(_) => panic!("negative token id must not be admitted"),
}
let (ingress, detok_rx, _consumer, _tm_tx) = make_ingress();
let mut req = generate_req(23, SamplingParams::default());
if let RequestKind::Generate(g) = &mut req.kind {
g.token_ids_logprob = Some(vec![999_999]);
}
ingress.drive(req);
match detok_rx.try_recv() {
Err(_) | Ok(DetokMsg::Deregister { .. }) => {}
Ok(_) => panic!("out-of-vocab token_ids_logprob must not be admitted"),
}
}
/// A valid request is registered and handed onward — never deregistered.
#[test]
fn admitted_request_keeps_registration() {
let (ingress, detok_rx, _consumer, _tm_tx) = make_ingress();
// Empty map → all sampling defaults, passes normalization.
ingress.drive(generate_req(9, SamplingParams::default()));
assert!(
matches!(detok_rx.try_recv(), Ok(DetokMsg::Register { rid, .. }) if rid.as_str() == "9"),
"expected Register for rid 9",
);
assert!(
detok_rx.try_recv().is_err(),
"admitted request must not be deregistered",
);
}
/// A pool return in `Failed` state (failed encode) is rejected via the same
/// path and deregistered, not leaked.
#[test]
fn tokenize_failure_deregisters_via_ingress() {
let (ingress, detok_rx, _consumer, tm_tx) = make_ingress();
// The pool marks a failed encode as `Failed(err)` before returning it.
let mut req = generate_req(11, SamplingParams::default());
let _ = req
.state
.apply(Event::Error(Error::Tokenize("boom".into())));
tm_tx.send(TmEvent::Tokenized(req)).unwrap();
// Close the inbox so the run loop returns after draining the one event.
drop(tm_tx);
ingress.run();
assert!(
matches!(detok_rx.try_recv(), Ok(DetokMsg::Deregister { rid }) if rid.as_str() == "11"),
"tokenize failure must deregister rid 11",
);
assert!(detok_rx.try_recv().is_err(), "no further shard messages");
}
/// An abort deregisters (by the id hashed from the rid string), so a request
/// aborted before any terminal chunk can't leak.
#[test]
fn abort_deregisters_from_shard() {
// Aborts arrive on their own unbounded lane now, not the request inbox.
let (abort_tx, abort_rx) = flume::unbounded::<AbortSource>();
let (ingress, detok_rx, _consumer, tm_tx) = make_ingress_with_abort(abort_rx);
abort_tx.send(AbortSource::Guard("rid-13".into())).unwrap();
drop(abort_tx);
drop(tm_tx);
ingress.run();
assert!(
matches!(detok_rx.try_recv(), Ok(DetokMsg::Deregister { rid }) if rid.as_str() == "rid-13"),
"abort must deregister by rid",
);
assert!(detok_rx.try_recv().is_err(), "no further shard messages");
}
/// A successful pool return (Queued, ids filled) is pushed to the ring, not
/// rejected; its registration is untouched.
#[test]
fn tokenized_return_pushes_without_deregister() {
let (ingress, detok_rx, _consumer, tm_tx) = make_ingress();
let mut req = generate_req(15, SamplingParams::default());
// Simulate a successful pool return: ids filled, PreSendValidating.
if let RequestKind::Generate(g) = &mut req.kind {
g.input_ids = Some(vec![1, 2, 3]);
}
req.state = RequestState::PreSendValidating;
tm_tx.send(TmEvent::Tokenized(req)).unwrap();
drop(tm_tx);
ingress.run();
// Pushed to the ring; the shard sees nothing.
assert!(
detok_rx.try_recv().is_err(),
"a queued pool-return must be pushed, not touch the shard",
);
}
/// If the pool is gone, a request needing tokenization is rejected +
/// deregistered, not silently dropped.
#[test]
fn tokenize_pool_gone_deregisters() {
// `make_ingress` drops the tok receiver, so `tok.send` fails.
let (ingress, detok_rx, _consumer, _tm_tx) = make_ingress();
// No ids → NeedsTokenize → Tokenizing branch.
let mut req = generate_req(21, SamplingParams::default());
if let RequestKind::Generate(g) = &mut req.kind {
g.input_ids = None;
}
ingress.drive(req);
assert!(
matches!(detok_rx.try_recv(), Ok(DetokMsg::Register { rid, .. }) if rid.as_str() == "21"),
"expected Register for rid 21",
);
assert!(
matches!(detok_rx.try_recv(), Ok(DetokMsg::Deregister { rid }) if rid.as_str() == "21"),
"pool-gone hand-off must deregister rid 21",
);
assert!(detok_rx.try_recv().is_err(), "no further shard messages");
}
}