#![cfg(feature = "hdf5")]
#![allow(dead_code)]
use nir_rs::nodes::{Input, Output};
use nir_rs::{NirError, NirGraph, NirNode};
use tempfile::TempDir;
pub fn assert_err<T>(
result: Result<T, NirError>,
expected_variant: fn(String) -> NirError,
needles: &[&str],
) {
let Err(err) = result else {
panic!("expected an error, got Ok");
};
let dummy = expected_variant(String::new());
assert_eq!(
std::mem::discriminant(&err),
std::mem::discriminant(&dummy),
"expected {dummy:?} variant, got {err:?}"
);
let msg = err.to_string();
for needle in needles {
assert!(
!needle.is_empty(),
"empty needle is vacuous; pass &[] to skip the message check"
);
assert!(
msg.contains(needle),
"expected message to contain {needle:?}, got {msg:?}"
);
}
}
pub fn input(shape: Vec<usize>) -> NirNode {
NirNode::Input(Input {
shape,
metadata: Default::default(),
})
}
pub fn write_then(
dir: &TempDir,
name: &str,
mutate: impl FnOnce(&hdf5::File),
) -> std::path::PathBuf {
let mut graph = NirGraph::new();
graph.insert_node("input", input(vec![1])).unwrap();
graph
.insert_node(
"output",
NirNode::Output(Output {
shape: vec![1],
metadata: Default::default(),
}),
)
.unwrap();
graph.add_edge("input", "output");
let path = dir.path().join(name);
nir_rs::io::write(&path, &graph).unwrap();
let file = hdf5::File::open_rw(&path).unwrap();
mutate(&file);
drop(file);
path
}