use std::{
collections::HashMap,
net::Ipv4Addr,
sync::Mutex,
time::{Duration, Instant},
};
use kynos::{
Router,
http::forwarded::TrustedProxies,
middleware::rate_limit::{
RateLimit,
decision::{Decision, QuotaPolicy, QuotaUnit, RateLimitPolicy, ServiceLimit},
key::{And, ByClientAddress, ByRoute, RateLimitKey},
},
prelude::*,
response::status::NoContent,
router::operation::Route,
server::Server,
};
#[allow(clippy::cast_possible_truncation, clippy::cast_sign_loss)]
fn whole_tokens(tokens: f64) -> u64 {
tokens as u64
}
trait Clock: Send + Sync + 'static {
fn now(&self) -> Instant;
}
struct SystemClock;
impl Clock for SystemClock {
fn now(&self) -> Instant {
Instant::now()
}
}
#[derive(Clone, Copy)]
struct Bucket {
tokens: f64,
at: Instant,
}
struct Buckets {
per_key: HashMap<String, Bucket>,
swept_at: Instant,
}
struct TokenBucket<K, C: Clock> {
capacity: f64,
refill_per_second: f64,
idle_before_full: Duration,
exempt: &'static [&'static str],
key: K,
buckets: Mutex<Buckets>,
advertised: Vec<QuotaPolicy>,
clock: C,
}
impl<K, C: Clock> TokenBucket<K, C> {
fn new(capacity: u32, refill_per_second: f64, key: K, clock: C) -> Self {
let idle_before_full = Duration::from_secs_f64(f64::from(capacity) / refill_per_second);
let swept_at = clock.now();
Self {
capacity: f64::from(capacity),
refill_per_second,
idle_before_full,
exempt: &[],
key,
buckets: Mutex::new(Buckets {
per_key: HashMap::new(),
swept_at,
}),
advertised: vec![QuotaPolicy {
name: "burst".into(),
quota: u64::from(capacity),
window: Some(idle_before_full),
unit: QuotaUnit::Requests,
}],
clock,
}
}
fn exempting(mut self, operations: &'static [&'static str]) -> Self {
self.exempt = operations;
self
}
#[allow(clippy::cast_possible_truncation, clippy::cast_sign_loss)]
fn ceiling(&self) -> u64 {
self.capacity as u64
}
fn full(&self) -> ServiceLimit {
ServiceLimit {
name: "burst".into(),
quota: self.ceiling(),
remaining: self.ceiling(),
reset: Duration::ZERO,
}
}
fn spend(&self, key: &str) -> Result<ServiceLimit, (Duration, ServiceLimit)> {
let now = self.clock.now();
let mut buckets = self.buckets.lock().expect("no holder of this lock panics");
if now.duration_since(buckets.swept_at) >= self.idle_before_full {
let idle_limit = self.idle_before_full.saturating_mul(2);
buckets
.per_key
.retain(|_, bucket| now.duration_since(bucket.at) < idle_limit);
buckets.swept_at = now;
}
let bucket = buckets.per_key.entry(key.to_owned()).or_insert(Bucket {
tokens: self.capacity,
at: now,
});
let elapsed = now.duration_since(bucket.at).as_secs_f64();
bucket.tokens = (bucket.tokens + elapsed * self.refill_per_second).min(self.capacity);
bucket.at = now;
let ceiling = self.ceiling();
if bucket.tokens < 1.0 {
let wait = Duration::from_secs_f64((1.0 - bucket.tokens) / self.refill_per_second);
return Err((
wait,
ServiceLimit {
name: "burst".into(),
quota: ceiling,
remaining: 0,
reset: wait,
},
));
}
bucket.tokens -= 1.0;
let remaining = whole_tokens(bucket.tokens);
Ok(ServiceLimit {
name: "burst".into(),
quota: ceiling,
remaining,
reset: Duration::from_secs_f64(
(self.capacity - bucket.tokens) / self.refill_per_second,
),
})
}
}
impl<Ctx: Sync + 'static, K: RateLimitKey<Ctx>, C: Clock> RateLimitPolicy<Ctx>
for TokenBucket<K, C>
{
fn advertised(&self) -> &[QuotaPolicy] {
&self.advertised
}
async fn check(
&self,
request: &kynos::http::Request,
route: Route<'_>,
context: &Ctx,
) -> Decision {
if self.exempt.contains(&route.operation_id()) {
return Decision::allow(self.full());
}
let Some(key) = self.key.partition(request, route, context) else {
return Decision::allow(self.full());
};
match self.spend(&key) {
Ok(limit) => Decision::allow(limit),
Err((wait, limit)) => Decision::deny(wait, limit),
}
}
}
#[kynos::get("/reports")]
async fn reports() -> NoContent {
NoContent
}
#[kynos::get("/healthz")]
async fn healthz() -> NoContent {
NoContent
}
#[tokio::main]
async fn main() -> kynos::Result<()> {
let router = Router::<()>::new()
.trusted_proxies(TrustedProxies::hops(1))
.intercept(
RateLimit::new(
TokenBucket::new(10, 5.0, And(ByClientAddress, ByRoute), SystemClock)
.exempting(&["healthz"]),
)
.standard_fields(),
)
.mount(kynos::routes![reports, healthz]);
println!("{}", router.openapi()?.to_json()?);
Server::new(router.build(())?)
.bind((Ipv4Addr::UNSPECIFIED, 3000))
.serve()
.await
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use super::*;
#[derive(Clone)]
struct TestClock(Arc<Mutex<Instant>>);
impl TestClock {
fn new() -> Self {
Self(Arc::new(Mutex::new(Instant::now())))
}
fn advance(&self, step: Duration) {
let mut at = self.0.lock().expect("no holder of this lock panics");
*at += step;
}
}
impl Clock for TestClock {
fn now(&self) -> Instant {
*self.0.lock().expect("no holder of this lock panics")
}
}
impl<K, C: Clock> TokenBucket<K, C> {
fn tracked(&self) -> usize {
self.buckets
.lock()
.expect("no holder of this lock panics")
.per_key
.len()
}
}
fn bucket(
capacity: u32,
refill_per_second: f64,
clock: &TestClock,
) -> TokenBucket<(), TestClock> {
TokenBucket::new(capacity, refill_per_second, (), clock.clone())
}
#[test]
fn a_burst_is_spent_down_to_the_refusal() {
let clock = TestClock::new();
let limiter = bucket(3, 1.0, &clock);
for expected in [2, 1, 0] {
let limit = limiter.spend("client").expect("the burst is unspent");
assert_eq!(limit.quota, 3);
assert_eq!(limit.remaining, expected);
}
let (_, limit) = limiter.spend("client").expect_err("the burst is spent");
assert_eq!(limit.remaining, 0);
}
#[test]
fn an_empty_bucket_refills_with_no_test_sleeping() {
let clock = TestClock::new();
let limiter = bucket(2, 2.0, &clock);
limiter.spend("client").expect("the burst is unspent");
limiter.spend("client").expect("the burst is unspent");
limiter.spend("client").expect_err("the burst is spent");
clock.advance(Duration::from_millis(500));
let limit = limiter.spend("client").expect("a token has accrued");
assert_eq!(limit.remaining, 0);
}
#[test]
fn the_wait_is_one_token_rather_than_the_window() {
let clock = TestClock::new();
let limiter = bucket(4, 1.0, &clock);
for _ in 0..4 {
limiter.spend("client").expect("the burst is unspent");
}
let (wait, limit) = limiter.spend("client").expect_err("the burst is spent");
assert_eq!(limiter.advertised[0].window, Some(Duration::from_secs(4)));
assert_eq!(wait, Duration::from_secs(1));
assert_eq!(limit.reset, wait);
}
#[test]
fn a_bucket_idle_past_the_threshold_is_dropped() {
let clock = TestClock::new();
let limiter = bucket(2, 1.0, &clock);
limiter.spend("first").expect("the burst is unspent");
assert_eq!(limiter.tracked(), 1);
clock.advance(Duration::from_secs(5));
limiter.spend("second").expect("a fresh bucket is full");
assert_eq!(limiter.tracked(), 1);
}
#[test]
fn a_bucket_inside_the_threshold_survives_the_sweep() {
let clock = TestClock::new();
let limiter = bucket(2, 1.0, &clock);
limiter.spend("first").expect("the burst is unspent");
clock.advance(Duration::from_secs(3));
limiter.spend("second").expect("a fresh bucket is full");
assert_eq!(limiter.tracked(), 2);
}
#[test]
fn a_bucket_outlives_its_threshold_by_at_most_one_sweep() {
let clock = TestClock::new();
let limiter = bucket(2, 1.0, &clock);
limiter.spend("first").expect("the burst is unspent");
clock.advance(Duration::from_secs(3));
limiter.spend("second").expect("a fresh bucket is full");
clock.advance(Duration::from_millis(1500));
limiter.spend("second").expect("a token has accrued");
assert_eq!(limiter.tracked(), 2);
clock.advance(Duration::from_millis(500));
limiter.spend("second").expect("a token has accrued");
assert_eq!(limiter.tracked(), 1);
}
}