use super::super::{N_FFT, N_MELS, PRED_HIDDEN};
use super::*;
fn word(text: &str, start: f64, end: f64) -> WordInfo {
WordInfo::new(text, start, end, 1.0, None)
}
#[test]
fn test_decoder_state_new_zeros() {
let blank_id = 1024;
let state = DecoderState::new(blank_id);
assert!(state.h.iter().all(|&v| v == 0.0));
assert!(state.c.iter().all(|&v| v == 0.0));
assert_eq!(state.prev_token, blank_id as i64);
}
#[test]
fn test_decoder_state_dimensions() {
let state = DecoderState::new(1024);
assert_eq!(state.h.len(), PRED_HIDDEN);
assert_eq!(state.c.len(), PRED_HIDDEN);
}
#[test]
fn test_decoder_state_custom_blank_id() {
let state = DecoderState::new(42);
assert_eq!(state.prev_token, 42);
}
#[test]
fn test_feature_extractor_default() {
let _fe = FeatureExtractor::default();
}
#[test]
fn test_transcript_assembler_default() {
let ta = TranscriptAssembler::default();
assert!(ta.is_empty());
}
#[test]
fn test_transcript_assembler_append_and_finalize() {
let mut asm = TranscriptAssembler::new();
assert!(asm.is_empty());
asm.append(vec![
WordInfo {
word: "hello".into(),
start: 0.0,
end: 0.5,
confidence: 0.9,
speaker: None,
},
WordInfo {
word: "world".into(),
start: 0.6,
end: 1.0,
confidence: 0.85,
speaker: None,
},
]);
assert!(!asm.is_empty());
let seg = asm.finalize(1.0);
assert_eq!(seg.text, "hello world");
assert_eq!(seg.words.len(), 2);
assert!(seg.is_final);
assert!(seg.speech_final);
assert_eq!(seg.endpoint_reason, Some(EndpointReason::Stop));
assert_eq!(seg.timestamp, 1.0);
assert!(asm.is_empty());
}
#[test]
fn test_transcript_assembler_commit_live_preserves_prefix_across_set_words() {
let mut asm = TranscriptAssembler::new();
asm.append(vec![WordInfo::new("hello", 0.0, 0.5, 0.9, None)]);
asm.commit_live();
asm.set_words(vec![WordInfo::new("world", 0.6, 1.0, 0.8, None)]);
let partial = asm.partial(1.0);
assert!(!partial.is_final);
assert!(!partial.speech_final);
assert_eq!(partial.text, "hello world");
assert_eq!(partial.words.len(), 2);
let fin = asm.finalize_with_reason(2.0, EndpointReason::Vad);
assert!(fin.speech_final);
assert_eq!(fin.endpoint_reason, Some(EndpointReason::Vad));
assert_eq!(fin.text, "hello world");
assert!(asm.is_empty());
}
#[test]
fn test_endpoint_mode_parse_token() {
assert_eq!(EndpointMode::parse_token("auto"), Some(EndpointMode::Auto));
assert_eq!(
EndpointMode::parse_token("ASSISTANT"),
Some(EndpointMode::Assistant)
);
assert_eq!(
EndpointMode::parse_token("manual"),
Some(EndpointMode::Manual)
);
assert_eq!(EndpointMode::parse_token("nope"), None);
}
#[test]
fn test_transcript_assembler_partial() {
let mut asm = TranscriptAssembler::new();
asm.append(vec![WordInfo {
word: "partial".into(),
start: 0.0,
end: 0.3,
confidence: 0.8,
speaker: None,
}]);
let seg = asm.partial(0.3);
assert_eq!(seg.text, "partial");
assert!(!seg.is_final);
assert!(!asm.is_empty());
}
#[test]
fn test_aggregate_confidence_duration_weighted() {
let words = vec![
WordInfo::new("short", 0.0, 1.0, 0.9, None),
WordInfo::new("long", 1.0, 4.0, 0.5, None),
];
let c = aggregate_confidence(&words).expect("non-empty words give a score");
assert!((c - 0.6).abs() < 1e-6, "duration-weighted mean, got {c}");
}
#[test]
fn test_aggregate_confidence_zero_duration_plain_mean() {
let words = vec![
WordInfo::new("a", 1.0, 1.0, 0.8, None),
WordInfo::new("b", 2.0, 2.0, 0.6, None),
];
let c = aggregate_confidence(&words).expect("non-empty words give a score");
assert!((c - 0.7).abs() < 1e-6, "plain mean, got {c}");
}
#[test]
fn test_aggregate_confidence_empty_words_none() {
assert_eq!(aggregate_confidence(&[]), None);
}
#[test]
fn test_assembler_fills_segment_confidence() {
let mut asm = TranscriptAssembler::new();
asm.append(vec![
WordInfo::new("hello", 0.0, 0.5, 0.9, None),
WordInfo::new("world", 0.5, 1.0, 0.8, None),
]);
let partial = asm.partial(0.5);
let c = partial.confidence.expect("partial carries the aggregate");
assert!(
(c - 0.85).abs() < 1e-6,
"equal durations → plain mean, got {c}"
);
let final_seg = asm.finalize(1.0);
let c = final_seg.confidence.expect("final carries the aggregate");
assert!((c - 0.85).abs() < 1e-6, "got {c}");
let empty = asm.finalize(1.0);
assert_eq!(empty.confidence, None);
}
#[test]
fn test_segment_confidence_omitted_from_json_when_none() {
let seg = TranscriptSegment::empty_final();
let v = serde_json::to_value(&seg).unwrap();
assert!(v.get("confidence").is_none());
}
#[test]
fn test_segment_old_json_without_confidence_still_deserializes() {
#[derive(serde::Deserialize)]
struct ClientSegmentView {
#[serde(default)]
confidence: Option<f32>,
}
let old_json = r#"{"text":"привет","words":[],"is_final":true,"timestamp":1.0}"#;
let view: ClientSegmentView = serde_json::from_str(old_json).unwrap();
assert_eq!(view.confidence, None);
}
#[test]
fn test_feature_extractor_compute_empty() {
let fe = FeatureExtractor::new();
let (mel, frames) = fe.compute(&[]);
assert_eq!(mel.len(), N_MELS);
assert_eq!(frames, 1);
assert!(mel.iter().all(|&v| v == 0.0));
}
#[test]
fn test_transcript_assembler_set_words_overwrites() {
let mut asm = TranscriptAssembler::new();
asm.set_words(vec![word("alpha", 0.0, 0.4), word("beta", 0.5, 0.9)]);
let p = asm.partial(0.0);
assert_eq!(p.text, "alpha beta");
assert_eq!(p.words.len(), 2);
asm.set_words(vec![word("gamma", 1.0, 1.4)]);
let p = asm.partial(0.0);
assert_eq!(p.text, "gamma", "set_words must overwrite, not append");
assert_eq!(p.words.len(), 1);
}
#[test]
fn test_transcript_assembler_set_words_empty_resets_text() {
let mut asm = TranscriptAssembler::new();
asm.set_words(vec![word("x", 0.0, 0.4)]);
assert!(!asm.is_empty());
asm.set_words(vec![]);
assert!(asm.is_empty(), "empty set_words clears the accumulation");
}
#[test]
fn test_feature_extractor_prepare_buffer_accumulates() {
let fe = FeatureExtractor::new();
let mut buf: Vec<f32> = Vec::new();
let usable = fe.prepare_buffer(&[0.1; 10], &mut buf);
assert_eq!(usable, None, "sub-frame input is buffered, not yet usable");
assert_eq!(buf.len(), 10, "samples are retained in the buffer");
let usable = fe.prepare_buffer(&vec![0.2; N_FFT], &mut buf);
assert!(
usable.is_some(),
"crossing a frame boundary yields a usable count"
);
}
#[test]
#[cfg_attr(miri, ignore = "mel FFT over 1s of audio is too slow under Miri")]
fn test_feature_extractor_compute_mel_reuses_buffers() {
let fe = FeatureExtractor::new();
let samples = vec![0.0f32; 16000];
let mut fft_buf = Vec::new();
let mut power_buf = Vec::new();
let mut out_buf = Vec::new();
let frames = fe.compute_mel(&samples, &mut fft_buf, &mut power_buf, &mut out_buf);
assert!(frames > 0, "1s of audio yields at least one mel frame");
assert_eq!(
out_buf.len(),
frames * N_MELS,
"output buffer holds frames * N_MELS values"
);
}