#![deny(clippy::all)]
#![deny(keyword_idents)]
#![deny(missing_docs)]
#![deny(non_ascii_idents)]
#![deny(unreachable_pub)]
#![deny(unsafe_code)]
#![deny(unused_crate_dependencies)]
#![deny(unused_qualifications)]
#![deny(warnings)]
use governor::clock::{Clock, DefaultClock};
pub use governor::Quota;
use lazy_static::lazy_static;
pub use limit_error::LimitError;
#[cfg(feature = "limit_info")]
pub use limit_header_gen::LimitHeaderGen;
use logger::{error, info, trace};
use registry::Registry;
#[cfg(feature = "limit_info")]
pub use req_state::ReqState;
pub use rocket::http::Method;
use rocket::{
async_trait, catch,
http::Status,
request::{FromRequest, Outcome},
Request,
};
pub use rocket_governable::RocketGovernable;
use std::marker::PhantomData;
pub use std::num::NonZeroU32;
pub mod header;
mod limit_error;
#[cfg(feature = "limit_info")]
mod limit_header_gen;
mod logger;
mod registry;
#[cfg(feature = "limit_info")]
mod req_state;
mod rocket_governable;
pub struct RocketGovernor<'r, T>
where
T: RocketGovernable<'r>,
{
_phantom: PhantomData<&'r T>,
}
lazy_static! {
static ref CLOCK: DefaultClock = DefaultClock::default();
}
#[doc(hidden)]
impl<'r, T> RocketGovernor<'r, T>
where
T: RocketGovernable<'r>,
{
#[inline(always)]
pub fn handle_from_request(request: &'r Request) -> Outcome<Self, LimitError> {
let res = request.local_cache(|| {
if let Some(route) = request.route() {
if let Some(route_name) = &route.name {
let limiter = Registry::get_or_insert::<T>(
route.method,
route_name,
T::quota(route.method, route_name),
);
if let Some(client_ip) = request.client_ip() {
let limit_check_res = limiter.check_key(&client_ip);
match limit_check_res {
Ok(state) => {
#[allow(unused_variables)] let request_capacity = state.remaining_burst_capacity();
trace!(
"not governed ip {} method {} route {}: remaining request capacity {}",
&client_ip,
&route.method,
route_name,
request_capacity
);
#[cfg(feature = "limit_info")] {
let req_state = ReqState::new(state.quota(), request_capacity);
let is_req_state_allowed = T::limit_info_allow(Some(route.method), Some(route_name), &req_state);
if is_req_state_allowed {
let _ = request.local_cache(|| req_state);
}
}
Ok(()) }
Err(notuntil) => {
let wait_time = notuntil.wait_time_from(CLOCK.now()).as_secs();
info!(
"ip {} method {} route {} limited {} sec",
&client_ip, &route.method, route_name, &wait_time
);
Err(LimitError::GovernedRequest(wait_time, notuntil.quota()))
}
}
} else {
error!(
"missing ip - method {} route {}: request: {:?}",
&route.method, route_name, request
);
Err(LimitError::MissingClientIpAddr)
}
} else {
error!("route without name: request: {:?}", request);
Err(LimitError::MissingRouteName)
}
} else {
error!("routing failure: request: {:?}", request);
Err(LimitError::MissingRoute)
}
});
match res {
Ok(_) => {
#[cfg(feature = "limit_info")]
{
let state_opt = ReqState::get_or_default(request);
#[allow(unused_variables)] if let Some(state) = state_opt {
trace!(
"request_capacity: {} rate-limit: {}",
state.request_capacity,
state.quota.burst_size().get()
);
}
}
Outcome::Success(Self::default())
}
Err(e) => {
let e = e.clone();
match e {
LimitError::GovernedRequest(_, _) => {
Outcome::Error((Status::TooManyRequests, e))
}
_ => Outcome::Error((Status::BadRequest, e)),
}
}
}
}
}
#[doc(hidden)]
impl<'r, T> Default for RocketGovernor<'r, T>
where
T: RocketGovernable<'r>,
{
fn default() -> Self {
Self {
_phantom: PhantomData,
}
}
}
#[doc(hidden)]
#[async_trait]
impl<'r, T> FromRequest<'r> for RocketGovernor<'r, T>
where
T: RocketGovernable<'r>,
{
type Error = LimitError;
async fn from_request(request: &'r Request<'_>) -> Outcome<Self, LimitError> {
Self::handle_from_request(request)
}
}
#[catch(429)]
pub fn rocket_governor_catcher<'r>(request: &'r Request) -> &'r LimitError {
let cached_res: &Result<(), LimitError> = request.local_cache(|| Err(LimitError::Error));
if let Err(limit_err) = cached_res {
limit_err
} else {
&LimitError::Error
}
}