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