1use super::*;
3
4#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
5pub struct DataMaterializationRequest {
6 pub run_id: RunId,
7 pub node_id: NodeId,
8 pub input_name: String,
9 pub phase: Phase,
10 pub variant_id: Option<VariantId>,
11 pub fold_id: Option<FoldId>,
12 pub binding: crate::data::DataBinding,
13}
14
15#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
16pub struct DataProviderViewSpec {
17 #[serde(default)]
18 pub sample_ids: Option<Vec<SampleId>>,
19 pub partition: DataRequestPartition,
20 #[serde(default)]
21 pub fold_id: Option<FoldId>,
22 #[serde(default)]
23 pub source_ids: Option<Vec<String>>,
24 #[serde(default)]
25 pub columns: Option<Vec<String>>,
26 pub include_augmented: bool,
27 pub include_excluded: bool,
28 #[serde(default, skip_serializing_if = "Option::is_none")]
29 pub branch_view: Option<crate::data::BranchViewPlan>,
30 #[serde(default)]
31 pub extra: BTreeMap<String, serde_json::Value>,
32}
33
34pub const DATA_OUTPUT_PROVENANCE_KEY: &str = "dag_ml_output";
35pub const DATA_OUTPUT_PROVENANCE_SCHEMA_VERSION: u32 = 1;
36pub const DATA_OUTPUT_PROVENANCE_SCHEMA_ID: &str =
37 "https://github.com/GBeurier/dag-ml/schemas/data_output_provenance.v1.schema.json";
38pub const NODE_TASK_SCHEMA_VERSION: u32 = 1;
39pub const NODE_TASK_SCHEMA_ID: &str =
40 "https://github.com/GBeurier/dag-ml/schemas/node_task.v1.schema.json";
41pub const NODE_RESULT_SCHEMA_VERSION: u32 = 1;
42pub const NODE_RESULT_SCHEMA_ID: &str =
43 "https://github.com/GBeurier/dag-ml/schemas/node_result.v1.schema.json";
44
45pub(crate) fn default_data_output_provenance_schema_version() -> u32 {
46 DATA_OUTPUT_PROVENANCE_SCHEMA_VERSION
47}
48
49impl DataProviderViewSpec {
50 pub fn validate(&self) -> Result<()> {
51 validate_optional_ids("sample id", &self.sample_ids)?;
52 validate_optional_strings("source id", &self.source_ids)?;
53 validate_optional_strings("column", &self.columns)?;
54 match self.partition {
55 DataRequestPartition::FoldTrain | DataRequestPartition::FoldValidation => {
56 if self.sample_ids.is_some() && self.fold_id.is_none() {
57 return Err(DagMlError::RuntimeValidation(format!(
58 "data provider view {:?} with explicit sample ids requires a fold id",
59 self.partition
60 )));
61 }
62 }
63 DataRequestPartition::FullTrain | DataRequestPartition::Predict => {
64 if self.fold_id.is_some() {
65 return Err(DagMlError::RuntimeValidation(format!(
66 "data provider view {:?} must not carry a fold id",
67 self.partition
68 )));
69 }
70 }
71 }
72 for key in self.extra.keys() {
73 if key.trim().is_empty() {
74 return Err(DagMlError::RuntimeValidation(
75 "data provider view extra contains an empty key".to_string(),
76 ));
77 }
78 }
79 if let Some(branch_view) = &self.branch_view {
80 branch_view.validate()?;
81 }
82 self.output_provenance()?;
83 Ok(())
84 }
85
86 pub fn output_provenance(&self) -> Result<Option<DataOutputProvenance>> {
87 let Some(value) = self.extra.get(DATA_OUTPUT_PROVENANCE_KEY) else {
88 return Ok(None);
89 };
90 let provenance: DataOutputProvenance = serde_json::from_value(value.clone())?;
91 provenance.validate()?;
92 Ok(Some(provenance))
93 }
94}
95
96#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
97pub struct DataOutputProvenance {
98 #[serde(default = "default_data_output_provenance_schema_version")]
99 pub schema_version: u32,
100 pub producer_node: NodeId,
101 pub producer_port: String,
102 pub producer_phase: Phase,
103 #[serde(default)]
104 pub variant_id: Option<VariantId>,
105 #[serde(default)]
106 pub fold_id: Option<FoldId>,
107 #[serde(default)]
108 pub shape_plan_fingerprint: Option<String>,
109 #[serde(default)]
110 pub aggregation_policy_fingerprint: Option<String>,
111 #[serde(default)]
112 pub feature_namespace: Option<String>,
113 #[serde(default)]
114 pub feature_schema_fingerprint: Option<String>,
115 #[serde(default, skip_serializing_if = "Option::is_none")]
116 pub representation_plan: Option<RepresentationPlan>,
117 #[serde(default, skip_serializing_if = "Option::is_none")]
118 pub representation_replay_manifest: Option<RepresentationReplayManifest>,
119 #[serde(default, skip_serializing_if = "Option::is_none")]
120 pub representation_compatibility: Option<RepresentationCompatibilityReport>,
121 #[serde(default, skip_serializing_if = "Option::is_none")]
122 pub relation_delta_fingerprint: Option<String>,
123 #[serde(default)]
124 pub shape_deltas: Vec<ShapeDelta>,
125}
126
127impl DataOutputProvenance {
128 pub fn validate(&self) -> Result<()> {
129 if self.schema_version != DATA_OUTPUT_PROVENANCE_SCHEMA_VERSION {
130 return Err(DagMlError::RuntimeValidation(format!(
131 "data output provenance for `{}` uses unsupported schema_version {}, expected {}",
132 self.producer_node, self.schema_version, DATA_OUTPUT_PROVENANCE_SCHEMA_VERSION
133 )));
134 }
135 if self.producer_port.trim().is_empty() {
136 return Err(DagMlError::RuntimeValidation(format!(
137 "data output provenance for `{}` has empty producer_port",
138 self.producer_node
139 )));
140 }
141 validate_optional_fingerprint(
142 "shape_plan_fingerprint",
143 &self.shape_plan_fingerprint,
144 &self.producer_node,
145 )?;
146 validate_optional_fingerprint(
147 "aggregation_policy_fingerprint",
148 &self.aggregation_policy_fingerprint,
149 &self.producer_node,
150 )?;
151 validate_optional_fingerprint(
152 "feature_schema_fingerprint",
153 &self.feature_schema_fingerprint,
154 &self.producer_node,
155 )?;
156 validate_optional_fingerprint(
157 "relation_delta_fingerprint",
158 &self.relation_delta_fingerprint,
159 &self.producer_node,
160 )?;
161 if let Some(representation_plan) = &self.representation_plan {
162 representation_plan.validate().map_err(|error| {
163 DagMlError::RuntimeValidation(format!(
164 "data output provenance for `{}` has invalid representation_plan: {error}",
165 self.producer_node
166 ))
167 })?;
168 }
169 if let Some(replay_manifest) = &self.representation_replay_manifest {
170 replay_manifest.validate().map_err(|error| {
171 DagMlError::RuntimeValidation(format!(
172 "data output provenance for `{}` has invalid representation_replay_manifest: {error}",
173 self.producer_node
174 ))
175 })?;
176 }
177 if let Some(report) = &self.representation_compatibility {
178 report.validate().map_err(|error| {
179 DagMlError::RuntimeValidation(format!(
180 "data output provenance for `{}` has invalid representation_compatibility: {error}",
181 self.producer_node
182 ))
183 })?;
184 }
185 if self
186 .feature_namespace
187 .as_ref()
188 .is_some_and(|namespace| namespace.trim().is_empty())
189 {
190 return Err(DagMlError::RuntimeValidation(format!(
191 "data output provenance for `{}` has empty feature_namespace",
192 self.producer_node
193 )));
194 }
195 for delta in &self.shape_deltas {
196 delta.validate()?;
197 if delta.node_id != self.producer_node {
198 return Err(DagMlError::RuntimeValidation(format!(
199 "data output provenance for `{}` contains shape delta for `{}`",
200 self.producer_node, delta.node_id
201 )));
202 }
203 }
204 if let Some(feature_schema_fingerprint) = &self.feature_schema_fingerprint {
205 if let Some(last_feature_delta) = self
206 .shape_deltas
207 .iter()
208 .rev()
209 .find(|delta| delta.kind == ShapeDeltaKind::Feature)
210 {
211 if &last_feature_delta.after_fingerprint != feature_schema_fingerprint {
212 return Err(DagMlError::RuntimeValidation(format!(
213 "data output provenance for `{}` has feature_schema_fingerprint `{feature_schema_fingerprint}` but last feature delta ends at `{}`",
214 self.producer_node, last_feature_delta.after_fingerprint
215 )));
216 }
217 }
218 }
219 Ok(())
220 }
221}
222
223pub(crate) fn validate_optional_fingerprint(
224 label: &str,
225 fingerprint: &Option<String>,
226 producer_node: &NodeId,
227) -> Result<()> {
228 let Some(fingerprint) = fingerprint else {
229 return Ok(());
230 };
231 if fingerprint.len() != 64 || !fingerprint.bytes().all(|byte| byte.is_ascii_hexdigit()) {
232 return Err(DagMlError::RuntimeValidation(format!(
233 "data output provenance for `{producer_node}` has invalid {label}"
234 )));
235 }
236 Ok(())
237}
238
239pub(crate) fn validate_optional_ids<T>(label: &str, values: &Option<Vec<T>>) -> Result<()>
240where
241 T: Ord + ToString,
242{
243 let Some(values) = values else {
244 return Ok(());
245 };
246 if values.is_empty() {
247 return Err(DagMlError::RuntimeValidation(format!(
248 "data provider view {label} list is empty"
249 )));
250 }
251 let mut seen = BTreeSet::new();
252 for value in values {
253 if !seen.insert(value) {
254 return Err(DagMlError::RuntimeValidation(format!(
255 "data provider view has duplicate {label} `{}`",
256 value.to_string()
257 )));
258 }
259 }
260 Ok(())
261}
262
263pub(crate) fn validate_optional_strings(label: &str, values: &Option<Vec<String>>) -> Result<()> {
264 let Some(values) = values else {
265 return Ok(());
266 };
267 if values.is_empty() {
268 return Err(DagMlError::RuntimeValidation(format!(
269 "data provider view {label} list is empty"
270 )));
271 }
272 let mut seen = BTreeSet::new();
273 for value in values {
274 if value.trim().is_empty() {
275 return Err(DagMlError::RuntimeValidation(format!(
276 "data provider view contains an empty {label}"
277 )));
278 }
279 if !seen.insert(value.as_str()) {
280 return Err(DagMlError::RuntimeValidation(format!(
281 "data provider view has duplicate {label} `{value}`"
282 )));
283 }
284 }
285 Ok(())
286}
287
288#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
289pub struct DataViewRequest {
290 pub run_id: RunId,
291 pub node_id: NodeId,
292 pub input_name: String,
293 pub phase: Phase,
294 pub variant_id: Option<VariantId>,
295 pub fold_id: Option<FoldId>,
296 pub binding: crate::data::DataBinding,
297 pub data_handle: HandleRef,
298 pub view: DataProviderViewSpec,
299}
300
301pub trait RuntimeDataProvider {
302 fn materialize(&self, request: &DataMaterializationRequest) -> Result<HandleRef>;
303 fn make_view(&self, request: &DataViewRequest) -> Result<HandleRef>;
304 fn training_data_identity(
310 &self,
311 _binding: &DataBinding,
312 ) -> Result<Option<crate::training::TrainingDataIdentity>> {
313 Ok(None)
314 }
315 fn coordinator_relations(&self, _binding: &DataBinding) -> Result<Option<SampleRelationSet>> {
316 Ok(None)
317 }
318}
319
320#[derive(Debug)]
321struct EnvelopeAttestation {
322 binding: DataBinding,
323 envelope: ExternalDataPlanEnvelope,
324 identity: crate::training::TrainingDataIdentity,
325}
326
327#[derive(Debug)]
335pub struct EnvelopeAttestedRuntimeDataProvider<P> {
336 inner: P,
337 attestations: BTreeMap<String, EnvelopeAttestation>,
338}
339
340impl<P> EnvelopeAttestedRuntimeDataProvider<P> {
341 pub fn new<I>(
342 inner: P,
343 bindings: I,
344 mut envelopes: BTreeMap<String, ExternalDataPlanEnvelope>,
345 ) -> Result<Self>
346 where
347 I: IntoIterator<Item = DataBinding>,
348 {
349 let mut bindings_by_key: BTreeMap<String, DataBinding> = BTreeMap::new();
350 for binding in bindings {
351 binding.validate()?;
352 let key = data_binding_requirement_key(&binding.node_id, &binding.input_name);
353 if let Some(previous) = bindings_by_key.get(&key) {
354 let detail = if previous.node_id == binding.node_id
355 && previous.input_name == binding.input_name
356 {
357 "duplicates the same coordinates"
358 } else {
359 "uses distinct coordinates that collide under the V1 node.input spelling"
360 };
361 return Err(DagMlError::RuntimeValidation(format!(
362 "data binding requirement key `{key}` {detail}"
363 )));
364 }
365 bindings_by_key.insert(key, binding);
366 }
367
368 let expected_keys = bindings_by_key.keys().cloned().collect::<BTreeSet<_>>();
369 let actual_keys = envelopes.keys().cloned().collect::<BTreeSet<_>>();
370 if expected_keys != actual_keys {
371 let missing = expected_keys
372 .difference(&actual_keys)
373 .cloned()
374 .collect::<Vec<_>>();
375 let unexpected = actual_keys
376 .difference(&expected_keys)
377 .cloned()
378 .collect::<Vec<_>>();
379 return Err(DagMlError::RuntimeValidation(format!(
380 "attested data envelopes must exactly cover runtime bindings (missing: [{}]; unexpected: [{}])",
381 missing.join(", "),
382 unexpected.join(", ")
383 )));
384 }
385
386 let mut attestations = BTreeMap::new();
387 for (key, binding) in bindings_by_key {
388 let envelope = envelopes
389 .remove(&key)
390 .expect("exact key coverage was checked above");
391 let identity =
392 crate::training::TrainingDataIdentity::from_binding_envelope(&binding, &envelope)?;
393 attestations.insert(
394 key,
395 EnvelopeAttestation {
396 binding,
397 envelope,
398 identity,
399 },
400 );
401 }
402
403 Ok(Self {
404 inner,
405 attestations,
406 })
407 }
408
409 pub fn inner(&self) -> &P {
410 &self.inner
411 }
412
413 pub fn into_inner(self) -> P {
414 self.inner
415 }
416
417 fn attestation_for_binding(&self, binding: &DataBinding) -> Result<&EnvelopeAttestation> {
418 binding.validate()?;
419 let key = data_binding_requirement_key(&binding.node_id, &binding.input_name);
420 let attestation = self.attestations.get(&key).ok_or_else(|| {
421 DagMlError::RuntimeValidation(format!(
422 "runtime data binding `{key}` has no registered envelope attestation"
423 ))
424 })?;
425 if attestation.binding != *binding {
426 return Err(DagMlError::RuntimeValidation(format!(
427 "runtime data binding `{key}` does not exactly match its attested binding"
428 )));
429 }
430 Ok(attestation)
431 }
432
433 fn validate_request_binding(
434 &self,
435 node_id: &NodeId,
436 input_name: &str,
437 binding: &DataBinding,
438 ) -> Result<()> {
439 if node_id != &binding.node_id || input_name != binding.input_name {
440 return Err(DagMlError::RuntimeValidation(format!(
441 "runtime data request coordinates `{node_id}.{input_name}` do not match binding `{}`",
442 data_binding_requirement_key(&binding.node_id, &binding.input_name)
443 )));
444 }
445 self.attestation_for_binding(binding)?;
446 Ok(())
447 }
448}
449
450impl<P: RuntimeDataProvider> RuntimeDataProvider for EnvelopeAttestedRuntimeDataProvider<P> {
451 fn materialize(&self, request: &DataMaterializationRequest) -> Result<HandleRef> {
452 self.validate_request_binding(&request.node_id, &request.input_name, &request.binding)?;
453 self.inner.materialize(request)
454 }
455
456 fn make_view(&self, request: &DataViewRequest) -> Result<HandleRef> {
457 request.view.validate()?;
458 self.validate_request_binding(&request.node_id, &request.input_name, &request.binding)?;
459 self.inner.make_view(request)
460 }
461
462 fn training_data_identity(
463 &self,
464 binding: &DataBinding,
465 ) -> Result<Option<crate::training::TrainingDataIdentity>> {
466 Ok(Some(
467 self.attestation_for_binding(binding)?.identity.clone(),
468 ))
469 }
470
471 fn coordinator_relations(&self, binding: &DataBinding) -> Result<Option<SampleRelationSet>> {
472 Ok(self
473 .attestation_for_binding(binding)?
474 .envelope
475 .coordinator_relations
476 .clone())
477 }
478}
479
480pub trait RuntimeController: Send + Sync {
481 fn controller_id(&self) -> &ControllerId;
482 fn invoke(&self, task: &NodeTask) -> Result<NodeResult>;
483
484 fn invoke_aggregation(
485 &self,
486 task: &AggregationControllerTask,
487 ) -> Result<AggregationControllerResult> {
488 Err(DagMlError::RuntimeValidation(format!(
489 "runtime controller `{}` does not implement aggregation task `{}`",
490 self.controller_id(),
491 task.task_id
492 )))
493 }
494}
495pub(crate) struct CollectedInputs {
496 pub(crate) handles: BTreeMap<String, HandleRef>,
497 pub(crate) data_views: BTreeMap<String, DataProviderViewSpec>,
498 pub(crate) prediction_inputs: BTreeMap<String, PredictionInputSpec>,
499 pub(crate) skip_node: bool,
500}
501
502pub(crate) fn data_view_key(input_name: &str) -> String {
503 format!("data:{input_name}")
504}
505
506pub(crate) fn validation_data_view_key(input_name: &str) -> String {
507 format!("{input_name}:validation")
508}
509
510pub(crate) fn derive_output_data_views(
511 plan: &ExecutionPlan,
512 task: &NodeTask,
513 result: &NodeResult,
514) -> Result<BTreeMap<String, DataProviderViewSpec>> {
515 let node = plan
516 .graph_plan
517 .graph
518 .nodes
519 .iter()
520 .find(|node| node.id == task.node_plan.node_id)
521 .expect("execution plan was validated");
522 let mut views = BTreeMap::new();
523 for port in node
524 .ports
525 .outputs
526 .iter()
527 .filter(|port| port.kind == PortKind::Data)
528 {
529 let Some(handle) = result.outputs.get(&port.name) else {
530 continue;
531 };
532 if !matches!(handle.kind, HandleKind::Data | HandleKind::DataView) {
533 return Err(DagMlError::RuntimeValidation(format!(
534 "node `{}` emitted data output `{}` with non-data/data-view handle kind {:?}",
535 task.node_plan.node_id, port.name, handle.kind
536 )));
537 }
538 if let Some(view) = primary_output_data_view(task) {
539 views.insert(
540 port.name.clone(),
541 output_data_view_for_port(task, result, &port.name, view)?,
542 );
543 }
544 if let Some(validation_view) = validation_output_data_view(task) {
545 views.insert(
546 validation_data_view_key(&port.name),
547 output_data_view_for_port(task, result, &port.name, validation_view)?,
548 );
549 }
550 }
551 Ok(views)
552}
553
554pub(crate) fn output_data_view_for_port(
555 task: &NodeTask,
556 result: &NodeResult,
557 port_name: &str,
558 base_view: &DataProviderViewSpec,
559) -> Result<DataProviderViewSpec> {
560 let mut view = base_view.clone();
561 if let Some(upstream_provenance) = view.extra.remove(DATA_OUTPUT_PROVENANCE_KEY) {
562 let provenance: DataOutputProvenance =
563 serde_json::from_value(upstream_provenance).map_err(|error| {
564 DagMlError::RuntimeValidation(format!(
565 "node `{}` cannot propagate data output `{port_name}` because upstream data output provenance is invalid JSON: {error}",
566 task.node_plan.node_id
567 ))
568 })?;
569 provenance.validate().map_err(|error| {
570 DagMlError::RuntimeValidation(format!(
571 "node `{}` cannot propagate data output `{port_name}` because upstream data output provenance is invalid: {error}",
572 task.node_plan.node_id
573 ))
574 })?;
575 }
576 let shape_deltas = result
577 .shape_deltas
578 .iter()
579 .filter(|delta| delta.node_id == task.node_plan.node_id)
580 .cloned()
581 .collect::<Vec<_>>();
582 let mut provenance = DataOutputProvenance {
583 schema_version: DATA_OUTPUT_PROVENANCE_SCHEMA_VERSION,
584 producer_node: task.node_plan.node_id.clone(),
585 producer_port: port_name.to_string(),
586 producer_phase: task.phase,
587 variant_id: task.variant_id.clone(),
588 fold_id: task.fold_id.clone(),
589 shape_plan_fingerprint: None,
590 aggregation_policy_fingerprint: None,
591 feature_namespace: None,
592 feature_schema_fingerprint: None,
593 representation_plan: None,
594 representation_replay_manifest: None,
595 representation_compatibility: None,
596 relation_delta_fingerprint: None,
597 shape_deltas,
598 };
599 if let Some(shape_plan) = &task.node_plan.shape_plan {
600 provenance.shape_plan_fingerprint = Some(stable_json_fingerprint(shape_plan)?);
601 provenance.aggregation_policy_fingerprint =
602 Some(stable_json_fingerprint(&shape_plan.aggregation_policy)?);
603 provenance.feature_namespace = shape_plan.feature_namespace.clone();
604 provenance.feature_schema_fingerprint =
605 output_feature_schema_fingerprint(shape_plan, result);
606 }
607 provenance.validate()?;
608
609 view.extra.insert(
610 DATA_OUTPUT_PROVENANCE_KEY.to_string(),
611 serde_json::to_value(provenance)?,
612 );
613 view.validate()?;
614 Ok(view)
615}
616
617pub(crate) fn output_feature_schema_fingerprint(
618 shape_plan: &crate::policy::DataModelShapePlan,
619 result: &NodeResult,
620) -> Option<String> {
621 result
622 .shape_deltas
623 .iter()
624 .rev()
625 .find(|delta| delta.kind == ShapeDeltaKind::Feature)
626 .map(|delta| delta.after_fingerprint.clone())
627 .or_else(|| shape_plan.feature_schema_fingerprint.clone())
628}
629
630pub(crate) fn primary_output_data_view(task: &NodeTask) -> Option<&DataProviderViewSpec> {
631 task.data_views
632 .values()
633 .find(|view| view.partition != DataRequestPartition::FoldValidation)
634 .or_else(|| task.data_views.values().next())
635}
636
637pub(crate) fn validation_output_data_view(task: &NodeTask) -> Option<&DataProviderViewSpec> {
638 task.data_views
639 .values()
640 .find(|view| view.partition == DataRequestPartition::FoldValidation)
641}
642
643pub(crate) fn make_data_view_handle(
644 data_provider: &dyn RuntimeDataProvider,
645 ctx: &RunContext,
646 node_plan: &NodePlan,
647 scope: &PhaseScope,
648 binding: &DataBinding,
649 data_handle: &HandleRef,
650 view: &DataProviderViewSpec,
651) -> Result<HandleRef> {
652 view.validate()?;
653 let view_handle = data_provider.make_view(&DataViewRequest {
654 run_id: ctx.run_id.clone(),
655 node_id: node_plan.node_id.clone(),
656 input_name: binding.input_name.clone(),
657 phase: scope.phase,
658 variant_id: scope.variant_id.clone(),
659 fold_id: scope.fold_id.clone(),
660 binding: binding.clone(),
661 data_handle: data_handle.clone(),
662 view: view.clone(),
663 })?;
664 if !matches!(view_handle.kind, HandleKind::Data | HandleKind::DataView) {
668 return Err(DagMlError::RuntimeValidation(format!(
669 "node `{}` data view `{}` resolved to a non-data/data-view handle kind {:?}",
670 node_plan.node_id, binding.input_name, view_handle.kind
671 )));
672 }
673 Ok(view_handle)
674}
675
676pub(crate) fn data_view_for_scope(
677 binding: &DataBinding,
678 fold_set: Option<&FoldSet>,
679 scope: &PhaseScope,
680 branch_view: Option<&crate::data::BranchViewPlan>,
681 excluded_samples: &BTreeSet<SampleId>,
682) -> Result<DataProviderViewSpec> {
683 let partition = data_partition_for_scope(binding, scope);
684 let role = match scope.phase {
687 Phase::FitCv | Phase::Refit => DataViewRole::Fit,
688 _ => DataViewRole::NonFit,
689 };
690 data_view_for_partition(
691 binding,
692 fold_set,
693 scope,
694 partition,
695 branch_view,
696 role,
697 excluded_samples,
698 )
699}
700
701pub(crate) fn validation_data_view_for_scope(
702 binding: &DataBinding,
703 fold_set: Option<&FoldSet>,
704 scope: &PhaseScope,
705 branch_view: Option<&crate::data::BranchViewPlan>,
706 excluded_samples: &BTreeSet<SampleId>,
707) -> Result<Option<DataProviderViewSpec>> {
708 if scope.phase != Phase::FitCv || scope.fold_id.is_none() {
709 return Ok(None);
710 }
711 let partition = binding.view_policy.predict_partition;
712 if partition == data_partition_for_scope(binding, scope) {
713 return Ok(None);
714 }
715 data_view_for_partition(
717 binding,
718 fold_set,
719 scope,
720 partition,
721 branch_view,
722 DataViewRole::NonFit,
723 excluded_samples,
724 )
725 .map(Some)
726}
727
728#[cfg(test)]
729mod envelope_attested_provider_tests {
730 use std::cell::Cell;
731
732 use super::*;
733
734 #[derive(Debug, Default)]
735 struct ProbeProvider {
736 materialize_calls: Cell<usize>,
737 make_view_calls: Cell<usize>,
738 }
739
740 impl RuntimeDataProvider for ProbeProvider {
741 fn materialize(&self, _request: &DataMaterializationRequest) -> Result<HandleRef> {
742 self.materialize_calls.set(self.materialize_calls.get() + 1);
743 Ok(HandleRef {
744 handle: 41,
745 kind: HandleKind::Data,
746 owner_controller: ControllerId::new("controller:data.probe").unwrap(),
747 })
748 }
749
750 fn make_view(&self, _request: &DataViewRequest) -> Result<HandleRef> {
751 self.make_view_calls.set(self.make_view_calls.get() + 1);
752 Ok(HandleRef {
753 handle: 42,
754 kind: HandleKind::DataView,
755 owner_controller: ControllerId::new("controller:data.probe").unwrap(),
756 })
757 }
758 }
759
760 fn complete_envelope() -> ExternalDataPlanEnvelope {
761 let mut envelope: ExternalDataPlanEnvelope = serde_json::from_str(include_str!(
762 "../../../../examples/fixtures/data/coordinator_data_plan_envelope_sample12.json"
763 ))
764 .unwrap();
765 envelope.data_content_fingerprint = Some("a".repeat(64));
766 envelope.target_content_fingerprint = Some("b".repeat(64));
767 envelope
768 }
769
770 fn binding_for(
771 node_id: &str,
772 input_name: &str,
773 envelope: &ExternalDataPlanEnvelope,
774 ) -> DataBinding {
775 DataBinding {
776 node_id: NodeId::new(node_id).unwrap(),
777 input_name: input_name.to_string(),
778 request_id: "request:data.probe".to_string(),
779 schema_fingerprint: envelope.schema_fingerprint.clone(),
780 plan_fingerprint: envelope.plan_fingerprint.clone(),
781 relation_fingerprint: envelope.relation_fingerprint.clone(),
782 output_representation: "tabular_numeric".to_string(),
783 feature_set_id: Some(input_name.to_string()),
784 source_ids: vec!["source:probe".to_string()],
785 require_relations: true,
786 view_policy: Default::default(),
787 metadata: BTreeMap::new(),
788 }
789 }
790
791 fn envelopes_for(
792 binding: &DataBinding,
793 envelope: ExternalDataPlanEnvelope,
794 ) -> BTreeMap<String, ExternalDataPlanEnvelope> {
795 BTreeMap::from([(
796 data_binding_requirement_key(&binding.node_id, &binding.input_name),
797 envelope,
798 )])
799 }
800
801 fn materialization_request(binding: &DataBinding) -> DataMaterializationRequest {
802 DataMaterializationRequest {
803 run_id: RunId::new("run:attested.provider").unwrap(),
804 node_id: binding.node_id.clone(),
805 input_name: binding.input_name.clone(),
806 phase: Phase::Refit,
807 variant_id: None,
808 fold_id: None,
809 binding: binding.clone(),
810 }
811 }
812
813 #[test]
814 fn envelope_attested_provider_delegates_and_returns_exact_attestations() {
815 let envelope = complete_envelope();
816 let binding = binding_for("model:base", "x", &envelope);
817 let expected_identity =
818 crate::training::TrainingDataIdentity::from_binding_envelope(&binding, &envelope)
819 .unwrap();
820 let expected_relations = envelope.coordinator_relations.clone();
821 let provider = EnvelopeAttestedRuntimeDataProvider::new(
822 ProbeProvider::default(),
823 vec![binding.clone()],
824 envelopes_for(&binding, envelope),
825 )
826 .unwrap();
827
828 assert_eq!(
829 provider.training_data_identity(&binding).unwrap(),
830 Some(expected_identity)
831 );
832 assert_eq!(
833 provider.coordinator_relations(&binding).unwrap(),
834 expected_relations
835 );
836
837 let materialization = materialization_request(&binding);
838 let data_handle = provider.materialize(&materialization).unwrap();
839 assert_eq!(data_handle.handle, 41);
840 let view_handle = provider
841 .make_view(&DataViewRequest {
842 run_id: materialization.run_id,
843 node_id: binding.node_id.clone(),
844 input_name: binding.input_name.clone(),
845 phase: Phase::Refit,
846 variant_id: None,
847 fold_id: None,
848 binding: binding.clone(),
849 data_handle,
850 view: DataProviderViewSpec {
851 sample_ids: None,
852 partition: DataRequestPartition::FullTrain,
853 fold_id: None,
854 source_ids: None,
855 columns: None,
856 include_augmented: true,
857 include_excluded: false,
858 branch_view: None,
859 extra: BTreeMap::new(),
860 },
861 })
862 .unwrap();
863 assert_eq!(view_handle.handle, 42);
864 assert_eq!(provider.inner().materialize_calls.get(), 1);
865 assert_eq!(provider.inner().make_view_calls.get(), 1);
866
867 let inner = provider.into_inner();
868 assert_eq!(inner.materialize_calls.get(), 1);
869 assert_eq!(inner.make_view_calls.get(), 1);
870 }
871
872 #[test]
873 fn envelope_attested_provider_requires_exact_envelope_coverage() {
874 let envelope = complete_envelope();
875 let binding = binding_for("model:base", "x", &envelope);
876
877 let missing = EnvelopeAttestedRuntimeDataProvider::new(
878 ProbeProvider::default(),
879 vec![binding.clone()],
880 BTreeMap::new(),
881 )
882 .unwrap_err();
883 assert!(missing.to_string().contains("exactly cover"));
884 assert!(missing.to_string().contains("model:base.x"));
885
886 let mut unexpected = envelopes_for(&binding, envelope.clone());
887 unexpected.insert("model:other.x".to_string(), envelope);
888 let extra = EnvelopeAttestedRuntimeDataProvider::new(
889 ProbeProvider::default(),
890 vec![binding],
891 unexpected,
892 )
893 .unwrap_err();
894 assert!(extra.to_string().contains("exactly cover"));
895 assert!(extra.to_string().contains("model:other.x"));
896 }
897
898 #[test]
899 fn envelope_attested_provider_rejects_rendered_key_collisions() {
900 let envelope = complete_envelope();
901 let left = binding_for("a.b", "c", &envelope);
902 let right = binding_for("a", "b.c", &envelope);
903 assert_eq!(
904 data_binding_requirement_key(&left.node_id, &left.input_name),
905 data_binding_requirement_key(&right.node_id, &right.input_name)
906 );
907
908 let error = EnvelopeAttestedRuntimeDataProvider::new(
909 ProbeProvider::default(),
910 vec![left.clone(), right],
911 envelopes_for(&left, envelope),
912 )
913 .unwrap_err();
914 assert!(error.to_string().contains("distinct coordinates"));
915 assert!(error.to_string().contains("a.b.c"));
916 }
917
918 #[test]
919 fn envelope_attested_provider_refuses_unattested_binding_before_delegation() {
920 let envelope = complete_envelope();
921 let binding = binding_for("model:base", "x", &envelope);
922 let provider = EnvelopeAttestedRuntimeDataProvider::new(
923 ProbeProvider::default(),
924 vec![binding.clone()],
925 envelopes_for(&binding, envelope),
926 )
927 .unwrap();
928 let mut changed = binding;
929 changed.request_id = "request:data.changed".to_string();
930
931 let error = provider
932 .materialize(&materialization_request(&changed))
933 .unwrap_err();
934 assert!(error.to_string().contains("does not exactly match"));
935 assert_eq!(provider.inner().materialize_calls.get(), 0);
936 }
937
938 #[test]
939 fn envelope_attested_provider_refuses_incomplete_training_envelope() {
940 let mut envelope = complete_envelope();
941 envelope.data_content_fingerprint = None;
942 let binding = binding_for("model:base", "x", &envelope);
943 let error = EnvelopeAttestedRuntimeDataProvider::new(
944 ProbeProvider::default(),
945 vec![binding.clone()],
946 envelopes_for(&binding, envelope),
947 )
948 .unwrap_err();
949 assert!(error.to_string().contains("data content fingerprint"));
950 }
951}