Skip to main content

dag_ml_core/runtime/
dataview.rs

1// Auto-split from the former monolithic `runtime.rs` (pure refactor).
2use 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    /// Attest the exact feature and target content bound to one training input.
305    ///
306    /// Legacy phase execution may return `None`; the native W1 training
307    /// operation requires `Some` and compares it byte-for-byte with the signed
308    /// [`TrainingDataIdentity`](crate::training::TrainingDataIdentity).
309    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/// Owns a host data provider while supplying exact, envelope-backed training
328/// attestations at the runtime trust boundary.
329///
330/// Construction validates the complete binding/envelope set before the inner
331/// provider can be invoked. Runtime calls are delegated only when their full
332/// [`DataBinding`] is field-for-field equal to the binding registered for the
333/// rendered V1 requirement key.
334#[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    // A data view is delivered to the controller as a data input, so the
665    // provider must return a data-bearing handle. Refuse a model / artifact /
666    // prediction / relation handle masquerading as a view across the ABI.
667    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    // During FIT_CV and REFIT this primary view IS the training input; during
685    // PREDICT/EXPLAIN (and the planning phases) it is a non-fit read.
686    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    // This is the validation companion read, never the training input.
716    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}