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 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 ®istry,
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}