Skip to main content

tower_rate_tier/
tier.rs

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/// Error returned by [`RateTier::check()`].
16///
17/// Distinguishes between an unknown tier name (a configuration/logic error)
18/// and a storage backend failure (e.g., Redis connection lost).
19#[derive(Debug)]
20#[non_exhaustive]
21pub enum CheckError {
22    /// The tier name passed to `check()` does not exist in the configured tiers.
23    UnknownTier(String),
24    /// The storage backend failed during the rate limit check.
25    Storage(StorageError),
26    /// The request costs more than the tier allows in a whole window, so it
27    /// can never be allowed. Nothing was consumed.
28    #[non_exhaustive]
29    CostExceedsLimit {
30        /// The cost of the rejected request.
31        cost: u32,
32        /// The tier's maximum burst.
33        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
68/// Tier-based rate limiter configuration.
69///
70/// Maps tier names to quotas and provides a programmatic `check()` API.
71///
72/// # Examples
73///
74/// ```
75/// # #[tokio::main]
76/// # async fn main() {
77/// use tower_rate_tier::{RateTier, Quota};
78///
79/// let limiter = RateTier::builder()
80///     .tier("free", Quota::per_hour(100))
81///     .tier("pro", Quota::per_hour(5_000))
82///     .tier("enterprise", Quota::unlimited())
83///     .default_tier("free")
84///     .build();
85/// # }
86/// ```
87pub 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
97/// Garbage collector for the built-in `MemoryStorage`.
98///
99/// The task needs a Tokio runtime, so it is spawned by `build()` when one is
100/// available and otherwise by the first use of the storage inside one.
101struct LazyGc {
102    storage: Arc<MemoryStorage>,
103    interval: Duration,
104    handle: OnceLock<GcHandle>,
105}
106
107impl LazyGc {
108    /// Spawn the GC task once, if a Tokio runtime is available here.
109    ///
110    /// Without a runtime this does nothing, and a later call tries again.
111    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    /// Returns a new [`RateTierBuilder`] for configuring tiers and quotas.
134    pub fn builder() -> RateTierBuilder {
135        RateTierBuilder::default()
136    }
137
138    /// Look up the quota for a tier name.
139    pub fn get_quota(&self, tier_name: &str) -> Option<&Quota> {
140        self.tiers.get(tier_name)
141    }
142
143    /// Get the on_missing policy.
144    pub fn on_missing(&self) -> OnMissing {
145        self.on_missing
146    }
147
148    /// Get the policy for tiers that are not configured.
149    pub fn on_unknown_tier(&self) -> OnUnknownTier {
150        self.on_unknown_tier
151    }
152
153    /// Get the default tier name, if set.
154    pub fn default_tier(&self) -> Option<&str> {
155        self.default_tier.as_deref()
156    }
157
158    /// Get a reference to the clock.
159    pub fn clock(&self) -> &dyn Clock {
160        self.clock.as_ref()
161    }
162
163    /// Get a reference to the storage backend.
164    ///
165    /// Also starts the garbage collector of the built-in storage if it is not
166    /// running yet and a Tokio runtime is available (see
167    /// [`RateTierBuilder::build`]).
168    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    /// Programmatic rate limit check (non-HTTP).
176    ///
177    /// Returns `Ok(Ok(info))` if allowed, `Ok(Err(limited))` if denied,
178    /// or `Err(CheckError)` if the tier is unknown or the storage backend fails.
179    /// Unlimited tiers always return `Ok(Ok(...))` without touching storage.
180    ///
181    /// # Errors
182    ///
183    /// - [`CheckError::UnknownTier`] — the tier name does not exist in the configured tiers.
184    /// - [`CheckError::Storage`] — the storage backend failed (e.g., Redis connection lost).
185    /// - [`CheckError::CostExceedsLimit`] — `cost` is above the tier's maximum
186    ///   burst, so the request could never be allowed.
187    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
222/// Builder for `RateTier`.
223pub 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    /// Define a tier with the given name and quota.
266    pub fn tier(mut self, name: impl Into<String>, quota: Quota) -> Self {
267        self.tiers.insert(name.into(), quota);
268        self
269    }
270
271    /// Set the default tier name.
272    ///
273    /// Used by [`OnMissing::UseDefault`] for unidentified requests and by
274    /// [`OnUnknownTier::UseDefault`] for tiers that are not configured.
275    /// Without one, `OnUnknownTier::UseDefault` answers 403.
276    pub fn default_tier(mut self, name: impl Into<String>) -> Self {
277        self.default_tier = Some(name.into());
278        self
279    }
280
281    /// Set the behavior when the identifier returns `None`.
282    pub fn on_missing(mut self, policy: OnMissing) -> Self {
283        self.on_missing = policy;
284        self
285    }
286
287    /// Set the behavior when the identifier returns a tier that is not
288    /// configured. Default: [`OnUnknownTier::UseDefault`].
289    pub fn on_unknown_tier(mut self, policy: OnUnknownTier) -> Self {
290        self.on_unknown_tier = policy;
291        self
292    }
293
294    /// Set a custom clock (useful for testing with `FakeClock`).
295    pub fn clock(mut self, clock: impl Clock) -> Self {
296        self.clock = Some(Arc::new(clock));
297        self
298    }
299
300    /// Set a custom storage backend.
301    ///
302    /// Garbage collection never runs for a custom storage, even if
303    /// [`gc_interval`](Self::gc_interval) is called afterwards: custom backends
304    /// are expected to manage their own expiry (e.g., Redis TTL). This also
305    /// applies to a [`MemoryStorage`] passed here; spawn
306    /// [`GcHandle::spawn`] for it yourself if you need cleanup.
307    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    /// Set the garbage collection interval for expired entries.
314    ///
315    /// Default: 60 seconds. Only applies to the built-in storage; it has no
316    /// effect once [`storage`](Self::storage) has been called.
317    ///
318    /// # Panics
319    ///
320    /// Panics if `interval` is zero.
321    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    /// Disable automatic garbage collection of expired entries.
329    ///
330    /// Useful when using a storage backend that manages its own expiry
331    /// (e.g., Redis with TTL).
332    pub fn disable_gc(mut self) -> Self {
333        self.gc_enabled = false;
334        self
335    }
336
337    /// Build the `RateTier` configuration.
338    ///
339    /// Inside a Tokio runtime, the garbage collector of the built-in storage
340    /// starts here. Building also works outside a runtime: the collector then
341    /// starts on the first use of the storage (a check, the middleware or
342    /// [`RateTier::storage`]) that happens inside one, and expired entries are
343    /// kept until then.
344    ///
345    /// # Panics
346    ///
347    /// - If no tiers are defined.
348    /// - If `default_tier` references a non-existent tier.
349    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        // Use provided storage or default to MemoryStorage.
363        // GC only runs for the default MemoryStorage and when gc_enabled is true.
364        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            // The first tick fires at once, while the entry is still live.
428            tokio::task::yield_now().await;
429            assert_eq!(gc.storage.len(), 1);
430
431            // Expire the entry, then move Tokio's paused clock to the next tick.
432            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}