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 survey_inline_bounded(doc, node, seen, 0, &mut inline_budget(doc));
88}
89
90const MAX_INLINE_DEPTH: usize = 64;
95
96fn inline_budget(doc: &Document) -> u64 {
105 (doc.node_count() as u64).saturating_mul(8).max(1024)
106}
107
108fn survey_inline_bounded(
109 doc: &Document,
110 node: NodeId,
111 seen: &mut InlineTypes,
112 depth: usize,
113 budget: &mut u64,
114) {
115 if depth > MAX_INLINE_DEPTH || *budget == 0 {
116 return;
117 }
118 *budget -= 1;
119
120 let resolved = doc.resolve(node);
121 if let Some(items) = doc.sequence_items(resolved).map(<[_]>::to_vec) {
122 for item in items {
123 survey_inline_bounded(doc, item, seen, depth + 1, budget);
124 }
125 return;
126 }
127
128 let Some(text) = doc.resolved(resolved).as_str() else {
129 return;
130 };
131 let style = match &doc.resolved(resolved).data {
132 NodeData::Scalar { style, .. } => *style,
133 _ => return,
134 };
135
136 match asdf_yaml::resolve(text, style, asdf_yaml::Schema::Libasdf) {
137 asdf_yaml::Resolved::Uint(v, _) => seen.uint_max = seen.uint_max.max(v),
138 asdf_yaml::Resolved::Int(v, _) => {
139 seen.has_signed = true;
140 seen.int_min = seen.int_min.min(v);
141 if v > 0 {
142 seen.uint_max = seen.uint_max.max(v as u64);
143 }
144 }
145 asdf_yaml::Resolved::Double(_) => seen.has_float = true,
146 asdf_yaml::Resolved::String => seen.has_string = true,
147 _ => {}
148 }
149}
150
151#[derive(Clone, PartialEq, Debug)]
153pub enum Mask {
154 Value(String),
156 Array(NodeId),
158}
159
160#[derive(Clone, PartialEq, Debug)]
162pub struct Ndarray {
163 pub source: Source,
165 pub shape: Vec<Option<u64>>,
168 pub datatype: Datatype,
170 pub byteorder: ByteOrder,
172 pub offset: u64,
174 pub strides: Option<Vec<i64>>,
176 pub mask: Option<Mask>,
178}
179
180impl Ndarray {
181 pub fn parse(doc: &Document, id: NodeId) -> Result<Self> {
183 let node = doc.resolved(id);
184
185 if matches!(node.data, NodeData::Sequence { .. }) {
187 let data = doc.resolve(id);
188 return Ok(Ndarray {
189 source: Source::Inline(data),
190 shape: infer_inline_shape(doc, data),
191 datatype: Datatype::scalar(infer_inline_datatype(doc, data)),
194 byteorder: ByteOrder::Default,
195 offset: 0,
196 strides: None,
197 mask: None,
198 });
199 }
200
201 if !matches!(node.data, NodeData::Mapping { .. }) {
202 return Err(err!(InvalidArgument, "ndarray must be a mapping or a sequence"));
203 }
204
205 let source = match (doc.mapping_get(id, "source"), doc.mapping_get(id, "data")) {
206 (Some(src), _) => parse_source(doc, src)?,
207 (None, Some(data)) => Source::Inline(doc.resolve(data)),
208 (None, None) => {
209 return Err(err!(
210 InvalidArgument,
211 "ndarray has neither a 'source' nor a 'data' key"
212 ));
213 }
214 };
215
216 let shape = match doc.mapping_get(id, "shape") {
217 Some(s) => parse_shape_with_star(doc, s)?,
218 None => match &source {
219 Source::Inline(node) => infer_inline_shape(doc, *node),
221 _ => Vec::new(),
222 },
223 };
224
225 let datatype = match doc.mapping_get(id, "datatype") {
226 Some(d) => Datatype::parse(doc, d)?,
227 None => match &source {
230 Source::Inline(node) => Datatype::scalar(infer_inline_datatype(doc, *node)),
231 _ => Datatype::default(),
232 },
233 };
234
235 let byteorder = doc
236 .mapping_get(id, "byteorder")
237 .and_then(|b| doc.resolved(b).as_str().map(ByteOrder::from_name))
238 .unwrap_or(ByteOrder::Default);
239
240 let offset = doc
241 .mapping_get(id, "offset")
242 .and_then(|o| doc.resolved(o).as_str().and_then(|s| s.parse().ok()))
243 .unwrap_or(0);
244
245 let strides = match doc.mapping_get(id, "strides") {
246 None => None,
247 Some(s) => {
248 let items = doc
249 .sequence_items(s)
250 .ok_or_else(|| err!(InvalidArgument, "strides must be a sequence"))?;
251 let mut out = Vec::with_capacity(items.len());
252 for item in items {
253 let text = doc
254 .resolved(*item)
255 .as_str()
256 .ok_or_else(|| err!(InvalidArgument, "stride entry is not a scalar"))?;
257 out.push(text.parse::<i64>().map_err(|_| {
258 err!(InvalidArgument, "stride entry is not an integer: {text}")
259 })?);
260 }
261 Some(out)
262 }
263 };
264
265 let mask = doc.mapping_get(id, "mask").map(|m| {
266 let n = doc.resolved(m);
267 match n.data {
268 NodeData::Mapping { .. } | NodeData::Sequence { .. } => Mask::Array(doc.resolve(m)),
269 _ => Mask::Value(n.as_str().unwrap_or_default().to_string()),
270 }
271 });
272
273 Ok(Ndarray { source, shape, datatype, byteorder, offset, strides, mask })
274 }
275
276 #[deny(clippy::arithmetic_side_effects)]
281 pub fn resolved_shape(&self, block_bytes: Option<u64>) -> Result<Vec<u64>> {
282 let item = self.datatype.item_size();
283 let mut out = Vec::with_capacity(self.shape.len());
284
285 for (idx, dim) in self.shape.iter().enumerate() {
286 match dim {
287 Some(d) => out.push(*d),
288 None => {
289 let bytes = block_bytes.ok_or_else(|| {
290 err!(
291 InvalidArgument,
292 "shape dimension {idx} is '*' but no block size is available"
293 )
294 })?;
295 #[allow(
300 clippy::arithmetic_side_effects,
301 reason = "idx indexes self.shape, so idx + 1 is at most its length"
302 )]
303 let tail = &self.shape[idx + 1..];
304 let mut row: u64 = 1;
305 for d in tail {
306 row = row.checked_mul(d.unwrap_or(1)).ok_or_else(|| {
307 err!(
308 OverLimit,
309 "shape {:?} describes a row too large to size a '*' dimension \
310 against",
311 self.shape
312 )
313 })?;
314 }
315 let row = row.max(1);
316 let row_bytes = row.checked_mul(item).filter(|b| *b != 0).ok_or_else(|| {
317 err!(InvalidArgument, "cannot size a '*' dimension with a zero-width row")
318 })?;
319 #[allow(
320 clippy::arithmetic_side_effects,
321 reason = "row_bytes was filtered non-zero just above"
322 )]
323 out.push(bytes / row_bytes);
324 }
325 }
326 }
327 Ok(out)
328 }
329
330 pub fn len(&self, block_bytes: Option<u64>) -> Result<u64> {
337 element_count(&self.resolved_shape(block_bytes)?)
338 }
339
340 pub fn is_empty(&self, block_bytes: Option<u64>) -> Result<bool> {
342 Ok(self.len(block_bytes)? == 0)
343 }
344
345 pub fn nbytes(&self, block_bytes: Option<u64>) -> Result<u64> {
347 self.len(block_bytes)?
348 .checked_mul(self.datatype.item_size())
349 .ok_or_else(|| err!(OverLimit, "array's size in bytes does not fit in 64 bits"))
350 }
351
352 #[deny(clippy::arithmetic_side_effects)]
358 pub fn c_strides(shape: &[u64], item_size: u64) -> Option<Vec<i64>> {
359 let mut strides = vec![0i64; shape.len()];
360 let mut acc = i64::try_from(item_size).ok()?;
361 for idx in (0..shape.len()).rev() {
362 strides[idx] = acc;
363 acc = acc.checked_mul(i64::try_from(shape[idx]).ok()?)?;
364 }
365 Some(strides)
366 }
367}
368
369fn parse_source(doc: &Document, id: NodeId) -> Result<Source> {
371 let node = doc.resolved(id);
372 let text =
373 node.as_str().ok_or_else(|| err!(InvalidArgument, "ndarray source must be a scalar"))?;
374
375 let quoted = node.scalar_style().is_some_and(|s| s.is_quoted());
377 if !quoted && let Ok(index) = text.parse::<i64>() {
378 return Ok(if index == -1 {
379 Source::LastBlock
380 } else if index < 0 {
381 return Err(err!(
385 InvalidArgument,
386 "negative ndarray source {index} other than -1 is not supported"
387 ));
388 } else {
389 Source::Block(index as usize)
390 });
391 }
392 Ok(Source::External(text.to_string()))
393}
394
395#[deny(clippy::arithmetic_side_effects)]
403pub fn element_count(shape: &[u64]) -> Result<u64> {
404 let mut count: u64 = 1;
405 for dim in shape {
406 count = count.checked_mul(*dim).ok_or_else(|| {
407 err!(OverLimit, "shape {shape:?} has more elements than 64 bits hold")
408 })?;
409 }
410 Ok(count)
411}
412
413fn infer_inline_shape(doc: &Document, id: NodeId) -> Vec<Option<u64>> {
415 let mut shape = Vec::new();
416 let mut current = id;
417 while let Some(items) = doc.sequence_items(current) {
425 if shape.len() >= MAX_INLINE_DEPTH {
426 break;
427 }
428 shape.push(Some(items.len() as u64));
429 match items.first() {
430 Some(first) => {
431 let next = doc.resolve(*first);
432 if next == current {
433 break;
434 }
435 current = next;
436 }
437 None => break,
438 }
439 }
440 shape
441}
442
443#[cfg(test)]
444mod tests {
445 use super::*;
446
447 #[test]
450 fn an_inline_arrays_datatype_is_inferred_from_its_values() {
451 let cases = [
452 ("[[0, 1, 2], [3, 4, 5]]", ScalarType::Uint8),
453 ("[0, 255]", ScalarType::Uint8),
454 ("[0, 256]", ScalarType::Uint16),
455 ("[0, 70000]", ScalarType::Uint32),
456 ("[0, 5000000000]", ScalarType::Uint64),
457 ("[-1, 1]", ScalarType::Int8),
458 ("[-200, 1]", ScalarType::Int16),
459 ("[-70000, 1]", ScalarType::Int32),
460 ("[-5000000000, 1]", ScalarType::Int64),
461 ("[1, 2.5]", ScalarType::Float64),
463 ("[-1, 200]", ScalarType::Int16),
465 ("['a', 'b']", ScalarType::Unknown),
467 ("[true, false]", ScalarType::Bool8),
468 ];
469
470 for (data, expected) in cases {
471 let doc = asdf_yaml::parse_document(&format!("a: {data}\n")).unwrap();
472 let root = doc.root().unwrap();
473 let node = doc.mapping_get(root, "a").unwrap();
474 assert_eq!(infer_inline_datatype(&doc, node), expected, "{data}");
475 }
476 }
477
478 #[test]
480 fn the_shorthand_form_infers_both_shape_and_type() {
481 let doc = asdf_yaml::parse_document("a: [[0, 1, 2], [3, 4, 5], [6, 7, 8]]\n").unwrap();
482 let root = doc.root().unwrap();
483 let nd = Ndarray::parse(&doc, doc.mapping_get(root, "a").unwrap()).unwrap();
484
485 assert_eq!(nd.resolved_shape(None).unwrap(), vec![3, 3]);
486 assert_eq!(nd.datatype.scalar, ScalarType::Uint8);
487 assert!(matches!(nd.source, Source::Inline(_)));
488 }
489 use crate::core::datatype::ScalarType;
490 use asdf_yaml::parse_document;
491
492 fn parse_nd(yaml: &str) -> Result<Ndarray> {
493 let doc = parse_document(yaml).unwrap();
494 let root = doc.root().unwrap();
495 let nd = doc.mapping_get(root, "a").unwrap();
496 Ndarray::parse(&doc, nd)
497 }
498
499 #[test]
500 fn parses_a_block_backed_array() {
501 let nd = parse_nd(
502 "a:\n source: 0\n datatype: float64\n shape: [1024, 1024]\n byteorder: little\n",
503 )
504 .unwrap();
505 assert_eq!(nd.source, Source::Block(0));
506 assert_eq!(nd.datatype.scalar, ScalarType::Float64);
507 assert_eq!(nd.byteorder, ByteOrder::Little);
508 assert_eq!(nd.resolved_shape(None).unwrap(), vec![1024, 1024]);
509 assert_eq!(nd.len(None).unwrap(), 1024 * 1024);
510 assert_eq!(nd.nbytes(None).unwrap(), 1024 * 1024 * 8);
511 }
512
513 #[test]
514 fn parses_a_view_with_offset_and_strides() {
515 let nd = parse_nd(
517 "a:\n source: 0\n shape: [256, 256]\n datatype: float64\n \
518 byteorder: little\n strides: [8192, 8]\n offset: 2099200\n",
519 )
520 .unwrap();
521 assert_eq!(nd.offset, 2099200);
522 assert_eq!(nd.strides, Some(vec![8192, 8]));
523 }
524
525 #[test]
526 fn parses_inline_data_under_a_data_key() {
527 let nd = parse_nd("a:\n data: [1, 2, 3, 4]\n datatype: int64\n shape: [4]\n").unwrap();
528 assert!(matches!(nd.source, Source::Inline(_)));
529 assert_eq!(nd.resolved_shape(None).unwrap(), vec![4]);
530 }
531
532 #[test]
533 fn parses_the_bare_sequence_shorthand() {
534 let nd = parse_nd("a: [[1, 0, 0], [0, 1, 0], [0, 0, 1]]\n").unwrap();
536 assert!(matches!(nd.source, Source::Inline(_)));
537 assert_eq!(nd.resolved_shape(None).unwrap(), vec![3, 3]);
538 }
539
540 #[test]
541 fn infers_nested_inline_shape() {
542 let nd = parse_nd("a:\n data: [[1, 2, 3], [4, 5, 6]]\n").unwrap();
543 assert_eq!(nd.resolved_shape(None).unwrap(), vec![2, 3]);
544 }
545
546 #[test]
547 fn an_external_source_is_a_uri() {
548 let nd = parse_nd(
549 "a:\n source: external.asdf\n shape: [4]\n datatype: int8\n byteorder: little\n",
550 )
551 .unwrap();
552 assert_eq!(nd.source, Source::External("external.asdf".into()));
553 }
554
555 #[test]
556 fn a_quoted_numeric_source_is_still_a_uri() {
557 let nd =
559 parse_nd("a:\n source: '0'\n shape: [4]\n datatype: int8\n byteorder: little\n")
560 .unwrap();
561 assert_eq!(nd.source, Source::External("0".into()));
562 }
563
564 #[test]
565 fn source_minus_one_is_the_last_block() {
566 let nd =
567 parse_nd("a:\n source: -1\n shape: ['*']\n datatype: int64\n byteorder: little\n")
568 .unwrap();
569 assert_eq!(nd.source, Source::LastBlock);
570 }
571
572 #[test]
573 fn a_star_dimension_is_sized_from_the_block() {
574 let nd = parse_nd(
575 "a:\n source: -1\n shape: ['*', 4]\n datatype: int64\n byteorder: little\n",
576 )
577 .unwrap();
578 assert_eq!(nd.shape, vec![None, Some(4)]);
579
580 assert_eq!(nd.resolved_shape(Some(320)).unwrap(), vec![10, 4]);
582 assert_eq!(nd.resolved_shape(Some(330)).unwrap(), vec![10, 4]);
584 assert!(nd.resolved_shape(None).is_err());
586 }
587
588 #[test]
589 fn parses_both_mask_forms() {
590 let nd = parse_nd(
591 "a:\n source: 0\n shape: [4]\n datatype: float64\n byteorder: little\n mask: -999\n",
592 )
593 .unwrap();
594 assert_eq!(nd.mask, Some(Mask::Value("-999".into())));
595
596 let nd = parse_nd(
597 "a:\n source: 0\n shape: [4]\n datatype: float64\n byteorder: little\n \
598 mask:\n source: 1\n shape: [4]\n datatype: bool8\n",
599 )
600 .unwrap();
601 assert!(matches!(nd.mask, Some(Mask::Array(_))));
602 }
603
604 #[test]
605 fn rejects_an_ndarray_with_no_data_at_all() {
606 assert!(parse_nd("a:\n shape: [4]\n datatype: int8\n").is_err());
607 }
608
609 #[test]
610 fn c_strides_are_row_major() {
611 assert_eq!(Ndarray::c_strides(&[2, 3], 8), Some(vec![24, 8]));
613 assert_eq!(Ndarray::c_strides(&[4], 4), Some(vec![4]));
614 assert_eq!(Ndarray::c_strides(&[2, 3, 4], 1), Some(vec![12, 4, 1]));
615 assert_eq!(Ndarray::c_strides(&[u64::MAX / 2, 4, 4], 8), None);
618 }
619
620 #[test]
621 fn compound_arrays_size_by_record() {
622 let nd = parse_nd(
623 "a:\n source: 0\n shape: [64]\n byteorder: little\n \
624 datatype:\n - name: x\n datatype: float64\n \
625 - name: y\n datatype: float64\n",
626 )
627 .unwrap();
628 assert!(nd.datatype.is_structured());
629 assert_eq!(nd.datatype.item_size(), 16);
630 assert_eq!(nd.nbytes(None).unwrap(), 64 * 16);
631 }
632}