tower_rate_limiter/limiter/
service.rs1use std::sync::Arc;
4use std::task::{Context, Poll};
5
6use http::{Request, Response};
7use tower_service::Service;
8
9use super::{
10 builder::{RateLimitConfig, check_skip_predicate},
11 future::ResponseFuture,
12 key_extractor::KeyExtractor,
13 limit::LimitProvider,
14 response::ResponseFactory,
15 store::Store,
16};
17
18#[must_use]
20#[derive(Clone, Debug)]
21pub struct RateLimit<Inner, K, S, P, F> {
22 pub(crate) inner: Inner,
23 pub(crate) key_extractor: K,
24 pub(crate) store: S,
25 pub(crate) limit_provider: P,
26 pub(crate) response_factory: F,
27 pub(crate) config: Arc<RateLimitConfig>,
28}
29
30impl<Inner, K, S, P, F> RateLimit<Inner, K, S, P, F> {
31 pub const fn get_ref(&self) -> &Inner {
33 &self.inner
34 }
35
36 pub fn get_mut(&mut self) -> &mut Inner {
38 &mut self.inner
39 }
40
41 pub fn into_inner(self) -> Inner {
43 self.inner
44 }
45}
46
47impl<Inner, K, S, P, F, ReqBody, ResBody> tower_service::Service<Request<ReqBody>> for RateLimit<Inner, K, S, P, F>
48where
49 Inner: Service<Request<ReqBody>, Response = Response<ResBody>> + Clone,
50 K: KeyExtractor,
51 S: Store,
52 P: LimitProvider,
53 F: ResponseFactory<ReqBody, ResBody>,
54{
55 type Response = Inner::Response;
56 type Error = Inner::Error;
57 type Future = ResponseFuture<ReqBody, Inner, S, P, F>;
58
59 fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
60 self.inner.poll_ready(cx)
61 }
62
63 fn call(&mut self, request: Request<ReqBody>) -> Self::Future {
64 let replacement = self.inner.clone();
65 let inner = std::mem::replace(&mut self.inner, replacement);
66
67 let request = match check_skip_predicate(self.config.skip_predicate.as_ref(), request) {
69 (true, request) => {
70 return ResponseFuture::skipped(
71 request,
72 inner,
73 self.store.clone(),
74 Arc::clone(&self.config),
75 self.response_factory.clone(),
76 );
77 },
78 (false, request) => request,
79 };
80
81 let key = match self.key_extractor.extract(&request) {
83 Ok(key) => key,
84 Err(error) => {
85 return ResponseFuture::error(
86 request,
87 error,
88 inner,
89 self.store.clone(),
90 Arc::clone(&self.config),
91 self.response_factory.clone(),
92 );
93 },
94 };
95
96 let limit_future = self.limit_provider.limit(&request);
97 ResponseFuture::new(
98 request,
99 inner,
100 self.store.clone(),
101 key.to_string(),
102 limit_future,
103 Arc::clone(&self.config),
104 self.response_factory.clone(),
105 )
106 }
107}