Skip to main content

little_durable_objects/
sandbox.rs

1use std::{
2    collections::HashMap,
3    process::Stdio,
4    time::{Duration, Instant},
5};
6
7use anyhow::{Context, Result, ensure};
8use async_trait::async_trait;
9use serde::{Deserialize, Serialize};
10use tokio::{
11    io::{AsyncRead, AsyncReadExt, AsyncWriteExt},
12    process::Command,
13};
14
15use crate::host::HostId;
16
17const PROVIDER_REQUEST_TIMEOUT: Duration = Duration::from_secs(120);
18const MAX_PROVIDER_OUTPUT_BYTES: usize = 1024 * 1024;
19
20#[derive(Clone, Debug, PartialEq, Eq, Serialize)]
21#[serde(rename_all = "camelCase")]
22pub struct EnsureHostRequest {
23    pub namespace_id: String,
24    pub code_revision: String,
25    pub canonical_region: String,
26    pub host_id: HostId,
27    pub session_id: String,
28    pub host_token: String,
29    pub jwt_public_keys: String,
30    pub control_plane_url: String,
31    pub jwt_issuer: String,
32    pub invocation_jwt_audience: String,
33    pub image_ref: String,
34    pub working_directory: String,
35    pub actor_entrypoint: Option<String>,
36    pub actor_idle_timeout_ms: u64,
37    pub host_idle_timeout_ms: u64,
38}
39
40#[derive(Clone, Debug, PartialEq, Eq, Deserialize)]
41#[serde(rename_all = "camelCase")]
42pub struct ActorHostHandle {
43    pub host_id: HostId,
44    pub route: String,
45    pub canonical_region: String,
46    pub provisioning: Option<ActorHostProvisioning>,
47}
48
49#[derive(Clone, Debug, PartialEq, Eq, Deserialize)]
50#[serde(rename_all = "camelCase")]
51pub struct ActorHostProvisioning {
52    pub provider: String,
53    pub resource_id: String,
54    pub reused: bool,
55    pub started_at_ms: u64,
56    pub input_parsed_at_ms: Option<u64>,
57    pub sdk_loaded_at_ms: Option<u64>,
58    pub resources_resolved_at_ms: Option<u64>,
59    pub existing_host_checked_at_ms: Option<u64>,
60    pub sandbox_scheduled_at_ms: Option<u64>,
61    pub host_ready_observed_at_ms: Option<u64>,
62    pub route_read_at_ms: Option<u64>,
63    pub metadata_written_at_ms: Option<u64>,
64    pub completed_at_ms: u64,
65    #[serde(default)]
66    pub command_spawned_at_ms: Option<u64>,
67    #[serde(default)]
68    pub request_written_at_ms: Option<u64>,
69    #[serde(default)]
70    pub process_completed_at_ms: Option<u64>,
71    #[serde(default)]
72    pub response_decoded_at_ms: Option<u64>,
73}
74
75#[derive(Clone, Debug, PartialEq, Eq, Serialize)]
76#[serde(rename_all = "camelCase")]
77pub struct WarmImageRequest {
78    pub namespace_id: String,
79    pub code_revision: String,
80    pub canonical_region: String,
81    pub image_ref: String,
82}
83
84#[derive(Clone, Debug, PartialEq, Eq, Deserialize)]
85#[serde(rename_all = "camelCase")]
86pub struct ImageWarmup {
87    pub provider: String,
88    pub resource_id: String,
89    pub total_ms: u64,
90}
91
92#[derive(Clone, Debug, PartialEq, Eq, Serialize)]
93#[serde(rename_all = "camelCase")]
94pub struct TerminateHostsRequest {
95    pub namespace_id: String,
96    pub code_revision: String,
97    pub canonical_regions: Vec<String>,
98}
99
100#[derive(Clone, Debug, PartialEq, Eq, Deserialize)]
101#[serde(rename_all = "camelCase")]
102pub struct HostTermination {
103    pub provider: String,
104    pub resource_ids: Vec<String>,
105}
106
107#[derive(Clone, Debug, PartialEq, Eq, Serialize)]
108#[serde(rename_all = "camelCase")]
109pub struct PublicHostRouteRequest {
110    pub namespace_id: String,
111    pub code_revision: String,
112    pub canonical_region: String,
113}
114
115#[derive(Clone, Debug, PartialEq, Eq, Deserialize)]
116#[serde(rename_all = "camelCase")]
117pub struct PublicHostRoute {
118    pub route: String,
119}
120
121#[async_trait]
122pub trait SandboxProvider: Send + Sync {
123    async fn ensure_host(&self, request: &EnsureHostRequest) -> Result<ActorHostHandle>;
124    async fn public_host_route(&self, request: &PublicHostRouteRequest) -> Result<PublicHostRoute>;
125    async fn warm_image(&self, request: &WarmImageRequest) -> Result<ImageWarmup>;
126    async fn terminate_hosts(&self, request: &TerminateHostsRequest) -> Result<HostTermination>;
127}
128
129#[derive(Clone)]
130pub struct HostSandboxRuntimeConfig {
131    pub control_plane_url: String,
132    pub jwt_issuer: String,
133    pub invocation_jwt_audience: String,
134    pub actor_idle_timeout_ms: u64,
135    pub host_idle_timeout_ms: u64,
136}
137
138pub struct CommandSandboxProvider {
139    provider_name: String,
140    command: String,
141    environment: HashMap<String, String>,
142}
143
144impl CommandSandboxProvider {
145    pub fn new(
146        provider_name: String,
147        command: String,
148        mut environment: HashMap<String, String>,
149    ) -> Result<Self> {
150        ensure!(
151            !provider_name.is_empty() && provider_name.trim() == provider_name,
152            "sandbox provider name must be non-empty without surrounding whitespace"
153        );
154        ensure!(
155            !command.is_empty() && command.trim() == command,
156            "DURABLE_OBJECT_SANDBOX_COMMAND must be non-empty without surrounding whitespace"
157        );
158        if let Ok(path) = std::env::var("PATH") {
159            environment.entry("PATH".into()).or_insert(path);
160        }
161        Ok(Self {
162            provider_name,
163            command,
164            environment,
165        })
166    }
167}
168
169#[async_trait]
170impl SandboxProvider for CommandSandboxProvider {
171    async fn ensure_host(&self, request: &EnsureHostRequest) -> Result<ActorHostHandle> {
172        let (mut response, command): (ActorHostHandle, _) =
173            self.execute_timed("ensure_host", request).await?;
174        if let Some(provisioning) = &mut response.provisioning {
175            provisioning.command_spawned_at_ms = command.spawned_at_ms;
176            provisioning.request_written_at_ms = command.request_written_at_ms;
177            provisioning.process_completed_at_ms = command.process_completed_at_ms;
178            provisioning.response_decoded_at_ms = command.response_decoded_at_ms;
179        }
180        ensure!(
181            response.canonical_region == request.canonical_region,
182            "{} sandbox command returned a host in the wrong canonical region",
183            self.provider_name
184        );
185        ensure!(
186            !response.host_id.as_str().is_empty() && !response.route.is_empty(),
187            "{} sandbox command returned an invalid host",
188            self.provider_name
189        );
190        Ok(response)
191    }
192
193    async fn public_host_route(&self, request: &PublicHostRouteRequest) -> Result<PublicHostRoute> {
194        let response: PublicHostRoute = self.execute("public_host_route", request).await?;
195        ensure!(
196            !response.route.is_empty(),
197            "{} sandbox command returned an invalid public host route",
198            self.provider_name
199        );
200        Ok(response)
201    }
202
203    async fn warm_image(&self, request: &WarmImageRequest) -> Result<ImageWarmup> {
204        self.execute("warm_image", request).await
205    }
206
207    async fn terminate_hosts(&self, request: &TerminateHostsRequest) -> Result<HostTermination> {
208        self.execute("terminate_hosts", request).await
209    }
210}
211
212impl CommandSandboxProvider {
213    async fn execute<Request: Serialize, Reply: for<'de> Deserialize<'de>>(
214        &self,
215        operation: &str,
216        request: &Request,
217    ) -> Result<Reply> {
218        Ok(self.execute_timed(operation, request).await?.0)
219    }
220
221    async fn execute_timed<Request: Serialize, Reply: for<'de> Deserialize<'de>>(
222        &self,
223        operation: &str,
224        request: &Request,
225    ) -> Result<(Reply, ProviderCommandTimings)> {
226        let started_at = Instant::now();
227        let mut timings = ProviderCommandTimings::default();
228        match self
229            .execute_timed_inner(operation, request, started_at, &mut timings)
230            .await
231        {
232            Ok(response) => Ok((response, timings)),
233            Err(source) => Err(ProviderCommandFailure { source, timings }.into()),
234        }
235    }
236
237    async fn execute_timed_inner<Request: Serialize, Reply: for<'de> Deserialize<'de>>(
238        &self,
239        operation: &str,
240        request: &Request,
241        started_at: Instant,
242        timings: &mut ProviderCommandTimings,
243    ) -> Result<Reply> {
244        let mut child = Command::new(&self.command)
245            .env_clear()
246            .envs(&self.environment)
247            .stdin(Stdio::piped())
248            .stdout(Stdio::piped())
249            .stderr(Stdio::piped())
250            .kill_on_drop(true)
251            .spawn()
252            .with_context(|| {
253                format!(
254                    "start {} sandbox command {:?}",
255                    self.provider_name, self.command
256                )
257            })?;
258        timings.spawned_at_ms = Some(elapsed_ms(started_at));
259        let document = serde_json::to_vec(&ProviderCommand { operation, request })?;
260        let mut stdin = child
261            .stdin
262            .take()
263            .with_context(|| format!("open {} sandbox command stdin", self.provider_name))?;
264        stdin
265            .write_all(&document)
266            .await
267            .with_context(|| format!("write {} sandbox command request", self.provider_name))?;
268        stdin
269            .shutdown()
270            .await
271            .with_context(|| format!("close {} sandbox command stdin", self.provider_name))?;
272        timings.request_written_at_ms = Some(elapsed_ms(started_at));
273        drop(stdin);
274        let stdout = child
275            .stdout
276            .take()
277            .with_context(|| format!("open {} sandbox command stdout", self.provider_name))?;
278        let stderr = child
279            .stderr
280            .take()
281            .with_context(|| format!("open {} sandbox command stderr", self.provider_name))?;
282        let execution = async {
283            let wait = async {
284                child
285                    .wait()
286                    .await
287                    .with_context(|| format!("wait for {} sandbox command", self.provider_name))
288            };
289            tokio::try_join!(
290                wait,
291                read_bounded_output(stdout, MAX_PROVIDER_OUTPUT_BYTES, "stdout"),
292                read_bounded_output(stderr, MAX_PROVIDER_OUTPUT_BYTES, "stderr"),
293            )
294        };
295        let (status, stdout, stderr) = tokio::time::timeout(PROVIDER_REQUEST_TIMEOUT, execution)
296            .await
297            .with_context(|| format!("{} sandbox command timed out", self.provider_name))??;
298        timings.process_completed_at_ms = Some(elapsed_ms(started_at));
299        ensure!(
300            status.success(),
301            "{} sandbox command failed with {}: {}",
302            self.provider_name,
303            status,
304            String::from_utf8_lossy(&stderr).trim()
305        );
306        let response = serde_json::from_slice(&stdout)
307            .with_context(|| format!("decode {} sandbox command response", self.provider_name))?;
308        timings.response_decoded_at_ms = Some(elapsed_ms(started_at));
309        Ok(response)
310    }
311}
312
313#[derive(Debug, Default)]
314struct ProviderCommandTimings {
315    spawned_at_ms: Option<u64>,
316    request_written_at_ms: Option<u64>,
317    process_completed_at_ms: Option<u64>,
318    response_decoded_at_ms: Option<u64>,
319}
320
321#[derive(Debug)]
322pub(crate) struct ProviderCommandFailure {
323    source: anyhow::Error,
324    timings: ProviderCommandTimings,
325}
326
327impl ProviderCommandFailure {
328    pub(crate) fn spawned_at_ms(&self) -> Option<u64> {
329        self.timings.spawned_at_ms
330    }
331
332    pub(crate) fn request_written_at_ms(&self) -> Option<u64> {
333        self.timings.request_written_at_ms
334    }
335
336    pub(crate) fn process_completed_at_ms(&self) -> Option<u64> {
337        self.timings.process_completed_at_ms
338    }
339
340    pub(crate) fn response_decoded_at_ms(&self) -> Option<u64> {
341        self.timings.response_decoded_at_ms
342    }
343}
344
345impl std::fmt::Display for ProviderCommandFailure {
346    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
347        self.source.fmt(formatter)
348    }
349}
350
351impl std::error::Error for ProviderCommandFailure {
352    fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
353        self.source.source()
354    }
355}
356
357fn elapsed_ms(started_at: Instant) -> u64 {
358    u64::try_from(started_at.elapsed().as_millis()).unwrap_or(u64::MAX)
359}
360
361async fn read_bounded_output(
362    reader: impl AsyncRead + Unpin,
363    max_bytes: usize,
364    label: &str,
365) -> Result<Vec<u8>> {
366    let mut output = Vec::new();
367    reader
368        .take((max_bytes + 1) as u64)
369        .read_to_end(&mut output)
370        .await
371        .with_context(|| format!("read sandbox command {label}"))?;
372    ensure!(
373        output.len() <= max_bytes,
374        "sandbox command {label} exceeds {max_bytes} bytes"
375    );
376    Ok(output)
377}
378
379#[derive(Serialize)]
380struct ProviderCommand<'a, Request> {
381    operation: &'a str,
382    request: &'a Request,
383}
384
385#[cfg(test)]
386mod tests {
387    use super::*;
388
389    #[test]
390    fn decodes_provider_provisioning_timings() {
391        let handle: ActorHostHandle = serde_json::from_value(serde_json::json!({
392            "hostId": "host.v1.namespace.revision.session",
393            "route": "https://host.example.com",
394            "canonicalRegion": "north-america-east",
395            "provisioning": {
396                "provider": "modal",
397                "resourceId": "sb-actor",
398                "reused": false,
399                "startedAtMs": 0,
400                "resourcesResolvedAtMs": 12,
401                "existingHostCheckedAtMs": 34,
402                "sandboxScheduledAtMs": 56,
403                "hostReadyObservedAtMs": 123,
404                "routeReadAtMs": 125,
405                "metadataWrittenAtMs": 129,
406                "completedAtMs": 130
407            }
408        }))
409        .expect("actor host handle");
410
411        let provisioning = handle.provisioning.expect("provisioning timings");
412        assert_eq!(provisioning.resource_id, "sb-actor");
413        assert_eq!(provisioning.sandbox_scheduled_at_ms, Some(56));
414        assert_eq!(provisioning.completed_at_ms, 130);
415    }
416
417    #[test]
418    fn rejects_ambiguous_command_configuration() {
419        assert!(CommandSandboxProvider::new("".into(), "modal".into(), HashMap::new()).is_err());
420        assert!(CommandSandboxProvider::new("modal".into(), "".into(), HashMap::new()).is_err());
421    }
422
423    #[tokio::test]
424    async fn provider_output_is_bounded_while_reading() {
425        let error = read_bounded_output(tokio::io::repeat(b'x'), 32, "stdout")
426            .await
427            .expect_err("unbounded provider output should fail");
428        assert!(error.to_string().contains("stdout exceeds 32 bytes"));
429    }
430}