1use std::collections::{BTreeMap, BTreeSet};
9use std::sync::Mutex;
10
11use serde::{Deserialize, Serialize};
12
13use crate::bundle::{ExecutionBundle, MethodsHpoResumeState, ReplayPhaseRequest};
14use crate::campaign::stable_json_fingerprint;
15use crate::canonical::{deserialize_external_contract, parse_typed_json};
16use crate::conformal::{ConformalMultiTargetPolicy, ConformalSmallSamplePolicy};
17use crate::conformal_runtime::{
18 ConformalCalibration, ConformalCalibrationCohort, ConformalCalibrationContext,
19 ConformalCalibrationTruth, ConformalIntervalBlock,
20};
21use crate::data::ExternalDataPlanEnvelope;
22use crate::error::{DagMlError, Result};
23use crate::fold::fold_set_fingerprint;
24use crate::ids::{ArtifactId, BundleId, ControllerId, RunId};
25use crate::phase::Phase;
26use crate::plan::ExecutionPlan;
27use crate::relation::SampleRelationSet;
28use crate::runtime::{
29 ArtifactMaterializationRequest, BundleReplayExecution, ExplanationBlock, HandleRef,
30 LineageRecord, RunContext, RuntimeArtifactStore, RuntimeControllerRegistry,
31 RuntimeDataProvider, SequentialScheduler,
32};
33use crate::training::{LoadedPredictor, PortablePredictorPackage, TrainingOutcomeRef};
34use crate::training_runtime::{
35 BoundTrainingOutput, PortableRefitPackageV3, TrainingOutcome,
36 BOUND_TRAINING_OUTPUT_SCHEMA_VERSION,
37};
38
39pub const TRAINING_REPLAY_REQUEST_SCHEMA_VERSION: u32 = 1;
40pub const TRAINING_REPLAY_OUTCOME_SCHEMA_VERSION: u32 = 3;
45pub const LEGACY_TRAINING_REPLAY_OUTCOME_SCHEMA_VERSION: u32 = 1;
46pub const CONFORMAL_TRAINING_REPLAY_OUTCOME_SCHEMA_VERSION: u32 = 2;
47pub const MIN_READABLE_TRAINING_REPLAY_OUTCOME_SCHEMA_VERSION: u32 = 1;
48
49pub fn methods_hpo_resume_state_from_json(json: &str) -> Result<MethodsHpoResumeState> {
56 parse_typed_json(json).map_err(|error| {
57 DagMlError::RuntimeValidation(format!(
58 "Methods HPO resume state is not strict TCV1 JSON: {error}"
59 ))
60 })?;
61 let state: MethodsHpoResumeState = deserialize_external_contract(
62 json,
63 "Methods HPO resume state",
64 DagMlError::RuntimeValidation,
65 )?;
66 state.validate()?;
67 Ok(state)
68}
69
70pub fn methods_hpo_resume_state_from_package_json(json: &str) -> Result<MethodsHpoResumeState> {
75 let package = PortablePredictorPackage::from_json(json)?;
76 let state = package
77 .execution_bundle
78 .methods_hpo_resume_state
79 .ok_or_else(|| {
80 DagMlError::RuntimeValidation(
81 "portable predictor package has no typed Methods HPO resume state; legacy checkpoint fields are not resumable"
82 .to_string(),
83 )
84 })?;
85 state.validate()?;
89 Ok(state)
90}
91
92pub struct AttachedTrainingReplayInput<'a> {
93 pub source: &'a TrainingOutcome,
94 pub request: &'a TrainingReplayRequest,
95 pub outcome_id: String,
96 pub run_id: RunId,
97 pub controllers: &'a RuntimeControllerRegistry,
98 pub data_provider: &'a dyn RuntimeDataProvider,
99 pub artifact_store: &'a dyn RuntimeArtifactStore,
100 pub data_envelopes: &'a BTreeMap<String, ExternalDataPlanEnvelope>,
101 pub warnings: Vec<String>,
102 pub diagnostics: BTreeMap<String, serde_json::Value>,
103}
104
105pub struct LoadedPredictorReplayInput<'a> {
106 pub predictor: &'a LoadedPredictor<HandleRef>,
107 pub request: &'a TrainingReplayRequest,
108 pub outcome_id: String,
109 pub run_id: RunId,
110 pub controllers: &'a RuntimeControllerRegistry,
111 pub data_provider: &'a dyn RuntimeDataProvider,
112 pub data_envelopes: &'a BTreeMap<String, ExternalDataPlanEnvelope>,
113 pub warnings: Vec<String>,
114 pub diagnostics: BTreeMap<String, serde_json::Value>,
115}
116
117pub struct LoadedPortableRefitReplayInputV3<'a> {
121 pub package: &'a PortableRefitPackageV3,
122 pub request: &'a TrainingReplayRequest,
123 pub outcome_id: String,
124 pub run_id: RunId,
125 pub controllers: &'a RuntimeControllerRegistry,
126 pub data_provider: &'a dyn RuntimeDataProvider,
127 pub data_envelopes: &'a BTreeMap<String, ExternalDataPlanEnvelope>,
128 pub warnings: Vec<String>,
129 pub diagnostics: BTreeMap<String, serde_json::Value>,
130}
131
132#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
135#[serde(deny_unknown_fields)]
136pub struct PortableRefitReplayOutcomeV3 {
137 pub schema_version: u32,
138 pub outcome_id: String,
139 pub run_id: RunId,
140 pub source_package_fingerprint: String,
141 pub source_refit_outcome_fingerprint: String,
142 pub replay_request_id: String,
143 pub replay_request_fingerprint: String,
144 pub input_data_identities: Vec<ReplayDataIdentity>,
145 pub bundle_id: BundleId,
146 pub plan_id: String,
147 pub phase: Phase,
148 pub outputs: Vec<BoundTrainingOutput>,
149 pub explanations: Vec<ExplanationBlock>,
150 pub lineage: Vec<LineageRecord>,
151 pub warnings: Vec<String>,
152 pub diagnostics: BTreeMap<String, serde_json::Value>,
153 pub outcome_fingerprint: String,
154}
155
156impl PortableRefitReplayOutcomeV3 {
157 pub fn compute_fingerprint(&self) -> Result<String> {
158 tcv1_fingerprint_without(
159 self,
160 "outcome_fingerprint",
161 "portable refit replay outcome V3",
162 )
163 }
164
165 pub fn from_json_for_package(
169 json: &str,
170 package: &PortableRefitPackageV3,
171 request: &TrainingReplayRequest,
172 ) -> Result<Self> {
173 let raw_fingerprint = strict_tcv1_fingerprint_without(
174 json,
175 "outcome_fingerprint",
176 "portable refit replay outcome V3",
177 )?;
178 let outcome: Self = serde_json::from_str(json)?;
179 if outcome.outcome_fingerprint != raw_fingerprint {
180 return contract_error(
181 "portable refit replay outcome fingerprint does not match original TCV1 JSON"
182 .to_string(),
183 );
184 }
185 outcome.validate_against(package, request)?;
186 Ok(outcome)
187 }
188
189 pub fn validate_against(
190 &self,
191 package: &PortableRefitPackageV3,
192 request: &TrainingReplayRequest,
193 ) -> Result<()> {
194 if self.schema_version != 3 {
195 return contract_error(
196 "portable refit replay outcome V3 has unsupported schema_version".to_string(),
197 );
198 }
199 package.validate()?;
200 request.validate()?;
201 validate_replay_phase(request.phase)?;
202 validate_identifier("portable refit replay outcome_id", &self.outcome_id)?;
203 validate_sha256(
204 "portable refit replay source package",
205 &self.source_package_fingerprint,
206 )?;
207 validate_sha256(
208 "portable refit replay source outcome",
209 &self.source_refit_outcome_fingerprint,
210 )?;
211 validate_sha256(
212 "portable refit replay request",
213 &self.replay_request_fingerprint,
214 )?;
215 validate_sha256("portable refit replay", &self.outcome_fingerprint)?;
216 if self.source_package_fingerprint != package.package_fingerprint
217 || self.source_refit_outcome_fingerprint != package.outcome.outcome_fingerprint
218 || request.source_outcome_fingerprint != package.outcome.outcome_fingerprint
219 || self.replay_request_id != request.request_id
220 || self.replay_request_fingerprint != request.request_fingerprint
221 || self.bundle_id != package.outcome.execution_bundle.bundle_id
222 || self.plan_id != package.outcome.effective_plan.id
223 || self.phase != request.phase
224 {
225 return contract_error(
226 "portable refit replay outcome does not exactly bind its V3 package and request"
227 .to_string(),
228 );
229 }
230 let identity_keys = self
231 .input_data_identities
232 .iter()
233 .map(|identity| identity.requirement_key.clone())
234 .collect::<Vec<_>>();
235 if identity_keys != request.data_envelope_keys {
236 return contract_error(
237 "portable refit replay identities do not exactly cover replay request envelopes"
238 .to_string(),
239 );
240 }
241 for identity in &self.input_data_identities {
242 identity.validate()?;
243 }
244 validate_sorted_unique_text("portable refit replay warnings", &self.warnings, false)?;
245 validate_diagnostics(&self.diagnostics)?;
246 validate_output_order_and_version(&self.outputs)?;
247 let bindings = package
248 .outcome
249 .output_bindings
250 .iter()
251 .map(|binding| (binding.binding_id.as_str(), binding))
252 .collect::<BTreeMap<_, _>>();
253 let emitted_binding_ids = self
254 .outputs
255 .iter()
256 .map(|output| output.binding.binding_id.clone())
257 .collect::<Vec<_>>();
258 if self.phase == Phase::Predict && emitted_binding_ids != request.output_binding_ids {
259 return contract_error(
260 "portable refit PREDICT outputs do not exactly cover replay request bindings"
261 .to_string(),
262 );
263 }
264 for output in &self.outputs {
265 let source = bindings
266 .get(output.binding.binding_id.as_str())
267 .ok_or_else(|| {
268 DagMlError::RuntimeValidation(
269 "portable refit replay emits an output absent from its V3 package"
270 .to_string(),
271 )
272 })?;
273 if &output.binding != *source {
274 return contract_error(
275 "portable refit replay output binding differs from V3 package".to_string(),
276 );
277 }
278 output.validate(&package.outcome.effective_plan)?;
279 validate_replay_bound_output_blocks(output)?;
280 }
281 for explanation in &self.explanations {
282 explanation.validate()?;
283 }
284 for record in &self.lineage {
285 record.validate()?;
286 }
287 match self.phase {
288 Phase::Predict if self.outputs.is_empty() => {
289 return contract_error("portable refit PREDICT requires outputs".to_string());
290 }
291 Phase::Predict if !self.explanations.is_empty() => {
292 return contract_error(
293 "portable refit PREDICT cannot emit explanations".to_string(),
294 );
295 }
296 Phase::Explain if self.explanations.is_empty() => {
297 return contract_error("portable refit EXPLAIN requires explanations".to_string());
298 }
299 _ => {}
300 }
301 if self.outcome_fingerprint != self.compute_fingerprint()? {
302 return contract_error(
303 "portable refit replay outcome fingerprint does not match TCV1 content".to_string(),
304 );
305 }
306 Ok(())
307 }
308}
309
310#[derive(Clone, Debug, Eq, PartialEq, Serialize, Deserialize)]
311#[serde(deny_unknown_fields)]
312pub struct TrainingReplayRequest {
313 pub schema_version: u32,
314 pub request_id: String,
315 pub source_outcome_fingerprint: String,
316 pub phase: Phase,
317 pub data_envelope_keys: Vec<String>,
318 pub output_binding_ids: Vec<String>,
319 pub request_fingerprint: String,
320}
321
322impl TrainingReplayRequest {
323 pub fn from_json(json: &str) -> Result<Self> {
324 let raw_fingerprint = strict_tcv1_fingerprint_without(
325 json,
326 "request_fingerprint",
327 "training replay request",
328 )?;
329 let request: Self = serde_json::from_str(json)?;
330 if request.request_fingerprint != raw_fingerprint {
331 return contract_error(
332 "training replay request fingerprint does not match original TCV1 JSON",
333 );
334 }
335 request.validate()?;
336 Ok(request)
337 }
338
339 pub fn compute_fingerprint(&self) -> Result<String> {
340 tcv1_fingerprint_without(self, "request_fingerprint", "training replay request")
341 }
342
343 pub fn validate(&self) -> Result<()> {
344 if self.schema_version != TRAINING_REPLAY_REQUEST_SCHEMA_VERSION {
345 return unsupported_version(
346 "training replay request",
347 self.schema_version,
348 TRAINING_REPLAY_REQUEST_SCHEMA_VERSION,
349 );
350 }
351 validate_identifier("training replay request_id", &self.request_id)?;
352 validate_sha256(
353 "training replay source outcome",
354 &self.source_outcome_fingerprint,
355 )?;
356 validate_replay_phase(self.phase)?;
357 validate_sorted_unique_text(
358 "training replay data_envelope_keys",
359 &self.data_envelope_keys,
360 true,
361 )?;
362 validate_sorted_unique_identifiers(
363 "training replay output_binding_ids",
364 &self.output_binding_ids,
365 true,
366 )?;
367 validate_sha256("training replay request", &self.request_fingerprint)?;
368 if self.request_fingerprint != self.compute_fingerprint()? {
369 return contract_error(
370 "training replay request fingerprint does not match TCV1 content",
371 );
372 }
373 Ok(())
374 }
375}
376
377#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
378#[serde(deny_unknown_fields)]
379pub struct TrainingReplayOutcome {
380 pub schema_version: u32,
381 pub outcome_id: String,
382 pub run_id: RunId,
383 pub source_training_outcome: TrainingOutcomeRef,
384 pub replay_request_id: String,
385 pub replay_request_fingerprint: String,
386 pub input_data_identities: Vec<ReplayDataIdentity>,
387 pub bundle_id: BundleId,
388 pub plan_id: String,
389 pub phase: Phase,
390 pub result_count: usize,
391 pub lineage_record_count: usize,
392 pub prediction_block_count: usize,
393 pub observation_prediction_block_count: usize,
394 pub aggregated_prediction_block_count: usize,
395 pub explanation_block_count: usize,
396 pub controller_count: usize,
397 pub prediction_cache_store: bool,
398 pub outputs: Vec<BoundTrainingOutput>,
399 #[serde(default, skip_serializing_if = "Vec::is_empty")]
400 pub conformal_intervals: Vec<ConformalIntervalBlock>,
401 pub explanations: Vec<ExplanationBlock>,
402 pub lineage: Vec<LineageRecord>,
403 pub warnings: Vec<String>,
404 pub diagnostics: BTreeMap<String, serde_json::Value>,
405 pub outcome_fingerprint: String,
406}
407
408#[derive(Clone, Debug, Eq, PartialEq, Serialize, Deserialize)]
417#[serde(deny_unknown_fields)]
418pub struct ReplayDataIdentity {
419 pub requirement_key: String,
420 pub schema_fingerprint: String,
421 pub plan_fingerprint: String,
422 pub relation_fingerprint: String,
423 pub data_content_fingerprint: String,
424 #[serde(default)]
425 pub target_content_fingerprint: Option<String>,
426 pub identity_fingerprint: String,
427}
428
429impl ReplayDataIdentity {
430 pub fn compute_fingerprint(&self) -> Result<String> {
433 tcv1_fingerprint_without(self, "identity_fingerprint", "replay data identity")
434 }
435
436 fn validate(&self) -> Result<()> {
437 validate_non_empty("replay data requirement_key", &self.requirement_key)?;
438 for (label, value) in [
439 ("replay data schema", &self.schema_fingerprint),
440 ("replay data plan", &self.plan_fingerprint),
441 ("replay data relation", &self.relation_fingerprint),
442 ("replay data content", &self.data_content_fingerprint),
443 ("replay data identity", &self.identity_fingerprint),
444 ] {
445 validate_sha256(label, value)?;
446 }
447 if let Some(target_content_fingerprint) = &self.target_content_fingerprint {
448 validate_sha256("replay target content", target_content_fingerprint)?;
449 }
450 if self.identity_fingerprint != self.compute_fingerprint()? {
451 return contract_error("replay data identity fingerprint does not match TCV1 content");
452 }
453 Ok(())
454 }
455}
456
457impl TrainingReplayOutcome {
458 pub fn from_json(json: &str) -> Result<Self> {
459 let raw_fingerprint = strict_tcv1_fingerprint_without(
460 json,
461 "outcome_fingerprint",
462 "training replay outcome",
463 )?;
464 let outcome: Self = serde_json::from_str(json)?;
465 if outcome.outcome_fingerprint != raw_fingerprint {
466 return contract_error(
467 "training replay outcome fingerprint does not match original TCV1 JSON",
468 );
469 }
470 outcome.validate()?;
471 Ok(outcome)
472 }
473
474 pub fn compute_fingerprint(&self) -> Result<String> {
475 tcv1_fingerprint_without(self, "outcome_fingerprint", "training replay outcome")
476 }
477
478 pub fn validate(&self) -> Result<()> {
479 if self.schema_version < MIN_READABLE_TRAINING_REPLAY_OUTCOME_SCHEMA_VERSION
480 || self.schema_version > TRAINING_REPLAY_OUTCOME_SCHEMA_VERSION
481 {
482 return unsupported_version(
483 "training replay outcome",
484 self.schema_version,
485 TRAINING_REPLAY_OUTCOME_SCHEMA_VERSION,
486 );
487 }
488 validate_identifier("training replay outcome_id", &self.outcome_id)?;
489 validate_identifier("training replay request_id", &self.replay_request_id)?;
490 validate_sha256("training replay request", &self.replay_request_fingerprint)?;
491 validate_non_empty("training replay plan_id", &self.plan_id)?;
492 validate_replay_phase(self.phase)?;
493 self.source_training_outcome.validate()?;
494 for identity in &self.input_data_identities {
495 identity.validate()?;
496 }
497 if self.schema_version < TRAINING_REPLAY_OUTCOME_SCHEMA_VERSION
498 && self
499 .input_data_identities
500 .iter()
501 .any(|identity| identity.target_content_fingerprint.is_none())
502 {
503 return contract_error(
504 "training replay outcome V1/V2 requires target-bound input identities; migrate target-free PREDICT/EXPLAIN evidence to V3",
505 );
506 }
507 validate_sorted_unique_keys(
508 "training replay input_data_identities",
509 self.input_data_identities
510 .iter()
511 .map(|identity| identity.requirement_key.as_str()),
512 true,
513 )?;
514 if self.prediction_cache_store {
515 return contract_error("training replay outcome cannot persist a prediction cache");
516 }
517 validate_sorted_unique_text("training replay warnings", &self.warnings, false)?;
518 validate_diagnostics(&self.diagnostics)?;
519 validate_output_order_and_version(&self.outputs)?;
520 if self.schema_version == LEGACY_TRAINING_REPLAY_OUTCOME_SCHEMA_VERSION
521 && (!self.conformal_intervals.is_empty()
522 || self
523 .source_training_outcome
524 .pre_conformal_outcome_fingerprint
525 .is_some())
526 {
527 return contract_error(
528 "training replay outcome V1 cannot carry conformal state; migrate to V2 or V3",
529 );
530 }
531 for output in &self.outputs {
532 validate_replay_bound_output_blocks(output)?;
533 }
534 for interval in &self.conformal_intervals {
535 let output = self
536 .outputs
537 .iter()
538 .find(|output| output.binding.binding_id == interval.binding_id)
539 .ok_or_else(|| {
540 DagMlError::RuntimeValidation(
541 "conformal interval references an absent replay output binding".to_string(),
542 )
543 })?;
544 let point = output
545 .predictions
546 .iter()
547 .find(|block| block.sample_ids == interval.sample_ids)
548 .ok_or_else(|| {
549 DagMlError::RuntimeValidation(
550 "conformal interval has no matching replay point prediction block"
551 .to_string(),
552 )
553 })?;
554 interval.validate()?;
555 if interval.point_prediction_fingerprint
556 != crate::conformal_runtime::point_prediction_fingerprint_for_runtime(point)?
557 {
558 return Err(DagMlError::RuntimeValidation(
559 "conformal interval point prediction fingerprint does not match replay output"
560 .to_string(),
561 ));
562 }
563 }
564 for explanation in &self.explanations {
565 explanation.validate()?;
566 validate_optional_port(
567 "training replay explanation producer_port",
568 &explanation.producer_port,
569 )?;
570 }
571 for record in &self.lineage {
572 record.validate()?;
573 }
574 match self.phase {
575 Phase::Predict if self.outputs.is_empty() => {
576 return contract_error("training replay PREDICT requires at least one output");
577 }
578 Phase::Predict if !self.explanations.is_empty() => {
579 return contract_error("training replay PREDICT cannot emit explanations");
580 }
581 Phase::Explain if self.explanations.is_empty() => {
582 return contract_error("training replay EXPLAIN requires at least one explanation");
583 }
584 _ => {}
585 }
586 self.validate_counters()?;
587 validate_sha256("training replay outcome", &self.outcome_fingerprint)?;
588 if self.outcome_fingerprint != self.compute_fingerprint()? {
589 return contract_error(
590 "training replay outcome fingerprint does not match TCV1 content",
591 );
592 }
593 Ok(())
594 }
595
596 pub fn validate_against(
597 &self,
598 source: &TrainingOutcome,
599 request: &TrainingReplayRequest,
600 ) -> Result<()> {
601 self.validate()?;
602 source.validate()?;
603 request.validate()?;
604 if request.source_outcome_fingerprint != source.outcome_fingerprint {
605 return contract_error("training replay request does not target source outcome");
606 }
607 if !source.replayable_phases.contains(&request.phase) {
608 return contract_error("training replay phase is not replayable by source outcome");
609 }
610 if self.source_training_outcome != source.to_reference()? {
611 return contract_error(
612 "training replay outcome source reference does not match source outcome",
613 );
614 }
615 if self.replay_request_id != request.request_id {
616 return contract_error(
617 "training replay outcome request id does not match ReplayRequest",
618 );
619 }
620 if self.replay_request_fingerprint != request.request_fingerprint {
621 return contract_error(
622 "training replay outcome request fingerprint does not match ReplayRequest",
623 );
624 }
625 if self.phase != request.phase {
626 return contract_error("training replay outcome phase does not match ReplayRequest");
627 }
628 if self.bundle_id != source.execution_bundle.bundle_id {
629 return contract_error("training replay outcome bundle does not match source outcome");
630 }
631 if self.plan_id != source.effective_plan.id {
632 return contract_error("training replay outcome plan does not match source outcome");
633 }
634 let identity_keys = self
635 .input_data_identities
636 .iter()
637 .map(|identity| identity.requirement_key.clone())
638 .collect::<Vec<_>>();
639 if identity_keys != request.data_envelope_keys {
640 return contract_error(
641 "training replay outcome identities do not exactly cover ReplayRequest envelopes",
642 );
643 }
644 let source_bindings = source
645 .outputs
646 .iter()
647 .map(|output| (output.binding.binding_id.as_str(), &output.binding))
648 .collect::<BTreeMap<_, _>>();
649 for binding_id in &request.output_binding_ids {
650 if !source_bindings.contains_key(binding_id.as_str()) {
651 return contract_error(
652 "training replay request references absent source output binding",
653 );
654 }
655 }
656 let emitted_binding_ids = self
657 .outputs
658 .iter()
659 .map(|output| output.binding.binding_id.clone())
660 .collect::<Vec<_>>();
661 if self.phase == Phase::Predict && emitted_binding_ids != request.output_binding_ids {
662 return contract_error(
663 "training replay PREDICT outputs do not exactly cover ReplayRequest bindings",
664 );
665 }
666 if self.phase == Phase::Explain
667 && !emitted_binding_ids
668 .iter()
669 .all(|binding_id| request.output_binding_ids.contains(binding_id))
670 {
671 return contract_error(
672 "training replay EXPLAIN outputs include a binding outside ReplayRequest",
673 );
674 }
675 for output in &self.outputs {
676 let Some(source_binding) = source_bindings.get(output.binding.binding_id.as_str())
677 else {
678 return contract_error(
679 "training replay output binding is absent from source outcome",
680 );
681 };
682 if &output.binding != *source_binding {
683 return contract_error(
684 "training replay output binding does not match source outcome binding",
685 );
686 }
687 output.validate(&source.effective_plan)?;
688 }
689 match &source.conformal_calibration {
690 None if !self.conformal_intervals.is_empty() => {
691 return contract_error(
692 "training replay intervals require source calibration context",
693 )
694 }
695 Some(calibration) => validate_replay_interval_closure(
696 calibration,
697 &self.outputs,
698 &self.conformal_intervals,
699 )?,
700 None => {}
701 }
702 Ok(())
703 }
704
705 pub fn validate_against_package(
706 &self,
707 package: &PortablePredictorPackage,
708 request: &TrainingReplayRequest,
709 ) -> Result<()> {
710 self.validate()?;
711 package.validate()?;
712 request.validate()?;
713 validate_replay_phase(request.phase)?;
714 if request.source_outcome_fingerprint != package.training_outcome.outcome_fingerprint {
715 return contract_error(
716 "training replay request does not target package source outcome",
717 );
718 }
719 if self.source_training_outcome != package.training_outcome {
720 return contract_error(
721 "training replay outcome source reference does not match package source outcome",
722 );
723 }
724 if self.replay_request_id != request.request_id {
725 return contract_error(
726 "training replay outcome request id does not match ReplayRequest",
727 );
728 }
729 if self.replay_request_fingerprint != request.request_fingerprint {
730 return contract_error(
731 "training replay outcome request fingerprint does not match ReplayRequest",
732 );
733 }
734 if self.phase != request.phase {
735 return contract_error("training replay outcome phase does not match ReplayRequest");
736 }
737 if self.bundle_id != package.execution_bundle.bundle_id {
738 return contract_error("training replay outcome bundle does not match package");
739 }
740 if self.plan_id != package.effective_plan.id {
741 return contract_error("training replay outcome plan does not match package");
742 }
743 let identity_keys = self
744 .input_data_identities
745 .iter()
746 .map(|identity| identity.requirement_key.clone())
747 .collect::<Vec<_>>();
748 if identity_keys != request.data_envelope_keys {
749 return contract_error(
750 "training replay outcome identities do not exactly cover ReplayRequest envelopes",
751 );
752 }
753 let package_bindings = package
754 .output_bindings
755 .iter()
756 .map(|binding| (binding.binding_id.as_str(), binding))
757 .collect::<BTreeMap<_, _>>();
758 for binding_id in &request.output_binding_ids {
759 if !package_bindings.contains_key(binding_id.as_str()) {
760 return contract_error(
761 "training replay request references absent package output binding",
762 );
763 }
764 }
765 let emitted_binding_ids = self
766 .outputs
767 .iter()
768 .map(|output| output.binding.binding_id.clone())
769 .collect::<Vec<_>>();
770 if self.phase == Phase::Predict && emitted_binding_ids != request.output_binding_ids {
771 return contract_error(
772 "training replay PREDICT outputs do not exactly cover ReplayRequest bindings",
773 );
774 }
775 if self.phase == Phase::Explain
776 && !emitted_binding_ids
777 .iter()
778 .all(|binding_id| request.output_binding_ids.contains(binding_id))
779 {
780 return contract_error(
781 "training replay EXPLAIN outputs include a binding outside ReplayRequest",
782 );
783 }
784 for output in &self.outputs {
785 let Some(package_binding) = package_bindings.get(output.binding.binding_id.as_str())
786 else {
787 return contract_error("training replay output binding is absent from package");
788 };
789 if &output.binding != *package_binding {
790 return contract_error(
791 "training replay output binding does not match package binding",
792 );
793 }
794 output.validate(&package.effective_plan)?;
795 }
796 match &package.conformal_calibration {
797 None if !self.conformal_intervals.is_empty() => {
798 return contract_error(
799 "training replay intervals require package calibration context",
800 )
801 }
802 Some(calibration) => validate_replay_interval_closure(
803 calibration,
804 &self.outputs,
805 &self.conformal_intervals,
806 )?,
807 None => {}
808 }
809 Ok(())
810 }
811
812 fn validate_counters(&self) -> Result<()> {
813 require_count(
814 "training replay result_count",
815 self.result_count,
816 self.lineage.len(),
817 )?;
818 require_count(
819 "training replay lineage_record_count",
820 self.lineage_record_count,
821 self.lineage.len(),
822 )?;
823 require_count(
824 "training replay prediction_block_count",
825 self.prediction_block_count,
826 self.outputs
827 .iter()
828 .map(|output| output.predictions.len())
829 .sum(),
830 )?;
831 require_count(
832 "training replay observation_prediction_block_count",
833 self.observation_prediction_block_count,
834 self.outputs
835 .iter()
836 .map(|output| output.observation_predictions.len())
837 .sum(),
838 )?;
839 require_count(
840 "training replay aggregated_prediction_block_count",
841 self.aggregated_prediction_block_count,
842 self.outputs
843 .iter()
844 .map(|output| output.aggregated_predictions.len())
845 .sum(),
846 )?;
847 require_count(
848 "training replay explanation_block_count",
849 self.explanation_block_count,
850 self.explanations.len(),
851 )?;
852 let controller_count = self
853 .lineage
854 .iter()
855 .map(|record| record.controller_id.as_str())
856 .collect::<BTreeSet<_>>()
857 .len();
858 require_count(
859 "training replay controller_count",
860 self.controller_count,
861 controller_count,
862 )?;
863 Ok(())
864 }
865}
866
867pub(crate) fn replay_request_from_outcome(replay: &TrainingReplayOutcome) -> TrainingReplayRequest {
871 TrainingReplayRequest {
872 schema_version: TRAINING_REPLAY_REQUEST_SCHEMA_VERSION,
873 request_id: replay.replay_request_id.clone(),
874 source_outcome_fingerprint: replay.source_training_outcome.outcome_fingerprint.clone(),
875 phase: replay.phase,
876 data_envelope_keys: replay
877 .input_data_identities
878 .iter()
879 .map(|identity| identity.requirement_key.clone())
880 .collect(),
881 output_binding_ids: replay
882 .outputs
883 .iter()
884 .map(|output| output.binding.binding_id.clone())
885 .collect(),
886 request_fingerprint: replay.replay_request_fingerprint.clone(),
887 }
888}
889
890struct LoadedPredictorArtifactStore<'a> {
891 predictor: &'a LoadedPredictor<HandleRef>,
892 records: BTreeMap<ArtifactId, crate::bundle::RefitArtifactRecord>,
893}
894
895struct BundlePayloadArtifactStore<'a> {
903 bundle: &'a ExecutionBundle,
904 controllers: &'a RuntimeControllerRegistry,
905 fallback: &'a dyn RuntimeArtifactStore,
906 hydrated_handles: Mutex<Vec<(ControllerId, HandleRef)>>,
907}
908
909impl RuntimeArtifactStore for BundlePayloadArtifactStore<'_> {
910 fn materialize(&self, request: &ArtifactMaterializationRequest) -> Result<HandleRef> {
911 let Some(payload) = self.bundle.raw_artifact_payloads.get(&request.artifact.id) else {
912 return self.fallback.materialize(request);
913 };
914 let controller = self
915 .controllers
916 .get(&request.controller_id)
917 .ok_or_else(|| {
918 DagMlError::RuntimeValidation(format!(
919 "bundle `{}` has no registered controller `{}` to hydrate raw artifact `{}`",
920 self.bundle.bundle_id, request.controller_id, request.artifact.id
921 ))
922 })?;
923 let handle = controller.hydrate_artifact_payload(request, payload)?;
924 self.hydrated_handles
925 .lock()
926 .map_err(|_| {
927 DagMlError::RuntimeValidation(
928 "bundle payload hydrated-handle registry lock poisoned".to_string(),
929 )
930 })?
931 .push((request.controller_id.clone(), handle.clone()));
932 Ok(handle)
933 }
934}
935
936impl BundlePayloadArtifactStore<'_> {
937 fn release_hydrated_handles(&self) -> Result<()> {
938 let handles = {
939 let mut handles = self.hydrated_handles.lock().map_err(|_| {
940 DagMlError::RuntimeValidation(
941 "bundle payload hydrated-handle registry lock poisoned".to_string(),
942 )
943 })?;
944 std::mem::take(&mut *handles)
945 };
946 let mut failures = Vec::new();
947 for (controller_id, handle) in handles.into_iter().rev() {
948 let release = self
949 .controllers
950 .get(&controller_id)
951 .ok_or_else(|| {
952 DagMlError::RuntimeValidation(format!(
953 "hydrated artifact owner controller `{controller_id}` is no longer registered"
954 ))
955 })
956 .and_then(|controller| controller.release_hydrated_artifact_payload(&handle));
957 if let Err(error) = release {
958 failures.push(error.to_string());
959 }
960 }
961 if failures.is_empty() {
962 Ok(())
963 } else {
964 Err(DagMlError::RuntimeValidation(format!(
965 "failed to release replay-hydrated artifact handles: {}",
966 failures.join("; ")
967 )))
968 }
969 }
970}
971
972fn finish_bundle_payload_replay<T>(
973 execution: Result<T>,
974 artifact_store: &BundlePayloadArtifactStore<'_>,
975) -> Result<T> {
976 let cleanup = artifact_store.release_hydrated_handles();
977 match (execution, cleanup) {
978 (Ok(value), Ok(())) => Ok(value),
979 (Err(error), Ok(())) => Err(error),
980 (Ok(_), Err(cleanup_error)) => Err(cleanup_error),
981 (Err(error), Err(cleanup_error)) => Err(DagMlError::RuntimeValidation(format!(
982 "{error}; replay hydration cleanup also failed: {cleanup_error}"
983 ))),
984 }
985}
986
987struct RejectRefitSidecarStore;
991
992impl RuntimeArtifactStore for RejectRefitSidecarStore {
993 fn materialize(&self, request: &ArtifactMaterializationRequest) -> Result<HandleRef> {
994 Err(DagMlError::RuntimeValidation(format!(
995 "portable refit V3 replay refuses non-raw or missing artifact `{}`",
996 request.artifact.id
997 )))
998 }
999}
1000
1001impl<'a> LoadedPredictorArtifactStore<'a> {
1002 fn new(predictor: &'a LoadedPredictor<HandleRef>) -> Result<Self> {
1003 predictor.package().validate()?;
1004 let records = predictor
1005 .package()
1006 .execution_bundle
1007 .refit_artifacts
1008 .iter()
1009 .map(|record| {
1010 record.validate()?;
1011 Ok((record.artifact.id.clone(), record.clone()))
1012 })
1013 .collect::<Result<BTreeMap<_, _>>>()?;
1014 Ok(Self { predictor, records })
1015 }
1016}
1017
1018impl RuntimeArtifactStore for LoadedPredictorArtifactStore<'_> {
1019 fn materialize(&self, request: &ArtifactMaterializationRequest) -> Result<HandleRef> {
1020 let record = self.records.get(&request.artifact.id).ok_or_else(|| {
1021 DagMlError::RuntimeValidation(format!(
1022 "loaded predictor is missing refit artifact `{}` for bundle `{}`",
1023 request.artifact.id, request.bundle_id
1024 ))
1025 })?;
1026 if record.node_id != request.node_id {
1027 return Err(DagMlError::RuntimeValidation(format!(
1028 "artifact `{}` is registered for node `{}` but requested for `{}`",
1029 request.artifact.id, record.node_id, request.node_id
1030 )));
1031 }
1032 if record.controller_id != request.controller_id {
1033 return Err(DagMlError::RuntimeValidation(format!(
1034 "artifact `{}` is registered for controller `{}` but requested for `{}`",
1035 request.artifact.id, record.controller_id, request.controller_id
1036 )));
1037 }
1038 if record.artifact != request.artifact {
1039 return Err(DagMlError::RuntimeValidation(format!(
1040 "artifact `{}` metadata does not match package bundle record",
1041 request.artifact.id
1042 )));
1043 }
1044 if record.params_fingerprint != request.params_fingerprint {
1045 return Err(DagMlError::RuntimeValidation(format!(
1046 "artifact `{}` params fingerprint does not match package bundle record",
1047 request.artifact.id
1048 )));
1049 }
1050 let handle = self
1051 .predictor
1052 .artifact(&request.artifact.id)
1053 .ok_or_else(|| {
1054 DagMlError::RuntimeValidation(format!(
1055 "loaded predictor has no process-local handle for `{}`",
1056 request.artifact.id
1057 ))
1058 })?;
1059 Ok(handle.clone())
1060 }
1061}
1062
1063pub fn execute_attached_training_replay(
1064 input: AttachedTrainingReplayInput<'_>,
1065) -> Result<TrainingReplayOutcome> {
1066 input.source.validate()?;
1067 input.request.validate()?;
1068 validate_sorted_unique_text("training replay execution warnings", &input.warnings, false)?;
1069 validate_diagnostics(&input.diagnostics)?;
1070 if input.request.source_outcome_fingerprint != input.source.outcome_fingerprint {
1071 return contract_error("training replay request does not target source outcome");
1072 }
1073 if !input
1074 .source
1075 .replayable_phases
1076 .contains(&input.request.phase)
1077 {
1078 return contract_error("training replay phase is not replayable by source outcome");
1079 }
1080 for node_plan in input.source.effective_plan.node_plans.values() {
1081 if input.controllers.get(&node_plan.controller_id).is_none() {
1082 return Err(DagMlError::RuntimeValidation(format!(
1083 "attached training replay controller `{}` for node `{}` is not registered",
1084 node_plan.controller_id, node_plan.node_id
1085 )));
1086 }
1087 }
1088
1089 let input_data_identities = replay_input_data_identities(
1090 &input.source.execution_bundle,
1091 input.request,
1092 input.data_envelopes,
1093 )?;
1094 let (replay_plan, replay_bundle) = replay_plan_and_bundle_for_current_cohort(
1095 &input.source.effective_plan,
1096 &input.source.execution_bundle,
1097 input.request,
1098 input.data_envelopes,
1099 )?;
1100 let phase_request = ReplayPhaseRequest {
1101 bundle_id: replay_bundle.bundle_id.clone(),
1102 phase: input.request.phase,
1103 data_envelope_keys: input.request.data_envelope_keys.clone(),
1104 };
1105 let artifact_store = BundlePayloadArtifactStore {
1106 bundle: &replay_bundle,
1107 controllers: input.controllers,
1108 fallback: input.artifact_store,
1109 hydrated_handles: Mutex::new(Vec::new()),
1110 };
1111 let mut ctx = RunContext::new(input.run_id.clone(), None);
1112 let execution = SequentialScheduler.execute_bundle_replay(
1113 BundleReplayExecution {
1114 plan: &replay_plan,
1115 bundle: &replay_bundle,
1116 replay_request: &phase_request,
1117 prediction_cache_store: None,
1118 controllers: input.controllers,
1119 data_provider: input.data_provider,
1120 artifact_store: &artifact_store,
1121 data_envelopes: input.data_envelopes,
1122 },
1123 &mut ctx,
1124 );
1125 let results = finish_bundle_payload_replay(execution, &artifact_store)?;
1126 if results
1127 .iter()
1128 .any(|result| !result.artifacts.is_empty() || !result.artifact_handles.is_empty())
1129 {
1130 return contract_error("attached training replay PREDICT/EXPLAIN cannot emit artifacts");
1131 }
1132
1133 let outputs = bind_attached_replay_outputs(input.source, input.request, &results)?;
1134 let conformal_intervals = input
1135 .source
1136 .conformal_calibration
1137 .as_ref()
1138 .map(|calibration| apply_replay_conformal_intervals(calibration, &outputs))
1139 .transpose()?
1140 .unwrap_or_default();
1141 if let Some(calibration) = input.source.conformal_calibration.as_ref() {
1142 validate_replay_interval_closure(calibration, &outputs, &conformal_intervals)?;
1143 }
1144 let explanations = bind_attached_replay_explanations(input.request, &results)?;
1145 let mut lineage = ctx.lineage.records().cloned().collect::<Vec<_>>();
1146 for record in &mut lineage {
1147 record.input_lineage.sort();
1148 record
1149 .artifact_refs
1150 .sort_by(|left, right| left.id.cmp(&right.id));
1151 }
1152 lineage.sort_by(|left, right| left.record_id.cmp(&right.record_id));
1153
1154 let mut outcome = TrainingReplayOutcome {
1155 schema_version: replay_outcome_schema_version(&input_data_identities),
1156 outcome_id: input.outcome_id,
1157 run_id: input.run_id,
1158 source_training_outcome: input.source.to_reference()?,
1159 replay_request_id: input.request.request_id.clone(),
1160 replay_request_fingerprint: input.request.request_fingerprint.clone(),
1161 input_data_identities,
1162 bundle_id: input.source.execution_bundle.bundle_id.clone(),
1163 plan_id: input.source.effective_plan.id.clone(),
1164 phase: input.request.phase,
1165 result_count: lineage.len(),
1166 lineage_record_count: lineage.len(),
1167 prediction_block_count: outputs.iter().map(|output| output.predictions.len()).sum(),
1168 observation_prediction_block_count: outputs
1169 .iter()
1170 .map(|output| output.observation_predictions.len())
1171 .sum(),
1172 aggregated_prediction_block_count: outputs
1173 .iter()
1174 .map(|output| output.aggregated_predictions.len())
1175 .sum(),
1176 explanation_block_count: explanations.len(),
1177 controller_count: lineage
1178 .iter()
1179 .map(|record| record.controller_id.as_str())
1180 .collect::<BTreeSet<_>>()
1181 .len(),
1182 prediction_cache_store: false,
1183 outputs,
1184 conformal_intervals,
1185 explanations,
1186 lineage,
1187 warnings: input.warnings,
1188 diagnostics: input.diagnostics,
1189 outcome_fingerprint: zero_fingerprint(),
1190 };
1191 outcome.outcome_fingerprint = outcome.compute_fingerprint()?;
1192 outcome.validate_against(input.source, input.request)?;
1193 Ok(outcome)
1194}
1195
1196pub fn execute_loaded_predictor_replay(
1197 input: LoadedPredictorReplayInput<'_>,
1198) -> Result<TrainingReplayOutcome> {
1199 let package = input.predictor.package();
1200 package.validate()?;
1201 input.request.validate()?;
1202 validate_sorted_unique_text("training replay execution warnings", &input.warnings, false)?;
1203 validate_diagnostics(&input.diagnostics)?;
1204 validate_replay_phase(input.request.phase)?;
1205 if input.request.source_outcome_fingerprint != package.training_outcome.outcome_fingerprint {
1206 return contract_error("training replay request does not target package source outcome");
1207 }
1208 for node_plan in package.effective_plan.node_plans.values() {
1209 if input.controllers.get(&node_plan.controller_id).is_none() {
1210 return Err(DagMlError::RuntimeValidation(format!(
1211 "loaded predictor replay controller `{}` for node `{}` is not registered",
1212 node_plan.controller_id, node_plan.node_id
1213 )));
1214 }
1215 }
1216
1217 let input_data_identities = replay_input_data_identities(
1218 &package.execution_bundle,
1219 input.request,
1220 input.data_envelopes,
1221 )?;
1222 let (replay_plan, replay_bundle) = replay_plan_and_bundle_for_current_cohort(
1223 &package.effective_plan,
1224 &package.execution_bundle,
1225 input.request,
1226 input.data_envelopes,
1227 )?;
1228 let phase_request = ReplayPhaseRequest {
1229 bundle_id: replay_bundle.bundle_id.clone(),
1230 phase: input.request.phase,
1231 data_envelope_keys: input.request.data_envelope_keys.clone(),
1232 };
1233 let loaded_artifact_store = LoadedPredictorArtifactStore::new(input.predictor)?;
1234 let artifact_store = BundlePayloadArtifactStore {
1235 bundle: &replay_bundle,
1236 controllers: input.controllers,
1237 fallback: &loaded_artifact_store,
1238 hydrated_handles: Mutex::new(Vec::new()),
1239 };
1240 let mut ctx = RunContext::new(input.run_id.clone(), None);
1241 let execution = SequentialScheduler.execute_bundle_replay(
1242 BundleReplayExecution {
1243 plan: &replay_plan,
1244 bundle: &replay_bundle,
1245 replay_request: &phase_request,
1246 prediction_cache_store: None,
1247 controllers: input.controllers,
1248 data_provider: input.data_provider,
1249 artifact_store: &artifact_store,
1250 data_envelopes: input.data_envelopes,
1251 },
1252 &mut ctx,
1253 );
1254 let results = finish_bundle_payload_replay(execution, &artifact_store)?;
1255 if results
1256 .iter()
1257 .any(|result| !result.artifacts.is_empty() || !result.artifact_handles.is_empty())
1258 {
1259 return contract_error("loaded predictor replay PREDICT/EXPLAIN cannot emit artifacts");
1260 }
1261
1262 let outputs = bind_package_replay_outputs(package, input.request, &results)?;
1263 let conformal_intervals = package
1264 .conformal_calibration
1265 .as_ref()
1266 .map(|calibration| apply_replay_conformal_intervals(calibration, &outputs))
1267 .transpose()?
1268 .unwrap_or_default();
1269 if let Some(calibration) = package.conformal_calibration.as_ref() {
1270 validate_replay_interval_closure(calibration, &outputs, &conformal_intervals)?;
1271 }
1272 let explanations = bind_attached_replay_explanations(input.request, &results)?;
1273 let mut lineage = ctx.lineage.records().cloned().collect::<Vec<_>>();
1274 for record in &mut lineage {
1275 record.input_lineage.sort();
1276 record
1277 .artifact_refs
1278 .sort_by(|left, right| left.id.cmp(&right.id));
1279 }
1280 lineage.sort_by(|left, right| left.record_id.cmp(&right.record_id));
1281
1282 let mut outcome = TrainingReplayOutcome {
1283 schema_version: replay_outcome_schema_version(&input_data_identities),
1284 outcome_id: input.outcome_id,
1285 run_id: input.run_id,
1286 source_training_outcome: package.training_outcome.clone(),
1287 replay_request_id: input.request.request_id.clone(),
1288 replay_request_fingerprint: input.request.request_fingerprint.clone(),
1289 input_data_identities,
1290 bundle_id: package.execution_bundle.bundle_id.clone(),
1291 plan_id: package.effective_plan.id.clone(),
1292 phase: input.request.phase,
1293 result_count: lineage.len(),
1294 lineage_record_count: lineage.len(),
1295 prediction_block_count: outputs.iter().map(|output| output.predictions.len()).sum(),
1296 observation_prediction_block_count: outputs
1297 .iter()
1298 .map(|output| output.observation_predictions.len())
1299 .sum(),
1300 aggregated_prediction_block_count: outputs
1301 .iter()
1302 .map(|output| output.aggregated_predictions.len())
1303 .sum(),
1304 explanation_block_count: explanations.len(),
1305 controller_count: lineage
1306 .iter()
1307 .map(|record| record.controller_id.as_str())
1308 .collect::<BTreeSet<_>>()
1309 .len(),
1310 prediction_cache_store: false,
1311 outputs,
1312 conformal_intervals,
1313 explanations,
1314 lineage,
1315 warnings: input.warnings,
1316 diagnostics: input.diagnostics,
1317 outcome_fingerprint: zero_fingerprint(),
1318 };
1319 outcome.outcome_fingerprint = outcome.compute_fingerprint()?;
1320 outcome.validate_against_package(package, input.request)?;
1321 Ok(outcome)
1322}
1323
1324pub fn execute_loaded_portable_refit_replay_v3(
1330 input: LoadedPortableRefitReplayInputV3<'_>,
1331) -> Result<PortableRefitReplayOutcomeV3> {
1332 input.package.validate()?;
1333 input.request.validate()?;
1334 validate_replay_phase(input.request.phase)?;
1335 validate_sorted_unique_text("portable refit replay warnings", &input.warnings, false)?;
1336 validate_diagnostics(&input.diagnostics)?;
1337 if input.request.source_outcome_fingerprint != input.package.outcome.outcome_fingerprint {
1338 return contract_error(
1339 "portable refit replay request does not target the V3 child outcome".to_string(),
1340 );
1341 }
1342 for node_plan in input.package.outcome.effective_plan.node_plans.values() {
1343 if input.controllers.get(&node_plan.controller_id).is_none() {
1344 return Err(DagMlError::RuntimeValidation(format!(
1345 "portable refit replay controller `{}` for node `{}` is not registered",
1346 node_plan.controller_id, node_plan.node_id
1347 )));
1348 }
1349 }
1350 let runtime_bundle = input.package.outcome.to_runtime_replay_bundle()?;
1351 let input_data_identities =
1352 replay_input_data_identities(&runtime_bundle, input.request, input.data_envelopes)?;
1353 let (replay_plan, replay_bundle) = replay_plan_and_bundle_for_current_cohort(
1354 &input.package.outcome.effective_plan,
1355 &runtime_bundle,
1356 input.request,
1357 input.data_envelopes,
1358 )?;
1359 let phase_request = ReplayPhaseRequest {
1360 bundle_id: replay_bundle.bundle_id.clone(),
1361 phase: input.request.phase,
1362 data_envelope_keys: input.request.data_envelope_keys.clone(),
1363 };
1364 let fallback = RejectRefitSidecarStore;
1365 let artifact_store = BundlePayloadArtifactStore {
1366 bundle: &replay_bundle,
1367 controllers: input.controllers,
1368 fallback: &fallback,
1369 hydrated_handles: Mutex::new(Vec::new()),
1370 };
1371 let mut ctx = RunContext::new(input.run_id.clone(), None);
1372 let execution = SequentialScheduler.execute_bundle_replay(
1373 BundleReplayExecution {
1374 plan: &replay_plan,
1375 bundle: &replay_bundle,
1376 replay_request: &phase_request,
1377 prediction_cache_store: None,
1378 controllers: input.controllers,
1379 data_provider: input.data_provider,
1380 artifact_store: &artifact_store,
1381 data_envelopes: input.data_envelopes,
1382 },
1383 &mut ctx,
1384 );
1385 let results = finish_bundle_payload_replay(execution, &artifact_store)?;
1386 if results
1387 .iter()
1388 .any(|result| !result.artifacts.is_empty() || !result.artifact_handles.is_empty())
1389 {
1390 return contract_error(
1391 "portable refit replay PREDICT/EXPLAIN cannot emit artifacts".to_string(),
1392 );
1393 }
1394 let outputs = bind_refit_package_replay_outputs(input.package, input.request, &results)?;
1395 let explanations = bind_attached_replay_explanations(input.request, &results)?;
1396 let mut lineage = ctx.lineage.records().cloned().collect::<Vec<_>>();
1397 for record in &mut lineage {
1398 record.input_lineage.sort();
1399 record
1400 .artifact_refs
1401 .sort_by(|left, right| left.id.cmp(&right.id));
1402 }
1403 lineage.sort_by(|left, right| left.record_id.cmp(&right.record_id));
1404 let mut outcome = PortableRefitReplayOutcomeV3 {
1405 schema_version: 3,
1406 outcome_id: input.outcome_id,
1407 run_id: input.run_id,
1408 source_package_fingerprint: input.package.package_fingerprint.clone(),
1409 source_refit_outcome_fingerprint: input.package.outcome.outcome_fingerprint.clone(),
1410 replay_request_id: input.request.request_id.clone(),
1411 replay_request_fingerprint: input.request.request_fingerprint.clone(),
1412 input_data_identities,
1413 bundle_id: input.package.outcome.execution_bundle.bundle_id.clone(),
1414 plan_id: input.package.outcome.effective_plan.id.clone(),
1415 phase: input.request.phase,
1416 outputs,
1417 explanations,
1418 lineage,
1419 warnings: input.warnings,
1420 diagnostics: input.diagnostics,
1421 outcome_fingerprint: zero_fingerprint(),
1422 };
1423 outcome.outcome_fingerprint = outcome.compute_fingerprint()?;
1424 outcome.validate_against(input.package, input.request)?;
1425 Ok(outcome)
1426}
1427
1428#[allow(clippy::too_many_arguments)]
1433pub fn calibrate_attached_training_replay(
1434 source: &mut TrainingOutcome,
1435 replay: &TrainingReplayOutcome,
1436 binding_id: &str,
1437 calibration_relations: &SampleRelationSet,
1438 truth: ConformalCalibrationTruth,
1439 context: ConformalCalibrationContext,
1440 coverages: Vec<f64>,
1441 multi_target_policy: ConformalMultiTargetPolicy,
1442 small_sample_policy: ConformalSmallSamplePolicy,
1443) -> Result<ConformalCalibration> {
1444 if replay.phase != Phase::Predict {
1445 return Err(DagMlError::RuntimeValidation(
1446 "conformal calibration requires a PREDICT replay outcome".to_string(),
1447 ));
1448 }
1449 if replay
1450 .input_data_identities
1451 .iter()
1452 .any(|identity| identity.target_content_fingerprint.is_none())
1453 {
1454 return Err(DagMlError::RuntimeValidation(
1455 "conformal calibration requires target-bound replay input identities; run the calibration replay with its authoritative truth cohort"
1456 .to_string(),
1457 ));
1458 }
1459 replay.validate_against(source, &replay_request_from_outcome(replay))?;
1460 let output = replay
1461 .outputs
1462 .iter()
1463 .find(|output| output.binding.binding_id == binding_id)
1464 .ok_or_else(|| {
1465 DagMlError::RuntimeValidation(
1466 "calibration replay has no requested output binding".to_string(),
1467 )
1468 })?;
1469 let [point] = output.predictions.as_slice() else {
1470 return Err(DagMlError::RuntimeValidation(
1471 "calibration replay requires exactly one sample point-prediction block".to_string(),
1472 ));
1473 };
1474 validate_calibration_context(
1475 source,
1476 replay,
1477 output,
1478 calibration_relations,
1479 &truth,
1480 &context,
1481 )?;
1482 let calibration = ConformalCalibration::calibrate_with_truth(
1483 binding_id,
1484 output.binding.target_names.clone(),
1485 point,
1486 &truth,
1487 context,
1488 coverages,
1489 multi_target_policy,
1490 small_sample_policy,
1491 )?;
1492 source.attach_conformal_calibration(calibration.clone(), replay.clone())?;
1493 Ok(calibration)
1494}
1495
1496pub fn derive_attached_conformal_calibration_context(
1504 source: &TrainingOutcome,
1505 replay: &TrainingReplayOutcome,
1506 binding_id: &str,
1507 calibration_relations: &SampleRelationSet,
1508) -> Result<ConformalCalibrationContext> {
1509 if replay.phase != Phase::Predict {
1510 return Err(DagMlError::RuntimeValidation(
1511 "conformal calibration requires a PREDICT replay outcome".to_string(),
1512 ));
1513 }
1514 if replay
1515 .input_data_identities
1516 .iter()
1517 .any(|identity| identity.target_content_fingerprint.is_none())
1518 {
1519 return Err(DagMlError::RuntimeValidation(
1520 "conformal calibration requires target-bound replay input identities; run the calibration replay with its authoritative truth cohort"
1521 .to_string(),
1522 ));
1523 }
1524 replay.validate_against(source, &replay_request_from_outcome(replay))?;
1525 calibration_relations.validate()?;
1526 let output = replay
1527 .outputs
1528 .iter()
1529 .find(|output| output.binding.binding_id == binding_id)
1530 .ok_or_else(|| {
1531 DagMlError::RuntimeValidation(
1532 "calibration replay has no requested output binding".to_string(),
1533 )
1534 })?;
1535 let [point] = output.predictions.as_slice() else {
1536 return Err(DagMlError::RuntimeValidation(
1537 "calibration replay requires exactly one sample point-prediction block".to_string(),
1538 ));
1539 };
1540 let relation_fingerprint = calibration_relations.fingerprint()?;
1541 if replay
1542 .input_data_identities
1543 .iter()
1544 .any(|identity| identity.relation_fingerprint != relation_fingerprint)
1545 {
1546 return Err(DagMlError::RuntimeValidation(
1547 "conformal calibration relation authority does not match replay provenance".to_string(),
1548 ));
1549 }
1550 let origin_sample_ids = calibration_origin_closure(&point.sample_ids, calibration_relations)?;
1551 let mut calibration_cohort = ConformalCalibrationCohort {
1552 role: "calibration".to_string(),
1553 physical_sample_ids: point.sample_ids.clone(),
1554 origin_sample_ids,
1555 target_names: output.binding.target_names.clone(),
1556 manifest_fingerprint: String::new(),
1557 };
1558 calibration_cohort.manifest_fingerprint = calibration_cohort.compute_fingerprint()?;
1559 let fold_set = source.effective_plan.fold_set.as_ref().ok_or_else(|| {
1560 DagMlError::RuntimeValidation("conformal calibration requires a source FoldSet".to_string())
1561 })?;
1562 let mut context = ConformalCalibrationContext {
1563 predictor_binding_fingerprint: output.binding.binding_fingerprint.clone(),
1564 source_training_outcome_fingerprint: source.outcome_fingerprint.clone(),
1565 calibration_replay_outcome_fingerprint: replay.outcome_fingerprint.clone(),
1566 data_identities_fingerprint: source.data_identities_fingerprint()?,
1567 fold_set_fingerprint: fold_set_fingerprint(fold_set)?,
1568 training_influence_fingerprint: source.training_influence.manifest_fingerprint.clone(),
1569 relation_fingerprint,
1570 calibration_cohort,
1571 context_fingerprint: String::new(),
1572 };
1573 context.context_fingerprint = context.compute_fingerprint()?;
1574 Ok(context)
1575}
1576
1577#[allow(clippy::too_many_arguments)]
1580pub fn calibrate_attached_training_replay_with_derived_context(
1581 source: &mut TrainingOutcome,
1582 replay: &TrainingReplayOutcome,
1583 binding_id: &str,
1584 calibration_relations: &SampleRelationSet,
1585 truth: ConformalCalibrationTruth,
1586 coverages: Vec<f64>,
1587 multi_target_policy: ConformalMultiTargetPolicy,
1588 small_sample_policy: ConformalSmallSamplePolicy,
1589) -> Result<ConformalCalibration> {
1590 let context = derive_attached_conformal_calibration_context(
1591 source,
1592 replay,
1593 binding_id,
1594 calibration_relations,
1595 )?;
1596 calibrate_attached_training_replay(
1597 source,
1598 replay,
1599 binding_id,
1600 calibration_relations,
1601 truth,
1602 context,
1603 coverages,
1604 multi_target_policy,
1605 small_sample_policy,
1606 )
1607}
1608
1609fn validate_calibration_context(
1610 source: &TrainingOutcome,
1611 replay: &TrainingReplayOutcome,
1612 output: &BoundTrainingOutput,
1613 calibration_relations: &SampleRelationSet,
1614 truth: &ConformalCalibrationTruth,
1615 context: &ConformalCalibrationContext,
1616) -> Result<()> {
1617 context.validate_for_truth(truth, &output.binding.target_names)?;
1618 calibration_relations.validate()?;
1619 let relation_fingerprint = calibration_relations.fingerprint()?;
1620 if context.relation_fingerprint != relation_fingerprint
1621 || replay
1622 .input_data_identities
1623 .iter()
1624 .any(|identity| identity.relation_fingerprint != relation_fingerprint)
1625 {
1626 return Err(DagMlError::RuntimeValidation(
1627 "conformal calibration relation authority does not match context/replay provenance"
1628 .to_string(),
1629 ));
1630 }
1631 let origin_sample_ids = calibration_origin_closure(
1632 &context.calibration_cohort.physical_sample_ids,
1633 calibration_relations,
1634 )?;
1635 if context.calibration_cohort.origin_sample_ids != origin_sample_ids {
1636 return Err(DagMlError::RuntimeValidation(
1637 "conformal calibration cohort origin closure does not match relation authority"
1638 .to_string(),
1639 ));
1640 }
1641 let fold_set = source.effective_plan.fold_set.as_ref().ok_or_else(|| {
1642 DagMlError::RuntimeValidation("conformal calibration requires a source FoldSet".to_string())
1643 })?;
1644 let expected_fold = fold_set_fingerprint(fold_set)?;
1645 if context.predictor_binding_fingerprint != output.binding.binding_fingerprint
1646 || context.source_training_outcome_fingerprint != source.outcome_fingerprint
1647 || context.calibration_replay_outcome_fingerprint != replay.outcome_fingerprint
1648 || context.data_identities_fingerprint != source.data_identities_fingerprint()?
1649 || context.fold_set_fingerprint != expected_fold
1650 || context.training_influence_fingerprint != source.training_influence.manifest_fingerprint
1651 {
1652 return Err(DagMlError::RuntimeValidation(
1653 "conformal calibration context does not exactly match source/replay provenance"
1654 .to_string(),
1655 ));
1656 }
1657 let cohort = &context.calibration_cohort;
1658 let training: std::collections::BTreeSet<_> = source
1659 .training_influence
1660 .entries
1661 .iter()
1662 .flat_map(|entry| {
1663 entry
1664 .physical_sample_ids
1665 .iter()
1666 .chain(entry.origin_sample_ids.iter())
1667 })
1668 .collect();
1669 if cohort
1670 .physical_sample_ids
1671 .iter()
1672 .chain(cohort.origin_sample_ids.iter())
1673 .any(|id| training.contains(id))
1674 {
1675 return Err(DagMlError::RuntimeValidation(
1676 "conformal calibration cohort overlaps training influence closure".to_string(),
1677 ));
1678 }
1679 if relation_fingerprint == source.training_influence.relation_fingerprint {
1680 return Err(DagMlError::RuntimeValidation(
1681 "conformal calibration relation authority must be distinct from development relations"
1682 .to_string(),
1683 ));
1684 }
1685 Ok(())
1686}
1687
1688fn calibration_origin_closure(
1689 physical_sample_ids: &[crate::ids::SampleId],
1690 relations: &SampleRelationSet,
1691) -> Result<Vec<crate::ids::SampleId>> {
1692 let requested = physical_sample_ids.iter().collect::<BTreeSet<_>>();
1693 let mut by_sample = BTreeMap::new();
1694 for relation in &relations.records {
1695 if !requested.contains(&relation.sample_id) {
1696 continue;
1697 }
1698 match by_sample.get(&relation.sample_id) {
1699 Some(origin) if origin != &relation.origin_sample_id => {
1700 return Err(DagMlError::RuntimeValidation(format!(
1701 "conformal calibration sample `{}` has ambiguous origin relations",
1702 relation.sample_id
1703 )));
1704 }
1705 Some(_) => {}
1706 None => {
1707 by_sample.insert(
1708 relation.sample_id.clone(),
1709 relation.origin_sample_id.clone(),
1710 );
1711 }
1712 }
1713 }
1714 if let Some(missing) = physical_sample_ids
1715 .iter()
1716 .find(|sample_id| !by_sample.contains_key(*sample_id))
1717 {
1718 return Err(DagMlError::RuntimeValidation(format!(
1719 "conformal calibration sample `{missing}` is absent from relation authority"
1720 )));
1721 }
1722 Ok(by_sample
1723 .into_values()
1724 .flatten()
1725 .collect::<BTreeSet<_>>()
1726 .into_iter()
1727 .collect())
1728}
1729
1730fn replay_input_data_identities(
1731 bundle: &ExecutionBundle,
1732 request: &TrainingReplayRequest,
1733 envelopes: &BTreeMap<String, ExternalDataPlanEnvelope>,
1734) -> Result<Vec<ReplayDataIdentity>> {
1735 request
1736 .data_envelope_keys
1737 .iter()
1738 .map(|key| {
1739 let requirement = bundle
1740 .data_requirements
1741 .iter()
1742 .find(|requirement| requirement.key() == *key)
1743 .ok_or_else(|| {
1744 DagMlError::RuntimeValidation(format!(
1745 "training replay request references unknown data envelope key `{key}`"
1746 ))
1747 })?;
1748 let envelope = envelopes.get(key).ok_or_else(|| {
1749 DagMlError::RuntimeValidation(format!(
1750 "training replay is missing external data envelope for `{key}`"
1751 ))
1752 })?;
1753 envelope.validate()?;
1754 if requirement.schema_fingerprint != envelope.schema_fingerprint
1755 || requirement.plan_fingerprint != envelope.plan_fingerprint
1756 {
1757 return Err(DagMlError::RuntimeValidation(format!(
1758 "training replay envelope for `{key}` changes schema or representation plan"
1759 )));
1760 }
1761 let relation_fingerprint = envelope.relation_fingerprint.clone().ok_or_else(|| {
1762 DagMlError::RuntimeValidation(format!(
1763 "training replay envelope for `{key}` requires a relation fingerprint"
1764 ))
1765 })?;
1766 let data_content_fingerprint =
1767 envelope.data_content_fingerprint.clone().ok_or_else(|| {
1768 DagMlError::RuntimeValidation(format!(
1769 "training replay envelope for `{key}` requires a data content fingerprint"
1770 ))
1771 })?;
1772 let mut identity = ReplayDataIdentity {
1773 requirement_key: key.clone(),
1774 schema_fingerprint: envelope.schema_fingerprint.clone(),
1775 plan_fingerprint: envelope.plan_fingerprint.clone(),
1776 relation_fingerprint,
1777 data_content_fingerprint,
1778 target_content_fingerprint: envelope.target_content_fingerprint.clone(),
1779 identity_fingerprint: zero_fingerprint(),
1780 };
1781 identity.identity_fingerprint = identity.compute_fingerprint()?;
1782 identity.validate()?;
1783 Ok(identity)
1784 })
1785 .collect()
1786}
1787
1788fn replay_outcome_schema_version(_identities: &[ReplayDataIdentity]) -> u32 {
1789 TRAINING_REPLAY_OUTCOME_SCHEMA_VERSION
1790}
1791
1792fn replay_plan_and_bundle_for_current_cohort(
1793 plan: &ExecutionPlan,
1794 bundle: &ExecutionBundle,
1795 request: &TrainingReplayRequest,
1796 envelopes: &BTreeMap<String, ExternalDataPlanEnvelope>,
1797) -> Result<(ExecutionPlan, ExecutionBundle)> {
1798 let mut replay_plan = plan.clone();
1799 let mut replay_bundle = bundle.clone();
1800 replay_bundle.methods_hpo_resume_state = None;
1806 for requirement in &mut replay_bundle.data_requirements {
1807 let key = requirement.key();
1808 if request.data_envelope_keys.contains(&key) {
1809 let envelope = envelopes.get(&key).ok_or_else(|| {
1810 DagMlError::RuntimeValidation(format!(
1811 "training replay is missing external data envelope for `{key}`"
1812 ))
1813 })?;
1814 requirement.relation_fingerprint = envelope.relation_fingerprint.clone();
1815 for bindings in replay_plan.campaign.data_bindings.values_mut() {
1816 for binding in bindings {
1817 if crate::data::data_binding_requirement_key(
1818 &binding.node_id,
1819 &binding.input_name,
1820 ) == key
1821 {
1822 binding.relation_fingerprint = envelope.relation_fingerprint.clone();
1823 }
1824 }
1825 }
1826 for node_plan in replay_plan.node_plans.values_mut() {
1827 for binding in &mut node_plan.data_bindings {
1828 if crate::data::data_binding_requirement_key(
1829 &binding.node_id,
1830 &binding.input_name,
1831 ) == key
1832 {
1833 binding.relation_fingerprint = envelope.relation_fingerprint.clone();
1834 }
1835 }
1836 }
1837 }
1838 }
1839 replay_plan.graph_fingerprint = stable_json_fingerprint(&replay_plan.graph_plan.graph)?;
1840 replay_plan.campaign_fingerprint = stable_json_fingerprint(&replay_plan.campaign)?;
1841 replay_plan.controller_fingerprint =
1842 stable_json_fingerprint(&replay_plan.controller_manifests)?;
1843 replay_plan.validate()?;
1844 replay_bundle.graph_fingerprint = replay_plan.graph_fingerprint.clone();
1845 replay_bundle.campaign_fingerprint = replay_plan.campaign_fingerprint.clone();
1846 replay_bundle.controller_fingerprint = replay_plan.controller_fingerprint.clone();
1847 replay_bundle.validate_against_plan(&replay_plan)?;
1848 Ok((replay_plan, replay_bundle))
1849}
1850
1851fn bind_attached_replay_outputs(
1852 source: &TrainingOutcome,
1853 request: &TrainingReplayRequest,
1854 results: &[crate::runtime::NodeResult],
1855) -> Result<Vec<BoundTrainingOutput>> {
1856 let mut outputs = Vec::new();
1857 for binding_id in &request.output_binding_ids {
1858 let source_output = source
1859 .outputs
1860 .iter()
1861 .find(|output| output.binding.binding_id == *binding_id)
1862 .ok_or_else(|| {
1863 DagMlError::RuntimeValidation(format!(
1864 "training replay request references absent binding `{binding_id}`"
1865 ))
1866 })?;
1867 let binding = source_output.binding.clone();
1868 let mut output = BoundTrainingOutput {
1869 schema_version: Some(BOUND_TRAINING_OUTPUT_SCHEMA_VERSION),
1870 binding: binding.clone(),
1871 predictions: Vec::new(),
1872 observation_predictions: Vec::new(),
1873 aggregated_predictions: Vec::new(),
1874 };
1875 for result in results {
1876 output.predictions.extend(
1877 result
1878 .predictions
1879 .iter()
1880 .filter(|block| {
1881 block.producer_node == binding.node_id
1882 && block.producer_port.as_deref() == Some(binding.port_name.as_str())
1883 && block.partition == crate::oof::PredictionPartition::Final
1884 && block.fold_id.is_none()
1885 })
1886 .cloned(),
1887 );
1888 output.observation_predictions.extend(
1889 result
1890 .observation_predictions
1891 .iter()
1892 .filter(|block| {
1893 block.producer_node == binding.node_id
1894 && block.producer_port.as_deref() == Some(binding.port_name.as_str())
1895 && block.partition == crate::oof::PredictionPartition::Final
1896 && block.fold_id.is_none()
1897 })
1898 .cloned(),
1899 );
1900 output.aggregated_predictions.extend(
1901 result
1902 .aggregated_predictions
1903 .iter()
1904 .filter(|block| {
1905 block.producer_node == binding.node_id
1906 && block.producer_port.as_deref() == Some(binding.port_name.as_str())
1907 && block.partition == crate::oof::PredictionPartition::Final
1908 && block.fold_id.is_none()
1909 })
1910 .cloned(),
1911 );
1912 }
1913 if !output.predictions.is_empty()
1914 || !output.observation_predictions.is_empty()
1915 || !output.aggregated_predictions.is_empty()
1916 {
1917 output.validate(&source.effective_plan)?;
1918 outputs.push(output);
1919 }
1920 }
1921 outputs.sort_by(|left, right| left.binding.binding_id.cmp(&right.binding.binding_id));
1922 Ok(outputs)
1923}
1924
1925fn bind_package_replay_outputs(
1926 package: &PortablePredictorPackage,
1927 request: &TrainingReplayRequest,
1928 results: &[crate::runtime::NodeResult],
1929) -> Result<Vec<BoundTrainingOutput>> {
1930 bind_replay_outputs(
1931 &package.effective_plan,
1932 &package.output_bindings,
1933 request,
1934 results,
1935 )
1936}
1937
1938fn bind_refit_package_replay_outputs(
1939 package: &PortableRefitPackageV3,
1940 request: &TrainingReplayRequest,
1941 results: &[crate::runtime::NodeResult],
1942) -> Result<Vec<BoundTrainingOutput>> {
1943 bind_replay_outputs(
1944 &package.outcome.effective_plan,
1945 &package.outcome.output_bindings,
1946 request,
1947 results,
1948 )
1949}
1950
1951fn bind_replay_outputs(
1952 plan: &ExecutionPlan,
1953 bindings: &[crate::training::OutputBinding],
1954 request: &TrainingReplayRequest,
1955 results: &[crate::runtime::NodeResult],
1956) -> Result<Vec<BoundTrainingOutput>> {
1957 let mut outputs = Vec::new();
1958 for binding_id in &request.output_binding_ids {
1959 let binding = bindings
1960 .iter()
1961 .find(|binding| binding.binding_id == *binding_id)
1962 .ok_or_else(|| {
1963 DagMlError::RuntimeValidation(format!(
1964 "training replay request references absent package binding `{binding_id}`"
1965 ))
1966 })?
1967 .clone();
1968 let mut output = BoundTrainingOutput {
1969 schema_version: Some(BOUND_TRAINING_OUTPUT_SCHEMA_VERSION),
1970 binding: binding.clone(),
1971 predictions: Vec::new(),
1972 observation_predictions: Vec::new(),
1973 aggregated_predictions: Vec::new(),
1974 };
1975 for result in results {
1976 output.predictions.extend(
1977 result
1978 .predictions
1979 .iter()
1980 .filter(|block| {
1981 block.producer_node == binding.node_id
1982 && block.producer_port.as_deref() == Some(binding.port_name.as_str())
1983 && block.partition == crate::oof::PredictionPartition::Final
1984 && block.fold_id.is_none()
1985 })
1986 .cloned(),
1987 );
1988 output.observation_predictions.extend(
1989 result
1990 .observation_predictions
1991 .iter()
1992 .filter(|block| {
1993 block.producer_node == binding.node_id
1994 && block.producer_port.as_deref() == Some(binding.port_name.as_str())
1995 && block.partition == crate::oof::PredictionPartition::Final
1996 && block.fold_id.is_none()
1997 })
1998 .cloned(),
1999 );
2000 output.aggregated_predictions.extend(
2001 result
2002 .aggregated_predictions
2003 .iter()
2004 .filter(|block| {
2005 block.producer_node == binding.node_id
2006 && block.producer_port.as_deref() == Some(binding.port_name.as_str())
2007 && block.partition == crate::oof::PredictionPartition::Final
2008 && block.fold_id.is_none()
2009 })
2010 .cloned(),
2011 );
2012 }
2013 if !output.predictions.is_empty()
2014 || !output.observation_predictions.is_empty()
2015 || !output.aggregated_predictions.is_empty()
2016 {
2017 output.validate(plan)?;
2018 outputs.push(output);
2019 }
2020 }
2021 outputs.sort_by(|left, right| left.binding.binding_id.cmp(&right.binding.binding_id));
2022 Ok(outputs)
2023}
2024
2025fn apply_replay_conformal_intervals(
2026 calibration: &ConformalCalibration,
2027 outputs: &[BoundTrainingOutput],
2028) -> Result<Vec<ConformalIntervalBlock>> {
2029 let Some(output) = outputs
2030 .iter()
2031 .find(|output| output.binding.binding_id == calibration.binding_id)
2032 else {
2033 return Ok(Vec::new());
2036 };
2037 if output.binding.target_names != calibration.target_names {
2038 return contract_error("replay conformal binding target order differs from calibration");
2039 }
2040 output
2041 .predictions
2042 .iter()
2043 .map(|block| calibration.apply(block))
2044 .collect()
2045}
2046
2047fn validate_replay_interval_closure(
2048 calibration: &ConformalCalibration,
2049 outputs: &[BoundTrainingOutput],
2050 intervals: &[ConformalIntervalBlock],
2051) -> Result<()> {
2052 let expected = apply_replay_conformal_intervals(calibration, outputs)?;
2053 if intervals != expected {
2054 return Err(DagMlError::RuntimeValidation(
2055 "conformal intervals do not exactly cover replay point predictions".to_string(),
2056 ));
2057 }
2058 for interval in intervals {
2059 let output = outputs
2060 .iter()
2061 .find(|output| output.binding.binding_id == interval.binding_id)
2062 .ok_or_else(|| {
2063 DagMlError::RuntimeValidation(
2064 "conformal interval references an absent replay output binding".to_string(),
2065 )
2066 })?;
2067 let point = output
2068 .predictions
2069 .iter()
2070 .find(|point| point.sample_ids == interval.sample_ids)
2071 .ok_or_else(|| {
2072 DagMlError::RuntimeValidation(
2073 "conformal interval has no matching replay point block".to_string(),
2074 )
2075 })?;
2076 interval.validate_against(calibration, point)?;
2077 }
2078 Ok(())
2079}
2080
2081fn bind_attached_replay_explanations(
2082 request: &TrainingReplayRequest,
2083 results: &[crate::runtime::NodeResult],
2084) -> Result<Vec<ExplanationBlock>> {
2085 if request.phase != Phase::Explain {
2086 return Ok(Vec::new());
2087 }
2088 let mut explanations = results
2089 .iter()
2090 .flat_map(|result| result.explanations.iter().cloned())
2091 .filter(|block| block.producer_port.is_some())
2092 .collect::<Vec<_>>();
2093 explanations.sort_by(|left, right| {
2094 (
2095 left.producer_node.as_str(),
2096 left.producer_port.as_deref().unwrap_or_default(),
2097 left.method.as_str(),
2098 left.target_name.as_deref().unwrap_or_default(),
2099 )
2100 .cmp(&(
2101 right.producer_node.as_str(),
2102 right.producer_port.as_deref().unwrap_or_default(),
2103 right.method.as_str(),
2104 right.target_name.as_deref().unwrap_or_default(),
2105 ))
2106 });
2107 Ok(explanations)
2108}
2109
2110fn validate_output_order_and_version(outputs: &[BoundTrainingOutput]) -> Result<()> {
2111 let mut previous: Option<&str> = None;
2112 for output in outputs {
2113 match output.schema_version {
2114 Some(BOUND_TRAINING_OUTPUT_SCHEMA_VERSION) => {}
2115 Some(version) => {
2116 return contract_error(format!(
2117 "training replay output schema_version {version} is unsupported; current {BOUND_TRAINING_OUTPUT_SCHEMA_VERSION}"
2118 ));
2119 }
2120 None => {
2121 return contract_error(
2122 "training replay output requires bound_training_output schema_version",
2123 );
2124 }
2125 }
2126 let binding_id = output.binding.binding_id.as_str();
2127 if previous.is_some_and(|previous| previous >= binding_id) {
2128 return contract_error("training replay outputs must be strictly sorted by binding_id");
2129 }
2130 previous = Some(binding_id);
2131 }
2132 Ok(())
2133}
2134
2135fn validate_replay_bound_output_blocks(output: &BoundTrainingOutput) -> Result<()> {
2136 for block in &output.predictions {
2137 validate_optional_port(
2138 "training replay prediction producer_port",
2139 &block.producer_port,
2140 )?;
2141 if block.partition != crate::oof::PredictionPartition::Final || block.fold_id.is_some() {
2142 return contract_error(
2143 "training replay prediction blocks must use final partition without fold",
2144 );
2145 }
2146 }
2147 for block in &output.observation_predictions {
2148 validate_optional_port(
2149 "training replay observation prediction producer_port",
2150 &block.producer_port,
2151 )?;
2152 if block.partition != crate::oof::PredictionPartition::Final || block.fold_id.is_some() {
2153 return contract_error(
2154 "training replay observation prediction blocks must use final partition without fold",
2155 );
2156 }
2157 }
2158 for block in &output.aggregated_predictions {
2159 validate_optional_port(
2160 "training replay aggregated prediction producer_port",
2161 &block.producer_port,
2162 )?;
2163 if block.partition != crate::oof::PredictionPartition::Final || block.fold_id.is_some() {
2164 return contract_error(
2165 "training replay aggregated prediction blocks must use final partition without fold",
2166 );
2167 }
2168 }
2169 Ok(())
2170}
2171
2172fn validate_replay_phase(phase: Phase) -> Result<()> {
2173 if matches!(phase, Phase::Predict | Phase::Explain) {
2174 Ok(())
2175 } else {
2176 contract_error("training replay supports only PREDICT and EXPLAIN")
2177 }
2178}
2179
2180fn validate_sha256(label: &str, value: &str) -> Result<()> {
2181 if value.len() == 64
2182 && value
2183 .bytes()
2184 .all(|byte| byte.is_ascii_digit() || (b'a'..=b'f').contains(&byte))
2185 {
2186 Ok(())
2187 } else {
2188 contract_error(format!(
2189 "{label} fingerprint must be 64 lowercase hexadecimal characters"
2190 ))
2191 }
2192}
2193
2194fn validate_identifier(label: &str, value: &str) -> Result<()> {
2195 if !value.is_empty()
2196 && value.len() <= 128
2197 && value
2198 .bytes()
2199 .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'_' | b'-' | b'.' | b':'))
2200 {
2201 Ok(())
2202 } else {
2203 contract_error(format!("{label} is not a valid DAG-ML identifier"))
2204 }
2205}
2206
2207fn validate_non_empty(label: &str, value: &str) -> Result<()> {
2208 if value.trim().is_empty() {
2209 contract_error(format!("{label} must be non-empty"))
2210 } else {
2211 Ok(())
2212 }
2213}
2214
2215fn validate_sorted_unique_identifiers(
2216 label: &str,
2217 values: &[String],
2218 require_non_empty: bool,
2219) -> Result<()> {
2220 validate_sorted_unique_text(label, values, require_non_empty)?;
2221 for value in values {
2222 validate_identifier(label, value)?;
2223 }
2224 Ok(())
2225}
2226
2227fn validate_sorted_unique_text(
2228 label: &str,
2229 values: &[String],
2230 require_non_empty: bool,
2231) -> Result<()> {
2232 if require_non_empty && values.is_empty() {
2233 return contract_error(format!("{label} must be non-empty"));
2234 }
2235 let mut previous: Option<&str> = None;
2236 for value in values {
2237 validate_non_empty(label, value)?;
2238 if previous.is_some_and(|previous| previous >= value.as_str()) {
2239 return contract_error(format!("{label} must be strictly sorted and unique"));
2240 }
2241 previous = Some(value.as_str());
2242 }
2243 Ok(())
2244}
2245
2246fn validate_sorted_unique_keys<'a>(
2247 label: &str,
2248 values: impl Iterator<Item = &'a str>,
2249 require_non_empty: bool,
2250) -> Result<()> {
2251 let values = values.collect::<Vec<_>>();
2252 if require_non_empty && values.is_empty() {
2253 return contract_error(format!("{label} must be non-empty"));
2254 }
2255 let mut previous: Option<&str> = None;
2256 for value in values {
2257 validate_non_empty(label, value)?;
2258 if previous.is_some_and(|previous| previous >= value) {
2259 return contract_error(format!("{label} must be strictly sorted and unique"));
2260 }
2261 previous = Some(value);
2262 }
2263 Ok(())
2264}
2265
2266fn validate_optional_port(label: &str, value: &Option<String>) -> Result<()> {
2267 match value {
2268 Some(value) if !value.trim().is_empty() => Ok(()),
2269 _ => contract_error(format!("{label} must be present and non-empty")),
2270 }
2271}
2272
2273fn validate_diagnostics(diagnostics: &BTreeMap<String, serde_json::Value>) -> Result<()> {
2274 for (key, value) in diagnostics {
2275 validate_non_empty("training replay diagnostic key", key)?;
2276 if !matches!(
2277 value,
2278 serde_json::Value::Null
2279 | serde_json::Value::Bool(_)
2280 | serde_json::Value::Number(_)
2281 | serde_json::Value::String(_)
2282 ) {
2283 return contract_error("training replay diagnostics must be scalar JSON values");
2284 }
2285 }
2286 Ok(())
2287}
2288
2289fn require_count(label: &str, actual: usize, expected: usize) -> Result<()> {
2290 if actual == expected {
2291 Ok(())
2292 } else {
2293 contract_error(format!("{label} does not match replay payload"))
2294 }
2295}
2296
2297fn zero_fingerprint() -> String {
2298 "0".repeat(64)
2299}
2300
2301fn tcv1_fingerprint_without<T: Serialize>(value: &T, field: &str, label: &str) -> Result<String> {
2302 let json = serde_json::to_string(value)?;
2303 strict_tcv1_fingerprint_without(&json, field, label)
2304}
2305
2306fn strict_tcv1_fingerprint_without(json: &str, field: &str, label: &str) -> Result<String> {
2307 parse_typed_json(json)
2308 .and_then(|value| value.fingerprint_without(field))
2309 .map_err(|error| {
2310 DagMlError::RuntimeValidation(format!("{label} is outside strict TCV1: {error}"))
2311 })
2312}
2313
2314fn unsupported_version<T>(label: &str, actual: u32, expected: u32) -> Result<T> {
2315 contract_error(format!(
2316 "{label} uses unsupported schema_version {actual}, expected {expected}"
2317 ))
2318}
2319
2320fn contract_error<T>(message: impl Into<String>) -> Result<T> {
2321 Err(DagMlError::CampaignValidation(message.into()))
2322}
2323
2324#[cfg(test)]
2325mod methods_hpo_resume_state_tests {
2326 use super::*;
2327
2328 #[test]
2329 fn methods_hpo_resume_state_reader_refuses_unknown_legacy_and_noncanonical_json() {
2330 let unknown = r#"{"schema_version":1,"unknown_resume_side_channel":true}"#;
2331 assert!(methods_hpo_resume_state_from_json(unknown).is_err());
2332
2333 let duplicate = r#"{"schema_version":1,"schema_version":1}"#;
2336 let error = methods_hpo_resume_state_from_json(duplicate)
2337 .unwrap_err()
2338 .to_string();
2339 assert!(error.contains("duplicate JSON object key"), "{error}");
2340
2341 let legacy = r#"{"schema_version":1,"tuner_node_id":"tuner:legacy"}"#;
2345 assert!(methods_hpo_resume_state_from_json(legacy).is_err());
2346 }
2347
2348 #[test]
2349 fn methods_hpo_resume_state_reader_refuses_tampered_schema_version() {
2350 let tampered = r#"{"schema_version":2,"operation_id":"campaign:hpo"}"#;
2351 assert!(methods_hpo_resume_state_from_json(tampered).is_err());
2352 }
2353}
2354
2355#[cfg(test)]
2356mod replay_identity_tests {
2357 use super::*;
2358
2359 fn replay_identity(target_content_fingerprint: Option<&str>) -> ReplayDataIdentity {
2360 let mut identity = ReplayDataIdentity {
2361 requirement_key: "model:base.X".to_string(),
2362 schema_fingerprint: "1".repeat(64),
2363 plan_fingerprint: "2".repeat(64),
2364 relation_fingerprint: "3".repeat(64),
2365 data_content_fingerprint: "4".repeat(64),
2366 target_content_fingerprint: target_content_fingerprint.map(str::to_string),
2367 identity_fingerprint: zero_fingerprint(),
2368 };
2369 identity.identity_fingerprint = identity.compute_fingerprint().unwrap();
2370 identity
2371 }
2372
2373 #[test]
2374 fn replay_identity_permits_an_unlabeled_predict_cohort_without_a_sentinel() {
2375 let identity = replay_identity(None);
2376 identity.validate().unwrap();
2377 assert!(identity.target_content_fingerprint.is_none());
2378 assert!(serde_json::to_value(&identity)
2379 .unwrap()
2380 .get("target_content_fingerprint")
2381 .unwrap()
2382 .is_null());
2383 }
2384
2385 #[test]
2386 fn replay_identity_rejects_a_resigned_target_fingerprint_tamper() {
2387 let mut identity = replay_identity(None);
2388 identity.target_content_fingerprint = Some("not-a-fingerprint".to_string());
2389 identity.identity_fingerprint = identity.compute_fingerprint().unwrap();
2390 let error = identity.validate().unwrap_err().to_string();
2391 assert!(error.contains("target content"), "{error}");
2392 }
2393
2394 #[test]
2395 fn replay_schema_version_is_v3_for_target_bound_and_target_free_cohorts() {
2396 assert_eq!(
2397 replay_outcome_schema_version(&[replay_identity(None)]),
2398 TRAINING_REPLAY_OUTCOME_SCHEMA_VERSION
2399 );
2400 assert_eq!(
2401 replay_outcome_schema_version(&[replay_identity(Some(&"5".repeat(64)))]),
2402 TRAINING_REPLAY_OUTCOME_SCHEMA_VERSION
2403 );
2404 }
2405}
2406
2407#[cfg(test)]
2408mod tests {
2409 #[cfg(dag_ml_workspace_contract_fixtures)]
2410 use std::fs;
2411 #[cfg(dag_ml_workspace_contract_fixtures)]
2412 use std::path::PathBuf;
2413
2414 #[cfg(dag_ml_workspace_contract_fixtures)]
2415 use super::*;
2416
2417 #[cfg(dag_ml_workspace_contract_fixtures)]
2418 fn root() -> PathBuf {
2419 PathBuf::from(env!("CARGO_MANIFEST_DIR"))
2420 .parent()
2421 .and_then(|path| path.parent())
2422 .expect("core crate is under crates/dag-ml-core")
2423 .to_path_buf()
2424 }
2425
2426 #[cfg(dag_ml_workspace_contract_fixtures)]
2427 fn fixture(name: &str) -> String {
2428 fs::read_to_string(
2429 root()
2430 .join("examples")
2431 .join("fixtures")
2432 .join("training")
2433 .join("replay")
2434 .join(name),
2435 )
2436 .expect(name)
2437 }
2438
2439 #[cfg(dag_ml_workspace_contract_fixtures)]
2440 fn training_fixture(name: &str) -> String {
2441 fs::read_to_string(
2442 root()
2443 .join("examples")
2444 .join("fixtures")
2445 .join("training")
2446 .join(name),
2447 )
2448 .expect(name)
2449 }
2450
2451 #[cfg(dag_ml_workspace_contract_fixtures)]
2452 #[test]
2453 fn training_replay_contract_fixtures_parse_and_cross_validate() {
2454 let predict_source =
2455 TrainingOutcome::from_json(&training_fixture("training_outcome_refit.v1.json"))
2456 .expect("predict source training outcome");
2457 let explain_source =
2458 TrainingOutcome::from_json(&fixture("training_replay_source_outcome_explain.v1.json"))
2459 .expect("explain source training outcome");
2460 let predict_request =
2461 TrainingReplayRequest::from_json(&fixture("training_replay_request_predict.v1.json"))
2462 .expect("predict request");
2463 let predict_outcome =
2464 TrainingReplayOutcome::from_json(&fixture("training_replay_outcome_predict.v1.json"))
2465 .expect("predict outcome");
2466 predict_outcome
2467 .validate_against(&predict_source, &predict_request)
2468 .expect("predict cross-links");
2469
2470 let explain_request =
2471 TrainingReplayRequest::from_json(&fixture("training_replay_request_explain.v1.json"))
2472 .expect("explain request");
2473 let explain_outcome =
2474 TrainingReplayOutcome::from_json(&fixture("training_replay_outcome_explain.v1.json"))
2475 .expect("explain outcome");
2476 explain_outcome
2477 .validate_against(&explain_source, &explain_request)
2478 .expect("explain cross-links");
2479
2480 let explain_only = TrainingReplayOutcome::from_json(&fixture(
2481 "training_replay_outcome_explain_only.v1.json",
2482 ))
2483 .expect("explain-only outcome");
2484 explain_only
2485 .validate_against(&explain_source, &explain_request)
2486 .expect("explain-only cross-links");
2487 }
2488
2489 #[cfg(dag_ml_workspace_contract_fixtures)]
2490 #[test]
2491 fn training_replay_request_rejects_refit_and_unsorted_bindings() {
2492 let mut request: serde_json::Value =
2493 serde_json::from_str(&fixture("training_replay_request_predict.v1.json")).unwrap();
2494 request["phase"] = serde_json::Value::String("REFIT".to_string());
2495 let err = serde_json::from_value::<TrainingReplayRequest>(request)
2496 .unwrap()
2497 .validate()
2498 .unwrap_err()
2499 .to_string();
2500 assert!(err.contains("PREDICT and EXPLAIN"));
2501
2502 let mut request: TrainingReplayRequest =
2503 TrainingReplayRequest::from_json(&fixture("training_replay_request_predict.v1.json"))
2504 .unwrap();
2505 request.output_binding_ids = vec!["z".to_string(), "a".to_string()];
2506 request.request_fingerprint = request.compute_fingerprint().unwrap();
2507 let err = request.validate().unwrap_err().to_string();
2508 assert!(err.contains("strictly sorted"));
2509 }
2510
2511 #[cfg(dag_ml_workspace_contract_fixtures)]
2512 #[test]
2513 fn training_replay_outcome_rejects_counter_and_source_transplants() {
2514 let source =
2515 TrainingOutcome::from_json(&training_fixture("training_outcome_refit.v1.json"))
2516 .unwrap();
2517 let request =
2518 TrainingReplayRequest::from_json(&fixture("training_replay_request_predict.v1.json"))
2519 .unwrap();
2520 let mut outcome =
2521 TrainingReplayOutcome::from_json(&fixture("training_replay_outcome_predict.v1.json"))
2522 .unwrap();
2523 outcome.prediction_block_count += 1;
2524 outcome.outcome_fingerprint = outcome.compute_fingerprint().unwrap();
2525 let err = outcome.validate().unwrap_err().to_string();
2526 assert!(err.contains("prediction_block_count"));
2527
2528 let mut outcome =
2529 TrainingReplayOutcome::from_json(&fixture("training_replay_outcome_predict.v1.json"))
2530 .unwrap();
2531 outcome.source_training_outcome.outcome_fingerprint = "f".repeat(64);
2532 outcome.outcome_fingerprint = outcome.compute_fingerprint().unwrap();
2533 let err = outcome
2534 .validate_against(&source, &request)
2535 .unwrap_err()
2536 .to_string();
2537 assert!(err.contains("source reference"));
2538 }
2539
2540 #[cfg(dag_ml_workspace_contract_fixtures)]
2541 #[test]
2542 fn target_free_predict_evidence_is_v3_only() {
2543 let mut outcome =
2544 TrainingReplayOutcome::from_json(&fixture("training_replay_outcome_predict.v1.json"))
2545 .unwrap();
2546 outcome.schema_version = TRAINING_REPLAY_OUTCOME_SCHEMA_VERSION;
2547 for identity in &mut outcome.input_data_identities {
2548 identity.target_content_fingerprint = None;
2549 identity.identity_fingerprint = identity.compute_fingerprint().unwrap();
2550 }
2551 outcome.outcome_fingerprint = outcome.compute_fingerprint().unwrap();
2552 outcome.validate().unwrap();
2553
2554 outcome.schema_version = CONFORMAL_TRAINING_REPLAY_OUTCOME_SCHEMA_VERSION;
2555 outcome.outcome_fingerprint = outcome.compute_fingerprint().unwrap();
2556 let error = outcome.validate().unwrap_err().to_string();
2557 assert!(error.contains("V1/V2 requires target-bound"), "{error}");
2558 }
2559}