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