fraiseql-server 2.15.0

HTTP server for FraiseQL v2 GraphQL engine
//! Global middleware layers: metrics, tracing, CORS, body limits, header limits,
//! timeout, and rate limiting.

use std::sync::Arc;

use axum::{Router, extract::DefaultBodyLimit, middleware};
use tracing::info;

use super::super::{Server, metrics_middleware, trace_layer};
use crate::{
    middleware::{
        cors::cors_layer_restricted_with, rate_limit::VerifiedSubject, security_headers_middleware,
    },
    routes::graphql::AppState,
};

impl Server {
    /// Apply global middleware layers to the router.
    pub(super) fn apply_middleware(&self, mut app: Router, state: &AppState) -> Router {
        let metrics = state.metrics.clone();

        // Add HTTP metrics middleware (tracks requests and response status codes)
        // This runs on ALL routes, even when metrics endpoints are disabled
        app = app.layer(middleware::from_fn_with_state(metrics, metrics_middleware));

        // Add security response headers (nosniff/XFO/HSTS/Referrer/CSP/XSS) to every
        // response (M-sec-headers). Set-if-absent, so a handler that needs a different
        // policy (e.g. the playground's relaxed CSP) is preserved.
        app = app.layer(middleware::from_fn(security_headers_middleware));

        // Add middleware
        if self.config.tracing_enabled {
            app = app.layer(trace_layer());
        }

        if self.config.cors_enabled {
            let origins = if self.config.cors_origins.is_empty() {
                tracing::warn!(
                    "CORS enabled but no origins configured. Using localhost:3000 as default. \
                     Set cors_origins in config for production."
                );
                vec!["http://localhost:3000".to_string()]
            } else {
                self.config.cors_origins.clone()
            };
            app = app.layer(cors_layer_restricted_with(&origins, self.config.enable_http_query));
        }

        // Add request body size limit (default 1 MB -- prevents memory exhaustion)
        if self.config.max_request_body_bytes > 0 {
            info!(
                max_bytes = self.config.max_request_body_bytes,
                "Request body size limit enabled"
            );
            app = app.layer(DefaultBodyLimit::max(self.config.max_request_body_bytes));
        }

        // Add HTTP header count and size limits (prevents header-flooding DoS)
        {
            let max_header_count = self.config.max_header_count;
            let max_header_bytes = self.config.max_header_bytes;
            info!(max_header_count, max_header_bytes, "HTTP header limits enabled");
            app = app.layer(axum::middleware::from_fn(move |req, next| {
                crate::middleware::header_limits_middleware(
                    req,
                    next,
                    max_header_count,
                    max_header_bytes,
                )
            }));
        }

        // Add per-request timeout (optional -- defence against runaway DB queries).
        if let Some(timeout_secs) = self.config.request_timeout_secs {
            use std::time::Duration;

            use tower_http::timeout::TimeoutLayer;

            info!(timeout_secs, "Request timeout enabled");
            app = app.layer(TimeoutLayer::with_status_code(
                axum::http::StatusCode::REQUEST_TIMEOUT,
                Duration::from_secs(timeout_secs),
            ));
        }

        // Add rate limiting middleware if configured.
        if let Some(ref limiter) = self.rate_limiter {
            use axum::Extension;

            use crate::middleware::rate_limit::rate_limit_middleware;

            info!("Enabling rate limiting middleware");
            app = app
                .layer(middleware::from_fn(rate_limit_middleware))
                .layer(Extension(limiter.clone()));

            // #1171: the per-user bucket needs a subject this deployment's own validator
            // accepts. Layered outside the limiter so it is in the extensions by the time
            // the limiter reads them; absent when no authentication is configured, in
            // which case every request buckets on its address as before.
            if let Some(subject) = self.rate_limit_subject() {
                info!("Per-user rate limiting enabled on a signature-verified subject");
                app = app.layer(Extension(subject));
            }
        }

        app
    }

    /// The validator the rate limiter verifies a subject with, or `None`.
    ///
    /// Mirrors [`Server::attach_auth`]'s precedence — OIDC first, then HS256 — so the
    /// limiter and the transports never disagree about which credential is the real one.
    /// A deployment with no authentication has no verified subject to key on, and its
    /// requests keep bucketing on the client address.
    fn rate_limit_subject(&self) -> Option<Arc<VerifiedSubject>> {
        if let Some(ref validator) = self.oidc_validator {
            return Some(Arc::new(VerifiedSubject::Oidc(Arc::clone(validator))));
        }
        if let Some(ref validator) = self.hs256_auth {
            return Some(Arc::new(VerifiedSubject::Hs256(Arc::clone(validator))));
        }
        None
    }
}