use std::{
sync::Arc,
time::{SystemTime, UNIX_EPOCH},
};
use axum::{
body::Body,
extract::State,
http::{Request, StatusCode, header},
middleware::Next,
response::{IntoResponse, Response},
};
use dashmap::DashMap;
use subtle::ConstantTimeEq as _;
const ADMIN_AUTH_WINDOW_SECS: u64 = 60;
#[derive(Clone)]
struct FailureRecord {
count: u32,
window_start: u64,
}
#[derive(Clone)]
pub(crate) struct FailureLimiter {
records: Arc<DashMap<String, FailureRecord>>,
max_failures: u32,
}
impl FailureLimiter {
#[cfg(test)]
pub(super) const WINDOW_SECS: u64 = ADMIN_AUTH_WINDOW_SECS;
pub(crate) fn new(max_failures: u32) -> Self {
Self {
records: Arc::new(DashMap::new()),
max_failures,
}
}
fn now_secs() -> u64 {
SystemTime::now().duration_since(UNIX_EPOCH).map_or(0, |d| d.as_secs())
}
pub(super) fn evict_expired(&self, now: u64) {
const EVICTION_THRESHOLD: usize = 1024;
const RETENTION_WINDOWS: u64 = 2;
if self.records.len() < EVICTION_THRESHOLD {
return;
}
let cutoff = ADMIN_AUTH_WINDOW_SECS * RETENTION_WINDOWS;
self.records
.retain(|_, record| now.saturating_sub(record.window_start) < cutoff);
}
pub(crate) fn record_failure(&self, ip: &str) -> bool {
let now = Self::now_secs();
self.evict_expired(now);
let mut entry = self.records.entry(ip.to_string()).or_insert_with(|| FailureRecord {
count: 0,
window_start: now,
});
if now >= entry.window_start + ADMIN_AUTH_WINDOW_SECS {
entry.count = 1;
entry.window_start = now;
false
} else {
entry.count = entry.count.saturating_add(1);
entry.count >= self.max_failures
}
}
pub(crate) fn is_blocked(&self, ip: &str) -> bool {
let now = Self::now_secs();
if let Some(entry) = self.records.get(ip) {
if now < entry.window_start + ADMIN_AUTH_WINDOW_SECS {
return entry.count >= self.max_failures;
}
}
false
}
pub(crate) fn record_success(&self, ip: &str) {
self.records.remove(ip);
}
#[cfg(test)]
pub(crate) fn failure_count(&self, ip: &str) -> u32 {
self.records.get(ip).map_or(0, |e| e.count)
}
#[cfg(test)]
pub(super) fn record_count(&self) -> usize {
self.records.len()
}
#[cfg(test)]
pub(super) fn record_failure_at(&self, ip: &str, now: u64) {
self.records.insert(
ip.to_string(),
FailureRecord {
count: 1,
window_start: now,
},
);
}
}
#[derive(Clone)]
pub struct BearerAuthState {
pub token: Arc<String>,
failure_limiter: FailureLimiter,
}
impl BearerAuthState {
#[must_use]
pub fn new(token: String) -> Self {
Self::with_max_failures(token, 10)
}
#[must_use]
pub fn with_max_failures(token: String, max_failures: u32) -> Self {
Self {
token: Arc::new(token),
failure_limiter: FailureLimiter::new(max_failures),
}
}
}
pub async fn bearer_auth_middleware(
State(auth_state): State<BearerAuthState>,
request: Request<Body>,
next: Next,
) -> Response {
use std::net::SocketAddr;
use axum::extract::ConnectInfo;
let peer_key = request
.extensions()
.get::<ConnectInfo<SocketAddr>>()
.map_or_else(|| "unknown".to_string(), |ci| ci.0.ip().to_string());
if auth_state.failure_limiter.is_blocked(&peer_key) {
return (StatusCode::TOO_MANY_REQUESTS, "Too many failed auth attempts").into_response();
}
let auth_header = request
.headers()
.get(header::AUTHORIZATION)
.and_then(|value| value.to_str().ok());
match auth_header {
None => {
return (
StatusCode::UNAUTHORIZED,
[(header::WWW_AUTHENTICATE, "Bearer")],
"Missing Authorization header",
)
.into_response();
},
Some(header_value) => {
if !header_value.starts_with("Bearer ") {
return (
StatusCode::UNAUTHORIZED,
[(header::WWW_AUTHENTICATE, "Bearer")],
"Invalid Authorization header format. Expected: Bearer <token>",
)
.into_response();
}
let token = &header_value[7..];
if !constant_time_compare(token, &auth_state.token) {
if auth_state.failure_limiter.record_failure(&peer_key) {
return (StatusCode::TOO_MANY_REQUESTS, "Too many failed auth attempts")
.into_response();
}
return (StatusCode::FORBIDDEN, "Invalid token").into_response();
}
auth_state.failure_limiter.record_success(&peer_key);
},
}
next.run(request).await
}
#[derive(Clone)]
pub struct AdminPrincipalState {
platform_token: Arc<String>,
tokens: Option<Arc<crate::api::admin_principal::PgAdminTokenStore>>,
failure_limiter: FailureLimiter,
}
impl AdminPrincipalState {
#[must_use]
pub fn new(
platform_token: String,
tokens: Option<Arc<crate::api::admin_principal::PgAdminTokenStore>>,
max_failures: u32,
) -> Self {
Self {
platform_token: Arc::new(platform_token),
tokens,
failure_limiter: FailureLimiter::new(max_failures),
}
}
}
pub async fn admin_principal_middleware(
State(state): State<AdminPrincipalState>,
mut request: Request<Body>,
next: Next,
) -> Response {
use std::net::SocketAddr;
use axum::extract::ConnectInfo;
use crate::api::admin_principal::AdminPrincipal;
let peer_key = request
.extensions()
.get::<ConnectInfo<SocketAddr>>()
.map_or_else(|| "unknown".to_string(), |ci| ci.0.ip().to_string());
if state.failure_limiter.is_blocked(&peer_key) {
return (StatusCode::TOO_MANY_REQUESTS, "Too many failed auth attempts").into_response();
}
let Some(header_value) =
request.headers().get(header::AUTHORIZATION).and_then(|v| v.to_str().ok())
else {
return (
StatusCode::UNAUTHORIZED,
[(header::WWW_AUTHENTICATE, "Bearer")],
"Missing Authorization header",
)
.into_response();
};
let Some(token) = extract_bearer_token(header_value) else {
return (
StatusCode::UNAUTHORIZED,
[(header::WWW_AUTHENTICATE, "Bearer")],
"Invalid Authorization header format. Expected: Bearer <token>",
)
.into_response();
};
let principal = if constant_time_compare(token, &state.platform_token) {
Some(AdminPrincipal::Platform)
} else if let Some(tokens) = state.tokens.as_ref() {
match tokens.authenticate(token).await {
Ok(tenant) => tenant.map(AdminPrincipal::Tenant),
Err(e) => {
tracing::error!(error = %e, "tenant admin credential lookup failed");
return (StatusCode::SERVICE_UNAVAILABLE, "Admin credential store unavailable")
.into_response();
},
}
} else {
None
};
let Some(principal) = principal else {
if state.failure_limiter.record_failure(&peer_key) {
return (StatusCode::TOO_MANY_REQUESTS, "Too many failed auth attempts")
.into_response();
}
return (StatusCode::FORBIDDEN, "Invalid token").into_response();
};
state.failure_limiter.record_success(&peer_key);
request.extensions_mut().insert(principal);
next.run(request).await
}
#[must_use]
pub fn extract_bearer_token(header_value: &str) -> Option<&str> {
header_value.strip_prefix("Bearer ")
}
pub(crate) fn constant_time_compare(a: &str, b: &str) -> bool {
a.as_bytes().ct_eq(b.as_bytes()).into()
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum AdminPrivilege {
ReadOnly,
ReadWrite,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct AdminCaller {
pub privilege: AdminPrivilege,
pub peer_ip: String,
}
#[derive(Clone)]
pub struct AdminDualAuthState {
write_token: Arc<String>,
readonly_token: Option<Arc<String>>,
failure_limiter: FailureLimiter,
}
impl AdminDualAuthState {
#[must_use]
pub fn new(write_token: String, readonly_token: Option<String>, max_failures: u32) -> Self {
Self {
write_token: Arc::new(write_token),
readonly_token: readonly_token.map(Arc::new),
failure_limiter: FailureLimiter::new(max_failures),
}
}
fn classify(&self, presented: &str) -> Option<AdminPrivilege> {
let is_write = constant_time_compare(presented, &self.write_token);
let is_readonly = self
.readonly_token
.as_ref()
.is_some_and(|t| constant_time_compare(presented, t));
match (is_write, is_readonly) {
(true, _) => Some(AdminPrivilege::ReadWrite),
(false, true) => Some(AdminPrivilege::ReadOnly),
(false, false) => None,
}
}
}
pub async fn admin_dual_auth_middleware(
State(auth_state): State<AdminDualAuthState>,
mut request: Request<Body>,
next: Next,
) -> Response {
use std::net::SocketAddr;
use axum::extract::ConnectInfo;
let peer_key = request
.extensions()
.get::<ConnectInfo<SocketAddr>>()
.map_or_else(|| "unknown".to_string(), |ci| ci.0.ip().to_string());
if auth_state.failure_limiter.is_blocked(&peer_key) {
return (StatusCode::TOO_MANY_REQUESTS, "Too many failed auth attempts").into_response();
}
let Some(header_value) = request
.headers()
.get(header::AUTHORIZATION)
.and_then(|value| value.to_str().ok())
else {
return (
StatusCode::UNAUTHORIZED,
[(header::WWW_AUTHENTICATE, "Bearer")],
"Missing Authorization header",
)
.into_response();
};
let Some(token) = extract_bearer_token(header_value) else {
return (
StatusCode::UNAUTHORIZED,
[(header::WWW_AUTHENTICATE, "Bearer")],
"Invalid Authorization header format. Expected: Bearer <token>",
)
.into_response();
};
let Some(privilege) = auth_state.classify(token) else {
if auth_state.failure_limiter.record_failure(&peer_key) {
return (StatusCode::TOO_MANY_REQUESTS, "Too many failed auth attempts")
.into_response();
}
return (StatusCode::FORBIDDEN, "Invalid token").into_response();
};
auth_state.failure_limiter.record_success(&peer_key);
request.extensions_mut().insert(AdminCaller {
privilege,
peer_ip: peer_key,
});
next.run(request).await
}
#[cfg(test)]
mod xff_tests {
#![allow(clippy::unwrap_used)]
use axum::{
Router,
body::Body,
http::{Request, StatusCode},
middleware,
routing::get,
};
use tower::ServiceExt as _;
use super::{BearerAuthState, bearer_auth_middleware};
async fn protected() -> &'static str {
"ok"
}
fn wrong_token_request(xff: &str) -> Request<Body> {
Request::builder()
.uri("/")
.header("authorization", "Bearer wrong-token")
.header("x-forwarded-for", xff)
.body(Body::empty())
.unwrap()
}
#[tokio::test]
async fn rotating_x_forwarded_for_does_not_refresh_the_failure_budget() {
let state = BearerAuthState::with_max_failures("correct-token".to_string(), 2);
let app = Router::new()
.route("/", get(protected))
.layer(middleware::from_fn_with_state(state, bearer_auth_middleware));
let mut statuses = Vec::new();
for i in 0..5 {
let resp = app
.clone()
.oneshot(wrong_token_request(&format!("203.0.113.{i}")))
.await
.unwrap();
statuses.push(resp.status());
}
assert!(
statuses.contains(&StatusCode::TOO_MANY_REQUESTS),
"rotating X-Forwarded-For must still hit the shared rate limit, got {statuses:?}"
);
}
}