1use 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
19pub(crate) const REFRESH_COOKIE: &str = "koan_refresh";
23const REFRESH_COOKIE_PATH: &str = "/auth";
24const STALE_REFRESH_COOKIE_PATH: &str = "/auth/refresh";
28
29const LOGIN_WINDOW_SECS: u64 = 60;
36const LOGIN_MAX_PER_WINDOW: u32 = 10;
37const 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 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
80pub(crate) fn client_ip(request: &axum::extract::Request) -> IpAddr {
94 address(request.extensions(), request.headers())
95}
96
97pub(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
129pub(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#[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 pub cookie_secure: bool,
161 pub login_limiter: Arc<RateLimiter>,
162 pub users: Arc<super::password::PasswordVerifier>,
164}
165
166impl AuthRouteState {
167 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 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 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
218pub(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
225pub(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
246pub(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
269pub(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#[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
344pub 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
362const AUTH_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(10);
364
365fn 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
389pub(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 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 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
534const REPLAY_GRACE_SECS: i64 = 30;
537
538pub(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 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 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#[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 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 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 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 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 let r = request_from("10.42.0.7", None);
792 assert_eq!(client_ip(&r), "10.42.0.7".parse::<IpAddr>().unwrap());
793 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 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 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 let refresh = state(false).refresh_cookie("tok");
900 assert!(refresh.contains("Path=/auth;"));
901 assert!(refresh.contains("HttpOnly"));
902 }
903}