#[path = "../../support/models_lock_manifest.rs"]
#[allow(dead_code)]
pub mod models_lock_manifest;
#[path = "../../support/workspace_root.rs"]
#[allow(dead_code)]
mod workspace_root;
#[allow(unused_imports)]
pub use workspace_root::{checkout_parent, models_root, workspace_root};
use std::path::{Path, PathBuf};
#[allow(dead_code)]
pub const EMBEDKIT_REVISION: &str = "a61241cb";
#[allow(dead_code)]
pub const VENDOR_DIR: &str = "embedkit-granite";
#[allow(dead_code)]
pub const EMBEDKIT_LOCK_REVISION: &str = "a61241cb18689de1a8b83c315875807030e2878f";
#[allow(dead_code)]
pub const BUNDLE_PATH: &str = "granite-97m-multilingual-r2/granite_97m_512.mlmodelc";
#[allow(dead_code)]
pub const TOKENIZER_PATH: &str = "granite-97m-multilingual-r2/tokenizer.json";
#[allow(dead_code)]
pub fn artifact_sha256() -> Vec<(String, String)> {
models_lock_manifest::bundle_manifest(
&workspace_root::workspace_root(),
VENDOR_DIR,
EMBEDKIT_LOCK_REVISION,
BUNDLE_PATH,
)
}
#[allow(dead_code)]
pub fn models_dir() -> PathBuf {
std::env::var_os("EMBEDKIT_TEST_MODELS").map_or_else(
|| workspace_root::models_root().join("embedkit-granite"),
PathBuf::from,
)
}
#[allow(dead_code)]
pub fn model_root() -> PathBuf {
models_dir().join("granite-97m-multilingual-r2")
}
#[allow(dead_code)]
pub fn model_path() -> PathBuf {
model_root().join("granite_97m_512.mlmodelc")
}
#[allow(dead_code)]
pub fn fixture_path(relative: &str) -> PathBuf {
PathBuf::from(env!("CARGO_MANIFEST_DIR"))
.join("tests")
.join("granite")
.join("fixtures")
.join(relative)
}
#[derive(Debug, serde::Deserialize)]
#[allow(dead_code)]
pub struct GoldenEntry {
pub id: String,
pub text: String,
pub token_ids: Vec<u32>,
pub n_tokens: usize,
pub embedding: Vec<f32>,
}
#[derive(Debug, serde::Deserialize)]
struct Corpus {
entries: Vec<GoldenEntry>,
}
#[allow(dead_code)]
pub fn golden_corpus() -> Vec<GoldenEntry> {
let path = fixture_path("goldens/corpus.json");
let bytes = std::fs::read(&path).unwrap_or_else(|e| panic!("read {}: {e}", path.display()));
let corpus: Corpus =
serde_json::from_slice(&bytes).unwrap_or_else(|e| panic!("parse {}: {e}", path.display()));
assert_eq!(
corpus.entries.len(),
16,
"the committed granite golden corpus must have 16 entries"
);
corpus.entries
}
#[derive(Debug, serde::Deserialize)]
#[allow(dead_code)]
pub struct CrosscheckEntry {
pub id: String,
pub cosine_canonical_vs_driver: f64,
pub max_abs_component_delta: f64,
}
#[derive(Debug, serde::Deserialize)]
#[allow(dead_code)]
pub struct DriverCrosscheck {
pub worst_cosine_canonical_vs_driver: f64,
pub min_max_abs_component_delta: f64,
pub stop_threshold_divergence: f64,
pub verdict: String,
pub corpus_sha256: String,
pub per_entry: Vec<CrosscheckEntry>,
}
#[allow(dead_code)]
pub fn driver_crosscheck() -> DriverCrosscheck {
let path = fixture_path("goldens/driver_crosscheck.json");
let bytes = std::fs::read(&path).unwrap_or_else(|e| panic!("read {}: {e}", path.display()));
serde_json::from_slice(&bytes).unwrap_or_else(|e| panic!("parse {}: {e}", path.display()))
}
#[allow(dead_code)]
pub fn corpus_sha256() -> String {
let path = fixture_path("goldens/corpus.json");
let bytes = std::fs::read(&path).unwrap_or_else(|e| panic!("read {}: {e}", path.display()));
sha256_hex(&bytes)
}
#[allow(dead_code)]
pub fn sha256_hex(bytes: &[u8]) -> String {
use sha2::{Digest, Sha256};
Sha256::digest(bytes)
.iter()
.map(|b| format!("{b:02x}"))
.collect()
}
#[allow(dead_code)]
pub fn cosine_checked(a: &[f32], b: &[f32]) -> f32 {
assert_eq!(a.len(), b.len(), "cosine operands differ in length");
assert!(!a.is_empty(), "cosine of empty vectors is undefined");
let mut dot = 0.0f64;
let mut na = 0.0f64;
let mut nb = 0.0f64;
for (i, (&x, &y)) in a.iter().zip(b.iter()).enumerate() {
assert!(x.is_finite(), "left operand non-finite at {i}");
assert!(y.is_finite(), "right operand non-finite at {i}");
dot += (x as f64) * (y as f64);
na += (x as f64) * (x as f64);
nb += (y as f64) * (y as f64);
}
assert!(na > 0.0 && nb > 0.0, "cosine of a zero vector is undefined");
(dot / (na.sqrt() * nb.sqrt())) as f32
}
#[allow(dead_code)]
pub fn collect_files_rel(root: &Path, dir: &Path, out: &mut std::collections::BTreeSet<String>) {
for entry in std::fs::read_dir(dir).unwrap_or_else(|e| panic!("read_dir {}: {e}", dir.display()))
{
let entry = entry.expect("read dir entry");
let name = entry.file_name();
let name = name.to_string_lossy();
if name.starts_with("._") || name == ".DS_Store" {
continue;
}
let path = entry.path();
if entry.file_type().expect("file type").is_dir() {
collect_files_rel(root, &path, out);
} else {
let rel = path
.strip_prefix(root)
.expect("walked path is under root")
.to_str()
.expect("utf-8 path")
.replace('\\', "/");
out.insert(rel);
}
}
}
#[path = "../../support/model_gate_report.rs"]
mod model_gate_report;
#[test]
fn model_gate_report() {
model_gate_report::report(&[("EMBEDKIT_TEST_MODELS", models_dir())]);
}