use axum::{
Json, Router,
extract::{DefaultBodyLimit, State},
response::{
IntoResponse, Response,
sse::{Event, KeepAlive, Sse},
},
routing::{get, post},
};
use futures::StreamExt as _;
use lattice_inference::Tokenizer;
use lattice_inference::forward::metal_qwen35::ChatMessage;
#[cfg(test)]
use lattice_inference::forward::metal_qwen35::format_chat_template;
#[cfg(feature = "metal-gpu")]
use lattice_inference::model::qwen35_config::GenerateConfig;
use lattice_inference::model::qwen35_config::{GenerateOutput, TokenLogprob};
use lattice_inference::serve::contract::{
ChatRequest as ChatCompletionRequest, GenerationDefaults, ServeProfile,
ValidatedChatRequest as ContractValidatedChatRequest,
normalize_request_with_context_and_budget, validate_context_window_with_budget,
};
#[cfg(test)]
use lattice_inference::serve::contract::{
ContentPart, Message, MessageContent, ResponseFormat, normalize_request,
};
use lattice_inference::serve::{format_normalized_chat_template, into_engine_chat_messages};
use serde::Serialize;
use serde_json::Value;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::{SystemTime, UNIX_EPOCH};
use lattice_inference::serve::REQUEST_BODY_LIMIT_BYTES;
#[cfg(feature = "metal-gpu")]
#[derive(Clone)]
pub struct MetalHandle {
client: lattice_inference::serve::metal_worker::MetalWorkerClient,
}
#[cfg(feature = "metal-gpu")]
impl MetalHandle {
fn supports_vision(&self) -> bool {
self.client.supports_vision()
}
fn normalize_cancelled(
ev: lattice_inference::serve::metal_worker::WorkerEvent,
) -> lattice_inference::serve::metal_worker::WorkerEvent {
use lattice_inference::serve::metal_worker::WorkerEvent;
match ev {
WorkerEvent::Cancelled => WorkerEvent::Complete(GenerateOutput {
text: String::new(),
token_ids: vec![],
prompt_tokens: 0,
generated_tokens: 0,
stopped: false,
stop_reason: Some(lattice_inference::StopReason::Interrupt),
token_logprobs: vec![],
}),
other => other,
}
}
async fn generate_streaming(
&self,
messages: Vec<ChatMessage>,
gen_cfg: GenerateConfig,
on_token: impl FnMut(&str) -> bool + Send + 'static,
) -> Result<GenerateOutput, ApiError> {
let (_never_cancels, cancel_rx) = tokio::sync::watch::channel(false);
self.generate_streaming_with_cancel(messages, gen_cfg, on_token, cancel_rx)
.await
}
async fn generate_streaming_with_cancel(
&self,
messages: Vec<ChatMessage>,
gen_cfg: GenerateConfig,
on_token: impl FnMut(&str) -> bool + Send + 'static,
cancel: tokio::sync::watch::Receiver<bool>,
) -> Result<GenerateOutput, ApiError> {
let mut rx = self.submit(messages, gen_cfg, cancel)?;
Self::drain(&mut rx, on_token).await
}
fn submit(
&self,
messages: Vec<ChatMessage>,
gen_cfg: GenerateConfig,
cancel: tokio::sync::watch::Receiver<bool>,
) -> Result<
tokio::sync::mpsc::UnboundedReceiver<lattice_inference::serve::metal_worker::WorkerEvent>,
ApiError,
> {
self.client.submit(messages, gen_cfg, cancel)
}
async fn drain(
rx: &mut tokio::sync::mpsc::UnboundedReceiver<
lattice_inference::serve::metal_worker::WorkerEvent,
>,
mut on_token: impl FnMut(&str) -> bool + Send + 'static,
) -> Result<GenerateOutput, ApiError> {
use lattice_inference::serve::metal_worker::WorkerEvent;
let mut deliver_deltas = true;
loop {
let Some(ev) = rx.recv().await else {
return Err(ApiError::Internal {
message: "inference worker unavailable".to_string(),
});
};
match Self::normalize_cancelled(ev) {
WorkerEvent::Delta(delta) => {
if deliver_deltas && !on_token(&delta) {
deliver_deltas = false;
}
}
WorkerEvent::Complete(output) => return Ok(output),
WorkerEvent::Rejected(api_err) => return Err(api_err),
WorkerEvent::Failed(message) | WorkerEvent::ConstraintBlocked(message) => {
return Err(ApiError::Internal {
message: format!("generation failed: {message}"),
});
}
WorkerEvent::Cancelled => {
unreachable!("normalize_cancelled already rewrote Cancelled into Complete")
}
_ => {
return Err(ApiError::Internal {
message: "generation failed: unrecognized worker event".to_string(),
});
}
}
}
}
}
#[cfg(feature = "metal-gpu")]
fn map_metal_generation_error(error: ApiError) -> ApiError {
match error {
error @ (ApiError::ServiceUnavailable { .. } | ApiError::BadRequest { .. }) => error,
other => {
eprintln!("generation error (metal): {other:?}");
ApiError::Internal {
message: "inference failed".to_string(),
}
}
}
}
#[derive(Clone)]
pub enum ModelBackend {
Cpu(Arc<lattice_inference::model::qwen35::Qwen35Model>),
#[cfg(feature = "metal-gpu")]
Metal {
handle: MetalHandle,
tokenizer: Arc<lattice_inference::tokenizer::bpe::BpeTokenizer>,
max_context: usize,
},
#[cfg(all(feature = "test-utils", test))]
CpuFakeGenerate {
model: Arc<lattice_inference::model::qwen35::Qwen35Model>,
#[allow(clippy::type_complexity)]
generate: Arc<
dyn Fn(
&str,
&lattice_inference::model::qwen35_config::GenerateConfig,
&mut dyn FnMut(&str) -> bool,
&mut dyn FnMut() -> bool,
)
-> Result<GenerateOutput, lattice_inference::error::InferenceError>
+ Send
+ Sync,
>,
},
}
impl ModelBackend {
pub fn tokenize_len(&self, text: &str) -> usize {
match self {
ModelBackend::Cpu(m) => m.tokenizer().tokenize(text).real_length,
#[cfg(feature = "metal-gpu")]
ModelBackend::Metal { tokenizer, .. } => tokenizer.tokenize(text).real_length,
#[cfg(all(feature = "test-utils", test))]
ModelBackend::CpuFakeGenerate { model, .. } => {
model.tokenizer().tokenize(text).real_length
}
}
}
pub fn max_context(&self) -> usize {
match self {
ModelBackend::Cpu(m) => m.max_context(),
#[cfg(feature = "metal-gpu")]
ModelBackend::Metal { max_context, .. } => *max_context,
#[cfg(all(feature = "test-utils", test))]
ModelBackend::CpuFakeGenerate { model, .. } => model.max_context(),
}
}
pub fn tokenizer(&self) -> &lattice_inference::tokenizer::bpe::BpeTokenizer {
match self {
ModelBackend::Cpu(m) => m.tokenizer(),
#[cfg(feature = "metal-gpu")]
ModelBackend::Metal { tokenizer, .. } => tokenizer,
#[cfg(all(feature = "test-utils", test))]
ModelBackend::CpuFakeGenerate { model, .. } => model.tokenizer(),
}
}
pub fn supports_vision(&self) -> bool {
match self {
ModelBackend::Cpu(_) => false,
#[cfg(feature = "metal-gpu")]
ModelBackend::Metal { handle, .. } => handle.supports_vision(),
#[cfg(all(feature = "test-utils", test))]
ModelBackend::CpuFakeGenerate { .. } => false,
}
}
#[cfg(feature = "metal-gpu")]
pub fn spawn_metal(
model_dir: std::path::PathBuf,
tokenizer_dir: Option<std::path::PathBuf>,
max_pending: usize,
preload_vision: bool,
) -> Result<(Self, usize), String> {
use lattice_inference::serve::metal_worker::{
ContextWindowPolicy, MetalWorker, StartupError, VisionRuntime, WorkerMetadata,
};
let tokenizer_path = tokenizer_dir
.as_deref()
.unwrap_or(&model_dir)
.join("tokenizer.json");
let tokenizer = Arc::new(
lattice_inference::tokenizer::bpe::BpeTokenizer::from_tokenizer_json(&tokenizer_path)
.map_err(|e| format!("tokenizer load failed ({}): {e}", tokenizer_path.display()))?,
);
let tokenizer_for_worker = (*tokenizer).clone();
let max_context = crate::chat::chat_max_cache_len();
let vision_config = crate::chat::load_q4_config(&model_dir)?;
let mut vision_runtime =
VisionRuntime::from_model_config(model_dir.clone(), &vision_config);
if preload_vision {
if let Err(err) = vision_runtime.preload() {
eprintln!(
"Warning: --preload-vision failed, falling back to lazy vision loading: {err}"
);
}
}
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,
max_pending,
)
.map_err(|e| match e {
StartupError::Load(msg) => msg,
StartupError::ThreadExited => {
"Metal worker thread exited before loading finished".to_string()
}
err @ StartupError::InvalidMaxPending { .. } => err.to_string(),
err => err.to_string(),
})?;
drop(owner);
Ok((
ModelBackend::Metal {
handle: MetalHandle { client },
tokenizer,
max_context,
},
max_context,
))
}
}
#[derive(Clone)]
pub struct AppState {
pub model: ModelBackend,
pub default_max_tokens: usize,
pub max_tokens_cap: usize,
pub model_id: String,
pub request_counter: Arc<AtomicU64>,
pub embedding_model: Option<Arc<lattice_inference::serve::embeddings::EmbeddingModel>>,
}
use lattice_inference::serve::ApiError;
#[derive(Serialize)]
pub struct ChatCompletionResponse {
pub id: String,
pub object: String,
pub created: u64,
pub model: String,
pub choices: Vec<Choice>,
pub usage: Usage,
}
#[derive(Serialize)]
pub struct Choice {
pub index: usize,
pub message: ResponseMessage,
pub finish_reason: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub logprobs: Option<ChoiceLogprobs>,
}
#[derive(Serialize)]
pub struct ChoiceLogprobs {
pub content: Vec<TokenLogprobEntry>,
}
#[derive(Serialize)]
pub struct TokenLogprobEntry {
pub token: String,
pub logprob: f32,
pub bytes: Option<Vec<u8>>,
pub top_logprobs: Vec<TopLogprobEntry>,
}
#[derive(Serialize)]
pub struct TopLogprobEntry {
pub token: String,
pub logprob: f32,
pub bytes: Option<Vec<u8>>,
}
#[derive(Serialize)]
pub struct ResponseMessage {
pub role: String,
pub content: String,
}
#[derive(Serialize)]
pub struct Usage {
pub prompt_tokens: usize,
pub completion_tokens: usize,
pub total_tokens: usize,
}
#[derive(Serialize)]
pub struct HealthResponse {
pub status: &'static str,
}
pub enum StreamMsg {
Delta(String),
Done { finish_reason: &'static str },
Failed,
}
#[derive(Serialize)]
pub struct ChatCompletionChunk {
pub id: String,
pub object: &'static str,
pub created: u64,
pub model: String,
pub choices: Vec<ChunkChoice>,
}
#[derive(Serialize)]
pub struct ChunkChoice {
pub index: usize,
pub delta: ChunkDelta,
#[serde(skip_serializing_if = "Option::is_none")]
pub finish_reason: Option<&'static str>,
}
#[derive(Serialize)]
pub struct ChunkDelta {
#[serde(skip_serializing_if = "Option::is_none")]
pub role: Option<&'static str>,
#[serde(skip_serializing_if = "Option::is_none")]
pub content: Option<String>,
}
#[cfg(test)]
fn validate_max_tokens(
req_max: Option<usize>,
req_max_completion: Option<usize>,
default_max_tokens: usize,
max_tokens_cap: usize,
) -> Result<usize, ApiError> {
let effective = match (req_max, req_max_completion) {
(None, None) => default_max_tokens,
(Some(a), None) => a,
(None, Some(b)) => b,
(Some(a), Some(b)) if a == b => a,
(Some(a), Some(b)) => {
return Err(ApiError::BadRequest {
message: format!(
"max_tokens ({a}) and max_completion_tokens ({b}) differ; supply only one"
),
code: "invalid_request",
});
}
};
lattice_inference::serve::reject_zero_max_tokens(effective)?;
if effective > max_tokens_cap {
return Err(ApiError::BadRequest {
message: format!("max_tokens {effective} exceeds server limit {max_tokens_cap}"),
code: "max_tokens_exceeds_limit",
});
}
Ok(effective)
}
#[cfg(test)]
fn validate_temperature(value: Option<f32>) -> Result<f32, ApiError> {
lattice_inference::serve::contract::validate_temperature(
value.unwrap_or(GenerationDefaults::standard(1).temperature),
)
}
#[cfg(test)]
fn validate_top_p(value: Option<f32>) -> Result<f32, ApiError> {
lattice_inference::serve::contract::validate_top_p(
value.unwrap_or(GenerationDefaults::standard(1).top_p),
)
}
#[cfg(test)]
fn validate_logprobs(
logprobs: Option<bool>,
top_logprobs: Option<usize>,
) -> Result<Option<usize>, ApiError> {
if !logprobs.unwrap_or(false) {
if top_logprobs.is_some() {
return Err(ApiError::BadRequest {
message: "top_logprobs requires logprobs: true".to_string(),
code: "invalid_request",
});
}
return Ok(None);
}
let top_n = top_logprobs.unwrap_or(0);
if top_n > 20 {
return Err(ApiError::BadRequest {
message: format!("top_logprobs {top_n} exceeds the maximum of 20"),
code: "invalid_top_logprobs",
});
}
Ok(Some(top_n))
}
#[cfg(test)]
fn parse_stop_strings(stop: &Option<Value>) -> Result<Vec<String>, ApiError> {
lattice_inference::serve::contract::parse_stop_strings(stop)
}
#[cfg(test)]
fn reject_unsupported(req: &ChatCompletionRequest) -> Result<(), ApiError> {
if req.tools.is_some() || req.tool_choice.is_some() {
return Err(ApiError::BadRequest {
message: "tools and tool_choice are not supported by this server".to_string(),
code: "unsupported_feature",
});
}
if req.stream == Some(true) && req.logprobs.unwrap_or(false) {
return Err(ApiError::BadRequest {
message: "logprobs is not supported together with stream: true".to_string(),
code: "unsupported_feature",
});
}
if req.n.unwrap_or(1) > 1 {
return Err(ApiError::BadRequest {
message: "n > 1 is not supported".to_string(),
code: "unsupported_feature",
});
}
if let Some(fmt) = &req.response_format
&& fmt.r#type != "text"
{
return Err(ApiError::BadRequest {
message: format!(
"response_format.type '{}' is not supported; use 'text'",
fmt.r#type
),
code: "unsupported_feature",
});
}
Ok(())
}
#[cfg(test)]
fn to_chat_messages(messages: &[Message]) -> Result<Vec<ChatMessage>, ApiError> {
lattice_inference::serve::contract::normalize_messages(messages)
.and_then(into_engine_chat_messages)
}
pub(super) fn finish_reason_for(
output: &lattice_inference::model::qwen35_config::GenerateOutput,
) -> &'static str {
lattice_inference::serve::finish_reason(output.stopped)
}
fn decode_token_budget(max_tokens: usize, reasoning_budget: Option<usize>) -> usize {
match reasoning_budget {
Some(rb) if rb > 0 => rb.saturating_add(max_tokens).saturating_add(1),
_ => max_tokens,
}
}
fn render_token_logprob(
tokenizer: &lattice_inference::tokenizer::bpe::BpeTokenizer,
token_id: u32,
) -> (String, Option<Vec<u8>>) {
match tokenizer.token_bytes_for_id(token_id) {
Some(bytes) => (String::from_utf8_lossy(&bytes).into_owned(), Some(bytes)),
None => (format!("<|unresolved_token_{token_id}|>"), None),
}
}
fn build_choice_logprobs(
tokenizer: &lattice_inference::tokenizer::bpe::BpeTokenizer,
token_logprobs: &[TokenLogprob],
) -> ChoiceLogprobs {
let content = token_logprobs
.iter()
.map(|tl| {
let (token, bytes) = render_token_logprob(tokenizer, tl.token_id);
let top_logprobs = tl
.top
.iter()
.map(|alt| {
let (token, bytes) = render_token_logprob(tokenizer, alt.token_id);
TopLogprobEntry {
token,
logprob: alt.logprob,
bytes,
}
})
.collect();
TokenLogprobEntry {
token,
logprob: tl.logprob,
bytes,
top_logprobs,
}
})
.collect();
ChoiceLogprobs { content }
}
pub async fn health() -> Json<HealthResponse> {
Json(HealthResponse { status: "ok" })
}
pub async fn root() -> Json<Value> {
let mut body = lattice_inference::serve::root_body();
if let Some(endpoints) = body.get_mut("endpoints").and_then(Value::as_array_mut) {
endpoints.push(Value::String("/v1/embeddings".to_string()));
}
Json(body)
}
pub async fn list_models(State(state): State<AppState>) -> Json<Value> {
let created = SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0);
Json(lattice_inference::serve::models_list_body(
&state.model_id,
created,
))
}
#[cfg(test)]
fn validate_chat_request(
req: &ChatCompletionRequest,
model_id: &str,
default_max_tokens: usize,
max_tokens_cap: usize,
) -> Result<ContractValidatedChatRequest, ApiError> {
normalize_request(
req,
GenerationDefaults::standard(default_max_tokens),
ServeProfile::lattice(model_id, max_tokens_cap),
)
}
#[derive(Debug)]
struct PreparedChatRequest {
messages: Vec<ChatMessage>,
max_tokens: usize,
temperature: f32,
top_p: f32,
logprobs: Option<usize>,
prompt: String,
stop_strings: Vec<String>,
reasoning_budget: Option<usize>,
seed: Option<u64>,
stream: bool,
}
fn prepare_chat_request(
req: &ChatCompletionRequest,
model_id: &str,
default_max_tokens: usize,
max_tokens_cap: usize,
vision_supported: bool,
tokenize_len: impl FnOnce(&str) -> usize,
max_context: impl FnOnce() -> usize,
) -> Result<PreparedChatRequest, ApiError> {
let (validated, prompt) = normalize_request_with_context_and_budget(
req,
GenerationDefaults::standard(default_max_tokens),
ServeProfile::lattice(model_id, max_tokens_cap).with_vision_support(vision_supported),
|messages, max_tokens, reasoning_budget| {
let prompt = format_normalized_chat_template(messages);
let prompt_token_count = tokenize_len(&prompt);
validate_context_window_with_budget(
prompt_token_count,
max_tokens,
reasoning_budget,
max_context(),
)?;
Ok(prompt)
},
)?;
let ContractValidatedChatRequest {
messages,
max_tokens,
temperature,
top_p,
logprobs,
stop_strings,
reasoning_budget,
seed,
stream,
..
} = validated;
let messages = into_engine_chat_messages(messages)?;
Ok(PreparedChatRequest {
messages,
max_tokens,
temperature,
top_p,
logprobs,
prompt,
stop_strings,
reasoning_budget,
seed,
stream,
})
}
#[allow(clippy::type_complexity)]
fn spawn_cpu_style_streaming_generation(
tx: futures::channel::mpsc::UnboundedSender<StreamMsg>,
cancel_rx: tokio::sync::watch::Receiver<bool>,
prompt: String,
gen_cfg: lattice_inference::model::qwen35_config::GenerateConfig,
generate: Arc<
dyn Fn(
&str,
&lattice_inference::model::qwen35_config::GenerateConfig,
&mut dyn FnMut(&str) -> bool,
&mut dyn FnMut() -> bool,
) -> Result<GenerateOutput, lattice_inference::error::InferenceError>
+ Send
+ Sync,
>,
finish_streaming: impl FnOnce(GenerateOutput) + Send + 'static,
) {
tokio::task::spawn_blocking(move || {
let tx_delta = tx.clone();
let mut on_token = move |delta: &str| {
tx_delta
.unbounded_send(StreamMsg::Delta(delta.to_string()))
.is_ok()
};
let mut should_cancel = move || *cancel_rx.borrow();
let result = generate(&prompt, &gen_cfg, &mut on_token, &mut should_cancel);
match result {
Ok(output) => finish_streaming(output),
Err(e) => {
eprintln!("generation error (streaming): {e}");
let _ = tx.unbounded_send(StreamMsg::Failed);
}
}
});
}
pub async fn chat_completions(
State(state): State<AppState>,
headers: axum::http::HeaderMap,
body: axum::body::Body,
) -> Result<Response, ApiError> {
lattice_inference::serve::require_json_content_type(&headers)?;
let bytes = axum::body::to_bytes(body, REQUEST_BODY_LIMIT_BYTES)
.await
.map_err(|err| {
let is_length_limit = std::error::Error::source(&err)
.is_some_and(<dyn std::error::Error>::is::<http_body_util::LengthLimitError>);
if is_length_limit {
return ApiError::PayloadTooLarge {
message: "request body exceeds 1 MiB limit".to_string(),
};
}
eprintln!("invalid request body: {err}");
ApiError::BadRequest {
message: "invalid JSON request body".to_string(),
code: "invalid_request_body",
}
})?;
let req: ChatCompletionRequest = serde_json::from_slice(&bytes).map_err(|err| {
if lattice_inference::serve::contract::is_message_flood_error(&err) {
return ApiError::BadRequest {
message: lattice_inference::serve::contract::message_flood_text(),
code: "invalid_request_body",
};
}
eprintln!("invalid request body: {err}");
ApiError::BadRequest {
message: "invalid JSON request body".to_string(),
code: "invalid_request_body",
}
})?;
chat_completions_with_request(State(state), req).await
}
async fn chat_completions_with_request(
State(state): State<AppState>,
req: ChatCompletionRequest,
) -> Result<Response, ApiError> {
let PreparedChatRequest {
messages: _normalized_messages,
max_tokens,
temperature,
top_p,
logprobs,
prompt,
stop_strings,
reasoning_budget,
seed,
stream,
} = 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_inference::model::qwen35_config::GenerateConfig {
max_new_tokens: max_tokens,
temperature,
top_p,
seed,
stop_strings,
reasoning_budget,
logprobs,
..Default::default()
};
#[cfg(feature = "metal-gpu")]
let chat_messages = _normalized_messages;
#[cfg(not(feature = "metal-gpu"))]
drop(_normalized_messages);
let model = state.model.clone();
let created = SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0);
let seq = state.request_counter.fetch_add(1, Ordering::Relaxed);
let response_id = format!("chatcmpl-{created}-{seq}");
if stream {
let (tx, rx) = futures::channel::mpsc::unbounded::<StreamMsg>();
let (cancel_guard, cancel_rx) = lattice_inference::serve::cancel_pair();
let stream_id = response_id.clone();
let stream_model = state.model_id.clone();
let finish_streaming = {
let tx = tx.clone();
move |output: GenerateOutput| {
let budget = decode_token_budget(max_tokens, reasoning_budget);
if output.generated_tokens > budget {
eprintln!(
"generation invariant violation: generated_tokens={} max_tokens={} reasoning_budget={:?}",
output.generated_tokens, max_tokens, reasoning_budget
);
let _ = tx.unbounded_send(StreamMsg::Failed);
} else {
let finish_reason = finish_reason_for(&output);
let _ = tx.unbounded_send(StreamMsg::Done { finish_reason });
}
}
};
match model {
ModelBackend::Cpu(cpu_model) => {
#[allow(clippy::type_complexity)]
let generate: Arc<
dyn Fn(
&str,
&lattice_inference::model::qwen35_config::GenerateConfig,
&mut dyn FnMut(&str) -> bool,
&mut dyn FnMut() -> bool,
)
-> Result<GenerateOutput, lattice_inference::error::InferenceError>
+ Send
+ Sync,
> = Arc::new(move |p, c, on_token, should_cancel| {
cpu_model.generate_streaming_with_cancel(p, c, on_token, should_cancel)
});
spawn_cpu_style_streaming_generation(
tx,
cancel_rx,
prompt,
gen_cfg,
generate,
finish_streaming,
);
}
#[cfg(feature = "metal-gpu")]
ModelBackend::Metal { handle, .. } => {
let mut rx = handle.submit(chat_messages, gen_cfg, cancel_rx)?;
tokio::spawn(async move {
let tx_delta = tx.clone();
let result = MetalHandle::drain(&mut rx, move |delta| {
tx_delta
.unbounded_send(StreamMsg::Delta(delta.to_string()))
.is_ok()
})
.await;
match result {
Ok(output) => finish_streaming(output),
Err(e) => {
eprintln!("generation error (streaming, metal): {e:?}");
let _ = tx.unbounded_send(StreamMsg::Failed);
}
}
});
}
#[cfg(all(feature = "test-utils", test))]
ModelBackend::CpuFakeGenerate { generate, .. } => {
spawn_cpu_style_streaming_generation(
tx,
cancel_rx,
prompt,
gen_cfg,
generate,
finish_streaming,
);
}
}
let role_chunk = {
let chunk = ChatCompletionChunk {
id: stream_id.clone(),
object: "chat.completion.chunk",
created,
model: stream_model.clone(),
choices: vec![ChunkChoice {
index: 0,
delta: ChunkDelta {
role: Some("assistant"),
content: None,
},
finish_reason: None,
}],
};
let data = serde_json::to_string(&chunk).unwrap_or_default();
Ok::<Event, std::convert::Infallible>(Event::default().data(data))
};
let body_stream = rx.flat_map(move |msg| {
let _cancel_guard_tied_to_stream_lifetime = &cancel_guard;
let id = stream_id.clone();
let mdl = stream_model.clone();
match msg {
StreamMsg::Delta(text) => {
let chunk = ChatCompletionChunk {
id,
object: "chat.completion.chunk",
created,
model: mdl,
choices: vec![ChunkChoice {
index: 0,
delta: ChunkDelta {
role: None,
content: Some(text),
},
finish_reason: None,
}],
};
let data = serde_json::to_string(&chunk).unwrap_or_default();
let events: Vec<Result<Event, std::convert::Infallible>> =
vec![Ok(Event::default().data(data))];
futures::stream::iter(events)
}
StreamMsg::Done { finish_reason } => {
let chunk = ChatCompletionChunk {
id,
object: "chat.completion.chunk",
created,
model: mdl,
choices: vec![ChunkChoice {
index: 0,
delta: ChunkDelta {
role: None,
content: None,
},
finish_reason: Some(finish_reason),
}],
};
let data = serde_json::to_string(&chunk).unwrap_or_default();
let events: Vec<Result<Event, std::convert::Infallible>> = vec![
Ok(Event::default().data(data)),
Ok(Event::default().data("[DONE]")),
];
futures::stream::iter(events)
}
StreamMsg::Failed => {
let data = serde_json::json!({
"error": {
"message": "inference failed",
"type": "server_error",
"code": "internal_error",
"param": null,
}
})
.to_string();
let events: Vec<Result<Event, std::convert::Infallible>> = vec![
Ok(Event::default().data(data)),
Ok(Event::default().data("[DONE]")),
];
futures::stream::iter(events)
}
}
});
let sse_stream = futures::stream::once(async move { role_chunk }).chain(body_stream);
Ok(Sse::new(sse_stream)
.keep_alive(KeepAlive::default())
.into_response())
} else {
let output = match model {
ModelBackend::Cpu(cpu_model) => {
tokio::task::spawn_blocking(move || cpu_model.generate(&prompt, &gen_cfg))
.await
.map_err(|e| {
eprintln!("task join error: {e}");
ApiError::Internal {
message: "inference failed".to_string(),
}
})?
.map_err(|e| {
eprintln!("generation error: {e}");
ApiError::Internal {
message: "inference failed".to_string(),
}
})?
}
#[cfg(feature = "metal-gpu")]
ModelBackend::Metal { handle, .. } => handle
.generate_streaming(chat_messages, gen_cfg, |_delta| true)
.await
.map_err(map_metal_generation_error)?,
#[cfg(all(feature = "test-utils", test))]
ModelBackend::CpuFakeGenerate { generate, .. } => {
tokio::task::spawn_blocking(move || {
generate(&prompt, &gen_cfg, &mut |_delta: &str| true, &mut || false)
})
.await
.map_err(|e| {
eprintln!("task join error: {e}");
ApiError::Internal {
message: "inference failed".to_string(),
}
})?
.map_err(|e| {
eprintln!("generation error: {e}");
ApiError::Internal {
message: "inference failed".to_string(),
}
})?
}
};
let budget = decode_token_budget(max_tokens, reasoning_budget);
if output.generated_tokens > budget {
eprintln!(
"generation invariant violation: generated_tokens={} max_tokens={} reasoning_budget={:?}",
output.generated_tokens, max_tokens, reasoning_budget
);
return Err(ApiError::Internal {
message: "inference failed".to_string(),
});
}
let finish_reason = finish_reason_for(&output);
let choice_logprobs = logprobs
.is_some()
.then(|| build_choice_logprobs(state.model.tokenizer(), &output.token_logprobs));
let response = ChatCompletionResponse {
id: response_id,
object: "chat.completion".to_string(),
created,
model: state.model_id.clone(),
choices: vec![Choice {
index: 0,
message: ResponseMessage {
role: "assistant".to_string(),
content: output.text.clone(),
},
finish_reason: finish_reason.to_string(),
logprobs: choice_logprobs,
}],
usage: Usage {
prompt_tokens: output.prompt_tokens,
completion_tokens: output.generated_tokens,
total_tokens: output.prompt_tokens + output.generated_tokens,
},
};
Ok(Json(response).into_response())
}
}
pub async fn embeddings(
State(state): State<AppState>,
headers: axum::http::HeaderMap,
body: axum::body::Body,
) -> Result<Response, ApiError> {
use lattice_inference::serve::embeddings::{
EmbeddingsRequest, embed_items, normalize_embedding_items, parse_pooling,
};
lattice_inference::serve::require_json_content_type(&headers)?;
let bytes = axum::body::to_bytes(body, REQUEST_BODY_LIMIT_BYTES)
.await
.map_err(|err| {
let is_length_limit = std::error::Error::source(&err)
.is_some_and(<dyn std::error::Error>::is::<http_body_util::LengthLimitError>);
if is_length_limit {
return ApiError::PayloadTooLarge {
message: "request body exceeds 1 MiB limit".to_string(),
};
}
eprintln!("invalid request body: {err}");
ApiError::BadRequest {
message: "invalid JSON request body".to_string(),
code: "invalid_request_body",
}
})?;
let req: EmbeddingsRequest = serde_json::from_slice(&bytes).map_err(|err| {
eprintln!("invalid request body: {err}");
ApiError::BadRequest {
message: "invalid JSON request body".to_string(),
code: "invalid_request_body",
}
})?;
let pooling = parse_pooling(req.pooling.as_deref())?;
let items = normalize_embedding_items(req.input.into_items())?;
let Some(embedder) = state.embedding_model.clone() else {
return Err(ApiError::BadRequest {
message: "embeddings require a loaded vision-language checkpoint; restart this \
server with `--model` pointed at a vision-language checkpoint directory \
to enable this route"
.to_string(),
code: "vision_unsupported",
});
};
let model_id = state.model_id.clone();
let (data, usage) = tokio::task::spawn_blocking(move || embed_items(&embedder, items, pooling))
.await
.map_err(|e| {
eprintln!("task join error: {e}");
ApiError::Internal {
message: "inference failed".to_string(),
}
})??;
Ok(
Json(lattice_inference::serve::embeddings::EmbeddingsResponse {
object: "list",
data,
model: model_id,
usage,
})
.into_response(),
)
}
pub fn router(state: AppState) -> Router {
Router::new()
.route("/", get(root))
.route("/health", get(health))
.route("/v1/models", get(list_models))
.route("/v1/chat/completions", post(chat_completions))
.route("/v1/embeddings", post(embeddings))
.layer(DefaultBodyLimit::max(REQUEST_BODY_LIMIT_BYTES))
.with_state(state)
}
#[cfg(test)]
mod tests {
use super::*;
use axum::http::StatusCode;
use lattice_inference::forward::metal_qwen35::ChatRole;
#[cfg(all(feature = "metal-gpu", feature = "test-utils"))]
#[test]
fn metal_handle_is_backed_by_the_shared_metal_worker_client() {
fn build_from_shared_client(
client: lattice_inference::serve::metal_worker::MetalWorkerClient,
) -> MetalHandle {
MetalHandle { client }
}
let (client, _jobs_rx) = lattice_inference::serve::metal_worker::test_client_and_jobs();
let _handle: MetalHandle = build_from_shared_client(client);
}
#[cfg(all(feature = "metal-gpu", feature = "test-utils"))]
mod admission_cap_939 {
use super::*;
use axum::body::Body;
use axum::http::StatusCode;
use lattice_inference::serve::metal_worker::{ContextWindowPolicy, spawn_fake_with_cap};
use std::sync::mpsc as std_mpsc;
use tower::ServiceExt as _;
fn tiny_tokenizer() -> lattice_inference::tokenizer::bpe::BpeTokenizer {
lattice_inference::model::qwen35::test_support::tiny_zero_model()
.tokenizer()
.clone()
}
fn single_slot_blocking_state() -> (AppState, std_mpsc::Sender<()>, std_mpsc::Receiver<()>)
{
let (unblock_tx, unblock_rx) = std_mpsc::channel::<()>();
let unblock_rx = std::sync::Mutex::new(unblock_rx);
let (started_tx, started_rx) = std_mpsc::channel::<()>();
let client = spawn_fake_with_cap(
1,
ContextWindowPolicy::PromptAndMaxTokens,
4096,
tiny_tokenizer(),
move |_messages, _cfg, prompt_tokens, _on_token, _should_cancel| {
let _ = started_tx.send(());
let _ = unblock_rx.lock().unwrap().recv();
Ok(GenerateOutput {
text: String::new(),
token_ids: vec![],
prompt_tokens,
generated_tokens: 0,
stopped: true,
stop_reason: None,
token_logprobs: vec![],
})
},
);
let state = AppState {
model: ModelBackend::Metal {
handle: MetalHandle { client },
tokenizer: Arc::new(tiny_tokenizer()),
max_context: 4096,
},
default_max_tokens: 16,
max_tokens_cap: 4096,
model_id: "test-model".to_string(),
request_counter: Arc::new(AtomicU64::new(0)),
embedding_model: None,
};
(state, unblock_tx, started_rx)
}
fn chat_request(stream: bool) -> axum::http::Request<Body> {
let body = if stream {
r#"{"model":"test-model","messages":[{"role":"user","content":"hi"}],"stream":true}"#
} else {
r#"{"model":"test-model","messages":[{"role":"user","content":"hi"}]}"#
};
axum::http::Request::builder()
.method("POST")
.uri("/v1/chat/completions")
.header("content-type", "application/json")
.body(Body::from(body.to_string()))
.expect("fixture request must build")
}
async fn assert_503_server_busy_no_sse(response: axum::http::Response<Body>, case: &str) {
assert_eq!(
response.status(),
StatusCode::SERVICE_UNAVAILABLE,
"{case}: a request submitted while the admission cap is full must be \
rejected with HTTP 503, fail-fast, before any response -- streaming \
or not -- is committed"
);
let content_type = response
.headers()
.get(axum::http::header::CONTENT_TYPE)
.and_then(|v| v.to_str().ok())
.unwrap_or_default()
.to_string();
assert!(
!content_type.contains("text/event-stream"),
"{case}: a 503 rejection must never carry an SSE content-type -- that \
would mean a streaming response had already committed before \
admission was checked: {content_type}"
);
let bytes = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.expect("503 response body must be readable");
let value: serde_json::Value = serde_json::from_slice(&bytes).unwrap_or_else(|e| {
panic!("{case}: 503 response must be JSON, not an SSE body: {e}")
});
assert_eq!(value["error"]["code"], "server_busy", "{case}: {value}");
assert_eq!(value["error"]["type"], "server_error", "{case}: {value}");
assert!(value["error"]["param"].is_null(), "{case}: {value}");
assert!(
!value["error"]["message"]
.as_str()
.unwrap_or_default()
.is_empty(),
"{case}: 503 envelope must carry a human-readable message: {value}"
);
}
#[tokio::test]
async fn chat_completions_streaming_returns_503_before_sse_commit_at_cap() {
let (state, unblock_tx, started_rx) = single_slot_blocking_state();
let app = router(state);
let app1 = app.clone();
let handle1 = tokio::spawn(async move { app1.oneshot(chat_request(false)).await });
tokio::task::spawn_blocking(move || started_rx.recv())
.await
.expect("blocking wait must not panic")
.expect("request 1's worker-thread generate() must signal it started");
let response2 = app
.clone()
.oneshot(chat_request(true))
.await
.expect("router must produce a response, not a transport error");
assert_503_server_busy_no_sse(response2, "stream:true").await;
unblock_tx.send(()).expect("unblock send must succeed");
let response1 = handle1
.await
.expect("request 1's task must not panic")
.expect("router must produce a response, not a transport error");
assert_eq!(
response1.status(),
StatusCode::OK,
"request 1 itself was never over any cap and must succeed normally"
);
}
#[tokio::test]
async fn chat_completions_non_streaming_returns_503_server_busy_not_500_at_cap() {
let (state, unblock_tx, started_rx) = single_slot_blocking_state();
let app = router(state);
let app1 = app.clone();
let handle1 = tokio::spawn(async move { app1.oneshot(chat_request(false)).await });
tokio::task::spawn_blocking(move || started_rx.recv())
.await
.expect("blocking wait must not panic")
.expect("request 1's worker-thread generate() must signal it started");
let response2 = app
.clone()
.oneshot(chat_request(false))
.await
.expect("router must produce a response, not a transport error");
assert_503_server_busy_no_sse(response2, "non-streaming").await;
unblock_tx.send(()).expect("unblock send must succeed");
let response1 = handle1
.await
.expect("request 1's task must not panic")
.expect("router must produce a response, not a transport error");
assert_eq!(
response1.status(),
StatusCode::OK,
"request 1 itself was never over any cap and must succeed normally"
);
}
}
#[test]
fn validate_max_tokens_rejects_zero() {
let err = validate_max_tokens(Some(0), None, 256, 4096).unwrap_err();
assert!(matches!(
err,
ApiError::BadRequest {
code: "invalid_max_tokens",
..
}
));
}
#[test]
fn validate_max_tokens_rejects_above_cap() {
let err = validate_max_tokens(Some(9999), None, 256, 4096).unwrap_err();
assert!(matches!(
err,
ApiError::BadRequest {
code: "max_tokens_exceeds_limit",
..
}
));
}
#[test]
fn validate_max_tokens_uses_default_when_absent() {
assert_eq!(validate_max_tokens(None, None, 128, 4096).unwrap(), 128);
}
#[test]
fn validate_max_tokens_alias_agrees() {
assert_eq!(
validate_max_tokens(Some(512), Some(512), 256, 4096).unwrap(),
512
);
}
#[test]
fn validate_max_tokens_alias_conflict_rejected() {
let err = validate_max_tokens(Some(100), Some(200), 256, 4096).unwrap_err();
assert!(matches!(
err,
ApiError::BadRequest {
code: "invalid_request",
..
}
));
}
#[test]
fn validate_temperature_rejects_negative() {
let err = validate_temperature(Some(-0.1)).unwrap_err();
assert!(matches!(
err,
ApiError::BadRequest {
code: "invalid_temperature",
..
}
));
}
#[test]
fn validate_temperature_rejects_above_two() {
let err = validate_temperature(Some(2.1)).unwrap_err();
assert!(matches!(
err,
ApiError::BadRequest {
code: "invalid_temperature",
..
}
));
}
#[test]
fn validate_temperature_accepts_boundary() {
assert_eq!(validate_temperature(Some(0.0)).unwrap(), 0.0);
assert_eq!(validate_temperature(Some(2.0)).unwrap(), 2.0);
}
#[test]
fn validate_top_p_rejects_zero() {
let err = validate_top_p(Some(0.0)).unwrap_err();
assert!(matches!(
err,
ApiError::BadRequest {
code: "invalid_top_p",
..
}
));
}
#[test]
fn validate_top_p_rejects_above_one() {
let err = validate_top_p(Some(1.1)).unwrap_err();
assert!(matches!(
err,
ApiError::BadRequest {
code: "invalid_top_p",
..
}
));
}
#[test]
fn validate_top_p_accepts_one() {
assert_eq!(validate_top_p(Some(1.0)).unwrap(), 1.0);
}
#[test]
fn chat_template_multi_message_chatml() {
let messages = vec![
Message {
role: "system".to_string(),
content: MessageContent::Text("Be helpful.".to_string()),
},
Message {
role: "user".to_string(),
content: MessageContent::Text("Hello".to_string()),
},
];
let prompt = format_chat_template(&to_chat_messages(&messages).unwrap());
assert!(prompt.contains("<|im_start|>system\nBe helpful.<|im_end|>"));
assert!(prompt.contains("<|im_start|>user\nHello<|im_end|>"));
assert!(prompt.ends_with("<|im_start|>assistant\n"));
}
#[test]
fn finish_reason_length_only_at_cap() {
use lattice_inference::model::qwen35_config::GenerateOutput;
let cap = GenerateOutput {
text: String::new(),
token_ids: vec![],
prompt_tokens: 10,
generated_tokens: 64,
stopped: false,
stop_reason: Some(lattice_inference::StopReason::Length),
token_logprobs: vec![],
};
assert_eq!(super::finish_reason_for(&cap), "length");
let natural = GenerateOutput {
text: "hello".into(),
token_ids: vec![1, 2, 3],
prompt_tokens: 10,
generated_tokens: 3,
stopped: true,
stop_reason: Some(lattice_inference::StopReason::Eos),
token_logprobs: vec![],
};
assert_eq!(super::finish_reason_for(&natural), "stop");
}
#[test]
fn finish_reason_stop_string_at_cap_is_stop_not_length() {
use lattice_inference::model::qwen35_config::GenerateOutput;
let max_tokens: usize = 4;
let output = GenerateOutput {
text: "hi".into(),
token_ids: vec![1, 2, 3, 4],
prompt_tokens: 5,
generated_tokens: max_tokens,
stopped: true,
stop_reason: Some(lattice_inference::StopReason::Eos),
token_logprobs: vec![],
};
assert_eq!(
super::finish_reason_for(&output),
"stop",
"stop-string hit at cap must yield finish_reason=stop, not length"
);
}
#[test]
fn finish_reason_natural_length_cap_is_length() {
use lattice_inference::model::qwen35_config::GenerateOutput;
let output = GenerateOutput {
text: "hi".into(),
token_ids: vec![1, 2, 3, 4],
prompt_tokens: 5,
generated_tokens: 4,
stopped: false,
stop_reason: Some(lattice_inference::StopReason::Length),
token_logprobs: vec![],
};
assert_eq!(super::finish_reason_for(&output), "length");
}
#[test]
fn reject_unsupported_stream_true_ok() {
let req = ChatCompletionRequest {
model: Some("m".to_string()),
messages: vec![],
max_tokens: None,
max_completion_tokens: None,
temperature: None,
top_p: None,
top_k: None,
repetition_penalty: None,
reasoning_budget: None,
stream: Some(true),
stop: None,
seed: None,
response_format: None,
tools: None,
tool_choice: None,
logprobs: None,
top_logprobs: None,
n: None,
};
assert!(reject_unsupported(&req).is_ok());
}
#[test]
fn reject_unsupported_stream_and_logprobs_rejected() {
let req = ChatCompletionRequest {
stream: Some(true),
logprobs: Some(true),
..bare_req()
};
let err = reject_unsupported(&req).unwrap_err();
assert!(matches!(
err,
ApiError::BadRequest {
code: "unsupported_feature",
..
}
));
}
#[test]
fn chunk_content_delta_serializes_correctly() {
let chunk = ChatCompletionChunk {
id: "chatcmpl-1-0".to_string(),
object: "chat.completion.chunk",
created: 1_000_000,
model: "test-model".to_string(),
choices: vec![ChunkChoice {
index: 0,
delta: ChunkDelta {
role: None,
content: Some("Hello".to_string()),
},
finish_reason: None,
}],
};
let json = serde_json::to_string(&chunk).unwrap();
assert!(
json.contains("\"object\":\"chat.completion.chunk\""),
"must contain object field"
);
assert!(
json.contains("\"delta\":{\"content\":\"Hello\"}"),
"delta must contain only content when role is None"
);
assert!(
!json.contains("finish_reason"),
"finish_reason must be omitted when None"
);
}
#[test]
fn chunk_finish_delta_serializes_correctly() {
let chunk = ChatCompletionChunk {
id: "chatcmpl-1-0".to_string(),
object: "chat.completion.chunk",
created: 1_000_000,
model: "test-model".to_string(),
choices: vec![ChunkChoice {
index: 0,
delta: ChunkDelta {
role: None,
content: None,
},
finish_reason: Some("stop"),
}],
};
let json = serde_json::to_string(&chunk).unwrap();
assert!(
json.contains("\"finish_reason\":\"stop\""),
"finish chunk must include finish_reason"
);
assert!(
json.contains("\"delta\":{}"),
"finish chunk delta must be empty object"
);
}
#[test]
fn reject_unsupported_n_gt_1() {
let req = ChatCompletionRequest {
model: Some("m".to_string()),
messages: vec![],
max_tokens: None,
max_completion_tokens: None,
temperature: None,
top_p: None,
top_k: None,
repetition_penalty: None,
reasoning_budget: None,
stream: None,
stop: None,
seed: None,
response_format: None,
tools: None,
tool_choice: None,
logprobs: None,
top_logprobs: None,
n: Some(3),
};
let err = reject_unsupported(&req).unwrap_err();
assert!(matches!(
err,
ApiError::BadRequest {
code: "unsupported_feature",
..
}
));
}
#[test]
fn reject_unsupported_response_format_json() {
let req = ChatCompletionRequest {
model: Some("m".to_string()),
messages: vec![],
max_tokens: None,
max_completion_tokens: None,
temperature: None,
top_p: None,
top_k: None,
repetition_penalty: None,
reasoning_budget: None,
stream: None,
stop: None,
seed: None,
response_format: Some(ResponseFormat {
r#type: "json_object".to_string(),
json_schema: None,
}),
tools: None,
tool_choice: None,
logprobs: None,
top_logprobs: None,
n: None,
};
let err = reject_unsupported(&req).unwrap_err();
assert!(matches!(
err,
ApiError::BadRequest {
code: "unsupported_feature",
..
}
));
}
fn bare_req() -> ChatCompletionRequest {
ChatCompletionRequest {
model: Some("m".to_string()),
messages: vec![],
max_tokens: None,
max_completion_tokens: None,
temperature: None,
top_p: None,
top_k: None,
repetition_penalty: None,
reasoning_budget: None,
stream: None,
stop: None,
seed: None,
response_format: None,
tools: None,
tool_choice: None,
logprobs: None,
top_logprobs: None,
n: None,
}
}
#[test]
fn reject_unsupported_tools_rejected() {
let req = ChatCompletionRequest {
tools: Some(serde_json::json!([])),
..bare_req()
};
let err = reject_unsupported(&req).unwrap_err();
assert!(matches!(
err,
ApiError::BadRequest {
code: "unsupported_feature",
..
}
));
}
#[test]
fn reject_unsupported_tool_choice_rejected() {
let req = ChatCompletionRequest {
tool_choice: Some(serde_json::json!("auto")),
..bare_req()
};
let err = reject_unsupported(&req).unwrap_err();
assert!(matches!(
err,
ApiError::BadRequest {
code: "unsupported_feature",
..
}
));
}
#[test]
fn reject_unsupported_logprobs_true_ok() {
let req = ChatCompletionRequest {
logprobs: Some(true),
..bare_req()
};
assert!(reject_unsupported(&req).is_ok());
}
#[test]
fn reject_unsupported_stop_now_accepted() {
let req = ChatCompletionRequest {
stop: Some(serde_json::json!("</s>")),
..bare_req()
};
assert!(reject_unsupported(&req).is_ok());
}
#[test]
fn parse_stop_strings_null_gives_empty() {
assert_eq!(parse_stop_strings(&None).unwrap(), Vec::<String>::new());
assert_eq!(
parse_stop_strings(&Some(serde_json::Value::Null)).unwrap(),
Vec::<String>::new()
);
}
#[test]
fn parse_stop_strings_single_string_gives_vec_of_one() {
let v = parse_stop_strings(&Some(serde_json::json!("</s>"))).unwrap();
assert_eq!(v, vec!["</s>".to_string()]);
}
#[test]
fn parse_stop_strings_array_of_two_accepted() {
let v = parse_stop_strings(&Some(serde_json::json!(["</s>", "\nUser:"]))).unwrap();
assert_eq!(v, vec!["</s>".to_string(), "\nUser:".to_string()]);
}
#[test]
fn parse_stop_strings_empty_array_rejected() {
let err = parse_stop_strings(&Some(serde_json::json!([]))).unwrap_err();
assert!(matches!(
err,
ApiError::BadRequest {
code: "invalid_stop",
..
}
));
}
#[test]
fn parse_stop_strings_array_over_four_rejected() {
let err =
parse_stop_strings(&Some(serde_json::json!(["a", "b", "c", "d", "e"]))).unwrap_err();
assert!(matches!(
err,
ApiError::BadRequest {
code: "invalid_stop",
..
}
));
}
#[test]
fn parse_stop_strings_array_with_number_rejected() {
let err = parse_stop_strings(&Some(serde_json::json!(["ok", 42]))).unwrap_err();
assert!(matches!(
err,
ApiError::BadRequest {
code: "invalid_stop",
..
}
));
}
#[test]
fn parse_stop_strings_empty_string_element_rejected() {
let err = parse_stop_strings(&Some(serde_json::json!(["ok", ""]))).unwrap_err();
assert!(matches!(
err,
ApiError::BadRequest {
code: "invalid_stop",
..
}
));
}
#[test]
fn parse_stop_strings_empty_string_scalar_rejected() {
let err = parse_stop_strings(&Some(serde_json::json!(""))).unwrap_err();
assert!(matches!(
err,
ApiError::BadRequest {
code: "invalid_stop",
..
}
));
}
#[test]
fn parse_stop_strings_array_exactly_four_accepted() {
let v = parse_stop_strings(&Some(serde_json::json!(["a", "b", "c", "d"]))).unwrap();
assert_eq!(v.len(), 4);
}
#[test]
fn reject_unsupported_stream_false_ok() {
let req = ChatCompletionRequest {
stream: Some(false),
..bare_req()
};
assert!(reject_unsupported(&req).is_ok());
}
#[test]
fn reject_unsupported_n_1_ok() {
let req = ChatCompletionRequest {
n: Some(1),
..bare_req()
};
assert!(reject_unsupported(&req).is_ok());
}
#[test]
fn reject_unsupported_response_format_text_ok() {
let req = ChatCompletionRequest {
response_format: Some(ResponseFormat {
r#type: "text".to_string(),
json_schema: None,
}),
..bare_req()
};
assert!(reject_unsupported(&req).is_ok());
}
#[test]
fn reject_unsupported_logprobs_false_ok() {
let req = ChatCompletionRequest {
logprobs: Some(false),
..bare_req()
};
assert!(reject_unsupported(&req).is_ok());
}
#[test]
fn validate_max_tokens_at_exactly_cap_ok() {
assert_eq!(
validate_max_tokens(Some(4096), None, 256, 4096).unwrap(),
4096
);
}
#[test]
fn validate_max_tokens_max_completion_only_ok() {
assert_eq!(
validate_max_tokens(None, Some(512), 256, 4096).unwrap(),
512
);
}
#[test]
fn validate_temperature_none_uses_default() {
assert_eq!(
validate_temperature(None).unwrap(),
GenerationDefaults::standard(1).temperature
);
}
#[test]
fn validate_top_p_none_uses_default() {
assert_eq!(
validate_top_p(None).unwrap(),
GenerationDefaults::standard(1).top_p
);
}
#[test]
fn chat_template_user_only() {
let msgs = vec![Message {
role: "user".to_string(),
content: MessageContent::Text("hi".to_string()),
}];
let prompt = format_chat_template(&to_chat_messages(&msgs).unwrap());
assert_eq!(
prompt,
"<|im_start|>user\nhi<|im_end|>\n<|im_start|>assistant\n"
);
}
#[test]
fn chat_template_multi_turn_assistant() {
let msgs = vec![
Message {
role: "user".to_string(),
content: MessageContent::Text("q1".to_string()),
},
Message {
role: "assistant".to_string(),
content: MessageContent::Text("a1".to_string()),
},
Message {
role: "user".to_string(),
content: MessageContent::Text("q2".to_string()),
},
];
let prompt = format_chat_template(&to_chat_messages(&msgs).unwrap());
assert!(prompt.contains("<|im_start|>user\nq1<|im_end|>"));
assert!(prompt.contains("<|im_start|>assistant\na1<|im_end|>"));
assert!(prompt.contains("<|im_start|>user\nq2<|im_end|>"));
assert!(prompt.ends_with("<|im_start|>assistant\n"));
}
#[test]
fn chat_template_content_parts_text_ok() {
let msgs = vec![Message {
role: "user".to_string(),
content: MessageContent::Parts(vec![
ContentPart::Text {
text: "hello".to_string(),
},
ContentPart::Text {
text: " world".to_string(),
},
]),
}];
let prompt = format_chat_template(&to_chat_messages(&msgs).unwrap());
assert!(prompt.contains("<|im_start|>user\nhello world<|im_end|>"));
}
#[test]
fn to_chat_messages_rejects_invalid_role() {
let messages = vec![Message {
role: "function".to_string(),
content: MessageContent::Text("data".to_string()),
}];
let err = to_chat_messages(&messages).unwrap_err();
assert!(matches!(
err,
ApiError::BadRequest {
code: "invalid_role",
..
}
));
}
#[test]
fn to_chat_messages_rejects_tool_role() {
let messages = vec![Message {
role: "tool".to_string(),
content: MessageContent::Text("result".to_string()),
}];
let err = to_chat_messages(&messages).unwrap_err();
assert!(matches!(
err,
ApiError::BadRequest {
code: "unsupported_feature",
..
}
));
}
#[test]
fn to_chat_messages_rejects_developer_role() {
let messages = vec![Message {
role: "developer".to_string(),
content: MessageContent::Text("system prompt".to_string()),
}];
let err = to_chat_messages(&messages).unwrap_err();
assert!(matches!(
err,
ApiError::BadRequest {
code: "unsupported_feature",
..
}
));
}
#[test]
fn to_chat_messages_rejects_non_text_content_part() {
let messages = vec![Message {
role: "user".to_string(),
content: MessageContent::Parts(vec![ContentPart::ImageUrl {
image_url: lattice_inference::serve::contract::ImageUrl {
url: "https://example.com/image.png".to_string(),
detail: None,
},
}]),
}];
let err = to_chat_messages(&messages).unwrap_err();
assert!(matches!(
err,
ApiError::BadRequest {
code: "vision_unsupported",
..
}
));
}
#[test]
fn to_chat_messages_accepts_valid_roles() {
let messages = vec![
Message {
role: "system".to_string(),
content: MessageContent::Text("Be helpful.".to_string()),
},
Message {
role: "user".to_string(),
content: MessageContent::Text("q1".to_string()),
},
Message {
role: "assistant".to_string(),
content: MessageContent::Text("a1".to_string()),
},
];
let chat_messages = to_chat_messages(&messages).unwrap();
assert_eq!(chat_messages.len(), 3);
assert_eq!(chat_messages[0].role, ChatRole::System);
assert_eq!(chat_messages[0].content, "Be helpful.");
assert_eq!(chat_messages[1].role, ChatRole::User);
assert_eq!(chat_messages[2].role, ChatRole::Assistant);
}
async fn response_json(response: axum::response::Response) -> serde_json::Value {
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.expect("response body reads");
serde_json::from_slice(&body).expect("response body is valid JSON")
}
#[tokio::test]
async fn error_envelope_bad_request_shape() {
let err = ApiError::BadRequest {
message: "test error".to_string(),
code: "invalid_request",
};
assert!(matches!(
err,
ApiError::BadRequest {
code: "invalid_request",
..
}
));
let response = ApiError::BadRequest {
message: "test error".to_string(),
code: "invalid_request",
}
.into_response();
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
let v = response_json(response).await;
assert!(v["error"].is_object(), "top-level key must be 'error'");
assert_eq!(v["error"]["message"], "test error");
assert_eq!(v["error"]["type"], "invalid_request_error");
assert_eq!(v["error"]["code"], "invalid_request");
assert!(v["error"]["param"].is_null());
}
#[tokio::test]
async fn error_envelope_payload_too_large_shape() {
let response = ApiError::PayloadTooLarge {
message: "request body exceeds 1 MiB limit".to_string(),
}
.into_response();
assert_eq!(response.status(), StatusCode::PAYLOAD_TOO_LARGE);
let v = response_json(response).await;
assert_eq!(v["error"]["code"], "request_body_too_large");
}
#[tokio::test]
async fn error_envelope_internal_shape() {
let response = ApiError::Internal {
message: "inference failed".to_string(),
}
.into_response();
assert_eq!(response.status(), StatusCode::INTERNAL_SERVER_ERROR);
let v = response_json(response).await;
assert_eq!(v["error"]["type"], "server_error");
assert_eq!(v["error"]["code"], "internal_error");
}
#[test]
fn message_text_plain_string() {
let messages = [Message {
role: "user".to_string(),
content: MessageContent::Text("hello".to_string()),
}];
assert_eq!(to_chat_messages(&messages).unwrap()[0].content, "hello");
}
#[test]
fn message_text_parts_concatenates() {
let messages = [Message {
role: "user".to_string(),
content: MessageContent::Parts(vec![
ContentPart::Text {
text: "foo".to_string(),
},
ContentPart::Text {
text: "bar".to_string(),
},
]),
}];
assert_eq!(to_chat_messages(&messages).unwrap()[0].content, "foobar");
}
#[test]
fn message_text_parts_rejects_image() {
let messages = [Message {
role: "user".to_string(),
content: MessageContent::Parts(vec![ContentPart::ImageUrl {
image_url: lattice_inference::serve::contract::ImageUrl {
url: "https://example.com/image.png".to_string(),
detail: None,
},
}]),
}];
let err = to_chat_messages(&messages).unwrap_err();
match err {
ApiError::BadRequest { message, code } => {
assert_eq!(code, "vision_unsupported");
assert_eq!(message, "image input requires a vision-capable model");
}
other => panic!("expected BadRequest, got {other:?}"),
}
}
#[test]
fn message_text_parts_rejects_unknown_part_type() {
let messages = [Message {
role: "user".to_string(),
content: MessageContent::Parts(vec![ContentPart::Unsupported {
kind: "file".to_string(),
}]),
}];
let err = to_chat_messages(&messages).unwrap_err();
match err {
ApiError::BadRequest { message, code } => {
assert_eq!(code, "unsupported_feature");
assert_eq!(
message,
"content part type 'file' is not supported; only 'text' and 'image_url' \
parts are accepted"
);
}
other => panic!("expected BadRequest, got {other:?}"),
}
}
#[test]
fn validate_logprobs_absent_disables_capture() {
assert_eq!(validate_logprobs(None, None).unwrap(), None);
}
#[test]
fn validate_logprobs_false_disables_capture() {
assert_eq!(validate_logprobs(Some(false), None).unwrap(), None);
}
#[test]
fn validate_logprobs_true_no_top_logprobs_defaults_to_zero() {
assert_eq!(validate_logprobs(Some(true), None).unwrap(), Some(0));
}
#[test]
fn validate_logprobs_true_with_top_logprobs_ok() {
assert_eq!(validate_logprobs(Some(true), Some(5)).unwrap(), Some(5));
}
#[test]
fn validate_logprobs_top_logprobs_at_boundary_twenty_ok() {
assert_eq!(validate_logprobs(Some(true), Some(20)).unwrap(), Some(20));
}
#[test]
fn validate_logprobs_top_logprobs_over_twenty_rejected() {
let err = validate_logprobs(Some(true), Some(21)).unwrap_err();
assert!(matches!(
err,
ApiError::BadRequest {
code: "invalid_top_logprobs",
..
}
));
}
#[test]
fn validate_logprobs_top_logprobs_without_logprobs_true_rejected() {
let err = validate_logprobs(None, Some(5)).unwrap_err();
assert!(matches!(
err,
ApiError::BadRequest {
code: "invalid_request",
..
}
));
}
#[test]
fn validate_logprobs_top_logprobs_with_logprobs_false_rejected() {
let err = validate_logprobs(Some(false), Some(5)).unwrap_err();
assert!(matches!(
err,
ApiError::BadRequest {
code: "invalid_request",
..
}
));
}
fn logprob_test_tokenizer() -> lattice_inference::tokenizer::bpe::BpeTokenizer {
let vocab: std::collections::HashMap<String, u32> =
[("Hello".to_string(), 0u32), ("world".to_string(), 1u32)]
.into_iter()
.collect();
lattice_inference::tokenizer::bpe::BpeTokenizer::from_vocab_and_merges(vocab, vec![])
.expect("in-memory test vocab must construct")
}
#[test]
fn render_token_logprob_resolves_known_token() {
let tokenizer = logprob_test_tokenizer();
let (token, bytes) = render_token_logprob(&tokenizer, 0);
assert_eq!(token, "Hello");
assert_eq!(bytes, Some(b"Hello".to_vec()));
}
#[test]
fn render_token_logprob_unresolved_id_fails_closed() {
let tokenizer = logprob_test_tokenizer();
let (token, bytes) = render_token_logprob(&tokenizer, 999);
assert_eq!(token, "<|unresolved_token_999|>");
assert_eq!(bytes, None);
}
#[test]
fn build_choice_logprobs_shapes_content_and_alternatives() {
let tokenizer = logprob_test_tokenizer();
let token_logprobs = vec![
TokenLogprob {
token_id: 0,
logprob: -0.1,
top: vec![
lattice_inference::model::qwen35_config::TopLogprob {
token_id: 0,
logprob: -0.1,
},
lattice_inference::model::qwen35_config::TopLogprob {
token_id: 1,
logprob: -2.3,
},
],
},
TokenLogprob {
token_id: 1,
logprob: -0.05,
top: vec![],
},
];
let choice_logprobs = build_choice_logprobs(&tokenizer, &token_logprobs);
assert_eq!(choice_logprobs.content.len(), 2);
assert_eq!(choice_logprobs.content[0].token, "Hello");
assert_eq!(choice_logprobs.content[0].logprob, -0.1);
assert_eq!(choice_logprobs.content[0].top_logprobs.len(), 2);
assert_eq!(choice_logprobs.content[0].top_logprobs[0].token, "Hello");
assert_eq!(choice_logprobs.content[0].top_logprobs[1].token, "world");
assert_eq!(choice_logprobs.content[1].token, "world");
assert_eq!(choice_logprobs.content[1].logprob, -0.05);
assert!(choice_logprobs.content[1].top_logprobs.is_empty());
}
#[test]
fn choice_logprobs_omitted_from_json_when_none() {
let choice = Choice {
index: 0,
message: ResponseMessage {
role: "assistant".to_string(),
content: "hi".to_string(),
},
finish_reason: "stop".to_string(),
logprobs: None,
};
let json = serde_json::to_string(&choice).unwrap();
assert!(
!json.contains("logprobs"),
"logprobs key must be entirely absent when None, got: {json}"
);
}
#[test]
fn choice_logprobs_present_when_requested() {
let choice = Choice {
index: 0,
message: ResponseMessage {
role: "assistant".to_string(),
content: "hi".to_string(),
},
finish_reason: "stop".to_string(),
logprobs: Some(ChoiceLogprobs {
content: vec![TokenLogprobEntry {
token: "hi".to_string(),
logprob: -0.2,
bytes: Some(b"hi".to_vec()),
top_logprobs: vec![],
}],
}),
};
let json = serde_json::to_string(&choice).unwrap();
let v: serde_json::Value = serde_json::from_str(&json).unwrap();
assert_eq!(v["logprobs"]["content"][0]["token"], "hi");
assert_eq!(v["logprobs"]["content"][0]["logprob"], -0.2);
assert_eq!(
v["logprobs"]["content"][0]["bytes"],
serde_json::json!([104, 105])
);
assert_eq!(
v["logprobs"]["content"][0]["top_logprobs"],
serde_json::json!([])
);
}
fn user_msg(text: &str) -> Message {
Message {
role: "user".to_string(),
content: MessageContent::Text(text.to_string()),
}
}
#[test]
fn cm_serve_model_mismatch_rejected() {
let req = ChatCompletionRequest {
model: Some("some-other-model".to_string()),
messages: vec![user_msg("hi")],
..bare_req()
};
let err = validate_chat_request(&req, "served-model", 256, 4096).unwrap_err();
assert!(matches!(
err,
ApiError::BadRequest {
code: "model_not_found",
..
}
));
}
#[test]
fn cm_serve_model_match_passes_model_check() {
let req = ChatCompletionRequest {
model: Some("served-model".to_string()),
messages: vec![user_msg("hi")],
..bare_req()
};
assert!(validate_chat_request(&req, "served-model", 256, 4096).is_ok());
}
#[test]
fn cm_serve_empty_messages_rejected() {
let req = ChatCompletionRequest {
model: Some("served-model".to_string()),
messages: vec![],
..bare_req()
};
let err = validate_chat_request(&req, "served-model", 256, 4096).unwrap_err();
assert!(matches!(
err,
ApiError::BadRequest {
code: "invalid_messages",
..
}
));
}
#[test]
fn cm_serve_last_message_not_user_rejected() {
let req = ChatCompletionRequest {
model: Some("served-model".to_string()),
messages: vec![
user_msg("hi"),
Message {
role: "assistant".to_string(),
content: MessageContent::Text("hello".to_string()),
},
],
..bare_req()
};
let err = validate_chat_request(&req, "served-model", 256, 4096).unwrap_err();
assert!(matches!(
err,
ApiError::BadRequest {
code: "invalid_messages",
..
}
));
}
#[test]
fn cm_serve_unsupported_feature_rejected_before_model_check() {
let req = ChatCompletionRequest {
model: Some("some-other-model".to_string()),
messages: vec![user_msg("hi")],
tools: Some(serde_json::json!([{"type": "function"}])),
..bare_req()
};
let err = validate_chat_request(&req, "served-model", 256, 4096).unwrap_err();
assert!(matches!(
err,
ApiError::BadRequest {
code: "unsupported_feature",
..
}
));
}
#[test]
fn cm_serve_stop_sequences_accepted_end_to_end() {
let req = ChatCompletionRequest {
model: Some("served-model".to_string()),
messages: vec![user_msg("hi")],
stop: Some(serde_json::json!(["\n\n"])),
..bare_req()
};
let prepared =
prepare_chat_request(&req, "served-model", 256, 4096, false, |_| 1, || 4096).unwrap();
assert_eq!(prepared.stop_strings, vec!["\n\n".to_string()]);
}
#[test]
fn cm_serve_context_window_checked_before_stop_parsing() {
let req = ChatCompletionRequest {
model: Some("served-model".to_string()),
messages: vec![user_msg("hi")],
stop: Some(serde_json::json!([])), ..bare_req()
};
let err = prepare_chat_request(&req, "served-model", 256, 4096, false, |_| 4096, || 4096)
.unwrap_err();
assert!(matches!(
err,
ApiError::BadRequest {
code: "context_length_exceeded",
..
}
));
}
#[test]
fn cm_serve_logprobs_resolved_end_to_end() {
let req = ChatCompletionRequest {
model: Some("served-model".to_string()),
messages: vec![user_msg("hi")],
logprobs: Some(true),
top_logprobs: Some(3),
..bare_req()
};
let validated = validate_chat_request(&req, "served-model", 256, 4096).unwrap();
assert_eq!(validated.logprobs, Some(3));
}
#[cfg(feature = "test-utils")]
fn tiny_state(max_tokens_cap: usize) -> AppState {
let model = lattice_inference::model::qwen35::test_support::tiny_zero_model();
AppState {
model: ModelBackend::Cpu(Arc::new(model)),
default_max_tokens: max_tokens_cap,
max_tokens_cap,
model_id: "test-model".to_string(),
request_counter: Arc::new(AtomicU64::new(0)),
embedding_model: None,
}
}
#[cfg(all(feature = "metal-gpu", feature = "test-utils"))]
mod vision_content_parts_1135 {
use super::*;
use axum::body::Body;
use base64::Engine as _;
use lattice_inference::serve::metal_worker::{ContextWindowPolicy, spawn_fake_with_vision};
use std::sync::atomic::{AtomicBool, Ordering};
use tower::ServiceExt as _;
fn inline_png_data_uri() -> String {
let image = image::RgbImage::new(32, 32);
let mut bytes = Vec::new();
image
.write_to(
&mut std::io::Cursor::new(&mut bytes),
image::ImageFormat::Png,
)
.expect("PNG fixture must encode");
format!(
"data:image/png;base64,{}",
base64::engine::general_purpose::STANDARD.encode(bytes)
)
}
fn request(data_uri: &str) -> axum::http::Request<Body> {
let body = serde_json::json!({
"model": "test-model",
"messages": [{
"role": "user",
"content": [
{"type": "text", "text": "before"},
{"type": "image_url", "image_url": {"url": data_uri}},
{"type": "text", "text": "after"}
]
}]
});
axum::http::Request::builder()
.method("POST")
.uri("/v1/chat/completions")
.header("content-type", "application/json")
.body(Body::from(body.to_string()))
.expect("request fixture must build")
}
#[test]
fn worker_bad_request_remains_a_client_error() {
let error = map_metal_generation_error(ApiError::BadRequest {
message: "image geometry is unsupported".to_string(),
code: "invalid_image",
});
assert!(matches!(
error,
ApiError::BadRequest {
code: "invalid_image",
..
}
));
}
#[tokio::test]
async fn vision_model_accepts_image_and_enqueues_it_on_the_shared_worker() {
let tokenizer = lattice_inference::model::qwen35::test_support::tiny_zero_model()
.tokenizer()
.clone();
let image_seen = Arc::new(AtomicBool::new(false));
let seen = Arc::clone(&image_seen);
let client = spawn_fake_with_vision(
ContextWindowPolicy::PromptAndMaxTokens,
4096,
tokenizer.clone(),
move |messages, _cfg, prompt_tokens, _on_token, _should_cancel| {
let image = messages[0]
.image
.as_ref()
.expect("image must reach the common worker job");
assert!(image.bytes.starts_with(b"\x89PNG\r\n\x1a\n"));
assert_eq!(image.text_offset, "before".len());
assert_eq!(messages[0].content, "beforeafter");
seen.store(true, Ordering::SeqCst);
Ok(GenerateOutput {
text: "ok".to_string(),
token_ids: vec![0],
prompt_tokens,
generated_tokens: 1,
stopped: true,
stop_reason: None,
token_logprobs: vec![],
})
},
);
let state = AppState {
model: ModelBackend::Metal {
handle: MetalHandle { client },
tokenizer: Arc::new(tokenizer),
max_context: 4096,
},
default_max_tokens: 16,
max_tokens_cap: 64,
model_id: "test-model".to_string(),
request_counter: Arc::new(AtomicU64::new(0)),
embedding_model: None,
};
let response = router(state)
.oneshot(request(&inline_png_data_uri()))
.await
.expect("router must return a response");
assert_eq!(response.status(), StatusCode::OK);
assert!(image_seen.load(Ordering::SeqCst));
}
#[tokio::test]
async fn text_only_model_rejects_the_same_image_with_capability_code() {
let response = router(tiny_state(64))
.oneshot(request(&inline_png_data_uri()))
.await
.expect("router must return a response");
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.expect("error body must be readable");
let value: serde_json::Value =
serde_json::from_slice(&body).expect("error body must be JSON");
assert_eq!(value["error"]["code"], "vision_unsupported");
assert_eq!(
value["error"]["message"],
"image input requires a vision-capable model"
);
}
}
#[cfg(feature = "test-utils")]
mod embeddings_route {
use super::*;
use axum::body::Body;
use lattice_inference::serve::embeddings::test_support::{
tiny_embedding_model, tiny_png_data_uri,
};
use tower::ServiceExt as _;
fn post_embeddings(body: serde_json::Value) -> axum::http::Request<Body> {
axum::http::Request::builder()
.method("POST")
.uri("/v1/embeddings")
.header("content-type", "application/json")
.body(Body::from(body.to_string()))
.expect("request fixture must build")
}
fn state_with_embedder() -> AppState {
let mut state = tiny_state(64);
state.embedding_model = Some(Arc::new(tiny_embedding_model()));
state
}
async fn json_body(response: Response) -> serde_json::Value {
let bytes = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.expect("response body must be readable");
serde_json::from_slice(&bytes).expect("response body must be JSON")
}
#[tokio::test]
async fn no_embedding_model_loaded_fails_closed() {
let response = router(tiny_state(64))
.oneshot(post_embeddings(serde_json::json!({"input": "a"})))
.await
.expect("router must return a response");
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
let value = json_body(response).await;
assert_eq!(value["error"]["code"], "vision_unsupported");
}
#[tokio::test]
async fn happy_path_text_only() {
let response = router(state_with_embedder())
.oneshot(post_embeddings(serde_json::json!({"input": "a"})))
.await
.expect("router must return a response");
assert_eq!(response.status(), StatusCode::OK);
let value = json_body(response).await;
assert_eq!(value["object"], "list");
assert_eq!(value["data"][0]["object"], "embedding");
assert_eq!(value["data"][0]["index"], 0);
assert_eq!(value["data"][0]["embedding"].as_array().unwrap().len(), 8);
}
#[tokio::test]
async fn happy_path_image() {
let response = router(state_with_embedder())
.oneshot(post_embeddings(serde_json::json!({
"input": {"type": "image_url", "image_url": {"url": tiny_png_data_uri(0)}},
})))
.await
.expect("router must return a response");
assert_eq!(response.status(), StatusCode::OK);
let value = json_body(response).await;
let embedding = value["data"][0]["embedding"].as_array().unwrap();
let norm: f64 = embedding
.iter()
.map(|x| x.as_f64().unwrap().powi(2))
.sum::<f64>()
.sqrt();
assert!((norm - 1.0).abs() < 1e-3, "expected unit norm, got {norm}");
}
#[tokio::test]
async fn mixed_batch_preserves_input_order() {
let response = router(state_with_embedder())
.oneshot(post_embeddings(serde_json::json!({
"input": [
"a",
{"type": "image_url", "image_url": {"url": tiny_png_data_uri(1)}},
"b",
],
})))
.await
.expect("router must return a response");
assert_eq!(response.status(), StatusCode::OK);
let value = json_body(response).await;
let data = value["data"].as_array().unwrap();
assert_eq!(data.len(), 3);
assert_eq!(data[0]["index"], 0);
assert_eq!(data[1]["index"], 1);
assert_eq!(data[2]["index"], 2);
}
#[tokio::test]
async fn remote_url_rejected() {
let response = router(state_with_embedder())
.oneshot(post_embeddings(serde_json::json!({
"input": {"type": "image_url", "image_url": {"url": "https://example.com/cat.png"}},
})))
.await
.expect("router must return a response");
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
let value = json_body(response).await;
assert_eq!(value["error"]["code"], "unsupported_image_url_scheme");
}
#[tokio::test]
async fn malformed_data_uri_rejected() {
let response = router(state_with_embedder())
.oneshot(post_embeddings(serde_json::json!({
"input": {"type": "image_url", "image_url": {"url": "data:image/png,not-base64-marked"}},
})))
.await
.expect("router must return a response");
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
}
#[tokio::test]
async fn invalid_pooling_rejected() {
let response = router(state_with_embedder())
.oneshot(post_embeddings(
serde_json::json!({"input": "a", "pooling": "max"}),
))
.await
.expect("router must return a response");
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
let value = json_body(response).await;
assert_eq!(value["error"]["code"], "invalid_pooling");
}
#[tokio::test]
async fn empty_input_rejected() {
let response = router(state_with_embedder())
.oneshot(post_embeddings(serde_json::json!({"input": []})))
.await
.expect("router must return a response");
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
let value = json_body(response).await;
assert_eq!(value["error"]["code"], "invalid_input");
}
#[tokio::test]
async fn embeddings_route_advertised_at_root() {
let response = router(tiny_state(64))
.oneshot(
axum::http::Request::builder()
.method("GET")
.uri("/")
.body(Body::empty())
.unwrap(),
)
.await
.expect("router must return a response");
let value = json_body(response).await;
assert!(
value["endpoints"]
.as_array()
.unwrap()
.iter()
.any(|e| e == "/v1/embeddings")
);
}
}
#[cfg(feature = "test-utils")]
mod disconnect_cancellation {
use super::*;
use http_body_util::BodyExt;
use std::time::Duration;
const NEAR_MAX_CONTEXT_TOKENS: usize = 900;
fn tiny_state() -> AppState {
super::tiny_state(NEAR_MAX_CONTEXT_TOKENS)
}
#[tokio::test]
async fn chat_completions_streaming_disconnect_stops_generation() {
let req = ChatCompletionRequest {
model: Some("test-model".to_string()),
messages: vec![user_msg("hi")],
max_tokens: Some(NEAR_MAX_CONTEXT_TOKENS),
stream: Some(true),
..bare_req()
};
let response = chat_completions_with_request(State(tiny_state()), req)
.await
.expect("streaming request must be accepted");
assert_eq!(response.status(), StatusCode::OK);
let mut body = response.into_body();
tokio::time::timeout(Duration::from_secs(10), body.frame())
.await
.expect("role chunk frame must arrive quickly")
.expect("role chunk frame must exist")
.expect("role chunk frame must be Ok");
let frame = tokio::time::timeout(Duration::from_secs(10), body.frame())
.await
.expect(
"a second frame must arrive quickly -- if this times out, \
nothing at all was sent after the role chunk",
)
.expect("second frame must exist")
.expect("second frame must be Ok");
let data = frame
.into_data()
.expect("second frame must carry data (not trailers)");
let text = std::str::from_utf8(&data).expect("SSE frame must be UTF-8");
let payload = text
.strip_prefix("data: ")
.unwrap_or(text)
.trim_end_matches(['\n', '\r']);
let chunk: serde_json::Value =
serde_json::from_str(payload).expect("SSE frame must carry a JSON chunk");
let choice = &chunk["choices"][0];
assert!(
choice["finish_reason"].is_null(),
"second frame must be a content delta, not the finish chunk \
(generation was interrupted before producing any output, \
i.e. cancel_guard fired before the stream had a chance to \
run): {chunk}"
);
assert!(
choice["delta"]["content"].is_string(),
"second frame must carry delta.content (a real generated \
token), got: {chunk}"
);
drop(body);
}
}
#[cfg(feature = "test-utils")]
mod post_drop_cancellation_probe {
use super::*;
use http_body_util::BodyExt;
use std::sync::Mutex;
use std::sync::mpsc::RecvTimeoutError;
use std::time::Duration;
const MAX_POLLS: usize = 4000;
const POLL_INTERVAL: Duration = Duration::from_millis(5);
const PROBE_POLL_BUDGET: Duration = POLL_INTERVAL.saturating_mul(MAX_POLLS as u32);
const PROBE_COMPLETION_HEADROOM: Duration = Duration::from_secs(20);
const PROBE_COMPLETION_TIMEOUT: Duration =
PROBE_POLL_BUDGET.saturating_add(PROBE_COMPLETION_HEADROOM);
const HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(10);
fn await_checkpoint(
rx: &std::sync::mpsc::Receiver<()>,
checkpoint: &str,
timeout: Duration,
) {
match rx.recv_timeout(timeout) {
Ok(()) => {}
Err(RecvTimeoutError::Timeout) => {
panic!(
"fake generator timed out waiting for the {checkpoint} \
checkpoint after {timeout:?}"
);
}
Err(RecvTimeoutError::Disconnected) => {
panic!(
"fake generator's {checkpoint} checkpoint sender \
disconnected before signaling"
);
}
}
}
fn probe_completion_timeout_message() -> String {
format!(
"observed no ProbeOutcome within {PROBE_COMPLETION_TIMEOUT:?} \
after the client disconnect; possible causes include a \
generator scheduling delay or stall, or cancellation \
propagation failing to complete within the \
{PROBE_POLL_BUDGET:?} probe poll budget"
)
}
#[test]
fn completion_timeout_exceeds_probe_poll_budget() {
assert_eq!(PROBE_POLL_BUDGET, Duration::from_secs(20));
assert!(
PROBE_COMPLETION_TIMEOUT > PROBE_POLL_BUDGET,
"completion timeout {PROBE_COMPLETION_TIMEOUT:?} must exceed \
the complete probe poll budget {PROBE_POLL_BUDGET:?}"
);
}
#[test]
#[should_panic(expected = "post-handler-return checkpoint")]
fn post_handler_return_checkpoint_timeout_fails_loudly() {
let (tx, rx) = std::sync::mpsc::channel();
await_checkpoint(&rx, "post-handler-return", Duration::ZERO);
drop(tx);
}
#[test]
#[should_panic(expected = "post-body-drop checkpoint")]
fn post_body_drop_checkpoint_timeout_fails_loudly() {
let (tx, rx) = std::sync::mpsc::channel();
await_checkpoint(&rx, "post-body-drop", Duration::ZERO);
drop(tx);
}
#[test]
fn completion_timeout_message_reports_observation_and_candidate_causes() {
let message = probe_completion_timeout_message();
assert!(message.contains("observed no ProbeOutcome"));
assert!(message.contains("generator scheduling delay or stall"));
assert!(message.contains("cancellation propagation"));
}
const NEAR_MAX_CONTEXT_TOKENS: usize = 900;
#[derive(Debug, PartialEq)]
struct ProbeOutcome {
pre_drop_cancelled: bool,
post_drop: PostDropOutcome,
}
#[derive(Debug, PartialEq)]
enum PostDropOutcome {
CancelledAfterPolls(usize),
ExhaustedWithoutCancel,
}
fn tiny_state_with_fake_cpu_generate(
max_tokens_cap: usize,
generate: impl Fn(
&str,
&lattice_inference::model::qwen35_config::GenerateConfig,
&mut dyn FnMut(&str) -> bool,
&mut dyn FnMut() -> bool,
)
-> Result<GenerateOutput, lattice_inference::error::InferenceError>
+ Send
+ Sync
+ 'static,
) -> AppState {
let model = lattice_inference::model::qwen35::test_support::tiny_zero_model();
AppState {
model: ModelBackend::CpuFakeGenerate {
model: Arc::new(model),
generate: Arc::new(generate),
},
default_max_tokens: max_tokens_cap,
max_tokens_cap,
model_id: "test-model".to_string(),
request_counter: Arc::new(AtomicU64::new(0)),
embedding_model: None,
}
}
#[tokio::test]
async fn chat_completions_streaming_failure_emits_error_event() {
let state =
tiny_state_with_fake_cpu_generate(64, |_prompt, _cfg, on_token, _should_cancel| {
let _ = on_token("partial");
Err(lattice_inference::error::InferenceError::InvalidInput(
"blocked by grammar".to_string(),
))
});
let req = ChatCompletionRequest {
model: Some("test-model".to_string()),
messages: vec![user_msg("hi")],
max_tokens: Some(64),
stream: Some(true),
..bare_req()
};
let response = chat_completions_with_request(State(state), req)
.await
.expect("streaming request must be accepted");
assert_eq!(response.status(), StatusCode::OK);
let bytes = response
.into_body()
.collect()
.await
.expect("SSE response body must be readable")
.to_bytes();
let text = String::from_utf8(bytes.to_vec()).expect("SSE body must be valid UTF-8");
assert!(
text.contains("\"content\":\"partial\""),
"partial output must precede the error event; got: {text}"
);
let error_payload = text
.lines()
.filter_map(|line| line.strip_prefix("data: "))
.find(|payload| payload.contains("\"error\""))
.expect("a failed generation must emit an SSE error payload");
let error: serde_json::Value =
serde_json::from_str(error_payload).expect("SSE error payload must be valid JSON");
assert_eq!(error["error"]["type"], "server_error");
assert_eq!(error["error"]["code"], "internal_error");
assert!(
!text.contains("\"finish_reason\":\"stop\""),
"generation failure must not masquerade as a clean stop; got: {text}"
);
}
#[tokio::test]
async fn chat_completions_streaming_disconnect_cancellation_reaches_generator_post_drop() {
let (outcome_tx, outcome_rx) = tokio::sync::oneshot::channel::<ProbeOutcome>();
let outcome_tx = Mutex::new(Some(outcome_tx));
let (checkpoint1_tx, checkpoint1_rx) = std::sync::mpsc::channel::<()>();
let checkpoint1_rx = Mutex::new(Some(checkpoint1_rx));
let (pre_drop_observed_tx, pre_drop_observed_rx) =
tokio::sync::oneshot::channel::<bool>();
let pre_drop_observed_tx = Mutex::new(Some(pre_drop_observed_tx));
let (go_tx, go_rx) = std::sync::mpsc::channel::<()>();
let go_rx = Mutex::new(Some(go_rx));
let state = tiny_state_with_fake_cpu_generate(
NEAR_MAX_CONTEXT_TOKENS,
move |_prompt, _cfg, on_token, should_cancel| {
on_token("probe");
if let Some(rx) = checkpoint1_rx.lock().unwrap().take() {
await_checkpoint(&rx, "post-handler-return", HANDSHAKE_TIMEOUT);
}
let pre_drop_cancelled = should_cancel();
if let Some(tx) = pre_drop_observed_tx.lock().unwrap().take() {
let _ = tx.send(pre_drop_cancelled);
}
if let Some(rx) = go_rx.lock().unwrap().take() {
await_checkpoint(&rx, "post-body-drop", HANDSHAKE_TIMEOUT);
}
let mut polls = 0usize;
let post_drop = loop {
if should_cancel() {
break PostDropOutcome::CancelledAfterPolls(polls);
}
if polls >= MAX_POLLS {
break PostDropOutcome::ExhaustedWithoutCancel;
}
polls += 1;
std::thread::sleep(POLL_INTERVAL);
};
if let Some(tx) = outcome_tx.lock().unwrap().take() {
let _ = tx.send(ProbeOutcome {
pre_drop_cancelled,
post_drop,
});
}
Ok(GenerateOutput {
text: "probe".to_string(),
token_ids: vec![],
prompt_tokens: 1,
generated_tokens: 1,
stopped: false,
stop_reason: Some(lattice_inference::stop_reason::StopReason::Interrupt),
token_logprobs: vec![],
})
},
);
let req = ChatCompletionRequest {
model: Some("test-model".to_string()),
messages: vec![user_msg("hi")],
max_tokens: Some(NEAR_MAX_CONTEXT_TOKENS),
stream: Some(true),
..bare_req()
};
let response = chat_completions_with_request(State(state), req)
.await
.expect("streaming request must be accepted");
checkpoint1_tx
.send(())
.expect("generator must still be waiting for the post-return checkpoint");
let observed_pre_drop_cancelled =
tokio::time::timeout(HANDSHAKE_TIMEOUT, pre_drop_observed_rx)
.await
.expect(
"generator must acknowledge its pre-drop cancellation \
reading before the handshake timeout",
)
.expect(
"generator must not drop the pre-drop observation \
sender without acknowledging its reading",
);
assert!(
!observed_pre_drop_cancelled,
"should_cancel must read false while the response body is \
still alive"
);
let mut body = response.into_body();
tokio::time::timeout(Duration::from_secs(10), body.frame())
.await
.expect("role chunk frame must arrive quickly")
.expect("role chunk frame must exist")
.expect("role chunk frame must be Ok");
tokio::time::timeout(Duration::from_secs(10), body.frame())
.await
.expect("delta frame must arrive quickly")
.expect("delta frame must exist")
.expect("delta frame must be Ok");
drop(body);
go_tx
.send(())
.expect("generator must still be waiting for the post-drop checkpoint");
let outcome = tokio::time::timeout(PROBE_COMPLETION_TIMEOUT, outcome_rx)
.await
.unwrap_or_else(|_| panic!("{}", probe_completion_timeout_message()))
.expect("probe outcome sender must not be dropped without sending");
assert!(
!outcome.pre_drop_cancelled,
"should_cancel must read false before the test has dropped the \
response body -- a true reading here means cancel_guard had \
already dropped (e.g. because it is no longer tied to the \
stream's lifetime) well before any disconnect happened: \
{outcome:?}"
);
assert!(
matches!(outcome.post_drop, PostDropOutcome::CancelledAfterPolls(_)),
"should_cancel must return true within the poll budget after \
the test drops the response body -- exhausting the budget \
means the disconnect signal never reached the generator at \
all: {outcome:?}"
);
}
}
#[cfg(feature = "test-utils")]
mod streaming_context_overflow {
use super::*;
use lattice_inference::serve::{
OVERFLOW_PARITY_MAX_TOKENS_CAP, OVERFLOW_PARITY_REQUEST_BODY,
};
use tower::ServiceExt as _;
#[tokio::test]
async fn chat_completions_streaming_context_overflow_returns_400_before_committing_sse() {
let body = axum::body::Body::from(OVERFLOW_PARITY_REQUEST_BODY.to_string());
let request = axum::http::Request::builder()
.method("POST")
.uri("/v1/chat/completions")
.header("content-type", "application/json")
.body(body)
.expect("fixture request must build");
let response = router(tiny_state(OVERFLOW_PARITY_MAX_TOKENS_CAP))
.oneshot(request)
.await
.expect("router must produce a response, not a transport error");
assert_eq!(
response.status(),
StatusCode::BAD_REQUEST,
"an over-context stream:true request must be rejected with HTTP \
400 before any SSE stream is committed, matching \
lattice_serve.rs's equivalent preflight for the identical \
shared-fixture request body"
);
let bytes = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.expect("error response body must be readable");
let value: serde_json::Value =
serde_json::from_slice(&bytes).expect("error response must be JSON");
assert_eq!(value["error"]["code"], "context_length_exceeded");
}
}
#[cfg(feature = "test-utils")]
mod message_flood {
use super::*;
use lattice_inference::serve::contract::MAX_MESSAGE_COUNT;
use tower::ServiceExt as _;
#[tokio::test]
async fn chat_completions_rejects_message_flood() {
let messages: Vec<String> = (0..MAX_MESSAGE_COUNT + 1)
.map(|_| r#"{"role":"user","content":""}"#.to_string())
.collect();
let body = format!(
r#"{{"model":"test-model","messages":[{}]}}"#,
messages.join(",")
);
let request = axum::http::Request::builder()
.method("POST")
.uri("/v1/chat/completions")
.header("content-type", "application/json")
.body(axum::body::Body::from(body))
.expect("fixture request must build");
let response = router(tiny_state(64))
.oneshot(request)
.await
.expect("router must produce a response, not a transport error");
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
let bytes = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.expect("error response body must be readable");
let value: serde_json::Value =
serde_json::from_slice(&bytes).expect("error response must be JSON");
assert_eq!(value["error"]["code"], "invalid_request_body");
}
}
#[cfg(feature = "test-utils")]
mod content_type_precedence {
use super::*;
use tower::ServiceExt as _;
#[tokio::test]
async fn invalid_content_type_rejected_before_body_is_parsed() {
let request = axum::http::Request::builder()
.method("POST")
.uri("/v1/chat/completions")
.header("content-type", "text/plain")
.body(axum::body::Body::from("this is not json"))
.expect("fixture request must build");
let response = router(tiny_state(64))
.oneshot(request)
.await
.expect("router must produce a response, not a transport error");
assert_eq!(response.status(), StatusCode::UNSUPPORTED_MEDIA_TYPE);
let bytes = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.expect("error response body must be readable");
let value: serde_json::Value =
serde_json::from_slice(&bytes).expect("error response must be JSON");
assert_eq!(value["error"]["code"], "unsupported_media_type");
}
}
#[cfg(feature = "test-utils")]
mod parity_table {
use super::*;
use lattice_inference::StopReason;
use lattice_inference::serve::{
BASELINE_CANNED_COMPLETION_TOKENS, BASELINE_CANNED_PROMPT_TOKENS, BASELINE_CANNED_TEXT,
Binary, CHAT_COMPLETIONS_PARITY_CASES, ExpectedResponse, check_sse_events,
};
use tower::ServiceExt as _;
const CAP: usize = 64;
fn baseline_fake_state(max_tokens_cap: usize) -> AppState {
let model = lattice_inference::model::qwen35::test_support::tiny_zero_model();
#[allow(clippy::type_complexity)]
let generate: Arc<
dyn Fn(
&str,
&lattice_inference::model::qwen35_config::GenerateConfig,
&mut dyn FnMut(&str) -> bool,
&mut dyn FnMut() -> bool,
)
-> Result<GenerateOutput, lattice_inference::error::InferenceError>
+ Send
+ Sync,
> = Arc::new(|_prompt, _cfg, on_token, _should_cancel| {
for chunk in ["hello", " world"] {
if !on_token(chunk) {
break;
}
}
Ok(GenerateOutput {
text: BASELINE_CANNED_TEXT.to_string(),
token_ids: vec![1, 2],
prompt_tokens: BASELINE_CANNED_PROMPT_TOKENS as usize,
generated_tokens: BASELINE_CANNED_COMPLETION_TOKENS as usize,
stopped: true,
stop_reason: Some(StopReason::Eos),
token_logprobs: vec![],
})
});
AppState {
model: ModelBackend::CpuFakeGenerate {
model: Arc::new(model),
generate,
},
default_max_tokens: max_tokens_cap,
max_tokens_cap,
model_id: "test-model".to_string(),
request_counter: Arc::new(AtomicU64::new(0)),
embedding_model: None,
}
}
#[tokio::test]
async fn chat_completions_matches_shared_parity_table() {
for case in CHAT_COMPLETIONS_PARITY_CASES {
let expected = case.expected(Binary::Lattice);
let app = match expected {
ExpectedResponse::Error { .. } => router(tiny_state(CAP)),
ExpectedResponse::Json { .. } | ExpectedResponse::Sse { .. } => {
router(baseline_fake_state(CAP))
}
};
let request = axum::http::Request::builder()
.method(case.method)
.uri(case.path)
.header("content-type", "application/json")
.body(axum::body::Body::from(case.body.build()))
.expect("fixture request must build");
let response = app
.oneshot(request)
.await
.expect("router must produce a response, not a transport error");
let status = response.status().as_u16();
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.expect("response body reads");
let text = String::from_utf8_lossy(&body);
assert_eq!(
status,
expected.status(),
"case '{}': expected status {}, got {status} (body: {text})",
case.name,
expected.status(),
);
match expected {
ExpectedResponse::Error { code, .. } => {
let value: serde_json::Value = serde_json::from_slice(&body)
.unwrap_or_else(|e| {
panic!(
"case '{}': non-2xx response body must be the shared \
error envelope JSON: {e} (body: {text})",
case.name,
)
});
assert_eq!(
value["error"]["code"], code,
"case '{}': expected error code '{code}', got {} \
(full body: {value})",
case.name, value["error"]["code"]
);
}
ExpectedResponse::Json { fields, .. } => {
let value: serde_json::Value = serde_json::from_slice(&body)
.unwrap_or_else(|e| {
panic!(
"case '{}': 2xx response body must be JSON: {e} \
(body: {text})",
case.name,
)
});
for field in fields {
field.check(&value).unwrap_or_else(|e| {
panic!("case '{}': field check failed: {e}", case.name)
});
}
}
ExpectedResponse::Sse { events, .. } => {
check_sse_events(&text, events).unwrap_or_else(|e| {
panic!("case '{}': SSE check failed: {e}", case.name)
});
}
}
}
}
}
#[cfg(feature = "test-utils")]
mod production_adapter_observation {
use super::*;
use lattice_inference::serve::{
ExpectedObservation, GenerateConfigSnapshot, OBSERVATION_GOLDEN_USER_HI_THERE_CHATML,
ProductionAdapterObservation, assert_observation_matches,
};
use std::sync::Mutex;
use tower::ServiceExt as _;
const OBSERVATION_GOLDEN_REQUEST_BODY: &str = r#"{"model":"test-model","messages":[{"role":"user","content":"hi there"}],"temperature":1.3,"top_p":0.55,"seed":7,"max_tokens":9}"#;
async fn run_observed(stopped: bool, body: &str) -> ProductionAdapterObservation {
let model = lattice_inference::model::qwen35::test_support::tiny_zero_model();
let tokenizer = model.tokenizer().clone();
let observed: Arc<Mutex<Option<ProductionAdapterObservation>>> =
Arc::new(Mutex::new(None));
let observed_for_closure = Arc::clone(&observed);
#[allow(clippy::type_complexity)]
let generate: Arc<
dyn Fn(
&str,
&lattice_inference::model::qwen35_config::GenerateConfig,
&mut dyn FnMut(&str) -> bool,
&mut dyn FnMut() -> bool,
)
-> Result<GenerateOutput, lattice_inference::error::InferenceError>
+ Send
+ Sync,
> = Arc::new(move |prompt, cfg, _on_token, _should_cancel| {
let prompt_tokens = tokenizer.tokenize(prompt).real_length;
*observed_for_closure
.lock()
.expect("observation mutex poisoned") = Some(ProductionAdapterObservation {
rendered_prompt: Some(prompt.to_string()),
messages: None,
gen_cfg: GenerateConfigSnapshot::from(cfg),
prompt_tokens,
stopped,
});
Ok(GenerateOutput {
text: "ok".to_string(),
token_ids: vec![1],
prompt_tokens,
generated_tokens: 1,
stopped,
stop_reason: if stopped {
Some(lattice_inference::StopReason::Eos)
} else {
None
},
token_logprobs: vec![],
})
});
let state = AppState {
model: ModelBackend::CpuFakeGenerate {
model: Arc::new(model),
generate,
},
default_max_tokens: 64,
max_tokens_cap: 64,
model_id: "test-model".to_string(),
request_counter: Arc::new(AtomicU64::new(0)),
embedding_model: None,
};
let request = axum::http::Request::builder()
.method("POST")
.uri("/v1/chat/completions")
.header("content-type", "application/json")
.body(axum::body::Body::from(body.to_string()))
.expect("fixture request must build");
let response = router(state)
.oneshot(request)
.await
.expect("router must produce a response, not a transport error");
assert_eq!(response.status(), StatusCode::OK);
observed
.lock()
.expect("observation mutex poisoned")
.clone()
.expect("the injected generate closure must have recorded an observation")
}
fn expected_gen_cfg() -> GenerateConfigSnapshot {
GenerateConfigSnapshot::from(&lattice_inference::model::qwen35_config::GenerateConfig {
max_new_tokens: 9,
temperature: 1.3,
top_p: 0.55,
seed: Some(7),
stop_strings: vec![],
logprobs: None,
..Default::default()
})
}
#[tokio::test]
async fn chat_completions_non_streaming_observation_captures_real_config_and_prompt() {
let obs = run_observed(true, OBSERVATION_GOLDEN_REQUEST_BODY).await;
let expected_prompt_tokens = {
let tokenizer = lattice_inference::model::qwen35::test_support::tiny_zero_model()
.tokenizer()
.clone();
tokenizer
.tokenize(OBSERVATION_GOLDEN_USER_HI_THERE_CHATML)
.real_length
};
assert_observation_matches(
&obs,
&ExpectedObservation {
gen_cfg: expected_gen_cfg(),
rendered_prompt: Some(OBSERVATION_GOLDEN_USER_HI_THERE_CHATML),
messages: None,
prompt_tokens: expected_prompt_tokens,
stopped: true,
},
);
}
#[tokio::test]
async fn chat_completions_non_streaming_observation_captures_real_stopped_false() {
let obs = run_observed(false, OBSERVATION_GOLDEN_REQUEST_BODY).await;
assert!(
!obs.stopped,
"observation must report the seam's actual stopped=false, not a hardcoded true"
);
}
#[tokio::test]
async fn chat_completions_non_streaming_observation_captures_real_reasoning_budget() {
let body = r#"{"model":"test-model","messages":[{"role":"user","content":"hi there"}],"temperature":1.3,"top_p":0.55,"seed":7,"max_tokens":9,"reasoning_budget":5}"#;
let obs = run_observed(true, body).await;
assert_eq!(
obs.gen_cfg.reasoning_budget,
Some(5),
"reasoning_budget from the request must reach the real GenerateConfig"
);
}
}
#[cfg(feature = "test-utils")]
mod completion_length_invariant_1334 {
use super::*;
use http_body_util::BodyExt as _;
use tower::ServiceExt as _;
fn state_returning(max_tokens_cap: usize, generated_tokens: usize) -> AppState {
let model = lattice_inference::model::qwen35::test_support::tiny_zero_model();
AppState {
model: ModelBackend::CpuFakeGenerate {
model: Arc::new(model),
generate: Arc::new(move |_prompt, _cfg, _on_token, _should_cancel| {
Ok(GenerateOutput {
text: "x".repeat(generated_tokens),
token_ids: vec![0; generated_tokens],
prompt_tokens: 1,
generated_tokens,
stopped: true,
stop_reason: Some(lattice_inference::StopReason::Length),
token_logprobs: vec![],
})
}),
},
default_max_tokens: max_tokens_cap,
max_tokens_cap,
model_id: "test-model".to_string(),
request_counter: Arc::new(AtomicU64::new(0)),
embedding_model: None,
}
}
fn request_body(
max_tokens: usize,
reasoning_budget: Option<usize>,
stream: bool,
) -> String {
let mut body = format!(
r#"{{"model":"test-model","messages":[{{"role":"user","content":"hi"}}],"max_tokens":{max_tokens}"#
);
if let Some(budget) = reasoning_budget {
body.push_str(&format!(r#","reasoning_budget":{budget}"#));
}
if stream {
body.push_str(r#","stream":true"#);
}
body.push('}');
body
}
async fn post(state: AppState, body: String) -> axum::http::Response<axum::body::Body> {
let request = axum::http::Request::builder()
.method("POST")
.uri("/v1/chat/completions")
.header("content-type", "application/json")
.body(axum::body::Body::from(body))
.expect("fixture request must build");
router(state)
.oneshot(request)
.await
.expect("router must produce a response, not a transport error")
}
const MAX_TOKENS: usize = 10;
const REASONING_BUDGET: usize = 4;
async fn assert_non_streaming(
reasoning_budget: Option<usize>,
generated_tokens: usize,
expect_success: bool,
case: &str,
) {
let state = state_returning(64, generated_tokens);
let body = request_body(MAX_TOKENS, reasoning_budget, false);
let response = post(state, body).await;
if expect_success {
assert_eq!(
response.status(),
StatusCode::OK,
"{case}: generated_tokens={generated_tokens} at/under the decode budget \
must be accepted, not reported as an inference failure"
);
} else {
assert_eq!(
response.status(),
StatusCode::INTERNAL_SERVER_ERROR,
"{case}: generated_tokens={generated_tokens}, one past the decode budget, \
must still be rejected by the invariant"
);
}
}
#[tokio::test]
async fn non_streaming_exact_acceptance_positive_budget() {
assert_non_streaming(
Some(REASONING_BUDGET),
MAX_TOKENS + REASONING_BUDGET + 1,
true,
"non-streaming, positive reasoning_budget, exact budget",
)
.await;
}
#[tokio::test]
async fn non_streaming_one_past_rejection_positive_budget() {
assert_non_streaming(
Some(REASONING_BUDGET),
MAX_TOKENS + REASONING_BUDGET + 2,
false,
"non-streaming, positive reasoning_budget, one past budget",
)
.await;
}
#[tokio::test]
async fn non_streaming_exact_acceptance_unset_budget() {
assert_non_streaming(
None,
MAX_TOKENS,
true,
"non-streaming, unset reasoning_budget, exact budget",
)
.await;
}
#[tokio::test]
async fn non_streaming_one_past_rejection_unset_budget() {
assert_non_streaming(
None,
MAX_TOKENS + 1,
false,
"non-streaming, unset reasoning_budget, one past budget",
)
.await;
}
async fn assert_streaming(
reasoning_budget: Option<usize>,
generated_tokens: usize,
expect_success: bool,
case: &str,
) {
let state = state_returning(64, generated_tokens);
let body = request_body(MAX_TOKENS, reasoning_budget, true);
let response = post(state, body).await;
assert_eq!(
response.status(),
StatusCode::OK,
"{case}: an SSE response commits with 200 regardless of how generation \
concludes -- failure surfaces as an in-stream error event"
);
let bytes = response
.into_body()
.collect()
.await
.expect("SSE response body must be readable")
.to_bytes();
let text = String::from_utf8(bytes.to_vec()).expect("SSE body must be valid UTF-8");
let has_error_event = text
.lines()
.filter_map(|line| line.strip_prefix("data: "))
.any(|payload| payload.contains("\"error\""));
let has_finish_reason = text.contains("\"finish_reason\":\"stop\"");
if expect_success {
assert!(
!has_error_event,
"{case}: generated_tokens={generated_tokens} at/under the decode budget \
must not emit an SSE error event; got: {text}"
);
assert!(
has_finish_reason,
"{case}: a successful completion must emit finish_reason \"stop\"; \
got: {text}"
);
} else {
assert!(
has_error_event,
"{case}: generated_tokens={generated_tokens}, one past the decode budget, \
must still emit an SSE error event; got: {text}"
);
assert!(
!has_finish_reason,
"{case}: an invariant violation must not also emit a clean finish_reason; \
got: {text}"
);
}
}
#[tokio::test]
async fn streaming_exact_acceptance_positive_budget() {
assert_streaming(
Some(REASONING_BUDGET),
MAX_TOKENS + REASONING_BUDGET + 1,
true,
"streaming, positive reasoning_budget, exact budget",
)
.await;
}
#[tokio::test]
async fn streaming_one_past_rejection_positive_budget() {
assert_streaming(
Some(REASONING_BUDGET),
MAX_TOKENS + REASONING_BUDGET + 2,
false,
"streaming, positive reasoning_budget, one past budget",
)
.await;
}
#[tokio::test]
async fn streaming_exact_acceptance_unset_budget() {
assert_streaming(
None,
MAX_TOKENS,
true,
"streaming, unset reasoning_budget, exact budget",
)
.await;
}
#[tokio::test]
async fn streaming_one_past_rejection_unset_budget() {
assert_streaming(
None,
MAX_TOKENS + 1,
false,
"streaming, unset reasoning_budget, one past budget",
)
.await;
}
}
}