microsandbox_protocol/control/
message.rs1use 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#[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 #[strum(serialize = "control.hello")]
28 Hello,
29 #[strum(serialize = "control.welcome")]
31 Welcome,
32 #[strum(serialize = "control.capabilities")]
34 Capabilities,
35 #[strum(serialize = "control.capabilities.result")]
37 CapabilitiesResult,
38 #[strum(serialize = "control.memory.state")]
40 MemoryState,
41 #[strum(serialize = "control.memory.target")]
43 MemoryTarget,
44 #[strum(serialize = "control.cpu.state")]
46 CpuState,
47 #[strum(serialize = "control.cpu.target")]
49 CpuTarget,
50 #[strum(serialize = "control.secrets.update")]
52 SecretsUpdate,
53 #[strum(serialize = "control.secrets.result")]
55 SecretsResult,
56 #[strum(serialize = "control.error")]
58 Error,
59}
60
61impl ControlMessageType {
66 pub fn as_str(self) -> &'static str {
68 self.into()
69 }
70
71 pub fn from_wire_str(name: &str) -> Option<Self> {
73 name.parse().ok()
74 }
75
76 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 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 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 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 #[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 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
209impl 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}