#![doc = include_str!("../README.md")]
#![deny(unsafe_code, missing_docs, clippy::unwrap_used)]
mod key;
use axum_core::extract::{FromRef, FromRequestParts};
use axum_core::response::{IntoResponse, Response};
use dashmap::DashMap;
use http::request::Parts;
use http::StatusCode;
use std::error::Error;
use std::fmt::Display;
use std::hash::Hash;
use std::ops::{Deref, DerefMut};
use std::sync::Arc;
use std::time::{Duration, Instant};
#[derive(Debug, Clone, Copy, Default)]
pub struct Limit<const COUNT: usize, const PER: u64, K>(pub K::Extractor)
where
K: Key;
pub type LimitPerSecond<const COUNT: usize, K> = Limit<COUNT, 1000, K>;
pub type LimitPerMinute<const COUNT: usize, K> = Limit<COUNT, 60_000, K>;
pub type LimitPerHour<const COUNT: usize, K> = Limit<COUNT, 3_600_000, K>;
pub type LimitPerDay<const COUNT: usize, K> = Limit<COUNT, 86_400_000, K>;
impl<const COUNT: usize, const PER: u64, K> AsRef<K::Extractor> for Limit<COUNT, PER, K>
where
K: Key,
{
fn as_ref(&self) -> &K::Extractor {
&self.0
}
}
impl<const COUNT: usize, const PER: u64, K> AsMut<K::Extractor> for Limit<COUNT, PER, K>
where
K: Key,
{
fn as_mut(&mut self) -> &mut K::Extractor {
&mut self.0
}
}
impl<const COUNT: usize, const PER: u64, K> Deref for Limit<COUNT, PER, K>
where
K: Key,
{
type Target = K::Extractor;
fn deref(&self) -> &Self::Target {
&self.0
}
}
impl<const COUNT: usize, const PER: u64, K> DerefMut for Limit<COUNT, PER, K>
where
K: Key,
{
fn deref_mut(&mut self) -> &mut Self::Target {
&mut self.0
}
}
impl<const COUNT: usize, const PER: u64, K> Display for Limit<COUNT, PER, K>
where
K: Key,
K::Extractor: Display,
{
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
self.0.fmt(f)
}
}
impl<const COUNT: usize, const PER: u64, K> Limit<COUNT, PER, K>
where
K: Key,
{
pub const fn count() -> usize {
COUNT
}
pub const fn per() -> u64 {
PER
}
pub fn into_inner(self) -> K::Extractor {
self.0
}
}
#[async_trait::async_trait]
pub trait Key: Eq + Hash + Send + Sync {
type Extractor;
fn from_extractor(extractor: &Self::Extractor) -> Self;
}
struct TokenBucket {
tokens: usize,
last_refill_time: Instant,
refill_duration: Duration,
}
impl TokenBucket {
fn new(tokens: impl Into<usize>, per: impl Into<u64>) -> Self {
Self {
tokens: tokens.into(),
last_refill_time: Instant::now(),
refill_duration: Duration::from_secs(per.into()),
}
}
fn try_acquire(&mut self) -> bool {
self.refill();
if self.tokens > 0 {
self.tokens -= 1;
true
} else {
false
}
}
fn refill(&mut self) {
let now = Instant::now();
let elapsed = now.duration_since(self.last_refill_time);
if elapsed >= self.refill_duration {
let new_tokens = (elapsed.as_secs() / self.refill_duration.as_secs()) as usize;
self.tokens += new_tokens;
self.last_refill_time =
now - Duration::from_secs(elapsed.as_secs() % self.refill_duration.as_secs());
}
}
}
#[derive(Clone)]
pub struct LimitState<K>
where
K: Key,
{
rate_limits: Arc<DashMap<K, TokenBucket>>,
}
impl<K> Default for LimitState<K>
where
K: Key,
{
fn default() -> Self {
Self {
rate_limits: Arc::new(DashMap::new()),
}
}
}
impl<K> LimitState<K>
where
K: Key,
{
pub fn check(&self, key: K, count: usize, per: u64) -> bool {
let mut bucket = self
.rate_limits
.entry(key)
.or_insert_with(|| TokenBucket::new(count, per));
bucket.try_acquire()
}
}
#[async_trait::async_trait]
impl<const C: usize, const P: u64, K, S> FromRequestParts<S> for Limit<C, P, K>
where
LimitState<K>: FromRef<S>,
S: Send + Sync,
K: Key,
K::Extractor: FromRequestParts<S>,
{
type Rejection = LimitRejection<<<K as Key>::Extractor as FromRequestParts<S>>::Rejection>;
async fn from_request_parts(parts: &mut Parts, state: &S) -> Result<Self, Self::Rejection> {
let key_extractor = match K::Extractor::from_request_parts(parts, state).await {
Ok(ke) => ke,
Err(rejection) => return Err(LimitRejection::KeyExtractionFailure(rejection)),
};
let limit_state: LimitState<K> = FromRef::from_ref(state);
let key = K::from_extractor(&key_extractor);
if limit_state.check(key, C, P) {
Ok(Self(key_extractor))
} else {
Err(LimitRejection::RateLimitExceeded)
}
}
}
#[derive(Debug)]
pub enum LimitRejection<R> {
KeyExtractionFailure(R),
RateLimitExceeded,
}
impl<R: Display> Display for LimitRejection<R> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
LimitRejection::KeyExtractionFailure(r) => write!(f, "{r}"),
LimitRejection::RateLimitExceeded => write!(f, "Rate limit exceeded."),
}
}
}
impl<R: Error + 'static> Error for LimitRejection<R> {
fn source(&self) -> Option<&(dyn Error + 'static)> {
match self {
LimitRejection::KeyExtractionFailure(ve) => Some(ve),
LimitRejection::RateLimitExceeded => None,
}
}
}
impl<R: IntoResponse> IntoResponse for LimitRejection<R> {
fn into_response(self) -> Response {
match self {
LimitRejection::KeyExtractionFailure(rejection) => rejection.into_response(),
LimitRejection::RateLimitExceeded => {
(StatusCode::TOO_MANY_REQUESTS, "Rate limit exceeded.").into_response()
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use axum::routing::get;
use axum::Router;
use axum_test::TestServer;
use http::Uri;
use std::future::IntoFuture;
#[tokio::test]
async fn limit() {
const TEST_ROUTE0: &str = "/limit0";
const TEST_ROUTE1: &str = "/limit1";
async fn handler0(Limit(_uri): Limit<1, 1, Uri>) -> impl IntoResponse {}
async fn handler1(Limit(_uri): Limit<3, 1, Uri>) -> impl IntoResponse {}
let my_app = Router::new()
.route(TEST_ROUTE0, get(handler0))
.route(TEST_ROUTE1, get(handler1))
.with_state(LimitState::default());
let server = TestServer::new(my_app).expect("Failed to create test server");
let response = server.get(TEST_ROUTE0).await;
assert_eq!(response.status_code(), StatusCode::OK);
let response = server.get(TEST_ROUTE0).await;
assert_eq!(response.status_code(), StatusCode::TOO_MANY_REQUESTS);
tokio::time::sleep(Duration::from_secs(1)).await;
let response = server.get(TEST_ROUTE0).await;
assert_eq!(response.status_code(), StatusCode::OK);
let gets = vec![
server.get(TEST_ROUTE1).into_future(),
server.get(TEST_ROUTE1).into_future(),
server.get(TEST_ROUTE1).into_future(),
];
let resp = futures::future::join_all(gets).await;
assert!(!resp.iter().any(|r| !r.status_code().is_success()));
assert_eq!(
server.get(TEST_ROUTE1).await.status_code(),
StatusCode::TOO_MANY_REQUESTS
);
tokio::time::sleep(Duration::from_secs(1)).await;
let response = server.get(TEST_ROUTE1).await;
assert_eq!(response.status_code(), StatusCode::OK);
}
}