Skip to main content

koan_server/auth/
routes.rs

1//! Auth HTTP routes: login, refresh, logout.
2
3use std::collections::HashMap;
4use std::net::{IpAddr, SocketAddr};
5use std::sync::{Arc, Mutex};
6
7use axum::Json;
8use axum::extract::{ConnectInfo, State};
9use axum::http::StatusCode;
10use axum::http::header::{COOKIE, SET_COOKIE};
11use axum::response::{AppendHeaders, IntoResponse, Response};
12use axum::routing::post;
13use serde::{Deserialize, Serialize};
14
15use koan_core::auth;
16use koan_core::db::pool::{Handle, Pool};
17use koan_core::db::queries::auth as auth_queries;
18
19/// Name of the cookie carrying the refresh token. Scoped to `/auth` so it
20/// reaches refresh and logout but is never attached to an API call, and
21/// `HttpOnly` so script cannot read it.
22pub(crate) const REFRESH_COOKIE: &str = "koan_refresh";
23const REFRESH_COOKIE_PATH: &str = "/auth";
24/// A cookie left at this narrower path is sent ahead of the one at
25/// `REFRESH_COOKIE_PATH` and shadows it, so every response that sets or clears
26/// the refresh cookie clears this one too.
27const STALE_REFRESH_COOKIE_PATH: &str = "/auth/refresh";
28
29/// Fixed-window per-IP cap on requests. By default, the cap on login attempts:
30///
31/// Argon2 is tuned to cost ~19MiB and real CPU per verification, which is
32/// correct for resisting cracking and ruinous when anyone may trigger it at
33/// will: a few hundred concurrent logins exhaust memory and starve every other
34/// request. The window is coarse on purpose — it bounds cost, it is not a quota.
35const LOGIN_WINDOW_SECS: u64 = 60;
36const LOGIN_MAX_PER_WINDOW: u32 = 10;
37/// Above this many tracked IPs, drop stale windows before inserting more.
38const TRACKED_IPS_MAX: usize = 4096;
39
40pub struct RateLimiter {
41    windows: Mutex<HashMap<IpAddr, (u64, u32)>>,
42    window_secs: u64,
43    max: u32,
44}
45
46impl Default for RateLimiter {
47    fn default() -> Self {
48        Self::new(LOGIN_WINDOW_SECS, LOGIN_MAX_PER_WINDOW)
49    }
50}
51
52impl RateLimiter {
53    pub fn new(window_secs: u64, max: u32) -> Self {
54        Self {
55            windows: Mutex::default(),
56            window_secs,
57            max,
58        }
59    }
60
61    /// Returns false when `ip` has spent its allowance for the current window.
62    fn allow(&self, ip: IpAddr) -> bool {
63        let now = auth::now_unix();
64        let mut windows = self.windows.lock().unwrap_or_else(|e| e.into_inner());
65
66        if windows.len() > TRACKED_IPS_MAX {
67            windows.retain(|_, (start, _)| now.saturating_sub(*start) < self.window_secs);
68        }
69
70        let entry = windows.entry(ip).or_insert((now, 0));
71        if now.saturating_sub(entry.0) >= self.window_secs {
72            *entry = (now, 0);
73        }
74        entry.1 += 1;
75        entry.1 <= self.max
76    }
77}
78
79/// The address a request came from.
80///
81/// Behind a reverse proxy the TCP peer is the proxy, for every client, so a
82/// limit keyed on it is one bucket for everyone. When the peer is on a private
83/// or loopback address — a proxy in the cluster or on the host — the last
84/// `X-Forwarded-For` entry is the one that proxy appended, and is the client.
85/// Earlier entries are whatever the client sent, and are ignored. A public
86/// peer is the client, and its header is not believed. Nor is it when the
87/// peer is unknown: a listener served without its connection info would
88/// otherwise let every client pick its own address.
89pub(crate) fn client_ip(request: &axum::extract::Request) -> IpAddr {
90    let Some(ConnectInfo(peer)) = request.extensions().get::<ConnectInfo<SocketAddr>>() else {
91        return IpAddr::V4(std::net::Ipv4Addr::UNSPECIFIED);
92    };
93    let peer = peer.ip();
94    if !is_internal(peer) {
95        return peer;
96    }
97    request
98        .headers()
99        .get_all("x-forwarded-for")
100        .iter()
101        .filter_map(|v| v.to_str().ok())
102        .flat_map(|v| v.split(','))
103        .filter_map(|ip| ip.trim().parse::<IpAddr>().ok())
104        .next_back()
105        .unwrap_or(peer)
106}
107
108fn is_internal(ip: IpAddr) -> bool {
109    match ip {
110        IpAddr::V4(v4) => v4.is_private() || v4.is_loopback(),
111        IpAddr::V6(v6) => v6.is_loopback() || (v6.segments()[0] & 0xfe00) == 0xfc00,
112    }
113}
114
115// ---------------------------------------------------------------------------
116// Shared state
117// ---------------------------------------------------------------------------
118
119#[derive(Clone)]
120pub struct AuthRouteState {
121    pub pool: Arc<Pool>,
122    pub private_pem: Arc<Vec<u8>>,
123    pub public_pem: Arc<Vec<u8>>,
124    pub access_ttl_secs: u64,
125    pub refresh_ttl_secs: u64,
126    /// Mark cookies `Secure`. Only when clients reach koan over HTTPS —
127    /// a browser discards a `Secure` cookie delivered over plain `http://`, so
128    /// setting this on a LAN deployment silently breaks cookie auth entirely.
129    pub cookie_secure: bool,
130    pub login_limiter: Arc<RateLimiter>,
131}
132
133impl AuthRouteState {
134    /// `SameSite=Lax` keeps the cookie off cross-site requests, which is what
135    /// takes the WebSocket and safelisted-content-type CSRF paths off the table.
136    fn cookie(&self, name: &str, value: &str, path: &str, max_age: u64) -> String {
137        let secure = if self.cookie_secure { "; Secure" } else { "" };
138        format!("{name}={value}; HttpOnly; SameSite=Lax; Path={path}; Max-Age={max_age}{secure}")
139    }
140
141    fn access_cookie(&self, token: &str) -> String {
142        self.cookie("koan_access", token, "/", self.access_ttl_secs)
143    }
144
145    fn refresh_cookie(&self, token: &str) -> String {
146        self.cookie(
147            REFRESH_COOKIE,
148            token,
149            REFRESH_COOKIE_PATH,
150            self.refresh_ttl_secs,
151        )
152    }
153
154    fn stale_refresh_cookie(&self) -> String {
155        self.cookie(REFRESH_COOKIE, "", STALE_REFRESH_COOKIE_PATH, 0)
156    }
157
158    /// The cookies that open a session. Appended rather than inserted: a
159    /// header array as a response part keeps only the last value per name.
160    pub(crate) fn session_cookies(
161        &self,
162        access_token: &str,
163        refresh_token: &str,
164    ) -> AppendHeaders<[(axum::http::HeaderName, String); 3]> {
165        AppendHeaders([
166            (SET_COOKIE, self.access_cookie(access_token)),
167            (SET_COOKIE, self.refresh_cookie(refresh_token)),
168            (SET_COOKIE, self.stale_refresh_cookie()),
169        ])
170    }
171
172    /// Clears the access cookie and both refresh cookies: what signing out sets.
173    pub(crate) fn cleared_cookies(&self) -> AppendHeaders<[(axum::http::HeaderName, String); 3]> {
174        AppendHeaders([
175            (SET_COOKIE, self.cookie("koan_access", "", "/", 0)),
176            (
177                SET_COOKIE,
178                self.cookie(REFRESH_COOKIE, "", REFRESH_COOKIE_PATH, 0),
179            ),
180            (SET_COOKIE, self.stale_refresh_cookie()),
181        ])
182    }
183}
184
185/// A hash with the same parameters as a real one, to verify unknown usernames
186/// against.
187pub(crate) fn dummy_password_hash() -> &'static str {
188    static HASH: std::sync::OnceLock<String> = std::sync::OnceLock::new();
189    HASH.get_or_init(|| auth::hash_password("koan-dummy-password").unwrap_or_default())
190}
191
192/// Read the refresh token from the request body, falling back to the cookie so a
193/// browser client never has to keep one in script-reachable storage.
194pub(crate) fn refresh_token_from(
195    body: Option<&str>,
196    headers: &axum::http::HeaderMap,
197) -> Option<String> {
198    if let Some(t) = body.filter(|t| !t.is_empty()) {
199        return Some(t.to_owned());
200    }
201    headers
202        .get(COOKIE)
203        .and_then(|v| v.to_str().ok())
204        .and_then(|cookies| {
205            cookies.split(';').find_map(|c| {
206                c.trim()
207                    .strip_prefix(&format!("{REFRESH_COOKIE}="))
208                    .map(str::to_owned)
209            })
210        })
211}
212
213/// Reject login attempts once an IP has spent its window.
214///
215/// A middleware rather than an extractor so it runs before the request body is
216/// read and before the database is touched.
217pub(crate) async fn login_rate_limit(
218    State(state): State<AuthRouteState>,
219    request: axum::extract::Request,
220    next: axum::middleware::Next,
221) -> Response {
222    let ip = client_ip(&request);
223
224    if !state.login_limiter.allow(ip) {
225        return (
226            StatusCode::TOO_MANY_REQUESTS,
227            Json(MessageResponse {
228                message: "too many login attempts".into(),
229            }),
230        )
231            .into_response();
232    }
233    next.run(request).await
234}
235
236/// `RateLimiter` as a middleware, for routes other than sign-in.
237pub(crate) async fn rate_limit(
238    State(limiter): State<Arc<RateLimiter>>,
239    request: axum::extract::Request,
240    next: axum::middleware::Next,
241) -> Response {
242    if !limiter.allow(client_ip(&request)) {
243        return (StatusCode::TOO_MANY_REQUESTS, "too many requests").into_response();
244    }
245    next.run(request).await
246}
247
248impl AuthRouteState {
249    fn open_db(&self) -> Result<Handle<'_>, (StatusCode, String)> {
250        self.pool.get().map_err(|e| {
251            log::error!("auth db open error: {}", e);
252            (
253                StatusCode::INTERNAL_SERVER_ERROR,
254                "internal error".to_string(),
255            )
256        })
257    }
258}
259
260// ---------------------------------------------------------------------------
261// Request/response types
262// ---------------------------------------------------------------------------
263
264#[derive(Deserialize)]
265pub struct LoginRequest {
266    pub username: String,
267    pub password: String,
268}
269
270#[derive(Serialize)]
271pub struct LoginResponse {
272    pub access_token: String,
273    pub refresh_token: String,
274    pub token_type: String,
275    pub expires_in: u64,
276    pub user: UserInfo,
277}
278
279#[derive(Serialize)]
280pub struct UserInfo {
281    pub id: i64,
282    pub username: String,
283    pub role: String,
284}
285
286#[derive(Deserialize, Default)]
287#[serde(default)]
288pub struct RefreshRequest {
289    pub refresh_token: Option<String>,
290}
291
292#[derive(Serialize)]
293pub struct RefreshResponse {
294    pub access_token: String,
295    pub refresh_token: String,
296    pub token_type: String,
297    pub expires_in: u64,
298}
299
300#[derive(Deserialize, Default)]
301#[serde(default)]
302pub struct LogoutRequest {
303    pub refresh_token: Option<String>,
304}
305
306#[derive(Serialize)]
307pub struct MessageResponse {
308    pub message: String,
309}
310
311// ---------------------------------------------------------------------------
312// Router
313// ---------------------------------------------------------------------------
314
315pub fn auth_router(state: AuthRouteState) -> axum::Router {
316    let router = axum::Router::new()
317        .route(
318            "/auth/login",
319            post(login).layer(axum::middleware::from_fn_with_state(
320                state.clone(),
321                login_rate_limit,
322            )),
323        )
324        .route("/auth/refresh", post(refresh))
325        .route("/auth/logout", post(logout));
326    auth_perimeter(router, AUTH_TIMEOUT).with_state(state)
327}
328
329/// How long a sign-in, refresh or sign-out may take, reading its body included.
330const AUTH_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(10);
331
332/// These routes are unauthenticated by definition and the work behind them is
333/// deliberately expensive, so they get their own ceiling rather than sharing
334/// the GraphQL one. It sheds rather than queues, and each request has a
335/// deadline: the permit is held while the body is read, so two clients that
336/// never finish sending one would otherwise hold a route shut.
337fn auth_perimeter<S>(router: axum::Router<S>, timeout: std::time::Duration) -> axum::Router<S>
338where
339    S: Clone + Send + Sync + 'static,
340{
341    router
342        .layer(tower_http::timeout::TimeoutLayer::with_status_code(
343            StatusCode::REQUEST_TIMEOUT,
344            timeout,
345        ))
346        .layer(
347            tower::ServiceBuilder::new()
348                .layer(axum::error_handling::HandleErrorLayer::new(
349                    |_: tower::BoxError| async { (StatusCode::SERVICE_UNAVAILABLE, "busy") },
350                ))
351                .load_shed()
352                .concurrency_limit(2),
353        )
354}
355
356// ---------------------------------------------------------------------------
357// Handlers
358// ---------------------------------------------------------------------------
359
360/// Check a username and password and open a session: a fresh access token and
361/// a stored refresh token. The error is the response to send. Shared by the
362/// JSON login and the web UI's sign-in form, so both are one implementation.
363///
364/// On the blocking pool as a whole: the queries, argon2 and the sealing all
365/// block, and on a runtime worker they stall every other request the server
366/// is handling.
367pub(crate) async fn authenticate(
368    state: &AuthRouteState,
369    username: &str,
370    password: &str,
371) -> Result<(auth_queries::UserRow, String, String), Box<Response>> {
372    let state = state.clone();
373    let (username, password) = (username.to_owned(), password.to_owned());
374    tokio::task::spawn_blocking(move || authenticate_blocking(&state, &username, &password))
375        .await
376        .unwrap_or_else(|_| Err(Box::new(StatusCode::INTERNAL_SERVER_ERROR.into_response())))
377}
378
379fn authenticate_blocking(
380    state: &AuthRouteState,
381    username: &str,
382    password: &str,
383) -> Result<(auth_queries::UserRow, String, String), Box<Response>> {
384    let db = match state.open_db() {
385        Ok(db) => db,
386        Err((status, msg)) => return Err(Box::new((status, msg).into_response())),
387    };
388
389    let user = match auth_queries::get_user_by_username(&db.conn, username) {
390        Ok(Some(u)) => u,
391        Ok(None) => {
392            // Pay for a verify anyway, so response time doesn't say which
393            // usernames exist.
394            let _ = auth::verify_password(password, dummy_password_hash());
395            return Err(Box::new(
396                (
397                    StatusCode::UNAUTHORIZED,
398                    Json(MessageResponse {
399                        message: "invalid username or password".into(),
400                    }),
401                )
402                    .into_response(),
403            ));
404        }
405        Err(e) => {
406            log::error!("auth login db error: {}", e);
407            return Err(Box::new(
408                (StatusCode::INTERNAL_SERVER_ERROR, "internal error").into_response(),
409            ));
410        }
411    };
412
413    if auth::verify_password(password, &user.password_hash).is_err() {
414        return Err(Box::new(
415            (
416                StatusCode::UNAUTHORIZED,
417                Json(MessageResponse {
418                    message: "invalid username or password".into(),
419                }),
420            )
421                .into_response(),
422        ));
423    }
424
425    // Lets the account use Subsonic token auth; see `auth::seal_password`.
426    if let Err(e) = auth_queries::remember_password(&db.conn, username, password) {
427        log::warn!("could not seal the password for Subsonic token auth: {e}");
428    }
429
430    let access_token = match auth::mint_access_token(
431        &state.private_pem,
432        user.id,
433        &user.username,
434        user.role,
435        state.access_ttl_secs,
436    ) {
437        Ok(t) => t,
438        Err(e) => {
439            log::error!("auth mint token error: {}", e);
440            return Err(Box::new(
441                (StatusCode::INTERNAL_SERVER_ERROR, "token error").into_response(),
442            ));
443        }
444    };
445
446    let refresh_token_id = match auth::random_token() {
447        Ok(t) => t,
448        Err(e) => {
449            log::error!("auth refresh token generation error: {}", e);
450            return Err(Box::new(
451                (StatusCode::INTERNAL_SERVER_ERROR, "token error").into_response(),
452            ));
453        }
454    };
455    let refresh_expires = auth::now_unix() as i64 + state.refresh_ttl_secs as i64;
456    if let Err(e) =
457        auth_queries::store_refresh_token(&db.conn, &refresh_token_id, user.id, refresh_expires)
458    {
459        log::error!("auth store refresh token error: {}", e);
460        return Err(Box::new(
461            (StatusCode::INTERNAL_SERVER_ERROR, "token error").into_response(),
462        ));
463    }
464
465    // Clear out expired tokens; a failure here does not fail the sign-in.
466    let _ = auth_queries::cleanup_expired_tokens(&db.conn);
467
468    Ok((user, access_token, refresh_token_id))
469}
470
471async fn login(State(state): State<AuthRouteState>, Json(req): Json<LoginRequest>) -> Response {
472    let (user, access_token, refresh_token_id) =
473        match authenticate(&state, &req.username, &req.password).await {
474            Ok(session) => session,
475            Err(resp) => return *resp,
476        };
477
478    let cookies = state.session_cookies(&access_token, &refresh_token_id);
479
480    let resp = LoginResponse {
481        access_token,
482        // Also in the body: the CLI and other non-browser clients have no cookie
483        // jar and store this in config.local.toml.
484        refresh_token: refresh_token_id,
485        token_type: "Bearer".into(),
486        expires_in: state.access_ttl_secs,
487        user: UserInfo {
488            id: user.id,
489            username: user.username,
490            role: user.role.as_str().into(),
491        },
492    };
493
494    (StatusCode::OK, cookies, Json(resp)).into_response()
495}
496
497/// How long after its refresh a spent OAuth refresh token may come back as a
498/// retry rather than a replay.
499const REPLAY_GRACE_SECS: i64 = 30;
500
501/// Spend a refresh token for a new access token and a new refresh token. The
502/// error is the response to send. Shared by the JSON refresh and the web UI's
503/// session resume.
504pub(crate) fn rotate(
505    state: &AuthRouteState,
506    supplied: &str,
507) -> Result<(String, String), Box<Response>> {
508    let db = match state.open_db() {
509        Ok(db) => db,
510        Err((status, msg)) => return Err(Box::new((status, msg).into_response())),
511    };
512
513    // Atomically consume (validate + revoke) the refresh token in a single
514    // statement to prevent TOCTOU races during token rotation.
515    let token = match auth_queries::consume_refresh_token(&db.conn, supplied) {
516        Ok(Some(t)) => t,
517        Ok(None) => {
518            match auth_queries::revoke_replayed_grant(&db.conn, supplied, REPLAY_GRACE_SECS) {
519                Ok(0) => {}
520                Ok(n) => log::warn!("a spent OAuth refresh token came back: revoked {n} tokens"),
521                Err(e) => log::error!("auth replay check error: {e}"),
522            }
523            return Err(Box::new(
524                (
525                    StatusCode::UNAUTHORIZED,
526                    Json(MessageResponse {
527                        message: "invalid or expired refresh token".into(),
528                    }),
529                )
530                    .into_response(),
531            ));
532        }
533        Err(e) => {
534            log::error!("auth refresh db error: {}", e);
535            return Err(Box::new(
536                (StatusCode::INTERNAL_SERVER_ERROR, "internal error").into_response(),
537            ));
538        }
539    };
540
541    let user = match auth_queries::get_user_by_id(&db.conn, token.user_id) {
542        Ok(Some(u)) => u,
543        Ok(None) => {
544            return Err(Box::new(
545                (
546                    StatusCode::UNAUTHORIZED,
547                    Json(MessageResponse {
548                        message: "user not found".into(),
549                    }),
550                )
551                    .into_response(),
552            ));
553        }
554        Err(e) => {
555            log::error!("auth refresh user lookup error: {}", e);
556            return Err(Box::new(
557                (StatusCode::INTERNAL_SERVER_ERROR, "internal error").into_response(),
558            ));
559        }
560    };
561
562    // A grant's tokens stay as narrow as the grant.
563    let access_token = match auth::mint_scoped_token(
564        &state.private_pem,
565        user.id,
566        &user.username,
567        user.role,
568        state.access_ttl_secs,
569        token.grant.as_ref().map(|_| auth::MCP_SCOPE),
570    ) {
571        Ok(t) => t,
572        Err(e) => {
573            log::error!("auth mint token error: {}", e);
574            return Err(Box::new(
575                (StatusCode::INTERNAL_SERVER_ERROR, "token error").into_response(),
576            ));
577        }
578    };
579
580    let new_refresh_id = match auth::random_token() {
581        Ok(t) => t,
582        Err(e) => {
583            log::error!("auth refresh token generation error: {}", e);
584            return Err(Box::new(
585                (StatusCode::INTERNAL_SERVER_ERROR, "token error").into_response(),
586            ));
587        }
588    };
589    let refresh_expires = auth::now_unix() as i64 + state.refresh_ttl_secs as i64;
590    if let Err(e) = auth_queries::store_grant_token(
591        &db.conn,
592        &new_refresh_id,
593        user.id,
594        refresh_expires,
595        token.grant.as_ref(),
596    ) {
597        log::error!("auth store refresh token error: {}", e);
598        return Err(Box::new(
599            (StatusCode::INTERNAL_SERVER_ERROR, "token error").into_response(),
600        ));
601    }
602
603    Ok((access_token, new_refresh_id))
604}
605
606async fn refresh(
607    State(state): State<AuthRouteState>,
608    headers: axum::http::HeaderMap,
609    body: Option<Json<RefreshRequest>>,
610) -> Response {
611    let supplied = body.and_then(|Json(req)| req.refresh_token);
612    let Some(supplied) = refresh_token_from(supplied.as_deref(), &headers) else {
613        return (
614            StatusCode::UNAUTHORIZED,
615            Json(MessageResponse {
616                message: "missing refresh token".into(),
617            }),
618        )
619            .into_response();
620    };
621
622    let rotating = state.clone();
623    let rotated = tokio::task::spawn_blocking(move || rotate(&rotating, &supplied))
624        .await
625        .unwrap_or_else(|_| Err(Box::new(StatusCode::INTERNAL_SERVER_ERROR.into_response())));
626    let (access_token, new_refresh_id) = match rotated {
627        Ok(pair) => pair,
628        Err(resp) => return *resp,
629    };
630
631    let cookies = state.session_cookies(&access_token, &new_refresh_id);
632
633    let resp = RefreshResponse {
634        access_token,
635        refresh_token: new_refresh_id,
636        token_type: "Bearer".into(),
637        expires_in: state.access_ttl_secs,
638    };
639
640    (StatusCode::OK, cookies, Json(resp)).into_response()
641}
642
643async fn logout(
644    State(state): State<AuthRouteState>,
645    headers: axum::http::HeaderMap,
646    body: Option<Json<LogoutRequest>>,
647) -> Response {
648    let supplied = body.and_then(|Json(req)| req.refresh_token);
649    let token = refresh_token_from(supplied.as_deref(), &headers);
650    let revoking = state.clone();
651    let revoked = tokio::task::spawn_blocking(move || {
652        let db = revoking.open_db()?;
653        if let Some(token) = token {
654            let _ = auth_queries::revoke_refresh_token(&db.conn, &token);
655        }
656        Ok(())
657    })
658    .await
659    .unwrap_or_else(|_| {
660        Err((
661            StatusCode::INTERNAL_SERVER_ERROR,
662            "internal error".to_string(),
663        ))
664    });
665    if let Err((status, msg)) = revoked {
666        return (status, msg).into_response();
667    }
668
669    let cookies = state.cleared_cookies();
670
671    (
672        StatusCode::OK,
673        cookies,
674        Json(MessageResponse {
675            message: "logged out".into(),
676        }),
677    )
678        .into_response()
679}
680
681// ---------------------------------------------------------------------------
682// Tests
683// ---------------------------------------------------------------------------
684
685#[cfg(test)]
686mod tests {
687    use super::*;
688
689    #[test]
690    fn login_limiter_caps_a_single_ip() {
691        let limiter = RateLimiter::default();
692        let ip: IpAddr = "10.0.0.5".parse().unwrap();
693        for _ in 0..LOGIN_MAX_PER_WINDOW {
694            assert!(limiter.allow(ip));
695        }
696        assert!(!limiter.allow(ip));
697
698        // Other callers are unaffected.
699        assert!(limiter.allow("10.0.0.6".parse().unwrap()));
700    }
701
702    fn request_from(peer: &str, forwarded: Option<&str>) -> axum::extract::Request {
703        let mut request = axum::http::Request::new(axum::body::Body::empty());
704        request.extensions_mut().insert(ConnectInfo(
705            format!("{peer}:1234").parse::<SocketAddr>().unwrap(),
706        ));
707        if let Some(f) = forwarded {
708            request
709                .headers_mut()
710                .insert("x-forwarded-for", f.parse().unwrap());
711        }
712        request
713    }
714
715    #[test]
716    fn client_ip_believes_only_an_internal_proxy() {
717        // Behind the cluster's proxy: the entry it appended, not the client's.
718        let r = request_from("10.42.0.7", Some("6.6.6.6, 203.0.113.9"));
719        assert_eq!(client_ip(&r), "203.0.113.9".parse::<IpAddr>().unwrap());
720        // A public peer is the client, whatever it claims.
721        let r = request_from("198.51.100.4", Some("10.0.0.1"));
722        assert_eq!(client_ip(&r), "198.51.100.4".parse::<IpAddr>().unwrap());
723        // An internal peer with no header is itself.
724        let r = request_from("10.42.0.7", None);
725        assert_eq!(client_ip(&r), "10.42.0.7".parse::<IpAddr>().unwrap());
726        // No peer known: the header is the client's own claim.
727        let mut r = axum::http::Request::new(axum::body::Body::empty());
728        r.headers_mut()
729            .insert("x-forwarded-for", "203.0.113.9".parse().unwrap());
730        assert_eq!(client_ip(&r), IpAddr::V4(std::net::Ipv4Addr::UNSPECIFIED));
731    }
732
733    #[tokio::test]
734    async fn stalled_bodies_are_shed_then_timed_out() {
735        use tower::ServiceExt as _;
736        async fn read(_: axum::body::Bytes) -> StatusCode {
737            StatusCode::OK
738        }
739        // With its state, as `auth_router` does: that is what builds each
740        // route's layers once, rather than per request.
741        let app = auth_perimeter(
742            axum::Router::new().route("/auth/login", post(read)),
743            std::time::Duration::from_millis(200),
744        )
745        .with_state(());
746        let stalled = || {
747            axum::http::Request::post("/auth/login")
748                .body(axum::body::Body::from_stream(tokio_stream::pending::<
749                    Result<axum::body::Bytes, std::io::Error>,
750                >()))
751                .unwrap()
752        };
753        let held: Vec<_> = (0..2)
754            .map(|_| tokio::spawn(app.clone().oneshot(stalled())))
755            .collect();
756        tokio::time::sleep(std::time::Duration::from_millis(50)).await;
757        let shed = app.clone().oneshot(stalled()).await.unwrap();
758        assert_eq!(shed.status(), StatusCode::SERVICE_UNAVAILABLE);
759        for h in held {
760            assert_eq!(
761                h.await.unwrap().unwrap().status(),
762                StatusCode::REQUEST_TIMEOUT
763            );
764        }
765        let ok = axum::http::Request::post("/auth/login")
766            .body(axum::body::Body::empty())
767            .unwrap();
768        assert_eq!(app.oneshot(ok).await.unwrap().status(), StatusCode::OK);
769    }
770
771    #[test]
772    fn refresh_token_falls_back_to_the_cookie() {
773        let mut headers = axum::http::HeaderMap::new();
774        headers.insert(
775            COOKIE,
776            format!("a=1; {REFRESH_COOKIE}=from-cookie; b=2")
777                .parse()
778                .unwrap(),
779        );
780
781        assert_eq!(
782            refresh_token_from(None, &headers).as_deref(),
783            Some("from-cookie")
784        );
785        assert_eq!(
786            refresh_token_from(Some("from-body"), &headers).as_deref(),
787            Some("from-body")
788        );
789        assert_eq!(
790            refresh_token_from(None, &axum::http::HeaderMap::new()),
791            None
792        );
793    }
794
795    #[test]
796    fn cookies_are_lax_and_only_secure_when_tls_is_in_play() {
797        let state = |cookie_secure| AuthRouteState {
798            pool: Arc::new(Pool::new("/nonexistent".into())),
799            private_pem: Arc::new(Vec::new()),
800            public_pem: Arc::new(Vec::new()),
801            access_ttl_secs: 900,
802            refresh_ttl_secs: 60,
803            cookie_secure,
804            login_limiter: Arc::new(RateLimiter::default()),
805        };
806
807        let plain = state(false).access_cookie("tok");
808        assert!(plain.contains("SameSite=Lax"));
809        assert!(plain.contains("HttpOnly"));
810        assert!(!plain.contains("Secure"));
811
812        assert!(state(true).access_cookie("tok").contains("; Secure"));
813
814        // Every cookie reaches the browser, not only the last one set.
815        let resp = (StatusCode::OK, state(false).session_cookies("a", "r")).into_response();
816        let set: Vec<_> = resp.headers().get_all(SET_COOKIE).iter().collect();
817        assert_eq!(set.len(), 3);
818        assert_eq!(
819            (StatusCode::OK, state(false).cleared_cookies())
820                .into_response()
821                .headers()
822                .get_all(SET_COOKIE)
823                .iter()
824                .count(),
825            3
826        );
827
828        // The refresh cookie never rides along on an API call.
829        let refresh = state(false).refresh_cookie("tok");
830        assert!(refresh.contains("Path=/auth;"));
831        assert!(refresh.contains("HttpOnly"));
832    }
833}