Skip to main content

tower_rate_limiter/limiter/
future.rs

1//! The unboxed request-execution state machine for the rate-limit service.
2
3use std::{
4    future::Future,
5    pin::Pin,
6    sync::Arc,
7    task::{Context, Poll, ready},
8};
9
10use http::{Request, Response};
11use pin_project_lite::pin_project;
12
13use super::{
14    LimitProvider, RateLimitConfig, RateLimitError,
15    policy::{Policy, ResponseMetadata, append_context, make_key},
16    response::{MiddlewareResponse, ResponseFactory, append_inner_response_headers},
17    store::{Store, StoreFailureMode},
18};
19
20pin_project! {
21    #[project = StateProj]
22    #[project_replace = StateProjReplace]
23    enum State<ReqBody, LimitFut, StoreFut, InnerFut> {
24        Limit {
25            #[pin]
26            future: LimitFut,
27            request: Request<ReqBody>,
28            key: String,
29        },
30        Store {
31            #[pin]
32            future: StoreFut,
33            request: Request<ReqBody>,
34            limit: u64,
35        },
36        Inner {
37            #[pin]
38            future: InnerFut,
39            metadata: Option<ResponseMetadata>,
40        },
41        Ready {
42            response: MiddlewareResponse<ReqBody>,
43        },
44        Done,
45    }
46}
47
48pin_project! {
49    /// Response future for [`super::RateLimit`].
50    pub struct ResponseFuture<ReqBody, Inner, S, P, F>
51        where
52        Inner: tower_service::Service<Request<ReqBody>>,
53        P: LimitProvider,
54        S: Store,
55    {
56        #[pin]
57        state: State<ReqBody, P::Future, S::Future, Inner::Future>,
58        inner: Inner,
59        store: S,
60        config: Arc<RateLimitConfig>,
61        factory: F,
62    }
63}
64
65impl<ReqBody, Inner, S, P, F> ResponseFuture<ReqBody, Inner, S, P, F>
66where
67    Inner: tower_service::Service<Request<ReqBody>>,
68    S: Store,
69    P: LimitProvider,
70{
71    pub(crate) fn new(
72        request: Request<ReqBody>,
73        inner: Inner,
74        store: S,
75        key: String,
76        future: P::Future,
77        config: Arc<RateLimitConfig>,
78        factory: F,
79    ) -> Self {
80        Self {
81            state: State::Limit { request, key, future },
82            inner,
83            store,
84            config,
85            factory,
86        }
87    }
88
89    pub(crate) fn error(
90        request: Request<ReqBody>,
91        error: RateLimitError,
92        inner: Inner,
93        store: S,
94        config: Arc<RateLimitConfig>,
95        factory: F,
96    ) -> Self {
97        Self {
98            state: State::Ready {
99                response: MiddlewareResponse::Error(request, error),
100            },
101            inner,
102            store,
103            config,
104            factory,
105        }
106    }
107
108    pub(crate) fn skipped(
109        request: Request<ReqBody>,
110        mut inner: Inner,
111        store: S,
112        config: Arc<RateLimitConfig>,
113        factory: F,
114    ) -> Self {
115        let future = inner.call(request);
116        Self {
117            state: State::Inner { future, metadata: None },
118            inner,
119            store,
120            config,
121            factory,
122        }
123    }
124}
125
126impl<ReqBody, ResBody, Inner, S, P, F> Future for ResponseFuture<ReqBody, Inner, S, P, F>
127where
128    Inner: tower_service::Service<Request<ReqBody>, Response = Response<ResBody>>,
129    S: Store,
130    P: LimitProvider,
131    F: ResponseFactory<ReqBody, ResBody>,
132{
133    type Output = Result<Response<ResBody>, Inner::Error>;
134
135    fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
136        let mut this = self.project();
137
138        loop {
139            match this.state.as_mut().project() {
140                StateProj::Limit { future, .. } => {
141                    let result = ready!(future.poll(cx));
142                    let StateProjReplace::Limit { request, key, .. } = this.state.as_mut().project_replace(State::Done)
143                    else {
144                        unreachable!("rate-limit future state changed while polling Limit")
145                    };
146
147                    match result {
148                        Err(error) => {
149                            let response = MiddlewareResponse::Error(request, error);
150                            this.state.as_mut().project_replace(State::Ready { response });
151                        },
152                        Ok(limit) => {
153                            let mut key = make_key(&this.config.policy_name, &key);
154                            if let Some(encoder) = this.config.key_encoder.as_ref() {
155                                key = encoder(&key);
156                            }
157                            let store_future = this.store.increment(&key, this.config.window);
158                            this.state.as_mut().project_replace(State::Store {
159                                request,
160                                future: store_future,
161                                limit,
162                            });
163                        },
164                    }
165                },
166                StateProj::Store { future, .. } => {
167                    let result = ready!(future.poll(cx));
168                    let StateProjReplace::Store { request, limit, .. } =
169                        this.state.as_mut().project_replace(State::Done)
170                    else {
171                        unreachable!("rate-limit future state changed while polling Store")
172                    };
173
174                    let policy = result.and_then(|usage| {
175                        Policy::from_usage(this.config.policy_name.clone(), this.config.window, limit, usage)
176                    });
177
178                    let next_state = match policy {
179                        Err(error) => {
180                            #[cfg(feature = "tracing")]
181                            trace_store_failure(
182                                &error,
183                                this.config.store_failure_mode,
184                                &this.config.policy_name,
185                                this.config.store_failure_tracing_level,
186                            );
187                            if this.config.store_failure_mode == StoreFailureMode::Allow {
188                                State::Inner {
189                                    future: this.inner.call(request),
190                                    metadata: None,
191                                }
192                            } else {
193                                State::Ready {
194                                    response: MiddlewareResponse::Error(request, error),
195                                }
196                            }
197                        },
198                        Ok(policy) => {
199                            let metadata = ResponseMetadata::new(policy, this.config.rate_limit_fields);
200                            if metadata.policy.is_rate_limited() {
201                                State::Ready {
202                                    response: MiddlewareResponse::RateLimited(request, metadata),
203                                }
204                            } else {
205                                let mut request = request;
206                                append_context(&mut request, &metadata);
207                                State::Inner {
208                                    future: this.inner.call(request),
209                                    metadata: Some(metadata),
210                                }
211                            }
212                        },
213                    };
214                    this.state.as_mut().project_replace(next_state);
215                },
216                StateProj::Inner { future, .. } => {
217                    let result = ready!(future.poll(cx));
218                    let StateProjReplace::Inner { metadata, .. } = this.state.as_mut().project_replace(State::Done)
219                    else {
220                        unreachable!("rate-limit future state changed while polling inner service")
221                    };
222                    return Poll::Ready(result.map(|response| append_inner_response_headers(response, metadata)));
223                },
224                StateProj::Ready { .. } => {
225                    let StateProjReplace::Ready { response } = this.state.as_mut().project_replace(State::Done) else {
226                        unreachable!("rate-limit future state changed while returning response")
227                    };
228                    return Poll::Ready(Ok(response.finalize(this.factory)));
229                },
230                StateProj::Done { .. } => {
231                    panic!("rate-limit response future polled after completion")
232                },
233            }
234        }
235    }
236}
237
238#[cfg(feature = "tracing")]
239fn trace_store_failure(
240    error: &RateLimitError,
241    failure_mode: StoreFailureMode,
242    policy_name: &str,
243    level: tracing::Level,
244) {
245    let failure_mode = match failure_mode {
246        StoreFailureMode::Reject => "reject",
247        StoreFailureMode::Allow => "allow",
248    };
249
250    macro_rules! emit {
251        ($level:expr) => {
252            tracing::event!(
253                target: "tower_rate_limiter::store",
254                $level,
255                event = "store_failure",
256                policy_name,
257                failure_mode,
258                error_code = error.code(),
259                "rate-limit Store failed"
260            )
261        };
262    }
263
264    match level {
265        tracing::Level::ERROR => emit!(tracing::Level::ERROR),
266        tracing::Level::WARN => emit!(tracing::Level::WARN),
267        tracing::Level::INFO => emit!(tracing::Level::INFO),
268        tracing::Level::DEBUG => emit!(tracing::Level::DEBUG),
269        tracing::Level::TRACE => emit!(tracing::Level::TRACE),
270    }
271}