fraiseql-server 2.16.0

HTTP server for FraiseQL v2 GraphQL engine
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
//! Authentication middleware.
//!
//! Provides bearer token authentication for protected endpoints.

use std::{
    sync::Arc,
    time::{SystemTime, UNIX_EPOCH},
};

use axum::{
    body::Body,
    extract::State,
    http::{Request, StatusCode, header},
    middleware::Next,
    response::{IntoResponse, Response},
};
use dashmap::DashMap;
use subtle::ConstantTimeEq as _;

/// Window length (in seconds) for the admin brute-force rate limiter.
const ADMIN_AUTH_WINDOW_SECS: u64 = 60;

/// Per-IP failure record for the admin brute-force guard.
#[derive(Clone)]
struct FailureRecord {
    count:        u32,
    window_start: u64,
}

/// Per-IP sliding-window counter for failed bearer token attempts.
///
/// Shared inside `BearerAuthState` via an `Arc`-wrapped `DashMap` so that
/// the state can be `Clone`d cheaply across requests.
#[derive(Clone)]
pub(crate) struct FailureLimiter {
    records:      Arc<DashMap<String, FailureRecord>>,
    max_failures: u32,
}

impl FailureLimiter {
    /// The window length, exposed so tests age records by a real multiple of it.
    #[cfg(test)]
    pub(super) const WINDOW_SECS: u64 = ADMIN_AUTH_WINDOW_SECS;

    pub(crate) fn new(max_failures: u32) -> Self {
        Self {
            records: Arc::new(DashMap::new()),
            max_failures,
        }
    }

    fn now_secs() -> u64 {
        SystemTime::now().duration_since(UNIX_EPOCH).map_or(0, |d| d.as_secs())
    }

    /// Drop records whose window expired long enough ago to be irrelevant.
    ///
    /// The map is keyed by client IP and only ever had entries removed on a
    /// *successful* auth, so a stream of failed attempts from changing source
    /// addresses grew it without bound — a slow memory leak reachable by any
    /// unauthenticated caller (#731). Sweeping is amortised: it runs on insert,
    /// and only once the map is large enough for the scan to be worth it.
    pub(super) fn evict_expired(&self, now: u64) {
        /// Only sweep once the map is big enough that a scan is worth its cost.
        const EVICTION_THRESHOLD: usize = 1024;
        /// Keep an expired window around for one further window, so a burst that
        /// straddles the boundary is still counted against the same IP.
        const RETENTION_WINDOWS: u64 = 2;

        if self.records.len() < EVICTION_THRESHOLD {
            return;
        }
        let cutoff = ADMIN_AUTH_WINDOW_SECS * RETENTION_WINDOWS;
        self.records
            .retain(|_, record| now.saturating_sub(record.window_start) < cutoff);
    }

    /// Record a failed attempt and return `true` if the IP is now rate-limited.
    pub(crate) fn record_failure(&self, ip: &str) -> bool {
        let now = Self::now_secs();
        // Before inserting a potentially new key, drop the dead ones.
        self.evict_expired(now);
        let mut entry = self.records.entry(ip.to_string()).or_insert_with(|| FailureRecord {
            count:        0,
            window_start: now,
        });

        if now >= entry.window_start + ADMIN_AUTH_WINDOW_SECS {
            // Window expired — start fresh
            entry.count = 1;
            entry.window_start = now;
            false
        } else {
            entry.count = entry.count.saturating_add(1);
            entry.count >= self.max_failures
        }
    }

    /// Return `true` if the IP is already rate-limited (without recording a new failure).
    pub(crate) fn is_blocked(&self, ip: &str) -> bool {
        let now = Self::now_secs();
        if let Some(entry) = self.records.get(ip) {
            if now < entry.window_start + ADMIN_AUTH_WINDOW_SECS {
                return entry.count >= self.max_failures;
            }
        }
        false
    }

    /// Reset the failure counter for an IP after a successful authentication.
    pub(crate) fn record_success(&self, ip: &str) {
        self.records.remove(ip);
    }

    /// Return the current failure count for an IP (used in tests).
    #[cfg(test)]
    pub(crate) fn failure_count(&self, ip: &str) -> u32 {
        self.records.get(ip).map_or(0, |e| e.count)
    }

