use super::tensor::Mat;
use crate::error::{FocrError, FocrResult};
pub const VOCAB_SIZE: usize = 129_280;
pub const DEFAULT_EOS_TOKEN_ID: u32 = 1;
pub const DEFAULT_NO_REPEAT_NGRAM_SIZE: usize = 35;
pub const NGRAM_WINDOW_SINGLE: usize = 128;
pub const NGRAM_WINDOW_MULTI: usize = 1024;
pub const DEFAULT_MAX_LENGTH: usize = 32_768;
pub const RUNAWAY_GUARD_ENV: &str = "FOCR_RUNAWAY_GUARD";
pub const RUNAWAY_GUARD_MIN_TOKENS: usize = 8_192;
pub const RUNAWAY_GUARD_CHECK_INTERVAL: usize = 256;
pub const RUNAWAY_GUARD_WINDOW_TOKENS: usize = 2_048;
pub const RUNAWAY_GUARD_MAX_PERIOD: usize = 256;
pub const RUNAWAY_GUARD_MIN_PERIOD_CYCLES: usize = 8;
pub const RUNAWAY_GUARD_MATCH_NUMERATOR: usize = 15;
pub const RUNAWAY_GUARD_MATCH_DENOMINATOR: usize = 16;
pub const RUNAWAY_GUARD_NGRAM_ORDER: usize = 4;
pub const RUNAWAY_GUARD_NOVELTY_NUMERATOR: usize = 1;
pub const RUNAWAY_GUARD_NOVELTY_DENOMINATOR: usize = 4;
pub const RUNAWAY_GUARD_REQUIRED_HITS: usize = 3;
#[derive(Debug, Clone)]
pub struct DecodeParams {
pub temperature: f32,
pub eos_token_id: u32,
pub max_length: usize,
pub no_repeat_ngram_size: usize,
pub ngram_window: usize,
}
impl Default for DecodeParams {
fn default() -> Self {
Self {
temperature: 0.0,
eos_token_id: DEFAULT_EOS_TOKEN_ID,
max_length: DEFAULT_MAX_LENGTH,
no_repeat_ngram_size: DEFAULT_NO_REPEAT_NGRAM_SIZE,
ngram_window: NGRAM_WINDOW_SINGLE,
}
}
}
impl DecodeParams {
#[must_use]
pub fn single_image() -> Self {
Self::default()
}
#[must_use]
pub fn multi_image() -> Self {
Self {
ngram_window: NGRAM_WINDOW_MULTI,
..Self::default()
}
}
#[must_use]
pub fn is_greedy(&self) -> bool {
#[allow(clippy::neg_cmp_op_on_partial_ord)]
!(self.temperature > 0.0)
}
#[must_use]
pub fn sliding_ngram_active(&self) -> bool {
self.no_repeat_ngram_size > 0 && self.ngram_window > 0
}
#[must_use]
pub fn matches_frozen_spec_ban(&self) -> bool {
self.no_repeat_ngram_size == DEFAULT_NO_REPEAT_NGRAM_SIZE
&& self.ngram_window == NGRAM_WINDOW_SINGLE
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct RunawayMetrics {
pub emitted_tokens: usize,
pub window_tokens: usize,
pub best_period: usize,
pub period_matches: usize,
pub period_comparisons: usize,
pub unique_ngrams: usize,
pub total_ngrams: usize,
}
impl RunawayMetrics {
#[must_use]
pub fn period_match_ppm(self) -> u32 {
ratio_ppm(self.period_matches, self.period_comparisons)
}
#[must_use]
pub fn ngram_novelty_ppm(self) -> u32 {
ratio_ppm(self.unique_ngrams, self.total_ngrams)
}
#[must_use]
pub fn is_suspicious(self) -> bool {
(self.period_matches as u128) * (RUNAWAY_GUARD_MATCH_DENOMINATOR as u128)
>= (self.period_comparisons as u128) * (RUNAWAY_GUARD_MATCH_NUMERATOR as u128)
&& (self.unique_ngrams as u128) * (RUNAWAY_GUARD_NOVELTY_DENOMINATOR as u128)
<= (self.total_ngrams as u128) * (RUNAWAY_GUARD_NOVELTY_NUMERATOR as u128)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct RunawayEvidence {
pub metrics: RunawayMetrics,
pub consecutive_hits: usize,
}
impl RunawayEvidence {
#[must_use]
pub fn timeout_error(self) -> FocrError {
FocrError::Timeout(format!(
"runaway token guard triggered at {} emitted tokens after {} consecutive \
checkpoints: period={} match={}/{} ({} ppm), token-4gram novelty={}/{} \
({} ppm); output rejected rather than silently truncated",
self.metrics.emitted_tokens,
self.consecutive_hits,
self.metrics.best_period,
self.metrics.period_matches,
self.metrics.period_comparisons,
self.metrics.period_match_ppm(),
self.metrics.unique_ngrams,
self.metrics.total_ngrams,
self.metrics.ngram_novelty_ppm(),
))
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RunawayDecision {
Continue,
Abort(RunawayEvidence),
}
#[derive(Debug, Clone)]
pub struct RunawayGuard {
enabled: bool,
next_checkpoint: Option<usize>,
consecutive_hits: usize,
last_metrics: Option<RunawayMetrics>,
terminal_evidence: Option<RunawayEvidence>,
}
impl RunawayGuard {
#[must_use]
pub fn new(enabled: bool) -> Self {
Self {
enabled,
next_checkpoint: Some(RUNAWAY_GUARD_MIN_TOKENS),
consecutive_hits: 0,
last_metrics: None,
terminal_evidence: None,
}
}
pub fn from_env_value(raw: Option<&str>) -> FocrResult<Self> {
let enabled = match raw {
None | Some("0") => false,
Some("1") => true,
Some(other) => {
return Err(FocrError::Usage(format!(
"{RUNAWAY_GUARD_ENV} must be exactly 0 or 1, got {other:?}"
)));
}
};
Ok(Self::new(enabled))
}
pub fn from_env() -> FocrResult<Self> {
match std::env::var(RUNAWAY_GUARD_ENV) {
Ok(raw) => Self::from_env_value(Some(&raw)),
Err(std::env::VarError::NotPresent) => Self::from_env_value(None),
Err(std::env::VarError::NotUnicode(_)) => Err(FocrError::Usage(format!(
"{RUNAWAY_GUARD_ENV} must be valid UTF-8 containing exactly 0 or 1"
))),
}
}
#[must_use]
pub fn enabled(&self) -> bool {
self.enabled
}
#[must_use]
pub fn last_metrics(&self) -> Option<RunawayMetrics> {
self.last_metrics
}
#[must_use]
pub fn observe(&mut self, emitted: &[u32]) -> RunawayDecision {
if !self.enabled {
return RunawayDecision::Continue;
}
if let Some(evidence) = self.terminal_evidence {
return RunawayDecision::Abort(evidence);
}
while let Some(checkpoint) = self.next_checkpoint {
if emitted.len() < checkpoint {
break;
}
self.next_checkpoint = checkpoint.checked_add(RUNAWAY_GUARD_CHECK_INTERVAL);
let Some(metrics) = analyze_runaway_suffix(&emitted[..checkpoint]) else {
continue;
};
self.last_metrics = Some(metrics);
if metrics.is_suspicious() {
self.consecutive_hits += 1;
if self.consecutive_hits >= RUNAWAY_GUARD_REQUIRED_HITS {
let evidence = RunawayEvidence {
metrics,
consecutive_hits: self.consecutive_hits,
};
self.terminal_evidence = Some(evidence);
return RunawayDecision::Abort(evidence);
}
} else {
self.consecutive_hits = 0;
}
}
RunawayDecision::Continue
}
pub fn check_after_emit(&mut self, emitted: &[u32], is_eos: bool) -> FocrResult<()> {
if let Some(evidence) = self.terminal_evidence {
return Err(evidence.timeout_error());
}
if is_eos {
return Ok(());
}
match self.observe(emitted) {
RunawayDecision::Continue => Ok(()),
RunawayDecision::Abort(evidence) => Err(evidence.timeout_error()),
}
}
}
impl Default for RunawayGuard {
fn default() -> Self {
Self::new(false)
}
}
#[must_use]
pub fn analyze_runaway_suffix(emitted: &[u32]) -> Option<RunawayMetrics> {
if emitted.len() < RUNAWAY_GUARD_WINDOW_TOKENS {
return None;
}
let window = &emitted[emitted.len() - RUNAWAY_GUARD_WINDOW_TOKENS..];
let max_period = RUNAWAY_GUARD_MAX_PERIOD.min(window.len() / RUNAWAY_GUARD_MIN_PERIOD_CYCLES);
let mut best_period = 1usize;
let mut best_matches = 0usize;
let mut best_comparisons = window.len() - 1;
for period in 1..=max_period {
let comparisons = window.len() - period;
let matches = (period..window.len())
.filter(|&i| window[i] == window[i - period])
.count();
let lhs = (matches as u128) * (best_comparisons as u128);
let rhs = (best_matches as u128) * (comparisons as u128);
if lhs > rhs || (lhs == rhs && period < best_period) {
best_period = period;
best_matches = matches;
best_comparisons = comparisons;
}
}
let total_ngrams = window.len() - RUNAWAY_GUARD_NGRAM_ORDER + 1;
let mut unique = std::collections::BTreeSet::new();
for ngram in window.windows(RUNAWAY_GUARD_NGRAM_ORDER) {
unique.insert(ngram);
}
Some(RunawayMetrics {
emitted_tokens: emitted.len(),
window_tokens: window.len(),
best_period,
period_matches: best_matches,
period_comparisons: best_comparisons,
unique_ngrams: unique.len(),
total_ngrams,
})
}
fn ratio_ppm(numerator: usize, denominator: usize) -> u32 {
if denominator == 0 {
return 0;
}
u32::try_from(((numerator as u128) * 1_000_000) / (denominator as u128)).unwrap_or(u32::MAX)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct DecodeOutput {
pub token_id: u32,
pub is_eos: bool,
}
impl DecodeOutput {
#[must_use]
pub fn new(token_id: u32, params: &DecodeParams) -> Self {
Self {
token_id,
is_eos: token_id == params.eos_token_id,
}
}
}
pub(crate) fn argmax_row(logits: &[f32]) -> FocrResult<u32> {
if logits.is_empty() {
return Err(FocrError::Other(anyhow::anyhow!(
"sampler::argmax_row: empty logits row"
)));
}
let mut best: Option<(usize, f32)> = None;
for (i, &v) in logits.iter().enumerate() {
if v.is_nan() {
return Ok(i as u32);
}
match best {
Some((_, best_val)) if v <= best_val => {}
_ => best = Some((i, v)),
}
}
let best_idx = best.map_or(0, |(i, _)| i);
Ok(best_idx as u32)
}
fn for_each_sliding_window_ngram_ban(
sequence: &[u32],
ngram_size: usize,
window: usize,
whitelist: &[u32],
vocab: usize,
mut visit: impl FnMut(usize),
) {
if ngram_size == 0 {
return;
}
let len = sequence.len();
if len < ngram_size {
return;
}
let search_start = if window == 0 {
0
} else {
len.saturating_sub(window)
};
let search_end = len - ngram_size + 1;
if search_end <= search_start {
return;
}
let prefix_len = ngram_size - 1;
let current_prefix = &sequence[len - prefix_len..];
for idx in search_start..search_end {
let ngram = &sequence[idx..idx + ngram_size];
let prefix_matches = ngram_size == 1 || &ngram[..prefix_len] == current_prefix;
if prefix_matches {
let banned = ngram[ngram_size - 1];
if whitelist.contains(&banned) {
continue;
}
let bi = banned as usize;
if bi < vocab {
visit(bi);
}
}
}
}
pub(crate) fn masked_sliding_window_logits_if_needed(
row: &[f32],
sequence: &[u32],
ngram_size: usize,
window: usize,
whitelist: &[u32],
) -> Option<Vec<f32>> {
let mut masked: Option<Vec<f32>> = None;
for_each_sliding_window_ngram_ban(sequence, ngram_size, window, whitelist, row.len(), |bi| {
let row = masked.get_or_insert_with(|| row.to_vec());
row[bi] = f32::NEG_INFINITY;
});
masked
}
pub(crate) fn collect_sliding_window_ngram_bans(
sequence: &[u32],
ngram_size: usize,
window: usize,
whitelist: &[u32],
vocab: usize,
) -> Vec<u32> {
let mut banned = Vec::new();
for_each_sliding_window_ngram_ban(sequence, ngram_size, window, whitelist, vocab, |bi| {
banned.push(bi as u32);
});
banned
}
#[cfg(test)]
fn apply_sliding_window_ngram_block(
logits: &mut [f32],
sequence: &[u32],
ngram_size: usize,
window: usize,
whitelist: &[u32],
) {
let vocab = logits.len();
for_each_sliding_window_ngram_ban(sequence, ngram_size, window, whitelist, vocab, |bi| {
logits[bi] = f32::NEG_INFINITY;
});
}
pub fn sample(logits: &Mat, generated: &[u32], params: &DecodeParams) -> FocrResult<u32> {
if logits.rows != 1 {
return Err(FocrError::Other(anyhow::anyhow!(
"sampler::sample expects a single [1, vocab] logits row, got [{}, {}]",
logits.rows,
logits.cols
)));
}
if !params.is_greedy() {
return Err(FocrError::NotImplemented(
"native_engine::sampler::sample — temperature>0 sampling is outside the greedy fp32 spine"
.into(),
));
}
let expected_len = logits.rows.checked_mul(logits.cols).ok_or_else(|| {
FocrError::Other(anyhow::anyhow!(
"sampler::sample: logits shape product overflow for [{}, {}]",
logits.rows,
logits.cols
))
})?;
if logits.data.len() != expected_len {
return Err(FocrError::Other(anyhow::anyhow!(
"sampler::sample: logits data len {} != rows*cols {} for shape [{}, {}]",
logits.data.len(),
expected_len,
logits.rows,
logits.cols
)));
}
let row = logits.row(0);
if params.no_repeat_ngram_size == 0 || generated.len() < params.no_repeat_ngram_size {
return argmax_row(row);
}
if let Some(masked) = masked_sliding_window_logits_if_needed(
row,
generated,
params.no_repeat_ngram_size,
params.ngram_window,
&[],
) {
return argmax_row(&masked);
}
argmax_row(row)
}
pub fn decode_step(
logits: &Mat,
generated: &[u32],
params: &DecodeParams,
) -> FocrResult<DecodeOutput> {
let token_id = sample(logits, generated, params)?;
Ok(DecodeOutput::new(token_id, params))
}
pub fn decode_step_premasked(logits: &Mat, params: &DecodeParams) -> FocrResult<DecodeOutput> {
if logits.rows != 1 {
return Err(FocrError::Other(anyhow::anyhow!(
"sampler::decode_step_premasked expects a single [1, vocab] logits row, got [{}, {}]",
logits.rows,
logits.cols
)));
}
if !params.is_greedy() {
return Err(FocrError::NotImplemented(
"native_engine::sampler::decode_step_premasked — temperature>0 sampling is outside the greedy fp32 spine"
.into(),
));
}
let expected_len = logits.rows.checked_mul(logits.cols).ok_or_else(|| {
FocrError::Other(anyhow::anyhow!(
"sampler::decode_step_premasked: logits shape product overflow for [{}, {}]",
logits.rows,
logits.cols
))
})?;
if logits.data.len() != expected_len {
return Err(FocrError::Other(anyhow::anyhow!(
"sampler::decode_step_premasked: logits data len {} != rows*cols {} for shape [{}, {}]",
logits.data.len(),
expected_len,
logits.rows,
logits.cols
)));
}
let token_id = argmax_row(logits.row(0))?;
Ok(DecodeOutput::new(token_id, params))
}
pub fn batched_sample(
logits: &Mat,
histories: &[&[u32]],
params: &DecodeParams,
) -> FocrResult<Vec<u32>> {
if histories.len() != logits.rows {
return Err(FocrError::Other(anyhow::anyhow!(
"sampler::batched_sample: {} histories for {} logits rows (need one history per stream)",
histories.len(),
logits.rows
)));
}
let expected_len = logits.rows.checked_mul(logits.cols).ok_or_else(|| {
FocrError::Other(anyhow::anyhow!(
"sampler::batched_sample: logits shape product overflow for [{}, {}]",
logits.rows,
logits.cols
))
})?;
if logits.data.len() != expected_len {
return Err(FocrError::Other(anyhow::anyhow!(
"sampler::batched_sample: logits data len {} != rows*cols {} for shape [{}, {}]",
logits.data.len(),
expected_len,
logits.rows,
logits.cols
)));
}
let mut tokens = Vec::with_capacity(logits.rows);
for (s, hist) in histories.iter().enumerate() {
let row = Mat::from_vec(1, logits.cols, logits.row(s).to_vec());
tokens.push(sample(&row, hist, params)?);
}
Ok(tokens)
}
pub fn batched_decode_step(
logits: &Mat,
histories: &[&[u32]],
params: &DecodeParams,
) -> FocrResult<Vec<DecodeOutput>> {
let tokens = batched_sample(logits, histories, params)?;
Ok(tokens
.into_iter()
.map(|token_id| DecodeOutput::new(token_id, params))
.collect())
}
#[cfg(test)]
mod tests {
use super::*;
fn row(v: Vec<f32>) -> Mat {
let n = v.len();
Mat::from_vec(1, n, v)
}
fn logits_preferring_35gram_banned_token() -> Mat {
let mut logits = vec![0.0; 128];
logits[7] = 10.0; logits[6] = 9.0; row(logits)
}
fn repeat_35gram_sequence(total_len: usize) -> Vec<u32> {
const NGRAM: usize = 35;
const PREFIX_LEN: usize = NGRAM - 1;
const BANNED: u32 = 7;
let prefix: Vec<u32> = (20..20 + PREFIX_LEN as u32).collect();
let min_len = PREFIX_LEN + 1 + PREFIX_LEN;
assert!(total_len >= min_len);
let mut seq = Vec::with_capacity(total_len);
seq.extend_from_slice(&prefix);
seq.push(BANNED);
seq.extend(std::iter::repeat_n(99, total_len - min_len));
seq.extend_from_slice(&prefix);
seq
}
fn params_with_window(window: usize) -> DecodeParams {
DecodeParams {
no_repeat_ngram_size: 35,
ngram_window: window,
..DecodeParams::default()
}
}
fn periodic_tokens(len: usize, period: usize) -> Vec<u32> {
(0..len).map(|i| (i % period) as u32).collect()
}
fn near_periodic_tokens(len: usize, period: usize) -> Vec<u32> {
(0..len)
.map(|i| {
if i % period == 0 {
10_000 + (i / period) as u32
} else {
(i % period) as u32
}
})
.collect()
}
#[test]
fn defaults_match_frozen_contract() {
let p = DecodeParams::default();
assert_eq!(p.temperature, 0.0);
assert_eq!(p.eos_token_id, 1);
assert_eq!(p.max_length, 32_768);
assert_eq!(p.no_repeat_ngram_size, 35);
assert_eq!(p.ngram_window, 128);
assert!(p.is_greedy());
assert!(p.sliding_ngram_active());
}
#[test]
fn runaway_guard_is_default_off_and_strictly_opt_in() {
let default_guard = RunawayGuard::default();
assert!(!default_guard.enabled());
assert!(!RunawayGuard::from_env_value(None).unwrap().enabled());
assert!(!RunawayGuard::from_env_value(Some("0")).unwrap().enabled());
assert!(RunawayGuard::from_env_value(Some("1")).unwrap().enabled());
for invalid in ["", "true", "on", "2", " 1 "] {
assert!(matches!(
RunawayGuard::from_env_value(Some(invalid)),
Err(FocrError::Usage(message))
if message.contains("must be exactly 0 or 1")
));
}
}
#[test]
fn runaway_metrics_are_exact_token_level_witnesses() {
let tokens = periodic_tokens(RUNAWAY_GUARD_WINDOW_TOKENS, 32);
let metrics = analyze_runaway_suffix(&tokens).expect("full analysis window");
assert_eq!(metrics.emitted_tokens, RUNAWAY_GUARD_WINDOW_TOKENS);
assert_eq!(metrics.window_tokens, RUNAWAY_GUARD_WINDOW_TOKENS);
assert_eq!(metrics.best_period, 32);
assert_eq!(metrics.period_matches, RUNAWAY_GUARD_WINDOW_TOKENS - 32);
assert_eq!(metrics.period_comparisons, RUNAWAY_GUARD_WINDOW_TOKENS - 32);
assert_eq!(metrics.period_match_ppm(), 1_000_000);
assert_eq!(metrics.unique_ngrams, 32);
assert_eq!(metrics.total_ngrams, RUNAWAY_GUARD_WINDOW_TOKENS - 3);
assert_eq!(metrics.ngram_novelty_ppm(), 15_647);
assert!(metrics.is_suspicious());
}
#[test]
fn runaway_guard_requires_three_consecutive_checkpoints() {
let tokens = periodic_tokens(
RUNAWAY_GUARD_MIN_TOKENS + 2 * RUNAWAY_GUARD_CHECK_INTERVAL,
32,
);
let mut guard = RunawayGuard::new(true);
assert_eq!(
guard.observe(&tokens[..RUNAWAY_GUARD_MIN_TOKENS]),
RunawayDecision::Continue
);
assert_eq!(
guard.observe(&tokens[..RUNAWAY_GUARD_MIN_TOKENS + RUNAWAY_GUARD_CHECK_INTERVAL]),
RunawayDecision::Continue
);
let RunawayDecision::Abort(evidence) = guard.observe(&tokens) else {
unreachable!("third suspicious checkpoint must abort");
};
assert_eq!(evidence.consecutive_hits, RUNAWAY_GUARD_REQUIRED_HITS);
assert_eq!(evidence.metrics.emitted_tokens, tokens.len());
assert!(evidence.metrics.is_suspicious());
}
#[test]
fn runaway_guard_catches_repeated_template_with_changing_field() {
let tokens = near_periodic_tokens(
RUNAWAY_GUARD_MIN_TOKENS + 2 * RUNAWAY_GUARD_CHECK_INTERVAL,
32,
);
let metrics = analyze_runaway_suffix(&tokens).expect("full analysis window");
assert_eq!(metrics.best_period, 32);
assert!(metrics.period_match_ppm() >= 968_000);
assert!(metrics.ngram_novelty_ppm() < 250_000);
assert!(metrics.is_suspicious());
let mut guard = RunawayGuard::new(true);
assert!(matches!(guard.observe(&tokens), RunawayDecision::Abort(_)));
}
#[test]
fn runaway_guard_rejects_with_typed_timeout_not_synthetic_eos() {
let tokens = periodic_tokens(
RUNAWAY_GUARD_MIN_TOKENS + 2 * RUNAWAY_GUARD_CHECK_INTERVAL,
19,
);
let mut guard = RunawayGuard::new(true);
let RunawayDecision::Abort(evidence) = guard.observe(&tokens) else {
unreachable!("periodic stream must produce an evidence witness");
};
let error = evidence.timeout_error();
assert_eq!(error.kind(), "timeout");
assert_eq!(error.exit_code(), crate::error::EXIT_TIMEOUT);
let message = error.to_string();
assert!(message.contains("output rejected rather than silently truncated"));
assert!(message.contains("period=19"));
assert!(message.contains("token-4gram novelty"));
}
#[test]
fn shared_commit_hook_propagates_timeout_but_never_overrides_real_eos() {
let tokens = periodic_tokens(
RUNAWAY_GUARD_MIN_TOKENS + 2 * RUNAWAY_GUARD_CHECK_INTERVAL,
29,
);
let mut eos_guard = RunawayGuard::new(true);
assert!(
eos_guard
.check_after_emit(&tokens[..RUNAWAY_GUARD_MIN_TOKENS], false)
.is_ok()
);
assert!(
eos_guard
.check_after_emit(
&tokens[..RUNAWAY_GUARD_MIN_TOKENS + RUNAWAY_GUARD_CHECK_INTERVAL],
false,
)
.is_ok()
);
assert!(eos_guard.check_after_emit(&tokens, true).is_ok());
let mut runaway_guard = RunawayGuard::new(true);
let error = runaway_guard
.check_after_emit(&tokens, false)
.expect_err("third no-EOS checkpoint must fail");
assert_eq!(error.kind(), "timeout");
assert_eq!(error.exit_code(), crate::error::EXIT_TIMEOUT);
}
#[test]
fn runaway_abort_is_terminal_and_preserves_its_first_witness() {
let mut tokens = periodic_tokens(
RUNAWAY_GUARD_MIN_TOKENS + 2 * RUNAWAY_GUARD_CHECK_INTERVAL,
23,
);
let mut guard = RunawayGuard::new(true);
let first = guard.observe(&tokens);
let RunawayDecision::Abort(first_evidence) = first else {
unreachable!("periodic stream must abort");
};
assert_eq!(guard.observe(&tokens), first);
tokens.extend(periodic_tokens(RUNAWAY_GUARD_CHECK_INTERVAL, 23));
assert_eq!(
guard.observe(&tokens),
RunawayDecision::Abort(first_evidence)
);
let sticky_error = guard
.check_after_emit(&tokens, true)
.expect_err("a terminal Abort cannot transition back to Continue on later EOS");
let first_error = first_evidence.timeout_error();
assert_eq!(sticky_error.kind(), first_error.kind());
assert_eq!(sticky_error.to_string(), first_error.to_string());
}
#[test]
fn runaway_guard_does_not_fire_on_high_novelty_tokens() {
let tokens: Vec<u32> = (0..DEFAULT_MAX_LENGTH as u32).collect();
let mut guard = RunawayGuard::new(true);
assert_eq!(guard.observe(&tokens), RunawayDecision::Continue);
let metrics = guard.last_metrics().expect("long stream was analyzed");
assert_eq!(metrics.ngram_novelty_ppm(), 1_000_000);
assert!(!metrics.is_suspicious());
}
#[test]
fn one_suspicious_checkpoint_cannot_survive_a_normal_checkpoint() {
let mut tokens = periodic_tokens(RUNAWAY_GUARD_MIN_TOKENS, 32);
let mut guard = RunawayGuard::new(true);
assert_eq!(guard.observe(&tokens), RunawayDecision::Continue);
assert!(guard.last_metrics().unwrap().is_suspicious());
tokens.extend((0..RUNAWAY_GUARD_CHECK_INTERVAL).map(|i| 1_000_000 + i as u32));
assert_eq!(guard.observe(&tokens), RunawayDecision::Continue);
assert!(!guard.last_metrics().unwrap().is_suspicious());
}
#[test]
fn runaway_decision_is_invariant_to_observer_polling_cadence() {
let tokens = near_periodic_tokens(
RUNAWAY_GUARD_MIN_TOKENS + 2 * RUNAWAY_GUARD_CHECK_INTERVAL,
47,
);
let mut one_shot = RunawayGuard::new(true);
let one_shot_decision = one_shot.observe(&tokens);
let mut incremental = RunawayGuard::new(true);
let mut incremental_decision = RunawayDecision::Continue;
for end in 1..=tokens.len() {
incremental_decision = incremental.observe(&tokens[..end]);
if matches!(incremental_decision, RunawayDecision::Abort(_)) {
break;
}
}
assert_eq!(incremental_decision, one_shot_decision);
}
#[test]
fn disabled_runaway_guard_never_inspects_or_changes_tokens() {
let tokens = periodic_tokens(DEFAULT_MAX_LENGTH, 7);
let original = tokens.clone();
let mut guard = RunawayGuard::default();
assert_eq!(guard.observe(&tokens), RunawayDecision::Continue);
assert!(guard.last_metrics().is_none());
assert_eq!(tokens, original);
}
#[test]
fn single_and_multi_windows() {
assert_eq!(DecodeParams::single_image().ngram_window, 128);
assert_eq!(DecodeParams::multi_image().ngram_window, 1024);
assert_eq!(DecodeParams::multi_image().no_repeat_ngram_size, 35);
assert!(DecodeParams::multi_image().is_greedy());
}
#[test]
fn vocab_size_constant() {
assert_eq!(VOCAB_SIZE, 129_280);
}
#[test]
fn argmax_picks_max() {
let r = row(vec![0.1, -2.0, 3.5, 3.4, 0.0]);
assert_eq!(sample(&r, &[], &DecodeParams::default()).unwrap(), 2);
}
#[test]
fn argmax_ties_pick_lowest_index() {
let r = row(vec![0.0, 5.0, 1.0, 5.0]);
let p = DecodeParams {
no_repeat_ngram_size: 0,
ngram_window: 0,
..DecodeParams::default()
};
assert_eq!(sample(&r, &[], &p).unwrap(), 1);
}
#[test]
fn argmax_selects_first_nan_like_pinned_torch() {
let r = row(vec![7.0, f32::NAN, f32::NEG_INFINITY, f32::NAN, 9.0]);
let p = DecodeParams {
no_repeat_ngram_size: 0,
ngram_window: 0,
..DecodeParams::default()
};
assert_eq!(sample(&r, &[], &p).unwrap(), 1);
}
#[test]
fn argmax_all_nan_selects_first_index() {
let r = row(vec![f32::NAN, f32::NAN, f32::NAN]);
let p = DecodeParams {
no_repeat_ngram_size: 0,
ngram_window: 0,
..DecodeParams::default()
};
assert_eq!(sample(&r, &[], &p).unwrap(), 0);
}
#[test]
fn rejects_multi_row_logits() {
let m = Mat::zeros(2, 4);
assert!(sample(&m, &[], &DecodeParams::default()).is_err());
}
#[test]
fn rejects_empty_row() {
let m = Mat::from_vec(1, 0, vec![]);
assert!(sample(&m, &[], &DecodeParams::default()).is_err());
}
#[test]
fn rejects_malformed_logits_backing_data_without_panicking() {
let m = Mat {
rows: 1,
cols: 4,
data: vec![0.0, 1.0, 2.0],
};
assert!(matches!(
sample(&m, &[], &DecodeParams::default()),
Err(err) if err.to_string().contains("logits data len 3 != rows*cols 4")
));
}
#[test]
fn temperature_sampling_not_implemented() {
let r = row(vec![1.0, 2.0, 3.0]);
let p = DecodeParams {
temperature: 0.7,
..DecodeParams::default()
};
let e = sample(&r, &[], &p).unwrap_err();
assert!(matches!(e, FocrError::NotImplemented(_)));
}
#[test]
fn decode_step_flags_eos() {
let r = row(vec![0.0, 9.0, 0.0]);
let p = DecodeParams {
no_repeat_ngram_size: 0,
ngram_window: 0,
..DecodeParams::default()
};
let out = decode_step(&r, &[], &p).unwrap();
assert_eq!(out.token_id, 1);
assert!(out.is_eos);
}
#[test]
fn decode_step_non_eos() {
let r = row(vec![0.0, 0.0, 9.0]);
let p = DecodeParams {
no_repeat_ngram_size: 0,
ngram_window: 0,
..DecodeParams::default()
};
let out = decode_step(&r, &[], &p).unwrap();
assert_eq!(out.token_id, 2);
assert!(!out.is_eos);
}
#[test]
fn ngram_size_one_bans_window_tokens() {
let r = row(vec![10.0, 5.0, 5.0]); let p = DecodeParams {
no_repeat_ngram_size: 1,
ngram_window: 8,
..DecodeParams::default()
};
let got = sample(&r, &[0, 0], &p).unwrap();
assert_eq!(got, 1);
}
#[test]
fn ngram_size_two_bans_repeat_completion() {
let r = row(vec![0.0, 0.0, 0.0, 9.0, 1.0]);
let p = DecodeParams {
no_repeat_ngram_size: 2,
ngram_window: 16,
..DecodeParams::default()
};
let got = sample(&r, &[7, 3, 7], &p).unwrap();
assert_eq!(got, 4);
}
#[test]
fn ngram_window_zero_uses_global_no_repeat_fallback() {
let r = row(vec![9.0, 1.0, 0.0]);
let p = DecodeParams {
no_repeat_ngram_size: 2,
ngram_window: 0,
..DecodeParams::default()
};
assert!(!p.sliding_ngram_active());
let got = sample(&r, &[5, 0, 5], &p).unwrap();
assert_eq!(got, 1);
}
#[test]
fn ngram_35_single_window_boundary_127_128_129() {
let r = logits_preferring_35gram_banned_token();
let p = params_with_window(NGRAM_WINDOW_SINGLE);
for (total_len, expected) in [(127usize, 6u32), (128, 6), (129, 7)] {
let seq = repeat_35gram_sequence(total_len);
assert_eq!(
sample(&r, &seq, &p).unwrap(),
expected,
"total_len={total_len} should map to token {expected}"
);
}
}
#[test]
fn ngram_35_multi_window_boundary_1023_1024_1025() {
let r = logits_preferring_35gram_banned_token();
let p = params_with_window(NGRAM_WINDOW_MULTI);
for (total_len, expected) in [(1023usize, 6u32), (1024, 6), (1025, 7)] {
let seq = repeat_35gram_sequence(total_len);
assert_eq!(
sample(&r, &seq, &p).unwrap(),
expected,
"total_len={total_len} should map to token {expected}"
);
}
}
#[test]
fn ngram_all_banned_falls_back_to_lowest_token() {
let r = row(vec![3.0, 2.0, 1.0]);
let p = DecodeParams {
no_repeat_ngram_size: 1,
ngram_window: 8,
..DecodeParams::default()
};
assert_eq!(sample(&r, &[0, 1, 2], &p).unwrap(), 0);
}
#[test]
fn sampler_boundary_masking_is_deterministic() {
let r = logits_preferring_35gram_banned_token();
let p = params_with_window(NGRAM_WINDOW_SINGLE);
let seq = repeat_35gram_sequence(128);
let first = sample(&r, &seq, &p).unwrap();
for _ in 0..8 {
assert_eq!(sample(&r, &seq, &p).unwrap(), first);
}
}
#[test]
fn ngram_two_no_ban_when_prefix_differs() {
let r = row(vec![0.0, 0.0, 9.0, 0.0]); let p = DecodeParams {
no_repeat_ngram_size: 2,
ngram_window: 16,
..DecodeParams::default()
};
let got = sample(&r, &[1, 2, 9], &p).unwrap();
assert_eq!(got, 2);
}
#[test]
fn ngram_mask_is_absent_when_scan_bans_nothing() {
let r = row(vec![0.0, 0.0, 9.0, 0.0]);
let masked = masked_sliding_window_logits_if_needed(r.row(0), &[1, 2, 9], 2, 16, &[]);
assert!(masked.is_none());
assert_eq!(sample(&r, &[1, 2, 9], &DecodeParams::default()).unwrap(), 2);
}
#[test]
fn ngram_mask_materializes_on_first_real_ban() {
let r = row(vec![0.0, 0.0, 9.0, 1.0]);
let masked = masked_sliding_window_logits_if_needed(r.row(0), &[0, 2, 0], 2, 16, &[])
.expect("token 2 should be banned");
assert_eq!(masked[2], f32::NEG_INFINITY);
assert_eq!(masked[3], 1.0);
let p = DecodeParams {
no_repeat_ngram_size: 2,
ngram_window: 16,
..DecodeParams::default()
};
assert_eq!(sample(&r, &[0, 2, 0], &p).unwrap(), 3);
}
#[test]
fn ngram_respects_window_lookback() {
let r = row(vec![0.0, 0.0, 0.0, 0.0, 0.0, 9.0]); let p = DecodeParams {
no_repeat_ngram_size: 2,
ngram_window: 2,
..DecodeParams::default()
};
let got = sample(&r, &[5, 0, 5], &p).unwrap();
assert_eq!(got, 5);
}
#[test]
fn ngram_skips_when_sequence_too_short() {
let r = row(vec![9.0, 0.0, 0.0]);
let p = DecodeParams {
no_repeat_ngram_size: 35,
ngram_window: 128,
..DecodeParams::default()
};
let got = sample(&r, &[0, 0, 0], &p).unwrap();
assert_eq!(got, 0);
}
#[test]
fn ngram_block_ignores_out_of_range_ban() {
let mut logits = vec![1.0, 2.0, 3.0];
apply_sliding_window_ngram_block(&mut logits, &[99, 99], 1, 8, &[]);
assert_eq!(logits, vec![1.0, 2.0, 3.0]);
}
#[test]
fn ngram_block_respects_whitelist() {
let mut logits = vec![1.0, 2.0, 3.0];
apply_sliding_window_ngram_block(&mut logits, &[1, 1], 1, 8, &[1]);
assert_eq!(logits, vec![1.0, 2.0, 3.0]);
}
#[test]
fn ngram_block_sets_neg_inf_on_banned() {
let mut logits = vec![0.0, 0.0, 0.0];
apply_sliding_window_ngram_block(&mut logits, &[0, 2, 0], 2, 16, &[]);
assert_eq!(logits[2], f32::NEG_INFINITY);
assert_eq!(logits[0], 0.0);
assert_eq!(logits[1], 0.0);
}
}
#[cfg(test)]
mod spec_gate_fault_injection {
use super::{DecodeParams, decode_step, sample};
use crate::native_engine::spec::{SPEC_DRAFT_MAX, accept_longest, resolve_round};
use crate::native_engine::tensor::Mat;
const V: usize = 128;
const EOS: u32 = 1;
fn peak_row(token: u32) -> Mat {
let mut r = vec![0.0f32; V];
r[token as usize] = 10.0;
Mat::from_vec(1, V, r)
}
fn row_peaked(peak: u32, runner_up: u32) -> Mat {
let mut r = vec![0.0f32; V];
r[peak as usize] = 10.0;
r[runner_up as usize] = 9.0;
Mat::from_vec(1, V, r)
}
fn params(max_length: usize) -> DecodeParams {
let mut p = DecodeParams::single_image();
p.max_length = max_length;
p
}
fn xs(s: &mut u64) -> u64 {
let mut x = *s;
x ^= x << 13;
x ^= x >> 7;
x ^= x << 17;
*s = x;
x
}
fn content_oracle(seq: &[u32]) -> Mat {
let start = seq.len().saturating_sub(3);
let mut h: u64 = 0x9E37_79B9_7F4A_7C15;
for &t in &seq[start..] {
h ^= u64::from(t).wrapping_add(0x517C_C1B7_2722_0A95);
h = h.rotate_left(23).wrapping_mul(0x2545_F491_4F6C_DD1D);
}
let pick = if seq.len() >= 5 && (h & 7) == 0 {
EOS
} else {
2 + (h % 5) as u32
};
peak_row(pick)
}
fn seq_generate(
oracle: &dyn Fn(&[u32]) -> Mat,
prompt: &[u32],
params: &DecodeParams,
) -> Vec<u32> {
let mut generated = prompt.to_vec();
let mut emitted = Vec::new();
while emitted.len() < params.max_length {
let logits = oracle(&generated);
let step = decode_step(&logits, &generated, params).expect("seq decode_step");
generated.push(step.token_id);
emitted.push(step.token_id);
if step.is_eos {
break;
}
}
emitted
}
#[derive(Default)]
struct RoundStats {
rounds: usize,
fallbacks: usize,
accepted_tokens: usize,
rejected_rounds: usize,
}
fn spec_generate_with_drafter(
oracle: &dyn Fn(&[u32]) -> Mat,
prompt: &[u32],
params: &DecodeParams,
drafter: &mut dyn FnMut(&[u32]) -> Vec<u32>,
stats: &mut RoundStats,
) -> Vec<u32> {
let mut generated = prompt.to_vec();
let mut emitted = Vec::new();
while emitted.len() < params.max_length {
let draft = drafter(&generated);
if draft.is_empty() {
stats.fallbacks += 1;
let logits = oracle(&generated);
let step = decode_step(&logits, &generated, params).expect("spec fallback step");
generated.push(step.token_id);
emitted.push(step.token_id);
if step.is_eos {
break;
}
continue;
}
let mut verify_logits: Vec<Mat> = Vec::with_capacity(draft.len() + 1);
for i in 0..=draft.len() {
let mut ctx = generated.clone();
ctx.extend_from_slice(&draft[..i]);
verify_logits.push(oracle(&ctx));
}
let emit =
resolve_round(&generated, &draft, &verify_logits, params).expect("resolve_round");
stats.rounds += 1;
stats.accepted_tokens += emit.accepted;
if emit.accepted < draft.len() {
stats.rejected_rounds += 1;
}
let mut stopped = false;
for &token in &draft[..emit.accepted] {
generated.push(token);
emitted.push(token);
if params.eos_token_id == token {
stopped = true;
break;
}
if emitted.len() >= params.max_length {
stopped = true;
break;
}
}
if stopped {
break;
}
match emit.correction {
None => break,
Some(c) => {
generated.push(c.token_id);
emitted.push(c.token_id);
if c.is_eos {
break;
}
}
}
}
emitted
}
fn greedy_lookahead(
oracle: &dyn Fn(&[u32]) -> Mat,
seq: &[u32],
k: usize,
params: &DecodeParams,
) -> Vec<u32> {
let mut ctx = seq.to_vec();
let mut out = Vec::new();
for _ in 0..k {
let step = decode_step(&oracle(&ctx), &ctx, params).expect("lookahead step");
out.push(step.token_id);
if step.is_eos {
break;
}
ctx.push(step.token_id);
}
out
}
fn assert_drafter_harmless(
label: &str,
oracle: &dyn Fn(&[u32]) -> Mat,
prompt: &[u32],
max_length: usize,
drafter: &mut dyn FnMut(&[u32]) -> Vec<u32>,
stats: &mut RoundStats,
) {
let p = params(max_length);
let seq = seq_generate(oracle, prompt, &p);
let spec = spec_generate_with_drafter(oracle, prompt, &p, drafter, stats);
assert_eq!(
spec, seq,
"{label}: adversarial drafter changed the emitted stream \
(prompt={prompt:?} ml={max_length})"
);
}
#[test]
fn adversarial_drafters_never_change_the_emitted_stream() {
let oracle: fn(&[u32]) -> Mat = content_oracle;
let prompts: [&[u32]; 3] = [&[2, 3, 4], &[5, 5, 5, 5], &[2, 6, 2, 6, 3]];
let mut garbage = RoundStats::default();
let mut wild = RoundStats::default();
let mut spam = RoundStats::default();
let mut oversized = RoundStats::default();
let mut empty = RoundStats::default();
let mut echo = RoundStats::default();
let mut seed: u64 = 0x5EC6_A7E0_D00D_F00D;
for prompt in prompts {
for ml in [8usize, 20] {
let mut s = xs(&mut seed);
assert_drafter_harmless(
"garbage",
&oracle,
prompt,
ml,
&mut |_: &[u32]| {
let len = 1 + (xs(&mut s) % 6) as usize;
(0..len).map(|_| (xs(&mut s) % 8) as u32).collect()
},
&mut garbage,
);
assert_drafter_harmless(
"wild-ids",
&oracle,
prompt,
ml,
&mut |_: &[u32]| vec![u32::MAX, V as u32, 0x7FFF_FFFF],
&mut wild,
);
assert_drafter_harmless(
"eos-spam",
&oracle,
prompt,
ml,
&mut |_: &[u32]| vec![EOS; 4],
&mut spam,
);
assert_drafter_harmless(
"oversized",
&oracle,
prompt,
ml,
&mut |_: &[u32]| vec![30u32; 64],
&mut oversized,
);
assert_drafter_harmless(
"empty",
&oracle,
prompt,
ml,
&mut |_: &[u32]| Vec::new(),
&mut empty,
);
let p_look = params(ml);
assert_drafter_harmless(
"echo",
&oracle,
prompt,
ml,
&mut |g: &[u32]| greedy_lookahead(&oracle, g, 4, &p_look),
&mut echo,
);
}
}
assert!(
garbage.rounds > 0,
"garbage drafter never reached the verifier"
);
assert!(
wild.rounds > 0,
"wild-id drafter never reached the verifier"
);
assert_eq!(
wild.rejected_rounds, wild.rounds,
"an out-of-vocab id can never equal a greedy token"
);
assert_eq!(wild.accepted_tokens, 0, "wild ids must never be accepted");
assert!(
spam.rounds > 0,
"EOS-spam drafter never reached the verifier"
);
assert!(
oversized.rounds > 0,
"oversized drafter never reached the verifier"
);
assert_eq!(
oversized.accepted_tokens, 0,
"an id the oracle never emits must never be accepted"
);
assert_eq!(
empty.rounds, 0,
"an empty draft must not reach the verifier"
);
assert!(empty.fallbacks > 0, "empty-draft fallback never exercised");
assert!(echo.rounds > 0, "echo drafter never reached the verifier");
assert!(
echo.accepted_tokens > 0,
"full agreement was never accepted — the harness has no teeth"
);
assert_eq!(
echo.rejected_rounds, 0,
"the true greedy continuation must never be rejected"
);
}
#[test]
fn oversized_draft_is_verified_position_by_position_never_trusted() {
const K: usize = 64;
let p = params(1000);
let target: Vec<u32> = (0..=K).map(|i| i as u32 + 2).collect();
let rows: Vec<Mat> = target.iter().map(|&t| peak_row(t)).collect();
let mut draft: Vec<u32> = target[..K].to_vec();
draft[5] = 99;
assert!(
draft.len() > SPEC_DRAFT_MAX,
"the fault draft must dwarf the live proposal budget"
);
let emit = resolve_round(&[], &draft, &rows, &p).unwrap();
assert_eq!(
emit.accepted, 5,
"first divergence truncates; budget ignored"
);
assert_eq!(
emit.correction.expect("correction at divergence").token_id,
target[5]
);
let emit = resolve_round(&[], &target[..K], &rows, &p).unwrap();
assert_eq!(emit.accepted, K, "genuine agreement is accepted in full");
assert_eq!(
emit.correction.expect("bonus after full accept").token_id,
target[K]
);
}
fn history_banning_token_7() -> Vec<u32> {
let prefix: Vec<u32> = (20u32..54).collect();
assert_eq!(prefix.len(), 34);
let mut h = Vec::with_capacity(69);
h.extend_from_slice(&prefix); h.push(7); h.extend_from_slice(&prefix); h
}
#[test]
fn drafter_cannot_smuggle_a_banned_token() {
let history = history_banning_token_7();
let rows = vec![row_peaked(7, 6), peak_row(8)];
let p = params(1000);
let g = sample(&rows[0], &history, &p).expect("production chooser");
assert_eq!(g, 6, "the 35-gram ban must flip greedy from 7 to 6");
let emit = resolve_round(&history, &[7], &rows, &p).unwrap();
assert_eq!(emit.accepted, 0, "the banned token must not be accepted");
assert_eq!(
emit.correction.expect("ban-aware correction").token_id,
6,
"the correction must be the ban-aware greedy token"
);
}
#[test]
fn drafter_cannot_forge_eos_termination() {
let rows = vec![peak_row(5), peak_row(6)];
let p = params(1000);
let emit = resolve_round(&[], &[EOS], &rows, &p).unwrap();
assert_eq!(emit.accepted, 0, "a forged EOS must be rejected");
let c = emit.correction.expect("correction after forged EOS");
assert_eq!(c.token_id, 5);
assert!(!c.is_eos, "the stream must not terminate on a forged EOS");
}
#[test]
fn empty_draft_resolves_to_the_pure_sequential_step() {
let rows = vec![peak_row(9)];
let p = params(1000);
let emit = resolve_round(&[2, 3], &[], &rows, &p).unwrap();
assert_eq!(emit.accepted, 0);
let c = emit
.correction
.expect("the round still yields the sequential token");
assert_eq!(c.token_id, 9);
assert_eq!(
c.token_id,
sample(&rows[0], &[2, 3], &p).expect("production chooser"),
"the empty-draft round must equal the sequential chooser"
);
}
#[test]
fn short_verify_rows_fail_closed() {
let p = params(1000);
let draft = [3u32, 4, 2];
let rows = vec![peak_row(3), peak_row(4)];
let emit = resolve_round(&[], &draft, &rows, &p).unwrap();
assert_eq!(emit.accepted, 2, "the unverifiable tail is not accepted");
assert!(
emit.correction.is_none(),
"no verify row to correct from -> no token"
);
}
#[test]
fn malformed_verify_row_never_emits_unverified_tokens() {
let p = params(1000);
let r = resolve_round(&[], &[3], &[Mat::from_vec(1, 0, vec![]), peak_row(4)], &p);
assert!(r.is_err(), "a malformed correction row must fail closed");
let good = peak_row(3);
let bad = Mat::from_vec(1, 0, vec![]);
let rows: Vec<&[f32]> = vec![good.row(0), bad.row(0)];
assert_eq!(
accept_longest(&[], &[3, 4], &rows, EOS),
1,
"acceptance must stop at the first unverifiable position"
);
}
}