1use serde::{Deserialize, Serialize};
16use std::{
17 collections::{HashMap, VecDeque},
18 sync::{Arc, Mutex, MutexGuard, PoisonError},
19 time::Duration,
20};
21
22#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
24#[serde(rename_all = "snake_case")]
25pub enum RateLimitPolicy {
26 #[default]
29 Off,
30 Wait {
34 max_wait: Duration,
36 },
37 Fail,
40}
41
42const DEFAULT_WINDOW_MS: i64 = 900_000;
44const DEFAULT_COST: i64 = 2;
46const WAIT_MARGIN_MS: i64 = 50;
48
49#[derive(Debug)]
51struct Bucket {
52 window_ms: i64,
53 remaining: i64,
55 cost: i64,
57 last_response_ms: i64,
58 spent: VecDeque<(i64, i64)>,
60 in_flight: i64,
62 blocked_until_ms: i64,
64 seen: bool,
66}
67
68impl Default for Bucket {
69 fn default() -> Self {
70 Bucket {
71 window_ms: DEFAULT_WINDOW_MS,
72 remaining: 0,
73 cost: DEFAULT_COST,
74 last_response_ms: 0,
75 spent: VecDeque::new(),
76 in_flight: 0,
77 blocked_until_ms: 0,
78 seen: false,
79 }
80 }
81}
82
83impl Bucket {
84 fn wait_ms(&self, now: i64) -> i64 {
86 if !self.seen {
87 return 0;
88 }
89 let needed = (self.in_flight + 1) * self.cost;
90 let mut available = self.remaining;
91 let mut ready_at = now;
92 if available < needed {
93 ready_at = self.last_response_ms + self.window_ms;
95 for (at, cost) in &self.spent {
96 let returns_at = at + self.window_ms;
97 if returns_at <= self.last_response_ms {
98 continue;
99 }
100 available += cost;
101 if available >= needed {
102 ready_at = returns_at;
103 break;
104 }
105 }
106 }
107 let wait = (ready_at - now).max(self.blocked_until_ms - now).max(0);
108 if wait > 0 {
109 wait + WAIT_MARGIN_MS
110 } else {
111 0
112 }
113 }
114}
115
116#[derive(Debug, PartialEq, Eq)]
118pub(crate) enum Acquire {
119 Granted,
121 Wait(i64),
123}
124
125#[derive(Debug, Default)]
127pub(crate) struct RateLimiter {
128 buckets: Mutex<HashMap<String, Bucket>>,
129}
130
131impl RateLimiter {
132 fn lock(&self) -> MutexGuard<'_, HashMap<String, Bucket>> {
133 self.buckets.lock().unwrap_or_else(PoisonError::into_inner)
134 }
135
136 pub(crate) fn key(group: &str, token: Option<&str>) -> String {
138 format!("{group}|{}", token.unwrap_or(""))
139 }
140
141 pub(crate) fn group_of(key: &str) -> &str {
143 key.split('|').next().unwrap_or(key)
144 }
145
146 pub(crate) fn try_acquire(&self, key: &str, now: i64) -> Acquire {
148 let mut buckets = self.lock();
149 let bucket = buckets.entry(key.to_owned()).or_default();
150 match bucket.wait_ms(now) {
151 0 => {
152 bucket.in_flight += 1;
153 Acquire::Granted
154 }
155 wait => Acquire::Wait(wait),
156 }
157 }
158
159 pub(crate) fn release(&self, key: &str) {
161 if let Some(bucket) = self.lock().get_mut(key) {
162 bucket.in_flight = (bucket.in_flight - 1).max(0);
163 }
164 }
165
166 pub(crate) fn record(
169 &self,
170 key: &str,
171 remaining: i64,
172 used: i64,
173 window_ms: i64,
174 success: bool,
175 now: i64,
176 ) {
177 let mut buckets = self.lock();
178 let bucket = buckets.entry(key.to_owned()).or_default();
179 if window_ms > 0 {
180 bucket.window_ms = window_ms;
181 }
182 bucket.remaining = remaining;
183 bucket.last_response_ms = now;
184 bucket.seen = true;
185 if success && used > 0 {
186 bucket.cost = used;
187 }
188 if used > 0 {
189 bucket.spent.push_back((now, used));
190 }
191 let window = bucket.window_ms;
192 while bucket
193 .spent
194 .front()
195 .is_some_and(|(at, _)| at + window <= now)
196 {
197 bucket.spent.pop_front();
198 }
199 buckets.retain(|_, b| b.in_flight > 0 || b.last_response_ms + b.window_ms > now);
200 }
201
202 pub(crate) fn block_until(&self, key: &str, until_ms: i64) {
204 let mut buckets = self.lock();
205 let bucket = buckets.entry(key.to_owned()).or_default();
206 bucket.blocked_until_ms = bucket.blocked_until_ms.max(until_ms);
207 bucket.seen = true;
208 }
209}
210
211pub(crate) struct Permit {
213 limiter: Arc<RateLimiter>,
214 key: String,
215}
216
217impl Permit {
218 pub(crate) fn new(limiter: Arc<RateLimiter>, key: &str) -> Self {
219 Permit {
220 limiter,
221 key: key.to_owned(),
222 }
223 }
224}
225
226impl Drop for Permit {
227 fn drop(&mut self) {
228 self.limiter.release(&self.key);
229 }
230}
231
232#[cfg(test)]
233mod tests {
234 use super::*;
235
236 const KEY: &str = "g|t";
237
238 #[test]
239 fn test_unseen_groups_are_not_throttled() {
240 let limiter = RateLimiter::default();
241 for _ in 0..10 {
242 assert_eq!(limiter.try_acquire(KEY, 0), Acquire::Granted);
243 }
244 }
245
246 #[test]
247 fn test_requests_that_fit_are_granted_and_in_flight_ones_count() {
248 let limiter = RateLimiter::default();
249 limiter.record(KEY, 5, 2, 10_000, true, 0);
251 assert_eq!(limiter.try_acquire(KEY, 1), Acquire::Granted);
252 assert_eq!(limiter.try_acquire(KEY, 1), Acquire::Granted);
253 assert!(matches!(limiter.try_acquire(KEY, 1), Acquire::Wait(_)));
254 limiter.release(KEY);
256 assert_eq!(limiter.try_acquire(KEY, 1), Acquire::Granted);
257 }
258
259 #[test]
260 fn test_tokens_return_when_their_window_ends() {
261 let limiter = RateLimiter::default();
262 limiter.record(KEY, 0, 2, 1_000, true, 0);
263 let Acquire::Wait(wait) = limiter.try_acquire(KEY, 100) else {
265 panic!("should wait");
266 };
267 assert_eq!(wait, 900 + WAIT_MARGIN_MS);
268 assert_eq!(limiter.try_acquire(KEY, 1_000), Acquire::Granted);
269 }
270
271 #[test]
272 fn test_the_wait_ends_at_the_first_return_that_is_enough() {
273 let limiter = RateLimiter::default();
274 for at in [0, 400, 800] {
277 limiter.record(KEY, 0, 2, 1_000, true, at);
278 }
279 assert_eq!(
280 limiter.try_acquire(KEY, 800),
281 Acquire::Wait(200 + WAIT_MARGIN_MS)
282 );
283 let limiter = RateLimiter::default();
285 for at in [0, 400, 800] {
286 limiter.record(KEY, 0, 2, 1_000, true, at);
287 }
288 limiter.record(KEY, 0, 0, 1_000, true, 800);
289 assert_eq!(limiter.try_acquire(KEY, 1_000), Acquire::Granted);
290 assert_eq!(
291 limiter.try_acquire(KEY, 1_000),
292 Acquire::Wait(400 + WAIT_MARGIN_MS)
293 );
294 }
295
296 #[test]
297 fn test_retry_after_blocks_the_group() {
298 let limiter = RateLimiter::default();
299 limiter.record(KEY, 100, 2, 10_000, true, 0);
300 limiter.block_until(KEY, 5_000);
301 assert_eq!(
302 limiter.try_acquire(KEY, 1_000),
303 Acquire::Wait(4_000 + WAIT_MARGIN_MS)
304 );
305 assert_eq!(limiter.try_acquire(KEY, 5_000), Acquire::Granted);
306 }
307
308 #[test]
309 fn test_buckets_are_independent_per_group_and_token() {
310 let limiter = RateLimiter::default();
311 limiter.record("a|x", 0, 2, 10_000, true, 0);
312 assert!(matches!(limiter.try_acquire("a|x", 1), Acquire::Wait(_)));
313 assert_eq!(limiter.try_acquire("a|y", 1), Acquire::Granted);
314 assert_eq!(limiter.try_acquire("b|x", 1), Acquire::Granted);
315 assert_eq!(RateLimiter::key("a", Some("x")), "a|x");
316 assert_eq!(RateLimiter::group_of("a|x"), "a");
317 }
318
319 #[test]
320 fn test_idle_buckets_are_dropped() {
321 let limiter = RateLimiter::default();
322 limiter.record("old|x", 10, 2, 1_000, true, 0);
323 limiter.record("new|x", 10, 2, 1_000, true, 5_000);
324 assert_eq!(limiter.lock().len(), 1);
325 }
326
327 #[test]
328 fn test_a_permit_releases_on_drop() {
329 let limiter = Arc::new(RateLimiter::default());
330 limiter.record(KEY, 2, 2, 10_000, true, 0);
331 assert_eq!(limiter.try_acquire(KEY, 1), Acquire::Granted);
332 let permit = Permit::new(Arc::clone(&limiter), KEY);
333 assert!(matches!(limiter.try_acquire(KEY, 1), Acquire::Wait(_)));
334 drop(permit);
335 assert_eq!(limiter.try_acquire(KEY, 1), Acquire::Granted);
336 }
337
338 #[test]
339 fn test_the_policy_serializes_in_snake_case() {
340 assert_eq!(
341 serde_json::to_string(&RateLimitPolicy::Off).unwrap(),
342 "\"off\""
343 );
344 let wait = RateLimitPolicy::Wait {
345 max_wait: Duration::from_secs(3),
346 };
347 let json = serde_json::to_string(&wait).unwrap();
348 assert_eq!(
349 serde_json::from_str::<RateLimitPolicy>(&json).unwrap(),
350 wait
351 );
352 }
353}