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
64#[derive(Clone)]
69pub struct AuthRouteState {
70 pub pool: Arc<Pool>,
71 pub private_pem: Arc<Vec<u8>>,
72 pub public_pem: Arc<Vec<u8>>,
73 pub access_ttl_secs: u64,
74 pub refresh_ttl_secs: u64,
75 pub cookie_secure: bool,
79 pub login_limiter: Arc<LoginRateLimiter>,
80}
81
82impl AuthRouteState {
83 fn cookie(&self, name: &str, value: &str, path: &str, max_age: u64) -> String {
86 let secure = if self.cookie_secure { "; Secure" } else { "" };
87 format!("{name}={value}; HttpOnly; SameSite=Lax; Path={path}; Max-Age={max_age}{secure}")
88 }
89
90 fn access_cookie(&self, token: &str) -> String {
91 self.cookie("koan_access", token, "/", self.access_ttl_secs)
92 }
93
94 fn refresh_cookie(&self, token: &str) -> String {
95 self.cookie(
96 REFRESH_COOKIE,
97 token,
98 REFRESH_COOKIE_PATH,
99 self.refresh_ttl_secs,
100 )
101 }
102
103 fn stale_refresh_cookie(&self) -> String {
104 self.cookie(REFRESH_COOKIE, "", STALE_REFRESH_COOKIE_PATH, 0)
105 }
106
107 pub(crate) fn session_cookies(
110 &self,
111 access_token: &str,
112 refresh_token: &str,
113 ) -> AppendHeaders<[(axum::http::HeaderName, String); 3]> {
114 AppendHeaders([
115 (SET_COOKIE, self.access_cookie(access_token)),
116 (SET_COOKIE, self.refresh_cookie(refresh_token)),
117 (SET_COOKIE, self.stale_refresh_cookie()),
118 ])
119 }
120
121 pub(crate) fn cleared_cookies(&self) -> AppendHeaders<[(axum::http::HeaderName, String); 3]> {
123 AppendHeaders([
124 (SET_COOKIE, self.cookie("koan_access", "", "/", 0)),
125 (
126 SET_COOKIE,
127 self.cookie(REFRESH_COOKIE, "", REFRESH_COOKIE_PATH, 0),
128 ),
129 (SET_COOKIE, self.stale_refresh_cookie()),
130 ])
131 }
132}
133
134pub(crate) fn dummy_password_hash() -> &'static str {
137 static HASH: std::sync::OnceLock<String> = std::sync::OnceLock::new();
138 HASH.get_or_init(|| auth::hash_password("koan-dummy-password").unwrap_or_default())
139}
140
141pub(crate) fn refresh_token_from(
144 body: Option<&str>,
145 headers: &axum::http::HeaderMap,
146) -> Option<String> {
147 if let Some(t) = body.filter(|t| !t.is_empty()) {
148 return Some(t.to_owned());
149 }
150 headers
151 .get(COOKIE)
152 .and_then(|v| v.to_str().ok())
153 .and_then(|cookies| {
154 cookies.split(';').find_map(|c| {
155 c.trim()
156 .strip_prefix(&format!("{REFRESH_COOKIE}="))
157 .map(str::to_owned)
158 })
159 })
160}
161
162pub(crate) async fn login_rate_limit(
167 State(state): State<AuthRouteState>,
168 request: axum::extract::Request,
169 next: axum::middleware::Next,
170) -> Response {
171 let ip = request
172 .extensions()
173 .get::<ConnectInfo<SocketAddr>>()
174 .map(|ConnectInfo(addr)| addr.ip())
175 .unwrap_or(IpAddr::V4(std::net::Ipv4Addr::UNSPECIFIED));
176
177 if !state.login_limiter.allow(ip) {
178 return (
179 StatusCode::TOO_MANY_REQUESTS,
180 Json(MessageResponse {
181 message: "too many login attempts".into(),
182 }),
183 )
184 .into_response();
185 }
186 next.run(request).await
187}
188
189impl AuthRouteState {
190 fn open_db(&self) -> Result<Handle<'_>, (StatusCode, String)> {
191 self.pool.get().map_err(|e| {
192 log::error!("auth db open error: {}", e);
193 (
194 StatusCode::INTERNAL_SERVER_ERROR,
195 "internal error".to_string(),
196 )
197 })
198 }
199}
200
201#[derive(Deserialize)]
206pub struct LoginRequest {
207 pub username: String,
208 pub password: String,
209}
210
211#[derive(Serialize)]
212pub struct LoginResponse {
213 pub access_token: String,
214 pub refresh_token: String,
215 pub token_type: String,
216 pub expires_in: u64,
217 pub user: UserInfo,
218}
219
220#[derive(Serialize)]
221pub struct UserInfo {
222 pub id: i64,
223 pub username: String,
224 pub role: String,
225}
226
227#[derive(Deserialize, Default)]
228#[serde(default)]
229pub struct RefreshRequest {
230 pub refresh_token: Option<String>,
231}
232
233#[derive(Serialize)]
234pub struct RefreshResponse {
235 pub access_token: String,
236 pub refresh_token: String,
237 pub token_type: String,
238 pub expires_in: u64,
239}
240
241#[derive(Deserialize, Default)]
242#[serde(default)]
243pub struct LogoutRequest {
244 pub refresh_token: Option<String>,
245}
246
247#[derive(Serialize)]
248pub struct MessageResponse {
249 pub message: String,
250}
251
252pub fn auth_router(state: AuthRouteState) -> axum::Router {
257 axum::Router::new()
258 .route(
259 "/auth/login",
260 post(login).layer(axum::middleware::from_fn_with_state(
261 state.clone(),
262 login_rate_limit,
263 )),
264 )
265 .route("/auth/refresh", post(refresh))
266 .route("/auth/logout", post(logout))
267 .layer(tower::limit::ConcurrencyLimitLayer::new(2))
271 .with_state(state)
272}
273
274pub(crate) async fn authenticate(
286 state: &AuthRouteState,
287 username: &str,
288 password: &str,
289) -> Result<(auth_queries::UserRow, String, String), Box<Response>> {
290 let state = state.clone();
291 let (username, password) = (username.to_owned(), password.to_owned());
292 tokio::task::spawn_blocking(move || authenticate_blocking(&state, &username, &password))
293 .await
294 .unwrap_or_else(|_| Err(Box::new(StatusCode::INTERNAL_SERVER_ERROR.into_response())))
295}
296
297fn authenticate_blocking(
298 state: &AuthRouteState,
299 username: &str,
300 password: &str,
301) -> Result<(auth_queries::UserRow, String, String), Box<Response>> {
302 let db = match state.open_db() {
303 Ok(db) => db,
304 Err((status, msg)) => return Err(Box::new((status, msg).into_response())),
305 };
306
307 let user = match auth_queries::get_user_by_username(&db.conn, username) {
309 Ok(Some(u)) => u,
310 Ok(None) => {
311 let _ = auth::verify_password(password, dummy_password_hash());
314 return Err(Box::new(
315 (
316 StatusCode::UNAUTHORIZED,
317 Json(MessageResponse {
318 message: "invalid username or password".into(),
319 }),
320 )
321 .into_response(),
322 ));
323 }
324 Err(e) => {
325 log::error!("auth login db error: {}", e);
326 return Err(Box::new(
327 (StatusCode::INTERNAL_SERVER_ERROR, "internal error").into_response(),
328 ));
329 }
330 };
331
332 if auth::verify_password(password, &user.password_hash).is_err() {
333 return Err(Box::new(
334 (
335 StatusCode::UNAUTHORIZED,
336 Json(MessageResponse {
337 message: "invalid username or password".into(),
338 }),
339 )
340 .into_response(),
341 ));
342 }
343
344 if let Err(e) = auth_queries::remember_password(&db.conn, username, password) {
346 log::warn!("could not seal the password for Subsonic token auth: {e}");
347 }
348
349 let access_token = match auth::mint_access_token(
351 &state.private_pem,
352 user.id,
353 &user.username,
354 user.role,
355 state.access_ttl_secs,
356 ) {
357 Ok(t) => t,
358 Err(e) => {
359 log::error!("auth mint token error: {}", e);
360 return Err(Box::new(
361 (StatusCode::INTERNAL_SERVER_ERROR, "token error").into_response(),
362 ));
363 }
364 };
365
366 let refresh_token_id = match auth::random_token() {
368 Ok(t) => t,
369 Err(e) => {
370 log::error!("auth refresh token generation error: {}", e);
371 return Err(Box::new(
372 (StatusCode::INTERNAL_SERVER_ERROR, "token error").into_response(),
373 ));
374 }
375 };
376 let refresh_expires = auth::now_unix() as i64 + state.refresh_ttl_secs as i64;
377 if let Err(e) =
378 auth_queries::store_refresh_token(&db.conn, &refresh_token_id, user.id, refresh_expires)
379 {
380 log::error!("auth store refresh token error: {}", e);
381 return Err(Box::new(
382 (StatusCode::INTERNAL_SERVER_ERROR, "token error").into_response(),
383 ));
384 }
385
386 let _ = auth_queries::cleanup_expired_tokens(&db.conn);
388
389 Ok((user, access_token, refresh_token_id))
390}
391
392async fn login(State(state): State<AuthRouteState>, Json(req): Json<LoginRequest>) -> Response {
393 let (user, access_token, refresh_token_id) =
394 match authenticate(&state, &req.username, &req.password).await {
395 Ok(session) => session,
396 Err(resp) => return *resp,
397 };
398
399 let cookies = state.session_cookies(&access_token, &refresh_token_id);
400
401 let resp = LoginResponse {
402 access_token,
403 refresh_token: refresh_token_id,
406 token_type: "Bearer".into(),
407 expires_in: state.access_ttl_secs,
408 user: UserInfo {
409 id: user.id,
410 username: user.username,
411 role: user.role.as_str().into(),
412 },
413 };
414
415 (StatusCode::OK, cookies, Json(resp)).into_response()
416}
417
418pub(crate) fn rotate(
422 state: &AuthRouteState,
423 supplied: &str,
424) -> Result<(String, String), Box<Response>> {
425 let db = match state.open_db() {
426 Ok(db) => db,
427 Err((status, msg)) => return Err(Box::new((status, msg).into_response())),
428 };
429
430 let token = match auth_queries::consume_refresh_token(&db.conn, supplied) {
433 Ok(Some(t)) => t,
434 Ok(None) => {
435 return Err(Box::new(
436 (
437 StatusCode::UNAUTHORIZED,
438 Json(MessageResponse {
439 message: "invalid or expired refresh token".into(),
440 }),
441 )
442 .into_response(),
443 ));
444 }
445 Err(e) => {
446 log::error!("auth refresh db error: {}", e);
447 return Err(Box::new(
448 (StatusCode::INTERNAL_SERVER_ERROR, "internal error").into_response(),
449 ));
450 }
451 };
452
453 let user = match auth_queries::get_user_by_id(&db.conn, token.user_id) {
455 Ok(Some(u)) => u,
456 Ok(None) => {
457 return Err(Box::new(
458 (
459 StatusCode::UNAUTHORIZED,
460 Json(MessageResponse {
461 message: "user not found".into(),
462 }),
463 )
464 .into_response(),
465 ));
466 }
467 Err(e) => {
468 log::error!("auth refresh user lookup error: {}", e);
469 return Err(Box::new(
470 (StatusCode::INTERNAL_SERVER_ERROR, "internal error").into_response(),
471 ));
472 }
473 };
474
475 let access_token = match auth::mint_access_token(
477 &state.private_pem,
478 user.id,
479 &user.username,
480 user.role,
481 state.access_ttl_secs,
482 ) {
483 Ok(t) => t,
484 Err(e) => {
485 log::error!("auth mint token error: {}", e);
486 return Err(Box::new(
487 (StatusCode::INTERNAL_SERVER_ERROR, "token error").into_response(),
488 ));
489 }
490 };
491
492 let new_refresh_id = match auth::random_token() {
494 Ok(t) => t,
495 Err(e) => {
496 log::error!("auth refresh token generation error: {}", e);
497 return Err(Box::new(
498 (StatusCode::INTERNAL_SERVER_ERROR, "token error").into_response(),
499 ));
500 }
501 };
502 let refresh_expires = auth::now_unix() as i64 + state.refresh_ttl_secs as i64;
503 if let Err(e) =
504 auth_queries::store_refresh_token(&db.conn, &new_refresh_id, user.id, refresh_expires)
505 {
506 log::error!("auth store refresh token error: {}", e);
507 return Err(Box::new(
508 (StatusCode::INTERNAL_SERVER_ERROR, "token error").into_response(),
509 ));
510 }
511
512 Ok((access_token, new_refresh_id))
513}
514
515async fn refresh(
516 State(state): State<AuthRouteState>,
517 headers: axum::http::HeaderMap,
518 body: Option<Json<RefreshRequest>>,
519) -> Response {
520 let supplied = body.and_then(|Json(req)| req.refresh_token);
521 let Some(supplied) = refresh_token_from(supplied.as_deref(), &headers) else {
522 return (
523 StatusCode::UNAUTHORIZED,
524 Json(MessageResponse {
525 message: "missing refresh token".into(),
526 }),
527 )
528 .into_response();
529 };
530
531 let rotating = state.clone();
532 let rotated = tokio::task::spawn_blocking(move || rotate(&rotating, &supplied))
533 .await
534 .unwrap_or_else(|_| Err(Box::new(StatusCode::INTERNAL_SERVER_ERROR.into_response())));
535 let (access_token, new_refresh_id) = match rotated {
536 Ok(pair) => pair,
537 Err(resp) => return *resp,
538 };
539
540 let cookies = state.session_cookies(&access_token, &new_refresh_id);
541
542 let resp = RefreshResponse {
543 access_token,
544 refresh_token: new_refresh_id,
545 token_type: "Bearer".into(),
546 expires_in: state.access_ttl_secs,
547 };
548
549 (StatusCode::OK, cookies, Json(resp)).into_response()
550}
551
552async fn logout(
553 State(state): State<AuthRouteState>,
554 headers: axum::http::HeaderMap,
555 body: Option<Json<LogoutRequest>>,
556) -> Response {
557 let supplied = body.and_then(|Json(req)| req.refresh_token);
558 let token = refresh_token_from(supplied.as_deref(), &headers);
559 let revoking = state.clone();
560 let revoked = tokio::task::spawn_blocking(move || {
561 let db = revoking.open_db()?;
562 if let Some(token) = token {
563 let _ = auth_queries::revoke_refresh_token(&db.conn, &token);
564 }
565 Ok(())
566 })
567 .await
568 .unwrap_or_else(|_| {
569 Err((
570 StatusCode::INTERNAL_SERVER_ERROR,
571 "internal error".to_string(),
572 ))
573 });
574 if let Err((status, msg)) = revoked {
575 return (status, msg).into_response();
576 }
577
578 let cookies = state.cleared_cookies();
579
580 (
581 StatusCode::OK,
582 cookies,
583 Json(MessageResponse {
584 message: "logged out".into(),
585 }),
586 )
587 .into_response()
588}
589
590#[cfg(test)]
595mod tests {
596 use super::*;
597
598 #[test]
599 fn login_limiter_caps_a_single_ip() {
600 let limiter = LoginRateLimiter::default();
601 let ip: IpAddr = "10.0.0.5".parse().unwrap();
602 for _ in 0..LOGIN_MAX_PER_WINDOW {
603 assert!(limiter.allow(ip));
604 }
605 assert!(!limiter.allow(ip));
606
607 assert!(limiter.allow("10.0.0.6".parse().unwrap()));
609 }
610
611 #[test]
612 fn refresh_token_falls_back_to_the_cookie() {
613 let mut headers = axum::http::HeaderMap::new();
614 headers.insert(
615 COOKIE,
616 format!("a=1; {REFRESH_COOKIE}=from-cookie; b=2")
617 .parse()
618 .unwrap(),
619 );
620
621 assert_eq!(
622 refresh_token_from(None, &headers).as_deref(),
623 Some("from-cookie")
624 );
625 assert_eq!(
626 refresh_token_from(Some("from-body"), &headers).as_deref(),
627 Some("from-body")
628 );
629 assert_eq!(
630 refresh_token_from(None, &axum::http::HeaderMap::new()),
631 None
632 );
633 }
634
635 #[test]
636 fn cookies_are_lax_and_only_secure_when_tls_is_in_play() {
637 let state = |cookie_secure| AuthRouteState {
638 pool: Arc::new(Pool::new("/nonexistent".into())),
639 private_pem: Arc::new(Vec::new()),
640 public_pem: Arc::new(Vec::new()),
641 access_ttl_secs: 900,
642 refresh_ttl_secs: 60,
643 cookie_secure,
644 login_limiter: Arc::new(LoginRateLimiter::default()),
645 };
646
647 let plain = state(false).access_cookie("tok");
648 assert!(plain.contains("SameSite=Lax"));
649 assert!(plain.contains("HttpOnly"));
650 assert!(!plain.contains("Secure"));
651
652 assert!(state(true).access_cookie("tok").contains("; Secure"));
653
654 let resp = (StatusCode::OK, state(false).session_cookies("a", "r")).into_response();
656 let set: Vec<_> = resp.headers().get_all(SET_COOKIE).iter().collect();
657 assert_eq!(set.len(), 3);
658 assert_eq!(
659 (StatusCode::OK, state(false).cleared_cookies())
660 .into_response()
661 .headers()
662 .get_all(SET_COOKIE)
663 .iter()
664 .count(),
665 3
666 );
667
668 let refresh = state(false).refresh_cookie("tok");
670 assert!(refresh.contains("Path=/auth;"));
671 assert!(refresh.contains("HttpOnly"));
672 }
673}