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}