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, 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
1195pub fn derive_attached_conformal_calibration_context(
1203 source: &TrainingOutcome,
1204 replay: &TrainingReplayOutcome,
1205 binding_id: &str,
1206 calibration_relations: &SampleRelationSet,
1207) -> Result<ConformalCalibrationContext> {
1208 if replay.phase != Phase::Predict {
1209 return Err(DagMlError::RuntimeValidation(
1210 "conformal calibration requires a PREDICT replay outcome".to_string(),
1211 ));
1212 }
1213 if replay
1214 .input_data_identities
1215 .iter()
1216 .any(|identity| identity.target_content_fingerprint.is_none())
1217 {
1218 return Err(DagMlError::RuntimeValidation(
1219 "conformal calibration requires target-bound replay input identities; run the calibration replay with its authoritative truth cohort"
1220 .to_string(),
1221 ));
1222 }
1223 replay.validate_against(source, &replay_request_from_outcome(replay))?;
1224 calibration_relations.validate()?;
1225 let output = replay
1226 .outputs
1227 .iter()
1228 .find(|output| output.binding.binding_id == binding_id)
1229 .ok_or_else(|| {
1230 DagMlError::RuntimeValidation(
1231 "calibration replay has no requested output binding".to_string(),
1232 )
1233 })?;
1234 let [point] = output.predictions.as_slice() else {
1235 return Err(DagMlError::RuntimeValidation(
1236 "calibration replay requires exactly one sample point-prediction block".to_string(),
1237 ));
1238 };
1239 let relation_fingerprint = calibration_relations.fingerprint()?;
1240 if replay
1241 .input_data_identities
1242 .iter()
1243 .any(|identity| identity.relation_fingerprint != relation_fingerprint)
1244 {
1245 return Err(DagMlError::RuntimeValidation(
1246 "conformal calibration relation authority does not match replay provenance".to_string(),
1247 ));
1248 }
1249 let origin_sample_ids = calibration_origin_closure(&point.sample_ids, calibration_relations)?;
1250 let mut calibration_cohort = ConformalCalibrationCohort {
1251 role: "calibration".to_string(),
1252 physical_sample_ids: point.sample_ids.clone(),
1253 origin_sample_ids,
1254 target_names: output.binding.target_names.clone(),
1255 manifest_fingerprint: String::new(),
1256 };
1257 calibration_cohort.manifest_fingerprint = calibration_cohort.compute_fingerprint()?;
1258 let fold_set = source.effective_plan.fold_set.as_ref().ok_or_else(|| {
1259 DagMlError::RuntimeValidation("conformal calibration requires a source FoldSet".to_string())
1260 })?;
1261 let mut context = ConformalCalibrationContext {
1262 predictor_binding_fingerprint: output.binding.binding_fingerprint.clone(),
1263 source_training_outcome_fingerprint: source.outcome_fingerprint.clone(),
1264 calibration_replay_outcome_fingerprint: replay.outcome_fingerprint.clone(),
1265 data_identities_fingerprint: source.data_identities_fingerprint()?,
1266 fold_set_fingerprint: fold_set_fingerprint(fold_set)?,
1267 training_influence_fingerprint: source.training_influence.manifest_fingerprint.clone(),
1268 relation_fingerprint,
1269 calibration_cohort,
1270 context_fingerprint: String::new(),
1271 };
1272 context.context_fingerprint = context.compute_fingerprint()?;
1273 Ok(context)
1274}
1275
1276#[allow(clippy::too_many_arguments)]
1279pub fn calibrate_attached_training_replay_with_derived_context(
1280 source: &mut TrainingOutcome,
1281 replay: &TrainingReplayOutcome,
1282 binding_id: &str,
1283 calibration_relations: &SampleRelationSet,
1284 truth: ConformalCalibrationTruth,
1285 coverages: Vec<f64>,
1286 multi_target_policy: ConformalMultiTargetPolicy,
1287 small_sample_policy: ConformalSmallSamplePolicy,
1288) -> Result<ConformalCalibration> {
1289 let context = derive_attached_conformal_calibration_context(
1290 source,
1291 replay,
1292 binding_id,
1293 calibration_relations,
1294 )?;
1295 calibrate_attached_training_replay(
1296 source,
1297 replay,
1298 binding_id,
1299 calibration_relations,
1300 truth,
1301 context,
1302 coverages,
1303 multi_target_policy,
1304 small_sample_policy,
1305 )
1306}
1307
1308fn validate_calibration_context(
1309 source: &TrainingOutcome,
1310 replay: &TrainingReplayOutcome,
1311 output: &BoundTrainingOutput,
1312 calibration_relations: &SampleRelationSet,
1313 truth: &ConformalCalibrationTruth,
1314 context: &ConformalCalibrationContext,
1315) -> Result<()> {
1316 context.validate_for_truth(truth, &output.binding.target_names)?;
1317 calibration_relations.validate()?;
1318 let relation_fingerprint = calibration_relations.fingerprint()?;
1319 if context.relation_fingerprint != relation_fingerprint
1320 || replay
1321 .input_data_identities
1322 .iter()
1323 .any(|identity| identity.relation_fingerprint != relation_fingerprint)
1324 {
1325 return Err(DagMlError::RuntimeValidation(
1326 "conformal calibration relation authority does not match context/replay provenance"
1327 .to_string(),
1328 ));
1329 }
1330 let origin_sample_ids = calibration_origin_closure(
1331 &context.calibration_cohort.physical_sample_ids,
1332 calibration_relations,
1333 )?;
1334 if context.calibration_cohort.origin_sample_ids != origin_sample_ids {
1335 return Err(DagMlError::RuntimeValidation(
1336 "conformal calibration cohort origin closure does not match relation authority"
1337 .to_string(),
1338 ));
1339 }
1340 let fold_set = source.effective_plan.fold_set.as_ref().ok_or_else(|| {
1341 DagMlError::RuntimeValidation("conformal calibration requires a source FoldSet".to_string())
1342 })?;
1343 let expected_fold = fold_set_fingerprint(fold_set)?;
1344 if context.predictor_binding_fingerprint != output.binding.binding_fingerprint
1345 || context.source_training_outcome_fingerprint != source.outcome_fingerprint
1346 || context.calibration_replay_outcome_fingerprint != replay.outcome_fingerprint
1347 || context.data_identities_fingerprint != source.data_identities_fingerprint()?
1348 || context.fold_set_fingerprint != expected_fold
1349 || context.training_influence_fingerprint != source.training_influence.manifest_fingerprint
1350 {
1351 return Err(DagMlError::RuntimeValidation(
1352 "conformal calibration context does not exactly match source/replay provenance"
1353 .to_string(),
1354 ));
1355 }
1356 let cohort = &context.calibration_cohort;
1357 let training: std::collections::BTreeSet<_> = source
1358 .training_influence
1359 .entries
1360 .iter()
1361 .flat_map(|entry| {
1362 entry
1363 .physical_sample_ids
1364 .iter()
1365 .chain(entry.origin_sample_ids.iter())
1366 })
1367 .collect();
1368 if cohort
1369 .physical_sample_ids
1370 .iter()
1371 .chain(cohort.origin_sample_ids.iter())
1372 .any(|id| training.contains(id))
1373 {
1374 return Err(DagMlError::RuntimeValidation(
1375 "conformal calibration cohort overlaps training influence closure".to_string(),
1376 ));
1377 }
1378 if relation_fingerprint == source.training_influence.relation_fingerprint {
1379 return Err(DagMlError::RuntimeValidation(
1380 "conformal calibration relation authority must be distinct from development relations"
1381 .to_string(),
1382 ));
1383 }
1384 Ok(())
1385}
1386
1387fn calibration_origin_closure(
1388 physical_sample_ids: &[crate::ids::SampleId],
1389 relations: &SampleRelationSet,
1390) -> Result<Vec<crate::ids::SampleId>> {
1391 let requested = physical_sample_ids.iter().collect::<BTreeSet<_>>();
1392 let mut by_sample = BTreeMap::new();
1393 for relation in &relations.records {
1394 if !requested.contains(&relation.sample_id) {
1395 continue;
1396 }
1397 match by_sample.get(&relation.sample_id) {
1398 Some(origin) if origin != &relation.origin_sample_id => {
1399 return Err(DagMlError::RuntimeValidation(format!(
1400 "conformal calibration sample `{}` has ambiguous origin relations",
1401 relation.sample_id
1402 )));
1403 }
1404 Some(_) => {}
1405 None => {
1406 by_sample.insert(
1407 relation.sample_id.clone(),
1408 relation.origin_sample_id.clone(),
1409 );
1410 }
1411 }
1412 }
1413 if let Some(missing) = physical_sample_ids
1414 .iter()
1415 .find(|sample_id| !by_sample.contains_key(*sample_id))
1416 {
1417 return Err(DagMlError::RuntimeValidation(format!(
1418 "conformal calibration sample `{missing}` is absent from relation authority"
1419 )));
1420 }
1421 Ok(by_sample
1422 .into_values()
1423 .flatten()
1424 .collect::<BTreeSet<_>>()
1425 .into_iter()
1426 .collect())
1427}
1428
1429fn replay_input_data_identities(
1430 bundle: &ExecutionBundle,
1431 request: &TrainingReplayRequest,
1432 envelopes: &BTreeMap<String, ExternalDataPlanEnvelope>,
1433) -> Result<Vec<ReplayDataIdentity>> {
1434 request
1435 .data_envelope_keys
1436 .iter()
1437 .map(|key| {
1438 let requirement = bundle
1439 .data_requirements
1440 .iter()
1441 .find(|requirement| requirement.key() == *key)
1442 .ok_or_else(|| {
1443 DagMlError::RuntimeValidation(format!(
1444 "training replay request references unknown data envelope key `{key}`"
1445 ))
1446 })?;
1447 let envelope = envelopes.get(key).ok_or_else(|| {
1448 DagMlError::RuntimeValidation(format!(
1449 "training replay is missing external data envelope for `{key}`"
1450 ))
1451 })?;
1452 envelope.validate()?;
1453 if requirement.schema_fingerprint != envelope.schema_fingerprint
1454 || requirement.plan_fingerprint != envelope.plan_fingerprint
1455 {
1456 return Err(DagMlError::RuntimeValidation(format!(
1457 "training replay envelope for `{key}` changes schema or representation plan"
1458 )));
1459 }
1460 let relation_fingerprint = envelope.relation_fingerprint.clone().ok_or_else(|| {
1461 DagMlError::RuntimeValidation(format!(
1462 "training replay envelope for `{key}` requires a relation fingerprint"
1463 ))
1464 })?;
1465 let data_content_fingerprint =
1466 envelope.data_content_fingerprint.clone().ok_or_else(|| {
1467 DagMlError::RuntimeValidation(format!(
1468 "training replay envelope for `{key}` requires a data content fingerprint"
1469 ))
1470 })?;
1471 let mut identity = ReplayDataIdentity {
1472 requirement_key: key.clone(),
1473 schema_fingerprint: envelope.schema_fingerprint.clone(),
1474 plan_fingerprint: envelope.plan_fingerprint.clone(),
1475 relation_fingerprint,
1476 data_content_fingerprint,
1477 target_content_fingerprint: envelope.target_content_fingerprint.clone(),
1478 identity_fingerprint: zero_fingerprint(),
1479 };
1480 identity.identity_fingerprint = identity.compute_fingerprint()?;
1481 identity.validate()?;
1482 Ok(identity)
1483 })
1484 .collect()
1485}
1486
1487fn replay_outcome_schema_version(_identities: &[ReplayDataIdentity]) -> u32 {
1488 TRAINING_REPLAY_OUTCOME_SCHEMA_VERSION
1489}
1490
1491fn replay_plan_and_bundle_for_current_cohort(
1492 plan: &ExecutionPlan,
1493 bundle: &ExecutionBundle,
1494 request: &TrainingReplayRequest,
1495 envelopes: &BTreeMap<String, ExternalDataPlanEnvelope>,
1496) -> Result<(ExecutionPlan, ExecutionBundle)> {
1497 let mut replay_plan = plan.clone();
1498 let mut replay_bundle = bundle.clone();
1499 replay_bundle.methods_hpo_resume_state = None;
1505 for requirement in &mut replay_bundle.data_requirements {
1506 let key = requirement.key();
1507 if request.data_envelope_keys.contains(&key) {
1508 let envelope = envelopes.get(&key).ok_or_else(|| {
1509 DagMlError::RuntimeValidation(format!(
1510 "training replay is missing external data envelope for `{key}`"
1511 ))
1512 })?;
1513 requirement.relation_fingerprint = envelope.relation_fingerprint.clone();
1514 for bindings in replay_plan.campaign.data_bindings.values_mut() {
1515 for binding in bindings {
1516 if crate::data::data_binding_requirement_key(
1517 &binding.node_id,
1518 &binding.input_name,
1519 ) == key
1520 {
1521 binding.relation_fingerprint = envelope.relation_fingerprint.clone();
1522 }
1523 }
1524 }
1525 for node_plan in replay_plan.node_plans.values_mut() {
1526 for binding in &mut node_plan.data_bindings {
1527 if crate::data::data_binding_requirement_key(
1528 &binding.node_id,
1529 &binding.input_name,
1530 ) == key
1531 {
1532 binding.relation_fingerprint = envelope.relation_fingerprint.clone();
1533 }
1534 }
1535 }
1536 }
1537 }
1538 replay_plan.graph_fingerprint = stable_json_fingerprint(&replay_plan.graph_plan.graph)?;
1539 replay_plan.campaign_fingerprint = stable_json_fingerprint(&replay_plan.campaign)?;
1540 replay_plan.controller_fingerprint =
1541 stable_json_fingerprint(&replay_plan.controller_manifests)?;
1542 replay_plan.validate()?;
1543 replay_bundle.graph_fingerprint = replay_plan.graph_fingerprint.clone();
1544 replay_bundle.campaign_fingerprint = replay_plan.campaign_fingerprint.clone();
1545 replay_bundle.controller_fingerprint = replay_plan.controller_fingerprint.clone();
1546 replay_bundle.validate_against_plan(&replay_plan)?;
1547 Ok((replay_plan, replay_bundle))
1548}
1549
1550fn bind_attached_replay_outputs(
1551 source: &TrainingOutcome,
1552 request: &TrainingReplayRequest,
1553 results: &[crate::runtime::NodeResult],
1554) -> Result<Vec<BoundTrainingOutput>> {
1555 let mut outputs = Vec::new();
1556 for binding_id in &request.output_binding_ids {
1557 let source_output = source
1558 .outputs
1559 .iter()
1560 .find(|output| output.binding.binding_id == *binding_id)
1561 .ok_or_else(|| {
1562 DagMlError::RuntimeValidation(format!(
1563 "training replay request references absent binding `{binding_id}`"
1564 ))
1565 })?;
1566 let binding = source_output.binding.clone();
1567 let mut output = BoundTrainingOutput {
1568 schema_version: Some(BOUND_TRAINING_OUTPUT_SCHEMA_VERSION),
1569 binding: binding.clone(),
1570 predictions: Vec::new(),
1571 observation_predictions: Vec::new(),
1572 aggregated_predictions: Vec::new(),
1573 };
1574 for result in results {
1575 output.predictions.extend(
1576 result
1577 .predictions
1578 .iter()
1579 .filter(|block| {
1580 block.producer_node == binding.node_id
1581 && block.producer_port.as_deref() == Some(binding.port_name.as_str())
1582 && block.partition == crate::oof::PredictionPartition::Final
1583 && block.fold_id.is_none()
1584 })
1585 .cloned(),
1586 );
1587 output.observation_predictions.extend(
1588 result
1589 .observation_predictions
1590 .iter()
1591 .filter(|block| {
1592 block.producer_node == binding.node_id
1593 && block.producer_port.as_deref() == Some(binding.port_name.as_str())
1594 && block.partition == crate::oof::PredictionPartition::Final
1595 && block.fold_id.is_none()
1596 })
1597 .cloned(),
1598 );
1599 output.aggregated_predictions.extend(
1600 result
1601 .aggregated_predictions
1602 .iter()
1603 .filter(|block| {
1604 block.producer_node == binding.node_id
1605 && block.producer_port.as_deref() == Some(binding.port_name.as_str())
1606 && block.partition == crate::oof::PredictionPartition::Final
1607 && block.fold_id.is_none()
1608 })
1609 .cloned(),
1610 );
1611 }
1612 if !output.predictions.is_empty()
1613 || !output.observation_predictions.is_empty()
1614 || !output.aggregated_predictions.is_empty()
1615 {
1616 output.validate(&source.effective_plan)?;
1617 outputs.push(output);
1618 }
1619 }
1620 outputs.sort_by(|left, right| left.binding.binding_id.cmp(&right.binding.binding_id));
1621 Ok(outputs)
1622}
1623
1624fn bind_package_replay_outputs(
1625 package: &PortablePredictorPackage,
1626 request: &TrainingReplayRequest,
1627 results: &[crate::runtime::NodeResult],
1628) -> Result<Vec<BoundTrainingOutput>> {
1629 let mut outputs = Vec::new();
1630 for binding_id in &request.output_binding_ids {
1631 let binding = package
1632 .output_bindings
1633 .iter()
1634 .find(|binding| binding.binding_id == *binding_id)
1635 .ok_or_else(|| {
1636 DagMlError::RuntimeValidation(format!(
1637 "training replay request references absent package binding `{binding_id}`"
1638 ))
1639 })?
1640 .clone();
1641 let mut output = BoundTrainingOutput {
1642 schema_version: Some(BOUND_TRAINING_OUTPUT_SCHEMA_VERSION),
1643 binding: binding.clone(),
1644 predictions: Vec::new(),
1645 observation_predictions: Vec::new(),
1646 aggregated_predictions: Vec::new(),
1647 };
1648 for result in results {
1649 output.predictions.extend(
1650 result
1651 .predictions
1652 .iter()
1653 .filter(|block| {
1654 block.producer_node == binding.node_id
1655 && block.producer_port.as_deref() == Some(binding.port_name.as_str())
1656 && block.partition == crate::oof::PredictionPartition::Final
1657 && block.fold_id.is_none()
1658 })
1659 .cloned(),
1660 );
1661 output.observation_predictions.extend(
1662 result
1663 .observation_predictions
1664 .iter()
1665 .filter(|block| {
1666 block.producer_node == binding.node_id
1667 && block.producer_port.as_deref() == Some(binding.port_name.as_str())
1668 && block.partition == crate::oof::PredictionPartition::Final
1669 && block.fold_id.is_none()
1670 })
1671 .cloned(),
1672 );
1673 output.aggregated_predictions.extend(
1674 result
1675 .aggregated_predictions
1676 .iter()
1677 .filter(|block| {
1678 block.producer_node == binding.node_id
1679 && block.producer_port.as_deref() == Some(binding.port_name.as_str())
1680 && block.partition == crate::oof::PredictionPartition::Final
1681 && block.fold_id.is_none()
1682 })
1683 .cloned(),
1684 );
1685 }
1686 if !output.predictions.is_empty()
1687 || !output.observation_predictions.is_empty()
1688 || !output.aggregated_predictions.is_empty()
1689 {
1690 output.validate(&package.effective_plan)?;
1691 outputs.push(output);
1692 }
1693 }
1694 outputs.sort_by(|left, right| left.binding.binding_id.cmp(&right.binding.binding_id));
1695 Ok(outputs)
1696}
1697
1698fn apply_replay_conformal_intervals(
1699 calibration: &ConformalCalibration,
1700 outputs: &[BoundTrainingOutput],
1701) -> Result<Vec<ConformalIntervalBlock>> {
1702 let Some(output) = outputs
1703 .iter()
1704 .find(|output| output.binding.binding_id == calibration.binding_id)
1705 else {
1706 return Ok(Vec::new());
1709 };
1710 if output.binding.target_names != calibration.target_names {
1711 return contract_error("replay conformal binding target order differs from calibration");
1712 }
1713 output
1714 .predictions
1715 .iter()
1716 .map(|block| calibration.apply(block))
1717 .collect()
1718}
1719
1720fn validate_replay_interval_closure(
1721 calibration: &ConformalCalibration,
1722 outputs: &[BoundTrainingOutput],
1723 intervals: &[ConformalIntervalBlock],
1724) -> Result<()> {
1725 let expected = apply_replay_conformal_intervals(calibration, outputs)?;
1726 if intervals != expected {
1727 return Err(DagMlError::RuntimeValidation(
1728 "conformal intervals do not exactly cover replay point predictions".to_string(),
1729 ));
1730 }
1731 for interval in intervals {
1732 let output = outputs
1733 .iter()
1734 .find(|output| output.binding.binding_id == interval.binding_id)
1735 .ok_or_else(|| {
1736 DagMlError::RuntimeValidation(
1737 "conformal interval references an absent replay output binding".to_string(),
1738 )
1739 })?;
1740 let point = output
1741 .predictions
1742 .iter()
1743 .find(|point| point.sample_ids == interval.sample_ids)
1744 .ok_or_else(|| {
1745 DagMlError::RuntimeValidation(
1746 "conformal interval has no matching replay point block".to_string(),
1747 )
1748 })?;
1749 interval.validate_against(calibration, point)?;
1750 }
1751 Ok(())
1752}
1753
1754fn bind_attached_replay_explanations(
1755 request: &TrainingReplayRequest,
1756 results: &[crate::runtime::NodeResult],
1757) -> Result<Vec<ExplanationBlock>> {
1758 if request.phase != Phase::Explain {
1759 return Ok(Vec::new());
1760 }
1761 let mut explanations = results
1762 .iter()
1763 .flat_map(|result| result.explanations.iter().cloned())
1764 .filter(|block| block.producer_port.is_some())
1765 .collect::<Vec<_>>();
1766 explanations.sort_by(|left, right| {
1767 (
1768 left.producer_node.as_str(),
1769 left.producer_port.as_deref().unwrap_or_default(),
1770 left.method.as_str(),
1771 left.target_name.as_deref().unwrap_or_default(),
1772 )
1773 .cmp(&(
1774 right.producer_node.as_str(),
1775 right.producer_port.as_deref().unwrap_or_default(),
1776 right.method.as_str(),
1777 right.target_name.as_deref().unwrap_or_default(),
1778 ))
1779 });
1780 Ok(explanations)
1781}
1782
1783fn validate_output_order_and_version(outputs: &[BoundTrainingOutput]) -> Result<()> {
1784 let mut previous: Option<&str> = None;
1785 for output in outputs {
1786 match output.schema_version {
1787 Some(BOUND_TRAINING_OUTPUT_SCHEMA_VERSION) => {}
1788 Some(version) => {
1789 return contract_error(format!(
1790 "training replay output schema_version {version} is unsupported; current {BOUND_TRAINING_OUTPUT_SCHEMA_VERSION}"
1791 ));
1792 }
1793 None => {
1794 return contract_error(
1795 "training replay output requires bound_training_output schema_version",
1796 );
1797 }
1798 }
1799 let binding_id = output.binding.binding_id.as_str();
1800 if previous.is_some_and(|previous| previous >= binding_id) {
1801 return contract_error("training replay outputs must be strictly sorted by binding_id");
1802 }
1803 previous = Some(binding_id);
1804 }
1805 Ok(())
1806}
1807
1808fn validate_replay_bound_output_blocks(output: &BoundTrainingOutput) -> Result<()> {
1809 for block in &output.predictions {
1810 validate_optional_port(
1811 "training replay prediction producer_port",
1812 &block.producer_port,
1813 )?;
1814 if block.partition != crate::oof::PredictionPartition::Final || block.fold_id.is_some() {
1815 return contract_error(
1816 "training replay prediction blocks must use final partition without fold",
1817 );
1818 }
1819 }
1820 for block in &output.observation_predictions {
1821 validate_optional_port(
1822 "training replay observation prediction producer_port",
1823 &block.producer_port,
1824 )?;
1825 if block.partition != crate::oof::PredictionPartition::Final || block.fold_id.is_some() {
1826 return contract_error(
1827 "training replay observation prediction blocks must use final partition without fold",
1828 );
1829 }
1830 }
1831 for block in &output.aggregated_predictions {
1832 validate_optional_port(
1833 "training replay aggregated prediction producer_port",
1834 &block.producer_port,
1835 )?;
1836 if block.partition != crate::oof::PredictionPartition::Final || block.fold_id.is_some() {
1837 return contract_error(
1838 "training replay aggregated prediction blocks must use final partition without fold",
1839 );
1840 }
1841 }
1842 Ok(())
1843}
1844
1845fn validate_replay_phase(phase: Phase) -> Result<()> {
1846 if matches!(phase, Phase::Predict | Phase::Explain) {
1847 Ok(())
1848 } else {
1849 contract_error("training replay supports only PREDICT and EXPLAIN")
1850 }
1851}
1852
1853fn validate_sha256(label: &str, value: &str) -> Result<()> {
1854 if value.len() == 64
1855 && value
1856 .bytes()
1857 .all(|byte| byte.is_ascii_digit() || (b'a'..=b'f').contains(&byte))
1858 {
1859 Ok(())
1860 } else {
1861 contract_error(format!(
1862 "{label} fingerprint must be 64 lowercase hexadecimal characters"
1863 ))
1864 }
1865}
1866
1867fn validate_identifier(label: &str, value: &str) -> Result<()> {
1868 if !value.is_empty()
1869 && value.len() <= 128
1870 && value
1871 .bytes()
1872 .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'_' | b'-' | b'.' | b':'))
1873 {
1874 Ok(())
1875 } else {
1876 contract_error(format!("{label} is not a valid DAG-ML identifier"))
1877 }
1878}
1879
1880fn validate_non_empty(label: &str, value: &str) -> Result<()> {
1881 if value.trim().is_empty() {
1882 contract_error(format!("{label} must be non-empty"))
1883 } else {
1884 Ok(())
1885 }
1886}
1887
1888fn validate_sorted_unique_identifiers(
1889 label: &str,
1890 values: &[String],
1891 require_non_empty: bool,
1892) -> Result<()> {
1893 validate_sorted_unique_text(label, values, require_non_empty)?;
1894 for value in values {
1895 validate_identifier(label, value)?;
1896 }
1897 Ok(())
1898}
1899
1900fn validate_sorted_unique_text(
1901 label: &str,
1902 values: &[String],
1903 require_non_empty: bool,
1904) -> Result<()> {
1905 if require_non_empty && values.is_empty() {
1906 return contract_error(format!("{label} must be non-empty"));
1907 }
1908 let mut previous: Option<&str> = None;
1909 for value in values {
1910 validate_non_empty(label, value)?;
1911 if previous.is_some_and(|previous| previous >= value.as_str()) {
1912 return contract_error(format!("{label} must be strictly sorted and unique"));
1913 }
1914 previous = Some(value.as_str());
1915 }
1916 Ok(())
1917}
1918
1919fn validate_sorted_unique_keys<'a>(
1920 label: &str,
1921 values: impl Iterator<Item = &'a str>,
1922 require_non_empty: bool,
1923) -> Result<()> {
1924 let values = values.collect::<Vec<_>>();
1925 if require_non_empty && values.is_empty() {
1926 return contract_error(format!("{label} must be non-empty"));
1927 }
1928 let mut previous: Option<&str> = None;
1929 for value in values {
1930 validate_non_empty(label, value)?;
1931 if previous.is_some_and(|previous| previous >= value) {
1932 return contract_error(format!("{label} must be strictly sorted and unique"));
1933 }
1934 previous = Some(value);
1935 }
1936 Ok(())
1937}
1938
1939fn validate_optional_port(label: &str, value: &Option<String>) -> Result<()> {
1940 match value {
1941 Some(value) if !value.trim().is_empty() => Ok(()),
1942 _ => contract_error(format!("{label} must be present and non-empty")),
1943 }
1944}
1945
1946fn validate_diagnostics(diagnostics: &BTreeMap<String, serde_json::Value>) -> Result<()> {
1947 for (key, value) in diagnostics {
1948 validate_non_empty("training replay diagnostic key", key)?;
1949 if !matches!(
1950 value,
1951 serde_json::Value::Null
1952 | serde_json::Value::Bool(_)
1953 | serde_json::Value::Number(_)
1954 | serde_json::Value::String(_)
1955 ) {
1956 return contract_error("training replay diagnostics must be scalar JSON values");
1957 }
1958 }
1959 Ok(())
1960}
1961
1962fn require_count(label: &str, actual: usize, expected: usize) -> Result<()> {
1963 if actual == expected {
1964 Ok(())
1965 } else {
1966 contract_error(format!("{label} does not match replay payload"))
1967 }
1968}
1969
1970fn zero_fingerprint() -> String {
1971 "0".repeat(64)
1972}
1973
1974fn tcv1_fingerprint_without<T: Serialize>(value: &T, field: &str, label: &str) -> Result<String> {
1975 let json = serde_json::to_string(value)?;
1976 strict_tcv1_fingerprint_without(&json, field, label)
1977}
1978
1979fn strict_tcv1_fingerprint_without(json: &str, field: &str, label: &str) -> Result<String> {
1980 parse_typed_json(json)
1981 .and_then(|value| value.fingerprint_without(field))
1982 .map_err(|error| {
1983 DagMlError::RuntimeValidation(format!("{label} is outside strict TCV1: {error}"))
1984 })
1985}
1986
1987fn unsupported_version<T>(label: &str, actual: u32, expected: u32) -> Result<T> {
1988 contract_error(format!(
1989 "{label} uses unsupported schema_version {actual}, expected {expected}"
1990 ))
1991}
1992
1993fn contract_error<T>(message: impl Into<String>) -> Result<T> {
1994 Err(DagMlError::CampaignValidation(message.into()))
1995}
1996
1997#[cfg(test)]
1998mod methods_hpo_resume_state_tests {
1999 use super::*;
2000
2001 #[test]
2002 fn methods_hpo_resume_state_reader_refuses_unknown_legacy_and_noncanonical_json() {
2003 let unknown = r#"{"schema_version":1,"unknown_resume_side_channel":true}"#;
2004 assert!(methods_hpo_resume_state_from_json(unknown).is_err());
2005
2006 let duplicate = r#"{"schema_version":1,"schema_version":1}"#;
2009 let error = methods_hpo_resume_state_from_json(duplicate)
2010 .unwrap_err()
2011 .to_string();
2012 assert!(error.contains("duplicate JSON object key"), "{error}");
2013
2014 let legacy = r#"{"schema_version":1,"tuner_node_id":"tuner:legacy"}"#;
2018 assert!(methods_hpo_resume_state_from_json(legacy).is_err());
2019 }
2020
2021 #[test]
2022 fn methods_hpo_resume_state_reader_refuses_tampered_schema_version() {
2023 let tampered = r#"{"schema_version":2,"operation_id":"campaign:hpo"}"#;
2024 assert!(methods_hpo_resume_state_from_json(tampered).is_err());
2025 }
2026}
2027
2028#[cfg(test)]
2029mod replay_identity_tests {
2030 use super::*;
2031
2032 fn replay_identity(target_content_fingerprint: Option<&str>) -> ReplayDataIdentity {
2033 let mut identity = ReplayDataIdentity {
2034 requirement_key: "model:base.X".to_string(),
2035 schema_fingerprint: "1".repeat(64),
2036 plan_fingerprint: "2".repeat(64),
2037 relation_fingerprint: "3".repeat(64),
2038 data_content_fingerprint: "4".repeat(64),
2039 target_content_fingerprint: target_content_fingerprint.map(str::to_string),
2040 identity_fingerprint: zero_fingerprint(),
2041 };
2042 identity.identity_fingerprint = identity.compute_fingerprint().unwrap();
2043 identity
2044 }
2045
2046 #[test]
2047 fn replay_identity_permits_an_unlabeled_predict_cohort_without_a_sentinel() {
2048 let identity = replay_identity(None);
2049 identity.validate().unwrap();
2050 assert!(identity.target_content_fingerprint.is_none());
2051 assert!(serde_json::to_value(&identity)
2052 .unwrap()
2053 .get("target_content_fingerprint")
2054 .unwrap()
2055 .is_null());
2056 }
2057
2058 #[test]
2059 fn replay_identity_rejects_a_resigned_target_fingerprint_tamper() {
2060 let mut identity = replay_identity(None);
2061 identity.target_content_fingerprint = Some("not-a-fingerprint".to_string());
2062 identity.identity_fingerprint = identity.compute_fingerprint().unwrap();
2063 let error = identity.validate().unwrap_err().to_string();
2064 assert!(error.contains("target content"), "{error}");
2065 }
2066
2067 #[test]
2068 fn replay_schema_version_is_v3_for_target_bound_and_target_free_cohorts() {
2069 assert_eq!(
2070 replay_outcome_schema_version(&[replay_identity(None)]),
2071 TRAINING_REPLAY_OUTCOME_SCHEMA_VERSION
2072 );
2073 assert_eq!(
2074 replay_outcome_schema_version(&[replay_identity(Some(&"5".repeat(64)))]),
2075 TRAINING_REPLAY_OUTCOME_SCHEMA_VERSION
2076 );
2077 }
2078}
2079
2080#[cfg(test)]
2081mod tests {
2082 #[cfg(dag_ml_workspace_contract_fixtures)]
2083 use std::fs;
2084 #[cfg(dag_ml_workspace_contract_fixtures)]
2085 use std::path::PathBuf;
2086
2087 #[cfg(dag_ml_workspace_contract_fixtures)]
2088 use super::*;
2089
2090 #[cfg(dag_ml_workspace_contract_fixtures)]
2091 fn root() -> PathBuf {
2092 PathBuf::from(env!("CARGO_MANIFEST_DIR"))
2093 .parent()
2094 .and_then(|path| path.parent())
2095 .expect("core crate is under crates/dag-ml-core")
2096 .to_path_buf()
2097 }
2098
2099 #[cfg(dag_ml_workspace_contract_fixtures)]
2100 fn fixture(name: &str) -> String {
2101 fs::read_to_string(
2102 root()
2103 .join("examples")
2104 .join("fixtures")
2105 .join("training")
2106 .join("replay")
2107 .join(name),
2108 )
2109 .expect(name)
2110 }
2111
2112 #[cfg(dag_ml_workspace_contract_fixtures)]
2113 fn training_fixture(name: &str) -> String {
2114 fs::read_to_string(
2115 root()
2116 .join("examples")
2117 .join("fixtures")
2118 .join("training")
2119 .join(name),
2120 )
2121 .expect(name)
2122 }
2123
2124 #[cfg(dag_ml_workspace_contract_fixtures)]
2125 #[test]
2126 fn training_replay_contract_fixtures_parse_and_cross_validate() {
2127 let predict_source =
2128 TrainingOutcome::from_json(&training_fixture("training_outcome_refit.v1.json"))
2129 .expect("predict source training outcome");
2130 let explain_source =
2131 TrainingOutcome::from_json(&fixture("training_replay_source_outcome_explain.v1.json"))
2132 .expect("explain source training outcome");
2133 let predict_request =
2134 TrainingReplayRequest::from_json(&fixture("training_replay_request_predict.v1.json"))
2135 .expect("predict request");
2136 let predict_outcome =
2137 TrainingReplayOutcome::from_json(&fixture("training_replay_outcome_predict.v1.json"))
2138 .expect("predict outcome");
2139 predict_outcome
2140 .validate_against(&predict_source, &predict_request)
2141 .expect("predict cross-links");
2142
2143 let explain_request =
2144 TrainingReplayRequest::from_json(&fixture("training_replay_request_explain.v1.json"))
2145 .expect("explain request");
2146 let explain_outcome =
2147 TrainingReplayOutcome::from_json(&fixture("training_replay_outcome_explain.v1.json"))
2148 .expect("explain outcome");
2149 explain_outcome
2150 .validate_against(&explain_source, &explain_request)
2151 .expect("explain cross-links");
2152
2153 let explain_only = TrainingReplayOutcome::from_json(&fixture(
2154 "training_replay_outcome_explain_only.v1.json",
2155 ))
2156 .expect("explain-only outcome");
2157 explain_only
2158 .validate_against(&explain_source, &explain_request)
2159 .expect("explain-only cross-links");
2160 }
2161
2162 #[cfg(dag_ml_workspace_contract_fixtures)]
2163 #[test]
2164 fn training_replay_request_rejects_refit_and_unsorted_bindings() {
2165 let mut request: serde_json::Value =
2166 serde_json::from_str(&fixture("training_replay_request_predict.v1.json")).unwrap();
2167 request["phase"] = serde_json::Value::String("REFIT".to_string());
2168 let err = serde_json::from_value::<TrainingReplayRequest>(request)
2169 .unwrap()
2170 .validate()
2171 .unwrap_err()
2172 .to_string();
2173 assert!(err.contains("PREDICT and EXPLAIN"));
2174
2175 let mut request: TrainingReplayRequest =
2176 TrainingReplayRequest::from_json(&fixture("training_replay_request_predict.v1.json"))
2177 .unwrap();
2178 request.output_binding_ids = vec!["z".to_string(), "a".to_string()];
2179 request.request_fingerprint = request.compute_fingerprint().unwrap();
2180 let err = request.validate().unwrap_err().to_string();
2181 assert!(err.contains("strictly sorted"));
2182 }
2183
2184 #[cfg(dag_ml_workspace_contract_fixtures)]
2185 #[test]
2186 fn training_replay_outcome_rejects_counter_and_source_transplants() {
2187 let source =
2188 TrainingOutcome::from_json(&training_fixture("training_outcome_refit.v1.json"))
2189 .unwrap();
2190 let request =
2191 TrainingReplayRequest::from_json(&fixture("training_replay_request_predict.v1.json"))
2192 .unwrap();
2193 let mut outcome =
2194 TrainingReplayOutcome::from_json(&fixture("training_replay_outcome_predict.v1.json"))
2195 .unwrap();
2196 outcome.prediction_block_count += 1;
2197 outcome.outcome_fingerprint = outcome.compute_fingerprint().unwrap();
2198 let err = outcome.validate().unwrap_err().to_string();
2199 assert!(err.contains("prediction_block_count"));
2200
2201 let mut outcome =
2202 TrainingReplayOutcome::from_json(&fixture("training_replay_outcome_predict.v1.json"))
2203 .unwrap();
2204 outcome.source_training_outcome.outcome_fingerprint = "f".repeat(64);
2205 outcome.outcome_fingerprint = outcome.compute_fingerprint().unwrap();
2206 let err = outcome
2207 .validate_against(&source, &request)
2208 .unwrap_err()
2209 .to_string();
2210 assert!(err.contains("source reference"));
2211 }
2212
2213 #[cfg(dag_ml_workspace_contract_fixtures)]
2214 #[test]
2215 fn target_free_predict_evidence_is_v3_only() {
2216 let mut outcome =
2217 TrainingReplayOutcome::from_json(&fixture("training_replay_outcome_predict.v1.json"))
2218 .unwrap();
2219 outcome.schema_version = TRAINING_REPLAY_OUTCOME_SCHEMA_VERSION;
2220 for identity in &mut outcome.input_data_identities {
2221 identity.target_content_fingerprint = None;
2222 identity.identity_fingerprint = identity.compute_fingerprint().unwrap();
2223 }
2224 outcome.outcome_fingerprint = outcome.compute_fingerprint().unwrap();
2225 outcome.validate().unwrap();
2226
2227 outcome.schema_version = CONFORMAL_TRAINING_REPLAY_OUTCOME_SCHEMA_VERSION;
2228 outcome.outcome_fingerprint = outcome.compute_fingerprint().unwrap();
2229 let error = outcome.validate().unwrap_err().to_string();
2230 assert!(error.contains("V1/V2 requires target-bound"), "{error}");
2231 }
2232}