Skip to main content

dag_ml_core/runtime/
scheduler.rs

1// Auto-split from the former monolithic `runtime.rs` (pure refactor).
2use super::*;
3
4#[derive(Clone, Debug, Default)]
5pub struct SequentialScheduler;
6
7#[derive(Clone, Debug)]
8pub struct ParallelScheduler {
9    max_workers: usize,
10}
11
12impl ParallelScheduler {
13    pub fn new(max_workers: usize) -> Result<Self> {
14        if max_workers == 0 {
15            return Err(DagMlError::RuntimeValidation(
16                "parallel scheduler max_workers must be at least 1".to_string(),
17            ));
18        }
19        Ok(Self { max_workers })
20    }
21
22    pub fn max_workers(&self) -> usize {
23        self.max_workers
24    }
25}
26
27#[derive(Clone, Debug)]
28pub(crate) struct PhaseScope {
29    pub(crate) phase: Phase,
30    pub(crate) variant_id: Option<VariantId>,
31    pub(crate) variant: Option<VariantExecutionSpec>,
32    pub(crate) fold_id: Option<FoldId>,
33    pub(crate) seed_root: Option<u64>,
34}
35
36#[derive(Clone, Debug)]
37pub(crate) struct ReplayPredictionCacheContract {
38    pub(crate) requirement: BundlePredictionRequirement,
39    pub(crate) cache: BundlePredictionCacheRecord,
40}
41
42pub(crate) struct MaterializedReplayArtifacts {
43    pub(crate) handles: BTreeMap<NodeId, BTreeMap<String, HandleRef>>,
44    pub(crate) inputs: BTreeMap<NodeId, BTreeMap<String, ArtifactInputSpec>>,
45}
46
47fn prediction_output_ports_for_node(plan: &ExecutionPlan, node_id: &NodeId) -> Result<Vec<String>> {
48    let node = plan
49        .graph_plan
50        .graph
51        .nodes
52        .iter()
53        .find(|node| node.id == *node_id)
54        .ok_or_else(|| {
55            DagMlError::RuntimeValidation(format!(
56                "node `{node_id}` is absent from the execution graph"
57            ))
58        })?;
59    let mut ports = node
60        .ports
61        .outputs
62        .iter()
63        .filter(|port| port.kind == PortKind::Prediction)
64        .map(|port| port.name.clone())
65        .collect::<Vec<_>>();
66    ports.sort();
67    Ok(ports)
68}
69
70fn normalize_prediction_result_port(
71    node_id: &NodeId,
72    block_kind: &str,
73    producer_port: &mut Option<String>,
74    prediction_ports: &[String],
75) -> Result<()> {
76    if let Some(port) = producer_port.as_ref() {
77        if port.trim().is_empty() {
78            return Err(DagMlError::RuntimeValidation(format!(
79                "node `{node_id}` emitted {block_kind} with blank producer_port"
80            )));
81        }
82        if !prediction_ports.iter().any(|candidate| candidate == port) {
83            return Err(DagMlError::RuntimeValidation(format!(
84                "node `{node_id}` emitted {block_kind} for undeclared or non-prediction output port `{port}`; declared prediction ports are {:?}",
85                prediction_ports
86            )));
87        }
88        return Ok(());
89    }
90    match prediction_ports {
91        [only] => {
92            *producer_port = Some(only.clone());
93            Ok(())
94        }
95        [] => Err(DagMlError::RuntimeValidation(format!(
96            "node `{node_id}` emitted {block_kind} without producer_port but declares no prediction output port"
97        ))),
98        _ => Err(DagMlError::RuntimeValidation(format!(
99            "node `{node_id}` emitted {block_kind} without producer_port but declares {} prediction output ports {:?}; multi-output controllers must emit producer_port explicitly",
100            prediction_ports.len(),
101            prediction_ports
102        ))),
103    }
104}
105
106pub(crate) fn normalize_result_prediction_ports(
107    plan: &ExecutionPlan,
108    task: &NodeTask,
109    result: &mut NodeResult,
110) -> Result<()> {
111    if result.predictions.is_empty()
112        && result.observation_predictions.is_empty()
113        && result.aggregated_predictions.is_empty()
114        && result.explanations.is_empty()
115    {
116        return Ok(());
117    }
118    let prediction_ports = prediction_output_ports_for_node(plan, &task.node_plan.node_id)?;
119    for block in &mut result.predictions {
120        normalize_prediction_result_port(
121            &task.node_plan.node_id,
122            "prediction block",
123            &mut block.producer_port,
124            &prediction_ports,
125        )?;
126    }
127    for block in &mut result.observation_predictions {
128        normalize_prediction_result_port(
129            &task.node_plan.node_id,
130            "observation prediction block",
131            &mut block.producer_port,
132            &prediction_ports,
133        )?;
134    }
135    for block in &mut result.aggregated_predictions {
136        normalize_prediction_result_port(
137            &task.node_plan.node_id,
138            "aggregated prediction block",
139            &mut block.producer_port,
140            &prediction_ports,
141        )?;
142    }
143    for block in &mut result.explanations {
144        normalize_prediction_result_port(
145            &task.node_plan.node_id,
146            "explanation block",
147            &mut block.producer_port,
148            &prediction_ports,
149        )?;
150    }
151    Ok(())
152}
153
154#[derive(Default)]
155pub(crate) struct PhaseScopeResources<'a> {
156    pub(crate) data_provider: Option<&'a dyn RuntimeDataProvider>,
157    pub(crate) replay_artifact_handles: Option<&'a BTreeMap<NodeId, BTreeMap<String, HandleRef>>>,
158    pub(crate) replay_artifact_inputs:
159        Option<&'a BTreeMap<NodeId, BTreeMap<String, ArtifactInputSpec>>>,
160    pub(crate) replay_bundle_id: Option<&'a BundleId>,
161    pub(crate) data_envelopes: Option<&'a BTreeMap<String, ExternalDataPlanEnvelope>>,
162    pub(crate) prediction_cache_store: Option<&'a dyn RuntimePredictionCacheStore>,
163    pub(crate) prediction_cache_contracts:
164        Option<&'a BTreeMap<String, ReplayPredictionCacheContract>>,
165    pub(crate) artifact_store: Option<&'a mut InMemoryArtifactStore>,
166}
167
168impl SequentialScheduler {
169    pub fn execute_phase(
170        &self,
171        plan: &ExecutionPlan,
172        controllers: &RuntimeControllerRegistry,
173        ctx: &mut RunContext,
174        phase: Phase,
175    ) -> Result<Vec<NodeResult>> {
176        plan.validate()?;
177        let variant_id = ctx.variant_id.clone();
178        let seed_root = ctx.root_seed;
179        self.execute_phase_scope(
180            plan,
181            controllers,
182            ctx,
183            PhaseScope {
184                phase,
185                variant_id,
186                variant: None,
187                fold_id: None,
188                seed_root,
189            },
190            PhaseScopeResources::default(),
191        )
192    }
193
194    pub fn execute_phase_with_data_provider(
195        &self,
196        plan: &ExecutionPlan,
197        controllers: &RuntimeControllerRegistry,
198        data_provider: &dyn RuntimeDataProvider,
199        ctx: &mut RunContext,
200        phase: Phase,
201    ) -> Result<Vec<NodeResult>> {
202        plan.validate()?;
203        let variant_id = ctx.variant_id.clone();
204        let seed_root = ctx.root_seed;
205        self.execute_phase_scope(
206            plan,
207            controllers,
208            ctx,
209            PhaseScope {
210                phase,
211                variant_id,
212                variant: None,
213                fold_id: None,
214                seed_root,
215            },
216            PhaseScopeResources {
217                data_provider: Some(data_provider),
218                ..Default::default()
219            },
220        )
221    }
222
223    pub fn execute_campaign_phase(
224        &self,
225        plan: &ExecutionPlan,
226        controllers: &RuntimeControllerRegistry,
227        ctx: &mut RunContext,
228        phase: Phase,
229    ) -> Result<Vec<NodeResult>> {
230        plan.validate()?;
231        let mut results = Vec::new();
232        let fold_ids = if phase == Phase::FitCv {
233            plan.fold_set
234                .as_ref()
235                .map(|fold_set| {
236                    fold_set
237                        .folds
238                        .iter()
239                        .map(|fold| Some(fold.fold_id.clone()))
240                        .collect::<Vec<_>>()
241                })
242                .unwrap_or_else(|| vec![None])
243        } else {
244            vec![None]
245        };
246        for variant in &plan.variants {
247            if ctx
248                .variant_id
249                .as_ref()
250                .is_some_and(|requested| requested != &variant.variant_id)
251            {
252                continue;
253            }
254            for fold_id in &fold_ids {
255                let seed_root = variant.seed.or(ctx.root_seed);
256                results.extend(self.execute_phase_scope(
257                    plan,
258                    controllers,
259                    ctx,
260                    PhaseScope {
261                        phase,
262                        variant_id: Some(variant.variant_id.clone()),
263                        variant: Some(VariantExecutionSpec::from_plan(variant)),
264                        fold_id: fold_id.clone(),
265                        seed_root,
266                    },
267                    PhaseScopeResources::default(),
268                )?);
269            }
270        }
271        Ok(results)
272    }
273
274    pub fn execute_campaign_phase_with_data_provider(
275        &self,
276        plan: &ExecutionPlan,
277        controllers: &RuntimeControllerRegistry,
278        data_provider: &dyn RuntimeDataProvider,
279        ctx: &mut RunContext,
280        phase: Phase,
281    ) -> Result<Vec<NodeResult>> {
282        plan.validate()?;
283        let mut results = Vec::new();
284        let fold_ids = if phase == Phase::FitCv {
285            plan.fold_set
286                .as_ref()
287                .map(|fold_set| {
288                    fold_set
289                        .folds
290                        .iter()
291                        .map(|fold| Some(fold.fold_id.clone()))
292                        .collect::<Vec<_>>()
293                })
294                .unwrap_or_else(|| vec![None])
295        } else {
296            vec![None]
297        };
298        for variant in &plan.variants {
299            if ctx
300                .variant_id
301                .as_ref()
302                .is_some_and(|requested| requested != &variant.variant_id)
303            {
304                continue;
305            }
306            for fold_id in &fold_ids {
307                let seed_root = variant.seed.or(ctx.root_seed);
308                results.extend(self.execute_phase_scope(
309                    plan,
310                    controllers,
311                    ctx,
312                    PhaseScope {
313                        phase,
314                        variant_id: Some(variant.variant_id.clone()),
315                        variant: Some(VariantExecutionSpec::from_plan(variant)),
316                        fold_id: fold_id.clone(),
317                        seed_root,
318                    },
319                    PhaseScopeResources {
320                        data_provider: Some(data_provider),
321                        ..Default::default()
322                    },
323                )?);
324            }
325        }
326        Ok(results)
327    }
328
329    pub fn execute_campaign_phase_with_data_provider_and_artifact_store(
330        &self,
331        plan: &ExecutionPlan,
332        controllers: &RuntimeControllerRegistry,
333        data_provider: &dyn RuntimeDataProvider,
334        artifact_store: &mut InMemoryArtifactStore,
335        ctx: &mut RunContext,
336        phase: Phase,
337    ) -> Result<Vec<NodeResult>> {
338        plan.validate()?;
339        let mut results = Vec::new();
340        let fold_ids = if phase == Phase::FitCv {
341            plan.fold_set
342                .as_ref()
343                .map(|fold_set| {
344                    fold_set
345                        .folds
346                        .iter()
347                        .map(|fold| Some(fold.fold_id.clone()))
348                        .collect::<Vec<_>>()
349                })
350                .unwrap_or_else(|| vec![None])
351        } else {
352            vec![None]
353        };
354        for variant in &plan.variants {
355            if ctx
356                .variant_id
357                .as_ref()
358                .is_some_and(|requested| requested != &variant.variant_id)
359            {
360                continue;
361            }
362            for fold_id in &fold_ids {
363                let seed_root = variant.seed.or(ctx.root_seed);
364                results.extend(self.execute_phase_scope(
365                    plan,
366                    controllers,
367                    ctx,
368                    PhaseScope {
369                        phase,
370                        variant_id: Some(variant.variant_id.clone()),
371                        variant: Some(VariantExecutionSpec::from_plan(variant)),
372                        fold_id: fold_id.clone(),
373                        seed_root,
374                    },
375                    PhaseScopeResources {
376                        data_provider: Some(data_provider),
377                        artifact_store: Some(&mut *artifact_store),
378                        ..Default::default()
379                    },
380                )?);
381            }
382        }
383        Ok(results)
384    }
385
386    pub fn execute_bundle_replay(
387        &self,
388        replay: BundleReplayExecution<'_>,
389        ctx: &mut RunContext,
390    ) -> Result<Vec<NodeResult>> {
391        replay.bundle.validate_against_plan(replay.plan)?;
392        replay
393            .replay_request
394            .validate_for_bundle_with_prediction_cache_store(
395                replay.bundle,
396                replay.prediction_cache_store.is_some(),
397            )?;
398        replay
399            .bundle
400            .validate_replay_envelopes(replay.data_envelopes)?;
401        let prediction_cache_contracts = if replay.replay_request.phase == Phase::Refit {
402            Some(replay_prediction_cache_contracts(replay.bundle)?)
403        } else {
404            None
405        };
406        if replay.replay_request.phase == Phase::Refit {
407            preload_replay_prediction_cache_store(
408                replay.bundle,
409                replay.prediction_cache_store,
410                ctx,
411            )?;
412        }
413        let replay_artifacts = materialize_replay_artifact_handles(
414            replay.plan,
415            replay.bundle,
416            replay.replay_request,
417            replay.artifact_store,
418            ctx,
419        )?;
420        let selected_variant = replay
421            .bundle
422            .selected_variant_id
423            .as_ref()
424            .map(|selected| {
425                replay
426                    .plan
427                    .variants
428                    .iter()
429                    .find(|variant| &variant.variant_id == selected)
430                    .map(VariantExecutionSpec::from_plan)
431                    .ok_or_else(|| {
432                        DagMlError::RuntimeValidation(format!(
433                            "bundle `{}` selected unknown variant `{selected}`",
434                            replay.bundle.bundle_id
435                        ))
436                    })
437            })
438            .transpose()?;
439        let seed_root = selected_variant
440            .as_ref()
441            .and_then(|variant| variant.seed)
442            .or(ctx.root_seed);
443
444        self.execute_phase_scope(
445            replay.plan,
446            replay.controllers,
447            ctx,
448            PhaseScope {
449                phase: replay.replay_request.phase,
450                variant_id: replay.bundle.selected_variant_id.clone(),
451                variant: selected_variant,
452                fold_id: None,
453                seed_root,
454            },
455            PhaseScopeResources {
456                data_provider: Some(replay.data_provider),
457                replay_artifact_handles: Some(&replay_artifacts.handles),
458                replay_artifact_inputs: Some(&replay_artifacts.inputs),
459                replay_bundle_id: Some(&replay.bundle.bundle_id),
460                data_envelopes: Some(replay.data_envelopes),
461                prediction_cache_store: replay.prediction_cache_store,
462                prediction_cache_contracts: prediction_cache_contracts.as_ref(),
463                ..Default::default()
464            },
465        )
466    }
467
468    fn execute_phase_scope(
469        &self,
470        plan: &ExecutionPlan,
471        controllers: &RuntimeControllerRegistry,
472        ctx: &mut RunContext,
473        scope: PhaseScope,
474        mut resources: PhaseScopeResources<'_>,
475    ) -> Result<Vec<NodeResult>> {
476        let _phase_span = crate::observability::phase_span(
477            ctx.run_id.as_str(),
478            plan.id.as_str(),
479            scope.phase.as_str(),
480            scope.variant_id.as_ref().map(VariantId::as_str),
481            scope.fold_id.as_ref().map(FoldId::as_str),
482        )
483        .entered();
484        let mut results = Vec::new();
485        let mut output_handles = BTreeMap::<NodeId, BTreeMap<String, HandleRef>>::new();
486        let mut output_data_views =
487            BTreeMap::<NodeId, BTreeMap<String, DataProviderViewSpec>>::new();
488        let mut input_lineage = BTreeMap::<NodeId, LineageId>::new();
489
490        for level in plan.node_parallel_levels_for_phase(scope.phase)? {
491            for node_id in &level {
492                let node_plan = plan
493                    .node_plans
494                    .get(node_id)
495                    .expect("execution plan was validated");
496                // Cross-branch merge reassembly (concat or late-fusion) is a
497                // scheduler/runtime handler, not a controller call: it reads the
498                // upstream branch OOF blocks from the prediction store and emits
499                // one merged per-sample OOF block. Intercept it before the
500                // controller path (and before the `requires_oof` edge collection,
501                // which is a stacking contract the branch inputs do not satisfy).
502                if let Some(reduction) = merge_reduction_mode(plan, node_plan) {
503                    if let Some(mut result) =
504                        reassemble_branch_merge(plan, node_plan, ctx, &scope, reduction)?
505                    {
506                        let task_node_plan = effective_node_plan_for_scope(node_plan, &scope)?;
507                        let task = NodeTask {
508                            inner_fold_set: None,
509                            run_id: ctx.run_id.clone(),
510                            node_plan: task_node_plan.clone(),
511                            phase: scope.phase,
512                            variant_id: scope.variant_id.clone(),
513                            variant: scope.variant.clone(),
514                            fold_id: scope.fold_id.clone(),
515                            branch_path: Vec::new(),
516                            input_handles: BTreeMap::new(),
517                            data_views: BTreeMap::new(),
518                            prediction_inputs: BTreeMap::new(),
519                            artifact_inputs: BTreeMap::new(),
520                            required_loss_attestations: NodeTask::required_loss_attestations_for(
521                                &task_node_plan,
522                                scope.phase,
523                            )?,
524                            fit_influence: FitInfluenceTask::default(),
525                            seed: None,
526                        };
527                        normalize_result_prediction_ports(plan, &task, &mut result)?;
528                        result.validate_for_task(&task)?;
529                        for prediction in &result.predictions {
530                            ctx.prediction_store.append(prediction.clone())?;
531                        }
532                        apply_result_scoring(
533                            &result,
534                            &mut ctx.score_collector,
535                            &mut ctx.regression_target_records,
536                        )?;
537                        ctx.lineage.record(result.lineage.clone())?;
538                        output_handles.insert(node_id.clone(), result.outputs.clone());
539                        input_lineage.insert(node_id.clone(), result.lineage.record_id.clone());
540                        results.push(result);
541                    }
542                    continue;
543                }
544                let controller = controllers.get(&node_plan.controller_id).ok_or_else(|| {
545                    DagMlError::RuntimeValidation(format!(
546                        "runtime controller `{}` is not registered",
547                        node_plan.controller_id
548                    ))
549                })?;
550                let collected_inputs = collect_input_handles(
551                    plan,
552                    node_plan,
553                    &output_handles,
554                    &output_data_views,
555                    &resources,
556                    ctx,
557                    &scope,
558                )?;
559                if collected_inputs.skip_node {
560                    continue;
561                }
562                let mut input_handles = collected_inputs.handles;
563                let mut artifact_inputs = BTreeMap::new();
564                if let Some(node_artifact_handles) = resources
565                    .replay_artifact_handles
566                    .and_then(|handles| handles.get(node_id))
567                {
568                    for (key, handle) in node_artifact_handles {
569                        if input_handles.insert(key.clone(), handle.clone()).is_some() {
570                            return Err(DagMlError::RuntimeValidation(format!(
571                                "node `{node_id}` received duplicate replay artifact input `{key}`"
572                            )));
573                        }
574                    }
575                }
576                if let Some(node_artifact_inputs) = resources
577                    .replay_artifact_inputs
578                    .and_then(|inputs| inputs.get(node_id))
579                {
580                    for (key, spec) in node_artifact_inputs {
581                        if artifact_inputs.insert(key.clone(), spec.clone()).is_some() {
582                            return Err(DagMlError::RuntimeValidation(format!(
583                                "node `{node_id}` received duplicate replay artifact metadata `{key}`"
584                            )));
585                        }
586                    }
587                }
588                let task_node_plan = effective_node_plan_for_scope(node_plan, &scope)?;
589                let inner_fold_set = inner_fold_set_for_scope(
590                    &plan.campaign,
591                    plan.fold_set.as_ref(),
592                    node_plan,
593                    &scope,
594                )?;
595                let fit_influence = fit_influence_task_for_node(
596                    plan,
597                    &task_node_plan,
598                    &collected_inputs.data_views,
599                )?;
600                let task = NodeTask {
601                    inner_fold_set,
602                    run_id: ctx.run_id.clone(),
603                    node_plan: task_node_plan.clone(),
604                    phase: scope.phase,
605                    variant_id: scope.variant_id.clone(),
606                    variant: scope.variant.clone(),
607                    fold_id: scope.fold_id.clone(),
608                    branch_path: Vec::new(),
609                    input_handles,
610                    data_views: collected_inputs.data_views,
611                    prediction_inputs: collected_inputs.prediction_inputs,
612                    artifact_inputs,
613                    required_loss_attestations: NodeTask::required_loss_attestations_for(
614                        &task_node_plan,
615                        scope.phase,
616                    )?,
617                    fit_influence,
618                    seed: derive_task_seed(
619                        scope.seed_root,
620                        scope.variant_id.as_ref(),
621                        scope.fold_id.as_ref(),
622                        &task_node_plan,
623                        scope.phase,
624                    ),
625                };
626                let _node_span = crate::observability::node_span(
627                    task.run_id.as_str(),
628                    plan.id.as_str(),
629                    task.phase.as_str(),
630                    task.node_plan.node_id.as_str(),
631                    task.node_plan.controller_id.as_str(),
632                )
633                .entered();
634                let mut result = controller.invoke(&task)?;
635                record_fit_influence_diagnostic(&task, &mut result);
636                normalize_result_prediction_ports(plan, &task, &mut result)?;
637                result.validate_for_task(&task)?;
638                apply_result_prediction_aggregation(
639                    plan,
640                    controllers,
641                    &task,
642                    &mut result,
643                    &resources,
644                )?;
645                attach_coordinator_input_lineage(
646                    &mut result,
647                    plan,
648                    &task.node_plan.node_id,
649                    &input_lineage,
650                )?;
651                if let Some(store) = resources.artifact_store.as_deref_mut() {
652                    if scope.phase == Phase::Refit {
653                        store.capture_refit_artifacts(&task, &result)?;
654                    }
655                }
656                for prediction in &result.predictions {
657                    ctx.prediction_store.append(prediction.clone())?;
658                }
659                for prediction in &result.aggregated_predictions {
660                    ctx.aggregated_prediction_store.append(prediction.clone())?;
661                }
662                apply_result_scoring(
663                    &result,
664                    &mut ctx.score_collector,
665                    &mut ctx.regression_target_records,
666                )?;
667                ctx.lineage.record(result.lineage.clone())?;
668                let data_views = derive_output_data_views(plan, &task, &result)?;
669                output_handles.insert(node_id.clone(), result.outputs.clone());
670                output_data_views.insert(node_id.clone(), data_views);
671                input_lineage.insert(node_id.clone(), result.lineage.record_id.clone());
672                results.push(result);
673            }
674        }
675
676        Ok(results)
677    }
678}
679
680impl ParallelScheduler {
681    pub fn execute_phase(
682        &self,
683        plan: &ExecutionPlan,
684        controllers: &RuntimeControllerRegistry,
685        ctx: &mut RunContext,
686        phase: Phase,
687    ) -> Result<Vec<NodeResult>> {
688        plan.validate()?;
689        let variant_id = ctx.variant_id.clone();
690        let seed_root = ctx.root_seed;
691        self.execute_phase_scope(
692            plan,
693            controllers,
694            ctx,
695            PhaseScope {
696                phase,
697                variant_id,
698                variant: None,
699                fold_id: None,
700                seed_root,
701            },
702            PhaseScopeResources::default(),
703        )
704    }
705
706    pub fn execute_phase_with_data_provider(
707        &self,
708        plan: &ExecutionPlan,
709        controllers: &RuntimeControllerRegistry,
710        data_provider: &dyn RuntimeDataProvider,
711        ctx: &mut RunContext,
712        phase: Phase,
713    ) -> Result<Vec<NodeResult>> {
714        plan.validate()?;
715        let variant_id = ctx.variant_id.clone();
716        let seed_root = ctx.root_seed;
717        self.execute_phase_scope(
718            plan,
719            controllers,
720            ctx,
721            PhaseScope {
722                phase,
723                variant_id,
724                variant: None,
725                fold_id: None,
726                seed_root,
727            },
728            PhaseScopeResources {
729                data_provider: Some(data_provider),
730                ..Default::default()
731            },
732        )
733    }
734
735    pub fn execute_campaign_phase(
736        &self,
737        plan: &ExecutionPlan,
738        controllers: &RuntimeControllerRegistry,
739        ctx: &mut RunContext,
740        phase: Phase,
741    ) -> Result<Vec<NodeResult>> {
742        plan.validate()?;
743        let mut results = Vec::new();
744        let fold_ids = if phase == Phase::FitCv {
745            plan.fold_set
746                .as_ref()
747                .map(|fold_set| {
748                    fold_set
749                        .folds
750                        .iter()
751                        .map(|fold| Some(fold.fold_id.clone()))
752                        .collect::<Vec<_>>()
753                })
754                .unwrap_or_else(|| vec![None])
755        } else {
756            vec![None]
757        };
758        for variant in &plan.variants {
759            if ctx
760                .variant_id
761                .as_ref()
762                .is_some_and(|requested| requested != &variant.variant_id)
763            {
764                continue;
765            }
766            for fold_id in &fold_ids {
767                let seed_root = variant.seed.or(ctx.root_seed);
768                results.extend(self.execute_phase_scope(
769                    plan,
770                    controllers,
771                    ctx,
772                    PhaseScope {
773                        phase,
774                        variant_id: Some(variant.variant_id.clone()),
775                        variant: Some(VariantExecutionSpec::from_plan(variant)),
776                        fold_id: fold_id.clone(),
777                        seed_root,
778                    },
779                    PhaseScopeResources::default(),
780                )?);
781            }
782        }
783        Ok(results)
784    }
785
786    pub fn execute_campaign_phase_with_data_provider(
787        &self,
788        plan: &ExecutionPlan,
789        controllers: &RuntimeControllerRegistry,
790        data_provider: &dyn RuntimeDataProvider,
791        ctx: &mut RunContext,
792        phase: Phase,
793    ) -> Result<Vec<NodeResult>> {
794        plan.validate()?;
795        let mut results = Vec::new();
796        let fold_ids = if phase == Phase::FitCv {
797            plan.fold_set
798                .as_ref()
799                .map(|fold_set| {
800                    fold_set
801                        .folds
802                        .iter()
803                        .map(|fold| Some(fold.fold_id.clone()))
804                        .collect::<Vec<_>>()
805                })
806                .unwrap_or_else(|| vec![None])
807        } else {
808            vec![None]
809        };
810        for variant in &plan.variants {
811            if ctx
812                .variant_id
813                .as_ref()
814                .is_some_and(|requested| requested != &variant.variant_id)
815            {
816                continue;
817            }
818            for fold_id in &fold_ids {
819                let seed_root = variant.seed.or(ctx.root_seed);
820                results.extend(self.execute_phase_scope(
821                    plan,
822                    controllers,
823                    ctx,
824                    PhaseScope {
825                        phase,
826                        variant_id: Some(variant.variant_id.clone()),
827                        variant: Some(VariantExecutionSpec::from_plan(variant)),
828                        fold_id: fold_id.clone(),
829                        seed_root,
830                    },
831                    PhaseScopeResources {
832                        data_provider: Some(data_provider),
833                        ..Default::default()
834                    },
835                )?);
836            }
837        }
838        Ok(results)
839    }
840
841    pub fn execute_campaign_phase_with_data_provider_and_artifact_store(
842        &self,
843        plan: &ExecutionPlan,
844        controllers: &RuntimeControllerRegistry,
845        data_provider: &dyn RuntimeDataProvider,
846        artifact_store: &mut InMemoryArtifactStore,
847        ctx: &mut RunContext,
848        phase: Phase,
849    ) -> Result<Vec<NodeResult>> {
850        plan.validate()?;
851        let mut results = Vec::new();
852        let fold_ids = if phase == Phase::FitCv {
853            plan.fold_set
854                .as_ref()
855                .map(|fold_set| {
856                    fold_set
857                        .folds
858                        .iter()
859                        .map(|fold| Some(fold.fold_id.clone()))
860                        .collect::<Vec<_>>()
861                })
862                .unwrap_or_else(|| vec![None])
863        } else {
864            vec![None]
865        };
866        for variant in &plan.variants {
867            if ctx
868                .variant_id
869                .as_ref()
870                .is_some_and(|requested| requested != &variant.variant_id)
871            {
872                continue;
873            }
874            for fold_id in &fold_ids {
875                let seed_root = variant.seed.or(ctx.root_seed);
876                results.extend(self.execute_phase_scope(
877                    plan,
878                    controllers,
879                    ctx,
880                    PhaseScope {
881                        phase,
882                        variant_id: Some(variant.variant_id.clone()),
883                        variant: Some(VariantExecutionSpec::from_plan(variant)),
884                        fold_id: fold_id.clone(),
885                        seed_root,
886                    },
887                    PhaseScopeResources {
888                        data_provider: Some(data_provider),
889                        artifact_store: Some(&mut *artifact_store),
890                        ..Default::default()
891                    },
892                )?);
893            }
894        }
895        Ok(results)
896    }
897
898    pub fn execute_bundle_replay(
899        &self,
900        replay: BundleReplayExecution<'_>,
901        ctx: &mut RunContext,
902    ) -> Result<Vec<NodeResult>> {
903        replay.bundle.validate_against_plan(replay.plan)?;
904        replay
905            .replay_request
906            .validate_for_bundle_with_prediction_cache_store(
907                replay.bundle,
908                replay.prediction_cache_store.is_some(),
909            )?;
910        replay
911            .bundle
912            .validate_replay_envelopes(replay.data_envelopes)?;
913        let prediction_cache_contracts = if replay.replay_request.phase == Phase::Refit {
914            Some(replay_prediction_cache_contracts(replay.bundle)?)
915        } else {
916            None
917        };
918        if replay.replay_request.phase == Phase::Refit {
919            preload_replay_prediction_cache_store(
920                replay.bundle,
921                replay.prediction_cache_store,
922                ctx,
923            )?;
924        }
925        let replay_artifacts = materialize_replay_artifact_handles(
926            replay.plan,
927            replay.bundle,
928            replay.replay_request,
929            replay.artifact_store,
930            ctx,
931        )?;
932        let selected_variant = replay
933            .bundle
934            .selected_variant_id
935            .as_ref()
936            .map(|selected| {
937                replay
938                    .plan
939                    .variants
940                    .iter()
941                    .find(|variant| &variant.variant_id == selected)
942                    .map(VariantExecutionSpec::from_plan)
943                    .ok_or_else(|| {
944                        DagMlError::RuntimeValidation(format!(
945                            "bundle `{}` selected unknown variant `{selected}`",
946                            replay.bundle.bundle_id
947                        ))
948                    })
949            })
950            .transpose()?;
951        let seed_root = selected_variant
952            .as_ref()
953            .and_then(|variant| variant.seed)
954            .or(ctx.root_seed);
955
956        self.execute_phase_scope(
957            replay.plan,
958            replay.controllers,
959            ctx,
960            PhaseScope {
961                phase: replay.replay_request.phase,
962                variant_id: replay.bundle.selected_variant_id.clone(),
963                variant: selected_variant,
964                fold_id: None,
965                seed_root,
966            },
967            PhaseScopeResources {
968                data_provider: Some(replay.data_provider),
969                replay_artifact_handles: Some(&replay_artifacts.handles),
970                replay_artifact_inputs: Some(&replay_artifacts.inputs),
971                replay_bundle_id: Some(&replay.bundle.bundle_id),
972                data_envelopes: Some(replay.data_envelopes),
973                prediction_cache_store: replay.prediction_cache_store,
974                prediction_cache_contracts: prediction_cache_contracts.as_ref(),
975                ..Default::default()
976            },
977        )
978    }
979
980    fn execute_phase_scope(
981        &self,
982        plan: &ExecutionPlan,
983        controllers: &RuntimeControllerRegistry,
984        ctx: &mut RunContext,
985        scope: PhaseScope,
986        mut resources: PhaseScopeResources<'_>,
987    ) -> Result<Vec<NodeResult>> {
988        // Hold the phase span on the scheduler thread, and clone it into each
989        // worker so worker-thread telemetry nests under the phase (tracing spans
990        // are thread-local and do not auto-propagate across `thread::scope`).
991        let phase_span = crate::observability::phase_span(
992            ctx.run_id.as_str(),
993            plan.id.as_str(),
994            scope.phase.as_str(),
995            scope.variant_id.as_ref().map(VariantId::as_str),
996            scope.fold_id.as_ref().map(FoldId::as_str),
997        );
998        let _phase_entered = phase_span.clone().entered();
999        // Borrowed for the `thread::scope` below; workers join before it ends.
1000        let plan_id = plan.id.as_str();
1001        plan.validate_parallel_controller_capabilities(self.max_workers, scope.phase)?;
1002        let mut results = Vec::new();
1003        let mut output_handles = BTreeMap::<NodeId, BTreeMap<String, HandleRef>>::new();
1004        let mut output_data_views =
1005            BTreeMap::<NodeId, BTreeMap<String, DataProviderViewSpec>>::new();
1006        let mut input_lineage = BTreeMap::<NodeId, LineageId>::new();
1007
1008        for level in plan.node_parallel_levels_for_phase(scope.phase)? {
1009            let mut prepared = Vec::<PreparedNodeTask>::new();
1010            // Cross-branch merge nodes (concat or late-fusion) are not controller
1011            // tasks: they read the upstream branch OOF blocks from the prediction
1012            // store and reassemble them on the scheduler thread (no worker), AFTER
1013            // this level's worker tasks have populated the store. They are in a
1014            // later level than their branches, so the store already holds the
1015            // branch OOF by the time we reassemble — see `reassemble_branch_merge`.
1016            let mut merge_nodes = Vec::<(NodeId, MergeReduction)>::new();
1017            for node_id in &level {
1018                let node_plan = plan
1019                    .node_plans
1020                    .get(node_id)
1021                    .expect("execution plan was validated");
1022                if let Some(reduction) = merge_reduction_mode(plan, node_plan) {
1023                    merge_nodes.push((node_id.clone(), reduction));
1024                    continue;
1025                }
1026                let collected_inputs = collect_input_handles(
1027                    plan,
1028                    node_plan,
1029                    &output_handles,
1030                    &output_data_views,
1031                    &resources,
1032                    ctx,
1033                    &scope,
1034                )?;
1035                if collected_inputs.skip_node {
1036                    continue;
1037                }
1038                let mut input_handles = collected_inputs.handles;
1039                let mut artifact_inputs = BTreeMap::new();
1040                if let Some(node_artifact_handles) = resources
1041                    .replay_artifact_handles
1042                    .and_then(|handles| handles.get(node_id))
1043                {
1044                    for (key, handle) in node_artifact_handles {
1045                        if input_handles.insert(key.clone(), handle.clone()).is_some() {
1046                            return Err(DagMlError::RuntimeValidation(format!(
1047                                "node `{node_id}` received duplicate replay artifact input `{key}`"
1048                            )));
1049                        }
1050                    }
1051                }
1052                if let Some(node_artifact_inputs) = resources
1053                    .replay_artifact_inputs
1054                    .and_then(|inputs| inputs.get(node_id))
1055                {
1056                    for (key, spec) in node_artifact_inputs {
1057                        if artifact_inputs.insert(key.clone(), spec.clone()).is_some() {
1058                            return Err(DagMlError::RuntimeValidation(format!(
1059                                "node `{node_id}` received duplicate replay artifact metadata `{key}`"
1060                            )));
1061                        }
1062                    }
1063                }
1064                let task_node_plan = effective_node_plan_for_scope(node_plan, &scope)?;
1065                let inner_fold_set = inner_fold_set_for_scope(
1066                    &plan.campaign,
1067                    plan.fold_set.as_ref(),
1068                    node_plan,
1069                    &scope,
1070                )?;
1071                let fit_influence = fit_influence_task_for_node(
1072                    plan,
1073                    &task_node_plan,
1074                    &collected_inputs.data_views,
1075                )?;
1076                prepared.push(PreparedNodeTask {
1077                    node_id: node_id.clone(),
1078                    task: NodeTask {
1079                        inner_fold_set,
1080                        run_id: ctx.run_id.clone(),
1081                        node_plan: task_node_plan.clone(),
1082                        phase: scope.phase,
1083                        variant_id: scope.variant_id.clone(),
1084                        variant: scope.variant.clone(),
1085                        fold_id: scope.fold_id.clone(),
1086                        branch_path: Vec::new(),
1087                        input_handles,
1088                        data_views: collected_inputs.data_views,
1089                        prediction_inputs: collected_inputs.prediction_inputs,
1090                        artifact_inputs,
1091                        required_loss_attestations: NodeTask::required_loss_attestations_for(
1092                            &task_node_plan,
1093                            scope.phase,
1094                        )?,
1095                        fit_influence,
1096                        seed: derive_task_seed(
1097                            scope.seed_root,
1098                            scope.variant_id.as_ref(),
1099                            scope.fold_id.as_ref(),
1100                            &task_node_plan,
1101                            scope.phase,
1102                        ),
1103                    },
1104                });
1105            }
1106
1107            for chunk in prepared.chunks(self.max_workers) {
1108                let chunk_results =
1109                    std::thread::scope(|thread_scope| -> Result<Vec<NodeResult>> {
1110                        let mut handles = Vec::with_capacity(chunk.len());
1111                        for prepared_task in chunk {
1112                            let controller = controllers
1113                                .get(&prepared_task.task.node_plan.controller_id)
1114                                .ok_or_else(|| {
1115                                    DagMlError::RuntimeValidation(format!(
1116                                        "runtime controller `{}` is not registered",
1117                                        prepared_task.task.node_plan.controller_id
1118                                    ))
1119                                })?;
1120                            let worker_span = phase_span.clone();
1121                            handles.push(thread_scope.spawn(move || {
1122                                let _worker_span = worker_span.entered();
1123                                let _node_span = crate::observability::node_span(
1124                                    prepared_task.task.run_id.as_str(),
1125                                    plan_id,
1126                                    prepared_task.task.phase.as_str(),
1127                                    prepared_task.task.node_plan.node_id.as_str(),
1128                                    prepared_task.task.node_plan.controller_id.as_str(),
1129                                )
1130                                .entered();
1131                                let mut result = controller.invoke(&prepared_task.task)?;
1132                                record_fit_influence_diagnostic(&prepared_task.task, &mut result);
1133                                normalize_result_prediction_ports(
1134                                    plan,
1135                                    &prepared_task.task,
1136                                    &mut result,
1137                                )?;
1138                                result.validate_for_task(&prepared_task.task)?;
1139                                Ok(result)
1140                            }));
1141                        }
1142                        handles
1143                            .into_iter()
1144                            .map(|handle| {
1145                                handle.join().map_err(|_| {
1146                                    DagMlError::RuntimeValidation(
1147                                        "parallel scheduler worker panicked".to_string(),
1148                                    )
1149                                })?
1150                            })
1151                            .collect()
1152                    })?;
1153
1154                for (prepared_task, mut result) in chunk.iter().zip(chunk_results) {
1155                    apply_result_prediction_aggregation(
1156                        plan,
1157                        controllers,
1158                        &prepared_task.task,
1159                        &mut result,
1160                        &resources,
1161                    )?;
1162                    attach_coordinator_input_lineage(
1163                        &mut result,
1164                        plan,
1165                        &prepared_task.task.node_plan.node_id,
1166                        &input_lineage,
1167                    )?;
1168                    if let Some(store) = resources.artifact_store.as_deref_mut() {
1169                        if scope.phase == Phase::Refit {
1170                            store.capture_refit_artifacts(&prepared_task.task, &result)?;
1171                        }
1172                    }
1173                    for prediction in &result.predictions {
1174                        ctx.prediction_store.append(prediction.clone())?;
1175                    }
1176                    for prediction in &result.aggregated_predictions {
1177                        ctx.aggregated_prediction_store.append(prediction.clone())?;
1178                    }
1179                    apply_result_scoring(
1180                        &result,
1181                        &mut ctx.score_collector,
1182                        &mut ctx.regression_target_records,
1183                    )?;
1184                    ctx.lineage.record(result.lineage.clone())?;
1185                    let data_views = derive_output_data_views(plan, &prepared_task.task, &result)?;
1186                    output_handles.insert(prepared_task.node_id.clone(), result.outputs.clone());
1187                    output_data_views.insert(prepared_task.node_id.clone(), data_views);
1188                    input_lineage.insert(
1189                        prepared_task.node_id.clone(),
1190                        result.lineage.record_id.clone(),
1191                    );
1192                    results.push(result);
1193                }
1194            }
1195
1196            // Reassemble any cross-branch merge nodes in this level now that the
1197            // level's worker tasks have populated the prediction store. Merge nodes
1198            // sit in a later level than the branches they consume, so the upstream
1199            // branch OOF is already present.
1200            for (node_id, reduction) in &merge_nodes {
1201                let node_plan = plan
1202                    .node_plans
1203                    .get(node_id)
1204                    .expect("execution plan was validated");
1205                if let Some(mut result) =
1206                    reassemble_branch_merge(plan, node_plan, ctx, &scope, *reduction)?
1207                {
1208                    let task_node_plan = effective_node_plan_for_scope(node_plan, &scope)?;
1209                    let task = NodeTask {
1210                        inner_fold_set: None,
1211                        run_id: ctx.run_id.clone(),
1212                        node_plan: task_node_plan.clone(),
1213                        phase: scope.phase,
1214                        variant_id: scope.variant_id.clone(),
1215                        variant: scope.variant.clone(),
1216                        fold_id: scope.fold_id.clone(),
1217                        branch_path: Vec::new(),
1218                        input_handles: BTreeMap::new(),
1219                        data_views: BTreeMap::new(),
1220                        prediction_inputs: BTreeMap::new(),
1221                        artifact_inputs: BTreeMap::new(),
1222                        required_loss_attestations: NodeTask::required_loss_attestations_for(
1223                            &task_node_plan,
1224                            scope.phase,
1225                        )?,
1226                        fit_influence: FitInfluenceTask::default(),
1227                        seed: None,
1228                    };
1229                    normalize_result_prediction_ports(plan, &task, &mut result)?;
1230                    result.validate_for_task(&task)?;
1231                    for prediction in &result.predictions {
1232                        ctx.prediction_store.append(prediction.clone())?;
1233                    }
1234                    apply_result_scoring(
1235                        &result,
1236                        &mut ctx.score_collector,
1237                        &mut ctx.regression_target_records,
1238                    )?;
1239                    ctx.lineage.record(result.lineage.clone())?;
1240                    output_handles.insert(node_id.clone(), result.outputs.clone());
1241                    input_lineage.insert(node_id.clone(), result.lineage.record_id.clone());
1242                    results.push(result);
1243                }
1244            }
1245        }
1246
1247        Ok(results)
1248    }
1249}
1250
1251pub(crate) struct PreparedNodeTask {
1252    pub(crate) node_id: NodeId,
1253    pub(crate) task: NodeTask,
1254}
1255
1256pub(crate) fn attach_coordinator_input_lineage(
1257    result: &mut NodeResult,
1258    plan: &ExecutionPlan,
1259    node_id: &NodeId,
1260    upstream_lineage: &BTreeMap<NodeId, LineageId>,
1261) -> Result<()> {
1262    let inferred = inferred_input_lineage_for_node(plan, node_id, upstream_lineage);
1263    if result.lineage.input_lineage.is_empty() {
1264        result.lineage.input_lineage = inferred;
1265        return Ok(());
1266    }
1267
1268    let declared = result
1269        .lineage
1270        .input_lineage
1271        .iter()
1272        .cloned()
1273        .collect::<BTreeSet<_>>()
1274        .into_iter()
1275        .collect::<Vec<_>>();
1276    if declared != inferred {
1277        return Err(DagMlError::RuntimeValidation(format!(
1278            "lineage for node `{}` declared input lineage {:?}, expected {:?}",
1279            result.node_id, declared, inferred
1280        )));
1281    }
1282    result.lineage.input_lineage = declared;
1283    Ok(())
1284}
1285
1286pub(crate) fn inferred_input_lineage_for_node(
1287    plan: &ExecutionPlan,
1288    node_id: &NodeId,
1289    upstream_lineage: &BTreeMap<NodeId, LineageId>,
1290) -> Vec<LineageId> {
1291    plan.graph_plan
1292        .graph
1293        .edges
1294        .iter()
1295        .filter(|edge| &edge.target.node_id == node_id && edge.contract.propagates_lineage)
1296        .filter_map(|edge| upstream_lineage.get(&edge.source.node_id).cloned())
1297        .collect::<BTreeSet<_>>()
1298        .into_iter()
1299        .collect()
1300}
1301pub(crate) fn collect_input_handles(
1302    plan: &ExecutionPlan,
1303    node_plan: &NodePlan,
1304    output_handles: &BTreeMap<NodeId, BTreeMap<String, HandleRef>>,
1305    output_data_views: &BTreeMap<NodeId, BTreeMap<String, DataProviderViewSpec>>,
1306    resources: &PhaseScopeResources<'_>,
1307    ctx: &RunContext,
1308    scope: &PhaseScope,
1309) -> Result<CollectedInputs> {
1310    let mut inputs = BTreeMap::new();
1311    let mut data_views = BTreeMap::new();
1312    let mut prediction_inputs = BTreeMap::new();
1313    let training_oof_edges = incoming_training_oof_edges(plan, node_plan, scope)?;
1314    // An OOF edge replaces exactly one raw producer port. Do not hide sibling
1315    // outputs from the same producer: a meta-node may legally consume both an
1316    // OOF prediction port and an auxiliary non-OOF port. PREDICT has no
1317    // Validation-OOF input, but its raw prediction port must still be masked so
1318    // only the explicit `:predict` off-fold input reaches the controller.
1319    let masked_oof_source_ports = if scope.phase == Phase::Predict {
1320        incoming_oof_edges(plan, node_plan)?
1321    } else {
1322        training_oof_edges.clone()
1323    }
1324    .into_iter()
1325    .map(|edge| (edge.source.node_id.clone(), edge.source.port_name.clone()))
1326    .collect::<BTreeSet<_>>();
1327    let bound_data_inputs = node_plan
1328        .data_bindings
1329        .iter()
1330        .map(|binding| binding.input_name.clone())
1331        .collect::<BTreeSet<_>>();
1332    // Only forward upstream handles for ports this node DECLARES an edge to.
1333    // A controller must never see a handle outside its declared port contract,
1334    // so a sibling consumer of the same producer cannot expose extra ports here.
1335    let declared_source_ports = plan
1336        .graph_plan
1337        .graph
1338        .edges
1339        .iter()
1340        .filter(|edge| edge.target.node_id == node_plan.node_id)
1341        .map(|edge| (edge.source.node_id.clone(), edge.source.port_name.clone()))
1342        .collect::<BTreeSet<_>>();
1343    for upstream in &node_plan.input_nodes {
1344        if let Some(handles) = output_handles.get(upstream) {
1345            for (port, handle) in handles {
1346                if !declared_source_ports.contains(&(upstream.clone(), port.clone())) {
1347                    continue;
1348                }
1349                if masked_oof_source_ports.contains(&(upstream.clone(), port.clone())) {
1350                    continue;
1351                }
1352                inputs.insert(format!("{upstream}.{port}"), handle.clone());
1353            }
1354        }
1355    }
1356    for edge in plan
1357        .graph_plan
1358        .graph
1359        .edges
1360        .iter()
1361        .filter(|edge| edge.target.node_id == node_plan.node_id)
1362        .filter(|edge| edge.contract.kind == PortKind::Data && !edge.contract.requires_oof)
1363    {
1364        if bound_data_inputs.contains(&edge.target.port_name) {
1365            continue;
1366        }
1367        let Some(handles) = output_handles.get(&edge.source.node_id) else {
1368            continue;
1369        };
1370        let Some(handle) = handles.get(&edge.source.port_name) else {
1371            continue;
1372        };
1373        let key = data_view_key(&edge.target.port_name);
1374        if inputs.insert(key.clone(), handle.clone()).is_some() {
1375            return Err(DagMlError::RuntimeValidation(format!(
1376                "node `{}` received duplicate data edge input `{key}`",
1377                node_plan.node_id
1378            )));
1379        }
1380        if let Some(source_views) = output_data_views.get(&edge.source.node_id) {
1381            if let Some(view) = source_views.get(&edge.source.port_name) {
1382                if data_views.insert(key.clone(), view.clone()).is_some() {
1383                    return Err(DagMlError::RuntimeValidation(format!(
1384                        "node `{}` received duplicate data edge view `{key}`",
1385                        node_plan.node_id
1386                    )));
1387                }
1388            }
1389            let source_validation_key = validation_data_view_key(&edge.source.port_name);
1390            if let Some(view) = source_views.get(&source_validation_key) {
1391                let validation_key = format!("{key}:validation");
1392                if data_views
1393                    .insert(validation_key.clone(), view.clone())
1394                    .is_some()
1395                {
1396                    return Err(DagMlError::RuntimeValidation(format!(
1397                        "node `{}` received duplicate data edge validation view `{validation_key}`",
1398                        node_plan.node_id
1399                    )));
1400                }
1401            }
1402        }
1403    }
1404    for edge in training_oof_edges {
1405        let key = format!("{}.{}", edge.source.node_id, edge.source.port_name);
1406        let Some(input) = collect_oof_prediction_input(plan, edge, ctx, scope, resources)? else {
1407            return Ok(CollectedInputs {
1408                handles: BTreeMap::new(),
1409                data_views: BTreeMap::new(),
1410                prediction_inputs: BTreeMap::new(),
1411                skip_node: true,
1412            });
1413        };
1414        if inputs.insert(key.clone(), input.handle).is_some() {
1415            return Err(DagMlError::RuntimeValidation(format!(
1416                "node `{}` received duplicate OOF prediction input `{key}`",
1417                node_plan.node_id
1418            )));
1419        }
1420        if prediction_inputs.insert(key.clone(), input.spec).is_some() {
1421            return Err(DagMlError::RuntimeValidation(format!(
1422                "node `{}` received duplicate OOF prediction spec `{key}`",
1423                node_plan.node_id
1424            )));
1425        }
1426    }
1427    // REFIT / PREDICT: deliver each base producer's off-fold (test / predict)
1428    // predictions to the stacking meta-node as a SEPARATE prediction input (suffixed
1429    // `:test` / `:predict`) so the host meta-model predicts from them. The FIT_CV
1430    // Validation-OOF input above is the meta-features the meta-model trains on; this
1431    // off-fold input is used ONLY for REFIT/PREDICT scoring/prediction, never FIT_CV
1432    // training — keeping the leakage invariant intact.
1433    if matches!(scope.phase, Phase::Refit | Phase::Predict) {
1434        let off_fold_suffix = scope.phase.as_str().to_ascii_lowercase();
1435        for edge in incoming_oof_edges(plan, node_plan)? {
1436            let Some(input) = collect_off_fold_prediction_input(plan, edge, ctx, scope)? else {
1437                continue;
1438            };
1439            let key = format!(
1440                "{}.{}:{off_fold_suffix}",
1441                edge.source.node_id, edge.source.port_name
1442            );
1443            if inputs.insert(key.clone(), input.handle).is_some() {
1444                return Err(DagMlError::RuntimeValidation(format!(
1445                    "node `{}` received duplicate off-fold prediction input `{key}`",
1446                    node_plan.node_id
1447                )));
1448            }
1449            if prediction_inputs.insert(key.clone(), input.spec).is_some() {
1450                return Err(DagMlError::RuntimeValidation(format!(
1451                    "node `{}` received duplicate off-fold prediction spec `{key}`",
1452                    node_plan.node_id
1453                )));
1454            }
1455        }
1456    }
1457    if !node_plan.data_bindings.is_empty() && resources.data_provider.is_none() {
1458        return Err(DagMlError::RuntimeValidation(format!(
1459            "node `{}` requires {} data binding(s) but no runtime data provider is registered",
1460            node_plan.node_id,
1461            node_plan.data_bindings.len()
1462        )));
1463    }
1464    if let Some(data_provider) = resources.data_provider {
1465        // Samples excluded from training (sample-local) for this node, derived
1466        // from its coordinator relations. Used to filter FIT view specs so the
1467        // spec, the materialized view, and fit-influence row_weights agree.
1468        let excluded_samples = coordinator_relations_for_node(node_plan, resources)?
1469            .map(|relations| relations.excluded_sample_ids())
1470            .unwrap_or_default();
1471        for binding in &node_plan.data_bindings {
1472            let materialized = data_provider.materialize(&DataMaterializationRequest {
1473                run_id: ctx.run_id.clone(),
1474                node_id: node_plan.node_id.clone(),
1475                input_name: binding.input_name.clone(),
1476                phase: scope.phase,
1477                variant_id: scope.variant_id.clone(),
1478                fold_id: scope.fold_id.clone(),
1479                binding: binding.clone(),
1480            })?;
1481            let branch_view_for_node = branch_view_from_node_metadata(plan, &node_plan.node_id)?;
1482            let view = data_view_for_scope(
1483                binding,
1484                plan.fold_set.as_ref(),
1485                scope,
1486                branch_view_for_node.as_ref(),
1487                &excluded_samples,
1488            )?;
1489            let key = data_view_key(&binding.input_name);
1490            let view_handle = make_data_view_handle(
1491                data_provider,
1492                ctx,
1493                node_plan,
1494                scope,
1495                binding,
1496                &materialized,
1497                &view,
1498            )?;
1499            if data_views.insert(key.clone(), view).is_some() {
1500                return Err(DagMlError::RuntimeValidation(format!(
1501                    "node `{}` received duplicate data view `{key}`",
1502                    node_plan.node_id
1503                )));
1504            }
1505            if inputs.insert(key.clone(), view_handle).is_some() {
1506                return Err(DagMlError::RuntimeValidation(format!(
1507                    "node `{}` received duplicate data input `{key}`",
1508                    node_plan.node_id
1509                )));
1510            }
1511
1512            if let Some(validation_view) = validation_data_view_for_scope(
1513                binding,
1514                plan.fold_set.as_ref(),
1515                scope,
1516                branch_view_for_node.as_ref(),
1517                &excluded_samples,
1518            )? {
1519                let validation_key = format!("{key}:validation");
1520                let validation_handle = make_data_view_handle(
1521                    data_provider,
1522                    ctx,
1523                    node_plan,
1524                    scope,
1525                    binding,
1526                    &materialized,
1527                    &validation_view,
1528                )?;
1529                if data_views
1530                    .insert(validation_key.clone(), validation_view)
1531                    .is_some()
1532                {
1533                    return Err(DagMlError::RuntimeValidation(format!(
1534                        "node `{}` received duplicate validation data view `{validation_key}`",
1535                        node_plan.node_id
1536                    )));
1537                }
1538                if inputs
1539                    .insert(validation_key.clone(), validation_handle)
1540                    .is_some()
1541                {
1542                    return Err(DagMlError::RuntimeValidation(format!(
1543                        "node `{}` received duplicate validation data input `{validation_key}`",
1544                        node_plan.node_id
1545                    )));
1546                }
1547            }
1548        }
1549    }
1550    Ok(CollectedInputs {
1551        handles: inputs,
1552        data_views,
1553        prediction_inputs,
1554        skip_node: false,
1555    })
1556}
1557pub(crate) fn preload_replay_prediction_cache_store(
1558    bundle: &ExecutionBundle,
1559    prediction_cache_store: Option<&dyn RuntimePredictionCacheStore>,
1560    ctx: &mut RunContext,
1561) -> Result<()> {
1562    if bundle.prediction_requirements.is_empty() {
1563        return Ok(());
1564    }
1565    let store = prediction_cache_store.ok_or_else(|| {
1566        DagMlError::RuntimeValidation(format!(
1567            "bundle `{}` cannot preload OOF prediction caches without a prediction cache store",
1568            bundle.bundle_id
1569        ))
1570    })?;
1571    if !ctx.prediction_store.blocks().is_empty() {
1572        return Err(DagMlError::RuntimeValidation(format!(
1573            "bundle `{}` cannot preload OOF prediction caches into a non-empty prediction store",
1574            bundle.bundle_id
1575        )));
1576    }
1577    let contracts = replay_prediction_cache_contracts(bundle)?;
1578    for contract in contracts.values() {
1579        if contract.requirement.prediction_level == PredictionLevel::Sample {
1580            let blocks = store.load_blocks(&contract.cache.requirement_key)?;
1581            if blocks.iter().any(|block| {
1582                block.producer_node != contract.requirement.producer_node
1583                    || block.partition != contract.requirement.partition
1584            }) {
1585                return Err(DagMlError::RuntimeValidation(format!(
1586                    "prediction cache store returned blocks outside requirement `{}`",
1587                    contract.cache.requirement_key
1588                )));
1589            }
1590            let mut payload = build_prediction_cache_payload(&contract.requirement, &blocks)?;
1591            payload.cache_namespace_fingerprints =
1592                contract.cache.cache_namespace_fingerprints.clone();
1593            validate_prediction_cache_payload_matches_record(&payload, &contract.cache)?;
1594            for block in &payload.blocks {
1595                ctx.prediction_store.append(block.clone())?;
1596            }
1597        } else {
1598            let blocks = store.load_aggregated_blocks(&contract.cache.requirement_key)?;
1599            if blocks.iter().any(|block| {
1600                block.producer_node != contract.requirement.producer_node
1601                    || block.partition != contract.requirement.partition
1602                    || block.level != contract.requirement.prediction_level
1603            }) {
1604                return Err(DagMlError::RuntimeValidation(format!(
1605                    "prediction cache store returned aggregated blocks outside requirement `{}`",
1606                    contract.cache.requirement_key
1607                )));
1608            }
1609            let mut payload =
1610                build_aggregated_prediction_cache_payload(&contract.requirement, &blocks)?;
1611            payload.cache_namespace_fingerprints =
1612                contract.cache.cache_namespace_fingerprints.clone();
1613            validate_prediction_cache_payload_matches_record(&payload, &contract.cache)?;
1614        }
1615    }
1616    Ok(())
1617}
1618
1619pub(crate) fn replay_prediction_cache_contracts(
1620    bundle: &ExecutionBundle,
1621) -> Result<BTreeMap<String, ReplayPredictionCacheContract>> {
1622    bundle.validate()?;
1623    let requirements = bundle
1624        .prediction_requirements
1625        .iter()
1626        .map(|requirement| (requirement.key(), requirement))
1627        .collect::<BTreeMap<_, _>>();
1628    let mut contracts = BTreeMap::new();
1629    for cache in &bundle.prediction_caches {
1630        let requirement = requirements.get(&cache.requirement_key).ok_or_else(|| {
1631            DagMlError::RuntimeValidation(format!(
1632                "prediction cache `{}` references unknown prediction requirement `{}`",
1633                cache.cache_id, cache.requirement_key
1634            ))
1635        })?;
1636        contracts.insert(
1637            cache.requirement_key.clone(),
1638            ReplayPredictionCacheContract {
1639                requirement: (*requirement).clone(),
1640                cache: cache.clone(),
1641            },
1642        );
1643    }
1644    Ok(contracts)
1645}
1646
1647pub(crate) fn materialize_replay_artifact_handles(
1648    plan: &ExecutionPlan,
1649    bundle: &ExecutionBundle,
1650    replay_request: &ReplayPhaseRequest,
1651    artifact_store: &dyn RuntimeArtifactStore,
1652    ctx: &RunContext,
1653) -> Result<MaterializedReplayArtifacts> {
1654    let mut handles = BTreeMap::<NodeId, BTreeMap<String, HandleRef>>::new();
1655    let mut inputs = BTreeMap::<NodeId, BTreeMap<String, ArtifactInputSpec>>::new();
1656    for artifact in &bundle.refit_artifacts {
1657        artifact.validate()?;
1658        let node_plan = plan.node_plans.get(&artifact.node_id).ok_or_else(|| {
1659            DagMlError::RuntimeValidation(format!(
1660                "bundle `{}` artifact references unknown node `{}`",
1661                bundle.bundle_id, artifact.node_id
1662            ))
1663        })?;
1664        if !node_plan.supported_phases.contains(&replay_request.phase) {
1665            return Err(DagMlError::RuntimeValidation(format!(
1666                "bundle `{}` artifact node `{}` does not support replay phase {:?}",
1667                bundle.bundle_id, artifact.node_id, replay_request.phase
1668            )));
1669        }
1670        let handle = artifact_store.materialize(&ArtifactMaterializationRequest {
1671            run_id: ctx.run_id.clone(),
1672            bundle_id: bundle.bundle_id.clone(),
1673            node_id: artifact.node_id.clone(),
1674            phase: replay_request.phase,
1675            variant_id: bundle.selected_variant_id.clone(),
1676            controller_id: artifact.controller_id.clone(),
1677            artifact: artifact.artifact.clone(),
1678            params_fingerprint: artifact.params_fingerprint.clone(),
1679            training_loss_fingerprint: artifact.training_loss_fingerprint.clone(),
1680        })?;
1681        if !matches!(handle.kind, HandleKind::Model | HandleKind::Artifact) {
1682            return Err(DagMlError::RuntimeValidation(format!(
1683                "artifact `{}` materialized as unsupported handle kind {:?}",
1684                artifact.artifact.id, handle.kind
1685            )));
1686        }
1687        if handle.owner_controller != artifact.controller_id {
1688            return Err(DagMlError::RuntimeValidation(format!(
1689                "artifact `{}` handle owner `{}` does not match controller `{}`",
1690                artifact.artifact.id, handle.owner_controller, artifact.controller_id
1691            )));
1692        }
1693        let key = refit_artifact_input_key(&artifact.artifact.id);
1694        if handles
1695            .entry(artifact.node_id.clone())
1696            .or_default()
1697            .insert(key.clone(), handle)
1698            .is_some()
1699        {
1700            return Err(DagMlError::RuntimeValidation(format!(
1701                "duplicate replay artifact input `{key}` for node `{}`",
1702                artifact.node_id
1703            )));
1704        }
1705        if inputs
1706            .entry(artifact.node_id.clone())
1707            .or_default()
1708            .insert(key.clone(), ArtifactInputSpec::from_refit_record(artifact)?)
1709            .is_some()
1710        {
1711            return Err(DagMlError::RuntimeValidation(format!(
1712                "duplicate replay artifact metadata `{key}` for node `{}`",
1713                artifact.node_id
1714            )));
1715        }
1716    }
1717    Ok(MaterializedReplayArtifacts { handles, inputs })
1718}
1719
1720pub(crate) fn derive_task_seed(
1721    root_seed: Option<u64>,
1722    variant_id: Option<&VariantId>,
1723    fold_id: Option<&FoldId>,
1724    node_plan: &NodePlan,
1725    phase: Phase,
1726) -> Option<u64> {
1727    root_seed.map(|root| {
1728        let mut context = SeedContext::root(root);
1729        if let Some(variant_id) = variant_id {
1730            context = context.child(format!("variant:{variant_id}"));
1731        }
1732        if let Some(fold_id) = fold_id {
1733            context = context.child(format!("fold:{fold_id}"));
1734        }
1735        context
1736            .child(format!("node:{}", node_plan.node_id))
1737            .child(format!("phase:{phase:?}"))
1738            .derive_u64("task")
1739    })
1740}