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                        error!("Failed to start observer runtime: {}", e);
212                        warn!("Server will continue without observers");
213                    },
214                }
215                drop(guard);
216            }
217            handle
218        };
219
220        // Explicitly enable TCP_NODELAY (disable Nagle's algorithm) on every
221        // accepted connection to minimise latency for small GraphQL responses.
222        let listener = TcpListener::bind(self.config.bind_addr)
223            .await
224            .map_err(|e| ServerError::BindError(e.to_string()))?
225            .tap_io(|tcp_stream| {
226                if let Err(err) = tcp_stream.set_nodelay(true) {
227                    warn!("failed to set TCP_NODELAY: {err:#}");
228                }
229            });
230
231        // Warn if the process file descriptor limit is below the recommended minimum.
232        // A low limit causes "too many open files" errors under load.
233        #[cfg(target_os = "linux")]
234        {
235            if let Ok(limits) = std::fs::read_to_string("/proc/self/limits") {
236                for line in limits.lines() {
237                    if line.starts_with("Max open files") {
238                        let parts: Vec<&str> = line.split_whitespace().collect();
239                        if let Some(soft) = parts.get(3) {
240                            if let Ok(n) = soft.parse::<u64>() {
241                                if n < 65_536 {
242                                    warn!(
243                                        current_fd_limit = n,
244                                        recommended = 65_536,
245                                        "File descriptor limit is low; consider raising ulimit -n"
246                                    );
247                                }
248                            }
249                        }
250                        break;
251                    }
252                }
253            }
254        }
255
256        // Log TLS configuration
257        if tls_setup.is_tls_enabled() {
258            // Verify TLS setup is valid (will error if certificates are missing/invalid)
259            let _ = tls_setup.create_rustls_config()?;
260            info!(
261                cert_path = ?tls_setup.cert_path(),
262                key_path = ?tls_setup.key_path(),
263                mtls_required = tls_setup.is_mtls_required(),
264                "Server TLS configuration loaded (note: use reverse proxy for server-side TLS termination)"
265            );
266        }
267
268        // Log database TLS configuration
269        info!(
270            postgres_ssl_mode = tls_setup.postgres_ssl_mode(),
271            redis_ssl = tls_setup.redis_ssl_enabled(),
272            clickhouse_https = tls_setup.clickhouse_https_enabled(),
273            elasticsearch_https = tls_setup.elasticsearch_https_enabled(),
274            "Database connection TLS configuration applied"
275        );
276
277        info!("Server listening on http://{}", self.config.bind_addr);
278
279        // Start both HTTP and gRPC servers concurrently if Arrow Flight is enabled
280        #[cfg(feature = "arrow")]
281        if let Some(flight_service) = self.flight_service.take() {
282            let flight_addr = self.config.flight_bind_addr;
283            info!("Arrow Flight server listening on grpc://{}", flight_addr);
284
285            // Spawn Flight server in background, registered on the server's
286            // JoinSet. The set's `shutdown` step abort-then-awaits the gRPC
287            // server when the HTTP server exits.
288            self.tasks.spawn(async move {
289                if let Err(e) = tonic::transport::Server::builder()
290                    .add_service(flight_service.into_server())
291                    .serve(flight_addr)
292                    .await
293                {
294                    error!(error = %e, "Arrow Flight server terminated with error");
295                }
296            });
297
298            // Wrap the user-supplied shutdown future so we can also stop observer runtime
299            #[cfg(feature = "observers")]
300            let observer_runtime = self.observer_runtime.clone();
301
302            let shutdown_with_cleanup = async move {
303                shutdown.await;
304                #[cfg(feature = "observers")]
305                if let Some(ref runtime) = observer_runtime {
306                    info!("Shutting down observer runtime");
307                    let mut guard = runtime.write().await;
308                    if let Err(e) = guard.stop().await {
309                        #[cfg(feature = "observers")]
310                        error!("Error stopping runtime: {}", e);
311                    } else {
312                        info!("Runtime stopped cleanly");
313                    }
314                }
315            };
316
317            // Run HTTP server with graceful shutdown
318            axum::serve(listener, app.into_make_service_with_connect_info::<SocketAddr>())
319                .with_graceful_shutdown(shutdown_with_cleanup)
320                .await
321                .map_err(|e| ServerError::IoError(std::io::Error::other(e)))?;
322
323            // Abort and await every lifecycle task (Flight server, SIGUSR1
324            // handler, PKCE cleanup, trusted-docs reload, usage flush, …).
325            drain_lifecycle_tasks(self.tasks, self.config.shutdown_timeout_secs).await;
326        }
327
328        // HTTP-only server (when arrow feature not enabled)
329        #[cfg(not(feature = "arrow"))]
330        {
331            axum::serve(listener, app.into_make_service_with_connect_info::<SocketAddr>())
332                .with_graceful_shutdown(shutdown)
333                .await
334                .map_err(|e| ServerError::IoError(std::io::Error::other(e)))?;
335
336            let shutdown_timeout =
337                std::time::Duration::from_secs(self.config.shutdown_timeout_secs);
338            info!(
339                timeout_secs = self.config.shutdown_timeout_secs,
340                "HTTP server stopped, draining remaining work"
341            );
342
343            let drain = tokio::time::timeout(shutdown_timeout, async {
344                #[cfg(feature = "observers")]
345                if let Some(ref runtime) = self.observer_runtime {
346                    let mut guard = runtime.write().await;
347                    match guard.stop().await {
348                        Ok(()) => info!("Observer runtime stopped cleanly"),
349                        Err(e) => warn!("Observer runtime shutdown error: {e}"),
350                    }
351                }
352            })
353            .await;
354
355            if drain.is_err() {
356                warn!(
357                    timeout_secs = self.config.shutdown_timeout_secs,
358                    "Shutdown drain timed out; forcing exit"
359                );
360            } else {
361                info!("Graceful shutdown complete");
362            }
363
364            // Abort and await every lifecycle task (SIGUSR1 handler, PKCE
365            // cleanup, trusted-docs reload, usage flush, …).
366            drain_lifecycle_tasks(self.tasks, self.config.shutdown_timeout_secs).await;
367        }
368
369        Ok(())
370    }
371
372    /// Start server on an externally created listener.
373    ///
374    /// Used in tests to discover the bound port before serving.
375    /// Skips TLS, Flight, and observer startup — suitable for unit/integration tests only.
376    ///
377    /// # Errors
378    ///
379    /// Returns error if the server encounters a runtime error.
380    pub async fn serve_on_listener<F>(self, listener: TcpListener, shutdown: F) -> Result<()>
381    where
382        F: std::future::Future<Output = ()> + Send + 'static,
383    {
384        let (app, _app_state) = self.build_router();
385        axum::serve(listener, app.into_make_service_with_connect_info::<SocketAddr>())
386            .with_graceful_shutdown(shutdown)
387            .await
388            .map_err(|e| ServerError::IoError(std::io::Error::other(e)))?;
389        // Abort and await any lifecycle tasks spawned during construction
390        // (e.g. PKCE cleanup, trusted-docs reload) so the test path doesn't
391        // leak background work into the next test.
392        drain_lifecycle_tasks(self.tasks, self.config.shutdown_timeout_secs).await;
393        Ok(())
394    }
395
396    /// Listen for shutdown signals (Ctrl+C or SIGTERM)
397    pub async fn shutdown_signal() {
398        use tokio::signal;
399
400        let ctrl_c = async {
401            match signal::ctrl_c().await {
402                Ok(()) => {},
403                Err(e) => {
404                    warn!(error = %e, "Failed to install Ctrl+C handler");
405                    std::future::pending::<()>().await;
406                },
407            }
408        };
409
410        #[cfg(unix)]
411        let terminate = async {
412            match signal::unix::signal(signal::unix::SignalKind::terminate()) {
413                Ok(mut s) => {
414                    s.recv().await;
415                },
416                Err(e) => {
417                    warn!(error = %e, "Failed to install SIGTERM handler");
418                    std::future::pending::<()>().await;
419                },
420            }
421        };
422
423        #[cfg(not(unix))]
424        let terminate = std::future::pending::<()>();
425
426        tokio::select! {
427            () = ctrl_c => info!("Received Ctrl+C"),
428            () = terminate => info!("Received SIGTERM"),
429        }
430    }
431}
432
433/// Abort every lifecycle task on the supplied [`tokio::task::JoinSet`] and await
434/// the resulting `JoinError`s so the runtime is fully drained before
435/// `serve_with_shutdown` returns.
436///
437/// Tasks are awaited under an outer timeout so a stuck task cannot prevent
438/// process exit. A `JoinError::is_cancelled()` after `JoinSet::abort_all` is
439/// the expected case — only unexpected panics are logged.
440pub(super) async fn drain_lifecycle_tasks(
441    mut tasks: tokio::task::JoinSet<()>,
442    shutdown_timeout_secs: u64,
443) {
444    if tasks.is_empty() {
445        return;
446    }
447
448    tasks.abort_all();
449    let timeout = std::time::Duration::from_secs(shutdown_timeout_secs);
450    let drained = tokio::time::timeout(timeout, async {
451        while let Some(res) = tasks.join_next().await {
452            if let Err(e) = res {
453                if !e.is_cancelled() {
454                    warn!(error = %e, "Lifecycle task terminated with a non-cancellation error");
455                }
456            }
457        }
458    })
459    .await;
460    if drained.is_err() {
461        warn!(
462            timeout_secs = shutdown_timeout_secs,
463            "Lifecycle task drain timed out; some background tasks did not stop in time"
464        );
465    } else {
466        info!("All lifecycle background tasks drained");
467    }
468}