1use std::collections::HashMap;
2use std::fmt;
3use std::sync::{Arc, OnceLock};
4use std::time::Duration;
5
6use crate::clock::{Clock, SystemClock};
7use crate::gc::GcHandle;
8use crate::gcra::{RateLimitInfo, RateLimited};
9use crate::on_missing::OnMissing;
10use crate::on_unknown_tier::OnUnknownTier;
11use crate::quota::Quota;
12use crate::storage::memory::MemoryStorage;
13use crate::storage::{Storage, StorageError, StorageKey};
14
15#[derive(Debug)]
20#[non_exhaustive]
21pub enum CheckError {
22 UnknownTier(String),
24 Storage(StorageError),
26 #[non_exhaustive]
29 CostExceedsLimit {
30 cost: u32,
32 limit: u32,
34 },
35}
36
37impl fmt::Display for CheckError {
38 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
39 match self {
40 CheckError::UnknownTier(name) => write!(f, "unknown tier: {}", name),
41 CheckError::Storage(err) => write!(f, "{}", err),
42 CheckError::CostExceedsLimit { cost, limit } => {
43 write!(
44 f,
45 "request cost {} exceeds the tier limit of {}",
46 cost, limit
47 )
48 }
49 }
50 }
51}
52
53impl std::error::Error for CheckError {
54 fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
55 match self {
56 CheckError::UnknownTier(_) | CheckError::CostExceedsLimit { .. } => None,
57 CheckError::Storage(err) => Some(err),
58 }
59 }
60}
61
62impl From<StorageError> for CheckError {
63 fn from(err: StorageError) -> Self {
64 CheckError::Storage(err)
65 }
66}
67
68pub struct RateTier {
88 tiers: HashMap<String, Quota>,
89 default_tier: Option<String>,
90 on_missing: OnMissing,
91 on_unknown_tier: OnUnknownTier,
92 storage: Arc<dyn Storage>,
93 clock: Arc<dyn Clock>,
94 gc: Option<LazyGc>,
95}
96
97struct LazyGc {
102 storage: Arc<MemoryStorage>,
103 interval: Duration,
104 handle: OnceLock<GcHandle>,
105}
106
107impl LazyGc {
108 fn ensure_started(&self, clock: &Arc<dyn Clock>) {
112 if self.handle.get().is_none() && tokio::runtime::Handle::try_current().is_ok() {
113 self.handle.get_or_init(|| {
114 GcHandle::spawn(self.storage.clone(), clock.clone(), self.interval)
115 });
116 }
117 }
118}
119
120impl fmt::Debug for RateTier {
121 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
122 f.debug_struct("RateTier")
123 .field("tiers", &self.tiers)
124 .field("default_tier", &self.default_tier)
125 .field("on_missing", &self.on_missing)
126 .field("on_unknown_tier", &self.on_unknown_tier)
127 .field("gc_enabled", &self.gc.is_some())
128 .finish_non_exhaustive()
129 }
130}
131
132impl RateTier {
133 pub fn builder() -> RateTierBuilder {
135 RateTierBuilder::default()
136 }
137
138 pub fn get_quota(&self, tier_name: &str) -> Option<&Quota> {
140 self.tiers.get(tier_name)
141 }
142
143 pub fn on_missing(&self) -> OnMissing {
145 self.on_missing
146 }
147
148 pub fn on_unknown_tier(&self) -> OnUnknownTier {
150 self.on_unknown_tier
151 }
152
153 pub fn default_tier(&self) -> Option<&str> {
155 self.default_tier.as_deref()
156 }
157
158 pub fn clock(&self) -> &dyn Clock {
160 self.clock.as_ref()
161 }
162
163 pub fn storage(&self) -> &dyn Storage {
169 if let Some(gc) = &self.gc {
170 gc.ensure_started(&self.clock);
171 }
172 self.storage.as_ref()
173 }
174
175 pub async fn check(
188 &self,
189 user_id: &str,
190 tier_name: &str,
191 cost: u32,
192 ) -> Result<Result<RateLimitInfo, RateLimited>, CheckError> {
193 let quota = self
194 .tiers
195 .get(tier_name)
196 .ok_or_else(|| CheckError::UnknownTier(tier_name.to_string()))?;
197
198 if quota.is_unlimited() {
199 return Ok(Ok(RateLimitInfo {
200 limit: 0,
201 remaining: 0,
202 reset_after: Duration::ZERO,
203 }));
204 }
205
206 if cost > quota.max_burst() {
207 return Err(CheckError::CostExceedsLimit {
208 cost,
209 limit: quota.max_burst(),
210 });
211 }
212
213 let now = self.clock.now();
214 let key = StorageKey::new(user_id, tier_name);
215 Ok(self
216 .storage()
217 .check_and_update(key, quota, cost, now)
218 .await?)
219 }
220}
221
222pub struct RateTierBuilder {
224 tiers: HashMap<String, Quota>,
225 default_tier: Option<String>,
226 on_missing: OnMissing,
227 on_unknown_tier: OnUnknownTier,
228 clock: Option<Arc<dyn Clock>>,
229 storage: Option<Arc<dyn Storage>>,
230 gc_interval: Duration,
231 gc_enabled: bool,
232}
233
234impl Default for RateTierBuilder {
235 fn default() -> Self {
236 Self {
237 tiers: HashMap::new(),
238 default_tier: None,
239 on_missing: OnMissing::default(),
240 on_unknown_tier: OnUnknownTier::default(),
241 clock: None,
242 storage: None,
243 gc_interval: Duration::from_secs(60),
244 gc_enabled: true,
245 }
246 }
247}
248
249impl fmt::Debug for RateTierBuilder {
250 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
251 f.debug_struct("RateTierBuilder")
252 .field("tiers", &self.tiers)
253 .field("default_tier", &self.default_tier)
254 .field("on_missing", &self.on_missing)
255 .field("on_unknown_tier", &self.on_unknown_tier)
256 .field("custom_clock", &self.clock.is_some())
257 .field("custom_storage", &self.storage.is_some())
258 .field("gc_interval", &self.gc_interval)
259 .field("gc_enabled", &self.gc_enabled)
260 .finish()
261 }
262}
263
264impl RateTierBuilder {
265 pub fn tier(mut self, name: impl Into<String>, quota: Quota) -> Self {
267 self.tiers.insert(name.into(), quota);
268 self
269 }
270
271 pub fn default_tier(mut self, name: impl Into<String>) -> Self {
277 self.default_tier = Some(name.into());
278 self
279 }
280
281 pub fn on_missing(mut self, policy: OnMissing) -> Self {
283 self.on_missing = policy;
284 self
285 }
286
287 pub fn on_unknown_tier(mut self, policy: OnUnknownTier) -> Self {
290 self.on_unknown_tier = policy;
291 self
292 }
293
294 pub fn clock(mut self, clock: impl Clock) -> Self {
296 self.clock = Some(Arc::new(clock));
297 self
298 }
299
300 pub fn storage(mut self, storage: Arc<dyn Storage>) -> Self {
308 self.storage = Some(storage);
309 self.gc_enabled = false;
310 self
311 }
312
313 pub fn gc_interval(mut self, interval: Duration) -> Self {
322 assert!(!interval.is_zero(), "gc interval must be non-zero");
323 self.gc_interval = interval;
324 self.gc_enabled = true;
325 self
326 }
327
328 pub fn disable_gc(mut self) -> Self {
333 self.gc_enabled = false;
334 self
335 }
336
337 pub fn build(self) -> RateTier {
350 assert!(!self.tiers.is_empty(), "at least one tier must be defined");
351
352 if let Some(ref default) = self.default_tier {
353 assert!(
354 self.tiers.contains_key(default),
355 "default tier '{}' does not exist in defined tiers",
356 default
357 );
358 }
359
360 let clock: Arc<dyn Clock> = self.clock.unwrap_or_else(|| Arc::new(SystemClock::new()));
361
362 let (storage, gc): (Arc<dyn Storage>, Option<LazyGc>) = match self.storage {
365 Some(custom) => (custom, None),
366 None => {
367 let memory = Arc::new(MemoryStorage::new());
368 let gc = self.gc_enabled.then(|| LazyGc {
369 storage: memory.clone(),
370 interval: self.gc_interval,
371 handle: OnceLock::new(),
372 });
373 (memory, gc)
374 }
375 };
376
377 let rate_tier = RateTier {
378 tiers: self.tiers,
379 default_tier: self.default_tier,
380 on_missing: self.on_missing,
381 on_unknown_tier: self.on_unknown_tier,
382 storage,
383 clock,
384 gc,
385 };
386 if let Some(gc) = &rate_tier.gc {
387 gc.ensure_started(&rate_tier.clock);
388 }
389 rate_tier
390 }
391}
392
393#[cfg(test)]
394mod tests {
395 use super::*;
396 use crate::clock::FakeClock;
397
398 #[tokio::test]
399 async fn gc_starts_at_build_inside_a_runtime() {
400 let limiter = RateTier::builder()
401 .tier("free", Quota::per_second(1))
402 .build();
403 let gc = limiter.gc.as_ref().expect("built-in storage enables GC");
404 assert!(gc.handle.get().is_some());
405 }
406
407 #[test]
408 fn gc_built_outside_a_runtime_starts_on_first_use_and_repeats() {
409 let clock = FakeClock::new();
410 let limiter = RateTier::builder()
411 .tier("free", Quota::per_second(1))
412 .clock(clock.clone())
413 .gc_interval(Duration::from_secs(1))
414 .build();
415 let gc = limiter.gc.as_ref().expect("built-in storage enables GC");
416 assert!(gc.handle.get().is_none(), "no runtime, so no GC yet");
417
418 let rt = tokio::runtime::Builder::new_current_thread()
419 .enable_time()
420 .start_paused(true)
421 .build()
422 .unwrap();
423 rt.block_on(async {
424 limiter.check("u1", "free", 1).await.unwrap().unwrap();
425 assert!(gc.handle.get().is_some(), "first use starts the GC");
426
427 tokio::task::yield_now().await;
429 assert_eq!(gc.storage.len(), 1);
430
431 clock.advance(Duration::from_secs(10));
433 tokio::time::advance(Duration::from_secs(1)).await;
434 tokio::task::yield_now().await;
435 assert_eq!(gc.storage.len(), 0, "a later tick must collect it");
436 });
437 }
438}