use super::{Acquire, Rate, TokenBucket};
use parking_lot::Mutex;
use std::{
future::Future,
pin::Pin,
sync::Arc,
task::{Context, Poll},
time::Duration,
};
use tokio::sync::{Notify, futures::OwnedNotified};
use tokio::time::Instant;
#[derive(Debug, Clone)]
pub struct RateLimiter {
inner: Arc<Inner>,
}
#[derive(Debug)]
#[must_use = "futures do nothing unless awaited or polled"]
pub struct RefundWait(Pin<Box<OwnedNotified>>);
impl RefundWait {
fn new(notify: Arc<Notify>) -> Self {
let mut notified = Box::pin(notify.notified_owned());
notified.as_mut().enable();
Self(notified)
}
}
impl Future for RefundWait {
type Output = ();
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
self.0.as_mut().poll(cx)
}
}
#[derive(Debug)]
struct Inner {
bucket: Mutex<TokenBucket>,
epoch: Instant,
rate: Rate,
burst: u64,
refunds: Arc<Notify>,
}
impl RateLimiter {
#[must_use]
pub fn new(rate: Rate, burst: u64) -> Self {
Self::from_bucket(TokenBucket::new(rate, burst))
}
#[must_use]
pub fn from_rate(rate: Rate) -> Self {
Self::from_bucket(TokenBucket::from_rate(rate))
}
fn from_bucket(bucket: TokenBucket) -> Self {
Self {
inner: Arc::new(Inner {
rate: bucket.rate(),
burst: bucket.burst(),
bucket: Mutex::new(bucket),
epoch: Instant::now(),
refunds: Arc::new(Notify::new()),
}),
}
}
#[must_use]
pub fn rate(&self) -> Rate {
self.inner.rate
}
#[must_use]
pub fn burst(&self) -> u64 {
self.inner.burst
}
fn now_nanos(&self) -> u64 {
u64::try_from(
Instant::now()
.saturating_duration_since(self.inner.epoch)
.as_nanos(),
)
.unwrap_or(u64::MAX)
}
pub fn try_acquire(&self, n: u64) -> Acquire {
self.inner.bucket.lock().try_acquire(self.now_nanos(), n)
}
pub async fn acquire(&self, n: u64) {
let mut remaining = n;
while remaining > 0 {
let want = remaining.min(self.inner.burst);
let refunded = self.inner.refunds.clone().notified_owned();
tokio::pin!(refunded);
refunded.as_mut().enable();
match self.try_acquire(want) {
Acquire::Granted => {
remaining -= want;
}
Acquire::RetryAt(at) => {
tokio::select! {
() = tokio::time::sleep_until(self.deadline(at)) => {}
() = refunded.as_mut() => {}
}
}
Acquire::Never => {
debug_assert!(false, "burst-clamped chunk reported Acquire::Never");
break;
}
}
}
}
pub fn refund(&self, n: u64) {
self.inner.bucket.lock().refund(n);
if n > 0 {
self.inner.refunds.notify_waiters();
}
}
pub fn notified_on_refund(&self) -> RefundWait {
RefundWait::new(self.inner.refunds.clone())
}
#[must_use]
pub fn deadline(&self, retry_at_nanos: u64) -> Instant {
self.inner
.epoch
.checked_add(Duration::from_nanos(retry_at_nanos))
.unwrap_or_else(|| {
Instant::now()
.checked_add(Duration::from_hours(8_760))
.unwrap_or_else(Instant::now)
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test(start_paused = true)]
async fn acquire_paces_exactly() {
let limiter = RateLimiter::from_rate(Rate::per_sec(10));
limiter.acquire(10).await;
let start = Instant::now();
for i in 1..=5u64 {
limiter.acquire(1).await;
assert_eq!(start.elapsed(), Duration::from_millis(i * 100));
}
}
#[tokio::test(start_paused = true)]
async fn acquire_chunks_oversized_requests() {
let limiter = RateLimiter::new(Rate::per_sec(10), 10);
let start = Instant::now();
limiter.acquire(25).await;
assert_eq!(start.elapsed(), Duration::from_millis(1_500));
}
#[tokio::test(start_paused = true)]
async fn try_acquire_never_waits() {
let limiter = RateLimiter::new(Rate::per_sec(2), 2);
assert_eq!(limiter.try_acquire(2), Acquire::Granted);
let retry_at = match limiter.try_acquire(1) {
Acquire::RetryAt(at) => at,
other => panic!("expected retry-at, got {other:?}"),
};
assert_eq!(
limiter.deadline(retry_at),
Instant::now() + Duration::from_millis(500)
);
assert_eq!(limiter.try_acquire(3), Acquire::Never);
}
#[tokio::test(start_paused = true)]
async fn clones_share_the_budget() {
let limiter = RateLimiter::new(Rate::per_sec(10), 2);
let clone = limiter.clone();
assert_eq!(limiter.rate(), Rate::per_sec(10));
assert_eq!(limiter.burst(), 2);
assert_eq!(limiter.try_acquire(1), Acquire::Granted);
assert_eq!(clone.try_acquire(1), Acquire::Granted);
assert!(matches!(clone.try_acquire(1), Acquire::RetryAt(_)));
limiter.refund(1);
assert_eq!(clone.try_acquire(1), Acquire::Granted);
}
#[tokio::test(start_paused = true)]
async fn a_refund_wakes_a_waiter_before_its_old_deadline() {
let limiter = RateLimiter::new(Rate::per_sec(10), 10);
limiter.acquire(10).await;
let waiting_limiter = limiter.clone();
let start = Instant::now();
let waiting = tokio::spawn(async move { waiting_limiter.acquire(10).await });
tokio::task::yield_now().await;
limiter.refund(10);
waiting.await.unwrap();
assert_eq!(start.elapsed(), Duration::ZERO);
}
#[tokio::test]
async fn refund_of_zero_does_not_wake_listeners() {
let limiter = RateLimiter::new(Rate::per_sec(10), 10);
let waiting_limiter = limiter.clone();
let waiting = tokio::spawn(async move { waiting_limiter.notified_on_refund().await });
tokio::task::yield_now().await;
limiter.refund(0);
tokio::task::yield_now().await;
assert!(!waiting.is_finished());
limiter.refund(1);
waiting.await.unwrap();
}
#[tokio::test(start_paused = true)]
async fn cancelling_oversized_acquire_does_not_mint_shared_budget() {
let limiter = RateLimiter::new(Rate::per_sec(10), 10);
limiter.acquire(10).await;
let acquire_limiter = limiter.clone();
let task = tokio::spawn(async move { acquire_limiter.acquire(20).await });
tokio::task::yield_now().await;
tokio::time::advance(Duration::from_secs(1)).await;
tokio::task::yield_now().await;
tokio::time::advance(Duration::from_millis(500)).await;
tokio::task::yield_now().await;
assert_eq!(limiter.try_acquire(5), Acquire::Granted);
task.abort();
let _cancelled = task.await;
assert!(matches!(limiter.try_acquire(1), Acquire::RetryAt(_)));
}
#[tokio::test(start_paused = true)]
async fn acquire_within_burst_is_cancel_safe() {
let limiter = RateLimiter::new(Rate::per_sec(10), 10);
limiter.acquire(10).await;
let cancelled = tokio::time::timeout(Duration::from_millis(500), limiter.acquire(10)).await;
cancelled.unwrap_err();
assert_eq!(limiter.try_acquire(5), Acquire::Granted);
}
#[tokio::test(start_paused = true)]
async fn concurrent_waiters_all_proceed() {
let limiter = RateLimiter::from_rate(Rate::per_sec(10));
limiter.acquire(10).await;
let start = Instant::now();
let mut set = tokio::task::JoinSet::new();
for _ in 0..5 {
let limiter = limiter.clone();
set.spawn(async move { limiter.acquire(1).await });
}
while let Some(res) = set.join_next().await {
res.unwrap();
}
assert_eq!(start.elapsed(), Duration::from_millis(500));
}
}