1use std::collections::{BTreeMap, BTreeSet};
17
18use serde::{Deserialize, Serialize};
19
20use crate::content_hash::StreamingHasher;
21use crate::coordinator::CoordinatorRelationSet;
22use crate::error::{DataError, Result};
23use crate::ids::{ObservationId, RepresentationId, SampleId, SourceId};
24
25pub const ND_TENSOR_MANIFEST_SCHEMA_VERSION: u32 = 1;
27
28pub const ND_TENSOR_MAX_RANK: usize = 16;
31
32#[derive(Clone, Copy, Debug, Eq, PartialEq, Serialize, Deserialize)]
36#[serde(rename_all = "snake_case")]
37pub enum NdTensorDType {
38 F64,
39 F32,
40 U8,
41 I32,
42 Bool,
43}
44
45impl NdTensorDType {
46 pub fn element_size(self) -> usize {
48 match self {
49 NdTensorDType::F64 => 8,
50 NdTensorDType::F32 => 4,
51 NdTensorDType::U8 => 1,
52 NdTensorDType::I32 => 4,
53 NdTensorDType::Bool => 1,
54 }
55 }
56
57 fn fingerprint_tag(self) -> u64 {
64 match self {
65 NdTensorDType::F64 => 1,
66 NdTensorDType::F32 => 2,
67 NdTensorDType::U8 => 3,
68 NdTensorDType::I32 => 4,
69 NdTensorDType::Bool => 5,
70 }
71 }
72}
73
74#[derive(Clone, Debug, PartialEq)]
77pub struct NdTensorInput {
78 pub tensor_id: String,
79 pub representation_id: RepresentationId,
80 pub container: String,
81 pub dtype: NdTensorDType,
82 pub shape: Vec<usize>,
84 pub observation_ids: Vec<ObservationId>,
86 pub sample_ids: Option<Vec<SampleId>>,
89 pub data: Vec<u8>,
99 pub row_presence: Option<Vec<bool>>,
101}
102
103#[derive(Clone, Debug, PartialEq)]
105pub struct NdTensor {
106 tensor_id: String,
107 representation_id: RepresentationId,
108 container: String,
109 dtype: NdTensorDType,
110 shape: Vec<usize>,
111 observation_ids: Vec<ObservationId>,
112 data: Vec<u8>,
113 row_presence: Option<Vec<bool>>,
114 row_index_by_observation: BTreeMap<ObservationId, usize>,
115 row_stride_bytes: usize,
117}
118
119#[derive(Clone, Debug, Eq, PartialEq, Serialize, Deserialize)]
121pub struct NdTensorManifest {
122 pub schema_version: u32,
123 pub tensor_id: String,
124 pub representation_id: RepresentationId,
125 pub container: String,
126 pub dtype: NdTensorDType,
127 pub shape: Vec<usize>,
128 pub observation_ids: Vec<ObservationId>,
129 pub row_count: usize,
130 pub element_bytes: usize,
131 pub data_bytes: usize,
132 pub tensor_fingerprint: String,
133}
134
135#[derive(Clone, Debug, Eq, PartialEq, Serialize, Deserialize)]
137pub struct NdTensorBinding {
138 pub tensor_id: String,
139 pub representation_id: RepresentationId,
140 pub container: String,
141 pub dtype: NdTensorDType,
142 pub source_ids: Vec<SourceId>,
143 pub shape: Vec<usize>,
144 pub row_count: usize,
145 pub tensor_fingerprint: String,
146}
147
148#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
151pub struct NdTensorBlock {
152 pub tensor_id: String,
153 pub representation_id: RepresentationId,
154 pub container: String,
155 pub dtype: NdTensorDType,
156 pub shape: Vec<usize>,
157 pub observation_ids: Vec<ObservationId>,
158 pub sample_ids: Vec<SampleId>,
159 pub data: Vec<u8>,
160 #[serde(default, skip_serializing_if = "Option::is_none")]
161 pub row_presence: Option<Vec<bool>>,
162}
163
164#[derive(Clone, Debug, Default, PartialEq)]
166pub struct NdTensorStore {
167 tensors: BTreeMap<String, NdTensor>,
168}
169
170#[derive(Clone, Debug, Default, PartialEq)]
172pub struct NdTensorArena {
173 store: NdTensorStore,
174 data_bindings: BTreeMap<u64, BTreeMap<String, NdTensorBinding>>,
175}
176
177fn checked_shape_product(tensor_id: &str, shape: &[usize]) -> Result<usize> {
178 let mut product: usize = 1;
179 for dim in shape {
180 product = product.checked_mul(*dim).ok_or_else(|| {
181 DataError::Validation(format!(
182 "nd tensor `{tensor_id}` shape product overflows usize"
183 ))
184 })?;
185 }
186 Ok(product)
187}
188
189impl NdTensor {
190 pub fn from_input(input: NdTensorInput) -> Result<Self> {
192 let NdTensorInput {
193 tensor_id,
194 representation_id,
195 container,
196 dtype,
197 shape,
198 observation_ids,
199 sample_ids,
200 data,
201 row_presence,
202 } = input;
203
204 if tensor_id.trim().is_empty() {
205 return Err(DataError::Validation(
206 "nd tensor has an empty tensor id".to_string(),
207 ));
208 }
209 if container.trim().is_empty() {
210 return Err(DataError::Validation(format!(
211 "nd tensor `{tensor_id}` has an empty container"
212 )));
213 }
214 if shape.is_empty() || shape.len() > ND_TENSOR_MAX_RANK {
215 return Err(DataError::Validation(format!(
216 "nd tensor `{tensor_id}` rank {} is not in 1..={ND_TENSOR_MAX_RANK}",
217 shape.len()
218 )));
219 }
220 if shape.contains(&0) {
221 return Err(DataError::Validation(format!(
222 "nd tensor `{tensor_id}` has a zero dimension in shape {shape:?}"
223 )));
224 }
225 let row_count = shape[0];
226 if observation_ids.len() != row_count {
227 return Err(DataError::Validation(format!(
228 "nd tensor `{tensor_id}` has {} observation ids for axis-0 size {row_count}",
229 observation_ids.len()
230 )));
231 }
232 if observation_ids.is_empty() {
233 return Err(DataError::Validation(format!(
234 "nd tensor `{tensor_id}` has no observations"
235 )));
236 }
237 let mut row_index_by_observation = BTreeMap::new();
238 for (idx, observation_id) in observation_ids.iter().enumerate() {
239 if row_index_by_observation
240 .insert(observation_id.clone(), idx)
241 .is_some()
242 {
243 return Err(DataError::Validation(format!(
244 "nd tensor `{tensor_id}` has duplicate observation `{observation_id}`"
245 )));
246 }
247 }
248 if let Some(sample_ids) = &sample_ids {
249 if sample_ids.len() != row_count {
250 return Err(DataError::Validation(format!(
251 "nd tensor `{tensor_id}` has {} sample ids for axis-0 size {row_count}",
252 sample_ids.len()
253 )));
254 }
255 }
256
257 let element_size = dtype.element_size();
258 let total_elements = checked_shape_product(&tensor_id, &shape)?;
259 let expected_bytes = total_elements.checked_mul(element_size).ok_or_else(|| {
260 DataError::Validation(format!("nd tensor `{tensor_id}` byte size overflows usize"))
261 })?;
262 if data.len() != expected_bytes {
263 return Err(DataError::Validation(format!(
264 "nd tensor `{tensor_id}` has {} data bytes for shape {shape:?} dtype {dtype:?} ({expected_bytes} expected)",
265 data.len()
266 )));
267 }
268 if dtype == NdTensorDType::Bool && data.iter().any(|byte| *byte > 1) {
269 return Err(DataError::Validation(format!(
270 "nd tensor `{tensor_id}` bool payload contains a byte that is not 0 or 1"
271 )));
272 }
273 if let Some(row_presence) = &row_presence {
274 if row_presence.len() != row_count {
275 return Err(DataError::Validation(format!(
276 "nd tensor `{tensor_id}` row presence has {} flags for axis-0 size {row_count}",
277 row_presence.len()
278 )));
279 }
280 }
281
282 let row_stride_bytes = expected_bytes / row_count;
285
286 Ok(Self {
287 tensor_id,
288 representation_id,
289 container,
290 dtype,
291 shape,
292 observation_ids,
293 data,
294 row_presence,
295 row_index_by_observation,
296 row_stride_bytes,
297 })
298 }
299
300 fn contains_observation(&self, observation_id: &ObservationId) -> bool {
301 self.row_index_by_observation.contains_key(observation_id)
302 }
303
304 pub fn project_relations(
308 &self,
309 relations: &CoordinatorRelationSet,
310 source_id: Option<&SourceId>,
311 ) -> Result<NdTensorBlock> {
312 relations.validate()?;
313 let selected = relations.records.iter().filter(|relation| {
314 source_id
315 .map(|source_id| relation.source_id.as_ref() == Some(source_id))
316 .unwrap_or(true)
317 });
318
319 let mut observation_ids = Vec::new();
320 let mut sample_ids = Vec::new();
321 let mut data = Vec::new();
322 let mut presence = Vec::new();
323 for relation in selected {
324 let row_idx = *self
325 .row_index_by_observation
326 .get(&relation.observation_id)
327 .ok_or_else(|| {
328 DataError::Validation(format!(
329 "nd tensor `{}` has no row for observation `{}`",
330 self.tensor_id, relation.observation_id
331 ))
332 })?;
333 let start = row_idx * self.row_stride_bytes;
334 data.extend_from_slice(&self.data[start..start + self.row_stride_bytes]);
335 if let Some(row_presence) = &self.row_presence {
336 presence.push(row_presence[row_idx]);
337 }
338 observation_ids.push(relation.observation_id.clone());
339 sample_ids.push(relation.sample_id.clone());
340 }
341
342 let mut shape = self.shape.clone();
343 shape[0] = observation_ids.len();
344
345 Ok(NdTensorBlock {
346 tensor_id: self.tensor_id.clone(),
347 representation_id: self.representation_id.clone(),
348 container: self.container.clone(),
349 dtype: self.dtype,
350 shape,
351 observation_ids,
352 sample_ids,
353 data,
354 row_presence: self.row_presence.as_ref().map(|_| presence),
355 })
356 }
357
358 fn fingerprint(&self) -> Result<String> {
389 let mut hasher = StreamingHasher::new(b"dag-ml-data.nd-tensor.v2\0");
390 hasher.absorb_str(&self.tensor_id);
391 hasher.absorb_str(self.representation_id.as_str());
392 hasher.absorb_str(&self.container);
393 hasher.absorb_u64(self.dtype.fingerprint_tag());
394 hasher.absorb_len(self.shape.len());
395 for dim in &self.shape {
396 hasher.absorb_len(*dim);
397 }
398 hasher.absorb_str_collection(self.observation_ids.iter().map(ObservationId::as_str));
399 hasher.absorb_len(self.data.len());
400 hasher.absorb_raw(&self.data);
401 match &self.row_presence {
402 None => hasher.absorb_u64(0),
403 Some(presence) => {
404 hasher.absorb_u64(1);
405 hasher.absorb_len(presence.len());
406 for present in presence {
407 hasher.absorb_raw(&[u8::from(*present)]);
408 }
409 }
410 }
411 Ok(hasher.finalize_hex())
412 }
413
414 fn manifest(&self) -> Result<NdTensorManifest> {
415 Ok(NdTensorManifest {
416 schema_version: ND_TENSOR_MANIFEST_SCHEMA_VERSION,
417 tensor_id: self.tensor_id.clone(),
418 representation_id: self.representation_id.clone(),
419 container: self.container.clone(),
420 dtype: self.dtype,
421 shape: self.shape.clone(),
422 observation_ids: self.observation_ids.clone(),
423 row_count: self.shape[0],
424 element_bytes: self.dtype.element_size(),
425 data_bytes: self.data.len(),
426 tensor_fingerprint: self.fingerprint()?,
427 })
428 }
429
430 fn binding_for_sources(&self, source_ids: Vec<SourceId>) -> Result<NdTensorBinding> {
431 Ok(NdTensorBinding {
432 tensor_id: self.tensor_id.clone(),
433 representation_id: self.representation_id.clone(),
434 container: self.container.clone(),
435 dtype: self.dtype,
436 source_ids,
437 shape: self.shape.clone(),
438 row_count: self.shape[0],
439 tensor_fingerprint: self.fingerprint()?,
440 })
441 }
442}
443
444impl NdTensorStore {
445 pub fn from_inputs(inputs: Vec<NdTensorInput>) -> Result<Self> {
447 let mut tensors = BTreeMap::new();
448 for input in inputs {
449 let tensor = NdTensor::from_input(input)?;
450 let tensor_id = tensor.tensor_id.clone();
451 if tensors.insert(tensor_id.clone(), tensor).is_some() {
452 return Err(DataError::Validation(format!(
453 "duplicate nd tensor `{tensor_id}`"
454 )));
455 }
456 }
457 Ok(Self { tensors })
458 }
459
460 pub fn is_empty(&self) -> bool {
461 self.tensors.is_empty()
462 }
463
464 pub fn manifests(&self) -> Result<Vec<NdTensorManifest>> {
466 self.tensors.values().map(NdTensor::manifest).collect()
467 }
468
469 pub fn bindings_for_relations(
472 &self,
473 relations: &CoordinatorRelationSet,
474 representation_id: &RepresentationId,
475 ) -> Result<Vec<NdTensorBinding>> {
476 relations.validate()?;
477 let source_ids = relations
478 .records
479 .iter()
480 .filter_map(|relation| relation.source_id.as_ref())
481 .collect::<BTreeSet<_>>();
482
483 let mut bindings = Vec::new();
484 for tensor in self.tensors.values() {
485 if &tensor.representation_id != representation_id {
486 continue;
487 }
488 if source_ids.is_empty() {
489 if relations
490 .records
491 .iter()
492 .all(|relation| tensor.contains_observation(&relation.observation_id))
493 {
494 bindings.push(tensor.binding_for_sources(Vec::new())?);
495 }
496 continue;
497 }
498 let mut covered_sources = Vec::new();
499 for source_id in &source_ids {
500 let covers_source = relations
501 .records
502 .iter()
503 .filter(|relation| relation.source_id.as_ref() == Some(*source_id))
504 .all(|relation| tensor.contains_observation(&relation.observation_id));
505 if covers_source {
506 covered_sources.push((*source_id).clone());
507 }
508 }
509 if !covered_sources.is_empty() {
510 bindings.push(tensor.binding_for_sources(covered_sources)?);
511 }
512 }
513 Ok(bindings)
514 }
515
516 pub fn project_relations(
518 &self,
519 tensor_id: &str,
520 relations: &CoordinatorRelationSet,
521 source_id: Option<&SourceId>,
522 ) -> Result<NdTensorBlock> {
523 let tensor = self.tensors.get(tensor_id).ok_or_else(|| {
524 DataError::Validation(format!("nd tensor `{tensor_id}` is not present"))
525 })?;
526 tensor.project_relations(relations, source_id)
527 }
528}
529
530impl NdTensorArena {
531 pub fn new(store: NdTensorStore) -> Self {
532 Self {
533 store,
534 data_bindings: BTreeMap::new(),
535 }
536 }
537
538 pub fn bind_data_handle(
541 &mut self,
542 data_handle: u64,
543 relations: &CoordinatorRelationSet,
544 representation_id: &RepresentationId,
545 ) -> Result<Vec<NdTensorBinding>> {
546 let bindings = self
547 .store
548 .bindings_for_relations(relations, representation_id)?;
549 self.data_bindings.insert(
550 data_handle,
551 bindings
552 .iter()
553 .cloned()
554 .map(|binding| (binding.tensor_id.clone(), binding))
555 .collect(),
556 );
557 Ok(bindings)
558 }
559
560 pub fn release_data_handle(&mut self, data_handle: u64) -> bool {
561 self.data_bindings.remove(&data_handle).is_some()
562 }
563
564 pub fn manifests(&self) -> Result<Vec<NdTensorManifest>> {
565 self.store.manifests()
566 }
567
568 pub fn bindings_for_data_handle(&self, data_handle: u64) -> Result<Vec<NdTensorBinding>> {
569 let bindings = self.data_bindings.get(&data_handle).ok_or_else(|| {
570 DataError::Validation(format!(
571 "data handle `{data_handle}` has no nd tensor bindings"
572 ))
573 })?;
574 Ok(bindings.values().cloned().collect())
575 }
576
577 pub fn project_bound_relations(
582 &self,
583 data_handle: u64,
584 tensor_id: &str,
585 relations: &CoordinatorRelationSet,
586 source_id: Option<&SourceId>,
587 ) -> Result<NdTensorBlock> {
588 relations.validate()?;
589 let binding = self
590 .data_bindings
591 .get(&data_handle)
592 .and_then(|bindings| bindings.get(tensor_id))
593 .ok_or_else(|| {
594 DataError::Validation(format!(
595 "nd tensor `{tensor_id}` is not bound to data handle `{data_handle}`"
596 ))
597 })?;
598 let view_source_ids = relations
599 .records
600 .iter()
601 .filter_map(|relation| relation.source_id.as_ref())
602 .cloned()
603 .collect::<BTreeSet<_>>();
604 let required_source_ids: Vec<SourceId> = if let Some(source_id) = source_id {
605 if view_source_ids.is_empty() || !view_source_ids.contains(source_id) {
606 return Err(DataError::Validation(format!(
607 "nd tensor `{tensor_id}` source `{source_id}` is not present in view for data handle `{data_handle}`"
608 )));
609 }
610 vec![source_id.clone()]
611 } else {
612 view_source_ids.into_iter().collect()
613 };
614 for required in &required_source_ids {
615 if !binding.source_ids.contains(required) {
616 return Err(DataError::Validation(format!(
617 "nd tensor `{tensor_id}` is not bound to source `{required}` for data handle `{data_handle}`"
618 )));
619 }
620 }
621 self.store
622 .project_relations(tensor_id, relations, source_id)
623 }
624}
625
626#[cfg(test)]
627mod tests {
628 use super::*;
629 use crate::coordinator::CoordinatorRelation;
630 use crate::ids::SampleId;
631
632 fn relation(observation: &str, sample: &str, source: &str) -> CoordinatorRelation {
633 CoordinatorRelation {
634 unit_level: crate::CoordinatorEntityUnitLevel::Observation,
635 unit_id: None,
636 rep_id: None,
637 derived_unit_id: None,
638 component_observation_ids: Vec::new(),
639 sample_influence_weight: None,
640 quality_flag: None,
641 observation_id: ObservationId::new(observation).unwrap(),
642 sample_id: SampleId::new(sample).unwrap(),
643 target_id: None,
644 group_id: None,
645 origin_sample_id: None,
646 source_id: Some(SourceId::new(source).unwrap()),
647 is_augmented: false,
648 excluded: false,
649 metadata: BTreeMap::new(),
650 tags: Vec::new(),
651 }
652 }
653
654 fn rgb_input() -> NdTensorInput {
656 NdTensorInput {
657 tensor_id: "rgb".to_string(),
658 representation_id: RepresentationId::new("rgb_image").unwrap(),
659 container: "pil_image_batch".to_string(),
660 dtype: NdTensorDType::U8,
661 shape: vec![3, 2, 2],
662 observation_ids: vec![
663 ObservationId::new("obs.s1").unwrap(),
664 ObservationId::new("obs.s2").unwrap(),
665 ObservationId::new("obs.s3").unwrap(),
666 ],
667 sample_ids: None,
668 data: (0u8..12).collect(),
669 row_presence: None,
670 }
671 }
672
673 #[test]
674 fn from_input_validates_and_projects_in_relation_order() {
675 let store = NdTensorStore::from_inputs(vec![rgb_input()]).unwrap();
676 let relations = CoordinatorRelationSet {
677 records: vec![
678 relation("obs.s3", "s3", "cam"),
679 relation("obs.s1", "s1", "cam"),
680 ],
681 };
682 let block = store.project_relations("rgb", &relations, None).unwrap();
683 assert_eq!(block.shape, vec![2, 2, 2]);
684 assert_eq!(block.dtype, NdTensorDType::U8);
685 assert_eq!(block.data, vec![8, 9, 10, 11, 0, 1, 2, 3]);
687 assert_eq!(
688 block.sample_ids,
689 vec![SampleId::new("s3").unwrap(), SampleId::new("s1").unwrap()]
690 );
691 }
692
693 #[test]
694 fn rejects_wrong_data_len() {
695 let mut input = rgb_input();
696 input.data.pop();
697 let error = NdTensor::from_input(input).unwrap_err();
698 assert!(format!("{error}").contains("data bytes"));
699 }
700
701 #[test]
702 fn rejects_observation_count_mismatch() {
703 let mut input = rgb_input();
704 input.observation_ids.pop();
705 let error = NdTensor::from_input(input).unwrap_err();
706 assert!(format!("{error}").contains("observation ids"));
707 }
708
709 #[test]
710 fn rejects_rank_zero_and_over_max() {
711 let mut zero = rgb_input();
712 zero.shape = vec![];
713 assert!(NdTensor::from_input(zero).is_err());
714 let mut huge = rgb_input();
715 huge.shape = vec![3; ND_TENSOR_MAX_RANK + 1];
716 assert!(NdTensor::from_input(huge).is_err());
717 }
718
719 #[test]
720 fn rejects_non_binary_bool_payload() {
721 let input = NdTensorInput {
722 tensor_id: "mask".to_string(),
723 representation_id: RepresentationId::new("mask_image").unwrap(),
724 container: "ndarray".to_string(),
725 dtype: NdTensorDType::Bool,
726 shape: vec![2, 2],
727 observation_ids: vec![
728 ObservationId::new("obs.s1").unwrap(),
729 ObservationId::new("obs.s2").unwrap(),
730 ],
731 sample_ids: None,
732 data: vec![1, 0, 2, 1],
733 row_presence: None,
734 };
735 let error = NdTensor::from_input(input).unwrap_err();
736 assert!(format!("{error}").contains("not 0 or 1"));
737 }
738
739 #[test]
740 fn arena_binds_and_refuses_unbound_or_wrong_source() {
741 let mut arena = NdTensorArena::new(NdTensorStore::from_inputs(vec![rgb_input()]).unwrap());
742 let relations = CoordinatorRelationSet {
743 records: vec![
744 relation("obs.s1", "s1", "cam"),
745 relation("obs.s2", "s2", "cam"),
746 relation("obs.s3", "s3", "cam"),
747 ],
748 };
749 let representation = RepresentationId::new("rgb_image").unwrap();
750 let bindings = arena
751 .bind_data_handle(1, &relations, &representation)
752 .unwrap();
753 assert_eq!(bindings.len(), 1);
754 assert_eq!(bindings[0].source_ids, vec![SourceId::new("cam").unwrap()]);
755
756 let block = arena
758 .project_bound_relations(1, "rgb", &relations, Some(&SourceId::new("cam").unwrap()))
759 .unwrap();
760 assert_eq!(block.shape, vec![3, 2, 2]);
761
762 assert!(arena
764 .project_bound_relations(1, "rgb", &relations, Some(&SourceId::new("nope").unwrap()))
765 .is_err());
766 assert!(arena
767 .project_bound_relations(2, "rgb", &relations, None)
768 .is_err());
769 assert!(arena.release_data_handle(1));
770 assert!(arena.bindings_for_data_handle(1).is_err());
771 }
772
773 #[test]
774 fn rejects_empty_tensor_id() {
775 let mut input = rgb_input();
776 input.tensor_id = " ".to_string();
777 let error = NdTensor::from_input(input).unwrap_err();
778 assert!(format!("{error}").contains("empty tensor id"));
779 }
780
781 #[test]
782 fn rejects_zero_dimension() {
783 let mut input = rgb_input();
784 input.shape = vec![3, 0];
785 input.data = Vec::new();
786 let error = NdTensor::from_input(input).unwrap_err();
787 assert!(format!("{error}").contains("zero dimension"));
788 }
789
790 #[test]
791 fn arena_refuses_unscoped_export_when_a_view_source_is_unbound() {
792 let input = NdTensorInput {
794 tensor_id: "multi".to_string(),
795 representation_id: RepresentationId::new("rgb_image").unwrap(),
796 container: "ndarray".to_string(),
797 dtype: NdTensorDType::U8,
798 shape: vec![1, 2],
799 observation_ids: vec![ObservationId::new("obs.a1").unwrap()],
800 sample_ids: None,
801 data: vec![1, 2],
802 row_presence: None,
803 };
804 let mut arena = NdTensorArena::new(NdTensorStore::from_inputs(vec![input]).unwrap());
805 let relations = CoordinatorRelationSet {
807 records: vec![relation("obs.a1", "a1", "a"), relation("obs.b1", "b1", "b")],
808 };
809 let representation = RepresentationId::new("rgb_image").unwrap();
810 let bindings = arena
811 .bind_data_handle(1, &relations, &representation)
812 .unwrap();
813 assert_eq!(bindings[0].source_ids, vec![SourceId::new("a").unwrap()]);
815
816 assert!(arena
818 .project_bound_relations(1, "multi", &relations, None)
819 .is_err());
820 assert!(arena
822 .project_bound_relations(1, "multi", &relations, Some(&SourceId::new("b").unwrap()))
823 .is_err());
824 let a_only = CoordinatorRelationSet {
826 records: vec![relation("obs.a1", "a1", "a")],
827 };
828 let block = arena
829 .project_bound_relations(1, "multi", &a_only, Some(&SourceId::new("a").unwrap()))
830 .unwrap();
831 assert_eq!(block.shape, vec![1, 2]);
832 }
833
834 #[test]
835 fn manifest_carries_shape_and_fingerprint() {
836 let store = NdTensorStore::from_inputs(vec![rgb_input()]).unwrap();
837 let manifests = store.manifests().unwrap();
838 assert_eq!(manifests.len(), 1);
839 assert_eq!(manifests[0].shape, vec![3, 2, 2]);
840 assert_eq!(manifests[0].data_bytes, 12);
841 assert_eq!(manifests[0].element_bytes, 1);
842 assert_eq!(manifests[0].tensor_fingerprint.len(), 64);
843 }
844
845 fn fp(input: NdTensorInput) -> String {
846 NdTensor::from_input(input).unwrap().fingerprint().unwrap()
847 }
848
849 #[test]
850 fn tensor_fingerprint_is_64_lowercase_hex() {
851 let fingerprint = fp(rgb_input());
852 assert_eq!(fingerprint.len(), 64);
853 assert!(fingerprint
854 .chars()
855 .all(|c| c.is_ascii_hexdigit() && !c.is_ascii_uppercase()));
856 }
857
858 #[test]
859 fn tensor_fingerprint_is_deterministic_across_calls_and_clone() {
860 let tensor = NdTensor::from_input(rgb_input()).unwrap();
861 let once = tensor.fingerprint().unwrap();
862 assert_eq!(once, tensor.fingerprint().unwrap());
863 assert_eq!(once, tensor.clone().fingerprint().unwrap());
864 }
865
866 #[test]
867 fn tensor_fingerprint_changes_when_a_single_data_byte_flips() {
868 let baseline = fp(rgb_input());
869 let mut flipped = rgb_input();
870 flipped.data[0] ^= 0xFF;
871 assert_ne!(baseline, fp(flipped));
872 }
873
874 #[test]
875 fn tensor_fingerprint_changes_when_tensor_id_is_renamed() {
876 let baseline = fp(rgb_input());
877 let mut renamed = rgb_input();
878 renamed.tensor_id = "rgb_renamed".to_string();
879 assert_ne!(baseline, fp(renamed));
880 }
881
882 #[test]
883 fn tensor_fingerprint_changes_when_observation_ids_are_reordered() {
884 let baseline = fp(rgb_input());
887 let mut reordered = rgb_input();
888 reordered.observation_ids.swap(0, 2);
889 let (head, tail) = reordered.data.split_at_mut(8);
891 head[0..4].swap_with_slice(&mut tail[0..4]);
892 assert_ne!(baseline, fp(reordered));
893 }
894
895 #[test]
896 fn tensor_fingerprint_distinguishes_transposed_shapes_with_identical_bytes() {
897 let base = NdTensorInput {
900 tensor_id: "t".to_string(),
901 representation_id: RepresentationId::new("rgb_image").unwrap(),
902 container: "ndarray".to_string(),
903 dtype: NdTensorDType::U8,
904 shape: vec![3, 2, 2],
905 observation_ids: vec![
906 ObservationId::new("obs.s1").unwrap(),
907 ObservationId::new("obs.s2").unwrap(),
908 ObservationId::new("obs.s3").unwrap(),
909 ],
910 sample_ids: None,
911 data: (0u8..12).collect(),
912 row_presence: None,
913 };
914 let mut reshaped = base.clone();
915 reshaped.shape = vec![3, 4];
916 assert_ne!(fp(base), fp(reshaped));
917 }
918
919 #[test]
920 fn tensor_fingerprint_distinguishes_dtype_with_identical_bytes() {
921 let as_i32 = NdTensorInput {
924 tensor_id: "t".to_string(),
925 representation_id: RepresentationId::new("rgb_image").unwrap(),
926 container: "ndarray".to_string(),
927 dtype: NdTensorDType::I32,
928 shape: vec![1, 1],
929 observation_ids: vec![ObservationId::new("obs.s1").unwrap()],
930 sample_ids: None,
931 data: vec![1, 2, 3, 4],
932 row_presence: None,
933 };
934 let mut as_u8 = as_i32.clone();
935 as_u8.dtype = NdTensorDType::U8;
936 as_u8.shape = vec![1, 4];
937 assert_ne!(fp(as_i32), fp(as_u8));
938 }
939
940 #[test]
941 fn tensor_fingerprint_hashes_data_bytes_verbatim_in_declared_order() {
942 let f32_one_le: [u8; 4] = 1.0f32.to_le_bytes(); let le = NdTensorInput {
949 tensor_id: "t".to_string(),
950 representation_id: RepresentationId::new("hyperspectral").unwrap(),
951 container: "ndarray".to_string(),
952 dtype: NdTensorDType::F32,
953 shape: vec![1, 1],
954 observation_ids: vec![ObservationId::new("obs.s1").unwrap()],
955 sample_ids: None,
956 data: f32_one_le.to_vec(),
957 row_presence: None,
958 };
959 let mut be = le.clone();
960 let mut reversed = f32_one_le;
961 reversed.reverse(); be.data = reversed.to_vec();
963 assert_eq!(fp(le.clone()), fp(le.clone()));
965 assert_ne!(fp(le), fp(be));
966 }
967
968 #[test]
969 fn tensor_fingerprint_distinguishes_row_presence_states() {
970 let baseline = fp(rgb_input());
973 let mut with_presence = rgb_input();
974 with_presence.row_presence = Some(vec![true, true, true]);
975 let present_fp = fp(with_presence);
976 assert_ne!(baseline, present_fp);
977
978 let mut one_absent = rgb_input();
979 one_absent.row_presence = Some(vec![true, false, true]);
980 assert_ne!(present_fp, fp(one_absent));
981 }
982
983 #[test]
984 #[ignore = "perf sanity probe; run with --release --ignored --nocapture"]
985 fn tensor_fingerprint_large_payload_under_500ms() {
986 let rows = 3021usize;
993 let cols = 1050usize;
994 let element_size = NdTensorDType::F32.element_size();
995 let data = vec![0x3Cu8; rows * cols * element_size];
996 let input = NdTensorInput {
997 tensor_id: "big".to_string(),
998 representation_id: RepresentationId::new("hyperspectral").unwrap(),
999 container: "ndarray".to_string(),
1000 dtype: NdTensorDType::F32,
1001 shape: vec![rows, cols],
1002 observation_ids: (0..rows)
1003 .map(|r| ObservationId::new(format!("obs.{r}")).unwrap())
1004 .collect(),
1005 sample_ids: None,
1006 data,
1007 row_presence: None,
1008 };
1009 let tensor = NdTensor::from_input(input).unwrap();
1010 let start = std::time::Instant::now();
1011 let fingerprint = tensor.fingerprint().unwrap();
1012 let elapsed = start.elapsed();
1013 println!(
1014 "nd tensor fingerprint({rows}x{cols} f32) = {:.3} ms (fp={fingerprint})",
1015 elapsed.as_secs_f64() * 1e3
1016 );
1017 assert_eq!(fingerprint.len(), 64);
1018 if !cfg!(debug_assertions) {
1019 assert!(
1020 elapsed.as_millis() < 500,
1021 "tensor fingerprint took {} ms (>= 500 ms budget)",
1022 elapsed.as_millis()
1023 );
1024 }
1025 }
1026}