use super::{NonZeroU32, Quota};
use rocket::Request;
#[derive(Debug)]
pub struct ReqState {
is_default: bool,
pub(crate) quota: Quota,
pub(crate) request_capacity: u32,
}
impl ReqState {
pub(crate) fn new(quota: Quota, request_capacity: u32) -> Self {
Self {
is_default: false,
quota,
request_capacity,
}
}
pub(crate) fn get_or_default<'r>(request: &'r Request) -> Option<&'r Self> {
let state: &ReqState = request.local_cache(|| {
ReqState::default() });
if !state.is_default {
Some(state)
} else {
None
}
}
pub fn quota(&self) -> &Quota {
&self.quota
}
pub fn request_capacity(&self) -> u32 {
self.request_capacity
}
}
impl Default for ReqState {
fn default() -> Self {
Self {
is_default: true,
quota: Quota::per_second(NonZeroU32::new(1).unwrap()),
request_capacity: 0,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use rocket::{
get,
http::{Header, Status},
local::blocking::Client,
routes, Build, Rocket,
};
#[get("/")]
fn route_test() -> Status {
Status::Ok
}
fn launch_rocket() -> Rocket<Build> {
rocket::build().mount("/", routes![route_test])
}
#[test]
fn test_req_state() {
let client = Client::untracked(launch_rocket()).expect("no rocket instance");
let mut req = client.get("/");
req.add_header(Header::new("X-Real-IP", "127.1.1.1"));
let request = req.inner_mut();
let _ = request
.local_cache(|| ReqState::new(Quota::per_second(NonZeroU32::new(1).unwrap()), 10));
let _ = request.real_ip();
let req_state = ReqState::get_or_default(request);
assert!(req_state.is_some());
assert_eq!(req_state.unwrap().request_capacity, 10);
let req_state = ReqState::get_or_default(request);
assert!(req_state.is_some());
assert_eq!(req_state.unwrap().request_capacity, 10);
let mut req = client.get("/");
req.add_header(Header::new("X-Real-IP", "127.1.1.2"));
let request = req.inner_mut();
let req_state = ReqState::get_or_default(request);
assert!(req_state.is_none());
}
}