mod common;
use std::collections::BTreeSet;
use coremlit::{
ComputeUnits, DataType, Model,
audio::ced::{CedModel, Classifier, ClassifierOptions, NUM_CLASSES},
};
const MINI_SHA256: &[(&str, &str)] = &[
(
"analytics/coremldata.bin",
"4b2fd2952ef76ebf73a70a7b91d8f6694941967e614aa062c08dd35e311cd82f",
),
(
"coremldata.bin",
"c5f1c8b0b078d89880a15d5271e77e5c9ca478704e0b14cb6e3601dbcb8b57ab",
),
(
"metadata.json",
"91bbae280e8777bc464064c885ccf6d81fb50eed2bc696738565ad3371cbfb42",
),
(
"model.mil",
"98605d30a89dc6ab352cc9679c6c6c14a7def969c5f7738baf67d9eae09ce153",
),
(
"weights/weight.bin",
"daf9f1fa64c8eb2a00fc5325ecfc6c45670b4e15c69a0277a16cacf0dba44c6c",
),
];
const SMALL_SHA256: &[(&str, &str)] = &[
(
"analytics/coremldata.bin",
"c4d275dca741a0b7358202d62e4fe3f3ebdf60fa35d68ffdd412f709384dc337",
),
(
"coremldata.bin",
"7ca0aac05c69e3315f88590b8870e783130987a4f7ad95669efdee0e11cacc48",
),
(
"metadata.json",
"61c2dc29a14c0b15989c8f9e5841ca2aec0ab7c95bfd5e844cd7a88aa87989f0",
),
(
"model.mil",
"e7974466d1a373964bbf99414623a7900f62a944c7baa00aa8950ad9255e73c4",
),
(
"weights/weight.bin",
"af2c05ff6f7de533aa649f0d765499fd6a7bc43b1e3237e25f6a7a8cf35c44a0",
),
];
const BASE_SHA256: &[(&str, &str)] = &[
(
"analytics/coremldata.bin",
"260b83886ef274a7f785281019722d8f4c7a201475144bc8b4edad704016945e",
),
(
"coremldata.bin",
"5b26521f2b699040b12fe239d016cd5783680b9a1b6e33f49afba3cbdad50ea4",
),
(
"metadata.json",
"aed58c97c54c58b7968c8fe35e7d1ffb0e291f9dcca22775f3d11db9fb9d6cd4",
),
(
"model.mil",
"ffb76dd2cac55c0b571c97565b9874c1804854c79f5b9969a4f62276ceebef2f",
),
(
"weights/weight.bin",
"bc80c08ef907a07ab153083a31bef836154e83bf88c919747157d36b86e6b5f2",
),
];
fn artifact_sha256(model: CedModel) -> Vec<(String, String)> {
let table: &[(&str, &str)] = match model {
CedModel::Tiny => {
return common::models_lock_manifest::bundle_manifest(
&common::workspace_root(),
common::VENDOR_DIR,
common::ARTIFACT_LOCK_REVISION,
common::TINY_BUNDLE_PATH,
);
}
CedModel::Mini => MINI_SHA256,
CedModel::Small => SMALL_SHA256,
CedModel::Base => BASE_SHA256,
};
table
.iter()
.map(|(rel, sha)| ((*rel).to_string(), (*sha).to_string()))
.collect()
}
fn io_contract(model: CedModel) {
let m = Model::load(common::model_path(model), ComputeUnits::CpuOnly).unwrap();
let d = m.description();
let mel = d.input("mel").expect("mel input");
assert_eq!(
mel.shape(),
&[1, 64, 1001],
"believed [1, n_mels, T] layout"
);
assert_eq!(mel.data_type(), Some(DataType::F32));
let logits = d.output("logits").expect("logits output");
assert_eq!(logits.shape(), &[1, NUM_CLASSES]);
assert_eq!(logits.data_type(), Some(DataType::F32));
let input_names: BTreeSet<&str> = d.inputs().iter().map(|f| f.name()).collect();
assert_eq!(input_names, BTreeSet::from(["mel"]), "exactly one input");
let output_names: BTreeSet<&str> = d.outputs().iter().map(|f| f.name()).collect();
assert_eq!(
output_names,
BTreeSet::from(["logits"]),
"exactly one output"
);
Classifier::load(
common::model_path(model),
ClassifierOptions::new().with_compute(ComputeUnits::CpuOnly),
)
.expect("the staged artifact must satisfy this door's load contract");
}
fn artifact_manifest(model: CedModel) {
assert!(
!artifact_sha256(model).is_empty(),
"Wave B must pin {model}'s artifact manifest before this gate can pass"
);
common::assert_exact_sha_manifest(&common::model_path(model), &artifact_sha256(model));
}
macro_rules! per_model_gates {
($($m:ident => $v:expr),+ $(,)?) => {$(
mod $m {
use super::CedModel;
#[test]
#[ignore = "requires staged CED model (CED_TEST_MODELS) — Wave B"]
fn io_contract_matches_the_believed_spec() {
super::io_contract($v);
}
#[test]
#[ignore = "requires staged CED model (CED_TEST_MODELS) — Wave B"]
fn artifact_bytes_match_pinned_sha256() {
super::artifact_manifest($v);
}
}
)+};
}
per_model_gates!(
tiny => CedModel::Tiny,
mini => CedModel::Mini,
small => CedModel::Small,
base => CedModel::Base,
);
#[test]
fn collect_files_rel_skips_sidecars_but_surfaces_real_extras() {
let dir = tempfile::tempdir().unwrap();
std::fs::write(dir.path().join("model.mil"), b"mil").unwrap();
std::fs::write(dir.path().join(".DS_Store"), b"junk").unwrap();
std::fs::write(dir.path().join("._model.mil"), b"appledouble").unwrap();
std::fs::create_dir(dir.path().join("weights")).unwrap();
std::fs::write(dir.path().join("weights/weight.bin"), b"w").unwrap();
let mut found = Vec::new();
common::collect_files_rel(dir.path(), "", &mut found);
found.sort();
assert_eq!(found, vec!["model.mil", "weights/weight.bin"]);
}
#[test]
fn model_path_composes_per_size_under_the_family_root() {
for m in CedModel::ALL {
assert_eq!(
common::model_path(m),
common::models_dir()
.join(m.dir_name())
.join(m.mlmodelc_name()),
);
}
}
#[test]
#[should_panic(expected = "does not match")]
fn golden_corpus_rejects_a_cross_size_oracle() {
let json = br#"{
"oracle": {"repo": "mispeech/ced-tiny", "revision": "r", "file": "model.safetensors", "sha256": "00"},
"clips": []
}"#;
common::parse_golden_corpus(json, CedModel::Small);
}
#[test]
fn golden_corpus_accepts_a_matching_oracle() {
let json = br#"{
"oracle": {"repo": "mispeech/ced-small", "revision": "r", "file": "model.safetensors", "sha256": "00"},
"clips": [{"id": "c0", "file": "clips/c0.wav", "n_samples": 160000, "logits": [0.0, 1.0]}]
}"#;
let corpus = common::parse_golden_corpus(json, CedModel::Small);
assert_eq!(corpus.oracle.repo, "mispeech/ced-small");
assert_eq!(corpus.clips.len(), 1);
}