use std::{path::PathBuf, sync::Mutex};
use super::*;
use crate::audio::whisper::{
backend::{ModelDims, mock::MockBackend},
options::DecodingOptions,
result::{TranscriptionProgress, TranscriptionSegment, TranscriptionTimings},
tokenizer::{SpecialTokens, WhisperTokenizer},
};
#[test]
fn state_change_callback_type_is_constructible_and_send() {
let old = StreamState::new();
let mut newer = StreamState::new();
newer.set_current_fallbacks(1);
let cb: StateChangeCallback<'_> = &|prev, next| {
assert_eq!(prev.current_fallbacks(), 0);
assert_eq!(next.current_fallbacks(), 1);
};
cb(&old, &newer);
fn assert_send<T: Send>(_: &T) {}
assert_send(&cb);
}
#[test]
fn stream_options_defaults_match_swift_init() {
let options = AudioStreamOptions::new();
assert_eq!(options.required_segments_for_confirmation(), 2);
assert_eq!(options.silence_threshold(), 0.3);
assert_eq!(options.compression_check_window(), 60);
assert!(options.use_vad());
assert_eq!(AudioStreamOptions::default(), AudioStreamOptions::new());
let options = options
.with_silence_threshold(0.5)
.with_required_segments_for_confirmation(3);
assert_eq!(options.silence_threshold(), 0.5);
assert_eq!(options.required_segments_for_confirmation(), 3);
}
#[test]
fn stream_update_vocabulary() {
assert_eq!(StreamUpdate::AwaitingVoice.as_str(), "awaiting_voice");
assert_eq!(StreamUpdate::Transcribed.to_string(), "transcribed");
assert!(StreamUpdate::AwaitingAudio.is_awaiting_audio());
}
fn progress_with(tokens: Vec<u32>, avg_logprob: Option<f32>) -> TranscriptionProgress {
let mut progress = TranscriptionProgress::new(TranscriptionTimings::new(), String::new(), tokens);
if let Some(avg) = avg_logprob {
progress.set_avg_logprob(avg);
}
progress
}
#[test]
fn should_stop_early_matches_swift_decision_table() {
let options = DecodingOptions::new();
assert_eq!(
should_stop_early(&progress_with(vec![42; 61], None), &options, 60),
Some(false)
);
assert_eq!(
should_stop_early(&progress_with(vec![42; 60], None), &options, 60),
None
);
assert_eq!(
should_stop_early(&progress_with(vec![1, 2, 3], Some(-2.0)), &options, 60),
Some(false)
);
assert_eq!(
should_stop_early(&progress_with(vec![1, 2, 3], Some(-0.1)), &options, 60),
None
);
let disabled = DecodingOptions::new().maybe_compression_ratio_threshold(None);
let varied: Vec<u32> = (0..61).collect();
assert_eq!(
should_stop_early(&progress_with(varied, None), &disabled, 60),
Some(false)
);
}
#[test]
fn energy_tracker_frames_and_first_frame_zero() {
let mut tracker = EnergyTracker::default();
let mut buffer = vec![0.001f32; 2 * ENERGY_FRAME_SAMPLES];
tracker.absorb(&buffer);
assert_eq!(tracker.relative_energies_from(0).len(), 2);
assert_eq!(tracker.relative_energies_from(0)[0], 0.0);
buffer.extend(std::iter::repeat_n(0.5, ENERGY_FRAME_SAMPLES));
tracker.absorb(&buffer);
let energies = tracker.relative_energies_from(0);
assert_eq!(energies.len(), 3);
assert!(
energies[2] > 0.5,
"loud-after-quiet is high relative energy, got {}",
energies[2]
);
buffer.extend(std::iter::repeat_n(0.5, 10));
tracker.absorb(&buffer);
assert_eq!(tracker.relative_energies_from(0).len(), 3);
}
#[test]
fn stream_state_defaults_and_pub_crate_mutation() {
let mut state = StreamState::new();
assert_eq!(state, StreamState::default());
assert_eq!(state.current_fallbacks(), 0);
assert_eq!(state.last_buffer_size(), 0);
assert_eq!(state.last_confirmed_segment_end_seconds(), 0.0);
assert!(state.buffer_energy_slice().is_empty());
assert!(state.current_text().is_empty());
assert!(state.confirmed_segments_slice().is_empty());
assert!(state.unconfirmed_segments_slice().is_empty());
assert!(state.unconfirmed_text_slice().is_empty());
state.set_current_fallbacks(2);
state.set_last_buffer_size(1_600);
state.set_last_confirmed_segment_end_seconds(3.5);
state.buffer_energy_mut().extend_from_slice(&[0.1, 0.2]);
state.set_current_text("hello");
let segment = TranscriptionSegment::new().with_text("hi");
state.confirmed_segments_mut().push(segment.clone());
state.set_unconfirmed_segments(vec![segment]);
state.set_unconfirmed_text(vec!["stale".to_string()]);
assert_eq!(state.current_fallbacks(), 2);
assert_eq!(state.last_buffer_size(), 1_600);
assert_eq!(state.last_confirmed_segment_end_seconds(), 3.5);
assert_eq!(state.buffer_energy_slice().to_vec(), vec![0.1, 0.2]);
assert_eq!(state.current_text(), "hello");
assert_eq!(state.confirmed_segments_slice().len(), 1);
assert_eq!(state.confirmed_segments_slice()[0].text(), "hi");
assert_eq!(state.unconfirmed_segments_slice().len(), 1);
assert_eq!(state.unconfirmed_segments_slice()[0].text(), "hi");
assert_eq!(
state.unconfirmed_text_slice().to_vec(),
vec!["stale".to_string()]
);
}
#[test]
fn contains_subsequence_true_false_and_edges() {
let a = TranscriptionSegment::new().with_text("a");
let b = TranscriptionSegment::new().with_text("b");
let c = TranscriptionSegment::new().with_text("c");
let haystack = [a.clone(), b.clone(), c.clone()];
assert!(contains_subsequence(&haystack, &[b.clone(), c.clone()]));
assert!(!contains_subsequence(&haystack, &[a.clone(), c.clone()]));
let z = TranscriptionSegment::new().with_text("z");
assert!(!contains_subsequence(&haystack, &[z]));
assert!(!contains_subsequence(std::slice::from_ref(&a), &haystack));
assert!(!contains_subsequence(&haystack, &[]));
}
#[cfg(feature = "serde")]
#[test]
fn stream_options_partial_config_falls_back_to_defaults() {
let partial: AudioStreamOptions = serde_json::from_str(r#"{"use_vad":false}"#).unwrap();
assert!(!partial.use_vad());
assert_eq!(
partial.required_segments_for_confirmation(),
DEFAULT_REQUIRED_SEGMENTS_FOR_CONFIRMATION
);
assert_eq!(partial.silence_threshold(), DEFAULT_SILENCE_THRESHOLD);
assert_eq!(
partial.compression_check_window(),
DEFAULT_COMPRESSION_CHECK_WINDOW
);
let round: AudioStreamOptions =
serde_json::from_str(&serde_json::to_string(&AudioStreamOptions::new()).unwrap()).unwrap();
assert_eq!(round, AudioStreamOptions::new());
}
fn tiny_tokenizer() -> WhisperTokenizer {
let root = std::env::var_os("WHISPERKIT_TEST_MODELS")
.map_or_else(crate::tests::models_root, PathBuf::from);
WhisperTokenizer::from_folder(root.join("tokenizers/whisper-tiny")).unwrap()
}
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn short_push_awaits_audio_with_waiting_text() {
let t = tiny_tokenizer();
let mock = MockBackend::new().with_dims(ModelDims::new().with_window_samples(16_000));
let fired = Mutex::new(0usize);
let callback: &(dyn Fn(&StreamState, &StreamState) + Sync) = &|_old, _new| {
*fired.lock().unwrap() += 1;
};
let mut streamer =
AudioStreamTranscriber::new(&mock, &t, DecodingOptions::new()).with_state_callback(callback);
let update = streamer.push_samples(&vec![0.5; 8_000]).unwrap();
assert!(update.is_awaiting_audio());
assert_eq!(streamer.state().current_text(), "Waiting for speech...");
assert!(
*fired.lock().unwrap() >= 2,
"buffer_energy + waiting-text assignments fired"
);
assert_eq!(mock.counters().encode_calls(), 0);
}
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn silent_audio_is_vad_skipped() {
let t = tiny_tokenizer();
let mock = MockBackend::new().with_dims(ModelDims::new().with_window_samples(16_000));
let mut streamer = AudioStreamTranscriber::new(&mock, &t, DecodingOptions::new());
let update = streamer.push_samples(&vec![0.001; 32_000]).unwrap();
assert!(update.is_awaiting_voice());
assert_eq!(
streamer.state().last_buffer_size(),
0,
"skipped runs do not consume the buffer"
);
assert_eq!(mock.counters().encode_calls(), 0);
}
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn voice_after_silence_transcribes_and_promotes_segments() {
let t = tiny_tokenizer();
let s = SpecialTokens::whisper_defaults();
let hello = t.encode(" Hello").unwrap()[0];
let mut mock = MockBackend::new().with_dims(ModelDims::new().with_window_samples(16_000));
mock.push_token_steps(&[
s.english_token(),
s.transcribe_token(),
s.time_token_begin(),
hello,
s.time_token_begin() + 50,
s.time_token_begin() + 50,
s.end_token(),
]);
let mut streamer = AudioStreamTranscriber::new(&mock, &t, DecodingOptions::new());
assert!(
streamer
.push_samples(&vec![0.001; 32_000])
.unwrap()
.is_awaiting_voice()
);
let update = streamer.push_samples(&vec![0.5; 32_000]).unwrap();
assert!(update.is_transcribed());
let state = streamer.state();
assert_eq!(state.confirmed_segments_slice().len(), 1);
assert_eq!(state.unconfirmed_segments_slice().len(), 2);
assert!((state.last_confirmed_segment_end_seconds() - 1.0).abs() < 1e-4);
assert_eq!(state.current_text(), "", "cleared after the run");
assert_eq!(state.last_buffer_size(), 64_000);
let update = streamer.push_samples(&vec![0.9; 32_000]).unwrap();
assert!(update.is_transcribed());
let state = streamer.state();
assert_eq!(state.confirmed_segments_slice().len(), 3);
assert!((state.last_confirmed_segment_end_seconds() - 3.0).abs() < 1e-4);
assert_eq!(state.unconfirmed_segments_slice().len(), 2);
}
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn segments_at_the_confirmation_boundary_stay_unconfirmed() {
let t = tiny_tokenizer();
let s = SpecialTokens::whisper_defaults();
let hello = t.encode(" Hello").unwrap()[0];
let mut mock = MockBackend::new().with_dims(ModelDims::new().with_window_samples(16_000));
mock.push_token_steps(&[
s.english_token(),
s.transcribe_token(),
s.time_token_begin(),
hello,
s.time_token_begin() + 50, s.time_token_begin() + 50,
s.end_token(),
]);
let mut stream_options = AudioStreamOptions::new();
stream_options.clear_use_vad();
let mut streamer = AudioStreamTranscriber::new(&mock, &t, DecodingOptions::new())
.with_stream_options(stream_options);
let update = streamer.push_samples(&vec![0.5; 48_000]).unwrap();
assert!(update.is_transcribed());
let state = streamer.state();
assert_eq!(
state.confirmed_segments_slice().len(),
0,
"2 <= required(2): nothing promoted"
);
assert_eq!(state.unconfirmed_segments_slice().len(), 2);
assert_eq!(
state.last_confirmed_segment_end_seconds(),
0.0,
"watermark untouched"
);
}
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn fallback_retry_stashes_superseded_text_as_unconfirmed() {
let t = tiny_tokenizer();
let s = SpecialTokens::whisper_defaults();
let hello = t.encode(" Hello").unwrap()[0];
let mut mock = MockBackend::new().with_dims(ModelDims::new().with_window_samples(16_000));
mock.push_token_steps(&[
s.english_token(),
s.transcribe_token(),
s.time_token_begin(),
hello,
s.time_token_begin() + 100, s.time_token_begin() + 100,
s.end_token(),
]);
let options = DecodingOptions::new()
.maybe_first_token_logprob_threshold(None)
.maybe_logprob_threshold(Some(-0.1))
.with_temperature_fallback_count(1);
let mut stream_options = AudioStreamOptions::new();
stream_options.clear_use_vad();
let history: Mutex<Vec<StreamState>> = Mutex::new(Vec::new());
let callback: &(dyn Fn(&StreamState, &StreamState) + Sync) =
&|_old, new| history.lock().unwrap().push(new.clone());
let mut streamer = AudioStreamTranscriber::new(&mock, &t, options)
.with_stream_options(stream_options)
.with_state_callback(callback);
let update = streamer.push_samples(&vec![0.5; 32_000]).unwrap();
assert!(update.is_transcribed());
assert_eq!(
mock.counters().resets(),
2,
"one fallback retry + one window reset -- confirms a real 2-attempt ladder ran"
);
let history = history.lock().unwrap();
let stashed = history
.iter()
.find(|snapshot| !snapshot.unconfirmed_text_slice().is_empty())
.unwrap_or_else(|| panic!("no snapshot ever recorded a stashed unconfirmed_text"));
assert_eq!(
stashed.unconfirmed_text_slice(),
&["<|startoftranscript|><|en|><|transcribe|><|0.00|> Hello".to_string()],
"attempt 0's cumulative text at its (should_stop_early-triggered) cutoff"
);
assert!(streamer.state().unconfirmed_text_slice().is_empty());
}
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn panicking_state_callback_preserves_accumulated_state() {
let t = tiny_tokenizer();
let s = SpecialTokens::whisper_defaults();
let hello = t.encode(" Hello").unwrap()[0];
let mut mock = MockBackend::new().with_dims(ModelDims::new().with_window_samples(16_000));
mock.push_token_steps(&[
s.english_token(),
s.transcribe_token(),
s.time_token_begin(),
hello,
s.time_token_begin() + 100, s.time_token_begin() + 100,
s.end_token(),
]);
let mut stream_options = AudioStreamOptions::new();
stream_options.clear_use_vad();
let fired = std::sync::atomic::AtomicUsize::new(0);
let callback: &(dyn Fn(&StreamState, &StreamState) + Sync) = &|_old, _new| {
if fired.fetch_add(1, std::sync::atomic::Ordering::SeqCst) + 1 == 3 {
panic!("scripted panic: the state callback must not panic (see StateChangeCallback's doc)");
}
};
let mut streamer = AudioStreamTranscriber::new(&mock, &t, DecodingOptions::new())
.with_stream_options(stream_options)
.with_state_callback(callback);
let samples = vec![0.5; 32_000];
let outcome = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
streamer.push_samples(&samples)
}));
assert!(
outcome.is_err(),
"the scripted callback panic must propagate out of push_samples, not get swallowed"
);
assert_eq!(
fired.load(std::sync::atomic::Ordering::SeqCst),
3,
"the run must not reach a 4th callback firing once the 3rd one panics"
);
assert!(
!streamer.state().buffer_energy_slice().is_empty(),
"pre-panic energy history survives the unwind"
);
assert!(
streamer.state().last_buffer_size() > 0,
"pre-panic buffer watermark survives the unwind"
);
let next = streamer
.push_samples(&[0.5; 1_600])
.expect("stream continues after the panic");
assert!(
next.is_awaiting_audio(),
"1 s of new audio since the surviving watermark is not enough to transcribe"
);
}
#[test]
fn incremental_energy_publication_equals_full_recompute() {
let mut incremental = EnergyTracker::default();
let mut oneshot = EnergyTracker::default();
let mut published: Vec<f32> = Vec::new();
let mut audio: Vec<f32> = Vec::new();
for chunk in [1_600usize, 2_400, 800, 4_000] {
audio.extend(std::iter::repeat_n(0.25, chunk));
incremental.absorb(&audio);
let tail = incremental.relative_energies_from(published.len());
published.extend_from_slice(&tail);
}
oneshot.absorb(&audio);
assert_eq!(published, oneshot.relative_energies_from(0));
assert_eq!(published.len(), audio.len() / ENERGY_FRAME_SAMPLES);
}