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