use super::*;
#[test]
fn test_calc_output_length_known_points() {
assert_eq!(calc_output_length(768), 192);
assert_eq!(calc_output_length(1500), 375);
assert_eq!(calc_output_length(3000), 750);
assert_eq!(calc_output_length(593), 149);
assert_eq!(calc_output_length(250), 63);
assert_eq!(calc_output_length(769), 193);
assert_eq!(calc_output_length(400), 100);
}
#[test]
fn test_select_bucket_fill_floor_cases() {
let buckets = [512usize, 768, 1536, 3000];
assert_eq!(select_bucket(400, &buckets, 0.5), Some(512));
assert_eq!(select_bucket(300, &buckets, 0.5), Some(512));
assert_eq!(select_bucket(200, &buckets, 0.5), None);
assert_eq!(select_bucket(600, &buckets, 0.5), Some(768));
assert_eq!(select_bucket(800, &buckets, 0.5), Some(1536));
assert_eq!(select_bucket(769, &buckets, 0.5), Some(1536));
assert_eq!(select_bucket(4000, &buckets, 0.5), None);
assert_eq!(select_bucket(768, &buckets, 0.5), Some(768));
assert_eq!(select_bucket(512, &buckets, 0.5), Some(512));
}
#[test]
fn test_select_bucket_fill_equal_to_floor_is_selected() {
assert_eq!(select_bucket(384, &[768], 0.5), Some(768));
assert_eq!(select_bucket(256, &[512], 0.5), Some(512));
}
#[test]
fn test_select_bucket_unsorted_input() {
let buckets = [3000usize, 512, 768, 1536];
assert_eq!(select_bucket(300, &buckets, 0.5), Some(512));
assert_eq!(select_bucket(400, &buckets, 0.5), Some(512));
assert_eq!(select_bucket(600, &buckets, 0.5), Some(768));
assert_eq!(select_bucket(800, &buckets, 0.5), Some(1536));
}
#[test]
fn test_select_bucket_streaming_floor_accepts_underfilled_window() {
let buckets = [512usize, 768, 1536, 3000];
const T: usize = 249;
assert_eq!(
select_bucket(T, &buckets, FILL_FLOOR),
None,
"file-mode floor must still reject underfilled streaming-sized T"
);
assert_eq!(
select_bucket(T, &buckets, STREAMING_FILL_FLOOR),
Some(512),
"streaming floor must select bucket 512 for a 249-frame window"
);
assert_eq!(select_bucket(4000, &buckets, STREAMING_FILL_FLOOR), None);
}
#[test]
fn test_pad_time_appends_zeros_per_channel() {
let mel = vec![1.0, 2.0, 3.0, 4.0];
let padded = pad_time(&mel, 2, 2, 4);
assert_eq!(padded, vec![1.0, 2.0, 0.0, 0.0, 3.0, 4.0, 0.0, 0.0]);
assert_eq!(padded.len(), 2 * 4);
}
#[test]
fn test_pad_time_noop_when_t_equals_n() {
let mel = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
let padded = pad_time(&mel, 2, 3, 3);
assert_eq!(padded, mel);
}
#[test]
fn test_trim_time_keeps_leading_frames_per_channel() {
let out = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
let trimmed = trim_time(&out, 2, 4, 2);
assert_eq!(trimmed, vec![1.0, 2.0, 5.0, 6.0]);
assert_eq!(trimmed.len(), 2 * 2);
}
#[test]
fn test_pad_then_trim_roundtrip_recovers_prefix() {
let mel = vec![10.0, 20.0, 30.0, 40.0, 50.0, 60.0]; let padded = pad_time(&mel, 2, 3, 5);
let back = trim_time(&padded, 2, 5, 3);
assert_eq!(back, mel);
}
#[test]
fn test_file_mode_underfilled_window_routes_to_ort_fallback() {
use crate::runtime::mock::MockSession;
const SHIPPED_BUCKETS: &[usize] = &[512, 768, 1536, 3000];
const T: usize = 250;
assert_eq!(
select_bucket(T, SHIPPED_BUCKETS, FILL_FLOOR),
None,
"a 250-frame window must not select any shipped bucket at FILL_FLOOR"
);
assert_eq!(
select_bucket(T, SHIPPED_BUCKETS, STREAMING_FILL_FLOOR),
Some(512)
);
let t_prime = calc_output_length(T);
let fallback = MockSession::new(
vec![Shape::new(vec![1, N_MELS, T]), Shape::new(vec![1])],
vec![
Tensor::new(
Shape::new(vec![1, ENC_DIM, t_prime]),
TensorData::F32(vec![0.0; ENC_DIM * t_prime]),
)
.unwrap(),
Tensor::new(Shape::new(vec![1]), TensorData::I64(vec![t_prime as i64])).unwrap(),
],
);
let session = AneEncoderSession::new(Vec::new(), Box::new(fallback));
assert!(session.is_ane_encoder());
let mel = Tensor::new(
Shape::new(vec![1, N_MELS, T]),
TensorData::F32(vec![0.0; N_MELS * T]),
)
.unwrap();
let len = Tensor::new(Shape::new(vec![1]), TensorData::I64(vec![T as i64])).unwrap();
let out = session.run(&[mel, len]).expect("fallback run succeeds");
assert_eq!(out.len(), 2, "encoder emits [encoded, encoded_len]");
assert_eq!(out[0].shape().dims(), &[1, ENC_DIM, t_prime]);
match out[1].view().data() {
crate::runtime::tensor::TensorDataView::I64(v) => {
assert_eq!(v[0], t_prime as i64, "fallback encoded_len passes through")
}
other => panic!("expected I64 encoded_len, got {other:?}"),
}
}