Skip to main content

microsandbox_protocol/control/
message.rs

1//! Control message inventory and checked dispatch decoding.
2
3use ciborium::Value;
4use serde::{Deserialize, Serialize};
5
6use super::{ControlRequest, CpuTarget, Empty, MemoryTarget, SecretsResult, SecretsUpdate};
7use crate::wire::{Envelope, WireError, decode_value, validate_record};
8
9//--------------------------------------------------------------------------------------------------
10// Types
11//--------------------------------------------------------------------------------------------------
12
13/// Known control wire names. The strum spelling is the authoritative mapping.
14#[derive(
15    Debug,
16    Clone,
17    Copy,
18    PartialEq,
19    Eq,
20    Hash,
21    strum::IntoStaticStr,
22    strum::EnumString,
23    strum::EnumIter,
24)]
25pub enum ControlMessageType {
26    /// Initial generation offer.
27    #[strum(serialize = "control.hello")]
28    Hello,
29    /// Selected generation and limits.
30    #[strum(serialize = "control.welcome")]
31    Welcome,
32    /// Facility query.
33    #[strum(serialize = "control.capabilities")]
34    Capabilities,
35    /// Facility response.
36    #[strum(serialize = "control.capabilities.result")]
37    CapabilitiesResult,
38    /// Memory query or terminal memory observation.
39    #[strum(serialize = "control.memory.state")]
40    MemoryState,
41    /// Set a memory target.
42    #[strum(serialize = "control.memory.target")]
43    MemoryTarget,
44    /// CPU query or terminal CPU observation.
45    #[strum(serialize = "control.cpu.state")]
46    CpuState,
47    /// Set a CPU target.
48    #[strum(serialize = "control.cpu.target")]
49    CpuTarget,
50    /// Ordered host secret modifications.
51    #[strum(serialize = "control.secrets.update")]
52    SecretsUpdate,
53    /// Complete or partial secret progress.
54    #[strum(serialize = "control.secrets.result")]
55    SecretsResult,
56    /// Recoverable operation or handshake error.
57    #[strum(serialize = "control.error")]
58    Error,
59}
60
61//--------------------------------------------------------------------------------------------------
62// Methods
63//--------------------------------------------------------------------------------------------------
64
65impl ControlMessageType {
66    /// Stable wire spelling.
67    pub fn as_str(self) -> &'static str {
68        self.into()
69    }
70
71    /// Resolve a known spelling, retaining unknown names in the envelope layer.
72    pub fn from_wire_str(name: &str) -> Option<Self> {
73        name.parse().ok()
74    }
75
76    /// Whether this is an application request permitted after the handshake.
77    pub fn is_request(self) -> bool {
78        matches!(
79            self,
80            Self::Capabilities
81                | Self::MemoryState
82                | Self::MemoryTarget
83                | Self::CpuState
84                | Self::CpuTarget
85                | Self::SecretsUpdate
86        )
87    }
88}
89
90impl ControlRequest {
91    /// Decode a known application request after the server validates its frame.
92    pub fn from_envelope(envelope: &Envelope) -> Result<Self, WireError> {
93        Ok(match ControlMessageType::from_wire_str(&envelope.t) {
94            Some(ControlMessageType::Capabilities) => {
95                envelope.payload::<Empty>()?;
96                Self::Capabilities
97            }
98            Some(ControlMessageType::MemoryState) => {
99                envelope.payload::<Empty>()?;
100                Self::MemoryState
101            }
102            Some(ControlMessageType::CpuState) => {
103                envelope.payload::<Empty>()?;
104                Self::CpuState
105            }
106            Some(ControlMessageType::MemoryTarget) => {
107                let payload: MemoryTarget = envelope.payload()?;
108                Self::MemoryTarget {
109                    total_mib: payload.total_mib,
110                }
111            }
112            Some(ControlMessageType::CpuTarget) => {
113                let payload: CpuTarget = envelope.payload()?;
114                Self::CpuTarget {
115                    online: payload.online,
116                }
117            }
118            Some(ControlMessageType::SecretsUpdate) => {
119                // Serde's tagged-enum buffering alone does not reject every
120                // duplicate. Check each entry before any host operation runs.
121                let value = decode_value(&envelope.p)?;
122                validate_record(&value)?;
123                let Value::Map(fields) = &value else {
124                    unreachable!()
125                };
126                if let Some((_, Value::Array(changes))) = fields
127                    .iter()
128                    .find(|(key, _)| key.as_text() == Some("changes"))
129                {
130                    for change in changes {
131                        validate_record(change)?;
132                    }
133                }
134                let payload: SecretsUpdate =
135                    value.deserialized().map_err(|_| WireError::InvalidRecord)?;
136                Self::SecretsUpdate {
137                    changes: payload.changes,
138                }
139            }
140            _ => return Err(WireError::InvalidRecord),
141        })
142    }
143
144    /// Encode the CBOR operation corresponding to this legacy-compatible value.
145    pub fn envelope(&self, generation: u8) -> Result<Envelope, WireError> {
146        match self {
147            Self::Capabilities => Envelope::new(
148                generation,
149                ControlMessageType::Capabilities.as_str(),
150                &Empty {},
151            ),
152            Self::MemoryState => Envelope::new(
153                generation,
154                ControlMessageType::MemoryState.as_str(),
155                &Empty {},
156            ),
157            Self::CpuState => {
158                Envelope::new(generation, ControlMessageType::CpuState.as_str(), &Empty {})
159            }
160            Self::MemoryTarget { total_mib } => Envelope::new(
161                generation,
162                ControlMessageType::MemoryTarget.as_str(),
163                &MemoryTarget {
164                    total_mib: *total_mib,
165                },
166            ),
167            Self::CpuTarget { online } => Envelope::new(
168                generation,
169                ControlMessageType::CpuTarget.as_str(),
170                &CpuTarget { online: *online },
171            ),
172            Self::SecretsUpdate { changes } => {
173                // Borrow secret entries instead of cloning their plaintext.
174                #[derive(Serialize)]
175                struct Payload<'a> {
176                    changes: &'a [super::SecretChange],
177                }
178                Envelope::new(
179                    generation,
180                    ControlMessageType::SecretsUpdate.as_str(),
181                    &Payload { changes },
182                )
183            }
184        }
185    }
186}
187
188impl SecretsResult {
189    /// Decode sequential progress, checking nested error keys and index equality.
190    pub fn decode(payload: &[u8]) -> Result<Self, WireError> {
191        let value = decode_value(payload)?;
192        validate_record(&value)?;
193        if let Value::Map(fields) = &value
194            && let Some((_, error)) = fields
195                .iter()
196                .find(|(key, _)| key.as_text() == Some("error"))
197        {
198            validate_record(error)?;
199        }
200        let result: Self = value.deserialized().map_err(|_| WireError::InvalidRecord)?;
201        if matches!(&result, Self::Failed { applied_count, failed_index, .. } if applied_count != failed_index)
202        {
203            return Err(WireError::InvalidRecord);
204        }
205        Ok(result)
206    }
207}
208
209//--------------------------------------------------------------------------------------------------
210// Trait Implementations
211//--------------------------------------------------------------------------------------------------
212
213impl AsRef<str> for ControlMessageType {
214    fn as_ref(&self) -> &str {
215        self.as_str()
216    }
217}
218
219impl Serialize for ControlMessageType {
220    fn serialize<S: serde::Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
221        serializer.serialize_str(self.as_str())
222    }
223}
224
225impl<'de> Deserialize<'de> for ControlMessageType {
226    fn deserialize<D: serde::Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
227        let name = String::deserialize(deserializer)?;
228        Self::from_wire_str(&name)
229            .ok_or_else(|| serde::de::Error::custom("unknown control message type"))
230    }
231}