1use asdf_yaml::{Document, NodeData, NodeId};
12
13use crate::core::datatype::{ByteOrder, Datatype, ScalarType, parse_shape_with_star};
14use crate::error::{Result, err};
15
16#[derive(Clone, PartialEq, Debug)]
18pub enum Source {
19 Block(usize),
21 LastBlock,
26 External(String),
28 Inline(NodeId),
30}
31
32#[derive(Default, Debug)]
34struct InlineTypes {
35 has_string: bool,
36 has_float: bool,
37 has_signed: bool,
38 int_min: i64,
39 uint_max: u64,
40}
41
42pub fn infer_inline_datatype(doc: &Document, node: NodeId) -> ScalarType {
49 let mut seen = InlineTypes::default();
50 survey_inline(doc, node, &mut seen);
51
52 if seen.has_string {
53 return ScalarType::Unknown;
54 }
55 if seen.has_float {
56 return ScalarType::Float64;
57 }
58 if !seen.has_signed && seen.uint_max == 0 && seen.int_min == 0 {
59 return ScalarType::Bool8;
61 }
62 if seen.has_signed {
63 if seen.int_min >= i64::from(i8::MIN) && seen.uint_max <= i8::MAX as u64 {
64 return ScalarType::Int8;
65 }
66 if seen.int_min >= i64::from(i16::MIN) && seen.uint_max <= i16::MAX as u64 {
67 return ScalarType::Int16;
68 }
69 if seen.int_min >= i64::from(i32::MIN) && seen.uint_max <= i32::MAX as u64 {
70 return ScalarType::Int32;
71 }
72 return ScalarType::Int64;
73 }
74 if seen.uint_max <= u64::from(u8::MAX) {
75 ScalarType::Uint8
76 } else if seen.uint_max <= u64::from(u16::MAX) {
77 ScalarType::Uint16
78 } else if seen.uint_max <= u64::from(u32::MAX) {
79 ScalarType::Uint32
80 } else {
81 ScalarType::Uint64
82 }
83}
84
85fn survey_inline(doc: &Document, node: NodeId, seen: &mut InlineTypes) {
87 let resolved = doc.resolve(node);
88 if let Some(items) = doc.sequence_items(resolved).map(<[_]>::to_vec) {
89 for item in items {
90 survey_inline(doc, item, seen);
91 }
92 return;
93 }
94
95 let Some(text) = doc.resolved(resolved).as_str() else {
96 return;
97 };
98 let style = match &doc.resolved(resolved).data {
99 NodeData::Scalar { style, .. } => *style,
100 _ => return,
101 };
102
103 match asdf_yaml::resolve(text, style, asdf_yaml::Schema::Libasdf) {
104 asdf_yaml::Resolved::Uint(v, _) => seen.uint_max = seen.uint_max.max(v),
105 asdf_yaml::Resolved::Int(v, _) => {
106 seen.has_signed = true;
107 seen.int_min = seen.int_min.min(v);
108 if v > 0 {
109 seen.uint_max = seen.uint_max.max(v as u64);
110 }
111 }
112 asdf_yaml::Resolved::Double(_) => seen.has_float = true,
113 asdf_yaml::Resolved::String => seen.has_string = true,
114 _ => {}
115 }
116}
117
118#[derive(Clone, PartialEq, Debug)]
120pub enum Mask {
121 Value(String),
123 Array(NodeId),
125}
126
127#[derive(Clone, PartialEq, Debug)]
129pub struct Ndarray {
130 pub source: Source,
132 pub shape: Vec<Option<u64>>,
135 pub datatype: Datatype,
137 pub byteorder: ByteOrder,
139 pub offset: u64,
141 pub strides: Option<Vec<i64>>,
143 pub mask: Option<Mask>,
145}
146
147impl Ndarray {
148 pub fn parse(doc: &Document, id: NodeId) -> Result<Self> {
150 let node = doc.resolved(id);
151
152 if matches!(node.data, NodeData::Sequence { .. }) {
154 let data = doc.resolve(id);
155 return Ok(Ndarray {
156 source: Source::Inline(data),
157 shape: infer_inline_shape(doc, data),
158 datatype: Datatype::scalar(infer_inline_datatype(doc, data)),
161 byteorder: ByteOrder::Default,
162 offset: 0,
163 strides: None,
164 mask: None,
165 });
166 }
167
168 if !matches!(node.data, NodeData::Mapping { .. }) {
169 return Err(err!(InvalidArgument, "ndarray must be a mapping or a sequence"));
170 }
171
172 let source = match (doc.mapping_get(id, "source"), doc.mapping_get(id, "data")) {
173 (Some(src), _) => parse_source(doc, src)?,
174 (None, Some(data)) => Source::Inline(doc.resolve(data)),
175 (None, None) => {
176 return Err(err!(
177 InvalidArgument,
178 "ndarray has neither a 'source' nor a 'data' key"
179 ));
180 }
181 };
182
183 let shape = match doc.mapping_get(id, "shape") {
184 Some(s) => parse_shape_with_star(doc, s)?,
185 None => match &source {
186 Source::Inline(node) => infer_inline_shape(doc, *node),
188 _ => Vec::new(),
189 },
190 };
191
192 let datatype = match doc.mapping_get(id, "datatype") {
193 Some(d) => Datatype::parse(doc, d)?,
194 None => match &source {
197 Source::Inline(node) => Datatype::scalar(infer_inline_datatype(doc, *node)),
198 _ => Datatype::default(),
199 },
200 };
201
202 let byteorder = doc
203 .mapping_get(id, "byteorder")
204 .and_then(|b| doc.resolved(b).as_str().map(ByteOrder::from_name))
205 .unwrap_or(ByteOrder::Default);
206
207 let offset = doc
208 .mapping_get(id, "offset")
209 .and_then(|o| doc.resolved(o).as_str().and_then(|s| s.parse().ok()))
210 .unwrap_or(0);
211
212 let strides = match doc.mapping_get(id, "strides") {
213 None => None,
214 Some(s) => {
215 let items = doc
216 .sequence_items(s)
217 .ok_or_else(|| err!(InvalidArgument, "strides must be a sequence"))?;
218 let mut out = Vec::with_capacity(items.len());
219 for item in items {
220 let text = doc
221 .resolved(*item)
222 .as_str()
223 .ok_or_else(|| err!(InvalidArgument, "stride entry is not a scalar"))?;
224 out.push(text.parse::<i64>().map_err(|_| {
225 err!(InvalidArgument, "stride entry is not an integer: {text}")
226 })?);
227 }
228 Some(out)
229 }
230 };
231
232 let mask = doc.mapping_get(id, "mask").map(|m| {
233 let n = doc.resolved(m);
234 match n.data {
235 NodeData::Mapping { .. } | NodeData::Sequence { .. } => Mask::Array(doc.resolve(m)),
236 _ => Mask::Value(n.as_str().unwrap_or_default().to_string()),
237 }
238 });
239
240 Ok(Ndarray { source, shape, datatype, byteorder, offset, strides, mask })
241 }
242
243 pub fn resolved_shape(&self, block_bytes: Option<u64>) -> Result<Vec<u64>> {
248 let item = self.datatype.item_size();
249 let mut out = Vec::with_capacity(self.shape.len());
250
251 for (idx, dim) in self.shape.iter().enumerate() {
252 match dim {
253 Some(d) => out.push(*d),
254 None => {
255 let bytes = block_bytes.ok_or_else(|| {
256 err!(
257 InvalidArgument,
258 "shape dimension {idx} is '*' but no block size is available"
259 )
260 })?;
261 let row: u64 = self.shape[idx + 1..]
262 .iter()
263 .map(|d| d.unwrap_or(1))
264 .product::<u64>()
265 .max(1);
266 let row_bytes = row.checked_mul(item).filter(|b| *b != 0).ok_or_else(|| {
267 err!(InvalidArgument, "cannot size a '*' dimension with a zero-width row")
268 })?;
269 out.push(bytes / row_bytes);
270 }
271 }
272 }
273 Ok(out)
274 }
275
276 pub fn len(&self, block_bytes: Option<u64>) -> Result<u64> {
278 Ok(self.resolved_shape(block_bytes)?.iter().product())
279 }
280
281 pub fn is_empty(&self, block_bytes: Option<u64>) -> Result<bool> {
283 Ok(self.len(block_bytes)? == 0)
284 }
285
286 pub fn nbytes(&self, block_bytes: Option<u64>) -> Result<u64> {
288 Ok(self.len(block_bytes)? * self.datatype.item_size())
289 }
290
291 pub fn c_strides(shape: &[u64], item_size: u64) -> Vec<i64> {
293 let mut strides = vec![0i64; shape.len()];
294 let mut acc = item_size as i64;
295 for idx in (0..shape.len()).rev() {
296 strides[idx] = acc;
297 acc *= shape[idx] as i64;
298 }
299 strides
300 }
301}
302
303fn parse_source(doc: &Document, id: NodeId) -> Result<Source> {
305 let node = doc.resolved(id);
306 let text =
307 node.as_str().ok_or_else(|| err!(InvalidArgument, "ndarray source must be a scalar"))?;
308
309 let quoted = node.scalar_style().is_some_and(|s| s.is_quoted());
311 if !quoted && let Ok(index) = text.parse::<i64>() {
312 return Ok(if index == -1 {
313 Source::LastBlock
314 } else if index < 0 {
315 return Err(err!(
319 InvalidArgument,
320 "negative ndarray source {index} other than -1 is not supported"
321 ));
322 } else {
323 Source::Block(index as usize)
324 });
325 }
326 Ok(Source::External(text.to_string()))
327}
328
329fn infer_inline_shape(doc: &Document, id: NodeId) -> Vec<Option<u64>> {
331 let mut shape = Vec::new();
332 let mut current = id;
333 while let Some(items) = doc.sequence_items(current) {
336 shape.push(Some(items.len() as u64));
337 match items.first() {
338 Some(first) => current = doc.resolve(*first),
339 None => break,
340 }
341 }
342 shape
343}
344
345#[cfg(test)]
346mod tests {
347 use super::*;
348
349 #[test]
352 fn an_inline_arrays_datatype_is_inferred_from_its_values() {
353 let cases = [
354 ("[[0, 1, 2], [3, 4, 5]]", ScalarType::Uint8),
355 ("[0, 255]", ScalarType::Uint8),
356 ("[0, 256]", ScalarType::Uint16),
357 ("[0, 70000]", ScalarType::Uint32),
358 ("[0, 5000000000]", ScalarType::Uint64),
359 ("[-1, 1]", ScalarType::Int8),
360 ("[-200, 1]", ScalarType::Int16),
361 ("[-70000, 1]", ScalarType::Int32),
362 ("[-5000000000, 1]", ScalarType::Int64),
363 ("[1, 2.5]", ScalarType::Float64),
365 ("[-1, 200]", ScalarType::Int16),
367 ("['a', 'b']", ScalarType::Unknown),
369 ("[true, false]", ScalarType::Bool8),
370 ];
371
372 for (data, expected) in cases {
373 let doc = asdf_yaml::parse_document(&format!("a: {data}\n")).unwrap();
374 let root = doc.root().unwrap();
375 let node = doc.mapping_get(root, "a").unwrap();
376 assert_eq!(infer_inline_datatype(&doc, node), expected, "{data}");
377 }
378 }
379
380 #[test]
382 fn the_shorthand_form_infers_both_shape_and_type() {
383 let doc = asdf_yaml::parse_document("a: [[0, 1, 2], [3, 4, 5], [6, 7, 8]]\n").unwrap();
384 let root = doc.root().unwrap();
385 let nd = Ndarray::parse(&doc, doc.mapping_get(root, "a").unwrap()).unwrap();
386
387 assert_eq!(nd.resolved_shape(None).unwrap(), vec![3, 3]);
388 assert_eq!(nd.datatype.scalar, ScalarType::Uint8);
389 assert!(matches!(nd.source, Source::Inline(_)));
390 }
391 use crate::core::datatype::ScalarType;
392 use asdf_yaml::parse_document;
393
394 fn parse_nd(yaml: &str) -> Result<Ndarray> {
395 let doc = parse_document(yaml).unwrap();
396 let root = doc.root().unwrap();
397 let nd = doc.mapping_get(root, "a").unwrap();
398 Ndarray::parse(&doc, nd)
399 }
400
401 #[test]
402 fn parses_a_block_backed_array() {
403 let nd = parse_nd(
404 "a:\n source: 0\n datatype: float64\n shape: [1024, 1024]\n byteorder: little\n",
405 )
406 .unwrap();
407 assert_eq!(nd.source, Source::Block(0));
408 assert_eq!(nd.datatype.scalar, ScalarType::Float64);
409 assert_eq!(nd.byteorder, ByteOrder::Little);
410 assert_eq!(nd.resolved_shape(None).unwrap(), vec![1024, 1024]);
411 assert_eq!(nd.len(None).unwrap(), 1024 * 1024);
412 assert_eq!(nd.nbytes(None).unwrap(), 1024 * 1024 * 8);
413 }
414
415 #[test]
416 fn parses_a_view_with_offset_and_strides() {
417 let nd = parse_nd(
419 "a:\n source: 0\n shape: [256, 256]\n datatype: float64\n \
420 byteorder: little\n strides: [8192, 8]\n offset: 2099200\n",
421 )
422 .unwrap();
423 assert_eq!(nd.offset, 2099200);
424 assert_eq!(nd.strides, Some(vec![8192, 8]));
425 }
426
427 #[test]
428 fn parses_inline_data_under_a_data_key() {
429 let nd = parse_nd("a:\n data: [1, 2, 3, 4]\n datatype: int64\n shape: [4]\n").unwrap();
430 assert!(matches!(nd.source, Source::Inline(_)));
431 assert_eq!(nd.resolved_shape(None).unwrap(), vec![4]);
432 }
433
434 #[test]
435 fn parses_the_bare_sequence_shorthand() {
436 let nd = parse_nd("a: [[1, 0, 0], [0, 1, 0], [0, 0, 1]]\n").unwrap();
438 assert!(matches!(nd.source, Source::Inline(_)));
439 assert_eq!(nd.resolved_shape(None).unwrap(), vec![3, 3]);
440 }
441
442 #[test]
443 fn infers_nested_inline_shape() {
444 let nd = parse_nd("a:\n data: [[1, 2, 3], [4, 5, 6]]\n").unwrap();
445 assert_eq!(nd.resolved_shape(None).unwrap(), vec![2, 3]);
446 }
447
448 #[test]
449 fn an_external_source_is_a_uri() {
450 let nd = parse_nd(
451 "a:\n source: external.asdf\n shape: [4]\n datatype: int8\n byteorder: little\n",
452 )
453 .unwrap();
454 assert_eq!(nd.source, Source::External("external.asdf".into()));
455 }
456
457 #[test]
458 fn a_quoted_numeric_source_is_still_a_uri() {
459 let nd =
461 parse_nd("a:\n source: '0'\n shape: [4]\n datatype: int8\n byteorder: little\n")
462 .unwrap();
463 assert_eq!(nd.source, Source::External("0".into()));
464 }
465
466 #[test]
467 fn source_minus_one_is_the_last_block() {
468 let nd =
469 parse_nd("a:\n source: -1\n shape: ['*']\n datatype: int64\n byteorder: little\n")
470 .unwrap();
471 assert_eq!(nd.source, Source::LastBlock);
472 }
473
474 #[test]
475 fn a_star_dimension_is_sized_from_the_block() {
476 let nd = parse_nd(
477 "a:\n source: -1\n shape: ['*', 4]\n datatype: int64\n byteorder: little\n",
478 )
479 .unwrap();
480 assert_eq!(nd.shape, vec![None, Some(4)]);
481
482 assert_eq!(nd.resolved_shape(Some(320)).unwrap(), vec![10, 4]);
484 assert_eq!(nd.resolved_shape(Some(330)).unwrap(), vec![10, 4]);
486 assert!(nd.resolved_shape(None).is_err());
488 }
489
490 #[test]
491 fn parses_both_mask_forms() {
492 let nd = parse_nd(
493 "a:\n source: 0\n shape: [4]\n datatype: float64\n byteorder: little\n mask: -999\n",
494 )
495 .unwrap();
496 assert_eq!(nd.mask, Some(Mask::Value("-999".into())));
497
498 let nd = parse_nd(
499 "a:\n source: 0\n shape: [4]\n datatype: float64\n byteorder: little\n \
500 mask:\n source: 1\n shape: [4]\n datatype: bool8\n",
501 )
502 .unwrap();
503 assert!(matches!(nd.mask, Some(Mask::Array(_))));
504 }
505
506 #[test]
507 fn rejects_an_ndarray_with_no_data_at_all() {
508 assert!(parse_nd("a:\n shape: [4]\n datatype: int8\n").is_err());
509 }
510
511 #[test]
512 fn c_strides_are_row_major() {
513 assert_eq!(Ndarray::c_strides(&[2, 3], 8), vec![24, 8]);
515 assert_eq!(Ndarray::c_strides(&[4], 4), vec![4]);
516 assert_eq!(Ndarray::c_strides(&[2, 3, 4], 1), vec![12, 4, 1]);
517 }
518
519 #[test]
520 fn compound_arrays_size_by_record() {
521 let nd = parse_nd(
522 "a:\n source: 0\n shape: [64]\n byteorder: little\n \
523 datatype:\n - name: x\n datatype: float64\n \
524 - name: y\n datatype: float64\n",
525 )
526 .unwrap();
527 assert!(nd.datatype.is_structured());
528 assert_eq!(nd.datatype.item_size(), 16);
529 assert_eq!(nd.nbytes(None).unwrap(), 64 * 16);
530 }
531}