postrust_proxy/ratelimit/
token_bucket.rs1use std::sync::atomic::{AtomicU64, Ordering};
4
5pub struct TokenBucket {
10 capacity: u64,
12 refill_rate: f64,
14 tokens: AtomicU64,
16 last_refill: AtomicU64,
18}
19
20impl TokenBucket {
21 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), last_refill: AtomicU64::new(now_ms),
33 }
34 }
35
36 pub fn try_acquire(&self) -> bool {
40 self.try_acquire_n(1)
41 }
42
43 pub fn try_acquire_n(&self, n: u64) -> bool {
47 let now_ms = Self::now_ms();
48 let cost = n * 1000; loop {
51 let last = self.last_refill.load(Ordering::Relaxed);
52 let current_tokens = self.tokens.load(Ordering::Relaxed);
53
54 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 let new_tokens = (current_tokens + tokens_to_add).min(self.capacity * 1000);
60
61 if new_tokens < cost {
63 return false;
64 }
65
66 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 let _ = self.last_refill.compare_exchange(
80 last,
81 now_ms,
82 Ordering::Relaxed,
83 Ordering::Relaxed,
84 );
85 return true;
86 }
87 }
89 }
90
91 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); for _ in 0..10 {
122 assert!(bucket.try_acquire());
123 }
124
125 assert!(!bucket.try_acquire());
127 }
128
129 #[test]
130 fn test_token_bucket_burst() {
131 let bucket = TokenBucket::new(5, 10.0); assert!(bucket.try_acquire_n(5));
135
136 assert!(!bucket.try_acquire());
138 }
139}