1use std::{collections::HashMap, env, future::Future, net::SocketAddr, sync::Arc, time::Duration};
2
3use anyhow::{Context, Result, ensure};
4use tracing::info;
5
6use crate::{
7 host_leases::PostgresHostLeaseStore,
8 placement::PostgresObjectPlacementStore,
9 postgres::PostgresDatabase,
10 sandbox::{CommandSandboxProvider, HostSandboxRuntimeConfig},
11 storage_urls::{GcsStorageUrlSigner, validate_buckets},
12};
13
14use super::{ActorJwtVerifier, ControlPlaneService};
15
16const DEFAULT_JWT_ISSUER: &str = "durable-object-control-plane";
17const DEFAULT_AUTHORITY_AUDIENCE: &str = "durable-object-authority";
18const DEFAULT_INVOCATION_AUDIENCE: &str = "durable-object-invoke";
19const DEFAULT_JWT_TTL_SECONDS: u64 = 1_800;
20const DEFAULT_ACTOR_IDLE_TIMEOUT_MS: u64 = 60_000;
21const DEFAULT_HOST_IDLE_TIMEOUT_MS: u64 = 300_000;
22const MAX_IDLE_TIMEOUT_MS: u64 = 86_400_000;
23
24pub struct ControlPlaneProcessConfig {
25 pub bind: SocketAddr,
26 pub jwt_signing_key: String,
27 pub jwt_key_id: String,
28 pub jwt_issuer: String,
29 pub authority_audience: String,
30 pub invocation_audience: String,
31 pub jwt_max_lifetime: Duration,
32 pub admin_token: String,
33 pub storage: ControlPlaneStorageConfig,
34 pub sandbox_provider: SandboxProviderConfig,
35 pub socket_event_sink: Option<SocketEventSinkConfig>,
36 pub socket_authenticator: Option<SocketAuthenticatorConfig>,
37}
38
39pub struct ControlPlaneStorageConfig {
40 pub postgres_url: String,
41 pub standard_buckets: HashMap<String, String>,
42}
43
44pub struct SandboxProviderConfig {
45 pub provider_name: String,
46 pub command: String,
47 pub environment: HashMap<String, String>,
48 pub runtime: HostSandboxRuntimeConfig,
49}
50
51pub struct SocketEventSinkConfig {
52 pub url: String,
53 pub token: String,
54}
55
56pub struct SocketAuthenticatorConfig {
57 pub url: String,
58 pub token: String,
59}
60
61impl ControlPlaneProcessConfig {
62 pub fn from_env() -> Result<Self> {
63 Self::from_lookup(|name| env::var(name).ok())
64 }
65}
66
67pub async fn serve_control_plane(
68 config: ControlPlaneProcessConfig,
69 shutdown: impl Future<Output = ()> + Send + 'static,
70) -> Result<()> {
71 let bind = config.bind;
72 let routes = control_plane_routes(config).await?;
73 info!(bind = %bind, "durable-object control plane is ready");
74 let listener = tokio::net::TcpListener::bind(bind)
75 .await
76 .context("bind durable-object control plane")?;
77 serve_routes(listener, routes, shutdown).await
78}
79
80async fn serve_routes(
81 listener: tokio::net::TcpListener,
82 routes: tonic::service::Routes,
83 shutdown: impl Future<Output = ()> + Send + 'static,
84) -> Result<()> {
85 axum::serve(listener, routes.into_axum_router())
86 .with_graceful_shutdown(shutdown)
87 .await
88 .context("serve durable-object control plane")
89}
90
91async fn control_plane_routes(config: ControlPlaneProcessConfig) -> Result<tonic::service::Routes> {
92 let issuer = super::ActorJwtIssuer::from_base64_pkcs8(
93 &config.jwt_signing_key,
94 config.jwt_key_id,
95 config.jwt_issuer.clone(),
96 config.authority_audience.clone(),
97 config.invocation_audience,
98 config.jwt_max_lifetime,
99 )?;
100 let auth = ActorJwtVerifier::for_scope(
101 issuer.verifier_keys_json()?,
102 config.jwt_issuer,
103 config.authority_audience,
104 super::ActorTokenPurpose::ControlPlane,
105 config.jwt_max_lifetime,
106 )?;
107 let database = PostgresDatabase::connect(&config.storage.postgres_url).await?;
108 let leases = Arc::new(PostgresHostLeaseStore::from_database(database.clone()));
109 let placements = Arc::new(PostgresObjectPlacementStore::from_database(
110 database.clone(),
111 ));
112 let registry = Arc::new(super::PostgresAdminRegistry::from_database(database));
113 let storage_urls = Arc::new(GcsStorageUrlSigner::from_adc(
114 config.storage.standard_buckets,
115 )?);
116 let provisioner = sandbox_provisioner(config.sandbox_provider, &issuer, &leases)?;
117 let socket_events = config
118 .socket_event_sink
119 .map(|sink| super::event_sink::HttpSocketMessageEventSink::new(sink.url, sink.token))
120 .transpose()?
121 .map(|sink| Arc::new(sink) as Arc<dyn super::event_sink::SocketMessageEventSink>);
122 let socket_authenticator = config
123 .socket_authenticator
124 .map(|auth| super::socket_auth::HttpSocketAuthenticator::new(auth.url, auth.token))
125 .transpose()?
126 .map(|auth| Arc::new(auth) as Arc<dyn super::socket_auth::SocketAuthenticator>);
127 let service = ControlPlaneService::new(
128 leases,
129 placements,
130 storage_urls,
131 auth,
132 registry.clone(),
133 issuer.clone(),
134 provisioner,
135 )
136 .with_socket_event_sink(socket_events)
137 .with_socket_authenticator(socket_authenticator);
138 let admin = super::admin::AdminService::new(config.admin_token, registry, issuer)?;
139 let public_api = super::public_api::router(service.clone(), admin);
140 let internal_api = service.into_internal_service();
141 Ok(tonic::service::Routes::from(public_api).add_service(internal_api))
142}
143
144fn sandbox_provisioner(
145 config: SandboxProviderConfig,
146 issuer: &super::ActorJwtIssuer,
147 leases: &Arc<PostgresHostLeaseStore>,
148) -> Result<Arc<dyn super::service::HostProvisioner>> {
149 let provider = Arc::new(CommandSandboxProvider::new(
150 config.provider_name,
151 config.command,
152 config.environment,
153 )?);
154 Ok(Arc::new(super::service::SandboxHostProvisioner::new(
155 provider,
156 config.runtime,
157 issuer.clone(),
158 leases.clone(),
159 )))
160}
161
162impl ControlPlaneProcessConfig {
163 fn from_lookup(mut get: impl FnMut(&str) -> Option<String>) -> Result<Self> {
164 let bind = get("DURABLE_OBJECT_CONTROL_PLANE_BIND")
165 .unwrap_or_else(|| "127.0.0.1:7100".into())
166 .parse()
167 .context("DURABLE_OBJECT_CONTROL_PLANE_BIND must be a socket address")?;
168 let jwt_signing_key = required(&mut get, "DURABLE_OBJECT_JWT_SIGNING_KEY")?;
169 let jwt_key_id = get("DURABLE_OBJECT_JWT_KEY_ID").unwrap_or_else(|| "primary".into());
170 let jwt_issuer =
171 get("DURABLE_OBJECT_JWT_ISSUER").unwrap_or_else(|| DEFAULT_JWT_ISSUER.into());
172 let authority_audience = get("DURABLE_OBJECT_AUTHORITY_JWT_AUDIENCE")
173 .unwrap_or_else(|| DEFAULT_AUTHORITY_AUDIENCE.into());
174 let invocation_audience = get("DURABLE_OBJECT_INVOKE_JWT_AUDIENCE")
175 .unwrap_or_else(|| DEFAULT_INVOCATION_AUDIENCE.into());
176 let jwt_max_lifetime = Duration::from_secs(
177 get("DURABLE_OBJECT_JWT_MAX_TTL_SECONDS")
178 .map(|value| value.parse())
179 .transpose()
180 .context("DURABLE_OBJECT_JWT_MAX_TTL_SECONDS must be an integer")?
181 .unwrap_or(DEFAULT_JWT_TTL_SECONDS),
182 );
183 ensure!(
184 !jwt_max_lifetime.is_zero(),
185 "DURABLE_OBJECT_JWT_MAX_TTL_SECONDS must be positive"
186 );
187 let admin_token = required(&mut get, "DURABLE_OBJECT_ADMIN_TOKEN")?;
188 ensure!(
189 admin_token.trim() == admin_token,
190 "DURABLE_OBJECT_ADMIN_TOKEN has surrounding whitespace"
191 );
192 let standard_buckets: HashMap<String, String> =
193 serde_json::from_str(&required(&mut get, "DURABLE_OBJECT_STANDARD_BUCKETS")?)
194 .context("DURABLE_OBJECT_STANDARD_BUCKETS must be a JSON region-to-bucket map")?;
195 validate_buckets(&standard_buckets)?;
196 let storage = ControlPlaneStorageConfig {
197 postgres_url: required(&mut get, "DURABLE_OBJECT_POSTGRES_URL")?,
198 standard_buckets,
199 };
200 let sandbox_provider =
201 sandbox_provider_config(&mut get, &jwt_issuer, &invocation_audience)?;
202 let socket_event_sink = socket_event_sink_config(&mut get)?;
203 let socket_authenticator = socket_authenticator_config(&mut get)?;
204 Ok(Self {
205 bind,
206 jwt_signing_key,
207 jwt_key_id,
208 jwt_issuer,
209 authority_audience,
210 invocation_audience,
211 jwt_max_lifetime,
212 admin_token,
213 storage,
214 sandbox_provider,
215 socket_event_sink,
216 socket_authenticator,
217 })
218 }
219}
220
221fn socket_authenticator_config(
222 get: &mut impl FnMut(&str) -> Option<String>,
223) -> Result<Option<SocketAuthenticatorConfig>> {
224 let url = get("DURABLE_OBJECT_SOCKET_AUTH_URL");
225 let token = get("DURABLE_OBJECT_SOCKET_AUTH_TOKEN");
226 match (url, token) {
227 (None, None) => Ok(None),
228 (Some(url), Some(token)) => Ok(Some(SocketAuthenticatorConfig {
229 url: validated_http_url(&url, "DURABLE_OBJECT_SOCKET_AUTH_URL")?,
230 token,
231 })),
232 _ => anyhow::bail!(
233 "DURABLE_OBJECT_SOCKET_AUTH_URL and DURABLE_OBJECT_SOCKET_AUTH_TOKEN must be configured together"
234 ),
235 }
236}
237
238fn socket_event_sink_config(
239 get: &mut impl FnMut(&str) -> Option<String>,
240) -> Result<Option<SocketEventSinkConfig>> {
241 let url = get("DURABLE_OBJECT_SOCKET_EVENT_URL");
242 let token = get("DURABLE_OBJECT_SOCKET_EVENT_TOKEN");
243 match (url, token) {
244 (None, None) => Ok(None),
245 (Some(url), Some(token)) => Ok(Some(SocketEventSinkConfig {
246 url: validated_http_url(&url, "DURABLE_OBJECT_SOCKET_EVENT_URL")?,
247 token,
248 })),
249 _ => anyhow::bail!(
250 "DURABLE_OBJECT_SOCKET_EVENT_URL and DURABLE_OBJECT_SOCKET_EVENT_TOKEN must be configured together"
251 ),
252 }
253}
254
255fn sandbox_provider_config(
256 get: &mut impl FnMut(&str) -> Option<String>,
257 jwt_issuer: &str,
258 invocation_audience: &str,
259) -> Result<SandboxProviderConfig> {
260 let provider_name = required(get, "DURABLE_OBJECT_SANDBOX_PROVIDER")?;
261 ensure!(
262 provider_name == "modal",
263 "unsupported sandbox provider {provider_name:?}"
264 );
265 let environment = HashMap::from([
266 (
267 "MODAL_TOKEN_ID".into(),
268 provider_credential(get, "MODAL_TOKEN_ID")?,
269 ),
270 (
271 "MODAL_TOKEN_SECRET".into(),
272 provider_credential(get, "MODAL_TOKEN_SECRET")?,
273 ),
274 ]);
275 let control_plane_url = validated_http_url(
276 &required(get, "DURABLE_OBJECT_CONTROL_PLANE_URL")?,
277 "DURABLE_OBJECT_CONTROL_PLANE_URL",
278 )?;
279 Ok(SandboxProviderConfig {
280 provider_name,
281 command: get("DURABLE_OBJECT_SANDBOX_COMMAND")
282 .unwrap_or_else(|| "little-durable-objects-modal".into()),
283 environment,
284 runtime: HostSandboxRuntimeConfig {
285 control_plane_url,
286 jwt_issuer: jwt_issuer.into(),
287 invocation_jwt_audience: invocation_audience.into(),
288 actor_idle_timeout_ms: idle_timeout(
289 get,
290 "DURABLE_OBJECT_ACTOR_IDLE_TIMEOUT_MS",
291 DEFAULT_ACTOR_IDLE_TIMEOUT_MS,
292 )?,
293 host_idle_timeout_ms: idle_timeout(
294 get,
295 "DURABLE_OBJECT_HOST_IDLE_TIMEOUT_MS",
296 DEFAULT_HOST_IDLE_TIMEOUT_MS,
297 )?,
298 },
299 })
300}
301
302fn provider_credential(get: &mut impl FnMut(&str) -> Option<String>, name: &str) -> Result<String> {
303 let value = required(get, name)?;
304 ensure!(value.trim() == value, "{name} has surrounding whitespace");
305 Ok(value)
306}
307
308fn required(get: &mut impl FnMut(&str) -> Option<String>, name: &str) -> Result<String> {
309 let value = get(name).with_context(|| format!("{name} is required"))?;
310 ensure!(!value.is_empty(), "{name} must not be empty");
311 Ok(value)
312}
313
314fn validated_http_url(value: &str, name: &str) -> Result<String> {
315 let url = reqwest::Url::parse(value).with_context(|| format!("{name} must be a URL"))?;
316 ensure!(
317 matches!(url.scheme(), "http" | "https") && url.host_str().is_some(),
318 "{name} must be HTTP or HTTPS"
319 );
320 Ok(url.to_string())
321}
322
323fn idle_timeout(
324 get: &mut impl FnMut(&str) -> Option<String>,
325 name: &str,
326 default: u64,
327) -> Result<u64> {
328 let value = get(name)
329 .map(|value| value.parse())
330 .transpose()
331 .with_context(|| format!("{name} must be an integer"))?
332 .unwrap_or(default);
333 ensure!(
334 (1..=MAX_IDLE_TIMEOUT_MS).contains(&value),
335 "{name} is outside the supported range"
336 );
337 Ok(value)
338}
339
340#[cfg(test)]
341mod tests {
342 use axum::{Router, extract::WebSocketUpgrade, response::Response, routing::get};
343 use futures_util::{SinkExt, StreamExt};
344 use tokio::sync::oneshot;
345 use tokio_tungstenite::{connect_async, tungstenite::Message};
346
347 use super::*;
348
349 #[tokio::test]
350 async fn server_carries_websocket_upgrades() -> Result<()> {
351 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await?;
352 let address = listener.local_addr()?;
353 let routes =
354 tonic::service::Routes::from(Router::new().route("/socket", get(echo_websocket)));
355 let (shutdown_tx, shutdown_rx) = oneshot::channel();
356 let server = tokio::spawn(serve_routes(listener, routes, async {
357 let _ = shutdown_rx.await;
358 }));
359
360 let (mut socket, _) = connect_async(format!("ws://{address}/socket")).await?;
361 socket.send(Message::Text("hello".into())).await?;
362 assert_eq!(
363 socket.next().await.transpose()?,
364 Some(Message::Text("hello".into()))
365 );
366 socket.close(None).await?;
367 let _ = shutdown_tx.send(());
368 server.await??;
369 Ok(())
370 }
371
372 #[test]
373 fn parses_the_minimal_storage_configuration() -> Result<()> {
374 let values = HashMap::from([
375 ("DURABLE_OBJECT_JWT_SIGNING_KEY", "c2lnbmluZw=="),
376 ("DURABLE_OBJECT_ADMIN_TOKEN", "admin-token"),
377 ("DURABLE_OBJECT_SANDBOX_PROVIDER", "modal"),
378 (
379 "DURABLE_OBJECT_CONTROL_PLANE_URL",
380 "https://objects.example.com",
381 ),
382 ("MODAL_TOKEN_ID", "modal-token-id"),
383 ("MODAL_TOKEN_SECRET", "modal-token-secret"),
384 (
385 "DURABLE_OBJECT_POSTGRES_URL",
386 "postgresql://localhost/actors",
387 ),
388 (
389 "DURABLE_OBJECT_STANDARD_BUCKETS",
390 "{\"us-east\":\"actor-state-test\"}",
391 ),
392 ]);
393 let config = ControlPlaneProcessConfig::from_lookup(|name| {
394 values.get(name).map(|value| (*value).into())
395 })?;
396 assert_eq!(
397 config.storage.standard_buckets["us-east"],
398 "actor-state-test"
399 );
400 Ok(())
401 }
402
403 #[test]
404 fn parses_socket_event_sink_only_when_url_and_token_are_present() -> Result<()> {
405 let mut complete = HashMap::from([
406 (
407 "DURABLE_OBJECT_SOCKET_EVENT_URL",
408 "https://api.example.com/events",
409 ),
410 ("DURABLE_OBJECT_SOCKET_EVENT_TOKEN", "event-token"),
411 ]);
412 let sink =
413 socket_event_sink_config(&mut |name| complete.get(name).map(|value| (*value).into()))?
414 .context("socket event sink was not configured")?;
415 assert_eq!(sink.url, "https://api.example.com/events");
416 assert_eq!(sink.token, "event-token");
417
418 complete.remove("DURABLE_OBJECT_SOCKET_EVENT_TOKEN");
419 assert!(
420 socket_event_sink_config(&mut |name| complete.get(name).map(|value| (*value).into()))
421 .is_err()
422 );
423 Ok(())
424 }
425
426 #[test]
427 fn parses_socket_authenticator_only_when_url_and_token_are_present() -> Result<()> {
428 let mut complete = HashMap::from([
429 (
430 "DURABLE_OBJECT_SOCKET_AUTH_URL",
431 "https://api.example.com/authorize",
432 ),
433 ("DURABLE_OBJECT_SOCKET_AUTH_TOKEN", "auth-token"),
434 ]);
435 let auth = socket_authenticator_config(&mut |name| {
436 complete.get(name).map(|value| (*value).into())
437 })?
438 .context("socket authenticator was not configured")?;
439 assert_eq!(auth.url, "https://api.example.com/authorize");
440 assert_eq!(auth.token, "auth-token");
441
442 complete.remove("DURABLE_OBJECT_SOCKET_AUTH_TOKEN");
443 assert!(
444 socket_authenticator_config(&mut |name| complete
445 .get(name)
446 .map(|value| (*value).into()))
447 .is_err()
448 );
449 Ok(())
450 }
451
452 async fn echo_websocket(upgrade: WebSocketUpgrade) -> Response {
453 upgrade.on_upgrade(async |mut socket| {
454 if let Some(Ok(message)) = socket.recv().await {
455 let _ = socket.send(message).await;
456 }
457 })
458 }
459}