use super::sampler::{decode_candidate_utf8, reject_candidates, GrammarCandidate, GrammarRuntime};
use std::cell::RefCell;
use std::sync::atomic::{AtomicU64, Ordering};
const GREEDY_PROBE_LIMIT: usize = 64;
thread_local! {
static GREEDY_CANDIDATES: RefCell<Vec<(usize, f32)>> = const { RefCell::new(Vec::new()) };
}
static ZT_LOG_PATH: std::sync::OnceLock<Option<std::path::PathBuf>> = std::sync::OnceLock::new();
static ZT_SEQ: AtomicU64 = AtomicU64::new(0);
fn zt_log_path() -> Option<&'static std::path::PathBuf> {
ZT_LOG_PATH
.get_or_init(|| {
std::env::var("HF2Q_ZT_LOG").ok().and_then(|v| {
let v = v.trim();
if v.is_empty() {
None
} else {
Some(std::path::PathBuf::from(v))
}
})
})
.as_ref()
}
fn zt_record(pre_finite_logits: &[(usize, f32)], post_logits: &[f32], masked: usize) {
let Some(path) = zt_log_path() else { return };
let max = pre_finite_logits
.iter()
.map(|(_, l)| *l)
.fold(f32::NEG_INFINITY, f32::max);
if !max.is_finite() {
return;
}
let denom: f64 = pre_finite_logits
.iter()
.map(|(_, l)| (*l - max) as f64)
.map(f64::exp)
.sum();
if denom <= 0.0 {
return;
}
let numer: f64 = pre_finite_logits
.iter()
.filter(|(i, _)| post_logits[*i].is_finite())
.map(|(_, l)| (*l - max) as f64)
.map(f64::exp)
.sum();
let z = numer / denom;
let seq = ZT_SEQ.fetch_add(1, Ordering::Relaxed);
let line = format!(
"{{\"seq\":{seq},\"z\":{z:.6e},\"masked\":{masked},\"finite\":{}}}\n",
pre_finite_logits.len()
);
if let Ok(mut f) = std::fs::OpenOptions::new().create(true).append(true).open(path) {
use std::io::Write;
let _ = f.write_all(line.as_bytes());
}
}
pub fn sample_greedy_valid_token(
logits: &mut [f32],
previous_tokens: &[u32],
repetition_penalty: f64,
grammar: &GrammarRuntime,
token_bytes: &[Vec<u8>],
eog_token_ids: &[u32],
) -> u32 {
if repetition_penalty != 1.0 && !previous_tokens.is_empty() {
crate::serve::sampler_pure::apply_repetition_penalty(
logits,
previous_tokens,
repetition_penalty,
);
}
if grammar.is_awaiting_trigger() {
return crate::serve::sampler_pure::sample_greedy(logits);
}
let selected = GREEDY_CANDIDATES.with(|cell| {
let mut candidates = cell.borrow_mut();
candidates.clear();
candidates.reserve(logits.len());
candidates.extend(
logits
.iter()
.copied()
.enumerate()
.filter(|(_, logit)| logit.is_finite()),
);
if candidates.is_empty() {
return None;
}
let compare = |left: &(usize, f32), right: &(usize, f32)| {
right
.1
.total_cmp(&left.1)
.then_with(|| left.0.cmp(&right.0))
};
let limit = GREEDY_PROBE_LIMIT.min(candidates.len());
if candidates.len() > limit {
candidates.select_nth_unstable_by(limit - 1, compare);
candidates.truncate(limit);
}
candidates.sort_unstable_by(compare);
for &(token, _) in candidates.iter() {
let Some(bytes) = token_bytes.get(token) else {
return Some(token as u32);
};
if eog_token_ids.contains(&(token as u32)) {
if grammar.is_terminally_accepted() {
return Some(token as u32);
}
logits[token] = f32::NEG_INFINITY;
continue;
}
if bytes.is_empty() || bytes.first() == Some(&0) {
logits[token] = f32::NEG_INFINITY;
continue;
}
let mut probe = grammar.clone();
if probe.accept_token(token as u32, bytes) {
return Some(token as u32);
}
logits[token] = f32::NEG_INFINITY;
}
None
});
if let Some(token) = selected {
return token;
}
mask_invalid_tokens_with_eog(grammar, token_bytes, eog_token_ids, logits);
crate::serve::sampler_pure::sample_greedy(logits)
}
pub fn mask_invalid_tokens(
grammar: &GrammarRuntime,
token_bytes: &[Vec<u8>],
logits: &mut [f32],
) -> usize {
mask_invalid_tokens_with_eog(grammar, token_bytes, &[], logits)
}
pub fn mask_invalid_tokens_with_eog(
grammar: &GrammarRuntime,
token_bytes: &[Vec<u8>],
eog_token_ids: &[u32],
logits: &mut [f32],
) -> usize {
if grammar.is_awaiting_trigger() {
return 0;
}
let pre_snapshot: Option<Vec<(usize, f32)>> = if zt_log_path().is_some() {
Some(
logits
.iter()
.copied()
.enumerate()
.filter(|(_, l)| l.is_finite())
.collect(),
)
} else {
None
};
let n = token_bytes.len().min(logits.len());
let mut code_points = Vec::with_capacity(n.saturating_mul(2));
let mut candidates = Vec::with_capacity(n);
let mut masked = 0usize;
for i in 0..n {
let bytes = &token_bytes[i];
if eog_token_ids.contains(&(i as u32)) {
if !grammar.is_terminally_accepted() && logits[i].is_finite() {
logits[i] = f32::NEG_INFINITY;
masked += 1;
}
continue;
}
if bytes.is_empty() || bytes.first() == Some(&0) {
if logits[i].is_finite() {
logits[i] = f32::NEG_INFINITY;
masked += 1;
}
continue;
}
if !logits[i].is_finite() {
continue;
}
let start = code_points.len();
let Some(partial_utf8) =
decode_candidate_utf8(bytes, grammar.partial_utf8, &mut code_points)
else {
logits[i] = f32::NEG_INFINITY;
masked += 1;
continue;
};
candidates.push(GrammarCandidate {
index: i,
token_id: i as u32,
cursor: start,
end: code_points.len(),
partial_utf8,
});
}
let rejects = reject_candidates(&grammar.grammar, &grammar.stacks, &candidates, &code_points);
for reject in rejects {
if logits[reject.index].is_finite() {
logits[reject.index] = f32::NEG_INFINITY;
masked += 1;
}
}
if let Some(pre) = pre_snapshot {
zt_record(&pre, logits, masked);
}
masked
}
#[cfg(test)]
fn mask_invalid_tokens_clone_oracle(
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 !logits[i].is_finite() {
continue;
}
if bytes.is_empty() || bytes.first() == Some(&0) {
logits[i] = f32::NEG_INFINITY;
masked += 1;
continue;
}
let mut runtime = grammar.clone();
if !runtime.accept_token(i as u32, bytes) {
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_token(i as u32, 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")
}
#[test]
fn token_terminal_masks_by_id_not_decoded_bytes() {
let runtime = rt("root ::= <[1]>\n", "root");
let token_bytes = vocab(&["same", "same", "different"]);
let mut logits = vec![1.0; token_bytes.len()];
mask_invalid_tokens_with_eog(&runtime, &token_bytes, &[], &mut logits);
assert!(!logits[0].is_finite());
assert!(logits[1].is_finite());
assert!(!logits[2].is_finite());
}
#[test]
fn token_any_and_exclusion_set_mask_by_id_without_enumerating_vocab() {
let token_bytes = vocab(&["same", "same", "same", "same", "same"]);
let any = rt("root ::= <[*]>\n", "root");
let mut any_logits = vec![1.0; token_bytes.len()];
assert_eq!(
mask_invalid_tokens_with_eog(&any, &token_bytes, &[], &mut any_logits),
0
);
let exclusion = rt("root ::= !<[1,3]>\n", "root");
let mut exclusion_logits = vec![1.0; token_bytes.len()];
assert_eq!(
mask_invalid_tokens_with_eog(&exclusion, &token_bytes, &[], &mut exclusion_logits,),
2
);
assert!(exclusion_logits[0].is_finite());
assert!(!exclusion_logits[1].is_finite());
assert!(exclusion_logits[2].is_finite());
assert!(!exclusion_logits[3].is_finite());
assert!(exclusion_logits[4].is_finite());
}
#[test]
fn eog_remains_masked_while_an_accepted_alternate_has_pending_utf8() {
let mut runtime = rt("root ::= \"\" | .\n", "root");
assert!(runtime.accept_bytes(&[0xCE]));
assert!(runtime.is_accepted());
assert!(!runtime.is_terminally_accepted());
let mut logits = vec![1.0];
assert_eq!(
mask_invalid_tokens_with_eog(&runtime, &[Vec::new()], &[0], &mut logits),
1
);
assert!(!logits[0].is_finite());
}
#[test]
fn eog_is_masked_until_accepted_and_never_satisfies_token_terminal() {
let eog = [2_u32];
let token_bytes = vocab(&["x", "y", "<eos>"]);
let token_runtime = rt("root ::= <[2]>\n", "root");
let mut token_logits = vec![1.0; token_bytes.len()];
mask_invalid_tokens_with_eog(&token_runtime, &token_bytes, &eog, &mut token_logits);
assert!(!token_logits[2].is_finite());
let mut accepted = rt("root ::= \"x\"\n", "root");
assert!(accepted.accept_token(0, b"x"));
let mut accepted_logits = vec![1.0; token_bytes.len()];
mask_invalid_tokens_with_eog(&accepted, &token_bytes, &eog, &mut accepted_logits);
assert!(accepted_logits[2].is_finite());
assert!(!accepted_logits[1].is_finite());
}
#[test]
fn empty_non_eog_piece_is_always_masked() {
let mut runtime = rt("root ::= \"x\"\n", "root");
assert!(runtime.accept_token(0, b"x"));
let token_bytes = vec![b"x".to_vec(), Vec::new()];
let mut logits = vec![1.0; token_bytes.len()];
mask_invalid_tokens_with_eog(&runtime, &token_bytes, &[], &mut logits);
assert!(!logits[1].is_finite());
}
#[test]
fn greedy_probe_obeys_eog_boundary() {
let token_bytes = vocab(&["x", "<eos>"]);
let eog = [1_u32];
let runtime = rt("root ::= \"x\"\n", "root");
let mut logits = vec![1.0_f32, 10.0];
assert_eq!(
sample_greedy_valid_token(&mut logits, &[], 1.0, &runtime, &token_bytes, &eog),
0
);
let mut accepted = rt("root ::= \"x\"\n", "root");
assert!(accepted.accept_token(0, b"x"));
let mut logits = vec![1.0_f32, 10.0];
assert_eq!(
sample_greedy_valid_token(&mut logits, &[], 1.0, &accepted, &token_bytes, &eog),
1
);
}
fn vocab(strings: &[&str]) -> Vec<Vec<u8>> {
strings.iter().map(|s| s.as_bytes().to_vec()).collect()
}
fn assert_candidate_set_matches_clone_oracle(
runtime: &GrammarRuntime,
token_bytes: &[Vec<u8>],
initial_logits: &[f32],
) {
let mut actual = initial_logits.to_vec();
let mut oracle = initial_logits.to_vec();
let actual_masked = mask_invalid_tokens(runtime, token_bytes, &mut actual);
let oracle_masked = mask_invalid_tokens_clone_oracle(runtime, token_bytes, &mut oracle);
assert_eq!(actual_masked, oracle_masked, "new/oracle mask count");
assert_eq!(
actual
.iter()
.map(|value| value.to_bits())
.collect::<Vec<_>>(),
oracle
.iter()
.map(|value| value.to_bits())
.collect::<Vec<_>>(),
"new candidate-set mask must be bit-identical to clone oracle"
);
}
#[test]
fn agentic_grammar_contract_candidate_set_matches_clone_oracle_across_runtime_states() {
let names = vocab(&[
"ruflo_call",
"ruflo_search",
"aqe_call",
"aqe_search",
"task",
"read",
"write",
"bash",
"brain",
"skill",
"agent",
"memory",
"workflow",
"hooks",
"routing",
"swarm",
"wrong",
"",
]);
let runtime = rt(
"root ::= \"ruflo_call\" | \"ruflo_search\" | \"aqe_call\" | \"aqe_search\" | \"task\" | \"read\" | \"write\" | \"bash\" | \"brain\" | \"skill\" | \"agent\" | \"memory\" | \"workflow\" | \"hooks\" | \"routing\" | \"swarm\"\n",
"root",
);
let mut logits = vec![1.0; names.len()];
logits[3] = f32::NEG_INFINITY;
assert_candidate_set_matches_clone_oracle(&runtime, &names, &logits);
let literal_runtime = rt("root ::= \"a\"\n", "root");
let byte_edge_tokens = vec![vec![0xCE], vec![0x80], b"a".to_vec(), b"x".to_vec()];
assert_candidate_set_matches_clone_oracle(
&literal_runtime,
&byte_edge_tokens,
&vec![1.0; byte_edge_tokens.len()],
);
let mut string_runtime = rt("root ::= \"ruflo_call:\\\"\" [^\\\"]* \"\\\"\"\n", "root");
assert!(string_runtime.accept_bytes(b"ruflo_call:\""));
let string_tokens = vec![
b"plain".to_vec(),
b"{}".to_vec(),
"α".as_bytes().to_vec(),
vec![0xCE],
vec![0x80],
b"\"".to_vec(),
Vec::new(),
];
let mut string_logits = vec![1.0; string_tokens.len()];
string_logits[1] = f32::NEG_INFINITY;
assert_candidate_set_matches_clone_oracle(&string_runtime, &string_tokens, &string_logits);
let mut partial_runtime = rt("root ::= \"α\"\n", "root");
assert!(partial_runtime.accept_bytes(&[0xCE]));
let partial_tokens = vec![vec![0xB1], vec![0xB2], b"a".to_vec(), Vec::new()];
assert_candidate_set_matches_clone_oracle(
&partial_runtime,
&partial_tokens,
&vec![1.0; partial_tokens.len()],
);
let mut accepted = rt("root ::= \"a\"\n", "root");
assert!(accepted.accept_bytes(b"a"));
assert!(accepted.is_accepted());
let terminal_tokens = vec![Vec::new(), b"a".to_vec(), b"x".to_vec(), vec![0xCE]];
assert_candidate_set_matches_clone_oracle(
&accepted,
&terminal_tokens,
&vec![1.0; terminal_tokens.len()],
);
let mut dead = rt("root ::= \"a\"\n", "root");
assert!(!dead.accept_bytes(b"x"));
assert!(dead.is_dead());
assert_candidate_set_matches_clone_oracle(
&dead,
&terminal_tokens,
&vec![1.0; terminal_tokens.len()],
);
}
#[test]
#[ignore = "diagnostic benchmark; set HF2Q_QWEN35_TOKENIZER to tokenizer.json"]
fn candidate_set_mask_qwen_vocab_benchmark() {
let Some(tokenizer_path) = std::env::var_os("HF2Q_QWEN35_TOKENIZER") else {
eprintln!("SKIP: set HF2Q_QWEN35_TOKENIZER to a tokenizer.json path");
return;
};
let tokenizer = tokenizers::Tokenizer::from_file(&tokenizer_path).unwrap_or_else(|error| {
panic!(
"failed to load {}: {error}",
tokenizer_path.to_string_lossy()
)
});
let vocab_size = tokenizer.get_vocab_size(true);
let token_bytes = (0..vocab_size as u32)
.map(|id| {
tokenizer
.decode(&[id], false)
.unwrap_or_default()
.into_bytes()
})
.collect::<Vec<_>>();
let mut runtime = rt("root ::= \"\\\"\" [^\\\"]* \"\\\"\"\n", "root");
assert!(runtime.accept_bytes(b"\""));
let initial = vec![1.0_f32; vocab_size];
let mut actual = initial.clone();
let mut oracle = initial;
let started = std::time::Instant::now();
let actual_masked = mask_invalid_tokens(&runtime, &token_bytes, &mut actual);
let candidate_set_elapsed = started.elapsed();
let started = std::time::Instant::now();
let oracle_masked = mask_invalid_tokens_clone_oracle(&runtime, &token_bytes, &mut oracle);
let clone_oracle_elapsed = started.elapsed();
assert_eq!(actual_masked, oracle_masked);
assert_eq!(
actual
.iter()
.map(|value| value.to_bits())
.collect::<Vec<_>>(),
oracle
.iter()
.map(|value| value.to_bits())
.collect::<Vec<_>>()
);
eprintln!(
"grammar mask benchmark: vocab={vocab_size} candidate_set={candidate_set_elapsed:?} clone_oracle={clone_oracle_elapsed:?} speedup={:.2}x",
clone_oracle_elapsed.as_secs_f64() / candidate_set_elapsed.as_secs_f64()
);
}
#[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_allows_declared_eog_only_after_grammar_acceptance() {
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_with_eog(&runtime, &token_bytes, &[1], &mut logits);
assert_eq!(masked, 2);
assert_eq!(logits[0], 1.0); assert!(logits[1].is_infinite()); assert!(logits[2].is_infinite());
let mut accepted = rt("root ::= \"a\"\n", "root");
assert!(accepted.accept_bytes(b"a"));
let mut terminal_logits = vec![2.0];
assert_eq!(
mask_invalid_tokens_with_eog(&accepted, &[Vec::new()], &[0], &mut terminal_logits,),
0
);
assert_eq!(terminal_logits[0], 2.0);
}
#[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 = super::super::test_fixtures::peer_grammar("json.gbnf");
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; }
#[test]
fn greedy_probe_returns_highest_logit_valid_token() {
let runtime = rt("root ::= \"a\" | \"b\"\n", "root");
let token_bytes = vocab(&["x", "b", "a"]);
let mut logits = vec![9.0, 8.0, 7.0];
let token = sample_greedy_valid_token(&mut logits, &[], 1.0, &runtime, &token_bytes, &[]);
assert_eq!(
token, 1,
"invalid top token must yield to highest valid token"
);
}
#[test]
fn greedy_probe_suspended_runtime_is_unconstrained() {
let mut runtime = rt("root ::= \"a\"\n", "root");
runtime.set_awaiting_trigger(true);
let token_bytes = vocab(&["x", "a"]);
let mut logits = vec![9.0, 8.0];
let token = sample_greedy_valid_token(&mut logits, &[], 1.0, &runtime, &token_bytes, &[]);
assert_eq!(token, 0);
assert_eq!(logits, vec![9.0, 8.0]);
}
#[test]
fn greedy_probe_applies_repetition_penalty_before_ranking() {
let runtime = rt("root ::= \"a\" | \"b\"\n", "root");
let token_bytes = vocab(&["a", "b"]);
let mut logits = vec![10.0, 9.0];
let token = sample_greedy_valid_token(&mut logits, &[0], 2.0, &runtime, &token_bytes, &[]);
assert_eq!(token, 1, "penalized prior token must no longer win");
}
}