Skip to main content

secure_exec_sidecar_core/
router.rs

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// Request payload variants intentionally vary widely in size (small acks next
25// to bulky create/exec payloads); boxing is a wire-adjacent refactor.
26#[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}