use std::collections::HashMap;
use std::fmt;
use std::str::{self, FromStr};
use std::sync::Arc;
use std::time::SystemTime;
use reqwest::header::HeaderMap;
use reqwest::{Client, Response, StatusCode};
use tokio::sync::{Mutex, RwLock};
use tokio::time::{sleep, Duration};
use tracing::{debug, instrument};
pub use super::routing::Route;
use super::routing::RouteInfo;
use super::{HttpError, LightMethod, Request};
use crate::internal::prelude::*;
#[derive(Clone, Debug)]
#[non_exhaustive]
pub struct RatelimitInfo {
pub timeout: std::time::Duration,
pub limit: i64,
pub method: LightMethod,
pub path: String,
pub global: bool,
}
pub struct Ratelimiter {
client: Client,
global: Arc<Mutex<()>>,
routes: Arc<RwLock<HashMap<Route, Arc<Mutex<Ratelimit>>>>>,
token: String,
ratelimit_callback: Box<dyn Fn(RatelimitInfo) + Send + Sync>,
}
impl fmt::Debug for Ratelimiter {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Ratelimiter")
.field("client", &self.client)
.field("global", &self.global)
.field("routes", &self.routes)
.finish()
}
}
impl Ratelimiter {
pub fn new(client: Client, token: impl Into<String>) -> Self {
Self::_new(client, token.into())
}
fn _new(client: Client, token: String) -> Self {
Self {
client,
global: Arc::default(),
routes: Arc::default(),
token,
ratelimit_callback: Box::new(|_| {}),
}
}
pub fn set_ratelimit_callback(
&mut self,
ratelimit_callback: Box<dyn Fn(RatelimitInfo) + Send + Sync>,
) {
self.ratelimit_callback = ratelimit_callback;
}
#[must_use]
pub fn routes(&self) -> Arc<RwLock<HashMap<Route, Arc<Mutex<Ratelimit>>>>> {
Arc::clone(&self.routes)
}
#[instrument]
pub async fn perform(&self, req: RatelimitedRequest<'_>) -> Result<Response> {
let RatelimitedRequest {
mut req,
} = req;
loop {
drop(self.global.lock().await);
let (method, route, path) = req.route.deconstruct();
let path = path.to_string();
let bucket = Arc::clone(self.routes.write().await.entry(route).or_default());
bucket.lock().await.pre_hook(&req.route, &self.ratelimit_callback).await;
let request = req.build(&self.client, &self.token, None).await?.build()?;
let response = self.client.execute(request).await?;
if route == Route::None {
return Ok(response);
}
let redo = if response.headers().get("x-ratelimit-global").is_some() {
drop(self.global.lock().await);
Ok(
if let Some(retry_after) =
parse_header::<f64>(response.headers(), "retry-after")?
{
debug!("Ratelimited on route {:?} for {:?}s", route, retry_after);
(self.ratelimit_callback)(RatelimitInfo {
timeout: Duration::from_secs_f64(retry_after),
limit: 50,
method,
path,
global: true,
});
sleep(Duration::from_secs_f64(retry_after)).await;
true
} else {
false
},
)
} else {
bucket.lock().await.post_hook(&response, &req.route, &self.ratelimit_callback).await
};
if !redo.unwrap_or(true) {
return Ok(response);
}
}
}
}
#[derive(Debug)]
pub struct Ratelimit {
limit: i64,
remaining: i64,
reset: Option<SystemTime>,
reset_after: Option<Duration>,
}
impl Ratelimit {
#[instrument(skip(ratelimit_callback))]
pub async fn pre_hook(
&mut self,
route: &RouteInfo<'_>,
ratelimit_callback: &(dyn Fn(RatelimitInfo) + Send + Sync),
) {
if self.limit() == 0 {
return;
}
let reset = if let Some(reset) = self.reset {
reset
} else {
self.remaining = self.limit;
return;
};
let delay = if let Ok(delay) = reset.duration_since(SystemTime::now()) {
delay
} else {
if self.remaining() != 0 {
self.remaining -= 1;
}
return;
};
if self.remaining() == 0 {
let (method, route, path) = route.deconstruct();
debug!("Pre-emptive ratelimit on route {:?} for {}ms", route, delay.as_millis(),);
ratelimit_callback(RatelimitInfo {
timeout: delay,
limit: self.limit,
method,
path: path.to_string(),
global: false,
});
sleep(delay).await;
return;
}
self.remaining -= 1;
}
#[instrument(skip(ratelimit_callback))]
pub async fn post_hook(
&mut self,
response: &Response,
route: &RouteInfo<'_>,
ratelimit_callback: &(dyn Fn(RatelimitInfo) + Send + Sync),
) -> Result<bool> {
if let Some(limit) = parse_header(response.headers(), "x-ratelimit-limit")? {
self.limit = limit;
}
if let Some(remaining) = parse_header(response.headers(), "x-ratelimit-remaining")? {
self.remaining = remaining;
}
#[cfg(feature = "absolute_ratelimits")]
if let Some(reset) = parse_header::<f64>(response.headers(), "x-ratelimit-reset")? {
self.reset = Some(std::time::UNIX_EPOCH + Duration::from_secs_f64(reset));
}
if let Some(reset_after) =
parse_header::<f64>(response.headers(), "x-ratelimit-reset-after")?
{
#[cfg(not(feature = "absolute_ratelimits"))]
{
self.reset = Some(SystemTime::now() + Duration::from_secs_f64(reset_after));
}
self.reset_after = Some(Duration::from_secs_f64(reset_after));
}
Ok(if response.status() != StatusCode::TOO_MANY_REQUESTS {
false
} else if let Some(retry_after) = parse_header::<f64>(response.headers(), "retry-after")? {
let (method, route, path) = route.deconstruct();
debug!("Ratelimited on route {:?} for {:?}s", route, retry_after);
ratelimit_callback(RatelimitInfo {
timeout: Duration::from_secs_f64(retry_after),
limit: self.limit,
method,
path: path.to_string(),
global: false,
});
sleep(Duration::from_secs_f64(retry_after)).await;
true
} else {
false
})
}
#[inline]
#[must_use]
pub fn limit(&self) -> i64 {
self.limit
}
#[inline]
#[must_use]
pub fn remaining(&self) -> i64 {
self.remaining
}
#[inline]
#[must_use]
pub fn reset(&self) -> Option<SystemTime> {
self.reset
}
#[inline]
#[must_use]
pub fn reset_after(&self) -> Option<Duration> {
self.reset_after
}
}
impl Default for Ratelimit {
fn default() -> Self {
Self {
limit: i64::MAX,
remaining: i64::MAX,
reset: None,
reset_after: None,
}
}
}
#[derive(Debug)]
pub struct RatelimitedRequest<'a> {
req: Request<'a>,
}
impl<'a> From<Request<'a>> for RatelimitedRequest<'a> {
fn from(req: Request<'a>) -> Self {
Self {
req,
}
}
}
fn parse_header<T: FromStr>(headers: &HeaderMap, header: &str) -> Result<Option<T>> {
let header = match headers.get(header) {
Some(v) => v,
None => return Ok(None),
};
let unicode =
str::from_utf8(header.as_bytes()).map_err(|_| Error::from(HttpError::RateLimitUtf8))?;
let num = unicode.parse().map_err(|_| Error::from(HttpError::RateLimitI64F64))?;
Ok(Some(num))
}
#[cfg(test)]
mod tests {
use std::error::Error as StdError;
use std::result::Result as StdResult;
use reqwest::header::{HeaderMap, HeaderName, HeaderValue};
use super::parse_header;
use crate::error::Error;
use crate::http::HttpError;
type Result<T> = StdResult<T, Box<dyn StdError>>;
fn headers() -> HeaderMap {
let pairs = &[
(HeaderName::from_static("x-ratelimit-limit"), HeaderValue::from_static("5")),
(HeaderName::from_static("x-ratelimit-remaining"), HeaderValue::from_static("4")),
(
HeaderName::from_static("x-ratelimit-reset"),
HeaderValue::from_static("1560704880.423"),
),
(HeaderName::from_static("x-bad-num"), HeaderValue::from_static("abc")),
(
HeaderName::from_static("x-bad-unicode"),
HeaderValue::from_bytes(&[255, 255, 255, 255]).unwrap(),
),
];
let mut map = HeaderMap::with_capacity(pairs.len());
for (name, val) in pairs {
map.insert(name, val.clone());
}
map
}
#[test]
#[allow(clippy::float_cmp)]
fn test_parse_header_good() -> Result<()> {
let headers = headers();
assert_eq!(parse_header::<i64>(&headers, "x-ratelimit-limit")?.unwrap(), 5);
assert_eq!(parse_header::<i64>(&headers, "x-ratelimit-remaining")?.unwrap(), 4,);
assert_eq!(parse_header::<f64>(&headers, "x-ratelimit-reset")?.unwrap(), 1_560_704_880.423);
Ok(())
}
#[test]
fn test_parse_header_errors() {
let headers = headers();
macro_rules! is_err {
($header:expr, $err:pat) => {
match parse_header::<i64>(&headers, $header).unwrap_err() {
Error::Http(x) => matches!(*x, $err),
_ => false,
}
};
}
assert!(is_err!("x-bad-num", HttpError::RateLimitI64F64));
assert!(is_err!("x-bad-unicode", HttpError::RateLimitUtf8));
}
}