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