Skip to main content

microsandbox_control_client/
request.rs

1//! Prepared checked operations. No SDK resource convergence policy lives here.
2
3use microsandbox_protocol::{
4    control::{
5        BranchCreate, BranchResult, Capabilities, CheckpointCreate, CheckpointResult, ControlError,
6        CpuState, CpuTarget, DiskCheckpointCreate, DiskCheckpointState, DiskCompact, Empty,
7        MemoryState, MemoryTarget, Pause, PauseState, RootDiskGrow, RootDiskState,
8        RuntimeCapabilities, SecretChange, SecretsResult,
9    },
10    wire,
11};
12use microsandbox_protocol_client::{EncodedMessage, Message, Request};
13use microsandbox_utils::size::Mebibytes;
14use serde::{Serialize, de::DeserializeOwned};
15
16use crate::{
17    CompatibleControlRequest, ControlClientError, ControlClientResult, ControlProtocol, JsonReply,
18};
19use zeroize::Zeroizing;
20
21//--------------------------------------------------------------------------------------------------
22// Types
23//--------------------------------------------------------------------------------------------------
24
25/// Query available host operations.
26#[derive(Debug, Clone, Copy, Default)]
27pub struct GetCapabilities;
28/// Query the complete generation-two runtime facility inventory.
29#[derive(Debug, Clone, Copy, Default)]
30pub struct GetRuntimeCapabilities;
31/// Read accepted and observed memory quantities.
32#[derive(Debug, Clone, Copy, Default)]
33pub struct GetMemoryState;
34/// Set a memory target without waiting for guest convergence.
35#[derive(Debug, Clone, Copy)]
36pub struct SetMemoryTarget {
37    /// Full-width wire quantity; direct construction avoids SDK input narrowing.
38    pub total_mib: u64,
39}
40/// Read CPU capacity, target, observation, and enforcement.
41#[derive(Debug, Clone, Copy, Default)]
42pub struct GetCpuState;
43/// Set a CPU target without waiting for guest convergence.
44#[derive(Debug, Clone, Copy)]
45pub struct SetCpuTarget {
46    /// Requested online CPUs.
47    pub online: u32,
48}
49/// Apply ordered secret changes, preserving partial completion in the result.
50#[derive(Debug, Clone)]
51pub struct UpdateSecrets {
52    /// Entries execute sequentially, stopping at the first operation failure.
53    pub changes: Vec<SecretChange>,
54}
55/// Create one full checkpoint.
56#[derive(Debug, Clone)]
57pub struct CreateCheckpoint(pub CheckpointCreate);
58/// Create one disk-only checkpoint.
59#[derive(Debug, Clone)]
60pub struct CreateDiskCheckpoint(pub DiskCheckpointCreate);
61/// Create one direct local branch without descriptor transfer.
62#[derive(Debug, Clone)]
63pub struct CreateBranch(pub BranchCreate);
64/// Pause the runtime, optionally requiring guest writeback.
65#[derive(Debug, Clone, Copy, Default)]
66pub struct PauseRuntime(pub Pause);
67/// Resume a resident pause.
68#[derive(Debug, Clone, Copy, Default)]
69pub struct ResumeRuntime;
70/// Inspect resident pause state.
71#[derive(Debug, Clone, Copy, Default)]
72pub struct GetPauseState;
73/// Grow the owned root disk and filesystem.
74#[derive(Debug, Clone, Copy)]
75pub struct GrowRootDisk(pub RootDiskGrow);
76/// Compact selected owned disk chains.
77#[derive(Debug, Clone)]
78pub struct CompactDisks(pub DiskCompact);
79
80//--------------------------------------------------------------------------------------------------
81// Methods
82//--------------------------------------------------------------------------------------------------
83
84impl SetMemoryTarget {
85    /// Accept the SDK's existing integer/MiB helpers and their conversion rules.
86    ///
87    /// `SetMemoryTarget::new(2048.mib())` uses the shared `SizeExt` owner. This
88    /// does not change its existing sub-MiB truncation or overflow semantics.
89    /// For full-width MiB input construct the public `total_mib` field directly.
90    pub fn new(size: impl Into<Mebibytes>) -> Self {
91        Self {
92            total_mib: u64::from(size.into().as_u32()),
93        }
94    }
95}
96
97impl SetCpuTarget {
98    /// Prepare an online CPU target without performing I/O.
99    pub fn new(online: u32) -> Self {
100        Self { online }
101    }
102}
103
104impl UpdateSecrets {
105    /// Prepare a caller-ordered batch without performing I/O.
106    pub fn new(changes: Vec<SecretChange>) -> Self {
107        Self { changes }
108    }
109}
110
111//--------------------------------------------------------------------------------------------------
112// Trait Implementations
113//--------------------------------------------------------------------------------------------------
114
115impl Request<ControlProtocol> for GetCapabilities {
116    type Response = Capabilities;
117    type Error = ControlClientError;
118    fn message(&self) -> ControlClientResult<EncodedMessage> {
119        prepared("control.capabilities", &Empty {})
120    }
121    fn decode(&self, response: Message) -> ControlClientResult<Self::Response> {
122        checked(response, 1, "control.capabilities.result")
123    }
124}
125
126impl Request<ControlProtocol> for GetRuntimeCapabilities {
127    type Response = RuntimeCapabilities;
128    type Error = ControlClientError;
129    fn message(&self) -> ControlClientResult<EncodedMessage> {
130        prepared("control.capabilities", &Empty {})
131    }
132    fn decode(&self, response: Message) -> ControlClientResult<Self::Response> {
133        checked(response, 2, "control.capabilities.result")
134    }
135}
136
137impl Request<ControlProtocol> for GetMemoryState {
138    type Response = MemoryState;
139    type Error = ControlClientError;
140    fn message(&self) -> ControlClientResult<EncodedMessage> {
141        prepared("control.memory.state", &Empty {})
142    }
143    fn decode(&self, response: Message) -> ControlClientResult<Self::Response> {
144        checked(response, 1, "control.memory.state")
145    }
146}
147
148impl Request<ControlProtocol> for SetMemoryTarget {
149    type Response = MemoryState;
150    type Error = ControlClientError;
151    fn message(&self) -> ControlClientResult<EncodedMessage> {
152        prepared(
153            "control.memory.target",
154            &MemoryTarget {
155                total_mib: self.total_mib,
156            },
157        )
158    }
159    fn decode(&self, response: Message) -> ControlClientResult<Self::Response> {
160        checked(response, 1, "control.memory.state")
161    }
162}
163
164impl Request<ControlProtocol> for GetCpuState {
165    type Response = CpuState;
166    type Error = ControlClientError;
167    fn message(&self) -> ControlClientResult<EncodedMessage> {
168        prepared("control.cpu.state", &Empty {})
169    }
170    fn decode(&self, response: Message) -> ControlClientResult<Self::Response> {
171        checked(response, 1, "control.cpu.state")
172    }
173}
174
175impl Request<ControlProtocol> for SetCpuTarget {
176    type Response = CpuState;
177    type Error = ControlClientError;
178    fn message(&self) -> ControlClientResult<EncodedMessage> {
179        prepared(
180            "control.cpu.target",
181            &CpuTarget {
182                online: self.online,
183            },
184        )
185    }
186    fn decode(&self, response: Message) -> ControlClientResult<Self::Response> {
187        checked(response, 1, "control.cpu.state")
188    }
189}
190
191impl Request<ControlProtocol> for UpdateSecrets {
192    type Response = SecretsResult;
193    type Error = ControlClientError;
194    fn message(&self) -> ControlClientResult<EncodedMessage> {
195        #[derive(Serialize)]
196        struct Payload<'a> {
197            changes: &'a [SecretChange],
198        }
199        prepared(
200            "control.secrets.update",
201            &Payload {
202                changes: &self.changes,
203            },
204        )
205    }
206    fn decode(&self, response: Message) -> ControlClientResult<Self::Response> {
207        checked_with(response, 1, "control.secrets.result", |bytes| {
208            let result = SecretsResult::decode(bytes)?;
209            // Completion cannot include entries the caller never sent. Preserve
210            // partial failure as an ordinary typed result, not a rollback claim.
211            let valid = match &result {
212                SecretsResult::Complete { applied_count } => {
213                    *applied_count as usize == self.changes.len()
214                }
215                SecretsResult::Failed { failed_index, .. } => {
216                    (*failed_index as usize) < self.changes.len()
217                }
218            };
219            if !valid {
220                return Err(wire::WireError::InvalidRecord);
221            }
222            Ok(result)
223        })
224    }
225}
226
227macro_rules! v2_request {
228    ($request:ty, $response:ty, $request_name:literal, $response_name:literal) => {
229        impl Request<ControlProtocol> for $request {
230            type Response = $response;
231            type Error = ControlClientError;
232
233            fn message(&self) -> ControlClientResult<EncodedMessage> {
234                prepared($request_name, &self.0)
235            }
236
237            fn decode(&self, response: Message) -> ControlClientResult<Self::Response> {
238                checked(response, 2, $response_name)
239            }
240        }
241    };
242}
243
244v2_request!(
245    CreateCheckpoint,
246    CheckpointResult,
247    "control.checkpoint.create",
248    "control.checkpoint.result"
249);
250v2_request!(
251    CreateDiskCheckpoint,
252    DiskCheckpointState,
253    "control.disk.checkpoint.create",
254    "control.disk.checkpoint.result"
255);
256v2_request!(
257    CreateBranch,
258    BranchResult,
259    "control.branch.create",
260    "control.branch.result"
261);
262v2_request!(
263    PauseRuntime,
264    PauseState,
265    "control.pause",
266    "control.pause.state"
267);
268v2_request!(
269    GrowRootDisk,
270    RootDiskState,
271    "control.root-disk.grow",
272    "control.root-disk.state"
273);
274v2_request!(
275    CompactDisks,
276    microsandbox_protocol::control::DiskCompactionResult,
277    "control.disk.compact",
278    "control.disk.compact.result"
279);
280
281impl Request<ControlProtocol> for ResumeRuntime {
282    type Response = PauseState;
283    type Error = ControlClientError;
284
285    fn message(&self) -> ControlClientResult<EncodedMessage> {
286        prepared("control.resume", &Empty {})
287    }
288
289    fn decode(&self, response: Message) -> ControlClientResult<Self::Response> {
290        checked(response, 2, "control.pause.state")
291    }
292}
293
294impl Request<ControlProtocol> for GetPauseState {
295    type Response = PauseState;
296    type Error = ControlClientError;
297
298    fn message(&self) -> ControlClientResult<EncodedMessage> {
299        prepared("control.pause.state", &Empty {})
300    }
301
302    fn decode(&self, response: Message) -> ControlClientResult<Self::Response> {
303        checked(response, 2, "control.pause.state")
304    }
305}
306
307macro_rules! compatible_v2 {
308    ($request:ty, $json:expr, $decode:expr) => {
309        impl CompatibleControlRequest for $request {
310            fn min_generation(&self) -> u8 {
311                2
312            }
313
314            fn compatibility_json_bytes(&self) -> ControlClientResult<Zeroizing<Vec<u8>>> {
315                let value = ($json)(self);
316                Ok(Zeroizing::new(serde_json::to_vec(&value).map_err(
317                    |_| {
318                        microsandbox_protocol_client::ClientError::new(
319                            microsandbox_protocol_client::ErrorKind::Encode,
320                        )
321                    },
322                )?))
323            }
324
325            fn decode_compatibility_json(
326                &self,
327                reply: JsonReply,
328            ) -> ControlClientResult<Self::Response> {
329                ($decode)(reply)
330            }
331        }
332    };
333}
334
335compatible_v2!(
336    GetRuntimeCapabilities,
337    |_: &GetRuntimeCapabilities| serde_json::json!({
338        "op": "capabilities",
339    }),
340    |reply| decode_json_field(reply, "capabilities")
341);
342
343compatible_v2!(
344    CreateCheckpoint,
345    |request: &CreateCheckpoint| serde_json::json!({
346        "op": "checkpoint_create",
347        "guest_flush": request.0.guest_flush,
348        "record_integrity": request.0.record_integrity,
349        "checkpoint_id": request.0.checkpoint_id,
350        "intent": request.0.intent,
351    }),
352    decode_checkpoint_json
353);
354compatible_v2!(
355    CreateDiskCheckpoint,
356    |request: &CreateDiskCheckpoint| serde_json::json!({
357        "op": "disk_checkpoint_create",
358        "guest_flush": request.0.guest_flush,
359        "checkpoint_id": request.0.checkpoint_id,
360    }),
361    |reply| decode_json_field(reply, "disk_checkpoint")
362);
363compatible_v2!(
364    CreateBranch,
365    |request: &CreateBranch| serde_json::json!({
366        "op": "branch_create",
367        "guest_flush": request.0.guest_flush,
368        "record_integrity": request.0.record_integrity,
369        "branch_id": request.0.branch_id,
370        "child_name": request.0.child_name,
371        "memory_cache_dir": request.0.memory_cache_dir,
372    }),
373    decode_branch_json
374);
375compatible_v2!(
376    PauseRuntime,
377    |request: &PauseRuntime| match request.0.guest_flush {
378        Some(policy) => serde_json::json!({"op": "pause_with_guest_flush", "guest_flush": policy}),
379        None => serde_json::json!({"op": "pause"}),
380    },
381    |reply| decode_json_field(reply, "pause")
382);
383compatible_v2!(
384    GrowRootDisk,
385    |request: &GrowRootDisk| serde_json::json!({
386        "op": "root_disk_grow", "size_bytes": request.0.size_bytes,
387    }),
388    |reply| decode_json_field(reply, "root_disk")
389);
390compatible_v2!(
391    CompactDisks,
392    |request: &CompactDisks| serde_json::json!({
393        "op": "disk_compact",
394        "target": request.0.target,
395        "layers": request.0.layers,
396        "dry_run": request.0.dry_run,
397    }),
398    |reply| decode_json_field(reply, "compaction")
399);
400
401impl CompatibleControlRequest for ResumeRuntime {
402    fn min_generation(&self) -> u8 {
403        2
404    }
405    fn compatibility_json_bytes(&self) -> ControlClientResult<Zeroizing<Vec<u8>>> {
406        json_operation("resume")
407    }
408    fn decode_compatibility_json(&self, reply: JsonReply) -> ControlClientResult<Self::Response> {
409        decode_json_field(reply, "pause")
410    }
411}
412
413impl CompatibleControlRequest for GetPauseState {
414    fn min_generation(&self) -> u8 {
415        2
416    }
417    fn compatibility_json_bytes(&self) -> ControlClientResult<Zeroizing<Vec<u8>>> {
418        json_operation("pause_state")
419    }
420    fn decode_compatibility_json(&self, reply: JsonReply) -> ControlClientResult<Self::Response> {
421        decode_json_field(reply, "pause")
422    }
423}
424
425//--------------------------------------------------------------------------------------------------
426// Functions
427//--------------------------------------------------------------------------------------------------
428
429fn prepared(name: &str, payload: &impl Serialize) -> ControlClientResult<EncodedMessage> {
430    Ok(EncodedMessage::new(name, wire::encode(payload)?))
431}
432
433fn checked<T: DeserializeOwned>(
434    response: Message,
435    minimum: u8,
436    expected: &str,
437) -> ControlClientResult<T> {
438    checked_with(response, minimum, expected, wire::decode_record)
439}
440
441fn checked_with<T>(
442    response: Message,
443    minimum: u8,
444    expected: &str,
445    decode: impl FnOnce(&[u8]) -> Result<T, wire::WireError>,
446) -> ControlClientResult<T> {
447    if !(minimum..=microsandbox_protocol::control::CONTROL_GENERATION).contains(&response.v)
448        || response.id == 0
449        || response.flags != 1
450    {
451        return Err(ControlClientError::InvalidResponse {
452            response: Box::new(response),
453        });
454    }
455    if response.t == "control.error" {
456        let Ok(error) = wire::decode_record::<ControlError>(&response.p) else {
457            return Err(ControlClientError::InvalidResponse {
458                response: Box::new(response),
459            });
460        };
461        return Err(ControlClientError::Peer {
462            error,
463            response: Box::new(response),
464        });
465    }
466    if response.t != expected {
467        return Err(ControlClientError::InvalidResponse {
468            response: Box::new(response),
469        });
470    }
471    decode(&response.p).map_err(|_| ControlClientError::InvalidResponse {
472        response: Box::new(response),
473    })
474}
475
476fn json_operation(operation: &str) -> ControlClientResult<Zeroizing<Vec<u8>>> {
477    Ok(Zeroizing::new(
478        serde_json::to_vec(&serde_json::json!({"op": operation})).map_err(|_| {
479            microsandbox_protocol_client::ClientError::new(
480                microsandbox_protocol_client::ErrorKind::Encode,
481            )
482        })?,
483    ))
484}
485
486fn decode_json_field<T: DeserializeOwned>(reply: JsonReply, field: &str) -> ControlClientResult<T> {
487    let decoded = serde_json::from_slice::<serde_json::Value>(reply.raw())
488        .ok()
489        .and_then(|value| value.get(field).cloned())
490        .and_then(|value| serde_json::from_value(value).ok());
491    reply.checked(|_| decoded)
492}
493
494fn decode_branch_json(reply: JsonReply) -> ControlClientResult<BranchResult> {
495    let path = serde_json::from_slice::<serde_json::Value>(reply.raw())
496        .ok()
497        .and_then(|value| value.get("branch").cloned())
498        .and_then(|value| serde_json::from_value(value).ok());
499    reply.checked(|_| path.map(|path| BranchResult { path }))
500}
501
502fn decode_checkpoint_json(reply: JsonReply) -> ControlClientResult<CheckpointResult> {
503    #[derive(serde::Deserialize)]
504    struct Response {
505        ok: bool,
506        #[serde(default)]
507        error: Option<String>,
508        #[serde(default)]
509        checkpoint: Option<microsandbox_protocol::control::CheckpointState>,
510    }
511
512    let decoded = serde_json::from_slice::<Response>(reply.raw()).ok();
513    if let Some(response) = decoded
514        && let Some(checkpoint) = response.checkpoint
515    {
516        return Ok(CheckpointResult {
517            checkpoint: Some(checkpoint),
518            recovery_error: (!response.ok).then_some(response.error).flatten(),
519        });
520    }
521    reply.checked(|_| None)
522}