use std::collections::HashMap;
use std::fmt;
use std::sync::{Arc, OnceLock};
use std::time::Duration;
use crate::clock::{Clock, SystemClock};
use crate::gc::GcHandle;
use crate::gcra::{RateLimitInfo, RateLimited};
use crate::on_missing::OnMissing;
use crate::on_unknown_tier::OnUnknownTier;
use crate::quota::Quota;
use crate::storage::memory::MemoryStorage;
use crate::storage::{Storage, StorageError, StorageKey};
#[derive(Debug)]
#[non_exhaustive]
pub enum CheckError {
UnknownTier(String),
Storage(StorageError),
#[non_exhaustive]
CostExceedsLimit {
cost: u32,
limit: u32,
},
}
impl fmt::Display for CheckError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
CheckError::UnknownTier(name) => write!(f, "unknown tier: {}", name),
CheckError::Storage(err) => write!(f, "{}", err),
CheckError::CostExceedsLimit { cost, limit } => {
write!(
f,
"request cost {} exceeds the tier limit of {}",
cost, limit
)
}
}
}
}
impl std::error::Error for CheckError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
CheckError::UnknownTier(_) | CheckError::CostExceedsLimit { .. } => None,
CheckError::Storage(err) => Some(err),
}
}
}
impl From<StorageError> for CheckError {
fn from(err: StorageError) -> Self {
CheckError::Storage(err)
}
}
pub struct RateTier {
tiers: HashMap<String, Quota>,
default_tier: Option<String>,
on_missing: OnMissing,
on_unknown_tier: OnUnknownTier,
storage: Arc<dyn Storage>,
clock: Arc<dyn Clock>,
gc: Option<LazyGc>,
}
struct LazyGc {
storage: Arc<MemoryStorage>,
interval: Duration,
handle: OnceLock<GcHandle>,
}
impl LazyGc {
fn ensure_started(&self, clock: &Arc<dyn Clock>) {
if self.handle.get().is_none() && tokio::runtime::Handle::try_current().is_ok() {
self.handle.get_or_init(|| {
GcHandle::spawn(self.storage.clone(), clock.clone(), self.interval)
});
}
}
}
impl fmt::Debug for RateTier {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("RateTier")
.field("tiers", &self.tiers)
.field("default_tier", &self.default_tier)
.field("on_missing", &self.on_missing)
.field("on_unknown_tier", &self.on_unknown_tier)
.field("gc_enabled", &self.gc.is_some())
.finish_non_exhaustive()
}
}
impl RateTier {
pub fn builder() -> RateTierBuilder {
RateTierBuilder::default()
}
pub fn get_quota(&self, tier_name: &str) -> Option<&Quota> {
self.tiers.get(tier_name)
}
pub fn on_missing(&self) -> OnMissing {
self.on_missing
}
pub fn on_unknown_tier(&self) -> OnUnknownTier {
self.on_unknown_tier
}
pub fn default_tier(&self) -> Option<&str> {
self.default_tier.as_deref()
}
pub fn clock(&self) -> &dyn Clock {
self.clock.as_ref()
}
pub fn storage(&self) -> &dyn Storage {
if let Some(gc) = &self.gc {
gc.ensure_started(&self.clock);
}
self.storage.as_ref()
}
pub async fn check(
&self,
user_id: &str,
tier_name: &str,
cost: u32,
) -> Result<Result<RateLimitInfo, RateLimited>, CheckError> {
let quota = self
.tiers
.get(tier_name)
.ok_or_else(|| CheckError::UnknownTier(tier_name.to_string()))?;
if quota.is_unlimited() {
return Ok(Ok(RateLimitInfo {
limit: 0,
remaining: 0,
reset_after: Duration::ZERO,
}));
}
if cost > quota.max_burst() {
return Err(CheckError::CostExceedsLimit {
cost,
limit: quota.max_burst(),
});
}
let now = self.clock.now();
let key = StorageKey::new(user_id, tier_name);
Ok(self
.storage()
.check_and_update(key, quota, cost, now)
.await?)
}
}
pub struct RateTierBuilder {
tiers: HashMap<String, Quota>,
default_tier: Option<String>,
on_missing: OnMissing,
on_unknown_tier: OnUnknownTier,
clock: Option<Arc<dyn Clock>>,
storage: Option<Arc<dyn Storage>>,
gc_interval: Duration,
gc_enabled: bool,
}
impl Default for RateTierBuilder {
fn default() -> Self {
Self {
tiers: HashMap::new(),
default_tier: None,
on_missing: OnMissing::default(),
on_unknown_tier: OnUnknownTier::default(),
clock: None,
storage: None,
gc_interval: Duration::from_secs(60),
gc_enabled: true,
}
}
}
impl fmt::Debug for RateTierBuilder {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("RateTierBuilder")
.field("tiers", &self.tiers)
.field("default_tier", &self.default_tier)
.field("on_missing", &self.on_missing)
.field("on_unknown_tier", &self.on_unknown_tier)
.field("custom_clock", &self.clock.is_some())
.field("custom_storage", &self.storage.is_some())
.field("gc_interval", &self.gc_interval)
.field("gc_enabled", &self.gc_enabled)
.finish()
}
}
impl RateTierBuilder {
pub fn tier(mut self, name: impl Into<String>, quota: Quota) -> Self {
self.tiers.insert(name.into(), quota);
self
}
pub fn default_tier(mut self, name: impl Into<String>) -> Self {
self.default_tier = Some(name.into());
self
}
pub fn on_missing(mut self, policy: OnMissing) -> Self {
self.on_missing = policy;
self
}
pub fn on_unknown_tier(mut self, policy: OnUnknownTier) -> Self {
self.on_unknown_tier = policy;
self
}
pub fn clock(mut self, clock: impl Clock) -> Self {
self.clock = Some(Arc::new(clock));
self
}
pub fn storage(mut self, storage: Arc<dyn Storage>) -> Self {
self.storage = Some(storage);
self.gc_enabled = false;
self
}
pub fn gc_interval(mut self, interval: Duration) -> Self {
assert!(!interval.is_zero(), "gc interval must be non-zero");
self.gc_interval = interval;
self.gc_enabled = true;
self
}
pub fn disable_gc(mut self) -> Self {
self.gc_enabled = false;
self
}
pub fn build(self) -> RateTier {
assert!(!self.tiers.is_empty(), "at least one tier must be defined");
if let Some(ref default) = self.default_tier {
assert!(
self.tiers.contains_key(default),
"default tier '{}' does not exist in defined tiers",
default
);
}
let clock: Arc<dyn Clock> = self.clock.unwrap_or_else(|| Arc::new(SystemClock::new()));
let (storage, gc): (Arc<dyn Storage>, Option<LazyGc>) = match self.storage {
Some(custom) => (custom, None),
None => {
let memory = Arc::new(MemoryStorage::new());
let gc = self.gc_enabled.then(|| LazyGc {
storage: memory.clone(),
interval: self.gc_interval,
handle: OnceLock::new(),
});
(memory, gc)
}
};
let rate_tier = RateTier {
tiers: self.tiers,
default_tier: self.default_tier,
on_missing: self.on_missing,
on_unknown_tier: self.on_unknown_tier,
storage,
clock,
gc,
};
if let Some(gc) = &rate_tier.gc {
gc.ensure_started(&rate_tier.clock);
}
rate_tier
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::clock::FakeClock;
#[tokio::test]
async fn gc_starts_at_build_inside_a_runtime() {
let limiter = RateTier::builder()
.tier("free", Quota::per_second(1))
.build();
let gc = limiter.gc.as_ref().expect("built-in storage enables GC");
assert!(gc.handle.get().is_some());
}
#[test]
fn gc_built_outside_a_runtime_starts_on_first_use_and_repeats() {
let clock = FakeClock::new();
let limiter = RateTier::builder()
.tier("free", Quota::per_second(1))
.clock(clock.clone())
.gc_interval(Duration::from_secs(1))
.build();
let gc = limiter.gc.as_ref().expect("built-in storage enables GC");
assert!(gc.handle.get().is_none(), "no runtime, so no GC yet");
let rt = tokio::runtime::Builder::new_current_thread()
.enable_time()
.start_paused(true)
.build()
.unwrap();
rt.block_on(async {
limiter.check("u1", "free", 1).await.unwrap().unwrap();
assert!(gc.handle.get().is_some(), "first use starts the GC");
tokio::task::yield_now().await;
assert_eq!(gc.storage.len(), 1);
clock.advance(Duration::from_secs(10));
tokio::time::advance(Duration::from_secs(1)).await;
tokio::task::yield_now().await;
assert_eq!(gc.storage.len(), 0, "a later tick must collect it");
});
}
}