    /// Number of per-IP records held (used in tests to pin eviction).
    #[cfg(test)]
    pub(super) fn record_count(&self) -> usize {
        self.records.len()
    }

    /// Record a failure with an explicit clock reading (used in tests to age
    /// records without sleeping).
    #[cfg(test)]
    pub(super) fn record_failure_at(&self, ip: &str, now: u64) {
        self.records.insert(
            ip.to_string(),
            FailureRecord {
                count:        1,
                window_start: now,
            },
        );
    }
}

/// Shared state for bearer token authentication.
#[derive(Clone)]
pub struct BearerAuthState {
    /// Expected bearer token.
    pub token:       Arc<String>,
    /// Per-IP brute-force guard.
    failure_limiter: FailureLimiter,
}

impl BearerAuthState {
    /// Create new bearer auth state with the default max-failures limit (10).
    #[must_use]
    pub fn new(token: String) -> Self {
        Self::with_max_failures(token, 10)
    }

    /// Create new bearer auth state with a custom max-failures limit.
    ///
    /// After `max_failures` failed attempts from the same IP within a 60-second
    /// window, further requests receive **429 Too Many Requests**.
    #[must_use]
    pub fn with_max_failures(token: String, max_failures: u32) -> Self {
        Self {
            token:           Arc::new(token),
            failure_limiter: FailureLimiter::new(max_failures),
        }
    }
}

/// Bearer token authentication middleware.
///
/// Validates that requests include a valid `Authorization: Bearer <token>` header.
///
/// # Response
///
/// - **401 Unauthorized**: Missing or malformed Authorization header
/// - **403 Forbidden**: Invalid token
///
/// # Example
///
/// ```text
/// // Requires: running Axum application with a route handler.
/// use axum::{Router, middleware};
/// use fraiseql_server::middleware::{bearer_auth_middleware, BearerAuthState};
///
/// let auth_state = BearerAuthState::new("my-secret-token".to_string());
///
/// let app = Router::new()
///     .route("/protected", get(handler))
///     .layer(middleware::from_fn_with_state(auth_state, bearer_auth_middleware));
/// ```
pub async fn bearer_auth_middleware(
    State(auth_state): State<BearerAuthState>,
    request: Request<Body>,
    next: Next,
) -> Response {
    // Derive the peer key for the brute-force limiter from the validated transport peer
    // only. ConnectInfo is the real socket address (present in the shipped binary, which
    // starts with `into_make_service_with_connect_info`). We deliberately do NOT fall back
    // to `X-Forwarded-For`: that header is attacker-controlled, so keying on it would let a
    // caller rotate it to mint a fresh failure budget per value, defeating the limiter
    // (M-xff-limiter). When ConnectInfo is absent (some library embeddings), all callers
    // share the single `unknown` bucket — fail-closed (more restrictive), not bypassable.
    use std::net::SocketAddr;

    use axum::extract::ConnectInfo;
    let peer_key = request
        .extensions()
        .get::<ConnectInfo<SocketAddr>>()
        .map_or_else(|| "unknown".to_string(), |ci| ci.0.ip().to_string());

    // Reject immediately if already rate-limited (avoids any header work).
    if auth_state.failure_limiter.is_blocked(&peer_key) {
        return (StatusCode::TOO_MANY_REQUESTS, "Too many failed auth attempts").into_response();
    }

    // Extract Authorization header
    let auth_header = request
        .headers()
        .get(header::AUTHORIZATION)
        .and_then(|value| value.to_str().ok());

    match auth_header {
        None => {
            return (
                StatusCode::UNAUTHORIZED,
                [(header::WWW_AUTHENTICATE, "Bearer")],
                "Missing Authorization header",
            )
                .into_response();
        },
        Some(header_value) => {
            // Check for "Bearer " prefix
            if !header_value.starts_with("Bearer ") {
                return (
                    StatusCode::UNAUTHORIZED,
                    [(header::WWW_AUTHENTICATE, "Bearer")],
                    "Invalid Authorization header format. Expected: Bearer <token>",
                )
                    .into_response();
            }

            // Extract token
            let token = &header_value[7..]; // Skip "Bearer "

            // Constant-time comparison to prevent timing attacks
            if !constant_time_compare(token, &auth_state.token) {
                // Record failure; return 429 once the limit is crossed.
                if auth_state.failure_limiter.record_failure(&peer_key) {
                    return (StatusCode::TOO_MANY_REQUESTS, "Too many failed auth attempts")
                        .into_response();
                }
                return (StatusCode::FORBIDDEN, "Invalid token").into_response();
            }

            // Successful auth — reset the failure counter.
            auth_state.failure_limiter.record_success(&peer_key);
        },
    }

    // Token is valid, proceed with request
    next.run(request).await
}

