diff --git a/crates/inference/examples/bench_serve_prepare.rs b/crates/inference/examples/bench_serve_prepare.rs index f1c86c6629..3c37ddc9cc 100644 --- a/crates/inference/examples/bench_serve_prepare.rs +++ b/crates/inference/examples/bench_serve_prepare.rs @@ -15,13 +15,15 @@ //! `Qwen35Model::generate_streaming_with_cancel`. This route //! has no admission gate. //! lattice-metal `lattice serve` on a Q4 checkpoint (Metal worker). -//! Handler: the same `prepare_chat_request` and mapping as -//! `cpu`, with the 4096-token context the binary uses. +//! Handler: the worker factory's `PreparationHandle` +//! (`prepare_lattice`, the same `prepare_chat_request` bound +//! to the worker's tokenizer and 4096-token context), then +//! `lattice_generate_config`. //! lattice-serve `lattice_serve` (Metal worker, Q4 or safetensors). -//! Handler: `serve::contract::normalize_request` with the -//! `lattice_serve` profile, `serve::prepare::build_cfg`, then -//! `serve::into_engine_chat_messages`. No render or tokenize -//! happens in this handler. +//! Handler: the handle's `normalize_standalone` with the +//! `lattice_serve` defaults, `standalone_generate_config`, +//! then `serve::into_engine_chat_messages`. No render or +//! tokenize happens in this handler. //! gemma-cpu Gemma 4 E2B text on a safetensors checkpoint, CPU. //! Preparation: `serve::prepare::prepare_gemma_chat_request` //! (validate, the Gemma prompt adapter's defaults, render @@ -273,24 +275,24 @@ fn lattice_prepare( )) } -metal_only! { - /// The `lattice_serve` handler's preparation: normalize, `build_cfg`, engine - /// messages. - fn lattice_serve_prepare( - body: &[u8], - model_max_context: usize, - ) -> Result, String> { - let req = parse_body(body)?; - Ok(normalize_request( - &req, - GenerationDefaults::standard(LATTICE_SERVE_DEFAULT_MAX_TOKENS), - ServeProfile::lattice_serve(MODEL_ID, model_max_context).with_vision_support(false), - ) - .and_then(|validated| { - let cfg = build_cfg(&validated); - into_engine_chat_messages(validated.messages).map(|messages| (messages, cfg)) - })) - } +/// The `lattice_serve` handler's preparation before the worker factory: +/// normalize, `build_cfg`, engine messages. Kept for the preparation tests +/// below; the Metal route now prepares through the worker's handle. +#[cfg_attr(not(test), allow(dead_code))] +fn lattice_serve_prepare( + body: &[u8], + model_max_context: usize, +) -> Result, String> { + let req = parse_body(body)?; + Ok(normalize_request( + &req, + GenerationDefaults::standard(LATTICE_SERVE_DEFAULT_MAX_TOKENS), + ServeProfile::lattice_serve(MODEL_ID, model_max_context).with_vision_support(false), + ) + .and_then(|validated| { + let cfg = build_cfg(&validated); + into_engine_chat_messages(validated.messages).map(|messages| (messages, cfg)) + })) } /// Gemma E2B text preparation, with the `lattice serve` defaults. @@ -655,6 +657,7 @@ fn run_metal( use lattice_inference::serve::metal_worker::{ ContextWindowPolicy, MetalWorker, VisionRuntime, WorkerEvent, WorkerMetadata, }; + use lattice_inference::serving_factory::ServingFactory; use lattice_inference::tokenizer::bpe::BpeTokenizer; use std::time::Instant; @@ -674,46 +677,48 @@ fn run_metal( }; let load = Instant::now(); - let (owner, client, meta) = MetalWorker::spawn_with_vision( - move || { - let tokenizer_path = loader_dir.join("tokenizer.json"); - // `lattice serve` allocates its fixed 4096-token context; - // `lattice_serve` requests the checkpoint's configured window. - let state = match (route, format) { - (Route::LatticeServe, ModelFormat::Safetensors) => { - let model = Qwen35Model::from_safetensors(&loader_dir) - .map_err(|e| format!("safetensors load failed: {e}"))?; - let cfg = model.config().clone(); - let context = cfg.max_position_embeddings; - MetalQwen35State::new(model.weights(), &cfg, context) - .map_err(|e| format!("Metal init failed: {e}"))? - } - _ => { - let cfg = Qwen35Config::from_model_dir(&loader_dir) - .map_err(|e| format!("config.json load failed: {e}"))?; - let context = match route { - Route::LatticeServe => cfg.max_position_embeddings, - _ => LATTICE_METAL_MAX_CONTEXT, - }; - MetalQwen35State::from_q4_dir(&loader_dir, &tokenizer_path, &cfg, context) - .map_err(|e| format!("Q4 model load failed: {e}"))? - } - }; - let model_max_context = match route { - Route::LatticeServe => state.max_context(), - _ => LATTICE_METAL_MAX_CONTEXT, - }; - Ok(( - state, - worker_tokenizer, - WorkerMetadata { - format: format!("{format:?}"), - model_max_context, - context_window_policy: policy, - }, - )) - }, - VisionRuntime::unsupported(), + let (owner, client, meta, preparation) = MetalWorker::spawn_with_vision( + ServingFactory::qwen_metal( + move || { + let tokenizer_path = loader_dir.join("tokenizer.json"); + // `lattice serve` allocates its fixed 4096-token context; + // `lattice_serve` requests the checkpoint's configured window. + let state = match (route, format) { + (Route::LatticeServe, ModelFormat::Safetensors) => { + let model = Qwen35Model::from_safetensors(&loader_dir) + .map_err(|e| format!("safetensors load failed: {e}"))?; + let cfg = model.config().clone(); + let context = cfg.max_position_embeddings; + MetalQwen35State::new(model.weights(), &cfg, context) + .map_err(|e| format!("Metal init failed: {e}"))? + } + _ => { + let cfg = Qwen35Config::from_model_dir(&loader_dir) + .map_err(|e| format!("config.json load failed: {e}"))?; + let context = match route { + Route::LatticeServe => cfg.max_position_embeddings, + _ => LATTICE_METAL_MAX_CONTEXT, + }; + MetalQwen35State::from_q4_dir(&loader_dir, &tokenizer_path, &cfg, context) + .map_err(|e| format!("Q4 model load failed: {e}"))? + } + }; + let model_max_context = match route { + Route::LatticeServe => state.max_context(), + _ => LATTICE_METAL_MAX_CONTEXT, + }; + Ok(( + state, + worker_tokenizer, + WorkerMetadata { + format: format!("{format:?}"), + model_max_context, + context_window_policy: policy, + }, + )) + }, + VisionRuntime::unsupported(), + ), 1, ResidencyLimits::default(), ) @@ -723,25 +728,39 @@ fn run_metal( println!("LOAD route={} load_ms={:.3}", route.name(), ms(load)); let prepare = |body: &[u8]| -> Result, String> { + let req = parse_body(body)?; match route { - Route::LatticeServe => lattice_serve_prepare(body, model_max_context), - _ => Ok(lattice_prepare( - body, - |p| tokenizer.tokenize(p).real_length, - LATTICE_METAL_MAX_CONTEXT, - )? - .map(|prepared| { - let cfg = lattice_gen_cfg( - prepared.max_tokens, - prepared.temperature, - prepared.top_p, - prepared.seed, - prepared.stop_strings.clone(), - prepared.reasoning_budget, - prepared.logprobs, - ); - (prepared.messages, cfg) - })), + Route::LatticeServe => Ok(preparation + .normalize_standalone( + &req, + GenerationDefaults::standard(LATTICE_SERVE_DEFAULT_MAX_TOKENS), + MODEL_ID, + false, + ) + .and_then(|validated| { + let cfg = preparation.standalone_generate_config(&validated); + into_engine_chat_messages(validated.messages).map(|messages| (messages, cfg)) + })), + _ => Ok(preparation + .prepare_lattice( + &req, + MODEL_ID, + LATTICE_DEFAULT_MAX_TOKENS, + LATTICE_MAX_TOKENS_CAP, + false, + ) + .map(|prepared| { + let cfg = preparation.lattice_generate_config( + prepared.max_tokens, + prepared.temperature, + prepared.top_p, + prepared.seed, + prepared.stop_strings.clone(), + prepared.reasoning_budget, + prepared.logprobs, + ); + (prepared.messages, cfg) + })), } }; diff --git a/crates/inference/src/bin/lattice/serve.rs b/crates/inference/src/bin/lattice/serve.rs index 47123877f6..f96ac23498 100644 --- a/crates/inference/src/bin/lattice/serve.rs +++ b/crates/inference/src/bin/lattice/serve.rs @@ -215,6 +215,72 @@ pub enum ModelBackend { } impl ModelBackend { + #[allow(clippy::too_many_arguments)] + fn generation_config( + &self, + max_tokens: usize, + temperature: f32, + top_p: f32, + seed: Option, + stop_strings: Vec, + reasoning_budget: Option, + logprobs: Option, + ) -> lattice_inference::GenerateConfig { + #[cfg(feature = "metal-gpu")] + if let ModelBackend::Metal { handle, .. } = self + && let Some(preparation) = handle.client.preparation() + { + return preparation.lattice_generate_config( + max_tokens, + temperature, + top_p, + seed, + stop_strings, + reasoning_budget, + logprobs, + ); + } + lattice_gen_cfg( + max_tokens, + temperature, + top_p, + seed, + stop_strings, + reasoning_budget, + logprobs, + ) + } + + fn prepare_chat_request( + &self, + req: &ChatCompletionRequest, + model_id: &str, + default_max_tokens: usize, + max_tokens_cap: usize, + ) -> Result { + #[cfg(feature = "metal-gpu")] + if let ModelBackend::Metal { handle, .. } = self + && let Some(preparation) = handle.client.preparation() + { + return preparation.prepare_lattice( + req, + model_id, + default_max_tokens, + max_tokens_cap, + self.supports_vision(), + ); + } + prepare_chat_request( + req, + model_id, + default_max_tokens, + max_tokens_cap, + self.supports_vision(), + |prompt| self.tokenize_len(prompt), + || self.max_context(), + ) + } + pub fn tokenize_len(&self, text: &str) -> usize { match self { ModelBackend::Cpu(m) => m.tokenizer().tokenize(text).pre_truncation_len, @@ -275,6 +341,7 @@ impl ModelBackend { use lattice_inference::serve::metal_worker::{ ContextWindowPolicy, MetalWorker, StartupError, VisionRuntime, WorkerMetadata, }; + use lattice_inference::serving_factory::ServingFactory; let tokenizer_path = tokenizer_dir .as_deref() @@ -318,28 +385,30 @@ impl ModelBackend { } let model_dir_for_loader = model_dir.clone(); let tokenizer_path_for_loader = tokenizer_path.clone(); - let (owner, client, _meta) = MetalWorker::spawn_with_vision( - move || { - let cfg = crate::chat::load_q4_config(&model_dir_for_loader)?; - let state = - lattice_inference::forward::metal_qwen35::MetalQwen35State::from_q4_dir( - &model_dir_for_loader, - &tokenizer_path_for_loader, - &cfg, - max_context, - ) - .map_err(|e| format!("Q4 model load failed: {e}"))?; - Ok(( - state, - tokenizer_for_worker, - WorkerMetadata { - format: "q4".to_string(), - model_max_context: max_context, - context_window_policy: ContextWindowPolicy::PromptAndMaxTokens, - }, - )) - }, - vision_runtime, + let (owner, client, _meta, _preparation) = MetalWorker::spawn_with_vision( + ServingFactory::qwen_metal( + move || { + let cfg = crate::chat::load_q4_config(&model_dir_for_loader)?; + let state = + lattice_inference::forward::metal_qwen35::MetalQwen35State::from_q4_dir( + &model_dir_for_loader, + &tokenizer_path_for_loader, + &cfg, + max_context, + ) + .map_err(|e| format!("Q4 model load failed: {e}"))?; + Ok(( + state, + tokenizer_for_worker, + WorkerMetadata { + format: "q4".to_string(), + model_max_context: max_context, + context_window_policy: ContextWindowPolicy::PromptAndMaxTokens, + }, + )) + }, + vision_runtime, + ), max_pending, residency_limits, ) @@ -1118,17 +1187,14 @@ async fn chat_completions_with_request( reasoning_budget, seed, stream, - } = prepare_chat_request( + } = state.model.prepare_chat_request( &req, &state.model_id, state.default_max_tokens, state.max_tokens_cap, - state.model.supports_vision(), - |p| state.model.tokenize_len(p), - || state.model.max_context(), )?; - let gen_cfg = lattice_gen_cfg( + let gen_cfg = state.model.generation_config( max_tokens, temperature, top_p, diff --git a/crates/inference/src/bin/lattice_serve.rs b/crates/inference/src/bin/lattice_serve.rs index 26ef8da530..e9203e7c20 100644 --- a/crates/inference/src/bin/lattice_serve.rs +++ b/crates/inference/src/bin/lattice_serve.rs @@ -123,6 +123,7 @@ mod imp { }; use lattice_inference::serve::metrics::ServeMetrics; use lattice_inference::serve::prepare::build_cfg; + use lattice_inference::serving_factory::ServingFactory; use lattice_inference::tokenizer::bpe::BpeTokenizer; use lattice_inference::{BertModel, BertPooling}; use serde_json::{Value, json}; @@ -1809,10 +1810,9 @@ mod imp { // all now live in `lattice_inference::serve::metal_worker` (issue #832), // replacing this binary's previous private `spawn_worker`/ // `run_worker_loop`/`check_prompt_fits_window`/`enforce_prompt_window`. - // `load_model`/`LoadedModel` below are unchanged: `run()` wraps - // `load_model` in a loader closure passed to `MetalWorker::spawn`, which - // runs it ON the worker thread it creates (the `!Send` `MetalQwen35State` - // this function returns never crosses a thread boundary). + // `load_model`/`LoadedModel` below are unchanged: `run()` gives the + // factory a loader closure, which it invokes on the worker thread + // (the `!Send` `MetalQwen35State` never crosses a thread boundary). /// Everything the worker thread needs after a successful model load, /// including the actual KV context (#551) so request clamping never @@ -2106,12 +2106,21 @@ mod imp { ); return err.into_response(); } - let validated = match normalize_request( - &req, - s.defaults, - ServeProfile::lattice_serve(s.model_id.as_ref(), s.model_max_context) - .with_vision_support(s.jobs.supports_vision()), - ) { + let normalized = match s.jobs.preparation() { + Some(preparation) => preparation.normalize_standalone( + &req, + s.defaults, + s.model_id.as_ref(), + s.jobs.supports_vision(), + ), + None => normalize_request( + &req, + s.defaults, + ServeProfile::lattice_serve(s.model_id.as_ref(), s.model_max_context) + .with_vision_support(s.jobs.supports_vision()), + ), + }; + let validated = match normalized { Ok(validated) => validated, Err(err) => { emit_serve_event( @@ -2155,7 +2164,10 @@ mod imp { return err_response(err.status(), err.message(), err.code()); } }; - let mut cfg = build_cfg(&validated); + let mut cfg = match s.jobs.preparation() { + Some(preparation) => preparation.standalone_generate_config(&validated), + None => build_cfg(&validated), + }; // Wire the compiled grammar into the worker config (design note // ยง"End-to-end execution", step 4) and force `enable_thinking` off // for strict requests regardless of server defaults: @@ -3608,25 +3620,29 @@ mod imp { model_max_context, .. }, + _preparation, ) = match MetalWorker::spawn_with_vision( - move || { - let LoadedModel { - metal, - tokenizer, - format, - model_max_context, - } = load_model(&model_dir_for_loader, &tokenizer_path, format)?; - Ok(( - metal, - tokenizer, - WorkerMetadata { + ServingFactory::qwen_metal( + move || { + let LoadedModel { + metal, + tokenizer, format, model_max_context, - context_window_policy: ContextWindowPolicy::PromptAndDecodeWithDelimiter, - }, - )) - }, - vision_runtime, + } = load_model(&model_dir_for_loader, &tokenizer_path, format)?; + Ok(( + metal, + tokenizer, + WorkerMetadata { + format, + model_max_context, + context_window_policy: + ContextWindowPolicy::PromptAndDecodeWithDelimiter, + }, + )) + }, + vision_runtime, + ), max_pending, residency_limits, ) { diff --git a/crates/inference/src/lib.rs b/crates/inference/src/lib.rs index 315c06c423..415ed09aa3 100644 --- a/crates/inference/src/lib.rs +++ b/crates/inference/src/lib.rs @@ -128,6 +128,11 @@ pub mod sampling; /// Requires the `serve` feature (axum/tokio/futures). #[cfg(feature = "serve")] pub mod serve; +/// Builds the Metal serving worker's runtime and request preparation for a +/// loaded model. Not a stable API. +#[cfg(all(target_os = "macos", feature = "metal-gpu", feature = "serve"))] +#[doc(hidden)] +pub mod serving_factory; /// N-gram prompt lookup speculative decoding. See [`sampling`] and [`model`]. pub mod speculative; /// Generation stop reason taxonomy; see [`StopReason`] and [`model`]. diff --git a/crates/inference/src/model/mod.rs b/crates/inference/src/model/mod.rs index 44cd703319..55e47eb7e2 100644 --- a/crates/inference/src/model/mod.rs +++ b/crates/inference/src/model/mod.rs @@ -15,6 +15,8 @@ pub mod paddleocr_vl; pub mod qwen; pub mod qwen35; pub mod qwen35_config; +#[cfg(all(target_os = "macos", feature = "metal-gpu", feature = "serve"))] +pub(crate) mod serving_runtime; // Re-export everything from bert (was top-level `model` module) pub use self::bert::*; diff --git a/crates/inference/src/model/serving_runtime/mod.rs b/crates/inference/src/model/serving_runtime/mod.rs new file mode 100644 index 0000000000..4b1c6a302c --- /dev/null +++ b/crates/inference/src/model/serving_runtime/mod.rs @@ -0,0 +1,54 @@ +//! Worker-local execution for model serving. + +mod qwen_metal; + +pub(crate) use qwen_metal::QwenMetalRuntime; + +use crate::forward::metal_qwen35::ChatMessage; +use crate::generation::{GenerateConfig, GenerateOutput}; +use crate::serve::lora::{ + AdapterControlError, AdapterControlResult, AdapterIndex, LoraSelection, ResidencyLimits, +}; +use crate::serve::metal_worker::{AdapterCommand, WorkerFailure, WorkerMetadata}; +use crate::serve::prepare::PreparationHandle; +use std::sync::atomic::AtomicBool; +use std::sync::{Arc, RwLock}; + +/// The execution state stays on the thread that built it. +pub(crate) trait ServingRuntime { + fn generate( + &mut self, + messages: &[ChatMessage], + cfg: &GenerateConfig, + lora: &[LoraSelection], + on_token: &mut dyn FnMut(&str, u32) -> bool, + should_cancel: &mut dyn FnMut() -> bool, + ) -> Result; + + fn control( + &mut self, + command: AdapterCommand, + ) -> Result; + + fn vision_supported(&self) -> Arc; +} + +/// A one-shot builder moved into the worker before any Metal state exists. +pub(crate) trait RuntimeFactory: Send + 'static { + fn build( + self: Box, + index: Arc>, + limits: ResidencyLimits, + ) -> Result<(Box, WorkerMetadata, PreparationHandle), String>; +} + +trait AmbiguousIfSend { + fn some_item() {} +} +impl AmbiguousIfSend<()> for T {} +#[allow(dead_code)] +struct IsSend; +impl AmbiguousIfSend for T {} +const _: fn() = || { + let _ = >::some_item; +}; diff --git a/crates/inference/src/model/serving_runtime/qwen_metal.rs b/crates/inference/src/model/serving_runtime/qwen_metal.rs new file mode 100644 index 0000000000..e7b8007c46 --- /dev/null +++ b/crates/inference/src/model/serving_runtime/qwen_metal.rs @@ -0,0 +1,199 @@ +//! Qwen Metal state and adapter residency confined to one serving worker. + +use super::ServingRuntime; +use crate::forward::metal_qwen35::{ChatMessage, MetalQwen35State}; +use crate::generation::{GenerateConfig, GenerateOutput}; +use crate::kv_cache::CrossTurnSlotId; +use crate::serve::lora::{ + AdapterControlError, AdapterControlResult, AdapterIndex, LoraSelection, ResidencyLimits, +}; +use crate::serve::lora_registry::ResidencyRegistry; +use crate::serve::metal_worker::{ + AdapterCommand, JobRoute, VisionRequestBuild, VisionRuntime, WorkerFailure, WorkerMetadata, + build_vision_request, cancelled_output, check_prompt_fits_window, classify_job, + render_text_prompt_within_window, +}; +use crate::tokenizer::bpe::BpeTokenizer; +use std::sync::atomic::AtomicBool; +use std::sync::{Arc, RwLock}; + +pub(crate) struct QwenMetalRuntime { + state: MetalQwen35State, + tokenizer: Arc, + vision: VisionRuntime, + registry: ResidencyRegistry, + metadata: WorkerMetadata, +} + +impl QwenMetalRuntime { + pub(crate) fn new( + state: MetalQwen35State, + tokenizer: Arc, + vision: VisionRuntime, + metadata: WorkerMetadata, + index: Arc>, + limits: ResidencyLimits, + ) -> Self { + Self { + state, + tokenizer, + vision, + registry: ResidencyRegistry::new(index, limits), + metadata, + } + } +} + +impl ServingRuntime for QwenMetalRuntime { + fn generate( + &mut self, + messages: &[ChatMessage], + cfg: &GenerateConfig, + lora: &[LoraSelection], + on_token: &mut dyn FnMut(&str, u32) -> bool, + should_cancel: &mut dyn FnMut() -> bool, + ) -> Result { + let state = &mut self.state; + let tokenizer = self.tokenizer.as_ref(); + let vision_runtime = &mut self.vision; + let registry = &mut self.registry; + let meta = &self.metadata; + 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( + 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}" + ); + registry + .apply(lora, state) + .map_err(WorkerFailure::Rejected)?; + 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); + } + + // Render the ChatML prompt exactly once (#828/#832: the prior + // `lattice_serve.rs` path rendered it a second time inside its own + // window preflight); reused for both the window check and generation. + let (prompt, _prompt_len) = render_text_prompt_within_window( + tokenizer, + messages, + meta.context_window_policy, + meta.model_max_context, + cfg, + ) + .map_err(WorkerFailure::Rejected)?; + + // Cache-aware + cancellation-aware call (#462/#744): + // reuses the previous turn's shared token prefix + // instead of a full re-prefill on every request, and + // observes client disconnect before prefill, + // immediately after prefill, and at the top of every + // decode iteration. This worker thread owns one + // `MetalQwen35State` for the whole process lifetime, so + // `CrossTurnSlotId::DEFAULT` is the only slot that + // exists; the planner re-verifies the retained prefix + // against this request's prompt on every call and + // falls back to `PrefixReuseMode::FullRefill` whenever + // they diverge, so correctness never depends on + // distinguishing clients. + // + // DEPLOYMENT ASSUMPTION, stated because it is currently + // true only by the accident that no multi-tenant consumer + // exists: this path assumes a single tenant, or clients + // that mutually trust one another. Reuse-versus-refill is + // externally visible as latency, so while no request can + // read another's content, a client CAN observe that some + // other request recently shared a prefix with its own. + // A shared inference endpoint serving mutually distrusting + // clients must key the slot per tenant via + // `CrossTurnSlotId::new`, not inherit `DEFAULT`. + if should_cancel() { + return Ok(cancelled_output()); + } + registry + .apply(lora, state) + .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) + } + + fn control( + &mut self, + command: AdapterCommand, + ) -> Result { + match command { + AdapterCommand::Load { + name, + path, + layers, + descriptor, + } => { + let id = self.registry.load(name, path, layers, *descriptor)?; + self.registry.metadata(id).map(AdapterControlResult::Loaded) + } + AdapterCommand::Unload { id } => self + .registry + .unload(id, &mut self.state) + .map(AdapterControlResult::Unloaded), + } + } + + fn vision_supported(&self) -> Arc { + self.vision.shared_capability() + } +} diff --git a/crates/inference/src/serve/lora_registry.rs b/crates/inference/src/serve/lora_registry.rs index 48e193c06d..3042c35b78 100644 --- a/crates/inference/src/serve/lora_registry.rs +++ b/crates/inference/src/serve/lora_registry.rs @@ -10,7 +10,7 @@ use lattice_fann::lora::LoraDescriptor; use std::collections::HashMap; use std::sync::{Arc, RwLock}; -pub(super) trait AdapterSlot { +pub(crate) trait AdapterSlot { fn load(&mut self, layers: Vec) -> Result<(), String>; fn unload(&mut self); } @@ -33,7 +33,7 @@ struct ResidentAdapter { payload_bytes: usize, } -pub(super) struct ResidencyRegistry { +pub(crate) struct ResidencyRegistry { residents: HashMap, identities: HashMap<(String, String), u32>, resident_bytes: usize, @@ -59,7 +59,7 @@ pub(super) struct ResidencyRegistry { } impl ResidencyRegistry { - pub(super) fn new(index: Arc>, limits: ResidencyLimits) -> Self { + pub(crate) fn new(index: Arc>, limits: ResidencyLimits) -> Self { Self { residents: HashMap::new(), identities: HashMap::new(), @@ -118,7 +118,7 @@ impl ResidencyRegistry { } } - pub(super) fn load( + pub(crate) fn load( &mut self, name: String, path: String, @@ -205,7 +205,7 @@ impl ResidencyRegistry { Ok(id) } - pub(super) fn unload( + pub(crate) fn unload( &mut self, id: u32, slot: &mut impl AdapterSlot, @@ -228,14 +228,14 @@ impl ResidencyRegistry { Ok(id) } - pub(super) fn metadata(&self, id: u32) -> Result { + pub(crate) fn metadata(&self, id: u32) -> Result { self.residents .get(&id) .map(|adapter| adapter.metadata.clone()) .ok_or(AdapterControlError::NotFound(id)) } - pub(super) fn apply( + pub(crate) fn apply( &mut self, selection: &[LoraSelection], slot: &mut impl AdapterSlot, diff --git a/crates/inference/src/serve/metal_worker.rs b/crates/inference/src/serve/metal_worker.rs index e466bc44a2..3225615264 100644 --- a/crates/inference/src/serve/metal_worker.rs +++ b/crates/inference/src/serve/metal_worker.rs @@ -11,7 +11,7 @@ //! dedicated OS thread (the Metal state can never cross a thread boundary), //! serve `Job`s FIFO from an unbounded channel, check a per-job //! disconnect-cancellation signal before paying for any prefill work, reuse -//! the single process-wide [`CrossTurnSlotId::DEFAULT`] cache slot, and +//! the single process-wide [`crate::kv_cache::CrossTurnSlotId::DEFAULT`] cache slot, and //! stream token deltas back to the HTTP handler. Only comments -- not //! shared code -- kept the two copies in sync, and they had already drifted: //! on dequeue-time cancellation, `lattice.rs`'s worker sent an empty @@ -36,10 +36,10 @@ //! //! [`run_worker_loop`] -- the FIFO/cancellation/terminal-event state //! machine -- is generic over an injected `generate` closure, exactly like -//! `lattice_serve.rs`'s pre-existing `run_worker_loop` was. [`MetalWorker::spawn`] -//! wires a REAL closure (calling `MetalQwen35State::generate_streaming_with_prefix_cache_and_cancel`) -//! into it for production; this module's own tests inject a fake generator -//! instead, so the state machine is fully covered without a Metal device. +//! `lattice_serve.rs`'s pre-existing `run_worker_loop` was. Production +//! supplies callbacks from a model-owned runtime; this module's own tests +//! inject a fake generator instead, so the state machine is covered without +//! a Metal device. //! `MetalWorker::spawn`'s `loader` failure path is also GPU-free: a loader //! that returns `Err` before ever constructing a `MetalQwen35State` //! typechecks and runs with no device involved. The real `spawn` -> real @@ -52,15 +52,16 @@ use super::lora::{ AdapterControlError, AdapterControlResult, AdapterIndex, LoraSelection, ResidencyLimits, }; -use super::lora_registry::ResidencyRegistry; +use crate::forward::metal_qwen35::format_chat_template; use crate::forward::metal_qwen35::{ - ChatMessage, LoraLayerData, MetalQwen35State, format_chat_template, push_chat_generation_open, - push_chat_turn_close, push_chat_turn_open, + ChatMessage, LoraLayerData, MetalQwen35State, push_chat_generation_open, push_chat_turn_close, + push_chat_turn_open, }; use crate::generation::{GenerateConfig, GenerateOutput}; -use crate::kv_cache::CrossTurnSlotId; use crate::model::qwen35_config::{Qwen35Config, VisionModelConfig}; use crate::serve::ApiError; +use crate::serve::prepare::PreparationHandle; +use crate::serving_factory::ServingFactory; use crate::tokenizer::Tokenizer as _; use crate::tokenizer::bpe::BpeTokenizer; use crate::vision::VisionError; @@ -75,7 +76,6 @@ use crate::vision::qwen35_vit_metal::qwen35_vit_forward_metal_with_cancel; use std::cell::RefCell; use std::io::Write as _; use std::path::PathBuf; -use std::rc::Rc; use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::{Arc, Mutex, MutexGuard, RwLock}; use std::time::{Duration, Instant}; @@ -171,7 +171,7 @@ pub enum WorkerEvent { /// instead of `lattice_serve.rs`'s prior string-prefix-sniffing convention /// (`PROMPT_EXCEEDS_WINDOW_PREFIX`). #[derive(Debug)] -enum WorkerFailure { +pub(crate) enum WorkerFailure { Rejected(ApiError), Failed(String), /// Mirrors [`WorkerEvent::ConstraintBlocked`] -- see that variant's doc @@ -488,6 +488,7 @@ pub struct MetalWorkerClient { control: Arc, adapters: Arc>, vision_supported: Arc, + preparation: Option, /// Keeps the worker join owner alive for exactly as long as the queue /// can accept jobs. Test-only clients without a worker carry an owner /// whose join slot is already empty. @@ -507,6 +508,7 @@ impl MetalWorkerClient { control: Arc::new(Semaphore::new(1)), adapters: Arc::new(RwLock::new(AdapterIndex::default())), vision_supported, + preparation: None, _owner: owner, } } @@ -667,6 +669,12 @@ impl MetalWorkerClient { pub fn supports_vision(&self) -> bool { self.vision_supported.load(Ordering::Acquire) } + + /// Model-bound preparation, present on a successfully loaded worker. + #[doc(hidden)] + pub fn preparation(&self) -> Option<&PreparationHandle> { + self.preparation.as_ref() + } } impl Drop for MetalWorkerClient { @@ -685,7 +693,7 @@ impl Drop for MetalWorkerClient { /// reasoning tokens and one delimiter slot. `lattice.rs` keeps its /// pre-existing HTTP formula, which accepts /// `prompt_tokens + max_tokens == max_context`. -fn check_prompt_fits_window( +pub(crate) fn check_prompt_fits_window( policy: ContextWindowPolicy, model_max_context: usize, prompt_len: usize, @@ -742,7 +750,7 @@ fn check_prompt_fits_window( /// prompt with this same tokenizer, so a lower cap would silently drop the /// prompt's tail (including the open assistant turn) for any prompt that /// fits the window but exceeds the cap. Only ever raises the cap. -fn serving_tokenizer(tokenizer: BpeTokenizer, model_max_context: usize) -> BpeTokenizer { +pub(crate) fn serving_tokenizer(tokenizer: BpeTokenizer, model_max_context: usize) -> BpeTokenizer { if tokenizer.max_seq_len() < model_max_context { tokenizer.with_max_seq_len(model_max_context) } else { @@ -755,7 +763,7 @@ fn serving_tokenizer(tokenizer: BpeTokenizer, model_max_context: usize) -> BpeTo /// Admission uses the pre-truncation token count: a truncated count can never /// exceed the tokenizer's cap, so an over-window prompt would pass the check /// and then be generated from a shortened prefix. -fn render_text_prompt_within_window( +pub(crate) fn render_text_prompt_within_window( tokenizer: &BpeTokenizer, messages: &[ChatMessage], policy: ContextWindowPolicy, @@ -1047,7 +1055,7 @@ impl VisionRuntime { self.vision_supported.load(Ordering::Acquire) } - fn shared_capability(&self) -> Arc { + pub(crate) fn shared_capability(&self) -> Arc { self.vision_supported.clone() } @@ -1195,7 +1203,7 @@ fn build_vision_prompt_ids( } #[derive(Debug)] -enum VisionRequestBuild { +pub(crate) enum VisionRequestBuild { Ready { request: Qwen35VisionRequest, metal_dispatches: usize, @@ -1204,7 +1212,7 @@ enum VisionRequestBuild { Cancelled, } -fn build_vision_request( +pub(crate) fn build_vision_request( runtime: &mut VisionRuntime, config: &Qwen35Config, tokenizer: &BpeTokenizer, @@ -1348,7 +1356,7 @@ fn build_vision_request( }) } -fn cancelled_output() -> GenerateOutput { +pub(crate) fn cancelled_output() -> GenerateOutput { GenerateOutput { text: String::new(), token_ids: Vec::new(), @@ -1361,12 +1369,12 @@ fn cancelled_output() -> GenerateOutput { } #[derive(Debug, Clone, Copy, PartialEq, Eq)] -enum JobRoute { +pub(crate) enum JobRoute { Text, Vision { message_index: usize }, } -fn classify_job(messages: &[ChatMessage]) -> Result { +pub(crate) fn classify_job(messages: &[ChatMessage]) -> Result { let mut image_positions = messages .iter() .enumerate() @@ -1428,33 +1436,34 @@ impl MetalWorker { max_pending: usize, residency_limits: ResidencyLimits, ) -> Result<(MetalWorkerOwner, MetalWorkerClient, WorkerMetadata), StartupError> { - Self::spawn_with_vision( - loader, - VisionRuntime::unsupported(), + let (owner, client, metadata, _preparation) = Self::spawn_with_vision( + ServingFactory::legacy_qwen(loader), max_pending, residency_limits, - ) + )?; + Ok((owner, client, metadata)) } - /// Vision-capable sibling of [`Self::spawn`]. - /// - /// `vision_runtime` is derived from the same concrete checkpoint config - /// as `loader`. It remains worker-local and loads vision tensors only - /// when the first image-bearing job is actually dispatched. + /// Spawn from a one-shot model factory. Its Metal state is built on the + /// dedicated worker thread after the spawn. pub fn spawn_with_vision( - loader: impl FnOnce() -> Result<(MetalQwen35State, BpeTokenizer, WorkerMetadata), String> - + Send - + 'static, - mut vision_runtime: VisionRuntime, + factory: ServingFactory, max_pending: usize, residency_limits: ResidencyLimits, - ) -> Result<(MetalWorkerOwner, MetalWorkerClient, WorkerMetadata), StartupError> { + ) -> Result< + ( + MetalWorkerOwner, + MetalWorkerClient, + WorkerMetadata, + PreparationHandle, + ), + StartupError, + > { if residency_limits.max_adapters == 0 || residency_limits.max_bytes == 0 { return Err(StartupError::InvalidResidencyLimits { limits: residency_limits, }); } - let vision_supported = vision_runtime.shared_capability(); // #939: validate BEFORE `Semaphore::new`, which panics outright for // `max_pending > Semaphore::MAX_PERMITS` and would otherwise let // `max_pending == 0` silently build a worker that admits nothing. @@ -1463,189 +1472,43 @@ impl MetalWorker { } let (job_tx, job_rx) = mpsc::unbounded_channel::(); let admission = Arc::new(Semaphore::new(max_pending)); - let (ready_tx, ready_rx) = std::sync::mpsc::channel::>(); + let (ready_tx, ready_rx) = std::sync::mpsc::channel::< + Result<(WorkerMetadata, PreparationHandle, Arc), String>, + >(); let adapters = Arc::new(RwLock::new(AdapterIndex::default())); let worker_index = Arc::clone(&adapters); - let join_handle = std::thread::spawn(move || match loader() { - Ok((state, tokenizer, meta)) => { - let tokenizer = serving_tokenizer(tokenizer, meta.model_max_context); - let _ = ready_tx.send(Ok(meta.clone())); - // `Rc`/`RefCell`, not `Arc`/`Mutex`: both handles are created - // here, inside the spawned thread and after `loader()` has run - // on it, and neither ever leaves. `!Send` is the correct - // property for a handle to `!Send` state -- it is the compiler - // checking the confinement this whole module exists to - // maintain, not a cost being paid to work around it. - // - // The two borrows below can never overlap: the loop holds - // exactly one message at a time and calls exactly one of the - // two closures per message, so `borrow_mut` is uncontended by - // construction rather than by convention. - let state_rc = Rc::new(RefCell::new(state)); - let state_for_control = Rc::clone(&state_rc); - let registry = Rc::new(RefCell::new(ResidencyRegistry::new( - worker_index, - residency_limits, - ))); - let control_registry = Rc::clone(®istry); - run_worker_loop_with_lora( - job_rx, - move |messages, cfg, lora, on_token, should_cancel| { - let mut guard = state_rc.borrow_mut(); - let state = &mut *guard; - 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}" - ); - registry - .borrow_mut() - .apply(lora, state) - .map_err(WorkerFailure::Rejected)?; - 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); - } - - // Render the ChatML prompt exactly once (#828/#832: the - // prior `lattice_serve.rs` path rendered it a second - // time inside its own window preflight); reused for - // both the window check and the generation call below. - let (prompt, _prompt_len) = render_text_prompt_within_window( - &tokenizer, - messages, - meta.context_window_policy, - meta.model_max_context, - cfg, - ) - .map_err(WorkerFailure::Rejected)?; - - // Cache-aware + cancellation-aware call (#462/#744): - // reuses the previous turn's shared token prefix - // instead of a full re-prefill on every request, and - // observes client disconnect before prefill, - // immediately after prefill, and at the top of every - // decode iteration. This worker thread owns one - // `MetalQwen35State` for the whole process lifetime, so - // `CrossTurnSlotId::DEFAULT` is the only slot that - // exists; the planner re-verifies the retained prefix - // against this request's prompt on every call and - // falls back to `PrefixReuseMode::FullRefill` whenever - // they diverge, so correctness never depends on - // distinguishing clients. - // - // DEPLOYMENT ASSUMPTION, stated because it is currently - // true only by the accident that no multi-tenant consumer - // exists: this path assumes a single tenant, or clients - // that mutually trust one another. Reuse-versus-refill is - // externally visible as latency, so while no request can - // read another's content, a client CAN observe that some - // other request recently shared a prefix with its own. - // A shared inference endpoint serving mutually distrusting - // clients must key the slot per tenant via - // `CrossTurnSlotId::new`, not inherit `DEFAULT`. - if should_cancel() { - return Ok(cancelled_output()); - } - registry - .borrow_mut() - .apply(lora, state) - .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) - }, - move |command| { - let mut guard = state_for_control.borrow_mut(); - let state = &mut *guard; - match command { - AdapterCommand::Load { - name, - path, - layers, - descriptor, - } => { - let mut registry = control_registry.borrow_mut(); - let id = registry.load(name, path, layers, *descriptor)?; - registry.metadata(id).map(AdapterControlResult::Loaded) - } - AdapterCommand::Unload { id } => control_registry - .borrow_mut() - .unload(id, state) - .map(AdapterControlResult::Unloaded), - } - }, - ); - } - Err(e) => { - let _ = ready_tx.send(Err(e)); + let join_handle = std::thread::spawn(move || { + match factory.build(worker_index, residency_limits) { + Ok((runtime, metadata, preparation)) => { + let vision_supported = runtime.vision_supported(); + let _ = ready_tx.send(Ok((metadata, preparation, vision_supported))); + // The non-Send runtime was built after the spawn and never + // crosses back. The loop lends it to one callback at a + // time, so these mutable borrows cannot overlap. + let runtime = RefCell::new(runtime); + run_worker_loop_with_lora( + job_rx, + |messages, cfg, lora, on_token, should_cancel| { + runtime.borrow_mut().generate( + messages, + cfg, + lora, + on_token, + should_cancel, + ) + }, + |command| runtime.borrow_mut().control(command), + ); + } + Err(error) => { + let _ = ready_tx.send(Err(error)); + } } }); - let owner = MetalWorkerOwner::from_handle(join_handle); match ready_rx.recv() { - Ok(Ok(meta)) => { + Ok(Ok((meta, preparation, vision_supported))) => { let mut client = MetalWorkerClient::with_owner( job_tx, admission, @@ -1653,7 +1516,8 @@ impl MetalWorker { owner.clone(), ); client.adapters = adapters; - Ok((owner, client, meta)) + client.preparation = Some(preparation.clone()); + Ok((owner, client, meta, preparation)) } Ok(Err(e)) => Err(StartupError::Load(e)), Err(_) => Err(StartupError::ThreadExited), diff --git a/crates/inference/src/serve/mod.rs b/crates/inference/src/serve/mod.rs index cea378dd1f..4c940d8653 100644 --- a/crates/inference/src/serve/mod.rs +++ b/crates/inference/src/serve/mod.rs @@ -98,7 +98,7 @@ pub fn format_normalized_chat_template(messages: &[contract::NormalizedChatMessa pub mod lora; #[cfg(all(target_os = "macos", feature = "metal-gpu"))] -mod lora_registry; +pub(crate) mod lora_registry; /// Shared Metal GPU worker owner (issue #832, ADR-080 cluster C2/C3): /// the single dedicated thread that owns the `!Send` `MetalQwen35State` for /// the whole process lifetime, used by both the `lattice` unified server and diff --git a/crates/inference/src/serve/prepare.rs b/crates/inference/src/serve/prepare.rs index 24a017c89c..eac93bf90b 100644 --- a/crates/inference/src/serve/prepare.rs +++ b/crates/inference/src/serve/prepare.rs @@ -19,18 +19,108 @@ use crate::generation::GenerateConfig; use crate::serve::ApiError; use crate::serve::contract::{ ChatRequest as ChatCompletionRequest, GenerationDefaults, MessageContent, ServeProfile, - ValidatedChatRequest as ContractValidatedChatRequest, + ValidatedChatRequest as ContractValidatedChatRequest, normalize_request, normalize_request_with_context_and_budget, normalize_requested_options, validate_context_window_with_budget, }; use crate::serve::into_engine_chat_messages; use crate::serve::prompt_adapter::{PromptAdapter as _, QwenPromptAdapter}; +use crate::tokenizer::Tokenizer as _; +use crate::tokenizer::bpe::BpeTokenizer; +use std::sync::Arc; pub use crate::serve::prompt_adapter::GemmaPromptAdapter; /// The `lattice_serve` handler's name for the validated request type. type ValidatedChatRequest = ContractValidatedChatRequest; +/// Opaque, model-bound preparation for the serving binaries. +#[doc(hidden)] +#[derive(Debug, Clone)] +pub struct PreparationHandle { + tokenizer: Arc, + model_max_context: usize, +} + +impl PreparationHandle { + #[cfg(any(test, all(target_os = "macos", feature = "metal-gpu")))] + pub(crate) fn qwen(tokenizer: Arc, model_max_context: usize) -> Self { + Self { + tokenizer, + model_max_context, + } + } + + /// Tokenize with the same tokenizer used by worker execution. + pub fn tokenize_len(&self, prompt: &str) -> usize { + self.tokenizer.tokenize(prompt).pre_truncation_len + } + + /// Run the CLI's render, tokenize and context check before stop parsing. + pub fn prepare_lattice( + &self, + req: &ChatCompletionRequest, + model_id: &str, + default_max_tokens: usize, + max_tokens_cap: usize, + vision_supported: bool, + ) -> Result { + prepare_chat_request( + req, + model_id, + default_max_tokens, + max_tokens_cap, + vision_supported, + |prompt| self.tokenize_len(prompt), + || self.model_max_context, + ) + } + + /// Apply the standalone server's existing normalization profile. + pub fn normalize_standalone( + &self, + req: &ChatCompletionRequest, + defaults: GenerationDefaults, + model_id: &str, + vision_supported: bool, + ) -> Result { + normalize_request( + req, + defaults, + ServeProfile::lattice_serve(model_id, self.model_max_context) + .with_vision_support(vision_supported), + ) + } + + /// Map validated standalone options through the model's prompt adapter. + pub fn standalone_generate_config(&self, req: &ValidatedChatRequest) -> GenerateConfig { + QwenPromptAdapter.generate_config(req) + } + + /// Map prepared CLI sampling options through the model's prompt adapter. + #[allow(clippy::too_many_arguments)] + pub fn lattice_generate_config( + &self, + max_tokens: usize, + temperature: f32, + top_p: f32, + seed: Option, + stop_strings: Vec, + reasoning_budget: Option, + logprobs: Option, + ) -> GenerateConfig { + lattice_gen_cfg( + max_tokens, + temperature, + top_p, + seed, + stop_strings, + reasoning_budget, + logprobs, + ) + } +} + /// Output of the full pre-generation validation cascade, ready for /// `gen_cfg` construction. #[doc(hidden)] @@ -208,3 +298,64 @@ pub fn prepare_gemma_chat_request( prompt, }) } + +#[cfg(test)] +mod tests { + use super::{ChatCompletionRequest, PreparationHandle}; + use crate::serve::ApiError; + use crate::tokenizer::bpe::BpeTokenizer; + use std::collections::HashMap; + use std::sync::Arc; + + fn handle(model_max_context: usize) -> PreparationHandle { + let tokenizer = BpeTokenizer::from_vocab_and_merges( + HashMap::from([("a".to_string(), 0), ("b".to_string(), 1)]), + Vec::new(), + ) + .expect("tiny tokenizer must construct"); + PreparationHandle::qwen(Arc::new(tokenizer), model_max_context) + } + + fn request(stop: serde_json::Value) -> ChatCompletionRequest { + serde_json::from_value(serde_json::json!({ + "model": "served-model", + "messages": [{"role": "user", "content": "a".repeat(64)}], + "max_tokens": 1, + "stop": stop, + })) + .expect("chat request body") + } + + #[test] + fn prepare_lattice_checks_context_with_its_tokenizer_before_stop() { + let err = handle(8) + .prepare_lattice( + &request(serde_json::json!([])), + "served-model", + 1, + 4096, + false, + ) + .unwrap_err(); + assert!( + matches!( + err, + ApiError::BadRequest { + code: "context_length_exceeded", + .. + } + ), + "{err:?}" + ); + + handle(4096) + .prepare_lattice( + &request(serde_json::Value::Null), + "served-model", + 1, + 4096, + false, + ) + .expect("a prompt inside the window is admitted"); + } +} diff --git a/crates/inference/src/serving_factory.rs b/crates/inference/src/serving_factory.rs new file mode 100644 index 0000000000..3c706f7a2a --- /dev/null +++ b/crates/inference/src/serving_factory.rs @@ -0,0 +1,102 @@ +//! Select the serving family before building worker-local execution state. + +use crate::forward::metal_qwen35::MetalQwen35State; +use crate::model::serving_runtime::{QwenMetalRuntime, RuntimeFactory, ServingRuntime}; +use crate::serve::lora::{AdapterIndex, ResidencyLimits}; +use crate::serve::metal_worker::{VisionRuntime, WorkerMetadata, serving_tokenizer}; +use crate::serve::prepare::PreparationHandle; +use crate::tokenizer::bpe::BpeTokenizer; +use std::sync::{Arc, RwLock}; + +type QwenLoader = + Box Result<(MetalQwen35State, BpeTokenizer, WorkerMetadata), String> + Send>; + +/// A one-shot, thread-safe entry into worker-local model construction. +#[doc(hidden)] +pub struct ServingFactory { + inner: Box, +} + +impl ServingFactory { + /// Bind the existing Qwen loader and lazy vision state. + /// Loading still occurs on the worker thread. + pub fn qwen_metal( + loader: impl FnOnce() -> Result<(MetalQwen35State, BpeTokenizer, WorkerMetadata), String> + + Send + + 'static, + vision: VisionRuntime, + ) -> Self { + Self { + inner: Box::new(QwenFactory { + loader: Box::new(loader), + vision, + }), + } + } + + pub(crate) fn legacy_qwen( + loader: impl FnOnce() -> Result<(MetalQwen35State, BpeTokenizer, WorkerMetadata), String> + + Send + + 'static, + ) -> Self { + Self::qwen_metal(loader, VisionRuntime::unsupported()) + } + + pub(crate) fn build( + self, + index: Arc>, + limits: ResidencyLimits, + ) -> Result<(Box, WorkerMetadata, PreparationHandle), String> { + fn require_send_sync(_: &T) {} + let (runtime, metadata, preparation) = self.inner.build(index, limits)?; + require_send_sync(&preparation); + Ok((runtime, metadata, preparation)) + } +} + +struct QwenFactory { + loader: QwenLoader, + vision: VisionRuntime, +} + +impl RuntimeFactory for QwenFactory { + fn build( + self: Box, + index: Arc>, + limits: ResidencyLimits, + ) -> Result<(Box, WorkerMetadata, PreparationHandle), String> { + let Self { loader, vision } = *self; + let (state, tokenizer, metadata) = loader()?; + let tokenizer = serving_tokenizer(tokenizer, metadata.model_max_context); + let tokenizer = Arc::new(tokenizer); + let preparation = + PreparationHandle::qwen(Arc::clone(&tokenizer), metadata.model_max_context); + let runtime = + QwenMetalRuntime::new(state, tokenizer, vision, metadata.clone(), index, limits); + Ok((Box::new(runtime), metadata, preparation)) + } +} + +#[cfg(test)] +mod tests { + use super::ServingFactory; + use crate::serve::lora::{AdapterIndex, ResidencyLimits}; + use crate::serve::metal_worker::VisionRuntime; + use std::sync::{Arc, RwLock}; + + #[test] + fn qwen_loader_error_is_returned_unchanged() { + let factory = ServingFactory::qwen_metal( + || Err("legacy load error".to_string()), + VisionRuntime::unsupported(), + ); + let error = factory + .build( + Arc::new(RwLock::new(AdapterIndex::default())), + ResidencyLimits::default(), + ) + .err() + .unwrap(); + assert_eq!(error, "legacy load error"); + } +} diff --git a/crates/inference/tests/metal_measurement_lock_contract.rs b/crates/inference/tests/metal_measurement_lock_contract.rs index ec224c9194..2393ed542d 100644 --- a/crates/inference/tests/metal_measurement_lock_contract.rs +++ b/crates/inference/tests/metal_measurement_lock_contract.rs @@ -495,7 +495,7 @@ const CONSTRUCTION_EXEMPTIONS: &[ConstructionExemption] = &[ }, ConstructionExemption { site: "bin:lattice:src/bin/lattice/main.rs=>src/bin/lattice/serve.rs::ModelBackend::spawn_metal::MetalQwen35State::from_q4_dir()#1", - recorded_position: "src/bin/lattice/serve.rs:325:81", + recorded_position: "src/bin/lattice/serve.rs:393:85", reason: "ModelBackend::spawn_metal initializes a long-running server worker outside the bounded measurement-harness contract", }, ConstructionExemption {