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 {
73 let peer = request
74 .extensions()
75 .get::<ConnectInfo<SocketAddr>>()
76 .map(|ConnectInfo(addr)| addr.ip())
77 .unwrap_or(IpAddr::V4(std::net::Ipv4Addr::UNSPECIFIED));
78 if !is_internal(peer) {
79 return peer;
80 }
81 request
82 .headers()
83 .get_all("x-forwarded-for")
84 .iter()
85 .filter_map(|v| v.to_str().ok())
86 .flat_map(|v| v.split(','))
87 .filter_map(|ip| ip.trim().parse::<IpAddr>().ok())
88 .next_back()
89 .unwrap_or(peer)
90}
91
92fn is_internal(ip: IpAddr) -> bool {
93 match ip {
94 IpAddr::V4(v4) => v4.is_private() || v4.is_loopback() || v4.is_unspecified(),
95 IpAddr::V6(v6) => {
96 v6.is_loopback() || v6.is_unspecified() || (v6.segments()[0] & 0xfe00) == 0xfc00
97 }
98 }
99}
100
101#[derive(Clone)]
106pub struct AuthRouteState {
107 pub pool: Arc<Pool>,
108 pub private_pem: Arc<Vec<u8>>,
109 pub public_pem: Arc<Vec<u8>>,
110 pub access_ttl_secs: u64,
111 pub refresh_ttl_secs: u64,
112 pub cookie_secure: bool,
116 pub login_limiter: Arc<LoginRateLimiter>,
117}
118
119impl AuthRouteState {
120 fn cookie(&self, name: &str, value: &str, path: &str, max_age: u64) -> String {
123 let secure = if self.cookie_secure { "; Secure" } else { "" };
124 format!("{name}={value}; HttpOnly; SameSite=Lax; Path={path}; Max-Age={max_age}{secure}")
125 }
126
127 fn access_cookie(&self, token: &str) -> String {
128 self.cookie("koan_access", token, "/", self.access_ttl_secs)
129 }
130
131 fn refresh_cookie(&self, token: &str) -> String {
132 self.cookie(
133 REFRESH_COOKIE,
134 token,
135 REFRESH_COOKIE_PATH,
136 self.refresh_ttl_secs,
137 )
138 }
139
140 fn stale_refresh_cookie(&self) -> String {
141 self.cookie(REFRESH_COOKIE, "", STALE_REFRESH_COOKIE_PATH, 0)
142 }
143
144 pub(crate) fn session_cookies(
147 &self,
148 access_token: &str,
149 refresh_token: &str,
150 ) -> AppendHeaders<[(axum::http::HeaderName, String); 3]> {
151 AppendHeaders([
152 (SET_COOKIE, self.access_cookie(access_token)),
153 (SET_COOKIE, self.refresh_cookie(refresh_token)),
154 (SET_COOKIE, self.stale_refresh_cookie()),
155 ])
156 }
157
158 pub(crate) fn cleared_cookies(&self) -> AppendHeaders<[(axum::http::HeaderName, String); 3]> {
160 AppendHeaders([
161 (SET_COOKIE, self.cookie("koan_access", "", "/", 0)),
162 (
163 SET_COOKIE,
164 self.cookie(REFRESH_COOKIE, "", REFRESH_COOKIE_PATH, 0),
165 ),
166 (SET_COOKIE, self.stale_refresh_cookie()),
167 ])
168 }
169}
170
171pub(crate) fn dummy_password_hash() -> &'static str {
174 static HASH: std::sync::OnceLock<String> = std::sync::OnceLock::new();
175 HASH.get_or_init(|| auth::hash_password("koan-dummy-password").unwrap_or_default())
176}
177
178pub(crate) fn refresh_token_from(
181 body: Option<&str>,
182 headers: &axum::http::HeaderMap,
183) -> Option<String> {
184 if let Some(t) = body.filter(|t| !t.is_empty()) {
185 return Some(t.to_owned());
186 }
187 headers
188 .get(COOKIE)
189 .and_then(|v| v.to_str().ok())
190 .and_then(|cookies| {
191 cookies.split(';').find_map(|c| {
192 c.trim()
193 .strip_prefix(&format!("{REFRESH_COOKIE}="))
194 .map(str::to_owned)
195 })
196 })
197}
198
199pub(crate) async fn login_rate_limit(
204 State(state): State<AuthRouteState>,
205 request: axum::extract::Request,
206 next: axum::middleware::Next,
207) -> Response {
208 let ip = client_ip(&request);
209
210 if !state.login_limiter.allow(ip) {
211 return (
212 StatusCode::TOO_MANY_REQUESTS,
213 Json(MessageResponse {
214 message: "too many login attempts".into(),
215 }),
216 )
217 .into_response();
218 }
219 next.run(request).await
220}
221
222impl AuthRouteState {
223 fn open_db(&self) -> Result<Handle<'_>, (StatusCode, String)> {
224 self.pool.get().map_err(|e| {
225 log::error!("auth db open error: {}", e);
226 (
227 StatusCode::INTERNAL_SERVER_ERROR,
228 "internal error".to_string(),
229 )
230 })
231 }
232}
233
234#[derive(Deserialize)]
239pub struct LoginRequest {
240 pub username: String,
241 pub password: String,
242}
243
244#[derive(Serialize)]
245pub struct LoginResponse {
246 pub access_token: String,
247 pub refresh_token: String,
248 pub token_type: String,
249 pub expires_in: u64,
250 pub user: UserInfo,
251}
252
253#[derive(Serialize)]
254pub struct UserInfo {
255 pub id: i64,
256 pub username: String,
257 pub role: String,
258}
259
260#[derive(Deserialize, Default)]
261#[serde(default)]
262pub struct RefreshRequest {
263 pub refresh_token: Option<String>,
264}
265
266#[derive(Serialize)]
267pub struct RefreshResponse {
268 pub access_token: String,
269 pub refresh_token: String,
270 pub token_type: String,
271 pub expires_in: u64,
272}
273
274#[derive(Deserialize, Default)]
275#[serde(default)]
276pub struct LogoutRequest {
277 pub refresh_token: Option<String>,
278}
279
280#[derive(Serialize)]
281pub struct MessageResponse {
282 pub message: String,
283}
284
285pub fn auth_router(state: AuthRouteState) -> axum::Router {
290 axum::Router::new()
291 .route(
292 "/auth/login",
293 post(login).layer(axum::middleware::from_fn_with_state(
294 state.clone(),
295 login_rate_limit,
296 )),
297 )
298 .route("/auth/refresh", post(refresh))
299 .route("/auth/logout", post(logout))
300 .layer(tower::limit::ConcurrencyLimitLayer::new(2))
304 .with_state(state)
305}
306
307pub(crate) async fn authenticate(
319 state: &AuthRouteState,
320 username: &str,
321 password: &str,
322) -> Result<(auth_queries::UserRow, String, String), Box<Response>> {
323 let state = state.clone();
324 let (username, password) = (username.to_owned(), password.to_owned());
325 tokio::task::spawn_blocking(move || authenticate_blocking(&state, &username, &password))
326 .await
327 .unwrap_or_else(|_| Err(Box::new(StatusCode::INTERNAL_SERVER_ERROR.into_response())))
328}
329
330fn authenticate_blocking(
331 state: &AuthRouteState,
332 username: &str,
333 password: &str,
334) -> Result<(auth_queries::UserRow, String, String), Box<Response>> {
335 let db = match state.open_db() {
336 Ok(db) => db,
337 Err((status, msg)) => return Err(Box::new((status, msg).into_response())),
338 };
339
340 let user = match auth_queries::get_user_by_username(&db.conn, username) {
342 Ok(Some(u)) => u,
343 Ok(None) => {
344 let _ = auth::verify_password(password, dummy_password_hash());
347 return Err(Box::new(
348 (
349 StatusCode::UNAUTHORIZED,
350 Json(MessageResponse {
351 message: "invalid username or password".into(),
352 }),
353 )
354 .into_response(),
355 ));
356 }
357 Err(e) => {
358 log::error!("auth login db error: {}", e);
359 return Err(Box::new(
360 (StatusCode::INTERNAL_SERVER_ERROR, "internal error").into_response(),
361 ));
362 }
363 };
364
365 if auth::verify_password(password, &user.password_hash).is_err() {
366 return Err(Box::new(
367 (
368 StatusCode::UNAUTHORIZED,
369 Json(MessageResponse {
370 message: "invalid username or password".into(),
371 }),
372 )
373 .into_response(),
374 ));
375 }
376
377 if let Err(e) = auth_queries::remember_password(&db.conn, username, password) {
379 log::warn!("could not seal the password for Subsonic token auth: {e}");
380 }
381
382 let access_token = match auth::mint_access_token(
384 &state.private_pem,
385 user.id,
386 &user.username,
387 user.role,
388 state.access_ttl_secs,
389 ) {
390 Ok(t) => t,
391 Err(e) => {
392 log::error!("auth mint token error: {}", e);
393 return Err(Box::new(
394 (StatusCode::INTERNAL_SERVER_ERROR, "token error").into_response(),
395 ));
396 }
397 };
398
399 let refresh_token_id = match auth::random_token() {
401 Ok(t) => t,
402 Err(e) => {
403 log::error!("auth refresh token generation error: {}", e);
404 return Err(Box::new(
405 (StatusCode::INTERNAL_SERVER_ERROR, "token error").into_response(),
406 ));
407 }
408 };
409 let refresh_expires = auth::now_unix() as i64 + state.refresh_ttl_secs as i64;
410 if let Err(e) =
411 auth_queries::store_refresh_token(&db.conn, &refresh_token_id, user.id, refresh_expires)
412 {
413 log::error!("auth store refresh token error: {}", e);
414 return Err(Box::new(
415 (StatusCode::INTERNAL_SERVER_ERROR, "token error").into_response(),
416 ));
417 }
418
419 let _ = auth_queries::cleanup_expired_tokens(&db.conn);
421
422 Ok((user, access_token, refresh_token_id))
423}
424
425async fn login(State(state): State<AuthRouteState>, Json(req): Json<LoginRequest>) -> Response {
426 let (user, access_token, refresh_token_id) =
427 match authenticate(&state, &req.username, &req.password).await {
428 Ok(session) => session,
429 Err(resp) => return *resp,
430 };
431
432 let cookies = state.session_cookies(&access_token, &refresh_token_id);
433
434 let resp = LoginResponse {
435 access_token,
436 refresh_token: refresh_token_id,
439 token_type: "Bearer".into(),
440 expires_in: state.access_ttl_secs,
441 user: UserInfo {
442 id: user.id,
443 username: user.username,
444 role: user.role.as_str().into(),
445 },
446 };
447
448 (StatusCode::OK, cookies, Json(resp)).into_response()
449}
450
451pub(crate) fn rotate(
455 state: &AuthRouteState,
456 supplied: &str,
457) -> Result<(String, String), Box<Response>> {
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 token = match auth_queries::consume_refresh_token(&db.conn, supplied) {
466 Ok(Some(t)) => t,
467 Ok(None) => {
468 return Err(Box::new(
469 (
470 StatusCode::UNAUTHORIZED,
471 Json(MessageResponse {
472 message: "invalid or expired refresh token".into(),
473 }),
474 )
475 .into_response(),
476 ));
477 }
478 Err(e) => {
479 log::error!("auth refresh db error: {}", e);
480 return Err(Box::new(
481 (StatusCode::INTERNAL_SERVER_ERROR, "internal error").into_response(),
482 ));
483 }
484 };
485
486 let user = match auth_queries::get_user_by_id(&db.conn, token.user_id) {
488 Ok(Some(u)) => u,
489 Ok(None) => {
490 return Err(Box::new(
491 (
492 StatusCode::UNAUTHORIZED,
493 Json(MessageResponse {
494 message: "user not found".into(),
495 }),
496 )
497 .into_response(),
498 ));
499 }
500 Err(e) => {
501 log::error!("auth refresh user lookup error: {}", e);
502 return Err(Box::new(
503 (StatusCode::INTERNAL_SERVER_ERROR, "internal error").into_response(),
504 ));
505 }
506 };
507
508 let access_token = match auth::mint_access_token(
510 &state.private_pem,
511 user.id,
512 &user.username,
513 user.role,
514 state.access_ttl_secs,
515 ) {
516 Ok(t) => t,
517 Err(e) => {
518 log::error!("auth mint token error: {}", e);
519 return Err(Box::new(
520 (StatusCode::INTERNAL_SERVER_ERROR, "token error").into_response(),
521 ));
522 }
523 };
524
525 let new_refresh_id = match auth::random_token() {
527 Ok(t) => t,
528 Err(e) => {
529 log::error!("auth refresh token generation error: {}", e);
530 return Err(Box::new(
531 (StatusCode::INTERNAL_SERVER_ERROR, "token error").into_response(),
532 ));
533 }
534 };
535 let refresh_expires = auth::now_unix() as i64 + state.refresh_ttl_secs as i64;
536 if let Err(e) =
537 auth_queries::store_refresh_token(&db.conn, &new_refresh_id, user.id, refresh_expires)
538 {
539 log::error!("auth store refresh token error: {}", e);
540 return Err(Box::new(
541 (StatusCode::INTERNAL_SERVER_ERROR, "token error").into_response(),
542 ));
543 }
544
545 Ok((access_token, new_refresh_id))
546}
547
548async fn refresh(
549 State(state): State<AuthRouteState>,
550 headers: axum::http::HeaderMap,
551 body: Option<Json<RefreshRequest>>,
552) -> Response {
553 let supplied = body.and_then(|Json(req)| req.refresh_token);
554 let Some(supplied) = refresh_token_from(supplied.as_deref(), &headers) else {
555 return (
556 StatusCode::UNAUTHORIZED,
557 Json(MessageResponse {
558 message: "missing refresh token".into(),
559 }),
560 )
561 .into_response();
562 };
563
564 let rotating = state.clone();
565 let rotated = tokio::task::spawn_blocking(move || rotate(&rotating, &supplied))
566 .await
567 .unwrap_or_else(|_| Err(Box::new(StatusCode::INTERNAL_SERVER_ERROR.into_response())));
568 let (access_token, new_refresh_id) = match rotated {
569 Ok(pair) => pair,
570 Err(resp) => return *resp,
571 };
572
573 let cookies = state.session_cookies(&access_token, &new_refresh_id);
574
575 let resp = RefreshResponse {
576 access_token,
577 refresh_token: new_refresh_id,
578 token_type: "Bearer".into(),
579 expires_in: state.access_ttl_secs,
580 };
581
582 (StatusCode::OK, cookies, Json(resp)).into_response()
583}
584
585async fn logout(
586 State(state): State<AuthRouteState>,
587 headers: axum::http::HeaderMap,
588 body: Option<Json<LogoutRequest>>,
589) -> Response {
590 let supplied = body.and_then(|Json(req)| req.refresh_token);
591 let token = refresh_token_from(supplied.as_deref(), &headers);
592 let revoking = state.clone();
593 let revoked = tokio::task::spawn_blocking(move || {
594 let db = revoking.open_db()?;
595 if let Some(token) = token {
596 let _ = auth_queries::revoke_refresh_token(&db.conn, &token);
597 }
598 Ok(())
599 })
600 .await
601 .unwrap_or_else(|_| {
602 Err((
603 StatusCode::INTERNAL_SERVER_ERROR,
604 "internal error".to_string(),
605 ))
606 });
607 if let Err((status, msg)) = revoked {
608 return (status, msg).into_response();
609 }
610
611 let cookies = state.cleared_cookies();
612
613 (
614 StatusCode::OK,
615 cookies,
616 Json(MessageResponse {
617 message: "logged out".into(),
618 }),
619 )
620 .into_response()
621}
622
623#[cfg(test)]
628mod tests {
629 use super::*;
630
631 #[test]
632 fn login_limiter_caps_a_single_ip() {
633 let limiter = LoginRateLimiter::default();
634 let ip: IpAddr = "10.0.0.5".parse().unwrap();
635 for _ in 0..LOGIN_MAX_PER_WINDOW {
636 assert!(limiter.allow(ip));
637 }
638 assert!(!limiter.allow(ip));
639
640 assert!(limiter.allow("10.0.0.6".parse().unwrap()));
642 }
643
644 fn request_from(peer: &str, forwarded: Option<&str>) -> axum::extract::Request {
645 let mut request = axum::http::Request::new(axum::body::Body::empty());
646 request.extensions_mut().insert(ConnectInfo(
647 format!("{peer}:1234").parse::<SocketAddr>().unwrap(),
648 ));
649 if let Some(f) = forwarded {
650 request
651 .headers_mut()
652 .insert("x-forwarded-for", f.parse().unwrap());
653 }
654 request
655 }
656
657 #[test]
658 fn client_ip_believes_only_an_internal_proxy() {
659 let r = request_from("10.42.0.7", Some("6.6.6.6, 203.0.113.9"));
661 assert_eq!(client_ip(&r), "203.0.113.9".parse::<IpAddr>().unwrap());
662 let r = request_from("198.51.100.4", Some("10.0.0.1"));
664 assert_eq!(client_ip(&r), "198.51.100.4".parse::<IpAddr>().unwrap());
665 let r = request_from("10.42.0.7", None);
667 assert_eq!(client_ip(&r), "10.42.0.7".parse::<IpAddr>().unwrap());
668 }
669
670 #[test]
671 fn refresh_token_falls_back_to_the_cookie() {
672 let mut headers = axum::http::HeaderMap::new();
673 headers.insert(
674 COOKIE,
675 format!("a=1; {REFRESH_COOKIE}=from-cookie; b=2")
676 .parse()
677 .unwrap(),
678 );
679
680 assert_eq!(
681 refresh_token_from(None, &headers).as_deref(),
682 Some("from-cookie")
683 );
684 assert_eq!(
685 refresh_token_from(Some("from-body"), &headers).as_deref(),
686 Some("from-body")
687 );
688 assert_eq!(
689 refresh_token_from(None, &axum::http::HeaderMap::new()),
690 None
691 );
692 }
693
694 #[test]
695 fn cookies_are_lax_and_only_secure_when_tls_is_in_play() {
696 let state = |cookie_secure| AuthRouteState {
697 pool: Arc::new(Pool::new("/nonexistent".into())),
698 private_pem: Arc::new(Vec::new()),
699 public_pem: Arc::new(Vec::new()),
700 access_ttl_secs: 900,
701 refresh_ttl_secs: 60,
702 cookie_secure,
703 login_limiter: Arc::new(LoginRateLimiter::default()),
704 };
705
706 let plain = state(false).access_cookie("tok");
707 assert!(plain.contains("SameSite=Lax"));
708 assert!(plain.contains("HttpOnly"));
709 assert!(!plain.contains("Secure"));
710
711 assert!(state(true).access_cookie("tok").contains("; Secure"));
712
713 let resp = (StatusCode::OK, state(false).session_cookies("a", "r")).into_response();
715 let set: Vec<_> = resp.headers().get_all(SET_COOKIE).iter().collect();
716 assert_eq!(set.len(), 3);
717 assert_eq!(
718 (StatusCode::OK, state(false).cleared_cookies())
719 .into_response()
720 .headers()
721 .get_all(SET_COOKIE)
722 .iter()
723 .count(),
724 3
725 );
726
727 let refresh = state(false).refresh_cookie("tok");
729 assert!(refresh.contains("Path=/auth;"));
730 assert!(refresh.contains("HttpOnly"));
731 }
732}