Skip to main content

fraiseql_server/server/
lifecycle.rs

1//! Server lifecycle: serve, `serve_with_shutdown`, and `shutdown_signal`.
2
3use std::net::SocketAddr;
4
5use axum::serve::ListenerExt;
6use tokio::net::TcpListener;
7use tracing::{error, info, warn};
8
9use super::{DatabaseAdapter, Result, Server, ServerError, TlsSetup};
10#[cfg(feature = "observers")]
11use crate::subscriptions::event_bridge::{EventBridge, EventBridgeConfig};
12
13impl<A: DatabaseAdapter + Clone + Send + Sync + 'static> Server<A> {
14    /// Start server and listen for requests.
15    ///
16    /// Uses SIGUSR1-aware shutdown signal when a schema path is configured,
17    /// enabling zero-downtime schema reloads via `kill -USR1 <pid>`.
18    ///
19    /// # Errors
20    ///
21    /// Returns error if server fails to bind or encounters runtime errors.
22    pub async fn serve(self) -> Result<()> {
23        self.serve_with_shutdown(Self::shutdown_signal()).await
24    }
25
26    /// Start server with a custom shutdown future.
27    ///
28    /// Enables programmatic shutdown (e.g., for `--watch` hot-reload) by accepting any
29    /// future that resolves when the server should stop.
30    ///
31    /// # Errors
32    ///
33    /// Returns error if server fails to bind or encounters runtime errors.
34    #[allow(clippy::cognitive_complexity)] // Reason: server lifecycle with TLS/non-TLS binding, signal handling, and graceful shutdown
35    pub async fn serve_with_shutdown<F>(mut self, shutdown: F) -> Result<()>
36    where
37        F: std::future::Future<Output = ()> + Send + 'static,
38    {
39        // Ensure RBAC schema exists before the router mounts RBAC endpoints.
40        // Must run here (async context) rather than inside build_router() (sync).
41        #[cfg(feature = "observers")]
42        if let Some(ref db_pool) = self.db_pool {
43            if self.config.admin_token.is_some() {
44                let rbac_backend =
45                    crate::api::rbac_management::db_backend::RbacDbBackend::new(db_pool.clone());
46                rbac_backend.ensure_schema().await.map_err(|e| {
47                    ServerError::ConfigError(format!("Failed to initialize RBAC schema: {e}"))
48                })?;
49            }
50        }
51
52        // Initialize usage persistence backend if configured.
53        // Must run before build_router() so the aggregator is populated before
54        // serving requests, but after the DB pool is available (async context).
55        if let Some(ref usage_cfg) = self.config.usage.clone() {
56            use std::time::Duration;
57
58            use sqlx::postgres::PgPoolOptions;
59            use tokio::time::MissedTickBehavior;
60
61            use crate::usage::aggregator::{PostgresBackend, global_aggregator};
62
63            match PgPoolOptions::new()
64                .max_connections(2) // small dedicated pool — only used for periodic flushes
65                .connect(&self.config.database_url)
66                .await
67            {
68                Ok(pool) => {
69                    match PostgresBackend::new(pool).await {
70                        Ok(backend) => {
71                            let backend = std::sync::Arc::new(backend);
72                            // Upgrade global aggregator's backend from NoopBackend.
73                            global_aggregator().set_backend(backend.clone());
74                            // Restore persisted counters before serving requests.
75                            if let Err(e) = global_aggregator().load_from_backend().await {
76                                warn!(error = %e, "Usage persistence: startup load failed — continuing with in-memory counters");
77                            } else {
78                                info!("Usage persistence: loaded counters from PostgreSQL");
79                            }
80                            // Spawn background flush task on the server's JoinSet
81                            // so graceful shutdown can await its termination.
82                            let flush_interval = Duration::from_secs(usage_cfg.flush_interval_secs);
83                            let agg = std::sync::Arc::clone(global_aggregator());
84                            self.tasks.spawn(async move {
85                                let mut ticker = tokio::time::interval(flush_interval);
86                                ticker.set_missed_tick_behavior(MissedTickBehavior::Skip);
87                                ticker.tick().await; // skip immediate first tick
88                                loop {
89                                    ticker.tick().await;
90                                    if let Err(e) = agg.flush_to_backend().await {
91                                        warn!(error = %e, "Usage persistence: background flush failed");
92                                    }
93                                }
94                            });
95                            info!(
96                                flush_interval_secs = usage_cfg.flush_interval_secs,
97                                "Usage persistence: PostgreSQL backend active"
98                            );
99                        },
100                        Err(e) => {
101                            warn!(
102                                error = %e,
103                                "Usage persistence: PostgresBackend initialization failed — \
104                                 continuing with in-memory (NoopBackend)"
105                            );
106                        },
107                    }
108                },
109                Err(e) => {
110                    warn!(
111                        error = %e,
112                        "Usage persistence: failed to connect to PostgreSQL — \
113                         continuing with in-memory (NoopBackend)"
114                    );
115                },
116            }
117        }
118
119        let (app, app_state) = self.build_router();
120
121        // Spawn SIGUSR1 schema reload handler when running on Unix.
122        // The handler loops forever, reloading on each signal, until the
123        // server process exits — tracked on the server's JoinSet so graceful
124        // shutdown awaits its termination.
125        #[cfg(unix)]
126        if let Some(ref schema_path) = app_state.schema_path {
127            let reload_state = app_state.clone();
128            let reload_path = schema_path.clone();
129            self.tasks.spawn(async move {
130                let mut sigusr1 = match tokio::signal::unix::signal(
131                    tokio::signal::unix::SignalKind::user_defined1(),
132                ) {
133                    Ok(s) => s,
134                    Err(e) => {
135                        warn!(error = %e, "Failed to install SIGUSR1 handler — schema hot-reload disabled");
136                        return;
137                    },
138                };
139                loop {
140                    sigusr1.recv().await;
141                    info!(
142                        path = %reload_path.display(),
143                        "Received SIGUSR1 — reloading schema"
144                    );
145                    match reload_state.reload_schema(&reload_path).await {
146                        Ok(()) => {
147                            let hash = reload_state.executor().schema().content_hash();
148                            reload_state
149                                .metrics
150                                .schema_reloads_total
151                                .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
152                            info!(schema_hash = %hash, "Schema reloaded successfully via SIGUSR1");
153                        },
154                        Err(e) => {
155                            reload_state
156                                .metrics
157                                .schema_reload_errors_total
158                                .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
159                            error!(
160                                error = %e,
161                                path = %reload_path.display(),
162                                "Schema reload failed via SIGUSR1 — keeping previous schema"
163                            );
164                        },
165                    }
166                }
167            });
168            info!(
169                path = %schema_path.display(),
170                "SIGUSR1 schema reload handler installed"
171            );
172        }
173
174        // Initialize TLS setup
175        let tls_setup = TlsSetup::new(self.config.tls.clone(), self.config.database_tls.clone())?;
176
177        info!(
178            bind_addr = %self.config.bind_addr,
179            graphql_path = %self.config.graphql_path,
180            tls_enabled = tls_setup.is_tls_enabled(),
181            "Starting FraiseQL server"
182        );
183
184        // Start observer runtime if configured, wiring CDC events to EventBridge
185        #[cfg(feature = "observers")]
186        #[allow(unused_variables)]
187        // Reason: _bridge_handle is kept alive to prevent task cancellation
188        let _bridge_handle = {
189            let mut handle: Option<tokio::task::JoinHandle<()>> = None;
190            if let Some(ref runtime) = self.observer_runtime {
191                info!("Starting observer runtime...");
192
193                // Create EventBridge to forward CDC events to GraphQL subscriptions
194                let bridge =
195                    EventBridge::new(self.subscription_manager.clone(), EventBridgeConfig::new());
196                let sender = bridge.sender();
197
198                let mut guard = runtime.write().await;
199                guard.set_event_bridge_sender(sender);
200
201                match guard.start().await {
202                    Ok(()) => {
203                        info!("Observer runtime started");
204                        // Spawn EventBridge after observer runtime is running
205                        handle = Some(bridge.spawn());
206                        info!(
207                            "EventBridge started — CDC events will be forwarded to subscriptions"
208                        );
209                    },
210                    Err(e) => {
211                        // A broker-backed transport (NATS) was an explicit operator
212                        // choice; refusing to boot in production rather than silently
213                        // coming up without it is the #350 dead-broker contract. The
214                        // default PostgreSQL transport keeps the resilient
215                        // log-and-continue behaviour (and development downgrades the
216                        // NATS failure to the same warning).
217                        if guard.transport_requires_broker()
218                            && crate::ServerConfig::is_production_mode()
219                        {
220                            error!(
221                                error = %e,
222                                "Observer runtime failed to start on its configured \
223                                 transport; refusing to boot (set FRAISEQL_ENV=development \
224                                 to downgrade to a warning)"
225                            );
226                            return Err(e);
227                        }
228                        error!("Failed to start observer runtime: {}", e);
229                        warn!("Server will continue without observers");
230                    },
231                }
232                drop(guard);
233            }
234            handle
235        };
236
237        // Explicitly enable TCP_NODELAY (disable Nagle's algorithm) on every
238        // accepted connection to minimise latency for small GraphQL responses.
239        let listener = TcpListener::bind(self.config.bind_addr)
240            .await
241            .map_err(|e| ServerError::BindError(e.to_string()))?
242            .tap_io(|tcp_stream| {
243                if let Err(err) = tcp_stream.set_nodelay(true) {
244                    warn!("failed to set TCP_NODELAY: {err:#}");
245                }
246            });
247
248        // Warn if the process file descriptor limit is below the recommended minimum.
249        // A low limit causes "too many open files" errors under load.
250        #[cfg(target_os = "linux")]
251        {
252            if let Ok(limits) = std::fs::read_to_string("/proc/self/limits") {
253                for line in limits.lines() {
254                    if line.starts_with("Max open files") {
255                        let parts: Vec<&str> = line.split_whitespace().collect();
256                        if let Some(soft) = parts.get(3) {
257                            if let Ok(n) = soft.parse::<u64>() {
258                                if n < 65_536 {
259                                    warn!(
260                                        current_fd_limit = n,
261                                        recommended = 65_536,
262                                        "File descriptor limit is low; consider raising ulimit -n"
263                                    );
264                                }
265                            }
266                        }
267                        break;
268                    }
269                }
270            }
271        }
272
273        // Log TLS configuration
274        if tls_setup.is_tls_enabled() {
275            // Verify TLS setup is valid (will error if certificates are missing/invalid)
276            let _ = tls_setup.create_rustls_config()?;
277            info!(
278                cert_path = ?tls_setup.cert_path(),
279                key_path = ?tls_setup.key_path(),
280                mtls_required = tls_setup.is_mtls_required(),
281                "Server TLS configuration loaded (note: use reverse proxy for server-side TLS termination)"
282            );
283        }
284
285        // Log database TLS configuration
286        info!(
287            postgres_ssl_mode = tls_setup.postgres_ssl_mode(),
288            redis_ssl = tls_setup.redis_ssl_enabled(),
289            clickhouse_https = tls_setup.clickhouse_https_enabled(),
290            elasticsearch_https = tls_setup.elasticsearch_https_enabled(),
291            "Database connection TLS configuration applied"
292        );
293
294        info!("Server listening on http://{}", self.config.bind_addr);
295
296        // Start both HTTP and gRPC servers concurrently if Arrow Flight is enabled
297        #[cfg(feature = "arrow")]
298        if let Some(flight_service) = self.flight_service.take() {
299            let flight_addr = self.config.flight_bind_addr;
300            info!("Arrow Flight server listening on grpc://{}", flight_addr);
301
302            // Spawn Flight server in background, registered on the server's
303            // JoinSet. The set's `shutdown` step abort-then-awaits the gRPC
304            // server when the HTTP server exits.
305            self.tasks.spawn(async move {
306                if let Err(e) = tonic::transport::Server::builder()
307                    .add_service(flight_service.into_server())
308                    .serve(flight_addr)
309                    .await
310                {
311                    error!(error = %e, "Arrow Flight server terminated with error");
312                }
313            });
314
315            // Wrap the user-supplied shutdown future so we can also stop observer runtime
316            #[cfg(feature = "observers")]
317            let observer_runtime = self.observer_runtime.clone();
318
319            let shutdown_with_cleanup = async move {
320                shutdown.await;
321                #[cfg(feature = "observers")]
322                if let Some(ref runtime) = observer_runtime {
323                    info!("Shutting down observer runtime");
324                    let mut guard = runtime.write().await;
325                    if let Err(e) = guard.stop().await {
326                        #[cfg(feature = "observers")]
327                        error!("Error stopping runtime: {}", e);
328                    } else {
329                        info!("Runtime stopped cleanly");
330                    }
331                }
332            };
333
334            // Run HTTP server with graceful shutdown
335            axum::serve(listener, app.into_make_service_with_connect_info::<SocketAddr>())
336                .with_graceful_shutdown(shutdown_with_cleanup)
337                .await
338                .map_err(|e| ServerError::IoError(std::io::Error::other(e)))?;
339
340            // Abort and await every lifecycle task (Flight server, SIGUSR1
341            // handler, PKCE cleanup, trusted-docs reload, usage flush, …).
342            drain_lifecycle_tasks(self.tasks, self.config.shutdown_timeout_secs).await;
343        }
344
345        // HTTP-only server (when arrow feature not enabled)
346        #[cfg(not(feature = "arrow"))]
347        {
348            axum::serve(listener, app.into_make_service_with_connect_info::<SocketAddr>())
349                .with_graceful_shutdown(shutdown)
350                .await
351                .map_err(|e| ServerError::IoError(std::io::Error::other(e)))?;
352
353            let shutdown_timeout =
354                std::time::Duration::from_secs(self.config.shutdown_timeout_secs);
355            info!(
356                timeout_secs = self.config.shutdown_timeout_secs,
357                "HTTP server stopped, draining remaining work"
358            );
359
360            let drain = tokio::time::timeout(shutdown_timeout, async {
361                #[cfg(feature = "observers")]
362                if let Some(ref runtime) = self.observer_runtime {
363                    let mut guard = runtime.write().await;
364                    match guard.stop().await {
365                        Ok(()) => info!("Observer runtime stopped cleanly"),
366                        Err(e) => warn!("Observer runtime shutdown error: {e}"),
367                    }
368                }
369            })
370            .await;
371
372            if drain.is_err() {
373                warn!(
374                    timeout_secs = self.config.shutdown_timeout_secs,
375                    "Shutdown drain timed out; forcing exit"
376                );
377            } else {
378                info!("Graceful shutdown complete");
379            }
380
381            // Abort and await every lifecycle task (SIGUSR1 handler, PKCE
382            // cleanup, trusted-docs reload, usage flush, …).
383            drain_lifecycle_tasks(self.tasks, self.config.shutdown_timeout_secs).await;
384        }
385
386        Ok(())
387    }
388
389    /// Start server on an externally created listener.
390    ///
391    /// Used in tests to discover the bound port before serving.
392    /// Skips TLS, Flight, and observer startup — suitable for unit/integration tests only.
393    ///
394    /// # Errors
395    ///
396    /// Returns error if the server encounters a runtime error.
397    pub async fn serve_on_listener<F>(self, listener: TcpListener, shutdown: F) -> Result<()>
398    where
399        F: std::future::Future<Output = ()> + Send + 'static,
400    {
401        let (app, _app_state) = self.build_router();
402        axum::serve(listener, app.into_make_service_with_connect_info::<SocketAddr>())
403            .with_graceful_shutdown(shutdown)
404            .await
405            .map_err(|e| ServerError::IoError(std::io::Error::other(e)))?;
406        // Abort and await any lifecycle tasks spawned during construction
407        // (e.g. PKCE cleanup, trusted-docs reload) so the test path doesn't
408        // leak background work into the next test.
409        drain_lifecycle_tasks(self.tasks, self.config.shutdown_timeout_secs).await;
410        Ok(())
411    }
412
413    /// Listen for shutdown signals (Ctrl+C or SIGTERM)
414    pub async fn shutdown_signal() {
415        use tokio::signal;
416
417        let ctrl_c = async {
418            match signal::ctrl_c().await {
419                Ok(()) => {},
420                Err(e) => {
421                    warn!(error = %e, "Failed to install Ctrl+C handler");
422                    std::future::pending::<()>().await;
423                },
424            }
425        };
426
427        #[cfg(unix)]
428        let terminate = async {
429            match signal::unix::signal(signal::unix::SignalKind::terminate()) {
430                Ok(mut s) => {
431                    s.recv().await;
432                },
433                Err(e) => {
434                    warn!(error = %e, "Failed to install SIGTERM handler");
435                    std::future::pending::<()>().await;
436                },
437            }
438        };
439
440        #[cfg(not(unix))]
441        let terminate = std::future::pending::<()>();
442
443        tokio::select! {
444            () = ctrl_c => info!("Received Ctrl+C"),
445            () = terminate => info!("Received SIGTERM"),
446        }
447    }
448}
449
450/// Abort every lifecycle task on the supplied [`tokio::task::JoinSet`] and await
451/// the resulting `JoinError`s so the runtime is fully drained before
452/// `serve_with_shutdown` returns.
453///
454/// Tasks are awaited under an outer timeout so a stuck task cannot prevent
455/// process exit. A `JoinError::is_cancelled()` after `JoinSet::abort_all` is
456/// the expected case — only unexpected panics are logged.
457pub(super) async fn drain_lifecycle_tasks(
458    mut tasks: tokio::task::JoinSet<()>,
459    shutdown_timeout_secs: u64,
460) {
461    if tasks.is_empty() {
462        return;
463    }
464
465    tasks.abort_all();
466    let timeout = std::time::Duration::from_secs(shutdown_timeout_secs);
467    let drained = tokio::time::timeout(timeout, async {
468        while let Some(res) = tasks.join_next().await {
469            if let Err(e) = res {
470                if !e.is_cancelled() {
471                    warn!(error = %e, "Lifecycle task terminated with a non-cancellation error");
472                }
473            }
474        }
475    })
476    .await;
477    if drained.is_err() {
478        warn!(
479            timeout_secs = shutdown_timeout_secs,
480            "Lifecycle task drain timed out; some background tasks did not stop in time"
481        );
482    } else {
483        info!("All lifecycle background tasks drained");
484    }
485}