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::{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
20const 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
109fn dummy_password_hash() -> &'static str {
112 static HASH: std::sync::OnceLock<String> = std::sync::OnceLock::new();
113 HASH.get_or_init(|| auth::hash_password("koan-dummy-password").unwrap_or_default())
114}
115
116fn refresh_token_from(body: Option<&str>, headers: &axum::http::HeaderMap) -> Option<String> {
119 if let Some(t) = body.filter(|t| !t.is_empty()) {
120 return Some(t.to_owned());
121 }
122 headers
123 .get(COOKIE)
124 .and_then(|v| v.to_str().ok())
125 .and_then(|cookies| {
126 cookies.split(';').find_map(|c| {
127 c.trim()
128 .strip_prefix(&format!("{REFRESH_COOKIE}="))
129 .map(str::to_owned)
130 })
131 })
132}
133
134async fn login_rate_limit(
139 State(state): State<AuthRouteState>,
140 request: axum::extract::Request,
141 next: axum::middleware::Next,
142) -> Response {
143 let ip = request
144 .extensions()
145 .get::<ConnectInfo<SocketAddr>>()
146 .map(|ConnectInfo(addr)| addr.ip())
147 .unwrap_or(IpAddr::V4(std::net::Ipv4Addr::UNSPECIFIED));
148
149 if !state.login_limiter.allow(ip) {
150 return (
151 StatusCode::TOO_MANY_REQUESTS,
152 Json(MessageResponse {
153 message: "too many login attempts".into(),
154 }),
155 )
156 .into_response();
157 }
158 next.run(request).await
159}
160
161impl AuthRouteState {
162 fn open_db(&self) -> Result<Database, (StatusCode, String)> {
163 Database::open(&self.db_path).map_err(|e| {
164 log::error!("auth db open error: {}", e);
165 (
166 StatusCode::INTERNAL_SERVER_ERROR,
167 "internal error".to_string(),
168 )
169 })
170 }
171}
172
173#[derive(Deserialize)]
178pub struct LoginRequest {
179 pub username: String,
180 pub password: String,
181}
182
183#[derive(Serialize)]
184pub struct LoginResponse {
185 pub access_token: String,
186 pub refresh_token: String,
187 pub token_type: String,
188 pub expires_in: u64,
189 pub user: UserInfo,
190}
191
192#[derive(Serialize)]
193pub struct UserInfo {
194 pub id: i64,
195 pub username: String,
196 pub role: String,
197}
198
199#[derive(Deserialize, Default)]
200#[serde(default)]
201pub struct RefreshRequest {
202 pub refresh_token: Option<String>,
203}
204
205#[derive(Serialize)]
206pub struct RefreshResponse {
207 pub access_token: String,
208 pub refresh_token: String,
209 pub token_type: String,
210 pub expires_in: u64,
211}
212
213#[derive(Deserialize, Default)]
214#[serde(default)]
215pub struct LogoutRequest {
216 pub refresh_token: Option<String>,
217}
218
219#[derive(Serialize)]
220pub struct MessageResponse {
221 pub message: String,
222}
223
224pub fn auth_router(state: AuthRouteState) -> axum::Router {
229 axum::Router::new()
230 .route(
231 "/auth/login",
232 post(login).layer(axum::middleware::from_fn_with_state(
233 state.clone(),
234 login_rate_limit,
235 )),
236 )
237 .route("/auth/refresh", post(refresh))
238 .route("/auth/logout", post(logout))
239 .layer(tower::limit::ConcurrencyLimitLayer::new(2))
243 .with_state(state)
244}
245
246async fn login(State(state): State<AuthRouteState>, Json(req): Json<LoginRequest>) -> Response {
251 let db = match state.open_db() {
252 Ok(db) => db,
253 Err((status, msg)) => return (status, msg).into_response(),
254 };
255
256 let user = match auth_queries::get_user_by_username(&db.conn, &req.username) {
258 Ok(Some(u)) => u,
259 Ok(None) => {
260 let password = req.password.clone();
263 let _ = tokio::task::spawn_blocking(move || {
264 auth::verify_password(&password, dummy_password_hash())
265 })
266 .await;
267 return (
268 StatusCode::UNAUTHORIZED,
269 Json(MessageResponse {
270 message: "invalid username or password".into(),
271 }),
272 )
273 .into_response();
274 }
275 Err(e) => {
276 log::error!("auth login db error: {}", e);
277 return (StatusCode::INTERNAL_SERVER_ERROR, "internal error").into_response();
278 }
279 };
280
281 let hash = user.password_hash.clone();
284 let password = req.password.clone();
285 let verified = tokio::task::spawn_blocking(move || auth::verify_password(&password, &hash))
286 .await
287 .map(|r| r.is_ok())
288 .unwrap_or(false);
289
290 if !verified {
291 return (
292 StatusCode::UNAUTHORIZED,
293 Json(MessageResponse {
294 message: "invalid username or password".into(),
295 }),
296 )
297 .into_response();
298 }
299
300 let access_token = match auth::mint_access_token(
302 &state.private_pem,
303 user.id,
304 &user.username,
305 user.role,
306 state.access_ttl_secs,
307 ) {
308 Ok(t) => t,
309 Err(e) => {
310 log::error!("auth mint token error: {}", e);
311 return (StatusCode::INTERNAL_SERVER_ERROR, "token error").into_response();
312 }
313 };
314
315 let refresh_token_id = match auth::random_token() {
317 Ok(t) => t,
318 Err(e) => {
319 log::error!("auth refresh token generation error: {}", e);
320 return (StatusCode::INTERNAL_SERVER_ERROR, "token error").into_response();
321 }
322 };
323 let refresh_expires = auth::now_unix() as i64 + state.refresh_ttl_secs as i64;
324 if let Err(e) =
325 auth_queries::store_refresh_token(&db.conn, &refresh_token_id, user.id, refresh_expires)
326 {
327 log::error!("auth store refresh token error: {}", e);
328 return (StatusCode::INTERNAL_SERVER_ERROR, "token error").into_response();
329 }
330
331 let _ = auth_queries::cleanup_expired_tokens(&db.conn);
333
334 let cookies = [
335 (SET_COOKIE, state.access_cookie(&access_token)),
336 (SET_COOKIE, state.refresh_cookie(&refresh_token_id)),
337 (SET_COOKIE, state.stale_refresh_cookie()),
338 ];
339
340 let resp = LoginResponse {
341 access_token,
342 refresh_token: refresh_token_id,
345 token_type: "Bearer".into(),
346 expires_in: state.access_ttl_secs,
347 user: UserInfo {
348 id: user.id,
349 username: user.username,
350 role: user.role.as_str().into(),
351 },
352 };
353
354 (StatusCode::OK, cookies, Json(resp)).into_response()
355}
356
357async fn refresh(
358 State(state): State<AuthRouteState>,
359 headers: axum::http::HeaderMap,
360 body: Option<Json<RefreshRequest>>,
361) -> Response {
362 let supplied = body.and_then(|Json(req)| req.refresh_token);
363 let Some(supplied) = refresh_token_from(supplied.as_deref(), &headers) else {
364 return (
365 StatusCode::UNAUTHORIZED,
366 Json(MessageResponse {
367 message: "missing refresh token".into(),
368 }),
369 )
370 .into_response();
371 };
372
373 let db = match state.open_db() {
374 Ok(db) => db,
375 Err((status, msg)) => return (status, msg).into_response(),
376 };
377
378 let token = match auth_queries::consume_refresh_token(&db.conn, &supplied) {
381 Ok(Some(t)) => t,
382 Ok(None) => {
383 return (
384 StatusCode::UNAUTHORIZED,
385 Json(MessageResponse {
386 message: "invalid or expired refresh token".into(),
387 }),
388 )
389 .into_response();
390 }
391 Err(e) => {
392 log::error!("auth refresh db error: {}", e);
393 return (StatusCode::INTERNAL_SERVER_ERROR, "internal error").into_response();
394 }
395 };
396
397 let user = match auth_queries::get_user_by_id(&db.conn, token.user_id) {
399 Ok(Some(u)) => u,
400 Ok(None) => {
401 return (
402 StatusCode::UNAUTHORIZED,
403 Json(MessageResponse {
404 message: "user not found".into(),
405 }),
406 )
407 .into_response();
408 }
409 Err(e) => {
410 log::error!("auth refresh user lookup error: {}", e);
411 return (StatusCode::INTERNAL_SERVER_ERROR, "internal error").into_response();
412 }
413 };
414
415 let access_token = match auth::mint_access_token(
417 &state.private_pem,
418 user.id,
419 &user.username,
420 user.role,
421 state.access_ttl_secs,
422 ) {
423 Ok(t) => t,
424 Err(e) => {
425 log::error!("auth mint token error: {}", e);
426 return (StatusCode::INTERNAL_SERVER_ERROR, "token error").into_response();
427 }
428 };
429
430 let new_refresh_id = match auth::random_token() {
432 Ok(t) => t,
433 Err(e) => {
434 log::error!("auth refresh token generation error: {}", e);
435 return (StatusCode::INTERNAL_SERVER_ERROR, "token error").into_response();
436 }
437 };
438 let refresh_expires = auth::now_unix() as i64 + state.refresh_ttl_secs as i64;
439 if let Err(e) =
440 auth_queries::store_refresh_token(&db.conn, &new_refresh_id, user.id, refresh_expires)
441 {
442 log::error!("auth store refresh token error: {}", e);
443 return (StatusCode::INTERNAL_SERVER_ERROR, "token error").into_response();
444 }
445
446 let cookies = [
447 (SET_COOKIE, state.access_cookie(&access_token)),
448 (SET_COOKIE, state.refresh_cookie(&new_refresh_id)),
449 (SET_COOKIE, state.stale_refresh_cookie()),
450 ];
451
452 let resp = RefreshResponse {
453 access_token,
454 refresh_token: new_refresh_id,
455 token_type: "Bearer".into(),
456 expires_in: state.access_ttl_secs,
457 };
458
459 (StatusCode::OK, cookies, Json(resp)).into_response()
460}
461
462async fn logout(
463 State(state): State<AuthRouteState>,
464 headers: axum::http::HeaderMap,
465 body: Option<Json<LogoutRequest>>,
466) -> Response {
467 let db = match state.open_db() {
468 Ok(db) => db,
469 Err((status, msg)) => return (status, msg).into_response(),
470 };
471
472 let supplied = body.and_then(|Json(req)| req.refresh_token);
473 if let Some(token) = refresh_token_from(supplied.as_deref(), &headers) {
474 let _ = auth_queries::revoke_refresh_token(&db.conn, &token);
475 }
476
477 let cookies = [
478 (SET_COOKIE, state.cookie("koan_access", "", "/", 0)),
479 (
480 SET_COOKIE,
481 state.cookie(REFRESH_COOKIE, "", REFRESH_COOKIE_PATH, 0),
482 ),
483 (SET_COOKIE, state.stale_refresh_cookie()),
484 ];
485
486 (
487 StatusCode::OK,
488 cookies,
489 Json(MessageResponse {
490 message: "logged out".into(),
491 }),
492 )
493 .into_response()
494}
495
496#[cfg(test)]
501mod tests {
502 use super::*;
503
504 #[test]
505 fn login_limiter_caps_a_single_ip() {
506 let limiter = LoginRateLimiter::default();
507 let ip: IpAddr = "10.0.0.5".parse().unwrap();
508 for _ in 0..LOGIN_MAX_PER_WINDOW {
509 assert!(limiter.allow(ip));
510 }
511 assert!(!limiter.allow(ip));
512
513 assert!(limiter.allow("10.0.0.6".parse().unwrap()));
515 }
516
517 #[test]
518 fn refresh_token_falls_back_to_the_cookie() {
519 let mut headers = axum::http::HeaderMap::new();
520 headers.insert(
521 COOKIE,
522 format!("a=1; {REFRESH_COOKIE}=from-cookie; b=2")
523 .parse()
524 .unwrap(),
525 );
526
527 assert_eq!(
528 refresh_token_from(None, &headers).as_deref(),
529 Some("from-cookie")
530 );
531 assert_eq!(
532 refresh_token_from(Some("from-body"), &headers).as_deref(),
533 Some("from-body")
534 );
535 assert_eq!(
536 refresh_token_from(None, &axum::http::HeaderMap::new()),
537 None
538 );
539 }
540
541 #[test]
542 fn cookies_are_lax_and_only_secure_when_tls_is_in_play() {
543 let state = |cookie_secure| AuthRouteState {
544 db_path: PathBuf::from("/nonexistent"),
545 private_pem: Arc::new(Vec::new()),
546 public_pem: Arc::new(Vec::new()),
547 access_ttl_secs: 900,
548 refresh_ttl_secs: 60,
549 cookie_secure,
550 login_limiter: Arc::new(LoginRateLimiter::default()),
551 };
552
553 let plain = state(false).access_cookie("tok");
554 assert!(plain.contains("SameSite=Lax"));
555 assert!(plain.contains("HttpOnly"));
556 assert!(!plain.contains("Secure"));
557
558 assert!(state(true).access_cookie("tok").contains("; Secure"));
559
560 let refresh = state(false).refresh_cookie("tok");
562 assert!(refresh.contains("Path=/auth;"));
563 assert!(refresh.contains("HttpOnly"));
564 }
565}