use serde::Deserialize;
use serde_json::Value;
use serde_json::value::RawValue;
use super::{ApiError, REQUEST_BODY_LIMIT_BYTES};
pub const MAX_MESSAGE_COUNT: usize = 4096;
pub const MAX_CUMULATIVE_CONTENT_BYTES: usize = REQUEST_BODY_LIMIT_BYTES;
pub const MAX_STOP_STRING_BYTES: usize = 4096;
pub const MAX_CUMULATIVE_STOP_BYTES: usize = 2 * MAX_STOP_STRING_BYTES;
#[derive(Debug, Deserialize)]
pub struct ChatRequest {
#[serde(default)]
pub model: Option<String>,
#[serde(default, deserialize_with = "deserialize_bounded_messages")]
pub messages: Vec<Message>,
#[serde(default)]
pub max_tokens: Option<usize>,
#[serde(default)]
pub max_completion_tokens: Option<usize>,
#[serde(default)]
pub temperature: Option<f32>,
#[serde(default)]
pub top_p: Option<f32>,
#[serde(default)]
pub top_k: Option<Box<RawValue>>,
#[serde(default)]
pub repetition_penalty: Option<Box<RawValue>>,
#[serde(default)]
pub seed: Option<u64>,
#[serde(default)]
pub stream: Option<bool>,
#[serde(default)]
pub stop: Option<Value>,
#[serde(default)]
pub reasoning_budget: Option<Box<RawValue>>,
#[serde(default)]
pub response_format: Option<ResponseFormat>,
#[serde(default)]
pub tools: Option<Value>,
#[serde(default)]
pub tool_choice: Option<Value>,
#[serde(default)]
pub logprobs: Option<bool>,
#[serde(default)]
pub top_logprobs: Option<usize>,
#[serde(default)]
pub n: Option<usize>,
}
#[derive(Debug, Deserialize)]
pub struct Message {
pub role: String,
pub content: MessageContent,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum NormalizedChatRole {
System,
User,
Assistant,
}
impl NormalizedChatRole {
pub(crate) const fn as_str(self) -> &'static str {
match self {
Self::System => "system",
Self::User => "user",
Self::Assistant => "assistant",
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct NormalizedChatMessage {
pub role: NormalizedChatRole,
pub content: String,
}
#[derive(Debug, Clone, Deserialize, PartialEq, Eq)]
#[serde(untagged)]
pub enum MessageContent {
Text(String),
Parts(Vec<ContentPart>),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ContentPart {
Text { text: String },
ImageUrl { image_url: ImageUrl },
Unsupported { kind: String },
}
#[derive(Debug, Clone, Deserialize, PartialEq, Eq)]
pub struct ImageUrl {
pub url: String,
#[serde(default)]
pub detail: Option<String>,
}
#[derive(Deserialize)]
struct RawContentPart {
#[serde(rename = "type")]
kind: Option<String>,
#[serde(default)]
text: Option<String>,
#[serde(default)]
image_url: Option<ImageUrl>,
}
impl<'de> Deserialize<'de> for ContentPart {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
let raw = RawContentPart::deserialize(deserializer)?;
match raw.kind.as_deref() {
Some("text") => raw.text.map(|text| Self::Text { text }).ok_or_else(|| {
serde::de::Error::custom("text content part must include string field 'text'")
}),
Some("image_url") => raw
.image_url
.map(|image_url| Self::ImageUrl { image_url })
.ok_or_else(|| {
serde::de::Error::custom(
"image_url content part must include object field 'image_url'",
)
}),
Some(kind) => Ok(Self::Unsupported {
kind: kind.to_string(),
}),
None => Ok(Self::Unsupported {
kind: "<missing>".to_string(),
}),
}
}
}
#[derive(Debug, Deserialize)]
pub struct ResponseFormat {
pub r#type: String,
#[serde(default)]
pub json_schema: Option<JsonSchemaFormat>,
}
#[derive(Debug, Deserialize)]
pub struct JsonSchemaFormat {
#[serde(default)]
pub name: Option<String>,
#[serde(default)]
pub strict: Option<bool>,
#[serde(default)]
pub schema: Option<Box<RawValue>>,
}
#[derive(Debug, Clone, Copy)]
pub struct GenerationDefaults {
pub max_tokens: usize,
pub temperature: f32,
pub top_k: usize,
pub top_p: f32,
pub repetition_penalty: f32,
pub reasoning_budget: Option<usize>,
}
impl GenerationDefaults {
pub const fn standard(max_tokens: usize) -> Self {
Self {
max_tokens,
temperature: 0.7,
top_k: 50,
top_p: 0.9,
repetition_penalty: 1.1,
reasoning_budget: None,
}
}
}
#[derive(Debug, Clone, Copy)]
enum ModelNamePolicy<'a> {
RequiredExact(&'a str),
OptionalExact(&'a str),
}
#[derive(Debug, Clone, Copy)]
enum MaxTokensPolicy {
RejectAbove { limit: usize },
ClampToContext { context: usize },
}
#[derive(Debug, Clone, Copy)]
pub struct ServeProfile<'a> {
model_name: ModelNamePolicy<'a>,
max_tokens: MaxTokensPolicy,
require_last_user: bool,
stop_supported: bool,
sampling_extensions_supported: bool,
reasoning_budget_supported: bool,
structured_output_supported: bool,
logprobs_supported: bool,
max_tokens_conflict_checked_early: bool,
}
impl<'a> ServeProfile<'a> {
pub fn lattice(model_id: &'a str, max_tokens_cap: usize) -> Self {
Self {
model_name: ModelNamePolicy::RequiredExact(model_id),
max_tokens: MaxTokensPolicy::RejectAbove {
limit: max_tokens_cap,
},
require_last_user: true,
stop_supported: true,
sampling_extensions_supported: false,
reasoning_budget_supported: false,
structured_output_supported: false,
logprobs_supported: true,
max_tokens_conflict_checked_early: false,
}
}
pub fn lattice_serve(model_id: &'a str, model_max_context: usize) -> Self {
Self {
model_name: ModelNamePolicy::OptionalExact(model_id),
max_tokens: MaxTokensPolicy::ClampToContext {
context: model_max_context,
},
require_last_user: false,
stop_supported: false,
sampling_extensions_supported: true,
reasoning_budget_supported: true,
structured_output_supported: true,
logprobs_supported: false,
max_tokens_conflict_checked_early: true,
}
}
}
#[derive(Debug)]
pub struct ValidatedChatRequest {
pub messages: Vec<NormalizedChatMessage>,
pub max_tokens: usize,
pub temperature: f32,
pub top_k: usize,
pub top_p: f32,
pub repetition_penalty: f32,
pub seed: Option<u64>,
pub stream: bool,
pub stop_strings: Vec<String>,
pub reasoning_budget: Option<usize>,
pub logprobs: Option<usize>,
}
pub fn normalize_request(
req: &ChatRequest,
defaults: GenerationDefaults,
profile: ServeProfile<'_>,
) -> Result<ValidatedChatRequest, ApiError> {
normalize_request_inner(req, defaults, profile, |_, _| Ok(())).map(|(validated, ())| validated)
}
pub fn normalize_request_with_context(
req: &ChatRequest,
defaults: GenerationDefaults,
profile: ServeProfile<'_>,
check_context: impl FnOnce(&[NormalizedChatMessage], usize) -> Result<String, ApiError>,
) -> Result<(ValidatedChatRequest, String), ApiError> {
normalize_request_inner(req, defaults, profile, check_context)
}
fn normalize_request_inner<C>(
req: &ChatRequest,
defaults: GenerationDefaults,
profile: ServeProfile<'_>,
check_context: impl FnOnce(&[NormalizedChatMessage], usize) -> Result<C, ApiError>,
) -> Result<(ValidatedChatRequest, C), ApiError> {
reject_unsupported(req, profile)?;
validate_model_name(req.model.as_deref(), profile.model_name)?;
if req.messages.is_empty() {
return Err(ApiError::BadRequest {
message: "messages must not be empty".to_string(),
code: "invalid_messages",
});
}
check_message_bounds(&req.messages)?;
if profile.require_last_user
&& req.messages.last().map(|message| message.role.as_str()) != Some("user")
{
return Err(ApiError::BadRequest {
message: "the last message must have role 'user'".to_string(),
code: "invalid_messages",
});
}
let max_tokens = normalize_max_tokens(
req,
defaults.max_tokens,
profile.max_tokens,
profile.max_tokens_conflict_checked_early,
)?;
let temperature = validate_temperature(req.temperature.unwrap_or(defaults.temperature))?;
let top_p = validate_top_p(req.top_p.unwrap_or(defaults.top_p))?;
let logprobs = normalize_logprobs(req)?;
let messages = normalize_messages(&req.messages)?;
let context = check_context(&messages, max_tokens)?;
let stop_strings = if profile.stop_supported {
parse_stop_strings(&req.stop)?
} else {
Vec::new()
};
let top_k = if profile.sampling_extensions_supported {
parse_ignorable_field::<usize>(&req.top_k, "top_k")?.unwrap_or(defaults.top_k)
} else {
defaults.top_k
};
let repetition_penalty = if profile.sampling_extensions_supported {
parse_ignorable_field::<f32>(&req.repetition_penalty, "repetition_penalty")?
.unwrap_or(defaults.repetition_penalty)
} else {
defaults.repetition_penalty
};
let mut reasoning_budget = if profile.reasoning_budget_supported {
parse_ignorable_field::<usize>(&req.reasoning_budget, "reasoning_budget")?
.filter(|&value| value > 0)
.or(defaults.reasoning_budget)
} else {
None
};
if let MaxTokensPolicy::ClampToContext { context } = profile.max_tokens {
let reasoning_room = context.saturating_sub(max_tokens).saturating_sub(1);
reasoning_budget = reasoning_budget
.map(|value| value.min(reasoning_room))
.filter(|&value| value > 0);
}
Ok((
ValidatedChatRequest {
messages,
max_tokens,
temperature,
top_k,
top_p,
repetition_penalty,
seed: req.seed,
stream: req.stream.unwrap_or(false),
stop_strings,
reasoning_budget,
logprobs,
},
context,
))
}
fn reject_unsupported(req: &ChatRequest, profile: ServeProfile<'_>) -> Result<(), ApiError> {
if req.tools.is_some() || req.tool_choice.is_some() {
return unsupported("tools and tool_choice are not supported by this server");
}
if req.n.unwrap_or(1) > 1 {
return unsupported("n > 1 is not supported");
}
if req.stop.is_some() && !profile.stop_supported {
return unsupported("stop is not supported by this server");
}
if !profile.logprobs_supported && (req.logprobs.unwrap_or(false) || req.top_logprobs.is_some())
{
return unsupported("logprobs/top_logprobs are not supported by this server");
}
if req.stream == Some(true) && req.logprobs.unwrap_or(false) {
return unsupported("logprobs is not supported together with stream: true");
}
if profile.max_tokens_conflict_checked_early {
reject_conflicting_max_tokens(req)?;
}
if let Some(format) = &req.response_format
&& format.r#type != "text"
&& !(profile.structured_output_supported && format.r#type == "json_schema")
{
return unsupported(format!(
"response_format.type '{}' is not supported; use 'text'",
format.r#type
));
}
Ok(())
}
fn reject_conflicting_max_tokens(req: &ChatRequest) -> Result<(), ApiError> {
if let (Some(max_tokens), Some(max_completion_tokens)) =
(req.max_tokens, req.max_completion_tokens)
&& max_tokens != max_completion_tokens
{
return Err(ApiError::BadRequest {
message: format!(
"max_tokens ({max_tokens}) and max_completion_tokens ({max_completion_tokens}) differ; supply only one"
),
code: "invalid_request",
});
}
Ok(())
}
const MESSAGE_FLOOD_SENTINEL: &str = "lattice_message_flood_exceeded";
fn deserialize_bounded_messages<'de, D>(deserializer: D) -> Result<Vec<Message>, D::Error>
where
D: serde::Deserializer<'de>,
{
struct Visitor;
impl<'de> serde::de::Visitor<'de> for Visitor {
type Value = Vec<Message>;
fn expecting(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("an array of messages")
}
fn visit_seq<A>(self, mut seq: A) -> Result<Self::Value, A::Error>
where
A: serde::de::SeqAccess<'de>,
{
let mut messages = Vec::new();
while let Some(message) = seq.next_element::<Message>()? {
messages.push(message);
if messages.len() > MAX_MESSAGE_COUNT {
return Err(serde::de::Error::custom(MESSAGE_FLOOD_SENTINEL));
}
}
Ok(messages)
}
}
deserializer.deserialize_seq(Visitor)
}
pub fn is_message_flood_error(err: &serde_json::Error) -> bool {
err.to_string().contains(MESSAGE_FLOOD_SENTINEL)
}
pub fn message_flood_text() -> String {
format!("messages has more than {MAX_MESSAGE_COUNT} entries; maximum is {MAX_MESSAGE_COUNT}")
}
fn check_message_bounds(messages: &[Message]) -> Result<(), ApiError> {
if messages.len() > MAX_MESSAGE_COUNT {
return Err(ApiError::BadRequest {
message: format!(
"messages has {} entries; maximum is {MAX_MESSAGE_COUNT}",
messages.len()
),
code: "invalid_request_body",
});
}
let cumulative: usize = messages
.iter()
.map(|message| content_byte_len(&message.content))
.sum();
if cumulative > MAX_CUMULATIVE_CONTENT_BYTES {
return Err(ApiError::BadRequest {
message: format!(
"messages content totals {cumulative} bytes; maximum is {MAX_CUMULATIVE_CONTENT_BYTES}"
),
code: "invalid_request_body",
});
}
Ok(())
}
fn content_byte_len(content: &MessageContent) -> usize {
match content {
MessageContent::Text(text) => text.len(),
MessageContent::Parts(parts) => parts
.iter()
.map(|part| match part {
ContentPart::Text { text } => text.len(),
ContentPart::ImageUrl { image_url } => image_url.url.len(),
ContentPart::Unsupported { kind } => kind.len(),
})
.sum(),
}
}
fn parse_ignorable_field<T: serde::de::DeserializeOwned>(
value: &Option<Box<RawValue>>,
field_name: &str,
) -> Result<Option<T>, ApiError> {
let Some(value) = value else {
return Ok(None);
};
let text = value.get().trim();
if text == "null" {
return Ok(None);
}
if text.starts_with('{') || text.starts_with('[') {
return Err(ApiError::BadRequest {
message: format!("{field_name} must be a scalar value"),
code: "invalid_request_body",
});
}
serde_json::from_str(text)
.map(Some)
.map_err(|err| ApiError::BadRequest {
message: format!("{field_name} is invalid: {err}"),
code: "invalid_request_body",
})
}
fn unsupported(message: impl Into<String>) -> Result<(), ApiError> {
Err(ApiError::BadRequest {
message: message.into(),
code: "unsupported_feature",
})
}
fn validate_model_name(
requested: Option<&str>,
policy: ModelNamePolicy<'_>,
) -> Result<(), ApiError> {
let expected = match policy {
ModelNamePolicy::RequiredExact(expected) => match requested {
None | Some("") => {
return Err(ApiError::BadRequest {
message: "model is required".to_string(),
code: "invalid_request",
});
}
Some(requested) if requested == expected => return Ok(()),
Some(_) => expected,
},
ModelNamePolicy::OptionalExact(_) if requested.is_none() => return Ok(()),
ModelNamePolicy::OptionalExact(expected) if requested == Some(expected) => return Ok(()),
ModelNamePolicy::OptionalExact(expected) => expected,
};
let requested = requested.unwrap_or_default();
Err(ApiError::BadRequest {
message: format!("model '{requested}' is not loaded; this server serves '{expected}'"),
code: "model_not_found",
})
}
fn normalize_max_tokens(
req: &ChatRequest,
default_max_tokens: usize,
policy: MaxTokensPolicy,
conflict_checked_early: bool,
) -> Result<usize, ApiError> {
if !conflict_checked_early {
reject_conflicting_max_tokens(req)?;
}
let requested = match (req.max_tokens, req.max_completion_tokens) {
(None, None) => default_max_tokens,
(Some(value), None) | (None, Some(value)) => value,
(Some(left), Some(_right)) => left,
};
super::reject_zero_max_tokens(requested)?;
match policy {
MaxTokensPolicy::RejectAbove { limit } if requested > limit => Err(ApiError::BadRequest {
message: format!("max_tokens {requested} exceeds server limit {limit}"),
code: "max_tokens_exceeds_limit",
}),
MaxTokensPolicy::RejectAbove { .. } => Ok(requested),
MaxTokensPolicy::ClampToContext { context } => Ok(requested.min(context.saturating_sub(1))),
}
}
pub fn validate_temperature(temperature: f32) -> Result<f32, ApiError> {
if !(0.0..=2.0).contains(&temperature) {
return Err(ApiError::BadRequest {
message: "temperature must be between 0 and 2".to_string(),
code: "invalid_temperature",
});
}
Ok(temperature)
}
pub fn validate_top_p(top_p: f32) -> Result<f32, ApiError> {
if !(top_p > 0.0 && top_p <= 1.0) {
return Err(ApiError::BadRequest {
message: "top_p must be greater than 0 and at most 1".to_string(),
code: "invalid_top_p",
});
}
Ok(top_p)
}
fn normalize_logprobs(req: &ChatRequest) -> Result<Option<usize>, ApiError> {
if !req.logprobs.unwrap_or(false) {
if req.top_logprobs.is_some() {
return Err(ApiError::BadRequest {
message: "top_logprobs requires logprobs: true".to_string(),
code: "invalid_request",
});
}
return Ok(None);
}
let top_logprobs = req.top_logprobs.unwrap_or(0);
if top_logprobs > 20 {
return Err(ApiError::BadRequest {
message: format!("top_logprobs {top_logprobs} exceeds the maximum of 20"),
code: "invalid_top_logprobs",
});
}
Ok(Some(top_logprobs))
}
pub fn normalize_messages(messages: &[Message]) -> Result<Vec<NormalizedChatMessage>, ApiError> {
messages
.iter()
.map(|message| {
let content = message_text(&message.content)?;
match message.role.as_str() {
"system" => Ok(NormalizedChatMessage {
role: NormalizedChatRole::System,
content,
}),
"user" => Ok(NormalizedChatMessage {
role: NormalizedChatRole::User,
content,
}),
"assistant" => Ok(NormalizedChatMessage {
role: NormalizedChatRole::Assistant,
content,
}),
"tool" | "developer" => Err(ApiError::BadRequest {
message: format!("role '{}' is not supported by this server", message.role),
code: "unsupported_feature",
}),
role => Err(ApiError::BadRequest {
message: format!(
"unsupported role '{role}'; must be 'system', 'user', or 'assistant'"
),
code: "invalid_role",
}),
}
})
.collect()
}
fn message_text(content: &MessageContent) -> Result<String, ApiError> {
match content {
MessageContent::Text(text) => Ok(text.clone()),
MessageContent::Parts(parts) => {
let mut output = String::new();
for part in parts {
match part {
ContentPart::Text { text } => output.push_str(text),
ContentPart::ImageUrl { .. } => {
return Err(ApiError::BadRequest {
message: "image input requires a vision-capable model".to_string(),
code: "unsupported_feature",
});
}
ContentPart::Unsupported { kind } => {
return Err(ApiError::BadRequest {
message: format!(
"content part type '{kind}' is not supported; only 'text' parts are accepted"
),
code: "unsupported_feature",
});
}
}
}
Ok(output)
}
}
}
fn check_stop_string_bytes(value: &str) -> Result<(), ApiError> {
if value.len() > MAX_STOP_STRING_BYTES {
return Err(ApiError::BadRequest {
message: format!(
"stop string has {} bytes; maximum is {MAX_STOP_STRING_BYTES}",
value.len()
),
code: "invalid_request_body",
});
}
Ok(())
}
pub fn parse_stop_strings(stop: &Option<Value>) -> Result<Vec<String>, ApiError> {
match stop {
None | Some(Value::Null) => Ok(Vec::new()),
Some(Value::String(value)) if value.is_empty() => Err(ApiError::BadRequest {
message: "stop string must not be empty".to_string(),
code: "invalid_stop",
}),
Some(Value::String(value)) => {
check_stop_string_bytes(value)?;
Ok(vec![value.clone()])
}
Some(Value::Array(values)) if values.is_empty() => Err(ApiError::BadRequest {
message: "stop array must not be empty".to_string(),
code: "invalid_stop",
}),
Some(Value::Array(values)) if values.len() > 4 => Err(ApiError::BadRequest {
message: format!("stop array has {} elements; maximum is 4", values.len()),
code: "invalid_stop",
}),
Some(Value::Array(values)) => {
let stops: Vec<String> = values
.iter()
.map(|value| match value {
Value::String(value) if value.is_empty() => Err(ApiError::BadRequest {
message: "stop string must not be empty".to_string(),
code: "invalid_stop",
}),
Value::String(value) => {
check_stop_string_bytes(value)?;
Ok(value.clone())
}
_ => Err(ApiError::BadRequest {
message: "each element of stop must be a string".to_string(),
code: "invalid_stop",
}),
})
.collect::<Result<_, _>>()?;
let cumulative: usize = stops.iter().map(String::len).sum();
if cumulative > MAX_CUMULATIVE_STOP_BYTES {
return Err(ApiError::BadRequest {
message: format!(
"stop strings total {cumulative} bytes; maximum is {MAX_CUMULATIVE_STOP_BYTES}"
),
code: "invalid_request_body",
});
}
Ok(stops)
}
Some(_) => Err(ApiError::BadRequest {
message: "stop must be a string or array of strings".to_string(),
code: "invalid_stop",
}),
}
}
pub fn validate_context_window(
prompt_tokens: usize,
max_tokens: usize,
max_context: usize,
) -> Result<(), ApiError> {
if prompt_tokens == 0 || prompt_tokens.saturating_add(max_tokens) > max_context {
return Err(ApiError::BadRequest {
message: format!(
"prompt ({prompt_tokens} tokens) plus max_tokens ({max_tokens}) exceeds model context window ({max_context})"
),
code: "context_length_exceeded",
});
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
fn request(body: &str) -> ChatRequest {
serde_json::from_str(body).unwrap()
}
fn defaults() -> GenerationDefaults {
GenerationDefaults::standard(16)
}
#[test]
fn standard_generation_defaults_snapshot() {
let defaults = GenerationDefaults::standard(123);
assert_eq!(defaults.max_tokens, 123);
assert_eq!(defaults.temperature, 0.7);
assert_eq!(defaults.top_k, 50);
assert_eq!(defaults.top_p, 0.9);
assert_eq!(defaults.repetition_penalty, 1.1);
assert_eq!(defaults.reasoning_budget, None);
}
#[test]
fn contract_source_has_no_backend_message_dependency() {
let source = include_str!("contract.rs");
assert!(!source.contains(concat!("crate::", "forward")));
}
#[test]
fn serving_adapters_do_not_restate_standard_defaults() {
for (name, source) in [
("lattice", include_str!("../bin/lattice.rs")),
("lattice_serve", include_str!("../bin/lattice_serve.rs")),
] {
for literal in [
"temperature: 0.7",
"top_k: 50",
"top_p: 0.9",
"repetition_penalty: 1.1",
"unwrap_or(512)",
"unwrap_or(0.7)",
"unwrap_or(50)",
"unwrap_or(0.9)",
"unwrap_or(1.1)",
] {
assert!(
!source.contains(literal),
"{name} restates canonical generation default {literal}"
);
}
}
}
#[test]
fn shared_sampling_bounds_reject_invalid_values() {
assert_eq!(
validate_temperature(-0.1).unwrap_err().code(),
"invalid_temperature"
);
assert_eq!(
validate_temperature(2.1).unwrap_err().code(),
"invalid_temperature"
);
assert_eq!(validate_top_p(0.0).unwrap_err().code(), "invalid_top_p");
assert_eq!(validate_top_p(1.1).unwrap_err().code(), "invalid_top_p");
}
#[test]
fn both_profiles_use_shared_sampling_bounds() {
let req = request(
r#"{"model":"model","messages":[{"role":"user","content":"hi"}],"temperature":2.5}"#,
);
for profile in [
ServeProfile::lattice("model", 32),
ServeProfile::lattice_serve("model", 32),
] {
assert_eq!(
normalize_request(&req, defaults(), profile)
.unwrap_err()
.code(),
"invalid_temperature"
);
}
}
#[test]
fn profiles_preserve_stop_and_max_token_policy() {
let req = request(
r#"{"model":"model","messages":[{"role":"user","content":"hi"}],"max_tokens":31,"stop":"done"}"#,
);
let lattice =
normalize_request(&req, defaults(), ServeProfile::lattice("model", 32)).unwrap();
assert_eq!(lattice.max_tokens, 31);
assert_eq!(lattice.stop_strings, ["done"]);
assert_eq!(
normalize_request(&req, defaults(), ServeProfile::lattice_serve("model", 16),)
.unwrap_err()
.code(),
"unsupported_feature"
);
let req = request(
r#"{"model":"model","messages":[{"role":"user","content":"hi"}],"max_tokens":31}"#,
);
let daemon =
normalize_request(&req, defaults(), ServeProfile::lattice_serve("model", 16)).unwrap();
assert_eq!(daemon.max_tokens, 15);
}
#[test]
fn profiles_preserve_sampling_extension_policy() {
let req = request(
r#"{"model":"model","messages":[{"role":"user","content":"hi"}],"top_k":20,"repetition_penalty":1.2}"#,
);
let lattice =
normalize_request(&req, defaults(), ServeProfile::lattice("model", 32)).unwrap();
assert_eq!(lattice.top_k, defaults().top_k);
assert_eq!(lattice.repetition_penalty, defaults().repetition_penalty);
let daemon =
normalize_request(&req, defaults(), ServeProfile::lattice_serve("model", 32)).unwrap();
assert_eq!(daemon.top_k, 20);
assert_eq!(daemon.repetition_penalty, 1.2);
}
#[test]
fn lattice_profile_accepts_and_ignores_reasoning_budget() {
let req = request(
r#"{"model":"model","messages":[{"role":"user","content":"hi"}],"reasoning_budget":40}"#,
);
let lattice =
normalize_request(&req, defaults(), ServeProfile::lattice("model", 32)).unwrap();
assert_eq!(lattice.reasoning_budget, None);
let daemon =
normalize_request(&req, defaults(), ServeProfile::lattice_serve("model", 100)).unwrap();
assert_eq!(daemon.reasoning_budget, Some(40));
}
#[test]
fn lattice_profile_tolerates_malformed_ignored_sampling_fields() {
let req = request(
r#"{"model":"model","messages":[{"role":"user","content":"hi"}],"top_k":"ignored","repetition_penalty":"x","reasoning_budget":"y"}"#,
);
let lattice =
normalize_request(&req, defaults(), ServeProfile::lattice("model", 32)).unwrap();
assert_eq!(lattice.top_k, defaults().top_k);
assert_eq!(lattice.repetition_penalty, defaults().repetition_penalty);
assert_eq!(lattice.reasoning_budget, None);
}
#[test]
fn honoring_profile_still_rejects_malformed_sampling_fields() {
for (field, bad_value) in [
("top_k", r#""ignored""#),
("repetition_penalty", r#""x""#),
("reasoning_budget", r#""y""#),
] {
let body = format!(
r#"{{"model":"model","messages":[{{"role":"user","content":"hi"}}],"{field}":{bad_value}}}"#
);
let req = request(&body);
assert_eq!(
normalize_request(&req, defaults(), ServeProfile::lattice_serve("model", 32),)
.unwrap_err()
.code(),
"invalid_request_body",
"field {field} should have been rejected on the honoring profile"
);
}
}
#[test]
fn message_count_over_limit_is_rejected_during_deserialization() {
let messages: Vec<String> = (0..MAX_MESSAGE_COUNT + 1)
.map(|_| r#"{"role":"user","content":""}"#.to_string())
.collect();
let body = format!(r#"{{"model":"model","messages":[{}]}}"#, messages.join(","));
let err = serde_json::from_str::<ChatRequest>(&body).unwrap_err();
assert!(is_message_flood_error(&err));
}
#[test]
fn message_flood_error_short_circuits_at_max_plus_one() {
let messages: Vec<String> = (0..MAX_MESSAGE_COUNT * 4)
.map(|_| r#"{"role":"user","content":""}"#.to_string())
.collect();
let body = format!(r#"{{"model":"model","messages":[{}]}}"#, messages.join(","));
let err = serde_json::from_str::<ChatRequest>(&body).unwrap_err();
assert!(is_message_flood_error(&err));
}
#[test]
fn non_flood_deserialize_errors_are_not_misclassified_as_message_flood() {
assert!(!is_message_flood_error(
&serde_json::from_str::<ChatRequest>("not json").unwrap_err()
));
assert!(!is_message_flood_error(
&serde_json::from_str::<ChatRequest>(
r#"{"model":"model","messages":[{"role":123,"content":"hi"}]}"#
)
.unwrap_err()
));
}
#[test]
fn cumulative_content_bytes_over_limit_is_rejected() {
let big_content = "x".repeat(MAX_CUMULATIVE_CONTENT_BYTES + 1);
let body = format!(
r#"{{"model":"model","messages":[{{"role":"user","content":"{big_content}"}}]}}"#
);
let req = request(&body);
assert_eq!(
normalize_request(&req, defaults(), ServeProfile::lattice("model", 32),)
.unwrap_err()
.code(),
"invalid_request_body"
);
}
#[test]
fn message_count_and_content_bytes_at_the_limit_are_accepted() {
let messages: Vec<String> = (0..MAX_MESSAGE_COUNT)
.map(|i| {
if i + 1 == MAX_MESSAGE_COUNT {
r#"{"role":"user","content":"hi"}"#.to_string()
} else {
r#"{"role":"user","content":""}"#.to_string()
}
})
.collect();
let body = format!(r#"{{"model":"model","messages":[{}]}}"#, messages.join(","));
let req = request(&body);
normalize_request(&req, defaults(), ServeProfile::lattice("model", 32)).unwrap();
}
#[test]
fn stop_string_over_byte_limit_is_rejected_before_matcher_construction() {
let big_stop = "x".repeat(MAX_STOP_STRING_BYTES + 1);
let body = format!(
r#"{{"model":"model","messages":[{{"role":"user","content":"hi"}}],"stop":"{big_stop}"}}"#
);
let req = request(&body);
let err =
normalize_request(&req, defaults(), ServeProfile::lattice("model", 32)).unwrap_err();
assert_eq!(err.code(), "invalid_request_body");
let body = format!(
r#"{{"model":"model","messages":[{{"role":"user","content":"hi"}}],"stop":["{big_stop}"]}}"#
);
let req = request(&body);
assert_eq!(
normalize_request(&req, defaults(), ServeProfile::lattice("model", 32),)
.unwrap_err()
.code(),
"invalid_request_body"
);
}
#[test]
fn stop_strings_over_cumulative_byte_limit_are_rejected() {
let each = "x".repeat(MAX_CUMULATIVE_STOP_BYTES / 4 + 1);
let stops = format!(r#""{each}","{each}","{each}","{each}""#);
let body = format!(
r#"{{"model":"model","messages":[{{"role":"user","content":"hi"}}],"stop":[{stops}]}}"#
);
let req = request(&body);
assert_eq!(
normalize_request(&req, defaults(), ServeProfile::lattice("model", 32),)
.unwrap_err()
.code(),
"invalid_request_body"
);
}
#[test]
fn stop_string_at_the_byte_limit_is_accepted() {
let ok_stop = "x".repeat(MAX_STOP_STRING_BYTES);
let body = format!(
r#"{{"model":"model","messages":[{{"role":"user","content":"hi"}}],"stop":"{ok_stop}"}}"#
);
let req = request(&body);
let validated =
normalize_request(&req, defaults(), ServeProfile::lattice("model", 32)).unwrap();
assert_eq!(validated.stop_strings, vec![ok_stop]);
}
#[test]
fn non_scalar_extension_field_rejected_without_cloning_large_payload() {
let big_array = format!("[{}]", vec!["1"; 100_000].join(","));
let body = format!(
r#"{{"model":"model","messages":[{{"role":"user","content":"hi"}}],"top_k":{big_array}}}"#
);
let req = request(&body);
let err = normalize_request(&req, defaults(), ServeProfile::lattice_serve("model", 32))
.unwrap_err();
assert_eq!(err.code(), "invalid_request_body");
assert_eq!(err.message(), "top_k must be a scalar value");
}
#[test]
fn wrong_model_with_conflicting_stream_and_logprobs_rejects_as_unsupported_feature() {
let req = request(
r#"{"model":"wrong-model","messages":[{"role":"user","content":"hi"}],"stream":true,"logprobs":true}"#,
);
assert_eq!(
normalize_request(&req, defaults(), ServeProfile::lattice("served-model", 32),)
.unwrap_err()
.code(),
"unsupported_feature"
);
}
#[test]
fn wrong_model_with_unsupported_logprobs_on_daemon_profile_rejects_as_unsupported_feature() {
let req = request(
r#"{"model":"wrong-model","messages":[{"role":"user","content":"hi"}],"logprobs":true}"#,
);
assert_eq!(
normalize_request(
&req,
defaults(),
ServeProfile::lattice_serve("served-model", 32),
)
.unwrap_err()
.code(),
"unsupported_feature"
);
}
#[test]
fn wrong_model_with_conflicting_max_tokens_alias_rejects_as_invalid_request() {
let req = request(
r#"{"model":"wrong-model","messages":[{"role":"user","content":"hi"}],"max_tokens":10,"max_completion_tokens":20}"#,
);
assert_eq!(
normalize_request(
&req,
defaults(),
ServeProfile::lattice_serve("served-model", 32),
)
.unwrap_err()
.code(),
"invalid_request"
);
}
#[test]
fn wrong_model_with_conflicting_max_tokens_alias_rejects_as_model_not_found_on_lattice_profile()
{
let req = request(
r#"{"model":"wrong-model","messages":[{"role":"user","content":"hi"}],"max_tokens":10,"max_completion_tokens":20}"#,
);
assert_eq!(
normalize_request(&req, defaults(), ServeProfile::lattice("served-model", 32),)
.unwrap_err()
.code(),
"model_not_found"
);
}
#[test]
fn lattice_profile_still_rejects_genuinely_unsupported_fields() {
let req = request(
r#"{"model":"model","messages":[{"role":"user","content":"hi"}],"tools":[{"type":"function"}]}"#,
);
assert_eq!(
normalize_request(&req, defaults(), ServeProfile::lattice("model", 32),)
.unwrap_err()
.code(),
"unsupported_feature"
);
}
#[test]
fn unmodeled_openai_top_level_fields_are_ignored_not_rejected() {
let req = serde_json::from_str::<ChatRequest>(
r#"{
"model": "model",
"messages": [{"role": "user", "content": "hi"}],
"presence_penalty": 1.0,
"frequency_penalty": 0.5,
"logit_bias": {"123": -100},
"user": "end-user-id"
}"#,
)
.expect("unmodeled top-level fields must be ignored, not rejected");
assert_eq!(req.model.as_deref(), Some("model"));
assert_eq!(req.messages.len(), 1);
}
#[test]
fn context_window_accepts_boundary_and_rejects_overflow() {
validate_context_window(8, 8, 16).unwrap();
assert_eq!(
validate_context_window(8, 9, 16).unwrap_err().code(),
"context_length_exceeded"
);
}
#[test]
fn optional_exact_rejects_explicit_empty_model() {
let err =
validate_model_name(Some(""), ModelNamePolicy::OptionalExact("served")).unwrap_err();
assert_eq!(err.code(), "model_not_found");
}
#[test]
fn optional_exact_accepts_omitted_model() {
validate_model_name(None, ModelNamePolicy::OptionalExact("served")).unwrap();
}
#[test]
fn optional_exact_matches_and_rejects_by_name() {
validate_model_name(Some("served"), ModelNamePolicy::OptionalExact("served")).unwrap();
assert_eq!(
validate_model_name(Some("other"), ModelNamePolicy::OptionalExact("served"))
.unwrap_err()
.code(),
"model_not_found"
);
}
#[test]
fn required_exact_rejects_missing_or_empty_and_accepts_match() {
assert_eq!(
validate_model_name(None, ModelNamePolicy::RequiredExact("served"))
.unwrap_err()
.code(),
"invalid_request"
);
assert_eq!(
validate_model_name(Some(""), ModelNamePolicy::RequiredExact("served"))
.unwrap_err()
.code(),
"invalid_request"
);
validate_model_name(Some("served"), ModelNamePolicy::RequiredExact("served")).unwrap();
}
fn message(body: &str) -> Message {
serde_json::from_str(body).unwrap()
}
#[test]
fn normalize_messages_accepts_system_user_assistant_roles() {
let messages = [
message(r#"{"role":"system","content":"sys"}"#),
message(r#"{"role":"user","content":"usr"}"#),
message(r#"{"role":"assistant","content":"asst"}"#),
];
let normalized = normalize_messages(&messages).unwrap();
assert_eq!(normalized[0].role, NormalizedChatRole::System);
assert_eq!(normalized[0].content, "sys");
assert_eq!(normalized[1].role, NormalizedChatRole::User);
assert_eq!(normalized[1].content, "usr");
assert_eq!(normalized[2].role, NormalizedChatRole::Assistant);
assert_eq!(normalized[2].content, "asst");
}
#[test]
fn normalize_messages_rejects_unrecognized_role() {
let messages = [message(r#"{"role":"moderator","content":"hi"}"#)];
assert_eq!(
normalize_messages(&messages).unwrap_err().code(),
"invalid_role"
);
}
#[test]
fn normalize_messages_rejects_known_but_unsupported_roles() {
for role in ["tool", "developer"] {
let messages = [message(&format!(r#"{{"role":"{role}","content":"hi"}}"#))];
assert_eq!(
normalize_messages(&messages).unwrap_err().code(),
"unsupported_feature"
);
}
}
#[test]
fn normalize_messages_checks_content_before_role() {
let messages = [message(
r#"{"role":"moderator","content":[{"type":"image_url","image_url":{"url":"https://example.com/x.png"}}]}"#,
)];
let err = normalize_messages(&messages).unwrap_err();
assert_eq!(err.message(), "image input requires a vision-capable model");
}
#[test]
fn honoring_profile_treats_explicit_null_sampling_fields_as_absent() {
let req = request(
r#"{"model":"model","messages":[{"role":"user","content":"hi"}],"top_k":null,"repetition_penalty":null,"reasoning_budget":null}"#,
);
let daemon =
normalize_request(&req, defaults(), ServeProfile::lattice_serve("model", 100)).unwrap();
assert_eq!(daemon.top_k, defaults().top_k);
assert_eq!(daemon.repetition_penalty, defaults().repetition_penalty);
assert_eq!(daemon.reasoning_budget, None);
}
#[test]
fn honoring_profile_rejects_near_body_cap_nested_extension_field() {
let big_array = format!("[{}]", vec!["1"; REQUEST_BODY_LIMIT_BYTES / 2].join(","));
let body = format!(
r#"{{"model":"model","messages":[{{"role":"user","content":"hi"}}],"top_k":{big_array}}}"#
);
let req = request(&body);
let err = normalize_request(&req, defaults(), ServeProfile::lattice_serve("model", 32))
.unwrap_err();
assert_eq!(err.code(), "invalid_request_body");
assert_eq!(err.message(), "top_k must be a scalar value");
}
#[test]
fn non_honoring_profile_ignores_near_body_cap_nested_extension_field() {
let big_array = format!("[{}]", vec!["1"; REQUEST_BODY_LIMIT_BYTES / 2].join(","));
let body = format!(
r#"{{"model":"model","messages":[{{"role":"user","content":"hi"}}],"top_k":{big_array}}}"#
);
let req = request(&body);
let lattice =
normalize_request(&req, defaults(), ServeProfile::lattice("model", 32)).unwrap();
assert_eq!(lattice.top_k, defaults().top_k);
}
#[test]
fn lattice_profile_rejects_json_schema_without_materializing_the_schema() {
let mut nested = String::new();
for _ in 0..200 {
nested.push('[');
}
nested.push('1');
for _ in 0..200 {
nested.push(']');
}
assert!(
serde_json::from_str::<Value>(&nested).is_err(),
"fixture must exceed serde_json's Value recursion limit"
);
let body = format!(
r#"{{"model":"model","messages":[{{"role":"user","content":"hi"}}],
"response_format":{{"type":"json_schema","json_schema":{{"name":"n","strict":true,"schema":{nested}}}}}}}"#
);
let req = request(&body);
let err =
normalize_request(&req, defaults(), ServeProfile::lattice("model", 32)).unwrap_err();
assert_eq!(err.code(), "unsupported_feature");
assert_eq!(
err.message(),
"response_format.type 'json_schema' is not supported; use 'text'"
);
}
#[test]
fn valid_request_body_is_parsed_by_a_single_deserialize_call() {
let messages: Vec<String> = (0..MAX_MESSAGE_COUNT)
.map(|_| r#"{"role":"user","content":""}"#.to_string())
.collect();
let body = format!(r#"{{"model":"model","messages":[{}]}}"#, messages.join(","));
let req = serde_json::from_str::<ChatRequest>(&body).unwrap();
assert_eq!(req.messages.len(), MAX_MESSAGE_COUNT);
}
#[test]
fn normalize_messages_passes_through_empty_and_whitespace_content() {
let messages = [
message(r#"{"role":"user","content":""}"#),
message(r#"{"role":"user","content":" "}"#),
];
let normalized = normalize_messages(&messages).unwrap();
assert_eq!(normalized[0].content, "");
assert_eq!(normalized[1].content, " ");
}
}