use std::sync::Arc;
use axum::{Router, http::Method, routing::MethodRouter};
use crate::auth::guard::login_required_middleware;
use crate::macros::routeur::register_url::register_pending;
use crate::middleware::rate_limit::{RateLimiter, rate_limit_middleware};
pub trait RouterExt {
fn login_required(
self,
path: impl Into<String>,
name: impl Into<String>,
handler: MethodRouter,
redirect_url: impl Into<String>,
) -> Self;
fn rate_limit(
self,
path: impl Into<String>,
name: impl Into<String>,
handler: MethodRouter,
max_requests: u32,
retry_after: u64,
methods: Vec<Method>,
) -> Self;
fn rate_limit_many(
self,
max_requests: u32,
retry_after: u64,
methods: Vec<Method>,
routes: Vec<(String, String, MethodRouter)>,
) -> Self;
}
impl RouterExt for Router {
fn login_required(
self,
path: impl Into<String>,
name: impl Into<String>,
handler: MethodRouter,
redirect_url: impl Into<String>,
) -> Self {
let path = path.into();
let name = name.into();
let redirect = Arc::new(redirect_url.into());
register_pending(&name, &path);
let protected =
Router::new()
.route(&path, handler)
.route_layer(axum::middleware::from_fn_with_state(
redirect,
login_required_middleware,
));
self.merge(protected)
}
fn rate_limit(
self,
path: impl Into<String>,
name: impl Into<String>,
handler: MethodRouter,
max_requests: u32,
retry_after: u64,
methods: Vec<Method>,
) -> Self {
self.rate_limit_many(
max_requests,
retry_after,
methods,
vec![(path.into(), name.into(), handler)],
)
}
fn rate_limit_many(
self,
max_requests: u32,
retry_after: u64,
methods: Vec<Method>,
routes: Vec<(String, String, MethodRouter)>,
) -> Self {
let mut limiter = RateLimiter::new()
.max_requests(max_requests)
.retry_after(retry_after);
if !methods.is_empty() {
limiter = limiter.only_methods(methods);
}
let limiter = Arc::new(limiter);
limiter.spawn_cleanup(tokio::time::Duration::from_secs(retry_after));
let mut r = self;
for (path, name, handler) in routes {
register_pending(&name, &path);
let limited = Router::new().route(&path, handler).route_layer(
axum::middleware::from_fn_with_state(limiter.clone(), rate_limit_middleware),
);
r = r.merge(limited);
}
r
}
}