Skip to main content

aion_server/
state.rs

1//! Shared server state constructed once at startup.
2
3use 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/// Cloneable shared state passed to all server transports.
26#[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    /// Build shared state from operator configuration.
46    ///
47    /// # Errors
48    ///
49    /// Returns [`ServerError`] if the store cannot connect or the engine cannot
50    /// be constructed.
51    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    /// Build shared state from an already-constructed store.
58    ///
59    /// # Errors
60    ///
61    /// Returns [`ServerError::EngineCall`] if the engine cannot be constructed.
62    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        // The server unconditionally mounts /events/stream, so the engine's
74        // broadcast channel must be installed and explicitly sized here —
75        // a mounted-but-unconfigured streaming endpoint is never acceptable.
76        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        // The server unconditionally mounts /workflows/query, so the engine's
84        // query seam must be installed with an explicitly configured reply
85        // deadline here — a mounted-but-unconfigured query surface is never
86        // acceptable.
87        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    /// Build shared state from explicit parts with a default worker registry.
162    #[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    /// Build shared state from explicit parts with a caller-supplied JWKS cache.
182    ///
183    /// Embedders that construct their own [`JwksCache`] (for example against a
184    /// private issuer) can install it here; transports then validate bearer
185    /// tokens against it exactly as with a [`Self::build`]-constructed state.
186    #[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    /// Build shared state from explicit parts with a caller-supplied registry.
210    #[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    /// Borrow the namespace guard shared by all transports.
234    #[must_use]
235    pub fn namespace_guard(&self) -> &NamespaceGuard {
236        &self.inner.namespace_guard
237    }
238
239    /// Build the deploy authorization guard over the shared resolver.
240    #[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    /// Borrow non-secret runtime settings needed by transports.
246    #[must_use]
247    pub fn runtime_config(&self) -> &RuntimeConfig {
248        &self.inner.runtime
249    }
250
251    /// Borrow the connected-worker registry shared by worker transports and dispatch.
252    #[must_use]
253    pub fn worker_registry(&self) -> &ConnectedWorkerRegistry {
254        &self.inner.worker_registry
255    }
256
257    /// Borrow the pending-activities tracker shared by the NIF bridge and worker stream handler.
258    #[must_use]
259    pub fn pending_activities(&self) -> &PendingActivities {
260        &self.inner.pending_activities
261    }
262
263    /// Borrow the heartbeat/liveness tracker shared by dispatch and worker streams.
264    #[must_use]
265    pub fn heartbeat_tracker(&self) -> &HeartbeatTracker {
266        &self.inner.heartbeat_tracker
267    }
268
269    /// Borrow the drain gate shared by transports and worker dispatch.
270    #[must_use]
271    pub fn drain_state(&self) -> &DrainState {
272        &self.inner.drain_state
273    }
274
275    /// Borrow the prometheus metrics handle when this state was built with a store.
276    #[must_use]
277    pub fn metrics(&self) -> Option<&Metrics> {
278        self.inner.metrics.as_ref()
279    }
280
281    /// Borrow health probe state when this state was built with a store.
282    #[must_use]
283    pub fn health(&self) -> Option<&HealthState> {
284        self.inner.health.as_ref()
285    }
286
287    /// Borrow the shared JWKS cache when authentication is enabled.
288    #[cfg(feature = "auth")]
289    #[must_use]
290    pub fn jwks_cache(&self) -> Option<&JwksCache> {
291        self.inner.jwks_cache.as_ref()
292    }
293
294    /// Shut down the embedded engine so in-flight durable appends can finish.
295    ///
296    /// # Errors
297    ///
298    /// Returns [`ServerError`] if the namespace resolver has no engine handle or the engine rejects
299    /// shutdown.
300    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}