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    /// Run one local tuner session and evaluate every proposal through the
170    /// ordinary FIT_CV scheduler.  The session remains on this thread; only a
171    /// portable [`RuntimeHpoProposal`] and OOF-derived scalar feedback cross
172    /// the controller boundary. SELECT and REFIT deliberately do not occur
173    /// here, so callers can make exactly one selection and one refit after the
174    /// returned report-grade candidate evidence has been audited.
175    pub fn execute_hpo_campaign(
176        &self,
177        plan: &ExecutionPlan,
178        controllers: &RuntimeControllerRegistry,
179        data_provider: &dyn RuntimeDataProvider,
180        ctx: &RunContext,
181        hpo: &RuntimeHpoExecutionContext,
182    ) -> Result<RuntimeHpoCampaignResult> {
183        plan.validate()?;
184        hpo.validate_for_plan(plan)?;
185        let controller = controllers.get(&hpo.controller_id).ok_or_else(|| {
186            DagMlError::RuntimeValidation(format!(
187                "runtime HPO campaign controller `{}` is not registered",
188                hpo.controller_id
189            ))
190        })?;
191        let task = RuntimeHpoCampaignTask {
192            run_id: ctx.run_id.clone(),
193            operation_id: hpo.operation_id.clone(),
194            controller_id: hpo.controller_id.clone(),
195            target_node_id: hpo.target_node_id.clone(),
196            seed: ctx.root_seed,
197        };
198        let mut session = controller.create_tuner_session(&task, hpo)?;
199        let history_at_start = session.trial_history_len()?;
200        if history_at_start > hpo.trial_budget_total {
201            return Err(DagMlError::RuntimeValidation(format!(
202                "runtime HPO restored native history ({history_at_start}) exceeds total trial budget ({})",
203                hpo.trial_budget_total
204            )));
205        }
206        let remaining_trials = hpo.trial_budget_total - history_at_start;
207        let mut candidates = Vec::new();
208        let mut proposed_variant_ids = BTreeSet::new();
209        // Fresh proposals are checkpointed by this call.  The native study can
210        // nevertheless retain an incumbent from a restored terminal trial, so
211        // keep its persisted trial->variant binding separate from the new
212        // checkpoint evidence and extend it as we ask new trials.
213        let mut trial_variants = BTreeMap::new();
214        let mut incumbent_variants = hpo.resume_variants.clone();
215        let mut terminal_trials = BTreeMap::new();
216        let mut completed_proposals = Vec::new();
217        let mut completed_reports = Vec::new();
218
219        for _ in 0..remaining_trials {
220            let Some(proposal) = session.ask()? else {
221                break;
222            };
223            if trial_variants
224                .insert(proposal.trial_id, proposal.variant.variant_id.clone())
225                .is_some()
226            {
227                return Err(DagMlError::RuntimeValidation(format!(
228                    "runtime HPO session proposed duplicate trial `{}`",
229                    proposal.trial_id
230                )));
231            }
232            if incumbent_variants
233                .insert(proposal.trial_id, proposal.variant.variant_id.clone())
234                .is_some()
235            {
236                return Err(DagMlError::RuntimeValidation(format!(
237                    "runtime HPO session reused restored trial `{}`",
238                    proposal.trial_id
239                )));
240            }
241            if !proposed_variant_ids.insert(proposal.variant.variant_id.clone()) {
242                return Err(DagMlError::RuntimeValidation(format!(
243                    "runtime HPO session proposed duplicate variant `{}`",
244                    proposal.variant.variant_id
245                )));
246            }
247            let mut candidate_plan = plan.clone();
248            candidate_plan.variants = vec![proposal.variant.clone()];
249            candidate_plan.validate()?;
250            let mut candidate_ctx =
251                RunContext::new(ctx.run_id.clone(), proposal.variant.seed.or(ctx.root_seed));
252            candidate_ctx.variant_id = Some(proposal.variant.variant_id.clone());
253
254            let evaluation = self.execute_hpo_candidate_fit_cv(
255                &candidate_plan,
256                controllers,
257                data_provider,
258                &mut candidate_ctx,
259            );
260            if let Err(error) = evaluation {
261                session.tell(
262                    proposal.trial_id,
263                    RuntimeHpoTerminal::Failed {
264                        failure: RuntimeHpoFailure {
265                            code: "DAGML_CV_ERROR".to_string(),
266                            message: error.to_string(),
267                            retryable: false,
268                        },
269                    },
270                )?;
271                terminal_trials.insert(proposal.trial_id, HpoTrialTerminalState::Failed);
272                continue;
273            }
274            if let Err(error) = candidate_ctx
275                .collect_cross_fold_validation_scores(plan_oof_partition_mode(&candidate_plan))
276            {
277                session.tell(
278                    proposal.trial_id,
279                    RuntimeHpoTerminal::Failed {
280                        failure: RuntimeHpoFailure {
281                            code: "DAGML_SCORE_ERROR".to_string(),
282                            message: error.to_string(),
283                            retryable: false,
284                        },
285                    },
286                )?;
287                terminal_trials.insert(proposal.trial_id, HpoTrialTerminalState::Failed);
288                continue;
289            }
290            let report = candidate_ctx
291                .score_collector
292                .iter()
293                .find(|report| {
294                    report.producer_node == hpo.selection.producer_node
295                        && report.producer_port.as_deref()
296                            == Some(hpo.selection.producer_port.as_str())
297                        && report.partition == PredictionPartition::Validation
298                        && report
299                            .fold_id
300                            .as_ref()
301                            .is_some_and(|fold| fold.as_str() == "avg")
302                })
303                .cloned();
304            let Some(mut report) = report else {
305                session.tell(
306                    proposal.trial_id,
307                    RuntimeHpoTerminal::Failed {
308                        failure: RuntimeHpoFailure {
309                            code: "DAGML_SCORE_MISSING".to_string(),
310                            message: format!(
311                                "runtime HPO trial `{}` emitted no target OOF average",
312                                proposal.trial_id
313                            ),
314                            retryable: false,
315                        },
316                    },
317                )?;
318                terminal_trials.insert(proposal.trial_id, HpoTrialTerminalState::Failed);
319                continue;
320            };
321            report.variant_id = Some(proposal.variant.variant_id.clone());
322            let score = report
323                .metrics
324                .get(hpo.selection.metric.name())
325                .copied()
326                .filter(|score| score.is_finite());
327            let Some(score) = score else {
328                session.tell(
329                    proposal.trial_id,
330                    RuntimeHpoTerminal::Failed {
331                        failure: RuntimeHpoFailure {
332                            code: "DAGML_SCORE_NONFINITE".to_string(),
333                            message: format!(
334                                "runtime HPO trial `{}` emitted no finite `{}` score",
335                                proposal.trial_id,
336                                hpo.selection.metric.name()
337                            ),
338                            retryable: false,
339                        },
340                    },
341                )?;
342                terminal_trials.insert(proposal.trial_id, HpoTrialTerminalState::Failed);
343                continue;
344            };
345            let intermediate = RuntimeHpoIntermediate {
346                trial_id: proposal.trial_id,
347                step: 0,
348                score,
349            };
350            if session.report_intermediate(intermediate)? == RuntimeHpoIntermediateOutcome::Pruned {
351                terminal_trials.insert(proposal.trial_id, HpoTrialTerminalState::Pruned);
352                continue;
353            }
354            session.tell(proposal.trial_id, RuntimeHpoTerminal::Completed { score })?;
355            terminal_trials.insert(proposal.trial_id, HpoTrialTerminalState::Completed);
356            completed_proposals.push(proposal.clone());
357            completed_reports.push(RuntimeHpoCompletedReport {
358                trial_id: proposal.trial_id,
359                variant_id: proposal.variant.variant_id.clone(),
360                report: report.clone(),
361            });
362
363            let mut validation_reports = candidate_ctx
364                .score_collector
365                .iter()
366                .filter(|item| item.partition == PredictionPartition::Validation)
367                .cloned()
368                .collect::<Vec<_>>();
369            for item in &mut validation_reports {
370                item.variant_id = Some(proposal.variant.variant_id.clone());
371            }
372            candidates.push(RuntimeHpoCandidateEvaluation {
373                validation_predictions: capture_variant_validation_predictions(
374                    &proposal.variant.variant_id,
375                    None,
376                    &candidate_ctx,
377                ),
378                lineage: candidate_ctx.lineage.records().cloned().collect(),
379                proposal,
380                score,
381                validation_reports,
382            });
383        }
384
385        let history_at_checkpoint = session.trial_history_len()?;
386        if history_at_checkpoint != hpo.trial_budget_total {
387            return Err(DagMlError::RuntimeValidation(format!(
388                "runtime HPO native history ended at {history_at_checkpoint}, expected total trial budget {}",
389                hpo.trial_budget_total
390            )));
391        }
392
393        let checkpoint = RuntimeHpoCheckpointResult {
394            artifact: session.checkpoint()?,
395            provenance: hpo.provenance.clone(),
396            operation_id: hpo.operation_id.clone(),
397            controller_id: hpo.controller_id.clone(),
398            target_node_id: hpo.target_node_id.clone(),
399            completed_proposals,
400            completed_reports,
401            trial_history_len: history_at_checkpoint,
402        };
403        validate_hpo_checkpoint_result(
404            &checkpoint,
405            hpo,
406            &trial_variants,
407            &terminal_trials,
408            history_at_start,
409        )?;
410        let incumbent = session.incumbent(&incumbent_variants)?.ok_or_else(|| {
411            DagMlError::RuntimeValidation(
412                "native HPO campaign has no completed native incumbent after terminalization"
413                    .to_string(),
414            )
415        })?;
416        if incumbent.metric != hpo.selection.metric.name()
417            || incumbent.direction != hpo.selection.direction
418            || incumbent_variants.get(&incumbent.trial_id) != Some(&incumbent.variant_id)
419            || !incumbent.score.is_finite()
420        {
421            return Err(DagMlError::RuntimeValidation(
422                "native HPO incumbent is not bound to this scheduler campaign's metric, direction, trial, and variant"
423                    .to_string(),
424            ));
425        }
426        let terminal_trials = session.terminal_trial_snapshots(&incumbent_variants)?;
427        if terminal_trials.len() != history_at_checkpoint as usize
428            || terminal_trials
429                .windows(2)
430                .any(|pair| pair[0].trial.id >= pair[1].trial.id)
431        {
432            return Err(DagMlError::RuntimeValidation(
433                "native HPO terminal ledger is not a complete strictly ordered history".to_string(),
434            ));
435        }
436        Ok(RuntimeHpoCampaignResult {
437            operation_id: hpo.operation_id.clone(),
438            controller_id: hpo.controller_id.clone(),
439            target_node_id: hpo.target_node_id.clone(),
440            candidates,
441            checkpoint,
442            incumbent,
443            terminal_trials,
444        })
445    }
446
447    fn execute_hpo_candidate_fit_cv(
448        &self,
449        plan: &ExecutionPlan,
450        controllers: &RuntimeControllerRegistry,
451        data_provider: &dyn RuntimeDataProvider,
452        ctx: &mut RunContext,
453    ) -> Result<Vec<NodeResult>> {
454        let candidate_plan = plan;
455        ctx.configure_global_oof_aggregation(candidate_plan, data_provider)?;
456        let fold_ids = candidate_plan
457            .fold_set
458            .as_ref()
459            .map(|fold_set| {
460                fold_set
461                    .folds
462                    .iter()
463                    .map(|fold| Some(fold.fold_id.clone()))
464                    .collect::<Vec<_>>()
465            })
466            .unwrap_or_else(|| vec![None]);
467        let variant = candidate_plan
468            .variants
469            .first()
470            .expect("candidate plan has exactly one variant");
471        let mut results = Vec::new();
472        for fold_id in fold_ids {
473            results.extend(self.execute_phase_scope(
474                candidate_plan,
475                controllers,
476                ctx,
477                PhaseScope {
478                    phase: Phase::FitCv,
479                    variant_id: Some(variant.variant_id.clone()),
480                    variant: Some(VariantExecutionSpec::from_plan(variant)),
481                    fold_id,
482                    seed_root: variant.seed.or(ctx.root_seed),
483                },
484                PhaseScopeResources {
485                    data_provider: Some(data_provider),
486                    ..Default::default()
487                },
488            )?);
489        }
490        Ok(results)
491    }
492
493    pub fn execute_phase(
494        &self,
495        plan: &ExecutionPlan,
496        controllers: &RuntimeControllerRegistry,
497        ctx: &mut RunContext,
498        phase: Phase,
499    ) -> Result<Vec<NodeResult>> {
500        plan.validate()?;
501        let variant_id = ctx.variant_id.clone();
502        let seed_root = ctx.root_seed;
503        self.execute_phase_scope(
504            plan,
505            controllers,
506            ctx,
507            PhaseScope {
508                phase,
509                variant_id,
510                variant: None,
511                fold_id: None,
512                seed_root,
513            },
514            PhaseScopeResources::default(),
515        )
516    }
517
518    pub fn execute_phase_with_data_provider(
519        &self,
520        plan: &ExecutionPlan,
521        controllers: &RuntimeControllerRegistry,
522        data_provider: &dyn RuntimeDataProvider,
523        ctx: &mut RunContext,
524        phase: Phase,
525    ) -> Result<Vec<NodeResult>> {
526        plan.validate()?;
527        let variant_id = ctx.variant_id.clone();
528        let seed_root = ctx.root_seed;
529        self.execute_phase_scope(
530            plan,
531            controllers,
532            ctx,
533            PhaseScope {
534                phase,
535                variant_id,
536                variant: None,
537                fold_id: None,
538                seed_root,
539            },
540            PhaseScopeResources {
541                data_provider: Some(data_provider),
542                ..Default::default()
543            },
544        )
545    }
546
547    pub fn execute_campaign_phase(
548        &self,
549        plan: &ExecutionPlan,
550        controllers: &RuntimeControllerRegistry,
551        ctx: &mut RunContext,
552        phase: Phase,
553    ) -> Result<Vec<NodeResult>> {
554        plan.validate()?;
555        let mut results = Vec::new();
556        let fold_ids = if phase == Phase::FitCv {
557            plan.fold_set
558                .as_ref()
559                .map(|fold_set| {
560                    fold_set
561                        .folds
562                        .iter()
563                        .map(|fold| Some(fold.fold_id.clone()))
564                        .collect::<Vec<_>>()
565                })
566                .unwrap_or_else(|| vec![None])
567        } else {
568            vec![None]
569        };
570        for variant in &plan.variants {
571            if ctx
572                .variant_id
573                .as_ref()
574                .is_some_and(|requested| requested != &variant.variant_id)
575            {
576                continue;
577            }
578            for fold_id in &fold_ids {
579                let seed_root = variant.seed.or(ctx.root_seed);
580                results.extend(self.execute_phase_scope(
581                    plan,
582                    controllers,
583                    ctx,
584                    PhaseScope {
585                        phase,
586                        variant_id: Some(variant.variant_id.clone()),
587                        variant: Some(VariantExecutionSpec::from_plan(variant)),
588                        fold_id: fold_id.clone(),
589                        seed_root,
590                    },
591                    PhaseScopeResources::default(),
592                )?);
593            }
594        }
595        Ok(results)
596    }
597
598    pub fn execute_campaign_phase_with_data_provider(
599        &self,
600        plan: &ExecutionPlan,
601        controllers: &RuntimeControllerRegistry,
602        data_provider: &dyn RuntimeDataProvider,
603        ctx: &mut RunContext,
604        phase: Phase,
605    ) -> Result<Vec<NodeResult>> {
606        plan.validate()?;
607        if phase == Phase::FitCv {
608            ctx.configure_global_oof_aggregation(plan, data_provider)?;
609        }
610        let mut results = Vec::new();
611        let fold_ids = if phase == Phase::FitCv {
612            plan.fold_set
613                .as_ref()
614                .map(|fold_set| {
615                    fold_set
616                        .folds
617                        .iter()
618                        .map(|fold| Some(fold.fold_id.clone()))
619                        .collect::<Vec<_>>()
620                })
621                .unwrap_or_else(|| vec![None])
622        } else {
623            vec![None]
624        };
625        for variant in &plan.variants {
626            if ctx
627                .variant_id
628                .as_ref()
629                .is_some_and(|requested| requested != &variant.variant_id)
630            {
631                continue;
632            }
633            for fold_id in &fold_ids {
634                let seed_root = variant.seed.or(ctx.root_seed);
635                results.extend(self.execute_phase_scope(
636                    plan,
637                    controllers,
638                    ctx,
639                    PhaseScope {
640                        phase,
641                        variant_id: Some(variant.variant_id.clone()),
642                        variant: Some(VariantExecutionSpec::from_plan(variant)),
643                        fold_id: fold_id.clone(),
644                        seed_root,
645                    },
646                    PhaseScopeResources {
647                        data_provider: Some(data_provider),
648                        ..Default::default()
649                    },
650                )?);
651            }
652        }
653        Ok(results)
654    }
655
656    pub fn execute_campaign_phase_with_data_provider_and_artifact_store(
657        &self,
658        plan: &ExecutionPlan,
659        controllers: &RuntimeControllerRegistry,
660        data_provider: &dyn RuntimeDataProvider,
661        artifact_store: &mut InMemoryArtifactStore,
662        ctx: &mut RunContext,
663        phase: Phase,
664    ) -> Result<Vec<NodeResult>> {
665        plan.validate()?;
666        if phase == Phase::FitCv {
667            ctx.configure_global_oof_aggregation(plan, data_provider)?;
668        }
669        let mut results = Vec::new();
670        let fold_ids = if phase == Phase::FitCv {
671            plan.fold_set
672                .as_ref()
673                .map(|fold_set| {
674                    fold_set
675                        .folds
676                        .iter()
677                        .map(|fold| Some(fold.fold_id.clone()))
678                        .collect::<Vec<_>>()
679                })
680                .unwrap_or_else(|| vec![None])
681        } else {
682            vec![None]
683        };
684        for variant in &plan.variants {
685            if ctx
686                .variant_id
687                .as_ref()
688                .is_some_and(|requested| requested != &variant.variant_id)
689            {
690                continue;
691            }
692            for fold_id in &fold_ids {
693                let seed_root = variant.seed.or(ctx.root_seed);
694                results.extend(self.execute_phase_scope(
695                    plan,
696                    controllers,
697                    ctx,
698                    PhaseScope {
699                        phase,
700                        variant_id: Some(variant.variant_id.clone()),
701                        variant: Some(VariantExecutionSpec::from_plan(variant)),
702                        fold_id: fold_id.clone(),
703                        seed_root,
704                    },
705                    PhaseScopeResources {
706                        data_provider: Some(data_provider),
707                        artifact_store: Some(&mut *artifact_store),
708                        ..Default::default()
709                    },
710                )?);
711            }
712        }
713        Ok(results)
714    }
715
716    pub fn execute_bundle_replay(
717        &self,
718        replay: BundleReplayExecution<'_>,
719        ctx: &mut RunContext,
720    ) -> Result<Vec<NodeResult>> {
721        replay.bundle.validate_against_plan(replay.plan)?;
722        replay
723            .replay_request
724            .validate_for_bundle_with_prediction_cache_store(
725                replay.bundle,
726                replay.prediction_cache_store.is_some(),
727            )?;
728        replay
729            .bundle
730            .validate_replay_envelopes(replay.data_envelopes)?;
731        let prediction_cache_contracts = if replay.replay_request.phase == Phase::Refit {
732            Some(replay_prediction_cache_contracts(replay.bundle)?)
733        } else {
734            None
735        };
736        if replay.replay_request.phase == Phase::Refit {
737            preload_replay_prediction_cache_store(
738                replay.bundle,
739                replay.prediction_cache_store,
740                ctx,
741            )?;
742        }
743        let replay_artifacts = materialize_replay_artifact_handles(
744            replay.plan,
745            replay.bundle,
746            replay.replay_request,
747            replay.artifact_store,
748            ctx,
749        )?;
750        let selected_variant = replay
751            .bundle
752            .selected_variant_id
753            .as_ref()
754            .map(|selected| {
755                replay
756                    .plan
757                    .variants
758                    .iter()
759                    .find(|variant| &variant.variant_id == selected)
760                    .map(VariantExecutionSpec::from_plan)
761                    .ok_or_else(|| {
762                        DagMlError::RuntimeValidation(format!(
763                            "bundle `{}` selected unknown variant `{selected}`",
764                            replay.bundle.bundle_id
765                        ))
766                    })
767            })
768            .transpose()?;
769        let seed_root = selected_variant
770            .as_ref()
771            .and_then(|variant| variant.seed)
772            .or(ctx.root_seed);
773
774        self.execute_phase_scope(
775            replay.plan,
776            replay.controllers,
777            ctx,
778            PhaseScope {
779                phase: replay.replay_request.phase,
780                variant_id: replay.bundle.selected_variant_id.clone(),
781                variant: selected_variant,
782                fold_id: None,
783                seed_root,
784            },
785            PhaseScopeResources {
786                data_provider: Some(replay.data_provider),
787                replay_artifact_handles: Some(&replay_artifacts.handles),
788                replay_artifact_inputs: Some(&replay_artifacts.inputs),
789                replay_bundle_id: Some(&replay.bundle.bundle_id),
790                data_envelopes: Some(replay.data_envelopes),
791                prediction_cache_store: replay.prediction_cache_store,
792                prediction_cache_contracts: prediction_cache_contracts.as_ref(),
793                ..Default::default()
794            },
795        )
796    }
797
798    fn execute_phase_scope(
799        &self,
800        plan: &ExecutionPlan,
801        controllers: &RuntimeControllerRegistry,
802        ctx: &mut RunContext,
803        scope: PhaseScope,
804        mut resources: PhaseScopeResources<'_>,
805    ) -> Result<Vec<NodeResult>> {
806        let _phase_span = crate::observability::phase_span(
807            ctx.run_id.as_str(),
808            plan.id.as_str(),
809            scope.phase.as_str(),
810            scope.variant_id.as_ref().map(VariantId::as_str),
811            scope.fold_id.as_ref().map(FoldId::as_str),
812        )
813        .entered();
814        let mut results = Vec::new();
815        let mut output_handles = BTreeMap::<NodeId, BTreeMap<String, HandleRef>>::new();
816        let mut output_data_views =
817            BTreeMap::<NodeId, BTreeMap<String, DataProviderViewSpec>>::new();
818        let mut input_lineage = BTreeMap::<NodeId, LineageId>::new();
819
820        for level in plan.node_parallel_levels_for_phase(scope.phase)? {
821            for node_id in &level {
822                let node_plan = plan
823                    .node_plans
824                    .get(node_id)
825                    .expect("execution plan was validated");
826                // Cross-branch merge reassembly (concat or late-fusion) is a
827                // scheduler/runtime handler, not a controller call: it reads the
828                // upstream branch OOF blocks from the prediction store and emits
829                // one merged per-sample OOF block. Intercept it before the
830                // controller path (and before the `requires_oof` edge collection,
831                // which is a stacking contract the branch inputs do not satisfy).
832                if let Some(reduction) = merge_reduction_mode(plan, node_plan) {
833                    if let Some(mut result) =
834                        reassemble_branch_merge(plan, node_plan, ctx, &scope, reduction)?
835                    {
836                        let task_node_plan = effective_node_plan_for_scope(node_plan, &scope)?;
837                        let task = NodeTask {
838                            inner_fold_set: None,
839                            run_id: ctx.run_id.clone(),
840                            node_plan: task_node_plan.clone(),
841                            phase: scope.phase,
842                            variant_id: scope.variant_id.clone(),
843                            variant: scope.variant.clone(),
844                            fold_id: scope.fold_id.clone(),
845                            branch_path: Vec::new(),
846                            input_handles: BTreeMap::new(),
847                            data_views: BTreeMap::new(),
848                            prediction_inputs: BTreeMap::new(),
849                            artifact_inputs: BTreeMap::new(),
850                            required_loss_attestations: NodeTask::required_loss_attestations_for(
851                                &task_node_plan,
852                                scope.phase,
853                            )?,
854                            fit_influence: FitInfluenceTask::default(),
855                            seed: None,
856                        };
857                        normalize_result_prediction_ports(plan, &task, &mut result)?;
858                        result.validate_for_task(&task)?;
859                        for prediction in &result.predictions {
860                            ctx.prediction_store.append(prediction.clone())?;
861                        }
862                        apply_result_scoring(
863                            &result,
864                            &mut ctx.score_collector,
865                            &mut ctx.regression_target_records,
866                        )?;
867                        ctx.lineage.record(result.lineage.clone())?;
868                        output_handles.insert(node_id.clone(), result.outputs.clone());
869                        input_lineage.insert(node_id.clone(), result.lineage.record_id.clone());
870                        results.push(result);
871                    }
872                    continue;
873                }
874                let controller = controllers.get(&node_plan.controller_id).ok_or_else(|| {
875                    DagMlError::RuntimeValidation(format!(
876                        "runtime controller `{}` is not registered",
877                        node_plan.controller_id
878                    ))
879                })?;
880                let collected_inputs = collect_input_handles(
881                    plan,
882                    node_plan,
883                    &output_handles,
884                    &output_data_views,
885                    &resources,
886                    ctx,
887                    &scope,
888                )?;
889                if collected_inputs.skip_node {
890                    continue;
891                }
892                let mut input_handles = collected_inputs.handles;
893                let mut artifact_inputs = BTreeMap::new();
894                if let Some(node_artifact_handles) = resources
895                    .replay_artifact_handles
896                    .and_then(|handles| handles.get(node_id))
897                {
898                    for (key, handle) in node_artifact_handles {
899                        if input_handles.insert(key.clone(), handle.clone()).is_some() {
900                            return Err(DagMlError::RuntimeValidation(format!(
901                                "node `{node_id}` received duplicate replay artifact input `{key}`"
902                            )));
903                        }
904                    }
905                }
906                if let Some(node_artifact_inputs) = resources
907                    .replay_artifact_inputs
908                    .and_then(|inputs| inputs.get(node_id))
909                {
910                    for (key, spec) in node_artifact_inputs {
911                        if artifact_inputs.insert(key.clone(), spec.clone()).is_some() {
912                            return Err(DagMlError::RuntimeValidation(format!(
913                                "node `{node_id}` received duplicate replay artifact metadata `{key}`"
914                            )));
915                        }
916                    }
917                }
918                let task_node_plan = effective_node_plan_for_scope(node_plan, &scope)?;
919                let inner_fold_set = inner_fold_set_for_scope(
920                    &plan.campaign,
921                    plan.fold_set.as_ref(),
922                    node_plan,
923                    &scope,
924                )?;
925                let fit_influence = fit_influence_task_for_node(
926                    plan,
927                    &task_node_plan,
928                    &collected_inputs.data_views,
929                )?;
930                let task = NodeTask {
931                    inner_fold_set,
932                    run_id: ctx.run_id.clone(),
933                    node_plan: task_node_plan.clone(),
934                    phase: scope.phase,
935                    variant_id: scope.variant_id.clone(),
936                    variant: scope.variant.clone(),
937                    fold_id: scope.fold_id.clone(),
938                    branch_path: Vec::new(),
939                    input_handles,
940                    data_views: collected_inputs.data_views,
941                    prediction_inputs: collected_inputs.prediction_inputs,
942                    artifact_inputs,
943                    required_loss_attestations: NodeTask::required_loss_attestations_for(
944                        &task_node_plan,
945                        scope.phase,
946                    )?,
947                    fit_influence,
948                    seed: derive_task_seed(
949                        scope.seed_root,
950                        scope.variant_id.as_ref(),
951                        scope.fold_id.as_ref(),
952                        &task_node_plan,
953                        scope.phase,
954                    ),
955                };
956                let _node_span = crate::observability::node_span(
957                    task.run_id.as_str(),
958                    plan.id.as_str(),
959                    task.phase.as_str(),
960                    task.node_plan.node_id.as_str(),
961                    task.node_plan.controller_id.as_str(),
962                )
963                .entered();
964                let mut result = if task.node_plan.kind == NodeKind::Tuner {
965                    return Err(DagMlError::RuntimeValidation(format!(
966                        "tuner node `{}` requires execute_hpo_campaign with an explicit RuntimeHpoExecutionContext",
967                        task.node_plan.node_id
968                    )));
969                } else {
970                    match resources.data_provider {
971                        Some(data_provider) => {
972                            controller.invoke_with_data_provider(&task, data_provider)?
973                        }
974                        None => controller.invoke(&task)?,
975                    }
976                };
977                record_fit_influence_diagnostic(&task, &mut result);
978                normalize_result_prediction_ports(plan, &task, &mut result)?;
979                result.validate_for_task(&task)?;
980                apply_result_prediction_aggregation(
981                    plan,
982                    controllers,
983                    &task,
984                    &mut result,
985                    &resources,
986                )?;
987                attach_coordinator_input_lineage(
988                    &mut result,
989                    plan,
990                    &task.node_plan.node_id,
991                    &input_lineage,
992                )?;
993                if let Some(store) = resources.artifact_store.as_deref_mut() {
994                    if scope.phase == Phase::Refit {
995                        store.capture_refit_artifacts(&task, &result)?;
996                    }
997                }
998                for prediction in &result.predictions {
999                    ctx.prediction_store.append(prediction.clone())?;
1000                }
1001                for prediction in &result.aggregated_predictions {
1002                    ctx.aggregated_prediction_store.append(prediction.clone())?;
1003                }
1004                apply_result_scoring(
1005                    &result,
1006                    &mut ctx.score_collector,
1007                    &mut ctx.regression_target_records,
1008                )?;
1009                ctx.lineage.record(result.lineage.clone())?;
1010                let data_views = derive_output_data_views(plan, &task, &result)?;
1011                output_handles.insert(node_id.clone(), result.outputs.clone());
1012                output_data_views.insert(node_id.clone(), data_views);
1013                input_lineage.insert(node_id.clone(), result.lineage.record_id.clone());
1014                results.push(result);
1015            }
1016        }
1017
1018        Ok(results)
1019    }
1020}
1021
1022impl ParallelScheduler {
1023    pub fn execute_phase(
1024        &self,
1025        plan: &ExecutionPlan,
1026        controllers: &RuntimeControllerRegistry,
1027        ctx: &mut RunContext,
1028        phase: Phase,
1029    ) -> Result<Vec<NodeResult>> {
1030        plan.validate()?;
1031        let variant_id = ctx.variant_id.clone();
1032        let seed_root = ctx.root_seed;
1033        self.execute_phase_scope(
1034            plan,
1035            controllers,
1036            ctx,
1037            PhaseScope {
1038                phase,
1039                variant_id,
1040                variant: None,
1041                fold_id: None,
1042                seed_root,
1043            },
1044            PhaseScopeResources::default(),
1045        )
1046    }
1047
1048    pub fn execute_phase_with_data_provider(
1049        &self,
1050        plan: &ExecutionPlan,
1051        controllers: &RuntimeControllerRegistry,
1052        data_provider: &dyn RuntimeDataProvider,
1053        ctx: &mut RunContext,
1054        phase: Phase,
1055    ) -> Result<Vec<NodeResult>> {
1056        plan.validate()?;
1057        let variant_id = ctx.variant_id.clone();
1058        let seed_root = ctx.root_seed;
1059        self.execute_phase_scope(
1060            plan,
1061            controllers,
1062            ctx,
1063            PhaseScope {
1064                phase,
1065                variant_id,
1066                variant: None,
1067                fold_id: None,
1068                seed_root,
1069            },
1070            PhaseScopeResources {
1071                data_provider: Some(data_provider),
1072                ..Default::default()
1073            },
1074        )
1075    }
1076
1077    pub fn execute_campaign_phase(
1078        &self,
1079        plan: &ExecutionPlan,
1080        controllers: &RuntimeControllerRegistry,
1081        ctx: &mut RunContext,
1082        phase: Phase,
1083    ) -> Result<Vec<NodeResult>> {
1084        plan.validate()?;
1085        let mut results = Vec::new();
1086        let fold_ids = if phase == Phase::FitCv {
1087            plan.fold_set
1088                .as_ref()
1089                .map(|fold_set| {
1090                    fold_set
1091                        .folds
1092                        .iter()
1093                        .map(|fold| Some(fold.fold_id.clone()))
1094                        .collect::<Vec<_>>()
1095                })
1096                .unwrap_or_else(|| vec![None])
1097        } else {
1098            vec![None]
1099        };
1100        for variant in &plan.variants {
1101            if ctx
1102                .variant_id
1103                .as_ref()
1104                .is_some_and(|requested| requested != &variant.variant_id)
1105            {
1106                continue;
1107            }
1108            for fold_id in &fold_ids {
1109                let seed_root = variant.seed.or(ctx.root_seed);
1110                results.extend(self.execute_phase_scope(
1111                    plan,
1112                    controllers,
1113                    ctx,
1114                    PhaseScope {
1115                        phase,
1116                        variant_id: Some(variant.variant_id.clone()),
1117                        variant: Some(VariantExecutionSpec::from_plan(variant)),
1118                        fold_id: fold_id.clone(),
1119                        seed_root,
1120                    },
1121                    PhaseScopeResources::default(),
1122                )?);
1123            }
1124        }
1125        Ok(results)
1126    }
1127
1128    pub fn execute_campaign_phase_with_data_provider(
1129        &self,
1130        plan: &ExecutionPlan,
1131        controllers: &RuntimeControllerRegistry,
1132        data_provider: &dyn RuntimeDataProvider,
1133        ctx: &mut RunContext,
1134        phase: Phase,
1135    ) -> Result<Vec<NodeResult>> {
1136        plan.validate()?;
1137        if phase == Phase::FitCv {
1138            ctx.configure_global_oof_aggregation(plan, data_provider)?;
1139        }
1140        let mut results = Vec::new();
1141        let fold_ids = if phase == Phase::FitCv {
1142            plan.fold_set
1143                .as_ref()
1144                .map(|fold_set| {
1145                    fold_set
1146                        .folds
1147                        .iter()
1148                        .map(|fold| Some(fold.fold_id.clone()))
1149                        .collect::<Vec<_>>()
1150                })
1151                .unwrap_or_else(|| vec![None])
1152        } else {
1153            vec![None]
1154        };
1155        for variant in &plan.variants {
1156            if ctx
1157                .variant_id
1158                .as_ref()
1159                .is_some_and(|requested| requested != &variant.variant_id)
1160            {
1161                continue;
1162            }
1163            for fold_id in &fold_ids {
1164                let seed_root = variant.seed.or(ctx.root_seed);
1165                results.extend(self.execute_phase_scope(
1166                    plan,
1167                    controllers,
1168                    ctx,
1169                    PhaseScope {
1170                        phase,
1171                        variant_id: Some(variant.variant_id.clone()),
1172                        variant: Some(VariantExecutionSpec::from_plan(variant)),
1173                        fold_id: fold_id.clone(),
1174                        seed_root,
1175                    },
1176                    PhaseScopeResources {
1177                        data_provider: Some(data_provider),
1178                        ..Default::default()
1179                    },
1180                )?);
1181            }
1182        }
1183        Ok(results)
1184    }
1185
1186    pub fn execute_campaign_phase_with_data_provider_and_artifact_store(
1187        &self,
1188        plan: &ExecutionPlan,
1189        controllers: &RuntimeControllerRegistry,
1190        data_provider: &dyn RuntimeDataProvider,
1191        artifact_store: &mut InMemoryArtifactStore,
1192        ctx: &mut RunContext,
1193        phase: Phase,
1194    ) -> Result<Vec<NodeResult>> {
1195        plan.validate()?;
1196        let mut results = Vec::new();
1197        let fold_ids = if phase == Phase::FitCv {
1198            plan.fold_set
1199                .as_ref()
1200                .map(|fold_set| {
1201                    fold_set
1202                        .folds
1203                        .iter()
1204                        .map(|fold| Some(fold.fold_id.clone()))
1205                        .collect::<Vec<_>>()
1206                })
1207                .unwrap_or_else(|| vec![None])
1208        } else {
1209            vec![None]
1210        };
1211        for variant in &plan.variants {
1212            if ctx
1213                .variant_id
1214                .as_ref()
1215                .is_some_and(|requested| requested != &variant.variant_id)
1216            {
1217                continue;
1218            }
1219            for fold_id in &fold_ids {
1220                let seed_root = variant.seed.or(ctx.root_seed);
1221                results.extend(self.execute_phase_scope(
1222                    plan,
1223                    controllers,
1224                    ctx,
1225                    PhaseScope {
1226                        phase,
1227                        variant_id: Some(variant.variant_id.clone()),
1228                        variant: Some(VariantExecutionSpec::from_plan(variant)),
1229                        fold_id: fold_id.clone(),
1230                        seed_root,
1231                    },
1232                    PhaseScopeResources {
1233                        data_provider: Some(data_provider),
1234                        artifact_store: Some(&mut *artifact_store),
1235                        ..Default::default()
1236                    },
1237                )?);
1238            }
1239        }
1240        Ok(results)
1241    }
1242
1243    pub fn execute_bundle_replay(
1244        &self,
1245        replay: BundleReplayExecution<'_>,
1246        ctx: &mut RunContext,
1247    ) -> Result<Vec<NodeResult>> {
1248        replay.bundle.validate_against_plan(replay.plan)?;
1249        replay
1250            .replay_request
1251            .validate_for_bundle_with_prediction_cache_store(
1252                replay.bundle,
1253                replay.prediction_cache_store.is_some(),
1254            )?;
1255        replay
1256            .bundle
1257            .validate_replay_envelopes(replay.data_envelopes)?;
1258        let prediction_cache_contracts = if replay.replay_request.phase == Phase::Refit {
1259            Some(replay_prediction_cache_contracts(replay.bundle)?)
1260        } else {
1261            None
1262        };
1263        if replay.replay_request.phase == Phase::Refit {
1264            preload_replay_prediction_cache_store(
1265                replay.bundle,
1266                replay.prediction_cache_store,
1267                ctx,
1268            )?;
1269        }
1270        let replay_artifacts = materialize_replay_artifact_handles(
1271            replay.plan,
1272            replay.bundle,
1273            replay.replay_request,
1274            replay.artifact_store,
1275            ctx,
1276        )?;
1277        let selected_variant = replay
1278            .bundle
1279            .selected_variant_id
1280            .as_ref()
1281            .map(|selected| {
1282                replay
1283                    .plan
1284                    .variants
1285                    .iter()
1286                    .find(|variant| &variant.variant_id == selected)
1287                    .map(VariantExecutionSpec::from_plan)
1288                    .ok_or_else(|| {
1289                        DagMlError::RuntimeValidation(format!(
1290                            "bundle `{}` selected unknown variant `{selected}`",
1291                            replay.bundle.bundle_id
1292                        ))
1293                    })
1294            })
1295            .transpose()?;
1296        let seed_root = selected_variant
1297            .as_ref()
1298            .and_then(|variant| variant.seed)
1299            .or(ctx.root_seed);
1300
1301        self.execute_phase_scope(
1302            replay.plan,
1303            replay.controllers,
1304            ctx,
1305            PhaseScope {
1306                phase: replay.replay_request.phase,
1307                variant_id: replay.bundle.selected_variant_id.clone(),
1308                variant: selected_variant,
1309                fold_id: None,
1310                seed_root,
1311            },
1312            PhaseScopeResources {
1313                data_provider: Some(replay.data_provider),
1314                replay_artifact_handles: Some(&replay_artifacts.handles),
1315                replay_artifact_inputs: Some(&replay_artifacts.inputs),
1316                replay_bundle_id: Some(&replay.bundle.bundle_id),
1317                data_envelopes: Some(replay.data_envelopes),
1318                prediction_cache_store: replay.prediction_cache_store,
1319                prediction_cache_contracts: prediction_cache_contracts.as_ref(),
1320                ..Default::default()
1321            },
1322        )
1323    }
1324
1325    fn execute_phase_scope(
1326        &self,
1327        plan: &ExecutionPlan,
1328        controllers: &RuntimeControllerRegistry,
1329        ctx: &mut RunContext,
1330        scope: PhaseScope,
1331        mut resources: PhaseScopeResources<'_>,
1332    ) -> Result<Vec<NodeResult>> {
1333        // Hold the phase span on the scheduler thread, and clone it into each
1334        // worker so worker-thread telemetry nests under the phase (tracing spans
1335        // are thread-local and do not auto-propagate across `thread::scope`).
1336        let phase_span = crate::observability::phase_span(
1337            ctx.run_id.as_str(),
1338            plan.id.as_str(),
1339            scope.phase.as_str(),
1340            scope.variant_id.as_ref().map(VariantId::as_str),
1341            scope.fold_id.as_ref().map(FoldId::as_str),
1342        );
1343        let _phase_entered = phase_span.clone().entered();
1344        // Borrowed for the `thread::scope` below; workers join before it ends.
1345        let plan_id = plan.id.as_str();
1346        plan.validate_parallel_controller_capabilities(self.max_workers, scope.phase)?;
1347        let mut results = Vec::new();
1348        let mut output_handles = BTreeMap::<NodeId, BTreeMap<String, HandleRef>>::new();
1349        let mut output_data_views =
1350            BTreeMap::<NodeId, BTreeMap<String, DataProviderViewSpec>>::new();
1351        let mut input_lineage = BTreeMap::<NodeId, LineageId>::new();
1352
1353        for level in plan.node_parallel_levels_for_phase(scope.phase)? {
1354            let mut prepared = Vec::<PreparedNodeTask>::new();
1355            // Cross-branch merge nodes (concat or late-fusion) are not controller
1356            // tasks: they read the upstream branch OOF blocks from the prediction
1357            // store and reassemble them on the scheduler thread (no worker), AFTER
1358            // this level's worker tasks have populated the store. They are in a
1359            // later level than their branches, so the store already holds the
1360            // branch OOF by the time we reassemble — see `reassemble_branch_merge`.
1361            let mut merge_nodes = Vec::<(NodeId, MergeReduction)>::new();
1362            for node_id in &level {
1363                let node_plan = plan
1364                    .node_plans
1365                    .get(node_id)
1366                    .expect("execution plan was validated");
1367                if let Some(reduction) = merge_reduction_mode(plan, node_plan) {
1368                    merge_nodes.push((node_id.clone(), reduction));
1369                    continue;
1370                }
1371                let collected_inputs = collect_input_handles(
1372                    plan,
1373                    node_plan,
1374                    &output_handles,
1375                    &output_data_views,
1376                    &resources,
1377                    ctx,
1378                    &scope,
1379                )?;
1380                if collected_inputs.skip_node {
1381                    continue;
1382                }
1383                let mut input_handles = collected_inputs.handles;
1384                let mut artifact_inputs = BTreeMap::new();
1385                if let Some(node_artifact_handles) = resources
1386                    .replay_artifact_handles
1387                    .and_then(|handles| handles.get(node_id))
1388                {
1389                    for (key, handle) in node_artifact_handles {
1390                        if input_handles.insert(key.clone(), handle.clone()).is_some() {
1391                            return Err(DagMlError::RuntimeValidation(format!(
1392                                "node `{node_id}` received duplicate replay artifact input `{key}`"
1393                            )));
1394                        }
1395                    }
1396                }
1397                if let Some(node_artifact_inputs) = resources
1398                    .replay_artifact_inputs
1399                    .and_then(|inputs| inputs.get(node_id))
1400                {
1401                    for (key, spec) in node_artifact_inputs {
1402                        if artifact_inputs.insert(key.clone(), spec.clone()).is_some() {
1403                            return Err(DagMlError::RuntimeValidation(format!(
1404                                "node `{node_id}` received duplicate replay artifact metadata `{key}`"
1405                            )));
1406                        }
1407                    }
1408                }
1409                let task_node_plan = effective_node_plan_for_scope(node_plan, &scope)?;
1410                let inner_fold_set = inner_fold_set_for_scope(
1411                    &plan.campaign,
1412                    plan.fold_set.as_ref(),
1413                    node_plan,
1414                    &scope,
1415                )?;
1416                let fit_influence = fit_influence_task_for_node(
1417                    plan,
1418                    &task_node_plan,
1419                    &collected_inputs.data_views,
1420                )?;
1421                prepared.push(PreparedNodeTask {
1422                    node_id: node_id.clone(),
1423                    task: NodeTask {
1424                        inner_fold_set,
1425                        run_id: ctx.run_id.clone(),
1426                        node_plan: task_node_plan.clone(),
1427                        phase: scope.phase,
1428                        variant_id: scope.variant_id.clone(),
1429                        variant: scope.variant.clone(),
1430                        fold_id: scope.fold_id.clone(),
1431                        branch_path: Vec::new(),
1432                        input_handles,
1433                        data_views: collected_inputs.data_views,
1434                        prediction_inputs: collected_inputs.prediction_inputs,
1435                        artifact_inputs,
1436                        required_loss_attestations: NodeTask::required_loss_attestations_for(
1437                            &task_node_plan,
1438                            scope.phase,
1439                        )?,
1440                        fit_influence,
1441                        seed: derive_task_seed(
1442                            scope.seed_root,
1443                            scope.variant_id.as_ref(),
1444                            scope.fold_id.as_ref(),
1445                            &task_node_plan,
1446                            scope.phase,
1447                        ),
1448                    },
1449                });
1450            }
1451
1452            for chunk in prepared.chunks(self.max_workers) {
1453                let chunk_results = std::thread::scope(
1454                    |thread_scope| -> Result<Vec<NodeResult>> {
1455                        let mut handles = Vec::with_capacity(chunk.len());
1456                        for prepared_task in chunk {
1457                            let controller = controllers
1458                                .get(&prepared_task.task.node_plan.controller_id)
1459                                .ok_or_else(|| {
1460                                    DagMlError::RuntimeValidation(format!(
1461                                        "runtime controller `{}` is not registered",
1462                                        prepared_task.task.node_plan.controller_id
1463                                    ))
1464                                })?;
1465                            let worker_span = phase_span.clone();
1466                            handles.push(thread_scope.spawn(move || {
1467                                let _worker_span = worker_span.entered();
1468                                let _node_span = crate::observability::node_span(
1469                                    prepared_task.task.run_id.as_str(),
1470                                    plan_id,
1471                                    prepared_task.task.phase.as_str(),
1472                                    prepared_task.task.node_plan.node_id.as_str(),
1473                                    prepared_task.task.node_plan.controller_id.as_str(),
1474                                )
1475                                .entered();
1476                                let mut result =
1477                                    if prepared_task.task.node_plan.kind == NodeKind::Tuner {
1478                                        return Err(DagMlError::RuntimeValidation(format!(
1479                                            "tuner node `{}` requires execute_hpo_campaign with an explicit RuntimeHpoExecutionContext",
1480                                            prepared_task.task.node_plan.node_id
1481                                        )));
1482                                    } else {
1483                                        // A provider-aware controller may require a
1484                                        // non-Sync host provider.  Parallel native
1485                                        // Methods PLS is deliberately refused by its
1486                                        // HPO preflight; ordinary controllers keep
1487                                        // their opaque-handle invocation here.
1488                                        controller.invoke(&prepared_task.task)?
1489                                    };
1490                                record_fit_influence_diagnostic(&prepared_task.task, &mut result);
1491                                normalize_result_prediction_ports(
1492                                    plan,
1493                                    &prepared_task.task,
1494                                    &mut result,
1495                                )?;
1496                                result.validate_for_task(&prepared_task.task)?;
1497                                Ok(result)
1498                            }));
1499                        }
1500                        handles
1501                            .into_iter()
1502                            .map(|handle| {
1503                                handle.join().map_err(|_| {
1504                                    DagMlError::RuntimeValidation(
1505                                        "parallel scheduler worker panicked".to_string(),
1506                                    )
1507                                })?
1508                            })
1509                            .collect()
1510                    },
1511                )?;
1512
1513                for (prepared_task, mut result) in chunk.iter().zip(chunk_results) {
1514                    apply_result_prediction_aggregation(
1515                        plan,
1516                        controllers,
1517                        &prepared_task.task,
1518                        &mut result,
1519                        &resources,
1520                    )?;
1521                    attach_coordinator_input_lineage(
1522                        &mut result,
1523                        plan,
1524                        &prepared_task.task.node_plan.node_id,
1525                        &input_lineage,
1526                    )?;
1527                    if let Some(store) = resources.artifact_store.as_deref_mut() {
1528                        if scope.phase == Phase::Refit {
1529                            store.capture_refit_artifacts(&prepared_task.task, &result)?;
1530                        }
1531                    }
1532                    for prediction in &result.predictions {
1533                        ctx.prediction_store.append(prediction.clone())?;
1534                    }
1535                    for prediction in &result.aggregated_predictions {
1536                        ctx.aggregated_prediction_store.append(prediction.clone())?;
1537                    }
1538                    apply_result_scoring(
1539                        &result,
1540                        &mut ctx.score_collector,
1541                        &mut ctx.regression_target_records,
1542                    )?;
1543                    ctx.lineage.record(result.lineage.clone())?;
1544                    let data_views = derive_output_data_views(plan, &prepared_task.task, &result)?;
1545                    output_handles.insert(prepared_task.node_id.clone(), result.outputs.clone());
1546                    output_data_views.insert(prepared_task.node_id.clone(), data_views);
1547                    input_lineage.insert(
1548                        prepared_task.node_id.clone(),
1549                        result.lineage.record_id.clone(),
1550                    );
1551                    results.push(result);
1552                }
1553            }
1554
1555            // Reassemble any cross-branch merge nodes in this level now that the
1556            // level's worker tasks have populated the prediction store. Merge nodes
1557            // sit in a later level than the branches they consume, so the upstream
1558            // branch OOF is already present.
1559            for (node_id, reduction) in &merge_nodes {
1560                let node_plan = plan
1561                    .node_plans
1562                    .get(node_id)
1563                    .expect("execution plan was validated");
1564                if let Some(mut result) =
1565                    reassemble_branch_merge(plan, node_plan, ctx, &scope, *reduction)?
1566                {
1567                    let task_node_plan = effective_node_plan_for_scope(node_plan, &scope)?;
1568                    let task = NodeTask {
1569                        inner_fold_set: None,
1570                        run_id: ctx.run_id.clone(),
1571                        node_plan: task_node_plan.clone(),
1572                        phase: scope.phase,
1573                        variant_id: scope.variant_id.clone(),
1574                        variant: scope.variant.clone(),
1575                        fold_id: scope.fold_id.clone(),
1576                        branch_path: Vec::new(),
1577                        input_handles: BTreeMap::new(),
1578                        data_views: BTreeMap::new(),
1579                        prediction_inputs: BTreeMap::new(),
1580                        artifact_inputs: BTreeMap::new(),
1581                        required_loss_attestations: NodeTask::required_loss_attestations_for(
1582                            &task_node_plan,
1583                            scope.phase,
1584                        )?,
1585                        fit_influence: FitInfluenceTask::default(),
1586                        seed: None,
1587                    };
1588                    normalize_result_prediction_ports(plan, &task, &mut result)?;
1589                    result.validate_for_task(&task)?;
1590                    for prediction in &result.predictions {
1591                        ctx.prediction_store.append(prediction.clone())?;
1592                    }
1593                    apply_result_scoring(
1594                        &result,
1595                        &mut ctx.score_collector,
1596                        &mut ctx.regression_target_records,
1597                    )?;
1598                    ctx.lineage.record(result.lineage.clone())?;
1599                    output_handles.insert(node_id.clone(), result.outputs.clone());
1600                    input_lineage.insert(node_id.clone(), result.lineage.record_id.clone());
1601                    results.push(result);
1602                }
1603            }
1604        }
1605
1606        Ok(results)
1607    }
1608}
1609
1610#[derive(Clone, Copy, Debug, Eq, PartialEq)]
1611enum HpoTrialTerminalState {
1612    Completed,
1613    Pruned,
1614    Failed,
1615}
1616
1617fn validate_hpo_checkpoint_result(
1618    checkpoint: &RuntimeHpoCheckpointResult,
1619    hpo: &RuntimeHpoExecutionContext,
1620    trial_variants: &BTreeMap<i64, VariantId>,
1621    terminal_trials: &BTreeMap<i64, HpoTrialTerminalState>,
1622    history_at_start: u32,
1623) -> Result<()> {
1624    checkpoint.artifact.validate().map_err(|error| {
1625        DagMlError::RuntimeValidation(format!(
1626            "runtime HPO checkpoint artifact is invalid: {error}"
1627        ))
1628    })?;
1629    if checkpoint.operation_id != hpo.operation_id
1630        || checkpoint.controller_id != hpo.controller_id
1631        || checkpoint.target_node_id != hpo.target_node_id
1632        || checkpoint.provenance != hpo.provenance
1633    {
1634        return Err(DagMlError::RuntimeValidation(
1635            "runtime HPO checkpoint provenance does not exactly match its execution context"
1636                .to_string(),
1637        ));
1638    }
1639    let proposed_count = u32::try_from(trial_variants.len()).map_err(|_| {
1640        DagMlError::RuntimeValidation(
1641            "runtime HPO scheduler proposal count does not fit u32".to_string(),
1642        )
1643    })?;
1644    if checkpoint.trial_history_len != hpo.trial_budget_total
1645        || checkpoint.trial_history_len < history_at_start
1646        || checkpoint.trial_history_len - history_at_start != proposed_count
1647    {
1648        return Err(DagMlError::RuntimeValidation(
1649            "runtime HPO checkpoint native history is inconsistent with scheduler-observed trials"
1650                .to_string(),
1651        ));
1652    }
1653    if checkpoint.artifact.binding.controller_id != hpo.controller_id.as_str()
1654        || checkpoint.artifact.binding.controller_id != hpo.study.controller_id
1655        || checkpoint.artifact.binding.study_id != hpo.study.study_id
1656        || checkpoint.artifact.methods_abi != hpo.study.methods_abi
1657    {
1658        return Err(DagMlError::RuntimeValidation(
1659            "runtime HPO checkpoint binding/controller/study does not match the active tuner"
1660                .to_string(),
1661        ));
1662    }
1663    let expected_search_space = hpo.study.search_space.fingerprint().map_err(|error| {
1664        DagMlError::RuntimeValidation(format!(
1665            "runtime HPO cannot fingerprint the configured search space: {error}"
1666        ))
1667    })?;
1668    if checkpoint.artifact.binding.search_space_fingerprint != expected_search_space {
1669        return Err(DagMlError::RuntimeValidation(
1670            "runtime HPO checkpoint search-space binding does not match the active study"
1671                .to_string(),
1672        ));
1673    }
1674
1675    let completed_trial_ids = terminal_trials
1676        .iter()
1677        .filter_map(|(trial_id, state)| {
1678            (*state == HpoTrialTerminalState::Completed).then_some(*trial_id)
1679        })
1680        .collect::<BTreeSet<_>>();
1681    let mut proposal_trial_ids = BTreeSet::new();
1682    for proposal in &checkpoint.completed_proposals {
1683        if !proposal_trial_ids.insert(proposal.trial_id) {
1684            return Err(DagMlError::RuntimeValidation(format!(
1685                "runtime HPO checkpoint has duplicate completed proposal for trial `{}`",
1686                proposal.trial_id
1687            )));
1688        }
1689        if trial_variants.get(&proposal.trial_id) != Some(&proposal.variant.variant_id) {
1690            return Err(DagMlError::RuntimeValidation(format!(
1691                "runtime HPO checkpoint proposal for trial `{}` does not exactly match its scheduler proposal",
1692                proposal.trial_id
1693            )));
1694        }
1695    }
1696    if proposal_trial_ids != completed_trial_ids {
1697        return Err(DagMlError::RuntimeValidation(
1698            "runtime HPO checkpoint proposals must cover exactly the completed trials".to_string(),
1699        ));
1700    }
1701
1702    let mut report_trial_ids = BTreeSet::new();
1703    for completed in &checkpoint.completed_reports {
1704        if !report_trial_ids.insert(completed.trial_id) {
1705            return Err(DagMlError::RuntimeValidation(format!(
1706                "runtime HPO checkpoint has duplicate completed report for trial `{}`",
1707                completed.trial_id
1708            )));
1709        }
1710        if trial_variants.get(&completed.trial_id) != Some(&completed.variant_id)
1711            || !proposal_trial_ids.contains(&completed.trial_id)
1712        {
1713            return Err(DagMlError::RuntimeValidation(format!(
1714                "runtime HPO checkpoint report for trial `{}` does not match a completed proposal",
1715                completed.trial_id
1716            )));
1717        }
1718        let report = &completed.report;
1719        if report.producer_node != hpo.selection.producer_node
1720            || report.producer_port.as_deref() != Some(hpo.selection.producer_port.as_str())
1721            || report.partition != PredictionPartition::Validation
1722            || report
1723                .fold_id
1724                .as_ref()
1725                .is_none_or(|fold| fold.as_str() != "avg")
1726            || report.variant_id.as_ref() != Some(&completed.variant_id)
1727            || !report
1728                .metrics
1729                .get(hpo.selection.metric.name())
1730                .is_some_and(|score| score.is_finite())
1731        {
1732            return Err(DagMlError::RuntimeValidation(format!(
1733                "runtime HPO checkpoint report for trial `{}` is not its one finite target OOF average",
1734                completed.trial_id
1735            )));
1736        }
1737    }
1738    if report_trial_ids != completed_trial_ids {
1739        return Err(DagMlError::RuntimeValidation(
1740            "runtime HPO checkpoint reports must cover exactly one OOF average per completed trial"
1741                .to_string(),
1742        ));
1743    }
1744    Ok(())
1745}
1746
1747pub(crate) struct PreparedNodeTask {
1748    pub(crate) node_id: NodeId,
1749    pub(crate) task: NodeTask,
1750}
1751
1752// This module stays adjacent to the scheduler-owned task preparation it
1753// exercises; the remaining helpers below are shared by both schedulers.
1754#[cfg(test)]
1755#[allow(clippy::items_after_test_module)]
1756mod hpo_scheduler_tests {
1757    use std::collections::{BTreeMap, BTreeSet};
1758    use std::sync::{Arc, Mutex};
1759
1760    use sha2::{Digest, Sha256};
1761
1762    use super::*;
1763    use crate::controller::{
1764        ArtifactPolicy, ControllerCapability, ControllerFitScope, ControllerManifest,
1765        ControllerRegistry, RngPolicy,
1766    };
1767    use crate::data::InMemoryDataProvider;
1768    use crate::fold::{FoldAssignment, FoldPartitionMode};
1769    use crate::graph::{GraphInterface, GraphSpec, NodeSpec, PortSchema, PortSpec};
1770    use crate::hpo::{
1771        HpoDirection, HpoMetric, HpoOptimizerConfig, HpoParameter, HpoPruner, HpoSampler,
1772        HpoSearchSpace, HpoStudyBinding, MethodsHpoStudyConfig, N4moptCheckpointArtifact,
1773        N4MOPT_ARTIFACT_KIND, N4MOPT_CHECKPOINT_SCHEMA_VERSION, N4MOPT_FORMAT,
1774    };
1775    use crate::metrics::RegressionTargetBlock;
1776    use crate::oof::PredictionBlock;
1777    use crate::plan::{build_execution_plan, SplitInvocation};
1778
1779    struct HpoTestModel {
1780        id: ControllerId,
1781        trace: Arc<Mutex<Vec<String>>>,
1782    }
1783
1784    impl RuntimeController for HpoTestModel {
1785        fn controller_id(&self) -> &ControllerId {
1786            &self.id
1787        }
1788
1789        fn invoke(&self, task: &NodeTask) -> Result<NodeResult> {
1790            self.trace.lock().unwrap().push("model_cv".to_string());
1791            let sample_id = match task.fold_id.as_ref().map(FoldId::as_str) {
1792                Some("fold:0") => SampleId::new("sample:one").unwrap(),
1793                Some("fold:1") => SampleId::new("sample:two").unwrap(),
1794                other => {
1795                    return Err(DagMlError::RuntimeValidation(format!(
1796                        "HPO test model received unexpected fold {other:?}"
1797                    )));
1798                }
1799            };
1800            Ok(NodeResult {
1801                schema_version: None,
1802                node_id: task.node_plan.node_id.clone(),
1803                outputs: BTreeMap::from([(
1804                    "prediction".to_string(),
1805                    HandleRef {
1806                        handle: 2,
1807                        kind: HandleKind::Prediction,
1808                        owner_controller: self.id.clone(),
1809                    },
1810                )]),
1811                predictions: vec![PredictionBlock {
1812                    prediction_id: Some(format!("prediction:{}", task.fold_id.as_ref().unwrap())),
1813                    producer_node: task.node_plan.node_id.clone(),
1814                    producer_port: None,
1815                    partition: PredictionPartition::Validation,
1816                    fold_id: task.fold_id.clone(),
1817                    sample_ids: vec![sample_id.clone()],
1818                    values: vec![vec![1.0]],
1819                    target_names: vec!["target".to_string()],
1820                }],
1821                observation_predictions: Vec::new(),
1822                aggregated_predictions: Vec::new(),
1823                explanations: Vec::new(),
1824                shape_deltas: Vec::new(),
1825                artifacts: Vec::new(),
1826                artifact_handles: BTreeMap::new(),
1827                fit_influence_diagnostics: Vec::new(),
1828                regression_targets: vec![RegressionTargetBlock {
1829                    level: PredictionLevel::Sample,
1830                    unit_ids: vec![PredictionUnitId::Sample(sample_id)],
1831                    values: vec![vec![1.0]],
1832                    target_names: vec!["target".to_string()],
1833                }],
1834                lineage: LineageRecord {
1835                    record_id: LineageId::new(format!(
1836                        "lineage:hpo-model:{}",
1837                        task.fold_id.as_ref().unwrap()
1838                    ))
1839                    .unwrap(),
1840                    run_id: task.run_id.clone(),
1841                    node_id: task.node_plan.node_id.clone(),
1842                    phase: task.phase,
1843                    controller_id: self.id.clone(),
1844                    controller_version: task.node_plan.controller_version.clone(),
1845                    variant_id: task.variant_id.clone(),
1846                    fold_id: task.fold_id.clone(),
1847                    branch_path: Vec::new(),
1848                    input_lineage: Vec::new(),
1849                    artifact_refs: Vec::new(),
1850                    params_fingerprint: task.node_plan.params_fingerprint.clone(),
1851                    data_model_shape_fingerprint: None,
1852                    aggregation_policy_fingerprint: None,
1853                    seed: task.seed,
1854                    unsafe_flags: BTreeSet::new(),
1855                    metrics: BTreeMap::new(),
1856                    loss_attestations: Vec::new(),
1857                    early_stopping_records: Vec::new(),
1858                },
1859            })
1860        }
1861    }
1862
1863    struct HpoTestTuner {
1864        id: ControllerId,
1865        trace: Arc<Mutex<Vec<String>>>,
1866        history_len: u32,
1867        proposal_count: u32,
1868    }
1869
1870    struct HpoTestSession {
1871        proposals: Vec<RuntimeHpoProposal>,
1872        trace: Arc<Mutex<Vec<String>>>,
1873        checkpoint: N4moptCheckpointArtifact,
1874        history_len: u32,
1875        completed: Option<(i64, f64)>,
1876    }
1877
1878    impl RuntimeController for HpoTestTuner {
1879        fn controller_id(&self) -> &ControllerId {
1880            &self.id
1881        }
1882
1883        fn invoke(&self, task: &NodeTask) -> Result<NodeResult> {
1884            Err(DagMlError::RuntimeValidation(format!(
1885                "HPO test tuner `{}` was dispatched through generic invoke",
1886                task.node_plan.node_id
1887            )))
1888        }
1889
1890        fn create_tuner_session(
1891            &self,
1892            task: &RuntimeHpoCampaignTask,
1893            context: &RuntimeHpoExecutionContext,
1894        ) -> Result<Box<dyn RuntimeTunerSession>> {
1895            assert_eq!(task.operation_id, context.operation_id);
1896            self.trace
1897                .lock()
1898                .unwrap()
1899                .push("session_factory".to_string());
1900            let payload = vec![7_u8];
1901            let proposals = (0..self.proposal_count)
1902                .map(|offset| {
1903                    let trial_id = i64::from(self.history_len + offset + 1);
1904                    let mut variant = context.base_variant.clone();
1905                    if self.history_len != 0 || self.proposal_count != 1 {
1906                        variant.variant_id = VariantId::new(format!("hpo:trial:{trial_id}"))
1907                            .map_err(|error| DagMlError::RuntimeValidation(error.to_string()))?;
1908                        variant.fingerprint = format!("hpo-test-{trial_id}");
1909                    }
1910                    Ok(RuntimeHpoProposal { trial_id, variant })
1911                })
1912                .collect::<Result<Vec<_>>>()?;
1913            Ok(Box::new(HpoTestSession {
1914                proposals: proposals.into_iter().rev().collect(),
1915                trace: Arc::clone(&self.trace),
1916                history_len: self.history_len,
1917                completed: None,
1918                checkpoint: N4moptCheckpointArtifact {
1919                    schema_version: N4MOPT_CHECKPOINT_SCHEMA_VERSION,
1920                    artifact_kind: N4MOPT_ARTIFACT_KIND.to_string(),
1921                    format: N4MOPT_FORMAT.to_string(),
1922                    binding: HpoStudyBinding {
1923                        controller_id: context.study.controller_id.clone(),
1924                        study_id: context.study.study_id.clone(),
1925                        search_space_fingerprint: context
1926                            .study
1927                            .search_space
1928                            .fingerprint()
1929                            .map_err(|error| DagMlError::RuntimeValidation(error.to_string()))?,
1930                        optimizer_fingerprint: "optimizer:test".to_string(),
1931                    },
1932                    methods_abi: context.study.methods_abi.clone(),
1933                    payload_sha256: format!("{:x}", Sha256::digest(&payload)),
1934                    opaque_payload: payload,
1935                },
1936            }))
1937        }
1938    }
1939
1940    impl RuntimeTunerSession for HpoTestSession {
1941        fn trial_history_len(&self) -> Result<u32> {
1942            Ok(self.history_len)
1943        }
1944
1945        fn ask(&mut self) -> Result<Option<RuntimeHpoProposal>> {
1946            self.trace.lock().unwrap().push("ask".to_string());
1947            let proposal = self.proposals.pop();
1948            if proposal.is_some() {
1949                self.history_len += 1;
1950            }
1951            Ok(proposal)
1952        }
1953
1954        fn report_intermediate(
1955            &mut self,
1956            intermediate: RuntimeHpoIntermediate,
1957        ) -> Result<RuntimeHpoIntermediateOutcome> {
1958            assert_eq!(intermediate.step, 0);
1959            assert!(intermediate.score.is_finite());
1960            self.trace.lock().unwrap().push("intermediate".to_string());
1961            Ok(RuntimeHpoIntermediateOutcome::Continue)
1962        }
1963
1964        fn tell(&mut self, trial_id: i64, terminal: RuntimeHpoTerminal) -> Result<()> {
1965            assert!(trial_id > 0);
1966            assert!(
1967                matches!(terminal, RuntimeHpoTerminal::Completed { score } if score.is_finite())
1968            );
1969            self.trace.lock().unwrap().push("tell".to_string());
1970            if let RuntimeHpoTerminal::Completed { score } = terminal {
1971                self.completed = Some((trial_id, score));
1972            }
1973            Ok(())
1974        }
1975
1976        fn checkpoint(&self) -> Result<N4moptCheckpointArtifact> {
1977            self.trace.lock().unwrap().push("checkpoint".to_string());
1978            Ok(self.checkpoint.clone())
1979        }
1980
1981        fn incumbent(
1982            &self,
1983            variants: &BTreeMap<i64, VariantId>,
1984        ) -> Result<Option<RuntimeHpoIncumbent>> {
1985            let Some((trial_id, score)) = self.completed else {
1986                return Ok(None);
1987            };
1988            Ok(Some(RuntimeHpoIncumbent {
1989                trial_id,
1990                score,
1991                metric: "rmse".to_string(),
1992                direction: HpoDirection::Minimize,
1993                variant_id: variants.get(&trial_id).cloned().unwrap(),
1994            }))
1995        }
1996
1997        fn terminal_trial_snapshots(
1998            &self,
1999            variants: &BTreeMap<i64, VariantId>,
2000        ) -> Result<Vec<RuntimeHpoTerminalSnapshot>> {
2001            let (trial_id, score) = self.completed.ok_or_else(|| {
2002                DagMlError::RuntimeValidation("test HPO session has no completed trial".to_string())
2003            })?;
2004            Ok((1..=i64::from(self.history_len))
2005                .map(|id| {
2006                    let completed = id == trial_id;
2007                    RuntimeHpoTerminalSnapshot {
2008                        trial: crate::hpo::HpoTrial {
2009                            id,
2010                            ask_sequence: id,
2011                            terminal_sequence: Some(id),
2012                            parameters: BTreeMap::new(),
2013                            parameter_order: Vec::new(),
2014                            status: if completed {
2015                                crate::hpo::HpoTrialStatus::Completed
2016                            } else {
2017                                crate::hpo::HpoTrialStatus::Failed
2018                            },
2019                            score: completed.then_some(score),
2020                            rung: 0,
2021                            duration: 0.0,
2022                            intermediates: Vec::new(),
2023                            failure: (!completed).then(|| crate::hpo::HpoFailure {
2024                                code: "RESTORED_TEST_FAILURE".to_string(),
2025                                message: "synthetic restored terminal".to_string(),
2026                                retryable: false,
2027                            }),
2028                        },
2029                        variant_id: variants.get(&id).cloned(),
2030                    }
2031                })
2032                .collect())
2033        }
2034    }
2035
2036    fn node(id: &str, kind: NodeKind, outputs: Vec<PortSpec>) -> NodeSpec {
2037        NodeSpec {
2038            id: NodeId::new(id).unwrap(),
2039            kind,
2040            operator: None,
2041            params: BTreeMap::new(),
2042            ports: PortSchema {
2043                inputs: Vec::new(),
2044                outputs,
2045            },
2046            metadata: BTreeMap::new(),
2047            seed_label: None,
2048        }
2049    }
2050
2051    fn manifest(id: &str, kind: NodeKind) -> ControllerManifest {
2052        ControllerManifest {
2053            controller_id: ControllerId::new(id).unwrap(),
2054            controller_version: "test".to_string(),
2055            operator_kind: kind,
2056            priority: 0,
2057            supported_phases: BTreeSet::from([Phase::FitCv]),
2058            input_ports: Vec::new(),
2059            output_ports: Vec::new(),
2060            data_requirements: None,
2061            capabilities: BTreeSet::from([
2062                ControllerCapability::Deterministic,
2063                ControllerCapability::EmitsPredictions,
2064            ]),
2065            operator_selectors: Vec::new(),
2066            fit_scope: ControllerFitScope::FoldTrain,
2067            rng_policy: RngPolicy::UsesCoreSeed,
2068            artifact_policy: ArtifactPolicy::Serializable,
2069        }
2070    }
2071
2072    #[test]
2073    fn hpo_campaign_invokes_registered_session_and_routes_oof_feedback() {
2074        let target = NodeId::new("model:score").unwrap();
2075        let graph = GraphSpec {
2076            id: "graph:hpo.scheduler".to_string(),
2077            interface: GraphInterface::default(),
2078            nodes: vec![node(
2079                "model:score",
2080                NodeKind::Model,
2081                vec![PortSpec {
2082                    name: "prediction".to_string(),
2083                    kind: PortKind::Prediction,
2084                    representation: None,
2085                    cardinality: crate::graph::PortCardinality::One,
2086                    unit_level: None,
2087                    alignment_key: None,
2088                    target_level: None,
2089                    description: String::new(),
2090                }],
2091            )],
2092            edges: Vec::new(),
2093            search_space_fingerprint: None,
2094            metadata: BTreeMap::new(),
2095        };
2096        let fold_set = FoldSet {
2097            id: "folds:hpo".to_string(),
2098            sample_ids: vec![
2099                SampleId::new("sample:one").unwrap(),
2100                SampleId::new("sample:two").unwrap(),
2101            ],
2102            folds: vec![
2103                FoldAssignment {
2104                    fold_id: FoldId::new("fold:0").unwrap(),
2105                    train_sample_ids: vec![SampleId::new("sample:two").unwrap()],
2106                    validation_sample_ids: vec![SampleId::new("sample:one").unwrap()],
2107                    metadata: BTreeMap::new(),
2108                },
2109                FoldAssignment {
2110                    fold_id: FoldId::new("fold:1").unwrap(),
2111                    train_sample_ids: vec![SampleId::new("sample:one").unwrap()],
2112                    validation_sample_ids: vec![SampleId::new("sample:two").unwrap()],
2113                    metadata: BTreeMap::new(),
2114                },
2115            ],
2116            sample_groups: BTreeMap::new(),
2117            partition_mode: FoldPartitionMode::Partition,
2118        };
2119        let mut registry = ControllerRegistry::new();
2120        registry
2121            .register(manifest("controller:model", NodeKind::Model))
2122            .unwrap();
2123        let plan = build_execution_plan(
2124            "plan:hpo.scheduler",
2125            graph,
2126            CampaignSpec {
2127                inner_cv: None,
2128                id: "campaign:hpo.scheduler".to_string(),
2129                root_seed: Some(13),
2130                leakage_policy: Default::default(),
2131                aggregation_policy: Default::default(),
2132                split_invocation: Some(SplitInvocation {
2133                    id: "split:hpo".to_string(),
2134                    controller_id: None,
2135                    leakage_policy: Default::default(),
2136                    params: BTreeMap::new(),
2137                    fold_set: Some(fold_set),
2138                }),
2139                generation: Default::default(),
2140                shape_plans: BTreeMap::new(),
2141                data_bindings: BTreeMap::new(),
2142                branch_view_plans: Vec::new(),
2143                metadata: BTreeMap::new(),
2144            },
2145            &registry,
2146        )
2147        .unwrap();
2148        let trace = Arc::new(Mutex::new(Vec::new()));
2149        let mut controllers = RuntimeControllerRegistry::new();
2150        controllers
2151            .register(Box::new(HpoTestTuner {
2152                id: ControllerId::new("controller:tuner").unwrap(),
2153                trace: Arc::clone(&trace),
2154                history_len: 0,
2155                proposal_count: 1,
2156            }))
2157            .unwrap();
2158        controllers
2159            .register(Box::new(HpoTestModel {
2160                id: ControllerId::new("controller:model").unwrap(),
2161                trace: Arc::clone(&trace),
2162            }))
2163            .unwrap();
2164        let hpo = RuntimeHpoExecutionContext {
2165            operation_id: "hpo:test".to_string(),
2166            controller_id: ControllerId::new("controller:tuner").unwrap(),
2167            target_node_id: target.clone(),
2168            base_variant: plan.variants[0].clone(),
2169            trial_budget_total: 1,
2170            study: MethodsHpoStudyConfig {
2171                controller_id: "controller:tuner".to_string(),
2172                study_id: "study:hpo.scheduler".to_string(),
2173                methods_abi: "test-abi".to_string(),
2174                search_space: HpoSearchSpace {
2175                    parameters: vec![HpoParameter::Int {
2176                        name: "n_components".to_string(),
2177                        low: 1,
2178                        high: 1,
2179                        step: 1,
2180                        log: false,
2181                    }],
2182                },
2183                optimizer: HpoOptimizerConfig {
2184                    sampler: HpoSampler::Random,
2185                    pruner: HpoPruner::None,
2186                    direction: HpoDirection::Minimize,
2187                    metric: HpoMetric::Rmse,
2188                    seed: 13,
2189                    n_startup_trials: 1,
2190                    max_resource: 0,
2191                    reduction_factor: 1,
2192                },
2193            },
2194            parameter_paths: BTreeMap::from([(
2195                "n_components".to_string(),
2196                "n_components".to_string(),
2197            )]),
2198            resume_checkpoint: None,
2199            resume_variants: BTreeMap::new(),
2200            resume_terminal_trials: Vec::new(),
2201            selection: RuntimeHpoSelectionTarget {
2202                producer_node: target,
2203                producer_port: "prediction".to_string(),
2204                metric: RegressionMetricKind::Rmse,
2205                direction: HpoDirection::Minimize,
2206            },
2207            provenance: RuntimeHpoProvenance {
2208                graph_fingerprint: plan.graph_fingerprint.clone(),
2209                campaign_fingerprint: plan.campaign_fingerprint.clone(),
2210                controller_fingerprint: plan.controller_fingerprint.clone(),
2211                data_identities_fingerprint: "identity:test".to_string(),
2212                fold_set_fingerprint: plan
2213                    .fold_set
2214                    .as_ref()
2215                    .map(stable_json_fingerprint)
2216                    .transpose()
2217                    .unwrap(),
2218                training_influence_fingerprint: "influence:test".to_string(),
2219                relation_fingerprint: "relation:test".to_string(),
2220            },
2221        };
2222        let provider = InMemoryDataProvider::new(ControllerId::new("controller:data").unwrap());
2223        let ctx = RunContext::new(RunId::new("run:hpo.scheduler").unwrap(), Some(13));
2224
2225        let result = SequentialScheduler
2226            .execute_hpo_campaign(&plan, &controllers, &provider, &ctx, &hpo)
2227            .unwrap();
2228
2229        assert_eq!(result.operation_id, "hpo:test");
2230        assert_eq!(result.candidates.len(), 1);
2231        assert_eq!(result.checkpoint.completed_proposals.len(), 1);
2232        assert_eq!(result.checkpoint.completed_reports.len(), 1);
2233        assert_eq!(result.candidates[0].lineage.len(), 2);
2234        assert_eq!(result.incumbent.variant_id, plan.variants[0].variant_id);
2235        let mut selected_ctx = RunContext::new(RunId::new("run:hpo.scheduler").unwrap(), Some(13));
2236        selected_ctx.variant_id = Some(plan.variants[0].variant_id.clone());
2237        let selected_results = SequentialScheduler
2238            .execute_campaign_phase_with_data_provider(
2239                &plan,
2240                &controllers,
2241                &provider,
2242                &mut selected_ctx,
2243                Phase::FitCv,
2244            )
2245            .unwrap();
2246        assert_eq!(selected_results.len(), 2);
2247        assert_eq!(selected_ctx.lineage.len(), 2);
2248        assert_eq!(
2249            trace.lock().unwrap().as_slice(),
2250            [
2251                "session_factory",
2252                "ask",
2253                "model_cv",
2254                "model_cv",
2255                "intermediate",
2256                "tell",
2257                "checkpoint",
2258                "model_cv",
2259                "model_cv"
2260            ]
2261        );
2262
2263        // The native study may restore failed/pruned history that has no
2264        // completed proposal evidence. Its local count, not the coordinator's
2265        // persisted completed list, determines the remaining global budget.
2266        let resumed_trace = Arc::new(Mutex::new(Vec::new()));
2267        let mut resumed_controllers = RuntimeControllerRegistry::new();
2268        resumed_controllers
2269            .register(Box::new(HpoTestTuner {
2270                id: ControllerId::new("controller:tuner").unwrap(),
2271                trace: Arc::clone(&resumed_trace),
2272                history_len: 2,
2273                proposal_count: 2,
2274            }))
2275            .unwrap();
2276        resumed_controllers
2277            .register(Box::new(HpoTestModel {
2278                id: ControllerId::new("controller:model").unwrap(),
2279                trace: Arc::clone(&resumed_trace),
2280            }))
2281            .unwrap();
2282        let mut resumed_hpo = hpo.clone();
2283        resumed_hpo.trial_budget_total = 4;
2284        let resumed_ctx = RunContext::new(RunId::new("run:hpo.resumed").unwrap(), Some(13));
2285        let resumed = SequentialScheduler
2286            .execute_hpo_campaign(
2287                &plan,
2288                &resumed_controllers,
2289                &provider,
2290                &resumed_ctx,
2291                &resumed_hpo,
2292            )
2293            .unwrap();
2294        assert_eq!(resumed.candidates.len(), 2);
2295        assert_eq!(resumed.checkpoint.trial_history_len, 4);
2296        assert_eq!(
2297            resumed_trace
2298                .lock()
2299                .unwrap()
2300                .iter()
2301                .filter(|event| event.as_str() == "ask")
2302                .count(),
2303            2
2304        );
2305
2306        let mut over_budget_controllers = RuntimeControllerRegistry::new();
2307        over_budget_controllers
2308            .register(Box::new(HpoTestTuner {
2309                id: ControllerId::new("controller:tuner").unwrap(),
2310                trace: Arc::new(Mutex::new(Vec::new())),
2311                history_len: 5,
2312                proposal_count: 0,
2313            }))
2314            .unwrap();
2315        let error = SequentialScheduler
2316            .execute_hpo_campaign(
2317                &plan,
2318                &over_budget_controllers,
2319                &provider,
2320                &resumed_ctx,
2321                &resumed_hpo,
2322            )
2323            .unwrap_err();
2324        assert!(error.to_string().contains("exceeds total trial budget"));
2325    }
2326}
2327
2328pub(crate) fn attach_coordinator_input_lineage(
2329    result: &mut NodeResult,
2330    plan: &ExecutionPlan,
2331    node_id: &NodeId,
2332    upstream_lineage: &BTreeMap<NodeId, LineageId>,
2333) -> Result<()> {
2334    let inferred = inferred_input_lineage_for_node(plan, node_id, upstream_lineage);
2335    if result.lineage.input_lineage.is_empty() {
2336        result.lineage.input_lineage = inferred;
2337        return Ok(());
2338    }
2339
2340    let declared = result
2341        .lineage
2342        .input_lineage
2343        .iter()
2344        .cloned()
2345        .collect::<BTreeSet<_>>()
2346        .into_iter()
2347        .collect::<Vec<_>>();
2348    if declared != inferred {
2349        return Err(DagMlError::RuntimeValidation(format!(
2350            "lineage for node `{}` declared input lineage {:?}, expected {:?}",
2351            result.node_id, declared, inferred
2352        )));
2353    }
2354    result.lineage.input_lineage = declared;
2355    Ok(())
2356}
2357
2358pub(crate) fn inferred_input_lineage_for_node(
2359    plan: &ExecutionPlan,
2360    node_id: &NodeId,
2361    upstream_lineage: &BTreeMap<NodeId, LineageId>,
2362) -> Vec<LineageId> {
2363    plan.graph_plan
2364        .graph
2365        .edges
2366        .iter()
2367        .filter(|edge| &edge.target.node_id == node_id && edge.contract.propagates_lineage)
2368        .filter_map(|edge| upstream_lineage.get(&edge.source.node_id).cloned())
2369        .collect::<BTreeSet<_>>()
2370        .into_iter()
2371        .collect()
2372}
2373pub(crate) fn collect_input_handles(
2374    plan: &ExecutionPlan,
2375    node_plan: &NodePlan,
2376    output_handles: &BTreeMap<NodeId, BTreeMap<String, HandleRef>>,
2377    output_data_views: &BTreeMap<NodeId, BTreeMap<String, DataProviderViewSpec>>,
2378    resources: &PhaseScopeResources<'_>,
2379    ctx: &RunContext,
2380    scope: &PhaseScope,
2381) -> Result<CollectedInputs> {
2382    let mut inputs = BTreeMap::new();
2383    let mut data_views = BTreeMap::new();
2384    let mut prediction_inputs = BTreeMap::new();
2385    let training_oof_edges = incoming_training_oof_edges(plan, node_plan, scope)?;
2386    // An OOF edge replaces exactly one raw producer port. Do not hide sibling
2387    // outputs from the same producer: a meta-node may legally consume both an
2388    // OOF prediction port and an auxiliary non-OOF port. PREDICT has no
2389    // Validation-OOF input, but its raw prediction port must still be masked so
2390    // only the explicit `:predict` off-fold input reaches the controller.
2391    let masked_oof_source_ports = if scope.phase == Phase::Predict {
2392        incoming_oof_edges(plan, node_plan)?
2393    } else {
2394        training_oof_edges.clone()
2395    }
2396    .into_iter()
2397    .map(|edge| (edge.source.node_id.clone(), edge.source.port_name.clone()))
2398    .collect::<BTreeSet<_>>();
2399    let bound_data_inputs = node_plan
2400        .data_bindings
2401        .iter()
2402        .map(|binding| binding.input_name.clone())
2403        .collect::<BTreeSet<_>>();
2404    // Only forward upstream handles for ports this node DECLARES an edge to.
2405    // A controller must never see a handle outside its declared port contract,
2406    // so a sibling consumer of the same producer cannot expose extra ports here.
2407    let declared_source_ports = plan
2408        .graph_plan
2409        .graph
2410        .edges
2411        .iter()
2412        .filter(|edge| edge.target.node_id == node_plan.node_id)
2413        .map(|edge| (edge.source.node_id.clone(), edge.source.port_name.clone()))
2414        .collect::<BTreeSet<_>>();
2415    for upstream in &node_plan.input_nodes {
2416        if let Some(handles) = output_handles.get(upstream) {
2417            for (port, handle) in handles {
2418                if !declared_source_ports.contains(&(upstream.clone(), port.clone())) {
2419                    continue;
2420                }
2421                if masked_oof_source_ports.contains(&(upstream.clone(), port.clone())) {
2422                    continue;
2423                }
2424                inputs.insert(format!("{upstream}.{port}"), handle.clone());
2425            }
2426        }
2427    }
2428    for edge in plan
2429        .graph_plan
2430        .graph
2431        .edges
2432        .iter()
2433        .filter(|edge| edge.target.node_id == node_plan.node_id)
2434        .filter(|edge| edge.contract.kind == PortKind::Data && !edge.contract.requires_oof)
2435    {
2436        if bound_data_inputs.contains(&edge.target.port_name) {
2437            continue;
2438        }
2439        let Some(handles) = output_handles.get(&edge.source.node_id) else {
2440            continue;
2441        };
2442        let Some(handle) = handles.get(&edge.source.port_name) else {
2443            continue;
2444        };
2445        let key = data_view_key(&edge.target.port_name);
2446        if inputs.insert(key.clone(), handle.clone()).is_some() {
2447            return Err(DagMlError::RuntimeValidation(format!(
2448                "node `{}` received duplicate data edge input `{key}`",
2449                node_plan.node_id
2450            )));
2451        }
2452        if let Some(source_views) = output_data_views.get(&edge.source.node_id) {
2453            if let Some(view) = source_views.get(&edge.source.port_name) {
2454                if data_views.insert(key.clone(), view.clone()).is_some() {
2455                    return Err(DagMlError::RuntimeValidation(format!(
2456                        "node `{}` received duplicate data edge view `{key}`",
2457                        node_plan.node_id
2458                    )));
2459                }
2460            }
2461            let source_validation_key = validation_data_view_key(&edge.source.port_name);
2462            if let Some(view) = source_views.get(&source_validation_key) {
2463                let validation_key = format!("{key}:validation");
2464                if data_views
2465                    .insert(validation_key.clone(), view.clone())
2466                    .is_some()
2467                {
2468                    return Err(DagMlError::RuntimeValidation(format!(
2469                        "node `{}` received duplicate data edge validation view `{validation_key}`",
2470                        node_plan.node_id
2471                    )));
2472                }
2473            }
2474        }
2475    }
2476    for edge in training_oof_edges {
2477        let key = format!("{}.{}", edge.source.node_id, edge.source.port_name);
2478        let Some(input) = collect_oof_prediction_input(plan, edge, ctx, scope, resources)? else {
2479            return Ok(CollectedInputs {
2480                handles: BTreeMap::new(),
2481                data_views: BTreeMap::new(),
2482                prediction_inputs: BTreeMap::new(),
2483                skip_node: true,
2484            });
2485        };
2486        if inputs.insert(key.clone(), input.handle).is_some() {
2487            return Err(DagMlError::RuntimeValidation(format!(
2488                "node `{}` received duplicate OOF prediction input `{key}`",
2489                node_plan.node_id
2490            )));
2491        }
2492        if prediction_inputs.insert(key.clone(), input.spec).is_some() {
2493            return Err(DagMlError::RuntimeValidation(format!(
2494                "node `{}` received duplicate OOF prediction spec `{key}`",
2495                node_plan.node_id
2496            )));
2497        }
2498    }
2499    // REFIT / PREDICT: deliver each base producer's off-fold (test / predict)
2500    // predictions to the stacking meta-node as a SEPARATE prediction input (suffixed
2501    // `:test` / `:predict`) so the host meta-model predicts from them. The FIT_CV
2502    // Validation-OOF input above is the meta-features the meta-model trains on; this
2503    // off-fold input is used ONLY for REFIT/PREDICT scoring/prediction, never FIT_CV
2504    // training — keeping the leakage invariant intact.
2505    if matches!(scope.phase, Phase::Refit | Phase::Predict) {
2506        let off_fold_suffix = scope.phase.as_str().to_ascii_lowercase();
2507        for edge in incoming_oof_edges(plan, node_plan)? {
2508            let Some(input) = collect_off_fold_prediction_input(plan, edge, ctx, scope)? else {
2509                continue;
2510            };
2511            let key = format!(
2512                "{}.{}:{off_fold_suffix}",
2513                edge.source.node_id, edge.source.port_name
2514            );
2515            if inputs.insert(key.clone(), input.handle).is_some() {
2516                return Err(DagMlError::RuntimeValidation(format!(
2517                    "node `{}` received duplicate off-fold prediction input `{key}`",
2518                    node_plan.node_id
2519                )));
2520            }
2521            if prediction_inputs.insert(key.clone(), input.spec).is_some() {
2522                return Err(DagMlError::RuntimeValidation(format!(
2523                    "node `{}` received duplicate off-fold prediction spec `{key}`",
2524                    node_plan.node_id
2525                )));
2526            }
2527        }
2528    }
2529    if !node_plan.data_bindings.is_empty() && resources.data_provider.is_none() {
2530        return Err(DagMlError::RuntimeValidation(format!(
2531            "node `{}` requires {} data binding(s) but no runtime data provider is registered",
2532            node_plan.node_id,
2533            node_plan.data_bindings.len()
2534        )));
2535    }
2536    if let Some(data_provider) = resources.data_provider {
2537        // Samples excluded from training (sample-local) for this node, derived
2538        // from its coordinator relations. Used to filter FIT view specs so the
2539        // spec, the materialized view, and fit-influence row_weights agree.
2540        let excluded_samples = coordinator_relations_for_node(node_plan, resources)?
2541            .map(|relations| relations.excluded_sample_ids())
2542            .unwrap_or_default();
2543        for binding in &node_plan.data_bindings {
2544            let materialized = data_provider.materialize(&DataMaterializationRequest {
2545                run_id: ctx.run_id.clone(),
2546                node_id: node_plan.node_id.clone(),
2547                input_name: binding.input_name.clone(),
2548                phase: scope.phase,
2549                variant_id: scope.variant_id.clone(),
2550                fold_id: scope.fold_id.clone(),
2551                binding: binding.clone(),
2552            })?;
2553            let branch_view_for_node = branch_view_from_node_metadata(plan, &node_plan.node_id)?;
2554            let view = data_view_for_scope(
2555                binding,
2556                plan.fold_set.as_ref(),
2557                scope,
2558                branch_view_for_node.as_ref(),
2559                &excluded_samples,
2560            )?;
2561            let key = data_view_key(&binding.input_name);
2562            let view_handle = make_data_view_handle(
2563                data_provider,
2564                ctx,
2565                node_plan,
2566                scope,
2567                binding,
2568                &materialized,
2569                &view,
2570            )?;
2571            if data_views.insert(key.clone(), view).is_some() {
2572                return Err(DagMlError::RuntimeValidation(format!(
2573                    "node `{}` received duplicate data view `{key}`",
2574                    node_plan.node_id
2575                )));
2576            }
2577            if inputs.insert(key.clone(), view_handle).is_some() {
2578                return Err(DagMlError::RuntimeValidation(format!(
2579                    "node `{}` received duplicate data input `{key}`",
2580                    node_plan.node_id
2581                )));
2582            }
2583
2584            if let Some(validation_view) = validation_data_view_for_scope(
2585                binding,
2586                plan.fold_set.as_ref(),
2587                scope,
2588                branch_view_for_node.as_ref(),
2589                &excluded_samples,
2590            )? {
2591                let validation_key = format!("{key}:validation");
2592                let validation_handle = make_data_view_handle(
2593                    data_provider,
2594                    ctx,
2595                    node_plan,
2596                    scope,
2597                    binding,
2598                    &materialized,
2599                    &validation_view,
2600                )?;
2601                if data_views
2602                    .insert(validation_key.clone(), validation_view)
2603                    .is_some()
2604                {
2605                    return Err(DagMlError::RuntimeValidation(format!(
2606                        "node `{}` received duplicate validation data view `{validation_key}`",
2607                        node_plan.node_id
2608                    )));
2609                }
2610                if inputs
2611                    .insert(validation_key.clone(), validation_handle)
2612                    .is_some()
2613                {
2614                    return Err(DagMlError::RuntimeValidation(format!(
2615                        "node `{}` received duplicate validation data input `{validation_key}`",
2616                        node_plan.node_id
2617                    )));
2618                }
2619            }
2620        }
2621    }
2622    Ok(CollectedInputs {
2623        handles: inputs,
2624        data_views,
2625        prediction_inputs,
2626        skip_node: false,
2627    })
2628}
2629pub(crate) fn preload_replay_prediction_cache_store(
2630    bundle: &ExecutionBundle,
2631    prediction_cache_store: Option<&dyn RuntimePredictionCacheStore>,
2632    ctx: &mut RunContext,
2633) -> Result<()> {
2634    if bundle.prediction_requirements.is_empty() {
2635        return Ok(());
2636    }
2637    let store = prediction_cache_store.ok_or_else(|| {
2638        DagMlError::RuntimeValidation(format!(
2639            "bundle `{}` cannot preload OOF prediction caches without a prediction cache store",
2640            bundle.bundle_id
2641        ))
2642    })?;
2643    if !ctx.prediction_store.blocks().is_empty() {
2644        return Err(DagMlError::RuntimeValidation(format!(
2645            "bundle `{}` cannot preload OOF prediction caches into a non-empty prediction store",
2646            bundle.bundle_id
2647        )));
2648    }
2649    let contracts = replay_prediction_cache_contracts(bundle)?;
2650    for contract in contracts.values() {
2651        if contract.requirement.prediction_level == PredictionLevel::Sample {
2652            let blocks = store.load_blocks(&contract.cache.requirement_key)?;
2653            if blocks.iter().any(|block| {
2654                block.producer_node != contract.requirement.producer_node
2655                    || block.partition != contract.requirement.partition
2656            }) {
2657                return Err(DagMlError::RuntimeValidation(format!(
2658                    "prediction cache store returned blocks outside requirement `{}`",
2659                    contract.cache.requirement_key
2660                )));
2661            }
2662            let mut payload = build_prediction_cache_payload(&contract.requirement, &blocks)?;
2663            payload.cache_namespace_fingerprints =
2664                contract.cache.cache_namespace_fingerprints.clone();
2665            validate_prediction_cache_payload_matches_record(&payload, &contract.cache)?;
2666            for block in &payload.blocks {
2667                ctx.prediction_store.append(block.clone())?;
2668            }
2669        } else {
2670            let blocks = store.load_aggregated_blocks(&contract.cache.requirement_key)?;
2671            if blocks.iter().any(|block| {
2672                block.producer_node != contract.requirement.producer_node
2673                    || block.partition != contract.requirement.partition
2674                    || block.level != contract.requirement.prediction_level
2675            }) {
2676                return Err(DagMlError::RuntimeValidation(format!(
2677                    "prediction cache store returned aggregated blocks outside requirement `{}`",
2678                    contract.cache.requirement_key
2679                )));
2680            }
2681            let mut payload =
2682                build_aggregated_prediction_cache_payload(&contract.requirement, &blocks)?;
2683            payload.cache_namespace_fingerprints =
2684                contract.cache.cache_namespace_fingerprints.clone();
2685            validate_prediction_cache_payload_matches_record(&payload, &contract.cache)?;
2686        }
2687    }
2688    Ok(())
2689}
2690
2691pub(crate) fn replay_prediction_cache_contracts(
2692    bundle: &ExecutionBundle,
2693) -> Result<BTreeMap<String, ReplayPredictionCacheContract>> {
2694    bundle.validate()?;
2695    let requirements = bundle
2696        .prediction_requirements
2697        .iter()
2698        .map(|requirement| (requirement.key(), requirement))
2699        .collect::<BTreeMap<_, _>>();
2700    let mut contracts = BTreeMap::new();
2701    for cache in &bundle.prediction_caches {
2702        let requirement = requirements.get(&cache.requirement_key).ok_or_else(|| {
2703            DagMlError::RuntimeValidation(format!(
2704                "prediction cache `{}` references unknown prediction requirement `{}`",
2705                cache.cache_id, cache.requirement_key
2706            ))
2707        })?;
2708        contracts.insert(
2709            cache.requirement_key.clone(),
2710            ReplayPredictionCacheContract {
2711                requirement: (*requirement).clone(),
2712                cache: cache.clone(),
2713            },
2714        );
2715    }
2716    Ok(contracts)
2717}
2718
2719pub(crate) fn materialize_replay_artifact_handles(
2720    plan: &ExecutionPlan,
2721    bundle: &ExecutionBundle,
2722    replay_request: &ReplayPhaseRequest,
2723    artifact_store: &dyn RuntimeArtifactStore,
2724    ctx: &RunContext,
2725) -> Result<MaterializedReplayArtifacts> {
2726    let mut handles = BTreeMap::<NodeId, BTreeMap<String, HandleRef>>::new();
2727    let mut inputs = BTreeMap::<NodeId, BTreeMap<String, ArtifactInputSpec>>::new();
2728    for artifact in &bundle.refit_artifacts {
2729        artifact.validate()?;
2730        let node_plan = plan.node_plans.get(&artifact.node_id).ok_or_else(|| {
2731            DagMlError::RuntimeValidation(format!(
2732                "bundle `{}` artifact references unknown node `{}`",
2733                bundle.bundle_id, artifact.node_id
2734            ))
2735        })?;
2736        if !node_plan.supported_phases.contains(&replay_request.phase) {
2737            return Err(DagMlError::RuntimeValidation(format!(
2738                "bundle `{}` artifact node `{}` does not support replay phase {:?}",
2739                bundle.bundle_id, artifact.node_id, replay_request.phase
2740            )));
2741        }
2742        let handle = artifact_store.materialize(&ArtifactMaterializationRequest {
2743            run_id: ctx.run_id.clone(),
2744            bundle_id: bundle.bundle_id.clone(),
2745            node_id: artifact.node_id.clone(),
2746            phase: replay_request.phase,
2747            variant_id: bundle.selected_variant_id.clone(),
2748            controller_id: artifact.controller_id.clone(),
2749            artifact: artifact.artifact.clone(),
2750            params_fingerprint: artifact.params_fingerprint.clone(),
2751            training_loss_fingerprint: artifact.training_loss_fingerprint.clone(),
2752        })?;
2753        if !matches!(handle.kind, HandleKind::Model | HandleKind::Artifact) {
2754            return Err(DagMlError::RuntimeValidation(format!(
2755                "artifact `{}` materialized as unsupported handle kind {:?}",
2756                artifact.artifact.id, handle.kind
2757            )));
2758        }
2759        if handle.owner_controller != artifact.controller_id {
2760            return Err(DagMlError::RuntimeValidation(format!(
2761                "artifact `{}` handle owner `{}` does not match controller `{}`",
2762                artifact.artifact.id, handle.owner_controller, artifact.controller_id
2763            )));
2764        }
2765        let key = refit_artifact_input_key(&artifact.artifact.id);
2766        if handles
2767            .entry(artifact.node_id.clone())
2768            .or_default()
2769            .insert(key.clone(), handle)
2770            .is_some()
2771        {
2772            return Err(DagMlError::RuntimeValidation(format!(
2773                "duplicate replay artifact input `{key}` for node `{}`",
2774                artifact.node_id
2775            )));
2776        }
2777        if inputs
2778            .entry(artifact.node_id.clone())
2779            .or_default()
2780            .insert(key.clone(), ArtifactInputSpec::from_refit_record(artifact)?)
2781            .is_some()
2782        {
2783            return Err(DagMlError::RuntimeValidation(format!(
2784                "duplicate replay artifact metadata `{key}` for node `{}`",
2785                artifact.node_id
2786            )));
2787        }
2788    }
2789    Ok(MaterializedReplayArtifacts { handles, inputs })
2790}
2791
2792pub(crate) fn derive_task_seed(
2793    root_seed: Option<u64>,
2794    variant_id: Option<&VariantId>,
2795    fold_id: Option<&FoldId>,
2796    node_plan: &NodePlan,
2797    phase: Phase,
2798) -> Option<u64> {
2799    root_seed.map(|root| {
2800        let mut context = SeedContext::root(root);
2801        if let Some(variant_id) = variant_id {
2802            context = context.child(format!("variant:{variant_id}"));
2803        }
2804        if let Some(fold_id) = fold_id {
2805            context = context.child(format!("fold:{fold_id}"));
2806        }
2807        context
2808            .child(format!("node:{}", node_plan.node_id))
2809            .child(format!("phase:{phase:?}"))
2810            .derive_u64("task")
2811    })
2812}