Skip to main content

asdf_core/core/
elements.rs

1//! Decoding an ndarray's bytes into typed elements, and inlining them back
2//! into the tree.
3//!
4//! Inlining is what the ASDF Standard's reference corpus asks a reader to do
5//! before comparing a `.asdf` file against its expected `.yaml`: turn every
6//! block-backed array into the nested sequences the expected output carries.
7
8use asdf_yaml::{CollectionStyle, Document, Node, NodeData, NodeId, ScalarStyle};
9
10use crate::core::datatype::{ByteOrder, Datatype, ScalarType};
11use crate::core::ndarray::{Ndarray, Source};
12use crate::error::{Result, err};
13
14/// One decoded array element.
15#[derive(Clone, PartialEq, Debug)]
16pub enum Element {
17    /// A signed integer.
18    Int(i64),
19    /// An unsigned integer.
20    Uint(u64),
21    /// A float, including the half- and single-precision types widened to
22    /// `f64`.
23    Float(f64),
24    /// A boolean, from `bool8`.
25    Bool(bool),
26    /// Fixed-length text, with trailing NULs trimmed.
27    Text(String),
28    /// A complex number.
29    Complex(f64, f64),
30    /// One record of a compound array.
31    Record(Vec<Element>),
32}
33
34/// Read a big- or little-endian integer of `n` bytes.
35fn read_uint(bytes: &[u8], order: ByteOrder) -> u64 {
36    let mut acc = 0u64;
37    if order == ByteOrder::Big {
38        for b in bytes {
39            acc = (acc << 8) | u64::from(*b);
40        }
41    } else {
42        for b in bytes.iter().rev() {
43            acc = (acc << 8) | u64::from(*b);
44        }
45    }
46    acc
47}
48
49/// Sign-extend an `n`-byte two's-complement value.
50fn sign_extend(value: u64, bytes: usize) -> i64 {
51    let bits = bytes * 8;
52    if bits >= 64 {
53        return value as i64;
54    }
55    let shift = 64 - bits;
56    ((value << shift) as i64) >> shift
57}
58
59/// The order to actually use for a field: its own if set, else the array's.
60fn effective_order(field: ByteOrder, array: ByteOrder) -> ByteOrder {
61    match field {
62        ByteOrder::Big | ByteOrder::Little => field,
63        // The schema's default when nothing says otherwise.
64        _ => match array {
65            ByteOrder::Big | ByteOrder::Little => array,
66            _ => ByteOrder::Little,
67        },
68    }
69}
70
71/// Decode a single element of `datatype` from the front of `bytes`.
72fn decode_one(datatype: &Datatype, bytes: &[u8], array_order: ByteOrder) -> Result<Element> {
73    if datatype.is_structured() {
74        let mut fields = Vec::with_capacity(datatype.fields.len());
75        let mut offset = 0usize;
76        for field in &datatype.fields {
77            let width = field.datatype.item_size() as usize;
78            let slice = bytes.get(offset..offset + width).ok_or_else(|| {
79                err!(UnexpectedEof, "compound element truncated at field offset {offset}")
80            })?;
81            fields.push(decode_one(&field.datatype, slice, array_order)?);
82            offset += width;
83        }
84        return Ok(Element::Record(fields));
85    }
86
87    let order = effective_order(datatype.byteorder, array_order);
88    let width = datatype.item_size() as usize;
89    let raw = bytes.get(..width).ok_or_else(|| {
90        err!(UnexpectedEof, "element needs {width} bytes, {} available", bytes.len())
91    })?;
92
93    Ok(match datatype.scalar {
94        ScalarType::Bool8 => Element::Bool(raw[0] != 0),
95
96        ScalarType::Uint8 | ScalarType::Uint16 | ScalarType::Uint32 | ScalarType::Uint64 => {
97            Element::Uint(read_uint(raw, order))
98        }
99
100        ScalarType::Int8 | ScalarType::Int16 | ScalarType::Int32 | ScalarType::Int64 => {
101            Element::Int(sign_extend(read_uint(raw, order), width))
102        }
103
104        ScalarType::Float16 => {
105            let bits = read_uint(raw, order) as u16;
106            Element::Float(f64::from(half::f16::from_bits(bits)))
107        }
108        ScalarType::Float32 => {
109            let bits = read_uint(raw, order) as u32;
110            Element::Float(f64::from(f32::from_bits(bits)))
111        }
112        ScalarType::Float64 => Element::Float(f64::from_bits(read_uint(raw, order))),
113
114        ScalarType::Complex64 => {
115            let re = f32::from_bits(read_uint(&raw[..4], order) as u32);
116            let im = f32::from_bits(read_uint(&raw[4..], order) as u32);
117            Element::Complex(f64::from(re), f64::from(im))
118        }
119        ScalarType::Complex128 => {
120            let re = f64::from_bits(read_uint(&raw[..8], order));
121            let im = f64::from_bits(read_uint(&raw[8..], order));
122            Element::Complex(re, im)
123        }
124
125        ScalarType::Ascii => {
126            // Fixed-length text is NUL-padded to its declared width.
127            let end = raw.iter().position(|b| *b == 0).unwrap_or(raw.len());
128            Element::Text(String::from_utf8_lossy(&raw[..end]).into_owned())
129        }
130        ScalarType::Ucs4 => {
131            let mut out = String::new();
132            let (quads, _) = raw.as_chunks::<4>();
133            for chunk in quads {
134                let cp = read_uint(chunk, order) as u32;
135                if cp == 0 {
136                    break;
137                }
138                out.push(char::from_u32(cp).unwrap_or(char::REPLACEMENT_CHARACTER));
139            }
140            Element::Text(out)
141        }
142
143        ScalarType::Unknown | ScalarType::Structured => {
144            return Err(err!(
145                InvalidArgument,
146                "cannot decode a {} element",
147                datatype.scalar.name()
148            ));
149        }
150    })
151}
152
153/// Decode every element of an array, in C order.
154///
155/// `shape` must be fully resolved. Strides, when present, are honoured, so a
156/// tile view or a Fortran-order array reads correctly.
157pub fn decode_all(nd: &Ndarray, shape: &[u64], bytes: &[u8]) -> Result<Vec<Element>> {
158    let item = nd.datatype.item_size();
159    if item == 0 {
160        return Err(err!(InvalidArgument, "cannot decode elements of zero width"));
161    }
162
163    let count: u64 = shape.iter().product();
164    let count = usize::try_from(count)
165        .map_err(|_| err!(OverLimit, "array has too many elements for this platform"))?;
166
167    let strides = match &nd.strides {
168        Some(s) if s.len() == shape.len() => s.clone(),
169        Some(s) => {
170            return Err(err!(
171                InvalidArgument,
172                "strides have {} entries but the shape has {}",
173                s.len(),
174                shape.len()
175            ));
176        }
177        None => Ndarray::c_strides(shape, item),
178    };
179
180    let base = usize::try_from(nd.offset)
181        .map_err(|_| err!(InvalidArgument, "ndarray offset overflows this platform"))?;
182
183    let mut out = Vec::with_capacity(count);
184    let mut index = vec![0u64; shape.len()];
185
186    for _ in 0..count {
187        // Byte position of this element, from the per-dimension strides.
188        let mut pos = base as i64;
189        for (dim, idx) in index.iter().enumerate() {
190            pos += strides[dim] * (*idx as i64);
191        }
192        let pos = usize::try_from(pos)
193            .map_err(|_| err!(InvalidArgument, "strides address a negative offset"))?;
194
195        let slice = bytes.get(pos..).ok_or_else(|| {
196            err!(UnexpectedEof, "element at byte {pos} is past the end of the block")
197        })?;
198        out.push(decode_one(&nd.datatype, slice, nd.byteorder)?);
199
200        // Odometer step, last dimension fastest.
201        for dim in (0..shape.len()).rev() {
202            index[dim] += 1;
203            if index[dim] < shape[dim] {
204                break;
205            }
206            index[dim] = 0;
207        }
208    }
209    Ok(out)
210}
211
212/// Decode an array whose data is inline in the tree rather than in a block.
213///
214/// The elements are already text in the document, so this is the mirror of
215/// [`inline_ndarray`]: walk the nested sequences under `data` in row-major
216/// order and read each scalar as the array's datatype.
217///
218/// The descent is driven by `shape`, not by how deeply the sequences happen
219/// to nest. That distinction matters for a compound array, whose innermost
220/// "sequence" is one record's fields rather than another dimension.
221///
222/// A `mask`ed value written as `null` decodes to the datatype's zero, since
223/// [`Element`] carries no missing-value marker; `mask` itself is available on
224/// the [`Ndarray`] for a caller that needs to know which those were.
225pub fn decode_inline(doc: &Document, array: &Ndarray, shape: &[u64]) -> Result<Vec<Element>> {
226    let Source::Inline(root) = array.source else {
227        return Err(err!(InvalidArgument, "this array's data is not inline"));
228    };
229
230    let expected: u64 = shape.iter().copied().product();
231    let mut out = Vec::with_capacity(usize::try_from(expected).unwrap_or(0));
232    collect_inline(doc, root, &array.datatype, shape, &mut out)?;
233
234    if out.len() as u64 != expected {
235        return Err(err!(
236            InvalidArgument,
237            "inline data holds {} elements but the shape calls for {expected}",
238            out.len()
239        ));
240    }
241    Ok(out)
242}
243
244/// Walk `shape.len()` levels of sequences, reading each leaf as `datatype`.
245fn collect_inline(
246    doc: &Document,
247    node: NodeId,
248    datatype: &Datatype,
249    shape: &[u64],
250    out: &mut Vec<Element>,
251) -> Result<()> {
252    let resolved = doc.resolve(node);
253
254    let Some((dim, rest)) = shape.split_first() else {
255        // Past the last dimension: this is one element.
256        out.push(leaf_element(doc, resolved, datatype)?);
257        return Ok(());
258    };
259
260    let items = doc.sequence_items(resolved).map(<[_]>::to_vec).ok_or_else(|| {
261        err!(InvalidArgument, "inline array data is not nested {} deep", shape.len())
262    })?;
263    if items.len() as u64 != *dim {
264        return Err(err!(
265            InvalidArgument,
266            "inline dimension holds {} entries but the shape calls for {dim}",
267            items.len()
268        ));
269    }
270    for item in items {
271        collect_inline(doc, item, datatype, rest, out)?;
272    }
273    Ok(())
274}
275
276/// Read one element: a scalar, or a record for a compound datatype.
277fn leaf_element(doc: &Document, node: NodeId, datatype: &Datatype) -> Result<Element> {
278    if !datatype.fields.is_empty() {
279        let items = doc.sequence_items(node).map(<[_]>::to_vec).ok_or_else(|| {
280            err!(InvalidArgument, "a compound element must be a sequence of its fields")
281        })?;
282        if items.len() != datatype.fields.len() {
283            return Err(err!(
284                InvalidArgument,
285                "a compound element holds {} values but the datatype has {} fields",
286                items.len(),
287                datatype.fields.len()
288            ));
289        }
290        let mut record = Vec::with_capacity(items.len());
291        for (item, field) in items.iter().zip(datatype.fields.iter()) {
292            record.push(leaf_element(doc, doc.resolve(*item), &field.datatype)?);
293        }
294        return Ok(Element::Record(record));
295    }
296
297    let text = doc
298        .resolved(node)
299        .as_str()
300        .ok_or_else(|| err!(InvalidArgument, "inline array data holds a non-scalar leaf"))?;
301    scalar_element(text, datatype.scalar)
302}
303
304/// Read one inline scalar as `scalar`.
305fn scalar_element(text: &str, scalar: ScalarType) -> Result<Element> {
306    use ScalarType as S;
307
308    // A masked element is written `null`; there is no missing marker in
309    // `Element`, so it reads as the type's zero.
310    if matches!(text, "null" | "~" | "") {
311        return Ok(match scalar {
312            S::Float16 | S::Float32 | S::Float64 => Element::Float(0.0),
313            S::Complex64 | S::Complex128 => Element::Complex(0.0, 0.0),
314            S::Bool8 => Element::Bool(false),
315            S::Ascii | S::Ucs4 => Element::Text(String::new()),
316            S::Uint8 | S::Uint16 | S::Uint32 | S::Uint64 => Element::Uint(0),
317            _ => Element::Int(0),
318        });
319    }
320
321    let bad = |what: &str| err!(InvalidArgument, "inline {what} value {text:?} does not parse");
322    match scalar {
323        S::Uint8 | S::Uint16 | S::Uint32 | S::Uint64 => {
324            Ok(Element::Uint(text.parse::<u64>().map_err(|_| bad("unsigned"))?))
325        }
326        S::Int8 | S::Int16 | S::Int32 | S::Int64 => {
327            Ok(Element::Int(text.parse::<i64>().map_err(|_| bad("integer"))?))
328        }
329        S::Float16 | S::Float32 | S::Float64 => Ok(Element::Float(parse_inline_float(text)?)),
330        S::Complex64 | S::Complex128 => {
331            let (re, im) = parse_inline_complex(text)?;
332            Ok(Element::Complex(re, im))
333        }
334        S::Bool8 => Ok(Element::Bool(matches!(text, "true" | "True" | "1"))),
335        S::Ascii | S::Ucs4 => Ok(Element::Text(text.to_string())),
336        S::Unknown | S::Structured => {
337            Err(err!(InvalidArgument, "inline data needs a known scalar datatype"))
338        }
339    }
340}
341
342/// Parse a float, accepting YAML's non-finite spellings and Python's.
343fn parse_inline_float(text: &str) -> Result<f64> {
344    match text {
345        ".nan" | ".NaN" | ".NAN" | "nan" => return Ok(f64::NAN),
346        ".inf" | ".Inf" | ".INF" | "inf" => return Ok(f64::INFINITY),
347        "-.inf" | "-.Inf" | "-.INF" | "-inf" => return Ok(f64::NEG_INFINITY),
348        _ => {}
349    }
350    text.parse::<f64>()
351        .map_err(|_| err!(InvalidArgument, "inline float value {text:?} does not parse"))
352}
353
354/// Parse a complex number in the spellings the `core/complex` schema allows.
355///
356/// Both `(1+2j)` and `1+2j` are valid, as is a pure imaginary `3j` or a pure
357/// real `1`. The imaginary unit may be `i` or `j`, either case.
358fn parse_inline_complex(text: &str) -> Result<(f64, f64)> {
359    let body = text.trim();
360    let body = body.strip_prefix('(').map_or(body, |rest| rest.strip_suffix(')').unwrap_or(rest));
361
362    let imaginary_unit = |c: char| matches!(c, 'i' | 'I' | 'j' | 'J');
363    let Some(unit) = body.char_indices().rev().find(|(_, c)| imaginary_unit(*c)) else {
364        // No imaginary part at all: a plain real number.
365        return Ok((parse_inline_float(body)?, 0.0));
366    };
367    // The unit must be last; anything after it is not a complex number.
368    if unit.0 + unit.1.len_utf8() != body.len() {
369        return Err(err!(InvalidArgument, "inline complex value {text:?} does not parse"));
370    }
371    let without_unit = &body[..unit.0];
372
373    // Split off the imaginary part at the sign that separates the two, which
374    // is the last `+`/`-` not part of an exponent and not the leading sign.
375    let split = without_unit
376        .char_indices()
377        .rev()
378        .find(|(index, c)| {
379            (*c == '+' || *c == '-')
380                && *index > 0
381                && !matches!(without_unit.as_bytes()[index - 1], b'e' | b'E')
382        })
383        .map(|(index, _)| index);
384
385    match split {
386        None => Ok((0.0, parse_inline_float(without_unit)?)),
387        Some(index) => {
388            let (real, imaginary) = without_unit.split_at(index);
389            // The imaginary part keeps its sign; a bare sign means one.
390            let imaginary = match imaginary {
391                "+" => "1",
392                "-" => "-1",
393                other => other,
394            };
395            Ok((parse_inline_float(real)?, parse_inline_float(imaginary)?))
396        }
397    }
398}
399
400/// The tag Python asdf puts on every inline complex value.
401const COMPLEX_TAG: &str = "tag:stsci.edu:asdf/core/complex-1.0.0";
402
403/// Format a float the way libasdf's emitter does.
404///
405/// `%.17g` for doubles, with YAML's own spellings for the non-finite values.
406pub fn format_float(value: f64) -> String {
407    if value.is_nan() {
408        return ".nan".to_string();
409    }
410    if value.is_infinite() {
411        return if value.is_sign_negative() { "-.inf".into() } else { ".inf".into() };
412    }
413    // Shortest representation that round-trips, which is what Rust's
414    // formatter gives and what a reader will parse back identically.
415    let mut s = format!("{value}");
416    if !s.contains('.') && !s.contains('e') && !s.contains("inf") && !s.contains("nan") {
417        s.push_str(".0");
418    }
419    s
420}
421
422/// Render one element as a tree node.
423fn element_to_node(doc: &mut Document, element: &Element) -> NodeId {
424    match element {
425        Element::Int(v) => doc.add_scalar(v.to_string()),
426        Element::Uint(v) => doc.add_scalar(v.to_string()),
427        Element::Bool(v) => doc.add_scalar(if *v { "true" } else { "false" }),
428        Element::Float(v) => doc.add_scalar(format_float(*v)),
429        // Text is quoted so it round-trips as a string rather than being
430        // re-resolved as a number.
431        Element::Text(s) => doc.add_scalar_styled(s.clone(), ScalarStyle::SingleQuoted),
432        Element::Complex(re, im) => {
433            // The core/complex schema leaves the spelling open; Python asdf
434            // writes CPython's `repr(complex)` and tags each value, so that
435            // is the canonical form to reproduce.
436            let node = Node::scalar(crate::core::pyrepr::repr_complex(*re, *im))
437                .with_tag(asdf_yaml::Tag::parse(COMPLEX_TAG));
438            doc.add(node)
439        }
440        Element::Record(fields) => {
441            let items: Vec<NodeId> = fields.iter().map(|f| element_to_node(doc, f)).collect();
442            doc.add_sequence(items)
443        }
444    }
445}
446
447/// Build nested sequences for `elements` laid out in `shape`.
448pub fn nest(doc: &mut Document, elements: &[Element], shape: &[u64]) -> NodeId {
449    fn build(
450        doc: &mut Document,
451        elements: &[Element],
452        shape: &[u64],
453        cursor: &mut usize,
454    ) -> NodeId {
455        match shape.split_first() {
456            None => {
457                let node = element_to_node(doc, &elements[*cursor]);
458                *cursor += 1;
459                node
460            }
461            Some((dim, rest)) => {
462                let mut items = Vec::with_capacity(*dim as usize);
463                for _ in 0..*dim {
464                    items.push(build(doc, elements, rest, cursor));
465                }
466                let id = doc.add_sequence(items);
467                // Inline data reads better in flow style, which is how both
468                // other implementations write it.
469                if let NodeData::Sequence { style, .. } = &mut doc.node_mut(id).data {
470                    *style = CollectionStyle::Flow;
471                }
472                id
473            }
474        }
475    }
476
477    let mut cursor = 0;
478    build(doc, elements, shape, &mut cursor)
479}
480
481/// Replace an ndarray's `source` with the inline `data` it stands for.
482///
483/// This is the transformation the reference corpus prescribes. The array's
484/// `byteorder`, `offset` and `strides` describe how bytes sit in a block and
485/// become meaningless once the data is inline, so they are removed too.
486pub fn inline_ndarray(
487    doc: &mut Document,
488    id: NodeId,
489    elements: &[Element],
490    shape: &[u64],
491) -> Result<()> {
492    let data = nest(doc, elements, shape);
493    let target = doc.resolve(id);
494
495    if !doc.node(target).is_mapping() {
496        // The bare-sequence shorthand is already inline.
497        return Ok(());
498    }
499
500    doc.mapping_remove(target, "source");
501    for key in ["byteorder", "offset", "strides"] {
502        doc.mapping_remove(target, key);
503    }
504    // A compound datatype's fields carry their own byteorder, which is just
505    // as meaningless once the data is inline.
506    if let Some(dt) = doc.mapping_get(target, "datatype")
507        && let Some(fields) = doc.sequence_items(dt).map(<[_]>::to_vec)
508    {
509        for field in fields {
510            let field = doc.resolve(field);
511            if doc.node(field).is_mapping() {
512                doc.mapping_remove(field, "byteorder");
513            }
514        }
515    }
516    doc.mapping_set(target, "data", data);
517
518    // Record the shape the data actually has, replacing any '*'.
519    let dims: Vec<NodeId> = shape.iter().map(|d| doc.add_scalar(d.to_string())).collect();
520    let shape_node = doc.add_sequence(dims);
521    if let NodeData::Sequence { style, .. } = &mut doc.node_mut(shape_node).data {
522        *style = CollectionStyle::Flow;
523    }
524    doc.mapping_set(target, "shape", shape_node);
525    Ok(())
526}
527
528/// Build a node holding an element, for tests and callers wanting one value.
529pub fn element_node(doc: &mut Document, element: &Element) -> NodeId {
530    element_to_node(doc, element)
531}
532
533/// A node with a tag applied.
534pub fn tagged(doc: &mut Document, node: Node, tag: asdf_yaml::Tag) -> NodeId {
535    doc.add(node.with_tag(tag))
536}
537
538#[cfg(test)]
539mod tests {
540    use super::*;
541    use asdf_yaml::parse_document;
542
543    fn ndarray(yaml: &str) -> Ndarray {
544        let doc = parse_document(yaml).unwrap();
545        let root = doc.root().unwrap();
546        Ndarray::parse(&doc, doc.mapping_get(root, "a").unwrap()).unwrap()
547    }
548
549    #[test]
550    fn inline_integers_decode_from_the_tree() {
551        let doc = parse_document(
552            "a:\n  data: [[1, 2, 3], [4, 5, 6]]\n  datatype: int32\n  shape: [2, 3]\n",
553        )
554        .unwrap();
555        let root = doc.root().unwrap();
556        let nd = Ndarray::parse(&doc, doc.mapping_get(root, "a").unwrap()).unwrap();
557        let shape = nd.resolved_shape(None).unwrap();
558        assert_eq!(shape, vec![2, 3]);
559
560        let els = decode_inline(&doc, &nd, &shape).unwrap();
561        assert_eq!(
562            els,
563            (1..=6).map(Element::Int).collect::<Vec<_>>(),
564            "row-major order, flattened"
565        );
566    }
567
568    #[test]
569    fn inline_floats_accept_yamls_non_finite_spellings() {
570        let doc = parse_document(
571            "a:\n  data: [1.5, .inf, -.inf, .nan]\n  datatype: float64\n  shape: [4]\n",
572        )
573        .unwrap();
574        let root = doc.root().unwrap();
575        let nd = Ndarray::parse(&doc, doc.mapping_get(root, "a").unwrap()).unwrap();
576        let els = decode_inline(&doc, &nd, &[4]).unwrap();
577
578        assert_eq!(els[0], Element::Float(1.5));
579        assert_eq!(els[1], Element::Float(f64::INFINITY));
580        assert_eq!(els[2], Element::Float(f64::NEG_INFINITY));
581        let Element::Float(nan) = els[3] else { panic!("{:?}", els[3]) };
582        assert!(nan.is_nan());
583    }
584
585    /// The schema allows a family of spellings; all of them must read.
586    #[test]
587    fn inline_complex_accepts_every_spelling_the_schema_allows() {
588        let cases = [
589            ("0j", (0.0, 0.0)),
590            ("(1+2j)", (1.0, 2.0)),
591            ("1+2j", (1.0, 2.0)),
592            ("(1-2j)", (1.0, -2.0)),
593            ("-1j", (0.0, -1.0)),
594            ("(-0+0j)", (-0.0, 0.0)),
595            ("3", (3.0, 0.0)),
596            ("2i", (0.0, 2.0)),
597            ("(1.5e-3+2.5e+4j)", (1.5e-3, 2.5e4)),
598            // A bare sign before the unit means one.
599            ("(1+j)", (1.0, 1.0)),
600            ("(1-j)", (1.0, -1.0)),
601        ];
602        for (text, (re, im)) in cases {
603            let got = parse_inline_complex(text).unwrap_or_else(|e| panic!("{text}: {e}"));
604            assert_eq!(got.0, re, "real part of {text}");
605            assert_eq!(got.1, im, "imaginary part of {text}");
606        }
607
608        // The non-finites, which the corpus does carry.
609        let (re, im) = parse_inline_complex("(nan-infj)").unwrap();
610        assert!(re.is_nan());
611        assert_eq!(im, f64::NEG_INFINITY);
612    }
613
614    /// Everything we write must read back, which is the property that
615    /// matters for a round trip through inline form.
616    #[test]
617    fn complex_spellings_round_trip_through_the_parser() {
618        let values = [
619            (0.0, 0.0),
620            (-0.0, 0.0),
621            (1.0, 2.0),
622            (1.0, -2.0),
623            (0.0, -1.0),
624            (1.5e-3, 2.5e4),
625            (f64::MAX, f64::MIN_POSITIVE),
626        ];
627        for (re, im) in values {
628            let text = crate::core::pyrepr::repr_complex(re, im);
629            let (back_re, back_im) = parse_inline_complex(&text).unwrap();
630            assert_eq!(back_re.to_bits(), re.to_bits(), "{text}");
631            assert_eq!(back_im.to_bits(), im.to_bits(), "{text}");
632        }
633    }
634
635    #[test]
636    fn inline_compound_records_stay_grouped() {
637        let doc = parse_document(
638            "a:\n  data: [[1, 2.5], [3, 4.5]]\n  shape: [2]\n  datatype:\n  \
639             - {name: n, datatype: int32}\n  - {name: x, datatype: float64}\n",
640        )
641        .unwrap();
642        let root = doc.root().unwrap();
643        let nd = Ndarray::parse(&doc, doc.mapping_get(root, "a").unwrap()).unwrap();
644        let els = decode_inline(&doc, &nd, &[2]).unwrap();
645        assert_eq!(
646            els,
647            vec![
648                Element::Record(vec![Element::Int(1), Element::Float(2.5)]),
649                Element::Record(vec![Element::Int(3), Element::Float(4.5)]),
650            ]
651        );
652    }
653
654    #[test]
655    fn inline_data_must_match_the_declared_shape() {
656        let doc =
657            parse_document("a:\n  data: [1, 2, 3]\n  datatype: int32\n  shape: [4]\n").unwrap();
658        let root = doc.root().unwrap();
659        let nd = Ndarray::parse(&doc, doc.mapping_get(root, "a").unwrap()).unwrap();
660        let err = decode_inline(&doc, &nd, &[4]).unwrap_err();
661        assert!(err.message().contains("shape calls for 4"), "{err}");
662    }
663
664    /// A block-backed array decoded and re-read through inline form must
665    /// come back the same, which is the property `tree_inlined` relies on.
666    #[test]
667    fn a_block_array_survives_a_trip_through_inline_form() {
668        let nd =
669            ndarray("a:\n  source: 0\n  shape: [5]\n  datatype: float64\n  byteorder: little\n");
670        let values = [1.5f64, -2.25, 0.0, f64::MAX, -0.125];
671        let bytes: Vec<u8> = values.iter().flat_map(|v| v.to_le_bytes()).collect();
672        let original = decode_all(&nd, &[5], &bytes).unwrap();
673
674        let mut doc = parse_document(
675            "a:\n  source: 0\n  shape: [5]\n  datatype: float64\n  byteorder: little\n",
676        )
677        .unwrap();
678        let root = doc.root().unwrap();
679        let node = doc.mapping_get(root, "a").unwrap();
680        inline_ndarray(&mut doc, node, &original, &[5]).unwrap();
681
682        let inlined = Ndarray::parse(&doc, node).unwrap();
683        let read_back = decode_inline(&doc, &inlined, &[5]).unwrap();
684        assert_eq!(read_back, original);
685    }
686
687    #[test]
688    fn decodes_little_endian_integers() {
689        let nd = ndarray("a:\n  source: 0\n  shape: [4]\n  datatype: int32\n  byteorder: little\n");
690        let mut bytes = Vec::new();
691        for v in [1i32, -1, 256, i32::MIN] {
692            bytes.extend_from_slice(&v.to_le_bytes());
693        }
694        let els = decode_all(&nd, &[4], &bytes).unwrap();
695        assert_eq!(
696            els,
697            vec![
698                Element::Int(1),
699                Element::Int(-1),
700                Element::Int(256),
701                Element::Int(i64::from(i32::MIN)),
702            ]
703        );
704    }
705
706    #[test]
707    fn decodes_big_endian_integers() {
708        let nd = ndarray("a:\n  source: 0\n  shape: [3]\n  datatype: int16\n  byteorder: big\n");
709        let mut bytes = Vec::new();
710        for v in [1i16, -2, 1000] {
711            bytes.extend_from_slice(&v.to_be_bytes());
712        }
713        let els = decode_all(&nd, &[3], &bytes).unwrap();
714        assert_eq!(els, vec![Element::Int(1), Element::Int(-2), Element::Int(1000)]);
715    }
716
717    #[test]
718    fn byte_order_actually_changes_the_value() {
719        let bytes = [0x01u8, 0x00];
720        let le =
721            ndarray("a:\n  source: 0\n  shape: [1]\n  datatype: uint16\n  byteorder: little\n");
722        let be = ndarray("a:\n  source: 0\n  shape: [1]\n  datatype: uint16\n  byteorder: big\n");
723        assert_eq!(decode_all(&le, &[1], &bytes).unwrap(), vec![Element::Uint(1)]);
724        assert_eq!(decode_all(&be, &[1], &bytes).unwrap(), vec![Element::Uint(256)]);
725    }
726
727    #[test]
728    fn decodes_floats_of_every_width() {
729        let nd =
730            ndarray("a:\n  source: 0\n  shape: [2]\n  datatype: float64\n  byteorder: little\n");
731        let mut bytes = Vec::new();
732        bytes.extend_from_slice(&1.5f64.to_le_bytes());
733        bytes.extend_from_slice(&(-0.25f64).to_le_bytes());
734        assert_eq!(
735            decode_all(&nd, &[2], &bytes).unwrap(),
736            vec![Element::Float(1.5), Element::Float(-0.25)]
737        );
738
739        let nd =
740            ndarray("a:\n  source: 0\n  shape: [1]\n  datatype: float32\n  byteorder: little\n");
741        assert_eq!(
742            decode_all(&nd, &[1], &2.5f32.to_le_bytes()).unwrap(),
743            vec![Element::Float(2.5)]
744        );
745
746        let nd =
747            ndarray("a:\n  source: 0\n  shape: [1]\n  datatype: float16\n  byteorder: little\n");
748        let h = half::f16::from_f32(0.5);
749        assert_eq!(
750            decode_all(&nd, &[1], &h.to_bits().to_le_bytes()).unwrap(),
751            vec![Element::Float(0.5)]
752        );
753    }
754
755    #[test]
756    fn decodes_bools_and_text() {
757        let nd = ndarray("a:\n  source: 0\n  shape: [2]\n  datatype: bool8\n  byteorder: little\n");
758        assert_eq!(
759            decode_all(&nd, &[2], &[0u8, 1]).unwrap(),
760            vec![Element::Bool(false), Element::Bool(true)]
761        );
762
763        // Fixed-length ASCII is NUL-padded and trimmed on the way out.
764        let nd = ndarray(
765            "a:\n  source: 0\n  shape: [2]\n  datatype: ['ascii', 4]\n  byteorder: little\n",
766        );
767        let bytes = b"M31\0Cas\0";
768        assert_eq!(
769            decode_all(&nd, &[2], bytes).unwrap(),
770            vec![Element::Text("M31".into()), Element::Text("Cas".into())]
771        );
772    }
773
774    #[test]
775    fn decodes_ucs4_text() {
776        let nd = ndarray(
777            "a:\n  source: 0\n  shape: [1]\n  datatype: ['ucs4', 3]\n  byteorder: little\n",
778        );
779        let mut bytes = Vec::new();
780        for cp in ['a' as u32, 0x00E9 /* é */, 0] {
781            bytes.extend_from_slice(&cp.to_le_bytes());
782        }
783        assert_eq!(decode_all(&nd, &[1], &bytes).unwrap(), vec![Element::Text("aé".into())]);
784    }
785
786    #[test]
787    fn honours_offset() {
788        let nd = ndarray(
789            "a:\n  source: 0\n  shape: [2]\n  datatype: uint8\n  byteorder: little\n  offset: 3\n",
790        );
791        let bytes = [9u8, 9, 9, 1, 2];
792        assert_eq!(
793            decode_all(&nd, &[2], &bytes).unwrap(),
794            vec![Element::Uint(1), Element::Uint(2)]
795        );
796    }
797
798    #[test]
799    fn honours_strides_for_a_fortran_order_array() {
800        // A 2x3 array stored column-major: strides are [1, 2] elements.
801        let nd = ndarray(
802            "a:\n  source: 0\n  shape: [2, 3]\n  datatype: uint8\n  byteorder: little\n  \
803             strides: [1, 2]\n",
804        );
805        // Column-major layout of [[1,2,3],[4,5,6]].
806        let bytes = [1u8, 4, 2, 5, 3, 6];
807        let els = decode_all(&nd, &[2, 3], &bytes).unwrap();
808        let values: Vec<u64> = els
809            .iter()
810            .map(|e| match e {
811                Element::Uint(v) => *v,
812                _ => unreachable!(),
813            })
814            .collect();
815        // Read back in C order it must be the logical array.
816        assert_eq!(values, vec![1, 2, 3, 4, 5, 6]);
817    }
818
819    #[test]
820    fn honours_strides_for_a_tile_view() {
821        // A 2x2 tile of a 4x4 uint8 image, starting at row 1 column 1.
822        let nd = ndarray(
823            "a:\n  source: 0\n  shape: [2, 2]\n  datatype: uint8\n  byteorder: little\n  \
824             strides: [4, 1]\n  offset: 5\n",
825        );
826        let bytes: Vec<u8> = (0..16).collect();
827        let els = decode_all(&nd, &[2, 2], &bytes).unwrap();
828        let values: Vec<u64> = els
829            .iter()
830            .map(|e| match e {
831                Element::Uint(v) => *v,
832                _ => unreachable!(),
833            })
834            .collect();
835        assert_eq!(values, vec![5, 6, 9, 10]);
836    }
837
838    #[test]
839    fn decodes_compound_records() {
840        let nd = ndarray(
841            "a:\n  source: 0\n  shape: [2]\n  byteorder: little\n  \
842             datatype:\n    - name: id\n      datatype: uint16\n    \
843             - name: value\n      datatype: float32\n",
844        );
845        let mut bytes = Vec::new();
846        for (id, value) in [(1u16, 1.5f32), (2, -2.5)] {
847            bytes.extend_from_slice(&id.to_le_bytes());
848            bytes.extend_from_slice(&value.to_le_bytes());
849        }
850        let els = decode_all(&nd, &[2], &bytes).unwrap();
851        assert_eq!(
852            els,
853            vec![
854                Element::Record(vec![Element::Uint(1), Element::Float(1.5)]),
855                Element::Record(vec![Element::Uint(2), Element::Float(-2.5)]),
856            ]
857        );
858    }
859
860    #[test]
861    fn truncated_data_is_an_error_not_a_panic() {
862        let nd = ndarray("a:\n  source: 0\n  shape: [4]\n  datatype: int64\n  byteorder: little\n");
863        assert!(decode_all(&nd, &[4], &[0u8; 8]).is_err());
864    }
865
866    #[test]
867    fn nesting_reproduces_the_shape() {
868        let mut doc = Document::new();
869        let els: Vec<Element> = (0..6).map(Element::Uint).collect();
870        let node = nest(&mut doc, &els, &[2, 3]);
871        doc.set_root(node);
872
873        assert_eq!(doc.container_len(node), Some(2));
874        let first = doc.sequence_get(node, 0).unwrap();
875        assert_eq!(doc.container_len(first), Some(3));
876        assert_eq!(doc.resolved(doc.sequence_get(first, 2).unwrap()).as_str(), Some("2"));
877    }
878
879    #[test]
880    fn float_formatting_uses_yaml_spellings() {
881        assert_eq!(format_float(f64::NAN), ".nan");
882        assert_eq!(format_float(f64::INFINITY), ".inf");
883        assert_eq!(format_float(f64::NEG_INFINITY), "-.inf");
884        // A whole float keeps a decimal point so it does not read as an int.
885        assert_eq!(format_float(1.0), "1.0");
886        assert_eq!(format_float(1.5), "1.5");
887    }
888
889    #[test]
890    fn inlining_replaces_source_with_data() {
891        let mut doc = parse_document(
892            "a:\n  source: 0\n  shape: [4]\n  datatype: uint8\n  byteorder: little\n  offset: 0\n",
893        )
894        .unwrap();
895        let root = doc.root().unwrap();
896        let nd_id = doc.mapping_get(root, "a").unwrap();
897
898        let els: Vec<Element> = (0..4).map(Element::Uint).collect();
899        inline_ndarray(&mut doc, nd_id, &els, &[4]).unwrap();
900
901        assert!(doc.mapping_get(nd_id, "source").is_none(), "source must be removed");
902        assert!(doc.mapping_get(nd_id, "byteorder").is_none(), "byteorder is meaningless inline");
903        assert!(doc.mapping_get(nd_id, "offset").is_none(), "offset is meaningless inline");
904
905        let data = doc.mapping_get(nd_id, "data").expect("data must be added");
906        assert_eq!(doc.container_len(data), Some(4));
907        // The datatype survives; it still describes the values.
908        assert!(doc.mapping_get(nd_id, "datatype").is_some());
909    }
910}