1use std::sync::atomic::{AtomicU64, Ordering};
6use std::sync::Arc;
7
8use crate::{RateLimitError, RateLimitResult, RateLimiter};
9
10#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
12pub enum CompositeStrategy {
13 AllMustAllow,
15 AnyCanAllow,
17 FirstRejectWins,
19 WeightedVote,
21}
22
23pub struct CompositeLimiter {
27 limiters: Vec<Arc<dyn RateLimiter>>,
28 strategy: CompositeStrategy,
29 weights: Vec<u32>,
30 total_calls: AtomicU64,
31}
32
33impl CompositeLimiter {
34 pub fn new(limiters: Vec<Arc<dyn RateLimiter>>, strategy: CompositeStrategy) -> Self {
36 let count = limiters.len();
37 Self {
38 limiters,
39 strategy,
40 weights: vec![1; count],
41 total_calls: AtomicU64::new(0),
42 }
43 }
44
45 pub fn with_weights(mut self, weights: Vec<u32>) -> Self {
47 if weights.len() == self.limiters.len() {
48 self.weights = weights;
49 }
50 self
51 }
52
53 pub fn with_limiter(mut self, limiter: Arc<dyn RateLimiter>) -> Self {
55 self.limiters.push(limiter);
56 self.weights.push(1);
57 self
58 }
59
60 pub fn limiter_count(&self) -> usize {
62 self.limiters.len()
63 }
64
65 pub fn strategy(&self) -> CompositeStrategy {
67 self.strategy
68 }
69
70 pub fn check(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
72 self.total_calls.fetch_add(1, Ordering::Relaxed);
73 if self.limiters.is_empty() {
74 return Err(RateLimitError::Internal(
75 "No limiters configured".to_string(),
76 ));
77 }
78 match self.strategy {
79 CompositeStrategy::AllMustAllow => self.check_all_must_allow(key),
80 CompositeStrategy::AnyCanAllow => self.check_any_can_allow(key),
81 CompositeStrategy::FirstRejectWins => self.check_first_reject_wins(key),
82 CompositeStrategy::WeightedVote => self.check_weighted_vote(key),
83 }
84 }
85
86 fn check_all_must_allow(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
87 let mut min_remaining = u64::MAX;
88 let mut max_reset = 0i64;
89 for limiter in &self.limiters {
90 let result = limiter.acquire(key)?;
91 if !result.allowed {
92 return Ok(result);
93 }
94 min_remaining = min_remaining.min(result.remaining);
95 max_reset = max_reset.max(result.reset_at);
96 }
97 Ok(RateLimitResult::allowed(min_remaining, max_reset))
98 }
99
100 fn check_any_can_allow(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
101 let mut best: Option<RateLimitResult> = None;
102 for limiter in &self.limiters {
103 let result = limiter.acquire(key)?;
104 if result.allowed {
105 return Ok(result);
106 }
107 match &best {
108 None => best = Some(result),
109 Some(current) => {
110 if result.remaining > current.remaining {
111 best = Some(result);
112 }
113 }
114 }
115 }
116 best.ok_or_else(|| RateLimitError::Internal("No results".to_string()))
117 }
118
119 fn check_first_reject_wins(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
120 let mut last_allowed: Option<RateLimitResult> = None;
121 for limiter in &self.limiters {
122 let result = limiter.acquire(key)?;
123 if !result.allowed {
124 return Ok(result);
125 }
126 last_allowed = Some(result);
127 }
128 last_allowed.ok_or_else(|| RateLimitError::Internal("No results".to_string()))
129 }
130
131 fn check_weighted_vote(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
132 let mut allow_weight = 0u32;
133 let mut reject_weight = 0u32;
134 let mut best_allowed: Option<RateLimitResult> = None;
135 let mut best_rejected: Option<RateLimitResult> = None;
136 for (i, limiter) in self.limiters.iter().enumerate() {
137 let result = limiter.acquire(key)?;
138 let weight = self.weights.get(i).copied().unwrap_or(1);
139 if result.allowed {
140 allow_weight += weight;
141 if best_allowed.is_none() {
142 best_allowed = Some(result);
143 }
144 } else {
145 reject_weight += weight;
146 if best_rejected.is_none() {
147 best_rejected = Some(result);
148 }
149 }
150 }
151 if allow_weight >= reject_weight {
152 best_allowed.ok_or_else(|| RateLimitError::Internal("No allowed".to_string()))
153 } else {
154 best_rejected.ok_or_else(|| RateLimitError::Internal("No rejected".to_string()))
155 }
156 }
157
158 pub fn total_calls(&self) -> u64 {
160 self.total_calls.load(Ordering::Relaxed)
161 }
162}
163
164pub struct FallbackLimiter {
168 primary: Arc<dyn RateLimiter>,
169 fallback: Arc<dyn RateLimiter>,
170 fallback_count: AtomicU64,
171}
172
173impl FallbackLimiter {
174 pub fn new(primary: Arc<dyn RateLimiter>, fallback: Arc<dyn RateLimiter>) -> Self {
176 Self {
177 primary,
178 fallback,
179 fallback_count: AtomicU64::new(0),
180 }
181 }
182
183 pub fn check(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
185 match self.primary.acquire(key) {
186 Ok(result) => Ok(result),
187 Err(_) => {
188 self.fallback_count.fetch_add(1, Ordering::Relaxed);
189 self.fallback.acquire(key)
190 }
191 }
192 }
193
194 pub fn fallback_count(&self) -> u64 {
196 self.fallback_count.load(Ordering::Relaxed)
197 }
198}
199
200#[derive(Debug, Clone)]
204pub struct LimitKeyBuilder {
205 parts: Vec<String>,
206 separator: String,
207}
208
209impl Default for LimitKeyBuilder {
210 fn default() -> Self {
211 Self::new()
212 }
213}
214
215impl LimitKeyBuilder {
216 pub fn new() -> Self {
218 Self {
219 parts: Vec::new(),
220 separator: ":".to_string(),
221 }
222 }
223
224 pub fn with_separator(mut self, sep: &str) -> Self {
226 self.separator = sep.to_string();
227 self
228 }
229
230 pub fn ip(mut self, ip: &str) -> Self {
232 self.parts.push(format!("ip:{}", ip));
233 self
234 }
235
236 pub fn user(mut self, user: &str) -> Self {
238 self.parts.push(format!("user:{}", user));
239 self
240 }
241
242 pub fn api(mut self, api: &str) -> Self {
244 self.parts.push(format!("api:{}", api));
245 self
246 }
247
248 pub fn dimension(mut self, name: &str, value: &str) -> Self {
250 self.parts.push(format!("{}:{}", name, value));
251 self
252 }
253
254 pub fn build(&self) -> String {
256 self.parts.join(&self.separator)
257 }
258
259 pub fn part_count(&self) -> usize {
261 self.parts.len()
262 }
263}
264
265#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
269pub struct RateLimitRule {
270 pub name: String,
272 pub key_prefix: String,
274 pub limiter_type: String,
276 pub capacity: u64,
278 pub window_ms: u64,
280 pub enabled: bool,
282}
283
284impl RateLimitRule {
285 pub fn new(name: &str, key_prefix: &str, limiter_type: &str, capacity: u64) -> Self {
287 Self {
288 name: name.to_string(),
289 key_prefix: key_prefix.to_string(),
290 limiter_type: limiter_type.to_string(),
291 capacity,
292 window_ms: 1000,
293 enabled: true,
294 }
295 }
296
297 pub fn with_window_ms(mut self, ms: u64) -> Self {
299 self.window_ms = ms;
300 self
301 }
302
303 pub fn disable(mut self) -> Self {
305 self.enabled = false;
306 self
307 }
308
309 pub fn matches(&self, key: &str) -> bool {
311 self.enabled && key.starts_with(&self.key_prefix)
312 }
313}
314
315pub struct RuleSet {
317 rules: Vec<RateLimitRule>,
318}
319
320impl RuleSet {
321 pub fn new() -> Self {
323 Self { rules: Vec::new() }
324 }
325
326 pub fn add(&mut self, rule: RateLimitRule) -> &mut Self {
328 self.rules.push(rule);
329 self
330 }
331
332 pub fn find_matches(&self, key: &str) -> Vec<&RateLimitRule> {
334 self.rules.iter().filter(|r| r.matches(key)).collect()
335 }
336
337 pub fn rule_count(&self) -> usize {
339 self.rules.len()
340 }
341
342 pub fn enabled_count(&self) -> usize {
344 self.rules.iter().filter(|r| r.enabled).count()
345 }
346
347 pub fn find_by_prefix(&self, prefix: &str) -> Option<&RateLimitRule> {
349 self.rules.iter().find(|r| r.key_prefix == prefix)
350 }
351
352 pub fn disable(&mut self, name: &str) -> bool {
354 for rule in &mut self.rules {
355 if rule.name == name {
356 rule.enabled = false;
357 return true;
358 }
359 }
360 false
361 }
362
363 pub fn enable(&mut self, name: &str) -> bool {
365 for rule in &mut self.rules {
366 if rule.name == name {
367 rule.enabled = true;
368 return true;
369 }
370 }
371 false
372 }
373}
374
375impl Default for RuleSet {
376 fn default() -> Self {
377 Self::new()
378 }
379}
380
381#[cfg(test)]
382mod tests {
383 use super::*;
384 use std::time::Duration;
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}