diff --git a/src/mirror.rs b/src/mirror.rs index 8c4c692..f8921ab 100644 --- a/src/mirror.rs +++ b/src/mirror.rs @@ -1,5 +1,6 @@ use crate::{GovernorMiddleware, logo}; use futures::TryStreamExt; +use governor::Quota; use poem::{ EndpointExt, Error, IntoResponse, Request, Response, Result, Route, Server, get, handler, http::StatusCode, @@ -127,7 +128,9 @@ pub async fn serve(upstream: &Url, plc: Url, bind: SocketAddr) -> std::io::Resul .with(AddData::new(state)) .with(Cors::new().allow_credentials(false)) .with(Compression::new()) - .with(GovernorMiddleware::per_minute(3000).unwrap()) + .with(GovernorMiddleware::new(Quota::per_minute( + 3000.try_into().unwrap(), + ))) .with(CatchPanic::new()) .with(Tracing); diff --git a/src/ratelimit.rs b/src/ratelimit.rs index 2fb140b..3dd4b8f 100644 --- a/src/ratelimit.rs +++ b/src/ratelimit.rs @@ -1,82 +1,85 @@ use crate::logo; use governor::{ - Quota, RateLimiter, + NotUntil, Quota, RateLimiter, clock::{Clock, DefaultClock}, state::keyed::DefaultKeyedStateStore, }; use poem::{Endpoint, Middleware, Request, Response, Result, http::StatusCode}; use std::{ - convert::TryInto, error::Error, net::IpAddr, num::NonZeroU32, sync::Arc, sync::LazyLock, + convert::TryInto, + net::{IpAddr, Ipv6Addr}, + sync::{Arc, LazyLock}, time::Duration, }; static CLOCK: LazyLock = LazyLock::new(DefaultClock::default); +const IP6_64_MASK: Ipv6Addr = Ipv6Addr::from_bits(0xFFFF_FFFF_FFFF_FFFF_0000_0000_0000_0000); +type IP6_56 = [u8; 7]; +type IP6_48 = [u8; 6]; + +fn scale_quota(quota: Quota, factor: u32) -> Option { + let period = quota.replenish_interval() / factor; + let burst = quota + .burst_size() + .checked_mul(factor.try_into().unwrap()) + .unwrap(); + Quota::with_period(period).map(|q| q.allow_burst(burst)) +} + +#[derive(Debug)] +struct IpLimiters { + per_ip: RateLimiter, DefaultClock>, + ip6_56: RateLimiter, DefaultClock>, + ip6_48: RateLimiter, DefaultClock>, +} + +impl IpLimiters { + pub fn new(quota: Quota) -> Self { + Self { + per_ip: RateLimiter::keyed(quota), + ip6_56: RateLimiter::keyed(scale_quota(quota, 8).unwrap()), + ip6_48: RateLimiter::keyed(scale_quota(quota, 256).unwrap()), + } + } + pub fn check_key(&self, ip: IpAddr) -> Result<(), Duration> { + let asdf = |n: NotUntil<_>| n.wait_time_from(CLOCK.now()); + match ip { + addr @ IpAddr::V4(_) => self.per_ip.check_key(&addr).map_err(asdf), + IpAddr::V6(a) => { + // always check all limiters + let check_ip = self + .per_ip + .check_key(&IpAddr::V6(a & IP6_64_MASK)) + .map_err(asdf); + let check_56 = self + .ip6_56 + .check_key(a.octets()[..7].try_into().unwrap()) + .map_err(asdf); + let check_48 = self + .ip6_48 + .check_key(a.octets()[..6].try_into().unwrap()) + .map_err(asdf); + check_ip.and(check_56).and(check_48) + } + } + } +} + /// Once the rate limit has been reached, the middleware will respond with /// status code 429 (too many requests) and a `Retry-After` header with the amount /// of time that needs to pass before another request will be allowed. #[derive(Debug, Clone)] pub struct GovernorMiddleware { - limiter: Arc, DefaultClock>>, + limiters: Arc, } impl GovernorMiddleware { - /// Constructs a rate-limiting middleware from a [`Duration`] that allows one request in the given time interval. - /// - /// If the time interval is zero, returns `None`. - #[must_use] - pub fn with_period(duration: Duration) -> Option { - Some(Self { - limiter: Arc::new(RateLimiter::::keyed(Quota::with_period( - duration, - )?)), - }) - } - - /// Constructs a rate-limiting middleware that allows a specified number of requests every second. - /// - /// Returns an error if `times` can't be converted into a [`NonZeroU32`]. - pub fn per_second(times: T) -> Result - where - T: TryInto, - T::Error: Error + Send + Sync + 'static, - { - Ok(Self { - limiter: Arc::new(RateLimiter::::keyed(Quota::per_second( - times.try_into().unwrap(), // TODO - ))), - }) - } - - /// Constructs a rate-limiting middleware that allows a specified number of requests every minute. - /// - /// Returns an error if `times` can't be converted into a [`NonZeroU32`]. - pub fn per_minute(times: T) -> Result - where - T: TryInto, - T::Error: Error + Send + Sync + 'static, - { - Ok(Self { - limiter: Arc::new(RateLimiter::::keyed(Quota::per_minute( - times.try_into().unwrap(), // TODO - ))), - }) - } - - /// Constructs a rate-limiting middleware that allows a specified number of requests every hour. - /// - /// Returns an error if `times` can't be converted into a [`NonZeroU32`]. - pub fn per_hour(times: T) -> Result - where - T: TryInto, - T::Error: Error + Send + Sync + 'static, - { - Ok(Self { - limiter: Arc::new(RateLimiter::::keyed(Quota::per_hour( - times.try_into().unwrap(), // TODO - ))), - }) + pub fn new(quota: Quota) -> Self { + Self { + limiters: Arc::new(IpLimiters::new(quota)), + } } } @@ -85,14 +88,14 @@ impl Middleware for GovernorMiddleware { fn transform(&self, ep: E) -> Self::Output { GovernorMiddlewareImpl { ep, - limiter: self.limiter.clone(), + limiters: self.limiters.clone(), } } } pub struct GovernorMiddlewareImpl { ep: E, - limiter: Arc, DefaultClock>>, + limiters: Arc, } impl Endpoint for GovernorMiddlewareImpl { @@ -107,13 +110,13 @@ impl Endpoint for GovernorMiddlewareImpl { log::trace!("remote: {remote}"); - match self.limiter.check_key(&remote) { + match self.limiters.check_key(remote) { Ok(_) => { log::debug!("allowing remote {remote}"); self.ep.call(req).await } - Err(negative) => { - let wait_time = negative.wait_time_from(CLOCK.now()).as_secs(); + Err(d) => { + let wait_time = d.as_secs(); log::debug!("rate limit exceeded for {remote}, quota reset in {wait_time}s");