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
80#[derive(Debug, Clone)]
82pub struct ManageJob(pub microsandbox_protocol::jobs::JobRequest);
83
84impl SetMemoryTarget {
89 pub fn new(size: impl Into<Mebibytes>) -> Self {
95 Self {
96 total_mib: u64::from(size.into().as_u32()),
97 }
98 }
99}
100
101impl SetCpuTarget {
102 pub fn new(online: u32) -> Self {
104 Self { online }
105 }
106}
107
108impl UpdateSecrets {
109 pub fn new(changes: Vec<SecretChange>) -> Self {
111 Self { changes }
112 }
113}
114
115impl Request<ControlProtocol> for ManageJob {
120 type Response = microsandbox_protocol::jobs::JobResponse;
121 type Error = ControlClientError;
122 fn message(&self) -> ControlClientResult<EncodedMessage> {
123 prepared(microsandbox_protocol::jobs::JOB_REQUEST, &self.0)
124 }
125 fn decode(&self, response: Message) -> ControlClientResult<Self::Response> {
126 checked(response, 2, microsandbox_protocol::jobs::JOB_RESPONSE)
127 }
128}
129
130impl CompatibleControlRequest for ManageJob {
131 fn min_generation(&self) -> u8 {
132 2
133 }
134 fn compatibility_json_bytes(&self) -> ControlClientResult<Zeroizing<Vec<u8>>> {
135 Err(ControlClientError::UnsupportedMode)
136 }
137 fn decode_compatibility_json(&self, _reply: JsonReply) -> ControlClientResult<Self::Response> {
138 Err(ControlClientError::UnsupportedMode)
139 }
140}
141
142impl Request<ControlProtocol> for GetCapabilities {
143 type Response = Capabilities;
144 type Error = ControlClientError;
145 fn message(&self) -> ControlClientResult<EncodedMessage> {
146 prepared("control.capabilities", &Empty {})
147 }
148 fn decode(&self, response: Message) -> ControlClientResult<Self::Response> {
149 checked(response, 1, "control.capabilities.result")
150 }
151}
152
153impl Request<ControlProtocol> for GetRuntimeCapabilities {
154 type Response = RuntimeCapabilities;
155 type Error = ControlClientError;
156 fn message(&self) -> ControlClientResult<EncodedMessage> {
157 prepared("control.capabilities", &Empty {})
158 }
159 fn decode(&self, response: Message) -> ControlClientResult<Self::Response> {
160 checked(response, 2, "control.capabilities.result")
161 }
162}
163
164impl Request<ControlProtocol> for GetMemoryState {
165 type Response = MemoryState;
166 type Error = ControlClientError;
167 fn message(&self) -> ControlClientResult<EncodedMessage> {
168 prepared("control.memory.state", &Empty {})
169 }
170 fn decode(&self, response: Message) -> ControlClientResult<Self::Response> {
171 checked(response, 1, "control.memory.state")
172 }
173}
174
175impl Request<ControlProtocol> for SetMemoryTarget {
176 type Response = MemoryState;
177 type Error = ControlClientError;
178 fn message(&self) -> ControlClientResult<EncodedMessage> {
179 prepared(
180 "control.memory.target",
181 &MemoryTarget {
182 total_mib: self.total_mib,
183 },
184 )
185 }
186 fn decode(&self, response: Message) -> ControlClientResult<Self::Response> {
187 checked(response, 1, "control.memory.state")
188 }
189}
190
191impl Request<ControlProtocol> for GetCpuState {
192 type Response = CpuState;
193 type Error = ControlClientError;
194 fn message(&self) -> ControlClientResult<EncodedMessage> {
195 prepared("control.cpu.state", &Empty {})
196 }
197 fn decode(&self, response: Message) -> ControlClientResult<Self::Response> {
198 checked(response, 1, "control.cpu.state")
199 }
200}
201
202impl Request<ControlProtocol> for SetCpuTarget {
203 type Response = CpuState;
204 type Error = ControlClientError;
205 fn message(&self) -> ControlClientResult<EncodedMessage> {
206 prepared(
207 "control.cpu.target",
208 &CpuTarget {
209 online: self.online,
210 },
211 )
212 }
213 fn decode(&self, response: Message) -> ControlClientResult<Self::Response> {
214 checked(response, 1, "control.cpu.state")
215 }
216}
217
218impl Request<ControlProtocol> for UpdateSecrets {
219 type Response = SecretsResult;
220 type Error = ControlClientError;
221 fn message(&self) -> ControlClientResult<EncodedMessage> {
222 #[derive(Serialize)]
223 struct Payload<'a> {
224 changes: &'a [SecretChange],
225 }
226 prepared(
227 "control.secrets.update",
228 &Payload {
229 changes: &self.changes,
230 },
231 )
232 }
233 fn decode(&self, response: Message) -> ControlClientResult<Self::Response> {
234 checked_with(response, 1, "control.secrets.result", |bytes| {
235 let result = SecretsResult::decode(bytes)?;
236 let valid = match &result {
239 SecretsResult::Complete { applied_count } => {
240 *applied_count as usize == self.changes.len()
241 }
242 SecretsResult::Failed { failed_index, .. } => {
243 (*failed_index as usize) < self.changes.len()
244 }
245 };
246 if !valid {
247 return Err(wire::WireError::InvalidRecord);
248 }
249 Ok(result)
250 })
251 }
252}
253
254macro_rules! v2_request {
255 ($request:ty, $response:ty, $request_name:literal, $response_name:literal) => {
256 impl Request<ControlProtocol> for $request {
257 type Response = $response;
258 type Error = ControlClientError;
259
260 fn message(&self) -> ControlClientResult<EncodedMessage> {
261 prepared($request_name, &self.0)
262 }
263
264 fn decode(&self, response: Message) -> ControlClientResult<Self::Response> {
265 checked(response, 2, $response_name)
266 }
267 }
268 };
269}
270
271v2_request!(
272 CreateCheckpoint,
273 CheckpointResult,
274 "control.checkpoint.create",
275 "control.checkpoint.result"
276);
277v2_request!(
278 CreateDiskCheckpoint,
279 DiskCheckpointState,
280 "control.disk.checkpoint.create",
281 "control.disk.checkpoint.result"
282);
283v2_request!(
284 CreateBranch,
285 BranchResult,
286 "control.branch.create",
287 "control.branch.result"
288);
289v2_request!(
290 PauseRuntime,
291 PauseState,
292 "control.pause",
293 "control.pause.state"
294);
295v2_request!(
296 GrowRootDisk,
297 RootDiskState,
298 "control.root-disk.grow",
299 "control.root-disk.state"
300);
301v2_request!(
302 CompactDisks,
303 microsandbox_protocol::control::DiskCompactionResult,
304 "control.disk.compact",
305 "control.disk.compact.result"
306);
307
308impl Request<ControlProtocol> for ResumeRuntime {
309 type Response = PauseState;
310 type Error = ControlClientError;
311
312 fn message(&self) -> ControlClientResult<EncodedMessage> {
313 prepared("control.resume", &Empty {})
314 }
315
316 fn decode(&self, response: Message) -> ControlClientResult<Self::Response> {
317 checked(response, 2, "control.pause.state")
318 }
319}
320
321impl Request<ControlProtocol> for GetPauseState {
322 type Response = PauseState;
323 type Error = ControlClientError;
324
325 fn message(&self) -> ControlClientResult<EncodedMessage> {
326 prepared("control.pause.state", &Empty {})
327 }
328
329 fn decode(&self, response: Message) -> ControlClientResult<Self::Response> {
330 checked(response, 2, "control.pause.state")
331 }
332}
333
334macro_rules! compatible_v2 {
335 ($request:ty, $json:expr, $decode:expr) => {
336 impl CompatibleControlRequest for $request {
337 fn min_generation(&self) -> u8 {
338 2
339 }
340
341 fn compatibility_json_bytes(&self) -> ControlClientResult<Zeroizing<Vec<u8>>> {
342 let value = ($json)(self);
343 Ok(Zeroizing::new(serde_json::to_vec(&value).map_err(
344 |_| {
345 microsandbox_protocol_client::ClientError::new(
346 microsandbox_protocol_client::ErrorKind::Encode,
347 )
348 },
349 )?))
350 }
351
352 fn decode_compatibility_json(
353 &self,
354 reply: JsonReply,
355 ) -> ControlClientResult<Self::Response> {
356 ($decode)(reply)
357 }
358 }
359 };
360}
361
362compatible_v2!(
363 GetRuntimeCapabilities,
364 |_: &GetRuntimeCapabilities| serde_json::json!({
365 "op": "capabilities",
366 }),
367 |reply| decode_json_field(reply, "capabilities")
368);
369
370compatible_v2!(
371 CreateCheckpoint,
372 |request: &CreateCheckpoint| serde_json::json!({
373 "op": "checkpoint_create",
374 "guest_flush": request.0.guest_flush,
375 "record_integrity": request.0.record_integrity,
376 "checkpoint_id": request.0.checkpoint_id,
377 "intent": request.0.intent,
378 }),
379 decode_checkpoint_json
380);
381compatible_v2!(
382 CreateDiskCheckpoint,
383 |request: &CreateDiskCheckpoint| serde_json::json!({
384 "op": "disk_checkpoint_create",
385 "guest_flush": request.0.guest_flush,
386 "checkpoint_id": request.0.checkpoint_id,
387 }),
388 |reply| decode_json_field(reply, "disk_checkpoint")
389);
390compatible_v2!(
391 CreateBranch,
392 |request: &CreateBranch| serde_json::json!({
393 "op": "branch_create",
394 "guest_flush": request.0.guest_flush,
395 "record_integrity": request.0.record_integrity,
396 "branch_id": request.0.branch_id,
397 "child_name": request.0.child_name,
398 "memory_cache_dir": request.0.memory_cache_dir,
399 }),
400 decode_branch_json
401);
402compatible_v2!(
403 PauseRuntime,
404 |request: &PauseRuntime| match request.0.guest_flush {
405 Some(policy) => serde_json::json!({"op": "pause_with_guest_flush", "guest_flush": policy}),
406 None => serde_json::json!({"op": "pause"}),
407 },
408 |reply| decode_json_field(reply, "pause")
409);
410compatible_v2!(
411 GrowRootDisk,
412 |request: &GrowRootDisk| serde_json::json!({
413 "op": "root_disk_grow", "size_bytes": request.0.size_bytes,
414 }),
415 |reply| decode_json_field(reply, "root_disk")
416);
417compatible_v2!(
418 CompactDisks,
419 |request: &CompactDisks| serde_json::json!({
420 "op": "disk_compact",
421 "target": request.0.target,
422 "layers": request.0.layers,
423 "dry_run": request.0.dry_run,
424 }),
425 |reply| decode_json_field(reply, "compaction")
426);
427
428impl CompatibleControlRequest for ResumeRuntime {
429 fn min_generation(&self) -> u8 {
430 2
431 }
432 fn compatibility_json_bytes(&self) -> ControlClientResult<Zeroizing<Vec<u8>>> {
433 json_operation("resume")
434 }
435 fn decode_compatibility_json(&self, reply: JsonReply) -> ControlClientResult<Self::Response> {
436 decode_json_field(reply, "pause")
437 }
438}
439
440impl CompatibleControlRequest for GetPauseState {
441 fn min_generation(&self) -> u8 {
442 2
443 }
444 fn compatibility_json_bytes(&self) -> ControlClientResult<Zeroizing<Vec<u8>>> {
445 json_operation("pause_state")
446 }
447 fn decode_compatibility_json(&self, reply: JsonReply) -> ControlClientResult<Self::Response> {
448 decode_json_field(reply, "pause")
449 }
450}
451
452fn prepared(name: &str, payload: &impl Serialize) -> ControlClientResult<EncodedMessage> {
457 Ok(EncodedMessage::new(name, wire::encode(payload)?))
458}
459
460fn checked<T: DeserializeOwned>(
461 response: Message,
462 minimum: u8,
463 expected: &str,
464) -> ControlClientResult<T> {
465 checked_with(response, minimum, expected, wire::decode_record)
466}
467
468fn checked_with<T>(
469 response: Message,
470 minimum: u8,
471 expected: &str,
472 decode: impl FnOnce(&[u8]) -> Result<T, wire::WireError>,
473) -> ControlClientResult<T> {
474 if !(minimum..=microsandbox_protocol::control::CONTROL_GENERATION).contains(&response.v)
475 || response.id == 0
476 || response.flags != 1
477 {
478 return Err(ControlClientError::InvalidResponse {
479 response: Box::new(response),
480 });
481 }
482 if response.t == "control.error" {
483 let Ok(error) = wire::decode_record::<ControlError>(&response.p) else {
484 return Err(ControlClientError::InvalidResponse {
485 response: Box::new(response),
486 });
487 };
488 return Err(ControlClientError::Peer {
489 error,
490 response: Box::new(response),
491 });
492 }
493 if response.t != expected {
494 return Err(ControlClientError::InvalidResponse {
495 response: Box::new(response),
496 });
497 }
498 decode(&response.p).map_err(|_| ControlClientError::InvalidResponse {
499 response: Box::new(response),
500 })
501}
502
503fn json_operation(operation: &str) -> ControlClientResult<Zeroizing<Vec<u8>>> {
504 Ok(Zeroizing::new(
505 serde_json::to_vec(&serde_json::json!({"op": operation})).map_err(|_| {
506 microsandbox_protocol_client::ClientError::new(
507 microsandbox_protocol_client::ErrorKind::Encode,
508 )
509 })?,
510 ))
511}
512
513fn decode_json_field<T: DeserializeOwned>(reply: JsonReply, field: &str) -> ControlClientResult<T> {
514 let decoded = serde_json::from_slice::<serde_json::Value>(reply.raw())
515 .ok()
516 .and_then(|value| value.get(field).cloned())
517 .and_then(|value| serde_json::from_value(value).ok());
518 reply.checked(|_| decoded)
519}
520
521fn decode_branch_json(reply: JsonReply) -> ControlClientResult<BranchResult> {
522 let path = serde_json::from_slice::<serde_json::Value>(reply.raw())
523 .ok()
524 .and_then(|value| value.get("branch").cloned())
525 .and_then(|value| serde_json::from_value(value).ok());
526 reply.checked(|_| path.map(|path| BranchResult { path }))
527}
528
529fn decode_checkpoint_json(reply: JsonReply) -> ControlClientResult<CheckpointResult> {
530 #[derive(serde::Deserialize)]
531 struct Response {
532 ok: bool,
533 #[serde(default)]
534 error: Option<String>,
535 #[serde(default)]
536 checkpoint: Option<microsandbox_protocol::control::CheckpointState>,
537 }
538
539 let decoded = serde_json::from_slice::<Response>(reply.raw()).ok();
540 if let Some(response) = decoded
541 && let Some(checkpoint) = response.checkpoint
542 {
543 return Ok(CheckpointResult {
544 checkpoint: Some(checkpoint),
545 recovery_error: (!response.ok).then_some(response.error).flatten(),
546 });
547 }
548 reply.checked(|_| None)
549}