1use std::collections::HashMap;
9use std::sync::atomic::{AtomicU64, Ordering};
10use std::sync::{Arc, RwLock};
11use std::time::{Duration, Instant};
12
13use crate::{now_timestamp, RateLimitError, RateLimitResult};
14
15pub struct LeakyBucketLimiter {
20 capacity: u64,
21 leak_rate: f64,
22 buckets: Arc<RwLock<HashMap<String, LeakyBucketEntry>>>,
23 max_keys: usize,
24 total_allowed: AtomicU64,
25 total_rejected: AtomicU64,
26}
27
28#[derive(Clone)]
29struct LeakyBucketEntry {
30 water: f64,
31 last_leak: Instant,
32}
33
34impl LeakyBucketLimiter {
35 pub fn new(capacity: u64, leak_rate: f64) -> Self {
40 Self {
41 capacity,
42 leak_rate,
43 buckets: Arc::new(RwLock::new(HashMap::new())),
44 max_keys: crate::DEFAULT_MAX_KEYS,
45 total_allowed: AtomicU64::new(0),
46 total_rejected: AtomicU64::new(0),
47 }
48 }
49
50 pub fn with_max_keys(mut self, max_keys: usize) -> Self {
52 self.max_keys = max_keys;
53 self
54 }
55
56 pub fn capacity(&self) -> u64 {
58 self.capacity
59 }
60
61 pub fn leak_rate(&self) -> f64 {
63 self.leak_rate
64 }
65
66 pub fn key_count(&self) -> usize {
68 self.buckets.read().map(|m| m.len()).unwrap_or(0)
69 }
70
71 pub fn water_level(&self, key: &str) -> f64 {
73 let buckets = self.buckets.read().map_err(|e| e.to_string());
74 match buckets {
75 Ok(map) => map.get(key).map(|e| e.water).unwrap_or(0.0),
76 Err(_) => 0.0,
77 }
78 }
79
80 fn leak(&self, entry: &mut LeakyBucketEntry) {
81 let now = Instant::now();
82 let elapsed = now.duration_since(entry.last_leak).as_secs_f64();
83 let leaked = if self.leak_rate > 0.0 {
84 elapsed * self.leak_rate
85 } else {
86 0.0
87 };
88 entry.water = (entry.water - leaked).max(0.0);
89 entry.last_leak = now;
90 }
91
92 fn enforce_max_keys(&self, buckets: &mut HashMap<String, LeakyBucketEntry>) {
93 while buckets.len() > self.max_keys {
94 let oldest = buckets
95 .iter()
96 .min_by_key(|(_, e)| e.last_leak)
97 .map(|(k, _)| k.clone());
98 match oldest {
99 Some(k) => {
100 buckets.remove(&k);
101 }
102 None => break,
103 }
104 }
105 }
106
107 pub fn acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
109 let mut buckets = self
110 .buckets
111 .write()
112 .map_err(|e| RateLimitError::Internal(e.to_string()))?;
113
114 if buckets.len() >= self.max_keys && !buckets.contains_key(key) {
115 self.enforce_max_keys(&mut buckets);
116 }
117
118 let entry = buckets
119 .entry(key.to_string())
120 .or_insert_with(|| LeakyBucketEntry {
121 water: 0.0,
122 last_leak: Instant::now(),
123 });
124
125 self.leak(entry);
126
127 if entry.water + 1.0 <= self.capacity as f64 {
128 entry.water += 1.0;
129 let remaining = (self.capacity as f64 - entry.water).floor() as u64;
130 self.total_allowed.fetch_add(1, Ordering::Relaxed);
131 Ok(RateLimitResult::allowed(remaining, now_timestamp() + 1000))
132 } else {
133 self.total_rejected.fetch_add(1, Ordering::Relaxed);
134 Ok(RateLimitResult::rejected(0, now_timestamp() + 1000))
135 }
136 }
137
138 pub fn reset(&self, key: &str) -> Result<(), RateLimitError> {
140 let mut buckets = self
141 .buckets
142 .write()
143 .map_err(|e| RateLimitError::Internal(e.to_string()))?;
144 buckets.remove(key);
145 Ok(())
146 }
147
148 pub fn stats(&self) -> LeakyBucketStats {
150 LeakyBucketStats {
151 capacity: self.capacity,
152 leak_rate: self.leak_rate,
153 key_count: self.key_count(),
154 total_allowed: self.total_allowed.load(Ordering::Relaxed),
155 total_rejected: self.total_rejected.load(Ordering::Relaxed),
156 }
157 }
158}
159
160#[derive(Debug, Clone, serde::Serialize)]
162pub struct LeakyBucketStats {
163 pub capacity: u64,
164 pub leak_rate: f64,
165 pub key_count: usize,
166 pub total_allowed: u64,
167 pub total_rejected: u64,
168}
169
170pub struct SlidingWindowLogLimiter {
175 max_requests: u64,
176 window_size: Duration,
177 entries: Arc<RwLock<HashMap<String, Vec<Instant>>>>,
178 max_keys: usize,
179}
180
181impl SlidingWindowLogLimiter {
182 pub fn new(max_requests: u64, window_size: Duration) -> Self {
184 Self {
185 max_requests,
186 window_size,
187 entries: Arc::new(RwLock::new(HashMap::new())),
188 max_keys: crate::DEFAULT_MAX_KEYS,
189 }
190 }
191
192 pub fn with_max_keys(mut self, max_keys: usize) -> Self {
194 self.max_keys = max_keys;
195 self
196 }
197
198 pub fn max_requests(&self) -> u64 {
200 self.max_requests
201 }
202
203 pub fn window_size(&self) -> Duration {
205 self.window_size
206 }
207
208 pub fn key_count(&self) -> usize {
210 self.entries.read().map(|m| m.len()).unwrap_or(0)
211 }
212
213 pub fn current_count(&self, key: &str) -> usize {
215 let entries = self.entries.read().map_err(|e| e.to_string());
216 match entries {
217 Ok(map) => {
218 if let Some(log) = map.get(key) {
219 let now = Instant::now();
220 log.iter()
221 .filter(|&&t| now.duration_since(t) < self.window_size)
222 .count()
223 } else {
224 0
225 }
226 }
227 Err(_) => 0,
228 }
229 }
230
231 fn enforce_max_keys(&self, entries: &mut HashMap<String, Vec<Instant>>) {
232 while entries.len() > self.max_keys {
233 let now = Instant::now();
234 let oldest = entries
235 .iter()
236 .min_by_key(|(_, log)| log.first().copied().unwrap_or(now))
237 .map(|(k, _)| k.clone());
238 match oldest {
239 Some(k) => {
240 entries.remove(&k);
241 }
242 None => break,
243 }
244 }
245 }
246
247 pub fn acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
249 let mut entries = self
250 .entries
251 .write()
252 .map_err(|e| RateLimitError::Internal(e.to_string()))?;
253
254 if entries.len() >= self.max_keys && !entries.contains_key(key) {
255 self.enforce_max_keys(&mut entries);
256 }
257
258 let log = entries.entry(key.to_string()).or_insert_with(Vec::new);
259
260 let now = Instant::now();
261 log.retain(|&t| now.duration_since(t) < self.window_size);
262
263 if log.len() < self.max_requests as usize {
264 log.push(now);
265 let remaining = self.max_requests - log.len() as u64;
266 Ok(RateLimitResult::allowed(remaining, now_timestamp() + 1000))
267 } else {
268 Ok(RateLimitResult::rejected(0, now_timestamp() + 1000))
269 }
270 }
271
272 pub fn reset(&self, key: &str) -> Result<(), RateLimitError> {
274 let mut entries = self
275 .entries
276 .write()
277 .map_err(|e| RateLimitError::Internal(e.to_string()))?;
278 entries.remove(key);
279 Ok(())
280 }
281}
282
283#[cfg(test)]
284mod tests {
285 use super::*;
286
287 #[test]
288 fn test_leaky_bucket_basic() {
289 let limiter = LeakyBucketLimiter::new(5, 1.0);
290 let r = limiter.acquire("k").unwrap();
291 assert!(r.allowed);
292 }
293
294 #[test]
295 fn test_leaky_bucket_full() {
296 let limiter = LeakyBucketLimiter::new(2, 0.0);
297 assert!(limiter.acquire("k").unwrap().allowed);
298 assert!(limiter.acquire("k").unwrap().allowed);
299 assert!(!limiter.acquire("k").unwrap().allowed);
300 }
301
302 #[test]
303 fn test_leaky_bucket_different_keys() {
304 let limiter = LeakyBucketLimiter::new(1, 0.0);
305 assert!(limiter.acquire("a").unwrap().allowed);
306 assert!(limiter.acquire("b").unwrap().allowed);
307 }
308
309 #[test]
310 fn test_leaky_bucket_water_level() {
311 let limiter = LeakyBucketLimiter::new(5, 0.0);
312 assert_eq!(limiter.water_level("k"), 0.0);
313 limiter.acquire("k").unwrap();
314 assert!(limiter.water_level("k") > 0.0);
315 }
316
317 #[test]
318 fn test_leaky_bucket_reset() {
319 let limiter = LeakyBucketLimiter::new(1, 0.0);
320 limiter.acquire("k").unwrap();
321 assert!(!limiter.acquire("k").unwrap().allowed);
322 limiter.reset("k").unwrap();
323 assert!(limiter.acquire("k").unwrap().allowed);
324 }
325
326 #[test]
327 fn test_leaky_bucket_stats() {
328 let limiter = LeakyBucketLimiter::new(3, 1.0);
329 limiter.acquire("k").unwrap();
330 limiter.acquire("k").unwrap();
331 let stats = limiter.stats();
332 assert_eq!(stats.capacity, 3);
333 assert_eq!(stats.total_allowed, 2);
334 assert_eq!(stats.total_rejected, 0);
335 }
336
337 #[test]
338 fn test_leaky_bucket_key_count() {
339 let limiter = LeakyBucketLimiter::new(5, 1.0);
340 assert_eq!(limiter.key_count(), 0);
341 limiter.acquire("a").unwrap();
342 limiter.acquire("b").unwrap();
343 assert_eq!(limiter.key_count(), 2);
344 }
345
346 #[test]
347 fn test_sliding_window_log_basic() {
348 let limiter = SlidingWindowLogLimiter::new(5, Duration::from_secs(60));
349 let r = limiter.acquire("k").unwrap();
350 assert!(r.allowed);
351 }
352
353 #[test]
354 fn test_sliding_window_log_full() {
355 let limiter = SlidingWindowLogLimiter::new(2, Duration::from_secs(60));
356 assert!(limiter.acquire("k").unwrap().allowed);
357 assert!(limiter.acquire("k").unwrap().allowed);
358 assert!(!limiter.acquire("k").unwrap().allowed);
359 }
360
361 #[test]
362 fn test_sliding_window_log_different_keys() {
363 let limiter = SlidingWindowLogLimiter::new(1, Duration::from_secs(60));
364 assert!(limiter.acquire("a").unwrap().allowed);
365 assert!(limiter.acquire("b").unwrap().allowed);
366 }
367
368 #[test]
369 fn test_sliding_window_log_current_count() {
370 let limiter = SlidingWindowLogLimiter::new(5, Duration::from_secs(60));
371 assert_eq!(limiter.current_count("k"), 0);
372 limiter.acquire("k").unwrap();
373 limiter.acquire("k").unwrap();
374 assert_eq!(limiter.current_count("k"), 2);
375 }
376
377 #[test]
378 fn test_sliding_window_log_reset() {
379 let limiter = SlidingWindowLogLimiter::new(1, Duration::from_secs(60));
380 limiter.acquire("k").unwrap();
381 assert!(!limiter.acquire("k").unwrap().allowed);
382 limiter.reset("k").unwrap();
383 assert!(limiter.acquire("k").unwrap().allowed);
384 }
385
386 #[test]
387 fn test_sliding_window_log_key_count() {
388 let limiter = SlidingWindowLogLimiter::new(5, Duration::from_secs(60));
389 assert_eq!(limiter.key_count(), 0);
390 limiter.acquire("a").unwrap();
391 assert_eq!(limiter.key_count(), 1);
392 }
393
394 #[test]
395 fn test_sliding_window_log_remaining() {
396 let limiter = SlidingWindowLogLimiter::new(3, Duration::from_secs(60));
397 let r1 = limiter.acquire("k").unwrap();
398 assert_eq!(r1.remaining, 2);
399 let r2 = limiter.acquire("k").unwrap();
400 assert_eq!(r2.remaining, 1);
401 }
402}