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