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