use std::net::SocketAddr;
use std::sync::Arc;
use axum::extract::{ConnectInfo, Request, State};
use axum::http::{HeaderMap, StatusCode};
use axum::middleware::Next;
use axum::response::Response;
use super::respond::error;
use super::AppState;
pub(super) async fn guard(
State(state): State<Arc<AppState>>,
req: Request,
next: Next,
) -> Response {
if state
.limiter
.limited(&client_ip(&req, &state.cfg.trusted_ip_header))
{
let mut resp = error(
StatusCode::TOO_MANY_REQUESTS,
"rate limit exceeded, try again later",
);
if let Ok(v) = state
.cfg
.rate_limit_window
.as_secs()
.to_string()
.parse::<axum::http::HeaderValue>()
{
resp.headers_mut().insert("retry-after", v);
}
return resp;
}
if !authorized(&state.cfg.token, req.headers()) {
return error(StatusCode::UNAUTHORIZED, "unauthorized");
}
next.run(req).await
}
fn authorized(token: &str, headers: &HeaderMap) -> bool {
let Some(value) = headers
.get(axum::http::header::AUTHORIZATION)
.and_then(|v| v.to_str().ok())
.and_then(|v| v.strip_prefix("Bearer "))
else {
return false;
};
!value.is_empty() && constant_time_eq(value.as_bytes(), token.as_bytes())
}
fn constant_time_eq(a: &[u8], b: &[u8]) -> bool {
if a.len() != b.len() {
return false;
}
let mut diff = 0u8;
for (x, y) in a.iter().zip(b) {
diff |= x ^ y;
}
std::hint::black_box(diff) == 0
}
fn client_ip(req: &Request, trusted_header: &str) -> String {
if !trusted_header.is_empty() {
if let Some(ip) = header_str(req.headers(), trusted_header) {
return ip.to_string();
}
}
req.extensions()
.get::<ConnectInfo<SocketAddr>>()
.map(|ConnectInfo(addr)| addr.ip().to_string())
.unwrap_or_else(|| "unknown".to_string())
}
fn header_str<'a>(headers: &'a HeaderMap, name: &str) -> Option<&'a str> {
headers
.get(name)
.and_then(|v| v.to_str().ok())
.map(str::trim)
.filter(|v| !v.is_empty())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn bearer_comparison_rejects_everything_but_the_exact_token() {
let mut h = HeaderMap::new();
assert!(!authorized("secret", &h), "no header");
h.insert("authorization", "secret".parse().unwrap());
assert!(!authorized("secret", &h), "missing Bearer scheme");
h.insert("authorization", "Bearer ".parse().unwrap());
assert!(!authorized("secret", &h), "empty token");
h.insert("authorization", "Bearer secre".parse().unwrap());
assert!(!authorized("secret", &h), "prefix of the token");
h.insert("authorization", "Basic secret".parse().unwrap());
assert!(!authorized("secret", &h), "wrong scheme");
h.insert("authorization", "Bearer secret".parse().unwrap());
assert!(authorized("secret", &h));
}
fn request_with(headers: Vec<(&str, &str)>) -> Request {
let mut req = Request::new(axum::body::Body::empty());
req.extensions_mut()
.insert(ConnectInfo(SocketAddr::from(([127, 0, 0, 1], 1234))));
for (k, v) in headers {
let name = axum::http::HeaderName::from_bytes(k.as_bytes()).unwrap();
req.headers_mut().insert(name, v.parse().unwrap());
}
req
}
#[test]
fn client_ip_reads_the_configured_header_then_the_socket() {
assert_eq!(
client_ip(
&request_with(vec![("cf-connecting-ip", "198.51.100.4")]),
"cf-connecting-ip"
),
"198.51.100.4"
);
assert_eq!(
client_ip(
&request_with(vec![("x-real-ip", "198.51.100.7")]),
"x-real-ip"
),
"198.51.100.7"
);
assert_eq!(
client_ip(&request_with(vec![]), "cf-connecting-ip"),
"127.0.0.1"
);
assert_eq!(
client_ip(
&request_with(vec![("cf-connecting-ip", "198.51.100.4")]),
""
),
"127.0.0.1"
);
}
#[test]
fn a_header_the_ingress_does_not_set_is_ignored() {
let attacker = request_with(vec![
("cf-connecting-ip", "1.1.1.1"),
("x-forwarded-for", "2.2.2.2"),
("true-client-ip", "3.3.3.3"),
("x-real-ip", "198.51.100.7"),
]);
assert_eq!(
client_ip(&attacker, "x-real-ip"),
"198.51.100.7",
"only the configured header may decide the bucket"
);
assert_eq!(client_ip(&attacker, "cf-connecting-ip"), "1.1.1.1");
}
#[test]
fn forwarded_for_is_no_longer_split_and_trusted() {
let req = request_with(vec![("x-forwarded-for", "203.0.113.9, 10.0.0.1")]);
assert_ne!(
client_ip(&req, "cf-connecting-ip"),
"203.0.113.9",
"x-forwarded-for must not be consulted when it is not the configured header"
);
assert_eq!(client_ip(&req, "cf-connecting-ip"), "127.0.0.1");
}
}