1use crate::frames::{reject, DispatchResult};
2use secure_exec_sidecar_protocol::protocol::{
3 AuthenticateRequest, BootstrapRootFilesystemRequest, CloseStdinRequest, ConfigureVmRequest,
4 CreateLayerRequest, CreateOverlayRequest, CreateVmRequest, DisposeVmRequest, ExecuteRequest,
5 ExportSnapshotRequest, ExtEnvelope, FindBoundUdpRequest, FindListenerRequest,
6 GetProcessSnapshotRequest, GetResourceSnapshotRequest, GetSignalStateRequest,
7 GetZombieTimerCountRequest, GuestFilesystemCallRequest, GuestKernelCallRequest,
8 ImportSnapshotRequest, KillProcessRequest, LinkPackageRequest, OpenSessionRequest,
9 OwnershipScope, RegisterHostCallbacksRequest, RequestFrame, RequestPayload, ResizePtyRequest,
10 SealLayerRequest, SnapshotRootFilesystemRequest, VmFetchRequest, WriteStdinRequest,
11};
12use secure_exec_sidecar_protocol::wire as generated_wire;
13
14pub const UNSUPPORTED_HOST_CALLBACK_DIRECTION_CODE: &str = "unsupported_direction";
15pub const UNSUPPORTED_HOST_CALLBACK_DIRECTION_MESSAGE: &str =
16 "host callback request categories are sidecar-to-host only in this scaffold";
17
18#[derive(Debug, Clone, Copy, PartialEq, Eq)]
19pub enum RequestDispatchMode {
20 Immediate,
21 Async,
22}
23
24#[allow(clippy::large_enum_variant)]
27#[derive(Debug, Clone)]
28pub enum RequestRoute {
29 Authenticate(AuthenticateRequest),
30 OpenSession(OpenSessionRequest),
31 CreateVm(CreateVmRequest),
32 DisposeVm(DisposeVmRequest),
33 BootstrapRootFilesystem(BootstrapRootFilesystemRequest),
34 ConfigureVm(ConfigureVmRequest),
35 RegisterHostCallbacks(RegisterHostCallbacksRequest),
36 CreateLayer(CreateLayerRequest),
37 SealLayer(SealLayerRequest),
38 ImportSnapshot(ImportSnapshotRequest),
39 ExportSnapshot(ExportSnapshotRequest),
40 CreateOverlay(CreateOverlayRequest),
41 GuestFilesystemCall(GuestFilesystemCallRequest),
42 GuestKernelCall(GuestKernelCallRequest),
43 SnapshotRootFilesystem(SnapshotRootFilesystemRequest),
44 Execute(ExecuteRequest),
45 WriteStdin(WriteStdinRequest),
46 ResizePty(ResizePtyRequest),
47 CloseStdin(CloseStdinRequest),
48 KillProcess(KillProcessRequest),
49 GetProcessSnapshot(GetProcessSnapshotRequest),
50 GetResourceSnapshot(GetResourceSnapshotRequest),
51 FindListener(FindListenerRequest),
52 FindBoundUdp(FindBoundUdpRequest),
53 VmFetch(VmFetchRequest),
54 GetSignalState(GetSignalStateRequest),
55 GetZombieTimerCount(GetZombieTimerCountRequest),
56 LinkPackage(LinkPackageRequest),
57 Ext(ExtEnvelope),
58 UnsupportedHostCallbackDirection,
59}
60
61#[derive(Debug, Clone, Copy, PartialEq, Eq)]
62pub enum BlockingExtensionInterrupt<'a> {
63 ExtensionPayload(&'a [u8]),
64 KillProcess,
65}
66
67pub fn route_request_payload(request: &RequestFrame) -> RequestRoute {
68 match request.payload.clone() {
69 RequestPayload::Authenticate(payload) => RequestRoute::Authenticate(payload),
70 RequestPayload::OpenSession(payload) => RequestRoute::OpenSession(payload),
71 RequestPayload::CreateVm(payload) => RequestRoute::CreateVm(payload),
72 RequestPayload::DisposeVm(payload) => RequestRoute::DisposeVm(payload),
73 RequestPayload::BootstrapRootFilesystem(payload) => {
74 RequestRoute::BootstrapRootFilesystem(payload)
75 }
76 RequestPayload::ConfigureVm(payload) => RequestRoute::ConfigureVm(payload),
77 RequestPayload::RegisterHostCallbacks(payload) => {
78 RequestRoute::RegisterHostCallbacks(payload)
79 }
80 RequestPayload::CreateLayer(payload) => RequestRoute::CreateLayer(payload),
81 RequestPayload::SealLayer(payload) => RequestRoute::SealLayer(payload),
82 RequestPayload::ImportSnapshot(payload) => RequestRoute::ImportSnapshot(payload),
83 RequestPayload::ExportSnapshot(payload) => RequestRoute::ExportSnapshot(payload),
84 RequestPayload::CreateOverlay(payload) => RequestRoute::CreateOverlay(payload),
85 RequestPayload::GuestFilesystemCall(payload) => RequestRoute::GuestFilesystemCall(payload),
86 RequestPayload::GuestKernelCall(payload) => RequestRoute::GuestKernelCall(payload),
87 RequestPayload::SnapshotRootFilesystem(payload) => {
88 RequestRoute::SnapshotRootFilesystem(payload)
89 }
90 RequestPayload::Execute(payload) => RequestRoute::Execute(payload),
91 RequestPayload::WriteStdin(payload) => RequestRoute::WriteStdin(payload),
92 RequestPayload::ResizePty(payload) => RequestRoute::ResizePty(payload),
93 RequestPayload::CloseStdin(payload) => RequestRoute::CloseStdin(payload),
94 RequestPayload::KillProcess(payload) => RequestRoute::KillProcess(payload),
95 RequestPayload::GetProcessSnapshot(payload) => RequestRoute::GetProcessSnapshot(payload),
96 RequestPayload::GetResourceSnapshot(payload) => RequestRoute::GetResourceSnapshot(payload),
97 RequestPayload::FindListener(payload) => RequestRoute::FindListener(payload),
98 RequestPayload::FindBoundUdp(payload) => RequestRoute::FindBoundUdp(payload),
99 RequestPayload::VmFetch(payload) => RequestRoute::VmFetch(payload),
100 RequestPayload::GetSignalState(payload) => RequestRoute::GetSignalState(payload),
101 RequestPayload::GetZombieTimerCount(payload) => RequestRoute::GetZombieTimerCount(payload),
102 RequestPayload::LinkPackage(payload) => RequestRoute::LinkPackage(payload),
103 RequestPayload::HostFilesystemCall(_)
104 | RequestPayload::PersistenceLoad(_)
105 | RequestPayload::PersistenceFlush(_) => RequestRoute::UnsupportedHostCallbackDirection,
106 RequestPayload::Ext(payload) => RequestRoute::Ext(payload),
107 }
108}
109
110pub fn generated_wire_blocking_extension_interrupt<'a>(
111 active_request: &generated_wire::RequestFrame,
112 blocking_namespace: &str,
113 interrupting_request: &'a generated_wire::RequestFrame,
114) -> Option<BlockingExtensionInterrupt<'a>> {
115 if interrupting_request.ownership != active_request.ownership {
116 return None;
117 }
118
119 match &interrupting_request.payload {
120 generated_wire::RequestPayload::ExtEnvelope(envelope)
121 if envelope.namespace == blocking_namespace =>
122 {
123 Some(BlockingExtensionInterrupt::ExtensionPayload(
124 &envelope.payload,
125 ))
126 }
127 generated_wire::RequestPayload::ExtEnvelope(_) => None,
128 generated_wire::RequestPayload::KillProcessRequest(_) => {
129 Some(BlockingExtensionInterrupt::KillProcess)
130 }
131 _ => None,
132 }
133}
134
135pub fn request_dispatch_mode(request: &RequestFrame) -> RequestDispatchMode {
136 match request.payload {
137 RequestPayload::DisposeVm(_) | RequestPayload::Ext(_) => RequestDispatchMode::Async,
138 RequestPayload::Authenticate(_)
139 | RequestPayload::OpenSession(_)
140 | RequestPayload::CreateVm(_)
141 | RequestPayload::BootstrapRootFilesystem(_)
142 | RequestPayload::ConfigureVm(_)
143 | RequestPayload::RegisterHostCallbacks(_)
144 | RequestPayload::CreateLayer(_)
145 | RequestPayload::SealLayer(_)
146 | RequestPayload::ImportSnapshot(_)
147 | RequestPayload::ExportSnapshot(_)
148 | RequestPayload::CreateOverlay(_)
149 | RequestPayload::GuestFilesystemCall(_)
150 | RequestPayload::GuestKernelCall(_)
151 | RequestPayload::SnapshotRootFilesystem(_)
152 | RequestPayload::Execute(_)
153 | RequestPayload::WriteStdin(_)
154 | RequestPayload::ResizePty(_)
155 | RequestPayload::CloseStdin(_)
156 | RequestPayload::KillProcess(_)
157 | RequestPayload::GetProcessSnapshot(_)
158 | RequestPayload::GetResourceSnapshot(_)
159 | RequestPayload::FindListener(_)
160 | RequestPayload::FindBoundUdp(_)
161 | RequestPayload::VmFetch(_)
162 | RequestPayload::GetSignalState(_)
163 | RequestPayload::GetZombieTimerCount(_)
164 | RequestPayload::LinkPackage(_)
165 | RequestPayload::HostFilesystemCall(_)
166 | RequestPayload::PersistenceLoad(_)
167 | RequestPayload::PersistenceFlush(_) => RequestDispatchMode::Immediate,
168 }
169}
170
171pub fn request_is_unsupported_host_callback_direction(request: &RequestFrame) -> bool {
172 matches!(
173 request.payload,
174 RequestPayload::HostFilesystemCall(_)
175 | RequestPayload::PersistenceLoad(_)
176 | RequestPayload::PersistenceFlush(_)
177 )
178}
179
180pub fn unsupported_host_callback_direction_dispatch(request: &RequestFrame) -> DispatchResult {
181 debug_assert!(request_is_unsupported_host_callback_direction(request));
182 DispatchResult {
183 response: reject(
184 request,
185 UNSUPPORTED_HOST_CALLBACK_DIRECTION_CODE,
186 UNSUPPORTED_HOST_CALLBACK_DIRECTION_MESSAGE,
187 ),
188 events: Vec::new(),
189 }
190}
191
192pub fn connection_id_of(ownership: &OwnershipScope) -> Option<String> {
193 match ownership {
194 OwnershipScope::ConnectionOwnership(ownership) => Some(ownership.connection_id.clone()),
195 OwnershipScope::SessionOwnership(ownership) => Some(ownership.connection_id.clone()),
196 OwnershipScope::VmOwnership(ownership) => Some(ownership.connection_id.clone()),
197 }
198}
199
200pub fn session_scope_of(ownership: &OwnershipScope) -> Option<(String, String)> {
201 match ownership {
202 OwnershipScope::SessionOwnership(ownership) => Some((
203 ownership.connection_id.clone(),
204 ownership.session_id.clone(),
205 )),
206 OwnershipScope::VmOwnership(ownership) => Some((
207 ownership.connection_id.clone(),
208 ownership.session_id.clone(),
209 )),
210 OwnershipScope::ConnectionOwnership(_) => None,
211 }
212}
213
214pub fn vm_id_of(ownership: &OwnershipScope) -> Option<String> {
215 match ownership {
216 OwnershipScope::VmOwnership(ownership) => Some(ownership.vm_id.clone()),
217 OwnershipScope::ConnectionOwnership(_) | OwnershipScope::SessionOwnership(_) => None,
218 }
219}
220
221#[cfg(test)]
222mod tests {
223 use super::*;
224 use secure_exec_sidecar_protocol::protocol::{
225 AuthenticateRequest, ExtEnvelope, FilesystemOperation, HostFilesystemCallRequest,
226 OwnershipScope, PersistenceFlushRequest, PersistenceLoadRequest, ResponsePayload,
227 PROTOCOL_VERSION,
228 };
229 use secure_exec_sidecar_protocol::wire as generated_wire;
230
231 fn request(payload: RequestPayload) -> RequestFrame {
232 RequestFrame::new(7, OwnershipScope::connection("conn"), payload)
233 }
234
235 fn generated_request(
236 request_id: i64,
237 ownership: generated_wire::OwnershipScope,
238 payload: generated_wire::RequestPayload,
239 ) -> generated_wire::RequestFrame {
240 generated_wire::RequestFrame {
241 schema: generated_wire::protocol_schema(),
242 request_id,
243 ownership,
244 payload,
245 }
246 }
247
248 fn reverse_host_callback_payloads() -> Vec<RequestPayload> {
249 vec![
250 RequestPayload::HostFilesystemCall(HostFilesystemCallRequest {
251 operation: FilesystemOperation::Read,
252 path: String::from("/state"),
253 payload_size_bytes: 0,
254 }),
255 RequestPayload::PersistenceLoad(PersistenceLoadRequest {
256 key: String::from("state"),
257 }),
258 RequestPayload::PersistenceFlush(PersistenceFlushRequest {
259 key: String::from("state"),
260 payload_size_bytes: 0,
261 }),
262 ]
263 }
264
265 #[test]
266 fn dispose_and_ext_requests_are_async() {
267 let ext = request(RequestPayload::Ext(ExtEnvelope {
268 namespace: String::from("test"),
269 payload: Vec::new(),
270 }));
271 assert_eq!(request_dispatch_mode(&ext), RequestDispatchMode::Async);
272 }
273
274 #[test]
275 fn normal_requests_are_immediate() {
276 let authenticate = request(RequestPayload::Authenticate(AuthenticateRequest {
277 client_name: String::from("test"),
278 auth_token: String::from("token"),
279 protocol_version: PROTOCOL_VERSION,
280 bridge_version: 1,
281 }));
282 assert_eq!(
283 request_dispatch_mode(&authenticate),
284 RequestDispatchMode::Immediate
285 );
286 }
287
288 #[test]
289 fn host_callback_requests_are_identified_as_reverse_direction_only() {
290 for payload in reverse_host_callback_payloads() {
291 let host_call = request(payload);
292
293 assert!(request_is_unsupported_host_callback_direction(&host_call));
294 assert_eq!(
295 request_dispatch_mode(&host_call),
296 RequestDispatchMode::Immediate
297 );
298 }
299 }
300
301 #[test]
302 fn routes_protocol_payloads_through_shared_enum() {
303 let authenticate = request(RequestPayload::Authenticate(AuthenticateRequest {
304 client_name: String::from("test"),
305 auth_token: String::from("token"),
306 protocol_version: PROTOCOL_VERSION,
307 bridge_version: 1,
308 }));
309 assert!(matches!(
310 route_request_payload(&authenticate),
311 RequestRoute::Authenticate(_)
312 ));
313
314 let extension = request(RequestPayload::Ext(ExtEnvelope {
315 namespace: String::from("test"),
316 payload: vec![1, 2, 3],
317 }));
318 assert!(matches!(
319 route_request_payload(&extension),
320 RequestRoute::Ext(_)
321 ));
322
323 for payload in reverse_host_callback_payloads() {
324 let host_call = request(payload);
325 assert!(matches!(
326 route_request_payload(&host_call),
327 RequestRoute::UnsupportedHostCallbackDirection
328 ));
329 }
330 }
331
332 #[test]
333 fn unsupported_host_callback_dispatch_rejects_with_shared_code() {
334 for payload in reverse_host_callback_payloads() {
335 let host_call = request(payload);
336
337 let dispatch = unsupported_host_callback_direction_dispatch(&host_call);
338
339 assert!(dispatch.events.is_empty());
340 assert_eq!(dispatch.response.request_id, host_call.request_id);
341 assert_eq!(dispatch.response.ownership, host_call.ownership);
342 match dispatch.response.payload {
343 ResponsePayload::Rejected(rejected) => {
344 assert_eq!(rejected.code, UNSUPPORTED_HOST_CALLBACK_DIRECTION_CODE);
345 assert_eq!(
346 rejected.message,
347 UNSUPPORTED_HOST_CALLBACK_DIRECTION_MESSAGE
348 );
349 }
350 other => panic!("unexpected response payload: {other:?}"),
351 }
352 }
353 }
354
355 #[test]
356 fn generated_wire_prompt_interrupt_classifier_matches_only_same_scope_interrupts() {
357 let ownership = generated_wire::OwnershipScope::VmOwnership(generated_wire::VmOwnership {
358 connection_id: String::from("conn"),
359 session_id: String::from("session"),
360 vm_id: String::from("vm"),
361 });
362 let active = generated_request(
363 1,
364 ownership.clone(),
365 generated_wire::RequestPayload::ExtEnvelope(generated_wire::ExtEnvelope {
366 namespace: String::from("prompt"),
367 payload: b"active".to_vec(),
368 }),
369 );
370
371 let same_namespace = generated_request(
372 2,
373 ownership.clone(),
374 generated_wire::RequestPayload::ExtEnvelope(generated_wire::ExtEnvelope {
375 namespace: String::from("prompt"),
376 payload: b"cancel".to_vec(),
377 }),
378 );
379 assert_eq!(
380 generated_wire_blocking_extension_interrupt(&active, "prompt", &same_namespace),
381 Some(BlockingExtensionInterrupt::ExtensionPayload(b"cancel"))
382 );
383
384 let kill = generated_request(
385 3,
386 ownership.clone(),
387 generated_wire::RequestPayload::KillProcessRequest(
388 generated_wire::KillProcessRequest {
389 process_id: String::from("proc"),
390 signal: String::from("SIGTERM"),
391 },
392 ),
393 );
394 assert_eq!(
395 generated_wire_blocking_extension_interrupt(&active, "prompt", &kill),
396 Some(BlockingExtensionInterrupt::KillProcess)
397 );
398
399 let other_namespace = generated_request(
400 4,
401 ownership.clone(),
402 generated_wire::RequestPayload::ExtEnvelope(generated_wire::ExtEnvelope {
403 namespace: String::from("other"),
404 payload: b"cancel".to_vec(),
405 }),
406 );
407 assert_eq!(
408 generated_wire_blocking_extension_interrupt(&active, "prompt", &other_namespace),
409 None
410 );
411
412 let other_scope = generated_request(
413 5,
414 generated_wire::OwnershipScope::VmOwnership(generated_wire::VmOwnership {
415 connection_id: String::from("conn"),
416 session_id: String::from("session"),
417 vm_id: String::from("other-vm"),
418 }),
419 generated_wire::RequestPayload::KillProcessRequest(
420 generated_wire::KillProcessRequest {
421 process_id: String::from("proc"),
422 signal: String::from("SIGTERM"),
423 },
424 ),
425 );
426 assert_eq!(
427 generated_wire_blocking_extension_interrupt(&active, "prompt", &other_scope),
428 None
429 );
430 }
431
432 #[test]
433 fn ownership_scope_helpers_extract_shared_ids() {
434 let connection = OwnershipScope::connection("conn-1");
435 let session = OwnershipScope::session("conn-1", "session-1");
436 let vm = OwnershipScope::vm("conn-1", "session-1", "vm-1");
437
438 assert_eq!(connection_id_of(&connection).as_deref(), Some("conn-1"));
439 assert_eq!(connection_id_of(&session).as_deref(), Some("conn-1"));
440 assert_eq!(connection_id_of(&vm).as_deref(), Some("conn-1"));
441 assert_eq!(
442 session_scope_of(&session),
443 Some((String::from("conn-1"), String::from("session-1")))
444 );
445 assert_eq!(
446 session_scope_of(&vm),
447 Some((String::from("conn-1"), String::from("session-1")))
448 );
449 assert_eq!(session_scope_of(&connection), None);
450 assert_eq!(vm_id_of(&vm).as_deref(), Some("vm-1"));
451 assert_eq!(vm_id_of(&session), None);
452 }
453}