use super::sampler::GrammarRuntime;
pub fn mask_invalid_tokens(
grammar: &GrammarRuntime,
token_bytes: &[Vec<u8>],
logits: &mut [f32],
) -> usize {
if grammar.is_awaiting_trigger() {
return 0;
}
let mut masked = 0usize;
let n = token_bytes.len().min(logits.len());
for i in 0..n {
let bytes = &token_bytes[i];
if bytes.is_empty() {
continue;
}
if !logits[i].is_finite() {
continue;
}
let mut rt = grammar.clone();
let alive = rt.accept_bytes(bytes);
if !alive {
logits[i] = f32::NEG_INFINITY;
masked += 1;
}
}
masked
}
#[cfg(test)]
pub fn surviving_token_ids(
grammar: &GrammarRuntime,
token_bytes: &[Vec<u8>],
logits: &[f32],
) -> Vec<u32> {
let mut out = Vec::new();
let n = token_bytes.len().min(logits.len());
for i in 0..n {
let bytes = &token_bytes[i];
if bytes.is_empty() || !logits[i].is_finite() {
if logits[i].is_finite() {
out.push(i as u32);
}
continue;
}
let mut rt = grammar.clone();
if rt.accept_bytes(bytes) {
out.push(i as u32);
}
}
out
}
#[cfg(test)]
mod tests {
use super::super::parser::parse;
use super::*;
fn rt(grammar_src: &str, start: &str) -> GrammarRuntime {
let g = parse(grammar_src).expect("parse");
let rid = g.rule_id(start).expect("start");
GrammarRuntime::new(g, rid).expect("runtime")
}
fn vocab(strings: &[&str]) -> Vec<Vec<u8>> {
strings.iter().map(|s| s.as_bytes().to_vec()).collect()
}
#[test]
fn mask_rejects_tokens_that_dont_match_literal() {
let runtime = rt("root ::= \"abc\"\n", "root");
let token_bytes = vocab(&["a", "b", "c", "x", "Z"]);
let mut logits = vec![1.0, 1.0, 1.0, 1.0, 1.0];
let masked = mask_invalid_tokens(&runtime, &token_bytes, &mut logits);
assert_eq!(masked, 4, "only 'a' should survive from {:?}", logits);
assert_eq!(logits[0], 1.0);
assert!(logits[1].is_infinite() && logits[1] < 0.0);
assert!(logits[2].is_infinite() && logits[2] < 0.0);
assert!(logits[3].is_infinite() && logits[3] < 0.0);
assert!(logits[4].is_infinite() && logits[4] < 0.0);
}
#[test]
fn mask_respects_char_class_range() {
let runtime = rt("root ::= [0-9]\n", "root");
let token_bytes = vocab(&["0", "5", "9", "a", "ZZ"]);
let mut logits = vec![1.0, 1.0, 1.0, 1.0, 1.0];
let masked = mask_invalid_tokens(&runtime, &token_bytes, &mut logits);
assert_eq!(masked, 2);
assert_eq!(logits[0], 1.0); assert_eq!(logits[1], 1.0); assert_eq!(logits[2], 1.0); assert!(logits[3].is_infinite());
assert!(logits[4].is_infinite()); }
#[test]
fn mask_accepts_multi_byte_utf8_token() {
let runtime = rt("root ::= \"α\"\n", "root");
let token_bytes = vec!["α".as_bytes().to_vec(), "β".as_bytes().to_vec()];
let mut logits = vec![1.0, 1.0];
let masked = mask_invalid_tokens(&runtime, &token_bytes, &mut logits);
assert_eq!(masked, 1);
assert_eq!(logits[0], 1.0);
assert!(logits[1].is_infinite());
}
#[test]
fn mask_skips_empty_token_strings() {
let runtime = rt("root ::= \"a\"\n", "root");
let token_bytes = vec![b"a".to_vec(), vec![], b"b".to_vec()];
let mut logits = vec![1.0, 2.0, 3.0];
let masked = mask_invalid_tokens(&runtime, &token_bytes, &mut logits);
assert_eq!(masked, 1); assert_eq!(logits[0], 1.0); assert_eq!(logits[1], 2.0); assert!(logits[2].is_infinite()); }
#[test]
fn mask_ignores_already_negative_infinity_tokens() {
let runtime = rt("root ::= \"a\" | \"b\"\n", "root");
let token_bytes = vocab(&["a", "b", "c"]);
let mut logits = vec![1.0, f32::NEG_INFINITY, 3.0];
let masked = mask_invalid_tokens(&runtime, &token_bytes, &mut logits);
assert_eq!(masked, 1); assert_eq!(logits[0], 1.0);
assert!(logits[1].is_infinite());
assert!(logits[2].is_infinite());
}
#[test]
fn mask_is_idempotent_after_running_twice() {
let runtime = rt("root ::= \"a\" | \"b\"\n", "root");
let token_bytes = vocab(&["a", "b", "c", "d"]);
let mut logits = vec![1.0; 4];
let m1 = mask_invalid_tokens(&runtime, &token_bytes, &mut logits);
let m2 = mask_invalid_tokens(&runtime, &token_bytes, &mut logits);
assert_eq!(m1, 2); assert_eq!(m2, 0); assert_eq!(logits[0], 1.0);
assert_eq!(logits[1], 1.0);
assert!(logits[2].is_infinite());
assert!(logits[3].is_infinite());
}
#[test]
fn mask_after_partial_decode_narrows_survivors() {
let mut runtime = rt("root ::= \"ab\"\n", "root");
let token_bytes = vocab(&["a", "b", "c"]);
let mut logits = vec![1.0, 1.0, 1.0];
mask_invalid_tokens(&runtime, &token_bytes, &mut logits);
assert_eq!(logits[0], 1.0);
assert!(logits[1].is_infinite());
assert!(logits[2].is_infinite());
assert!(runtime.accept_char('a' as u32));
let mut logits = vec![1.0, 1.0, 1.0];
mask_invalid_tokens(&runtime, &token_bytes, &mut logits);
assert!(logits[0].is_infinite());
assert_eq!(logits[1], 1.0);
assert!(logits[2].is_infinite());
}
#[test]
fn mask_with_json_grammar_accepts_opening_brace() {
let src = std::fs::read_to_string("/opt/llama.cpp/grammars/json.gbnf")
.expect("json.gbnf fixture");
let g = parse(&src).unwrap();
let rid = g.rule_id("root").unwrap();
let runtime = GrammarRuntime::new(g, rid).unwrap();
let token_bytes = vocab(&["{", "}", "[", "\"", "a", "1"]);
let mut logits = vec![1.0; 6];
let _ = mask_invalid_tokens(&runtime, &token_bytes, &mut logits);
assert_eq!(logits[0], 1.0, "'{{' must survive");
assert!(logits[1].is_infinite(), "'}}' must be masked");
assert!(logits[2].is_infinite(), "'[' must be masked");
assert!(logits[3].is_infinite(), "'\"' must be masked");
assert!(logits[4].is_infinite());
assert!(logits[5].is_infinite());
}
#[test]
fn surviving_token_ids_helper_matches_mask_counts() {
let runtime = rt("root ::= \"abc\"\n", "root");
let token_bytes = vocab(&["a", "b", "c", "x"]);
let logits = vec![1.0, 1.0, 1.0, 1.0];
let survivors = surviving_token_ids(&runtime, &token_bytes, &logits);
assert_eq!(survivors, vec![0u32]); }
#[test]
fn mask_does_not_exceed_logits_length() {
let runtime = rt("root ::= \"a\"\n", "root");
let token_bytes = vocab(&["a", "b", "c", "d", "e"]);
let mut logits = vec![1.0, 1.0, 1.0];
let masked = mask_invalid_tokens(&runtime, &token_bytes, &mut logits);
assert_eq!(masked, 2);
assert_eq!(logits.len(), 3);
}
#[test]
fn runtime_apply_noops_when_awaiting_trigger() {
let mut runtime = rt("root ::= \"a\"\n", "root");
runtime.set_awaiting_trigger(true);
let token_bytes = vocab(&["a", "b", "c", "x"]);
let mut logits = vec![1.0, 1.0, 1.0, 1.0];
let masked = mask_invalid_tokens(&runtime, &token_bytes, &mut logits);
assert_eq!(
masked, 0,
"suspended runtime MUST mask zero tokens (preamble freedom)"
);
for (i, &l) in logits.iter().enumerate() {
assert_eq!(l, 1.0, "logit {i} must be unchanged while awaiting trigger");
}
}
#[test]
fn runtime_apply_active_after_trigger() {
let mut runtime = rt("root ::= \"a\"\n", "root");
runtime.set_awaiting_trigger(true);
runtime.trigger();
assert!(!runtime.is_awaiting_trigger());
let token_bytes = vocab(&["a", "b", "c", "x"]);
let mut logits = vec![1.0, 1.0, 1.0, 1.0];
let masked = mask_invalid_tokens(&runtime, &token_bytes, &mut logits);
assert_eq!(
masked, 3,
"post-trigger runtime masks the 3 invalid tokens (only 'a' survives)"
);
assert!(logits[0].is_finite(), "'a' survives");
assert!(logits[1].is_infinite(), "'b' masked");
assert!(logits[2].is_infinite(), "'c' masked");
assert!(logits[3].is_infinite(), "'x' masked");
}
#[test]
fn runtime_response_format_never_awaits() {
let runtime = rt("root ::= \"a\"\n", "root");
assert!(
!runtime.is_awaiting_trigger(),
"default (ResponseFormat-equivalent) runtime MUST NOT await trigger"
);
let token_bytes = vocab(&["a", "b"]);
let mut logits = vec![1.0, 1.0];
let masked = mask_invalid_tokens(&runtime, &token_bytes, &mut logits);
assert_eq!(
masked, 1,
"ResponseFormat-kind runtime enforces from token 0 with no \
trigger flip required"
);
}
const GEMMA4_TOKENIZER_PATH: &str = "/opt/hf2q/models/gemma4/tokenizer.json";
fn load_gemma4_tokenizer_or_skip() -> Option<tokenizers::Tokenizer> {
if !std::path::Path::new(GEMMA4_TOKENIZER_PATH).exists() {
return None;
}
match tokenizers::Tokenizer::from_file(GEMMA4_TOKENIZER_PATH) {
Ok(t) => Some(t),
Err(e) => panic!(
"Tokenizer fixture exists at {} but failed to load: {}\n\
Fix or remove the fixture; do not silence this error.",
GEMMA4_TOKENIZER_PATH, e
),
}
}
fn token_bytes_table_for_range(tok: &tokenizers::Tokenizer, up_to: u32) -> Vec<Vec<u8>> {
let mut out: Vec<Vec<u8>> = Vec::with_capacity(up_to as usize);
for id in 0..up_to {
let s = tok.decode(&[id], false).unwrap_or_default();
out.push(s.into_bytes());
}
out
}
#[test]
fn tokenizer_backed_table_preserves_gemma_open_marker_bytes() {
let Some(tok) = load_gemma4_tokenizer_or_skip() else {
return;
};
let table = token_bytes_table_for_range(&tok, 256);
assert_eq!(table.len(), 256);
let id_48 = &table[48];
assert!(
!id_48.is_empty(),
"Gemma 4 id 48 (<|tool_call>) decoded to empty bytes through \
tok.decode(&[48], false); the mask's bytes.is_empty() skip \
at mask.rs:77-79 would leave the open marker un-maskable. \
This breaks the wave-2.7 Q-A eager-grammar contract."
);
assert_eq!(
id_48.as_slice(),
b"<|tool_call>",
"Gemma 4 id 48 must decode to the 12-byte literal '<|tool_call>'; \
got {:?}",
String::from_utf8_lossy(id_48)
);
assert_eq!(
id_48.len(),
12,
"Gemma 4 '<|tool_call>' is 12 ASCII bytes; got {} bytes",
id_48.len()
);
}
#[test]
fn mask_with_real_tokenizer_keeps_gemma_open_marker_alive() {
let Some(tok) = load_gemma4_tokenizer_or_skip() else {
return;
};
let token_bytes = token_bytes_table_for_range(&tok, 256);
assert_eq!(token_bytes[48], b"<|tool_call>");
let runtime = rt("root ::= \"<|tool_call>\"\n", "root");
let mut logits = vec![1.0_f32; token_bytes.len()];
let _ = mask_invalid_tokens(&runtime, &token_bytes, &mut logits);
assert!(
logits[48].is_finite(),
"Gemma id 48 (<|tool_call>) was masked to {}; the eager \
grammar's open-marker constraint must FUNNEL the model to \
this token, not mask it out",
logits[48]
);
let surviving: Vec<u32> = (0..token_bytes.len() as u32)
.filter(|&i| !token_bytes[i as usize].is_empty() && logits[i as usize].is_finite())
.collect();
assert!(
surviving.contains(&48),
"id 48 must be in surviving set; got {:?}",
surviving
);
let empty_byte_tokens: usize = token_bytes.iter().filter(|b| b.is_empty()).count();
let _ = empty_byte_tokens; }
}