microsandbox_protocol/control/
message.rs1use ciborium::Value;
4use serde::{Deserialize, Serialize};
5
6use super::{
7 BranchCreate, CheckpointCreate, ControlRequest, CpuTarget, DiskCheckpointCreate, DiskCompact,
8 Empty, MemoryTarget, Pause, RootDiskGrow, SecretsResult, SecretsUpdate,
9};
10use crate::wire::{Envelope, WireError, decode_value, validate_record};
11
12pub const CONTROL_GENERATION_TWO_MESSAGES: &[&str] = &[
21 "control.checkpoint.create",
22 "control.checkpoint.result",
23 "control.disk.checkpoint.create",
24 "control.disk.checkpoint.result",
25 "control.branch.create",
26 "control.branch.result",
27 "control.pause",
28 "control.resume",
29 "control.pause.state",
30 "control.root-disk.grow",
31 "control.root-disk.state",
32 "control.disk.compact",
33 "control.disk.compact.result",
34];
35
36#[derive(
42 Debug,
43 Clone,
44 Copy,
45 PartialEq,
46 Eq,
47 Hash,
48 strum::IntoStaticStr,
49 strum::EnumString,
50 strum::EnumIter,
51)]
52pub enum ControlMessageType {
53 #[strum(serialize = "control.hello")]
55 Hello,
56 #[strum(serialize = "control.welcome")]
58 Welcome,
59 #[strum(serialize = "control.capabilities")]
61 Capabilities,
62 #[strum(serialize = "control.capabilities.result")]
64 CapabilitiesResult,
65 #[strum(serialize = "control.memory.state")]
67 MemoryState,
68 #[strum(serialize = "control.memory.target")]
70 MemoryTarget,
71 #[strum(serialize = "control.cpu.state")]
73 CpuState,
74 #[strum(serialize = "control.cpu.target")]
76 CpuTarget,
77 #[strum(serialize = "control.secrets.update")]
79 SecretsUpdate,
80 #[strum(serialize = "control.secrets.result")]
82 SecretsResult,
83 #[strum(serialize = "control.error")]
85 Error,
86}
87
88#[derive(Debug, Clone)]
93pub enum ControlOperation {
94 GenerationOne(ControlRequest),
96 CheckpointCreate(CheckpointCreate),
98 DiskCheckpointCreate(DiskCheckpointCreate),
100 BranchCreate(BranchCreate),
102 Pause(Pause),
104 Resume,
106 PauseState,
108 RootDiskGrow(RootDiskGrow),
110 DiskCompact(DiskCompact),
112}
113
114impl ControlMessageType {
119 pub fn as_str(self) -> &'static str {
121 self.into()
122 }
123
124 pub fn from_wire_str(name: &str) -> Option<Self> {
126 name.parse().ok()
127 }
128
129 pub fn is_request(self) -> bool {
131 matches!(
132 self,
133 Self::Capabilities
134 | Self::MemoryState
135 | Self::MemoryTarget
136 | Self::CpuState
137 | Self::CpuTarget
138 | Self::SecretsUpdate
139 )
140 }
141}
142
143impl ControlOperation {
144 pub fn from_envelope(envelope: &Envelope, generation: u8) -> Result<Self, WireError> {
146 if envelope.v != generation {
147 return Err(WireError::InvalidRecord);
148 }
149 let operation = match envelope.t.as_str() {
150 "control.checkpoint.create" if generation >= 2 => {
151 Self::CheckpointCreate(envelope.payload()?)
152 }
153 "control.disk.checkpoint.create" if generation >= 2 => {
154 Self::DiskCheckpointCreate(envelope.payload()?)
155 }
156 "control.branch.create" if generation >= 2 => Self::BranchCreate(envelope.payload()?),
157 "control.pause" if generation >= 2 => Self::Pause(envelope.payload()?),
158 "control.resume" if generation >= 2 => {
159 envelope.payload::<Empty>()?;
160 Self::Resume
161 }
162 "control.pause.state" if generation >= 2 => {
163 envelope.payload::<Empty>()?;
164 Self::PauseState
165 }
166 "control.root-disk.grow" if generation >= 2 => Self::RootDiskGrow(envelope.payload()?),
167 "control.disk.compact" if generation >= 2 => {
168 let value = decode_value(&envelope.p)?;
171 validate_record(&value)?;
172 let Value::Map(fields) = &value else {
173 unreachable!()
174 };
175 if let Some((_, target)) = fields
176 .iter()
177 .find(|(key, _)| key.as_text() == Some("target"))
178 {
179 validate_record(target)?;
180 }
181 Self::DiskCompact(value.deserialized().map_err(|_| WireError::InvalidRecord)?)
182 }
183 _ => Self::GenerationOne(ControlRequest::from_envelope(envelope)?),
184 };
185 Ok(operation)
186 }
187}
188
189pub fn control_message_min_generation(name: &str) -> Option<u8> {
193 if CONTROL_GENERATION_TWO_MESSAGES.contains(&name) {
194 return Some(2);
195 }
196 Some(match name {
197 "control.capabilities"
198 | "control.capabilities.result"
199 | "control.memory.state"
200 | "control.memory.target"
201 | "control.cpu.state"
202 | "control.cpu.target"
203 | "control.secrets.update"
204 | "control.secrets.result"
205 | "control.error" => 1,
206 _ => return None,
207 })
208}
209
210impl ControlRequest {
211 pub fn from_envelope(envelope: &Envelope) -> Result<Self, WireError> {
213 Ok(match ControlMessageType::from_wire_str(&envelope.t) {
214 Some(ControlMessageType::Capabilities) => {
215 envelope.payload::<Empty>()?;
216 Self::Capabilities
217 }
218 Some(ControlMessageType::MemoryState) => {
219 envelope.payload::<Empty>()?;
220 Self::MemoryState
221 }
222 Some(ControlMessageType::CpuState) => {
223 envelope.payload::<Empty>()?;
224 Self::CpuState
225 }
226 Some(ControlMessageType::MemoryTarget) => {
227 let payload: MemoryTarget = envelope.payload()?;
228 Self::MemoryTarget {
229 total_mib: payload.total_mib,
230 }
231 }
232 Some(ControlMessageType::CpuTarget) => {
233 let payload: CpuTarget = envelope.payload()?;
234 Self::CpuTarget {
235 online: payload.online,
236 }
237 }
238 Some(ControlMessageType::SecretsUpdate) => {
239 let value = decode_value(&envelope.p)?;
242 validate_record(&value)?;
243 let Value::Map(fields) = &value else {
244 unreachable!()
245 };
246 if let Some((_, Value::Array(changes))) = fields
247 .iter()
248 .find(|(key, _)| key.as_text() == Some("changes"))
249 {
250 for change in changes {
251 validate_record(change)?;
252 }
253 }
254 let payload: SecretsUpdate =
255 value.deserialized().map_err(|_| WireError::InvalidRecord)?;
256 Self::SecretsUpdate {
257 changes: payload.changes,
258 }
259 }
260 _ => return Err(WireError::InvalidRecord),
261 })
262 }
263
264 pub fn envelope(&self, generation: u8) -> Result<Envelope, WireError> {
266 match self {
267 Self::Capabilities => Envelope::new(
268 generation,
269 ControlMessageType::Capabilities.as_str(),
270 &Empty {},
271 ),
272 Self::MemoryState => Envelope::new(
273 generation,
274 ControlMessageType::MemoryState.as_str(),
275 &Empty {},
276 ),
277 Self::CpuState => {
278 Envelope::new(generation, ControlMessageType::CpuState.as_str(), &Empty {})
279 }
280 Self::MemoryTarget { total_mib } => Envelope::new(
281 generation,
282 ControlMessageType::MemoryTarget.as_str(),
283 &MemoryTarget {
284 total_mib: *total_mib,
285 },
286 ),
287 Self::CpuTarget { online } => Envelope::new(
288 generation,
289 ControlMessageType::CpuTarget.as_str(),
290 &CpuTarget { online: *online },
291 ),
292 Self::SecretsUpdate { changes } => {
293 #[derive(Serialize)]
295 struct Payload<'a> {
296 changes: &'a [super::SecretChange],
297 }
298 Envelope::new(
299 generation,
300 ControlMessageType::SecretsUpdate.as_str(),
301 &Payload { changes },
302 )
303 }
304 }
305 }
306}
307
308impl SecretsResult {
309 pub fn decode(payload: &[u8]) -> Result<Self, WireError> {
311 let value = decode_value(payload)?;
312 validate_record(&value)?;
313 if let Value::Map(fields) = &value
314 && let Some((_, error)) = fields
315 .iter()
316 .find(|(key, _)| key.as_text() == Some("error"))
317 {
318 validate_record(error)?;
319 }
320 let result: Self = value.deserialized().map_err(|_| WireError::InvalidRecord)?;
321 if matches!(&result, Self::Failed { applied_count, failed_index, .. } if applied_count != failed_index)
322 {
323 return Err(WireError::InvalidRecord);
324 }
325 Ok(result)
326 }
327}
328
329impl AsRef<str> for ControlMessageType {
334 fn as_ref(&self) -> &str {
335 self.as_str()
336 }
337}
338
339impl Serialize for ControlMessageType {
340 fn serialize<S: serde::Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
341 serializer.serialize_str(self.as_str())
342 }
343}
344
345impl<'de> Deserialize<'de> for ControlMessageType {
346 fn deserialize<D: serde::Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
347 let name = String::deserialize(deserializer)?;
348 Self::from_wire_str(&name)
349 .ok_or_else(|| serde::de::Error::custom("unknown control message type"))
350 }
351}