1use std::collections::{BTreeMap, BTreeSet};
9
10use serde::{Deserialize, Serialize};
11
12use crate::bundle::{ExecutionBundle, ReplayPhaseRequest};
13use crate::campaign::stable_json_fingerprint;
14use crate::canonical::parse_typed_json;
15use crate::data::ExternalDataPlanEnvelope;
16use crate::error::{DagMlError, Result};
17use crate::ids::{ArtifactId, BundleId, RunId};
18use crate::phase::Phase;
19use crate::plan::ExecutionPlan;
20use crate::runtime::{
21 ArtifactMaterializationRequest, BundleReplayExecution, ExplanationBlock, HandleRef,
22 LineageRecord, RunContext, RuntimeArtifactStore, RuntimeControllerRegistry,
23 RuntimeDataProvider, SequentialScheduler,
24};
25use crate::training::{
26 LoadedPredictor, PortablePredictorPackage, TrainingDataIdentity, TrainingOutcomeRef,
27};
28use crate::training_runtime::{
29 BoundTrainingOutput, TrainingOutcome, BOUND_TRAINING_OUTPUT_SCHEMA_VERSION,
30};
31
32pub const TRAINING_REPLAY_REQUEST_SCHEMA_VERSION: u32 = 1;
33pub const TRAINING_REPLAY_OUTCOME_SCHEMA_VERSION: u32 = 1;
34
35pub struct AttachedTrainingReplayInput<'a> {
36 pub source: &'a TrainingOutcome,
37 pub request: &'a TrainingReplayRequest,
38 pub outcome_id: String,
39 pub run_id: RunId,
40 pub controllers: &'a RuntimeControllerRegistry,
41 pub data_provider: &'a dyn RuntimeDataProvider,
42 pub artifact_store: &'a dyn RuntimeArtifactStore,
43 pub data_envelopes: &'a BTreeMap<String, ExternalDataPlanEnvelope>,
44 pub warnings: Vec<String>,
45 pub diagnostics: BTreeMap<String, serde_json::Value>,
46}
47
48pub struct LoadedPredictorReplayInput<'a> {
49 pub predictor: &'a LoadedPredictor<HandleRef>,
50 pub request: &'a TrainingReplayRequest,
51 pub outcome_id: String,
52 pub run_id: RunId,
53 pub controllers: &'a RuntimeControllerRegistry,
54 pub data_provider: &'a dyn RuntimeDataProvider,
55 pub data_envelopes: &'a BTreeMap<String, ExternalDataPlanEnvelope>,
56 pub warnings: Vec<String>,
57 pub diagnostics: BTreeMap<String, serde_json::Value>,
58}
59
60#[derive(Clone, Debug, Eq, PartialEq, Serialize, Deserialize)]
61#[serde(deny_unknown_fields)]
62pub struct TrainingReplayRequest {
63 pub schema_version: u32,
64 pub request_id: String,
65 pub source_outcome_fingerprint: String,
66 pub phase: Phase,
67 pub data_envelope_keys: Vec<String>,
68 pub output_binding_ids: Vec<String>,
69 pub request_fingerprint: String,
70}
71
72impl TrainingReplayRequest {
73 pub fn from_json(json: &str) -> Result<Self> {
74 let raw_fingerprint = strict_tcv1_fingerprint_without(
75 json,
76 "request_fingerprint",
77 "training replay request",
78 )?;
79 let request: Self = serde_json::from_str(json)?;
80 if request.request_fingerprint != raw_fingerprint {
81 return contract_error(
82 "training replay request fingerprint does not match original TCV1 JSON",
83 );
84 }
85 request.validate()?;
86 Ok(request)
87 }
88
89 pub fn compute_fingerprint(&self) -> Result<String> {
90 tcv1_fingerprint_without(self, "request_fingerprint", "training replay request")
91 }
92
93 pub fn validate(&self) -> Result<()> {
94 if self.schema_version != TRAINING_REPLAY_REQUEST_SCHEMA_VERSION {
95 return unsupported_version(
96 "training replay request",
97 self.schema_version,
98 TRAINING_REPLAY_REQUEST_SCHEMA_VERSION,
99 );
100 }
101 validate_identifier("training replay request_id", &self.request_id)?;
102 validate_sha256(
103 "training replay source outcome",
104 &self.source_outcome_fingerprint,
105 )?;
106 validate_replay_phase(self.phase)?;
107 validate_sorted_unique_text(
108 "training replay data_envelope_keys",
109 &self.data_envelope_keys,
110 true,
111 )?;
112 validate_sorted_unique_identifiers(
113 "training replay output_binding_ids",
114 &self.output_binding_ids,
115 true,
116 )?;
117 validate_sha256("training replay request", &self.request_fingerprint)?;
118 if self.request_fingerprint != self.compute_fingerprint()? {
119 return contract_error(
120 "training replay request fingerprint does not match TCV1 content",
121 );
122 }
123 Ok(())
124 }
125}
126
127#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
128#[serde(deny_unknown_fields)]
129pub struct TrainingReplayOutcome {
130 pub schema_version: u32,
131 pub outcome_id: String,
132 pub run_id: RunId,
133 pub source_training_outcome: TrainingOutcomeRef,
134 pub replay_request_id: String,
135 pub replay_request_fingerprint: String,
136 pub input_data_identities: Vec<TrainingDataIdentity>,
137 pub bundle_id: BundleId,
138 pub plan_id: String,
139 pub phase: Phase,
140 pub result_count: usize,
141 pub lineage_record_count: usize,
142 pub prediction_block_count: usize,
143 pub observation_prediction_block_count: usize,
144 pub aggregated_prediction_block_count: usize,
145 pub explanation_block_count: usize,
146 pub controller_count: usize,
147 pub prediction_cache_store: bool,
148 pub outputs: Vec<BoundTrainingOutput>,
149 pub explanations: Vec<ExplanationBlock>,
150 pub lineage: Vec<LineageRecord>,
151 pub warnings: Vec<String>,
152 pub diagnostics: BTreeMap<String, serde_json::Value>,
153 pub outcome_fingerprint: String,
154}
155
156impl TrainingReplayOutcome {
157 pub fn from_json(json: &str) -> Result<Self> {
158 let raw_fingerprint = strict_tcv1_fingerprint_without(
159 json,
160 "outcome_fingerprint",
161 "training replay outcome",
162 )?;
163 let outcome: Self = serde_json::from_str(json)?;
164 if outcome.outcome_fingerprint != raw_fingerprint {
165 return contract_error(
166 "training replay outcome fingerprint does not match original TCV1 JSON",
167 );
168 }
169 outcome.validate()?;
170 Ok(outcome)
171 }
172
173 pub fn compute_fingerprint(&self) -> Result<String> {
174 tcv1_fingerprint_without(self, "outcome_fingerprint", "training replay outcome")
175 }
176
177 pub fn validate(&self) -> Result<()> {
178 if self.schema_version != TRAINING_REPLAY_OUTCOME_SCHEMA_VERSION {
179 return unsupported_version(
180 "training replay outcome",
181 self.schema_version,
182 TRAINING_REPLAY_OUTCOME_SCHEMA_VERSION,
183 );
184 }
185 validate_identifier("training replay outcome_id", &self.outcome_id)?;
186 validate_identifier("training replay request_id", &self.replay_request_id)?;
187 validate_sha256("training replay request", &self.replay_request_fingerprint)?;
188 validate_non_empty("training replay plan_id", &self.plan_id)?;
189 validate_replay_phase(self.phase)?;
190 self.source_training_outcome.validate()?;
191 for identity in &self.input_data_identities {
192 identity.validate()?;
193 }
194 validate_sorted_unique_keys(
195 "training replay input_data_identities",
196 self.input_data_identities
197 .iter()
198 .map(|identity| identity.requirement_key.as_str()),
199 true,
200 )?;
201 if self.prediction_cache_store {
202 return contract_error("training replay outcome cannot persist a prediction cache");
203 }
204 validate_sorted_unique_text("training replay warnings", &self.warnings, false)?;
205 validate_diagnostics(&self.diagnostics)?;
206 validate_output_order_and_version(&self.outputs)?;
207 for output in &self.outputs {
208 validate_replay_bound_output_blocks(output)?;
209 }
210 for explanation in &self.explanations {
211 explanation.validate()?;
212 validate_optional_port(
213 "training replay explanation producer_port",
214 &explanation.producer_port,
215 )?;
216 }
217 for record in &self.lineage {
218 record.validate()?;
219 }
220 match self.phase {
221 Phase::Predict if self.outputs.is_empty() => {
222 return contract_error("training replay PREDICT requires at least one output");
223 }
224 Phase::Predict if !self.explanations.is_empty() => {
225 return contract_error("training replay PREDICT cannot emit explanations");
226 }
227 Phase::Explain if self.explanations.is_empty() => {
228 return contract_error("training replay EXPLAIN requires at least one explanation");
229 }
230 _ => {}
231 }
232 self.validate_counters()?;
233 validate_sha256("training replay outcome", &self.outcome_fingerprint)?;
234 if self.outcome_fingerprint != self.compute_fingerprint()? {
235 return contract_error(
236 "training replay outcome fingerprint does not match TCV1 content",
237 );
238 }
239 Ok(())
240 }
241
242 pub fn validate_against(
243 &self,
244 source: &TrainingOutcome,
245 request: &TrainingReplayRequest,
246 ) -> Result<()> {
247 self.validate()?;
248 source.validate()?;
249 request.validate()?;
250 if request.source_outcome_fingerprint != source.outcome_fingerprint {
251 return contract_error("training replay request does not target source outcome");
252 }
253 if !source.replayable_phases.contains(&request.phase) {
254 return contract_error("training replay phase is not replayable by source outcome");
255 }
256 if self.source_training_outcome != source.to_reference()? {
257 return contract_error(
258 "training replay outcome source reference does not match source outcome",
259 );
260 }
261 if self.replay_request_id != request.request_id {
262 return contract_error(
263 "training replay outcome request id does not match ReplayRequest",
264 );
265 }
266 if self.replay_request_fingerprint != request.request_fingerprint {
267 return contract_error(
268 "training replay outcome request fingerprint does not match ReplayRequest",
269 );
270 }
271 if self.phase != request.phase {
272 return contract_error("training replay outcome phase does not match ReplayRequest");
273 }
274 if self.bundle_id != source.execution_bundle.bundle_id {
275 return contract_error("training replay outcome bundle does not match source outcome");
276 }
277 if self.plan_id != source.effective_plan.id {
278 return contract_error("training replay outcome plan does not match source outcome");
279 }
280 let identity_keys = self
281 .input_data_identities
282 .iter()
283 .map(|identity| identity.requirement_key.clone())
284 .collect::<Vec<_>>();
285 if identity_keys != request.data_envelope_keys {
286 return contract_error(
287 "training replay outcome identities do not exactly cover ReplayRequest envelopes",
288 );
289 }
290 let source_bindings = source
291 .outputs
292 .iter()
293 .map(|output| (output.binding.binding_id.as_str(), &output.binding))
294 .collect::<BTreeMap<_, _>>();
295 for binding_id in &request.output_binding_ids {
296 if !source_bindings.contains_key(binding_id.as_str()) {
297 return contract_error(
298 "training replay request references absent source output binding",
299 );
300 }
301 }
302 let emitted_binding_ids = self
303 .outputs
304 .iter()
305 .map(|output| output.binding.binding_id.clone())
306 .collect::<Vec<_>>();
307 if self.phase == Phase::Predict && emitted_binding_ids != request.output_binding_ids {
308 return contract_error(
309 "training replay PREDICT outputs do not exactly cover ReplayRequest bindings",
310 );
311 }
312 if self.phase == Phase::Explain
313 && !emitted_binding_ids
314 .iter()
315 .all(|binding_id| request.output_binding_ids.contains(binding_id))
316 {
317 return contract_error(
318 "training replay EXPLAIN outputs include a binding outside ReplayRequest",
319 );
320 }
321 for output in &self.outputs {
322 let Some(source_binding) = source_bindings.get(output.binding.binding_id.as_str())
323 else {
324 return contract_error(
325 "training replay output binding is absent from source outcome",
326 );
327 };
328 if &output.binding != *source_binding {
329 return contract_error(
330 "training replay output binding does not match source outcome binding",
331 );
332 }
333 output.validate(&source.effective_plan)?;
334 }
335 Ok(())
336 }
337
338 pub fn validate_against_package(
339 &self,
340 package: &PortablePredictorPackage,
341 request: &TrainingReplayRequest,
342 ) -> Result<()> {
343 self.validate()?;
344 package.validate()?;
345 request.validate()?;
346 validate_replay_phase(request.phase)?;
347 if request.source_outcome_fingerprint != package.training_outcome.outcome_fingerprint {
348 return contract_error(
349 "training replay request does not target package source outcome",
350 );
351 }
352 if self.source_training_outcome != package.training_outcome {
353 return contract_error(
354 "training replay outcome source reference does not match package source outcome",
355 );
356 }
357 if self.replay_request_id != request.request_id {
358 return contract_error(
359 "training replay outcome request id does not match ReplayRequest",
360 );
361 }
362 if self.replay_request_fingerprint != request.request_fingerprint {
363 return contract_error(
364 "training replay outcome request fingerprint does not match ReplayRequest",
365 );
366 }
367 if self.phase != request.phase {
368 return contract_error("training replay outcome phase does not match ReplayRequest");
369 }
370 if self.bundle_id != package.execution_bundle.bundle_id {
371 return contract_error("training replay outcome bundle does not match package");
372 }
373 if self.plan_id != package.effective_plan.id {
374 return contract_error("training replay outcome plan does not match package");
375 }
376 let identity_keys = self
377 .input_data_identities
378 .iter()
379 .map(|identity| identity.requirement_key.clone())
380 .collect::<Vec<_>>();
381 if identity_keys != request.data_envelope_keys {
382 return contract_error(
383 "training replay outcome identities do not exactly cover ReplayRequest envelopes",
384 );
385 }
386 let package_bindings = package
387 .output_bindings
388 .iter()
389 .map(|binding| (binding.binding_id.as_str(), binding))
390 .collect::<BTreeMap<_, _>>();
391 for binding_id in &request.output_binding_ids {
392 if !package_bindings.contains_key(binding_id.as_str()) {
393 return contract_error(
394 "training replay request references absent package output binding",
395 );
396 }
397 }
398 let emitted_binding_ids = self
399 .outputs
400 .iter()
401 .map(|output| output.binding.binding_id.clone())
402 .collect::<Vec<_>>();
403 if self.phase == Phase::Predict && emitted_binding_ids != request.output_binding_ids {
404 return contract_error(
405 "training replay PREDICT outputs do not exactly cover ReplayRequest bindings",
406 );
407 }
408 if self.phase == Phase::Explain
409 && !emitted_binding_ids
410 .iter()
411 .all(|binding_id| request.output_binding_ids.contains(binding_id))
412 {
413 return contract_error(
414 "training replay EXPLAIN outputs include a binding outside ReplayRequest",
415 );
416 }
417 for output in &self.outputs {
418 let Some(package_binding) = package_bindings.get(output.binding.binding_id.as_str())
419 else {
420 return contract_error("training replay output binding is absent from package");
421 };
422 if &output.binding != *package_binding {
423 return contract_error(
424 "training replay output binding does not match package binding",
425 );
426 }
427 output.validate(&package.effective_plan)?;
428 }
429 Ok(())
430 }
431
432 fn validate_counters(&self) -> Result<()> {
433 require_count(
434 "training replay result_count",
435 self.result_count,
436 self.lineage.len(),
437 )?;
438 require_count(
439 "training replay lineage_record_count",
440 self.lineage_record_count,
441 self.lineage.len(),
442 )?;
443 require_count(
444 "training replay prediction_block_count",
445 self.prediction_block_count,
446 self.outputs
447 .iter()
448 .map(|output| output.predictions.len())
449 .sum(),
450 )?;
451 require_count(
452 "training replay observation_prediction_block_count",
453 self.observation_prediction_block_count,
454 self.outputs
455 .iter()
456 .map(|output| output.observation_predictions.len())
457 .sum(),
458 )?;
459 require_count(
460 "training replay aggregated_prediction_block_count",
461 self.aggregated_prediction_block_count,
462 self.outputs
463 .iter()
464 .map(|output| output.aggregated_predictions.len())
465 .sum(),
466 )?;
467 require_count(
468 "training replay explanation_block_count",
469 self.explanation_block_count,
470 self.explanations.len(),
471 )?;
472 let controller_count = self
473 .lineage
474 .iter()
475 .map(|record| record.controller_id.as_str())
476 .collect::<BTreeSet<_>>()
477 .len();
478 require_count(
479 "training replay controller_count",
480 self.controller_count,
481 controller_count,
482 )?;
483 Ok(())
484 }
485}
486
487struct LoadedPredictorArtifactStore<'a> {
488 predictor: &'a LoadedPredictor<HandleRef>,
489 records: BTreeMap<ArtifactId, crate::bundle::RefitArtifactRecord>,
490}
491
492impl<'a> LoadedPredictorArtifactStore<'a> {
493 fn new(predictor: &'a LoadedPredictor<HandleRef>) -> Result<Self> {
494 predictor.package().validate()?;
495 let records = predictor
496 .package()
497 .execution_bundle
498 .refit_artifacts
499 .iter()
500 .map(|record| {
501 record.validate()?;
502 Ok((record.artifact.id.clone(), record.clone()))
503 })
504 .collect::<Result<BTreeMap<_, _>>>()?;
505 Ok(Self { predictor, records })
506 }
507}
508
509impl RuntimeArtifactStore for LoadedPredictorArtifactStore<'_> {
510 fn materialize(&self, request: &ArtifactMaterializationRequest) -> Result<HandleRef> {
511 let record = self.records.get(&request.artifact.id).ok_or_else(|| {
512 DagMlError::RuntimeValidation(format!(
513 "loaded predictor is missing refit artifact `{}` for bundle `{}`",
514 request.artifact.id, request.bundle_id
515 ))
516 })?;
517 if record.node_id != request.node_id {
518 return Err(DagMlError::RuntimeValidation(format!(
519 "artifact `{}` is registered for node `{}` but requested for `{}`",
520 request.artifact.id, record.node_id, request.node_id
521 )));
522 }
523 if record.controller_id != request.controller_id {
524 return Err(DagMlError::RuntimeValidation(format!(
525 "artifact `{}` is registered for controller `{}` but requested for `{}`",
526 request.artifact.id, record.controller_id, request.controller_id
527 )));
528 }
529 if record.artifact != request.artifact {
530 return Err(DagMlError::RuntimeValidation(format!(
531 "artifact `{}` metadata does not match package bundle record",
532 request.artifact.id
533 )));
534 }
535 if record.params_fingerprint != request.params_fingerprint {
536 return Err(DagMlError::RuntimeValidation(format!(
537 "artifact `{}` params fingerprint does not match package bundle record",
538 request.artifact.id
539 )));
540 }
541 if record.training_loss_fingerprint != request.training_loss_fingerprint {
542 return Err(DagMlError::RuntimeValidation(format!(
543 "artifact `{}` training loss fingerprint does not match package bundle record",
544 request.artifact.id
545 )));
546 }
547 let handle = self
548 .predictor
549 .artifact(&request.artifact.id)
550 .ok_or_else(|| {
551 DagMlError::RuntimeValidation(format!(
552 "loaded predictor has no process-local handle for `{}`",
553 request.artifact.id
554 ))
555 })?;
556 Ok(handle.clone())
557 }
558}
559
560pub fn execute_attached_training_replay(
561 input: AttachedTrainingReplayInput<'_>,
562) -> Result<TrainingReplayOutcome> {
563 input.source.validate()?;
564 input.request.validate()?;
565 validate_sorted_unique_text("training replay execution warnings", &input.warnings, false)?;
566 validate_diagnostics(&input.diagnostics)?;
567 if input.request.source_outcome_fingerprint != input.source.outcome_fingerprint {
568 return contract_error("training replay request does not target source outcome");
569 }
570 if !input
571 .source
572 .replayable_phases
573 .contains(&input.request.phase)
574 {
575 return contract_error("training replay phase is not replayable by source outcome");
576 }
577 for node_plan in input.source.effective_plan.node_plans.values() {
578 if input.controllers.get(&node_plan.controller_id).is_none() {
579 return Err(DagMlError::RuntimeValidation(format!(
580 "attached training replay controller `{}` for node `{}` is not registered",
581 node_plan.controller_id, node_plan.node_id
582 )));
583 }
584 }
585
586 let input_data_identities = replay_input_data_identities(
587 &input.source.execution_bundle,
588 input.request,
589 input.data_envelopes,
590 )?;
591 let (replay_plan, replay_bundle) = replay_plan_and_bundle_for_current_cohort(
592 &input.source.effective_plan,
593 &input.source.execution_bundle,
594 input.request,
595 input.data_envelopes,
596 )?;
597 let phase_request = ReplayPhaseRequest {
598 bundle_id: replay_bundle.bundle_id.clone(),
599 phase: input.request.phase,
600 data_envelope_keys: input.request.data_envelope_keys.clone(),
601 };
602 let mut ctx = RunContext::new(input.run_id.clone(), None);
603 let results = SequentialScheduler.execute_bundle_replay(
604 BundleReplayExecution {
605 plan: &replay_plan,
606 bundle: &replay_bundle,
607 replay_request: &phase_request,
608 prediction_cache_store: None,
609 controllers: input.controllers,
610 data_provider: input.data_provider,
611 artifact_store: input.artifact_store,
612 data_envelopes: input.data_envelopes,
613 },
614 &mut ctx,
615 )?;
616 if results
617 .iter()
618 .any(|result| !result.artifacts.is_empty() || !result.artifact_handles.is_empty())
619 {
620 return contract_error("attached training replay PREDICT/EXPLAIN cannot emit artifacts");
621 }
622
623 let outputs = bind_attached_replay_outputs(input.source, input.request, &results)?;
624 let explanations = bind_attached_replay_explanations(input.request, &results)?;
625 let mut lineage = ctx.lineage.records().cloned().collect::<Vec<_>>();
626 for record in &mut lineage {
627 record.input_lineage.sort();
628 record
629 .artifact_refs
630 .sort_by(|left, right| left.id.cmp(&right.id));
631 }
632 lineage.sort_by(|left, right| left.record_id.cmp(&right.record_id));
633
634 let mut outcome = TrainingReplayOutcome {
635 schema_version: TRAINING_REPLAY_OUTCOME_SCHEMA_VERSION,
636 outcome_id: input.outcome_id,
637 run_id: input.run_id,
638 source_training_outcome: input.source.to_reference()?,
639 replay_request_id: input.request.request_id.clone(),
640 replay_request_fingerprint: input.request.request_fingerprint.clone(),
641 input_data_identities,
642 bundle_id: input.source.execution_bundle.bundle_id.clone(),
643 plan_id: input.source.effective_plan.id.clone(),
644 phase: input.request.phase,
645 result_count: lineage.len(),
646 lineage_record_count: lineage.len(),
647 prediction_block_count: outputs.iter().map(|output| output.predictions.len()).sum(),
648 observation_prediction_block_count: outputs
649 .iter()
650 .map(|output| output.observation_predictions.len())
651 .sum(),
652 aggregated_prediction_block_count: outputs
653 .iter()
654 .map(|output| output.aggregated_predictions.len())
655 .sum(),
656 explanation_block_count: explanations.len(),
657 controller_count: lineage
658 .iter()
659 .map(|record| record.controller_id.as_str())
660 .collect::<BTreeSet<_>>()
661 .len(),
662 prediction_cache_store: false,
663 outputs,
664 explanations,
665 lineage,
666 warnings: input.warnings,
667 diagnostics: input.diagnostics,
668 outcome_fingerprint: zero_fingerprint(),
669 };
670 outcome.outcome_fingerprint = outcome.compute_fingerprint()?;
671 outcome.validate_against(input.source, input.request)?;
672 Ok(outcome)
673}
674
675pub fn execute_loaded_predictor_replay(
676 input: LoadedPredictorReplayInput<'_>,
677) -> Result<TrainingReplayOutcome> {
678 let package = input.predictor.package();
679 package.validate()?;
680 input.request.validate()?;
681 validate_sorted_unique_text("training replay execution warnings", &input.warnings, false)?;
682 validate_diagnostics(&input.diagnostics)?;
683 validate_replay_phase(input.request.phase)?;
684 if input.request.source_outcome_fingerprint != package.training_outcome.outcome_fingerprint {
685 return contract_error("training replay request does not target package source outcome");
686 }
687 for node_plan in package.effective_plan.node_plans.values() {
688 if input.controllers.get(&node_plan.controller_id).is_none() {
689 return Err(DagMlError::RuntimeValidation(format!(
690 "loaded predictor replay controller `{}` for node `{}` is not registered",
691 node_plan.controller_id, node_plan.node_id
692 )));
693 }
694 }
695
696 let input_data_identities = replay_input_data_identities(
697 &package.execution_bundle,
698 input.request,
699 input.data_envelopes,
700 )?;
701 let (replay_plan, replay_bundle) = replay_plan_and_bundle_for_current_cohort(
702 &package.effective_plan,
703 &package.execution_bundle,
704 input.request,
705 input.data_envelopes,
706 )?;
707 let phase_request = ReplayPhaseRequest {
708 bundle_id: replay_bundle.bundle_id.clone(),
709 phase: input.request.phase,
710 data_envelope_keys: input.request.data_envelope_keys.clone(),
711 };
712 let artifact_store = LoadedPredictorArtifactStore::new(input.predictor)?;
713 let mut ctx = RunContext::new(input.run_id.clone(), None);
714 let results = SequentialScheduler.execute_bundle_replay(
715 BundleReplayExecution {
716 plan: &replay_plan,
717 bundle: &replay_bundle,
718 replay_request: &phase_request,
719 prediction_cache_store: None,
720 controllers: input.controllers,
721 data_provider: input.data_provider,
722 artifact_store: &artifact_store,
723 data_envelopes: input.data_envelopes,
724 },
725 &mut ctx,
726 )?;
727 if results
728 .iter()
729 .any(|result| !result.artifacts.is_empty() || !result.artifact_handles.is_empty())
730 {
731 return contract_error("loaded predictor replay PREDICT/EXPLAIN cannot emit artifacts");
732 }
733
734 let outputs = bind_package_replay_outputs(package, input.request, &results)?;
735 let explanations = bind_attached_replay_explanations(input.request, &results)?;
736 let mut lineage = ctx.lineage.records().cloned().collect::<Vec<_>>();
737 for record in &mut lineage {
738 record.input_lineage.sort();
739 record
740 .artifact_refs
741 .sort_by(|left, right| left.id.cmp(&right.id));
742 }
743 lineage.sort_by(|left, right| left.record_id.cmp(&right.record_id));
744
745 let mut outcome = TrainingReplayOutcome {
746 schema_version: TRAINING_REPLAY_OUTCOME_SCHEMA_VERSION,
747 outcome_id: input.outcome_id,
748 run_id: input.run_id,
749 source_training_outcome: package.training_outcome.clone(),
750 replay_request_id: input.request.request_id.clone(),
751 replay_request_fingerprint: input.request.request_fingerprint.clone(),
752 input_data_identities,
753 bundle_id: package.execution_bundle.bundle_id.clone(),
754 plan_id: package.effective_plan.id.clone(),
755 phase: input.request.phase,
756 result_count: lineage.len(),
757 lineage_record_count: lineage.len(),
758 prediction_block_count: outputs.iter().map(|output| output.predictions.len()).sum(),
759 observation_prediction_block_count: outputs
760 .iter()
761 .map(|output| output.observation_predictions.len())
762 .sum(),
763 aggregated_prediction_block_count: outputs
764 .iter()
765 .map(|output| output.aggregated_predictions.len())
766 .sum(),
767 explanation_block_count: explanations.len(),
768 controller_count: lineage
769 .iter()
770 .map(|record| record.controller_id.as_str())
771 .collect::<BTreeSet<_>>()
772 .len(),
773 prediction_cache_store: false,
774 outputs,
775 explanations,
776 lineage,
777 warnings: input.warnings,
778 diagnostics: input.diagnostics,
779 outcome_fingerprint: zero_fingerprint(),
780 };
781 outcome.outcome_fingerprint = outcome.compute_fingerprint()?;
782 outcome.validate_against_package(package, input.request)?;
783 Ok(outcome)
784}
785
786fn replay_input_data_identities(
787 bundle: &ExecutionBundle,
788 request: &TrainingReplayRequest,
789 envelopes: &BTreeMap<String, ExternalDataPlanEnvelope>,
790) -> Result<Vec<TrainingDataIdentity>> {
791 request
792 .data_envelope_keys
793 .iter()
794 .map(|key| {
795 let requirement = bundle
796 .data_requirements
797 .iter()
798 .find(|requirement| requirement.key() == *key)
799 .ok_or_else(|| {
800 DagMlError::RuntimeValidation(format!(
801 "training replay request references unknown data envelope key `{key}`"
802 ))
803 })?;
804 let envelope = envelopes.get(key).ok_or_else(|| {
805 DagMlError::RuntimeValidation(format!(
806 "training replay is missing external data envelope for `{key}`"
807 ))
808 })?;
809 envelope.validate()?;
810 if requirement.schema_fingerprint != envelope.schema_fingerprint
811 || requirement.plan_fingerprint != envelope.plan_fingerprint
812 {
813 return Err(DagMlError::RuntimeValidation(format!(
814 "training replay envelope for `{key}` changes schema or representation plan"
815 )));
816 }
817 let relation_fingerprint = envelope.relation_fingerprint.clone().ok_or_else(|| {
818 DagMlError::RuntimeValidation(format!(
819 "training replay envelope for `{key}` requires a relation fingerprint"
820 ))
821 })?;
822 let data_content_fingerprint =
823 envelope.data_content_fingerprint.clone().ok_or_else(|| {
824 DagMlError::RuntimeValidation(format!(
825 "training replay envelope for `{key}` requires a data content fingerprint"
826 ))
827 })?;
828 let target_content_fingerprint =
829 envelope.target_content_fingerprint.clone().ok_or_else(|| {
830 DagMlError::RuntimeValidation(format!(
831 "training replay envelope for `{key}` requires a target content fingerprint"
832 ))
833 })?;
834 let mut identity = TrainingDataIdentity {
835 requirement_key: key.clone(),
836 schema_fingerprint: envelope.schema_fingerprint.clone(),
837 plan_fingerprint: envelope.plan_fingerprint.clone(),
838 relation_fingerprint,
839 data_content_fingerprint,
840 target_content_fingerprint,
841 identity_fingerprint: zero_fingerprint(),
842 };
843 identity.identity_fingerprint = identity.compute_fingerprint()?;
844 identity.validate()?;
845 Ok(identity)
846 })
847 .collect()
848}
849
850fn replay_plan_and_bundle_for_current_cohort(
851 plan: &ExecutionPlan,
852 bundle: &ExecutionBundle,
853 request: &TrainingReplayRequest,
854 envelopes: &BTreeMap<String, ExternalDataPlanEnvelope>,
855) -> Result<(ExecutionPlan, ExecutionBundle)> {
856 let mut replay_plan = plan.clone();
857 let mut replay_bundle = bundle.clone();
858 for requirement in &mut replay_bundle.data_requirements {
859 let key = requirement.key();
860 if request.data_envelope_keys.contains(&key) {
861 let envelope = envelopes.get(&key).ok_or_else(|| {
862 DagMlError::RuntimeValidation(format!(
863 "training replay is missing external data envelope for `{key}`"
864 ))
865 })?;
866 requirement.relation_fingerprint = envelope.relation_fingerprint.clone();
867 for bindings in replay_plan.campaign.data_bindings.values_mut() {
868 for binding in bindings {
869 if crate::data::data_binding_requirement_key(
870 &binding.node_id,
871 &binding.input_name,
872 ) == key
873 {
874 binding.relation_fingerprint = envelope.relation_fingerprint.clone();
875 }
876 }
877 }
878 for node_plan in replay_plan.node_plans.values_mut() {
879 for binding in &mut node_plan.data_bindings {
880 if crate::data::data_binding_requirement_key(
881 &binding.node_id,
882 &binding.input_name,
883 ) == key
884 {
885 binding.relation_fingerprint = envelope.relation_fingerprint.clone();
886 }
887 }
888 }
889 }
890 }
891 replay_plan.campaign_fingerprint = stable_json_fingerprint(&replay_plan.campaign)?;
892 replay_bundle.campaign_fingerprint = replay_plan.campaign_fingerprint.clone();
893 Ok((replay_plan, replay_bundle))
894}
895
896fn bind_attached_replay_outputs(
897 source: &TrainingOutcome,
898 request: &TrainingReplayRequest,
899 results: &[crate::runtime::NodeResult],
900) -> Result<Vec<BoundTrainingOutput>> {
901 let mut outputs = Vec::new();
902 for binding_id in &request.output_binding_ids {
903 let source_output = source
904 .outputs
905 .iter()
906 .find(|output| output.binding.binding_id == *binding_id)
907 .ok_or_else(|| {
908 DagMlError::RuntimeValidation(format!(
909 "training replay request references absent binding `{binding_id}`"
910 ))
911 })?;
912 let binding = source_output.binding.clone();
913 let mut output = BoundTrainingOutput {
914 schema_version: Some(BOUND_TRAINING_OUTPUT_SCHEMA_VERSION),
915 binding: binding.clone(),
916 predictions: Vec::new(),
917 observation_predictions: Vec::new(),
918 aggregated_predictions: Vec::new(),
919 };
920 for result in results {
921 output.predictions.extend(
922 result
923 .predictions
924 .iter()
925 .filter(|block| {
926 block.producer_node == binding.node_id
927 && block.producer_port.as_deref() == Some(binding.port_name.as_str())
928 && block.partition == crate::oof::PredictionPartition::Final
929 && block.fold_id.is_none()
930 })
931 .cloned(),
932 );
933 output.observation_predictions.extend(
934 result
935 .observation_predictions
936 .iter()
937 .filter(|block| {
938 block.producer_node == binding.node_id
939 && block.producer_port.as_deref() == Some(binding.port_name.as_str())
940 && block.partition == crate::oof::PredictionPartition::Final
941 && block.fold_id.is_none()
942 })
943 .cloned(),
944 );
945 output.aggregated_predictions.extend(
946 result
947 .aggregated_predictions
948 .iter()
949 .filter(|block| {
950 block.producer_node == binding.node_id
951 && block.producer_port.as_deref() == Some(binding.port_name.as_str())
952 && block.partition == crate::oof::PredictionPartition::Final
953 && block.fold_id.is_none()
954 })
955 .cloned(),
956 );
957 }
958 if !output.predictions.is_empty()
959 || !output.observation_predictions.is_empty()
960 || !output.aggregated_predictions.is_empty()
961 {
962 output.validate(&source.effective_plan)?;
963 outputs.push(output);
964 }
965 }
966 outputs.sort_by(|left, right| left.binding.binding_id.cmp(&right.binding.binding_id));
967 Ok(outputs)
968}
969
970fn bind_package_replay_outputs(
971 package: &PortablePredictorPackage,
972 request: &TrainingReplayRequest,
973 results: &[crate::runtime::NodeResult],
974) -> Result<Vec<BoundTrainingOutput>> {
975 let mut outputs = Vec::new();
976 for binding_id in &request.output_binding_ids {
977 let binding = package
978 .output_bindings
979 .iter()
980 .find(|binding| binding.binding_id == *binding_id)
981 .ok_or_else(|| {
982 DagMlError::RuntimeValidation(format!(
983 "training replay request references absent package binding `{binding_id}`"
984 ))
985 })?
986 .clone();
987 let mut output = BoundTrainingOutput {
988 schema_version: Some(BOUND_TRAINING_OUTPUT_SCHEMA_VERSION),
989 binding: binding.clone(),
990 predictions: Vec::new(),
991 observation_predictions: Vec::new(),
992 aggregated_predictions: Vec::new(),
993 };
994 for result in results {
995 output.predictions.extend(
996 result
997 .predictions
998 .iter()
999 .filter(|block| {
1000 block.producer_node == binding.node_id
1001 && block.producer_port.as_deref() == Some(binding.port_name.as_str())
1002 && block.partition == crate::oof::PredictionPartition::Final
1003 && block.fold_id.is_none()
1004 })
1005 .cloned(),
1006 );
1007 output.observation_predictions.extend(
1008 result
1009 .observation_predictions
1010 .iter()
1011 .filter(|block| {
1012 block.producer_node == binding.node_id
1013 && block.producer_port.as_deref() == Some(binding.port_name.as_str())
1014 && block.partition == crate::oof::PredictionPartition::Final
1015 && block.fold_id.is_none()
1016 })
1017 .cloned(),
1018 );
1019 output.aggregated_predictions.extend(
1020 result
1021 .aggregated_predictions
1022 .iter()
1023 .filter(|block| {
1024 block.producer_node == binding.node_id
1025 && block.producer_port.as_deref() == Some(binding.port_name.as_str())
1026 && block.partition == crate::oof::PredictionPartition::Final
1027 && block.fold_id.is_none()
1028 })
1029 .cloned(),
1030 );
1031 }
1032 if !output.predictions.is_empty()
1033 || !output.observation_predictions.is_empty()
1034 || !output.aggregated_predictions.is_empty()
1035 {
1036 output.validate(&package.effective_plan)?;
1037 outputs.push(output);
1038 }
1039 }
1040 outputs.sort_by(|left, right| left.binding.binding_id.cmp(&right.binding.binding_id));
1041 Ok(outputs)
1042}
1043
1044fn bind_attached_replay_explanations(
1045 request: &TrainingReplayRequest,
1046 results: &[crate::runtime::NodeResult],
1047) -> Result<Vec<ExplanationBlock>> {
1048 if request.phase != Phase::Explain {
1049 return Ok(Vec::new());
1050 }
1051 let mut explanations = results
1052 .iter()
1053 .flat_map(|result| result.explanations.iter().cloned())
1054 .filter(|block| block.producer_port.is_some())
1055 .collect::<Vec<_>>();
1056 explanations.sort_by(|left, right| {
1057 (
1058 left.producer_node.as_str(),
1059 left.producer_port.as_deref().unwrap_or_default(),
1060 left.method.as_str(),
1061 left.target_name.as_deref().unwrap_or_default(),
1062 )
1063 .cmp(&(
1064 right.producer_node.as_str(),
1065 right.producer_port.as_deref().unwrap_or_default(),
1066 right.method.as_str(),
1067 right.target_name.as_deref().unwrap_or_default(),
1068 ))
1069 });
1070 Ok(explanations)
1071}
1072
1073fn validate_output_order_and_version(outputs: &[BoundTrainingOutput]) -> Result<()> {
1074 let mut previous: Option<&str> = None;
1075 for output in outputs {
1076 match output.schema_version {
1077 Some(BOUND_TRAINING_OUTPUT_SCHEMA_VERSION) => {}
1078 Some(version) => {
1079 return contract_error(format!(
1080 "training replay output schema_version {version} is unsupported; current {BOUND_TRAINING_OUTPUT_SCHEMA_VERSION}"
1081 ));
1082 }
1083 None => {
1084 return contract_error(
1085 "training replay output requires bound_training_output schema_version",
1086 );
1087 }
1088 }
1089 let binding_id = output.binding.binding_id.as_str();
1090 if previous.is_some_and(|previous| previous >= binding_id) {
1091 return contract_error("training replay outputs must be strictly sorted by binding_id");
1092 }
1093 previous = Some(binding_id);
1094 }
1095 Ok(())
1096}
1097
1098fn validate_replay_bound_output_blocks(output: &BoundTrainingOutput) -> Result<()> {
1099 for block in &output.predictions {
1100 validate_optional_port(
1101 "training replay prediction producer_port",
1102 &block.producer_port,
1103 )?;
1104 if block.partition != crate::oof::PredictionPartition::Final || block.fold_id.is_some() {
1105 return contract_error(
1106 "training replay prediction blocks must use final partition without fold",
1107 );
1108 }
1109 }
1110 for block in &output.observation_predictions {
1111 validate_optional_port(
1112 "training replay observation prediction producer_port",
1113 &block.producer_port,
1114 )?;
1115 if block.partition != crate::oof::PredictionPartition::Final || block.fold_id.is_some() {
1116 return contract_error(
1117 "training replay observation prediction blocks must use final partition without fold",
1118 );
1119 }
1120 }
1121 for block in &output.aggregated_predictions {
1122 validate_optional_port(
1123 "training replay aggregated prediction producer_port",
1124 &block.producer_port,
1125 )?;
1126 if block.partition != crate::oof::PredictionPartition::Final || block.fold_id.is_some() {
1127 return contract_error(
1128 "training replay aggregated prediction blocks must use final partition without fold",
1129 );
1130 }
1131 }
1132 Ok(())
1133}
1134
1135fn validate_replay_phase(phase: Phase) -> Result<()> {
1136 if matches!(phase, Phase::Predict | Phase::Explain) {
1137 Ok(())
1138 } else {
1139 contract_error("training replay V1 supports only PREDICT and EXPLAIN")
1140 }
1141}
1142
1143fn validate_sha256(label: &str, value: &str) -> Result<()> {
1144 if value.len() == 64
1145 && value
1146 .bytes()
1147 .all(|byte| byte.is_ascii_digit() || (b'a'..=b'f').contains(&byte))
1148 {
1149 Ok(())
1150 } else {
1151 contract_error(format!(
1152 "{label} fingerprint must be 64 lowercase hexadecimal characters"
1153 ))
1154 }
1155}
1156
1157fn validate_identifier(label: &str, value: &str) -> Result<()> {
1158 if !value.is_empty()
1159 && value.len() <= 128
1160 && value
1161 .bytes()
1162 .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'_' | b'-' | b'.' | b':'))
1163 {
1164 Ok(())
1165 } else {
1166 contract_error(format!("{label} is not a valid DAG-ML identifier"))
1167 }
1168}
1169
1170fn validate_non_empty(label: &str, value: &str) -> Result<()> {
1171 if value.trim().is_empty() {
1172 contract_error(format!("{label} must be non-empty"))
1173 } else {
1174 Ok(())
1175 }
1176}
1177
1178fn validate_sorted_unique_identifiers(
1179 label: &str,
1180 values: &[String],
1181 require_non_empty: bool,
1182) -> Result<()> {
1183 validate_sorted_unique_text(label, values, require_non_empty)?;
1184 for value in values {
1185 validate_identifier(label, value)?;
1186 }
1187 Ok(())
1188}
1189
1190fn validate_sorted_unique_text(
1191 label: &str,
1192 values: &[String],
1193 require_non_empty: bool,
1194) -> Result<()> {
1195 if require_non_empty && values.is_empty() {
1196 return contract_error(format!("{label} must be non-empty"));
1197 }
1198 let mut previous: Option<&str> = None;
1199 for value in values {
1200 validate_non_empty(label, value)?;
1201 if previous.is_some_and(|previous| previous >= value.as_str()) {
1202 return contract_error(format!("{label} must be strictly sorted and unique"));
1203 }
1204 previous = Some(value.as_str());
1205 }
1206 Ok(())
1207}
1208
1209fn validate_sorted_unique_keys<'a>(
1210 label: &str,
1211 values: impl Iterator<Item = &'a str>,
1212 require_non_empty: bool,
1213) -> Result<()> {
1214 let values = values.collect::<Vec<_>>();
1215 if require_non_empty && values.is_empty() {
1216 return contract_error(format!("{label} must be non-empty"));
1217 }
1218 let mut previous: Option<&str> = None;
1219 for value in values {
1220 validate_non_empty(label, value)?;
1221 if previous.is_some_and(|previous| previous >= value) {
1222 return contract_error(format!("{label} must be strictly sorted and unique"));
1223 }
1224 previous = Some(value);
1225 }
1226 Ok(())
1227}
1228
1229fn validate_optional_port(label: &str, value: &Option<String>) -> Result<()> {
1230 match value {
1231 Some(value) if !value.trim().is_empty() => Ok(()),
1232 _ => contract_error(format!("{label} must be present and non-empty")),
1233 }
1234}
1235
1236fn validate_diagnostics(diagnostics: &BTreeMap<String, serde_json::Value>) -> Result<()> {
1237 for (key, value) in diagnostics {
1238 validate_non_empty("training replay diagnostic key", key)?;
1239 if !matches!(
1240 value,
1241 serde_json::Value::Null
1242 | serde_json::Value::Bool(_)
1243 | serde_json::Value::Number(_)
1244 | serde_json::Value::String(_)
1245 ) {
1246 return contract_error("training replay diagnostics must be scalar JSON values");
1247 }
1248 }
1249 Ok(())
1250}
1251
1252fn require_count(label: &str, actual: usize, expected: usize) -> Result<()> {
1253 if actual == expected {
1254 Ok(())
1255 } else {
1256 contract_error(format!("{label} does not match replay payload"))
1257 }
1258}
1259
1260fn zero_fingerprint() -> String {
1261 "0".repeat(64)
1262}
1263
1264fn tcv1_fingerprint_without<T: Serialize>(value: &T, field: &str, label: &str) -> Result<String> {
1265 let json = serde_json::to_string(value)?;
1266 strict_tcv1_fingerprint_without(&json, field, label)
1267}
1268
1269fn strict_tcv1_fingerprint_without(json: &str, field: &str, label: &str) -> Result<String> {
1270 parse_typed_json(json)
1271 .and_then(|value| value.fingerprint_without(field))
1272 .map_err(|error| {
1273 DagMlError::RuntimeValidation(format!("{label} is outside strict TCV1: {error}"))
1274 })
1275}
1276
1277fn unsupported_version<T>(label: &str, actual: u32, expected: u32) -> Result<T> {
1278 contract_error(format!(
1279 "{label} uses unsupported schema_version {actual}, expected {expected}"
1280 ))
1281}
1282
1283fn contract_error<T>(message: impl Into<String>) -> Result<T> {
1284 Err(DagMlError::CampaignValidation(message.into()))
1285}
1286
1287#[cfg(test)]
1288mod tests {
1289 use std::fs;
1290 use std::path::PathBuf;
1291
1292 use super::*;
1293
1294 fn root() -> PathBuf {
1295 PathBuf::from(env!("CARGO_MANIFEST_DIR"))
1296 .parent()
1297 .and_then(|path| path.parent())
1298 .expect("core crate is under crates/dag-ml-core")
1299 .to_path_buf()
1300 }
1301
1302 fn fixture(name: &str) -> String {
1303 fs::read_to_string(
1304 root()
1305 .join("examples")
1306 .join("fixtures")
1307 .join("training")
1308 .join("replay")
1309 .join(name),
1310 )
1311 .expect(name)
1312 }
1313
1314 fn training_fixture(name: &str) -> String {
1315 fs::read_to_string(
1316 root()
1317 .join("examples")
1318 .join("fixtures")
1319 .join("training")
1320 .join(name),
1321 )
1322 .expect(name)
1323 }
1324
1325 #[test]
1326 fn training_replay_contract_fixtures_parse_and_cross_validate() {
1327 let predict_source =
1328 TrainingOutcome::from_json(&training_fixture("training_outcome_refit.v1.json"))
1329 .expect("predict source training outcome");
1330 let explain_source =
1331 TrainingOutcome::from_json(&fixture("training_replay_source_outcome_explain.v1.json"))
1332 .expect("explain source training outcome");
1333 let predict_request =
1334 TrainingReplayRequest::from_json(&fixture("training_replay_request_predict.v1.json"))
1335 .expect("predict request");
1336 let predict_outcome =
1337 TrainingReplayOutcome::from_json(&fixture("training_replay_outcome_predict.v1.json"))
1338 .expect("predict outcome");
1339 predict_outcome
1340 .validate_against(&predict_source, &predict_request)
1341 .expect("predict cross-links");
1342
1343 let explain_request =
1344 TrainingReplayRequest::from_json(&fixture("training_replay_request_explain.v1.json"))
1345 .expect("explain request");
1346 let explain_outcome =
1347 TrainingReplayOutcome::from_json(&fixture("training_replay_outcome_explain.v1.json"))
1348 .expect("explain outcome");
1349 explain_outcome
1350 .validate_against(&explain_source, &explain_request)
1351 .expect("explain cross-links");
1352
1353 let explain_only = TrainingReplayOutcome::from_json(&fixture(
1354 "training_replay_outcome_explain_only.v1.json",
1355 ))
1356 .expect("explain-only outcome");
1357 explain_only
1358 .validate_against(&explain_source, &explain_request)
1359 .expect("explain-only cross-links");
1360 }
1361
1362 #[test]
1363 fn training_replay_request_rejects_refit_and_unsorted_bindings() {
1364 let mut request: serde_json::Value =
1365 serde_json::from_str(&fixture("training_replay_request_predict.v1.json")).unwrap();
1366 request["phase"] = serde_json::Value::String("REFIT".to_string());
1367 let err = serde_json::from_value::<TrainingReplayRequest>(request)
1368 .unwrap()
1369 .validate()
1370 .unwrap_err()
1371 .to_string();
1372 assert!(err.contains("PREDICT and EXPLAIN"));
1373
1374 let mut request: TrainingReplayRequest =
1375 TrainingReplayRequest::from_json(&fixture("training_replay_request_predict.v1.json"))
1376 .unwrap();
1377 request.output_binding_ids = vec!["z".to_string(), "a".to_string()];
1378 request.request_fingerprint = request.compute_fingerprint().unwrap();
1379 let err = request.validate().unwrap_err().to_string();
1380 assert!(err.contains("strictly sorted"));
1381 }
1382
1383 #[test]
1384 fn training_replay_outcome_rejects_counter_and_source_transplants() {
1385 let source =
1386 TrainingOutcome::from_json(&training_fixture("training_outcome_refit.v1.json"))
1387 .unwrap();
1388 let request =
1389 TrainingReplayRequest::from_json(&fixture("training_replay_request_predict.v1.json"))
1390 .unwrap();
1391 let mut outcome =
1392 TrainingReplayOutcome::from_json(&fixture("training_replay_outcome_predict.v1.json"))
1393 .unwrap();
1394 outcome.prediction_block_count += 1;
1395 outcome.outcome_fingerprint = outcome.compute_fingerprint().unwrap();
1396 let err = outcome.validate().unwrap_err().to_string();
1397 assert!(err.contains("prediction_block_count"));
1398
1399 let mut outcome =
1400 TrainingReplayOutcome::from_json(&fixture("training_replay_outcome_predict.v1.json"))
1401 .unwrap();
1402 outcome.source_training_outcome.outcome_fingerprint = "f".repeat(64);
1403 outcome.outcome_fingerprint = outcome.compute_fingerprint().unwrap();
1404 let err = outcome
1405 .validate_against(&source, &request)
1406 .unwrap_err()
1407 .to_string();
1408 assert!(err.contains("source reference"));
1409 }
1410}