/// State for [`admin_principal_middleware`] (#1089).
#[derive(Clone)]
pub struct AdminPrincipalState {
    /// The deployment `admin_token`, which authenticates as the platform.
    platform_token:  Arc<String>,
    /// Tenant admin credentials. `None` when the deployment has no database pool, in which
    /// case only the platform token authenticates.
    tokens:          Option<Arc<crate::api::admin_principal::PgAdminTokenStore>>,
    /// Per-IP brute-force guard, as for the bearer gate.
    failure_limiter: FailureLimiter,
}

impl AdminPrincipalState {
    /// Accept the platform token and, when `tokens` is set, tenant admin tokens.
    #[must_use]
    pub fn new(
        platform_token: String,
        tokens: Option<Arc<crate::api::admin_principal::PgAdminTokenStore>>,
        max_failures: u32,
    ) -> Self {
        Self {
            platform_token: Arc::new(platform_token),
            tokens,
            failure_limiter: FailureLimiter::new(max_failures),
        }
    }
}

/// Admin authentication for routers that accept tenant administrators (#1089).
///
/// Accepts the deployment `admin_token`, which authenticates as
/// [`AdminPrincipal::Platform`](crate::api::admin_principal::AdminPrincipal::Platform), or a
/// tenant admin token, which authenticates as
/// [`AdminPrincipal::Tenant`](crate::api::admin_principal::AdminPrincipal::Tenant). Inserts the
/// principal as a request extension. Status codes and the brute-force limiter match
/// [`bearer_auth_middleware`]: 401 missing or malformed, 403 unknown, 429 throttled, and 503
/// if the credential store cannot be read, so a store outage fails closed.
pub async fn admin_principal_middleware(
    State(state): State<AdminPrincipalState>,
    mut request: Request<Body>,
    next: Next,
) -> Response {
    use std::net::SocketAddr;

    use axum::extract::ConnectInfo;

    use crate::api::admin_principal::AdminPrincipal;

    // Same peer keying as the bearer gate: the transport peer only, never X-Forwarded-For.
    let peer_key = request
        .extensions()
        .get::<ConnectInfo<SocketAddr>>()
        .map_or_else(|| "unknown".to_string(), |ci| ci.0.ip().to_string());
    if state.failure_limiter.is_blocked(&peer_key) {
        return (StatusCode::TOO_MANY_REQUESTS, "Too many failed auth attempts").into_response();
    }

    let Some(header_value) =
        request.headers().get(header::AUTHORIZATION).and_then(|v| v.to_str().ok())
    else {
        return (
            StatusCode::UNAUTHORIZED,
            [(header::WWW_AUTHENTICATE, "Bearer")],
            "Missing Authorization header",
        )
            .into_response();
    };
    let Some(token) = extract_bearer_token(header_value) else {
        return (
            StatusCode::UNAUTHORIZED,
            [(header::WWW_AUTHENTICATE, "Bearer")],
            "Invalid Authorization header format. Expected: Bearer <token>",
        )
            .into_response();
    };

    let principal = if constant_time_compare(token, &state.platform_token) {
        Some(AdminPrincipal::Platform)
    } else if let Some(tokens) = state.tokens.as_ref() {
        match tokens.authenticate(token).await {
            Ok(tenant) => tenant.map(AdminPrincipal::Tenant),
            Err(e) => {
                tracing::error!(error = %e, "tenant admin credential lookup failed");
                return (StatusCode::SERVICE_UNAVAILABLE, "Admin credential store unavailable")
                    .into_response();
            },
        }
    } else {
        None
    };

    let Some(principal) = principal else {
        if state.failure_limiter.record_failure(&peer_key) {
            return (StatusCode::TOO_MANY_REQUESTS, "Too many failed auth attempts")
                .into_response();
        }
        return (StatusCode::FORBIDDEN, "Invalid token").into_response();
    };
    state.failure_limiter.record_success(&peer_key);
    request.extensions_mut().insert(principal);
    next.run(request).await
}

