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