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 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 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
236pub(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
243pub(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
264pub(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
287pub(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#[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
362pub 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
380const AUTH_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(10);
382
383fn 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
407pub(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 let _ = auth_queries::cleanup_expired_tokens(&db.conn);
518
519 Ok((user, access_token, refresh_token_id))
520}
521
522pub(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 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
577const REPLAY_GRACE_SECS: i64 = 30;
580
581pub(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 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 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#[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 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 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 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 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 let r = request_from("10.42.0.7", None);
835 assert_eq!(client_ip(&r), "10.42.0.7".parse::<IpAddr>().unwrap());
836 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 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 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 let refresh = state(false).refresh_cookie("tok");
943 assert!(refresh.contains("Path=/auth;"));
944 assert!(refresh.contains("HttpOnly"));
945 }
946}