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> {
283 element_count(&self.resolved_shape(block_bytes)?)
284 }
285
286 pub fn is_empty(&self, block_bytes: Option<u64>) -> Result<bool> {
288 Ok(self.len(block_bytes)? == 0)
289 }
290
291 pub fn nbytes(&self, block_bytes: Option<u64>) -> Result<u64> {
293 self.len(block_bytes)?
294 .checked_mul(self.datatype.item_size())
295 .ok_or_else(|| err!(OverLimit, "array's size in bytes does not fit in 64 bits"))
296 }
297
298 pub fn c_strides(shape: &[u64], item_size: u64) -> Option<Vec<i64>> {
304 let mut strides = vec![0i64; shape.len()];
305 let mut acc = i64::try_from(item_size).ok()?;
306 for idx in (0..shape.len()).rev() {
307 strides[idx] = acc;
308 acc = acc.checked_mul(i64::try_from(shape[idx]).ok()?)?;
309 }
310 Some(strides)
311 }
312}
313
314fn parse_source(doc: &Document, id: NodeId) -> Result<Source> {
316 let node = doc.resolved(id);
317 let text =
318 node.as_str().ok_or_else(|| err!(InvalidArgument, "ndarray source must be a scalar"))?;
319
320 let quoted = node.scalar_style().is_some_and(|s| s.is_quoted());
322 if !quoted && let Ok(index) = text.parse::<i64>() {
323 return Ok(if index == -1 {
324 Source::LastBlock
325 } else if index < 0 {
326 return Err(err!(
330 InvalidArgument,
331 "negative ndarray source {index} other than -1 is not supported"
332 ));
333 } else {
334 Source::Block(index as usize)
335 });
336 }
337 Ok(Source::External(text.to_string()))
338}
339
340pub fn element_count(shape: &[u64]) -> Result<u64> {
348 let mut count: u64 = 1;
349 for dim in shape {
350 count = count.checked_mul(*dim).ok_or_else(|| {
351 err!(OverLimit, "shape {shape:?} has more elements than 64 bits hold")
352 })?;
353 }
354 Ok(count)
355}
356
357fn infer_inline_shape(doc: &Document, id: NodeId) -> Vec<Option<u64>> {
359 let mut shape = Vec::new();
360 let mut current = id;
361 while let Some(items) = doc.sequence_items(current) {
364 shape.push(Some(items.len() as u64));
365 match items.first() {
366 Some(first) => current = doc.resolve(*first),
367 None => break,
368 }
369 }
370 shape
371}
372
373#[cfg(test)]
374mod tests {
375 use super::*;
376
377 #[test]
380 fn an_inline_arrays_datatype_is_inferred_from_its_values() {
381 let cases = [
382 ("[[0, 1, 2], [3, 4, 5]]", ScalarType::Uint8),
383 ("[0, 255]", ScalarType::Uint8),
384 ("[0, 256]", ScalarType::Uint16),
385 ("[0, 70000]", ScalarType::Uint32),
386 ("[0, 5000000000]", ScalarType::Uint64),
387 ("[-1, 1]", ScalarType::Int8),
388 ("[-200, 1]", ScalarType::Int16),
389 ("[-70000, 1]", ScalarType::Int32),
390 ("[-5000000000, 1]", ScalarType::Int64),
391 ("[1, 2.5]", ScalarType::Float64),
393 ("[-1, 200]", ScalarType::Int16),
395 ("['a', 'b']", ScalarType::Unknown),
397 ("[true, false]", ScalarType::Bool8),
398 ];
399
400 for (data, expected) in cases {
401 let doc = asdf_yaml::parse_document(&format!("a: {data}\n")).unwrap();
402 let root = doc.root().unwrap();
403 let node = doc.mapping_get(root, "a").unwrap();
404 assert_eq!(infer_inline_datatype(&doc, node), expected, "{data}");
405 }
406 }
407
408 #[test]
410 fn the_shorthand_form_infers_both_shape_and_type() {
411 let doc = asdf_yaml::parse_document("a: [[0, 1, 2], [3, 4, 5], [6, 7, 8]]\n").unwrap();
412 let root = doc.root().unwrap();
413 let nd = Ndarray::parse(&doc, doc.mapping_get(root, "a").unwrap()).unwrap();
414
415 assert_eq!(nd.resolved_shape(None).unwrap(), vec![3, 3]);
416 assert_eq!(nd.datatype.scalar, ScalarType::Uint8);
417 assert!(matches!(nd.source, Source::Inline(_)));
418 }
419 use crate::core::datatype::ScalarType;
420 use asdf_yaml::parse_document;
421
422 fn parse_nd(yaml: &str) -> Result<Ndarray> {
423 let doc = parse_document(yaml).unwrap();
424 let root = doc.root().unwrap();
425 let nd = doc.mapping_get(root, "a").unwrap();
426 Ndarray::parse(&doc, nd)
427 }
428
429 #[test]
430 fn parses_a_block_backed_array() {
431 let nd = parse_nd(
432 "a:\n source: 0\n datatype: float64\n shape: [1024, 1024]\n byteorder: little\n",
433 )
434 .unwrap();
435 assert_eq!(nd.source, Source::Block(0));
436 assert_eq!(nd.datatype.scalar, ScalarType::Float64);
437 assert_eq!(nd.byteorder, ByteOrder::Little);
438 assert_eq!(nd.resolved_shape(None).unwrap(), vec![1024, 1024]);
439 assert_eq!(nd.len(None).unwrap(), 1024 * 1024);
440 assert_eq!(nd.nbytes(None).unwrap(), 1024 * 1024 * 8);
441 }
442
443 #[test]
444 fn parses_a_view_with_offset_and_strides() {
445 let nd = parse_nd(
447 "a:\n source: 0\n shape: [256, 256]\n datatype: float64\n \
448 byteorder: little\n strides: [8192, 8]\n offset: 2099200\n",
449 )
450 .unwrap();
451 assert_eq!(nd.offset, 2099200);
452 assert_eq!(nd.strides, Some(vec![8192, 8]));
453 }
454
455 #[test]
456 fn parses_inline_data_under_a_data_key() {
457 let nd = parse_nd("a:\n data: [1, 2, 3, 4]\n datatype: int64\n shape: [4]\n").unwrap();
458 assert!(matches!(nd.source, Source::Inline(_)));
459 assert_eq!(nd.resolved_shape(None).unwrap(), vec![4]);
460 }
461
462 #[test]
463 fn parses_the_bare_sequence_shorthand() {
464 let nd = parse_nd("a: [[1, 0, 0], [0, 1, 0], [0, 0, 1]]\n").unwrap();
466 assert!(matches!(nd.source, Source::Inline(_)));
467 assert_eq!(nd.resolved_shape(None).unwrap(), vec![3, 3]);
468 }
469
470 #[test]
471 fn infers_nested_inline_shape() {
472 let nd = parse_nd("a:\n data: [[1, 2, 3], [4, 5, 6]]\n").unwrap();
473 assert_eq!(nd.resolved_shape(None).unwrap(), vec![2, 3]);
474 }
475
476 #[test]
477 fn an_external_source_is_a_uri() {
478 let nd = parse_nd(
479 "a:\n source: external.asdf\n shape: [4]\n datatype: int8\n byteorder: little\n",
480 )
481 .unwrap();
482 assert_eq!(nd.source, Source::External("external.asdf".into()));
483 }
484
485 #[test]
486 fn a_quoted_numeric_source_is_still_a_uri() {
487 let nd =
489 parse_nd("a:\n source: '0'\n shape: [4]\n datatype: int8\n byteorder: little\n")
490 .unwrap();
491 assert_eq!(nd.source, Source::External("0".into()));
492 }
493
494 #[test]
495 fn source_minus_one_is_the_last_block() {
496 let nd =
497 parse_nd("a:\n source: -1\n shape: ['*']\n datatype: int64\n byteorder: little\n")
498 .unwrap();
499 assert_eq!(nd.source, Source::LastBlock);
500 }
501
502 #[test]
503 fn a_star_dimension_is_sized_from_the_block() {
504 let nd = parse_nd(
505 "a:\n source: -1\n shape: ['*', 4]\n datatype: int64\n byteorder: little\n",
506 )
507 .unwrap();
508 assert_eq!(nd.shape, vec![None, Some(4)]);
509
510 assert_eq!(nd.resolved_shape(Some(320)).unwrap(), vec![10, 4]);
512 assert_eq!(nd.resolved_shape(Some(330)).unwrap(), vec![10, 4]);
514 assert!(nd.resolved_shape(None).is_err());
516 }
517
518 #[test]
519 fn parses_both_mask_forms() {
520 let nd = parse_nd(
521 "a:\n source: 0\n shape: [4]\n datatype: float64\n byteorder: little\n mask: -999\n",
522 )
523 .unwrap();
524 assert_eq!(nd.mask, Some(Mask::Value("-999".into())));
525
526 let nd = parse_nd(
527 "a:\n source: 0\n shape: [4]\n datatype: float64\n byteorder: little\n \
528 mask:\n source: 1\n shape: [4]\n datatype: bool8\n",
529 )
530 .unwrap();
531 assert!(matches!(nd.mask, Some(Mask::Array(_))));
532 }
533
534 #[test]
535 fn rejects_an_ndarray_with_no_data_at_all() {
536 assert!(parse_nd("a:\n shape: [4]\n datatype: int8\n").is_err());
537 }
538
539 #[test]
540 fn c_strides_are_row_major() {
541 assert_eq!(Ndarray::c_strides(&[2, 3], 8), Some(vec![24, 8]));
543 assert_eq!(Ndarray::c_strides(&[4], 4), Some(vec![4]));
544 assert_eq!(Ndarray::c_strides(&[2, 3, 4], 1), Some(vec![12, 4, 1]));
545 assert_eq!(Ndarray::c_strides(&[u64::MAX / 2, 4, 4], 8), None);
548 }
549
550 #[test]
551 fn compound_arrays_size_by_record() {
552 let nd = parse_nd(
553 "a:\n source: 0\n shape: [64]\n byteorder: little\n \
554 datatype:\n - name: x\n datatype: float64\n \
555 - name: y\n datatype: float64\n",
556 )
557 .unwrap();
558 assert!(nd.datatype.is_structured());
559 assert_eq!(nd.datatype.item_size(), 16);
560 assert_eq!(nd.nbytes(None).unwrap(), 64 * 16);
561 }
562}