reqkey 0.1.0

Official Rust SDK for ReqKey API key validation, credit metering, and analytics
Documentation
//! Rocket 0.5 request-guard and response-fairing integration.
//!
//! Rocket request fairings cannot short-circuit routing. Protected routes use
//! [`ReqKeyGuard`], while [`ReqKeyFairing`] records their responses.

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>>);

/// Fairing that installs shared state and records responses for [`ReqKeyGuard`].
pub struct ReqKeyFairing {
    middleware: Middleware,
}

impl ReqKeyFairing {
    /// Create the fairing from the shared middleware engine.
    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;
    }
}

/// Request guard required by each protected Rocket route.
#[derive(Clone, Debug)]
pub struct ReqKeyGuard {
    authorized: Option<AuthorizedRequest>,
}

impl ReqKeyGuard {
    /// Return the validation decision in validation/both mode.
    pub fn decision(&self) -> Option<&VerificationResult> {
        self.authorized
            .as_ref()
            .and_then(|authorized| authorized.decision.as_ref())
    }

    /// Return the fail-open service error, when present.
    pub fn failure(&self) -> Option<&ReqKeyFailure> {
        self.authorized
            .as_ref()
            .and_then(|authorized| authorized.failure.as_ref())
    }
}

/// Error returned by [`ReqKeyGuard`] before a protected route runs.
#[derive(Clone, Debug)]
pub enum ReqKeyGuardError {
    /// The fairing/state was not attached.
    MissingFairing,
    /// ReqKey denied this request.
    Denied(DeniedResponse),
}

impl ReqKeyGuardError {
    /// Return the detailed denial when the guard reached ReqKey policy.
    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),
                })
            }
        }
    }
}

/// Return JSON catchers for ReqKey's 401/402/403/429/503 denials.
///
/// Register them on the same mount scope as protected routes to preserve the
/// stable error body, custom message, and `Retry-After` header:
///
/// ```ignore
/// rocket::build().register("/api", reqkey::rocket::default_catchers())
/// ```
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
}