use axum::{
body::Body,
extract::State,
http::{
header::{AUTHORIZATION, RETRY_AFTER},
HeaderValue, Request, StatusCode,
},
middleware::Next,
response::{IntoResponse, Response},
Json,
};
use std::sync::{Arc, Mutex};
use std::time::Instant;
#[derive(Clone)]
pub struct AuthConfig {
pub api_key: Arc<String>,
}
pub async fn require_api_key(
State(config): State<AuthConfig>,
req: Request<Body>,
next: Next,
) -> Response {
let provided = req
.headers()
.get(AUTHORIZATION)
.and_then(|v| v.to_str().ok())
.and_then(|v| v.strip_prefix("Bearer "));
if provided == Some(config.api_key.as_str()) {
next.run(req).await
} else {
(
StatusCode::UNAUTHORIZED,
Json(serde_json::json!({"error": {"message": "invalid or missing API key"}})),
)
.into_response()
}
}
struct Bucket {
tokens: f64,
capacity: f64,
refill_per_sec: f64,
last_refill: Instant,
}
pub struct RateLimiter {
inner: Mutex<Bucket>,
}
impl RateLimiter {
pub fn new(capacity: u32, refill_per_sec: f64) -> Self {
RateLimiter {
inner: Mutex::new(Bucket {
tokens: capacity as f64,
capacity: capacity as f64,
refill_per_sec,
last_refill: Instant::now(),
}),
}
}
pub fn per_minute(requests_per_minute: u32) -> Self {
Self::new(requests_per_minute, requests_per_minute as f64 / 60.0)
}
pub fn try_acquire(&self) -> bool {
let mut b = self.inner.lock().unwrap_or_else(|p| p.into_inner());
let now = Instant::now();
let elapsed = now.duration_since(b.last_refill).as_secs_f64();
b.tokens = (b.tokens + elapsed * b.refill_per_sec).min(b.capacity);
b.last_refill = now;
if b.tokens >= 1.0 {
b.tokens -= 1.0;
true
} else {
false
}
}
}
pub async fn rate_limit(
State(limiter): State<Arc<RateLimiter>>,
req: Request<Body>,
next: Next,
) -> Response {
if limiter.try_acquire() {
next.run(req).await
} else {
(
StatusCode::TOO_MANY_REQUESTS,
Json(serde_json::json!({"error": {"message": "rate limit exceeded"}})),
)
.into_response()
}
}
pub const RETRY_AFTER_SECONDS: u64 = 1;
pub async fn retry_after(req: Request<Body>, next: Next) -> Response {
let mut response = next.run(req).await;
if response.status() == StatusCode::SERVICE_UNAVAILABLE
&& !response.headers().contains_key(RETRY_AFTER)
{
response
.headers_mut()
.insert(RETRY_AFTER, HeaderValue::from(RETRY_AFTER_SECONDS));
}
response
}
#[cfg(test)]
mod tests {
use super::*;
use axum::{routing::get, Router};
use tower::ServiceExt;
async fn call(app: Router, path: &str) -> Response {
app.oneshot(
Request::builder()
.uri(path)
.body(Body::empty())
.expect("request"),
)
.await
.expect("response")
}
fn retry_after_router() -> Router {
Router::new()
.route("/busy", get(|| async { StatusCode::SERVICE_UNAVAILABLE }))
.route("/fine", get(|| async { StatusCode::OK }))
.route(
"/busy-with-hint",
get(|| async { ([(RETRY_AFTER, "30")], StatusCode::SERVICE_UNAVAILABLE) }),
)
.layer(axum::middleware::from_fn(retry_after))
}
#[tokio::test]
async fn a_503_gets_a_retry_after_header() {
let response = call(retry_after_router(), "/busy").await;
assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE);
assert_eq!(
response.headers().get(RETRY_AFTER).expect("header present"),
&HeaderValue::from(RETRY_AFTER_SECONDS)
);
}
#[tokio::test]
async fn a_successful_response_is_left_alone() {
let response = call(retry_after_router(), "/fine").await;
assert_eq!(response.status(), StatusCode::OK);
assert!(response.headers().get(RETRY_AFTER).is_none());
}
#[tokio::test]
async fn an_existing_retry_after_is_not_overwritten() {
let response = call(retry_after_router(), "/busy-with-hint").await;
assert_eq!(
response.headers().get(RETRY_AFTER).expect("header present"),
"30",
"a handler that knows a better value keeps it"
);
}
#[test]
fn rate_limiter_allows_up_to_capacity_then_blocks() {
let limiter = RateLimiter::new(3, 0.0); assert!(limiter.try_acquire());
assert!(limiter.try_acquire());
assert!(limiter.try_acquire());
assert!(
!limiter.try_acquire(),
"a 4th request within the same instant must be rejected"
);
}
#[test]
fn rate_limiter_refills_over_time() {
let limiter = RateLimiter::new(1, 1000.0); assert!(limiter.try_acquire());
assert!(!limiter.try_acquire());
std::thread::sleep(std::time::Duration::from_millis(5));
assert!(
limiter.try_acquire(),
"tokens must refill over time, not stay exhausted forever"
);
}
#[test]
fn per_minute_constructor_sets_a_full_minute_of_burst_capacity() {
let limiter = RateLimiter::per_minute(60);
for _ in 0..60 {
assert!(limiter.try_acquire());
}
assert!(!limiter.try_acquire());
}
}