use core::num::NonZeroU32;
use mediatime::Timebase;
use super::*;
use crate::{
core::oov::default_oov_decisions,
runner::aligner::{
emissions_api::{SampleSpan, SpanError},
normalizer::{NormalizationError, NormalizedText, TextNormalizer},
},
};
const TOKENIZER_JSON: &str = r#"{
"version": "1.0",
"truncation": null,
"padding": null,
"added_tokens": [],
"normalizer": null,
"pre_tokenizer": {"type": "Split", "pattern": {"Regex": ""}, "behavior": "Isolated", "invert": false},
"post_processor": null,
"decoder": null,
"model": {
"type": "WordLevel",
"vocab": {
"<pad>": 0, "<s>": 1, "</s>": 2, "<unk>": 3, "|": 4,
"E": 5, "T": 6, "A": 7, "O": 8, "N": 9, "I": 10, "H": 11, "S": 12,
"R": 13, "D": 14, "L": 15, "U": 16, "M": 17, "W": 18, "C": 19, "F": 20,
"G": 21, "Y": 22, "P": 23, "B": 24, "V": 25, "K": 26, "'": 27, "X": 28,
"J": 29, "Q": 30, "Z": 31
},
"unk_token": "<unk>"
}
}"#;
const PERMUTED_TOKENIZER_JSON: &str = r#"{
"version": "1.0",
"truncation": null,
"padding": null,
"added_tokens": [],
"normalizer": null,
"pre_tokenizer": {"type": "Split", "pattern": {"Regex": ""}, "behavior": "Isolated", "invert": false},
"post_processor": null,
"decoder": null,
"model": {
"type": "WordLevel",
"vocab": {
"<pad>": 0, "<s>": 1, "</s>": 2, "<unk>": 3, "|": 4,
"T": 5, "E": 6, "A": 7, "O": 8, "N": 9, "I": 10, "H": 11, "S": 12,
"R": 13, "D": 14, "L": 15, "U": 16, "M": 17, "W": 18, "C": 19, "F": 20,
"G": 21, "Y": 22, "P": 23, "B": 24, "V": 25, "K": 26, "'": 27, "X": 28,
"J": 29, "Q": 30, "Z": 31
},
"unk_token": "<unk>"
}
}"#;
const VOCAB_SIZE: usize = 32;
fn aligner() -> EmissionsAligner {
EmissionsAligner::builder(Lang::En, TOKENIZER_JSON.as_bytes())
.build()
.expect("a wav2vec2-shape tokenizer must build")
}
fn analysis_tb() -> Timebase {
Timebase::new(1, NonZeroU32::new(16_000).expect("16000 != 0"))
}
fn fake_encoder(prepared: &PreparedChunk<'_>, hop: usize) -> (usize, Vec<f32>) {
let t = prepared.encoder_input().len() / hop;
let mut raw = vec![0.0_f32; t * VOCAB_SIZE];
for frame in 0..t {
raw[frame * VOCAB_SIZE] = 1.0;
let token = 5 + (frame % (VOCAB_SIZE - 5));
raw[frame * VOCAB_SIZE + token] = 2.0;
}
(t, raw)
}
#[test]
fn builder_runs_the_same_construction_guards_as_from_paths() {
let a = aligner();
assert_eq!(*a.language(), Lang::En);
assert_eq!(a.blank_token_id(), 0, "<pad> is the CTC blank");
assert_eq!(a.hop_samples().get(), 320);
assert_eq!(a.vocab_size().get(), VOCAB_SIZE);
assert_eq!(a.min_speech_coverage(), SpeechCoverage::DEFAULT);
}
#[test]
fn builder_rejects_a_tokenizer_missing_the_word_delimiter() {
let no_pipe = TOKENIZER_JSON.replace("\"|\": 4,", "");
let Err(err) = EmissionsAligner::builder(Lang::En, no_pipe.as_bytes()).build() else {
panic!("an English normalizer needs a `|` delimiter");
};
let EmissionsError::Config(f) = err else {
panic!("expected a Config error");
};
assert!(
f.message().contains("`|` word-delimiter"),
"diagnostic must name the missing delimiter; got {}",
f.message()
);
}
#[test]
fn builder_rejects_a_tokenizer_with_no_blank_token() {
let no_pad = TOKENIZER_JSON.replace("\"<pad>\": 0,", "");
let Err(err) = EmissionsAligner::builder(Lang::En, no_pad.as_bytes()).build() else {
panic!("no <pad> means no CTC blank");
};
assert!(matches!(err, EmissionsError::Config(_)));
}
#[test]
fn builder_accepts_an_explicit_blank_token_id() {
let a = EmissionsAligner::builder(Lang::En, TOKENIZER_JSON.as_bytes())
.blank_token_id(2)
.min_speech_coverage(SpeechCoverage::clamped(0.25))
.hop_samples(NonZeroU32::new(160).expect("160 != 0"))
.build()
.expect("build");
assert_eq!(a.blank_token_id(), 2);
assert_eq!(a.hop_samples().get(), 160);
assert_eq!(a.min_speech_coverage().get(), 0.25);
}
#[test]
fn prepare_pads_short_audio_to_the_receptive_field_and_zeroes_non_speech() {
let a = aligner();
let samples = vec![0.5_f32; 200];
let speech = SpeechSpans::new([SampleSpan::new(0, 100).expect("ok")]);
let prepared = a
.prepare(&samples, &speech, "hello", &[], &AtomicBool::new(false))
.expect("prepare must succeed");
let buf = prepared.encoder_input();
assert_eq!(buf.len(), 400, "padded to wav2vec2's receptive field");
assert!(
buf[..100].iter().all(|&s| s == 0.5),
"speech samples survive"
);
assert!(
buf[100..].iter().all(|&s| s == 0.0),
"non-speech AND padding are exactly zero"
);
}
#[test]
fn prepare_rejects_non_finite_audio_even_outside_the_speech_spans() {
let a = aligner();
let mut samples = vec![0.1_f32; 800];
samples[700] = f32::NAN; let speech = SpeechSpans::new([SampleSpan::new(0, 100).expect("ok")]);
let Err(err) = a.prepare(&samples, &speech, "hello", &[], &AtomicBool::new(false)) else {
panic!("a NaN anywhere in the raw audio is a hard error");
};
assert!(
matches!(err, EmissionsError::NonFiniteAudio(_)),
"must be classified as non-finite audio, NOT as 'invalid configuration'"
);
}
#[test]
fn trivial_chunks_skip_the_encoder() {
let a = aligner();
let samples = vec![0.1_f32; 1600];
let speech = SpeechSpans::all_speech();
let prepared = a
.prepare(&samples, &speech, "!!!...", &[], &AtomicBool::new(false))
.expect("punctuation-only normalises to empty; that is not a failure");
assert!(prepared.is_trivial());
assert!(prepared.encoder_input().is_empty());
let emissions = Emissions::from_log_probs(
1,
NonZeroUsize::new(VOCAB_SIZE).unwrap(),
vec![-1.0; VOCAB_SIZE],
)
.expect("ok");
let clock = OutputClock::new(0, analysis_tb(), 0).expect("1/16000 is a valid output timebase");
let result = a
.finish(prepared, &emissions, clock, &AtomicBool::new(false))
.expect("a trivial chunk finishes as an empty result, not an error");
assert!(result.words().is_empty());
}
#[test]
fn finish_rejects_a_vocab_dim_that_disagrees_with_the_tokenizer() {
let a = aligner();
let samples = vec![0.1_f32; 3200];
let speech = SpeechSpans::all_speech();
let prepared = a
.prepare(&samples, &speech, "hello", &[], &AtomicBool::new(false))
.expect("prepare");
let t = prepared.encoder_input().len() / 320;
let wrong_v = NonZeroUsize::new(29).expect("29 != 0");
let emissions =
Emissions::from_logits(t, wrong_v, vec![0.5_f32; t * 29]).expect("well-formed 29-wide logits");
let clock = OutputClock::new(0, analysis_tb(), 0).expect("1/16000 is a valid output timebase");
let err = a
.finish(prepared, &emissions, clock, &AtomicBool::new(false))
.expect_err("a V mismatch must be a hard error, not a corrupt alignment");
assert!(
matches!(err, EmissionsError::VocabMismatch(_)),
"must be VocabMismatch — NOT the undifferentiated 'invalid configuration' \
the pre-existing seam mapper would have produced; got {err:?}"
);
}
#[test]
fn finish_rejects_a_frame_count_that_cannot_match_the_audio() {
let a = aligner();
let samples = vec![0.1_f32; 3200]; let speech = SpeechSpans::all_speech();
let prepared = a
.prepare(&samples, &speech, "hello", &[], &AtomicBool::new(false))
.expect("prepare");
let t = 1500;
let v = NonZeroUsize::new(VOCAB_SIZE).expect("ok");
let emissions =
Emissions::from_logits(t, v, vec![0.5_f32; t * VOCAB_SIZE]).expect("well-formed logits");
let clock = OutputClock::new(0, analysis_tb(), 0).expect("1/16000 is a valid output timebase");
let err = a
.finish(prepared, &emissions, clock, &AtomicBool::new(false))
.expect_err("T * hop must land within the chunk's real extent");
assert!(
matches!(err, EmissionsError::StrideMismatch(_)),
"must be StrideMismatch, not Config; got {err:?}"
);
}
#[test]
fn prepare_rejects_oov_decisions_resolved_for_another_language() {
use crate::core::{OovDecision, OovEvent, OovKind, ResolvedOov};
let a = aligner(); let samples = vec![0.2_f32; 16_000];
let foreign = vec![ResolvedOov::new(
OovEvent::new(OovKind::Symbol('&'), 0, 0, Lang::Ko),
OovDecision::Wildcard,
)];
let Err(err) = a.prepare(
&samples,
&SpeechSpans::all_speech(),
"&",
&foreign,
&AtomicBool::new(false),
) else {
panic!("a Korean decision must not drive an English aligner's OOV policy");
};
let EmissionsError::Tokenization(f) = err else {
panic!("expected a Tokenization error; got {err:?}");
};
assert!(
f.message().contains("oov_decisions[0].event.language")
&& f.message().contains("Ko")
&& f.message().contains("En"),
"diagnostic must cite the offending index and both languages; got {}",
f.message()
);
}
#[test]
fn prepare_accepts_oov_decisions_resolved_for_its_own_language() {
use crate::core::oov::wildcard_all_decisions;
let a = aligner();
let samples = vec![0.2_f32; 16_000];
let decisions = wildcard_all_decisions(&a.detect_oov("hello & world").expect("detect_oov"));
a.prepare(
&samples,
&SpeechSpans::all_speech(),
"hello & world",
&decisions,
&AtomicBool::new(false),
)
.expect("decisions detected from THIS aligner carry its language");
}
#[test]
fn finish_rejects_a_prepared_chunk_from_a_different_aligner() {
let a = aligner();
let b = EmissionsAligner::builder(Lang::En, PERMUTED_TOKENIZER_JSON.as_bytes())
.build()
.expect("the permuted vocab is well-formed");
assert_eq!(a.vocab_size(), b.vocab_size(), "same width");
assert_eq!(a.blank_token_id(), b.blank_token_id(), "same blank id");
assert_eq!(a.hop_samples(), b.hop_samples(), "same hop");
let samples = vec![0.2_f32; 16_000];
let prepared_from_a = a
.prepare(
&samples,
&SpeechSpans::all_speech(),
"hello",
&[],
&AtomicBool::new(false),
)
.expect("prepare on A");
assert!(!prepared_from_a.is_trivial(), "'hello' has tokens to align");
let (t, logits) = fake_encoder(&prepared_from_a, 320);
let emissions_from_b = Emissions::from_logits(t, b.vocab_size(), logits).expect("well-formed");
let clock = OutputClock::new(0, analysis_tb(), 0).expect("1/16000 is a valid output timebase");
let err = b
.finish(
prepared_from_a,
&emissions_from_b,
clock,
&AtomicBool::new(false),
)
.expect_err("A's chunk must not be finishable on B");
assert!(
matches!(err, EmissionsError::AlignerMismatch(_)),
"must be AlignerMismatch — every dimension check PASSES here, which is \
exactly why an identity is required; got {err:?}"
);
}
#[test]
fn finish_rejects_a_foreign_trivial_chunk_too() {
let a = aligner();
let b = EmissionsAligner::builder(Lang::En, PERMUTED_TOKENIZER_JSON.as_bytes())
.build()
.expect("build");
let samples = vec![0.2_f32; 1600];
let prepared_from_a = a
.prepare(
&samples,
&SpeechSpans::all_speech(),
"!!!...",
&[],
&AtomicBool::new(false),
)
.expect("prepare on A");
assert!(prepared_from_a.is_trivial());
let emissions = Emissions::from_log_probs(
1,
NonZeroUsize::new(VOCAB_SIZE).expect("32 != 0"),
vec![-1.0; VOCAB_SIZE],
)
.expect("ok");
let clock = OutputClock::new(0, analysis_tb(), 0).expect("1/16000 is a valid output timebase");
let err = b
.finish(prepared_from_a, &emissions, clock, &AtomicBool::new(false))
.expect_err("even an empty chunk from another aligner is crossed wiring");
assert!(matches!(err, EmissionsError::AlignerMismatch(_)));
}
#[test]
fn finish_accepts_the_chunk_its_own_prepare_minted() {
let a = aligner();
let samples = vec![0.2_f32; 16_000];
let prepared = a
.prepare(
&samples,
&SpeechSpans::all_speech(),
"hello",
&[],
&AtomicBool::new(false),
)
.expect("prepare");
let (t, logits) = fake_encoder(&prepared, 320);
let emissions = Emissions::from_logits(t, a.vocab_size(), logits).expect("ok");
let clock = OutputClock::new(0, analysis_tb(), 0).expect("1/16000 is a valid output timebase");
a.finish(prepared, &emissions, clock, &AtomicBool::new(false))
.expect("an aligner finishes the chunk it prepared");
}
#[test]
fn alignkit_call_site_aligns_end_to_end() {
let aligner = EmissionsAligner::builder(Lang::En, TOKENIZER_JSON.as_bytes())
.hop_samples(NonZeroU32::new(320).expect("320 != 0"))
.min_speech_coverage(SpeechCoverage::DEFAULT)
.build()
.expect("build");
let vocab = aligner.vocab_size();
let coreml_head_dim = VOCAB_SIZE;
assert_eq!(vocab.get(), coreml_head_dim);
let transcript = "hello world";
let samples = vec![0.2_f32; 16_000]; let abort = AtomicBool::new(false);
let decisions = default_oov_decisions(&aligner.detect_oov(transcript).expect("detect_oov"));
let speech = SpeechSpans::all_speech();
let prepared = aligner
.prepare(&samples, &speech, transcript, &decisions, &abort)
.expect("prepare");
if prepared.is_trivial() {
panic!("'hello world' is not trivial");
}
let (t, logits) = fake_encoder(&prepared, 320);
let emissions = Emissions::from_logits(t, vocab, logits).expect("one door, all the guards");
let clock = OutputClock::new(0, analysis_tb(), 0).expect("1/16000 is a valid output timebase");
let result = aligner
.finish(prepared, &emissions, clock, &abort)
.expect("finish");
assert!(
!result.words().is_empty(),
"a 1 s chunk of speech with a two-word transcript must align to words"
);
for w in result.words() {
let s = w.score();
assert!(
!s.is_nan() && (0.0..=1.0).contains(&s),
"every emitted Word satisfies the [0,1] NaN-free score contract; got {s}"
);
assert_eq!(w.range().timebase(), analysis_tb());
}
}
#[test]
fn prepared_chunk_is_consumed_by_finish() {
let a = aligner();
let samples = vec![0.2_f32; 16_000];
let prepared = a
.prepare(
&samples,
&SpeechSpans::all_speech(),
"hello",
&[],
&AtomicBool::new(false),
)
.expect("prepare");
let (t, logits) = fake_encoder(&prepared, 320);
let emissions = Emissions::from_logits(t, a.vocab_size(), logits).expect("ok");
let clock = OutputClock::new(0, analysis_tb(), 0).expect("1/16000 is a valid output timebase");
let _first = a.finish(prepared, &emissions, clock, &AtomicBool::new(false));
}
#[test]
fn finish_honours_the_abort_flag() {
let a = aligner();
let samples = vec![0.2_f32; 16_000];
let prepared = a
.prepare(
&samples,
&SpeechSpans::all_speech(),
"hello",
&[],
&AtomicBool::new(false),
)
.expect("prepare");
let (t, logits) = fake_encoder(&prepared, 320);
let emissions = Emissions::from_logits(t, a.vocab_size(), logits).expect("ok");
let clock = OutputClock::new(0, analysis_tb(), 0).expect("1/16000 is a valid output timebase");
let aborted = AtomicBool::new(true);
let err = a
.finish(prepared, &emissions, clock, &aborted)
.expect_err("a set abort flag must stop the pipeline");
assert!(matches!(err, EmissionsError::Aborted(_)));
}
struct PanicNormalizer;
impl TextNormalizer for PanicNormalizer {
fn normalize<'a>(&self, _text: &'a str) -> Result<NormalizedText<'a>, NormalizationError> {
panic!(
"prepare invoked the custom normalizer despite an already-set abort flag — the \
prepare-stage cancellation guard is dead"
);
}
fn use_word_delimiter(&self) -> bool {
true
}
}
fn aligner_with_panic_normalizer() -> EmissionsAligner {
EmissionsAligner::builder(Lang::En, TOKENIZER_JSON.as_bytes())
.normalizer(Box::new(PanicNormalizer))
.build()
.expect("build with the sentinel normalizer")
}
#[test]
fn prepare_aborts_before_the_custom_normalizer_runs() {
let a = aligner_with_panic_normalizer();
let samples = vec![0.2_f32; 16_000];
let aborted = AtomicBool::new(true);
let Err(err) = a.prepare(
&samples,
&SpeechSpans::all_speech(),
"hello world",
&[],
&aborted,
) else {
panic!("an already-set abort flag must stop prepare before it does any work");
};
assert!(
matches!(err, EmissionsError::Aborted(_)),
"a set abort flag must abort prepare; got {err:?}"
);
}
#[test]
fn a_malformed_decision_wins_over_a_set_abort_flag() {
use crate::core::{OovDecision, OovEvent, OovKind, ResolvedOov};
let a = aligner_with_panic_normalizer();
let samples = vec![0.2_f32; 16_000];
let foreign = vec![ResolvedOov::new(
OovEvent::new(OovKind::Symbol('&'), 0, 0, Lang::Ko),
OovDecision::Wildcard,
)];
let aborted = AtomicBool::new(true);
let Err(err) = a.prepare(
&samples,
&SpeechSpans::all_speech(),
"&",
&foreign,
&aborted,
) else {
panic!("a cross-language decision must be rejected even under cancellation");
};
let EmissionsError::Tokenization(f) = err else {
panic!("cancellation must not mask the language error; got {err:?}");
};
assert!(
f.message().contains("oov_decisions[0].event.language")
&& f.message().contains("Ko")
&& f.message().contains("En"),
"the decision error must win over abort and name both languages; got {}",
f.message()
);
}
#[test]
fn rescaled_vad_spans_reach_prepare() {
use mediatime::TimeRange;
let ms = Timebase::new(1, NonZeroU32::new(1000).expect("ok"));
let err = SpeechSpans::from_time_ranges(&[TimeRange::new(0, 500, ms)])
.expect_err("the strict bridge rejects a foreign timebase");
assert!(matches!(err, SpanError::Timebase { .. }));
let spans = SpeechSpans::from_time_ranges_rescaled(&[TimeRange::new(0, 500, ms)])
.expect("the explicit opt-in converts");
assert_eq!(spans.as_slice()[0].end(), 8_000, "500 ms == 8000 samples");
let a = aligner();
let samples = vec![0.2_f32; 16_000];
let prepared = a
.prepare(&samples, &spans, "hello", &[], &AtomicBool::new(false))
.expect("prepare with rescaled spans");
let buf = prepared.encoder_input();
assert!(buf[..8_000].iter().all(|&s| s == 0.2), "speech survives");
assert!(buf[8_000..].iter().all(|&s| s == 0.0), "the rest is masked");
}