#![allow(zero_ptr)]
use chrono::Utc;
use hyper::client::{RequestBuilder, Response};
use hyper::header::Headers;
use hyper::status::StatusCode;
use std::collections::HashMap;
use std::sync::{Arc, Mutex};
use std::time::Duration;
use std::{i64, str, thread};
use super::{HttpError, LightMethod};
use ::internal::prelude::*;
lazy_static! {
pub static ref GLOBAL: Arc<Mutex<()>> = Arc::new(Mutex::new(()));
pub static ref ROUTES: Arc<Mutex<HashMap<Route, Arc<Mutex<RateLimit>>>>> = Arc::new(Mutex::new(HashMap::default()));
}
#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)]
pub enum Route {
ChannelsId(u64),
ChannelsIdInvites(u64),
ChannelsIdMessages(u64),
ChannelsIdMessagesBulkDelete(u64),
ChannelsIdMessagesId(LightMethod, u64),
ChannelsIdMessagesIdAck(u64),
ChannelsIdMessagesIdReactions(u64),
ChannelsIdMessagesIdReactionsUserIdType(u64),
ChannelsIdPermissionsOverwriteId(u64),
ChannelsIdPins(u64),
ChannelsIdPinsMessageId(u64),
ChannelsIdTyping(u64),
ChannelsIdWebhooks(u64),
Gateway,
GatewayBot,
Guilds,
GuildsId(u64),
GuildsIdBans(u64),
GuildsIdBansUserId(u64),
GuildsIdChannels(u64),
GuildsIdEmbed(u64),
GuildsIdEmojis(u64),
GuildsIdEmojisId(u64),
GuildsIdIntegrations(u64),
GuildsIdIntegrationsId(u64),
GuildsIdIntegrationsIdSync(u64),
GuildsIdInvites(u64),
GuildsIdMembers(u64),
GuildsIdMembersId(u64),
GuildsIdMembersIdRolesId(u64),
GuildsIdMembersMeNick(u64),
GuildsIdPrune(u64),
GuildsIdRegions(u64),
GuildsIdRoles(u64),
GuildsIdRolesId(u64),
GuildsIdWebhooks(u64),
InvitesCode,
UsersId,
UsersMe,
UsersMeChannels,
UsersMeGuilds,
UsersMeGuildsId,
VoiceRegions,
WebhooksId,
None,
}
pub(crate) fn perform<'a, F>(route: Route, f: F) -> Result<Response>
where F: Fn() -> RequestBuilder<'a> {
loop {
let _ = GLOBAL.lock().expect("global route lock poisoned");
let bucket = ROUTES
.lock()
.expect("routes poisoned")
.entry(route)
.or_insert_with(|| Arc::new(Mutex::new(RateLimit {
limit: i64::MAX,
remaining: i64::MAX,
reset: i64::MAX,
}))).clone();
let mut lock = bucket.lock().unwrap();
lock.pre_hook(&route);
let response = super::retry(&f)?;
if route == Route::None {
return Ok(response);
} else {
let redo = if response.headers.get_raw("x-ratelimit-global").is_some() {
let _ = GLOBAL.lock().expect("global route lock poisoned");
Ok(if let Some(retry_after) = parse_header(&response.headers, "retry-after")? {
debug!("Ratelimited on route {:?} for {:?}ms", route, retry_after);
thread::sleep(Duration::from_millis(retry_after as u64));
true
} else {
false
})
} else {
lock.post_hook(&response, &route)
};
if !redo.unwrap_or(true) {
return Ok(response);
}
}
}
}
#[derive(Clone, Debug, Default)]
pub struct RateLimit {
pub limit: i64,
pub remaining: i64,
pub reset: i64,
}
impl RateLimit {
pub(crate) fn pre_hook(&mut self, route: &Route) {
if self.limit == 0 {
return;
}
let current_time = Utc::now().timestamp();
if current_time > self.reset {
self.remaining = self.limit;
return;
}
let diff = (self.reset - current_time) as u64;
if self.remaining == 0 {
let delay = (diff * 1000) + 500;
debug!("Pre-emptive ratelimit on route {:?} for {:?}ms", route, delay);
thread::sleep(Duration::from_millis(delay));
return;
}
self.remaining -= 1;
}
pub(crate) fn post_hook(&mut self, response: &Response, route: &Route) -> 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;
}
if let Some(reset) = parse_header(&response.headers, "x-ratelimit-reset")? {
self.reset = reset;
}
Ok(if response.status != StatusCode::TooManyRequests {
false
} else if let Some(retry_after) = parse_header(&response.headers, "retry-after")? {
debug!("Ratelimited on route {:?} for {:?}ms", route, retry_after);
thread::sleep(Duration::from_millis(retry_after as u64));
true
} else {
false
})
}
}
fn parse_header(headers: &Headers, header: &str) -> Result<Option<i64>> {
match headers.get_raw(header) {
Some(header) => match str::from_utf8(&header[0]) {
Ok(v) => match v.parse::<i64>() {
Ok(v) => Ok(Some(v)),
Err(_) => Err(Error::Http(HttpError::RateLimitI64)),
},
Err(_) => Err(Error::Http(HttpError::RateLimitUtf8)),
},
None => Ok(None),
}
}