1use 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 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}