Skip to main content

sz_orm_limit/
composite.rs

1//! 限流策略组合器(Composite Limiter)
2//!
3//! 提供更丰富的限流器组合策略,包括链式、回退、权重等。
4
5use std::sync::atomic::{AtomicU64, Ordering};
6use std::sync::Arc;
7use std::time::Duration;
8
9use crate::{RateLimitError, RateLimitResult, RateLimiter};
10
11/// 组合策略
12#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
13pub enum CompositeStrategy {
14    /// 全部通过才允许(最严格)
15    AllMustAllow,
16    /// 任一通过即允许(最宽松)
17    AnyCanAllow,
18    /// 按优先级,第一个拒绝即拒绝
19    FirstRejectWins,
20    /// 按权重投票
21    WeightedVote,
22}
23
24/// 组合限流器
25///
26/// 将多个限流器按指定策略组合。
27pub struct CompositeLimiter {
28    limiters: Vec<Arc<dyn RateLimiter>>,
29    strategy: CompositeStrategy,
30    weights: Vec<u32>,
31    total_calls: AtomicU64,
32}
33
34impl CompositeLimiter {
35    /// 创建组合限流器
36    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    /// 设置权重(仅 WeightedVote 策略生效)
47    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    /// 添加限流器
55    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    /// 限流器数量
62    pub fn limiter_count(&self) -> usize {
63        self.limiters.len()
64    }
65
66    /// 策略
67    pub fn strategy(&self) -> CompositeStrategy {
68        self.strategy
69    }
70
71    /// 检查
72    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    /// 总调用次数
160    pub fn total_calls(&self) -> u64 {
161        self.total_calls.load(Ordering::Relaxed)
162    }
163}
164
165/// 回退限流器
166///
167/// 主限流器失败时回退到备用限流器。
168pub struct FallbackLimiter {
169    primary: Arc<dyn RateLimiter>,
170    fallback: Arc<dyn RateLimiter>,
171    fallback_count: AtomicU64,
172}
173
174impl FallbackLimiter {
175    /// 创建回退限流器
176    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    /// 检查,主限流器出错时回退
185    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    /// 回退次数
196    pub fn fallback_count(&self) -> u64 {
197        self.fallback_count.load(Ordering::Relaxed)
198    }
199}
200
201/// 限流键构建器
202///
203/// 从多个维度构建限流键,如 IP + 用户 + API。
204#[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    /// 创建构建器
218    pub fn new() -> Self {
219        Self {
220            parts: Vec::new(),
221            separator: ":".to_string(),
222        }
223    }
224
225    /// 设置分隔符
226    pub fn with_separator(mut self, sep: &str) -> Self {
227        self.separator = sep.to_string();
228        self
229    }
230
231    /// 添加 IP 维度
232    pub fn ip(mut self, ip: &str) -> Self {
233        self.parts.push(format!("ip:{}", ip));
234        self
235    }
236
237    /// 添加用户维度
238    pub fn user(mut self, user: &str) -> Self {
239        self.parts.push(format!("user:{}", user));
240        self
241    }
242
243    /// 添加 API 维度
244    pub fn api(mut self, api: &str) -> Self {
245        self.parts.push(format!("api:{}", api));
246        self
247    }
248
249    /// 添加自定义维度
250    pub fn dimension(mut self, name: &str, value: &str) -> Self {
251        self.parts.push(format!("{}:{}", name, value));
252        self
253    }
254
255    /// 构建限流键
256    pub fn build(&self) -> String {
257        self.parts.join(&self.separator)
258    }
259
260    /// 部分数量
261    pub fn part_count(&self) -> usize {
262        self.parts.len()
263    }
264}
265
266/// 限流规则
267///
268/// 描述一条限流规则,包括匹配条件和限流参数。
269#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
270pub struct RateLimitRule {
271    /// 规则名称
272    pub name: String,
273    /// 匹配的 key 前缀
274    pub key_prefix: String,
275    /// 限流器类型
276    pub limiter_type: String,
277    /// 容量
278    pub capacity: u64,
279    /// 窗口大小(毫秒)
280    pub window_ms: u64,
281    /// 是否启用
282    pub enabled: bool,
283}
284
285impl RateLimitRule {
286    /// 创建规则
287    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    /// 设置窗口大小
299    pub fn with_window_ms(mut self, ms: u64) -> Self {
300        self.window_ms = ms;
301        self
302    }
303
304    /// 禁用
305    pub fn disable(mut self) -> Self {
306        self.enabled = false;
307        self
308    }
309
310    /// 检查 key 是否匹配
311    pub fn matches(&self, key: &str) -> bool {
312        self.enabled && key.starts_with(&self.key_prefix)
313    }
314}
315
316/// 规则集
317pub struct RuleSet {
318    rules: Vec<RateLimitRule>,
319}
320
321impl RuleSet {
322    /// 创建规则集
323    pub fn new() -> Self {
324        Self { rules: Vec::new() }
325    }
326
327    /// 添加规则
328    pub fn add(&mut self, rule: RateLimitRule) -> &mut Self {
329        self.rules.push(rule);
330        self
331    }
332
333    /// 查找匹配的规则
334    pub fn find_matches(&self, key: &str) -> Vec<&RateLimitRule> {
335        self.rules.iter().filter(|r| r.matches(key)).collect()
336    }
337
338    /// 规则数量
339    pub fn rule_count(&self) -> usize {
340        self.rules.len()
341    }
342
343    /// 启用的规则数
344    pub fn enabled_count(&self) -> usize {
345        self.rules.iter().filter(|r| r.enabled).count()
346    }
347
348    /// 按前缀查找
349    pub fn find_by_prefix(&self, prefix: &str) -> Option<&RateLimitRule> {
350        self.rules.iter().find(|r| r.key_prefix == prefix)
351    }
352
353    /// 禁用规则
354    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    /// 启用规则
365    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,使 AnyCanAllow 第二次 check 时两个都满
421        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}