#[cfg(test)]
mod tests {
use crate::loader::namb::*;
use anyhow::Result;
fn build_valid_namb_v1(w_floats: &[f32]) -> Vec<u8> {
let header_size = std::mem::size_of::<NambHeader>();
let mut data = vec![0u8; header_size + w_floats.len() * 4];
let header = unsafe { &mut *data.as_mut_ptr().cast::<NambHeader>() };
header.magic = 0x4E414D42; header.version = 1;
header.weights_offset = header_size as u32;
header.sample_rate = 48000.0;
header.input_level_dbu = 12.0;
header.output_level_dbu = -6.0;
header.version_str[0..5].copy_from_slice(b"1.0.0");
for (i, &f) in w_floats.iter().enumerate() {
let offset = header_size + i * 4;
data[offset..offset + 4].copy_from_slice(&f.to_le_bytes());
}
header.crc32 = crc32_ieee(&data[header_size..]);
data
}
#[test]
fn test_parse_namb_v1() -> Result<()> {
let w = [0.1f32, -0.5f32, 1.0f32];
let data = build_valid_namb_v1(&w);
let parsed = parse_namb(&data)?;
assert_eq!(parsed.weights, w);
assert_eq!(parsed.weights_layout, WeightsLayout::Original);
assert_eq!(parsed.sample_rate, Some(48000.0));
Ok(())
}
#[test]
fn test_parse_namb_v2_gate_major() -> Result<()> {
let header_size = std::mem::size_of::<NambHeader>();
let w = [0.0f32; 4];
let mut data = vec![0u8; header_size + w.len() * 4];
let header = unsafe { &mut *data.as_mut_ptr().cast::<NambHeader>() };
header.magic = 0x4E414D42;
header.version = 2;
header.layout_type = 1; header.flags = FLAG_HAS_CRC32;
header.weights_offset = header_size as u32;
header.sample_rate = 48000.0;
header.input_level_dbu = 12.0;
header.output_level_dbu = -6.0;
for (i, &f) in w.iter().enumerate() {
let offset = header_size + i * 4;
data[offset..offset + 4].copy_from_slice(&f.to_le_bytes());
}
let crc = {
let mut crc_val = crc32_ieee_update(0xFFFFFFFFu32, &data[..24]);
crc_val = crc32_ieee_update(crc_val, &data[28..]);
crc_val ^ 0xFFFFFFFFu32
};
header.crc32 = crc;
let parsed = parse_namb(&data)?;
assert_eq!(parsed.weights_layout, WeightsLayout::GateMajorLstm);
Ok(())
}
#[test]
fn test_v2_missing_crc32_flag_rejected() {
let header_size = std::mem::size_of::<NambHeader>();
let mut data = vec![0u8; header_size];
let header = unsafe { &mut *data.as_mut_ptr().cast::<NambHeader>() };
header.magic = 0x4E414D42;
header.version = 2;
header.layout_type = 1;
header.flags = 0; header.weights_offset = header_size as u32;
header.crc32 = 0xDEADBEEF;
let err = parse_namb(&data).unwrap_err();
let namb_err = err
.downcast_ref::<NambError>()
.expect("Error should be NambError::CrcMissing");
assert!(
matches!(namb_err, NambError::CrcMissing { version: 2 }),
"Expected CrcMissing, got: {:?}",
namb_err
);
}
#[test]
fn test_v2_crc32_valid_passes() -> Result<()> {
let header_size = std::mem::size_of::<NambHeader>();
let mut data = vec![0u8; header_size];
let header = unsafe { &mut *data.as_mut_ptr().cast::<NambHeader>() };
header.magic = 0x4E414D42;
header.version = 2;
header.layout_type = 1;
header.flags = FLAG_HAS_CRC32;
header.weights_offset = header_size as u32;
header.sample_rate = 48000.0;
header.input_level_dbu = 12.0;
header.output_level_dbu = -6.0;
let crc = {
let mut crc_val = crc32_ieee_update(0xFFFFFFFFu32, &data[..24]);
crc_val = crc32_ieee_update(crc_val, &data[28..]);
crc_val ^ 0xFFFFFFFFu32
};
header.crc32 = crc;
let parsed = parse_namb(&data)?;
assert!(parsed.weights.is_empty());
Ok(())
}
#[test]
fn test_v1_crc32_zero_with_weights_rejected() {
let header_size = std::mem::size_of::<NambHeader>();
let mut data = vec![0u8; header_size + 4];
let header = unsafe { &mut *data.as_mut_ptr().cast::<NambHeader>() };
header.magic = 0x4E414D42;
header.version = 1;
header.weights_offset = header_size as u32;
header.sample_rate = 48000.0;
header.input_level_dbu = 12.0;
header.output_level_dbu = -6.0;
header.crc32 = 0;
let w = 0.5f32;
data[header_size..header_size + 4].copy_from_slice(&w.to_le_bytes());
let err = parse_namb(&data).unwrap_err();
let namb_err = err
.downcast_ref::<NambError>()
.expect("Error should be NambError::CrcMismatch");
assert!(
matches!(namb_err, NambError::CrcMismatch { .. }),
"Expected CrcMismatch, got: {:?}",
namb_err
);
}
#[test]
fn test_v1_crc32_zero_empty_weights_rejected() {
let header_size = std::mem::size_of::<NambHeader>();
let mut data = vec![0u8; header_size];
let header = unsafe { &mut *data.as_mut_ptr().cast::<NambHeader>() };
header.magic = 0x4E414D42;
header.version = 1;
header.weights_offset = header_size as u32;
header.sample_rate = 48000.0;
header.input_level_dbu = 12.0;
header.output_level_dbu = -6.0;
header.crc32 = 0;
let err = parse_namb(&data).unwrap_err();
let namb_err = err
.downcast_ref::<NambError>()
.expect("Error should be NambError::CrcMismatch");
assert!(
matches!(namb_err, NambError::CrcMismatch { .. }),
"Expected CrcMismatch, got: {:?}",
namb_err
);
}
#[test]
fn test_reject_magic_bman() {
let header_size = std::mem::size_of::<NambHeader>();
let mut data = vec![0u8; header_size];
let header = unsafe { &mut *data.as_mut_ptr().cast::<NambHeader>() };
header.magic = 0x424D414E; header.version = 1;
header.weights_offset = header_size as u32;
let err = parse_namb(&data).unwrap_err();
let namb_err = err
.downcast_ref::<NambError>()
.expect("Error should be NambError::InvalidMagic");
assert!(
matches!(namb_err, NambError::InvalidMagic(m) if *m == 0x424D414E),
"Expected InvalidMagic(0x424D414E), got: {:?}",
namb_err
);
}
#[test]
fn test_truncated_header() {
for len in 0..80 {
let data = vec![0u8; len];
let err = parse_namb(&data).unwrap_err();
let namb_err = err
.downcast_ref::<NambError>()
.expect("Error should be NambError::Truncated");
assert!(
matches!(namb_err, NambError::Truncated { got: _, need: _ }),
"Expected Truncated, got: {:?}",
namb_err
);
}
}
#[test]
fn test_weight_residue_rejected() {
for residue in 1..=3 {
let header_size = std::mem::size_of::<NambHeader>();
let weights_bytes = residue + 4;
let mut data = vec![0u8; header_size + weights_bytes];
let header = unsafe { &mut *data.as_mut_ptr().cast::<NambHeader>() };
header.magic = 0x4E414D42;
header.version = 1;
header.weights_offset = header_size as u32;
header.version_str[0..5].copy_from_slice(b"1.0.0");
let weight = 0.5f32;
data[header_size..header_size + 4].copy_from_slice(&weight.to_le_bytes());
header.crc32 = crc32_ieee(&data[header_size..]);
let err = parse_namb(&data).unwrap_err();
let namb_err = err
.downcast_ref::<NambError>()
.expect("Error should be NambError::Truncated");
assert!(
matches!(namb_err, NambError::Truncated { got: _, need: _ }),
"Residue of {} byte(s) should be rejected as Truncated, got: {:?}",
residue,
namb_err
);
}
}
#[test]
fn test_crc32_kat_canonical() {
assert_eq!(
crc32_ieee(b"123456789"),
0xCBF43926,
"CRC32 canonical KAT failed"
);
}
#[test]
fn test_crc32_kat_empty() {
assert_eq!(crc32_ieee(b""), 0x00000000, "CRC32 of empty slice != 0");
}
#[test]
fn test_crc32_kat_single_zero() {
assert_eq!(
crc32_ieee(&[0u8; 1]),
0xD202EF8D,
"CRC32 of single 0x00 byte"
);
}
#[test]
fn test_crc32_kat_four_zeros() {
assert_eq!(
crc32_ieee(&[0u8; 4]),
0x2144DF1C,
"CRC32 of four zero bytes"
);
}
#[test]
fn test_crc32_kat_thirtytwo_zeros() {
assert_eq!(crc32_ieee(&[0u8; 32]), 0x190A55AD, "CRC32 of 32 zero bytes");
}
#[test]
fn test_crc32_kat_sequential() {
let data: Vec<u8> = (0u8..32).collect();
assert_eq!(crc32_ieee(&data), 0x91267E8A, "CRC32 of 0x00..0x1F");
}
#[test]
fn test_crc32_kat_all_ff() {
assert_eq!(
crc32_ieee(&[0xFFu8; 32]),
0xFF6CAB0B,
"CRC32 of 32 bytes of 0xFF"
);
}
#[test]
fn test_crc32_kat_alternating() {
let pattern: Vec<u8> = std::iter::repeat_n([0x55u8, 0xAAu8], 16)
.flatten()
.collect();
assert_eq!(crc32_ieee(&pattern), 0x8BA7B8B6, "CRC32 of 0x55/0xAA Ă—16");
}
fn build_namb_v1_no_crc(w_floats: &[f32]) -> Vec<u8> {
let header_size = std::mem::size_of::<NambHeader>();
let mut data = vec![0u8; header_size + w_floats.len() * 4];
let header = unsafe { &mut *data.as_mut_ptr().cast::<NambHeader>() };
header.magic = 0x4E414D42;
header.version = 1;
header.weights_offset = header_size as u32;
header.version_str[0..5].copy_from_slice(b"1.0.0");
header.sample_rate = 48000.0;
header.input_level_dbu = 12.0;
header.output_level_dbu = -6.0;
for (i, &f) in w_floats.iter().enumerate() {
let offset = header_size + i * 4;
data[offset..offset + 4].copy_from_slice(&f.to_le_bytes());
}
header.crc32 = crc32_ieee(&data[header_size..]);
data
}
#[test]
fn test_truncation_header_boundaries() {
let weights = [0.1f32, 0.2f32, 0.3f32];
let full = build_namb_v1_no_crc(&weights);
let header_size = std::mem::size_of::<NambHeader>();
for len in 0..header_size {
let truncated = &full[..len];
let err = parse_namb(truncated).unwrap_err();
let namb_err = err
.downcast_ref::<NambError>()
.expect("Error should be NambError::Truncated");
assert!(
matches!(namb_err, NambError::Truncated { got: _, need: _ }),
"Truncation at {} bytes: expected Truncated, got {:?}",
len,
namb_err
);
}
}
#[test]
fn test_truncation_weight_boundaries() -> Result<()> {
let weights: Vec<f32> = (0..16).map(|i| i as f32).collect();
let full = build_namb_v1_no_crc(&weights);
let header_size = std::mem::size_of::<NambHeader>();
for num_weights in 0..=weights.len() {
let truncate_at = header_size + num_weights * 4;
let truncated = &full[..truncate_at];
let result = parse_namb(truncated);
if num_weights == weights.len() {
let parsed = result?;
assert_eq!(parsed.weights.len(), weights.len());
} else {
let err = result.unwrap_err();
let namb_err = err
.downcast_ref::<NambError>()
.expect("Error should be NambError variant");
assert!(
matches!(
namb_err,
NambError::CrcMismatch { .. } | NambError::Truncated { .. }
),
"Truncation at {} weights (offset {}): expected CrcMismatch or Truncated, got {:?}",
num_weights,
truncate_at,
namb_err
);
}
}
Ok(())
}
#[test]
fn test_truncation_residue_at_every_boundary() {
let weights: Vec<f32> = (0..8).map(|i| i as f32).collect();
let full = build_namb_v1_no_crc(&weights);
let header_size = std::mem::size_of::<NambHeader>();
for num_weights in 0..=weights.len() {
let base = header_size + num_weights * 4;
for residue in 1..=3 {
let truncate_at = base + residue;
if truncate_at > full.len() {
continue;
}
let truncated = &full[..truncate_at];
let err = parse_namb(truncated).unwrap_err();
let namb_err = err
.downcast_ref::<NambError>()
.expect("Error should be NambError variant");
assert!(
matches!(
namb_err,
NambError::CrcMismatch { .. } | NambError::Truncated { .. }
),
"Residue {} after {} weights (offset {}): expected CrcMismatch or Truncated, got {:?}",
residue,
num_weights,
truncate_at,
namb_err
);
}
}
}
#[test]
fn test_truncation_v2_header_boundary() {
let header_size = std::mem::size_of::<NambHeader>();
let mut data = vec![0u8; header_size];
let header = unsafe { &mut *data.as_mut_ptr().cast::<NambHeader>() };
header.magic = 0x4E414D42;
header.version = 2;
header.flags = FLAG_HAS_CRC32;
header.weights_offset = header_size as u32;
header.sample_rate = 48000.0;
header.input_level_dbu = 12.0;
header.output_level_dbu = -6.0;
let crc = {
let mut crc_val = crc32_ieee_update(0xFFFFFFFFu32, &data[..24]);
crc_val = crc32_ieee_update(crc_val, &data[28..]);
crc_val ^ 0xFFFFFFFFu32
};
header.crc32 = crc;
let parsed = parse_namb(&data).unwrap();
assert!(parsed.weights.is_empty());
}
#[test]
fn test_truncation_v2_post_header_residue() {
let header_size = std::mem::size_of::<NambHeader>();
let mut data = vec![0u8; header_size + 1];
let header = unsafe { &mut *data.as_mut_ptr().cast::<NambHeader>() };
header.magic = 0x4E414D42;
header.version = 2;
header.flags = FLAG_HAS_CRC32;
header.weights_offset = header_size as u32;
let crc = {
let mut crc_val = crc32_ieee_update(0xFFFFFFFFu32, &data[..24]);
crc_val = crc32_ieee_update(crc_val, &data[28..]);
crc_val ^ 0xFFFFFFFFu32
};
header.crc32 = crc;
let err = parse_namb(&data).unwrap_err();
let namb_err = err
.downcast_ref::<NambError>()
.expect("Error should be NambError::Truncated");
assert!(
matches!(namb_err, NambError::Truncated { .. }),
"v2 residue after header: expected Truncated, got {:?}",
namb_err
);
}
#[test]
fn test_non_finite_weight_nan_rejected() {
let header_size = std::mem::size_of::<NambHeader>();
let w = [f32::NAN];
let mut data = vec![0u8; header_size + w.len() * 4];
let header = unsafe { &mut *data.as_mut_ptr().cast::<NambHeader>() };
header.magic = 0x4E414D42;
header.version = 1;
header.weights_offset = header_size as u32;
header.sample_rate = 48000.0;
header.input_level_dbu = 12.0;
header.output_level_dbu = -6.0;
header.version_str[0..5].copy_from_slice(b"1.0.0");
data[header_size..header_size + 4].copy_from_slice(&w[0].to_le_bytes());
header.crc32 = crc32_ieee(&data[header_size..]);
let err = parse_namb(&data).unwrap_err();
let namb_err = err
.downcast_ref::<NambError>()
.expect("Error should be NambError::NonFiniteWeight");
assert!(
matches!(namb_err, NambError::NonFiniteWeight { index: 0, .. }),
"Expected NonFiniteWeight at index 0, got: {:?}",
namb_err
);
}
#[test]
fn test_non_finite_weight_inf_rejected() {
let header_size = std::mem::size_of::<NambHeader>();
let w = [0.5f32, f32::INFINITY];
let mut data = vec![0u8; header_size + w.len() * 4];
let header = unsafe { &mut *data.as_mut_ptr().cast::<NambHeader>() };
header.magic = 0x4E414D42;
header.version = 1;
header.weights_offset = header_size as u32;
header.sample_rate = 48000.0;
header.input_level_dbu = 12.0;
header.output_level_dbu = -6.0;
header.version_str[0..5].copy_from_slice(b"1.0.0");
for (i, &f) in w.iter().enumerate() {
let offset = header_size + i * 4;
data[offset..offset + 4].copy_from_slice(&f.to_le_bytes());
}
header.crc32 = crc32_ieee(&data[header_size..]);
let err = parse_namb(&data).unwrap_err();
let namb_err = err
.downcast_ref::<NambError>()
.expect("Error should be NambError::NonFiniteWeight");
assert!(
matches!(namb_err, NambError::NonFiniteWeight { index: 1, .. }),
"Expected NonFiniteWeight at index 1, got: {:?}",
namb_err
);
}
#[test]
fn test_non_finite_weight_neg_inf_rejected() {
let header_size = std::mem::size_of::<NambHeader>();
let w = [1.0f32, 2.0f32, f32::NEG_INFINITY];
let mut data = vec![0u8; header_size + w.len() * 4];
let header = unsafe { &mut *data.as_mut_ptr().cast::<NambHeader>() };
header.magic = 0x4E414D42;
header.version = 1;
header.weights_offset = header_size as u32;
header.sample_rate = 48000.0;
header.input_level_dbu = 12.0;
header.output_level_dbu = -6.0;
header.version_str[0..5].copy_from_slice(b"1.0.0");
for (i, &f) in w.iter().enumerate() {
let offset = header_size + i * 4;
data[offset..offset + 4].copy_from_slice(&f.to_le_bytes());
}
header.crc32 = crc32_ieee(&data[header_size..]);
let err = parse_namb(&data).unwrap_err();
let namb_err = err
.downcast_ref::<NambError>()
.expect("Error should be NambError::NonFiniteWeight");
assert!(
matches!(namb_err, NambError::NonFiniteWeight { index: 2, .. }),
"Expected NonFiniteWeight at index 2, got: {:?}",
namb_err
);
}
#[test]
fn test_invalid_header_sample_rate_nan() {
let header_size = std::mem::size_of::<NambHeader>();
let mut data = vec![0u8; header_size + 4];
let header = unsafe { &mut *data.as_mut_ptr().cast::<NambHeader>() };
header.magic = 0x4E414D42;
header.version = 1;
header.weights_offset = header_size as u32;
header.sample_rate = f32::NAN;
header.input_level_dbu = 12.0;
header.output_level_dbu = -6.0;
header.version_str[0..5].copy_from_slice(b"1.0.0");
let w = 1.0f32;
data[header_size..header_size + 4].copy_from_slice(&w.to_le_bytes());
header.crc32 = crc32_ieee(&data[header_size..]);
let err = parse_namb(&data).unwrap_err();
let namb_err = err
.downcast_ref::<NambError>()
.expect("Error should be NambError::InvalidHeaderField");
assert!(
matches!(
namb_err,
NambError::InvalidHeaderField {
field: "sample_rate",
..
}
),
"Expected InvalidHeaderField(sample_rate), got: {:?}",
namb_err
);
}
#[test]
fn test_invalid_header_sample_rate_negative() {
let header_size = std::mem::size_of::<NambHeader>();
let mut data = vec![0u8; header_size + 4];
let header = unsafe { &mut *data.as_mut_ptr().cast::<NambHeader>() };
header.magic = 0x4E414D42;
header.version = 1;
header.weights_offset = header_size as u32;
header.sample_rate = -44100.0;
header.input_level_dbu = 12.0;
header.output_level_dbu = -6.0;
header.version_str[0..5].copy_from_slice(b"1.0.0");
let w = 1.0f32;
data[header_size..header_size + 4].copy_from_slice(&w.to_le_bytes());
header.crc32 = crc32_ieee(&data[header_size..]);
let err = parse_namb(&data).unwrap_err();
let namb_err = err
.downcast_ref::<NambError>()
.expect("Error should be NambError::InvalidHeaderField");
assert!(
matches!(
namb_err,
NambError::InvalidHeaderField {
field: "sample_rate",
..
}
),
"Expected InvalidHeaderField(sample_rate), got: {:?}",
namb_err
);
}
#[test]
fn test_invalid_header_sample_rate_zero() {
let header_size = std::mem::size_of::<NambHeader>();
let mut data = vec![0u8; header_size + 4];
let header = unsafe { &mut *data.as_mut_ptr().cast::<NambHeader>() };
header.magic = 0x4E414D42;
header.version = 1;
header.weights_offset = header_size as u32;
header.sample_rate = 0.0;
header.input_level_dbu = 12.0;
header.output_level_dbu = -6.0;
header.version_str[0..5].copy_from_slice(b"1.0.0");
let w = 1.0f32;
data[header_size..header_size + 4].copy_from_slice(&w.to_le_bytes());
header.crc32 = crc32_ieee(&data[header_size..]);
let err = parse_namb(&data).unwrap_err();
let namb_err = err
.downcast_ref::<NambError>()
.expect("Error should be NambError::InvalidHeaderField");
assert!(
matches!(
namb_err,
NambError::InvalidHeaderField {
field: "sample_rate",
..
}
),
"Expected InvalidHeaderField(sample_rate), got: {:?}",
namb_err
);
}
#[test]
fn test_invalid_header_input_level_inf() {
let header_size = std::mem::size_of::<NambHeader>();
let mut data = vec![0u8; header_size + 4];
let header = unsafe { &mut *data.as_mut_ptr().cast::<NambHeader>() };
header.magic = 0x4E414D42;
header.version = 1;
header.weights_offset = header_size as u32;
header.sample_rate = 48000.0;
header.input_level_dbu = f32::INFINITY;
header.output_level_dbu = -6.0;
header.version_str[0..5].copy_from_slice(b"1.0.0");
let w = 1.0f32;
data[header_size..header_size + 4].copy_from_slice(&w.to_le_bytes());
header.crc32 = crc32_ieee(&data[header_size..]);
let err = parse_namb(&data).unwrap_err();
let namb_err = err
.downcast_ref::<NambError>()
.expect("Error should be NambError::InvalidHeaderField");
assert!(
matches!(
namb_err,
NambError::InvalidHeaderField {
field: "input_level_dbu",
..
}
),
"Expected InvalidHeaderField(input_level_dbu), got: {:?}",
namb_err
);
}
#[test]
fn test_invalid_header_output_level_neg_inf() {
let header_size = std::mem::size_of::<NambHeader>();
let mut data = vec![0u8; header_size + 4];
let header = unsafe { &mut *data.as_mut_ptr().cast::<NambHeader>() };
header.magic = 0x4E414D42;
header.version = 1;
header.weights_offset = header_size as u32;
header.sample_rate = 48000.0;
header.input_level_dbu = 12.0;
header.output_level_dbu = f32::NEG_INFINITY;
header.version_str[0..5].copy_from_slice(b"1.0.0");
let w = 1.0f32;
data[header_size..header_size + 4].copy_from_slice(&w.to_le_bytes());
header.crc32 = crc32_ieee(&data[header_size..]);
let err = parse_namb(&data).unwrap_err();
let namb_err = err
.downcast_ref::<NambError>()
.expect("Error should be NambError::InvalidHeaderField");
assert!(
matches!(
namb_err,
NambError::InvalidHeaderField {
field: "output_level_dbu",
..
}
),
"Expected InvalidHeaderField(output_level_dbu), got: {:?}",
namb_err
);
}
#[test]
fn test_v2_metadata_header_corruption_rejected() -> Result<()> {
let header_size = std::mem::size_of::<NambHeader>();
let w = [0.1f32, 0.2f32];
let mut data = vec![0u8; header_size + w.len() * 4];
let header = unsafe { &mut *data.as_mut_ptr().cast::<NambHeader>() };
header.magic = 0x4E414D42;
header.version = 2;
header.flags = FLAG_HAS_CRC32;
header.weights_offset = header_size as u32;
header.sample_rate = 48000.0;
header.input_level_dbu = 12.0;
header.output_level_dbu = -6.0;
for (i, &f) in w.iter().enumerate() {
let offset = header_size + i * 4;
data[offset..offset + 4].copy_from_slice(&f.to_le_bytes());
}
let crc = {
let mut crc_val = crc32_ieee_update(0xFFFFFFFFu32, &data[..24]);
crc_val = crc32_ieee_update(crc_val, &data[28..]);
crc_val ^ 0xFFFFFFFFu32
};
header.crc32 = crc;
let parsed = parse_namb(&data)?;
assert_eq!(parsed.weights.len(), 2);
let mut corrupted_header = data.clone();
corrupted_header[64] ^= 0xFF; let err = parse_namb(&corrupted_header).unwrap_err();
let namb_err = err.downcast_ref::<NambError>().unwrap();
assert!(
matches!(namb_err, NambError::CrcMismatch { .. }),
"Expected CrcMismatch for corrupted sample_rate in v2, got {:?}",
namb_err
);
let mut corrupted_weights = data.clone();
corrupted_weights[header_size] ^= 0xFF; let err2 = parse_namb(&corrupted_weights).unwrap_err();
let namb_err2 = err2.downcast_ref::<NambError>().unwrap();
assert!(
matches!(namb_err2, NambError::CrcMismatch { .. }),
"Expected CrcMismatch for corrupted weights in v2, got {:?}",
namb_err2
);
Ok(())
}
#[test]
fn test_v1_vs_v2_metadata_header_corruption() -> Result<()> {
{
let header_size = std::mem::size_of::<NambHeader>();
let w = [0.1f32, 0.2f32];
let mut data = vec![0u8; header_size + w.len() * 4];
let header = unsafe { &mut *data.as_mut_ptr().cast::<NambHeader>() };
header.magic = 0x4E414D42;
header.version = 1;
header.weights_offset = header_size as u32;
header.sample_rate = 48000.0;
header.input_level_dbu = 12.0;
header.output_level_dbu = -6.0;
header.version_str[0..5].copy_from_slice(b"1.0.0");
for (i, &f) in w.iter().enumerate() {
let offset = header_size + i * 4;
data[offset..offset + 4].copy_from_slice(&f.to_le_bytes());
}
header.crc32 = crc32_ieee(&data[header_size..]);
let mut corrupted = data.clone();
let corrupted_header = unsafe { &mut *corrupted.as_mut_ptr().cast::<NambHeader>() };
corrupted_header.input_level_dbu = 13.0;
let parsed = parse_namb(&corrupted)?;
assert_eq!(parsed.metadata.unwrap().input_level_dbu, Some(13.0));
}
{
let header_size = std::mem::size_of::<NambHeader>();
let w = [0.1f32, 0.2f32];
let mut data = vec![0u8; header_size + w.len() * 4];
let header = unsafe { &mut *data.as_mut_ptr().cast::<NambHeader>() };
header.magic = 0x4E414D42;
header.version = 2;
header.flags = FLAG_HAS_CRC32;
header.weights_offset = header_size as u32;
header.sample_rate = 48000.0;
header.input_level_dbu = 12.0;
header.output_level_dbu = -6.0;
for (i, &f) in w.iter().enumerate() {
let offset = header_size + i * 4;
data[offset..offset + 4].copy_from_slice(&f.to_le_bytes());
}
let crc = {
let mut crc_val = crc32_ieee_update(0xFFFFFFFFu32, &data[..24]);
crc_val = crc32_ieee_update(crc_val, &data[28..]);
crc_val ^ 0xFFFFFFFFu32
};
header.crc32 = crc;
let mut corrupted = data.clone();
let corrupted_header = unsafe { &mut *corrupted.as_mut_ptr().cast::<NambHeader>() };
corrupted_header.input_level_dbu = 13.0;
let err = parse_namb(&corrupted).unwrap_err();
let namb_err = err.downcast_ref::<NambError>().unwrap();
assert!(matches!(namb_err, NambError::CrcMismatch { .. }));
}
Ok(())
}
}