use super::*;
#[test]
fn frame_count_follows_the_centre_padded_hop_arithmetic() {
assert_eq!(frame_count(0), 1);
assert_eq!(frame_count(1), 1);
assert_eq!(frame_count(159), 1);
assert_eq!(frame_count(160), 2);
assert_eq!(frame_count(161), 2);
assert_eq!(frame_count(16_000), 101); assert_eq!(frame_count(48_000), 301); assert_eq!(frame_count(480_000), 3_001);
for extra in 0..HOP {
assert_eq!(
frame_count(480_000 + extra),
3_001,
"{extra} extra samples must not add a frame"
);
}
assert_eq!(frame_count(480_000 + HOP), 3_002);
}
#[test]
fn sample_bounds_round_trip_through_the_frame_bounds() {
assert_eq!(MIN_SAMPLES, 1_440);
assert_eq!(MAX_SAMPLES, 480_159);
assert_eq!(frame_count(MIN_SAMPLES), MIN_FRAMES);
assert_eq!(frame_count(MAX_SAMPLES), MAX_FRAMES);
assert_eq!(frame_count(MIN_SAMPLES - 1), MIN_FRAMES - 1);
assert_eq!(frame_count(MAX_SAMPLES + 1), MAX_FRAMES + 1);
let rate = f64::from(SAMPLE_RATE_HZ);
assert!((MIN_SAMPLES as f64 / rate - 0.09).abs() < 1e-9);
assert!((MAX_SAMPLES as f64 / rate - 30.0099).abs() < 1e-4);
}
#[test]
fn the_range_guard_accepts_and_rejects_at_both_boundaries() {
assert_eq!(
validate_frame_range(MIN_SAMPLES).expect("accepted"),
MIN_FRAMES
);
assert_eq!(
validate_frame_range(MIN_SAMPLES + 1).expect("accepted"),
MIN_FRAMES
);
assert_eq!(
validate_frame_range(MAX_SAMPLES).expect("accepted"),
MAX_FRAMES
);
assert_eq!(
validate_frame_range(MAX_SAMPLES - 1).expect("accepted"),
MAX_FRAMES
);
for rejected in [0, 1, MIN_SAMPLES - 1] {
let error = validate_frame_range(rejected).expect_err("must reject");
let Error::FrameCountOutOfRange(detail) = error else {
panic!("expected FrameCountOutOfRange for {rejected} samples, got {error:?}");
};
assert_eq!(detail.samples(), rejected);
assert_eq!(detail.frames(), frame_count(rejected));
assert!(detail.is_too_short(), "{rejected} samples is a short clip");
}
for rejected in [MAX_SAMPLES + 1, MAX_SAMPLES + HOP, 10_000_000] {
let error = validate_frame_range(rejected).expect_err("must reject");
let Error::FrameCountOutOfRange(detail) = error else {
panic!("expected FrameCountOutOfRange for {rejected} samples, got {error:?}");
};
assert_eq!(detail.samples(), rejected);
assert!(!detail.is_too_short(), "{rejected} samples is a long clip");
}
}
#[test]
fn empty_audio_is_rejected_as_a_short_clip() {
let error = validate_frame_range(0).expect_err("empty audio must be rejected");
assert!(matches!(&error, Error::FrameCountOutOfRange(d) if d.is_too_short()));
let rendered = error.to_string();
assert!(rendered.contains("0 samples"), "{rendered}");
assert!(rendered.contains("1440"), "{rendered}");
}
#[test]
fn non_finite_samples_are_reported_by_first_index() {
assert!(check_finite_samples(&[0.0, 1.0, -1.0]).is_ok());
assert!(check_finite_samples(&[]).is_ok());
for bad in [f32::NAN, f32::INFINITY, f32::NEG_INFINITY] {
let mut samples = vec![0.5f32; 16];
samples[7] = bad;
samples[11] = f32::NAN;
assert!(matches!(
check_finite_samples(&samples),
Err(Error::NonFiniteInput(7))
));
}
}
#[test]
fn options_share_one_default() {
assert_eq!(IdentifierOptions::new(), IdentifierOptions::default());
assert_eq!(IdentifierOptions::new().compute(), DEFAULT_COMPUTE);
assert_eq!(DEFAULT_COMPUTE, ComputeUnits::All);
let options = IdentifierOptions::new().with_compute(ComputeUnits::CpuAndGpu);
assert_eq!(options.compute(), ComputeUnits::CpuAndGpu);
let mut mutated = IdentifierOptions::new();
mutated.set_compute(ComputeUnits::CpuOnly);
assert_eq!(mutated.compute(), ComputeUnits::CpuOnly);
assert_ne!(mutated, IdentifierOptions::new());
}
#[test]
fn options_are_const_constructible() {
const PINNED: IdentifierOptions = IdentifierOptions::new().with_compute(ComputeUnits::CpuAndGpu);
assert_eq!(PINNED.compute(), ComputeUnits::CpuAndGpu);
}
#[cfg(feature = "serde")]
#[test]
fn options_round_trip_through_serde_by_compute_unit_name() {
let options = IdentifierOptions::new().with_compute(ComputeUnits::CpuAndNeuralEngine);
let json = serde_json::to_string(&options).expect("serialize");
assert!(
json.contains(ComputeUnits::CpuAndNeuralEngine.as_str()),
"the bridge must write the unit's own name: {json}"
);
assert_eq!(
serde_json::from_str::<IdentifierOptions>(&json).expect("deserialize"),
options
);
assert_eq!(
serde_json::from_str::<IdentifierOptions>("{}").expect("default"),
IdentifierOptions::new()
);
assert!(serde_json::from_str::<IdentifierOptions>(r#"{"compute":"quantum"}"#).is_err());
}
#[test]
fn tensor_names_are_pinned() {
assert_eq!(names::MEL_FEATURES, "mel_features");
assert_eq!(names::LOG_PROBABILITIES, "log_probabilities");
}
#[test]
fn the_long_guard_keeps_the_floor_and_drops_the_ceiling() {
assert!(validate_long_input(&vec![0.0; MIN_SAMPLES]).is_ok());
for accepted in [MAX_SAMPLES, MAX_SAMPLES + 1, 10 * MAX_SAMPLES] {
assert!(
validate_long_input(&vec![0.0; accepted]).is_ok(),
"{accepted} samples must be accepted by the long path"
);
}
for rejected in [0, 1, MIN_SAMPLES - 1] {
let error = validate_long_input(&vec![0.0; rejected]).expect_err("must reject");
let Error::FrameCountOutOfRange(detail) = error else {
panic!("expected FrameCountOutOfRange for {rejected} samples, got {error:?}");
};
assert_eq!(detail.samples(), rejected);
assert!(detail.is_too_short());
}
}
#[test]
fn the_long_guard_reports_a_clip_absolute_index() {
let mut samples = vec![0.5f32; 3 * MAX_SAMPLES];
let deep = 2 * MAX_SAMPLES + 12_345;
samples[deep] = f32::NAN;
assert!(matches!(
validate_long_input(&samples),
Err(Error::NonFiniteInput(index)) if index == deep
));
}
#[test]
fn prewarm_covers_the_default_plans_window() {
assert_eq!(
DEFAULT_WINDOW_SAMPLES as usize,
10 * SAMPLE_RATE_HZ as usize
);
assert_eq!(frame_count(DEFAULT_WINDOW_SAMPLES as usize), 1_001);
assert_eq!(
WindowPlan::new().window_samples(),
DEFAULT_WINDOW_SAMPLES,
"prewarm's clip length is the default plan's window"
);
}
fn row_with(index: usize, value: f32) -> Vec<f32> {
let mut row = vec![-14.0f32; NUM_LANGUAGES];
row[index] = value;
row
}
#[test]
fn a_positive_model_score_is_refused_at_the_door_not_ranked_above_probability_one() {
let row = row_with(94, 0.25);
let mut accumulator = aggregate::Accumulator::new(ScorePooling::default());
accumulator
.push(
&LogProbabilities::new(row.clone()),
DEFAULT_WINDOW_SAMPLES as usize,
)
.expect("a row with a finite maximum is normalizable");
let ranked = accumulator
.finish()
.expect("one window folds to itself")
.top_k(1)
.expect("top_k");
assert_eq!(ranked[0].index(), 94);
assert_eq!(ranked[0].log_probability(), 0.25);
assert!(
ranked[0].probability() > 1.0,
"the identity path returns its row verbatim, so a positive score reaches the caller as \
probability {} — an impossible confidence, which is why the door has to refuse it",
ranked[0].probability()
);
let error = validate_model_row(&row).expect_err("a positive score is not a log-probability");
assert!(
matches!(&error, Error::PositiveOutput(detail)
if detail.index() == 94 && detail.value() == 0.25),
"{error:?}"
);
assert!(error.to_string().contains("0.25"), "{error}");
}
#[test]
fn the_model_door_admits_exactly_zero_and_refuses_the_first_value_past_it() {
assert!(validate_model_row(&row_with(0, 0.0)).is_ok());
assert!(validate_model_row(&row_with(0, -0.0)).is_ok());
assert!(validate_model_row(&vec![0.0f32; NUM_LANGUAGES]).is_ok());
let smallest_positive = row_with(7, f32::MIN_POSITIVE);
assert!(matches!(
validate_model_row(&smallest_positive),
Err(Error::PositiveOutput(detail)) if detail.index() == 7
));
assert!(matches!(
validate_model_row(&row_with(7, 1e-45)),
Err(Error::PositiveOutput(_))
));
}
#[test]
fn the_model_door_still_reports_a_non_finite_score_as_it_did() {
for value in [f32::NAN, f32::INFINITY, f32::NEG_INFINITY] {
assert!(
matches!(
validate_model_row(&row_with(11, value)),
Err(Error::NonFiniteOutput(11))
),
"{value}"
);
}
let mut row = row_with(3, f32::NAN);
row[50] = 0.25;
assert!(matches!(
validate_model_row(&row),
Err(Error::NonFiniteOutput(3))
));
}
#[test]
fn the_model_door_is_the_callers_door_plus_finiteness() {
let values = [
0.0f32,
-0.0,
-1e-45,
-0.010_064,
-37.27,
f32::MIN,
f32::NEG_INFINITY,
f32::INFINITY,
f32::NAN,
f32::MIN_POSITIVE,
0.25,
22.86,
f32::MAX,
];
for value in values {
let row = row_with(5, value);
let caller_admits = LogProbabilities::try_from_slice(&row).is_ok();
let model_admits = validate_model_row(&row).is_ok();
assert_eq!(
model_admits,
caller_admits && value.is_finite(),
"{value:e}: caller door {caller_admits}, model door {model_admits}"
);
}
}
use crate::{AxisRange, FeatureInfo, ModelDescription, model::RawShapeConstraint};
fn ranged(name: &str, shape: &[usize], dtype: DataType, ranges: &[AxisRange]) -> FeatureInfo {
FeatureInfo::from_parts(
name.to_string(),
shape.to_vec(),
Some(dtype),
false,
Some(RawShapeConstraint::new(3, Vec::new(), ranges.to_vec())),
)
}
fn fixed(name: &str, shape: &[usize], dtype: DataType) -> FeatureInfo {
FeatureInfo::from_parts(
name.to_string(),
shape.to_vec(),
Some(dtype),
false,
Some(RawShapeConstraint::new(
2,
vec![shape.to_vec()],
shape.iter().map(|d| AxisRange::new(*d, 1)).collect(),
)),
)
}
const MEASURED_DEFAULT_FRAMES: usize = 301;
fn lid_description(time_axis: AxisRange) -> ModelDescription {
ModelDescription::from_parts(
vec![ranged(
names::MEL_FEATURES,
&[1, MEASURED_DEFAULT_FRAMES, N_MELS],
DataType::F32,
&[AxisRange::new(1, 1), time_axis, AxisRange::new(N_MELS, 1)],
)],
vec![fixed(
names::LOG_PROBABILITIES,
&[1, NUM_LANGUAGES],
DataType::F32,
)],
Vec::new(),
)
}
fn check(description: &ModelDescription) -> Result<()> {
crate::model::contract::check_load_contract(description, &lid_contract())
.map_err(contract_violation)
}
#[test]
fn the_contract_accepts_the_staged_artifacts_range() {
assert!(
check(&lid_description(AxisRange::inclusive(
MIN_FRAMES, MAX_FRAMES
)))
.is_ok(),
"the staged artifact's own range must satisfy the contract"
);
assert_eq!(
AxisRange::inclusive(MIN_FRAMES, MAX_FRAMES),
AxisRange::new(10, 2_992)
);
}
#[test]
fn the_contract_refuses_a_graph_whose_upper_bound_is_not_the_published_one() {
let description = lid_description(AxisRange::inclusive(MIN_FRAMES, 4_000));
assert_eq!(
description
.input(names::MEL_FEATURES)
.expect("mel_features")
.shape()[1],
MEASURED_DEFAULT_FRAMES
);
let err = check(&description).unwrap_err();
assert!(
matches!(&err, Error::ContractMismatch(m) if m.feature() == names::MEL_FEATURES),
"{err}"
);
assert!(err.to_string().contains("axis 1 10..=4000"), "{err}");
assert!(err.to_string().contains("axis 1 10..=3001"), "{err}");
}
#[test]
fn the_contract_refuses_a_graph_whose_lower_bound_is_not_the_published_one() {
let err = check(&lid_description(AxisRange::inclusive(1, MAX_FRAMES))).unwrap_err();
assert!(
matches!(&err, Error::ContractMismatch(m) if m.feature() == names::MEL_FEATURES),
"{err}"
);
assert!(err.to_string().contains("axis 1 1..=3001"), "{err}");
}
#[test]
fn the_contract_refuses_a_graph_that_pins_the_time_axis() {
let description = ModelDescription::from_parts(
vec![fixed(
names::MEL_FEATURES,
&[1, MEASURED_DEFAULT_FRAMES, N_MELS],
DataType::F32,
)],
vec![fixed(
names::LOG_PROBABILITIES,
&[1, NUM_LANGUAGES],
DataType::F32,
)],
Vec::new(),
);
let err = check(&description).unwrap_err();
assert!(
matches!(&err, Error::ContractMismatch(m)
if m.feature() == names::MEL_FEATURES && m.expected() == "range" && m.actual() == "fixed"),
"{err}"
);
}
#[test]
fn the_contract_refuses_a_different_mel_width() {
let description = ModelDescription::from_parts(
vec![ranged(
names::MEL_FEATURES,
&[1, MEASURED_DEFAULT_FRAMES, 80],
DataType::F32,
&[
AxisRange::new(1, 1),
AxisRange::inclusive(MIN_FRAMES, MAX_FRAMES),
AxisRange::new(80, 1),
],
)],
vec![fixed(
names::LOG_PROBABILITIES,
&[1, NUM_LANGUAGES],
DataType::F32,
)],
Vec::new(),
);
let err = check(&description).unwrap_err();
assert!(
matches!(&err, Error::ContractMismatch(m) if m.feature() == names::MEL_FEATURES),
"{err}"
);
}
#[test]
fn the_contract_refuses_a_different_language_count() {
let description = ModelDescription::from_parts(
lid_description(AxisRange::inclusive(MIN_FRAMES, MAX_FRAMES))
.inputs()
.to_vec(),
vec![fixed(
names::LOG_PROBABILITIES,
&[1, NUM_LANGUAGES + 1],
DataType::F32,
)],
Vec::new(),
);
let err = check(&description).unwrap_err();
assert!(
matches!(&err, Error::ContractMismatch(m) if m.feature() == names::LOG_PROBABILITIES),
"{err}"
);
}
#[test]
fn the_contract_refuses_an_extra_required_input() {
let base = lid_description(AxisRange::inclusive(MIN_FRAMES, MAX_FRAMES));
let mut inputs = base.inputs().to_vec();
inputs.push(fixed("language_prior", &[1, NUM_LANGUAGES], DataType::F32));
let description = ModelDescription::from_parts(inputs, base.outputs().to_vec(), Vec::new());
let err = check(&description).unwrap_err();
assert!(
matches!(&err, Error::UnsatisfiableInput(name) if name == "language_prior"),
"{err}"
);
}
#[test]
fn the_contract_refuses_a_graph_that_declares_state() {
let base = lid_description(AxisRange::inclusive(MIN_FRAMES, MAX_FRAMES));
let description = ModelDescription::from_parts(
base.inputs().to_vec(),
base.outputs().to_vec(),
vec![fixed("ecapa_state", &[1, 192], DataType::F32)],
);
let err = check(&description).unwrap_err();
assert!(
matches!(&err, Error::UnsatisfiableState(name) if name == "ecapa_state"),
"{err}"
);
}
#[test]
fn the_lid_contract_refuses_the_vendored_silero_bundle() {
let bundle = std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("../Models/vadkit/silero-vad-unified-256ms-v6.2.1.mlmodelc");
assert!(
bundle.is_dir(),
"the vendored silero bundle is committed, so this gate is NOT model-gated; looked for {}",
bundle.display()
);
let err = Identifier::from_file(&bundle).expect_err("silero is not this door's model");
assert!(
matches!(&err, Error::ContractMismatch(m)
if m.feature() == names::MEL_FEATURES && m.actual() == "missing"),
"{err}"
);
}
#[test]
fn both_batch_axes_are_pinned_to_one() {
let batched_input = ModelDescription::from_parts(
vec![ranged(
names::MEL_FEATURES,
&[2, MEASURED_DEFAULT_FRAMES, N_MELS],
DataType::F32,
&[
AxisRange::new(2, 1),
AxisRange::inclusive(MIN_FRAMES, MAX_FRAMES),
AxisRange::new(N_MELS, 1),
],
)],
vec![fixed(
names::LOG_PROBABILITIES,
&[1, NUM_LANGUAGES],
DataType::F32,
)],
Vec::new(),
);
assert!(matches!(check(&batched_input).unwrap_err(),
Error::ContractMismatch(m) if m.feature() == names::MEL_FEATURES));
let base = lid_description(AxisRange::inclusive(MIN_FRAMES, MAX_FRAMES));
let batched_output = ModelDescription::from_parts(
base.inputs().to_vec(),
vec![fixed(
names::LOG_PROBABILITIES,
&[2, NUM_LANGUAGES],
DataType::F32,
)],
Vec::new(),
);
assert!(matches!(check(&batched_output).unwrap_err(),
Error::ContractMismatch(m) if m.feature() == names::LOG_PROBABILITIES));
}
#[test]
fn both_named_features_element_types_are_pinned() {
let base = lid_description(AxisRange::inclusive(MIN_FRAMES, MAX_FRAMES));
let fp16_input = ModelDescription::from_parts(
vec![ranged(
names::MEL_FEATURES,
&[1, MEASURED_DEFAULT_FRAMES, N_MELS],
DataType::F16,
&[
AxisRange::new(1, 1),
AxisRange::inclusive(MIN_FRAMES, MAX_FRAMES),
AxisRange::new(N_MELS, 1),
],
)],
base.outputs().to_vec(),
Vec::new(),
);
assert!(matches!(check(&fp16_input).unwrap_err(),
Error::ContractMismatch(m) if m.feature() == names::MEL_FEATURES
&& m.expected() == "float32" && m.actual() == "float16"));
let fp16_output = ModelDescription::from_parts(
base.inputs().to_vec(),
vec![fixed(
names::LOG_PROBABILITIES,
&[1, NUM_LANGUAGES],
DataType::F16,
)],
Vec::new(),
);
assert!(matches!(check(&fp16_output).unwrap_err(),
Error::ContractMismatch(m) if m.feature() == names::LOG_PROBABILITIES));
}