tower_rate_limiter/store/
memory.rs1use 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#[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
31impl 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 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 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#[derive(Clone, Debug)]
79pub struct MemoryStore {
80 cache: Cache<String, Entry>,
81}
82
83impl Default for MemoryStore {
84 fn default() -> Self {
86 Self::new()
87 }
88}
89
90impl MemoryStore {
91 pub fn new() -> Self {
93 Self {
94 cache: Cache::builder().expire_after(EntryExpiry).build(),
95 }
96 }
97}
98
99impl 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 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 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}