1use std::sync::atomic::{AtomicBool, AtomicI64, AtomicU32, AtomicU64, Ordering};
14use std::sync::{Arc, RwLock};
15use std::time::Instant;
16
17use tokio::sync::Semaphore;
18use tokio::time::{self, Duration};
19
20use crate::spec::RateSpec;
21
22const MAX_ACTIVE_POOL: u32 = 1_000_000_000;
24
25const REFILL_INTERVAL_MS: u64 = 10;
27
28struct SharedState {
30 active_pool: Semaphore,
31 waiting_pool: AtomicI64,
32 running: AtomicBool,
33 last_refill_nanos: AtomicU64,
34 start_time: Instant,
35 blocks: AtomicU64,
36 ticks_per_op: AtomicU32,
40 refill_cfg: RwLock<RefillCfg>,
44 spec: RwLock<RateSpec>,
47}
48
49#[derive(Clone, Copy)]
50struct RefillCfg {
51 max_active: u32,
52 max_over_active: u32,
53 burst_pool_size: u32,
54 unit: crate::spec::TimeUnit,
55}
56
57impl RefillCfg {
58 fn from_spec(spec: &RateSpec) -> Self {
59 let max_over_active = (MAX_ACTIVE_POOL as f64 * spec.burst_ratio) as u32;
60 let burst_pool_size = max_over_active.saturating_sub(MAX_ACTIVE_POOL);
61 Self {
62 max_active: MAX_ACTIVE_POOL,
63 max_over_active,
64 burst_pool_size,
65 unit: spec.unit,
66 }
67 }
68}
69
70pub struct RateLimiter {
77 state: Arc<SharedState>,
78 refill_handle: Option<tokio::task::JoinHandle<()>>,
79}
80
81impl RateLimiter {
82 pub fn start(spec: RateSpec) -> Self {
84 let ticks_per_op = spec.ticks_per_op();
85 let refill_cfg = RefillCfg::from_spec(&spec);
86
87 let state = Arc::new(SharedState {
88 active_pool: Semaphore::new(ticks_per_op as usize), waiting_pool: AtomicI64::new(0),
90 running: AtomicBool::new(true),
91 last_refill_nanos: AtomicU64::new(0),
92 start_time: Instant::now(),
93 blocks: AtomicU64::new(0),
94 ticks_per_op: AtomicU32::new(ticks_per_op),
95 refill_cfg: RwLock::new(refill_cfg),
96 spec: RwLock::new(spec),
97 });
98
99 state.last_refill_nanos.store(
101 state.start_time.elapsed().as_nanos() as u64,
102 Ordering::Relaxed,
103 );
104
105 let refill_state = state.clone();
106 let refill_handle = tokio::spawn(async move {
107 refill_loop(refill_state).await;
108 });
109
110 Self {
111 state,
112 refill_handle: Some(refill_handle),
113 }
114 }
115
116 pub async fn acquire(&self) -> i64 {
123 self.state.blocks.fetch_add(1, Ordering::Relaxed);
124 let permits = self.state.ticks_per_op.load(Ordering::Relaxed);
125
126 let permit = self
131 .state
132 .active_pool
133 .acquire_many(permits)
134 .await
135 .expect("semaphore closed unexpectedly");
136 permit.forget();
137
138 self.state.waiting_pool.load(Ordering::Relaxed)
139 }
140
141 pub fn wait_time_nanos(&self) -> u64 {
143 let ticks = self.state.waiting_pool.load(Ordering::Relaxed);
144 if ticks <= 0 {
145 return 0;
146 }
147 let unit = self
148 .state
149 .refill_cfg
150 .read()
151 .unwrap_or_else(|e| e.into_inner())
152 .unit;
153 unit.ticks_to_nanos(ticks as u32)
154 }
155
156 pub fn total_blocks(&self) -> u64 {
158 self.state.blocks.load(Ordering::Relaxed)
159 }
160
161 pub fn rate(&self) -> f64 {
163 self.state
164 .spec
165 .read()
166 .unwrap_or_else(|e| e.into_inner())
167 .ops_per_sec
168 }
169
170 pub fn spec(&self) -> RateSpec {
172 self.state
173 .spec
174 .read()
175 .unwrap_or_else(|e| e.into_inner())
176 .clone()
177 }
178
179 pub fn reconfigure(&self, spec: RateSpec) -> Result<(), String> {
196 if spec.ops_per_sec <= 0.0 {
197 return Err(format!("rate must be > 0, got {}", spec.ops_per_sec,));
198 }
199 if spec.burst_ratio < 1.0 {
200 return Err(format!(
201 "burst_ratio must be >= 1.0, got {}",
202 spec.burst_ratio,
203 ));
204 }
205 let new_ticks_per_op = spec.ticks_per_op();
206 let new_cfg = RefillCfg::from_spec(&spec);
207
208 *self
215 .state
216 .refill_cfg
217 .write()
218 .unwrap_or_else(|e| e.into_inner()) = new_cfg;
219 self.state
220 .ticks_per_op
221 .store(new_ticks_per_op, Ordering::Relaxed);
222 *self.state.spec.write().unwrap_or_else(|e| e.into_inner()) = spec;
223 Ok(())
224 }
225
226 pub async fn stop(mut self) {
228 self.state.running.store(false, Ordering::Relaxed);
229 if let Some(handle) = self.refill_handle.take() {
230 let _ = handle.await;
231 }
232 }
233}
234
235impl Drop for RateLimiter {
236 fn drop(&mut self) {
237 self.state.running.store(false, Ordering::Relaxed);
238 }
239}
240
241async fn refill_loop(state: Arc<SharedState>) {
243 let mut interval = time::interval(Duration::from_millis(REFILL_INTERVAL_MS));
244 interval.set_missed_tick_behavior(time::MissedTickBehavior::Skip);
245
246 while state.running.load(Ordering::Relaxed) {
247 interval.tick().await;
248 refill(&state);
249 }
250}
251
252fn refill(state: &SharedState) {
254 let now_nanos = state.start_time.elapsed().as_nanos() as u64;
255 let last = state.last_refill_nanos.swap(now_nanos, Ordering::Relaxed);
256 let elapsed_nanos = now_nanos.saturating_sub(last);
257
258 if elapsed_nanos == 0 {
259 return;
260 }
261
262 let cfg = *state.refill_cfg.read().unwrap_or_else(|e| e.into_inner());
266
267 let new_ticks = cfg.unit.nanos_to_ticks(elapsed_nanos);
269 if new_ticks == 0 {
270 return;
271 }
272
273 let available = state.active_pool.available_permits() as u32;
275
276 let room = cfg.max_active.saturating_sub(available);
278 let to_active = new_ticks.min(room);
279 if to_active > 0 {
280 state.active_pool.add_permits(to_active as usize);
281 }
282
283 let overflow = new_ticks.saturating_sub(to_active);
285 if overflow > 0 {
286 state
287 .waiting_pool
288 .fetch_add(overflow as i64, Ordering::Relaxed);
289 }
290
291 let available_after = state.active_pool.available_permits() as u32;
293 let refill_factor = (new_ticks as f64 / cfg.max_active as f64).min(1.0);
294 let burst_allowed = (refill_factor * cfg.burst_pool_size as f64) as u32;
295 let burst_room = cfg.max_over_active.saturating_sub(available_after);
296 let burst_from_waiting = burst_allowed
297 .min(burst_room)
298 .min(state.waiting_pool.load(Ordering::Relaxed).max(0) as u32);
299
300 if burst_from_waiting > 0 {
301 state
302 .waiting_pool
303 .fetch_sub(burst_from_waiting as i64, Ordering::Relaxed);
304 state.active_pool.add_permits(burst_from_waiting as usize);
305 }
306}
307
308#[cfg(test)]
309mod tests {
310 use super::*;
311 use std::time::Instant;
312
313 #[tokio::test]
314 async fn limiter_starts_and_stops() {
315 let spec = RateSpec::new(1000.0);
316 let limiter = RateLimiter::start(spec);
317 assert_eq!(limiter.rate(), 1000.0);
318 limiter.stop().await;
319 }
320
321 #[tokio::test]
322 async fn limiter_acquire_returns() {
323 let spec = RateSpec::new(10000.0);
324 let limiter = RateLimiter::start(spec);
325
326 let _backlog = limiter.acquire().await;
328 assert_eq!(limiter.total_blocks(), 1);
329
330 limiter.stop().await;
331 }
332
333 #[tokio::test]
334 async fn limiter_rate_limits() {
335 let spec = RateSpec::new(100.0); let limiter = RateLimiter::start(spec);
337
338 tokio::time::sleep(Duration::from_millis(50)).await;
340
341 let start = Instant::now();
342 let mut count = 0u64;
343
344 for _ in 0..20 {
346 limiter.acquire().await;
347 count += 1;
348 }
349
350 let elapsed = start.elapsed();
351 limiter.stop().await;
352
353 assert_eq!(count, 20);
354 assert!(
357 elapsed.as_millis() >= 50,
358 "expected rate limiting, took {}ms for 20 ops at 100/s",
359 elapsed.as_millis()
360 );
361 }
362
363 #[tokio::test]
364 async fn limiter_high_rate_is_fast() {
365 let spec = RateSpec::new(1_000_000.0); let limiter = RateLimiter::start(spec);
367
368 tokio::time::sleep(Duration::from_millis(50)).await;
369
370 let start = Instant::now();
371 for _ in 0..100 {
372 limiter.acquire().await;
373 }
374 let elapsed = start.elapsed();
375 limiter.stop().await;
376
377 assert!(
379 elapsed.as_millis() < 500,
380 "high rate should be fast, took {}ms",
381 elapsed.as_millis()
382 );
383 }
384
385 #[tokio::test]
386 async fn limiter_wait_time_grows_under_load() {
387 let spec = RateSpec::new(100.0); let limiter = RateLimiter::start(spec);
389
390 tokio::time::sleep(Duration::from_millis(50)).await;
391
392 for _ in 0..50 {
394 limiter.acquire().await;
395 }
396
397 let _wt = limiter.wait_time_nanos();
399 assert!(limiter.total_blocks() == 50);
402
403 limiter.stop().await;
404 }
405
406 #[tokio::test]
407 async fn limiter_spec_parsing() {
408 let spec = RateSpec::parse("500, 1.5, start").unwrap();
409 let limiter = RateLimiter::start(spec);
410 assert_eq!(limiter.rate(), 500.0);
411 limiter.stop().await;
412 }
413
414 #[tokio::test]
417 async fn reconfigure_updates_rate_and_spec() {
418 let limiter = RateLimiter::start(RateSpec::new(1_000.0));
419 assert_eq!(limiter.rate(), 1_000.0);
420
421 limiter.reconfigure(RateSpec::new(5_000.0)).unwrap();
422 assert_eq!(limiter.rate(), 5_000.0);
423 assert_eq!(limiter.spec().ops_per_sec, 5_000.0);
424
425 limiter.stop().await;
426 }
427
428 #[tokio::test]
429 async fn reconfigure_rejects_nonpositive_rate() {
430 let limiter = RateLimiter::start(RateSpec::new(1_000.0));
431 let bad = RateSpec {
432 ops_per_sec: 0.0,
433 ..RateSpec::new(1_000.0)
434 };
435 assert!(limiter.reconfigure(bad).is_err());
436 assert_eq!(limiter.rate(), 1_000.0);
437 limiter.stop().await;
438 }
439
440 #[tokio::test]
441 async fn reconfigure_rejects_subunit_burst_ratio() {
442 let limiter = RateLimiter::start(RateSpec::new(1_000.0));
443 let bad = RateSpec {
444 burst_ratio: 0.5,
445 ..RateSpec::new(1_000.0)
446 };
447 assert!(limiter.reconfigure(bad).is_err());
448 limiter.stop().await;
449 }
450
451 #[tokio::test]
452 async fn reconfigure_preserves_total_blocks_and_keeps_running() {
453 let limiter = RateLimiter::start(RateSpec::new(10_000.0));
457 tokio::time::sleep(Duration::from_millis(30)).await;
458
459 for _ in 0..5 {
460 limiter.acquire().await;
461 }
462 let blocks_before = limiter.total_blocks();
463 assert!(blocks_before >= 5);
464
465 limiter.reconfigure(RateSpec::new(50_000.0)).unwrap();
466 tokio::time::sleep(Duration::from_millis(30)).await;
467
468 for _ in 0..5 {
469 limiter.acquire().await;
470 }
471 let blocks_after = limiter.total_blocks();
472 assert!(
473 blocks_after >= blocks_before + 5,
474 "acquire counter should continue across reconfigure"
475 );
476
477 limiter.stop().await;
478 }
479
480 #[tokio::test]
481 async fn reconfigure_changes_observed_throughput() {
482 let limiter = RateLimiter::start(RateSpec::new(200.0));
485 tokio::time::sleep(Duration::from_millis(30)).await;
486
487 let t0 = Instant::now();
488 for _ in 0..8 {
489 limiter.acquire().await;
490 }
491 let slow = t0.elapsed();
492
493 limiter.reconfigure(RateSpec::new(100_000.0)).unwrap();
494 tokio::time::sleep(Duration::from_millis(30)).await;
495
496 let t1 = Instant::now();
497 for _ in 0..8 {
498 limiter.acquire().await;
499 }
500 let fast = t1.elapsed();
501
502 assert!(
503 fast < slow,
504 "expected faster after reconfigure: slow={slow:?} fast={fast:?}",
505 );
506 limiter.stop().await;
507 }
508}