use std::sync::atomic::{AtomicU64, Ordering}; use std::sync::{Arc, Mutex}; use std::time::Duration; /// 令牌桶限速器(字节/秒,0=不限) /// /// 滑动窗口算法:1 秒窗口内累计消耗字节数,超出限制时 sleep 到下一窗口。 pub struct RateLimiter { /// 限速 bytes/s(0=不限),用 Arc 以便运行时动态调整 limit: Arc, /// 当前窗口已消耗字节 consumed: AtomicU64, /// 窗口起始时间 window_start: Mutex, } impl RateLimiter { pub fn new(limit: u64) -> Self { Self { limit: Arc::new(AtomicU64::new(limit)), consumed: AtomicU64::new(0), window_start: Mutex::new(std::time::Instant::now()), } } /// 动态设置限速(bytes/s,0=不限) pub fn set_limit(&self, limit: u64) { self.limit.store(limit, Ordering::Relaxed); } /// 消费指定字节数,若超出限速则 sleep 等待 pub async fn consume(&self, bytes: u64) { let limit = self.limit.load(Ordering::Relaxed); if limit == 0 || bytes == 0 { return; } // 尝试在当前窗口消费(用作用域确保 MutexGuard 在 await 前释放) let over_limit = { let now = std::time::Instant::now(); let mut start = self.window_start.lock().unwrap_or_else(|e| e.into_inner()); let elapsed = now.duration_since(*start); // 窗口过期,重置 if elapsed >= Duration::from_secs(1) { self.consumed.store(0, Ordering::Relaxed); *start = now; } let current = self.consumed.fetch_add(bytes, Ordering::Relaxed) + bytes; if current > limit { Some(current - limit) } else { None } }; // MutexGuard 在此释放 // 超出限速,计算需要等待的时间 if let Some(over) = over_limit { let wait_ms = (over as f64 / limit as f64 * 1000.0).ceil() as u64; let wait = Duration::from_millis(wait_ms.min(1000)); tokio::time::sleep(wait).await; } } } /// 克隆限速器的共享句柄(共享同一限速状态) impl Clone for RateLimiter { fn clone(&self) -> Self { // 注意:clone 的是限速配置,不是状态。新实例独立计数。 // 全局限速器应通过引用共享,而非 clone。 Self::new(self.limit.load(Ordering::Relaxed)) } } #[cfg(test)] mod rate_limiter_tests { use super::*; use std::time::Instant; #[tokio::test] async fn zero_limit_never_blocks() { let limiter = RateLimiter::new(0); let start = Instant::now(); limiter.consume(1024 * 1024).await; limiter.consume(u64::MAX).await; assert!(start.elapsed() < Duration::from_millis(50)); } #[tokio::test] async fn under_limit_returns_immediately() { let limiter = RateLimiter::new(100_000); // 100KB/s let start = Instant::now(); limiter.consume(1024).await; assert!(start.elapsed() < Duration::from_millis(50)); } #[tokio::test] async fn exceeding_limit_waits_proportionally() { // 限速 200 B/s:先消耗 100 未超限,再消耗 150 → 超限 50 → 等待约 250ms let limiter = RateLimiter::new(200); limiter.consume(100).await; let start = Instant::now(); limiter.consume(150).await; let elapsed = start.elapsed(); assert!( elapsed >= Duration::from_millis(200), "等待时间不足: {:?}", elapsed ); assert!(elapsed < Duration::from_millis(1100)); } #[tokio::test] async fn window_resets_after_one_second() { // 限速 100 B/s:第一窗口耗尽后,1.1s 窗口重置,再消耗 100 不应阻塞 let limiter = RateLimiter::new(100); limiter.consume(100).await; tokio::time::sleep(Duration::from_millis(1100)).await; let start = Instant::now(); limiter.consume(100).await; assert!(start.elapsed() < Duration::from_millis(100)); } #[tokio::test] async fn set_limit_takes_effect_dynamically() { let limiter = RateLimiter::new(0); limiter.consume(1024).await; // 不限速 limiter.set_limit(100); limiter.consume(100).await; let start = Instant::now(); limiter.consume(100).await; // 累计 200 > 100 → 等待 1000ms assert!(start.elapsed() >= Duration::from_millis(900)); } }