Skip to main content

dag_ml_data_core/
fingerprint.rs

1use sha2::{Digest, Sha256};
2
3use crate::error::Result;
4use crate::model::DatasetSchema;
5use crate::plan::DataPlan;
6use crate::relation::{FoldSet, SampleRelationTable};
7
8pub fn schema_fingerprint(schema: &DatasetSchema) -> Result<String> {
9    let mut canonical = schema.clone();
10    canonical.validate()?;
11    canonical.sample_ids.sort();
12    canonical
13        .sources
14        .sort_by(|left, right| left.id.cmp(&right.id));
15    canonical
16        .groups
17        .sort_by(|left, right| left.id.cmp(&right.id));
18    canonical
19        .folds
20        .sort_by(|left, right| left.id.cmp(&right.id));
21
22    let json = canonical_typed_json(&canonical)?;
23    let digest = Sha256::digest(json);
24    Ok(to_hex(&digest))
25}
26
27pub fn data_plan_fingerprint(plan: &DataPlan) -> Result<String> {
28    plan.validate()?;
29    let json = canonical_typed_json(plan)?;
30    let digest = Sha256::digest(json);
31    Ok(to_hex(&digest))
32}
33
34pub fn sample_relation_fingerprint(relations: &SampleRelationTable) -> Result<String> {
35    let mut canonical = relations.clone();
36    canonical.validate()?;
37    canonical.rows.sort_by(|left, right| {
38        left.observation_id
39            .cmp(&right.observation_id)
40            .then_with(|| left.sample_id.cmp(&right.sample_id))
41            .then_with(|| left.source_id.cmp(&right.source_id))
42    });
43
44    let json = canonical_typed_json(&canonical)?;
45    let digest = Sha256::digest(json);
46    Ok(to_hex(&digest))
47}
48
49pub fn fold_set_fingerprint(fold_set: &FoldSet) -> Result<String> {
50    let mut canonical = fold_set.clone();
51    canonical.validate()?;
52    canonical.sample_ids.sort();
53    canonical
54        .folds
55        .sort_by(|left, right| left.fold_id.cmp(&right.fold_id));
56    for fold in &mut canonical.folds {
57        fold.train_sample_ids.sort();
58        fold.validation_sample_ids.sort();
59    }
60
61    let mut value = serde_json::to_value(&canonical)?;
62    remove_empty_fold_set_maps(&mut value);
63    // Fold fingerprints historically serialize a Value (sorted object keys),
64    // unlike schema/plan/relation fingerprints which preserve struct order.
65    value.sort_all_objects();
66    let json = serde_json::to_vec(&value)?;
67    let digest = Sha256::digest(json);
68    Ok(to_hex(&digest))
69}
70
71/// Normalize nested JSON maps without changing the published order of typed
72/// struct fields. This is invariant under serde_json's additive preserve_order
73/// feature, including when enabled only by a downstream dependency.
74fn canonical_typed_json<T: serde::Serialize + serde::de::DeserializeOwned>(
75    value: &T,
76) -> Result<Vec<u8>> {
77    let mut json = serde_json::to_value(value)?;
78    json.sort_all_objects();
79    let canonical: T = serde_json::from_value(json)?;
80    Ok(serde_json::to_vec(&canonical)?)
81}
82
83pub(crate) fn typed_fingerprint<T: serde::Serialize + serde::de::DeserializeOwned>(
84    value: &T,
85) -> Result<String> {
86    Ok(to_hex(&Sha256::digest(canonical_typed_json(value)?)))
87}
88
89fn remove_empty_fold_set_maps(value: &mut serde_json::Value) {
90    let Some(object) = value.as_object_mut() else {
91        return;
92    };
93    if object
94        .get("sample_groups")
95        .and_then(serde_json::Value::as_object)
96        .is_some_and(serde_json::Map::is_empty)
97    {
98        object.remove("sample_groups");
99    }
100    let Some(folds) = object
101        .get_mut("folds")
102        .and_then(serde_json::Value::as_array_mut)
103    else {
104        return;
105    };
106    for fold in folds {
107        let Some(fold_object) = fold.as_object_mut() else {
108            continue;
109        };
110        if fold_object
111            .get("metadata")
112            .and_then(serde_json::Value::as_object)
113            .is_some_and(serde_json::Map::is_empty)
114        {
115            fold_object.remove("metadata");
116        }
117    }
118}
119
120fn to_hex(bytes: &[u8]) -> String {
121    let mut out = String::with_capacity(bytes.len() * 2);
122    for byte in bytes {
123        use std::fmt::Write;
124        write!(&mut out, "{byte:02x}").expect("writing to string cannot fail");
125    }
126    out
127}
128
129#[cfg(test)]
130mod tests {
131    use std::collections::BTreeMap;
132
133    use crate::ids::{GroupId, RepresentationId, SampleId, SourceId, TypeId};
134    use crate::model::{
135        AxisKind, AxisSpec, CoordinateDType, CoordinateSpec, CoordinateValues, DatasetSchema,
136        FoldSpec, GroupKind, GroupSpec, RepresentationSpec, SourceDescriptor, SourceGranularity,
137    };
138    use crate::plan::DataPlan;
139    use crate::relation::{FoldAssignment, FoldSet};
140
141    use super::{
142        data_plan_fingerprint, fold_set_fingerprint, sample_relation_fingerprint,
143        schema_fingerprint,
144    };
145
146    const SHARED_FOLD_SET_FINGERPRINT: &str =
147        "54d3185d6c628ef0df848828a8d8ae650222a283a78bbd3ab3bc2256f222c05c";
148
149    #[test]
150    fn nested_json_order_does_not_change_published_typed_fingerprints() {
151        let ascending: serde_json::Value =
152            serde_json::from_str(r#"{"a":{"a":1,"z":2},"z":[{"a":3,"z":4}]}"#).unwrap();
153        let descending: serde_json::Value =
154            serde_json::from_str(r#"{"z":[{"z":4,"a":3}],"a":{"z":2,"a":1}}"#).unwrap();
155        let mut schema: DatasetSchema =
156            serde_json::from_str(include_str!("../../../examples/minimal_schema.json")).unwrap();
157        assert_eq!(
158            schema_fingerprint(&schema).unwrap(),
159            "e1b5174cbd2b6282d9d4017ba3b1f8dc2ad829b020e4e3df3bc9689496164ba0"
160        );
161        schema.sources[0]
162            .schema
163            .insert("nested".into(), ascending.clone());
164        let left = schema_fingerprint(&schema).unwrap();
165        schema.sources[0]
166            .schema
167            .insert("nested".into(), descending.clone());
168        assert_eq!(left, schema_fingerprint(&schema).unwrap());
169
170        let mut plan: DataPlan = serde_json::from_str(include_str!(
171            "../../../examples/fixtures/oof_campaign/expected_data_plan_nir_to_tabular.json"
172        ))
173        .unwrap();
174        let mut expected: serde_json::Value = serde_json::from_str(include_str!(
175            "../../../examples/fixtures/oof_campaign/coordinator_data_plan_envelope_nir.json"
176        ))
177        .unwrap();
178        assert_eq!(
179            data_plan_fingerprint(&plan).unwrap(),
180            expected["plan_fingerprint"].take().as_str().unwrap()
181        );
182        plan.steps[0]
183            .metadata
184            .insert("nested".into(), ascending.clone());
185        let left = data_plan_fingerprint(&plan).unwrap();
186        plan.steps[0]
187            .metadata
188            .insert("nested".into(), descending.clone());
189        assert_eq!(left, data_plan_fingerprint(&plan).unwrap());
190
191        let mut relations: crate::SampleRelationTable = serde_json::from_str(include_str!(
192            "../../../examples/fixtures/oof_campaign/sample_relations_grouped_augmented.json"
193        ))
194        .unwrap();
195        relations.rows[0]
196            .metadata
197            .insert("nested".into(), ascending);
198        let left = sample_relation_fingerprint(&relations).unwrap();
199        relations.rows[0]
200            .metadata
201            .insert("nested".into(), descending);
202        assert_eq!(left, sample_relation_fingerprint(&relations).unwrap());
203    }
204
205    fn representation(id: &str) -> RepresentationSpec {
206        RepresentationSpec {
207            id: RepresentationId::new(id).unwrap(),
208            type_id: TypeId::new("table").unwrap(),
209            rank: Some(2),
210            axes: vec![
211                AxisSpec {
212                    name: "sample".to_string(),
213                    kind: AxisKind::Sample,
214                    unit: None,
215                    size: Some(2),
216                    variable: false,
217                    coordinate: None,
218                },
219                AxisSpec {
220                    name: "feature".to_string(),
221                    kind: AxisKind::Feature,
222                    unit: None,
223                    size: Some(1),
224                    variable: false,
225                    coordinate: None,
226                },
227            ],
228            container: "dataframe".to_string(),
229            dtype: Some("float32".to_string()),
230            sparse: false,
231            ragged: false,
232            signal_type: None,
233        }
234    }
235
236    fn source(id: &str) -> SourceDescriptor {
237        SourceDescriptor {
238            id: SourceId::new(id).unwrap(),
239            name: id.to_string(),
240            type_id: TypeId::new("table").unwrap(),
241            modality: "metadata".to_string(),
242            native_representation: representation("tabular"),
243            sample_key: "sample_id".to_string(),
244            granularity: SourceGranularity::PerSample,
245            schema: BTreeMap::new(),
246            tags: BTreeMap::new(),
247            shape_contract: None,
248        }
249    }
250
251    #[test]
252    fn fingerprint_is_independent_of_source_order() {
253        let mut left = DatasetSchema {
254            dataset_id: "d".to_string(),
255            sample_ids: vec![SampleId::new("s2").unwrap(), SampleId::new("s1").unwrap()],
256            sources: vec![source("b"), source("a")],
257            targets: BTreeMap::new(),
258            metadata: BTreeMap::new(),
259            metadata_schema: None,
260            groups: vec![
261                GroupSpec {
262                    id: GroupId::new("g.b").unwrap(),
263                    kind: GroupKind::Batch,
264                    column: "batch_b".to_string(),
265                    source_id: Some(SourceId::new("b").unwrap()),
266                    strict: false,
267                    metadata: BTreeMap::new(),
268                },
269                GroupSpec {
270                    id: GroupId::new("g.a").unwrap(),
271                    kind: GroupKind::RepetitionGroup,
272                    column: "sample_id".to_string(),
273                    source_id: Some(SourceId::new("a").unwrap()),
274                    strict: true,
275                    metadata: BTreeMap::new(),
276                },
277            ],
278            folds: vec![
279                FoldSpec {
280                    id: "fold.b".to_string(),
281                    group_id: Some(GroupId::new("g.b").unwrap()),
282                    split_column: Some("fold_b".to_string()),
283                    metadata: BTreeMap::new(),
284                },
285                FoldSpec {
286                    id: "fold.a".to_string(),
287                    group_id: Some(GroupId::new("g.a").unwrap()),
288                    split_column: Some("fold_a".to_string()),
289                    metadata: BTreeMap::new(),
290                },
291            ],
292        };
293        let mut right = left.clone();
294        right.sources.reverse();
295        right.sample_ids.reverse();
296        right.groups.reverse();
297        right.folds.reverse();
298
299        assert_eq!(
300            schema_fingerprint(&left).unwrap(),
301            schema_fingerprint(&right).unwrap()
302        );
303
304        left.dataset_id = "different".to_string();
305        assert_ne!(
306            schema_fingerprint(&left).unwrap(),
307            schema_fingerprint(&right).unwrap()
308        );
309    }
310
311    #[test]
312    fn data_plan_fingerprint_is_stable() {
313        let plan: DataPlan = serde_json::from_str(include_str!(
314            "../../../examples/fixtures/oof_campaign/expected_data_plan_nir_to_tabular.json"
315        ))
316        .unwrap();
317
318        assert_eq!(
319            data_plan_fingerprint(&plan).unwrap(),
320            data_plan_fingerprint(&plan).unwrap()
321        );
322    }
323
324    #[test]
325    fn sample_relation_fingerprint_is_stable() {
326        let relations: crate::relation::SampleRelationTable = serde_json::from_str(include_str!(
327            "../../../examples/fixtures/oof_campaign/sample_relations_grouped_augmented.json"
328        ))
329        .unwrap();
330
331        assert_eq!(
332            sample_relation_fingerprint(&relations).unwrap(),
333            sample_relation_fingerprint(&relations).unwrap()
334        );
335    }
336
337    #[test]
338    fn empty_tags_keep_relation_fingerprint_byte_identical() {
339        // `tags` is skip-serialized when empty, so attaching an empty Vec to a
340        // row must not change the source-relation fingerprint; a non-empty value
341        // does change it (it changes which samples a `by_tag` view selects).
342        let base: crate::relation::SampleRelationTable = serde_json::from_str(include_str!(
343            "../../../examples/fixtures/oof_campaign/sample_relations_grouped_augmented.json"
344        ))
345        .unwrap();
346
347        let mut explicit_empty = base.clone();
348        explicit_empty.rows[0].tags = Vec::new();
349        assert_eq!(
350            sample_relation_fingerprint(&base).unwrap(),
351            sample_relation_fingerprint(&explicit_empty).unwrap()
352        );
353
354        let mut with_tags = base.clone();
355        with_tags.rows[0].tags = vec!["clean".to_string()];
356        assert_ne!(
357            sample_relation_fingerprint(&base).unwrap(),
358            sample_relation_fingerprint(&with_tags).unwrap()
359        );
360    }
361
362    #[test]
363    fn fold_set_fingerprint_is_independent_of_ordering() {
364        let mut left = FoldSet {
365            id: "cv.partition".to_string(),
366            sample_ids: vec![
367                SampleId::new("s3").unwrap(),
368                SampleId::new("s2").unwrap(),
369                SampleId::new("s1").unwrap(),
370            ],
371            folds: vec![
372                FoldAssignment {
373                    fold_id: "fold1".to_string(),
374                    train_sample_ids: vec![
375                        SampleId::new("s2").unwrap(),
376                        SampleId::new("s1").unwrap(),
377                    ],
378                    validation_sample_ids: vec![SampleId::new("s3").unwrap()],
379                    metadata: BTreeMap::new(),
380                },
381                FoldAssignment {
382                    fold_id: "fold0".to_string(),
383                    train_sample_ids: vec![SampleId::new("s3").unwrap()],
384                    validation_sample_ids: vec![
385                        SampleId::new("s2").unwrap(),
386                        SampleId::new("s1").unwrap(),
387                    ],
388                    metadata: BTreeMap::new(),
389                },
390            ],
391            sample_groups: BTreeMap::new(),
392        };
393        let mut right = left.clone();
394        right.sample_ids.reverse();
395        right.folds.reverse();
396        for fold in &mut right.folds {
397            fold.train_sample_ids.reverse();
398            fold.validation_sample_ids.reverse();
399        }
400
401        assert_eq!(
402            fold_set_fingerprint(&left).unwrap(),
403            fold_set_fingerprint(&right).unwrap()
404        );
405
406        left.id = "cv.partition.changed".to_string();
407        assert_ne!(
408            fold_set_fingerprint(&left).unwrap(),
409            fold_set_fingerprint(&right).unwrap()
410        );
411    }
412
413    fn coordinate_schema(coordinate: Option<CoordinateSpec>) -> DatasetSchema {
414        let mut repr = representation("tabular");
415        // The feature axis (size 1) carries the coordinate.
416        repr.axes[1].coordinate = coordinate;
417        let mut descriptor = source("a");
418        descriptor.native_representation = repr;
419        DatasetSchema {
420            dataset_id: "coords".to_string(),
421            sample_ids: vec![SampleId::new("s1").unwrap()],
422            sources: vec![descriptor],
423            targets: BTreeMap::new(),
424            metadata: BTreeMap::new(),
425            metadata_schema: None,
426            groups: Vec::new(),
427            folds: Vec::new(),
428        }
429    }
430
431    #[test]
432    fn schema_fingerprint_reflects_axis_coordinates() {
433        let explicit = CoordinateSpec {
434            dtype: CoordinateDType::Categorical,
435            ordered: false,
436            values: CoordinateValues::Explicit {
437                values: vec![serde_json::Value::from("R")],
438            },
439        };
440        let grid = CoordinateSpec {
441            dtype: CoordinateDType::Numeric,
442            ordered: true,
443            values: CoordinateValues::RegularGrid {
444                start: 400.0,
445                step: 2.0,
446            },
447        };
448
449        let bare = schema_fingerprint(&coordinate_schema(None)).unwrap();
450        let with_explicit = schema_fingerprint(&coordinate_schema(Some(explicit.clone()))).unwrap();
451        let with_grid = schema_fingerprint(&coordinate_schema(Some(grid))).unwrap();
452
453        // Stable for a fixed coordinate spec.
454        assert_eq!(
455            with_explicit,
456            schema_fingerprint(&coordinate_schema(Some(explicit))).unwrap()
457        );
458        // Coordinates participate in the fingerprint and the two shapes differ.
459        assert_ne!(bare, with_explicit);
460        assert_ne!(bare, with_grid);
461        assert_ne!(with_explicit, with_grid);
462    }
463
464    #[test]
465    fn shared_fold_set_fixture_fingerprint_is_locked() {
466        let fixture = include_str!("../../../examples/fixtures/shared/fold_set_cv_partition.json");
467        let fold_set = serde_json::from_str::<FoldSet>(fixture).unwrap();
468
469        assert_eq!(
470            fold_set_fingerprint(&fold_set).unwrap(),
471            SHARED_FOLD_SET_FINGERPRINT
472        );
473    }
474}