use std::collections::HashMap;
use std::hash::Hash;
use std::time::{Duration, Instant};
#[derive(Debug)]
pub struct KeyedRateLimiter<K> {
limiters: HashMap<K, SlidingWindow>,
window: Duration,
max_requests: u32,
}
impl<K> KeyedRateLimiter<K>
where
K: Hash + Eq + Clone,
{
pub fn new(window: Duration, max_requests: u32) -> Self {
Self {
limiters: HashMap::new(),
window,
max_requests,
}
}
pub fn try_acquire(&mut self, key: K) -> Result<(), Duration> {
let limiter = self
.limiters
.entry(key)
.or_insert_with(|| SlidingWindow::new(self.window, self.max_requests));
limiter.try_acquire()
}
pub fn would_allow(&self, key: &K) -> bool {
self.limiters
.get(key)
.is_none_or(|limiter| limiter.would_allow())
}
pub fn remaining(&self, key: &K) -> u32 {
self.limiters
.get(key)
.map_or(self.max_requests, |limiter| limiter.remaining())
}
pub fn time_until_available(&self, key: &K) -> Option<Duration> {
self.limiters
.get(key)
.and_then(|limiter| limiter.time_until_available())
}
pub fn remove(&mut self, key: &K) {
self.limiters.remove(key);
}
pub fn cleanup(&mut self) {
self.limiters.retain(|_, limiter| !limiter.is_empty());
}
pub fn tracked_keys(&self) -> usize {
self.limiters.len()
}
pub fn clear(&mut self) {
self.limiters.clear();
}
}
impl<K> Default for KeyedRateLimiter<K>
where
K: Hash + Eq + Clone,
{
fn default() -> Self {
Self::new(Duration::from_secs(1), 1)
}
}
#[derive(Debug)]
pub struct SlidingWindow {
requests: Vec<Instant>,
window: Duration,
max_requests: u32,
}
impl SlidingWindow {
pub fn new(window: Duration, max_requests: u32) -> Self {
Self {
requests: Vec::with_capacity(max_requests as usize),
window,
max_requests,
}
}
pub fn try_acquire(&mut self) -> Result<(), Duration> {
self.cleanup_old();
if (self.requests.len() as u32) < self.max_requests {
self.requests.push(Instant::now());
Ok(())
} else {
let wait_time = self
.requests
.first()
.map(|oldest| self.window.saturating_sub(oldest.elapsed()))
.unwrap_or_default();
Err(wait_time)
}
}
pub fn would_allow(&self) -> bool {
let count = self
.requests
.iter()
.filter(|ts| ts.elapsed() < self.window)
.count();
(count as u32) < self.max_requests
}
pub fn remaining(&self) -> u32 {
let count = self
.requests
.iter()
.filter(|ts| ts.elapsed() < self.window)
.count() as u32;
self.max_requests.saturating_sub(count)
}
pub fn time_until_available(&self) -> Option<Duration> {
self.cleanup_check();
let count = self
.requests
.iter()
.filter(|ts| ts.elapsed() < self.window)
.count();
if (count as u32) < self.max_requests {
None
} else {
self.requests
.iter()
.find(|ts| ts.elapsed() < self.window)
.map(|oldest| self.window.saturating_sub(oldest.elapsed()))
}
}
pub fn is_empty(&self) -> bool {
self.requests.iter().all(|ts| ts.elapsed() >= self.window)
}
fn cleanup_old(&mut self) {
let window = self.window;
self.requests.retain(|ts| ts.elapsed() < window);
}
fn cleanup_check(&self) -> usize {
self.requests
.iter()
.filter(|ts| ts.elapsed() < self.window)
.count()
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::thread;
#[test]
fn test_sliding_window_allows_within_limit() {
let mut limiter = SlidingWindow::new(Duration::from_secs(1), 3);
assert!(limiter.try_acquire().is_ok());
assert!(limiter.try_acquire().is_ok());
assert!(limiter.try_acquire().is_ok());
assert!(limiter.try_acquire().is_err());
}
#[test]
fn test_sliding_window_resets_after_window() {
let mut limiter = SlidingWindow::new(Duration::from_millis(50), 2);
assert!(limiter.try_acquire().is_ok());
assert!(limiter.try_acquire().is_ok());
assert!(limiter.try_acquire().is_err());
thread::sleep(Duration::from_millis(60));
assert!(limiter.try_acquire().is_ok());
}
#[test]
fn test_remaining() {
let mut limiter = SlidingWindow::new(Duration::from_secs(1), 3);
assert_eq!(limiter.remaining(), 3);
limiter.try_acquire().ok();
assert_eq!(limiter.remaining(), 2);
limiter.try_acquire().ok();
assert_eq!(limiter.remaining(), 1);
}
#[test]
fn test_keyed_limiter() {
let mut limiter: KeyedRateLimiter<String> =
KeyedRateLimiter::new(Duration::from_secs(1), 2);
assert!(limiter.try_acquire("BTC/USD".to_string()).is_ok());
assert!(limiter.try_acquire("BTC/USD".to_string()).is_ok());
assert!(limiter.try_acquire("BTC/USD".to_string()).is_err());
assert!(limiter.try_acquire("ETH/USD".to_string()).is_ok());
assert!(limiter.try_acquire("ETH/USD".to_string()).is_ok());
assert!(limiter.try_acquire("ETH/USD".to_string()).is_err());
}
#[test]
fn test_keyed_limiter_cleanup() {
let mut limiter: KeyedRateLimiter<String> =
KeyedRateLimiter::new(Duration::from_millis(50), 1);
limiter.try_acquire("key1".to_string()).ok();
limiter.try_acquire("key2".to_string()).ok();
assert_eq!(limiter.tracked_keys(), 2);
thread::sleep(Duration::from_millis(60));
limiter.cleanup();
assert_eq!(limiter.tracked_keys(), 0);
}
}