Skip to main content

postrust_proxy/ratelimit/
token_bucket.rs

1//! Token bucket algorithm for rate limiting.
2
3use std::sync::atomic::{AtomicU64, Ordering};
4
5/// Token bucket rate limiter.
6///
7/// Implements the token bucket algorithm where tokens are added at a fixed rate
8/// up to a maximum capacity. Each request consumes one token.
9pub struct TokenBucket {
10    /// Maximum number of tokens (burst capacity)
11    capacity: u64,
12    /// Tokens added per second
13    refill_rate: f64,
14    /// Current token count (scaled by 1000 for precision)
15    tokens: AtomicU64,
16    /// Last refill time in milliseconds since epoch
17    last_refill: AtomicU64,
18}
19
20impl TokenBucket {
21    /// Create a new token bucket.
22    ///
23    /// # Arguments
24    /// * `capacity` - Maximum tokens (burst size)
25    /// * `refill_rate` - Tokens added per second
26    pub fn new(capacity: u64, refill_rate: f64) -> Self {
27        let now_ms = Self::now_ms();
28        Self {
29            capacity,
30            refill_rate,
31            tokens: AtomicU64::new(capacity * 1000), // Scale for precision
32            last_refill: AtomicU64::new(now_ms),
33        }
34    }
35
36    /// Try to consume a token.
37    ///
38    /// Returns `true` if a token was available and consumed, `false` otherwise.
39    pub fn try_acquire(&self) -> bool {
40        self.try_acquire_n(1)
41    }
42
43    /// Try to consume N tokens.
44    ///
45    /// Returns `true` if N tokens were available and consumed, `false` otherwise.
46    pub fn try_acquire_n(&self, n: u64) -> bool {
47        let now_ms = Self::now_ms();
48        let cost = n * 1000; // Scale for precision
49
50        loop {
51            let last = self.last_refill.load(Ordering::Relaxed);
52            let current_tokens = self.tokens.load(Ordering::Relaxed);
53
54            // Calculate tokens to add based on elapsed time
55            let elapsed_ms = now_ms.saturating_sub(last);
56            let tokens_to_add = (elapsed_ms as f64 * self.refill_rate).round() as u64;
57
58            // Calculate new token count (capped at capacity)
59            let new_tokens = (current_tokens + tokens_to_add).min(self.capacity * 1000);
60
61            // Check if we have enough tokens
62            if new_tokens < cost {
63                return false;
64            }
65
66            // Try to atomically update
67            let final_tokens = new_tokens - cost;
68            if self
69                .tokens
70                .compare_exchange_weak(
71                    current_tokens,
72                    final_tokens,
73                    Ordering::SeqCst,
74                    Ordering::Relaxed,
75                )
76                .is_ok()
77            {
78                // Update last refill time
79                let _ = self.last_refill.compare_exchange(
80                    last,
81                    now_ms,
82                    Ordering::Relaxed,
83                    Ordering::Relaxed,
84                );
85                return true;
86            }
87            // CAS failed, retry
88        }
89    }
90
91    /// Get current available tokens (approximate).
92    pub fn available(&self) -> u64 {
93        let now_ms = Self::now_ms();
94        let last = self.last_refill.load(Ordering::Relaxed);
95        let current_tokens = self.tokens.load(Ordering::Relaxed);
96
97        let elapsed_ms = now_ms.saturating_sub(last);
98        let tokens_to_add = (elapsed_ms as f64 * self.refill_rate).round() as u64;
99
100        (current_tokens + tokens_to_add).min(self.capacity * 1000) / 1000
101    }
102
103    fn now_ms() -> u64 {
104        use std::time::{SystemTime, UNIX_EPOCH};
105        SystemTime::now()
106            .duration_since(UNIX_EPOCH)
107            .unwrap_or_default()
108            .as_millis() as u64
109    }
110}
111
112#[cfg(test)]
113mod tests {
114    use super::*;
115
116    #[test]
117    fn test_token_bucket_basic() {
118        let bucket = TokenBucket::new(10, 1.0); // 10 capacity, 1 token/sec
119
120        // Should be able to acquire up to capacity
121        for _ in 0..10 {
122            assert!(bucket.try_acquire());
123        }
124
125        // Should fail after exhausting tokens
126        assert!(!bucket.try_acquire());
127    }
128
129    #[test]
130    fn test_token_bucket_burst() {
131        let bucket = TokenBucket::new(5, 10.0); // 5 burst, 10 tokens/sec
132
133        // Burst of 5 should succeed
134        assert!(bucket.try_acquire_n(5));
135
136        // Next request should fail
137        assert!(!bucket.try_acquire());
138    }
139}