use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "UPPERCASE")]
pub enum BatchInvariance {
Pass,
Fail,
Unmeasurable,
}
pub const DEFAULT_MAX_CONSTANT_RUN: u32 = 16;
impl BatchInvariance {
#[must_use]
pub fn wire_token(self) -> &'static str {
match self {
Self::Pass => "PASS",
Self::Fail => "FAIL",
Self::Unmeasurable => "UNMEASURABLE",
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct BatchInvarianceWitness {
pub batch_invariance: BatchInvariance,
pub divergence_at: Option<u32>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub intra_agree_to: Option<u32>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub max_constant_run: Option<u32>,
pub declared_min: u32,
pub m_formed: u32,
pub source: String,
}
impl BatchInvarianceWitness {
#[must_use]
pub fn compare(m1_tokens: &[u32], batched: &[u32], declared_min: u32) -> Self {
Self::compare_batch(
m1_tokens,
&[batched],
declared_min,
DEFAULT_MAX_CONSTANT_RUN,
)
}
#[must_use]
pub fn compare_batch(
m1_tokens: &[u32],
slots: &[&[u32]],
declared_min: u32,
max_constant_run: u32,
) -> Self {
let source = "client-side token comparison across the batch's slots (m=1 recorded)";
let Some(first) = slots.first() else {
return Self {
batch_invariance: BatchInvariance::Unmeasurable,
divergence_at: None,
intra_agree_to: None,
max_constant_run: None,
declared_min,
m_formed: 0,
source: source.to_string(),
};
};
let divergence_at = Self::first_difference(m1_tokens, first);
let intra = slots
.iter()
.map(|slot| Self::agreement(first, slot))
.min()
.unwrap_or(0);
let run = slots
.iter()
.map(|slot| Self::longest_constant_run(slot))
.max()
.unwrap_or(0);
let shortest = slots.iter().map(|slot| slot.len()).min().unwrap_or(0);
let shortest_u32 = u32::try_from(shortest).unwrap_or(u32::MAX);
let verdict = if slots.iter().any(|slot| slot.is_empty()) || shortest_u32 < declared_min {
if run >= max_constant_run && max_constant_run > 0 {
BatchInvariance::Fail
} else {
BatchInvariance::Unmeasurable
}
} else if (max_constant_run > 0 && run >= max_constant_run) || intra < declared_min {
BatchInvariance::Fail
} else {
BatchInvariance::Pass
};
Self {
batch_invariance: verdict,
divergence_at,
intra_agree_to: Some(intra),
max_constant_run: Some(run),
declared_min,
m_formed: 0,
source: source.to_string(),
}
}
fn agreement(a: &[u32], b: &[u32]) -> u32 {
let agree = a.iter().zip(b.iter()).take_while(|(x, y)| x == y).count();
u32::try_from(agree).unwrap_or(u32::MAX)
}
fn first_difference(a: &[u32], b: &[u32]) -> Option<u32> {
if a.is_empty() || b.is_empty() {
return None;
}
let agree = Self::agreement(a, b);
let shortest = u32::try_from(a.len().min(b.len())).unwrap_or(u32::MAX);
(agree < shortest).then_some(agree)
}
fn longest_constant_run(tokens: &[u32]) -> u32 {
let mut best = 0u32;
let mut run = 0u32;
let mut prev: Option<u32> = None;
for &t in tokens {
run = if prev == Some(t) { run + 1 } else { 1 };
prev = Some(t);
best = best.max(run);
}
best
}
#[must_use]
pub fn formed_at(mut self, m_formed: u32, source: impl Into<String>) -> Self {
self.m_formed = m_formed;
self.source = source.into();
self
}
#[must_use]
pub fn passed(&self) -> bool {
self.batch_invariance == BatchInvariance::Pass
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_constant_token_batch_is_invalid_correctness() {
let m1: Vec<u32> = (0..128).map(|i| 1000 + i).collect();
let batched: Vec<u32> = vec![474; 128];
let w = BatchInvarianceWitness::compare(&m1, &batched, 64)
.formed_at(3, "scripts/perf041_batched_parity_probe.py");
assert_eq!(w.batch_invariance, BatchInvariance::Fail);
assert_eq!(w.divergence_at, Some(0));
assert_eq!(w.m_formed, 3);
assert!(!w.passed());
}
#[test]
fn identical_128_token_prefixes_pass() {
let tokens: Vec<u32> = (0..128).map(|i| 2000 + i).collect();
let w = BatchInvarianceWitness::compare(&tokens, &tokens, 64).formed_at(4, "perf041");
assert_eq!(w.batch_invariance, BatchInvariance::Pass);
assert_eq!(w.divergence_at, None);
assert!(w.passed());
}
#[test]
fn divergence_after_the_declared_point_still_passes_and_is_recorded() {
let m1: Vec<u32> = (0..128).map(|i| 3000 + i).collect();
let mut batched = m1.clone();
batched[100] = 9;
let w = BatchInvarianceWitness::compare(&m1, &batched, 64);
assert_eq!(w.batch_invariance, BatchInvariance::Pass);
assert_eq!(w.divergence_at, Some(100));
}
#[test]
fn the_declared_minimum_is_an_inclusive_boundary() {
let m1: Vec<u32> = (0..128).map(|i| 4000 + i).collect();
let mut at_min = m1.clone();
at_min[64] = 7;
let mut below = m1.clone();
below[63] = 7;
let w = BatchInvarianceWitness::compare_batch(&m1, &[&m1, &at_min], 64, 16);
assert_eq!(w.batch_invariance, BatchInvariance::Pass);
assert_eq!(w.intra_agree_to, Some(64));
let w = BatchInvarianceWitness::compare_batch(&m1, &[&m1, &below], 64, 16);
assert_eq!(w.batch_invariance, BatchInvariance::Fail);
assert_eq!(w.intra_agree_to, Some(63));
}
#[test]
fn a_kernel_family_flip_passes_and_records_the_m1_agreement() {
let m1: Vec<u32> = (0..128).map(|i| 6000 + i).collect();
let mut flipped = m1.clone();
for t in flipped.iter_mut().skip(3) {
*t += 500;
}
let w = BatchInvarianceWitness::compare_batch(
&m1,
&[&flipped, &flipped, &flipped, &flipped],
64,
16,
)
.formed_at(4, "perf041");
assert_eq!(w.batch_invariance, BatchInvariance::Pass);
assert_eq!(w.divergence_at, Some(3), "the m=1 agreement is recorded");
assert_eq!(w.intra_agree_to, Some(128));
assert_eq!(w.max_constant_run, Some(1));
assert!(w.passed());
}
#[test]
fn a_frozen_slot_fails_even_below_the_declared_length() {
let m1: Vec<u32> = (0..20).map(|i| 7000 + i).collect();
let frozen = vec![474_u32; 20];
let w = BatchInvarianceWitness::compare_batch(&m1, &[&m1, &frozen], 64, 16);
assert_eq!(w.batch_invariance, BatchInvariance::Fail);
assert_eq!(w.max_constant_run, Some(20));
}
#[test]
fn a_v3_0_witness_still_deserialises_and_none_fields_stay_off_the_wire() {
let old = r#"{"batch_invariance":"PASS","divergence_at":null,"declared_min":64,"m_formed":4,"source":"perf041"}"#;
let w: BatchInvarianceWitness = serde_json::from_str(old).expect("v3.0 shape reads");
assert_eq!(w.intra_agree_to, None);
let j = serde_json::to_string(&w).expect("serialises");
assert!(!j.contains("intra_agree_to"), "{j}");
let new =
BatchInvarianceWitness::compare_batch(&[1, 2, 3], &[&[1, 2, 3], &[1, 2, 3]], 2, 16);
let j = serde_json::to_string(&new).expect("serialises");
assert!(j.contains("\"intra_agree_to\":3"), "{j}");
}
#[test]
fn agreement_short_of_the_declared_point_is_unmeasurable_not_a_pass() {
let short: Vec<u32> = (0..32).map(|i| 5000 + i).collect();
let w = BatchInvarianceWitness::compare(&short, &short, 64);
assert_eq!(w.batch_invariance, BatchInvariance::Unmeasurable);
assert_eq!(w.divergence_at, None);
assert!(!w.passed(), "Unmeasurable is on the failing side of P-4");
}
#[test]
fn an_empty_stream_is_unmeasurable() {
assert_eq!(
BatchInvarianceWitness::compare(&[], &[1, 2, 3], 64).batch_invariance,
BatchInvariance::Unmeasurable
);
assert_eq!(
BatchInvarianceWitness::compare(&[1, 2, 3], &[], 64).batch_invariance,
BatchInvariance::Unmeasurable
);
}
#[test]
fn the_verdict_wire_tokens_are_the_schema_spelling() {
assert_eq!(BatchInvariance::Pass.wire_token(), "PASS");
assert_eq!(BatchInvariance::Fail.wire_token(), "FAIL");
assert_eq!(BatchInvariance::Unmeasurable.wire_token(), "UNMEASURABLE");
let j = serde_json::to_string(&BatchInvariance::Fail).expect("serialises");
assert_eq!(j, "\"FAIL\"", "serde must spell it as wire_token does");
}
}