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