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}