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}