use std::sync::Mutex;
use std::time::{Duration, Instant};
pub struct RateLimiter {
bytes_per_sec: f64,
capacity: f64,
state: Mutex<Bucket>,
}
struct Bucket {
tokens: f64,
last: Instant,
}
const MIN_CAPACITY: f64 = 1.0 * 1024.0 * 1024.0;
impl RateLimiter {
pub fn new(bytes_per_sec: u64) -> Self {
let rate = bytes_per_sec as f64;
let capacity = (rate / 4.0).max(MIN_CAPACITY).min(rate.max(1.0));
Self {
bytes_per_sec: rate,
capacity,
state: Mutex::new(Bucket {
tokens: capacity,
last: Instant::now(),
}),
}
}
pub async fn acquire(&self, n: u64) {
let mut remaining = n as f64;
while remaining > 0.0 {
let wait = {
let mut b = self.state.lock().expect("rate limiter poisoned");
let now = Instant::now();
let elapsed = now.duration_since(b.last).as_secs_f64();
b.last = now;
b.tokens = (b.tokens + elapsed * self.bytes_per_sec).min(self.capacity);
if b.tokens >= remaining {
b.tokens -= remaining;
remaining = 0.0;
Duration::ZERO
} else {
remaining -= b.tokens.max(0.0);
b.tokens = 0.0;
Duration::from_secs_f64((remaining / self.bytes_per_sec).min(1.0))
}
};
if wait > Duration::ZERO {
tokio::time::sleep(wait).await;
}
}
}
}
pub fn parse_rate(input: &str) -> Result<u64, String> {
let s = input.trim().trim_end_matches("/s").trim_end_matches("/S");
let s = s.trim();
let split = s
.find(|c: char| !c.is_ascii_digit() && c != '.' && c != ',')
.unwrap_or(s.len());
let (num, unit) = s.split_at(split);
let num: f64 = num
.replace(',', "")
.parse()
.map_err(|_| format!("invalid rate `{input}`"))?;
if num <= 0.0 {
return Err(format!("rate must be positive, got `{input}`"));
}
let mult: f64 = match unit.trim().to_ascii_lowercase().as_str() {
"" | "b" => 1.0,
"k" | "kb" | "kib" => 1024.0,
"m" | "mb" | "mib" => 1024.0 * 1024.0,
"g" | "gb" | "gib" => 1024.0 * 1024.0 * 1024.0,
other => return Err(format!("unknown rate unit `{other}` in `{input}`")),
};
Ok((num * mult) as u64)
}
pub fn parse_duration(input: &str) -> Result<Duration, String> {
let s = input.trim();
let split = s
.find(|c: char| !c.is_ascii_digit() && c != '.')
.unwrap_or(s.len());
let (num, unit) = s.split_at(split);
let num: f64 = num
.parse()
.map_err(|_| format!("invalid duration `{input}`"))?;
let secs = match unit.trim().to_ascii_lowercase().as_str() {
"ms" => num / 1000.0,
"" | "s" | "sec" | "secs" => num,
"m" | "min" | "mins" => num * 60.0,
"h" | "hr" | "hrs" => num * 3600.0,
other => return Err(format!("unknown duration unit `{other}` in `{input}`")),
};
if secs <= 0.0 {
return Err(format!("duration must be positive, got `{input}`"));
}
Ok(Duration::from_secs_f64(secs))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parses_rates() {
assert_eq!(parse_rate("1024"), Ok(1024));
assert_eq!(parse_rate("1k"), Ok(1024));
assert_eq!(parse_rate("20MiB/s"), Ok(20 * 1024 * 1024));
assert_eq!(
parse_rate("20 MB / s".replace(' ', "").as_str()),
Ok(20 * 1024 * 1024)
);
assert_eq!(parse_rate("1.5m"), Ok(1_572_864));
assert_eq!(parse_rate("2G"), Ok(2 * 1024 * 1024 * 1024));
assert!(parse_rate("0").is_err());
assert!(parse_rate("-5m").is_err());
assert!(parse_rate("fast").is_err());
assert!(parse_rate("10furlongs").is_err());
}
#[test]
fn parses_durations() {
assert_eq!(parse_duration("30"), Ok(Duration::from_secs(30)));
assert_eq!(parse_duration("30s"), Ok(Duration::from_secs(30)));
assert_eq!(parse_duration("500ms"), Ok(Duration::from_millis(500)));
assert_eq!(parse_duration("2m"), Ok(Duration::from_secs(120)));
assert_eq!(parse_duration("1h"), Ok(Duration::from_secs(3600)));
assert!(parse_duration("0").is_err());
assert!(parse_duration("soon").is_err());
}
#[tokio::test(start_paused = true)]
async fn limits_throughput_globally() {
let rate = 8 * 1024 * 1024;
let limiter = RateLimiter::new(rate);
let start = tokio::time::Instant::now();
limiter.acquire(rate / 4).await;
assert_eq!(start.elapsed(), Duration::ZERO);
limiter.acquire(rate).await;
limiter.acquire(rate).await;
assert!(
start.elapsed() >= Duration::from_millis(1900),
"elapsed {:?}",
start.elapsed()
);
}
#[test]
fn burst_allowance_is_a_small_fraction_of_the_rate() {
let fast = RateLimiter::new(40 * 1024 * 1024);
assert_eq!(fast.capacity, 10.0 * 1024.0 * 1024.0);
let slow = RateLimiter::new(64 * 1024);
assert_eq!(slow.capacity, 64.0 * 1024.0);
let mid = RateLimiter::new(2 * 1024 * 1024);
assert_eq!(mid.capacity, MIN_CAPACITY);
}
}