tower_rate_limiter/limiter/
builder.rs1use std::{fmt, sync::Arc, time::Duration};
4
5use http::Request;
6
7use super::{
8 error::ConfigError,
9 layer::RateLimitLayer,
10 limit::LimitProvider,
11 response::{DefaultResponseFactory, RateLimitFields},
12 store::{Store, StoreFailureMode},
13};
14
15const MINIMUM_WINDOW: Duration = Duration::from_millis(1);
17
18pub(crate) type KeyEncoder = Box<dyn Fn(&str) -> String + Send + Sync>;
20
21pub(crate) type SkipPredicate = Box<dyn Fn(&Request<()>) -> bool + Send + Sync>;
23
24pub(crate) fn check_skip_predicate<B>(predicate: Option<&SkipPredicate>, request: Request<B>) -> (bool, Request<B>) {
26 let Some(predicate) = predicate else {
27 return (false, request);
28 };
29
30 let (parts, body) = request.into_parts();
32 let request_head = Request::from_parts(parts, ());
33 let should_skip = predicate(&request_head);
34 let (parts, ()) = request_head.into_parts();
35
36 (should_skip, Request::from_parts(parts, body))
37}
38
39#[derive(Debug)]
41#[must_use]
42pub struct RateLimitBuilder<K, S = (), P = u64, F = DefaultResponseFactory> {
43 key_extractor: K,
44 store: S,
45 limit_provider: P,
46 response_factory: F,
47 config: RateLimitConfig,
48}
49
50impl<K> RateLimitBuilder<K> {
51 pub(crate) fn new(key_extractor: K) -> Self {
52 Self {
53 key_extractor,
54 store: (),
55 limit_provider: 1,
56 response_factory: DefaultResponseFactory,
57 config: RateLimitConfig {
58 policy_name: String::from("default-policy"),
59 window: Duration::from_secs(60),
60 key_encoder: None,
61 skip_predicate: None,
62 store_failure_mode: StoreFailureMode::default(),
63 #[cfg(feature = "tracing")]
64 store_failure_tracing_level: tracing::Level::WARN,
65 rate_limit_fields: RateLimitFields::default(),
66 },
67 }
68 }
69}
70
71impl<K, S, P, F> RateLimitBuilder<K, S, P, F> {
72 pub fn with_store<S2>(self, store: S2) -> RateLimitBuilder<K, S2, P, F> {
74 let Self {
75 key_extractor,
76 limit_provider,
77 response_factory,
78 config,
79 ..
80 } = self;
81 RateLimitBuilder {
82 key_extractor,
83 store,
84 limit_provider,
85 response_factory,
86 config,
87 }
88 }
89
90 pub fn limit(self, limit: u64) -> RateLimitBuilder<K, S, u64, F> {
92 let Self {
93 key_extractor,
94 store,
95 response_factory,
96 config,
97 ..
98 } = self;
99 RateLimitBuilder {
100 key_extractor,
101 store,
102 limit_provider: limit,
103 response_factory,
104 config,
105 }
106 }
107
108 pub fn limit_provider<P2>(self, limit_provider: P2) -> RateLimitBuilder<K, S, P2, F> {
110 let Self {
111 key_extractor,
112 store,
113 response_factory,
114 config,
115 ..
116 } = self;
117 RateLimitBuilder {
118 key_extractor,
119 store,
120 limit_provider,
121 response_factory,
122 config,
123 }
124 }
125
126 pub fn response_factory<F2>(self, response_factory: F2) -> RateLimitBuilder<K, S, P, F2> {
128 let Self {
129 key_extractor,
130 store,
131 limit_provider,
132 config,
133 ..
134 } = self;
135 RateLimitBuilder {
136 key_extractor,
137 store,
138 limit_provider,
139 response_factory,
140 config,
141 }
142 }
143
144 pub fn window(mut self, window: Duration) -> Self {
146 self.config.window = window;
147 self
148 }
149
150 pub fn policy_name(mut self, policy_name: impl Into<String>) -> Self {
152 self.config.policy_name = policy_name.into();
153 self
154 }
155
156 pub fn with_key_encoder<E>(mut self, encoder: E) -> Self
163 where
164 E: Fn(&str) -> String + Send + Sync + 'static,
165 {
166 self.config.key_encoder = Some(Box::new(encoder));
167 self
168 }
169
170 pub fn skip<Predicate>(mut self, predicate: Predicate) -> Self
177 where
178 Predicate: Fn(&Request<()>) -> bool + Send + Sync + 'static,
179 {
180 self.config.skip_predicate = Some(Box::new(predicate));
181 self
182 }
183
184 pub fn store_failure_mode(mut self, mode: StoreFailureMode) -> Self {
186 self.config.store_failure_mode = mode;
187 self
188 }
189
190 #[cfg(feature = "tracing")]
196 pub fn store_failure_tracing_level(mut self, level: tracing::Level) -> Self {
197 self.config.store_failure_tracing_level = level;
198 self
199 }
200
201 pub fn rate_limit_fields(mut self, fields: RateLimitFields) -> Self {
203 self.config.rate_limit_fields = fields;
204 self
205 }
206
207 fn validate(&self) -> Result<(), ConfigError> {
208 if self.config.window < MINIMUM_WINDOW {
209 return Err(ConfigError::WindowTooShort(self.config.window, MINIMUM_WINDOW));
210 }
211 if self.config.policy_name.is_empty() {
212 return Err(ConfigError::EmptyPolicyName);
213 }
214 Ok(())
215 }
216}
217
218impl<K> RateLimitLayer<K, (), u64, DefaultResponseFactory> {
219 pub fn builder(key_extractor: K) -> RateLimitBuilder<K> {
221 RateLimitBuilder::new(key_extractor)
222 }
223}
224
225impl<K, S, P, F> RateLimitBuilder<K, S, P, F>
226where
227 S: Store,
228 P: LimitProvider,
229{
230 pub fn build(self) -> Result<RateLimitLayer<K, S, P, F>, ConfigError> {
232 self.validate()?;
233 Ok(RateLimitLayer {
234 key_extractor: self.key_extractor,
235 store: self.store,
236 limit_provider: self.limit_provider,
237 response_factory: self.response_factory,
238 config: Arc::new(self.config),
239 })
240 }
241}
242
243pub(crate) struct RateLimitConfig {
245 pub(crate) policy_name: String,
247 pub(crate) window: Duration,
249 pub(crate) key_encoder: Option<KeyEncoder>,
251 pub(crate) skip_predicate: Option<SkipPredicate>,
253 pub(crate) store_failure_mode: StoreFailureMode,
255 #[cfg(feature = "tracing")]
257 pub(crate) store_failure_tracing_level: tracing::Level,
258 pub(crate) rate_limit_fields: RateLimitFields,
260}
261
262impl fmt::Debug for RateLimitConfig {
263 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
264 let mut debug = f.debug_struct("RateLimitConfig");
265 debug
266 .field("policy_name", &self.policy_name)
267 .field("window", &self.window)
268 .field("has_key_encoder", &self.key_encoder.is_some())
269 .field("has_skip_predicate", &self.skip_predicate.is_some())
270 .field("store_failure_mode", &self.store_failure_mode);
271 #[cfg(feature = "tracing")]
272 debug.field("store_failure_tracing_level", &self.store_failure_tracing_level);
273 debug.field("rate_limit_fields", &self.rate_limit_fields).finish()
274 }
275}