use ::rocket::{
async_trait,
fairing::{Fairing, Info, Kind},
http::{ContentType, Header, Status},
request::{FromRequest, Outcome, Request},
response::{Responder, Response},
tokio::sync::Mutex,
Build, Catcher, Rocket,
};
use std::io::Cursor;
use crate::{
middleware::{
AuthorizationOutcome, AuthorizedRequest, DeniedResponse, Middleware, ReqKeyFailure,
RequestContext, ResponseContext,
},
VerificationResult,
};
struct RequestState(Mutex<Option<AuthorizedRequest>>);
struct DenialState(Mutex<Option<DeniedResponse>>);
pub struct ReqKeyFairing {
middleware: Middleware,
}
impl ReqKeyFairing {
pub const fn new(middleware: Middleware) -> Self {
Self { middleware }
}
}
#[async_trait]
impl Fairing for ReqKeyFairing {
fn info(&self) -> Info {
Info {
name: "ReqKey API validation and analytics",
kind: Kind::Ignite | Kind::Response,
}
}
async fn on_ignite(&self, rocket: Rocket<Build>) -> Result<Rocket<Build>, Rocket<Build>> {
Ok(rocket.manage(self.middleware.clone()))
}
async fn on_response<'r>(&self, request: &'r Request<'_>, response: &mut Response<'r>) {
let state = request.local_cache(|| RequestState(Mutex::new(None)));
let Some(authorized) = state.0.lock().await.take() else {
return;
};
let original_headers = rocket_headers_to_http(response.headers());
let decision_headers = authorized.response_headers();
for (name, value) in &decision_headers {
if let Ok(value) = value.to_str() {
response.set_header(Header::new(name.as_str().to_owned(), value.to_owned()));
}
}
let status = response.status().code;
self.middleware
.record(
authorized.clone(),
ResponseContext::new(status)
.with_headers(original_headers)
.with_latency_ms(authorized.elapsed_ms()),
)
.await;
}
}
#[derive(Clone, Debug)]
pub struct ReqKeyGuard {
authorized: Option<AuthorizedRequest>,
}
impl ReqKeyGuard {
pub fn decision(&self) -> Option<&VerificationResult> {
self.authorized
.as_ref()
.and_then(|authorized| authorized.decision.as_ref())
}
pub fn failure(&self) -> Option<&ReqKeyFailure> {
self.authorized
.as_ref()
.and_then(|authorized| authorized.failure.as_ref())
}
}
#[derive(Clone, Debug)]
pub enum ReqKeyGuardError {
MissingFairing,
Denied(DeniedResponse),
}
impl ReqKeyGuardError {
pub const fn denial(&self) -> Option<&DeniedResponse> {
match self {
Self::Denied(denial) => Some(denial),
Self::MissingFairing => None,
}
}
}
#[async_trait]
impl<'r> FromRequest<'r> for ReqKeyGuard {
type Error = ReqKeyGuardError;
async fn from_request(request: &'r Request<'_>) -> Outcome<Self, Self::Error> {
let Some(middleware) = request.rocket().state::<Middleware>() else {
return Outcome::Error((
Status::InternalServerError,
ReqKeyGuardError::MissingFairing,
));
};
match middleware.authorize(request_context(request)).await {
AuthorizationOutcome::Bypass => Outcome::Success(Self { authorized: None }),
AuthorizationOutcome::Denied(denial) => {
let status =
Status::from_code(denial.status_code).unwrap_or(Status::InternalServerError);
let state = request.local_cache(|| DenialState(Mutex::new(None)));
*state.0.lock().await = Some(denial.clone());
Outcome::Error((status, ReqKeyGuardError::Denied(denial)))
}
AuthorizationOutcome::Authorized(authorized) => {
let state = request.local_cache(|| RequestState(Mutex::new(None)));
*state.0.lock().await = Some(authorized.clone());
Outcome::Success(Self {
authorized: Some(authorized),
})
}
}
}
}
pub fn default_catchers() -> Vec<Catcher> {
::rocket::catchers![
unauthorized,
payment_required,
forbidden,
rate_limited,
unavailable
]
}
struct DenialJson(DeniedResponse);
impl<'r> Responder<'r, 'static> for DenialJson {
fn respond_to(self, _request: &'r Request<'_>) -> ::rocket::response::Result<'static> {
let body = self.0.json_body();
let mut response = Response::build();
response
.status(Status::from_code(self.0.status_code).unwrap_or(Status::InternalServerError))
.header(ContentType::JSON)
.sized_body(body.len(), Cursor::new(body));
if let Some(seconds) = self.0.retry_after {
response.raw_header("Retry-After", seconds.to_string());
}
response.ok()
}
}
async fn denial_from_request(request: &Request<'_>, status: u16) -> DenialJson {
let state = request.local_cache(|| DenialState(Mutex::new(None)));
let denial = state.0.lock().await.clone().unwrap_or_else(|| {
let (error, message) = match status {
401 => (
crate::middleware::DenialCode::InvalidApiKey,
"The API key is invalid or inactive.",
),
402 => (
crate::middleware::DenialCode::InsufficientCredits,
"The API key has insufficient credits.",
),
403 => (
crate::middleware::DenialCode::AccessDenied,
"The API key is not allowed to access this API.",
),
429 => (
crate::middleware::DenialCode::RateLimited,
"The API key has exceeded its rate limit.",
),
_ => (
crate::middleware::DenialCode::ReqKeyUnavailable,
"API key verification is temporarily unavailable.",
),
};
DeniedResponse {
status_code: status,
error,
message: message.into(),
retry_after: None,
}
});
DenialJson(denial)
}
#[::rocket::catch(401)]
async fn unauthorized(request: &Request<'_>) -> DenialJson {
denial_from_request(request, 401).await
}
#[::rocket::catch(402)]
async fn payment_required(request: &Request<'_>) -> DenialJson {
denial_from_request(request, 402).await
}
#[::rocket::catch(403)]
async fn forbidden(request: &Request<'_>) -> DenialJson {
denial_from_request(request, 403).await
}
#[::rocket::catch(429)]
async fn rate_limited(request: &Request<'_>) -> DenialJson {
denial_from_request(request, 429).await
}
#[::rocket::catch(503)]
async fn unavailable(request: &Request<'_>) -> DenialJson {
denial_from_request(request, 503).await
}
fn request_context(request: &Request<'_>) -> RequestContext {
let method =
http::Method::from_bytes(request.method().as_str().as_bytes()).unwrap_or(http::Method::GET);
let mut headers = http::HeaderMap::new();
for header in request.headers().iter() {
let Ok(name) = http::HeaderName::from_bytes(header.name().as_str().as_bytes()) else {
continue;
};
let Ok(value) = http::HeaderValue::from_str(header.value()) else {
continue;
};
headers.append(name, value);
}
let mut context =
RequestContext::new(method, request.uri().path().as_str()).with_headers(headers);
if let Some(query) = request.uri().query() {
context = context.with_query(query.as_str());
}
if let Some(ip) = request.client_ip() {
context = context.with_client_ip(ip.to_string());
}
context
}
fn rocket_headers_to_http(headers: &::rocket::http::HeaderMap<'_>) -> http::HeaderMap {
let mut converted = http::HeaderMap::new();
for header in headers.iter() {
let Ok(name) = http::HeaderName::from_bytes(header.name().as_str().as_bytes()) else {
continue;
};
let Ok(value) = http::HeaderValue::from_str(header.value()) else {
continue;
};
converted.append(name, value);
}
converted
}