1use std::future::Future;
2use std::pin::Pin;
3use std::time::Duration;
4
5use crate::protocol::{
6 CloseStdinRequest, EventFrame, EventPayload, ExecuteRequest, ExtEnvelope,
7 GuestFilesystemCallRequest, GuestFilesystemResultResponse, KillProcessRequest, OwnershipScope,
8 ProcessKilledResponse, ProcessStartedResponse, SidecarRequestPayload, SidecarResponsePayload,
9 StdinClosedResponse, StdinWrittenResponse, WriteStdinRequest,
10};
11use crate::state::{SharedEventSink, SharedSidecarRequestClient, SidecarError};
12
13pub type ExtensionFuture<'a, T> = Pin<Box<dyn Future<Output = Result<T, SidecarError>> + 'a>>;
14
15pub trait ExtensionHost {
16 fn spawn_process<'a>(
17 &'a mut self,
18 ownership: OwnershipScope,
19 request: ExecuteRequest,
20 ) -> ExtensionFuture<'a, ProcessStartedResponse>;
21
22 fn write_stdin<'a>(
23 &'a mut self,
24 ownership: OwnershipScope,
25 request: WriteStdinRequest,
26 ) -> ExtensionFuture<'a, StdinWrittenResponse>;
27
28 fn close_stdin<'a>(
29 &'a mut self,
30 ownership: OwnershipScope,
31 request: CloseStdinRequest,
32 ) -> ExtensionFuture<'a, StdinClosedResponse>;
33
34 fn kill_process<'a>(
35 &'a mut self,
36 ownership: OwnershipScope,
37 request: KillProcessRequest,
38 ) -> ExtensionFuture<'a, ProcessKilledResponse>;
39
40 fn poll_event<'a>(
41 &'a mut self,
42 ownership: OwnershipScope,
43 timeout: Duration,
44 ) -> ExtensionFuture<'a, Option<EventFrame>>;
45
46 fn guest_filesystem_call<'a>(
47 &'a mut self,
48 ownership: OwnershipScope,
49 request: GuestFilesystemCallRequest,
50 ) -> ExtensionFuture<'a, GuestFilesystemResultResponse>;
51
52 fn bind_process_to_session<'a>(
53 &'a mut self,
54 ownership: OwnershipScope,
55 namespace: String,
56 ext_session_id: String,
57 process_id: String,
58 ) -> ExtensionFuture<'a, ()>;
59
60 fn bind_vm_to_session<'a>(
61 &'a mut self,
62 ownership: OwnershipScope,
63 namespace: String,
64 ext_session_id: String,
65 ) -> ExtensionFuture<'a, ()>;
66
67 fn dispose_session_resources<'a>(
68 &'a mut self,
69 ownership: OwnershipScope,
70 namespace: String,
71 ext_session_id: String,
72 ) -> ExtensionFuture<'a, Vec<EventFrame>>;
73
74 fn start_buffering_process_output<'a>(
75 &'a mut self,
76 ownership: OwnershipScope,
77 process_id: String,
78 ) -> ExtensionFuture<'a, ()>;
79
80 fn handoff_buffered_process_output<'a>(
81 &'a mut self,
82 ownership: OwnershipScope,
83 namespace: String,
84 ext_session_id: String,
85 process_id: String,
86 timeout: Duration,
87 ) -> ExtensionFuture<'a, ExtensionBufferedProcessOutput>;
88}
89
90#[derive(Debug, Clone, Default, PartialEq, Eq)]
91pub struct ExtensionBufferedProcessOutput {
92 pub stdout: Vec<u8>,
93 pub stderr: Vec<u8>,
94 pub stdout_truncated: bool,
95 pub stderr_truncated: bool,
96}
97
98impl ExtensionBufferedProcessOutput {
99 pub(crate) fn append_stdout(&mut self, chunk: &[u8], cap: usize) {
100 self.stdout_truncated |= append_bounded_bytes(&mut self.stdout, chunk, cap);
101 }
102
103 pub(crate) fn append_stderr(&mut self, chunk: &[u8], cap: usize) {
104 self.stderr_truncated |= append_bounded_bytes(&mut self.stderr, chunk, cap);
105 }
106}
107
108fn append_bounded_bytes(buffer: &mut Vec<u8>, chunk: &[u8], cap: usize) -> bool {
109 buffer.extend_from_slice(chunk);
110 if buffer.len() <= cap {
111 return false;
112 }
113 let remove_len = buffer.len() - cap;
114 buffer.drain(..remove_len);
115 true
116}
117
118#[derive(Debug, Clone)]
119pub struct ExtensionResponse {
120 pub payload: Vec<u8>,
121 pub events: Vec<EventFrame>,
122}
123
124impl ExtensionResponse {
125 pub fn new(payload: Vec<u8>) -> Self {
126 Self {
127 payload,
128 events: Vec::new(),
129 }
130 }
131
132 pub fn with_events(payload: Vec<u8>, events: Vec<EventFrame>) -> Self {
133 Self { payload, events }
134 }
135
136 pub fn with_wire_events(
137 payload: Vec<u8>,
138 events: Vec<crate::wire::EventFrame>,
139 ) -> Result<Self, SidecarError> {
140 let events = events
141 .into_iter()
142 .map(crate::wire::event_frame_to_compat)
143 .collect::<Result<Vec<_>, _>>()
144 .map_err(wire_protocol_error)?;
145 Ok(Self { payload, events })
146 }
147}
148
149#[derive(Clone)]
150pub struct ExtensionSnapshot {
151 namespace: String,
152 ownership: OwnershipScope,
153 sidecar_requests: SharedSidecarRequestClient,
154 event_sink: SharedEventSink,
155}
156
157pub struct ExtensionContext<'a> {
158 snapshot: ExtensionSnapshot,
159 host: &'a mut dyn ExtensionHost,
160}
161
162impl ExtensionSnapshot {
163 pub(crate) fn new(
164 namespace: String,
165 ownership: OwnershipScope,
166 sidecar_requests: SharedSidecarRequestClient,
167 event_sink: SharedEventSink,
168 ) -> Self {
169 Self {
170 namespace,
171 ownership,
172 sidecar_requests,
173 event_sink,
174 }
175 }
176
177 pub fn namespace(&self) -> &str {
178 &self.namespace
179 }
180
181 pub fn ownership(&self) -> &OwnershipScope {
182 &self.ownership
183 }
184
185 pub fn ext_event(&self, payload: Vec<u8>) -> EventFrame {
186 EventFrame::new(
187 self.ownership.clone(),
188 EventPayload::Ext(ExtEnvelope {
189 namespace: self.namespace.clone(),
190 payload,
191 }),
192 )
193 }
194
195 pub fn ext_event_wire(
196 &self,
197 payload: Vec<u8>,
198 ) -> Result<crate::wire::EventFrame, SidecarError> {
199 crate::wire::event_frame_from_compat(self.ext_event(payload)).map_err(wire_protocol_error)
200 }
201
202 pub fn emit_event_wire(
208 &self,
209 event: crate::wire::EventFrame,
210 ) -> Result<Option<crate::wire::EventFrame>, SidecarError> {
211 self.event_sink.try_emit(event)
212 }
213
214 pub fn emit_ext_event(
217 &self,
218 payload: Vec<u8>,
219 ) -> Result<Option<crate::wire::EventFrame>, SidecarError> {
220 let event = self.ext_event_wire(payload)?;
221 self.emit_event_wire(event)
222 }
223
224 pub fn invoke_callback(
225 &self,
226 payload: Vec<u8>,
227 timeout: Duration,
228 ) -> Result<Vec<u8>, SidecarError> {
229 let response = self.sidecar_requests.invoke(
230 self.ownership.clone(),
231 SidecarRequestPayload::Ext(ExtEnvelope {
232 namespace: self.namespace.clone(),
233 payload,
234 }),
235 timeout,
236 )?;
237 extension_callback_response_payload(&self.namespace, response)
238 }
239}
240
241impl<'a> ExtensionContext<'a> {
242 pub(crate) fn new(snapshot: ExtensionSnapshot, host: &'a mut dyn ExtensionHost) -> Self {
243 Self { snapshot, host }
244 }
245
246 pub fn snapshot(&self) -> ExtensionSnapshot {
247 self.snapshot.clone()
248 }
249
250 pub fn namespace(&self) -> &str {
251 self.snapshot.namespace()
252 }
253
254 pub fn ownership(&self) -> &OwnershipScope {
255 self.snapshot.ownership()
256 }
257
258 pub fn ext_event(&self, payload: Vec<u8>) -> EventFrame {
259 self.snapshot.ext_event(payload)
260 }
261
262 pub fn ext_event_wire(
263 &self,
264 payload: Vec<u8>,
265 ) -> Result<crate::wire::EventFrame, SidecarError> {
266 self.snapshot.ext_event_wire(payload)
267 }
268
269 pub fn emit_event_wire(
272 &self,
273 event: crate::wire::EventFrame,
274 ) -> Result<Option<crate::wire::EventFrame>, SidecarError> {
275 self.snapshot.emit_event_wire(event)
276 }
277
278 pub fn emit_ext_event(
281 &self,
282 payload: Vec<u8>,
283 ) -> Result<Option<crate::wire::EventFrame>, SidecarError> {
284 self.snapshot.emit_ext_event(payload)
285 }
286
287 pub fn invoke_callback(
288 &self,
289 payload: Vec<u8>,
290 timeout: Duration,
291 ) -> Result<Vec<u8>, SidecarError> {
292 self.snapshot.invoke_callback(payload, timeout)
293 }
294
295 pub async fn spawn_process(
296 &mut self,
297 request: ExecuteRequest,
298 ) -> Result<ProcessStartedResponse, SidecarError> {
299 self.host
300 .spawn_process(self.snapshot.ownership.clone(), request)
301 .await
302 }
303
304 pub async fn spawn_process_wire(
305 &mut self,
306 request: crate::wire::ExecuteRequest,
307 ) -> Result<crate::wire::ProcessStartedResponse, SidecarError> {
308 let payload = crate::wire::request_payload_to_compat(
309 self.snapshot.ownership(),
310 crate::wire::RequestPayload::ExecuteRequest(request),
311 )
312 .map_err(wire_protocol_error)?;
313 let crate::protocol::RequestPayload::Execute(request) = payload else {
314 return Err(unexpected_wire_request_payload("execute"));
315 };
316 let response = self.spawn_process(request).await?;
317 let payload = crate::wire::response_payload_from_compat(
318 self.snapshot.ownership(),
319 crate::protocol::ResponsePayload::ProcessStarted(response),
320 )
321 .map_err(wire_protocol_error)?;
322 let crate::wire::ResponsePayload::ProcessStartedResponse(response) = payload else {
323 return Err(unexpected_wire_response_payload("process started"));
324 };
325 Ok(response)
326 }
327
328 pub async fn write_stdin(
329 &mut self,
330 request: WriteStdinRequest,
331 ) -> Result<StdinWrittenResponse, SidecarError> {
332 self.host
333 .write_stdin(self.snapshot.ownership.clone(), request)
334 .await
335 }
336
337 pub async fn write_stdin_wire(
338 &mut self,
339 request: crate::wire::WriteStdinRequest,
340 ) -> Result<crate::wire::StdinWrittenResponse, SidecarError> {
341 let payload = crate::wire::request_payload_to_compat(
342 self.snapshot.ownership(),
343 crate::wire::RequestPayload::WriteStdinRequest(request),
344 )
345 .map_err(wire_protocol_error)?;
346 let crate::protocol::RequestPayload::WriteStdin(request) = payload else {
347 return Err(unexpected_wire_request_payload("write stdin"));
348 };
349 let response = self.write_stdin(request).await?;
350 let payload = crate::wire::response_payload_from_compat(
351 self.snapshot.ownership(),
352 crate::protocol::ResponsePayload::StdinWritten(response),
353 )
354 .map_err(wire_protocol_error)?;
355 let crate::wire::ResponsePayload::StdinWrittenResponse(response) = payload else {
356 return Err(unexpected_wire_response_payload("stdin written"));
357 };
358 Ok(response)
359 }
360
361 pub async fn close_stdin(
362 &mut self,
363 request: CloseStdinRequest,
364 ) -> Result<StdinClosedResponse, SidecarError> {
365 self.host
366 .close_stdin(self.snapshot.ownership.clone(), request)
367 .await
368 }
369
370 pub async fn close_stdin_wire(
371 &mut self,
372 request: crate::wire::CloseStdinRequest,
373 ) -> Result<crate::wire::StdinClosedResponse, SidecarError> {
374 let payload = crate::wire::request_payload_to_compat(
375 self.snapshot.ownership(),
376 crate::wire::RequestPayload::CloseStdinRequest(request),
377 )
378 .map_err(wire_protocol_error)?;
379 let crate::protocol::RequestPayload::CloseStdin(request) = payload else {
380 return Err(unexpected_wire_request_payload("close stdin"));
381 };
382 let response = self.close_stdin(request).await?;
383 let payload = crate::wire::response_payload_from_compat(
384 self.snapshot.ownership(),
385 crate::protocol::ResponsePayload::StdinClosed(response),
386 )
387 .map_err(wire_protocol_error)?;
388 let crate::wire::ResponsePayload::StdinClosedResponse(response) = payload else {
389 return Err(unexpected_wire_response_payload("stdin closed"));
390 };
391 Ok(response)
392 }
393
394 pub async fn kill_process(
395 &mut self,
396 request: KillProcessRequest,
397 ) -> Result<ProcessKilledResponse, SidecarError> {
398 self.host
399 .kill_process(self.snapshot.ownership.clone(), request)
400 .await
401 }
402
403 pub async fn kill_process_wire(
404 &mut self,
405 request: crate::wire::KillProcessRequest,
406 ) -> Result<crate::wire::ProcessKilledResponse, SidecarError> {
407 let payload = crate::wire::request_payload_to_compat(
408 self.snapshot.ownership(),
409 crate::wire::RequestPayload::KillProcessRequest(request),
410 )
411 .map_err(wire_protocol_error)?;
412 let crate::protocol::RequestPayload::KillProcess(request) = payload else {
413 return Err(unexpected_wire_request_payload("kill process"));
414 };
415 let response = self.kill_process(request).await?;
416 let payload = crate::wire::response_payload_from_compat(
417 self.snapshot.ownership(),
418 crate::protocol::ResponsePayload::ProcessKilled(response),
419 )
420 .map_err(wire_protocol_error)?;
421 let crate::wire::ResponsePayload::ProcessKilledResponse(response) = payload else {
422 return Err(unexpected_wire_response_payload("process killed"));
423 };
424 Ok(response)
425 }
426
427 pub async fn poll_event(
428 &mut self,
429 timeout: Duration,
430 ) -> Result<Option<EventFrame>, SidecarError> {
431 self.host
432 .poll_event(self.snapshot.ownership.clone(), timeout)
433 .await
434 }
435
436 pub async fn poll_event_wire(
437 &mut self,
438 timeout: Duration,
439 ) -> Result<Option<crate::wire::EventFrame>, SidecarError> {
440 self.poll_event(timeout)
441 .await?
442 .map(crate::wire::event_frame_from_compat)
443 .transpose()
444 .map_err(wire_protocol_error)
445 }
446
447 pub async fn guest_filesystem_call(
448 &mut self,
449 request: GuestFilesystemCallRequest,
450 ) -> Result<GuestFilesystemResultResponse, SidecarError> {
451 self.host
452 .guest_filesystem_call(self.snapshot.ownership.clone(), request)
453 .await
454 }
455
456 pub async fn guest_filesystem_call_wire(
457 &mut self,
458 request: crate::wire::GuestFilesystemCallRequest,
459 ) -> Result<crate::wire::GuestFilesystemResultResponse, SidecarError> {
460 let payload = crate::wire::request_payload_to_compat(
461 self.snapshot.ownership(),
462 crate::wire::RequestPayload::GuestFilesystemCallRequest(request),
463 )
464 .map_err(wire_protocol_error)?;
465 let crate::protocol::RequestPayload::GuestFilesystemCall(request) = payload else {
466 return Err(unexpected_wire_request_payload("guest filesystem call"));
467 };
468 let response = self.guest_filesystem_call(request).await?;
469 let payload = crate::wire::response_payload_from_compat(
470 self.snapshot.ownership(),
471 crate::protocol::ResponsePayload::GuestFilesystemResult(response),
472 )
473 .map_err(wire_protocol_error)?;
474 let crate::wire::ResponsePayload::GuestFilesystemResultResponse(response) = payload else {
475 return Err(unexpected_wire_response_payload("guest filesystem result"));
476 };
477 Ok(response)
478 }
479
480 pub async fn bind_process_to_session(
481 &mut self,
482 ext_session_id: impl Into<String>,
483 process_id: impl Into<String>,
484 ) -> Result<(), SidecarError> {
485 self.host
486 .bind_process_to_session(
487 self.snapshot.ownership.clone(),
488 self.snapshot.namespace.clone(),
489 ext_session_id.into(),
490 process_id.into(),
491 )
492 .await
493 }
494
495 pub async fn bind_vm_to_session(
496 &mut self,
497 ext_session_id: impl Into<String>,
498 ) -> Result<(), SidecarError> {
499 self.host
500 .bind_vm_to_session(
501 self.snapshot.ownership.clone(),
502 self.snapshot.namespace.clone(),
503 ext_session_id.into(),
504 )
505 .await
506 }
507
508 pub async fn dispose_session_resources(
509 &mut self,
510 ext_session_id: impl Into<String>,
511 ) -> Result<Vec<EventFrame>, SidecarError> {
512 self.host
513 .dispose_session_resources(
514 self.snapshot.ownership.clone(),
515 self.snapshot.namespace.clone(),
516 ext_session_id.into(),
517 )
518 .await
519 }
520
521 pub async fn dispose_session_resources_wire(
522 &mut self,
523 ext_session_id: impl Into<String>,
524 ) -> Result<Vec<crate::wire::EventFrame>, SidecarError> {
525 self.dispose_session_resources(ext_session_id)
526 .await?
527 .into_iter()
528 .map(crate::wire::event_frame_from_compat)
529 .collect::<Result<Vec<_>, _>>()
530 .map_err(wire_protocol_error)
531 }
532
533 pub async fn start_buffering_process_output(
534 &mut self,
535 process_id: impl Into<String>,
536 ) -> Result<(), SidecarError> {
537 self.host
538 .start_buffering_process_output(self.snapshot.ownership.clone(), process_id.into())
539 .await
540 }
541
542 pub async fn handoff_buffered_process_output(
543 &mut self,
544 ext_session_id: impl Into<String>,
545 process_id: impl Into<String>,
546 timeout: Duration,
547 ) -> Result<ExtensionBufferedProcessOutput, SidecarError> {
548 self.host
549 .handoff_buffered_process_output(
550 self.snapshot.ownership.clone(),
551 self.snapshot.namespace.clone(),
552 ext_session_id.into(),
553 process_id.into(),
554 timeout,
555 )
556 .await
557 }
558}
559
560fn wire_protocol_error(error: crate::wire::ProtocolCodecError) -> SidecarError {
561 SidecarError::InvalidState(format!("invalid generated wire protocol frame: {error}"))
562}
563
564fn unexpected_wire_request_payload(operation: &str) -> SidecarError {
565 SidecarError::InvalidState(format!(
566 "generated wire {operation} request converted to the wrong compatibility payload"
567 ))
568}
569
570fn unexpected_wire_response_payload(operation: &str) -> SidecarError {
571 SidecarError::InvalidState(format!(
572 "compatibility {operation} response converted to the wrong generated wire payload"
573 ))
574}
575
576fn extension_callback_response_payload(
577 namespace: &str,
578 response: SidecarResponsePayload,
579) -> Result<Vec<u8>, SidecarError> {
580 match response {
581 SidecarResponsePayload::ExtResult(envelope) if envelope.namespace == namespace => {
582 Ok(envelope.payload)
583 }
584 SidecarResponsePayload::ExtResult(envelope) => Err(SidecarError::InvalidState(format!(
585 "extension callback response namespace {} did not match {}",
586 envelope.namespace, namespace
587 ))),
588 SidecarResponsePayload::HostCallbackResult(_)
589 | SidecarResponsePayload::JsBridgeResult(_) => Err(SidecarError::InvalidState(
590 String::from("extension callback received a non-extension response"),
591 )),
592 }
593}
594
595pub enum ExtensionInterruptRequest<'a> {
596 ExtensionPayload(&'a [u8]),
597 KillProcess,
598}
599
600#[derive(Debug, Clone)]
601pub struct ExtensionInterruptResponse {
602 pub interrupted_response_payload: Vec<u8>,
603 pub interrupting_response_payload: Option<Vec<u8>>,
604}
605
606pub trait Extension: Send + Sync {
607 fn namespace(&self) -> &str;
608
609 fn handle_request<'a>(
610 &'a self,
611 ctx: ExtensionContext<'a>,
612 payload: Vec<u8>,
613 ) -> ExtensionFuture<'a, ExtensionResponse>;
614
615 fn on_vm_created<'a>(&'a self, _ctx: ExtensionSnapshot) -> ExtensionFuture<'a, ()> {
616 Box::pin(async { Ok(()) })
617 }
618
619 fn on_session_disposed<'a>(&'a self, _ctx: ExtensionSnapshot) -> ExtensionFuture<'a, ()> {
628 Box::pin(async { Ok(()) })
629 }
630
631 fn is_blocking_request(&self, _payload: &[u8]) -> bool {
632 false
633 }
634
635 fn interrupt_blocking_request(
636 &self,
637 _blocking_payload: &[u8],
638 _interrupt: ExtensionInterruptRequest<'_>,
639 ) -> Option<ExtensionInterruptResponse> {
640 None
641 }
642
643 fn on_dispose<'a>(&'a self) -> ExtensionFuture<'a, ()> {
644 Box::pin(async { Ok(()) })
645 }
646}
647
648#[cfg(test)]
649mod live_event_tests {
650 use super::*;
651 use crate::state::EventSinkTransport;
652 use std::sync::{Arc, Mutex};
653
654 #[derive(Default)]
657 struct RecordingEventSink {
658 events: Arc<Mutex<Vec<crate::wire::EventFrame>>>,
659 }
660
661 impl EventSinkTransport for RecordingEventSink {
662 fn emit_event(&self, event: crate::wire::EventFrame) -> Result<(), SidecarError> {
663 self.events.lock().unwrap().push(event);
664 Ok(())
665 }
666 }
667
668 fn snapshot_with_sink(event_sink: SharedEventSink) -> ExtensionSnapshot {
669 ExtensionSnapshot::new(
670 String::from("dev.rivet.test.live-event"),
671 OwnershipScope::session("conn-live", "sess-live"),
672 SharedSidecarRequestClient::default(),
673 event_sink,
674 )
675 }
676
677 #[test]
680 fn emit_ext_event_streams_live_when_sink_configured() {
681 let recorded = Arc::new(Mutex::new(Vec::new()));
682 let mut sink = SharedEventSink::default();
683 sink.set_transport(Arc::new(RecordingEventSink {
684 events: recorded.clone(),
685 }));
686
687 let snapshot = snapshot_with_sink(sink);
688 let leftover = snapshot
689 .emit_ext_event(b"live-update".to_vec())
690 .expect("emit must succeed");
691
692 assert!(
693 leftover.is_none(),
694 "a configured sink consumes the event live, leaving nothing to batch"
695 );
696 let recorded = recorded.lock().unwrap();
697 assert_eq!(recorded.len(), 1, "the event must reach the live transport");
698 let compat = crate::wire::event_frame_to_compat(recorded[0].clone())
700 .expect("recorded frame converts back to compat");
701 match compat.payload {
702 EventPayload::Ext(envelope) => {
703 assert_eq!(envelope.namespace, "dev.rivet.test.live-event");
704 assert_eq!(envelope.payload, b"live-update");
705 }
706 other => panic!("unexpected live event payload: {other:?}"),
707 }
708 }
709
710 #[test]
713 fn emit_ext_event_falls_back_to_batch_without_sink() {
714 let snapshot = snapshot_with_sink(SharedEventSink::default());
715 let leftover = snapshot
716 .emit_ext_event(b"batched-update".to_vec())
717 .expect("emit must succeed");
718
719 let frame = leftover.expect("without a sink the event is returned for batching");
720 let compat = crate::wire::event_frame_to_compat(frame)
721 .expect("returned frame converts back to compat");
722 match compat.payload {
723 EventPayload::Ext(envelope) => assert_eq!(envelope.payload, b"batched-update"),
724 other => panic!("unexpected batched event payload: {other:?}"),
725 }
726 }
727}