Skip to main content

ferrum_interfaces/vnext/model/checkpoint/
inputs.rs

1use std::collections::BTreeSet;
2
3use serde::{Deserialize, Deserializer, Serialize};
4
5use crate::vnext::{ContractVersion, ProgramValueId, VNextError};
6
7pub const PROGRAM_CHECKPOINT_INPUTS_VERSION: ContractVersion = ContractVersion::new(1, 0);
8
9/// Explicit checkpoint input roles. Conditioning inputs are compared in full,
10/// including shape, dtype and byte contents. Output-only inputs may be omitted
11/// from matching only after the execution plan proves they cannot affect any
12/// state. Names and tensor widths do not establish an input's semantic role.
13#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
14pub struct ProgramCheckpointInputs {
15    contract_version: ContractVersion,
16    token_input: ProgramValueId,
17    conditioning_inputs: BTreeSet<ProgramValueId>,
18    #[serde(skip_serializing_if = "BTreeSet::is_empty")]
19    output_only_inputs: BTreeSet<ProgramValueId>,
20}
21
22#[cfg(test)]
23mod tests;
24
25#[derive(Deserialize)]
26#[serde(deny_unknown_fields)]
27struct Wire {
28    contract_version: ContractVersion,
29    token_input: ProgramValueId,
30    conditioning_inputs: Vec<ProgramValueId>,
31    #[serde(default)]
32    output_only_inputs: Vec<ProgramValueId>,
33}
34
35impl<'de> Deserialize<'de> for ProgramCheckpointInputs {
36    fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
37        let wire = Wire::deserialize(deserializer)?;
38        if wire.contract_version != PROGRAM_CHECKPOINT_INPUTS_VERSION
39            || wire
40                .conditioning_inputs
41                .windows(2)
42                .any(|pair| pair[0] >= pair[1])
43            || wire
44                .output_only_inputs
45                .windows(2)
46                .any(|pair| pair[0] >= pair[1])
47        {
48            return Err(serde::de::Error::custom(
49                "checkpoint input version or canonical input order is invalid",
50            ));
51        }
52        Self::new(
53            wire.token_input,
54            wire.conditioning_inputs.into_iter().collect(),
55        )
56        .and_then(|inputs| {
57            inputs.with_output_only_inputs(wire.output_only_inputs.into_iter().collect())
58        })
59        .map_err(serde::de::Error::custom)
60    }
61}
62
63impl ProgramCheckpointInputs {
64    pub fn new(
65        token_input: ProgramValueId,
66        conditioning_inputs: BTreeSet<ProgramValueId>,
67    ) -> Result<Self, VNextError> {
68        if conditioning_inputs.contains(&token_input) {
69            return Err(VNextError::InvalidExecutionPlan {
70                reason: "checkpoint token input cannot also be a conditioning input".to_owned(),
71            });
72        }
73        Ok(Self {
74            contract_version: PROGRAM_CHECKPOINT_INPUTS_VERSION,
75            token_input,
76            conditioning_inputs,
77            output_only_inputs: BTreeSet::new(),
78        })
79    }
80
81    /// Declare inputs that affect only current outputs, not continuation state.
82    /// This declaration alone is not authority to skip matching: plan derivation
83    /// checks all value, token-work and exact-alias paths to state effects.
84    pub fn with_output_only_inputs(
85        mut self,
86        output_only_inputs: BTreeSet<ProgramValueId>,
87    ) -> Result<Self, VNextError> {
88        if output_only_inputs.contains(&self.token_input)
89            || !output_only_inputs.is_disjoint(&self.conditioning_inputs)
90        {
91            return Err(VNextError::InvalidExecutionPlan {
92                reason: "checkpoint input roles must be disjoint".to_owned(),
93            });
94        }
95        self.output_only_inputs = output_only_inputs;
96        Ok(self)
97    }
98
99    pub fn token_input(&self) -> &ProgramValueId {
100        &self.token_input
101    }
102
103    pub fn conditioning_inputs(&self) -> &BTreeSet<ProgramValueId> {
104        &self.conditioning_inputs
105    }
106
107    pub fn output_only_inputs(&self) -> &BTreeSet<ProgramValueId> {
108        &self.output_only_inputs
109    }
110
111    pub(crate) fn covers(&self, inputs: &[ProgramValueId]) -> bool {
112        let declared = self
113            .conditioning_inputs
114            .iter()
115            .chain(self.output_only_inputs.iter())
116            .chain([&self.token_input])
117            .collect::<BTreeSet<_>>();
118        declared == inputs.iter().collect::<BTreeSet<_>>()
119    }
120}