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 LOGIN_TRACKED_IPS_MAX: usize = 4096;
39
40#[derive(Default)]
41pub struct LoginRateLimiter {
42 windows: Mutex<HashMap<IpAddr, (u64, u32)>>,
43}
44
45impl LoginRateLimiter {
46 fn allow(&self, ip: IpAddr) -> bool {
48 let now = auth::now_unix();
49 let mut windows = self.windows.lock().unwrap_or_else(|e| e.into_inner());
50
51 if windows.len() > LOGIN_TRACKED_IPS_MAX {
52 windows.retain(|_, (start, _)| now.saturating_sub(*start) < LOGIN_WINDOW_SECS);
53 }
54
55 let entry = windows.entry(ip).or_insert((now, 0));
56 if now.saturating_sub(entry.0) >= LOGIN_WINDOW_SECS {
57 *entry = (now, 0);
58 }
59 entry.1 += 1;
60 entry.1 <= LOGIN_MAX_PER_WINDOW
61 }
62}
63
64pub(crate) fn client_ip(request: &axum::extract::Request) -> IpAddr {
75 let Some(ConnectInfo(peer)) = request.extensions().get::<ConnectInfo<SocketAddr>>() else {
76 return IpAddr::V4(std::net::Ipv4Addr::UNSPECIFIED);
77 };
78 let peer = peer.ip();
79 if !is_internal(peer) {
80 return peer;
81 }
82 request
83 .headers()
84 .get_all("x-forwarded-for")
85 .iter()
86 .filter_map(|v| v.to_str().ok())
87 .flat_map(|v| v.split(','))
88 .filter_map(|ip| ip.trim().parse::<IpAddr>().ok())
89 .next_back()
90 .unwrap_or(peer)
91}
92
93fn is_internal(ip: IpAddr) -> bool {
94 match ip {
95 IpAddr::V4(v4) => v4.is_private() || v4.is_loopback(),
96 IpAddr::V6(v6) => v6.is_loopback() || (v6.segments()[0] & 0xfe00) == 0xfc00,
97 }
98}
99
100#[derive(Clone)]
105pub struct AuthRouteState {
106 pub pool: Arc<Pool>,
107 pub private_pem: Arc<Vec<u8>>,
108 pub public_pem: Arc<Vec<u8>>,
109 pub access_ttl_secs: u64,
110 pub refresh_ttl_secs: u64,
111 pub cookie_secure: bool,
115 pub login_limiter: Arc<LoginRateLimiter>,
116}
117
118impl AuthRouteState {
119 fn cookie(&self, name: &str, value: &str, path: &str, max_age: u64) -> String {
122 let secure = if self.cookie_secure { "; Secure" } else { "" };
123 format!("{name}={value}; HttpOnly; SameSite=Lax; Path={path}; Max-Age={max_age}{secure}")
124 }
125
126 fn access_cookie(&self, token: &str) -> String {
127 self.cookie("koan_access", token, "/", self.access_ttl_secs)
128 }
129
130 fn refresh_cookie(&self, token: &str) -> String {
131 self.cookie(
132 REFRESH_COOKIE,
133 token,
134 REFRESH_COOKIE_PATH,
135 self.refresh_ttl_secs,
136 )
137 }
138
139 fn stale_refresh_cookie(&self) -> String {
140 self.cookie(REFRESH_COOKIE, "", STALE_REFRESH_COOKIE_PATH, 0)
141 }
142
143 pub(crate) fn session_cookies(
146 &self,
147 access_token: &str,
148 refresh_token: &str,
149 ) -> AppendHeaders<[(axum::http::HeaderName, String); 3]> {
150 AppendHeaders([
151 (SET_COOKIE, self.access_cookie(access_token)),
152 (SET_COOKIE, self.refresh_cookie(refresh_token)),
153 (SET_COOKIE, self.stale_refresh_cookie()),
154 ])
155 }
156
157 pub(crate) fn cleared_cookies(&self) -> AppendHeaders<[(axum::http::HeaderName, String); 3]> {
159 AppendHeaders([
160 (SET_COOKIE, self.cookie("koan_access", "", "/", 0)),
161 (
162 SET_COOKIE,
163 self.cookie(REFRESH_COOKIE, "", REFRESH_COOKIE_PATH, 0),
164 ),
165 (SET_COOKIE, self.stale_refresh_cookie()),
166 ])
167 }
168}
169
170pub(crate) fn dummy_password_hash() -> &'static str {
173 static HASH: std::sync::OnceLock<String> = std::sync::OnceLock::new();
174 HASH.get_or_init(|| auth::hash_password("koan-dummy-password").unwrap_or_default())
175}
176
177pub(crate) fn refresh_token_from(
180 body: Option<&str>,
181 headers: &axum::http::HeaderMap,
182) -> Option<String> {
183 if let Some(t) = body.filter(|t| !t.is_empty()) {
184 return Some(t.to_owned());
185 }
186 headers
187 .get(COOKIE)
188 .and_then(|v| v.to_str().ok())
189 .and_then(|cookies| {
190 cookies.split(';').find_map(|c| {
191 c.trim()
192 .strip_prefix(&format!("{REFRESH_COOKIE}="))
193 .map(str::to_owned)
194 })
195 })
196}
197
198pub(crate) async fn login_rate_limit(
203 State(state): State<AuthRouteState>,
204 request: axum::extract::Request,
205 next: axum::middleware::Next,
206) -> Response {
207 let ip = client_ip(&request);
208
209 if !state.login_limiter.allow(ip) {
210 return (
211 StatusCode::TOO_MANY_REQUESTS,
212 Json(MessageResponse {
213 message: "too many login attempts".into(),
214 }),
215 )
216 .into_response();
217 }
218 next.run(request).await
219}
220
221impl AuthRouteState {
222 fn open_db(&self) -> Result<Handle<'_>, (StatusCode, String)> {
223 self.pool.get().map_err(|e| {
224 log::error!("auth db open error: {}", e);
225 (
226 StatusCode::INTERNAL_SERVER_ERROR,
227 "internal error".to_string(),
228 )
229 })
230 }
231}
232
233#[derive(Deserialize)]
238pub struct LoginRequest {
239 pub username: String,
240 pub password: String,
241}
242
243#[derive(Serialize)]
244pub struct LoginResponse {
245 pub access_token: String,
246 pub refresh_token: String,
247 pub token_type: String,
248 pub expires_in: u64,
249 pub user: UserInfo,
250}
251
252#[derive(Serialize)]
253pub struct UserInfo {
254 pub id: i64,
255 pub username: String,
256 pub role: String,
257}
258
259#[derive(Deserialize, Default)]
260#[serde(default)]
261pub struct RefreshRequest {
262 pub refresh_token: Option<String>,
263}
264
265#[derive(Serialize)]
266pub struct RefreshResponse {
267 pub access_token: String,
268 pub refresh_token: String,
269 pub token_type: String,
270 pub expires_in: u64,
271}
272
273#[derive(Deserialize, Default)]
274#[serde(default)]
275pub struct LogoutRequest {
276 pub refresh_token: Option<String>,
277}
278
279#[derive(Serialize)]
280pub struct MessageResponse {
281 pub message: String,
282}
283
284pub fn auth_router(state: AuthRouteState) -> axum::Router {
289 let router = axum::Router::new()
290 .route(
291 "/auth/login",
292 post(login).layer(axum::middleware::from_fn_with_state(
293 state.clone(),
294 login_rate_limit,
295 )),
296 )
297 .route("/auth/refresh", post(refresh))
298 .route("/auth/logout", post(logout));
299 auth_perimeter(router, AUTH_TIMEOUT).with_state(state)
300}
301
302const AUTH_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(10);
304
305fn auth_perimeter<S>(router: axum::Router<S>, timeout: std::time::Duration) -> axum::Router<S>
311where
312 S: Clone + Send + Sync + 'static,
313{
314 router
315 .layer(tower_http::timeout::TimeoutLayer::with_status_code(
316 StatusCode::REQUEST_TIMEOUT,
317 timeout,
318 ))
319 .layer(
320 tower::ServiceBuilder::new()
321 .layer(axum::error_handling::HandleErrorLayer::new(
322 |_: tower::BoxError| async { (StatusCode::SERVICE_UNAVAILABLE, "busy") },
323 ))
324 .load_shed()
325 .concurrency_limit(2),
326 )
327}
328
329pub(crate) async fn authenticate(
341 state: &AuthRouteState,
342 username: &str,
343 password: &str,
344) -> Result<(auth_queries::UserRow, String, String), Box<Response>> {
345 let state = state.clone();
346 let (username, password) = (username.to_owned(), password.to_owned());
347 tokio::task::spawn_blocking(move || authenticate_blocking(&state, &username, &password))
348 .await
349 .unwrap_or_else(|_| Err(Box::new(StatusCode::INTERNAL_SERVER_ERROR.into_response())))
350}
351
352fn authenticate_blocking(
353 state: &AuthRouteState,
354 username: &str,
355 password: &str,
356) -> Result<(auth_queries::UserRow, String, String), Box<Response>> {
357 let db = match state.open_db() {
358 Ok(db) => db,
359 Err((status, msg)) => return Err(Box::new((status, msg).into_response())),
360 };
361
362 let user = match auth_queries::get_user_by_username(&db.conn, username) {
363 Ok(Some(u)) => u,
364 Ok(None) => {
365 let _ = auth::verify_password(password, dummy_password_hash());
368 return Err(Box::new(
369 (
370 StatusCode::UNAUTHORIZED,
371 Json(MessageResponse {
372 message: "invalid username or password".into(),
373 }),
374 )
375 .into_response(),
376 ));
377 }
378 Err(e) => {
379 log::error!("auth login db error: {}", e);
380 return Err(Box::new(
381 (StatusCode::INTERNAL_SERVER_ERROR, "internal error").into_response(),
382 ));
383 }
384 };
385
386 if auth::verify_password(password, &user.password_hash).is_err() {
387 return Err(Box::new(
388 (
389 StatusCode::UNAUTHORIZED,
390 Json(MessageResponse {
391 message: "invalid username or password".into(),
392 }),
393 )
394 .into_response(),
395 ));
396 }
397
398 if let Err(e) = auth_queries::remember_password(&db.conn, username, password) {
400 log::warn!("could not seal the password for Subsonic token auth: {e}");
401 }
402
403 let access_token = match auth::mint_access_token(
404 &state.private_pem,
405 user.id,
406 &user.username,
407 user.role,
408 state.access_ttl_secs,
409 ) {
410 Ok(t) => t,
411 Err(e) => {
412 log::error!("auth mint token error: {}", e);
413 return Err(Box::new(
414 (StatusCode::INTERNAL_SERVER_ERROR, "token error").into_response(),
415 ));
416 }
417 };
418
419 let refresh_token_id = match auth::random_token() {
420 Ok(t) => t,
421 Err(e) => {
422 log::error!("auth refresh token generation error: {}", e);
423 return Err(Box::new(
424 (StatusCode::INTERNAL_SERVER_ERROR, "token error").into_response(),
425 ));
426 }
427 };
428 let refresh_expires = auth::now_unix() as i64 + state.refresh_ttl_secs as i64;
429 if let Err(e) =
430 auth_queries::store_refresh_token(&db.conn, &refresh_token_id, user.id, refresh_expires)
431 {
432 log::error!("auth store refresh token error: {}", e);
433 return Err(Box::new(
434 (StatusCode::INTERNAL_SERVER_ERROR, "token error").into_response(),
435 ));
436 }
437
438 let _ = auth_queries::cleanup_expired_tokens(&db.conn);
440
441 Ok((user, access_token, refresh_token_id))
442}
443
444async fn login(State(state): State<AuthRouteState>, Json(req): Json<LoginRequest>) -> Response {
445 let (user, access_token, refresh_token_id) =
446 match authenticate(&state, &req.username, &req.password).await {
447 Ok(session) => session,
448 Err(resp) => return *resp,
449 };
450
451 let cookies = state.session_cookies(&access_token, &refresh_token_id);
452
453 let resp = LoginResponse {
454 access_token,
455 refresh_token: refresh_token_id,
458 token_type: "Bearer".into(),
459 expires_in: state.access_ttl_secs,
460 user: UserInfo {
461 id: user.id,
462 username: user.username,
463 role: user.role.as_str().into(),
464 },
465 };
466
467 (StatusCode::OK, cookies, Json(resp)).into_response()
468}
469
470pub(crate) fn rotate(
474 state: &AuthRouteState,
475 supplied: &str,
476) -> Result<(String, String), Box<Response>> {
477 let db = match state.open_db() {
478 Ok(db) => db,
479 Err((status, msg)) => return Err(Box::new((status, msg).into_response())),
480 };
481
482 let token = match auth_queries::consume_refresh_token(&db.conn, supplied) {
485 Ok(Some(t)) => t,
486 Ok(None) => {
487 return Err(Box::new(
488 (
489 StatusCode::UNAUTHORIZED,
490 Json(MessageResponse {
491 message: "invalid or expired refresh token".into(),
492 }),
493 )
494 .into_response(),
495 ));
496 }
497 Err(e) => {
498 log::error!("auth refresh db error: {}", e);
499 return Err(Box::new(
500 (StatusCode::INTERNAL_SERVER_ERROR, "internal error").into_response(),
501 ));
502 }
503 };
504
505 let user = match auth_queries::get_user_by_id(&db.conn, token.user_id) {
506 Ok(Some(u)) => u,
507 Ok(None) => {
508 return Err(Box::new(
509 (
510 StatusCode::UNAUTHORIZED,
511 Json(MessageResponse {
512 message: "user not found".into(),
513 }),
514 )
515 .into_response(),
516 ));
517 }
518 Err(e) => {
519 log::error!("auth refresh user lookup error: {}", e);
520 return Err(Box::new(
521 (StatusCode::INTERNAL_SERVER_ERROR, "internal error").into_response(),
522 ));
523 }
524 };
525
526 let access_token = match auth::mint_access_token(
527 &state.private_pem,
528 user.id,
529 &user.username,
530 user.role,
531 state.access_ttl_secs,
532 ) {
533 Ok(t) => t,
534 Err(e) => {
535 log::error!("auth mint token error: {}", e);
536 return Err(Box::new(
537 (StatusCode::INTERNAL_SERVER_ERROR, "token error").into_response(),
538 ));
539 }
540 };
541
542 let new_refresh_id = match auth::random_token() {
543 Ok(t) => t,
544 Err(e) => {
545 log::error!("auth refresh token generation error: {}", e);
546 return Err(Box::new(
547 (StatusCode::INTERNAL_SERVER_ERROR, "token error").into_response(),
548 ));
549 }
550 };
551 let refresh_expires = auth::now_unix() as i64 + state.refresh_ttl_secs as i64;
552 if let Err(e) =
553 auth_queries::store_refresh_token(&db.conn, &new_refresh_id, user.id, refresh_expires)
554 {
555 log::error!("auth store refresh token error: {}", e);
556 return Err(Box::new(
557 (StatusCode::INTERNAL_SERVER_ERROR, "token error").into_response(),
558 ));
559 }
560
561 Ok((access_token, new_refresh_id))
562}
563
564async fn refresh(
565 State(state): State<AuthRouteState>,
566 headers: axum::http::HeaderMap,
567 body: Option<Json<RefreshRequest>>,
568) -> Response {
569 let supplied = body.and_then(|Json(req)| req.refresh_token);
570 let Some(supplied) = refresh_token_from(supplied.as_deref(), &headers) else {
571 return (
572 StatusCode::UNAUTHORIZED,
573 Json(MessageResponse {
574 message: "missing refresh token".into(),
575 }),
576 )
577 .into_response();
578 };
579
580 let rotating = state.clone();
581 let rotated = tokio::task::spawn_blocking(move || rotate(&rotating, &supplied))
582 .await
583 .unwrap_or_else(|_| Err(Box::new(StatusCode::INTERNAL_SERVER_ERROR.into_response())));
584 let (access_token, new_refresh_id) = match rotated {
585 Ok(pair) => pair,
586 Err(resp) => return *resp,
587 };
588
589 let cookies = state.session_cookies(&access_token, &new_refresh_id);
590
591 let resp = RefreshResponse {
592 access_token,
593 refresh_token: new_refresh_id,
594 token_type: "Bearer".into(),
595 expires_in: state.access_ttl_secs,
596 };
597
598 (StatusCode::OK, cookies, Json(resp)).into_response()
599}
600
601async fn logout(
602 State(state): State<AuthRouteState>,
603 headers: axum::http::HeaderMap,
604 body: Option<Json<LogoutRequest>>,
605) -> Response {
606 let supplied = body.and_then(|Json(req)| req.refresh_token);
607 let token = refresh_token_from(supplied.as_deref(), &headers);
608 let revoking = state.clone();
609 let revoked = tokio::task::spawn_blocking(move || {
610 let db = revoking.open_db()?;
611 if let Some(token) = token {
612 let _ = auth_queries::revoke_refresh_token(&db.conn, &token);
613 }
614 Ok(())
615 })
616 .await
617 .unwrap_or_else(|_| {
618 Err((
619 StatusCode::INTERNAL_SERVER_ERROR,
620 "internal error".to_string(),
621 ))
622 });
623 if let Err((status, msg)) = revoked {
624 return (status, msg).into_response();
625 }
626
627 let cookies = state.cleared_cookies();
628
629 (
630 StatusCode::OK,
631 cookies,
632 Json(MessageResponse {
633 message: "logged out".into(),
634 }),
635 )
636 .into_response()
637}
638
639#[cfg(test)]
644mod tests {
645 use super::*;
646
647 #[test]
648 fn login_limiter_caps_a_single_ip() {
649 let limiter = LoginRateLimiter::default();
650 let ip: IpAddr = "10.0.0.5".parse().unwrap();
651 for _ in 0..LOGIN_MAX_PER_WINDOW {
652 assert!(limiter.allow(ip));
653 }
654 assert!(!limiter.allow(ip));
655
656 assert!(limiter.allow("10.0.0.6".parse().unwrap()));
658 }
659
660 fn request_from(peer: &str, forwarded: Option<&str>) -> axum::extract::Request {
661 let mut request = axum::http::Request::new(axum::body::Body::empty());
662 request.extensions_mut().insert(ConnectInfo(
663 format!("{peer}:1234").parse::<SocketAddr>().unwrap(),
664 ));
665 if let Some(f) = forwarded {
666 request
667 .headers_mut()
668 .insert("x-forwarded-for", f.parse().unwrap());
669 }
670 request
671 }
672
673 #[test]
674 fn client_ip_believes_only_an_internal_proxy() {
675 let r = request_from("10.42.0.7", Some("6.6.6.6, 203.0.113.9"));
677 assert_eq!(client_ip(&r), "203.0.113.9".parse::<IpAddr>().unwrap());
678 let r = request_from("198.51.100.4", Some("10.0.0.1"));
680 assert_eq!(client_ip(&r), "198.51.100.4".parse::<IpAddr>().unwrap());
681 let r = request_from("10.42.0.7", None);
683 assert_eq!(client_ip(&r), "10.42.0.7".parse::<IpAddr>().unwrap());
684 let mut r = axum::http::Request::new(axum::body::Body::empty());
686 r.headers_mut()
687 .insert("x-forwarded-for", "203.0.113.9".parse().unwrap());
688 assert_eq!(client_ip(&r), IpAddr::V4(std::net::Ipv4Addr::UNSPECIFIED));
689 }
690
691 #[tokio::test]
692 async fn stalled_bodies_are_shed_then_timed_out() {
693 use tower::ServiceExt as _;
694 async fn read(_: axum::body::Bytes) -> StatusCode {
695 StatusCode::OK
696 }
697 let app = auth_perimeter(
700 axum::Router::new().route("/auth/login", post(read)),
701 std::time::Duration::from_millis(200),
702 )
703 .with_state(());
704 let stalled = || {
705 axum::http::Request::post("/auth/login")
706 .body(axum::body::Body::from_stream(tokio_stream::pending::<
707 Result<axum::body::Bytes, std::io::Error>,
708 >()))
709 .unwrap()
710 };
711 let held: Vec<_> = (0..2)
712 .map(|_| tokio::spawn(app.clone().oneshot(stalled())))
713 .collect();
714 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
715 let shed = app.clone().oneshot(stalled()).await.unwrap();
716 assert_eq!(shed.status(), StatusCode::SERVICE_UNAVAILABLE);
717 for h in held {
718 assert_eq!(
719 h.await.unwrap().unwrap().status(),
720 StatusCode::REQUEST_TIMEOUT
721 );
722 }
723 let ok = axum::http::Request::post("/auth/login")
724 .body(axum::body::Body::empty())
725 .unwrap();
726 assert_eq!(app.oneshot(ok).await.unwrap().status(), StatusCode::OK);
727 }
728
729 #[test]
730 fn refresh_token_falls_back_to_the_cookie() {
731 let mut headers = axum::http::HeaderMap::new();
732 headers.insert(
733 COOKIE,
734 format!("a=1; {REFRESH_COOKIE}=from-cookie; b=2")
735 .parse()
736 .unwrap(),
737 );
738
739 assert_eq!(
740 refresh_token_from(None, &headers).as_deref(),
741 Some("from-cookie")
742 );
743 assert_eq!(
744 refresh_token_from(Some("from-body"), &headers).as_deref(),
745 Some("from-body")
746 );
747 assert_eq!(
748 refresh_token_from(None, &axum::http::HeaderMap::new()),
749 None
750 );
751 }
752
753 #[test]
754 fn cookies_are_lax_and_only_secure_when_tls_is_in_play() {
755 let state = |cookie_secure| AuthRouteState {
756 pool: Arc::new(Pool::new("/nonexistent".into())),
757 private_pem: Arc::new(Vec::new()),
758 public_pem: Arc::new(Vec::new()),
759 access_ttl_secs: 900,
760 refresh_ttl_secs: 60,
761 cookie_secure,
762 login_limiter: Arc::new(LoginRateLimiter::default()),
763 };
764
765 let plain = state(false).access_cookie("tok");
766 assert!(plain.contains("SameSite=Lax"));
767 assert!(plain.contains("HttpOnly"));
768 assert!(!plain.contains("Secure"));
769
770 assert!(state(true).access_cookie("tok").contains("; Secure"));
771
772 let resp = (StatusCode::OK, state(false).session_cookies("a", "r")).into_response();
774 let set: Vec<_> = resp.headers().get_all(SET_COOKIE).iter().collect();
775 assert_eq!(set.len(), 3);
776 assert_eq!(
777 (StatusCode::OK, state(false).cleared_cookies())
778 .into_response()
779 .headers()
780 .get_all(SET_COOKIE)
781 .iter()
782 .count(),
783 3
784 );
785
786 let refresh = state(false).refresh_cookie("tok");
788 assert!(refresh.contains("Path=/auth;"));
789 assert!(refresh.contains("HttpOnly"));
790 }
791}