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