1use crate::error::{NirError, Result};
25use crate::nodes::Padding;
26
27pub const KEY_VERSION: &str = "version";
29pub const KEY_NODE: &str = "node";
31pub const KEY_NODES: &str = "nodes";
33pub const KEY_EDGES: &str = "edges";
35pub const KEY_METADATA: &str = "metadata";
37pub const KEY_TYPE: &str = "type";
39
40pub const WIRE_TYPES: [&str; 19] = [
48 "Input",
49 "Output",
50 "Affine",
51 "Linear",
52 "Scale",
53 "Conv1d",
54 "Conv2d",
55 "CubaLI",
56 "CubaLIF",
57 "Delay",
58 "Flatten",
59 "I",
60 "IF",
61 "LI",
62 "LIF",
63 "SumPool2d",
64 "AvgPool2d",
65 "Threshold",
66 "NIRGraph",
67];
68
69#[must_use]
71pub fn is_wire_type(name: &str) -> bool {
72 WIRE_TYPES.contains(&name)
73}
74
75#[must_use]
80pub fn padding_as_wire_str(padding: &Padding) -> Option<&'static str> {
81 match padding {
82 Padding::Same => Some("same"),
83 Padding::Valid => Some("valid"),
84 Padding::Explicit(_) => None,
85 }
86}
87
88pub fn padding_from_wire_str(s: &str) -> Result<Padding> {
95 match s {
96 "same" => Ok(Padding::Same),
97 "valid" => Ok(Padding::Valid),
98 other => Err(NirError::InvalidGraph(format!(
99 "padding must be \"same\", \"valid\", or integer extents, not {other:?}"
100 ))),
101 }
102}
103
104pub fn check_link_name(kind: &str, name: &str) -> Result<()> {
123 if name.is_empty() {
124 return Err(NirError::InvalidGraph(format!("{kind} must not be empty")));
125 }
126 if name.contains('/') {
127 return Err(NirError::InvalidGraph(format!(
128 "{kind} {name:?} must not contain '/' (HDF5 path separator)"
129 )));
130 }
131 if name.contains('\0') {
132 return Err(NirError::InvalidGraph(format!(
133 "{kind} {name:?} must not contain a NUL byte (HDF5 link names are C strings)"
134 )));
135 }
136 if name == "." || name == ".." {
137 return Err(NirError::InvalidGraph(format!(
138 "{kind} {name:?} is a reserved HDF5 path component"
139 )));
140 }
141 Ok(())
142}
143
144pub fn check_hdf5_string(kind: &str, value: &str) -> Result<()> {
154 if value.contains('\0') {
155 return Err(NirError::InvalidGraph(format!(
156 "{kind} {value:?} must not contain a NUL byte (HDF5 strings are C strings)"
157 )));
158 }
159 Ok(())
160}
161
162#[cfg(test)]
163mod tests {
164 use super::*;
165 use crate::graph::NirGraph;
166 use crate::nodes::{
167 Affine, AvgPool2d, Conv1d, Conv2d, CubaLi, CubaLif, Delay, Flatten, I, If, Input, Li, Lif,
168 Linear, NirNode, Output, Scale, SumPool2d, Threshold,
169 };
170 use crate::types::Tensor;
171
172 fn v() -> Tensor {
174 Tensor::from_f64([2], vec![1.0, 1.0]).unwrap()
175 }
176
177 fn pool() -> Tensor {
179 Tensor::from_i64([2], vec![2, 2]).unwrap()
180 }
181
182 fn port_and_linear_nodes() -> Vec<NirNode> {
183 let weight = || Tensor::from_f32(vec![2, 2], vec![1., 0., 0., 1.]).unwrap();
184 vec![
185 NirNode::Input(Input {
186 shape: vec![2],
187 metadata: Default::default(),
188 }),
189 NirNode::Output(Output {
190 shape: vec![2],
191 metadata: Default::default(),
192 }),
193 NirNode::Affine(Affine {
194 weight: weight(),
195 bias: Tensor::from_f32([2], vec![0., 0.]).unwrap(),
196 metadata: Default::default(),
197 }),
198 NirNode::Linear(Linear {
199 weight: weight(),
200 metadata: Default::default(),
201 }),
202 NirNode::Scale(Scale {
203 scale: v(),
204 metadata: Default::default(),
205 }),
206 ]
207 }
208
209 fn conv_nodes() -> Vec<NirNode> {
210 vec![
211 NirNode::Conv1d(Conv1d {
212 weight: Tensor::from_f32(vec![1, 1, 3], vec![1., 0., -1.]).unwrap(),
213 stride: vec![1],
214 padding: Padding::single(0),
215 dilation: vec![1],
216 groups: 1,
217 bias: Tensor::from_f32([1], vec![0.]).unwrap(),
218 input_shape: Some(10),
219 metadata: Default::default(),
220 }),
221 NirNode::Conv2d(Conv2d {
222 weight: Tensor::from_f32(vec![1, 1, 2, 2], vec![0.; 4]).unwrap(),
223 stride: vec![1, 1],
224 padding: Padding::Same,
225 dilation: vec![1, 1],
226 groups: 1,
227 bias: Tensor::from_f32([1], vec![0.]).unwrap(),
228 input_shape: Some(vec![8, 8]),
229 metadata: Default::default(),
230 }),
231 ]
232 }
233
234 fn cuba_nodes() -> Vec<NirNode> {
235 vec![
236 NirNode::CubaLi(CubaLi {
237 tau_syn: v(),
238 tau_mem: v(),
239 r: v(),
240 v_leak: v(),
241 w_in: None,
242 metadata: Default::default(),
243 }),
244 NirNode::CubaLif(CubaLif {
245 tau_syn: v(),
246 tau_mem: v(),
247 r: v(),
248 v_leak: v(),
249 v_threshold: v(),
250 v_reset: None,
251 w_in: None,
252 metadata: Default::default(),
253 }),
254 ]
255 }
256
257 fn neuron_nodes() -> Vec<NirNode> {
258 vec![
259 NirNode::I(I {
260 r: v(),
261 metadata: Default::default(),
262 }),
263 NirNode::If(If {
264 r: v(),
265 v_threshold: v(),
266 v_reset: None,
267 metadata: Default::default(),
268 }),
269 NirNode::Li(Li {
270 tau: v(),
271 r: v(),
272 v_leak: v(),
273 metadata: Default::default(),
274 }),
275 NirNode::Lif(Lif {
276 tau: v(),
277 r: v(),
278 v_leak: v(),
279 v_threshold: v(),
280 v_reset: None,
281 metadata: Default::default(),
282 }),
283 ]
284 }
285
286 fn pool_nodes() -> Vec<NirNode> {
287 let no_pad = || Tensor::from_i64([2], vec![0, 0]).unwrap();
288 vec![
289 NirNode::SumPool2d(SumPool2d {
290 kernel_size: pool(),
291 stride: pool(),
292 padding: no_pad(),
293 metadata: Default::default(),
294 }),
295 NirNode::AvgPool2d(AvgPool2d {
296 kernel_size: pool(),
297 stride: pool(),
298 padding: no_pad(),
299 metadata: Default::default(),
300 }),
301 ]
302 }
303
304 fn variant_index(node: &NirNode) -> usize {
322 match node {
323 NirNode::Input(_) => 0,
324 NirNode::Output(_) => 1,
325 NirNode::Affine(_) => 2,
326 NirNode::Linear(_) => 3,
327 NirNode::Scale(_) => 4,
328 NirNode::Conv1d(_) => 5,
329 NirNode::Conv2d(_) => 6,
330 NirNode::CubaLi(_) => 7,
331 NirNode::CubaLif(_) => 8,
332 NirNode::Delay(_) => 9,
333 NirNode::Flatten(_) => 10,
334 NirNode::I(_) => 11,
335 NirNode::If(_) => 12,
336 NirNode::Li(_) => 13,
337 NirNode::Lif(_) => 14,
338 NirNode::SumPool2d(_) => 15,
339 NirNode::AvgPool2d(_) => 16,
340 NirNode::Threshold(_) => 17,
341 NirNode::Graph(_) => 18,
342 }
343 }
344
345 fn one_of_each() -> Vec<NirNode> {
350 let mut nodes = port_and_linear_nodes();
351 nodes.extend(conv_nodes());
352 nodes.extend(cuba_nodes());
353 nodes.push(NirNode::Delay(Delay {
354 delay: v(),
355 metadata: Default::default(),
356 }));
357 nodes.push(NirNode::Flatten(Flatten {
358 start_dim: 1,
359 end_dim: -1,
360 input_type: None,
361 metadata: Default::default(),
362 }));
363 nodes.extend(neuron_nodes());
364 nodes.extend(pool_nodes());
365 nodes.push(NirNode::Threshold(Threshold {
366 threshold: v(),
367 metadata: Default::default(),
368 }));
369 nodes.push(NirNode::Graph(Box::new(NirGraph::new())));
370 nodes
371 }
372
373 #[test]
374 fn wire_types_matches_every_node_variant() {
375 const VARIANT_COUNT: usize = 19;
378 assert_eq!(WIRE_TYPES.len(), VARIANT_COUNT);
379
380 let nodes = one_of_each();
381 assert_eq!(nodes.len(), VARIANT_COUNT);
382
383 let mut seen = [false; VARIANT_COUNT];
384 for node in &nodes {
385 let i = variant_index(node);
386 assert!(
387 i < VARIANT_COUNT,
388 "variant_index {i} is outside VARIANT_COUNT"
389 );
390 assert!(!seen[i], "duplicate sample for variant index {i}");
391 seen[i] = true;
392 assert_eq!(
393 node.type_name(),
394 WIRE_TYPES[i],
395 "sample at index {i} must match WIRE_TYPES"
396 );
397 }
398 assert!(
399 seen.iter().all(|&s| s),
400 "one_of_each must cover every variant index"
401 );
402 }
403
404 #[test]
405 fn is_wire_type_rejects_marketing_aliases() {
406 for good in WIRE_TYPES {
407 assert!(is_wire_type(good), "{good} should be a wire type");
408 }
409 for bad in ["CurrLIF", "Convolution", "Integrator", "SumPooling", ""] {
410 assert!(!is_wire_type(bad), "{bad} must not be a wire type");
411 }
412 }
413
414 #[test]
415 fn padding_wire_strings_round_trip() {
416 assert_eq!(padding_as_wire_str(&Padding::Same), Some("same"));
417 assert_eq!(padding_as_wire_str(&Padding::Valid), Some("valid"));
418 assert_eq!(padding_as_wire_str(&Padding::pair(1, 1)), None);
419 assert_eq!(padding_from_wire_str("same").unwrap(), Padding::Same);
420 assert_eq!(padding_from_wire_str("valid").unwrap(), Padding::Valid);
421 }
422
423 #[test]
424 fn padding_from_unknown_string_is_rejected() {
425 let err = padding_from_wire_str("SAME").unwrap_err();
426 assert!(matches!(err, NirError::InvalidGraph(_)));
427 assert!(err.to_string().contains("\"SAME\""));
428 }
429
430 #[test]
431 fn node_names_with_dots_are_allowed() {
432 assert!(check_link_name("node name", "lif1.lif").is_ok());
434 assert!(check_link_name("node name", "0").is_ok());
435 }
436
437 #[test]
438 fn illegal_link_names_are_rejected() {
439 for bad in ["", "a/b", ".", "..", "nul\0inside"] {
440 let err = check_link_name("node name", bad).unwrap_err();
441 assert!(
442 matches!(err, NirError::InvalidGraph(_)),
443 "{bad:?} should be rejected"
444 );
445 }
446 }
447
448 #[test]
449 fn the_kind_label_appears_in_the_error() {
450 let err = check_link_name("metadata key", "a/b").unwrap_err();
451 assert!(err.to_string().contains("metadata key"), "got {err}");
452 }
453
454 #[test]
455 fn nul_bytes_in_string_values_are_rejected() {
456 let err = check_hdf5_string("version", "0.2\0.0").unwrap_err();
457 assert!(matches!(err, NirError::InvalidGraph(_)));
458 assert!(err.to_string().contains("version"), "got {err}");
459 assert!(err.to_string().contains("NUL"), "got {err}");
460 }
461}