use crate::forward::metal_qwen35::{ChatMessage, MetalQwen35State, format_chat_template};
use crate::kv_cache::CrossTurnSlotId;
use crate::model::qwen35_config::{GenerateConfig, GenerateOutput};
use crate::serve::ApiError;
use crate::tokenizer::Tokenizer as _;
use crate::tokenizer::bpe::BpeTokenizer;
use std::io::Write as _;
use std::sync::{Arc, Mutex, MutexGuard};
use std::time::{Duration, Instant};
use tokio::sync::{OwnedSemaphorePermit, Semaphore, mpsc, watch};
pub const DEFAULT_MAX_PENDING_JOBS: usize = 32;
const WORKER_SHUTDOWN_TIMEOUT: Duration = Duration::from_secs(2);
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[must_use]
enum WorkerShutdown {
Joined,
AlreadyStopped,
TimedOut,
Panicked,
ReaperUnavailable,
}
#[derive(Debug, Clone, Copy)]
pub enum ContextWindowPolicy {
PromptAndMaxTokens,
PromptAndDecodeWithDelimiter,
}
#[derive(Debug, Clone)]
pub struct WorkerMetadata {
pub format: String,
pub model_max_context: usize,
pub context_window_policy: ContextWindowPolicy,
}
#[derive(Debug)]
pub enum WorkerEvent {
Delta(String),
Complete(GenerateOutput),
Rejected(ApiError),
Failed(String),
ConstraintBlocked(String),
Cancelled,
}
enum WorkerFailure {
Rejected(ApiError),
Failed(String),
ConstraintBlocked(String),
}
impl From<crate::error::InferenceError> for WorkerFailure {
fn from(err: crate::error::InferenceError) -> Self {
match err {
crate::error::InferenceError::GrammarConstraintBlocked(message) => {
WorkerFailure::ConstraintBlocked(message)
}
other => WorkerFailure::Failed(other.to_string()),
}
}
}
#[derive(Debug)]
pub enum StartupError {
Load(String),
ThreadExited,
InvalidMaxPending {
max_pending: usize,
},
}
impl std::fmt::Display for StartupError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
StartupError::Load(message) => write!(f, "{message}"),
StartupError::ThreadExited => {
write!(f, "worker thread exited before loading finished")
}
StartupError::InvalidMaxPending { max_pending } => write!(
f,
"--max-pending must be between 1 and {} (got {max_pending})",
Semaphore::MAX_PERMITS
),
}
}
}
impl std::error::Error for StartupError {}
pub struct WorkerJob {
messages: Vec<ChatMessage>,
cfg: GenerateConfig,
tx: mpsc::UnboundedSender<WorkerEvent>,
cancel: watch::Receiver<bool>,
_admission_permit: OwnedSemaphorePermit,
}
#[derive(Debug, Clone)]
pub struct MetalWorkerOwner {
_inner: Arc<MetalWorkerOwnerInner>,
}
#[derive(Debug)]
struct MetalWorkerOwnerInner {
join_handle: Mutex<Option<std::thread::JoinHandle<()>>>,
drop_timeout: Duration,
}
fn lock_unpoisoned<T>(mutex: &Mutex<T>) -> MutexGuard<'_, T> {
match mutex.lock() {
Ok(guard) => guard,
Err(poisoned) => poisoned.into_inner(),
}
}
impl MetalWorkerOwnerInner {
fn wait_for_exit(&self, timeout: Duration) -> WorkerShutdown {
let handle = {
let mut join_handle = lock_unpoisoned(&self.join_handle);
let Some(handle) = join_handle.take() else {
return WorkerShutdown::AlreadyStopped;
};
handle
};
let started = Instant::now();
let (joined_tx, joined_rx) = std::sync::mpsc::sync_channel(1);
let reaper = std::thread::Builder::new()
.name("lattice-metal-worker-reaper".to_string())
.spawn(move || {
let _ = joined_tx.send(handle.join());
});
let Ok(reaper) = reaper else {
return WorkerShutdown::ReaperUnavailable;
};
drop(reaper);
let remaining = timeout.saturating_sub(started.elapsed());
match joined_rx.recv_timeout(remaining) {
Ok(Ok(())) => WorkerShutdown::Joined,
Ok(Err(_)) => WorkerShutdown::Panicked,
Err(std::sync::mpsc::RecvTimeoutError::Timeout) => WorkerShutdown::TimedOut,
Err(std::sync::mpsc::RecvTimeoutError::Disconnected) => {
WorkerShutdown::ReaperUnavailable
}
}
}
}
impl Drop for MetalWorkerOwnerInner {
fn drop(&mut self) {
match self.wait_for_exit(self.drop_timeout) {
WorkerShutdown::Joined | WorkerShutdown::AlreadyStopped => {}
WorkerShutdown::TimedOut => {
let _ = writeln!(
std::io::stderr().lock(),
"[metal-worker] shutdown timed out after {} ms; detaching worker join reaper",
self.drop_timeout.as_millis()
);
}
WorkerShutdown::Panicked => {
let _ = writeln!(
std::io::stderr().lock(),
"[metal-worker] worker thread panicked during shutdown"
);
}
WorkerShutdown::ReaperUnavailable => {
let _ = writeln!(
std::io::stderr().lock(),
"[metal-worker] worker join reaper unavailable; detaching worker thread"
);
}
}
}
}
impl MetalWorkerOwner {
fn from_handle(join_handle: std::thread::JoinHandle<()>) -> Self {
Self::from_handle_with_timeout(join_handle, WORKER_SHUTDOWN_TIMEOUT)
}
fn from_handle_with_timeout(
join_handle: std::thread::JoinHandle<()>,
drop_timeout: Duration,
) -> Self {
Self {
_inner: Arc::new(MetalWorkerOwnerInner {
join_handle: Mutex::new(Some(join_handle)),
drop_timeout,
}),
}
}
#[cfg(any(test, feature = "test-utils"))]
fn unattached_for_test() -> Self {
Self {
_inner: Arc::new(MetalWorkerOwnerInner {
join_handle: Mutex::new(None),
drop_timeout: WORKER_SHUTDOWN_TIMEOUT,
}),
}
}
}
#[derive(Debug, Clone)]
pub struct MetalWorkerClient {
jobs: Option<mpsc::UnboundedSender<WorkerJob>>,
admission: Arc<Semaphore>,
_owner: MetalWorkerOwner,
}
impl MetalWorkerClient {
fn with_owner(
jobs: mpsc::UnboundedSender<WorkerJob>,
admission: Arc<Semaphore>,
owner: MetalWorkerOwner,
) -> Self {
Self {
jobs: Some(jobs),
admission,
_owner: owner,
}
}
#[cfg(any(test, feature = "test-utils"))]
fn unattached_for_test(
jobs: mpsc::UnboundedSender<WorkerJob>,
admission: Arc<Semaphore>,
) -> Self {
Self::with_owner(jobs, admission, MetalWorkerOwner::unattached_for_test())
}
pub fn submit(
&self,
messages: Vec<ChatMessage>,
gen_cfg: GenerateConfig,
cancel: watch::Receiver<bool>,
) -> Result<mpsc::UnboundedReceiver<WorkerEvent>, ApiError> {
let permit = self.admission.clone().try_acquire_owned().map_err(|_| {
ApiError::ServiceUnavailable {
message: "too many outstanding requests; the inference worker's pending-job \
queue is full, retry shortly"
.to_string(),
}
})?;
let (tx, rx) = mpsc::unbounded_channel();
let job = WorkerJob {
messages,
cfg: gen_cfg,
tx,
cancel,
_admission_permit: permit,
};
if let Some(jobs) = self.jobs.as_ref() {
let _ = jobs.send(job);
}
Ok(rx)
}
pub fn available_permits(&self) -> usize {
self.admission.available_permits()
}
}
impl Drop for MetalWorkerClient {
fn drop(&mut self) {
drop(self.jobs.take());
}
}
fn check_prompt_fits_window(
policy: ContextWindowPolicy,
model_max_context: usize,
prompt_len: usize,
cfg: &GenerateConfig,
) -> Result<(), ApiError> {
if matches!(policy, ContextWindowPolicy::PromptAndMaxTokens) && prompt_len == 0 {
return Err(ApiError::BadRequest {
message: format!(
"prompt (0 tokens) plus max_tokens ({max_tokens}) exceeds model \
context window ({model_max_context})",
max_tokens = cfg.max_new_tokens,
),
code: "context_length_exceeded",
});
}
let (decode_cap, delimiter_tokens) = match policy {
ContextWindowPolicy::PromptAndMaxTokens => (cfg.max_new_tokens, 0),
ContextWindowPolicy::PromptAndDecodeWithDelimiter => (
cfg.max_new_tokens
.saturating_add(cfg.reasoning_budget.unwrap_or(0)),
1,
),
};
let required = prompt_len
.saturating_add(decode_cap)
.saturating_add(delimiter_tokens);
if required > model_max_context {
let available = model_max_context.saturating_sub(prompt_len);
let delimiter_clause = match delimiter_tokens {
0 => String::new(),
n => format!(" plus {n}"),
};
return Err(ApiError::BadRequest {
message: format!(
"prompt has {prompt_len} tokens, leaving {available} of the \
{model_max_context}-token context window for generation, but this \
request needs {decode_cap} generated tokens{delimiter_clause} (total {required}); \
reduce max_tokens/reasoning_budget or shorten the prompt"
),
code: "context_length_exceeded",
});
}
Ok(())
}
fn run_worker_loop(
mut job_rx: mpsc::UnboundedReceiver<WorkerJob>,
mut generate: impl FnMut(
&[ChatMessage],
&GenerateConfig,
&mut dyn FnMut(&str, u32) -> bool,
&mut dyn FnMut() -> bool,
) -> Result<GenerateOutput, WorkerFailure>,
) {
while let Some(job) = job_rx.blocking_recv() {
if *job.cancel.borrow() || job.tx.is_closed() {
let _ = job.tx.send(WorkerEvent::Cancelled);
continue;
}
let cb_tx = job.tx.clone();
let cancel_for_token = job.cancel.clone();
let mut on_token = move |delta: &str, _token_id: u32| {
if *cancel_for_token.borrow() {
return false;
}
cb_tx.send(WorkerEvent::Delta(delta.to_string())).is_ok()
};
let cancel_for_predicate = job.cancel.clone();
let tx_for_predicate = job.tx.clone();
let mut should_cancel =
move || *cancel_for_predicate.borrow() || tx_for_predicate.is_closed();
match generate(&job.messages, &job.cfg, &mut on_token, &mut should_cancel) {
Ok(output) => {
let _ = job.tx.send(WorkerEvent::Complete(output));
}
Err(WorkerFailure::Rejected(api_err)) => {
let _ = job.tx.send(WorkerEvent::Rejected(api_err));
}
Err(WorkerFailure::Failed(message)) => {
eprintln!("[metal-worker] generation error: {message}");
let _ = job.tx.send(WorkerEvent::Failed(message));
}
Err(WorkerFailure::ConstraintBlocked(message)) => {
eprintln!("[metal-worker] generation error: {message}");
let _ = job.tx.send(WorkerEvent::ConstraintBlocked(message));
}
}
}
}
pub struct MetalWorker;
impl MetalWorker {
pub fn spawn(
loader: impl FnOnce() -> Result<(MetalQwen35State, BpeTokenizer, WorkerMetadata), String>
+ Send
+ 'static,
max_pending: usize,
) -> Result<(MetalWorkerOwner, MetalWorkerClient, WorkerMetadata), StartupError> {
if max_pending == 0 || max_pending > Semaphore::MAX_PERMITS {
return Err(StartupError::InvalidMaxPending { max_pending });
}
let (job_tx, job_rx) = mpsc::unbounded_channel::<WorkerJob>();
let admission = Arc::new(Semaphore::new(max_pending));
let (ready_tx, ready_rx) = std::sync::mpsc::channel::<Result<WorkerMetadata, String>>();
let join_handle = std::thread::spawn(move || match loader() {
Ok((mut state, tokenizer, meta)) => {
let _ = ready_tx.send(Ok(meta.clone()));
run_worker_loop(job_rx, move |messages, cfg, on_token, should_cancel| {
let prompt = format_chat_template(messages);
let prompt_len = tokenizer.tokenize(&prompt).real_length;
check_prompt_fits_window(
meta.context_window_policy,
meta.model_max_context,
prompt_len,
cfg,
)
.map_err(WorkerFailure::Rejected)?;
let cached = state.generate_streaming_with_prefix_cache_and_cancel(
CrossTurnSlotId::DEFAULT,
&prompt,
&tokenizer,
cfg,
on_token,
should_cancel,
);
if let Ok(c) = &cached {
eprintln!(
"[metal-worker] cross-turn cache: mode={:?} reused={} \
prefetched={} prompt={}",
c.cache.mode,
c.cache.reused_tokens,
c.cache.prefetched_tokens,
c.cache.prompt_tokens,
);
}
cached.map(|c| c.output).map_err(WorkerFailure::from)
});
}
Err(e) => {
let _ = ready_tx.send(Err(e));
}
});
let owner = MetalWorkerOwner::from_handle(join_handle);
match ready_rx.recv() {
Ok(Ok(meta)) => {
let client = MetalWorkerClient::with_owner(job_tx, admission, owner.clone());
Ok((owner, client, meta))
}
Ok(Err(e)) => Err(StartupError::Load(e)),
Err(_) => Err(StartupError::ThreadExited),
}
}
}
#[cfg(any(test, feature = "test-utils"))]
impl WorkerJob {
pub fn reply(&self, event: WorkerEvent) -> bool {
self.tx.send(event).is_ok()
}
}
#[cfg(any(test, feature = "test-utils"))]
pub fn test_client_and_jobs() -> (MetalWorkerClient, mpsc::UnboundedReceiver<WorkerJob>) {
test_client_and_jobs_with_cap(TEST_EFFECTIVELY_UNBOUNDED_CAP)
}
#[cfg(any(test, feature = "test-utils"))]
pub fn test_client_and_jobs_with_cap(
max_pending: usize,
) -> (MetalWorkerClient, mpsc::UnboundedReceiver<WorkerJob>) {
let (job_tx, job_rx) = mpsc::unbounded_channel::<WorkerJob>();
(
MetalWorkerClient::unattached_for_test(job_tx, Arc::new(Semaphore::new(max_pending))),
job_rx,
)
}
#[cfg(any(test, feature = "test-utils"))]
const TEST_EFFECTIVELY_UNBOUNDED_CAP: usize = 1_000_000;
#[cfg(any(test, feature = "test-utils"))]
#[allow(clippy::type_complexity)]
pub fn spawn_fake(
context_window_policy: ContextWindowPolicy,
model_max_context: usize,
tokenizer: BpeTokenizer,
generate: impl FnMut(
&[ChatMessage],
&GenerateConfig,
usize,
&mut dyn FnMut(&str, u32) -> bool,
&mut dyn FnMut() -> bool,
) -> Result<GenerateOutput, String>
+ Send
+ 'static,
) -> MetalWorkerClient {
spawn_fake_with_cap(
TEST_EFFECTIVELY_UNBOUNDED_CAP,
context_window_policy,
model_max_context,
tokenizer,
generate,
)
}
#[cfg(any(test, feature = "test-utils"))]
#[allow(clippy::type_complexity)]
pub fn spawn_fake_with_cap(
max_pending: usize,
context_window_policy: ContextWindowPolicy,
model_max_context: usize,
tokenizer: BpeTokenizer,
mut generate: impl FnMut(
&[ChatMessage],
&GenerateConfig,
usize,
&mut dyn FnMut(&str, u32) -> bool,
&mut dyn FnMut() -> bool,
) -> Result<GenerateOutput, String>
+ Send
+ 'static,
) -> MetalWorkerClient {
let (job_tx, job_rx) = mpsc::unbounded_channel::<WorkerJob>();
let join_handle = std::thread::spawn(move || {
run_worker_loop(job_rx, move |messages, cfg, on_token, should_cancel| {
let prompt = format_chat_template(messages);
let prompt_tokens = tokenizer.tokenize(&prompt).real_length;
check_prompt_fits_window(context_window_policy, model_max_context, prompt_tokens, cfg)
.map_err(WorkerFailure::Rejected)?;
generate(messages, cfg, prompt_tokens, on_token, should_cancel)
.map_err(WorkerFailure::Failed)
});
});
let owner = MetalWorkerOwner::from_handle(join_handle);
MetalWorkerClient::with_owner(job_tx, Arc::new(Semaphore::new(max_pending)), owner)
}
#[cfg(test)]
mod tests {
use super::*;
use std::cell::RefCell;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::time::Duration;
struct BlockingThreadLocalDrop {
entered: std::sync::mpsc::SyncSender<()>,
release: std::sync::mpsc::Receiver<()>,
finished: std::sync::mpsc::SyncSender<()>,
}
impl Drop for BlockingThreadLocalDrop {
fn drop(&mut self) {
let _ = self.entered.send(());
let _ = self.release.recv();
let _ = self.finished.send(());
}
}
thread_local! {
static BLOCKING_THREAD_LOCAL_DROP: RefCell<Option<BlockingThreadLocalDrop>> =
const { RefCell::new(None) };
}
fn test_owner(
join_handle: std::thread::JoinHandle<()>,
drop_timeout: Duration,
) -> MetalWorkerOwner {
MetalWorkerOwner::from_handle_with_timeout(join_handle, drop_timeout)
}
#[allow(clippy::type_complexity)]
fn fake_generate(
cap: usize,
started: Arc<AtomicUsize>,
ran_tokens: Arc<AtomicUsize>,
) -> impl FnMut(
&[ChatMessage],
&GenerateConfig,
&mut dyn FnMut(&str, u32) -> bool,
&mut dyn FnMut() -> bool,
) -> Result<GenerateOutput, WorkerFailure> {
move |_messages, _cfg, on_token, should_cancel| {
started.fetch_add(1, Ordering::SeqCst);
let mut n = 0usize;
for i in 0..cap {
std::thread::sleep(Duration::from_millis(5));
if should_cancel() {
break;
}
if !on_token("x", i as u32) {
break;
}
n += 1;
ran_tokens.fetch_add(1, Ordering::SeqCst);
}
Ok(GenerateOutput {
text: "x".repeat(n),
token_ids: vec![0; n],
prompt_tokens: 1,
generated_tokens: n,
stopped: false,
stop_reason: None,
token_logprobs: vec![],
})
}
}
#[allow(clippy::type_complexity)]
fn fake_generate_with_prefill_gap(
prefill_steps: usize,
decode_cap: usize,
entered_decode: Arc<AtomicBool>,
) -> impl FnMut(
&[ChatMessage],
&GenerateConfig,
&mut dyn FnMut(&str, u32) -> bool,
&mut dyn FnMut() -> bool,
) -> Result<GenerateOutput, WorkerFailure> {
move |_messages, _cfg, on_token, should_cancel| {
for _ in 0..prefill_steps {
std::thread::sleep(Duration::from_millis(5));
if should_cancel() {
return Ok(GenerateOutput {
text: String::new(),
token_ids: vec![],
prompt_tokens: 1,
generated_tokens: 0,
stopped: false,
stop_reason: None,
token_logprobs: vec![],
});
}
}
entered_decode.store(true, Ordering::SeqCst);
let mut n = 0usize;
for i in 0..decode_cap {
std::thread::sleep(Duration::from_millis(5));
if should_cancel() {
break;
}
if !on_token("x", i as u32) {
break;
}
n += 1;
}
Ok(GenerateOutput {
text: "x".repeat(n),
token_ids: vec![0; n],
prompt_tokens: 1,
generated_tokens: n,
stopped: false,
stop_reason: None,
token_logprobs: vec![],
})
}
}
#[allow(clippy::type_complexity)]
fn fake_generate_fails_once_then_succeeds(
message: &'static str,
call_count: Arc<AtomicUsize>,
) -> impl FnMut(
&[ChatMessage],
&GenerateConfig,
&mut dyn FnMut(&str, u32) -> bool,
&mut dyn FnMut() -> bool,
) -> Result<GenerateOutput, WorkerFailure> {
move |_messages, _cfg, on_token, _should_cancel| {
if call_count.fetch_add(1, Ordering::SeqCst) == 0 {
return Err(WorkerFailure::Failed(message.to_string()));
}
let _ = on_token("x", 0);
Ok(GenerateOutput {
text: "x".to_string(),
token_ids: vec![0],
prompt_tokens: 1,
generated_tokens: 1,
stopped: true,
stop_reason: None,
token_logprobs: vec![],
})
}
}
fn make_job() -> (
WorkerJob,
mpsc::UnboundedReceiver<WorkerEvent>,
crate::serve::CancelOnDrop,
) {
let (tx, rx) = mpsc::unbounded_channel::<WorkerEvent>();
let (cancel_guard, cancel_rx) = crate::serve::cancel_pair();
let permit = Arc::new(Semaphore::new(1))
.try_acquire_owned()
.expect("fresh single-permit semaphore must have a permit available");
let job = WorkerJob {
messages: vec![ChatMessage::user("hi")],
cfg: GenerateConfig::default(),
tx,
cancel: cancel_rx,
_admission_permit: permit,
};
(job, rx, cancel_guard)
}
#[test]
fn queued_job_cancelled_before_dequeue_sends_exactly_one_cancelled_event() {
let (job_tx, job_rx) = mpsc::unbounded_channel::<WorkerJob>();
let started = Arc::new(AtomicUsize::new(0));
let ran_tokens = Arc::new(AtomicUsize::new(0));
let (job1, rx1, _guard1) = make_job();
job_tx.send(job1).unwrap();
let (job2, mut rx2, guard2) = make_job();
job_tx.send(job2).unwrap();
drop(guard2);
let (job3, rx3, _guard3) = make_job();
job_tx.send(job3).unwrap();
drop(job_tx);
let started2 = started.clone();
let ran2 = ran_tokens.clone();
let handle =
std::thread::spawn(move || run_worker_loop(job_rx, fake_generate(50, started2, ran2)));
let completion_tokens_of = |mut rx: mpsc::UnboundedReceiver<WorkerEvent>| -> Option<usize> {
let mut ct = None;
while let Some(ev) = rx.blocking_recv() {
if let WorkerEvent::Complete(output) = ev {
ct = Some(output.generated_tokens);
}
}
ct
};
assert_eq!(
completion_tokens_of(rx1),
Some(50),
"job 1 should run to completion undisturbed"
);
match rx2.blocking_recv() {
Some(WorkerEvent::Cancelled) => {}
other => panic!("expected exactly one Cancelled event, got {other:?}"),
}
assert!(
rx2.blocking_recv().is_none(),
"cancelled queued job must produce no further events after Cancelled"
);
assert_eq!(
completion_tokens_of(rx3),
Some(50),
"worker must survive cancelling job 2 and serve job 3 normally afterward"
);
handle.join().expect("worker thread must not panic");
assert_eq!(
started.load(Ordering::SeqCst),
2,
"generate() must run exactly twice (job 1, job 3) -- never for cancelled job 2"
);
assert_eq!(
ran_tokens.load(Ordering::SeqCst),
100,
"50 real fake-tokens each for job 1 and job 3, zero for cancelled job 2"
);
}
#[test]
fn job_whose_event_receiver_is_already_closed_is_cancelled_without_running_generate() {
let (job_tx, job_rx) = mpsc::unbounded_channel::<WorkerJob>();
let (tx, rx) = mpsc::unbounded_channel::<WorkerEvent>();
drop(rx);
let (_guard, cancel_rx) = crate::serve::cancel_pair();
let permit = Arc::new(Semaphore::new(1))
.try_acquire_owned()
.expect("fresh single-permit semaphore must have a permit available");
let job = WorkerJob {
messages: vec![ChatMessage::user("hi")],
cfg: GenerateConfig::default(),
tx,
cancel: cancel_rx,
_admission_permit: permit,
};
job_tx.send(job).unwrap();
drop(job_tx);
let started = Arc::new(AtomicUsize::new(0));
let ran_tokens = Arc::new(AtomicUsize::new(0));
let started2 = started.clone();
let ran2 = ran_tokens.clone();
let handle =
std::thread::spawn(move || run_worker_loop(job_rx, fake_generate(50, started2, ran2)));
handle.join().expect("worker thread must not panic");
assert_eq!(
started.load(Ordering::SeqCst),
0,
"generate() must never run for a job whose event receiver was already closed"
);
assert_eq!(ran_tokens.load(Ordering::SeqCst), 0);
}
#[test]
fn running_job_cancelled_midstream_stops_early_and_worker_survives() {
let (job_tx, job_rx) = mpsc::unbounded_channel::<WorkerJob>();
let started = Arc::new(AtomicUsize::new(0));
let ran_tokens = Arc::new(AtomicUsize::new(0));
let (job1, mut rx1, guard1) = make_job();
job_tx.send(job1).unwrap();
let mut guard1 = Some(guard1);
let (job2, mut rx2, _guard2) = make_job();
job_tx.send(job2).unwrap();
drop(job_tx);
let started2 = started.clone();
let ran2 = ran_tokens.clone();
let handle = std::thread::spawn(move || {
run_worker_loop(job_rx, fake_generate(2000, started2, ran2))
});
let mut seen = 0;
loop {
match rx1.blocking_recv() {
Some(WorkerEvent::Delta(_)) => {
seen += 1;
if seen == 5 {
guard1.take();
}
}
Some(WorkerEvent::Complete(output)) => {
assert!(
output.generated_tokens < 2000,
"job 1 must stop well short of its 2000-token cap after \
cancellation, got {}",
output.generated_tokens
);
assert!(
output.generated_tokens < 100,
"job 1 must stop within a handful of tokens of the client \
disconnecting, not run on regardless; got {}",
output.generated_tokens
);
break;
}
Some(WorkerEvent::Failed(message)) => {
panic!("fake_generate never fails; unexpected Failed: {message}")
}
Some(WorkerEvent::ConstraintBlocked(message)) => {
panic!(
"fake_generate never blocks on a grammar constraint; unexpected \
ConstraintBlocked: {message}"
)
}
Some(WorkerEvent::Rejected(err)) => {
panic!("fake_generate never rejects; unexpected Rejected: {err:?}")
}
Some(WorkerEvent::Cancelled) => {
panic!("job 1 was already running -- Cancelled is a dequeue-only event")
}
None => panic!("job 1's reply channel closed before a Complete event"),
}
}
let mut n2 = None;
while let Some(ev) = rx2.blocking_recv() {
if let WorkerEvent::Complete(output) = ev {
n2 = Some(output.generated_tokens);
}
}
assert_eq!(
n2,
Some(2000),
"worker must survive mid-stream cancellation and serve the next job to completion"
);
handle.join().expect("worker thread must not panic");
}
#[test]
fn running_job_cancelled_during_prefill_like_phase_never_calls_on_token() {
let (job_tx, job_rx) = mpsc::unbounded_channel::<WorkerJob>();
let entered_decode = Arc::new(AtomicBool::new(false));
let (job1, mut rx1, guard1) = make_job();
job_tx.send(job1).unwrap();
job_tx.send(make_job().0).unwrap_or(()); drop(job_tx);
let entered2 = entered_decode.clone();
let handle = std::thread::spawn(move || {
run_worker_loop(job_rx, fake_generate_with_prefill_gap(400, 50, entered2))
});
std::thread::sleep(Duration::from_millis(20));
drop(guard1);
match rx1.blocking_recv() {
Some(WorkerEvent::Delta(_)) => panic!(
"on_token must never be called: cancellation happened while the fake \
generator was still in its prefill-like phase, which does not call \
on_token at all"
),
Some(WorkerEvent::Complete(output)) => {
assert_eq!(
output.generated_tokens, 0,
"job cancelled during the prefill-like phase must produce zero tokens, \
got {}",
output.generated_tokens
);
}
Some(WorkerEvent::Failed(message)) => {
panic!("fake_generate_with_prefill_gap never fails; unexpected Failed: {message}")
}
Some(WorkerEvent::ConstraintBlocked(message)) => {
panic!(
"fake_generate_with_prefill_gap never blocks on a grammar constraint; \
unexpected ConstraintBlocked: {message}"
)
}
Some(WorkerEvent::Rejected(err)) => {
panic!("fake_generate_with_prefill_gap never rejects; unexpected Rejected: {err:?}")
}
Some(WorkerEvent::Cancelled) => {
panic!("job 1 was already dequeued and running -- not a dequeue-time cancel")
}
None => panic!("job 1's reply channel closed before a Complete event"),
}
handle.join().expect("worker thread must not panic");
assert!(
!entered_decode.load(Ordering::SeqCst),
"should_cancel alone (on_token is never called during this phase) must stop \
the job before the decode phase is ever reached"
);
}
#[test]
fn generation_failure_is_reported_as_failed_not_complete() {
let (job_tx, job_rx) = mpsc::unbounded_channel::<WorkerJob>();
let (job1, mut rx1, _guard1) = make_job();
job_tx.send(job1).unwrap();
let (job2, mut rx2, _guard2) = make_job();
job_tx.send(job2).unwrap();
drop(job_tx);
let call_count = Arc::new(AtomicUsize::new(0));
let handle = std::thread::spawn({
let call_count = call_count.clone();
move || {
run_worker_loop(
job_rx,
fake_generate_fails_once_then_succeeds(
"grammar constraint blocked every token; no legal continuation \
exists in the current grammar state",
call_count,
),
)
}
});
match rx1.blocking_recv() {
Some(WorkerEvent::Failed(message)) => {
assert!(
message.contains("grammar constraint blocked every token"),
"Failed must carry the underlying error message, got: {message}"
);
}
Some(WorkerEvent::Complete(_)) => panic!(
"a failed generation must never be reported as Complete -- that would \
silently hand the HTTP layer a fabricated result for a request that \
produced no legal output"
),
other => panic!("expected Failed as the first and only event, got {other:?}"),
}
let mut done = None;
while let Some(ev) = rx2.blocking_recv() {
if let WorkerEvent::Complete(output) = ev {
done = Some(output.generated_tokens);
}
}
assert_eq!(
done,
Some(1),
"worker thread must survive a failed generation and serve the next job \
normally afterward"
);
handle
.join()
.expect("worker thread must not panic on a generation error");
}
#[test]
fn queue_closure_lets_the_worker_thread_exit_and_join() {
let (job_tx, job_rx) = mpsc::unbounded_channel::<WorkerJob>();
let started = Arc::new(AtomicUsize::new(0));
let ran_tokens = Arc::new(AtomicUsize::new(0));
let started2 = started.clone();
let ran2 = ran_tokens.clone();
let handle =
std::thread::spawn(move || run_worker_loop(job_rx, fake_generate(1, started2, ran2)));
drop(job_tx);
handle
.join()
.expect("worker thread must exit and be joinable once every job sender drops");
assert_eq!(started.load(Ordering::SeqCst), 0);
}
#[test]
fn owner_shutdown_joins_cleanly_once_the_queue_closes() {
let (job_tx, job_rx) = mpsc::unbounded_channel::<WorkerJob>();
let started = Arc::new(AtomicUsize::new(0));
let ran_tokens = Arc::new(AtomicUsize::new(0));
let started2 = started.clone();
let ran2 = ran_tokens.clone();
let join_handle = std::thread::spawn(move || {
run_worker_loop(job_rx, fake_generate(1, started2, ran2));
});
let owner = test_owner(join_handle, Duration::from_secs(1));
drop(job_tx);
assert_eq!(
owner._inner.wait_for_exit(Duration::from_secs(1)),
WorkerShutdown::Joined
);
assert_eq!(
owner._inner.wait_for_exit(Duration::ZERO),
WorkerShutdown::AlreadyStopped,
"the join handle must be claimed exactly once"
);
assert_eq!(started.load(Ordering::SeqCst), 0);
}
#[test]
fn owner_shutdown_deadline_covers_blocking_thread_local_destructor() {
let (destructor_entered_tx, destructor_entered_rx) = std::sync::mpsc::sync_channel(1);
let (release_destructor_tx, release_destructor_rx) = std::sync::mpsc::sync_channel(1);
let (destructor_finished_tx, destructor_finished_rx) = std::sync::mpsc::sync_channel(1);
let join_handle = std::thread::spawn(move || {
BLOCKING_THREAD_LOCAL_DROP.with(|slot| {
*slot.borrow_mut() = Some(BlockingThreadLocalDrop {
entered: destructor_entered_tx,
release: release_destructor_rx,
finished: destructor_finished_tx,
});
});
});
destructor_entered_rx
.recv_timeout(Duration::from_secs(1))
.expect("the worker must enter its thread-local destructor");
let finished_deadline = Instant::now() + Duration::from_secs(1);
while !join_handle.is_finished() && Instant::now() < finished_deadline {
std::thread::yield_now();
}
assert!(
join_handle.is_finished(),
"the worker main function must finish while its thread-local destructor is blocked"
);
let owner = test_owner(join_handle, Duration::from_millis(20));
let (shutdown_done_tx, shutdown_done_rx) = std::sync::mpsc::sync_channel(1);
let shutdown_thread = std::thread::spawn(move || {
let started = Instant::now();
let result = owner._inner.wait_for_exit(Duration::from_millis(20));
let _ = shutdown_done_tx.send((result, started.elapsed()));
});
let before_watchdog = shutdown_done_rx
.recv_timeout(Duration::from_millis(250))
.ok();
release_destructor_tx
.send(())
.expect("the blocked thread-local destructor must still accept its release");
destructor_finished_rx
.recv_timeout(Duration::from_secs(1))
.expect("the thread-local destructor must finish after release");
let observed = match before_watchdog {
Some(result) => result,
None => shutdown_done_rx
.recv_timeout(Duration::from_secs(1))
.expect("shutdown must finish after the destructor is released"),
};
shutdown_thread
.join()
.expect("shutdown helper thread must not panic");
let Some((result, elapsed)) = before_watchdog else {
panic!(
"the configured deadline must include thread-local destructor cleanup; \
observed {observed:?}"
);
};
assert_eq!(
result,
WorkerShutdown::TimedOut,
"blocked thread-local cleanup must exhaust the configured deadline"
);
assert!(
elapsed < Duration::from_millis(250),
"shutdown exceeded the deadline watchdog: {elapsed:?}"
);
}
#[test]
fn final_client_drop_closes_queue_before_owner_joins() {
let (job_tx, mut job_rx) = mpsc::unbounded_channel::<WorkerJob>();
let (queue_closed_tx, queue_closed_rx) = std::sync::mpsc::sync_channel(1);
let (allow_exit_tx, allow_exit_rx) = std::sync::mpsc::sync_channel(1);
let join_handle = std::thread::spawn(move || {
while job_rx.blocking_recv().is_some() {}
let _ = queue_closed_tx.send(());
let _ = allow_exit_rx.recv();
});
let owner = test_owner(join_handle, Duration::from_secs(2));
let client =
MetalWorkerClient::with_owner(job_tx, Arc::new(Semaphore::new(1)), owner.clone());
drop(owner);
let (drop_done_tx, drop_done_rx) = std::sync::mpsc::sync_channel(1);
let drop_thread = std::thread::spawn(move || {
drop(client);
let _ = drop_done_tx.send(());
});
queue_closed_rx
.recv_timeout(Duration::from_millis(500))
.expect("dropping the last client must close the queue");
assert!(
matches!(
drop_done_rx.try_recv(),
Err(std::sync::mpsc::TryRecvError::Empty)
),
"last client drop must still be waiting while the worker is live"
);
allow_exit_tx
.send(())
.expect("the worker must still be waiting for the exit release");
drop_done_rx
.recv_timeout(Duration::from_secs(1))
.expect("last client drop must finish after the worker exits");
drop_thread
.join()
.expect("client drop thread must not panic");
}
#[test]
fn final_client_drop_timeout_detaches_instead_of_blocking() {
let (worker_started_tx, worker_started_rx) = std::sync::mpsc::sync_channel(1);
let (release_worker_tx, release_worker_rx) = std::sync::mpsc::sync_channel(1);
let (worker_done_tx, worker_done_rx) = std::sync::mpsc::sync_channel(1);
let join_handle = std::thread::spawn(move || {
let _ = worker_started_tx.send(());
let _ = release_worker_rx.recv();
let _ = worker_done_tx.send(());
});
worker_started_rx
.recv_timeout(Duration::from_secs(1))
.expect("the worker must reach its stuck-backend stand-in");
let owner = test_owner(join_handle, Duration::from_millis(20));
let (job_tx, _job_rx) = mpsc::unbounded_channel::<WorkerJob>();
let client =
MetalWorkerClient::with_owner(job_tx, Arc::new(Semaphore::new(1)), owner.clone());
drop(owner);
let (drop_done_tx, drop_done_rx) = std::sync::mpsc::sync_channel(1);
let drop_thread = std::thread::spawn(move || {
drop(client);
let _ = drop_done_tx.send(());
});
let returned_before_watchdog = drop_done_rx
.recv_timeout(Duration::from_millis(500))
.is_ok();
release_worker_tx
.send(())
.expect("detached worker must still accept the cleanup release");
worker_done_rx
.recv_timeout(Duration::from_secs(1))
.expect("detached worker must exit after the cleanup release");
drop_thread
.join()
.expect("timed-out client drop thread must not panic");
assert!(
returned_before_watchdog,
"last client Drop must honor its configured deadline instead of joining a stuck worker"
);
}
#[test]
fn owner_shutdown_reports_worker_panic_after_join() {
let join_handle = std::thread::spawn(move || {
panic!("simulated worker panic");
});
let owner = test_owner(join_handle, Duration::from_secs(1));
assert_eq!(
owner._inner.wait_for_exit(Duration::from_secs(1)),
WorkerShutdown::Panicked
);
}
#[test]
fn loader_failure_before_readiness_is_reported_without_touching_a_device() {
let result = MetalWorker::spawn(
|| Err("simulated load failure".to_string()),
DEFAULT_MAX_PENDING_JOBS,
);
match result {
Err(StartupError::Load(message)) => {
assert_eq!(message, "simulated load failure");
}
other => panic!("expected StartupError::Load, got {other:?}"),
}
}
#[test]
fn startup_error_display_matches_each_variant() {
assert_eq!(StartupError::Load("boom".to_string()).to_string(), "boom");
assert_eq!(
StartupError::ThreadExited.to_string(),
"worker thread exited before loading finished"
);
assert_eq!(
StartupError::InvalidMaxPending { max_pending: 0 }.to_string(),
format!(
"--max-pending must be between 1 and {} (got 0)",
Semaphore::MAX_PERMITS
)
);
}
#[test]
fn max_pending_zero_is_rejected_before_semaphore_new() {
let result = MetalWorker::spawn(
|| -> Result<(MetalQwen35State, BpeTokenizer, WorkerMetadata), String> {
panic!("loader must not run: max_pending=0 must be rejected first")
},
0,
);
match result {
Err(StartupError::InvalidMaxPending { max_pending: 0 }) => {}
other => panic!("expected InvalidMaxPending{{max_pending: 0}}, got {other:?}"),
}
}
#[test]
fn max_pending_above_max_permits_is_rejected_before_semaphore_new() {
let too_big = Semaphore::MAX_PERMITS + 1;
let result = MetalWorker::spawn(
|| -> Result<(MetalQwen35State, BpeTokenizer, WorkerMetadata), String> {
panic!("loader must not run: max_pending above MAX_PERMITS must be rejected first")
},
too_big,
);
match result {
Err(StartupError::InvalidMaxPending { max_pending }) => {
assert_eq!(max_pending, too_big);
}
other => panic!("expected InvalidMaxPending, got {other:?}"),
}
}
fn cfg_with(max_new_tokens: usize, reasoning_budget: Option<usize>) -> GenerateConfig {
GenerateConfig {
max_new_tokens,
reasoning_budget,
..Default::default()
}
}
#[test]
fn check_prompt_fits_window_rejects_when_prompt_plus_decode_overflows() {
let cfg = cfg_with(7, None);
let err = check_prompt_fits_window(
ContextWindowPolicy::PromptAndDecodeWithDelimiter,
8,
2,
&cfg,
)
.unwrap_err();
match err {
ApiError::BadRequest { message, code } => {
assert_eq!(code, "context_length_exceeded");
assert!(
message.contains("2 tokens") && message.contains("8-token"),
"error must name the actual prompt length and window: {message}"
);
}
other => panic!("expected BadRequest, got {other:?}"),
}
}
#[test]
fn lattice_context_boundary_accepts_exact_window_and_rejects_one_past() {
let cfg = cfg_with(7, None);
assert!(
check_prompt_fits_window(ContextWindowPolicy::PromptAndMaxTokens, 8, 1, &cfg).is_ok()
);
assert!(
check_prompt_fits_window(ContextWindowPolicy::PromptAndMaxTokens, 8, 2, &cfg).is_err()
);
}
#[test]
fn lattice_policy_rejects_zero_token_prompt_even_when_window_fits() {
let cfg = cfg_with(7, None);
let err = check_prompt_fits_window(ContextWindowPolicy::PromptAndMaxTokens, 8, 0, &cfg)
.unwrap_err();
match err {
ApiError::BadRequest { message, code } => {
assert_eq!(code, "context_length_exceeded");
assert!(
message.contains("0 tokens"),
"error must name the zero-length prompt: {message}"
);
}
other => panic!("expected BadRequest, got {other:?}"),
}
assert!(
check_prompt_fits_window(
ContextWindowPolicy::PromptAndDecodeWithDelimiter,
9,
0,
&cfg_with(7, None),
)
.is_ok()
);
}
#[test]
fn lattice_serve_context_boundary_accepts_exact_window_and_rejects_one_past() {
let at_boundary = cfg_with(5, Some(1));
assert!(
check_prompt_fits_window(
ContextWindowPolicy::PromptAndDecodeWithDelimiter,
8,
1,
&at_boundary,
)
.is_ok()
);
let one_past = cfg_with(6, Some(1));
assert!(
check_prompt_fits_window(
ContextWindowPolicy::PromptAndDecodeWithDelimiter,
8,
1,
&one_past,
)
.is_err()
);
}
#[test]
fn check_prompt_fits_window_accepts_ordinary_prompt_unclamped() {
let cfg = cfg_with(50, None);
assert!(
check_prompt_fits_window(
ContextWindowPolicy::PromptAndDecodeWithDelimiter,
4096,
100,
&cfg,
)
.is_ok()
);
}
#[test]
fn submit_rejects_once_admission_cap_reached() {
let cap = 2;
let (client, job_rx) = test_client_and_jobs_with_cap(cap);
let started = Arc::new(AtomicUsize::new(0));
let ran_tokens = Arc::new(AtomicUsize::new(0));
let started2 = started.clone();
let ran2 = ran_tokens.clone();
let handle = std::thread::spawn(move || {
run_worker_loop(job_rx, fake_generate(2000, started2, ran2))
});
let (guard1, cancel1) = crate::serve::cancel_pair();
let rx1 = client
.submit(
vec![ChatMessage::user("hi")],
GenerateConfig::default(),
cancel1,
)
.expect("job 1 must be admitted: cap=2, 0 outstanding");
std::thread::sleep(Duration::from_millis(30));
let (guard2, cancel2) = crate::serve::cancel_pair();
let rx2 = client
.submit(
vec![ChatMessage::user("hi")],
GenerateConfig::default(),
cancel2,
)
.expect("job 2 must be admitted: cap=2, 1 outstanding");
let (_guard3, cancel3) = crate::serve::cancel_pair();
let err = client
.submit(
vec![ChatMessage::user("hi")],
GenerateConfig::default(),
cancel3,
)
.expect_err("job 3 must be rejected once the cap is reached");
match err {
ApiError::ServiceUnavailable { message } => {
assert!(
message.contains("outstanding") || message.contains("pending"),
"rejection message should explain admission capacity: {message}"
);
}
other => panic!("expected ServiceUnavailable, got {other:?}"),
}
drop(guard1);
drop(guard2);
drop(rx1);
drop(rx2);
drop(client);
handle.join().expect("worker thread must not panic");
}
#[test]
fn admission_slot_is_released_when_a_job_completes() {
let cap = 1;
let (client, job_rx) = test_client_and_jobs_with_cap(cap);
let started = Arc::new(AtomicUsize::new(0));
let ran_tokens = Arc::new(AtomicUsize::new(0));
let started2 = started.clone();
let ran2 = ran_tokens.clone();
let handle =
std::thread::spawn(move || run_worker_loop(job_rx, fake_generate(5, started2, ran2)));
let (_guard1, cancel1) = crate::serve::cancel_pair();
let mut rx1 = client
.submit(
vec![ChatMessage::user("hi")],
GenerateConfig::default(),
cancel1,
)
.expect("job 1 must be admitted");
let mut completed = false;
while let Some(ev) = rx1.blocking_recv() {
if matches!(ev, WorkerEvent::Complete(_)) {
completed = true;
}
}
assert!(completed, "job 1 must complete normally");
let mut admitted = false;
for _ in 0..50 {
let (_guard2, cancel2) = crate::serve::cancel_pair();
match client.submit(
vec![ChatMessage::user("hi")],
GenerateConfig::default(),
cancel2,
) {
Ok(_rx2) => {
admitted = true;
break;
}
Err(_) => std::thread::sleep(Duration::from_millis(5)),
}
}
assert!(
admitted,
"slot must be released once job 1 completes, admitting job 2 at the same cap=1"
);
drop(client);
handle.join().expect("worker thread must not panic");
}
#[test]
fn admission_slot_is_released_when_a_queued_job_is_cancelled() {
let cap = 2;
let (client, job_rx) = test_client_and_jobs_with_cap(cap);
let started = Arc::new(AtomicUsize::new(0));
let ran_tokens = Arc::new(AtomicUsize::new(0));
let started2 = started.clone();
let ran2 = ran_tokens.clone();
let handle = std::thread::spawn(move || {
run_worker_loop(job_rx, fake_generate(2000, started2, ran2))
});
let (guard1, cancel1) = crate::serve::cancel_pair();
let rx1 = client
.submit(
vec![ChatMessage::user("hi")],
GenerateConfig::default(),
cancel1,
)
.expect("job 1 must be admitted");
std::thread::sleep(Duration::from_millis(30));
let (guard2, cancel2) = crate::serve::cancel_pair();
let mut rx2 = client
.submit(
vec![ChatMessage::user("hi")],
GenerateConfig::default(),
cancel2,
)
.expect("job 2 must be admitted");
drop(guard2);
let (_guard3, cancel3) = crate::serve::cancel_pair();
client
.submit(
vec![ChatMessage::user("hi")],
GenerateConfig::default(),
cancel3,
)
.expect_err("cap must still be full: job 2's slot isn't released until dequeued");
drop(guard1);
match rx2.blocking_recv() {
Some(WorkerEvent::Cancelled) => {}
other => panic!("expected job 2's exactly-one Cancelled event, got {other:?}"),
}
let mut permits_restored = false;
for _ in 0..50 {
if client.admission.available_permits() == cap {
permits_restored = true;
break;
}
std::thread::sleep(Duration::from_millis(5));
}
assert!(
permits_restored,
"both permits (job 1's own release AND job 2's queued-cancel release) must be \
free once job 2's Cancelled event has fired, got {} of {cap}",
client.admission.available_permits()
);
let (_guard4, cancel4) = crate::serve::cancel_pair();
client
.submit(
vec![ChatMessage::user("hi")],
GenerateConfig::default(),
cancel4,
)
.expect("job 2's slot must be released after its Cancelled event, not leaked");
drop(rx1);
drop(client);
handle.join().expect("worker thread must not panic");
}
}