use axum::{ extract::{Request, State}, http::header::{COOKIE, HeaderMap}, middleware::Next, response::Response, }; use crank_runtime::RateLimitRejection; use crate::{auth::SESSION_COOKIE_NAME, error::ApiError, state::AppState}; pub async fn apply_api_rate_limit( State(state): State, request: Request, next: Next, ) -> Result { let key = rate_limit_key(request.headers(), request.uri().path()); if let Err(rejection) = state.api_rate_limiter.check(&key).await { return Err(ApiError::rate_limited_with_context( "request rate limit exceeded", rejection_context(rejection), )); } Ok(next.run(request).await) } fn rejection_context(rejection: RateLimitRejection) -> serde_json::Value { serde_json::json!({ "retry_after_ms": rejection.retry_after_ms, }) } fn rate_limit_key(headers: &HeaderMap, path: &str) -> String { if let Some(session_id) = session_id_from_headers(headers) { return format!("session:{session_id}"); } if let Some(forwarded_for) = header_value(headers, "x-forwarded-for") { let ip = forwarded_for .split(',') .next() .map(str::trim) .filter(|value| !value.is_empty()) .unwrap_or("unknown"); return format!("ip:{ip}"); } if let Some(real_ip) = header_value(headers, "x-real-ip") { return format!("ip:{real_ip}"); } format!("anonymous:{path}") } fn session_id_from_headers(headers: &HeaderMap) -> Option { let cookies = headers.get(COOKIE)?.to_str().ok()?; for part in cookies.split(';') { let (name, value) = part.trim().split_once('=')?; if name != SESSION_COOKIE_NAME { continue; } let (session_id, _) = value.split_once('.')?; if !session_id.is_empty() { return Some(session_id.to_owned()); } } None } fn header_value<'a>(headers: &'a HeaderMap, name: &'static str) -> Option<&'a str> { headers.get(name)?.to_str().ok().map(str::trim) } #[cfg(test)] mod tests { use axum::http::{HeaderMap, HeaderValue, header::COOKIE}; use super::rate_limit_key; #[test] fn keys_by_session_cookie_first() { let mut headers = HeaderMap::new(); headers.insert( COOKIE, HeaderValue::from_static("theme=dark; crank_session=sess_123.secret_456"), ); headers.insert("x-forwarded-for", HeaderValue::from_static("10.0.0.5")); assert_eq!( rate_limit_key(&headers, "/api/auth/login"), "session:sess_123" ); } #[test] fn falls_back_to_forwarded_ip() { let mut headers = HeaderMap::new(); headers.insert( "x-forwarded-for", HeaderValue::from_static("10.0.0.5, 10.0.0.6"), ); assert_eq!(rate_limit_key(&headers, "/api/auth/login"), "ip:10.0.0.5"); } }