/// Extract the bearer token from an `Authorization` header value.
///
/// Returns `Some(token)` if the header has the `Bearer ` prefix (with trailing space),
/// `None` for all other formats (Basic, Digest, missing prefix, etc.).
///
/// Exposed as `pub` for property testing.
#[must_use]
pub fn extract_bearer_token(header_value: &str) -> Option<&str> {
    header_value.strip_prefix("Bearer ")
}

/// Constant-time string comparison to prevent timing attacks.
///
/// Uses [`subtle::ConstantTimeEq`] to compare the byte representations of
/// both strings, preventing the compiler from optimising the comparison into
/// an early-exit branch that would leak information about where the strings
/// differ (timing oracle, RFC 6749 §10.12).
///
/// Strings of different lengths return `false` without inspecting bytes;
/// token lengths are considered non-secret (administrators choose them).
pub(crate) fn constant_time_compare(a: &str, b: &str) -> bool {
    a.as_bytes().ct_eq(b.as_bytes()).into()
}

/// Which admin credential authenticated a request to the SQL console (#962).
///
/// Every other admin route is on one of two routers, each authenticated by one
/// token, so "which token" is answered by which router matched. The console is a
/// single route that both tokens may reach and that must behave *differently* for
/// each, so the answer has to travel from the middleware to the handler. It does
/// so as a request extension the middleware inserts and nothing else constructs
/// — a client cannot supply it, because it is a Rust type and not a header.
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum AdminPrivilege {
    /// Authenticated by `admin_readonly_token`. The transaction runs `READ ONLY`.
    ReadOnly,
    /// Authenticated by `admin_token`. Writes are permitted; committing is still
    /// opt-in per request.
    ReadWrite,
}

/// What [`admin_dual_auth_middleware`] establishes about a request.
///
/// Inserted as a request extension on success and constructed nowhere else, so a
/// handler reading it is reading the middleware's conclusion rather than
/// re-deriving one. The peer address travels with the privilege because the
/// middleware has already resolved it (from `ConnectInfo`, never from
/// `X-Forwarded-For`), and both belong in the same audit record.
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct AdminCaller {
    /// Which admin credential authenticated.
    pub privilege: AdminPrivilege,
    /// The transport peer's IP, or `"unknown"` when there is no `ConnectInfo`.
    pub peer_ip:   String,
}

/// Shared state for the dual-token admin authentication used by the SQL console.
///
/// Holds both tokens and **one** failure limiter, deliberately: the limiter
/// counts failed attempts per peer, and two limiters would let a caller spend a
/// fresh budget by guessing against each token in turn.
#[derive(Clone)]
pub struct AdminDualAuthState {
    write_token:     Arc<String>,
    readonly_token:  Option<Arc<String>>,
    failure_limiter: FailureLimiter,
}

impl AdminDualAuthState {
    /// Build the state from the two configured tokens.
    ///
    /// `readonly_token` is `None` in single-token mode, where `admin_token`
    /// grants everything — the console then has no read-only credential at all,
    /// which is the honest reading of "one token, all operations" and matches how
    /// the rest of the admin API already behaves.
    #[must_use]
    pub fn new(write_token: String, readonly_token: Option<String>, max_failures: u32) -> Self {
        Self {
            write_token:     Arc::new(write_token),
            readonly_token:  readonly_token.map(Arc::new),
            failure_limiter: FailureLimiter::new(max_failures),
        }
    }

    /// Classify a presented token.
    ///
    /// Both comparisons always run: returning early on the write-token match
    /// would make the response time depend on which token was presented, and the
    /// whole point of [`constant_time_compare`] is that it does not.
    fn classify(&self, presented: &str) -> Option<AdminPrivilege> {
        let is_write = constant_time_compare(presented, &self.write_token);
        let is_readonly = self
            .readonly_token
            .as_ref()
            .is_some_and(|t| constant_time_compare(presented, t));
        match (is_write, is_readonly) {
            (true, _) => Some(AdminPrivilege::ReadWrite),
            (false, true) => Some(AdminPrivilege::ReadOnly),
            (false, false) => None,
        }
    }
}

