Skip to main content

security_rust/throttle/
store.rs

1// Copyright (c) 2026 erik <erik@erik.xyz> — https://erik.xyz
2
3use super::StoreError;
4use std::collections::HashMap;
5use std::sync::Mutex;
6
7/// 限流状态后端。
8///
9/// 所有 `now` 均为 unix 秒,调用方须保证**单调不减**(系统时钟正常满足):
10/// 滑动窗口靠比较时间戳,时间倒流会让新记录被当成过期数据丢掉。
11pub trait ThrottleStore: Send + Sync {
12    /// 记录一次失败,返回**窗口内**的累计失败数(含本次)。
13    /// 实现须自行滑动窗口:只计 `now - window_secs` 之后的失败。
14    fn record_failure(&self, key: &str, now: u64, window_secs: u64) -> Result<u32, StoreError>;
15    /// 只读查询窗口内失败数,**不写入任何状态**。
16    ///
17    /// 供 `Throttle::check` 在放行时算剩余额度 —— 没有它,请求路径只能返回一个
18    /// 乐观的满额数字(`remaining` 会被写进 X-RateLimit 响应头,谎报等于误导调用方)。
19    fn failure_count(&self, key: &str, now: u64, window_secs: u64) -> Result<u32, StoreError>;
20    /// 该 key 是否处于封禁中;是则返回解封时刻。
21    ///
22    /// 返回**未过滤的原始值也是合法的**:`Throttle::check` 会自己拿 `now` 比对
23    /// `until` 再决定是否 `Banned`,后端不必代劳(返回 `Some(until)` 且
24    /// `until <= now` 不会被当成永久封禁)。
25    fn is_banned(&self, key: &str, now: u64) -> Result<Option<u64>, StoreError>;
26    /// 写入封禁截止时刻。**只能延长不能缩短**:已生效的封禁遇上更早的 `until`
27    /// 应被忽略(时钟回拨时 `now + ban_secs` 可能反而更小)。
28    fn ban(&self, key: &str, until: u64) -> Result<(), StoreError>;
29    /// 只清空失败计数,**保留封禁**。
30    ///
31    /// 认证成功走这里,人工解封走 `reset`:二者合并会让共享桶(如 `ip:`)的封禁
32    /// 被桶里任意另一个人一次成功认证解除。
33    fn clear_failures(&self, key: &str) -> Result<(), StoreError>;
34    /// 清零该 key 的失败计数与封禁。
35    fn reset(&self, key: &str) -> Result<(), StoreError>;
36    /// 清除已过期状态,返回清除条数。
37    fn purge_expired(&self, now: u64) -> Result<usize, StoreError>;
38}
39
40/// 某个 key 的限流状态。
41///
42/// 存**时间戳**而非计数:只有留下每次失败的坐标才能滑动窗口 ——
43/// 单纯自增的计数器没法判断「哪几次失败已经滑出窗口」,只能整段清零。
44#[derive(Debug, Default)]
45struct Entry {
46    /// 窗口内的失败时刻。时钟单调时才恰好是非递减的,回拨会破坏这个序 ——
47    /// 因此判定一律扫全量,不要用 `last()`。
48    failures: Vec<u64>,
49    /// 封禁截止时刻;`None` 表示从未封禁。只增不减:`ban` 取 max,时钟回拨
50    /// 不会把已生效的封禁缩短。
51    banned_until: Option<u64>,
52    /// 最近一次 `record_failure` 传入的窗口长度,仅供 `purge_expired` 判断陈旧。
53    /// store 不记住窗口,就无法区分「窗口内仍有效的失败」与「早该滑出的失败」,
54    /// 而 `purge_expired(now)` 的签名里没有窗口参数。
55    window_secs: u64,
56}
57
58/// 内存后端。无后台线程,且**只有写路径**会清理:`record_failure` 顺手
59/// `retain` 掉滑出窗口的失败,`failure_count` / `is_banned` 这类读路径不碰状态。
60///
61/// 于是有两条边界:
62/// - 每个 key 的 `failures` 向量有界(每次 `record_failure` 都 retain,上界是
63///   窗口内的失败数);
64/// - **map 本身的条目数无上限** —— key 只增不减,只有 `purge_expired`(和
65///   `reset`)会移除条目。
66///
67/// 长期运行的进程必须按定时器调 `purge_expired`(间隔取 `window_secs` 量级即可),
68/// 否则 key 基数的增长会一直吃内存。
69#[derive(Debug)]
70pub struct MemoryThrottleStore {
71    entries: Mutex<HashMap<String, Entry>>,
72}
73
74impl Default for MemoryThrottleStore {
75    fn default() -> Self {
76        Self::new()
77    }
78}
79
80impl MemoryThrottleStore {
81    pub fn new() -> Self {
82        Self {
83            entries: Mutex::new(HashMap::new()),
84        }
85    }
86
87    /// 互斥锁获取。线程 panic 导致锁中毒时恢复内部数据而非永久 Err:
88    /// 本 store 的每个操作都是单次 HashMap 读/写(外加一次与本次读写同批完成的
89    /// `Vec::retain`),不存在「改到一半」的不变量,恢复是安全的;
90    /// 而让一次 panic 永久锁死整个限流存储是一种自我 DoS。
91    fn lock<T>(m: &Mutex<T>) -> std::sync::MutexGuard<'_, T> {
92        m.lock().unwrap_or_else(|e| e.into_inner())
93    }
94
95    /// 窗口下界:早于或等于它的失败都算滑出。
96    fn cutoff(now: u64, window_secs: u64) -> u64 {
97        now.saturating_sub(window_secs)
98    }
99}
100
101impl ThrottleStore for MemoryThrottleStore {
102    fn record_failure(&self, key: &str, now: u64, window_secs: u64) -> Result<u32, StoreError> {
103        let cutoff = Self::cutoff(now, window_secs);
104        let mut g = Self::lock(&self.entries);
105        let e = g.entry(key.to_string()).or_default();
106        e.window_secs = window_secs;
107        // ponytail: 每次失败做一次 O(窗口内失败数) 的 retain,不需要后台清理线程。
108        // 天花板 = 单个 key 在窗口内的失败数;持续暴力破解下约等于 window_secs 条
109        // (默认 60),单个 key 几 KB。若把窗口/threshold 调到万级,改环形缓冲
110        // 或分桶计数(bucket 粒度 = window/10 的固定格子)。
111        e.failures.retain(|t| *t > cutoff);
112        e.failures.push(now);
113        Ok(e.failures.len() as u32)
114    }
115
116    fn failure_count(&self, key: &str, now: u64, window_secs: u64) -> Result<u32, StoreError> {
117        let cutoff = Self::cutoff(now, window_secs);
118        let g = Self::lock(&self.entries);
119        Ok(g.get(key)
120            .map(|e| e.failures.iter().filter(|t| **t > cutoff).count() as u32)
121            .unwrap_or(0))
122    }
123
124    fn is_banned(&self, key: &str, now: u64) -> Result<Option<u64>, StoreError> {
125        Ok(Self::lock(&self.entries)
126            .get(key)
127            .and_then(|e| e.banned_until)
128            // `until` 是解封时刻:到点即解封,故严格大于
129            .filter(|until| *until > now))
130    }
131
132    fn ban(&self, key: &str, until: u64) -> Result<(), StoreError> {
133        let mut g = Self::lock(&self.entries);
134        let e = g.entry(key.to_string()).or_default();
135        // 取 max:时钟回拨后 `now + ban_secs` 可能早于已写入的 until,
136        // 无条件覆盖等于让攻击者靠回拨提前解封。
137        e.banned_until = Some(e.banned_until.map_or(until, |existing| existing.max(until)));
138        Ok(())
139    }
140
141    fn clear_failures(&self, key: &str) -> Result<(), StoreError> {
142        if let Some(e) = Self::lock(&self.entries).get_mut(key) {
143            e.failures.clear();
144        }
145        Ok(())
146    }
147
148    fn reset(&self, key: &str) -> Result<(), StoreError> {
149        Self::lock(&self.entries).remove(key);
150        Ok(())
151    }
152
153    fn purge_expired(&self, now: u64) -> Result<usize, StoreError> {
154        let mut g = Self::lock(&self.entries);
155        let before = g.len();
156        g.retain(|_, e| {
157            // 封禁未到期 ⇒ 留着;否则只要窗口内还有失败也留着。
158            // 两者都不成立才叫「过期」—— 只清封禁记录而不看失败,会把还在
159            // 计数窗口内的对手顺手洗白,等于给攻击者一个免费的计数重置。
160            let banned = e.banned_until.is_some_and(|until| until > now);
161            // 扫全量而非取 `last()`:`failures` 只在时钟单调时才是非递减的,
162            // 回拨会让 `last()` 变成最小值,把窗口内的有效失败整条丢掉。
163            let fresh = e
164                .failures
165                .iter()
166                .any(|t| *t > Self::cutoff(now, e.window_secs));
167            banned || fresh
168        });
169        Ok(before - g.len())
170    }
171}
172
173#[cfg(test)]
174mod tests {
175    use super::*;
176
177    const NOW: u64 = 1_000_000;
178
179    fn store() -> MemoryThrottleStore {
180        MemoryThrottleStore::new()
181    }
182
183    #[test]
184    fn record_failure_counts_within_window() {
185        let s = store();
186        assert_eq!(s.record_failure("k", NOW, 60).unwrap(), 1);
187        assert_eq!(s.record_failure("k", NOW + 1, 60).unwrap(), 2);
188    }
189
190    #[test]
191    fn failures_at_window_edge_are_dropped() {
192        let s = store();
193        s.record_failure("k", NOW, 60).unwrap();
194        // 恰好 window_secs 秒前的失败算滑出(保留条件是 t > now - window)
195        assert_eq!(s.record_failure("k", NOW + 60, 60).unwrap(), 1);
196        // 窗口内一秒之差则仍计入
197        assert_eq!(s.record_failure("k", NOW + 61, 60).unwrap(), 2);
198    }
199
200    #[test]
201    fn failure_count_is_read_only() {
202        let s = store();
203        assert_eq!(s.failure_count("k", NOW, 60).unwrap(), 0);
204        s.record_failure("k", NOW, 60).unwrap();
205        assert_eq!(s.failure_count("k", NOW, 60).unwrap(), 1);
206        // 查十次也不该把计数查大
207        for _ in 0..10 {
208            assert_eq!(s.failure_count("k", NOW, 60).unwrap(), 1);
209        }
210        assert_eq!(s.record_failure("k", NOW, 60).unwrap(), 2);
211    }
212
213    #[test]
214    fn failure_count_of_unknown_key_is_zero() {
215        let s = store();
216        assert_eq!(s.failure_count("ghost", NOW, 60).unwrap(), 0);
217    }
218
219    #[test]
220    fn failures_vec_stays_bounded_to_window() {
221        // 有界性:窗口外的失败必须真被丢掉,否则 Vec 无上限增长
222        let s = store();
223        for i in 0..100 {
224            s.record_failure("k", NOW + i, 60).unwrap();
225        }
226        // 每秒一次、窗口 60 秒 ⇒ 至多留下 60 条,而不是 100 条
227        assert_eq!(len(&s, "k"), 60, "窗口外的失败必须真被丢弃");
228        // 跳到窗口之外再来一次,前面 60 条也应全部被清掉
229        s.record_failure("k", NOW + 1_000, 60).unwrap();
230        assert_eq!(len(&s, "k"), 1);
231    }
232
233    #[test]
234    fn banned_expires_at_until_exclusive() {
235        let s = store();
236        s.ban("k", NOW + 100).unwrap();
237        assert_eq!(s.is_banned("k", NOW).unwrap(), Some(NOW + 100));
238        assert_eq!(s.is_banned("k", NOW + 99).unwrap(), Some(NOW + 100));
239        assert_eq!(s.is_banned("k", NOW + 100).unwrap(), None);
240    }
241
242    #[test]
243    fn is_banned_unknown_key_is_none() {
244        let s = store();
245        assert_eq!(s.is_banned("ghost", NOW).unwrap(), None);
246    }
247
248    #[test]
249    fn reset_clears_failures_and_ban() {
250        let s = store();
251        s.record_failure("k", NOW, 60).unwrap();
252        s.ban("k", NOW + 100).unwrap();
253        s.reset("k").unwrap();
254        assert_eq!(s.failure_count("k", NOW, 60).unwrap(), 0);
255        assert_eq!(s.is_banned("k", NOW).unwrap(), None);
256    }
257
258    #[test]
259    fn reset_unknown_key_is_ok() {
260        let s = store();
261        assert!(s.reset("ghost").is_ok());
262    }
263
264    #[test]
265    fn purge_keeps_banned_and_fresh_entries() {
266        let s = store();
267        s.record_failure("stale", NOW, 60).unwrap();
268        s.record_failure("fresh", NOW + 990, 60).unwrap();
269        s.ban("banned", NOW + 5_000).unwrap();
270        assert_eq!(s.purge_expired(NOW + 1_000).unwrap(), 1);
271        assert_eq!(s.failure_count("stale", NOW + 1_000, 60).unwrap(), 0);
272        assert_eq!(s.failure_count("fresh", NOW + 1_000, 60).unwrap(), 1);
273        assert!(s.is_banned("banned", NOW + 1_000).unwrap().is_some());
274    }
275
276    #[test]
277    fn purge_with_expired_ban_drops_entry() {
278        let s = store();
279        s.record_failure("k", NOW, 60).unwrap();
280        s.ban("k", NOW + 10).unwrap();
281        // 封禁已到期且失败已滑出窗口 ⇒ 状态全过期
282        assert_eq!(s.purge_expired(NOW + 100).unwrap(), 1);
283        assert_eq!(s.is_banned("k", NOW + 100).unwrap(), None);
284    }
285
286    #[test]
287    fn purge_empty_store_is_zero() {
288        let s = store();
289        assert_eq!(s.purge_expired(NOW).unwrap(), 0);
290    }
291
292    #[test]
293    fn purge_keeps_in_window_failures_after_clock_rollback() {
294        // 回拨让 failures 乱序成 [NOW, NOW-200]:取 `last()` 会拿到最小值,
295        // 把仍在窗口内的 [NOW] 判成过期 —— 攻击者只要诱发一次回拨再等一次
296        // purge,计数就被免费重置。判定必须扫全量。
297        let s = store();
298        s.record_failure("k", NOW, 60).unwrap();
299        s.record_failure("k", NOW - 200, 60).unwrap();
300        assert_eq!(s.failure_count("k", NOW, 60).unwrap(), 1);
301        assert_eq!(
302            s.purge_expired(NOW).unwrap(),
303            0,
304            "窗口内仍有有效失败,不该清"
305        );
306        assert_eq!(
307            s.failure_count("k", NOW, 60).unwrap(),
308            1,
309            "计数不能被 purge 免费重置"
310        );
311    }
312
313    #[test]
314    fn ban_cannot_be_shortened_by_clock_rollback() {
315        let s = store();
316        s.ban("k", NOW + 900).unwrap();
317        // 回拨后重新封禁:天真的 `until = now + ban_secs` 会写进一个更早的时刻
318        s.ban("k", NOW - 5_000 + 900).unwrap();
319        assert_eq!(
320            s.is_banned("k", NOW).unwrap(),
321            Some(NOW + 900),
322            "封禁只能延长,回拨不能提前解封"
323        );
324    }
325
326    #[test]
327    fn empty_key_is_a_normal_key() {
328        // 空 key 不做特殊处理:调用方负责构造非空且不重名的 key
329        let s = store();
330        assert_eq!(s.record_failure("", NOW, 60).unwrap(), 1);
331        assert_eq!(s.record_failure("", NOW, 60).unwrap(), 2);
332        assert_eq!(s.record_failure("k", NOW, 60).unwrap(), 1);
333    }
334
335    #[test]
336    fn lock_recovers_from_poisoned_mutex() {
337        // 钉住 lock() 的恢复不变量:一次 panic 不能永久锁死限流存储
338        let m = Mutex::new(Entry {
339            failures: vec![NOW],
340            banned_until: Some(NOW + 1),
341            window_secs: 60,
342        });
343        std::panic::catch_unwind(|| {
344            let _guard = m.lock().unwrap();
345            panic!("poison");
346        })
347        .unwrap_err();
348        assert!(m.is_poisoned());
349
350        let g = MemoryThrottleStore::lock(&m);
351        assert_eq!(g.failures, vec![NOW]);
352        assert_eq!(g.banned_until, Some(NOW + 1));
353    }
354
355    /// 直接读私有字段,只有单测能这么干。
356    fn len(s: &MemoryThrottleStore, key: &str) -> usize {
357        MemoryThrottleStore::lock(&s.entries)
358            .get(key)
359            .map(|e| e.failures.len())
360            .unwrap_or(0)
361    }
362}