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}