136 lines
4.5 KiB
Rust
136 lines
4.5 KiB
Rust
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<AtomicU64> 以便运行时动态调整
|
||
limit: Arc<AtomicU64>,
|
||
/// 当前窗口已消耗字节
|
||
consumed: AtomicU64,
|
||
/// 窗口起始时间
|
||
window_start: Mutex<std::time::Instant>,
|
||
}
|
||
|
||
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));
|
||
}
|
||
}
|