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}