use super::ReadOptions;
use super::wire::{self, KEY_EDGES, KEY_METADATA, KEY_NODE, KEY_NODES, KEY_TYPE, KEY_VERSION};
use crate::error::{NirError, ReadLimitResource, Result};
use crate::graph::NirGraph;
use crate::nodes::{
Affine, AvgPool2d, Conv1d, Conv2d, CubaLi, CubaLif, Delay, Flatten, I, If, Input, Li, Lif,
Linear, NirNode, Output, Padding, Scale, SumPool2d, Threshold,
};
use crate::types::{MetadataMap, MetadataValue, Tensor, TensorData};
use hdf5::plist::dataset_create::Layout;
use hdf5::types::{
FixedAscii, FixedUnicode, FloatSize, IntSize, TypeDescriptor as Td, VarLenAscii, VarLenUnicode,
};
use hdf5::{Dataset, File, Group, LocationToken};
use std::cell::Cell;
use std::path::Path;
pub(super) const MAX_NESTED_GRAPHS: usize = 1024;
struct ReadBudget {
limit: Option<usize>,
used: Cell<usize>,
nodes: CountLedger,
edges: CountLedger,
graphs: CountLedger,
}
struct CountLedger {
limit: Option<usize>,
used: Cell<usize>,
}
impl CountLedger {
fn new(limit: Option<usize>) -> Self {
Self {
limit,
used: Cell::new(0),
}
}
fn charge(
&self,
resource: ReadLimitResource,
context: &str,
requested: Option<usize>,
) -> Result<()> {
let Some(limit) = self.limit else {
return Ok(());
};
let used = self.used.get();
let Some(requested) = requested else {
return Err(NirError::ReadCountLimitExceeded {
resource,
context: context.to_owned(),
limit,
used,
requested: usize::MAX,
});
};
let Some(next) = used.checked_add(requested) else {
return Err(NirError::ReadCountLimitExceeded {
resource,
context: context.to_owned(),
limit,
used,
requested,
});
};
if next > limit {
return Err(NirError::ReadCountLimitExceeded {
resource,
context: context.to_owned(),
limit,
used,
requested,
});
}
self.used.set(next);
Ok(())
}
}
impl ReadBudget {
fn new(opts: &ReadOptions) -> Self {
Self {
limit: opts.max_bytes,
used: Cell::new(0),
nodes: CountLedger::new(opts.max_nodes),
edges: CountLedger::new(opts.max_edges),
graphs: CountLedger::new(opts.max_nested_graphs),
}
}
fn charge_nodes(&self, context: &str, requested: Option<usize>) -> Result<()> {
self.nodes
.charge(ReadLimitResource::Nodes, context, requested)
}
fn charge_edges(&self, context: &str, requested: Option<usize>) -> Result<()> {
self.edges
.charge(ReadLimitResource::Edges, context, requested)
}
fn charge_graph(&self, context: &str) -> Result<()> {
self.graphs
.charge(ReadLimitResource::NestedGraphs, context, Some(1))
}
fn charge(&self, context: &str, requested: Option<usize>) -> Result<()> {
let Some(limit) = self.limit else {
return Ok(());
};
let used = self.used.get();
let Some(requested) = requested else {
return Err(NirError::ReadLimitExceeded {
context: context.to_owned(),
limit,
used,
requested: usize::MAX,
});
};
let Some(next) = used.checked_add(requested) else {
return Err(NirError::ReadLimitExceeded {
context: context.to_owned(),
limit,
used,
requested,
});
};
if next > limit {
return Err(NirError::ReadLimitExceeded {
context: context.to_owned(),
limit,
used,
requested,
});
}
self.used.set(next);
Ok(())
}
fn would_fit(&self, context: &str, requested: Option<usize>) -> Result<()> {
let Some(limit) = self.limit else {
return Ok(());
};
let used = self.used.get();
let Some(requested) = requested else {
return Err(NirError::ReadLimitExceeded {
context: context.to_owned(),
limit,
used,
requested: usize::MAX,
});
};
let Some(next) = used.checked_add(requested) else {
return Err(NirError::ReadLimitExceeded {
context: context.to_owned(),
limit,
used,
requested,
});
};
if next > limit {
return Err(NirError::ReadLimitExceeded {
context: context.to_owned(),
limit,
used,
requested,
});
}
Ok(())
}
}
const FIXED_STRING_CAPS: [usize; 3] = [64, 256, 4096];
const _: () = assert!(
FIXED_STRING_CAPS[0] == 64 && FIXED_STRING_CAPS[1] == 256 && FIXED_STRING_CAPS[2] == 4096,
"FIXED_STRING_CAPS and the read_fixed! rungs in read_strings_unchecked must match"
);
pub(super) fn read(path: &Path, opts: &ReadOptions) -> Result<NirGraph> {
let budget = ReadBudget::new(opts);
let file = open(path)?;
validate_group_links(&file, "/")?;
let version = load_version(&file, &budget)?;
super::version::enforce_version_policy(version.as_deref(), &opts.version_policy)?;
let root = file.group(KEY_NODE).map_err(|_| {
NirError::MissingField(format!(
"/{KEY_NODE} (not a NIR graph file: {})",
path.display()
))
})?;
validate_group_links(&root, &format!("/{KEY_NODE}"))?;
let root_type = read_string_scalar(
&root
.dataset(KEY_TYPE)
.map_err(|_| NirError::MissingField(format!("/{KEY_NODE}/{KEY_TYPE}")))?,
KEY_NODE,
&budget,
)?;
if root_type != "NIRGraph" {
return Err(NirError::InvalidGraph(format!(
"/{KEY_NODE} must be a NIRGraph, found {root_type:?}"
)));
}
let mut graph = read_graph_body(&root, &format!("/{KEY_NODE}"), &mut Vec::new(), &budget)?;
graph.metadata = read_metadata(&root, &budget)?;
graph.version = version;
Ok(graph)
}
fn version_dataset(file: &File) -> Result<Option<Dataset>> {
if !file.link_exists(KEY_VERSION) {
return Ok(None);
}
file.dataset(KEY_VERSION).map(Some).map_err(|e| {
NirError::Io(format!(
"/{KEY_VERSION}: expected a dataset, found another link kind: {e}"
))
})
}
fn load_version(file: &File, budget: &ReadBudget) -> Result<Option<String>> {
match version_dataset(file)? {
Some(ds) => Ok(Some(read_string_scalar(&ds, KEY_VERSION, budget)?)),
None => Ok(None),
}
}
pub(super) fn read_version(path: &Path, opts: &ReadOptions) -> Result<String> {
let budget = ReadBudget::new(opts);
let file = open(path)?;
validate_group_links(&file, "/")?;
let version = load_version(&file, &budget)?;
super::version::enforce_version_policy(version.as_deref(), &opts.version_policy)?;
version.ok_or_else(|| NirError::MissingField(format!("/{KEY_VERSION}")))
}
fn open(path: &Path) -> Result<File> {
File::open(path).map_err(|e| NirError::Io(format!("cannot open {}: {e}", path.display())))
}
fn read_graph_body(
group: &Group,
context: &str,
visited: &mut Vec<LocationToken>,
budget: &ReadBudget,
) -> Result<NirGraph> {
let token = group.loc_info()?.token;
if visited.contains(&token) {
return Err(NirError::InvalidGraph(format!(
"{context}: nested NIRGraph group is already being decoded (hard-link cycle or alias)"
)));
}
budget.charge_graph(context)?;
if visited.len() >= MAX_NESTED_GRAPHS {
return Err(NirError::InvalidGraph(format!(
"{context}: more than {MAX_NESTED_GRAPHS} nested NIRGraph groups"
)));
}
visited.push(token);
read_graph_body_inner(group, context, visited, budget)
}
fn read_graph_body_inner(
group: &Group,
context: &str,
visited: &mut Vec<LocationToken>,
budget: &ReadBudget,
) -> Result<NirGraph> {
let mut graph = NirGraph::new();
validate_group_links(group, context)?;
let nodes = group
.group(KEY_NODES)
.map_err(|_| NirError::MissingField(format!("{context}/{KEY_NODES}")))?;
let nodes_ctx = format!("{context}/{KEY_NODES}");
budget.charge_nodes(&nodes_ctx, group_nlinks(&nodes, &nodes_ctx)?)?;
validate_group_links(&nodes, &nodes_ctx)?;
let mut names = nodes.member_names()?;
names.sort();
if let Ok(n) = usize::try_from(nodes.len()) {
graph.nodes.reserve(n);
}
for name in names {
let node_group = nodes
.group(&name)
.map_err(|e| NirError::Io(format!("node {name:?} is not a group: {e}")))?;
validate_group_links(&node_group, &name)?;
let node = read_node(&node_group, &name, context, visited, budget)?;
graph.insert_node(name, node)?;
}
let edges = group
.dataset(KEY_EDGES)
.map_err(|_| NirError::MissingField(format!("{context}/{KEY_EDGES}")))?;
graph.edges = read_edges(&edges, &format!("{context}/{KEY_EDGES}"), budget)?;
Ok(graph)
}
fn read_edges(ds: &Dataset, context: &str, budget: &ReadBudget) -> Result<Vec<(String, String)>> {
validate_dataset_security(ds, context)?;
if ds.size() == 0 {
return Ok(Vec::new());
}
let shape = ds.shape();
if shape.len() != 2 || shape[1] != 2 {
return Err(NirError::InvalidGraph(format!(
"{context} must have shape (E, 2), found {shape:?}"
)));
}
let edge_count = shape[0];
budget.charge_edges(context, Some(edge_count))?;
let flat = read_strings_unchecked(ds, context, budget)?;
budget.charge(
context,
edge_count.checked_mul(std::mem::size_of::<(String, String)>()),
)?;
let mut strings = flat.into_iter();
Ok(std::iter::from_fn(|| Some((strings.next()?, strings.next()?))).collect())
}
fn read_node(
group: &Group,
name: &str,
parent_path: &str,
visited: &mut Vec<LocationToken>,
budget: &ReadBudget,
) -> Result<NirNode> {
let type_ds = group
.dataset(KEY_TYPE)
.map_err(|_| NirError::MissingField(format!("{name}.{KEY_TYPE}")))?;
let ty = read_string_scalar(&type_ds, name, budget)?;
let node = if ty == "NIRGraph" {
read_graph_node(group, name, parent_path, visited, budget)?
} else {
read_leaf_node(group, name, &ty, budget)?
};
debug_assert!(
wire::is_wire_type(node.type_name()),
"decoded a node whose type is not in WIRE_TYPES"
);
Ok(node)
}
fn read_graph_node(
group: &Group,
name: &str,
parent_path: &str,
visited: &mut Vec<LocationToken>,
budget: &ReadBudget,
) -> Result<NirNode> {
let nested_path = format!("{parent_path}/{KEY_NODES}/{name}");
let mut sub = read_graph_body(group, &nested_path, visited, budget)?;
sub.metadata = read_metadata(group, budget)?;
Ok(NirNode::Graph(Box::new(sub)))
}
fn read_leaf_node(group: &Group, name: &str, ty: &str, budget: &ReadBudget) -> Result<NirNode> {
let metadata = read_metadata(group, budget)?;
let r = NodeReader {
group,
name,
budget,
};
match ty {
"Input" | "Output" | "Affine" | "Linear" | "Scale" => read_map_leaf(&r, ty, metadata),
"Conv1d" | "Conv2d" => read_conv_leaf(&r, ty, metadata),
"CubaLI" | "CubaLIF" | "I" | "IF" | "LI" | "LIF" => read_neuron_leaf(&r, ty, metadata),
"SumPool2d" | "AvgPool2d" | "Delay" | "Flatten" | "Threshold" => {
read_window_leaf(&r, ty, metadata)
}
other => Err(NirError::UnknownNodeType(other.to_owned())),
}
}
fn read_map_leaf(r: &NodeReader, ty: &str, metadata: MetadataMap) -> Result<NirNode> {
match ty {
"Input" => Ok(NirNode::Input(read_input(r, metadata)?)),
"Output" => Ok(NirNode::Output(read_output(r, metadata)?)),
"Affine" => Ok(NirNode::Affine(read_affine(r, metadata)?)),
"Linear" => Ok(NirNode::Linear(read_linear(r, metadata)?)),
"Scale" => Ok(NirNode::Scale(read_scale(r, metadata)?)),
other => Err(NirError::UnknownNodeType(other.to_owned())),
}
}
fn read_conv_leaf(r: &NodeReader, ty: &str, metadata: MetadataMap) -> Result<NirNode> {
match ty {
"Conv1d" => Ok(NirNode::Conv1d(read_conv1d(r, metadata)?)),
"Conv2d" => Ok(NirNode::Conv2d(read_conv2d(r, metadata)?)),
other => Err(NirError::UnknownNodeType(other.to_owned())),
}
}
fn read_neuron_leaf(r: &NodeReader, ty: &str, metadata: MetadataMap) -> Result<NirNode> {
match ty {
"CubaLI" => Ok(NirNode::CubaLi(read_cuba_li(r, metadata)?)),
"CubaLIF" => Ok(NirNode::CubaLif(read_cuba_lif(r, metadata)?)),
"I" => Ok(NirNode::I(read_i(r, metadata)?)),
"IF" => Ok(NirNode::If(read_if(r, metadata)?)),
"LI" => Ok(NirNode::Li(read_li(r, metadata)?)),
"LIF" => Ok(NirNode::Lif(read_lif(r, metadata)?)),
other => Err(NirError::UnknownNodeType(other.to_owned())),
}
}
fn read_window_leaf(r: &NodeReader, ty: &str, metadata: MetadataMap) -> Result<NirNode> {
match ty {
"SumPool2d" => Ok(NirNode::SumPool2d(read_sum_pool2d(r, metadata)?)),
"AvgPool2d" => Ok(NirNode::AvgPool2d(read_avg_pool2d(r, metadata)?)),
"Delay" => Ok(NirNode::Delay(read_delay(r, metadata)?)),
"Flatten" => Ok(NirNode::Flatten(read_flatten(r, metadata)?)),
"Threshold" => Ok(NirNode::Threshold(read_threshold(r, metadata)?)),
other => Err(NirError::UnknownNodeType(other.to_owned())),
}
}
fn read_input(r: &NodeReader, metadata: MetadataMap) -> Result<Input> {
Ok(Input {
shape: r.usizes("shape")?,
metadata,
})
}
fn read_output(r: &NodeReader, metadata: MetadataMap) -> Result<Output> {
Ok(Output {
shape: r.usizes("shape")?,
metadata,
})
}
fn read_affine(r: &NodeReader, metadata: MetadataMap) -> Result<Affine> {
Ok(Affine {
weight: r.tensor("weight")?,
bias: r.tensor("bias")?,
metadata,
})
}
fn read_linear(r: &NodeReader, metadata: MetadataMap) -> Result<Linear> {
Ok(Linear {
weight: r.tensor("weight")?,
metadata,
})
}
fn read_scale(r: &NodeReader, metadata: MetadataMap) -> Result<Scale> {
Ok(Scale {
scale: r.tensor("scale")?,
metadata,
})
}
struct ConvGeometry {
weight: Tensor,
stride: Vec<i64>,
padding: Padding,
dilation: Vec<i64>,
groups: i64,
bias: Tensor,
}
impl ConvGeometry {
fn read(r: &NodeReader) -> Result<Self> {
Ok(Self {
weight: r.tensor("weight")?,
stride: r.ints("stride")?,
padding: r.padding()?,
dilation: r.ints("dilation")?,
groups: r.int_scalar("groups")?,
bias: r.tensor("bias")?,
})
}
fn into_conv1d(self, input_shape: Option<usize>, metadata: MetadataMap) -> Conv1d {
let Self {
weight,
stride,
padding,
dilation,
groups,
bias,
} = self;
Conv1d {
weight,
stride,
padding,
dilation,
groups,
bias,
input_shape,
metadata,
}
}
fn into_conv2d(self, input_shape: Option<Vec<usize>>, metadata: MetadataMap) -> Conv2d {
Conv2d {
weight: self.weight,
stride: self.stride,
padding: self.padding,
dilation: self.dilation,
groups: self.groups,
bias: self.bias,
input_shape,
metadata,
}
}
}
fn read_conv1d(r: &NodeReader, metadata: MetadataMap) -> Result<Conv1d> {
Ok(ConvGeometry::read(r)?.into_conv1d(r.opt_dim("input_shape")?, metadata))
}
fn read_conv2d(r: &NodeReader, metadata: MetadataMap) -> Result<Conv2d> {
Ok(ConvGeometry::read(r)?.into_conv2d(r.opt_usizes("input_shape")?, metadata))
}
fn threshold_and_reset(r: &NodeReader) -> Result<(Tensor, Option<Tensor>)> {
let v_threshold = r.tensor("v_threshold")?;
let v_reset = r.v_reset(&v_threshold)?;
Ok((v_threshold, v_reset))
}
struct LeakDynamics {
tau: Tensor,
r: Tensor,
v_leak: Tensor,
}
impl LeakDynamics {
fn read(r: &NodeReader) -> Result<Self> {
Ok(Self {
tau: r.tensor("tau")?,
r: r.tensor("r")?,
v_leak: r.tensor("v_leak")?,
})
}
fn into_li(self, metadata: MetadataMap) -> Li {
let Self { tau, r, v_leak } = self;
Li {
tau,
r,
v_leak,
metadata,
}
}
fn into_lif(self, v_threshold: Tensor, v_reset: Option<Tensor>, metadata: MetadataMap) -> Lif {
let Self { tau, r, v_leak } = self;
Lif {
tau,
r,
v_leak,
v_reset,
v_threshold,
metadata,
}
}
}
struct CubaDynamics {
tau_syn: Tensor,
tau_mem: Tensor,
r: Tensor,
w_in: Option<Tensor>,
v_leak: Tensor,
}
impl CubaDynamics {
fn read(r: &NodeReader) -> Result<Self> {
let v_leak = r.tensor("v_leak")?;
Ok(Self {
tau_syn: r.tensor("tau_syn")?,
tau_mem: r.tensor("tau_mem")?,
r: r.tensor("r")?,
w_in: r.w_in(&v_leak)?,
v_leak,
})
}
fn into_li(self, metadata: MetadataMap) -> CubaLi {
let Self {
tau_syn,
tau_mem,
r,
w_in,
v_leak,
} = self;
CubaLi {
tau_syn,
tau_mem,
r,
w_in,
v_leak,
metadata,
}
}
fn into_lif(
self,
v_threshold: Tensor,
v_reset: Option<Tensor>,
metadata: MetadataMap,
) -> CubaLif {
let Self {
tau_syn,
tau_mem,
r,
w_in,
v_leak,
} = self;
CubaLif {
tau_syn,
tau_mem,
r,
v_reset,
w_in,
v_leak,
v_threshold,
metadata,
}
}
}
fn read_cuba_li(r: &NodeReader, metadata: MetadataMap) -> Result<CubaLi> {
Ok(CubaDynamics::read(r)?.into_li(metadata))
}
fn read_cuba_lif(r: &NodeReader, metadata: MetadataMap) -> Result<CubaLif> {
let (v_threshold, v_reset) = threshold_and_reset(r)?;
Ok(CubaDynamics::read(r)?.into_lif(v_threshold, v_reset, metadata))
}
fn read_i(r: &NodeReader, metadata: MetadataMap) -> Result<I> {
Ok(I {
r: r.tensor("r")?,
metadata,
})
}
fn read_if(r: &NodeReader, metadata: MetadataMap) -> Result<If> {
let (v_threshold, v_reset) = threshold_and_reset(r)?;
Ok(If {
r: r.tensor("r")?,
v_reset,
v_threshold,
metadata,
})
}
fn read_li(r: &NodeReader, metadata: MetadataMap) -> Result<Li> {
Ok(LeakDynamics::read(r)?.into_li(metadata))
}
fn read_lif(r: &NodeReader, metadata: MetadataMap) -> Result<Lif> {
let (v_threshold, v_reset) = threshold_and_reset(r)?;
Ok(LeakDynamics::read(r)?.into_lif(v_threshold, v_reset, metadata))
}
fn read_pool_window(r: &NodeReader) -> Result<(Tensor, Tensor, Tensor)> {
Ok((
r.tensor("kernel_size")?,
r.tensor("stride")?,
r.tensor("padding")?,
))
}
fn read_sum_pool2d(r: &NodeReader, metadata: MetadataMap) -> Result<SumPool2d> {
let (kernel_size, stride, padding) = read_pool_window(r)?;
Ok(SumPool2d {
kernel_size,
stride,
padding,
metadata,
})
}
fn read_avg_pool2d(r: &NodeReader, metadata: MetadataMap) -> Result<AvgPool2d> {
let (kernel_size, stride, padding) = read_pool_window(r)?;
Ok(AvgPool2d {
kernel_size,
stride,
padding,
metadata,
})
}
fn read_delay(r: &NodeReader, metadata: MetadataMap) -> Result<Delay> {
Ok(Delay {
delay: r.tensor("delay")?,
metadata,
})
}
fn read_flatten(r: &NodeReader, metadata: MetadataMap) -> Result<Flatten> {
Ok(Flatten {
start_dim: r.opt_int_scalar("start_dim")?.unwrap_or(1),
end_dim: r.opt_int_scalar("end_dim")?.unwrap_or(-1),
input_type: r.opt_usizes("input_type")?,
metadata,
})
}
fn read_threshold(r: &NodeReader, metadata: MetadataMap) -> Result<Threshold> {
Ok(Threshold {
threshold: r.tensor("threshold")?,
metadata,
})
}
struct NodeReader<'a> {
group: &'a Group,
name: &'a str,
budget: &'a ReadBudget,
}
impl NodeReader<'_> {
fn context(&self, field: &str) -> String {
format!("{}.{field}", self.name)
}
fn optional(&self, field: &str) -> Result<Option<Dataset>> {
if !self.group.link_exists(field) {
return Ok(None);
}
self.group.dataset(field).map(Some).map_err(|e| {
NirError::Io(format!(
"{}: expected a dataset, found another link kind: {e}",
self.context(field)
))
})
}
fn required(&self, field: &str) -> Result<Dataset> {
self.optional(field)?
.ok_or_else(|| NirError::MissingField(self.context(field)))
}
fn tensor(&self, field: &str) -> Result<Tensor> {
read_tensor(&self.required(field)?, &self.context(field), self.budget)
}
fn opt_tensor(&self, field: &str) -> Result<Option<Tensor>> {
match self.optional(field)? {
Some(ds) => read_tensor(&ds, &self.context(field), self.budget).map(Some),
None => Ok(None),
}
}
fn ints(&self, field: &str) -> Result<Vec<i64>> {
read_ints(&self.required(field)?, &self.context(field), self.budget)
}
fn opt_ints(&self, field: &str) -> Result<Option<Vec<i64>>> {
match self.optional(field)? {
Some(ds) => read_ints(&ds, &self.context(field), self.budget).map(Some),
None => Ok(None),
}
}
fn int_scalar(&self, field: &str) -> Result<i64> {
single_int(self.ints(field)?, &self.context(field))
}
fn opt_int_scalar(&self, field: &str) -> Result<Option<i64>> {
match self.opt_ints(field)? {
Some(values) => single_int(values, &self.context(field)).map(Some),
None => Ok(None),
}
}
fn usizes(&self, field: &str) -> Result<Vec<usize>> {
to_usizes(self.ints(field)?, &self.context(field), self.budget)
}
fn opt_usizes(&self, field: &str) -> Result<Option<Vec<usize>>> {
match self.opt_ints(field)? {
Some(values) => to_usizes(values, &self.context(field), self.budget).map(Some),
None => Ok(None),
}
}
fn opt_dim(&self, field: &str) -> Result<Option<usize>> {
let Some(values) = self.opt_usizes(field)? else {
return Ok(None);
};
match values.as_slice() {
[only] => Ok(Some(*only)),
other => Err(NirError::InvalidTensor(format!(
"{}: expected a single extent, found {} values",
self.context(field),
other.len()
))),
}
}
fn padding(&self) -> Result<Padding> {
let ds = self.required("padding")?;
if is_string(&ds)? {
wire::padding_from_wire_str(&read_string_scalar(&ds, self.name, self.budget)?)
} else {
Ok(Padding::Explicit(read_ints(
&ds,
&self.context("padding"),
self.budget,
)?))
}
}
fn v_reset(&self, v_threshold: &Tensor) -> Result<Option<Tensor>> {
match self.opt_tensor("v_reset")? {
Some(tensor) => Ok(Some(tensor)),
None => {
self.budget.charge(
&self.context("v_reset (synthesized)"),
v_threshold
.data()
.len()
.checked_mul(v_threshold.dtype().size_of()),
)?;
Ok(Some(v_threshold.zeros_like()))
}
}
}
fn w_in(&self, v_leak: &Tensor) -> Result<Option<Tensor>> {
match self.opt_tensor("w_in")? {
Some(tensor) => Ok(Some(tensor)),
None => {
self.budget.charge(
&self.context("w_in (synthesized)"),
v_leak.data().len().checked_mul(v_leak.dtype().size_of()),
)?;
Ok(Some(v_leak.ones_like()))
}
}
}
}
fn to_usizes(values: Vec<i64>, context: &str, budget: &ReadBudget) -> Result<Vec<usize>> {
budget.charge(context, values.len().checked_mul(size_of::<usize>()))?;
values
.into_iter()
.map(|v| {
usize::try_from(v).map_err(|_| {
NirError::InvalidTensor(format!("{context}: negative axis length {v}"))
})
})
.collect()
}
fn single_int(values: Vec<i64>, context: &str) -> Result<i64> {
match values.as_slice() {
[only] => Ok(*only),
other => Err(NirError::InvalidTensor(format!(
"{context}: expected a single integer, found {} values",
other.len()
))),
}
}
fn group_nlinks(group: &Group, context: &str) -> Result<Option<usize>> {
let mut info = hdf5_sys::h5g::H5G_info_t::default();
let status = hdf5::sync::sync(|| {
unsafe { hdf5_sys::h5g::H5Gget_info(group.id(), &mut info) }
});
hdf5::h5check(status)
.map_err(|e| NirError::Io(format!("{context}: cannot count group members: {e}")))?;
Ok(usize::try_from(info.nlinks).ok())
}
fn validate_group_links(group: &Group, context: &str) -> Result<()> {
let bad_link = group
.iter_visit_default(
None,
|_group, name, info, found: &mut Option<(String, &'static str)>| match info.link_type {
hdf5::LinkType::External => {
*found = Some((name.to_owned(), "external"));
false
}
hdf5::LinkType::Soft => {
*found = Some((name.to_owned(), "soft"));
false
}
hdf5::LinkType::Hard => true,
},
)
.map_err(|e| NirError::Io(format!("{context}: cannot inspect group links: {e}")))?;
if let Some((name, kind)) = bad_link {
return Err(NirError::InvalidGraph(format!(
"{context}: {kind} link '{name}' is not allowed"
)));
}
Ok(())
}
fn validate_dataset_security(ds: &Dataset, context: &str) -> Result<()> {
let dcpl = ds.dcpl().map_err(|e| {
NirError::Io(format!(
"{context}: cannot read dataset creation property list: {e}"
))
})?;
let external = dcpl.external();
if !external.is_empty() {
return Err(NirError::InvalidGraph(format!(
"{context}: external storage is not allowed ({} external file(s))",
external.len()
)));
}
if !matches!(
ds.layout(),
Layout::Compact | Layout::Contiguous | Layout::Chunked
) {
return Err(NirError::InvalidGraph(format!(
"{context}: virtual dataset layouts are not allowed"
)));
}
Ok(())
}
fn read_metadata(group: &Group, budget: &ReadBudget) -> Result<MetadataMap> {
if !group.link_exists(KEY_METADATA) {
return Ok(MetadataMap::new());
}
let md = group.group(KEY_METADATA).map_err(|e| {
NirError::Io(format!(
"{KEY_METADATA}: expected a group, found another link kind: {e}"
))
})?;
validate_group_links(&md, KEY_METADATA)?;
let mut out = MetadataMap::new();
for key in md.member_names()? {
let ds = md.dataset(&key).map_err(|e| {
NirError::Io(format!(
"metadata {key:?} must be a dataset, not a group: {e}"
))
})?;
let value = read_metadata_value(&ds, &key, budget)?;
out.insert(key, value);
}
Ok(out)
}
fn read_string_metadata(
ds: &Dataset,
key: &str,
context: &str,
budget: &ReadBudget,
) -> Result<MetadataValue> {
match ds.shape().as_slice() {
[] => Ok(MetadataValue::String(read_string_scalar_validated(
ds, key, budget,
)?)),
[_] => Ok(MetadataValue::StringList(read_strings_unchecked(
ds, context, budget,
)?)),
shape => Err(NirError::InvalidTensor(format!(
"{context}: string metadata must be scalar or rank-1, found shape {shape:?}"
))),
}
}
fn read_metadata_value(ds: &Dataset, key: &str, budget: &ReadBudget) -> Result<MetadataValue> {
validate_dataset_security(ds, &format!("{KEY_METADATA}.{key}"))?;
let scalar = ds.shape().is_empty();
let context = format!("{KEY_METADATA}.{key}");
let value = match ds.dtype()?.to_descriptor()? {
Td::VarLenUnicode | Td::VarLenAscii | Td::FixedAscii(_) | Td::FixedUnicode(_) => {
read_string_metadata(ds, key, &context, budget)?
}
Td::Boolean if scalar => {
budget.charge(&context, Some(size_of::<bool>()))?;
MetadataValue::Bool(ds.read_scalar::<bool>()?)
}
Td::Float(_) if scalar => {
budget.charge(&context, Some(size_of::<f64>()))?;
MetadataValue::F64(ds.read_scalar::<f64>()?)
}
Td::Integer(_) if scalar => {
budget.charge(&context, Some(size_of::<i64>()))?;
MetadataValue::I64(ds.read_scalar::<i64>()?)
}
Td::Unsigned(IntSize::U8) if scalar => {
budget.charge(&context, Some(size_of::<i64>()))?;
let v = ds.read_scalar::<u64>()?;
MetadataValue::I64(i64::try_from(v).map_err(|_| {
NirError::InvalidTensor(format!("{context}: u64 value {v} does not fit in i64"))
})?)
}
Td::Unsigned(IntSize::U1 | IntSize::U2 | IntSize::U4) if scalar => {
budget.charge(&context, Some(size_of::<i64>()))?;
MetadataValue::I64(ds.read_scalar::<i64>()?)
}
_ => MetadataValue::Tensor(read_tensor(ds, &context, budget)?),
};
Ok(value)
}
fn read_tensor(ds: &Dataset, context: &str, budget: &ReadBudget) -> Result<Tensor> {
validate_dataset_security(ds, context)?;
let descriptor = ds.dtype()?.to_descriptor()?;
let data = match descriptor {
Td::Float(FloatSize::U4) => {
charge_elements(budget, ds, context, size_of::<f32>())?;
TensorData::F32(ds.read_raw::<f32>()?)
}
Td::Float(FloatSize::U8) => {
charge_elements(budget, ds, context, size_of::<f64>())?;
TensorData::F64(ds.read_raw::<f64>()?)
}
Td::Integer(_) | Td::Unsigned(IntSize::U1 | IntSize::U2 | IntSize::U4) => {
charge_elements(budget, ds, context, size_of::<i64>())?;
TensorData::I64(ds.read_raw::<i64>()?)
}
Td::Unsigned(IntSize::U8) => {
let two_buffers = size_of::<u64>().checked_add(size_of::<i64>());
budget.charge(
context,
two_buffers.and_then(|width| ds.size().checked_mul(width)),
)?;
TensorData::I64(read_u64_as_i64(ds, context)?)
}
Td::Boolean => {
charge_elements(budget, ds, context, size_of::<bool>())?;
TensorData::Bool(ds.read_raw::<bool>()?)
}
other => {
return Err(NirError::InvalidTensor(format!(
"{context}: element type {other} has no NIR dtype"
)));
}
};
Tensor::new(ds.shape(), data)
}
fn charge_elements(
budget: &ReadBudget,
ds: &Dataset,
context: &str,
decoded_width: usize,
) -> Result<()> {
budget.charge(context, ds.size().checked_mul(decoded_width))
}
fn read_u64_as_i64(ds: &Dataset, context: &str) -> Result<Vec<i64>> {
ds.read_raw::<u64>()?
.into_iter()
.map(|v| {
i64::try_from(v).map_err(|_| {
NirError::InvalidTensor(format!("{context}: u64 value {v} does not fit in i64"))
})
})
.collect()
}
fn read_ints(ds: &Dataset, context: &str, budget: &ReadBudget) -> Result<Vec<i64>> {
match read_tensor(ds, context, budget)?.into_data() {
TensorData::I64(values) => Ok(values),
other => Err(NirError::InvalidTensor(format!(
"{context}: expected integer data, found {:?}",
other.dtype()
))),
}
}
fn is_string(ds: &Dataset) -> Result<bool> {
Ok(matches!(
ds.dtype()?.to_descriptor()?,
Td::VarLenUnicode | Td::VarLenAscii | Td::FixedAscii(_) | Td::FixedUnicode(_)
))
}
fn read_string_scalar(ds: &Dataset, context: &str, budget: &ReadBudget) -> Result<String> {
validate_dataset_security(ds, context)?;
read_string_scalar_validated(ds, context, budget)
}
fn read_string_scalar_validated(
ds: &Dataset,
context: &str,
budget: &ReadBudget,
) -> Result<String> {
let size = ds.size();
if size != 1 {
return Err(NirError::Io(format!(
"{context}: expected a single string, found {size} elements"
)));
}
let mut values = read_strings_unchecked(ds, context, budget)?;
match values.len() {
1 => Ok(values.remove(0)),
n => Err(NirError::Io(format!(
"{context}: expected a single string, found {n}"
))),
}
}
fn read_strings_unchecked(ds: &Dataset, context: &str, budget: &ReadBudget) -> Result<Vec<String>> {
charge_strings(ds, context, budget)?;
macro_rules! read_fixed {
($ty:ident, $width:expr, $($cap:literal),+) => {
$(if $width <= $cap {
return Ok(ds
.read_raw::<$ty<$cap>>()?
.iter()
.map(ToString::to_string)
.collect());
})+
};
}
match ds.dtype()?.to_descriptor()? {
Td::VarLenUnicode => Ok(ds
.read_raw::<VarLenUnicode>()?
.iter()
.map(ToString::to_string)
.collect()),
Td::VarLenAscii => Ok(ds
.read_raw::<VarLenAscii>()?
.iter()
.map(ToString::to_string)
.collect()),
Td::FixedAscii(width) => {
read_fixed!(FixedAscii, width, 64, 256, 4096);
Err(too_wide(context, width))
}
Td::FixedUnicode(width) => {
read_fixed!(FixedUnicode, width, 64, 256, 4096);
Err(too_wide(context, width))
}
other => Err(NirError::Io(format!(
"{context}: expected a string dataset, found {other}"
))),
}
}
fn charge_strings(ds: &Dataset, context: &str, budget: &ReadBudget) -> Result<()> {
if budget.limit.is_none() {
return Ok(());
}
let count = ds.size();
let headers = count.checked_mul(size_of::<String>());
let requested = match ds.dtype()?.to_descriptor()? {
Td::VarLenUnicode => {
let descriptors = count.checked_mul(size_of::<VarLenUnicode>());
let min = checked_sum([descriptors, headers]);
budget.would_fit(context, min)?;
let payload = vlen_payload_bytes(ds, context)?;
checked_sum([descriptors, Some(payload), headers, Some(payload)])
}
Td::VarLenAscii => {
let descriptors = count.checked_mul(size_of::<VarLenAscii>());
let min = checked_sum([descriptors, headers]);
budget.would_fit(context, min)?;
let payload = vlen_payload_bytes(ds, context)?;
checked_sum([descriptors, Some(payload), headers, Some(payload)])
}
Td::FixedAscii(width) | Td::FixedUnicode(width) => {
let Some(capacity) = FIXED_STRING_CAPS.iter().copied().find(|cap| width <= *cap) else {
return Ok(());
};
checked_sum([
count.checked_mul(capacity),
headers,
count.checked_mul(width),
])
}
_ => return Ok(()),
};
budget.charge(context, requested)
}
fn checked_sum<const N: usize>(parts: [Option<usize>; N]) -> Option<usize> {
parts
.into_iter()
.try_fold(0usize, |sum, part| sum.checked_add(part?))
}
#[allow(deprecated)]
fn vlen_payload_bytes(ds: &Dataset, context: &str) -> Result<usize> {
if ds.shape().is_empty() {
return usize::try_from(ds.file()?.size()).map_err(|_| {
NirError::Io(format!(
"{context}: containing file size does not fit usize"
))
});
}
let dtype = ds.dtype()?;
let space = ds.space()?;
let mut bytes: hdf5_sys::h5::hsize_t = 0;
let status = hdf5::sync::sync(|| {
unsafe { hdf5_sys::h5d::H5Dvlen_get_buf_size(ds.id(), dtype.id(), space.id(), &mut bytes) }
});
hdf5::h5check(status).map_err(|e| {
NirError::Io(format!(
"{context}: cannot determine variable-length string allocation: {e}"
))
})?;
usize::try_from(bytes).map_err(|_| {
NirError::Io(format!(
"{context}: variable-length string allocation does not fit usize"
))
})
}
fn too_wide(context: &str, width: usize) -> NirError {
NirError::Io(format!(
"{context}: fixed-length string of {width} bytes exceeds the supported maximum of 4096"
))
}
#[cfg(test)]
mod tests {
use super::*;
use proptest::prelude::*;
fn count_samples() -> impl Strategy<Value = usize> {
prop_oneof![
Just(0usize),
Just(1usize),
Just(usize::MAX),
Just(usize::MAX - 1),
0usize..4096,
]
}
#[test]
fn read_limit_count_overflow_is_a_structured_error() {
let budget = ReadBudget::new(&ReadOptions::default().with_max_nodes(Some(usize::MAX)));
budget.nodes.used.set(usize::MAX);
let err = budget.charge_nodes("/node/nodes", Some(1)).unwrap_err();
match err {
NirError::ReadCountLimitExceeded {
resource: ReadLimitResource::Nodes,
context,
limit,
used,
requested,
} => {
assert_eq!(context, "/node/nodes");
assert_eq!(limit, usize::MAX);
assert_eq!(used, usize::MAX);
assert_eq!(requested, 1);
}
other => panic!("expected ReadCountLimitExceeded, got {other:?}"),
}
}
#[test]
fn read_limit_count_overflowing_request_is_rejected() {
let budget = ReadBudget::new(&ReadOptions::default().with_max_edges(Some(1)));
let err = budget.charge_edges("/node/edges", None).unwrap_err();
match err {
NirError::ReadCountLimitExceeded {
resource: ReadLimitResource::Edges,
requested,
..
} => assert_eq!(requested, usize::MAX),
other => panic!("expected ReadCountLimitExceeded, got {other:?}"),
}
}
#[test]
fn read_limit_exact_count_succeeds_and_next_fails() {
let budget = ReadBudget::new(&ReadOptions::default().with_max_nested_graphs(Some(2)));
budget.charge_graph("/node").unwrap();
budget.charge_graph("sub").unwrap();
let err = budget.charge_graph("too_deep").unwrap_err();
match err {
NirError::ReadCountLimitExceeded {
resource: ReadLimitResource::NestedGraphs,
context,
limit,
used,
requested,
} => {
assert_eq!(context, "too_deep");
assert_eq!(limit, 2);
assert_eq!(used, 2);
assert_eq!(requested, 1);
}
other => panic!("expected ReadCountLimitExceeded, got {other:?}"),
}
}
proptest! {
#![proptest_config(ProptestConfig::with_cases(64))]
#[test]
fn read_limit_count_charge_matches_checked_add(
used in count_samples(),
requested in count_samples(),
limit in count_samples(),
) {
let budget = ReadBudget::new(&ReadOptions::default().with_max_nodes(Some(limit)));
budget.nodes.used.set(used);
let result = budget.charge_nodes("ctx", Some(requested));
match used.checked_add(requested) {
None => {
prop_assert!(result.is_err());
prop_assert_eq!(budget.nodes.used.get(), used);
}
Some(next) if next > limit => {
prop_assert!(result.is_err());
prop_assert_eq!(budget.nodes.used.get(), used);
}
Some(next) => {
prop_assert!(result.is_ok());
prop_assert_eq!(budget.nodes.used.get(), next);
}
}
}
}
}