use crate::forward::metal_qwen35::{
ChatMessage, MetalQwen35State, format_chat_template, push_chat_turn_close, push_chat_turn_open,
};
use crate::kv_cache::CrossTurnSlotId;
use crate::model::qwen35_config::{
GenerateConfig, GenerateOutput, Qwen35Config, VisionModelConfig,
};
use crate::serve::ApiError;
use crate::tokenizer::Tokenizer as _;
use crate::tokenizer::bpe::BpeTokenizer;
use crate::vision::checkpoint::{
Qwen35VisionWeights, load_qwen35_vision_weights_with_cancel,
validate_qwen35_vision_weight_inventory,
};
use crate::vision::multimodal::Qwen35VisionRequest;
use crate::vision::qwen35_merger::qwen35_merger_forward_with_cancel;
use crate::vision::qwen35_vit::preprocess_qwen35_image_for_serve;
use crate::vision::qwen35_vit_metal::qwen35_vit_forward_metal_with_cancel;
use std::io::Write as _;
use std::path::PathBuf;
use std::sync::atomic::{AtomicBool, Ordering};
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,
}
#[derive(Debug)]
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>,
vision_supported: Arc<AtomicBool>,
_owner: MetalWorkerOwner,
}
impl MetalWorkerClient {
fn with_owner(
jobs: mpsc::UnboundedSender<WorkerJob>,
admission: Arc<Semaphore>,
vision_supported: Arc<AtomicBool>,
owner: MetalWorkerOwner,
) -> Self {
Self {
jobs: Some(jobs),
admission,
vision_supported,
_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,
Arc::new(AtomicBool::new(false)),
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()
}
pub fn supports_vision(&self) -> bool {
self.vision_supported.load(Ordering::Acquire)
}
}
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));
}
}
}
}
enum VisionState {
Unsupported,
Pending {
model_dir: PathBuf,
config: VisionModelConfig,
},
Loaded(Qwen35VisionWeights),
Failed(String),
}
#[derive(Debug)]
enum VisionRuntimeLoad<'a> {
Ready(&'a Qwen35VisionWeights),
Unsupported,
Cancelled,
}
pub struct VisionRuntime {
state: VisionState,
vision_supported: Arc<AtomicBool>,
}
impl VisionRuntime {
pub fn from_model_config(model_dir: PathBuf, config: &Qwen35Config) -> Self {
let token_metadata_present = match (
config.image_token_id,
config.vision_start_token_id,
config.vision_end_token_id,
) {
(Some(image), Some(start), Some(end)) => {
[image, start, end]
.into_iter()
.all(|token| (token as usize) < config.vocab_size)
&& image != start
&& image != end
&& start != end
}
_ => false,
};
let supported_weight_source = token_metadata_present
&& config.vision_config.as_ref().is_some_and(|vision_config| {
validate_qwen35_vision_weight_inventory(&model_dir, vision_config).is_ok()
});
let state = match (
&config.vision_config,
token_metadata_present,
supported_weight_source,
) {
(Some(vision_config), true, true) => VisionState::Pending {
model_dir,
config: vision_config.clone(),
},
_ => VisionState::Unsupported,
};
let vision_supported =
matches!(&state, VisionState::Pending { .. } | VisionState::Loaded(_));
Self {
state,
vision_supported: Arc::new(AtomicBool::new(vision_supported)),
}
}
pub fn unsupported() -> Self {
Self {
state: VisionState::Unsupported,
vision_supported: Arc::new(AtomicBool::new(false)),
}
}
pub fn is_supported(&self) -> bool {
self.vision_supported.load(Ordering::Acquire)
}
fn shared_capability(&self) -> Arc<AtomicBool> {
self.vision_supported.clone()
}
fn get_or_load(
&mut self,
should_cancel: &mut dyn FnMut() -> bool,
) -> Result<VisionRuntimeLoad<'_>, String> {
if let VisionState::Pending { model_dir, config } = &self.state {
let model_dir = model_dir.clone();
let config = config.clone();
match load_qwen35_vision_weights_with_cancel(&model_dir, &config, should_cancel) {
Ok(Some(weights)) => self.state = VisionState::Loaded(weights),
Ok(None) => return Ok(VisionRuntimeLoad::Cancelled),
Err(err) => {
let message = format!("vision weights failed to load: {err}");
self.state = VisionState::Failed(message.clone());
self.vision_supported.store(false, Ordering::Release);
return Err(message);
}
}
}
match &self.state {
VisionState::Unsupported => Ok(VisionRuntimeLoad::Unsupported),
VisionState::Loaded(weights) => Ok(VisionRuntimeLoad::Ready(weights)),
VisionState::Failed(message) => Err(message.clone()),
VisionState::Pending { .. } => {
self.vision_supported.store(false, Ordering::Release);
Err("vision weights remained pending after a load attempt".to_string())
}
}
}
}
#[cfg(test)]
fn tokenize_text(tokenizer: &BpeTokenizer, text: &str) -> Vec<u32> {
let encoded = tokenizer.tokenize(text);
encoded.input_ids[..encoded.real_length].to_vec()
}
#[allow(clippy::too_many_arguments)]
fn build_vision_prompt_ids(
messages: &[ChatMessage],
image_message_index: usize,
tokenizer: &BpeTokenizer,
vision_start_token_id: u32,
vision_end_token_id: u32,
image_token_id: u32,
image_pad_count: usize,
) -> Result<Vec<u32>, WorkerFailure> {
let image_message = &messages[image_message_index];
let image = image_message
.image
.as_ref()
.ok_or_else(|| WorkerFailure::Failed("vision dispatch lost its image payload".into()))?;
if image.text_offset > image_message.content.len()
|| !image_message.content.is_char_boundary(image.text_offset)
{
return Err(WorkerFailure::Failed(
"normalized image text offset is not a UTF-8 boundary".into(),
));
}
let mut before = String::new();
for message in &messages[..image_message_index] {
push_chat_turn_open(&mut before, message.role.as_str());
before.push_str(&message.content);
push_chat_turn_close(&mut before);
}
push_chat_turn_open(&mut before, image_message.role.as_str());
before.push_str(&image_message.content[..image.text_offset]);
let mut after = String::new();
after.push_str(&image_message.content[image.text_offset..]);
push_chat_turn_close(&mut after);
for message in &messages[image_message_index + 1..] {
push_chat_turn_open(&mut after, message.role.as_str());
after.push_str(&message.content);
push_chat_turn_close(&mut after);
}
after.push_str("<|im_start|>assistant\n");
let mut inserted_ids = Vec::with_capacity(image_pad_count.saturating_add(2));
inserted_ids.push(vision_start_token_id);
inserted_ids.extend(std::iter::repeat_n(image_token_id, image_pad_count));
inserted_ids.push(vision_end_token_id);
let ids = tokenizer.tokenize_fragments_with_inserted_ids(&before, &inserted_ids, &after);
if ids.iter().filter(|&&id| id == image_token_id).count() != image_pad_count {
return Err(WorkerFailure::Rejected(ApiError::BadRequest {
message: "message text must not contain the checkpoint's reserved image token"
.to_string(),
code: "invalid_messages",
}));
}
Ok(ids)
}
enum VisionRequestBuild {
Ready {
request: Qwen35VisionRequest,
metal_dispatches: usize,
gemm_calls: usize,
},
Cancelled,
}
fn build_vision_request(
runtime: &mut VisionRuntime,
config: &Qwen35Config,
tokenizer: &BpeTokenizer,
messages: &[ChatMessage],
image_message_index: usize,
should_cancel: &mut dyn FnMut() -> bool,
window_preflight: impl FnOnce(usize) -> Result<(), ApiError>,
) -> Result<VisionRequestBuild, WorkerFailure> {
if should_cancel() {
return Ok(VisionRequestBuild::Cancelled);
}
let vision_config = config.vision_config.as_ref().ok_or_else(|| {
WorkerFailure::Rejected(ApiError::BadRequest {
message: "image input requires a vision-capable model".to_string(),
code: "vision_unsupported",
})
})?;
let image_token_id = config.image_token_id.ok_or_else(|| {
WorkerFailure::Failed("vision checkpoint has no image_token_id".to_string())
})?;
let vision_start_token_id = config.vision_start_token_id.ok_or_else(|| {
WorkerFailure::Failed("vision checkpoint has no vision_start_token_id".to_string())
})?;
let vision_end_token_id = config.vision_end_token_id.ok_or_else(|| {
WorkerFailure::Failed("vision checkpoint has no vision_end_token_id".to_string())
})?;
let image = messages[image_message_index]
.image
.as_ref()
.ok_or_else(|| WorkerFailure::Failed("vision dispatch lost its image payload".into()))?;
let (pixel_values, grid) = preprocess_qwen35_image_for_serve(&image.bytes, vision_config)
.map_err(|err| {
WorkerFailure::Rejected(ApiError::BadRequest {
message: format!("image preprocessing failed: {err}"),
code: "invalid_image",
})
})?;
if should_cancel() {
return Ok(VisionRequestBuild::Cancelled);
}
let merge_area = vision_config
.spatial_merge_size
.checked_mul(vision_config.spatial_merge_size)
.filter(|&area| area > 0)
.ok_or_else(|| WorkerFailure::Failed("vision spatial_merge_size is invalid".to_string()))?;
if !grid.num_patches().is_multiple_of(merge_area) {
return Err(WorkerFailure::Rejected(ApiError::BadRequest {
message: "image patch grid is incompatible with the checkpoint merge size".to_string(),
code: "invalid_image",
}));
}
let image_pad_count = grid.num_patches() / merge_area;
let input_ids = build_vision_prompt_ids(
messages,
image_message_index,
tokenizer,
vision_start_token_id,
vision_end_token_id,
image_token_id,
image_pad_count,
)?;
if should_cancel() {
return Ok(VisionRequestBuild::Cancelled);
}
window_preflight(input_ids.len()).map_err(WorkerFailure::Rejected)?;
if should_cancel() {
return Ok(VisionRequestBuild::Cancelled);
}
let weights = match runtime
.get_or_load(should_cancel)
.map_err(WorkerFailure::Failed)?
{
VisionRuntimeLoad::Ready(weights) => weights,
VisionRuntimeLoad::Unsupported => {
return Err(WorkerFailure::Rejected(ApiError::BadRequest {
message: "image input requires a vision-capable model".to_string(),
code: "vision_unsupported",
}));
}
VisionRuntimeLoad::Cancelled => return Ok(VisionRequestBuild::Cancelled),
};
if should_cancel() {
return Ok(VisionRequestBuild::Cancelled);
}
let Some(vit_output) = qwen35_vit_forward_metal_with_cancel(
weights,
vision_config,
&pixel_values,
grid,
should_cancel,
)
.map_err(|err| WorkerFailure::Failed(format!("vision forward failed: {err}")))?
else {
return Ok(VisionRequestBuild::Cancelled);
};
let Some(post_merger) = qwen35_merger_forward_with_cancel(
&weights.merger,
vision_config,
&vit_output.hidden_states,
should_cancel,
)
.map_err(|err| WorkerFailure::Failed(format!("vision merger failed: {err}")))?
else {
return Ok(VisionRequestBuild::Cancelled);
};
if should_cancel() {
return Ok(VisionRequestBuild::Cancelled);
}
Ok(VisionRequestBuild::Ready {
request: Qwen35VisionRequest {
input_ids,
image_grids: vec![grid],
post_merger_rows: post_merger,
image_token_id,
spatial_merge_size: vision_config.spatial_merge_size,
decoder_hidden_size: config.hidden_size,
},
metal_dispatches: vit_output.metal_dispatches,
gemm_calls: vit_output.gemm_calls,
})
}
fn cancelled_output() -> GenerateOutput {
GenerateOutput {
text: String::new(),
token_ids: Vec::new(),
prompt_tokens: 0,
generated_tokens: 0,
stopped: false,
stop_reason: Some(crate::StopReason::Interrupt),
token_logprobs: Vec::new(),
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum JobRoute {
Text,
Vision { message_index: usize },
}
fn classify_job(messages: &[ChatMessage]) -> Result<JobRoute, WorkerFailure> {
let mut image_positions = messages
.iter()
.enumerate()
.filter(|(_, message)| message.image.is_some());
let first = image_positions.next();
if image_positions.next().is_some() {
return Err(WorkerFailure::Rejected(ApiError::BadRequest {
message: "only one image is supported per request".to_string(),
code: "multiple_images_unsupported",
}));
}
Ok(match first {
Some((_message_index, message))
if message.role != crate::forward::metal_qwen35::ChatRole::User =>
{
return Err(WorkerFailure::Rejected(ApiError::BadRequest {
message: "image content is supported only on user messages".to_string(),
code: "invalid_image_role",
}));
}
Some((message_index, _)) => JobRoute::Vision { message_index },
None => JobRoute::Text,
})
}
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> {
Self::spawn_with_vision(loader, VisionRuntime::unsupported(), max_pending)
}
pub fn spawn_with_vision(
loader: impl FnOnce() -> Result<(MetalQwen35State, BpeTokenizer, WorkerMetadata), String>
+ Send
+ 'static,
mut vision_runtime: VisionRuntime,
max_pending: usize,
) -> Result<(MetalWorkerOwner, MetalWorkerClient, WorkerMetadata), StartupError> {
let vision_supported = vision_runtime.shared_capability();
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| {
if let JobRoute::Vision {
message_index: image_message_index,
} = classify_job(messages)?
{
if should_cancel() {
return Ok(cancelled_output());
}
let config = state.engine.config.clone();
let (request, metal_dispatches, gemm_calls) = match build_vision_request(
&mut vision_runtime,
&config,
&tokenizer,
messages,
image_message_index,
should_cancel,
|prompt_len| {
check_prompt_fits_window(
meta.context_window_policy,
meta.model_max_context,
prompt_len,
cfg,
)
},
)? {
VisionRequestBuild::Ready {
request,
metal_dispatches,
gemm_calls,
} => (request, metal_dispatches, gemm_calls),
VisionRequestBuild::Cancelled => return Ok(cancelled_output()),
};
if should_cancel() {
return Ok(cancelled_output());
}
eprintln!(
"[metal-worker] route=vision dispatch=multimodal \
metal_gemm_dispatches={metal_dispatches} \
metal_gemm_calls={gemm_calls}"
);
let output = state
.generate_multimodal_vision_with_cancel(
&request,
&tokenizer,
cfg,
should_cancel,
)
.map_err(WorkerFailure::from)?;
if !output.text.is_empty() {
let _ = on_token(&output.text, 0);
}
return Ok(output);
}
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,
vision_supported,
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,
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_capability(
max_pending,
context_window_policy,
model_max_context,
tokenizer,
false,
generate,
)
}
#[cfg(any(test, feature = "test-utils"))]
#[allow(clippy::type_complexity)]
pub fn spawn_fake_with_vision(
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_capability(
TEST_EFFECTIVELY_UNBOUNDED_CAP,
context_window_policy,
model_max_context,
tokenizer,
true,
generate,
)
}
#[cfg(any(test, feature = "test-utils"))]
#[allow(clippy::type_complexity)]
fn spawn_fake_with_capability(
max_pending: usize,
context_window_policy: ContextWindowPolicy,
model_max_context: usize,
tokenizer: BpeTokenizer,
vision_supported: bool,
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)),
Arc::new(AtomicBool::new(vision_supported)),
owner,
)
}
#[cfg(test)]
mod tests {
use super::*;
use std::cell::RefCell;
use std::collections::HashMap;
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)
}
fn tiny_tokenizer() -> BpeTokenizer {
BpeTokenizer::from_vocab_and_merges(
HashMap::from([
("a".to_string(), 0),
("b".to_string(), 1),
("e".to_string(), 2),
("f".to_string(), 3),
("o".to_string(), 4),
("r".to_string(), 5),
]),
Vec::new(),
)
.expect("tiny tokenizer must construct")
}
fn tiny_vision_config(out_hidden_size: usize) -> VisionModelConfig {
VisionModelConfig {
depth: 1,
hidden_size: 8,
num_heads: 2,
patch_size: 2,
spatial_merge_size: 2,
out_hidden_size,
temporal_patch_size: 1,
num_position_embeddings: 16,
in_channels: 3,
deepstack_visual_indexes: Vec::new(),
intermediate_size: None,
}
}
fn make_test_png(width: u32, height: u32) -> Vec<u8> {
let mut image = image::RgbImage::new(width, height);
for y in 0..height {
for x in 0..width {
let value = ((x + y) % 256) as u8;
image.put_pixel(x, y, image::Rgb([value, value, value]));
}
}
let mut bytes = Vec::new();
image
.write_to(
&mut std::io::Cursor::new(&mut bytes),
image::ImageFormat::Png,
)
.expect("test PNG encode");
bytes
}
#[test]
fn image_job_classification_enforces_role_and_multiplicity() {
assert_eq!(
classify_job(&[ChatMessage::user("hello")]).expect("text job"),
JobRoute::Text
);
let image = ChatMessage::user_with_image("beforeafter", vec![1, 2, 3], 6);
assert_eq!(
classify_job(&[ChatMessage::system("policy"), image.clone()]).expect("one image job"),
JobRoute::Vision { message_index: 1 }
);
let err = classify_job(&[image.clone(), image]).expect_err("two images must fail");
assert!(matches!(
err,
WorkerFailure::Rejected(ApiError::BadRequest {
code: "multiple_images_unsupported",
..
})
));
for role in [
crate::forward::metal_qwen35::ChatRole::System,
crate::forward::metal_qwen35::ChatRole::Assistant,
] {
let mut non_user_image = ChatMessage::user_with_image("beforeafter", vec![1, 2, 3], 6);
non_user_image.role = role;
let err = classify_job(&[non_user_image]).expect_err("non-user image role must fail");
assert!(matches!(
err,
WorkerFailure::Rejected(ApiError::BadRequest {
code: "invalid_image_role",
..
})
));
}
}
#[test]
fn vision_prompt_splices_tokens_at_original_content_part_position() {
let tokenizer = tiny_tokenizer();
let messages = vec![
ChatMessage::system("policy"),
ChatMessage::user_with_image("beforeafter", vec![1], "before".len()),
ChatMessage::assistant("prior"),
];
let actual = build_vision_prompt_ids(&messages, 1, &tokenizer, 90, 91, 92, 3)
.expect("vision prompt");
let before = "<|im_start|>system\npolicy<|im_end|>\n<|im_start|>user\nbefore";
let after =
"after<|im_end|>\n<|im_start|>assistant\nprior<|im_end|>\n<|im_start|>assistant\n";
let mut expected = tokenize_text(&tokenizer, before);
expected.extend([90, 92, 92, 92, 91]);
expected.extend(tokenize_text(&tokenizer, after));
assert_eq!(actual, expected);
}
#[test]
fn vision_prompt_tokenizes_the_complete_sequence_before_window_validation() {
let capped = tiny_tokenizer().with_max_seq_len(4);
let unbounded = capped.with_max_seq_len(usize::MAX);
let before = "before".repeat(64);
let messages = vec![ChatMessage::user_with_image(
format!("{before}after"),
vec![1],
before.len(),
)];
let actual = build_vision_prompt_ids(&messages, 0, &capped, 90, 91, 92, 3)
.expect("capped tokenizer must not truncate prompt fragments");
let expected = build_vision_prompt_ids(&messages, 0, &unbounded, 90, 91, 92, 3)
.expect("unbounded tokenizer");
assert_eq!(actual, expected);
assert!(actual.len() > 4);
}
fn vision_build_config() -> Qwen35Config {
let mut config = Qwen35Config::qwen35_0_8b();
config.vision_config = Some(tiny_vision_config(config.hidden_size));
config.image_token_id = Some(90);
config.vision_start_token_id = Some(91);
config.vision_end_token_id = Some(92);
config
}
#[test]
fn vision_request_build_cancels_before_image_preprocessing() {
let mut runtime = VisionRuntime::unsupported();
let config = vision_build_config();
let messages = vec![ChatMessage::user_with_image(
"beforeafter",
b"not an image".to_vec(),
"before".len(),
)];
let mut polls = 0;
let result = build_vision_request(
&mut runtime,
&config,
&tiny_tokenizer(),
&messages,
0,
&mut || {
polls += 1;
true
},
|_| panic!("window preflight must not run after cancellation"),
)
.expect("cancellation is not a worker failure");
assert!(matches!(result, VisionRequestBuild::Cancelled));
assert_eq!(polls, 1);
}
#[test]
fn vision_request_build_cancels_after_preprocess_and_prompt_before_window_check() {
let mut runtime = VisionRuntime::unsupported();
let config = vision_build_config();
let messages = vec![ChatMessage::user_with_image(
"beforeafter",
make_test_png(8, 8),
"before".len(),
)];
let mut polls = 0;
let mut window_checked = false;
let result = build_vision_request(
&mut runtime,
&config,
&tiny_tokenizer(),
&messages,
0,
&mut || {
polls += 1;
polls == 3
},
|_| {
window_checked = true;
Ok(())
},
)
.expect("cancellation is not a worker failure");
assert!(matches!(result, VisionRequestBuild::Cancelled));
assert_eq!(polls, 3);
assert!(!window_checked);
}
#[test]
fn vision_request_build_cancels_after_window_check_before_lazy_load() {
let mut runtime = VisionRuntime::unsupported();
let config = vision_build_config();
let messages = vec![ChatMessage::user_with_image(
"beforeafter",
make_test_png(8, 8),
"before".len(),
)];
let mut polls = 0;
let mut window_checked = false;
let result = build_vision_request(
&mut runtime,
&config,
&tiny_tokenizer(),
&messages,
0,
&mut || {
polls += 1;
polls == 4
},
|_| {
window_checked = true;
Ok(())
},
)
.expect("cancellation is not a worker failure");
assert!(matches!(result, VisionRequestBuild::Cancelled));
assert_eq!(polls, 4);
assert!(window_checked);
}
#[test]
fn vision_runtime_capability_requires_config_token_metadata_and_weight_source() {
let temp = tempfile::tempdir().expect("tempdir");
std::fs::write(temp.path().join("quantize_index.json"), "[]")
.expect("weight-source marker");
let mut config = Qwen35Config::qwen35_0_8b();
assert!(
!VisionRuntime::from_model_config(temp.path().to_path_buf(), &config).is_supported()
);
config.vision_config = Some(tiny_vision_config(config.hidden_size));
config.image_token_id = Some(10);
config.vision_start_token_id = Some(11);
config.vision_end_token_id = Some(12);
assert!(
!VisionRuntime::from_model_config(temp.path().to_path_buf(), &config).is_supported(),
"an empty manifest must not advertise vision capability"
);
let mut names = vec![
"model.visual.patch_embed.proj.weight".to_string(),
"model.visual.patch_embed.proj.bias".to_string(),
"model.visual.pos_embed.weight".to_string(),
"model.visual.merger.linear_fc1.weight".to_string(),
"model.visual.merger.linear_fc1.bias".to_string(),
"model.visual.merger.linear_fc2.weight".to_string(),
"model.visual.merger.linear_fc2.bias".to_string(),
"model.visual.merger.norm.weight".to_string(),
"model.visual.merger.norm.bias".to_string(),
];
for suffix in [
"attn.qkv.weight",
"attn.qkv.bias",
"attn.proj.weight",
"attn.proj.bias",
"mlp.linear_fc1.weight",
"mlp.linear_fc1.bias",
"mlp.linear_fc2.weight",
"mlp.linear_fc2.bias",
"norm1.weight",
"norm1.bias",
"norm2.weight",
"norm2.bias",
] {
names.push(format!("model.visual.blocks.0.{suffix}"));
}
std::fs::write(temp.path().join("visual.bin"), b"inventory preflight")
.expect("visual tensor marker");
let entries: Vec<_> = names
.iter()
.map(|name| serde_json::json!({"name": name, "file": "visual.bin"}))
.collect();
std::fs::write(
temp.path().join("quantize_index.json"),
serde_json::to_vec(&entries).expect("manifest fixture"),
)
.expect("complete vision manifest");
let mut invalid_tokens = config.clone();
invalid_tokens.image_token_id = Some(config.vocab_size as u32);
assert!(
!VisionRuntime::from_model_config(temp.path().to_path_buf(), &invalid_tokens)
.is_supported(),
"out-of-vocabulary image metadata must not advertise capability"
);
invalid_tokens.image_token_id = invalid_tokens.vision_start_token_id;
assert!(
!VisionRuntime::from_model_config(temp.path().to_path_buf(), &invalid_tokens)
.is_supported(),
"aliased vision token metadata must not advertise capability"
);
let mut runtime = VisionRuntime::from_model_config(temp.path().to_path_buf(), &config);
assert!(runtime.is_supported());
let (job_tx, _job_rx) = mpsc::unbounded_channel();
let client = MetalWorkerClient::with_owner(
job_tx,
Arc::new(Semaphore::new(1)),
runtime.shared_capability(),
MetalWorkerOwner::unattached_for_test(),
);
assert!(client.supports_vision());
let cancelled = runtime
.get_or_load(&mut || true)
.expect("cancellation is not a lazy-load failure");
assert!(matches!(cancelled, VisionRuntimeLoad::Cancelled));
assert!(runtime.is_supported());
assert!(
client.supports_vision(),
"cancellation must preserve Pending capability for a later retry"
);
let mut never_cancel = || false;
let first_error = runtime
.get_or_load(&mut never_cancel)
.expect_err("junk tensor payload must fail its first lazy load");
assert!(first_error.contains("vision weights failed to load"));
assert!(!runtime.is_supported());
assert!(
!client.supports_vision(),
"terminal lazy-load failure must revoke the shared client capability"
);
std::fs::remove_file(temp.path().join("visual.bin"))
.expect("remove source after the first attempt");
let second_error = runtime
.get_or_load(&mut never_cancel)
.expect_err("terminal failure must be returned, not retried");
assert_eq!(second_error, first_error);
assert!(
!VisionRuntime::from_model_config(PathBuf::from("/unused"), &config).is_supported(),
"config metadata alone must not advertise vision without a supported weight source"
);
config.vision_end_token_id = None;
assert!(
!VisionRuntime::from_model_config(temp.path().to_path_buf(), &config).is_supported()
);
}
#[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)),
Arc::new(AtomicBool::new(false)),
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)),
Arc::new(AtomicBool::new(false)),
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");
}
}