Skip to main content

dag_ml_data_core/
planner.rs

1use std::collections::BTreeMap;
2
3use serde::{Deserialize, Serialize};
4
5use crate::adapter::{AdapterPath, AdapterRegistry, PlanningPolicy};
6use crate::alignment::{alignment_metadata, alignment_mode_from_fusion};
7use crate::error::{DataError, Result};
8use crate::ids::{RepresentationId, SourceId};
9use crate::model::{DatasetSchema, SourceDescriptor};
10use crate::plan::{DataPlan, DataPlanStep, DataPlanStepKind, FitScope, PlanIssue};
11use crate::ModelInputSpec;
12
13#[derive(Clone, Debug, Eq, PartialEq, Serialize, Deserialize)]
14pub struct DataPlanRequest {
15    pub id: String,
16    #[serde(default)]
17    pub source_ids: Option<Vec<SourceId>>,
18    #[serde(default)]
19    pub planning_policy: PlanningPolicy,
20}
21
22impl DataPlanRequest {
23    pub fn new(id: impl Into<String>) -> Self {
24        Self {
25            id: id.into(),
26            source_ids: None,
27            planning_policy: PlanningPolicy::default(),
28        }
29    }
30}
31
32struct ResolvedSource<'a> {
33    source: &'a SourceDescriptor,
34    target_representation: RepresentationId,
35    path: AdapterPath,
36    issues: Vec<PlanIssue>,
37}
38
39pub fn plan_model_input(
40    schema: &DatasetSchema,
41    model_input: &ModelInputSpec,
42    adapters: &AdapterRegistry,
43    request: &DataPlanRequest,
44) -> Result<DataPlan> {
45    schema.validate()?;
46    model_input.validate()?;
47    validate_request(request)?;
48
49    let candidate_sources = candidate_sources(schema, request)?;
50    let mut steps = Vec::new();
51    let mut issues = Vec::new();
52    let mut output_representation = None;
53    let mut adapt_idx = 0usize;
54    let mut align_idx = 0usize;
55
56    for port in &model_input.ports {
57        let mut resolved = candidate_sources
58            .iter()
59            .filter_map(|source| resolve_source(source, port, adapters, &request.planning_policy))
60            .collect::<Vec<_>>();
61
62        if port.multi_source && resolved.len() > 1 {
63            // Negotiate one accepted type/representation for every source
64            // eligible for this port, rather than resolving each independently.
65            let eligible = resolved.iter().map(|item| item.source).collect::<Vec<_>>();
66            let mut common = None;
67            'targets: for target_type in &port.accepted_types {
68                for representation in &port.accepted_representations {
69                    let mut target_port = port.clone();
70                    target_port.accepted_types = vec![target_type.clone()];
71                    target_port.accepted_representations = vec![representation.clone()];
72                    let plans = eligible
73                        .iter()
74                        .map(|source| {
75                            resolve_source(source, &target_port, adapters, &request.planning_policy)
76                        })
77                        .collect::<Option<Vec<_>>>();
78                    if let Some(plans) = plans {
79                        common = Some(plans);
80                        break 'targets;
81                    }
82                }
83            }
84            resolved = common.ok_or_else(|| {
85                DataError::Validation(format!(
86                    "model input port `{}` has no common accepted representation for its sources",
87                    port.name
88                ))
89            })?;
90        }
91
92        if resolved.is_empty() {
93            if port.optional {
94                continue;
95            }
96            return Err(DataError::Validation(format!(
97                "no source can satisfy required model input port `{}`",
98                port.name
99            )));
100        }
101        if resolved.len() > 1 && !port.multi_source {
102            return Err(DataError::Validation(format!(
103                "model input port `{}` accepts one source but {} sources resolved",
104                port.name,
105                resolved.len()
106            )));
107        }
108
109        let port_output_representation = resolved[0].target_representation.clone();
110        let mut join_inputs = Vec::new();
111        for source_plan in resolved {
112            issues.extend(source_plan.issues);
113            let source_output = source_output_id(&source_plan.source.id);
114            steps.push(DataPlanStep {
115                kind: DataPlanStepKind::Materialize,
116                source_id: Some(source_plan.source.id.clone()),
117                adapter_id: None,
118                input_representation: None,
119                output_representation: Some(source_plan.source.native_representation.id.clone()),
120                fit_scope: FitScope::Stateless,
121                requires_user_choice: false,
122                metadata: BTreeMap::from([(
123                    "output".to_string(),
124                    serde_json::Value::String(source_output.clone()),
125                )]),
126            });
127
128            let mut current_output = source_output;
129            let mut current_representation = source_plan.source.native_representation.id.clone();
130            for adapter in source_plan.path.adapters {
131                let step_output = format!("step:adapt:{adapt_idx}");
132                steps.push(DataPlanStep {
133                    kind: DataPlanStepKind::Adapt,
134                    source_id: Some(source_plan.source.id.clone()),
135                    adapter_id: Some(adapter.id),
136                    input_representation: Some(current_representation),
137                    output_representation: Some(adapter.output_representation.clone()),
138                    fit_scope: adapter.fit_scope,
139                    requires_user_choice: false,
140                    metadata: BTreeMap::from([
141                        (
142                            "adapter_version".to_string(),
143                            serde_json::Value::String(adapter.version),
144                        ),
145                        (
146                            "adapter_params".to_string(),
147                            serde_json::to_value(&adapter.params)?,
148                        ),
149                        (
150                            "input".to_string(),
151                            serde_json::Value::String(current_output),
152                        ),
153                        (
154                            "output".to_string(),
155                            serde_json::Value::String(step_output.clone()),
156                        ),
157                    ]),
158                });
159                current_output = step_output;
160                current_representation = adapter.output_representation;
161                adapt_idx += 1;
162            }
163            join_inputs.push(current_output);
164        }
165
166        let join_inputs = if join_inputs.len() > 1 {
167            let align_output = format!("step:align:{align_idx}");
168            let align_input_representation = port_output_representation.clone();
169            let alignment_mode = alignment_mode_from_fusion(model_input.default_fusion.as_ref())?;
170            steps.push(DataPlanStep {
171                kind: DataPlanStepKind::Align,
172                source_id: None,
173                adapter_id: None,
174                input_representation: Some(align_input_representation),
175                output_representation: Some(port_output_representation.clone()),
176                fit_scope: FitScope::Stateless,
177                requires_user_choice: false,
178                metadata: alignment_metadata(join_inputs, align_output.clone(), alignment_mode)?,
179            });
180            align_idx += 1;
181            vec![align_output]
182        } else {
183            join_inputs
184        };
185
186        steps.push(DataPlanStep {
187            kind: DataPlanStepKind::Join,
188            source_id: None,
189            adapter_id: None,
190            input_representation: Some(port_output_representation.clone()),
191            output_representation: Some(port_output_representation.clone()),
192            fit_scope: FitScope::Stateless,
193            requires_user_choice: false,
194            metadata: BTreeMap::from([
195                (
196                    "inputs".to_string(),
197                    serde_json::Value::Array(
198                        join_inputs
199                            .into_iter()
200                            .map(serde_json::Value::String)
201                            .collect(),
202                    ),
203                ),
204                (
205                    "output".to_string(),
206                    serde_json::Value::String(format!("port:{}", port.name)),
207                ),
208            ]),
209        });
210        output_representation = Some(port_output_representation);
211    }
212
213    let output_representation = output_representation
214        .ok_or_else(|| DataError::Validation("model input plan produced no outputs".to_string()))?;
215    let plan = DataPlan {
216        id: request.id.clone(),
217        steps,
218        output_representation,
219        issues,
220    };
221    plan.validate()?;
222    Ok(plan)
223}
224
225fn validate_request(request: &DataPlanRequest) -> Result<()> {
226    if request.id.trim().is_empty() {
227        return Err(DataError::Validation(
228            "data plan request id is empty".to_string(),
229        ));
230    }
231    if let Some(source_ids) = &request.source_ids {
232        let mut seen = std::collections::BTreeSet::new();
233        for source_id in source_ids {
234            if !seen.insert(source_id) {
235                return Err(DataError::Validation(format!(
236                    "data plan request contains duplicate source `{source_id}`"
237                )));
238            }
239        }
240    }
241    Ok(())
242}
243
244fn candidate_sources<'a>(
245    schema: &'a DatasetSchema,
246    request: &DataPlanRequest,
247) -> Result<Vec<&'a SourceDescriptor>> {
248    let mut sources = schema.sources.iter().collect::<Vec<_>>();
249    sources.sort_by(|left, right| left.id.cmp(&right.id));
250    if let Some(requested) = &request.source_ids {
251        let by_id = sources
252            .iter()
253            .map(|source| (&source.id, *source))
254            .collect::<BTreeMap<_, _>>();
255        requested
256            .iter()
257            .map(|source_id| {
258                by_id.get(source_id).copied().ok_or_else(|| {
259                    DataError::Validation(format!(
260                        "data plan request references unknown source `{source_id}`"
261                    ))
262                })
263            })
264            .collect()
265    } else {
266        Ok(sources)
267    }
268}
269
270fn resolve_source<'a>(
271    source: &'a SourceDescriptor,
272    port: &crate::InputPortSpec,
273    adapters: &AdapterRegistry,
274    policy: &PlanningPolicy,
275) -> Option<ResolvedSource<'a>> {
276    for target_type in &port.accepted_types {
277        for target_representation in &port.accepted_representations {
278            let known_rank = if target_representation == &source.native_representation.id {
279                source.native_representation.rank
280            } else {
281                crate::builtin_representations()
282                    .into_iter()
283                    .find(|spec| &spec.id == target_representation)
284                    .and_then(|spec| spec.rank)
285            };
286            if port
287                .rank
288                .zip(known_rank)
289                .is_some_and(|(required, actual)| required != actual)
290            {
291                continue;
292            }
293            let mut resolution = adapters.find_path(
294                &source.type_id,
295                &source.native_representation.id,
296                target_type,
297                target_representation,
298                policy,
299            );
300            if resolution.requires_user_choice {
301                // Keep an inspectable candidate and all choices; materialization
302                // already refuses unresolved plans. Never silently drop an
303                // ambiguous optional source or hide the solver diagnostics.
304                let mut candidate_policy = policy.clone();
305                candidate_policy.require_user_choice_on_ambiguity = false;
306                resolution.path = adapters
307                    .find_path(
308                        &source.type_id,
309                        &source.native_representation.id,
310                        target_type,
311                        target_representation,
312                        &candidate_policy,
313                    )
314                    .path;
315            }
316            if let Some(path) = resolution.path {
317                if port.rank.is_some() && known_rank.is_none() {
318                    resolution.issues.push(PlanIssue {
319                        code: "unverified_rank".into(),
320                        message: format!("input port `{}` requires rank {:?}, but representation `{target_representation}` has no declared rank", port.name, port.rank),
321                        choices: vec![],
322                    });
323                }
324                return Some(ResolvedSource {
325                    source,
326                    target_representation: target_representation.clone(),
327                    path,
328                    issues: resolution.issues,
329                });
330            }
331        }
332    }
333    None
334}
335
336fn source_output_id(source_id: &SourceId) -> String {
337    format!("src:{source_id}")
338}
339
340#[cfg(test)]
341mod tests {
342    use super::*;
343    use crate::adapter::AdapterRegistrySpec;
344    use crate::ids::{RepresentationId, SourceId, TypeId};
345    use crate::plan::DataPlan;
346
347    fn load_schema() -> DatasetSchema {
348        serde_json::from_str(include_str!(
349            "../../../examples/fixtures/oof_campaign/schema_nir_6_samples.json"
350        ))
351        .unwrap()
352    }
353
354    fn load_model_input() -> ModelInputSpec {
355        serde_json::from_str(include_str!(
356            "../../../examples/fixtures/oof_campaign/model_input_tabular_numeric.json"
357        ))
358        .unwrap()
359    }
360
361    fn load_registry() -> AdapterRegistry {
362        let spec: AdapterRegistrySpec = serde_json::from_str(include_str!(
363            "../../../examples/fixtures/oof_campaign/adapter_registry_signal_to_tabular.json"
364        ))
365        .unwrap();
366        AdapterRegistry::from_spec(spec).unwrap()
367    }
368
369    #[test]
370    fn fixture_plan_matches_expected_json() {
371        let plan = plan_model_input(
372            &load_schema(),
373            &load_model_input(),
374            &load_registry(),
375            &DataPlanRequest {
376                id: "nir-to-tabular".to_string(),
377                source_ids: Some(vec![SourceId::new("nir").unwrap()]),
378                planning_policy: PlanningPolicy::default(),
379            },
380        )
381        .unwrap();
382        let expected: DataPlan = serde_json::from_str(include_str!(
383            "../../../examples/fixtures/oof_campaign/expected_data_plan_nir_to_tabular_resolved.json"
384        ))
385        .unwrap();
386
387        assert_eq!(plan, expected);
388    }
389
390    #[test]
391    fn planner_refuses_missing_required_port_source() {
392        let err = plan_model_input(
393            &load_schema(),
394            &load_model_input(),
395            &load_registry(),
396            &DataPlanRequest {
397                id: "missing-source".to_string(),
398                source_ids: Some(vec![SourceId::new("other").unwrap()]),
399                planning_policy: PlanningPolicy::default(),
400            },
401        )
402        .unwrap_err();
403
404        assert!(err.to_string().contains("unknown source"));
405    }
406
407    #[test]
408    fn planner_checks_target_rank_after_adaptation() {
409        let mut model = load_model_input();
410        model.ports[0].rank = Some(999);
411        assert!(plan_model_input(
412            &load_schema(),
413            &model,
414            &load_registry(),
415            &DataPlanRequest::new("rank")
416        )
417        .is_err());
418        model.ports[0].rank = Some(2);
419        assert!(plan_model_input(
420            &load_schema(),
421            &model,
422            &load_registry(),
423            &DataPlanRequest::new("rank")
424        )
425        .unwrap()
426        .issues
427        .is_empty());
428    }
429
430    #[test]
431    fn planner_keeps_each_ports_representation_independent() {
432        let schema = load_schema();
433        let mut model = load_model_input();
434        let source = &schema.sources[0];
435        model.ports.insert(
436            0,
437            crate::InputPortSpec {
438                name: "raw".into(),
439                accepted_representations: vec![source.native_representation.id.clone()],
440                accepted_types: vec![source.type_id.clone()],
441                rank: source.native_representation.rank,
442                multi_source: false,
443                optional: false,
444            },
445        );
446        let plan = plan_model_input(
447            &schema,
448            &model,
449            &load_registry(),
450            &DataPlanRequest::new("multi-port"),
451        )
452        .unwrap();
453        plan.validate().unwrap();
454        let ports: Vec<_> = plan
455            .steps
456            .iter()
457            .filter(|step| step.kind == DataPlanStepKind::Join)
458            .collect();
459        assert_ne!(
460            ports[0].output_representation,
461            ports[1].output_representation
462        );
463        for port in ports {
464            assert_eq!(port.input_representation, port.output_representation);
465        }
466    }
467
468    #[test]
469    fn planner_preserves_ambiguous_adapter_choices() {
470        let mut spec: AdapterRegistrySpec = serde_json::from_str(include_str!(
471            "../../../examples/fixtures/oof_campaign/adapter_registry_signal_to_tabular.json"
472        ))
473        .unwrap();
474        let mut duplicate = spec.adapters[0].clone();
475        duplicate.id.push_str("_alternative");
476        spec.adapters.push(duplicate);
477        let registry = AdapterRegistry::from_spec(spec).unwrap();
478        let plan = plan_model_input(
479            &load_schema(),
480            &load_model_input(),
481            &registry,
482            &DataPlanRequest::new("ambiguous"),
483        )
484        .unwrap();
485        assert!(plan.requires_user_choice());
486        let issue = plan
487            .issues
488            .iter()
489            .find(|issue| issue.code == "ambiguous_path")
490            .unwrap();
491        assert_eq!(issue.choices.len(), 2);
492    }
493
494    #[test]
495    fn planner_emits_alignment_step_for_multi_source_port() {
496        let mut schema = load_schema();
497        let mut chem = schema.sources[0].clone();
498        chem.id = SourceId::new("chem").unwrap();
499        chem.name = "Chemistry".to_string();
500        chem.type_id = TypeId::new("table").unwrap();
501        chem.native_representation.id = RepresentationId::new("tabular_numeric").unwrap();
502        chem.native_representation.type_id = TypeId::new("table").unwrap();
503        chem.native_representation.container = "dataframe".to_string();
504        schema.sources.push(chem);
505
506        let plan = plan_model_input(
507            &schema,
508            &load_model_input(),
509            &load_registry(),
510            &DataPlanRequest {
511                id: "nir-chem-to-tabular".to_string(),
512                source_ids: Some(vec![
513                    SourceId::new("nir").unwrap(),
514                    SourceId::new("chem").unwrap(),
515                ]),
516                planning_policy: PlanningPolicy::default(),
517            },
518        )
519        .unwrap();
520
521        let align = plan
522            .steps
523            .iter()
524            .find(|step| step.kind == DataPlanStepKind::Align)
525            .expect("multi-source plan should expose an Align step");
526        assert_eq!(align.metadata["alignment"], serde_json::json!("left"));
527        assert_eq!(align.metadata["output"], serde_json::json!("step:align:0"));
528        assert_eq!(
529            align.metadata["inputs"].as_array().unwrap().len(),
530            2,
531            "alignment should receive both source feature streams"
532        );
533
534        let join = plan
535            .steps
536            .iter()
537            .find(|step| step.kind == DataPlanStepKind::Join)
538            .unwrap();
539        assert_eq!(join.metadata["inputs"], serde_json::json!(["step:align:0"]));
540    }
541}