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 the peer key for the brute-force limiter from the validated transport peer
151    // only. ConnectInfo is the real socket address (present in the shipped binary, which
152    // starts with `into_make_service_with_connect_info`). We deliberately do NOT fall back
153    // to `X-Forwarded-For`: that header is attacker-controlled, so keying on it would let a
154    // caller rotate it to mint a fresh failure budget per value, defeating the limiter
155    // (M-xff-limiter). When ConnectInfo is absent (some library embeddings), all callers
156    // share the single `unknown` bucket — fail-closed (more restrictive), not bypassable.
157    use std::net::SocketAddr;
158
159    use axum::extract::ConnectInfo;
160    let peer_key = request
161        .extensions()
162        .get::<ConnectInfo<SocketAddr>>()
163        .map_or_else(|| "unknown".to_string(), |ci| ci.0.ip().to_string());
164
165    // Reject immediately if already rate-limited (avoids any header work).
166    if auth_state.failure_limiter.is_blocked(&peer_key) {
167        return (StatusCode::TOO_MANY_REQUESTS, "Too many failed auth attempts").into_response();
168    }
169
170    // Extract Authorization header
171    let auth_header = request
172        .headers()
173        .get(header::AUTHORIZATION)
174        .and_then(|value| value.to_str().ok());
175
176    match auth_header {
177        None => {
178            return (
179                StatusCode::UNAUTHORIZED,
180                [(header::WWW_AUTHENTICATE, "Bearer")],
181                "Missing Authorization header",
182            )
183                .into_response();
184        },
185        Some(header_value) => {
186            // Check for "Bearer " prefix
187            if !header_value.starts_with("Bearer ") {
188                return (
189                    StatusCode::UNAUTHORIZED,
190                    [(header::WWW_AUTHENTICATE, "Bearer")],
191                    "Invalid Authorization header format. Expected: Bearer <token>",
192                )
193                    .into_response();
194            }
195
196            // Extract token
197            let token = &header_value[7..]; // Skip "Bearer "
198
199            // Constant-time comparison to prevent timing attacks
200            if !constant_time_compare(token, &auth_state.token) {
201                // Record failure; return 429 once the limit is crossed.
202                if auth_state.failure_limiter.record_failure(&peer_key) {
203                    return (StatusCode::TOO_MANY_REQUESTS, "Too many failed auth attempts")
204                        .into_response();
205                }
206                return (StatusCode::FORBIDDEN, "Invalid token").into_response();
207            }
208
209            // Successful auth — reset the failure counter.
210            auth_state.failure_limiter.record_success(&peer_key);
211        },
212    }
213
214    // Token is valid, proceed with request
215    next.run(request).await
216}
217
218/// Extract the bearer token from an `Authorization` header value.
219///
220/// Returns `Some(token)` if the header has the `Bearer ` prefix (with trailing space),
221/// `None` for all other formats (Basic, Digest, missing prefix, etc.).
222///
223/// Exposed as `pub` for property testing.
224#[must_use]
225pub fn extract_bearer_token(header_value: &str) -> Option<&str> {
226    header_value.strip_prefix("Bearer ")
227}
228
229/// Constant-time string comparison to prevent timing attacks.
230///
231/// Uses [`subtle::ConstantTimeEq`] to compare the byte representations of
232/// both strings, preventing the compiler from optimising the comparison into
233/// an early-exit branch that would leak information about where the strings
234/// differ (timing oracle, RFC 6749 §10.12).
235///
236/// Strings of different lengths return `false` without inspecting bytes;
237/// token lengths are considered non-secret (administrators choose them).
238pub(crate) fn constant_time_compare(a: &str, b: &str) -> bool {
239    a.as_bytes().ct_eq(b.as_bytes()).into()
240}
241
242#[cfg(test)]
243mod xff_tests {
244    //! M-xff-limiter: the brute-force limiter must not key on the attacker-controlled
245    //! `X-Forwarded-For` header — rotating it must not grant a fresh failure budget.
246    #![allow(clippy::unwrap_used)]
247
248    use axum::{
249        Router,
250        body::Body,
251        http::{Request, StatusCode},
252        middleware,
253        routing::get,
254    };
255    use tower::ServiceExt as _;
256
257    use super::{BearerAuthState, bearer_auth_middleware};
258
259    async fn protected() -> &'static str {
260        "ok"
261    }
262
263    fn wrong_token_request(xff: &str) -> Request<Body> {
264        Request::builder()
265            .uri("/")
266            .header("authorization", "Bearer wrong-token")
267            .header("x-forwarded-for", xff)
268            .body(Body::empty())
269            .unwrap()
270    }
271
272    #[tokio::test]
273    async fn rotating_x_forwarded_for_does_not_refresh_the_failure_budget() {
274        let state = BearerAuthState::with_max_failures("correct-token".to_string(), 2);
275        let app = Router::new()
276            .route("/", get(protected))
277            .layer(middleware::from_fn_with_state(state, bearer_auth_middleware));
278
279        // A oneshot sets no ConnectInfo, so the peer key falls back to "unknown" for every
280        // request. Each failed attempt carries a DIFFERENT X-Forwarded-For: if the limiter
281        // keyed on it, each value would get its own budget and never block. With the XFF
282        // fallback removed they share the single "unknown" bucket, so the limit is reached.
283        let mut statuses = Vec::new();
284        for i in 0..5 {
285            let resp = app
286                .clone()
287                .oneshot(wrong_token_request(&format!("203.0.113.{i}")))
288                .await
289                .unwrap();
290            statuses.push(resp.status());
291        }
292
293        assert!(
294            statuses.contains(&StatusCode::TOO_MANY_REQUESTS),
295            "rotating X-Forwarded-For must still hit the shared rate limit, got {statuses:?}"
296        );
297    }
298}