1use std::{path::PathBuf, sync::Arc};
4
5use aion::{EngineBuilder, RuntimeHandle, SignalRouter, signal::ConcreteSignalRouter};
6use aion_store::EventStore;
7use aion_store_libsql::LibSqlStore;
8
9#[cfg(feature = "auth")]
10use crate::auth::JwksCache;
11use crate::{
12 config::{RuntimeConfig, ServerConfig, StoreBackend, StoreConfig},
13 error::ServerError,
14 namespace::{NamespaceGuard, resolver::NamespaceResolver},
15 observability::{
16 Metrics, health::HealthState, instrumented_store::InstrumentedEventStore,
17 metrics::MetricsError,
18 },
19 shutdown::DrainState,
20 worker::{
21 ConnectedWorkerRegistry, HeartbeatTracker, PendingActivities, WorkerActivityDispatcher,
22 },
23};
24
25#[derive(Clone)]
27pub struct ServerState {
28 inner: Arc<ServerStateInner>,
29}
30
31struct ServerStateInner {
32 namespace_guard: NamespaceGuard,
33 runtime: RuntimeConfig,
34 worker_registry: ConnectedWorkerRegistry,
35 pending_activities: PendingActivities,
36 heartbeat_tracker: HeartbeatTracker,
37 drain_state: DrainState,
38 metrics: Option<Metrics>,
39 health: Option<HealthState>,
40 #[cfg(feature = "auth")]
41 jwks_cache: Option<JwksCache>,
42}
43
44impl ServerState {
45 pub async fn build(config: ServerConfig) -> Result<Self, ServerError> {
52 let (store_config, runtime) = config.into_parts();
53 let store = connect_store(store_config).await?;
54 Self::build_with_store_arc(store, runtime).await
55 }
56
57 pub async fn build_with_store<S>(store: S, runtime: RuntimeConfig) -> Result<Self, ServerError>
63 where
64 S: EventStore,
65 {
66 Self::build_with_store_arc(Arc::new(store), runtime).await
67 }
68
69 async fn build_with_store_arc(
70 store: Arc<dyn EventStore>,
71 runtime: RuntimeConfig,
72 ) -> Result<Self, ServerError> {
73 let event_broadcast_capacity = runtime
77 .websocket
78 .event_broadcast_capacity
79 .and_then(std::num::NonZeroUsize::new)
80 .ok_or_else(|| ServerError::Config {
81 message: crate::config::EVENT_BROADCAST_CAPACITY_REQUIRED.to_owned(),
82 })?;
83 let query_timeout = runtime
88 .query_timeout
89 .filter(|timeout| !timeout.is_zero())
90 .ok_or_else(|| ServerError::Config {
91 message: crate::config::QUERY_TIMEOUT_REQUIRED.to_owned(),
92 })?;
93 let metrics = Metrics::new().map_err(|error| metrics_config_error(&error))?;
94 let instrumented_store = Arc::new(InstrumentedEventStore::new(
95 store.clone(),
96 metrics.clone(),
97 runtime.default_namespace.clone(),
98 ));
99 let exported_metrics = runtime.metrics.enabled.then_some(metrics.clone());
100 let worker_registry = ConnectedWorkerRegistry::default();
101 let active_registry = Arc::new(aion::Registry::default());
102 let pending_activities = PendingActivities::default();
103 let heartbeat_tracker = HeartbeatTracker::new(runtime.worker.heartbeat_window);
104 let drain_state = DrainState::default();
105 let dispatcher = WorkerActivityDispatcher::new(
106 worker_registry.clone(),
107 runtime.default_namespace.clone(),
108 heartbeat_tracker.clone(),
109 )
110 .with_pending(pending_activities.clone())
111 .with_drain_state(drain_state.clone())
112 .with_workflow_registry(active_registry.clone())
113 .with_tokio_handle(tokio::runtime::Handle::current());
114 let dispatcher = Arc::new(dispatcher);
115
116 let mut search_attribute_schema = aion_core::SearchAttributeSchema::new();
117 search_attribute_schema
118 .register(
119 crate::namespace::NAMESPACE_ATTRIBUTE,
120 aion_core::SearchAttributeType::String,
121 )
122 .map_err(|error| ServerError::Config {
123 message: format!("failed to register namespace search attribute: {error}"),
124 })?;
125 let engine = EngineBuilder::new()
126 .store_arc(instrumented_store.clone())
127 .event_streaming(event_broadcast_capacity)
128 .in_memory_visibility()
129 .search_attribute_schema(search_attribute_schema)
130 .scheduler_threads(runtime.scheduler_threads)
131 .activity_dispatcher(dispatcher)
132 .active_registry(active_registry)
133 .production_recovery_seam()
134 .signal_router_factory(|runtime: Arc<RuntimeHandle>, handoff| {
135 Arc::new(ConcreteSignalRouter::new(runtime, handoff)) as Arc<dyn SignalRouter>
136 })
137 .query_timeout(query_timeout)
138 .load_workflow_sources(runtime.workflow_packages.iter().map(PathBuf::as_path))
139 .build()
140 .await?;
141 let engine = Arc::new(engine);
142 let namespace_resolver = NamespaceResolver::from_config(runtime.namespace.clone(), engine);
143 #[cfg(feature = "auth")]
144 let jwks_cache = build_jwks_cache(&runtime).await?;
145 Ok(Self {
146 inner: Arc::new(ServerStateInner {
147 namespace_guard: NamespaceGuard::new(namespace_resolver),
148 runtime,
149 worker_registry,
150 pending_activities,
151 heartbeat_tracker,
152 drain_state,
153 metrics: exported_metrics,
154 health: Some(HealthState::new(instrumented_store, true)),
155 #[cfg(feature = "auth")]
156 jwks_cache,
157 }),
158 })
159 }
160
161 #[must_use]
163 pub fn from_parts(namespace_resolver: NamespaceResolver, runtime: RuntimeConfig) -> Self {
164 let heartbeat_tracker = HeartbeatTracker::new(runtime.worker.heartbeat_window);
165 Self {
166 inner: Arc::new(ServerStateInner {
167 namespace_guard: NamespaceGuard::new(namespace_resolver),
168 runtime,
169 worker_registry: ConnectedWorkerRegistry::default(),
170 pending_activities: PendingActivities::default(),
171 heartbeat_tracker,
172 drain_state: DrainState::default(),
173 metrics: None,
174 health: None,
175 #[cfg(feature = "auth")]
176 jwks_cache: None,
177 }),
178 }
179 }
180
181 #[cfg(feature = "auth")]
187 #[must_use]
188 pub fn from_parts_with_jwks(
189 namespace_resolver: NamespaceResolver,
190 runtime: RuntimeConfig,
191 jwks_cache: JwksCache,
192 ) -> Self {
193 let heartbeat_tracker = HeartbeatTracker::new(runtime.worker.heartbeat_window);
194 Self {
195 inner: Arc::new(ServerStateInner {
196 namespace_guard: NamespaceGuard::new(namespace_resolver),
197 runtime,
198 worker_registry: ConnectedWorkerRegistry::default(),
199 pending_activities: PendingActivities::default(),
200 heartbeat_tracker,
201 drain_state: DrainState::default(),
202 metrics: None,
203 health: None,
204 jwks_cache: Some(jwks_cache),
205 }),
206 }
207 }
208
209 #[must_use]
211 pub fn from_parts_with_registry(
212 namespace_resolver: NamespaceResolver,
213 runtime: RuntimeConfig,
214 worker_registry: ConnectedWorkerRegistry,
215 ) -> Self {
216 let heartbeat_tracker = HeartbeatTracker::new(runtime.worker.heartbeat_window);
217 Self {
218 inner: Arc::new(ServerStateInner {
219 namespace_guard: NamespaceGuard::new(namespace_resolver),
220 runtime,
221 worker_registry,
222 pending_activities: PendingActivities::default(),
223 heartbeat_tracker,
224 drain_state: DrainState::default(),
225 metrics: None,
226 health: None,
227 #[cfg(feature = "auth")]
228 jwks_cache: None,
229 }),
230 }
231 }
232
233 #[must_use]
235 pub fn namespace_guard(&self) -> &NamespaceGuard {
236 &self.inner.namespace_guard
237 }
238
239 #[must_use]
241 pub fn deploy_guard(&self) -> crate::deploy::DeployGuard {
242 crate::deploy::DeployGuard::new(self.inner.namespace_guard.resolver().clone())
243 }
244
245 #[must_use]
247 pub fn runtime_config(&self) -> &RuntimeConfig {
248 &self.inner.runtime
249 }
250
251 #[must_use]
253 pub fn worker_registry(&self) -> &ConnectedWorkerRegistry {
254 &self.inner.worker_registry
255 }
256
257 #[must_use]
259 pub fn pending_activities(&self) -> &PendingActivities {
260 &self.inner.pending_activities
261 }
262
263 #[must_use]
265 pub fn heartbeat_tracker(&self) -> &HeartbeatTracker {
266 &self.inner.heartbeat_tracker
267 }
268
269 #[must_use]
271 pub fn drain_state(&self) -> &DrainState {
272 &self.inner.drain_state
273 }
274
275 #[must_use]
277 pub fn metrics(&self) -> Option<&Metrics> {
278 self.inner.metrics.as_ref()
279 }
280
281 #[must_use]
283 pub fn health(&self) -> Option<&HealthState> {
284 self.inner.health.as_ref()
285 }
286
287 #[cfg(feature = "auth")]
289 #[must_use]
290 pub fn jwks_cache(&self) -> Option<&JwksCache> {
291 self.inner.jwks_cache.as_ref()
292 }
293
294 pub fn shutdown(&self) -> Result<(), ServerError> {
301 self.inner.namespace_guard.resolver().shutdown_engine()
302 }
303}
304
305#[cfg(feature = "auth")]
306async fn build_jwks_cache(runtime: &RuntimeConfig) -> Result<Option<JwksCache>, ServerError> {
307 if !runtime.auth.enabled {
308 return Ok(None);
309 }
310 let Some(url) = runtime.auth.jwks_url.clone() else {
311 return Err(ServerError::Config {
312 message: "auth.jwks_url must not be empty when auth.enabled is true".to_owned(),
313 });
314 };
315 let interval = std::time::Duration::from_secs(runtime.auth.jwks_refresh_seconds);
316 let cache = JwksCache::new(url, interval)
317 .await
318 .map_err(|error| ServerError::Config {
319 message: format!("auth jwks initial fetch failed: {error}"),
320 })?;
321 Ok(Some(cache))
322}
323
324fn metrics_config_error(error: &MetricsError) -> ServerError {
325 ServerError::Config {
326 message: error.to_string(),
327 }
328}
329
330async fn connect_store(config: StoreConfig) -> Result<Arc<dyn EventStore>, ServerError> {
331 match config.backend {
332 StoreBackend::Memory => Ok(Arc::new(aion_store::InMemoryStore::default())),
333 StoreBackend::LibSql => {
334 let Some(url) = config.url else {
335 return Err(ServerError::Config {
336 message: "store.url must not be empty when store.backend is libsql".to_owned(),
337 });
338 };
339 let store = LibSqlStore::open(url.clone())
340 .await
341 .map_err(ServerError::from)?;
342 store
343 .validate_event_compatibility()
344 .await
345 .map_err(|error| match error {
346 aion_store::StoreError::Serialization(_) => ServerError::Config {
347 message: format!(
348 "Database schema mismatch — delete {url} and restart, or run migrations."
349 ),
350 },
351 other => ServerError::from(other),
352 })?;
353 Ok(Arc::new(store))
354 }
355 }
356}
357
358#[cfg(test)]
359mod tests {
360 use std::{net::SocketAddr, time::Duration};
361
362 use aion_store::InMemoryStore;
363
364 use super::ServerState;
365 use crate::config::{
366 AuthConfig, DashboardAssetSource, DashboardConfig, DeployConfig, ListenConfig,
367 MetricsConfig, NamespaceConfig, NamespaceMode, RuntimeConfig, WebSocketConfig,
368 WorkerConfig,
369 };
370
371 fn runtime_config() -> RuntimeConfig {
372 RuntimeConfig {
373 listen: ListenConfig {
374 grpc: SocketAddr::from(([127, 0, 0, 1], 50051)),
375 http: SocketAddr::from(([127, 0, 0, 1], 8080)),
376 },
377 tls: None,
378 auth: AuthConfig {
379 enabled: false,
380 jwks_url: None,
381 jwks_refresh_seconds: 300,
382 },
383 dashboard: DashboardConfig {
384 source: DashboardAssetSource::Embedded,
385 },
386 namespace: NamespaceConfig {
387 mode: NamespaceMode::SharedEngine,
388 },
389 worker: WorkerConfig {
390 heartbeat_window: Duration::from_millis(30_000),
391 },
392 websocket: WebSocketConfig {
393 outbound_buffer_bound: 32,
394 event_broadcast_capacity: Some(64),
395 },
396 workflow_packages: Vec::new(),
397 deploy: DeployConfig::default(),
398 scheduler_threads: 1,
399 query_timeout: Some(Duration::from_millis(10_000)),
400 default_namespace: "default".to_owned(),
401 drain_timeout: Duration::from_secs(30),
402 metrics: MetricsConfig { enabled: true },
403 }
404 }
405
406 #[tokio::test]
407 async fn builds_state_with_in_memory_store() -> Result<(), Box<dyn std::error::Error>> {
408 let state =
409 ServerState::build_with_store(InMemoryStore::default(), runtime_config()).await?;
410
411 std::hint::black_box(state.namespace_guard());
412 std::hint::black_box(state.worker_registry());
413
414 Ok(())
415 }
416
417 #[tokio::test]
418 async fn state_build_fails_without_event_broadcast_capacity()
419 -> Result<(), Box<dyn std::error::Error>> {
420 let mut runtime = runtime_config();
421 runtime.websocket.event_broadcast_capacity = None;
422
423 let error = ServerState::build_with_store(InMemoryStore::default(), runtime)
424 .await
425 .err()
426 .ok_or("state build must fail when event streaming is unsized")?;
427
428 assert!(error.is_config(), "expected a config error, got {error}");
429 assert!(
430 error
431 .to_string()
432 .contains("websocket.event_broadcast_capacity"),
433 "error must name the missing key: {error}"
434 );
435 Ok(())
436 }
437
438 #[tokio::test]
439 async fn state_build_fails_without_query_timeout() -> Result<(), Box<dyn std::error::Error>> {
440 let mut runtime = runtime_config();
441 runtime.query_timeout = None;
442
443 let error = ServerState::build_with_store(InMemoryStore::default(), runtime)
444 .await
445 .err()
446 .ok_or("state build must fail when the query reply deadline is unset")?;
447
448 assert!(error.is_config(), "expected a config error, got {error}");
449 assert!(
450 error.to_string().contains("runtime.query_timeout_ms"),
451 "error must name the missing key: {error}"
452 );
453 assert!(
454 error.to_string().contains("AION_RUNTIME_QUERY_TIMEOUT_MS"),
455 "error must name the environment override: {error}"
456 );
457 Ok(())
458 }
459
460 #[tokio::test]
461 async fn state_build_fails_with_zero_query_timeout() -> Result<(), Box<dyn std::error::Error>> {
462 let mut runtime = runtime_config();
463 runtime.query_timeout = Some(Duration::ZERO);
464
465 let error = ServerState::build_with_store(InMemoryStore::default(), runtime)
466 .await
467 .err()
468 .ok_or("state build must fail when the query reply deadline is zero")?;
469
470 assert!(error.is_config(), "expected a config error, got {error}");
471 assert!(
472 error.to_string().contains("runtime.query_timeout_ms"),
473 "error must name the zero-valued key: {error}"
474 );
475 Ok(())
476 }
477}