Skip to main content

dag_ml_data_core/
plan.rs

1use std::collections::BTreeMap;
2
3use serde::{Deserialize, Serialize};
4
5use crate::error::{DataError, Result};
6use crate::ids::{RepresentationId, SourceId};
7
8#[derive(Clone, Debug, Eq, PartialEq, Serialize, Deserialize)]
9#[serde(rename_all = "snake_case")]
10pub enum FitScope {
11    Stateless,
12    FoldTrain,
13    FullTrain,
14    InferenceOnly,
15}
16
17#[derive(Clone, Debug, Eq, PartialEq, Serialize, Deserialize)]
18#[serde(rename_all = "snake_case")]
19pub enum DataPlanStepKind {
20    Materialize,
21    Adapt,
22    Align,
23    Join,
24    Collate,
25}
26
27#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
28pub struct DataPlanStep {
29    pub kind: DataPlanStepKind,
30    pub source_id: Option<SourceId>,
31    pub adapter_id: Option<String>,
32    pub input_representation: Option<RepresentationId>,
33    pub output_representation: Option<RepresentationId>,
34    pub fit_scope: FitScope,
35    #[serde(default)]
36    pub requires_user_choice: bool,
37    #[serde(default)]
38    pub metadata: BTreeMap<String, serde_json::Value>,
39}
40
41#[derive(Clone, Debug, Eq, PartialEq, Serialize, Deserialize)]
42pub struct PlanIssue {
43    pub code: String,
44    pub message: String,
45    #[serde(default)]
46    pub choices: Vec<String>,
47}
48
49#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
50pub struct DataPlan {
51    pub id: String,
52    pub steps: Vec<DataPlanStep>,
53    pub output_representation: RepresentationId,
54    #[serde(default)]
55    pub issues: Vec<PlanIssue>,
56}
57
58impl DataPlan {
59    pub fn validate(&self) -> Result<()> {
60        if self.id.trim().is_empty() {
61            return Err(DataError::Validation("data plan id is empty".to_string()));
62        }
63        if self.steps.is_empty() {
64            return Err(DataError::Validation(format!(
65                "data plan `{}` contains no steps",
66                self.id
67            )));
68        }
69        let mut outputs = BTreeMap::new();
70        let mut source_outputs = BTreeMap::new();
71        for (idx, step) in self.steps.iter().enumerate() {
72            let invalid = |message: String| {
73                DataError::Validation(format!("data plan `{}` step {idx}: {message}", self.id))
74            };
75            match step.kind {
76                DataPlanStepKind::Materialize if step.source_id.is_none() => {
77                    return Err(DataError::Validation(format!(
78                        "data plan `{}` step {} materializes without source_id",
79                        self.id, idx
80                    )));
81                }
82                DataPlanStepKind::Adapt if step.adapter_id.is_none() => {
83                    return Err(DataError::Validation(format!(
84                        "data plan `{}` step {} adapts without adapter_id",
85                        self.id, idx
86                    )));
87                }
88                _ => {}
89            }
90            if step
91                .adapter_id
92                .as_ref()
93                .is_some_and(|id| id.trim().is_empty())
94            {
95                return Err(invalid("adapter_id must not be empty".into()));
96            }
97            let mut inputs = Vec::new();
98            if let Some(value) = step.metadata.get("input") {
99                inputs.push(
100                    value
101                        .as_str()
102                        .ok_or_else(|| invalid("input must be a string".into()))?,
103                );
104            }
105            if let Some(value) = step.metadata.get("inputs") {
106                let values = value
107                    .as_array()
108                    .ok_or_else(|| invalid("inputs must be an array".into()))?;
109                if values.is_empty() {
110                    return Err(invalid("inputs must not be empty".into()));
111                }
112                for value in values {
113                    inputs.push(
114                        value
115                            .as_str()
116                            .ok_or_else(|| invalid("inputs must contain strings".into()))?,
117                    );
118                }
119            }
120            if idx == 0 && step.kind != DataPlanStepKind::Materialize {
121                return Err(invalid(
122                    "must materialize a source before consuming data".into(),
123                ));
124            }
125            for input in inputs {
126                let representation = outputs.get(input).ok_or_else(|| {
127                    invalid(format!("input `{input}` references no earlier output"))
128                })?;
129                if let (Some(actual), Some(expected)) = (representation, &step.input_representation)
130                {
131                    if actual != expected {
132                        return Err(invalid(format!(
133                            "input `{input}` representation `{actual}` does not match `{expected}`"
134                        )));
135                    }
136                }
137            }
138            let output = if let Some(value) = step.metadata.get("output") {
139                let id = value
140                    .as_str()
141                    .filter(|id| !id.trim().is_empty())
142                    .ok_or_else(|| invalid("output must be a nonempty string".into()))?;
143                Some(id.to_owned())
144            } else if step.kind == DataPlanStepKind::Materialize {
145                step.source_id.as_ref().map(|id| format!("src:{id}"))
146            } else {
147                None // Published linear plans may omit explicit edge names.
148            };
149            if let Some(output) = output {
150                if let Some(previous) =
151                    outputs.insert(output.clone(), step.output_representation.clone())
152                {
153                    // A source can be materialized for several model ports.
154                    if step.kind != DataPlanStepKind::Materialize
155                        || previous != step.output_representation
156                        || source_outputs.get(&output) != Some(&step.source_id)
157                    {
158                        return Err(invalid(format!("duplicate output `{output}`")));
159                    }
160                }
161                if step.kind == DataPlanStepKind::Materialize {
162                    source_outputs.insert(output, step.source_id.clone());
163                }
164            }
165        }
166        if let Some(Some(output)) = self
167            .steps
168            .last()
169            .map(|step| step.output_representation.as_ref())
170        {
171            if output != &self.output_representation {
172                return Err(DataError::Validation(format!(
173                    "data plan `{}` output representation does not match its final step",
174                    self.id
175                )));
176            }
177        }
178        Ok(())
179    }
180
181    pub fn requires_user_choice(&self) -> bool {
182        self.steps.iter().any(|step| step.requires_user_choice) || !self.issues.is_empty()
183    }
184}
185
186#[cfg(test)]
187mod tests {
188    use super::*;
189    use crate::ids::{RepresentationId, SourceId};
190
191    #[test]
192    fn rejects_empty_plan() {
193        let plan = DataPlan {
194            id: "p".to_string(),
195            steps: vec![],
196            output_representation: RepresentationId::new("tabular").unwrap(),
197            issues: vec![],
198        };
199
200        assert!(plan.validate().is_err());
201    }
202
203    #[test]
204    fn flags_user_choice() {
205        let plan = DataPlan {
206            id: "p".to_string(),
207            steps: vec![DataPlanStep {
208                kind: DataPlanStepKind::Materialize,
209                source_id: Some(SourceId::new("nir").unwrap()),
210                adapter_id: None,
211                input_representation: None,
212                output_representation: Some(RepresentationId::new("signal").unwrap()),
213                fit_scope: FitScope::Stateless,
214                requires_user_choice: true,
215                metadata: BTreeMap::new(),
216            }],
217            output_representation: RepresentationId::new("signal").unwrap(),
218            issues: vec![],
219        };
220
221        assert!(plan.validate().is_ok());
222        assert!(plan.requires_user_choice());
223    }
224
225    #[test]
226    fn validates_explicit_edges_without_breaking_linear_plans() {
227        let original: DataPlan = serde_json::from_str(include_str!(
228            "../../../examples/fixtures/oof_campaign/expected_data_plan_nir_to_tabular.json"
229        ))
230        .unwrap();
231        original.validate().unwrap();
232        let mut plan = original.clone();
233        plan.steps[1]
234            .metadata
235            .insert("input".into(), serde_json::json!("step:missing"));
236        assert!(plan
237            .validate()
238            .unwrap_err()
239            .to_string()
240            .contains("no earlier output"));
241        let mut plan = original.clone();
242        plan.steps[1]
243            .metadata
244            .insert("input".into(), serde_json::json!(42));
245        assert!(plan.validate().is_err());
246        let mut plan = original.clone();
247        plan.steps[1].input_representation = Some(RepresentationId::new("wrong").unwrap());
248        assert!(plan.validate().is_err());
249        let mut plan = original.clone();
250        plan.steps[1]
251            .metadata
252            .insert("output".into(), serde_json::json!("src:nir"));
253        assert!(plan.validate().is_err());
254        let mut plan = original.clone();
255        plan.steps.remove(0);
256        assert!(plan.validate().is_err());
257        let mut plan = original.clone();
258        plan.output_representation = RepresentationId::new("wrong").unwrap();
259        assert!(plan.validate().is_err());
260        let mut plan = original;
261        for step in &mut plan.steps {
262            step.metadata.clear();
263        }
264        plan.validate().unwrap();
265    }
266}