1use 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#[derive(Debug, Clone, Copy, Default)]
27pub struct GetCapabilities;
28#[derive(Debug, Clone, Copy, Default)]
30pub struct GetRuntimeCapabilities;
31#[derive(Debug, Clone, Copy, Default)]
33pub struct GetMemoryState;
34#[derive(Debug, Clone, Copy)]
36pub struct SetMemoryTarget {
37 pub total_mib: u64,
39}
40#[derive(Debug, Clone, Copy, Default)]
42pub struct GetCpuState;
43#[derive(Debug, Clone, Copy)]
45pub struct SetCpuTarget {
46 pub online: u32,
48}
49#[derive(Debug, Clone)]
51pub struct UpdateSecrets {
52 pub changes: Vec<SecretChange>,
54}
55#[derive(Debug, Clone)]
57pub struct CreateCheckpoint(pub CheckpointCreate);
58#[derive(Debug, Clone)]
60pub struct CreateDiskCheckpoint(pub DiskCheckpointCreate);
61#[derive(Debug, Clone)]
63pub struct CreateBranch(pub BranchCreate);
64#[derive(Debug, Clone, Copy, Default)]
66pub struct PauseRuntime(pub Pause);
67#[derive(Debug, Clone, Copy, Default)]
69pub struct ResumeRuntime;
70#[derive(Debug, Clone, Copy, Default)]
72pub struct GetPauseState;
73#[derive(Debug, Clone, Copy)]
75pub struct GrowRootDisk(pub RootDiskGrow);
76#[derive(Debug, Clone)]
78pub struct CompactDisks(pub DiskCompact);
79
80impl SetMemoryTarget {
85 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 pub fn new(online: u32) -> Self {
100 Self { online }
101 }
102}
103
104impl UpdateSecrets {
105 pub fn new(changes: Vec<SecretChange>) -> Self {
107 Self { changes }
108 }
109}
110
111impl 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 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
425fn 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}