use std::collections::HashMap;
use std::fmt;
use std::sync::Arc;
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::quota::Quota;
use crate::storage::memory::MemoryStorage;
use crate::storage::{Storage, StorageError};
#[derive(Debug)]
pub enum CheckError {
UnknownTier(String),
Storage(StorageError),
}
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),
}
}
}
impl std::error::Error for CheckError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
CheckError::UnknownTier(_) => 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,
storage: Arc<dyn Storage>,
clock: Arc<dyn Clock>,
_gc: Option<GcHandle>,
}
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 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 {
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_at: 0,
}));
}
let now = self.clock.now();
let storage_key = format!("{}:{}", user_id, tier_name);
Ok(self
.storage
.check_and_update(&storage_key, quota, cost, now)
.await?)
}
}
pub struct RateTierBuilder {
tiers: HashMap<String, Quota>,
default_tier: Option<String>,
on_missing: OnMissing,
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(),
clock: None,
storage: None,
gc_interval: Duration::from_secs(60),
gc_enabled: true,
}
}
}
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 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 {
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<GcHandle>) = match self.storage {
Some(custom) => (custom, None),
None => {
let memory = Arc::new(MemoryStorage::new());
let gc = if self.gc_enabled {
Some(GcHandle::spawn(
memory.clone(),
clock.clone(),
self.gc_interval,
))
} else {
None
};
(memory, gc)
}
};
RateTier {
tiers: self.tiers,
default_tier: self.default_tier,
on_missing: self.on_missing,
storage,
clock,
_gc: gc,
}
}
}