Skip to main content

asdf_core/core/
ndarray.rs

1//! The `core/ndarray` schema.
2//!
3//! An array's data lives in one of three places, distinguished by the schema's
4//! `source` and `data` keys:
5//!
6//! - an **internal block**, `source: 0` naming a block by index (negative
7//!   indices count back from the last block);
8//! - an **external file**, `source: "other.asdf"`, used for exploded form;
9//! - **inline** in the tree, under `data`, as nested sequences.
10
11use asdf_yaml::{Document, NodeData, NodeId};
12
13use crate::core::datatype::{ByteOrder, Datatype, ScalarType, parse_shape_with_star};
14use crate::error::{Result, err};
15
16/// Where an array's data comes from.
17#[derive(Clone, PartialEq, Debug)]
18pub enum Source {
19    /// A binary block in this file, by index.
20    Block(usize),
21    /// The last block in this file, written as `source: -1`.
22    ///
23    /// Kept distinct from a resolved index because a streamed array is
24    /// written this way before the block count is known.
25    LastBlock,
26    /// The first block of another ASDF file, named by URI.
27    External(String),
28    /// Nested sequences in the tree itself.
29    Inline(NodeId),
30}
31
32/// What the values in an inline array look like, so a type can be chosen.
33#[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
42/// The narrowest scalar type that holds every value of an inline array.
43///
44/// Inline data may carry no `datatype`, in which case the type is whatever
45/// the values need: a float if any is fractional, a signed integer if any is
46/// negative, and the smallest width that fits otherwise. Strings are not
47/// supported inline, and an array of nothing but booleans is `bool8`.
48pub 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        // Nothing numeric at all: the values were booleans or nulls.
60        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
85/// Walk an inline array's values, recording what types they need.
86fn 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/// How missing values are marked.
119#[derive(Clone, PartialEq, Debug)]
120pub enum Mask {
121    /// A sentinel value; elements equal to it are missing.
122    Value(String),
123    /// Another array of the same shape, non-zero where this array is missing.
124    Array(NodeId),
125}
126
127/// A parsed `core/ndarray`.
128#[derive(Clone, PartialEq, Debug)]
129pub struct Ndarray {
130    /// Where the data lives.
131    pub source: Source,
132    /// The array's shape. A leading `None` means the dimension is determined
133    /// from the block's size, which the schema allows for streamed arrays.
134    pub shape: Vec<Option<u64>>,
135    /// The element type.
136    pub datatype: Datatype,
137    /// Byte order of the elements.
138    pub byteorder: ByteOrder,
139    /// Offset in bytes into the block where the data starts.
140    pub offset: u64,
141    /// Bytes to step per dimension. Absent means C-contiguous.
142    pub strides: Option<Vec<i64>>,
143    /// How missing values are marked, if at all.
144    pub mask: Option<Mask>,
145}
146
147impl Ndarray {
148    /// Parse an ndarray from a tree node.
149    pub fn parse(doc: &Document, id: NodeId) -> Result<Self> {
150        let node = doc.resolved(id);
151
152        // The schema's shorthand: the whole tagged value is the nested data.
153        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                // With no `datatype` key there is nothing to state one, so
159                // it is read off the values.
160                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                // Inline data carries its shape implicitly.
187                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            // Inline data with no declared type is read off the values, as
195            // the shorthand above is.
196            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    /// The shape with every dimension known, given the block's byte length.
244    ///
245    /// A streamed array's first dimension is `*` in the file and is derived
246    /// from how many whole rows the block holds.
247    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    /// The number of elements, for a fully-known shape.
277    pub fn len(&self, block_bytes: Option<u64>) -> Result<u64> {
278        Ok(self.resolved_shape(block_bytes)?.iter().product())
279    }
280
281    /// Whether the array has no elements.
282    pub fn is_empty(&self, block_bytes: Option<u64>) -> Result<bool> {
283        Ok(self.len(block_bytes)? == 0)
284    }
285
286    /// The number of bytes the elements occupy.
287    pub fn nbytes(&self, block_bytes: Option<u64>) -> Result<u64> {
288        Ok(self.len(block_bytes)? * self.datatype.item_size())
289    }
290
291    /// C-contiguous strides for a shape, in bytes.
292    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
303/// Parse the `source` key, which is either a block index or a URI.
304fn 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    // A quoted scalar is always a URI, even if it looks numeric.
310    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            // Other negative indices count back from the end; resolving them
316            // needs the block count, so they are rejected here rather than
317            // guessed at.
318            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
329/// Work out the shape of nested inline sequences.
330fn infer_inline_shape(doc: &Document, id: NodeId) -> Vec<Option<u64>> {
331    let mut shape = Vec::new();
332    let mut current = id;
333    // Follow the first element down; a ragged array is not valid ASDF, so the
334    // first branch describes the whole.
335    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    /// Inline data with no `datatype` takes the narrowest type that holds
350    /// every value, which is what libasdf and Python asdf both do.
351    #[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            // One fractional value makes the whole array a float.
364            ("[1, 2.5]", ScalarType::Float64),
365            // A signed type still has to hold the largest positive value.
366            ("[-1, 200]", ScalarType::Int16),
367            // Strings are not supported inline.
368            ("['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    /// The bare-sequence shorthand infers its type as well as its shape.
381    #[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        // The schema's own example: a tile of a larger image.
418        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        // The schema allows the whole tagged value to be the nested data.
437        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        // Quoting makes it a string, so it names a file rather than a block.
460        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        // Each row is 4 int64s, so 32 bytes; 320 bytes is 10 rows.
483        assert_eq!(nd.resolved_shape(Some(320)).unwrap(), vec![10, 4]);
484        // A partial trailing row is not counted.
485        assert_eq!(nd.resolved_shape(Some(330)).unwrap(), vec![10, 4]);
486        // Without a block size the dimension cannot be resolved.
487        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        // A 2x3 array of 8-byte elements: rows are 24 bytes, columns 8.
514        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}