107 lines
2.9 KiB
Rust
107 lines
2.9 KiB
Rust
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<AppState>,
|
|
request: Request,
|
|
next: Next,
|
|
) -> Result<Response, ApiError> {
|
|
let key = rate_limit_key(request.headers(), request.uri().path());
|
|
if let Err(rejection) = state.api_rate_limiter.check(&key) {
|
|
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<String> {
|
|
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");
|
|
}
|
|
}
|