Skip to main content

tower_rate_limiter/store/
memory.rs

1//! Memory fixed-window rate-limit store.
2
3use moka::{ops::compute::Op, policy::Expiry, sync::Cache};
4use std::{
5    future::{Ready, ready},
6    time::{Duration, Instant},
7};
8
9use crate::{RateLimitError, Store, Usage};
10
11/// Errors returned by the in-memory rate limit store.
12///
13/// # Examples
14///
15/// ```rust
16/// use tower_rate_limiter::MemoryStoreError;
17///
18/// let err = MemoryStoreError::InstantOutOfRange;
19/// // does something
20/// assert_eq!(err, MemoryStoreError::InstantOutOfRange);
21/// ```
22#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
23pub enum MemoryStoreError {
24    #[error("instant is out of range")]
25    InstantOutOfRange,
26
27    #[error("failed to compute entry")]
28    FailedToCompute,
29}
30
31/// Convert the MemoryStoreError to a RateLimitError.
32impl From<MemoryStoreError> for RateLimitError {
33    fn from(err: MemoryStoreError) -> Self {
34        RateLimitError::Store("memory_store_error".into(), err.to_string())
35    }
36}
37
38#[derive(Clone, Debug)]
39struct Entry {
40    used: u64,
41    expires_at: Instant,
42}
43
44#[derive(Clone, Copy, Debug)]
45struct EntryExpiry;
46
47impl Expiry<String, Entry> for EntryExpiry {
48    /// Returns the duration until the entry expires.
49    fn expire_after_create(&self, _key: &String, entry: &Entry, created_at: Instant) -> Option<Duration> {
50        Some(entry.expires_at.saturating_duration_since(created_at))
51    }
52
53    /// Returns the duration until the entry expires.
54    fn expire_after_update(
55        &self,
56        _key: &String,
57        entry: &Entry,
58        updated_at: Instant,
59        _duration_until_expiry: Option<Duration>,
60    ) -> Option<Duration> {
61        Some(entry.expires_at.saturating_duration_since(updated_at))
62    }
63}
64
65/// In-process Store whose clones share one counter set.
66///
67/// Entries expire with their fixed windows. Window durations use [`Duration`] and [`Instant`]
68/// directly, without reducing their precision before storage.
69///
70/// # Examples
71///
72/// ```rust
73/// use std::time::Duration;
74/// use tower_rate_limiter::MemoryStore;
75///
76/// let store = MemoryStore::new();
77/// ```
78#[derive(Clone, Debug)]
79pub struct MemoryStore {
80    cache: Cache<String, Entry>,
81}
82
83impl Default for MemoryStore {
84    /// Creates a new MemoryStore with default settings.
85    fn default() -> Self {
86        Self::new()
87    }
88}
89
90impl MemoryStore {
91    /// Creates a new MemoryStore.
92    pub fn new() -> Self {
93        Self {
94            cache: Cache::builder().expire_after(EntryExpiry).build(),
95        }
96    }
97}
98
99/// Implement the Store trait for the MemoryStore.
100impl Store for MemoryStore {
101    type Future = Ready<Result<Usage, RateLimitError>>;
102
103    fn increment(&self, key: &str, window: Duration) -> Self::Future {
104        ready(self.increment_usage(key, window).map_err(Into::into))
105    }
106}
107
108impl MemoryStore {
109    /// Increments the usage counter for the given key.
110    ///
111    /// Resets the counter when the fixed window expires.
112    ///
113    /// # Errors
114    /// Returns an error if the instant is out of range or the cache computation returns no entry.
115    fn increment_usage(&self, key: &str, window: Duration) -> Result<Usage, MemoryStoreError> {
116        let entry = self
117            .cache
118            .entry_by_ref(key)
119            .and_try_compute_with(|entry| {
120                let now = Instant::now();
121                let mut entry = match entry {
122                    Some(entry) => entry.into_value(),
123                    None => Entry {
124                        used: 0,
125                        expires_at: now.checked_add(window).ok_or(MemoryStoreError::InstantOutOfRange)?,
126                    },
127                };
128
129                if now >= entry.expires_at {
130                    entry.used = 0;
131                    entry.expires_at = now.checked_add(window).ok_or(MemoryStoreError::InstantOutOfRange)?;
132                }
133
134                entry.used = entry.used.saturating_add(1);
135                Ok(Op::Put(entry))
136            })?
137            .into_entry()
138            .map(|entry| entry.into_value())
139            .ok_or(MemoryStoreError::FailedToCompute)?;
140
141        let now = Instant::now();
142
143        Ok(Usage {
144            used: entry.used,
145            reset_after: entry.expires_at.saturating_duration_since(now),
146        })
147    }
148}
149
150#[cfg(test)]
151mod tests {
152    use super::*;
153    use std::{sync::Arc, thread, time::Duration};
154
155    #[test]
156    fn test_increment_usage() {
157        let key = "user1";
158        let store = MemoryStore::new();
159
160        let window = Duration::from_secs(60);
161
162        let usage = store.increment_usage(key, window).unwrap();
163
164        assert_eq!(usage.used, 1);
165
166        let usage = store.increment_usage(key, window).unwrap();
167
168        assert_eq!(usage.used, 2);
169    }
170
171    #[test]
172    fn test_fixed_window_reset() {
173        let store = MemoryStore::new();
174        let key = "user1";
175        let window = Duration::from_millis(20);
176
177        let usage = store.increment_usage(key, window).unwrap();
178
179        assert_eq!(usage.used, 1);
180
181        let usage = store.increment_usage(key, window).unwrap();
182
183        assert_eq!(usage.used, 2);
184
185        // wait until window expired
186        thread::sleep(Duration::from_millis(30));
187
188        let usage = store.increment_usage(key, window).unwrap();
189
190        assert_eq!(usage.used, 1);
191    }
192
193    #[test]
194    fn test_clone_share_cache() {
195        let store = MemoryStore::new();
196
197        let cloned = store.clone();
198
199        let window = Duration::from_secs(60);
200
201        let usage = store.increment_usage("user1", window).unwrap();
202
203        assert_eq!(usage.used, 1);
204
205        let usage = cloned.increment_usage("user1", window).unwrap();
206
207        assert_eq!(usage.used, 2);
208    }
209
210    #[test]
211    fn test_concurrent_increment() {
212        let store = Arc::new(MemoryStore::new());
213
214        let mut handles = Vec::new();
215
216        for _ in 0..10 {
217            let store = store.clone();
218
219            handles.push(thread::spawn(move || {
220                store.increment_usage("user1", Duration::from_secs(60)).unwrap();
221            }));
222        }
223
224        for handle in handles {
225            handle.join().unwrap();
226        }
227
228        let usage = store.increment_usage("user1", Duration::from_secs(60)).unwrap();
229
230        assert_eq!(usage.used, 11);
231    }
232
233    #[test]
234    fn new_window_starts_after_waiting_for_the_same_key_update() {
235        use std::sync::mpsc;
236
237        let store = MemoryStore::new();
238        let key = "contended-expiring";
239        let initial_window = Duration::from_millis(300);
240        let new_window = Duration::from_millis(500);
241
242        assert_eq!(store.increment_usage(key, initial_window).unwrap().used, 1);
243
244        let blocker = store.clone();
245        let (locked_tx, locked_rx) = mpsc::channel();
246        let (release_tx, release_rx) = mpsc::channel();
247        let holding_key = thread::spawn(move || {
248            blocker.cache.entry_by_ref(key).and_upsert_with(|entry| {
249                locked_tx.send(()).unwrap();
250                release_rx.recv().unwrap();
251                entry.expect("seeded entry").into_value()
252            });
253        });
254
255        locked_rx.recv().unwrap();
256        let incrementing_store = store.clone();
257        let incrementing = thread::spawn(move || incrementing_store.increment_usage(key, new_window).unwrap());
258
259        thread::sleep(Duration::from_millis(400));
260        release_tx.send(()).unwrap();
261
262        holding_key.join().unwrap();
263        let usage = incrementing.join().unwrap();
264
265        assert_eq!(usage.used, 1);
266        assert!(
267            usage.reset_after > Duration::from_millis(400),
268            "new window was shortened while waiting: {usage:?}"
269        );
270
271        thread::sleep(Duration::from_millis(200));
272        let next_usage = store.increment_usage(key, new_window).unwrap();
273        assert_eq!(
274            next_usage.used, 2,
275            "new window expired before its requested duration: {next_usage:?}"
276        );
277    }
278
279    #[test]
280    fn fixed_window_expiry_removes_the_cached_entry() {
281        let store = MemoryStore::new();
282
283        let window = Duration::from_millis(20);
284
285        assert_eq!(store.increment_usage("user-1", window).unwrap().used, 1);
286        thread::sleep(Duration::from_millis(30));
287
288        assert!(!store.cache.contains_key("user-1"));
289        let deadline = Instant::now() + Duration::from_secs(2);
290        while store.cache.entry_count() != 0 && Instant::now() < deadline {
291            store.cache.run_pending_tasks();
292            thread::sleep(Duration::from_millis(10));
293        }
294        assert_eq!(store.cache.entry_count(), 0);
295    }
296}