use std::collections::HashMap;
use std::convert::Infallible;
use std::fmt;
use std::pin::Pin;
use std::sync::{Arc, Mutex};
use std::task::{Context, Poll};
use std::time::{Duration, Instant};
use axum::http::{HeaderName, HeaderValue, Request, Response, header};
use tower::{Layer, Service};
use crate::api::{Problem, ProblemKind};
pub const RATELIMIT_LIMIT: HeaderName = HeaderName::from_static("ratelimit-limit");
pub const RATELIMIT_REMAINING: HeaderName = HeaderName::from_static("ratelimit-remaining");
pub const RATELIMIT_RESET: HeaderName = HeaderName::from_static("ratelimit-reset");
pub const UNIDENTIFIED_KEY: &str = "unidentified";
const SWEEP_AT: usize = 8192;
#[derive(Clone)]
pub enum KeySource {
Ip,
Header(HeaderName),
Global,
Custom(KeyFn),
}
pub type KeyFn = Arc<dyn Fn(&Request<axum::body::Body>) -> Option<String> + Send + Sync>;
impl fmt::Debug for KeySource {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Ip => f.write_str("Ip"),
Self::Header(name) => write!(f, "Header({name})"),
Self::Global => f.write_str("Global"),
Self::Custom(_) => f.write_str("Custom(..)"),
}
}
}
impl KeySource {
fn key_for(&self, request: &Request<axum::body::Body>) -> String {
let resolved = match self {
Self::Ip => request
.extensions()
.get::<crate::http::ClientIp>()
.map(|client| client.addr().to_string())
.or_else(|| {
request
.extensions()
.get::<axum::extract::ConnectInfo<std::net::SocketAddr>>()
.map(|info| info.0.ip().to_canonical().to_string())
}),
Self::Header(name) => request
.headers()
.get(name)
.and_then(|value| value.to_str().ok())
.filter(|value| !value.is_empty())
.map(str::to_string),
Self::Global => Some(String::from("global")),
Self::Custom(f) => f(request).filter(|key| !key.is_empty()),
};
resolved.unwrap_or_else(|| UNIDENTIFIED_KEY.to_string())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum OnBackendError {
Refuse,
Allow,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Decision {
pub allowed: bool,
pub remaining: u32,
pub reset_after: u64,
pub retry_after: u64,
}
#[derive(Debug, Clone, Copy)]
struct Quota {
limit: u32,
capacity: f64,
refill_per_sec: f64,
}
impl Quota {
fn settle(self, tokens: f64) -> (f64, Decision) {
let (left, allowed) = if tokens >= 1.0 {
(tokens - 1.0, true)
} else {
(tokens, false)
};
(left, self.describe(left, allowed))
}
fn describe(self, tokens_left: f64, allowed: bool) -> Decision {
let deficit = (self.capacity - tokens_left).max(0.0);
Decision {
allowed,
remaining: tokens_left.max(0.0) as u32,
reset_after: seconds_to_accrue(deficit, self.refill_per_sec),
retry_after: if allowed {
0
} else {
seconds_to_accrue(1.0 - tokens_left, self.refill_per_sec).max(1)
},
}
}
}
fn seconds_to_accrue(tokens: f64, rate: f64) -> u64 {
if tokens <= 0.0 || rate <= 0.0 {
return 0;
}
(tokens / rate).ceil() as u64
}
#[derive(Debug, Clone, Copy)]
struct Bucket {
tokens: f64,
updated: Instant,
}
#[derive(Debug, Default)]
struct MemoryBuckets {
buckets: Mutex<HashMap<String, Bucket>>,
}
impl MemoryBuckets {
fn check(&self, key: &str, quota: Quota, now: Instant) -> Decision {
let mut buckets = match self.buckets.lock() {
Ok(guard) => guard,
Err(poisoned) => poisoned.into_inner(),
};
if buckets.len() >= SWEEP_AT {
buckets.retain(|_, bucket| {
refilled(
bucket.tokens,
quota,
now.saturating_duration_since(bucket.updated),
) < quota.capacity
});
}
let bucket = buckets.entry(key.to_string()).or_insert(Bucket {
tokens: quota.capacity,
updated: now,
});
let tokens = refilled(
bucket.tokens,
quota,
now.saturating_duration_since(bucket.updated),
);
let (left, decision) = quota.settle(tokens);
bucket.tokens = left;
bucket.updated = now;
decision
}
}
fn refilled(tokens: f64, quota: Quota, elapsed: Duration) -> f64 {
(tokens + elapsed.as_secs_f64() * quota.refill_per_sec).min(quota.capacity)
}
#[cfg(feature = "cache")]
const BUCKET_SCRIPT: &str = r"
local capacity = tonumber(ARGV[1])
local refill_per_ms = tonumber(ARGV[2])
local now_ms = tonumber(ARGV[3])
local ttl_ms = tonumber(ARGV[4])
local state = redis.call('HMGET', KEYS[1], 't', 'u')
local tokens = tonumber(state[1])
local updated = tonumber(state[2])
if tokens == nil or updated == nil then
tokens = capacity
updated = now_ms
end
local elapsed = now_ms - updated
if elapsed < 0 then elapsed = 0 end
tokens = math.min(capacity, tokens + elapsed * refill_per_ms)
local allowed = 0
if tokens >= 1 then
tokens = tokens - 1
allowed = 1
end
redis.call('HSET', KEYS[1], 't', tokens, 'u', now_ms)
redis.call('PEXPIRE', KEYS[1], ttl_ms)
return {allowed, math.floor(tokens * 1000)}
";
#[cfg(feature = "cache")]
struct RedisBuckets {
cache: crate::cache::Cache,
}
#[cfg(feature = "cache")]
impl fmt::Debug for RedisBuckets {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("RedisBuckets").finish_non_exhaustive()
}
}
#[cfg(feature = "cache")]
impl RedisBuckets {
fn new(cache: crate::cache::Cache) -> Self {
Self { cache }
}
async fn check(&self, key: &str, quota: Quota) -> Result<Decision, ()> {
let full_key = self.cache.resolve_key(&format!("ratelimit:{key}"));
let now_ms = unix_millis();
let ttl_ms = (((quota.capacity / quota.refill_per_sec) * 2000.0) as u64).max(1000);
let mut connection = self.cache.connection_for_op();
let outcome: Result<(i64, i64), _> = redis::cmd("EVAL")
.arg(BUCKET_SCRIPT)
.arg(1_i64)
.arg(full_key)
.arg(quota.capacity)
.arg(quota.refill_per_sec / 1000.0)
.arg(now_ms)
.arg(ttl_ms)
.query_async(&mut connection)
.await;
match outcome {
Ok((allowed, milli_tokens)) => {
Ok(quota.describe(milli_tokens as f64 / 1000.0, allowed == 1))
}
Err(_) => Err(()),
}
}
}
#[cfg(feature = "cache")]
fn unix_millis() -> u64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_millis() as u64)
.unwrap_or_default()
}
#[derive(Debug, Clone)]
enum Backend {
Memory(Arc<MemoryBuckets>),
#[cfg(feature = "cache")]
Redis(Arc<RedisBuckets>),
}
#[derive(Debug, Clone)]
pub struct RateLimit {
quota: Quota,
key: KeySource,
backend: Backend,
on_backend_error: OnBackendError,
}
impl RateLimit {
#[must_use]
pub fn new(limit: u32, window: Duration) -> Self {
let window = window.max(Duration::from_millis(1));
Self {
quota: Quota {
limit,
capacity: f64::from(limit),
refill_per_sec: f64::from(limit) / window.as_secs_f64(),
},
key: KeySource::Ip,
backend: Backend::Memory(Arc::new(MemoryBuckets::default())),
on_backend_error: OnBackendError::Refuse,
}
}
#[must_use]
pub fn per_second(limit: u32) -> Self {
Self::new(limit, Duration::from_secs(1))
}
#[must_use]
pub fn per_minute(limit: u32) -> Self {
Self::new(limit, Duration::from_secs(60))
}
#[must_use]
pub fn per_hour(limit: u32) -> Self {
Self::new(limit, Duration::from_secs(3600))
}
#[must_use]
pub fn burst(mut self, burst: u32) -> Self {
self.quota.capacity = f64::from(burst);
self
}
#[must_use]
pub fn by(mut self, key: KeySource) -> Self {
self.key = key;
self
}
#[must_use]
pub fn by_fn<F>(self, f: F) -> Self
where
F: Fn(&Request<axum::body::Body>) -> Option<String> + Send + Sync + 'static,
{
self.by(KeySource::Custom(Arc::new(f)))
}
#[cfg(feature = "cache")]
#[must_use]
pub fn redis(mut self, cache: crate::cache::Cache) -> Self {
self.backend = Backend::Redis(Arc::new(RedisBuckets::new(cache)));
self
}
#[must_use]
pub fn on_backend_error(mut self, behaviour: OnBackendError) -> Self {
self.on_backend_error = behaviour;
self
}
#[must_use]
pub fn limit(&self) -> u32 {
self.quota.limit
}
#[must_use]
pub fn capacity(&self) -> u32 {
self.quota.capacity as u32
}
#[must_use]
pub fn refill_per_second(&self) -> f64 {
self.quota.refill_per_sec
}
}
enum Checked {
Decided(Decision),
#[cfg(feature = "cache")]
BackendDown,
}
impl RateLimit {
async fn check(&self, key: &str) -> Checked {
match &self.backend {
Backend::Memory(buckets) => {
Checked::Decided(buckets.check(key, self.quota, Instant::now()))
}
#[cfg(feature = "cache")]
Backend::Redis(buckets) => match buckets.check(key, self.quota).await {
Ok(decision) => Checked::Decided(decision),
Err(()) => Checked::BackendDown,
},
}
}
}
impl<S> Layer<S> for RateLimit {
type Service = RateLimitService<S>;
fn layer(&self, inner: S) -> Self::Service {
RateLimitService {
inner,
limit: self.clone(),
}
}
}
#[derive(Debug, Clone)]
pub struct RateLimitService<S> {
inner: S,
limit: RateLimit,
}
impl<S> Service<Request<axum::body::Body>> for RateLimitService<S>
where
S: Service<
Request<axum::body::Body>,
Response = Response<axum::body::Body>,
Error = Infallible,
> + Clone
+ Send
+ 'static,
S::Future: Send + 'static,
{
type Response = Response<axum::body::Body>;
type Error = Infallible;
type Future =
Pin<Box<dyn std::future::Future<Output = Result<Self::Response, Self::Error>> + Send>>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, request: Request<axum::body::Body>) -> Self::Future {
let limit = self.limit.clone();
let key = limit.key.key_for(&request);
let clone = self.inner.clone();
let mut inner = std::mem::replace(&mut self.inner, clone);
Box::pin(async move {
#[allow(
clippy::infallible_destructuring_match,
reason = "the second arm exists under the `cache` feature"
)]
let decision = match limit.check(&key).await {
Checked::Decided(decision) => decision,
#[cfg(feature = "cache")]
Checked::BackendDown => match limit.on_backend_error {
OnBackendError::Refuse => return Ok(backend_unavailable()),
OnBackendError::Allow => return inner.call(request).await,
},
};
if !decision.allowed {
return Ok(refused(limit.quota.limit, decision));
}
let mut response = inner.call(request).await?;
annotate(response.headers_mut(), limit.quota.limit, decision);
Ok(response)
})
}
}
fn annotate(headers: &mut axum::http::HeaderMap, limit: u32, decision: Decision) {
for (name, value) in [
(RATELIMIT_LIMIT, u64::from(limit)),
(RATELIMIT_REMAINING, u64::from(decision.remaining)),
(RATELIMIT_RESET, decision.reset_after),
] {
if let Ok(value) = HeaderValue::from_str(&value.to_string()) {
headers.insert(name, value);
}
}
}
fn refused(limit: u32, decision: Decision) -> Response<axum::body::Body> {
use axum::response::IntoResponse as _;
let mut response = Problem::of(ProblemKind::RateLimit)
.with_detail("Too many requests. Slow down and try again shortly.")
.into_response();
annotate(response.headers_mut(), limit, decision);
if let Ok(value) = HeaderValue::from_str(&decision.retry_after.to_string()) {
response.headers_mut().insert(header::RETRY_AFTER, value);
}
response
}
#[cfg(feature = "cache")]
fn backend_unavailable() -> Response<axum::body::Body> {
use axum::response::IntoResponse as _;
Problem::of(ProblemKind::Unavailable)
.with_detail("The rate limiter is unavailable. Please try again shortly.")
.into_response()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_per_minute_limit_refills_at_one_per_second() {
let limit = RateLimit::per_minute(60);
assert_eq!(limit.limit(), 60);
assert_eq!(limit.capacity(), 60);
assert!((limit.refill_per_second() - 1.0).abs() < f64::EPSILON);
}
#[test]
fn burst_raises_the_capacity_without_raising_the_rate() {
let limit = RateLimit::per_minute(60).burst(120);
assert_eq!(limit.limit(), 60);
assert_eq!(limit.capacity(), 120);
assert!((limit.refill_per_second() - 1.0).abs() < f64::EPSILON);
}
#[test]
fn a_zero_window_does_not_divide_by_zero() {
let limit = RateLimit::new(10, Duration::ZERO);
assert!(limit.refill_per_second().is_finite());
}
#[test]
fn a_bucket_empties_and_then_refuses() {
let buckets = MemoryBuckets::default();
let quota = RateLimit::per_second(3).quota;
let now = Instant::now();
for expected_remaining in [2u32, 1, 0] {
let decision = buckets.check("k", quota, now);
assert!(decision.allowed);
assert_eq!(decision.remaining, expected_remaining);
}
let decision = buckets.check("k", quota, now);
assert!(!decision.allowed);
assert_eq!(decision.remaining, 0);
assert_eq!(decision.retry_after, 1);
}
#[test]
fn a_bucket_refills_over_time() {
let buckets = MemoryBuckets::default();
let quota = RateLimit::per_second(2).quota;
let now = Instant::now();
assert!(buckets.check("k", quota, now).allowed);
assert!(buckets.check("k", quota, now).allowed);
assert!(!buckets.check("k", quota, now).allowed);
let later = now + Duration::from_secs(1);
assert!(buckets.check("k", quota, later).allowed);
assert!(buckets.check("k", quota, later).allowed);
assert!(!buckets.check("k", quota, later).allowed);
}
#[test]
fn buckets_do_not_leak_across_keys() {
let buckets = MemoryBuckets::default();
let quota = RateLimit::per_second(1).quota;
let now = Instant::now();
assert!(buckets.check("a", quota, now).allowed);
assert!(!buckets.check("a", quota, now).allowed);
assert!(buckets.check("b", quota, now).allowed);
}
#[test]
fn a_zero_limit_refuses_everything() {
let buckets = MemoryBuckets::default();
let quota = RateLimit::new(0, Duration::from_secs(1)).quota;
let decision = buckets.check("k", quota, Instant::now());
assert!(!decision.allowed);
assert_eq!(decision.retry_after, 1);
}
#[test]
fn reset_counts_down_to_a_full_bucket() {
let quota = RateLimit::per_second(10).quota;
let (_, decision) = quota.settle(10.0);
assert!(decision.allowed);
assert_eq!(decision.remaining, 9);
assert_eq!(decision.reset_after, 1);
assert_eq!(decision.retry_after, 0);
}
#[test]
fn an_unidentified_request_gets_the_shared_bucket() {
let request = Request::builder()
.uri("/")
.body(axum::body::Body::empty())
.expect("request builds");
assert_eq!(KeySource::Ip.key_for(&request), UNIDENTIFIED_KEY);
assert_eq!(
KeySource::Header(HeaderName::from_static("x-api-key")).key_for(&request),
UNIDENTIFIED_KEY
);
assert_eq!(KeySource::Global.key_for(&request), "global");
}
#[test]
fn a_header_key_source_reads_the_header() {
let request = Request::builder()
.uri("/")
.header("x-api-key", "abc")
.body(axum::body::Body::empty())
.expect("request builds");
assert_eq!(
KeySource::Header(HeaderName::from_static("x-api-key")).key_for(&request),
"abc"
);
}
fn served(peer: &str, forwarded: Option<&str>, trusted: &str) -> Request<axum::body::Body> {
let peer: std::net::SocketAddr = peer.parse().expect("a literal peer address");
let trusted: crate::http::TrustedProxies = trusted.parse().expect("a literal proxy list");
let mut builder = Request::builder().uri("/");
if let Some(forwarded) = forwarded {
builder = builder.header(crate::http::X_FORWARDED_FOR, forwarded);
}
let mut request = builder
.body(axum::body::Body::empty())
.expect("request builds");
let client = crate::http::ClientIp::resolve(peer.ip(), request.headers(), &trusted);
let extensions = request.extensions_mut();
extensions.insert(axum::extract::ConnectInfo(peer));
extensions.insert(client);
request
}
#[test]
fn two_addresses_get_two_buckets() {
let one = served("203.0.113.7:40000", None, "");
let two = served("203.0.113.8:40000", None, "");
assert_eq!(KeySource::Ip.key_for(&one), "203.0.113.7");
assert_eq!(KeySource::Ip.key_for(&two), "203.0.113.8");
let buckets = MemoryBuckets::default();
let quota = RateLimit::per_second(1).quota;
let now = Instant::now();
assert!(
buckets
.check(&KeySource::Ip.key_for(&one), quota, now)
.allowed
);
assert!(
!buckets
.check(&KeySource::Ip.key_for(&one), quota, now)
.allowed
);
assert!(
buckets
.check(&KeySource::Ip.key_for(&two), quota, now)
.allowed
);
}
#[test]
fn a_forged_forwarded_header_from_an_untrusted_peer_is_ignored() {
let request = served("203.0.113.7:40000", Some("198.51.100.23"), "");
assert_eq!(KeySource::Ip.key_for(&request), "203.0.113.7");
let rotated = served("203.0.113.7:40000", Some("198.51.100.24"), "");
assert_eq!(
KeySource::Ip.key_for(&request),
KeySource::Ip.key_for(&rotated)
);
}
#[test]
fn a_forwarded_header_from_a_trusted_peer_is_believed() {
let request = served("10.0.0.4:40000", Some("198.51.100.23"), "10.0.0.0/8");
assert_eq!(KeySource::Ip.key_for(&request), "198.51.100.23");
}
#[test]
fn the_peer_address_is_the_fallback_when_nothing_resolved_a_client() {
let peer: std::net::SocketAddr = "203.0.113.7:40000".parse().expect("a literal address");
let mut request = Request::builder()
.uri("/")
.body(axum::body::Body::empty())
.expect("request builds");
request
.extensions_mut()
.insert(axum::extract::ConnectInfo(peer));
assert_eq!(KeySource::Ip.key_for(&request), "203.0.113.7");
}
}