1use std::sync::atomic::{AtomicU64, Ordering};
6use std::sync::Arc;
7use std::time::Duration;
8
9use crate::{RateLimitError, RateLimitResult, RateLimiter};
10
11#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
13pub enum CompositeStrategy {
14 AllMustAllow,
16 AnyCanAllow,
18 FirstRejectWins,
20 WeightedVote,
22}
23
24pub struct CompositeLimiter {
28 limiters: Vec<Arc<dyn RateLimiter>>,
29 strategy: CompositeStrategy,
30 weights: Vec<u32>,
31 total_calls: AtomicU64,
32}
33
34impl CompositeLimiter {
35 pub fn new(limiters: Vec<Arc<dyn RateLimiter>>, strategy: CompositeStrategy) -> Self {
37 let count = limiters.len();
38 Self {
39 limiters,
40 strategy,
41 weights: vec![1; count],
42 total_calls: AtomicU64::new(0),
43 }
44 }
45
46 pub fn with_weights(mut self, weights: Vec<u32>) -> Self {
48 if weights.len() == self.limiters.len() {
49 self.weights = weights;
50 }
51 self
52 }
53
54 pub fn with_limiter(mut self, limiter: Arc<dyn RateLimiter>) -> Self {
56 self.limiters.push(limiter);
57 self.weights.push(1);
58 self
59 }
60
61 pub fn limiter_count(&self) -> usize {
63 self.limiters.len()
64 }
65
66 pub fn strategy(&self) -> CompositeStrategy {
68 self.strategy
69 }
70
71 pub fn check(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
73 self.total_calls.fetch_add(1, Ordering::Relaxed);
74 if self.limiters.is_empty() {
75 return Err(RateLimitError::Internal(
76 "No limiters configured".to_string(),
77 ));
78 }
79 match self.strategy {
80 CompositeStrategy::AllMustAllow => self.check_all_must_allow(key),
81 CompositeStrategy::AnyCanAllow => self.check_any_can_allow(key),
82 CompositeStrategy::FirstRejectWins => self.check_first_reject_wins(key),
83 CompositeStrategy::WeightedVote => self.check_weighted_vote(key),
84 }
85 }
86
87 fn check_all_must_allow(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
88 let mut min_remaining = u64::MAX;
89 let mut max_reset = 0i64;
90 for limiter in &self.limiters {
91 let result = limiter.acquire(key)?;
92 if !result.allowed {
93 return Ok(result);
94 }
95 min_remaining = min_remaining.min(result.remaining);
96 max_reset = max_reset.max(result.reset_at);
97 }
98 Ok(RateLimitResult::allowed(min_remaining, max_reset))
99 }
100
101 fn check_any_can_allow(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
102 let mut best: Option<RateLimitResult> = None;
103 for limiter in &self.limiters {
104 let result = limiter.acquire(key)?;
105 if result.allowed {
106 return Ok(result);
107 }
108 match &best {
109 None => best = Some(result),
110 Some(current) => {
111 if result.remaining > current.remaining {
112 best = Some(result);
113 }
114 }
115 }
116 }
117 best.ok_or_else(|| RateLimitError::Internal("No results".to_string()))
118 }
119
120 fn check_first_reject_wins(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
121 let mut last_allowed: Option<RateLimitResult> = None;
122 for limiter in &self.limiters {
123 let result = limiter.acquire(key)?;
124 if !result.allowed {
125 return Ok(result);
126 }
127 last_allowed = Some(result);
128 }
129 last_allowed.ok_or_else(|| RateLimitError::Internal("No results".to_string()))
130 }
131
132 fn check_weighted_vote(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
133 let mut allow_weight = 0u32;
134 let mut reject_weight = 0u32;
135 let mut best_allowed: Option<RateLimitResult> = None;
136 let mut best_rejected: Option<RateLimitResult> = None;
137 for (i, limiter) in self.limiters.iter().enumerate() {
138 let result = limiter.acquire(key)?;
139 let weight = self.weights.get(i).copied().unwrap_or(1);
140 if result.allowed {
141 allow_weight += weight;
142 if best_allowed.is_none() {
143 best_allowed = Some(result);
144 }
145 } else {
146 reject_weight += weight;
147 if best_rejected.is_none() {
148 best_rejected = Some(result);
149 }
150 }
151 }
152 if allow_weight >= reject_weight {
153 best_allowed.ok_or_else(|| RateLimitError::Internal("No allowed".to_string()))
154 } else {
155 best_rejected.ok_or_else(|| RateLimitError::Internal("No rejected".to_string()))
156 }
157 }
158
159 pub fn total_calls(&self) -> u64 {
161 self.total_calls.load(Ordering::Relaxed)
162 }
163}
164
165pub struct FallbackLimiter {
169 primary: Arc<dyn RateLimiter>,
170 fallback: Arc<dyn RateLimiter>,
171 fallback_count: AtomicU64,
172}
173
174impl FallbackLimiter {
175 pub fn new(primary: Arc<dyn RateLimiter>, fallback: Arc<dyn RateLimiter>) -> Self {
177 Self {
178 primary,
179 fallback,
180 fallback_count: AtomicU64::new(0),
181 }
182 }
183
184 pub fn check(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
186 match self.primary.acquire(key) {
187 Ok(result) => Ok(result),
188 Err(_) => {
189 self.fallback_count.fetch_add(1, Ordering::Relaxed);
190 self.fallback.acquire(key)
191 }
192 }
193 }
194
195 pub fn fallback_count(&self) -> u64 {
197 self.fallback_count.load(Ordering::Relaxed)
198 }
199}
200
201#[derive(Debug, Clone)]
205pub struct LimitKeyBuilder {
206 parts: Vec<String>,
207 separator: String,
208}
209
210impl Default for LimitKeyBuilder {
211 fn default() -> Self {
212 Self::new()
213 }
214}
215
216impl LimitKeyBuilder {
217 pub fn new() -> Self {
219 Self {
220 parts: Vec::new(),
221 separator: ":".to_string(),
222 }
223 }
224
225 pub fn with_separator(mut self, sep: &str) -> Self {
227 self.separator = sep.to_string();
228 self
229 }
230
231 pub fn ip(mut self, ip: &str) -> Self {
233 self.parts.push(format!("ip:{}", ip));
234 self
235 }
236
237 pub fn user(mut self, user: &str) -> Self {
239 self.parts.push(format!("user:{}", user));
240 self
241 }
242
243 pub fn api(mut self, api: &str) -> Self {
245 self.parts.push(format!("api:{}", api));
246 self
247 }
248
249 pub fn dimension(mut self, name: &str, value: &str) -> Self {
251 self.parts.push(format!("{}:{}", name, value));
252 self
253 }
254
255 pub fn build(&self) -> String {
257 self.parts.join(&self.separator)
258 }
259
260 pub fn part_count(&self) -> usize {
262 self.parts.len()
263 }
264}
265
266#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
270pub struct RateLimitRule {
271 pub name: String,
273 pub key_prefix: String,
275 pub limiter_type: String,
277 pub capacity: u64,
279 pub window_ms: u64,
281 pub enabled: bool,
283}
284
285impl RateLimitRule {
286 pub fn new(name: &str, key_prefix: &str, limiter_type: &str, capacity: u64) -> Self {
288 Self {
289 name: name.to_string(),
290 key_prefix: key_prefix.to_string(),
291 limiter_type: limiter_type.to_string(),
292 capacity,
293 window_ms: 1000,
294 enabled: true,
295 }
296 }
297
298 pub fn with_window_ms(mut self, ms: u64) -> Self {
300 self.window_ms = ms;
301 self
302 }
303
304 pub fn disable(mut self) -> Self {
306 self.enabled = false;
307 self
308 }
309
310 pub fn matches(&self, key: &str) -> bool {
312 self.enabled && key.starts_with(&self.key_prefix)
313 }
314}
315
316pub struct RuleSet {
318 rules: Vec<RateLimitRule>,
319}
320
321impl RuleSet {
322 pub fn new() -> Self {
324 Self { rules: Vec::new() }
325 }
326
327 pub fn add(&mut self, rule: RateLimitRule) -> &mut Self {
329 self.rules.push(rule);
330 self
331 }
332
333 pub fn find_matches(&self, key: &str) -> Vec<&RateLimitRule> {
335 self.rules.iter().filter(|r| r.matches(key)).collect()
336 }
337
338 pub fn rule_count(&self) -> usize {
340 self.rules.len()
341 }
342
343 pub fn enabled_count(&self) -> usize {
345 self.rules.iter().filter(|r| r.enabled).count()
346 }
347
348 pub fn find_by_prefix(&self, prefix: &str) -> Option<&RateLimitRule> {
350 self.rules.iter().find(|r| r.key_prefix == prefix)
351 }
352
353 pub fn disable(&mut self, name: &str) -> bool {
355 for rule in &mut self.rules {
356 if rule.name == name {
357 rule.enabled = false;
358 return true;
359 }
360 }
361 false
362 }
363
364 pub fn enable(&mut self, name: &str) -> bool {
366 for rule in &mut self.rules {
367 if rule.name == name {
368 rule.enabled = true;
369 return true;
370 }
371 }
372 false
373 }
374}
375
376impl Default for RuleSet {
377 fn default() -> Self {
378 Self::new()
379 }
380}
381
382#[cfg(test)]
383mod tests {
384 use super::*;
385 use crate::{SlidingWindowRateLimiter, TokenBucketRateLimiter};
386
387 #[test]
388 fn test_composite_all_must_allow() {
389 let l1 = Arc::new(SlidingWindowRateLimiter::new(10, Duration::from_secs(60)));
390 let l2 = Arc::new(TokenBucketRateLimiter::new(10, 1.0));
391 let composite = CompositeLimiter::new(vec![l1, l2], CompositeStrategy::AllMustAllow);
392 let r = composite.check("k").unwrap();
393 assert!(r.allowed);
394 }
395
396 #[test]
397 fn test_composite_all_must_allow_one_rejects() {
398 let l1 = Arc::new(SlidingWindowRateLimiter::new(10, Duration::from_secs(60)));
399 let l2 = Arc::new(TokenBucketRateLimiter::new(1, 1.0));
400 let composite = CompositeLimiter::new(vec![l1, l2], CompositeStrategy::AllMustAllow);
401 composite.check("k").unwrap();
402 let r = composite.check("k").unwrap();
403 assert!(!r.allowed);
404 }
405
406 #[test]
407 fn test_composite_any_can_allow() {
408 let l1 = Arc::new(SlidingWindowRateLimiter::new(1, Duration::from_secs(60)));
409 let l2 = Arc::new(TokenBucketRateLimiter::new(10, 1.0));
410 let composite = CompositeLimiter::new(vec![l1, l2], CompositeStrategy::AnyCanAllow);
411 composite.check("k").unwrap();
412 let r = composite.check("k").unwrap();
413 assert!(r.allowed);
414 }
415
416 #[test]
417 fn test_composite_any_can_allow_all_reject() {
418 let l1 = Arc::new(SlidingWindowRateLimiter::new(1, Duration::from_secs(60)));
419 let l2 = Arc::new(TokenBucketRateLimiter::new(1, 0.0));
420 l2.acquire("k").unwrap();
422 let composite = CompositeLimiter::new(vec![l1, l2], CompositeStrategy::AnyCanAllow);
423 composite.check("k").unwrap();
424 let r = composite.check("k").unwrap();
425 assert!(!r.allowed);
426 }
427
428 #[test]
429 fn test_composite_first_reject_wins() {
430 let l1 = Arc::new(SlidingWindowRateLimiter::new(1, Duration::from_secs(60)));
431 let l2 = Arc::new(TokenBucketRateLimiter::new(10, 1.0));
432 let composite = CompositeLimiter::new(vec![l1, l2], CompositeStrategy::FirstRejectWins);
433 composite.check("k").unwrap();
434 let r = composite.check("k").unwrap();
435 assert!(!r.allowed);
436 }
437
438 #[test]
439 fn test_composite_weighted_vote_allow() {
440 let l1 = Arc::new(SlidingWindowRateLimiter::new(10, Duration::from_secs(60)));
441 let l2 = Arc::new(TokenBucketRateLimiter::new(1, 1.0));
442 let composite = CompositeLimiter::new(vec![l1, l2], CompositeStrategy::WeightedVote)
443 .with_weights(vec![3, 1]);
444 composite.check("k").unwrap();
445 let r = composite.check("k").unwrap();
446 assert!(r.allowed);
447 }
448
449 #[test]
450 fn test_composite_empty_errors() {
451 let composite = CompositeLimiter::new(vec![], CompositeStrategy::AllMustAllow);
452 assert!(composite.check("k").is_err());
453 }
454
455 #[test]
456 fn test_composite_with_limiter() {
457 let l1 = Arc::new(SlidingWindowRateLimiter::new(10, Duration::from_secs(60)));
458 let composite =
459 CompositeLimiter::new(vec![], CompositeStrategy::AllMustAllow).with_limiter(l1);
460 assert_eq!(composite.limiter_count(), 1);
461 assert!(composite.check("k").unwrap().allowed);
462 }
463
464 #[test]
465 fn test_composite_total_calls() {
466 let l1 = Arc::new(SlidingWindowRateLimiter::new(10, Duration::from_secs(60)));
467 let composite = CompositeLimiter::new(vec![l1], CompositeStrategy::AllMustAllow);
468 composite.check("k").unwrap();
469 composite.check("k").unwrap();
470 assert_eq!(composite.total_calls(), 2);
471 }
472
473 #[test]
474 fn test_fallback_primary_ok() {
475 let primary = Arc::new(SlidingWindowRateLimiter::new(10, Duration::from_secs(60)));
476 let fallback = Arc::new(TokenBucketRateLimiter::new(10, 1.0));
477 let limiter = FallbackLimiter::new(primary, fallback);
478 let r = limiter.check("k").unwrap();
479 assert!(r.allowed);
480 assert_eq!(limiter.fallback_count(), 0);
481 }
482
483 #[test]
484 fn test_limit_key_builder_basic() {
485 let key = LimitKeyBuilder::new()
486 .ip("127.0.0.1")
487 .user("user-1")
488 .api("/query")
489 .build();
490 assert!(key.contains("ip:127.0.0.1"));
491 assert!(key.contains("user:user-1"));
492 assert!(key.contains("api:/query"));
493 }
494
495 #[test]
496 fn test_limit_key_builder_separator() {
497 let key = LimitKeyBuilder::new()
498 .with_separator("|")
499 .ip("127.0.0.1")
500 .user("user-1")
501 .build();
502 assert!(key.contains("|"));
503 }
504
505 #[test]
506 fn test_limit_key_builder_dimension() {
507 let key = LimitKeyBuilder::new().dimension("tenant", "acme").build();
508 assert!(key.contains("tenant:acme"));
509 }
510
511 #[test]
512 fn test_limit_key_builder_part_count() {
513 let builder = LimitKeyBuilder::new().ip("127.0.0.1").user("user-1");
514 assert_eq!(builder.part_count(), 2);
515 }
516
517 #[test]
518 fn test_rate_limit_rule_new() {
519 let rule = RateLimitRule::new("ip-limit", "ip:", "sliding_window", 100);
520 assert_eq!(rule.name, "ip-limit");
521 assert_eq!(rule.capacity, 100);
522 assert!(rule.enabled);
523 }
524
525 #[test]
526 fn test_rate_limit_rule_matches() {
527 let rule = RateLimitRule::new("ip-limit", "ip:", "sliding_window", 100);
528 assert!(rule.matches("ip:127.0.0.1"));
529 assert!(!rule.matches("user:1"));
530 }
531
532 #[test]
533 fn test_rate_limit_rule_disabled() {
534 let rule = RateLimitRule::new("ip-limit", "ip:", "sliding_window", 100).disable();
535 assert!(!rule.matches("ip:127.0.0.1"));
536 }
537
538 #[test]
539 fn test_rate_limit_rule_with_window() {
540 let rule =
541 RateLimitRule::new("ip-limit", "ip:", "sliding_window", 100).with_window_ms(5000);
542 assert_eq!(rule.window_ms, 5000);
543 }
544
545 #[test]
546 fn test_rule_set_add() {
547 let mut set = RuleSet::new();
548 set.add(RateLimitRule::new("r1", "ip:", "sw", 100));
549 assert_eq!(set.rule_count(), 1);
550 }
551
552 #[test]
553 fn test_rule_set_find_matches() {
554 let mut set = RuleSet::new();
555 set.add(RateLimitRule::new("r1", "ip:", "sw", 100));
556 set.add(RateLimitRule::new("r2", "user:", "sw", 200));
557 let matches = set.find_matches("ip:127.0.0.1");
558 assert_eq!(matches.len(), 1);
559 }
560
561 #[test]
562 fn test_rule_set_enabled_count() {
563 let mut set = RuleSet::new();
564 set.add(RateLimitRule::new("r1", "ip:", "sw", 100));
565 set.add(RateLimitRule::new("r2", "user:", "sw", 200).disable());
566 assert_eq!(set.enabled_count(), 1);
567 }
568
569 #[test]
570 fn test_rule_set_find_by_prefix() {
571 let mut set = RuleSet::new();
572 set.add(RateLimitRule::new("r1", "ip:", "sw", 100));
573 assert!(set.find_by_prefix("ip:").is_some());
574 assert!(set.find_by_prefix("user:").is_none());
575 }
576
577 #[test]
578 fn test_rule_set_disable_enable() {
579 let mut set = RuleSet::new();
580 set.add(RateLimitRule::new("r1", "ip:", "sw", 100));
581 assert!(set.disable("r1"));
582 assert!(!set.disable("nonexistent"));
583 assert_eq!(set.enabled_count(), 0);
584 assert!(set.enable("r1"));
585 assert_eq!(set.enabled_count(), 1);
586 }
587
588 #[test]
589 fn test_composite_strategy_serde() {
590 let s = CompositeStrategy::AllMustAllow;
591 let json = serde_json::to_string(&s).unwrap();
592 let back: CompositeStrategy = serde_json::from_str(&json).unwrap();
593 assert_eq!(s, back);
594 }
595}