/// Bearer authentication that reports *which* admin token authenticated (#962).
///
/// Same refusals and the same per-peer brute-force guard as
/// [`bearer_auth_middleware`]; the difference is that a success inserts an
/// [`AdminPrivilege`] into the request extensions instead of discarding the
/// distinction. Used only by the SQL console, whose behaviour depends on it.
///
/// # Response
///
/// - **401 Unauthorized**: missing or malformed `Authorization` header
/// - **403 Forbidden**: the token matches neither admin credential
/// - **429 Too Many Requests**: too many failures from this peer
pub async fn admin_dual_auth_middleware(
    State(auth_state): State<AdminDualAuthState>,
    mut request: Request<Body>,
    next: Next,
) -> Response {
    use std::net::SocketAddr;

    use axum::extract::ConnectInfo;

    // Keyed on the transport peer only, never on `X-Forwarded-For` — see
    // `bearer_auth_middleware` for why (M-xff-limiter).
    let peer_key = request
        .extensions()
        .get::<ConnectInfo<SocketAddr>>()
        .map_or_else(|| "unknown".to_string(), |ci| ci.0.ip().to_string());

    if auth_state.failure_limiter.is_blocked(&peer_key) {
        return (StatusCode::TOO_MANY_REQUESTS, "Too many failed auth attempts").into_response();
    }

    let Some(header_value) = request
        .headers()
        .get(header::AUTHORIZATION)
        .and_then(|value| value.to_str().ok())
    else {
        return (
            StatusCode::UNAUTHORIZED,
            [(header::WWW_AUTHENTICATE, "Bearer")],
            "Missing Authorization header",
        )
            .into_response();
    };

    let Some(token) = extract_bearer_token(header_value) else {
        return (
            StatusCode::UNAUTHORIZED,
            [(header::WWW_AUTHENTICATE, "Bearer")],
            "Invalid Authorization header format. Expected: Bearer <token>",
        )
            .into_response();
    };

    let Some(privilege) = auth_state.classify(token) else {
        if auth_state.failure_limiter.record_failure(&peer_key) {
            return (StatusCode::TOO_MANY_REQUESTS, "Too many failed auth attempts")
                .into_response();
        }
        return (StatusCode::FORBIDDEN, "Invalid token").into_response();
    };

    auth_state.failure_limiter.record_success(&peer_key);
    request.extensions_mut().insert(AdminCaller {
        privilege,
        peer_ip: peer_key,
    });
    next.run(request).await
}

#[cfg(test)]
mod xff_tests {
    //! M-xff-limiter: the brute-force limiter must not key on the attacker-controlled
    //! `X-Forwarded-For` header — rotating it must not grant a fresh failure budget.
    #![allow(clippy::unwrap_used)] // Reason: test code, panics are acceptable

    use axum::{
        Router,
        body::Body,
        http::{Request, StatusCode},
        middleware,
        routing::get,
    };
    use tower::ServiceExt as _;

    use super::{BearerAuthState, bearer_auth_middleware};

    async fn protected() -> &'static str {
        "ok"
    }

    fn wrong_token_request(xff: &str) -> Request<Body> {
        Request::builder()
            .uri("/")
            .header("authorization", "Bearer wrong-token")
            .header("x-forwarded-for", xff)
            .body(Body::empty())
            .unwrap()
    }

    #[tokio::test]
    async fn rotating_x_forwarded_for_does_not_refresh_the_failure_budget() {
        let state = BearerAuthState::with_max_failures("correct-token".to_string(), 2);
        let app = Router::new()
            .route("/", get(protected))
            .layer(middleware::from_fn_with_state(state, bearer_auth_middleware));

        // A oneshot sets no ConnectInfo, so the peer key falls back to "unknown" for every
        // request. Each failed attempt carries a DIFFERENT X-Forwarded-For: if the limiter
        // keyed on it, each value would get its own budget and never block. With the XFF
        // fallback removed they share the single "unknown" bucket, so the limit is reached.
        let mut statuses = Vec::new();
        for i in 0..5 {
            let resp = app
                .clone()
                .oneshot(wrong_token_request(&format!("203.0.113.{i}")))
                .await
                .unwrap();
            statuses.push(resp.status());
        }

        assert!(
            statuses.contains(&StatusCode::TOO_MANY_REQUESTS),
            "rotating X-Forwarded-For must still hit the shared rate limit, got {statuses:?}"
        );
    }
}