Skip to main content

fraiseql_server/middleware/
auth.rs

1//! Authentication middleware.
2//!
3//! Provides bearer token authentication for protected endpoints.
4
5use std::{
6    sync::Arc,
7    time::{SystemTime, UNIX_EPOCH},
8};
9
10use axum::{
11    body::Body,
12    extract::State,
13    http::{Request, StatusCode, header},
14    middleware::Next,
15    response::{IntoResponse, Response},
16};
17use dashmap::DashMap;
18use subtle::ConstantTimeEq as _;
19
20/// Window length (in seconds) for the admin brute-force rate limiter.
21const ADMIN_AUTH_WINDOW_SECS: u64 = 60;
22
23/// Per-IP failure record for the admin brute-force guard.
24#[derive(Clone)]
25struct FailureRecord {
26    count:        u32,
27    window_start: u64,
28}
29
30/// Per-IP sliding-window counter for failed bearer token attempts.
31///
32/// Shared inside `BearerAuthState` via an `Arc`-wrapped `DashMap` so that
33/// the state can be `Clone`d cheaply across requests.
34#[derive(Clone)]
35pub(crate) struct FailureLimiter {
36    records:      Arc<DashMap<String, FailureRecord>>,
37    max_failures: u32,
38}
39
40impl FailureLimiter {
41    pub(crate) fn new(max_failures: u32) -> Self {
42        Self {
43            records: Arc::new(DashMap::new()),
44            max_failures,
45        }
46    }
47
48    fn now_secs() -> u64 {
49        SystemTime::now().duration_since(UNIX_EPOCH).map(|d| d.as_secs()).unwrap_or(0)
50    }
51
52    /// Record a failed attempt and return `true` if the IP is now rate-limited.
53    pub(crate) fn record_failure(&self, ip: &str) -> bool {
54        let now = Self::now_secs();
55        let mut entry = self.records.entry(ip.to_string()).or_insert_with(|| FailureRecord {
56            count:        0,
57            window_start: now,
58        });
59
60        if now >= entry.window_start + ADMIN_AUTH_WINDOW_SECS {
61            // Window expired — start fresh
62            entry.count = 1;
63            entry.window_start = now;
64            false
65        } else {
66            entry.count = entry.count.saturating_add(1);
67            entry.count >= self.max_failures
68        }
69    }
70
71    /// Return `true` if the IP is already rate-limited (without recording a new failure).
72    pub(crate) fn is_blocked(&self, ip: &str) -> bool {
73        let now = Self::now_secs();
74        if let Some(entry) = self.records.get(ip) {
75            if now < entry.window_start + ADMIN_AUTH_WINDOW_SECS {
76                return entry.count >= self.max_failures;
77            }
78        }
79        false
80    }
81
82    /// Reset the failure counter for an IP after a successful authentication.
83    pub(crate) fn record_success(&self, ip: &str) {
84        self.records.remove(ip);
85    }
86
87    /// Return the current failure count for an IP (used in tests).
88    #[cfg(test)]
89    pub(crate) fn failure_count(&self, ip: &str) -> u32 {
90        self.records.get(ip).map_or(0, |e| e.count)
91    }
92}
93
94/// Shared state for bearer token authentication.
95#[derive(Clone)]
96pub struct BearerAuthState {
97    /// Expected bearer token.
98    pub token:       Arc<String>,
99    /// Per-IP brute-force guard.
100    failure_limiter: FailureLimiter,
101}
102
103impl BearerAuthState {
104    /// Create new bearer auth state with the default max-failures limit (10).
105    #[must_use]
106    pub fn new(token: String) -> Self {
107        Self::with_max_failures(token, 10)
108    }
109
110    /// Create new bearer auth state with a custom max-failures limit.
111    ///
112    /// After `max_failures` failed attempts from the same IP within a 60-second
113    /// window, further requests receive **429 Too Many Requests**.
114    #[must_use]
115    pub fn with_max_failures(token: String, max_failures: u32) -> Self {
116        Self {
117            token:           Arc::new(token),
118            failure_limiter: FailureLimiter::new(max_failures),
119        }
120    }
121}
122
123/// Bearer token authentication middleware.
124///
125/// Validates that requests include a valid `Authorization: Bearer <token>` header.
126///
127/// # Response
128///
129/// - **401 Unauthorized**: Missing or malformed Authorization header
130/// - **403 Forbidden**: Invalid token
131///
132/// # Example
133///
134/// ```text
135/// // Requires: running Axum application with a route handler.
136/// use axum::{Router, middleware};
137/// use fraiseql_server::middleware::{bearer_auth_middleware, BearerAuthState};
138///
139/// let auth_state = BearerAuthState::new("my-secret-token".to_string());
140///
141/// let app = Router::new()
142///     .route("/protected", get(handler))
143///     .layer(middleware::from_fn_with_state(auth_state, bearer_auth_middleware));
144/// ```
145pub async fn bearer_auth_middleware(
146    State(auth_state): State<BearerAuthState>,
147    request: Request<Body>,
148    next: Next,
149) -> Response {
150    // Derive a best-effort peer key for rate limiting.
151    // ConnectInfo is only available when the server was started with
152    // `into_make_service_with_connect_info`; fall back to a header-based key.
153    use std::net::SocketAddr;
154
155    use axum::extract::ConnectInfo;
156    let peer_key = request
157        .extensions()
158        .get::<ConnectInfo<SocketAddr>>()
159        .map(|ci| ci.0.ip().to_string())
160        .or_else(|| {
161            request
162                .headers()
163                .get("x-forwarded-for")
164                .and_then(|v| v.to_str().ok())
165                .map(|v| v.split(',').next().unwrap_or(v).trim().to_string())
166        })
167        .unwrap_or_else(|| "unknown".to_string());
168
169    // Reject immediately if already rate-limited (avoids any header work).
170    if auth_state.failure_limiter.is_blocked(&peer_key) {
171        return (StatusCode::TOO_MANY_REQUESTS, "Too many failed auth attempts").into_response();
172    }
173
174    // Extract Authorization header
175    let auth_header = request
176        .headers()
177        .get(header::AUTHORIZATION)
178        .and_then(|value| value.to_str().ok());
179
180    match auth_header {
181        None => {
182            return (
183                StatusCode::UNAUTHORIZED,
184                [(header::WWW_AUTHENTICATE, "Bearer")],
185                "Missing Authorization header",
186            )
187                .into_response();
188        },
189        Some(header_value) => {
190            // Check for "Bearer " prefix
191            if !header_value.starts_with("Bearer ") {
192                return (
193                    StatusCode::UNAUTHORIZED,
194                    [(header::WWW_AUTHENTICATE, "Bearer")],
195                    "Invalid Authorization header format. Expected: Bearer <token>",
196                )
197                    .into_response();
198            }
199
200            // Extract token
201            let token = &header_value[7..]; // Skip "Bearer "
202
203            // Constant-time comparison to prevent timing attacks
204            if !constant_time_compare(token, &auth_state.token) {
205                // Record failure; return 429 once the limit is crossed.
206                if auth_state.failure_limiter.record_failure(&peer_key) {
207                    return (StatusCode::TOO_MANY_REQUESTS, "Too many failed auth attempts")
208                        .into_response();
209                }
210                return (StatusCode::FORBIDDEN, "Invalid token").into_response();
211            }
212
213            // Successful auth — reset the failure counter.
214            auth_state.failure_limiter.record_success(&peer_key);
215        },
216    }
217
218    // Token is valid, proceed with request
219    next.run(request).await
220}
221
222/// Extract the bearer token from an `Authorization` header value.
223///
224/// Returns `Some(token)` if the header has the `Bearer ` prefix (with trailing space),
225/// `None` for all other formats (Basic, Digest, missing prefix, etc.).
226///
227/// Exposed as `pub` for property testing.
228#[must_use]
229pub fn extract_bearer_token(header_value: &str) -> Option<&str> {
230    header_value.strip_prefix("Bearer ")
231}
232
233/// Constant-time string comparison to prevent timing attacks.
234///
235/// Uses [`subtle::ConstantTimeEq`] to compare the byte representations of
236/// both strings, preventing the compiler from optimising the comparison into
237/// an early-exit branch that would leak information about where the strings
238/// differ (timing oracle, RFC 6749 §10.12).
239///
240/// Strings of different lengths return `false` without inspecting bytes;
241/// token lengths are considered non-secret (administrators choose them).
242pub(crate) fn constant_time_compare(a: &str, b: &str) -> bool {
243    a.as_bytes().ct_eq(b.as_bytes()).into()
244}