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    /// The one cookie a session vouched for by an authenticating proxy has:
206    /// an access token, and any refresh cookie the browser held cleared. Such
207    /// a session is derived again from the proxy's header on each page load,
208    /// so it never outlives what the proxy says.
209    pub(crate) fn proxied_cookies(
210        &self,
211        access_token: &str,
212    ) -> AppendHeaders<[(axum::http::HeaderName, String); 3]> {
213        AppendHeaders([
214            (SET_COOKIE, self.access_cookie(access_token)),
215            (
216                SET_COOKIE,
217                self.cookie(REFRESH_COOKIE, "", REFRESH_COOKIE_PATH, 0),
218            ),
219            (SET_COOKIE, self.stale_refresh_cookie()),
220        ])
221    }
222
223    /// Clears the access cookie and both refresh cookies: what signing out sets.
224    pub(crate) fn cleared_cookies(&self) -> AppendHeaders<[(axum::http::HeaderName, String); 3]> {
225        AppendHeaders([
226            (SET_COOKIE, self.cookie("koan_access", "", "/", 0)),
227            (
228                SET_COOKIE,
229                self.cookie(REFRESH_COOKIE, "", REFRESH_COOKIE_PATH, 0),
230            ),
231            (SET_COOKIE, self.stale_refresh_cookie()),
232        ])
233    }
234}
235
236/// A hash with the same parameters as a real one, to verify unknown usernames
237/// against.
238pub(crate) fn dummy_password_hash() -> &'static str {
239    static HASH: std::sync::OnceLock<String> = std::sync::OnceLock::new();
240    HASH.get_or_init(|| auth::hash_password("koan-dummy-password").unwrap_or_default())
241}
242
243/// Read the refresh token from the request body, falling back to the cookie so a
244/// browser client never has to keep one in script-reachable storage.
245pub(crate) fn refresh_token_from(
246    body: Option<&str>,
247    headers: &axum::http::HeaderMap,
248) -> Option<String> {
249    if let Some(t) = body.filter(|t| !t.is_empty()) {
250        return Some(t.to_owned());
251    }
252    headers
253        .get(COOKIE)
254        .and_then(|v| v.to_str().ok())
255        .and_then(|cookies| {
256            cookies.split(';').find_map(|c| {
257                c.trim()
258                    .strip_prefix(&format!("{REFRESH_COOKIE}="))
259                    .map(str::to_owned)
260            })
261        })
262}
263
264/// Reject login attempts once an IP has spent its window.
265///
266/// A middleware rather than an extractor so it runs before the request body is
267/// read and before the database is touched.
268pub(crate) async fn login_rate_limit(
269    State(state): State<AuthRouteState>,
270    request: axum::extract::Request,
271    next: axum::middleware::Next,
272) -> Response {
273    let ip = client_ip(&request);
274
275    if !state.login_limiter.allow(ip) {
276        return (
277            StatusCode::TOO_MANY_REQUESTS,
278            Json(MessageResponse {
279                message: "too many login attempts".into(),
280            }),
281        )
282            .into_response();
283    }
284    next.run(request).await
285}
286
287/// `RateLimiter` as a middleware, for routes other than sign-in.
288pub(crate) async fn rate_limit(
289    State(limiter): State<Arc<RateLimiter>>,
290    request: axum::extract::Request,
291    next: axum::middleware::Next,
292) -> Response {
293    if !limiter.allow(client_ip(&request)) {
294        return (StatusCode::TOO_MANY_REQUESTS, "too many requests").into_response();
295    }
296    next.run(request).await
297}
298
299impl AuthRouteState {
300    fn open_db(&self) -> Result<Handle<'_>, (StatusCode, String)> {
301        self.pool.get().map_err(|e| {
302            log::error!("auth db open error: {}", e);
303            (
304                StatusCode::INTERNAL_SERVER_ERROR,
305                "internal error".to_string(),
306            )
307        })
308    }
309}
310
311// ---------------------------------------------------------------------------
312// Request/response types
313// ---------------------------------------------------------------------------
314
315#[derive(Deserialize)]
316pub struct LoginRequest {
317    pub username: String,
318    pub password: String,
319}
320
321#[derive(Serialize)]
322pub struct LoginResponse {
323    pub access_token: String,
324    pub refresh_token: String,
325    pub token_type: String,
326    pub expires_in: u64,
327    pub user: UserInfo,
328}
329
330#[derive(Serialize)]
331pub struct UserInfo {
332    pub id: i64,
333    pub username: String,
334    pub role: String,
335}
336
337#[derive(Deserialize, Default)]
338#[serde(default)]
339pub struct RefreshRequest {
340    pub refresh_token: Option<String>,
341}
342
343#[derive(Serialize)]
344pub struct RefreshResponse {
345    pub access_token: String,
346    pub refresh_token: String,
347    pub token_type: String,
348    pub expires_in: u64,
349}
350
351#[derive(Deserialize, Default)]
352#[serde(default)]
353pub struct LogoutRequest {
354    pub refresh_token: Option<String>,
355}
356
357#[derive(Serialize)]
358pub struct MessageResponse {
359    pub message: String,
360}
361
362// ---------------------------------------------------------------------------
363// Router
364// ---------------------------------------------------------------------------
365
366pub fn auth_router(state: AuthRouteState) -> axum::Router {
367    let router = axum::Router::new()
368        .route(
369            "/auth/login",
370            post(login).layer(axum::middleware::from_fn_with_state(
371                state.clone(),
372                login_rate_limit,
373            )),
374        )
375        .route("/auth/refresh", post(refresh))
376        .route("/auth/logout", post(logout));
377    auth_perimeter(router, AUTH_TIMEOUT).with_state(state)
378}
379
380/// How long a sign-in, refresh or sign-out may take, reading its body included.
381const AUTH_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(10);
382
383/// These routes are unauthenticated by definition and the work behind them is
384/// deliberately expensive, so they get their own ceiling rather than sharing
385/// the GraphQL one. It sheds rather than queues, and each request has a
386/// deadline: the permit is held while the body is read, so two clients that
387/// never finish sending one would otherwise hold a route shut.
388fn auth_perimeter<S>(router: axum::Router<S>, timeout: std::time::Duration) -> axum::Router<S>
389where
390    S: Clone + Send + Sync + 'static,
391{
392    router
393        .layer(tower_http::timeout::TimeoutLayer::with_status_code(
394            StatusCode::REQUEST_TIMEOUT,
395            timeout,
396        ))
397        .layer(
398            tower::ServiceBuilder::new()
399                .layer(axum::error_handling::HandleErrorLayer::new(
400                    |_: tower::BoxError| async { (StatusCode::SERVICE_UNAVAILABLE, "busy") },
401                ))
402                .load_shed()
403                .concurrency_limit(2),
404        )
405}
406
407// ---------------------------------------------------------------------------
408// Handlers
409// ---------------------------------------------------------------------------
410
411/// Check a username and password and open a session: a fresh access token and
412/// a stored refresh token. The error is the response to send. Shared by the
413/// JSON login and the web UI's sign-in form, so both are one implementation.
414///
415/// 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
416/// is handling. The password goes through the server's one verifier, so the
417/// ceiling on argon2 and the per-username budget on failures hold here as on
418/// every other door.
419pub(crate) async fn authenticate(
420    state: &AuthRouteState,
421    username: &str,
422    password: &str,
423    from: IpAddr,
424) -> Result<(auth_queries::UserRow, String, String), Box<Response>> {
425    let state = state.clone();
426    let (username, password) = (username.to_owned(), password.to_owned());
427    tokio::task::spawn_blocking(move || authenticate_blocking(&state, &username, &password, from))
428        .await
429        .unwrap_or_else(|_| Err(Box::new(StatusCode::INTERNAL_SERVER_ERROR.into_response())))
430}
431
432fn authenticate_blocking(
433    state: &AuthRouteState,
434    username: &str,
435    password: &str,
436    from: IpAddr,
437) -> Result<(auth_queries::UserRow, String, String), Box<Response>> {
438    use super::password::Refused;
439    let refused = |status, message: &str| {
440        Box::new(
441            (
442                status,
443                Json(MessageResponse {
444                    message: message.into(),
445                }),
446            )
447                .into_response(),
448        )
449    };
450    if state.users.spent(username, from) {
451        return Err(refused(
452            StatusCode::TOO_MANY_REQUESTS,
453            "too many failed sign-ins for this account; try again in a minute",
454        ));
455    }
456    let user = match state.users.verify(username, password) {
457        Ok(user) => {
458            state.users.signed_in(username, from);
459            user
460        }
461        Err(Refused::Wrong) => {
462            state.users.failed(username);
463            return Err(refused(
464                StatusCode::UNAUTHORIZED,
465                "invalid username or password",
466            ));
467        }
468        Err(Refused::Busy) => {
469            return Err(refused(
470                StatusCode::SERVICE_UNAVAILABLE,
471                "busy; try again in a moment",
472            ));
473        }
474    };
475
476    let db = match state.open_db() {
477        Ok(db) => db,
478        Err((status, msg)) => return Err(Box::new((status, msg).into_response())),
479    };
480
481    let access_token = match auth::mint_access_token(
482        &state.private_pem,
483        user.id,
484        &user.username,
485        user.role,
486        state.access_ttl_secs,
487    ) {
488        Ok(t) => t,
489        Err(e) => {
490            log::error!("auth mint token error: {}", e);
491            return Err(Box::new(
492                (StatusCode::INTERNAL_SERVER_ERROR, "token error").into_response(),
493            ));
494        }
495    };
496
497    let refresh_token_id = match auth::random_token() {
498        Ok(t) => t,
499        Err(e) => {
500            log::error!("auth refresh token generation error: {}", e);
501            return Err(Box::new(
502                (StatusCode::INTERNAL_SERVER_ERROR, "token error").into_response(),
503            ));
504        }
505    };
506    let refresh_expires = auth::now_unix() as i64 + state.refresh_ttl_secs as i64;
507    if let Err(e) =
508        auth_queries::store_refresh_token(&db.conn, &refresh_token_id, user.id, refresh_expires)
509    {
510        log::error!("auth store refresh token error: {}", e);
511        return Err(Box::new(
512            (StatusCode::INTERNAL_SERVER_ERROR, "token error").into_response(),
513        ));
514    }
515
516    // Clear out expired tokens; a failure here does not fail the sign-in.
517    let _ = auth_queries::cleanup_expired_tokens(&db.conn);
518
519    Ok((user, access_token, refresh_token_id))
520}
521
522/// An access token for the account an authenticating proxy names, and no
523/// refresh token: see `proxied_cookies`. No password is checked; the caller
524/// has established that the request came through that proxy. `None` for a
525/// name the server has no account for.
526pub(crate) async fn proxied_access(state: &AuthRouteState, username: &str) -> Option<String> {
527    let state = state.clone();
528    let username = username.to_owned();
529    tokio::task::spawn_blocking(move || {
530        let db = state.open_db().ok()?;
531        let user = auth_queries::get_user_by_username(&db.conn, &username).ok()??;
532        auth::mint_access_token(
533            &state.private_pem,
534            user.id,
535            &user.username,
536            user.role,
537            state.access_ttl_secs,
538        )
539        .map_err(|e| log::error!("auth mint token error: {e}"))
540        .ok()
541    })
542    .await
543    .ok()
544    .flatten()
545}
546
547async fn login(
548    State(state): State<AuthRouteState>,
549    ClientIp(from): ClientIp,
550    Json(req): Json<LoginRequest>,
551) -> Response {
552    let (user, access_token, refresh_token_id) =
553        match authenticate(&state, &req.username, &req.password, from).await {
554            Ok(session) => session,
555            Err(resp) => return *resp,
556        };
557
558    let cookies = state.session_cookies(&access_token, &refresh_token_id);
559
560    let resp = LoginResponse {
561        access_token,
562        // Also in the body: the CLI and other non-browser clients have no cookie
563        // jar and store this in config.local.toml.
564        refresh_token: refresh_token_id,
565        token_type: "Bearer".into(),
566        expires_in: state.access_ttl_secs,
567        user: UserInfo {
568            id: user.id,
569            username: user.username,
570            role: user.role.as_str().into(),
571        },
572    };
573
574    (StatusCode::OK, cookies, Json(resp)).into_response()
575}
576
577/// How long after its refresh a spent OAuth refresh token may come back as a
578/// retry rather than a replay.
579const REPLAY_GRACE_SECS: i64 = 30;
580
581/// Spend a refresh token for a new access token and a new refresh token. The
582/// error is the response to send. Shared by the JSON refresh, the web UI's
583/// session resume and OAuth.
584///
585/// A refresh is the account signing in from `from` as surely as a password
586/// is, so it marks that network as the account's (see
587/// `PasswordVerifier::signed_in`): a browser that keeps its session is how an
588/// account's own network stays known. `None` for an OAuth client, whose
589/// address is a service's, not the account's.
590pub(crate) fn rotate(
591    state: &AuthRouteState,
592    supplied: &str,
593    from: Option<IpAddr>,
594) -> Result<(String, String), Box<Response>> {
595    let db = match state.open_db() {
596        Ok(db) => db,
597        Err((status, msg)) => return Err(Box::new((status, msg).into_response())),
598    };
599
600    // Atomically consume (validate + revoke) the refresh token in a single
601    // statement to prevent TOCTOU races during token rotation.
602    let token = match auth_queries::consume_refresh_token(&db.conn, supplied) {
603        Ok(Some(t)) => t,
604        Ok(None) => {
605            match auth_queries::revoke_replayed_grant(&db.conn, supplied, REPLAY_GRACE_SECS) {
606                Ok(0) => {}
607                Ok(n) => log::warn!("a spent OAuth refresh token came back: revoked {n} tokens"),
608                Err(e) => log::error!("auth replay check error: {e}"),
609            }
610            return Err(Box::new(
611                (
612                    StatusCode::UNAUTHORIZED,
613                    Json(MessageResponse {
614                        message: "invalid or expired refresh token".into(),
615                    }),
616                )
617                    .into_response(),
618            ));
619        }
620        Err(e) => {
621            log::error!("auth refresh db error: {}", e);
622            return Err(Box::new(
623                (StatusCode::INTERNAL_SERVER_ERROR, "internal error").into_response(),
624            ));
625        }
626    };
627
628    let user = match auth_queries::get_user_by_id(&db.conn, token.user_id) {
629        Ok(Some(u)) => u,
630        Ok(None) => {
631            return Err(Box::new(
632                (
633                    StatusCode::UNAUTHORIZED,
634                    Json(MessageResponse {
635                        message: "user not found".into(),
636                    }),
637                )
638                    .into_response(),
639            ));
640        }
641        Err(e) => {
642            log::error!("auth refresh user lookup error: {}", e);
643            return Err(Box::new(
644                (StatusCode::INTERNAL_SERVER_ERROR, "internal error").into_response(),
645            ));
646        }
647    };
648
649    // A grant's tokens stay as narrow as the grant.
650    let access_token = match auth::mint_scoped_token(
651        &state.private_pem,
652        user.id,
653        &user.username,
654        user.role,
655        state.access_ttl_secs,
656        token.grant.as_ref().map(|_| auth::MCP_SCOPE),
657    ) {
658        Ok(t) => t,
659        Err(e) => {
660            log::error!("auth mint token error: {}", e);
661            return Err(Box::new(
662                (StatusCode::INTERNAL_SERVER_ERROR, "token error").into_response(),
663            ));
664        }
665    };
666
667    let new_refresh_id = match auth::random_token() {
668        Ok(t) => t,
669        Err(e) => {
670            log::error!("auth refresh token generation error: {}", e);
671            return Err(Box::new(
672                (StatusCode::INTERNAL_SERVER_ERROR, "token error").into_response(),
673            ));
674        }
675    };
676    let refresh_expires = auth::now_unix() as i64 + state.refresh_ttl_secs as i64;
677    if let Err(e) = auth_queries::store_grant_token(
678        &db.conn,
679        &new_refresh_id,
680        user.id,
681        refresh_expires,
682        token.grant.as_ref(),
683    ) {
684        log::error!("auth store refresh token error: {}", e);
685        return Err(Box::new(
686            (StatusCode::INTERNAL_SERVER_ERROR, "token error").into_response(),
687        ));
688    }
689
690    if let Some(from) = from {
691        state.users.signed_in(&user.username, from);
692    }
693    Ok((access_token, new_refresh_id))
694}
695
696async fn refresh(
697    State(state): State<AuthRouteState>,
698    ClientIp(from): ClientIp,
699    headers: axum::http::HeaderMap,
700    body: Option<Json<RefreshRequest>>,
701) -> Response {
702    let supplied = body.and_then(|Json(req)| req.refresh_token);
703    let Some(supplied) = refresh_token_from(supplied.as_deref(), &headers) else {
704        return (
705            StatusCode::UNAUTHORIZED,
706            Json(MessageResponse {
707                message: "missing refresh token".into(),
708            }),
709        )
710            .into_response();
711    };
712
713    let rotating = state.clone();
714    let rotated = tokio::task::spawn_blocking(move || rotate(&rotating, &supplied, Some(from)))
715        .await
716        .unwrap_or_else(|_| Err(Box::new(StatusCode::INTERNAL_SERVER_ERROR.into_response())));
717    let (access_token, new_refresh_id) = match rotated {
718        Ok(pair) => pair,
719        Err(resp) => return *resp,
720    };
721
722    let cookies = state.session_cookies(&access_token, &new_refresh_id);
723
724    let resp = RefreshResponse {
725        access_token,
726        refresh_token: new_refresh_id,
727        token_type: "Bearer".into(),
728        expires_in: state.access_ttl_secs,
729    };
730
731    (StatusCode::OK, cookies, Json(resp)).into_response()
732}
733
734async fn logout(
735    State(state): State<AuthRouteState>,
736    headers: axum::http::HeaderMap,
737    body: Option<Json<LogoutRequest>>,
738) -> Response {
739    let supplied = body.and_then(|Json(req)| req.refresh_token);
740    let token = refresh_token_from(supplied.as_deref(), &headers);
741    let revoking = state.clone();
742    let revoked = tokio::task::spawn_blocking(move || {
743        let db = revoking.open_db()?;
744        if let Some(token) = token {
745            let _ = auth_queries::revoke_refresh_token(&db.conn, &token);
746        }
747        Ok(())
748    })
749    .await
750    .unwrap_or_else(|_| {
751        Err((
752            StatusCode::INTERNAL_SERVER_ERROR,
753            "internal error".to_string(),
754        ))
755    });
756    if let Err((status, msg)) = revoked {
757        return (status, msg).into_response();
758    }
759
760    let cookies = state.cleared_cookies();
761
762    (
763        StatusCode::OK,
764        cookies,
765        Json(MessageResponse {
766            message: "logged out".into(),
767        }),
768    )
769        .into_response()
770}
771
772// ---------------------------------------------------------------------------
773// Tests
774// ---------------------------------------------------------------------------
775
776#[cfg(test)]
777mod tests {
778    use super::*;
779
780    #[test]
781    fn login_limiter_caps_a_single_ip() {
782        let limiter = RateLimiter::default();
783        let ip: IpAddr = "10.0.0.5".parse().unwrap();
784        for _ in 0..LOGIN_MAX_PER_WINDOW {
785            assert!(limiter.allow(ip));
786        }
787        assert!(!limiter.allow(ip));
788
789        // Other callers are unaffected.
790        assert!(limiter.allow("10.0.0.6".parse().unwrap()));
791    }
792
793    #[test]
794    fn an_ipv6_client_is_limited_by_its_64() {
795        let limiter = RateLimiter::default();
796        for n in 0..LOGIN_MAX_PER_WINDOW {
797            assert!(limiter.allow(format!("2001:db8:1:2::{n:x}").parse().unwrap()));
798        }
799        assert!(!limiter.allow("2001:db8:1:2:ffff::1".parse().unwrap()));
800        assert!(limiter.allow("2001:db8:1:3::1".parse().unwrap()));
801    }
802
803    #[test]
804    fn a_mapped_ipv4_address_is_the_ipv4_address() {
805        let r = request_from("[::ffff:198.51.100.4]", None);
806        assert_eq!(client_ip(&r), "198.51.100.4".parse::<IpAddr>().unwrap());
807        // A mapped private address is a proxy like any other.
808        let r = request_from("[::ffff:10.42.0.7]", Some("::ffff:203.0.113.9"));
809        assert_eq!(client_ip(&r), "203.0.113.9".parse::<IpAddr>().unwrap());
810    }
811
812    fn request_from(peer: &str, forwarded: Option<&str>) -> axum::extract::Request {
813        let mut request = axum::http::Request::new(axum::body::Body::empty());
814        request.extensions_mut().insert(ConnectInfo(
815            format!("{peer}:1234").parse::<SocketAddr>().unwrap(),
816        ));
817        if let Some(f) = forwarded {
818            request
819                .headers_mut()
820                .insert("x-forwarded-for", f.parse().unwrap());
821        }
822        request
823    }
824
825    #[test]
826    fn client_ip_believes_only_an_internal_proxy() {
827        // Behind the cluster's proxy: the entry it appended, not the client's.
828        let r = request_from("10.42.0.7", Some("6.6.6.6, 203.0.113.9"));
829        assert_eq!(client_ip(&r), "203.0.113.9".parse::<IpAddr>().unwrap());
830        // A public peer is the client, whatever it claims.
831        let r = request_from("198.51.100.4", Some("10.0.0.1"));
832        assert_eq!(client_ip(&r), "198.51.100.4".parse::<IpAddr>().unwrap());
833        // An internal peer with no header is itself.
834        let r = request_from("10.42.0.7", None);
835        assert_eq!(client_ip(&r), "10.42.0.7".parse::<IpAddr>().unwrap());
836        // No peer known: the header is the client's own claim.
837        let mut r = axum::http::Request::new(axum::body::Body::empty());
838        r.headers_mut()
839            .insert("x-forwarded-for", "203.0.113.9".parse().unwrap());
840        assert_eq!(client_ip(&r), IpAddr::V4(std::net::Ipv4Addr::UNSPECIFIED));
841    }
842
843    #[tokio::test]
844    async fn stalled_bodies_are_shed_then_timed_out() {
845        use tower::ServiceExt as _;
846        async fn read(_: axum::body::Bytes) -> StatusCode {
847            StatusCode::OK
848        }
849        // With its state, as `auth_router` does: that is what builds each
850        // route's layers once, rather than per request.
851        let app = auth_perimeter(
852            axum::Router::new().route("/auth/login", post(read)),
853            std::time::Duration::from_millis(200),
854        )
855        .with_state(());
856        let stalled = || {
857            axum::http::Request::post("/auth/login")
858                .body(axum::body::Body::from_stream(tokio_stream::pending::<
859                    Result<axum::body::Bytes, std::io::Error>,
860                >()))
861                .unwrap()
862        };
863        let held: Vec<_> = (0..2)
864            .map(|_| tokio::spawn(app.clone().oneshot(stalled())))
865            .collect();
866        tokio::time::sleep(std::time::Duration::from_millis(50)).await;
867        let shed = app.clone().oneshot(stalled()).await.unwrap();
868        assert_eq!(shed.status(), StatusCode::SERVICE_UNAVAILABLE);
869        for h in held {
870            assert_eq!(
871                h.await.unwrap().unwrap().status(),
872                StatusCode::REQUEST_TIMEOUT
873            );
874        }
875        let ok = axum::http::Request::post("/auth/login")
876            .body(axum::body::Body::empty())
877            .unwrap();
878        assert_eq!(app.oneshot(ok).await.unwrap().status(), StatusCode::OK);
879    }
880
881    #[test]
882    fn refresh_token_falls_back_to_the_cookie() {
883        let mut headers = axum::http::HeaderMap::new();
884        headers.insert(
885            COOKIE,
886            format!("a=1; {REFRESH_COOKIE}=from-cookie; b=2")
887                .parse()
888                .unwrap(),
889        );
890
891        assert_eq!(
892            refresh_token_from(None, &headers).as_deref(),
893            Some("from-cookie")
894        );
895        assert_eq!(
896            refresh_token_from(Some("from-body"), &headers).as_deref(),
897            Some("from-body")
898        );
899        assert_eq!(
900            refresh_token_from(None, &axum::http::HeaderMap::new()),
901            None
902        );
903    }
904
905    #[test]
906    fn cookies_are_lax_and_only_secure_when_tls_is_in_play() {
907        let state = |cookie_secure| AuthRouteState {
908            pool: Arc::new(Pool::new("/nonexistent".into())),
909            private_pem: Arc::new(Vec::new()),
910            public_pem: Arc::new(Vec::new()),
911            access_ttl_secs: 900,
912            refresh_ttl_secs: 60,
913            cookie_secure,
914            login_limiter: Arc::new(RateLimiter::default()),
915            users: Arc::new(crate::auth::password::PasswordVerifier::new(Arc::new(
916                Pool::new("/nonexistent/koan.db".into()),
917            ))),
918        };
919
920        let plain = state(false).access_cookie("tok");
921        assert!(plain.contains("SameSite=Lax"));
922        assert!(plain.contains("HttpOnly"));
923        assert!(!plain.contains("Secure"));
924
925        assert!(state(true).access_cookie("tok").contains("; Secure"));
926
927        // Every cookie reaches the browser, not only the last one set.
928        let resp = (StatusCode::OK, state(false).session_cookies("a", "r")).into_response();
929        let set: Vec<_> = resp.headers().get_all(SET_COOKIE).iter().collect();
930        assert_eq!(set.len(), 3);
931        assert_eq!(
932            (StatusCode::OK, state(false).cleared_cookies())
933                .into_response()
934                .headers()
935                .get_all(SET_COOKIE)
936                .iter()
937                .count(),
938            3
939        );
940
941        // The refresh cookie never rides along on an API call.
942        let refresh = state(false).refresh_cookie("tok");
943        assert!(refresh.contains("Path=/auth;"));
944        assert!(refresh.contains("HttpOnly"));
945    }
946}