use super::*;
use bytes::Bytes;
fn fixture_tone_pcm() -> Vec<i16> {
let wav = include_bytes!("../../../../tests/fixtures/telephony/tone_src.wav");
let data = crate::inference::audio::telephony::find_riff_chunk(wav, b"data")
.expect("fixture data chunk");
data.chunks_exact(2)
.map(|b| i16::from_le_bytes([b[0], b[1]]))
.collect()
}
fn sweep_pcm(rate: u32, seconds: f32) -> Vec<i16> {
let n = (rate as f32 * seconds) as usize;
(0..n)
.map(|i| {
let t = i as f32 / rate as f32;
let f = 50.0 + (0.45 * rate as f32 - 50.0) * (t / seconds);
(0.8 * (std::f32::consts::PI * f * t).sin() * 32000.0) as i16
})
.collect()
}
fn tone_pcm(rate: u32, seconds: f64, freq: f64) -> Vec<i16> {
let n = (f64::from(rate) * seconds) as usize;
(0..n)
.map(|i| {
let t = i as f64 / f64::from(rate);
(0.8 * (std::f64::consts::TAU * freq * t).sin() * 32000.0) as i16
})
.collect()
}
fn max_phase_drift_16k(samples: &[f32], freq: f64) -> f64 {
const WINDOW: usize = 16_000;
let phase_of = |start: usize| {
let mut re = 0.0f64;
let mut im = 0.0f64;
for (i, &v) in samples[start..start + WINDOW].iter().enumerate() {
let w = std::f64::consts::TAU * freq * ((start + i) as f64 / 16_000.0);
re += f64::from(v) * w.cos();
im += f64::from(v) * w.sin();
}
im.atan2(re)
};
let first = phase_of(0);
let mut worst = 0.0f64;
for w in 0..samples.len() / WINDOW {
let mut d = phase_of(w * WINDOW) - first;
d -= std::f64::consts::TAU * (d / std::f64::consts::TAU).round();
worst = worst.max(d.abs());
}
worst
}
fn signal_to_error_db(reference: &[f32], candidate: &[f32]) -> f64 {
let mut err = 0.0f64;
let mut sig = 0.0f64;
for (&r, &c) in reference.iter().zip(candidate) {
let d = f64::from(r) - f64::from(c);
err += d * d;
sig += f64::from(r) * f64::from(r);
}
if err == 0.0 {
return f64::INFINITY;
}
10.0 * (sig / err).log10()
}
fn assert_matches_whole_buffer(streamed: &[f32], reference: &[f32], what: &str) {
let len_diff = streamed.len().abs_diff(reference.len());
assert!(
len_diff <= 1,
"{what}: length diverged, streaming {} vs whole-buffer {}",
streamed.len(),
reference.len()
);
let cmp = streamed.len().min(reference.len());
assert!(cmp > 0, "{what}: nothing to compare");
let mut max_diff = 0.0f32;
let mut max_at = 0usize;
for i in 0..cmp {
let d = (streamed[i] - reference[i]).abs();
if d > max_diff {
max_diff = d;
max_at = i;
}
}
assert!(
max_diff <= 1e-4,
"{what}: max |streaming - whole-buffer| = {max_diff} at sample {max_at}"
);
}
fn check_mono_equivalence(pcm: &[i16], rate: u32, what: &str) {
let at_source = decode_audio_bytes(&make_wav_bytes(pcm, 16000)).unwrap();
let reference = resample(&at_source, SampleRate(rate), SampleRate(16000)).unwrap();
let streamed = decode_audio_bytes(&make_wav_bytes(pcm, rate)).unwrap();
assert_matches_whole_buffer(&streamed, &reference, what);
}
fn check_stereo_equivalence(left: &[i16], right: &[i16], rate: u32, what: &str) {
let mixed_at_source = decode_audio_bytes(&make_stereo_wav_bytes(left, right, 16000)).unwrap();
let mixed_reference = resample(&mixed_at_source, SampleRate(rate), SampleRate(16000)).unwrap();
let mixed_streamed = decode_audio_bytes(&make_stereo_wav_bytes(left, right, rate)).unwrap();
assert_matches_whole_buffer(&mixed_streamed, &mixed_reference, &format!("{what} mixed"));
let split_at_source =
decode_audio_bytes_shared_channels(Bytes::from(make_stereo_wav_bytes(left, right, 16000)))
.unwrap();
let split_streamed =
decode_audio_bytes_shared_channels(Bytes::from(make_stereo_wav_bytes(left, right, rate)))
.unwrap();
assert_eq!(split_streamed.len(), split_at_source.len());
for (c, (streamed, source)) in split_streamed.iter().zip(&split_at_source).enumerate() {
let reference = resample(source, SampleRate(rate), SampleRate(16000)).unwrap();
assert_matches_whole_buffer(streamed, &reference, &format!("{what} channel {c}"));
}
}
#[test]
#[cfg_attr(miri, ignore = "rubato sinc resampler is too slow under Miri")]
fn test_streaming_decode_matches_whole_buffer_resample_48k() {
let pcm = sweep_pcm(48_000, 2.5);
check_mono_equivalence(&pcm, 48_000, "48k sweep mono");
}
#[test]
#[cfg_attr(miri, ignore = "rubato sinc resampler is too slow under Miri")]
fn test_streaming_decode_matches_whole_buffer_resample_44k1() {
let pcm = sweep_pcm(44_100, 2.5);
check_mono_equivalence(&pcm, 44_100, "44.1k sweep mono");
}
#[test]
#[cfg_attr(miri, ignore = "rubato sinc resampler is too slow under Miri")]
fn test_streaming_decode_matches_whole_buffer_resample_stereo_48k() {
let left = sweep_pcm(48_000, 2.5);
let right: Vec<i16> = left.iter().rev().copied().collect();
check_stereo_equivalence(&left, &right, 48_000, "48k sweep stereo");
}
#[test]
#[cfg_attr(miri, ignore = "rubato sinc resampler is too slow under Miri")]
fn test_streaming_decode_matches_whole_buffer_resample_stereo_44k1() {
let left = sweep_pcm(44_100, 2.5);
let right: Vec<i16> = left.iter().rev().copied().collect();
check_stereo_equivalence(&left, &right, 44_100, "44.1k sweep stereo");
}
#[test]
#[ignore = "~50 s in debug; long-duration numeric gate, run on main push"]
#[cfg_attr(miri, ignore = "rubato sinc resampler is too slow under Miri")]
fn test_streaming_decode_long_44k1_input_holds_phase_better_than_whole_buffer() {
const FREQ: f64 = 1_000.0;
let pcm = tone_pcm(44_100, 300.0, FREQ);
let at_source = decode_audio_bytes(&make_wav_bytes(&pcm, 16000)).unwrap();
let reference = resample(&at_source, SampleRate(44_100), SampleRate(16000)).unwrap();
let streamed = decode_audio_bytes(&make_wav_bytes(&pcm, 44_100)).unwrap();
assert_eq!(
streamed.len(),
reference.len(),
"long 44.1k input: length diverged"
);
let snr = signal_to_error_db(&reference, &streamed);
assert!(
snr >= 70.0,
"long 44.1k input: streaming vs whole-buffer SNR {snr:.1} dB below the 70 dB floor"
);
let streamed_drift = max_phase_drift_16k(&streamed, FREQ);
let reference_drift = max_phase_drift_16k(&reference, FREQ);
assert!(
streamed_drift <= reference_drift,
"long 44.1k input: staged path drifted {streamed_drift:.3e} rad, \
more than the whole-buffer reference's {reference_drift:.3e} rad"
);
assert!(
streamed_drift <= 2e-4,
"long 44.1k input: staged path phase drift {streamed_drift:.3e} rad exceeds 2e-4"
);
}
#[test]
#[cfg_attr(miri, ignore = "rubato sinc resampler is too slow under Miri")]
fn test_streaming_decode_matches_whole_buffer_resample_fixture() {
let pcm = [fixture_tone_pcm(), fixture_tone_pcm()].concat();
assert!(
pcm.len() > crate::inference::audio::resample::RESAMPLE_STAGING_FRAMES,
"fixture must span more than one staging flush, got {} samples",
pcm.len()
);
check_mono_equivalence(&pcm, 48_000, "fixture 48k");
check_mono_equivalence(&pcm, 44_100, "fixture 44.1k");
}
#[test]
fn test_streaming_decode_16k_input_is_bit_identical() {
let pcm = fixture_tone_pcm();
let mono = decode_audio_bytes(&make_wav_bytes(&pcm, 16000)).unwrap();
assert_eq!(mono.len(), pcm.len());
for (i, (&raw, &got)) in pcm.iter().zip(&mono).enumerate() {
let expected = f32::from(raw) / 32768.0;
assert_eq!(
got.to_bits(),
expected.to_bits(),
"sample {i} was filtered: {got} vs {expected}"
);
}
let right: Vec<i16> = pcm.iter().rev().copied().collect();
let channels =
decode_audio_bytes_shared_channels(Bytes::from(make_stereo_wav_bytes(&pcm, &right, 16000)))
.unwrap();
assert_eq!(channels.len(), 2);
for (c, raw) in [&pcm, &right].iter().enumerate() {
assert_eq!(channels[c].len(), raw.len());
for (i, (&r, &got)) in raw.iter().zip(&channels[c]).enumerate() {
let expected = f32::from(r) / 32768.0;
assert_eq!(
got.to_bits(),
expected.to_bits(),
"channel {c} sample {i} was filtered: {got} vs {expected}"
);
}
}
}
#[test]
#[cfg_attr(miri, ignore = "rubato sinc resampler is too slow under Miri")]
fn test_telephony_raw_streaming_matches_whole_buffer_resample() {
let pcm = sweep_pcm(8_000, 12.5);
assert!(
pcm.len() > crate::inference::audio::resample::RESAMPLE_STAGING_FRAMES,
"clip must span more than one staging flush, got {} samples",
pcm.len()
);
let mut encoder = audio_codec::pcmu::PcmuEncoder::new();
let encoded = audio_codec::Encoder::encode(&mut encoder, &pcm);
let mut decoder = audio_codec::pcmu::PcmuDecoder::new();
let round_tripped = audio_codec::Decoder::decode(&mut decoder, &encoded);
let at_source: Vec<f32> = round_tripped
.iter()
.map(|&s| f32::from(s) / 32768.0)
.collect();
let reference = resample(&at_source, SampleRate(8_000), SampleRate(16_000)).unwrap();
let streamed = decode_telephony_raw(&encoded, TelephonyCodec::Pcmu, 8_000).unwrap();
assert_matches_whole_buffer(&streamed, &reference, "pcmu 8k");
}
#[test]
fn test_audio_chunks_match_flat_decode_chunked() {
for &n in &[1usize, 999, 16_000, 16_001, 48_000, 120_000] {
for &chunk in &[16_000usize, 640, 7_000] {
let src: Vec<f32> = (0..n)
.map(|i| 0.4 * ((i as f32) * 0.017).sin() + 0.2 * ((i as f32) * 0.0031).sin())
.collect();
let wav = bytes::Bytes::from(encode_wav_pcm16(&src, 16000));
let flat = decode_audio_bytes(&wav).expect("flat decode");
let mut chunks = AudioChunks::from_bytes(wav, chunk, None).expect("open chunks");
let mut got: Vec<Vec<f32>> = Vec::new();
while let Some(c) = chunks.next_chunk().expect("chunk") {
got.push(c.to_vec());
}
let want: Vec<Vec<f32>> = flat.chunks(chunk).map(<[f32]>::to_vec).collect();
assert_eq!(got, want, "n={n} chunk={chunk}");
assert_eq!(
chunks.total_16k_samples(),
flat.len(),
"n={n} chunk={chunk}"
);
}
}
}
#[test]
fn test_audio_chunks_honour_max_audio_secs() {
let src = vec![0.1f32; 16_000 * 5];
let wav = bytes::Bytes::from(encode_wav_pcm16(&src, 16000));
let mut ok = AudioChunks::from_bytes(wav.clone(), 16_000, None).expect("open");
let mut total = 0;
while let Some(c) = ok.next_chunk().expect("chunk") {
total += c.len();
}
assert_eq!(total, src.len());
let mut capped = AudioChunks::from_bytes(wav, 16_000, Some(1.0)).expect("open");
let err = loop {
match capped.next_chunk() {
Ok(Some(_)) => continue,
Ok(None) => panic!("a 5 s clip must not drain under a 1 s limit"),
Err(e) => break e,
}
};
assert!(
matches!(
err.downcast_ref::<crate::error::GigasttError>(),
Some(crate::error::GigasttError::AudioTooLong { .. })
),
"expected a typed AudioTooLong, got: {err:#}"
);
}
#[test]
fn test_dual_mono_detector_matches_batch_correlation() {
let n = 40_000;
let base: Vec<f32> = (0..n)
.map(|i| 0.5 * ((i as f32) * 0.013).sin() + 0.2 * ((i as f32) * 0.0007).cos())
.collect();
let other: Vec<f32> = (0..n)
.map(|i| 0.5 * ((i as f32) * 0.031 + 1.7).sin() - 0.3 * ((i as f32) * 0.0021).cos())
.collect();
let cases: Vec<(&str, Vec<f32>, Vec<f32>)> = vec![
("identical", base.clone(), base.clone()),
(
"attenuated",
base.clone(),
base.iter().map(|v| v * 0.85).collect(),
),
(
"mostly-same",
base.clone(),
base.iter()
.zip(&other)
.map(|(l, r)| l * 0.9 + r * 0.1)
.collect(),
),
("independent", base.clone(), other.clone()),
("inverted", base.clone(), base.iter().map(|v| -v).collect()),
("right-silent", base.clone(), vec![0.0; n]),
("both-silent", vec![0.0; n], vec![0.0; n]),
(
"dc-offset",
base.iter().map(|v| v + 0.7).collect(),
base.iter().map(|v| v + 0.7).collect(),
),
];
for (label, left, right) in cases {
let batch = crate::inference::audio::decode::normalized_correlation_for_test(&left, &right);
let mut det = DualMonoDetector::new();
let mut i = 0;
let mut step = 1;
while i < left.len() {
let end = (i + step).min(left.len());
det.push(&left[i..end], &right[i..end]);
i = end;
step = step % 1_237 + 1;
}
assert!(
(det.correlation() - batch).abs() < 1e-6,
"{label}: streaming {} vs batch {batch}",
det.correlation()
);
let want = is_dual_mono(&[left, right]);
assert_eq!(det.is_dual_mono(), want, "{label}: verdict diverged");
}
}
#[test]
fn test_dual_mono_detector_empty_is_not_dual_mono() {
let det = DualMonoDetector::new();
assert!(!det.is_dual_mono());
assert_eq!(det.correlation(), 0.0);
let mut det = DualMonoDetector::new();
det.push(&[], &[]);
assert!(!det.is_dual_mono());
}
#[test]
fn test_dual_mono_detector_uses_the_overlap_only() {
let a = vec![0.3f32, -0.4, 0.5, 0.9, -0.1];
let b = vec![0.3f32, -0.4, 0.5];
let mut det = DualMonoDetector::new();
det.push(&a, &b);
let batch = crate::inference::audio::decode::normalized_correlation_for_test(&a[..3], &b);
assert!((det.correlation() - batch).abs() < 1e-9);
}
#[cfg(feature = "file-decode")]
fn stereo_wav(left: &[f32], right: &[f32], rate: u32) -> bytes::Bytes {
let frames = left.len().min(right.len());
let data_bytes = (frames * 4) as u32;
let mut w = Vec::with_capacity(44 + data_bytes as usize);
w.extend_from_slice(b"RIFF");
w.extend_from_slice(&(36 + data_bytes).to_le_bytes());
w.extend_from_slice(b"WAVE");
w.extend_from_slice(b"fmt ");
w.extend_from_slice(&16u32.to_le_bytes());
w.extend_from_slice(&1u16.to_le_bytes());
w.extend_from_slice(&2u16.to_le_bytes());
w.extend_from_slice(&rate.to_le_bytes());
w.extend_from_slice(&(rate * 4).to_le_bytes());
w.extend_from_slice(&4u16.to_le_bytes());
w.extend_from_slice(&16u16.to_le_bytes());
w.extend_from_slice(b"data");
w.extend_from_slice(&data_bytes.to_le_bytes());
let q = |s: f32| (s.clamp(-1.0, 1.0) * i16::MAX as f32) as i16;
for i in 0..frames {
w.extend_from_slice(&q(left[i]).to_le_bytes());
w.extend_from_slice(&q(right[i]).to_le_bytes());
}
bytes::Bytes::from(w)
}
#[cfg(feature = "file-decode")]
#[test]
fn test_scan_channels_matches_batch_dual_mono_verdict() {
for rate in [16_000u32, 48_000] {
let n = rate as usize * 2;
let a: Vec<f32> = (0..n)
.map(|i| 0.5 * ((i as f32) * 0.011).sin() + 0.15 * ((i as f32) * 0.0009).cos())
.collect();
let b: Vec<f32> = (0..n)
.map(|i| 0.45 * ((i as f32) * 0.029 + 0.9).sin())
.collect();
for (label, left, right) in [
("identical", a.clone(), a.clone()),
("attenuated", a.clone(), a.iter().map(|v| v * 0.8).collect()),
("stereo", a.clone(), b.clone()),
] {
let wav = stereo_wav(&left, &right, rate);
let batch = decode_audio_bytes_shared_channels(wav.clone()).expect("batch");
let want = is_dual_mono(&batch);
let scan = scan_channels(wav, None).expect("scan");
assert_eq!(scan.channels, 2, "{label} @{rate}");
assert_eq!(
scan.dual_mono, want,
"{label} @{rate}: scan verdict diverged from batch"
);
assert_eq!(
scan.mono_fallback_reason().is_some(),
want,
"{label} @{rate}"
);
}
}
}
#[cfg(feature = "file-decode")]
#[test]
fn test_scan_channels_non_stereo_is_header_only() {
let mono = encode_wav_pcm16(&vec![0.2f32; 16_000], 16000);
let scan = scan_channels(bytes::Bytes::from(mono), None).expect("scan mono");
assert_eq!(scan.channels, 1);
assert!(!scan.dual_mono);
assert_eq!(scan.mono_fallback_reason(), Some("mono audio"));
}
#[cfg(feature = "file-decode")]
#[test]
fn test_channel_scan_fallback_reasons() {
let r = |channels, dual_mono| {
ChannelScan {
channels,
dual_mono,
}
.mono_fallback_reason()
};
assert_eq!(r(0, false), Some("no channels"));
assert_eq!(r(1, false), Some("mono audio"));
assert_eq!(r(2, true), Some("dual-mono audio"));
assert_eq!(r(2, false), None);
assert_eq!(r(6, false), Some("more than two channels"));
}