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
17mod command_process;
18
19const PROVIDER_REQUEST_TIMEOUT: Duration = Duration::from_secs(120);
20const MAX_PROVIDER_OUTPUT_BYTES: usize = 1024 * 1024;
21
22#[derive(Clone, Debug, PartialEq, Eq, Serialize)]
23#[serde(rename_all = "camelCase")]
24pub struct EnsureHostRequest {
25    pub namespace_id: String,
26    pub code_revision: String,
27    pub canonical_region: String,
28    pub host_id: HostId,
29    pub session_id: String,
30    pub host_token: String,
31    pub jwt_public_keys: String,
32    pub control_plane_url: String,
33    pub jwt_issuer: String,
34    pub invocation_jwt_audience: String,
35    pub image_ref: String,
36    pub working_directory: String,
37    pub actor_entrypoint: Option<String>,
38    pub actor_idle_timeout_ms: u64,
39    pub host_idle_timeout_ms: u64,
40}
41
42#[derive(Clone, Debug, PartialEq, Eq, Deserialize)]
43#[serde(rename_all = "camelCase")]
44pub struct ActorHostHandle {
45    pub host_id: HostId,
46    pub route: String,
47    pub canonical_region: String,
48    pub provisioning: Option<ActorHostProvisioning>,
49}
50
51#[derive(Clone, Debug, PartialEq, Eq, Deserialize)]
52#[serde(rename_all = "camelCase")]
53pub struct ActorHostProvisioning {
54    pub provider: String,
55    pub resource_id: String,
56    pub reused: bool,
57    pub started_at_ms: u64,
58    pub input_parsed_at_ms: Option<u64>,
59    pub sdk_loaded_at_ms: Option<u64>,
60    pub resources_resolved_at_ms: Option<u64>,
61    pub existing_host_checked_at_ms: Option<u64>,
62    pub sandbox_scheduled_at_ms: Option<u64>,
63    pub host_ready_observed_at_ms: Option<u64>,
64    pub route_read_at_ms: Option<u64>,
65    pub metadata_written_at_ms: Option<u64>,
66    pub completed_at_ms: u64,
67    #[serde(default)]
68    pub command_spawned_at_ms: Option<u64>,
69    #[serde(default)]
70    pub request_written_at_ms: Option<u64>,
71    #[serde(default)]
72    pub process_completed_at_ms: Option<u64>,
73    #[serde(default)]
74    pub response_decoded_at_ms: Option<u64>,
75}
76
77#[derive(Clone, Debug, PartialEq, Eq, Serialize)]
78#[serde(rename_all = "camelCase")]
79pub struct WarmImageRequest {
80    pub namespace_id: String,
81    pub code_revision: String,
82    pub canonical_region: String,
83    pub image_ref: String,
84}
85
86#[derive(Clone, Debug, PartialEq, Eq, Deserialize)]
87#[serde(rename_all = "camelCase")]
88pub struct ImageWarmup {
89    pub provider: String,
90    pub resource_id: String,
91    pub total_ms: u64,
92}
93
94#[derive(Clone, Debug, PartialEq, Eq, Serialize)]
95#[serde(rename_all = "camelCase")]
96pub struct TerminateHostsRequest {
97    pub namespace_id: String,
98    pub code_revision: String,
99    pub canonical_regions: Vec<String>,
100}
101
102#[derive(Clone, Debug, PartialEq, Eq, Deserialize)]
103#[serde(rename_all = "camelCase")]
104pub struct HostTermination {
105    pub provider: String,
106    pub resource_ids: Vec<String>,
107}
108
109#[derive(Clone, Debug, PartialEq, Eq, Serialize)]
110#[serde(rename_all = "camelCase")]
111pub struct PublicHostRouteRequest {
112    pub namespace_id: String,
113    pub code_revision: String,
114    pub canonical_region: String,
115}
116
117#[derive(Clone, Debug, PartialEq, Eq, Deserialize)]
118#[serde(rename_all = "camelCase")]
119pub struct PublicHostRoute {
120    pub route: String,
121}
122
123#[async_trait]
124pub trait SandboxProvider: Send + Sync {
125    async fn ensure_host(&self, request: &EnsureHostRequest) -> Result<ActorHostHandle>;
126    async fn public_host_route(&self, request: &PublicHostRouteRequest) -> Result<PublicHostRoute>;
127    async fn warm_image(&self, request: &WarmImageRequest) -> Result<ImageWarmup>;
128    async fn terminate_hosts(&self, request: &TerminateHostsRequest) -> Result<HostTermination>;
129}
130
131#[derive(Clone)]
132pub struct HostSandboxRuntimeConfig {
133    pub control_plane_url: String,
134    pub jwt_issuer: String,
135    pub invocation_jwt_audience: String,
136    pub actor_idle_timeout_ms: u64,
137    pub host_idle_timeout_ms: u64,
138}
139
140pub struct CommandSandboxProvider {
141    provider_name: String,
142    command: String,
143    environment: HashMap<String, String>,
144    processes: Option<deadpool::managed::Pool<command_process::ProviderProcessManager>>,
145}
146
147impl CommandSandboxProvider {
148    pub fn new(
149        provider_name: String,
150        command: String,
151        mut environment: HashMap<String, String>,
152    ) -> Result<Self> {
153        ensure!(
154            !provider_name.is_empty() && provider_name.trim() == provider_name,
155            "sandbox provider name must be non-empty without surrounding whitespace"
156        );
157        ensure!(
158            !command.is_empty() && command.trim() == command,
159            "DURABLE_OBJECT_SANDBOX_COMMAND must be non-empty without surrounding whitespace"
160        );
161        if let Ok(path) = std::env::var("PATH") {
162            environment.entry("PATH".into()).or_insert(path);
163        }
164        let processes = if std::path::Path::new(&command)
165            .file_name()
166            .is_some_and(|name| name == "little-durable-objects-modal")
167        {
168            Some(command_process::pool(command.clone(), environment.clone())?)
169        } else {
170            None
171        };
172        Ok(Self {
173            provider_name,
174            command,
175            environment,
176            processes,
177        })
178    }
179}
180
181#[async_trait]
182impl SandboxProvider for CommandSandboxProvider {
183    async fn ensure_host(&self, request: &EnsureHostRequest) -> Result<ActorHostHandle> {
184        let (mut response, command): (ActorHostHandle, _) =
185            self.execute_timed("ensure_host", request).await?;
186        if let Some(provisioning) = &mut response.provisioning {
187            provisioning.command_spawned_at_ms = command.spawned_at_ms;
188            provisioning.request_written_at_ms = command.request_written_at_ms;
189            provisioning.process_completed_at_ms = command.process_completed_at_ms;
190            provisioning.response_decoded_at_ms = command.response_decoded_at_ms;
191        }
192        ensure!(
193            response.canonical_region == request.canonical_region,
194            "{} sandbox command returned a host in the wrong canonical region",
195            self.provider_name
196        );
197        ensure!(
198            !response.host_id.as_str().is_empty() && !response.route.is_empty(),
199            "{} sandbox command returned an invalid host",
200            self.provider_name
201        );
202        Ok(response)
203    }
204
205    async fn public_host_route(&self, request: &PublicHostRouteRequest) -> Result<PublicHostRoute> {
206        let response: PublicHostRoute = self.execute("public_host_route", request).await?;
207        ensure!(
208            !response.route.is_empty(),
209            "{} sandbox command returned an invalid public host route",
210            self.provider_name
211        );
212        Ok(response)
213    }
214
215    async fn warm_image(&self, request: &WarmImageRequest) -> Result<ImageWarmup> {
216        self.execute("warm_image", request).await
217    }
218
219    async fn terminate_hosts(&self, request: &TerminateHostsRequest) -> Result<HostTermination> {
220        self.execute("terminate_hosts", request).await
221    }
222}
223
224impl CommandSandboxProvider {
225    async fn execute<Request: Serialize, Reply: for<'de> Deserialize<'de>>(
226        &self,
227        operation: &str,
228        request: &Request,
229    ) -> Result<Reply> {
230        Ok(self.execute_timed(operation, request).await?.0)
231    }
232
233    async fn execute_timed<Request: Serialize, Reply: for<'de> Deserialize<'de>>(
234        &self,
235        operation: &str,
236        request: &Request,
237    ) -> Result<(Reply, ProviderCommandTimings)> {
238        let started_at = Instant::now();
239        let mut timings = ProviderCommandTimings::default();
240        match self
241            .execute_timed_inner(operation, request, started_at, &mut timings)
242            .await
243        {
244            Ok(response) => Ok((response, timings)),
245            Err(source) => Err(ProviderCommandFailure { source, timings }.into()),
246        }
247    }
248
249    async fn execute_timed_inner<Request: Serialize, Reply: for<'de> Deserialize<'de>>(
250        &self,
251        operation: &str,
252        request: &Request,
253        started_at: Instant,
254        timings: &mut ProviderCommandTimings,
255    ) -> Result<Reply> {
256        if let Some(processes) = &self.processes {
257            let command = ProviderCommand { operation, request };
258            let execution = command_process::exchange(processes, &command, started_at, timings);
259            return tokio::time::timeout(PROVIDER_REQUEST_TIMEOUT, execution)
260                .await
261                .context("sandbox provider command timed out; outcome may be unknown")?;
262        }
263        let mut child = Command::new(&self.command)
264            .env_clear()
265            .envs(&self.environment)
266            .stdin(Stdio::piped())
267            .stdout(Stdio::piped())
268            .stderr(Stdio::piped())
269            .kill_on_drop(true)
270            .spawn()
271            .with_context(|| {
272                format!(
273                    "start {} sandbox command {:?}",
274                    self.provider_name, self.command
275                )
276            })?;
277        timings.spawned_at_ms = Some(elapsed_ms(started_at));
278        let document = serde_json::to_vec(&ProviderCommand { operation, request })?;
279        let mut stdin = child
280            .stdin
281            .take()
282            .with_context(|| format!("open {} sandbox command stdin", self.provider_name))?;
283        stdin
284            .write_all(&document)
285            .await
286            .with_context(|| format!("write {} sandbox command request", self.provider_name))?;
287        stdin
288            .shutdown()
289            .await
290            .with_context(|| format!("close {} sandbox command stdin", self.provider_name))?;
291        timings.request_written_at_ms = Some(elapsed_ms(started_at));
292        drop(stdin);
293        let stdout = child
294            .stdout
295            .take()
296            .with_context(|| format!("open {} sandbox command stdout", self.provider_name))?;
297        let stderr = child
298            .stderr
299            .take()
300            .with_context(|| format!("open {} sandbox command stderr", self.provider_name))?;
301        let execution = async {
302            let wait = async {
303                child
304                    .wait()
305                    .await
306                    .with_context(|| format!("wait for {} sandbox command", self.provider_name))
307            };
308            tokio::try_join!(
309                wait,
310                read_bounded_output(stdout, MAX_PROVIDER_OUTPUT_BYTES, "stdout"),
311                read_bounded_output(stderr, MAX_PROVIDER_OUTPUT_BYTES, "stderr"),
312            )
313        };
314        let (status, stdout, stderr) = tokio::time::timeout(PROVIDER_REQUEST_TIMEOUT, execution)
315            .await
316            .with_context(|| format!("{} sandbox command timed out", self.provider_name))??;
317        timings.process_completed_at_ms = Some(elapsed_ms(started_at));
318        ensure!(
319            status.success(),
320            "{} sandbox command failed with {}: {}",
321            self.provider_name,
322            status,
323            String::from_utf8_lossy(&stderr).trim()
324        );
325        let response = serde_json::from_slice(&stdout)
326            .with_context(|| format!("decode {} sandbox command response", self.provider_name))?;
327        timings.response_decoded_at_ms = Some(elapsed_ms(started_at));
328        Ok(response)
329    }
330}
331
332#[derive(Debug, Default)]
333struct ProviderCommandTimings {
334    spawned_at_ms: Option<u64>,
335    request_written_at_ms: Option<u64>,
336    process_completed_at_ms: Option<u64>,
337    response_decoded_at_ms: Option<u64>,
338}
339
340#[derive(Debug)]
341pub(crate) struct ProviderCommandFailure {
342    source: anyhow::Error,
343    timings: ProviderCommandTimings,
344}
345
346impl ProviderCommandFailure {
347    pub(crate) fn spawned_at_ms(&self) -> Option<u64> {
348        self.timings.spawned_at_ms
349    }
350
351    pub(crate) fn request_written_at_ms(&self) -> Option<u64> {
352        self.timings.request_written_at_ms
353    }
354
355    pub(crate) fn process_completed_at_ms(&self) -> Option<u64> {
356        self.timings.process_completed_at_ms
357    }
358
359    pub(crate) fn response_decoded_at_ms(&self) -> Option<u64> {
360        self.timings.response_decoded_at_ms
361    }
362}
363
364impl std::fmt::Display for ProviderCommandFailure {
365    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
366        self.source.fmt(formatter)
367    }
368}
369
370impl std::error::Error for ProviderCommandFailure {
371    fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
372        self.source.source()
373    }
374}
375
376fn elapsed_ms(started_at: Instant) -> u64 {
377    u64::try_from(started_at.elapsed().as_millis()).unwrap_or(u64::MAX)
378}
379
380async fn read_bounded_output(
381    reader: impl AsyncRead + Unpin,
382    max_bytes: usize,
383    label: &str,
384) -> Result<Vec<u8>> {
385    let mut output = Vec::new();
386    reader
387        .take((max_bytes + 1) as u64)
388        .read_to_end(&mut output)
389        .await
390        .with_context(|| format!("read sandbox command {label}"))?;
391    ensure!(
392        output.len() <= max_bytes,
393        "sandbox command {label} exceeds {max_bytes} bytes"
394    );
395    Ok(output)
396}
397
398#[derive(Serialize)]
399struct ProviderCommand<'a, Request> {
400    operation: &'a str,
401    request: &'a Request,
402}
403
404#[cfg(test)]
405mod tests {
406    use super::*;
407
408    #[test]
409    fn decodes_provider_provisioning_timings() {
410        let handle: ActorHostHandle = serde_json::from_value(serde_json::json!({
411            "hostId": "host.v1.namespace.revision.session",
412            "route": "https://host.example.com",
413            "canonicalRegion": "north-america-east",
414            "provisioning": {
415                "provider": "modal",
416                "resourceId": "sb-actor",
417                "reused": false,
418                "startedAtMs": 0,
419                "resourcesResolvedAtMs": 12,
420                "existingHostCheckedAtMs": 34,
421                "sandboxScheduledAtMs": 56,
422                "hostReadyObservedAtMs": 123,
423                "routeReadAtMs": 125,
424                "metadataWrittenAtMs": 129,
425                "completedAtMs": 130
426            }
427        }))
428        .expect("actor host handle");
429
430        let provisioning = handle.provisioning.expect("provisioning timings");
431        assert_eq!(provisioning.resource_id, "sb-actor");
432        assert_eq!(provisioning.sandbox_scheduled_at_ms, Some(56));
433        assert_eq!(provisioning.completed_at_ms, 130);
434    }
435
436    #[test]
437    fn rejects_ambiguous_command_configuration() {
438        assert!(CommandSandboxProvider::new("".into(), "modal".into(), HashMap::new()).is_err());
439        assert!(CommandSandboxProvider::new("modal".into(), "".into(), HashMap::new()).is_err());
440    }
441
442    #[tokio::test]
443    async fn provider_output_is_bounded_while_reading() {
444        let error = read_bounded_output(tokio::io::repeat(b'x'), 32, "stdout")
445            .await
446            .expect_err("unbounded provider output should fail");
447        assert!(error.to_string().contains("stdout exceeds 32 bytes"));
448    }
449
450    #[tokio::test]
451    async fn builtin_provider_reuses_its_process_and_discards_cancelled_exchanges() -> Result<()> {
452        use std::os::unix::fs::PermissionsExt;
453        let directory = tempfile::tempdir()?;
454        let path = directory.path().join("little-durable-objects-modal");
455        std::fs::write(
456            &path,
457            r#"#!/usr/bin/env node
458const readline = require('node:readline');
459let sequence = 0;
460async function reply(command, persistent) {
461  if (command.request.delay) await new Promise(resolve => setTimeout(resolve, command.request.delay));
462  const result = { pid: process.pid, sequence: ++sequence };
463  process.stdout.write(JSON.stringify(persistent ? { status: 'success', result } : result) + '\n');
464}
465if (process.argv.includes('--serve')) {
466  (async () => { for await (const line of readline.createInterface({input: process.stdin})) await reply(JSON.parse(line), true); })();
467} else {
468  let input = ''; process.stdin.on('data', chunk => input += chunk);
469  process.stdin.on('end', () => reply(JSON.parse(input), false));
470}
471"#,
472        )?;
473        std::fs::set_permissions(&path, std::fs::Permissions::from_mode(0o700))?;
474        let provider = CommandSandboxProvider::new(
475            "modal".into(),
476            path.display().to_string(),
477            HashMap::new(),
478        )?;
479        let first: serde_json::Value = provider.execute("test", &serde_json::json!({})).await?;
480        let second: serde_json::Value = provider.execute("test", &serde_json::json!({})).await?;
481        assert_eq!(first["pid"], second["pid"]);
482        assert_eq!(second["sequence"], 2);
483        assert!(
484            tokio::time::timeout(
485                Duration::from_millis(50),
486                provider
487                    .execute::<_, serde_json::Value>("test", &serde_json::json!({"delay": 500}))
488            )
489            .await
490            .is_err()
491        );
492        let next: serde_json::Value = provider.execute("test", &serde_json::json!({})).await?;
493        assert_ne!(first["pid"], next["pid"]);
494        assert_eq!(next["sequence"], 1);
495        let request = serde_json::json!({"delay": 50});
496        let (one, two, three) = tokio::try_join!(
497            provider.execute::<_, serde_json::Value>("test", &request),
498            provider.execute::<_, serde_json::Value>("test", &request),
499            provider.execute::<_, serde_json::Value>("test", &request),
500        )?;
501        let pids: std::collections::HashSet<_> = [one, two, three]
502            .into_iter()
503            .map(|value| value["pid"].as_u64().unwrap())
504            .collect();
505        assert_eq!(
506            pids.len(),
507            2,
508            "provider concurrency stays within the process limit"
509        );
510        Ok(())
511    }
512}