Skip to main content

dag_ml_data_core/
alignment.rs

1use std::collections::{BTreeMap, BTreeSet};
2
3use serde::{Deserialize, Serialize};
4
5use crate::error::{DataError, Result};
6use crate::ids::{SampleId, SourceId};
7use crate::model::PresenceMask;
8
9#[derive(Clone, Copy, Debug, Default, Eq, PartialEq, Serialize, Deserialize)]
10#[serde(rename_all = "snake_case")]
11pub enum AlignmentMode {
12    #[default]
13    Inner,
14    Left,
15    Outer,
16}
17
18#[derive(Clone, Debug, Eq, PartialEq, Serialize, Deserialize)]
19pub struct AlignmentPolicy {
20    pub mode: AlignmentMode,
21}
22
23impl Default for AlignmentPolicy {
24    fn default() -> Self {
25        Self {
26            mode: AlignmentMode::Inner,
27        }
28    }
29}
30
31#[derive(Clone, Debug, Eq, PartialEq, Serialize, Deserialize)]
32pub struct SourceSampleSet {
33    pub source_id: SourceId,
34    pub sample_ids: Vec<SampleId>,
35}
36
37#[derive(Clone, Debug, Eq, PartialEq, Serialize, Deserialize)]
38pub struct SampleAlignmentPlan {
39    pub mode: AlignmentMode,
40    pub sample_ids: Vec<SampleId>,
41    pub masks: Vec<PresenceMask>,
42}
43
44impl SampleAlignmentPlan {
45    pub fn validate(&self) -> Result<()> {
46        if self.sample_ids.is_empty() {
47            return Err(DataError::Validation(
48                "alignment plan contains no samples".to_string(),
49            ));
50        }
51        let mut samples = BTreeSet::new();
52        for sample_id in &self.sample_ids {
53            if !samples.insert(sample_id) {
54                return Err(DataError::Validation(format!(
55                    "alignment plan contains duplicate sample `{sample_id}`"
56                )));
57            }
58        }
59        if self.masks.is_empty() {
60            return Err(DataError::Validation(
61                "alignment plan contains no presence masks".to_string(),
62            ));
63        }
64        let mut sources = BTreeSet::new();
65        for mask in &self.masks {
66            mask.validate()?;
67            if mask.sample_ids != self.sample_ids {
68                return Err(DataError::Validation(format!(
69                    "presence mask for `{}` does not use alignment sample order",
70                    mask.source_id
71                )));
72            }
73            if !sources.insert(&mask.source_id) {
74                return Err(DataError::Validation(format!(
75                    "alignment plan contains duplicate source mask `{}`",
76                    mask.source_id
77                )));
78            }
79        }
80        for (idx, sample_id) in self.sample_ids.iter().enumerate() {
81            if !self.masks.iter().any(|mask| mask.present[idx]) {
82                return Err(DataError::Validation(format!(
83                    "alignment sample `{sample_id}` is absent from every fused source"
84                )));
85            }
86            let valid = match self.mode {
87                AlignmentMode::Inner => self.masks.iter().all(|mask| mask.present[idx]),
88                AlignmentMode::Left => self.masks[0].present[idx],
89                AlignmentMode::Outer => self.masks.iter().any(|mask| mask.present[idx]),
90            };
91            if !valid {
92                return Err(DataError::Validation(format!(
93                    "alignment presence for sample `{sample_id}` violates {:?} mode",
94                    self.mode
95                )));
96            }
97        }
98        Ok(())
99    }
100}
101
102pub fn build_sample_alignment_plan(
103    sources: &[SourceSampleSet],
104    policy: &AlignmentPolicy,
105) -> Result<SampleAlignmentPlan> {
106    if sources.is_empty() {
107        return Err(DataError::Validation(
108            "alignment requires at least one source".to_string(),
109        ));
110    }
111    let mut source_ids = BTreeSet::new();
112    let mut samples_by_source = Vec::with_capacity(sources.len());
113    for source in sources {
114        if !source_ids.insert(&source.source_id) {
115            return Err(DataError::Validation(format!(
116                "alignment contains duplicate source `{}`",
117                source.source_id
118            )));
119        }
120        if source.sample_ids.is_empty() {
121            return Err(DataError::Validation(format!(
122                "alignment source `{}` contains no samples",
123                source.source_id
124            )));
125        }
126        let mut samples = BTreeSet::new();
127        for sample_id in &source.sample_ids {
128            if !samples.insert(sample_id) {
129                return Err(DataError::Validation(format!(
130                    "alignment source `{}` contains duplicate sample `{sample_id}`",
131                    source.source_id
132                )));
133            }
134        }
135        samples_by_source.push(samples);
136    }
137
138    let sample_ids = match policy.mode {
139        AlignmentMode::Inner => sources[0]
140            .sample_ids
141            .iter()
142            .filter(|sample_id| {
143                samples_by_source
144                    .iter()
145                    .all(|samples| samples.contains(*sample_id))
146            })
147            .cloned()
148            .collect::<Vec<_>>(),
149        AlignmentMode::Left => sources[0].sample_ids.clone(),
150        AlignmentMode::Outer => {
151            let mut seen = BTreeSet::new();
152            let mut ordered = Vec::new();
153            for source in sources {
154                for sample_id in &source.sample_ids {
155                    if seen.insert(sample_id) {
156                        ordered.push(sample_id.clone());
157                    }
158                }
159            }
160            ordered
161        }
162    };
163
164    if sample_ids.is_empty() {
165        return Err(DataError::Validation(format!(
166            "{:?} alignment produced no shared samples",
167            policy.mode
168        )));
169    }
170
171    let masks = sources
172        .iter()
173        .zip(samples_by_source.iter())
174        .map(|(source, samples)| PresenceMask {
175            sample_ids: sample_ids.clone(),
176            source_id: source.source_id.clone(),
177            present: sample_ids
178                .iter()
179                .map(|sample_id| samples.contains(sample_id))
180                .collect(),
181        })
182        .collect::<Vec<_>>();
183    let plan = SampleAlignmentPlan {
184        mode: policy.mode,
185        sample_ids,
186        masks,
187    };
188    plan.validate()?;
189    Ok(plan)
190}
191
192pub fn alignment_mode_from_fusion(value: Option<&serde_json::Value>) -> Result<AlignmentMode> {
193    let Some(value) = value else {
194        return Ok(AlignmentMode::Inner);
195    };
196    let Some(mode) = value.get("alignment") else {
197        return Ok(AlignmentMode::Inner);
198    };
199    serde_json::from_value(mode.clone())
200        .map_err(|error| DataError::Validation(format!("invalid fusion alignment policy: {error}")))
201}
202
203pub fn alignment_metadata(
204    inputs: Vec<String>,
205    output: String,
206    mode: AlignmentMode,
207) -> Result<BTreeMap<String, serde_json::Value>> {
208    Ok(BTreeMap::from([
209        (
210            "inputs".to_string(),
211            serde_json::Value::Array(inputs.into_iter().map(serde_json::Value::String).collect()),
212        ),
213        ("output".to_string(), serde_json::Value::String(output)),
214        ("alignment".to_string(), serde_json::to_value(mode)?),
215    ]))
216}
217
218#[cfg(test)]
219mod tests {
220    use super::*;
221
222    fn source(source_id: &str, sample_ids: &[&str]) -> SourceSampleSet {
223        SourceSampleSet {
224            source_id: SourceId::new(source_id).unwrap(),
225            sample_ids: sample_ids
226                .iter()
227                .map(|sample_id| SampleId::new(*sample_id).unwrap())
228                .collect(),
229        }
230    }
231
232    fn samples(plan: &SampleAlignmentPlan) -> Vec<String> {
233        plan.sample_ids.iter().map(ToString::to_string).collect()
234    }
235
236    #[test]
237    fn inner_alignment_keeps_first_source_order_for_shared_samples() {
238        let plan = build_sample_alignment_plan(
239            &[
240                source("nir", &["S002", "S001", "S003"]),
241                source("chem", &["S001", "S003"]),
242            ],
243            &AlignmentPolicy {
244                mode: AlignmentMode::Inner,
245            },
246        )
247        .unwrap();
248
249        assert_eq!(samples(&plan), vec!["S001", "S003"]);
250        assert_eq!(plan.masks[0].present, vec![true, true]);
251        assert_eq!(plan.masks[1].present, vec![true, true]);
252    }
253
254    #[test]
255    fn left_alignment_preserves_left_samples_and_marks_missing_sources() {
256        let plan = build_sample_alignment_plan(
257            &[
258                source("nir", &["S002", "S001", "S003"]),
259                source("chem", &["S001", "S003"]),
260            ],
261            &AlignmentPolicy {
262                mode: AlignmentMode::Left,
263            },
264        )
265        .unwrap();
266
267        assert_eq!(samples(&plan), vec!["S002", "S001", "S003"]);
268        assert_eq!(plan.masks[0].present, vec![true, true, true]);
269        assert_eq!(plan.masks[1].present, vec![false, true, true]);
270    }
271
272    #[test]
273    fn outer_alignment_appends_new_samples_in_source_order() {
274        let plan = build_sample_alignment_plan(
275            &[
276                source("nir", &["S002", "S001"]),
277                source("chem", &["S003", "S001"]),
278                source("image", &["S004", "S002"]),
279            ],
280            &AlignmentPolicy {
281                mode: AlignmentMode::Outer,
282            },
283        )
284        .unwrap();
285
286        assert_eq!(samples(&plan), vec!["S002", "S001", "S003", "S004"]);
287        assert_eq!(plan.masks[0].present, vec![true, true, false, false]);
288        assert_eq!(plan.masks[1].present, vec![false, true, true, false]);
289        assert_eq!(plan.masks[2].present, vec![true, false, false, true]);
290    }
291
292    #[test]
293    fn alignment_refuses_duplicate_samples_per_source() {
294        let err = build_sample_alignment_plan(
295            &[source("nir", &["S001", "S001"])],
296            &AlignmentPolicy::default(),
297        )
298        .unwrap_err();
299
300        assert!(err.to_string().contains("duplicate sample"));
301    }
302}