Skip to main content

tower_rate_limiter/limiter/
service.rs

1//! Tower service implementation.
2
3use 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/// Tower service produced by [`super::RateLimitLayer`].
19#[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    /// Borrow the wrapped service.
32    pub const fn get_ref(&self) -> &Inner {
33        &self.inner
34    }
35
36    /// Mutably borrow the wrapped service.
37    pub fn get_mut(&mut self) -> &mut Inner {
38        &mut self.inner
39    }
40
41    /// Consume this middleware and return the wrapped service.
42    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        // Check if the request should be skipped.
68        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        // Extract the key.
82        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}