cloudillo_core/rate_limit/
middleware.rs1use std::sync::Arc;
9use std::task::{Context, Poll};
10
11use axum::body::Body;
12use axum::response::IntoResponse;
13use futures::future::BoxFuture;
14use hyper::Request;
15use tower::{Layer, Service};
16
17use super::extractors::extract_client_ip;
18use super::limiter::RateLimitManager;
19use crate::app::ServerMode;
20
21#[derive(Clone)]
23pub struct RateLimitLayer {
24 manager: Arc<RateLimitManager>,
25 category: &'static str,
26 mode: ServerMode,
27 skip_ban: bool,
28}
29
30impl RateLimitLayer {
31 pub fn new(manager: Arc<RateLimitManager>, category: &'static str, mode: ServerMode) -> Self {
33 Self { manager, category, mode, skip_ban: false }
34 }
35
36 pub fn new_skip_ban(
42 manager: Arc<RateLimitManager>,
43 category: &'static str,
44 mode: ServerMode,
45 ) -> Self {
46 Self { manager, category, mode, skip_ban: true }
47 }
48}
49
50impl<S> Layer<S> for RateLimitLayer {
51 type Service = RateLimitService<S>;
52
53 fn layer(&self, inner: S) -> Self::Service {
54 RateLimitService {
55 inner,
56 manager: self.manager.clone(),
57 category: self.category,
58 mode: self.mode,
59 skip_ban: self.skip_ban,
60 }
61 }
62}
63
64#[derive(Clone)]
66pub struct RateLimitService<S> {
67 inner: S,
68 manager: Arc<RateLimitManager>,
69 category: &'static str,
70 mode: ServerMode,
71 skip_ban: bool,
72}
73
74impl<S> Service<Request<Body>> for RateLimitService<S>
75where
76 S: Service<Request<Body>, Response = axum::response::Response> + Clone + Send + 'static,
77 S::Future: Send + 'static,
78{
79 type Response = S::Response;
80 type Error = S::Error;
81 type Future = BoxFuture<'static, Result<Self::Response, Self::Error>>;
82
83 fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
84 self.inner.poll_ready(cx)
85 }
86
87 fn call(&mut self, req: Request<Body>) -> Self::Future {
88 let manager = self.manager.clone();
89 let category = self.category;
90 let mode = self.mode;
91 let skip_ban = self.skip_ban;
92 let mut inner = self.inner.clone();
93
94 Box::pin(async move {
95 let client_ip = extract_client_ip(&req, &mode);
97
98 if let Some(ip) = client_ip {
99 let result = if skip_ban {
101 manager.check_skip_ban(&ip, category)
102 } else {
103 manager.check(&ip, category)
104 };
105 if let Err(error) = result {
106 return Ok(error.into_response());
108 }
109 }
110
111 inner.call(req).await
113 })
114 }
115}
116
117