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, ConformalCalibrationContext, ConformalCalibrationTruth,
19 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, TrainingOutcome, BOUND_TRAINING_OUTPUT_SCHEMA_VERSION,
36};
37
38pub const TRAINING_REPLAY_REQUEST_SCHEMA_VERSION: u32 = 1;
39pub 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
117#[derive(Clone, Debug, Eq, PartialEq, Serialize, Deserialize)]
118#[serde(deny_unknown_fields)]
119pub struct TrainingReplayRequest {
120 pub schema_version: u32,
121 pub request_id: String,
122 pub source_outcome_fingerprint: String,
123 pub phase: Phase,
124 pub data_envelope_keys: Vec<String>,
125 pub output_binding_ids: Vec<String>,
126 pub request_fingerprint: String,
127}
128
129impl TrainingReplayRequest {
130 pub fn from_json(json: &str) -> Result<Self> {
131 let raw_fingerprint = strict_tcv1_fingerprint_without(
132 json,
133 "request_fingerprint",
134 "training replay request",
135 )?;
136 let request: Self = serde_json::from_str(json)?;
137 if request.request_fingerprint != raw_fingerprint {
138 return contract_error(
139 "training replay request fingerprint does not match original TCV1 JSON",
140 );
141 }
142 request.validate()?;
143 Ok(request)
144 }
145
146 pub fn compute_fingerprint(&self) -> Result<String> {
147 tcv1_fingerprint_without(self, "request_fingerprint", "training replay request")
148 }
149
150 pub fn validate(&self) -> Result<()> {
151 if self.schema_version != TRAINING_REPLAY_REQUEST_SCHEMA_VERSION {
152 return unsupported_version(
153 "training replay request",
154 self.schema_version,
155 TRAINING_REPLAY_REQUEST_SCHEMA_VERSION,
156 );
157 }
158 validate_identifier("training replay request_id", &self.request_id)?;
159 validate_sha256(
160 "training replay source outcome",
161 &self.source_outcome_fingerprint,
162 )?;
163 validate_replay_phase(self.phase)?;
164 validate_sorted_unique_text(
165 "training replay data_envelope_keys",
166 &self.data_envelope_keys,
167 true,
168 )?;
169 validate_sorted_unique_identifiers(
170 "training replay output_binding_ids",
171 &self.output_binding_ids,
172 true,
173 )?;
174 validate_sha256("training replay request", &self.request_fingerprint)?;
175 if self.request_fingerprint != self.compute_fingerprint()? {
176 return contract_error(
177 "training replay request fingerprint does not match TCV1 content",
178 );
179 }
180 Ok(())
181 }
182}
183
184#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
185#[serde(deny_unknown_fields)]
186pub struct TrainingReplayOutcome {
187 pub schema_version: u32,
188 pub outcome_id: String,
189 pub run_id: RunId,
190 pub source_training_outcome: TrainingOutcomeRef,
191 pub replay_request_id: String,
192 pub replay_request_fingerprint: String,
193 pub input_data_identities: Vec<ReplayDataIdentity>,
194 pub bundle_id: BundleId,
195 pub plan_id: String,
196 pub phase: Phase,
197 pub result_count: usize,
198 pub lineage_record_count: usize,
199 pub prediction_block_count: usize,
200 pub observation_prediction_block_count: usize,
201 pub aggregated_prediction_block_count: usize,
202 pub explanation_block_count: usize,
203 pub controller_count: usize,
204 pub prediction_cache_store: bool,
205 pub outputs: Vec<BoundTrainingOutput>,
206 #[serde(default, skip_serializing_if = "Vec::is_empty")]
207 pub conformal_intervals: Vec<ConformalIntervalBlock>,
208 pub explanations: Vec<ExplanationBlock>,
209 pub lineage: Vec<LineageRecord>,
210 pub warnings: Vec<String>,
211 pub diagnostics: BTreeMap<String, serde_json::Value>,
212 pub outcome_fingerprint: String,
213}
214
215#[derive(Clone, Debug, Eq, PartialEq, Serialize, Deserialize)]
224#[serde(deny_unknown_fields)]
225pub struct ReplayDataIdentity {
226 pub requirement_key: String,
227 pub schema_fingerprint: String,
228 pub plan_fingerprint: String,
229 pub relation_fingerprint: String,
230 pub data_content_fingerprint: String,
231 #[serde(default)]
232 pub target_content_fingerprint: Option<String>,
233 pub identity_fingerprint: String,
234}
235
236impl ReplayDataIdentity {
237 pub fn compute_fingerprint(&self) -> Result<String> {
240 tcv1_fingerprint_without(self, "identity_fingerprint", "replay data identity")
241 }
242
243 fn validate(&self) -> Result<()> {
244 validate_non_empty("replay data requirement_key", &self.requirement_key)?;
245 for (label, value) in [
246 ("replay data schema", &self.schema_fingerprint),
247 ("replay data plan", &self.plan_fingerprint),
248 ("replay data relation", &self.relation_fingerprint),
249 ("replay data content", &self.data_content_fingerprint),
250 ("replay data identity", &self.identity_fingerprint),
251 ] {
252 validate_sha256(label, value)?;
253 }
254 if let Some(target_content_fingerprint) = &self.target_content_fingerprint {
255 validate_sha256("replay target content", target_content_fingerprint)?;
256 }
257 if self.identity_fingerprint != self.compute_fingerprint()? {
258 return contract_error("replay data identity fingerprint does not match TCV1 content");
259 }
260 Ok(())
261 }
262}
263
264impl TrainingReplayOutcome {
265 pub fn from_json(json: &str) -> Result<Self> {
266 let raw_fingerprint = strict_tcv1_fingerprint_without(
267 json,
268 "outcome_fingerprint",
269 "training replay outcome",
270 )?;
271 let outcome: Self = serde_json::from_str(json)?;
272 if outcome.outcome_fingerprint != raw_fingerprint {
273 return contract_error(
274 "training replay outcome fingerprint does not match original TCV1 JSON",
275 );
276 }
277 outcome.validate()?;
278 Ok(outcome)
279 }
280
281 pub fn compute_fingerprint(&self) -> Result<String> {
282 tcv1_fingerprint_without(self, "outcome_fingerprint", "training replay outcome")
283 }
284
285 pub fn validate(&self) -> Result<()> {
286 if self.schema_version < MIN_READABLE_TRAINING_REPLAY_OUTCOME_SCHEMA_VERSION
287 || self.schema_version > TRAINING_REPLAY_OUTCOME_SCHEMA_VERSION
288 {
289 return unsupported_version(
290 "training replay outcome",
291 self.schema_version,
292 TRAINING_REPLAY_OUTCOME_SCHEMA_VERSION,
293 );
294 }
295 validate_identifier("training replay outcome_id", &self.outcome_id)?;
296 validate_identifier("training replay request_id", &self.replay_request_id)?;
297 validate_sha256("training replay request", &self.replay_request_fingerprint)?;
298 validate_non_empty("training replay plan_id", &self.plan_id)?;
299 validate_replay_phase(self.phase)?;
300 self.source_training_outcome.validate()?;
301 for identity in &self.input_data_identities {
302 identity.validate()?;
303 }
304 if self.schema_version < TRAINING_REPLAY_OUTCOME_SCHEMA_VERSION
305 && self
306 .input_data_identities
307 .iter()
308 .any(|identity| identity.target_content_fingerprint.is_none())
309 {
310 return contract_error(
311 "training replay outcome V1/V2 requires target-bound input identities; migrate target-free PREDICT/EXPLAIN evidence to V3",
312 );
313 }
314 validate_sorted_unique_keys(
315 "training replay input_data_identities",
316 self.input_data_identities
317 .iter()
318 .map(|identity| identity.requirement_key.as_str()),
319 true,
320 )?;
321 if self.prediction_cache_store {
322 return contract_error("training replay outcome cannot persist a prediction cache");
323 }
324 validate_sorted_unique_text("training replay warnings", &self.warnings, false)?;
325 validate_diagnostics(&self.diagnostics)?;
326 validate_output_order_and_version(&self.outputs)?;
327 if self.schema_version == LEGACY_TRAINING_REPLAY_OUTCOME_SCHEMA_VERSION
328 && (!self.conformal_intervals.is_empty()
329 || self
330 .source_training_outcome
331 .pre_conformal_outcome_fingerprint
332 .is_some())
333 {
334 return contract_error(
335 "training replay outcome V1 cannot carry conformal state; migrate to V2 or V3",
336 );
337 }
338 if self
339 .input_data_identities
340 .iter()
341 .any(|identity| identity.target_content_fingerprint.is_none())
342 && !self.conformal_intervals.is_empty()
343 {
344 return contract_error(
345 "target-free training replay evidence cannot carry conformal intervals",
346 );
347 }
348 for output in &self.outputs {
349 validate_replay_bound_output_blocks(output)?;
350 }
351 for interval in &self.conformal_intervals {
352 let output = self
353 .outputs
354 .iter()
355 .find(|output| output.binding.binding_id == interval.binding_id)
356 .ok_or_else(|| {
357 DagMlError::RuntimeValidation(
358 "conformal interval references an absent replay output binding".to_string(),
359 )
360 })?;
361 let point = output
362 .predictions
363 .iter()
364 .find(|block| block.sample_ids == interval.sample_ids)
365 .ok_or_else(|| {
366 DagMlError::RuntimeValidation(
367 "conformal interval has no matching replay point prediction block"
368 .to_string(),
369 )
370 })?;
371 interval.validate()?;
372 if interval.point_prediction_fingerprint
373 != crate::conformal_runtime::point_prediction_fingerprint_for_runtime(point)?
374 {
375 return Err(DagMlError::RuntimeValidation(
376 "conformal interval point prediction fingerprint does not match replay output"
377 .to_string(),
378 ));
379 }
380 }
381 for explanation in &self.explanations {
382 explanation.validate()?;
383 validate_optional_port(
384 "training replay explanation producer_port",
385 &explanation.producer_port,
386 )?;
387 }
388 for record in &self.lineage {
389 record.validate()?;
390 }
391 match self.phase {
392 Phase::Predict if self.outputs.is_empty() => {
393 return contract_error("training replay PREDICT requires at least one output");
394 }
395 Phase::Predict if !self.explanations.is_empty() => {
396 return contract_error("training replay PREDICT cannot emit explanations");
397 }
398 Phase::Explain if self.explanations.is_empty() => {
399 return contract_error("training replay EXPLAIN requires at least one explanation");
400 }
401 _ => {}
402 }
403 self.validate_counters()?;
404 validate_sha256("training replay outcome", &self.outcome_fingerprint)?;
405 if self.outcome_fingerprint != self.compute_fingerprint()? {
406 return contract_error(
407 "training replay outcome fingerprint does not match TCV1 content",
408 );
409 }
410 Ok(())
411 }
412
413 pub fn validate_against(
414 &self,
415 source: &TrainingOutcome,
416 request: &TrainingReplayRequest,
417 ) -> Result<()> {
418 self.validate()?;
419 source.validate()?;
420 request.validate()?;
421 if request.source_outcome_fingerprint != source.outcome_fingerprint {
422 return contract_error("training replay request does not target source outcome");
423 }
424 if !source.replayable_phases.contains(&request.phase) {
425 return contract_error("training replay phase is not replayable by source outcome");
426 }
427 if self.source_training_outcome != source.to_reference()? {
428 return contract_error(
429 "training replay outcome source reference does not match source outcome",
430 );
431 }
432 if self.replay_request_id != request.request_id {
433 return contract_error(
434 "training replay outcome request id does not match ReplayRequest",
435 );
436 }
437 if self.replay_request_fingerprint != request.request_fingerprint {
438 return contract_error(
439 "training replay outcome request fingerprint does not match ReplayRequest",
440 );
441 }
442 if self.phase != request.phase {
443 return contract_error("training replay outcome phase does not match ReplayRequest");
444 }
445 if self.bundle_id != source.execution_bundle.bundle_id {
446 return contract_error("training replay outcome bundle does not match source outcome");
447 }
448 if self.plan_id != source.effective_plan.id {
449 return contract_error("training replay outcome plan does not match source outcome");
450 }
451 let identity_keys = self
452 .input_data_identities
453 .iter()
454 .map(|identity| identity.requirement_key.clone())
455 .collect::<Vec<_>>();
456 if identity_keys != request.data_envelope_keys {
457 return contract_error(
458 "training replay outcome identities do not exactly cover ReplayRequest envelopes",
459 );
460 }
461 let source_bindings = source
462 .outputs
463 .iter()
464 .map(|output| (output.binding.binding_id.as_str(), &output.binding))
465 .collect::<BTreeMap<_, _>>();
466 for binding_id in &request.output_binding_ids {
467 if !source_bindings.contains_key(binding_id.as_str()) {
468 return contract_error(
469 "training replay request references absent source output binding",
470 );
471 }
472 }
473 let emitted_binding_ids = self
474 .outputs
475 .iter()
476 .map(|output| output.binding.binding_id.clone())
477 .collect::<Vec<_>>();
478 if self.phase == Phase::Predict && emitted_binding_ids != request.output_binding_ids {
479 return contract_error(
480 "training replay PREDICT outputs do not exactly cover ReplayRequest bindings",
481 );
482 }
483 if self.phase == Phase::Explain
484 && !emitted_binding_ids
485 .iter()
486 .all(|binding_id| request.output_binding_ids.contains(binding_id))
487 {
488 return contract_error(
489 "training replay EXPLAIN outputs include a binding outside ReplayRequest",
490 );
491 }
492 for output in &self.outputs {
493 let Some(source_binding) = source_bindings.get(output.binding.binding_id.as_str())
494 else {
495 return contract_error(
496 "training replay output binding is absent from source outcome",
497 );
498 };
499 if &output.binding != *source_binding {
500 return contract_error(
501 "training replay output binding does not match source outcome binding",
502 );
503 }
504 output.validate(&source.effective_plan)?;
505 }
506 match &source.conformal_calibration {
507 None if !self.conformal_intervals.is_empty() => {
508 return contract_error(
509 "training replay intervals require source calibration context",
510 )
511 }
512 Some(calibration) => validate_replay_interval_closure(
513 calibration,
514 &self.outputs,
515 &self.conformal_intervals,
516 )?,
517 None => {}
518 }
519 Ok(())
520 }
521
522 pub fn validate_against_package(
523 &self,
524 package: &PortablePredictorPackage,
525 request: &TrainingReplayRequest,
526 ) -> Result<()> {
527 self.validate()?;
528 package.validate()?;
529 request.validate()?;
530 validate_replay_phase(request.phase)?;
531 if request.source_outcome_fingerprint != package.training_outcome.outcome_fingerprint {
532 return contract_error(
533 "training replay request does not target package source outcome",
534 );
535 }
536 if self.source_training_outcome != package.training_outcome {
537 return contract_error(
538 "training replay outcome source reference does not match package source outcome",
539 );
540 }
541 if self.replay_request_id != request.request_id {
542 return contract_error(
543 "training replay outcome request id does not match ReplayRequest",
544 );
545 }
546 if self.replay_request_fingerprint != request.request_fingerprint {
547 return contract_error(
548 "training replay outcome request fingerprint does not match ReplayRequest",
549 );
550 }
551 if self.phase != request.phase {
552 return contract_error("training replay outcome phase does not match ReplayRequest");
553 }
554 if self.bundle_id != package.execution_bundle.bundle_id {
555 return contract_error("training replay outcome bundle does not match package");
556 }
557 if self.plan_id != package.effective_plan.id {
558 return contract_error("training replay outcome plan does not match package");
559 }
560 let identity_keys = self
561 .input_data_identities
562 .iter()
563 .map(|identity| identity.requirement_key.clone())
564 .collect::<Vec<_>>();
565 if identity_keys != request.data_envelope_keys {
566 return contract_error(
567 "training replay outcome identities do not exactly cover ReplayRequest envelopes",
568 );
569 }
570 let package_bindings = package
571 .output_bindings
572 .iter()
573 .map(|binding| (binding.binding_id.as_str(), binding))
574 .collect::<BTreeMap<_, _>>();
575 for binding_id in &request.output_binding_ids {
576 if !package_bindings.contains_key(binding_id.as_str()) {
577 return contract_error(
578 "training replay request references absent package output binding",
579 );
580 }
581 }
582 let emitted_binding_ids = self
583 .outputs
584 .iter()
585 .map(|output| output.binding.binding_id.clone())
586 .collect::<Vec<_>>();
587 if self.phase == Phase::Predict && emitted_binding_ids != request.output_binding_ids {
588 return contract_error(
589 "training replay PREDICT outputs do not exactly cover ReplayRequest bindings",
590 );
591 }
592 if self.phase == Phase::Explain
593 && !emitted_binding_ids
594 .iter()
595 .all(|binding_id| request.output_binding_ids.contains(binding_id))
596 {
597 return contract_error(
598 "training replay EXPLAIN outputs include a binding outside ReplayRequest",
599 );
600 }
601 for output in &self.outputs {
602 let Some(package_binding) = package_bindings.get(output.binding.binding_id.as_str())
603 else {
604 return contract_error("training replay output binding is absent from package");
605 };
606 if &output.binding != *package_binding {
607 return contract_error(
608 "training replay output binding does not match package binding",
609 );
610 }
611 output.validate(&package.effective_plan)?;
612 }
613 match &package.conformal_calibration {
614 None if !self.conformal_intervals.is_empty() => {
615 return contract_error(
616 "training replay intervals require package calibration context",
617 )
618 }
619 Some(calibration) => validate_replay_interval_closure(
620 calibration,
621 &self.outputs,
622 &self.conformal_intervals,
623 )?,
624 None => {}
625 }
626 Ok(())
627 }
628
629 fn validate_counters(&self) -> Result<()> {
630 require_count(
631 "training replay result_count",
632 self.result_count,
633 self.lineage.len(),
634 )?;
635 require_count(
636 "training replay lineage_record_count",
637 self.lineage_record_count,
638 self.lineage.len(),
639 )?;
640 require_count(
641 "training replay prediction_block_count",
642 self.prediction_block_count,
643 self.outputs
644 .iter()
645 .map(|output| output.predictions.len())
646 .sum(),
647 )?;
648 require_count(
649 "training replay observation_prediction_block_count",
650 self.observation_prediction_block_count,
651 self.outputs
652 .iter()
653 .map(|output| output.observation_predictions.len())
654 .sum(),
655 )?;
656 require_count(
657 "training replay aggregated_prediction_block_count",
658 self.aggregated_prediction_block_count,
659 self.outputs
660 .iter()
661 .map(|output| output.aggregated_predictions.len())
662 .sum(),
663 )?;
664 require_count(
665 "training replay explanation_block_count",
666 self.explanation_block_count,
667 self.explanations.len(),
668 )?;
669 let controller_count = self
670 .lineage
671 .iter()
672 .map(|record| record.controller_id.as_str())
673 .collect::<BTreeSet<_>>()
674 .len();
675 require_count(
676 "training replay controller_count",
677 self.controller_count,
678 controller_count,
679 )?;
680 Ok(())
681 }
682}
683
684pub(crate) fn replay_request_from_outcome(replay: &TrainingReplayOutcome) -> TrainingReplayRequest {
688 TrainingReplayRequest {
689 schema_version: TRAINING_REPLAY_REQUEST_SCHEMA_VERSION,
690 request_id: replay.replay_request_id.clone(),
691 source_outcome_fingerprint: replay.source_training_outcome.outcome_fingerprint.clone(),
692 phase: replay.phase,
693 data_envelope_keys: replay
694 .input_data_identities
695 .iter()
696 .map(|identity| identity.requirement_key.clone())
697 .collect(),
698 output_binding_ids: replay
699 .outputs
700 .iter()
701 .map(|output| output.binding.binding_id.clone())
702 .collect(),
703 request_fingerprint: replay.replay_request_fingerprint.clone(),
704 }
705}
706
707struct LoadedPredictorArtifactStore<'a> {
708 predictor: &'a LoadedPredictor<HandleRef>,
709 records: BTreeMap<ArtifactId, crate::bundle::RefitArtifactRecord>,
710}
711
712struct BundlePayloadArtifactStore<'a> {
720 bundle: &'a ExecutionBundle,
721 controllers: &'a RuntimeControllerRegistry,
722 fallback: &'a dyn RuntimeArtifactStore,
723 hydrated_handles: Mutex<Vec<(ControllerId, HandleRef)>>,
724}
725
726impl RuntimeArtifactStore for BundlePayloadArtifactStore<'_> {
727 fn materialize(&self, request: &ArtifactMaterializationRequest) -> Result<HandleRef> {
728 let Some(payload) = self.bundle.raw_artifact_payloads.get(&request.artifact.id) else {
729 return self.fallback.materialize(request);
730 };
731 let controller = self
732 .controllers
733 .get(&request.controller_id)
734 .ok_or_else(|| {
735 DagMlError::RuntimeValidation(format!(
736 "bundle `{}` has no registered controller `{}` to hydrate raw artifact `{}`",
737 self.bundle.bundle_id, request.controller_id, request.artifact.id
738 ))
739 })?;
740 let handle = controller.hydrate_artifact_payload(request, payload)?;
741 self.hydrated_handles
742 .lock()
743 .map_err(|_| {
744 DagMlError::RuntimeValidation(
745 "bundle payload hydrated-handle registry lock poisoned".to_string(),
746 )
747 })?
748 .push((request.controller_id.clone(), handle.clone()));
749 Ok(handle)
750 }
751}
752
753impl BundlePayloadArtifactStore<'_> {
754 fn release_hydrated_handles(&self) -> Result<()> {
755 let handles = {
756 let mut handles = self.hydrated_handles.lock().map_err(|_| {
757 DagMlError::RuntimeValidation(
758 "bundle payload hydrated-handle registry lock poisoned".to_string(),
759 )
760 })?;
761 std::mem::take(&mut *handles)
762 };
763 let mut failures = Vec::new();
764 for (controller_id, handle) in handles.into_iter().rev() {
765 let release = self
766 .controllers
767 .get(&controller_id)
768 .ok_or_else(|| {
769 DagMlError::RuntimeValidation(format!(
770 "hydrated artifact owner controller `{controller_id}` is no longer registered"
771 ))
772 })
773 .and_then(|controller| controller.release_hydrated_artifact_payload(&handle));
774 if let Err(error) = release {
775 failures.push(error.to_string());
776 }
777 }
778 if failures.is_empty() {
779 Ok(())
780 } else {
781 Err(DagMlError::RuntimeValidation(format!(
782 "failed to release replay-hydrated artifact handles: {}",
783 failures.join("; ")
784 )))
785 }
786 }
787}
788
789fn finish_bundle_payload_replay<T>(
790 execution: Result<T>,
791 artifact_store: &BundlePayloadArtifactStore<'_>,
792) -> Result<T> {
793 let cleanup = artifact_store.release_hydrated_handles();
794 match (execution, cleanup) {
795 (Ok(value), Ok(())) => Ok(value),
796 (Err(error), Ok(())) => Err(error),
797 (Ok(_), Err(cleanup_error)) => Err(cleanup_error),
798 (Err(error), Err(cleanup_error)) => Err(DagMlError::RuntimeValidation(format!(
799 "{error}; replay hydration cleanup also failed: {cleanup_error}"
800 ))),
801 }
802}
803
804impl<'a> LoadedPredictorArtifactStore<'a> {
805 fn new(predictor: &'a LoadedPredictor<HandleRef>) -> Result<Self> {
806 predictor.package().validate()?;
807 let records = predictor
808 .package()
809 .execution_bundle
810 .refit_artifacts
811 .iter()
812 .map(|record| {
813 record.validate()?;
814 Ok((record.artifact.id.clone(), record.clone()))
815 })
816 .collect::<Result<BTreeMap<_, _>>>()?;
817 Ok(Self { predictor, records })
818 }
819}
820
821impl RuntimeArtifactStore for LoadedPredictorArtifactStore<'_> {
822 fn materialize(&self, request: &ArtifactMaterializationRequest) -> Result<HandleRef> {
823 let record = self.records.get(&request.artifact.id).ok_or_else(|| {
824 DagMlError::RuntimeValidation(format!(
825 "loaded predictor is missing refit artifact `{}` for bundle `{}`",
826 request.artifact.id, request.bundle_id
827 ))
828 })?;
829 if record.node_id != request.node_id {
830 return Err(DagMlError::RuntimeValidation(format!(
831 "artifact `{}` is registered for node `{}` but requested for `{}`",
832 request.artifact.id, record.node_id, request.node_id
833 )));
834 }
835 if record.controller_id != request.controller_id {
836 return Err(DagMlError::RuntimeValidation(format!(
837 "artifact `{}` is registered for controller `{}` but requested for `{}`",
838 request.artifact.id, record.controller_id, request.controller_id
839 )));
840 }
841 if record.artifact != request.artifact {
842 return Err(DagMlError::RuntimeValidation(format!(
843 "artifact `{}` metadata does not match package bundle record",
844 request.artifact.id
845 )));
846 }
847 if record.params_fingerprint != request.params_fingerprint {
848 return Err(DagMlError::RuntimeValidation(format!(
849 "artifact `{}` params fingerprint does not match package bundle record",
850 request.artifact.id
851 )));
852 }
853 let handle = self
854 .predictor
855 .artifact(&request.artifact.id)
856 .ok_or_else(|| {
857 DagMlError::RuntimeValidation(format!(
858 "loaded predictor has no process-local handle for `{}`",
859 request.artifact.id
860 ))
861 })?;
862 Ok(handle.clone())
863 }
864}
865
866pub fn execute_attached_training_replay(
867 input: AttachedTrainingReplayInput<'_>,
868) -> Result<TrainingReplayOutcome> {
869 input.source.validate()?;
870 input.request.validate()?;
871 validate_sorted_unique_text("training replay execution warnings", &input.warnings, false)?;
872 validate_diagnostics(&input.diagnostics)?;
873 if input.request.source_outcome_fingerprint != input.source.outcome_fingerprint {
874 return contract_error("training replay request does not target source outcome");
875 }
876 if !input
877 .source
878 .replayable_phases
879 .contains(&input.request.phase)
880 {
881 return contract_error("training replay phase is not replayable by source outcome");
882 }
883 for node_plan in input.source.effective_plan.node_plans.values() {
884 if input.controllers.get(&node_plan.controller_id).is_none() {
885 return Err(DagMlError::RuntimeValidation(format!(
886 "attached training replay controller `{}` for node `{}` is not registered",
887 node_plan.controller_id, node_plan.node_id
888 )));
889 }
890 }
891
892 let input_data_identities = replay_input_data_identities(
893 &input.source.execution_bundle,
894 input.request,
895 input.data_envelopes,
896 )?;
897 let (replay_plan, replay_bundle) = replay_plan_and_bundle_for_current_cohort(
898 &input.source.effective_plan,
899 &input.source.execution_bundle,
900 input.request,
901 input.data_envelopes,
902 )?;
903 let phase_request = ReplayPhaseRequest {
904 bundle_id: replay_bundle.bundle_id.clone(),
905 phase: input.request.phase,
906 data_envelope_keys: input.request.data_envelope_keys.clone(),
907 };
908 let artifact_store = BundlePayloadArtifactStore {
909 bundle: &replay_bundle,
910 controllers: input.controllers,
911 fallback: input.artifact_store,
912 hydrated_handles: Mutex::new(Vec::new()),
913 };
914 let mut ctx = RunContext::new(input.run_id.clone(), None);
915 let execution = SequentialScheduler.execute_bundle_replay(
916 BundleReplayExecution {
917 plan: &replay_plan,
918 bundle: &replay_bundle,
919 replay_request: &phase_request,
920 prediction_cache_store: None,
921 controllers: input.controllers,
922 data_provider: input.data_provider,
923 artifact_store: &artifact_store,
924 data_envelopes: input.data_envelopes,
925 },
926 &mut ctx,
927 );
928 let results = finish_bundle_payload_replay(execution, &artifact_store)?;
929 if results
930 .iter()
931 .any(|result| !result.artifacts.is_empty() || !result.artifact_handles.is_empty())
932 {
933 return contract_error("attached training replay PREDICT/EXPLAIN cannot emit artifacts");
934 }
935
936 let outputs = bind_attached_replay_outputs(input.source, input.request, &results)?;
937 let conformal_intervals = input
938 .source
939 .conformal_calibration
940 .as_ref()
941 .map(|calibration| apply_replay_conformal_intervals(calibration, &outputs))
942 .transpose()?
943 .unwrap_or_default();
944 if let Some(calibration) = input.source.conformal_calibration.as_ref() {
945 validate_replay_interval_closure(calibration, &outputs, &conformal_intervals)?;
946 }
947 let explanations = bind_attached_replay_explanations(input.request, &results)?;
948 let mut lineage = ctx.lineage.records().cloned().collect::<Vec<_>>();
949 for record in &mut lineage {
950 record.input_lineage.sort();
951 record
952 .artifact_refs
953 .sort_by(|left, right| left.id.cmp(&right.id));
954 }
955 lineage.sort_by(|left, right| left.record_id.cmp(&right.record_id));
956
957 let mut outcome = TrainingReplayOutcome {
958 schema_version: replay_outcome_schema_version(&input_data_identities),
959 outcome_id: input.outcome_id,
960 run_id: input.run_id,
961 source_training_outcome: input.source.to_reference()?,
962 replay_request_id: input.request.request_id.clone(),
963 replay_request_fingerprint: input.request.request_fingerprint.clone(),
964 input_data_identities,
965 bundle_id: input.source.execution_bundle.bundle_id.clone(),
966 plan_id: input.source.effective_plan.id.clone(),
967 phase: input.request.phase,
968 result_count: lineage.len(),
969 lineage_record_count: lineage.len(),
970 prediction_block_count: outputs.iter().map(|output| output.predictions.len()).sum(),
971 observation_prediction_block_count: outputs
972 .iter()
973 .map(|output| output.observation_predictions.len())
974 .sum(),
975 aggregated_prediction_block_count: outputs
976 .iter()
977 .map(|output| output.aggregated_predictions.len())
978 .sum(),
979 explanation_block_count: explanations.len(),
980 controller_count: lineage
981 .iter()
982 .map(|record| record.controller_id.as_str())
983 .collect::<BTreeSet<_>>()
984 .len(),
985 prediction_cache_store: false,
986 outputs,
987 conformal_intervals,
988 explanations,
989 lineage,
990 warnings: input.warnings,
991 diagnostics: input.diagnostics,
992 outcome_fingerprint: zero_fingerprint(),
993 };
994 outcome.outcome_fingerprint = outcome.compute_fingerprint()?;
995 outcome.validate_against(input.source, input.request)?;
996 Ok(outcome)
997}
998
999pub fn execute_loaded_predictor_replay(
1000 input: LoadedPredictorReplayInput<'_>,
1001) -> Result<TrainingReplayOutcome> {
1002 let package = input.predictor.package();
1003 package.validate()?;
1004 input.request.validate()?;
1005 validate_sorted_unique_text("training replay execution warnings", &input.warnings, false)?;
1006 validate_diagnostics(&input.diagnostics)?;
1007 validate_replay_phase(input.request.phase)?;
1008 if input.request.source_outcome_fingerprint != package.training_outcome.outcome_fingerprint {
1009 return contract_error("training replay request does not target package source outcome");
1010 }
1011 for node_plan in package.effective_plan.node_plans.values() {
1012 if input.controllers.get(&node_plan.controller_id).is_none() {
1013 return Err(DagMlError::RuntimeValidation(format!(
1014 "loaded predictor replay controller `{}` for node `{}` is not registered",
1015 node_plan.controller_id, node_plan.node_id
1016 )));
1017 }
1018 }
1019
1020 let input_data_identities = replay_input_data_identities(
1021 &package.execution_bundle,
1022 input.request,
1023 input.data_envelopes,
1024 )?;
1025 let (replay_plan, replay_bundle) = replay_plan_and_bundle_for_current_cohort(
1026 &package.effective_plan,
1027 &package.execution_bundle,
1028 input.request,
1029 input.data_envelopes,
1030 )?;
1031 let phase_request = ReplayPhaseRequest {
1032 bundle_id: replay_bundle.bundle_id.clone(),
1033 phase: input.request.phase,
1034 data_envelope_keys: input.request.data_envelope_keys.clone(),
1035 };
1036 let loaded_artifact_store = LoadedPredictorArtifactStore::new(input.predictor)?;
1037 let artifact_store = BundlePayloadArtifactStore {
1038 bundle: &replay_bundle,
1039 controllers: input.controllers,
1040 fallback: &loaded_artifact_store,
1041 hydrated_handles: Mutex::new(Vec::new()),
1042 };
1043 let mut ctx = RunContext::new(input.run_id.clone(), None);
1044 let execution = SequentialScheduler.execute_bundle_replay(
1045 BundleReplayExecution {
1046 plan: &replay_plan,
1047 bundle: &replay_bundle,
1048 replay_request: &phase_request,
1049 prediction_cache_store: None,
1050 controllers: input.controllers,
1051 data_provider: input.data_provider,
1052 artifact_store: &artifact_store,
1053 data_envelopes: input.data_envelopes,
1054 },
1055 &mut ctx,
1056 );
1057 let results = finish_bundle_payload_replay(execution, &artifact_store)?;
1058 if results
1059 .iter()
1060 .any(|result| !result.artifacts.is_empty() || !result.artifact_handles.is_empty())
1061 {
1062 return contract_error("loaded predictor replay PREDICT/EXPLAIN cannot emit artifacts");
1063 }
1064
1065 let outputs = bind_package_replay_outputs(package, input.request, &results)?;
1066 let conformal_intervals = package
1067 .conformal_calibration
1068 .as_ref()
1069 .map(|calibration| apply_replay_conformal_intervals(calibration, &outputs))
1070 .transpose()?
1071 .unwrap_or_default();
1072 if let Some(calibration) = package.conformal_calibration.as_ref() {
1073 validate_replay_interval_closure(calibration, &outputs, &conformal_intervals)?;
1074 }
1075 let explanations = bind_attached_replay_explanations(input.request, &results)?;
1076 let mut lineage = ctx.lineage.records().cloned().collect::<Vec<_>>();
1077 for record in &mut lineage {
1078 record.input_lineage.sort();
1079 record
1080 .artifact_refs
1081 .sort_by(|left, right| left.id.cmp(&right.id));
1082 }
1083 lineage.sort_by(|left, right| left.record_id.cmp(&right.record_id));
1084
1085 let mut outcome = TrainingReplayOutcome {
1086 schema_version: replay_outcome_schema_version(&input_data_identities),
1087 outcome_id: input.outcome_id,
1088 run_id: input.run_id,
1089 source_training_outcome: package.training_outcome.clone(),
1090 replay_request_id: input.request.request_id.clone(),
1091 replay_request_fingerprint: input.request.request_fingerprint.clone(),
1092 input_data_identities,
1093 bundle_id: package.execution_bundle.bundle_id.clone(),
1094 plan_id: package.effective_plan.id.clone(),
1095 phase: input.request.phase,
1096 result_count: lineage.len(),
1097 lineage_record_count: lineage.len(),
1098 prediction_block_count: outputs.iter().map(|output| output.predictions.len()).sum(),
1099 observation_prediction_block_count: outputs
1100 .iter()
1101 .map(|output| output.observation_predictions.len())
1102 .sum(),
1103 aggregated_prediction_block_count: outputs
1104 .iter()
1105 .map(|output| output.aggregated_predictions.len())
1106 .sum(),
1107 explanation_block_count: explanations.len(),
1108 controller_count: lineage
1109 .iter()
1110 .map(|record| record.controller_id.as_str())
1111 .collect::<BTreeSet<_>>()
1112 .len(),
1113 prediction_cache_store: false,
1114 outputs,
1115 conformal_intervals,
1116 explanations,
1117 lineage,
1118 warnings: input.warnings,
1119 diagnostics: input.diagnostics,
1120 outcome_fingerprint: zero_fingerprint(),
1121 };
1122 outcome.outcome_fingerprint = outcome.compute_fingerprint()?;
1123 outcome.validate_against_package(package, input.request)?;
1124 Ok(outcome)
1125}
1126
1127#[allow(clippy::too_many_arguments)]
1132pub fn calibrate_attached_training_replay(
1133 source: &mut TrainingOutcome,
1134 replay: &TrainingReplayOutcome,
1135 binding_id: &str,
1136 calibration_relations: &SampleRelationSet,
1137 truth: ConformalCalibrationTruth,
1138 context: ConformalCalibrationContext,
1139 coverages: Vec<f64>,
1140 multi_target_policy: ConformalMultiTargetPolicy,
1141 small_sample_policy: ConformalSmallSamplePolicy,
1142) -> Result<ConformalCalibration> {
1143 if replay.phase != Phase::Predict {
1144 return Err(DagMlError::RuntimeValidation(
1145 "conformal calibration requires a PREDICT replay outcome".to_string(),
1146 ));
1147 }
1148 if replay
1149 .input_data_identities
1150 .iter()
1151 .any(|identity| identity.target_content_fingerprint.is_none())
1152 {
1153 return Err(DagMlError::RuntimeValidation(
1154 "conformal calibration requires target-bound replay input identities; run the calibration replay with its authoritative truth cohort"
1155 .to_string(),
1156 ));
1157 }
1158 replay.validate_against(source, &replay_request_from_outcome(replay))?;
1159 let output = replay
1160 .outputs
1161 .iter()
1162 .find(|output| output.binding.binding_id == binding_id)
1163 .ok_or_else(|| {
1164 DagMlError::RuntimeValidation(
1165 "calibration replay has no requested output binding".to_string(),
1166 )
1167 })?;
1168 let [point] = output.predictions.as_slice() else {
1169 return Err(DagMlError::RuntimeValidation(
1170 "calibration replay requires exactly one sample point-prediction block".to_string(),
1171 ));
1172 };
1173 validate_calibration_context(
1174 source,
1175 replay,
1176 output,
1177 calibration_relations,
1178 &truth,
1179 &context,
1180 )?;
1181 let calibration = ConformalCalibration::calibrate_with_truth(
1182 binding_id,
1183 output.binding.target_names.clone(),
1184 point,
1185 &truth,
1186 context,
1187 coverages,
1188 multi_target_policy,
1189 small_sample_policy,
1190 )?;
1191 source.attach_conformal_calibration(calibration.clone(), replay.clone())?;
1192 Ok(calibration)
1193}
1194
1195fn validate_calibration_context(
1196 source: &TrainingOutcome,
1197 replay: &TrainingReplayOutcome,
1198 output: &BoundTrainingOutput,
1199 calibration_relations: &SampleRelationSet,
1200 truth: &ConformalCalibrationTruth,
1201 context: &ConformalCalibrationContext,
1202) -> Result<()> {
1203 context.validate_for_truth(truth, &output.binding.target_names)?;
1204 calibration_relations.validate()?;
1205 let relation_fingerprint = calibration_relations.fingerprint()?;
1206 if context.relation_fingerprint != relation_fingerprint
1207 || replay
1208 .input_data_identities
1209 .iter()
1210 .any(|identity| identity.relation_fingerprint != relation_fingerprint)
1211 {
1212 return Err(DagMlError::RuntimeValidation(
1213 "conformal calibration relation authority does not match context/replay provenance"
1214 .to_string(),
1215 ));
1216 }
1217 let origin_sample_ids = calibration_origin_closure(
1218 &context.calibration_cohort.physical_sample_ids,
1219 calibration_relations,
1220 )?;
1221 if context.calibration_cohort.origin_sample_ids != origin_sample_ids {
1222 return Err(DagMlError::RuntimeValidation(
1223 "conformal calibration cohort origin closure does not match relation authority"
1224 .to_string(),
1225 ));
1226 }
1227 let fold_set = source.effective_plan.fold_set.as_ref().ok_or_else(|| {
1228 DagMlError::RuntimeValidation("conformal calibration requires a source FoldSet".to_string())
1229 })?;
1230 let expected_fold = fold_set_fingerprint(fold_set)?;
1231 if context.predictor_binding_fingerprint != output.binding.binding_fingerprint
1232 || context.source_training_outcome_fingerprint != source.outcome_fingerprint
1233 || context.calibration_replay_outcome_fingerprint != replay.outcome_fingerprint
1234 || context.data_identities_fingerprint != source.data_identities_fingerprint()?
1235 || context.fold_set_fingerprint != expected_fold
1236 || context.training_influence_fingerprint != source.training_influence.manifest_fingerprint
1237 {
1238 return Err(DagMlError::RuntimeValidation(
1239 "conformal calibration context does not exactly match source/replay provenance"
1240 .to_string(),
1241 ));
1242 }
1243 let cohort = &context.calibration_cohort;
1244 let training: std::collections::BTreeSet<_> = source
1245 .training_influence
1246 .entries
1247 .iter()
1248 .flat_map(|entry| {
1249 entry
1250 .physical_sample_ids
1251 .iter()
1252 .chain(entry.origin_sample_ids.iter())
1253 })
1254 .collect();
1255 if cohort
1256 .physical_sample_ids
1257 .iter()
1258 .chain(cohort.origin_sample_ids.iter())
1259 .any(|id| training.contains(id))
1260 {
1261 return Err(DagMlError::RuntimeValidation(
1262 "conformal calibration cohort overlaps training influence closure".to_string(),
1263 ));
1264 }
1265 if relation_fingerprint == source.training_influence.relation_fingerprint {
1266 return Err(DagMlError::RuntimeValidation(
1267 "conformal calibration relation authority must be distinct from development relations"
1268 .to_string(),
1269 ));
1270 }
1271 Ok(())
1272}
1273
1274fn calibration_origin_closure(
1275 physical_sample_ids: &[crate::ids::SampleId],
1276 relations: &SampleRelationSet,
1277) -> Result<Vec<crate::ids::SampleId>> {
1278 let requested = physical_sample_ids.iter().collect::<BTreeSet<_>>();
1279 let mut by_sample = BTreeMap::new();
1280 for relation in &relations.records {
1281 if !requested.contains(&relation.sample_id) {
1282 continue;
1283 }
1284 match by_sample.get(&relation.sample_id) {
1285 Some(origin) if origin != &relation.origin_sample_id => {
1286 return Err(DagMlError::RuntimeValidation(format!(
1287 "conformal calibration sample `{}` has ambiguous origin relations",
1288 relation.sample_id
1289 )));
1290 }
1291 Some(_) => {}
1292 None => {
1293 by_sample.insert(
1294 relation.sample_id.clone(),
1295 relation.origin_sample_id.clone(),
1296 );
1297 }
1298 }
1299 }
1300 if let Some(missing) = physical_sample_ids
1301 .iter()
1302 .find(|sample_id| !by_sample.contains_key(*sample_id))
1303 {
1304 return Err(DagMlError::RuntimeValidation(format!(
1305 "conformal calibration sample `{missing}` is absent from relation authority"
1306 )));
1307 }
1308 Ok(by_sample
1309 .into_values()
1310 .flatten()
1311 .collect::<BTreeSet<_>>()
1312 .into_iter()
1313 .collect())
1314}
1315
1316fn replay_input_data_identities(
1317 bundle: &ExecutionBundle,
1318 request: &TrainingReplayRequest,
1319 envelopes: &BTreeMap<String, ExternalDataPlanEnvelope>,
1320) -> Result<Vec<ReplayDataIdentity>> {
1321 request
1322 .data_envelope_keys
1323 .iter()
1324 .map(|key| {
1325 let requirement = bundle
1326 .data_requirements
1327 .iter()
1328 .find(|requirement| requirement.key() == *key)
1329 .ok_or_else(|| {
1330 DagMlError::RuntimeValidation(format!(
1331 "training replay request references unknown data envelope key `{key}`"
1332 ))
1333 })?;
1334 let envelope = envelopes.get(key).ok_or_else(|| {
1335 DagMlError::RuntimeValidation(format!(
1336 "training replay is missing external data envelope for `{key}`"
1337 ))
1338 })?;
1339 envelope.validate()?;
1340 if requirement.schema_fingerprint != envelope.schema_fingerprint
1341 || requirement.plan_fingerprint != envelope.plan_fingerprint
1342 {
1343 return Err(DagMlError::RuntimeValidation(format!(
1344 "training replay envelope for `{key}` changes schema or representation plan"
1345 )));
1346 }
1347 let relation_fingerprint = envelope.relation_fingerprint.clone().ok_or_else(|| {
1348 DagMlError::RuntimeValidation(format!(
1349 "training replay envelope for `{key}` requires a relation fingerprint"
1350 ))
1351 })?;
1352 let data_content_fingerprint =
1353 envelope.data_content_fingerprint.clone().ok_or_else(|| {
1354 DagMlError::RuntimeValidation(format!(
1355 "training replay envelope for `{key}` requires a data content fingerprint"
1356 ))
1357 })?;
1358 let mut identity = ReplayDataIdentity {
1359 requirement_key: key.clone(),
1360 schema_fingerprint: envelope.schema_fingerprint.clone(),
1361 plan_fingerprint: envelope.plan_fingerprint.clone(),
1362 relation_fingerprint,
1363 data_content_fingerprint,
1364 target_content_fingerprint: envelope.target_content_fingerprint.clone(),
1365 identity_fingerprint: zero_fingerprint(),
1366 };
1367 identity.identity_fingerprint = identity.compute_fingerprint()?;
1368 identity.validate()?;
1369 Ok(identity)
1370 })
1371 .collect()
1372}
1373
1374fn replay_outcome_schema_version(_identities: &[ReplayDataIdentity]) -> u32 {
1375 TRAINING_REPLAY_OUTCOME_SCHEMA_VERSION
1376}
1377
1378fn replay_plan_and_bundle_for_current_cohort(
1379 plan: &ExecutionPlan,
1380 bundle: &ExecutionBundle,
1381 request: &TrainingReplayRequest,
1382 envelopes: &BTreeMap<String, ExternalDataPlanEnvelope>,
1383) -> Result<(ExecutionPlan, ExecutionBundle)> {
1384 let mut replay_plan = plan.clone();
1385 let mut replay_bundle = bundle.clone();
1386 replay_bundle.methods_hpo_resume_state = None;
1392 for requirement in &mut replay_bundle.data_requirements {
1393 let key = requirement.key();
1394 if request.data_envelope_keys.contains(&key) {
1395 let envelope = envelopes.get(&key).ok_or_else(|| {
1396 DagMlError::RuntimeValidation(format!(
1397 "training replay is missing external data envelope for `{key}`"
1398 ))
1399 })?;
1400 requirement.relation_fingerprint = envelope.relation_fingerprint.clone();
1401 for bindings in replay_plan.campaign.data_bindings.values_mut() {
1402 for binding in bindings {
1403 if crate::data::data_binding_requirement_key(
1404 &binding.node_id,
1405 &binding.input_name,
1406 ) == key
1407 {
1408 binding.relation_fingerprint = envelope.relation_fingerprint.clone();
1409 }
1410 }
1411 }
1412 for node_plan in replay_plan.node_plans.values_mut() {
1413 for binding in &mut node_plan.data_bindings {
1414 if crate::data::data_binding_requirement_key(
1415 &binding.node_id,
1416 &binding.input_name,
1417 ) == key
1418 {
1419 binding.relation_fingerprint = envelope.relation_fingerprint.clone();
1420 }
1421 }
1422 }
1423 }
1424 }
1425 replay_plan.graph_fingerprint = stable_json_fingerprint(&replay_plan.graph_plan.graph)?;
1426 replay_plan.campaign_fingerprint = stable_json_fingerprint(&replay_plan.campaign)?;
1427 replay_plan.controller_fingerprint =
1428 stable_json_fingerprint(&replay_plan.controller_manifests)?;
1429 replay_plan.validate()?;
1430 replay_bundle.graph_fingerprint = replay_plan.graph_fingerprint.clone();
1431 replay_bundle.campaign_fingerprint = replay_plan.campaign_fingerprint.clone();
1432 replay_bundle.controller_fingerprint = replay_plan.controller_fingerprint.clone();
1433 replay_bundle.validate_against_plan(&replay_plan)?;
1434 Ok((replay_plan, replay_bundle))
1435}
1436
1437fn bind_attached_replay_outputs(
1438 source: &TrainingOutcome,
1439 request: &TrainingReplayRequest,
1440 results: &[crate::runtime::NodeResult],
1441) -> Result<Vec<BoundTrainingOutput>> {
1442 let mut outputs = Vec::new();
1443 for binding_id in &request.output_binding_ids {
1444 let source_output = source
1445 .outputs
1446 .iter()
1447 .find(|output| output.binding.binding_id == *binding_id)
1448 .ok_or_else(|| {
1449 DagMlError::RuntimeValidation(format!(
1450 "training replay request references absent binding `{binding_id}`"
1451 ))
1452 })?;
1453 let binding = source_output.binding.clone();
1454 let mut output = BoundTrainingOutput {
1455 schema_version: Some(BOUND_TRAINING_OUTPUT_SCHEMA_VERSION),
1456 binding: binding.clone(),
1457 predictions: Vec::new(),
1458 observation_predictions: Vec::new(),
1459 aggregated_predictions: Vec::new(),
1460 };
1461 for result in results {
1462 output.predictions.extend(
1463 result
1464 .predictions
1465 .iter()
1466 .filter(|block| {
1467 block.producer_node == binding.node_id
1468 && block.producer_port.as_deref() == Some(binding.port_name.as_str())
1469 && block.partition == crate::oof::PredictionPartition::Final
1470 && block.fold_id.is_none()
1471 })
1472 .cloned(),
1473 );
1474 output.observation_predictions.extend(
1475 result
1476 .observation_predictions
1477 .iter()
1478 .filter(|block| {
1479 block.producer_node == binding.node_id
1480 && block.producer_port.as_deref() == Some(binding.port_name.as_str())
1481 && block.partition == crate::oof::PredictionPartition::Final
1482 && block.fold_id.is_none()
1483 })
1484 .cloned(),
1485 );
1486 output.aggregated_predictions.extend(
1487 result
1488 .aggregated_predictions
1489 .iter()
1490 .filter(|block| {
1491 block.producer_node == binding.node_id
1492 && block.producer_port.as_deref() == Some(binding.port_name.as_str())
1493 && block.partition == crate::oof::PredictionPartition::Final
1494 && block.fold_id.is_none()
1495 })
1496 .cloned(),
1497 );
1498 }
1499 if !output.predictions.is_empty()
1500 || !output.observation_predictions.is_empty()
1501 || !output.aggregated_predictions.is_empty()
1502 {
1503 output.validate(&source.effective_plan)?;
1504 outputs.push(output);
1505 }
1506 }
1507 outputs.sort_by(|left, right| left.binding.binding_id.cmp(&right.binding.binding_id));
1508 Ok(outputs)
1509}
1510
1511fn bind_package_replay_outputs(
1512 package: &PortablePredictorPackage,
1513 request: &TrainingReplayRequest,
1514 results: &[crate::runtime::NodeResult],
1515) -> Result<Vec<BoundTrainingOutput>> {
1516 let mut outputs = Vec::new();
1517 for binding_id in &request.output_binding_ids {
1518 let binding = package
1519 .output_bindings
1520 .iter()
1521 .find(|binding| binding.binding_id == *binding_id)
1522 .ok_or_else(|| {
1523 DagMlError::RuntimeValidation(format!(
1524 "training replay request references absent package binding `{binding_id}`"
1525 ))
1526 })?
1527 .clone();
1528 let mut output = BoundTrainingOutput {
1529 schema_version: Some(BOUND_TRAINING_OUTPUT_SCHEMA_VERSION),
1530 binding: binding.clone(),
1531 predictions: Vec::new(),
1532 observation_predictions: Vec::new(),
1533 aggregated_predictions: Vec::new(),
1534 };
1535 for result in results {
1536 output.predictions.extend(
1537 result
1538 .predictions
1539 .iter()
1540 .filter(|block| {
1541 block.producer_node == binding.node_id
1542 && block.producer_port.as_deref() == Some(binding.port_name.as_str())
1543 && block.partition == crate::oof::PredictionPartition::Final
1544 && block.fold_id.is_none()
1545 })
1546 .cloned(),
1547 );
1548 output.observation_predictions.extend(
1549 result
1550 .observation_predictions
1551 .iter()
1552 .filter(|block| {
1553 block.producer_node == binding.node_id
1554 && block.producer_port.as_deref() == Some(binding.port_name.as_str())
1555 && block.partition == crate::oof::PredictionPartition::Final
1556 && block.fold_id.is_none()
1557 })
1558 .cloned(),
1559 );
1560 output.aggregated_predictions.extend(
1561 result
1562 .aggregated_predictions
1563 .iter()
1564 .filter(|block| {
1565 block.producer_node == binding.node_id
1566 && block.producer_port.as_deref() == Some(binding.port_name.as_str())
1567 && block.partition == crate::oof::PredictionPartition::Final
1568 && block.fold_id.is_none()
1569 })
1570 .cloned(),
1571 );
1572 }
1573 if !output.predictions.is_empty()
1574 || !output.observation_predictions.is_empty()
1575 || !output.aggregated_predictions.is_empty()
1576 {
1577 output.validate(&package.effective_plan)?;
1578 outputs.push(output);
1579 }
1580 }
1581 outputs.sort_by(|left, right| left.binding.binding_id.cmp(&right.binding.binding_id));
1582 Ok(outputs)
1583}
1584
1585fn apply_replay_conformal_intervals(
1586 calibration: &ConformalCalibration,
1587 outputs: &[BoundTrainingOutput],
1588) -> Result<Vec<ConformalIntervalBlock>> {
1589 let Some(output) = outputs
1590 .iter()
1591 .find(|output| output.binding.binding_id == calibration.binding_id)
1592 else {
1593 return Ok(Vec::new());
1596 };
1597 if output.binding.target_names != calibration.target_names {
1598 return contract_error("replay conformal binding target order differs from calibration");
1599 }
1600 output
1601 .predictions
1602 .iter()
1603 .map(|block| calibration.apply(block))
1604 .collect()
1605}
1606
1607fn validate_replay_interval_closure(
1608 calibration: &ConformalCalibration,
1609 outputs: &[BoundTrainingOutput],
1610 intervals: &[ConformalIntervalBlock],
1611) -> Result<()> {
1612 let expected = apply_replay_conformal_intervals(calibration, outputs)?;
1613 if intervals != expected {
1614 return Err(DagMlError::RuntimeValidation(
1615 "conformal intervals do not exactly cover replay point predictions".to_string(),
1616 ));
1617 }
1618 for interval in intervals {
1619 let output = outputs
1620 .iter()
1621 .find(|output| output.binding.binding_id == interval.binding_id)
1622 .ok_or_else(|| {
1623 DagMlError::RuntimeValidation(
1624 "conformal interval references an absent replay output binding".to_string(),
1625 )
1626 })?;
1627 let point = output
1628 .predictions
1629 .iter()
1630 .find(|point| point.sample_ids == interval.sample_ids)
1631 .ok_or_else(|| {
1632 DagMlError::RuntimeValidation(
1633 "conformal interval has no matching replay point block".to_string(),
1634 )
1635 })?;
1636 interval.validate_against(calibration, point)?;
1637 }
1638 Ok(())
1639}
1640
1641fn bind_attached_replay_explanations(
1642 request: &TrainingReplayRequest,
1643 results: &[crate::runtime::NodeResult],
1644) -> Result<Vec<ExplanationBlock>> {
1645 if request.phase != Phase::Explain {
1646 return Ok(Vec::new());
1647 }
1648 let mut explanations = results
1649 .iter()
1650 .flat_map(|result| result.explanations.iter().cloned())
1651 .filter(|block| block.producer_port.is_some())
1652 .collect::<Vec<_>>();
1653 explanations.sort_by(|left, right| {
1654 (
1655 left.producer_node.as_str(),
1656 left.producer_port.as_deref().unwrap_or_default(),
1657 left.method.as_str(),
1658 left.target_name.as_deref().unwrap_or_default(),
1659 )
1660 .cmp(&(
1661 right.producer_node.as_str(),
1662 right.producer_port.as_deref().unwrap_or_default(),
1663 right.method.as_str(),
1664 right.target_name.as_deref().unwrap_or_default(),
1665 ))
1666 });
1667 Ok(explanations)
1668}
1669
1670fn validate_output_order_and_version(outputs: &[BoundTrainingOutput]) -> Result<()> {
1671 let mut previous: Option<&str> = None;
1672 for output in outputs {
1673 match output.schema_version {
1674 Some(BOUND_TRAINING_OUTPUT_SCHEMA_VERSION) => {}
1675 Some(version) => {
1676 return contract_error(format!(
1677 "training replay output schema_version {version} is unsupported; current {BOUND_TRAINING_OUTPUT_SCHEMA_VERSION}"
1678 ));
1679 }
1680 None => {
1681 return contract_error(
1682 "training replay output requires bound_training_output schema_version",
1683 );
1684 }
1685 }
1686 let binding_id = output.binding.binding_id.as_str();
1687 if previous.is_some_and(|previous| previous >= binding_id) {
1688 return contract_error("training replay outputs must be strictly sorted by binding_id");
1689 }
1690 previous = Some(binding_id);
1691 }
1692 Ok(())
1693}
1694
1695fn validate_replay_bound_output_blocks(output: &BoundTrainingOutput) -> Result<()> {
1696 for block in &output.predictions {
1697 validate_optional_port(
1698 "training replay prediction producer_port",
1699 &block.producer_port,
1700 )?;
1701 if block.partition != crate::oof::PredictionPartition::Final || block.fold_id.is_some() {
1702 return contract_error(
1703 "training replay prediction blocks must use final partition without fold",
1704 );
1705 }
1706 }
1707 for block in &output.observation_predictions {
1708 validate_optional_port(
1709 "training replay observation prediction producer_port",
1710 &block.producer_port,
1711 )?;
1712 if block.partition != crate::oof::PredictionPartition::Final || block.fold_id.is_some() {
1713 return contract_error(
1714 "training replay observation prediction blocks must use final partition without fold",
1715 );
1716 }
1717 }
1718 for block in &output.aggregated_predictions {
1719 validate_optional_port(
1720 "training replay aggregated prediction producer_port",
1721 &block.producer_port,
1722 )?;
1723 if block.partition != crate::oof::PredictionPartition::Final || block.fold_id.is_some() {
1724 return contract_error(
1725 "training replay aggregated prediction blocks must use final partition without fold",
1726 );
1727 }
1728 }
1729 Ok(())
1730}
1731
1732fn validate_replay_phase(phase: Phase) -> Result<()> {
1733 if matches!(phase, Phase::Predict | Phase::Explain) {
1734 Ok(())
1735 } else {
1736 contract_error("training replay supports only PREDICT and EXPLAIN")
1737 }
1738}
1739
1740fn validate_sha256(label: &str, value: &str) -> Result<()> {
1741 if value.len() == 64
1742 && value
1743 .bytes()
1744 .all(|byte| byte.is_ascii_digit() || (b'a'..=b'f').contains(&byte))
1745 {
1746 Ok(())
1747 } else {
1748 contract_error(format!(
1749 "{label} fingerprint must be 64 lowercase hexadecimal characters"
1750 ))
1751 }
1752}
1753
1754fn validate_identifier(label: &str, value: &str) -> Result<()> {
1755 if !value.is_empty()
1756 && value.len() <= 128
1757 && value
1758 .bytes()
1759 .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'_' | b'-' | b'.' | b':'))
1760 {
1761 Ok(())
1762 } else {
1763 contract_error(format!("{label} is not a valid DAG-ML identifier"))
1764 }
1765}
1766
1767fn validate_non_empty(label: &str, value: &str) -> Result<()> {
1768 if value.trim().is_empty() {
1769 contract_error(format!("{label} must be non-empty"))
1770 } else {
1771 Ok(())
1772 }
1773}
1774
1775fn validate_sorted_unique_identifiers(
1776 label: &str,
1777 values: &[String],
1778 require_non_empty: bool,
1779) -> Result<()> {
1780 validate_sorted_unique_text(label, values, require_non_empty)?;
1781 for value in values {
1782 validate_identifier(label, value)?;
1783 }
1784 Ok(())
1785}
1786
1787fn validate_sorted_unique_text(
1788 label: &str,
1789 values: &[String],
1790 require_non_empty: bool,
1791) -> Result<()> {
1792 if require_non_empty && values.is_empty() {
1793 return contract_error(format!("{label} must be non-empty"));
1794 }
1795 let mut previous: Option<&str> = None;
1796 for value in values {
1797 validate_non_empty(label, value)?;
1798 if previous.is_some_and(|previous| previous >= value.as_str()) {
1799 return contract_error(format!("{label} must be strictly sorted and unique"));
1800 }
1801 previous = Some(value.as_str());
1802 }
1803 Ok(())
1804}
1805
1806fn validate_sorted_unique_keys<'a>(
1807 label: &str,
1808 values: impl Iterator<Item = &'a str>,
1809 require_non_empty: bool,
1810) -> Result<()> {
1811 let values = values.collect::<Vec<_>>();
1812 if require_non_empty && values.is_empty() {
1813 return contract_error(format!("{label} must be non-empty"));
1814 }
1815 let mut previous: Option<&str> = None;
1816 for value in values {
1817 validate_non_empty(label, value)?;
1818 if previous.is_some_and(|previous| previous >= value) {
1819 return contract_error(format!("{label} must be strictly sorted and unique"));
1820 }
1821 previous = Some(value);
1822 }
1823 Ok(())
1824}
1825
1826fn validate_optional_port(label: &str, value: &Option<String>) -> Result<()> {
1827 match value {
1828 Some(value) if !value.trim().is_empty() => Ok(()),
1829 _ => contract_error(format!("{label} must be present and non-empty")),
1830 }
1831}
1832
1833fn validate_diagnostics(diagnostics: &BTreeMap<String, serde_json::Value>) -> Result<()> {
1834 for (key, value) in diagnostics {
1835 validate_non_empty("training replay diagnostic key", key)?;
1836 if !matches!(
1837 value,
1838 serde_json::Value::Null
1839 | serde_json::Value::Bool(_)
1840 | serde_json::Value::Number(_)
1841 | serde_json::Value::String(_)
1842 ) {
1843 return contract_error("training replay diagnostics must be scalar JSON values");
1844 }
1845 }
1846 Ok(())
1847}
1848
1849fn require_count(label: &str, actual: usize, expected: usize) -> Result<()> {
1850 if actual == expected {
1851 Ok(())
1852 } else {
1853 contract_error(format!("{label} does not match replay payload"))
1854 }
1855}
1856
1857fn zero_fingerprint() -> String {
1858 "0".repeat(64)
1859}
1860
1861fn tcv1_fingerprint_without<T: Serialize>(value: &T, field: &str, label: &str) -> Result<String> {
1862 let json = serde_json::to_string(value)?;
1863 strict_tcv1_fingerprint_without(&json, field, label)
1864}
1865
1866fn strict_tcv1_fingerprint_without(json: &str, field: &str, label: &str) -> Result<String> {
1867 parse_typed_json(json)
1868 .and_then(|value| value.fingerprint_without(field))
1869 .map_err(|error| {
1870 DagMlError::RuntimeValidation(format!("{label} is outside strict TCV1: {error}"))
1871 })
1872}
1873
1874fn unsupported_version<T>(label: &str, actual: u32, expected: u32) -> Result<T> {
1875 contract_error(format!(
1876 "{label} uses unsupported schema_version {actual}, expected {expected}"
1877 ))
1878}
1879
1880fn contract_error<T>(message: impl Into<String>) -> Result<T> {
1881 Err(DagMlError::CampaignValidation(message.into()))
1882}
1883
1884#[cfg(test)]
1885mod methods_hpo_resume_state_tests {
1886 use super::*;
1887
1888 #[test]
1889 fn methods_hpo_resume_state_reader_refuses_unknown_legacy_and_noncanonical_json() {
1890 let unknown = r#"{"schema_version":1,"unknown_resume_side_channel":true}"#;
1891 assert!(methods_hpo_resume_state_from_json(unknown).is_err());
1892
1893 let duplicate = r#"{"schema_version":1,"schema_version":1}"#;
1896 let error = methods_hpo_resume_state_from_json(duplicate)
1897 .unwrap_err()
1898 .to_string();
1899 assert!(error.contains("duplicate JSON object key"), "{error}");
1900
1901 let legacy = r#"{"schema_version":1,"tuner_node_id":"tuner:legacy"}"#;
1905 assert!(methods_hpo_resume_state_from_json(legacy).is_err());
1906 }
1907
1908 #[test]
1909 fn methods_hpo_resume_state_reader_refuses_tampered_schema_version() {
1910 let tampered = r#"{"schema_version":2,"operation_id":"campaign:hpo"}"#;
1911 assert!(methods_hpo_resume_state_from_json(tampered).is_err());
1912 }
1913}
1914
1915#[cfg(test)]
1916mod replay_identity_tests {
1917 use super::*;
1918
1919 fn replay_identity(target_content_fingerprint: Option<&str>) -> ReplayDataIdentity {
1920 let mut identity = ReplayDataIdentity {
1921 requirement_key: "model:base.X".to_string(),
1922 schema_fingerprint: "1".repeat(64),
1923 plan_fingerprint: "2".repeat(64),
1924 relation_fingerprint: "3".repeat(64),
1925 data_content_fingerprint: "4".repeat(64),
1926 target_content_fingerprint: target_content_fingerprint.map(str::to_string),
1927 identity_fingerprint: zero_fingerprint(),
1928 };
1929 identity.identity_fingerprint = identity.compute_fingerprint().unwrap();
1930 identity
1931 }
1932
1933 #[test]
1934 fn replay_identity_permits_an_unlabeled_predict_cohort_without_a_sentinel() {
1935 let identity = replay_identity(None);
1936 identity.validate().unwrap();
1937 assert!(identity.target_content_fingerprint.is_none());
1938 assert!(serde_json::to_value(&identity)
1939 .unwrap()
1940 .get("target_content_fingerprint")
1941 .unwrap()
1942 .is_null());
1943 }
1944
1945 #[test]
1946 fn replay_identity_rejects_a_resigned_target_fingerprint_tamper() {
1947 let mut identity = replay_identity(None);
1948 identity.target_content_fingerprint = Some("not-a-fingerprint".to_string());
1949 identity.identity_fingerprint = identity.compute_fingerprint().unwrap();
1950 let error = identity.validate().unwrap_err().to_string();
1951 assert!(error.contains("target content"), "{error}");
1952 }
1953
1954 #[test]
1955 fn replay_schema_version_is_v3_for_target_bound_and_target_free_cohorts() {
1956 assert_eq!(
1957 replay_outcome_schema_version(&[replay_identity(None)]),
1958 TRAINING_REPLAY_OUTCOME_SCHEMA_VERSION
1959 );
1960 assert_eq!(
1961 replay_outcome_schema_version(&[replay_identity(Some(&"5".repeat(64)))]),
1962 TRAINING_REPLAY_OUTCOME_SCHEMA_VERSION
1963 );
1964 }
1965}
1966
1967#[cfg(test)]
1968mod tests {
1969 #[cfg(dag_ml_workspace_contract_fixtures)]
1970 use std::fs;
1971 #[cfg(dag_ml_workspace_contract_fixtures)]
1972 use std::path::PathBuf;
1973
1974 #[cfg(dag_ml_workspace_contract_fixtures)]
1975 use super::*;
1976
1977 #[cfg(dag_ml_workspace_contract_fixtures)]
1978 fn root() -> PathBuf {
1979 PathBuf::from(env!("CARGO_MANIFEST_DIR"))
1980 .parent()
1981 .and_then(|path| path.parent())
1982 .expect("core crate is under crates/dag-ml-core")
1983 .to_path_buf()
1984 }
1985
1986 #[cfg(dag_ml_workspace_contract_fixtures)]
1987 fn fixture(name: &str) -> String {
1988 fs::read_to_string(
1989 root()
1990 .join("examples")
1991 .join("fixtures")
1992 .join("training")
1993 .join("replay")
1994 .join(name),
1995 )
1996 .expect(name)
1997 }
1998
1999 #[cfg(dag_ml_workspace_contract_fixtures)]
2000 fn training_fixture(name: &str) -> String {
2001 fs::read_to_string(
2002 root()
2003 .join("examples")
2004 .join("fixtures")
2005 .join("training")
2006 .join(name),
2007 )
2008 .expect(name)
2009 }
2010
2011 #[cfg(dag_ml_workspace_contract_fixtures)]
2012 #[test]
2013 fn training_replay_contract_fixtures_parse_and_cross_validate() {
2014 let predict_source =
2015 TrainingOutcome::from_json(&training_fixture("training_outcome_refit.v1.json"))
2016 .expect("predict source training outcome");
2017 let explain_source =
2018 TrainingOutcome::from_json(&fixture("training_replay_source_outcome_explain.v1.json"))
2019 .expect("explain source training outcome");
2020 let predict_request =
2021 TrainingReplayRequest::from_json(&fixture("training_replay_request_predict.v1.json"))
2022 .expect("predict request");
2023 let predict_outcome =
2024 TrainingReplayOutcome::from_json(&fixture("training_replay_outcome_predict.v1.json"))
2025 .expect("predict outcome");
2026 predict_outcome
2027 .validate_against(&predict_source, &predict_request)
2028 .expect("predict cross-links");
2029
2030 let explain_request =
2031 TrainingReplayRequest::from_json(&fixture("training_replay_request_explain.v1.json"))
2032 .expect("explain request");
2033 let explain_outcome =
2034 TrainingReplayOutcome::from_json(&fixture("training_replay_outcome_explain.v1.json"))
2035 .expect("explain outcome");
2036 explain_outcome
2037 .validate_against(&explain_source, &explain_request)
2038 .expect("explain cross-links");
2039
2040 let explain_only = TrainingReplayOutcome::from_json(&fixture(
2041 "training_replay_outcome_explain_only.v1.json",
2042 ))
2043 .expect("explain-only outcome");
2044 explain_only
2045 .validate_against(&explain_source, &explain_request)
2046 .expect("explain-only cross-links");
2047 }
2048
2049 #[cfg(dag_ml_workspace_contract_fixtures)]
2050 #[test]
2051 fn training_replay_request_rejects_refit_and_unsorted_bindings() {
2052 let mut request: serde_json::Value =
2053 serde_json::from_str(&fixture("training_replay_request_predict.v1.json")).unwrap();
2054 request["phase"] = serde_json::Value::String("REFIT".to_string());
2055 let err = serde_json::from_value::<TrainingReplayRequest>(request)
2056 .unwrap()
2057 .validate()
2058 .unwrap_err()
2059 .to_string();
2060 assert!(err.contains("PREDICT and EXPLAIN"));
2061
2062 let mut request: TrainingReplayRequest =
2063 TrainingReplayRequest::from_json(&fixture("training_replay_request_predict.v1.json"))
2064 .unwrap();
2065 request.output_binding_ids = vec!["z".to_string(), "a".to_string()];
2066 request.request_fingerprint = request.compute_fingerprint().unwrap();
2067 let err = request.validate().unwrap_err().to_string();
2068 assert!(err.contains("strictly sorted"));
2069 }
2070
2071 #[cfg(dag_ml_workspace_contract_fixtures)]
2072 #[test]
2073 fn training_replay_outcome_rejects_counter_and_source_transplants() {
2074 let source =
2075 TrainingOutcome::from_json(&training_fixture("training_outcome_refit.v1.json"))
2076 .unwrap();
2077 let request =
2078 TrainingReplayRequest::from_json(&fixture("training_replay_request_predict.v1.json"))
2079 .unwrap();
2080 let mut outcome =
2081 TrainingReplayOutcome::from_json(&fixture("training_replay_outcome_predict.v1.json"))
2082 .unwrap();
2083 outcome.prediction_block_count += 1;
2084 outcome.outcome_fingerprint = outcome.compute_fingerprint().unwrap();
2085 let err = outcome.validate().unwrap_err().to_string();
2086 assert!(err.contains("prediction_block_count"));
2087
2088 let mut outcome =
2089 TrainingReplayOutcome::from_json(&fixture("training_replay_outcome_predict.v1.json"))
2090 .unwrap();
2091 outcome.source_training_outcome.outcome_fingerprint = "f".repeat(64);
2092 outcome.outcome_fingerprint = outcome.compute_fingerprint().unwrap();
2093 let err = outcome
2094 .validate_against(&source, &request)
2095 .unwrap_err()
2096 .to_string();
2097 assert!(err.contains("source reference"));
2098 }
2099
2100 #[cfg(dag_ml_workspace_contract_fixtures)]
2101 #[test]
2102 fn target_free_predict_evidence_is_v3_only() {
2103 let mut outcome =
2104 TrainingReplayOutcome::from_json(&fixture("training_replay_outcome_predict.v1.json"))
2105 .unwrap();
2106 outcome.schema_version = TRAINING_REPLAY_OUTCOME_SCHEMA_VERSION;
2107 for identity in &mut outcome.input_data_identities {
2108 identity.target_content_fingerprint = None;
2109 identity.identity_fingerprint = identity.compute_fingerprint().unwrap();
2110 }
2111 outcome.outcome_fingerprint = outcome.compute_fingerprint().unwrap();
2112 outcome.validate().unwrap();
2113
2114 outcome.schema_version = CONFORMAL_TRAINING_REPLAY_OUTCOME_SCHEMA_VERSION;
2115 outcome.outcome_fingerprint = outcome.compute_fingerprint().unwrap();
2116 let error = outcome.validate().unwrap_err().to_string();
2117 assert!(error.contains("V1/V2 requires target-bound"), "{error}");
2118 }
2119}