use std::sync::Arc;
use tokio::sync::{OwnedSemaphorePermit, Semaphore};
use tracing::{debug, warn};
#[derive(Clone)]
pub struct DriveRateLimiter {
semaphore: Arc<Semaphore>,
max_requests_per_period: u32,
period_seconds: u64,
}
impl DriveRateLimiter {
pub fn new() -> Self {
Self::with_limits(10, 800, 100)
}
pub fn with_limits(max_concurrent: usize, max_requests: u32, period_secs: u64) -> Self {
debug!(
"Criando rate limiter: {} concorrentes, {} req/{} sec",
max_concurrent, max_requests, period_secs
);
Self {
semaphore: Arc::new(Semaphore::new(max_concurrent)),
max_requests_per_period: max_requests,
period_seconds: period_secs,
}
}
pub async fn acquire(&self) -> RateLimitGuard {
let permit = self
.semaphore
.clone()
.acquire_owned()
.await
.expect("Semaphore should never be closed");
RateLimitGuard {
_permit: permit,
}
}
pub fn try_acquire(&self) -> Option<RateLimitGuard> {
self.semaphore.clone().try_acquire_owned().ok().map(|permit| RateLimitGuard {
_permit: permit,
})
}
pub fn available_permits(&self) -> usize {
self.semaphore.available_permits()
}
pub fn is_near_limit(&self) -> bool {
let available = self.available_permits();
let total = self.semaphore.available_permits() + 1; available < total / 5
}
}
impl Default for DriveRateLimiter {
fn default() -> Self {
Self::new()
}
}
pub struct RateLimitGuard {
_permit: OwnedSemaphorePermit,
}
pub struct RateLimitedClient<T> {
inner: T,
limiter: DriveRateLimiter,
}
impl<T> RateLimitedClient<T> {
pub fn new(client: T, limiter: DriveRateLimiter) -> Self {
Self {
inner: client,
limiter,
}
}
pub fn inner(&self) -> &T {
&self.inner
}
pub fn inner_mut(&mut self) -> &mut T {
&mut self.inner
}
pub fn limiter(&self) -> &DriveRateLimiter {
&self.limiter
}
pub async fn execute<F, Fut, R>(&self, operation: F) -> R
where
F: FnOnce(&T) -> Fut,
Fut: std::future::Future<Output = R>,
{
let _guard = self.limiter.acquire().await;
if self.limiter.is_near_limit() {
warn!("Rate limiter próximo do limite - considere reduzir a taxa de requisições");
}
operation(&self.inner).await
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::time::{Duration, Instant};
use tokio::time::sleep;
#[tokio::test]
async fn test_rate_limiter_basic() {
let limiter = DriveRateLimiter::with_limits(2, 10, 1);
let _guard1 = limiter.acquire().await;
let _guard2 = limiter.acquire().await;
assert_eq!(limiter.available_permits(), 0);
}
#[tokio::test]
async fn test_rate_limiter_release() {
let limiter = DriveRateLimiter::with_limits(1, 10, 1);
{
let _guard = limiter.acquire().await;
assert_eq!(limiter.available_permits(), 0);
}
assert_eq!(limiter.available_permits(), 1);
}
#[tokio::test]
async fn test_try_acquire() {
let limiter = DriveRateLimiter::with_limits(1, 10, 1);
let guard1 = limiter.try_acquire();
assert!(guard1.is_some());
let guard2 = limiter.try_acquire();
assert!(guard2.is_none());
drop(guard1);
let guard3 = limiter.try_acquire();
assert!(guard3.is_some());
}
#[tokio::test]
async fn test_concurrent_operations() {
let limiter = Arc::new(DriveRateLimiter::with_limits(5, 100, 1));
let start = Instant::now();
let tasks: Vec<_> = (0..10)
.map(|i| {
let limiter = limiter.clone();
tokio::spawn(async move {
let _guard = limiter.acquire().await;
sleep(Duration::from_millis(100)).await;
i
})
})
.collect();
let results: Vec<_> = futures::future::join_all(tasks)
.await
.into_iter()
.map(|r| r.unwrap())
.collect();
let duration = start.elapsed();
assert_eq!(results.len(), 10);
assert!(duration >= Duration::from_millis(200));
}
}