reqkey 0.1.0

Official Rust SDK for ReqKey API key validation, credit metering, and analytics
Documentation
//! Axum 0.8 middleware integration.

use std::{
    future::Future,
    net::SocketAddr,
    pin::Pin,
    task::{Context, Poll},
};

use ::axum::{
    body::{to_bytes, Body, HttpBody},
    extract::{ConnectInfo, Request},
    response::Response,
};
use bytes::Bytes;
use futures_util::stream;
use tower::{Layer, Service};

use crate::middleware::{
    AuthorizationOutcome, DeniedResponse, Middleware, RequestContext, ResponseContext,
};

const MAX_CAPTURE_BYTES: usize = 4_004;

/// Tower layer that applies ReqKey policy to an Axum router or route.
#[derive(Clone)]
pub struct ReqKeyLayer {
    middleware: Middleware,
}

impl ReqKeyLayer {
    /// Create a layer from the shared middleware engine.
    pub const fn new(middleware: Middleware) -> Self {
        Self { middleware }
    }
}

impl<S> Layer<S> for ReqKeyLayer {
    type Service = ReqKeyService<S>;

    fn layer(&self, inner: S) -> Self::Service {
        ReqKeyService {
            inner,
            middleware: self.middleware.clone(),
        }
    }
}

/// Tower service produced by [`ReqKeyLayer`].
#[derive(Clone)]
pub struct ReqKeyService<S> {
    inner: S,
    middleware: Middleware,
}

impl<S> Service<Request> for ReqKeyService<S>
where
    S: Service<Request, Response = Response> + Clone + Send + 'static,
    S::Future: Send + 'static,
    S::Error: Send + 'static,
{
    type Response = Response;
    type Error = S::Error;
    type Future = Pin<Box<dyn Future<Output = Result<Response, S::Error>> + Send>>;

    fn poll_ready(&mut self, context: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
        self.inner.poll_ready(context)
    }

    fn call(&mut self, mut request: Request) -> Self::Future {
        let clone = self.inner.clone();
        let mut inner = std::mem::replace(&mut self.inner, clone);
        let middleware = self.middleware.clone();
        let context = request_context(&request);

        Box::pin(async move {
            match middleware.authorize(context).await {
                AuthorizationOutcome::Bypass => inner.call(request).await,
                AuthorizationOutcome::Denied(denial) => Ok(denial_response(&denial)),
                AuthorizationOutcome::Authorized(authorized) => {
                    if let Some(decision) = &authorized.decision {
                        request.extensions_mut().insert(decision.clone());
                    }
                    if let Some(failure) = &authorized.failure {
                        request.extensions_mut().insert(failure.clone());
                    }

                    let mut response = inner.call(request).await?;
                    let original_headers = response.headers().clone();
                    let response_body = capture_response_body(
                        &mut response,
                        middleware.config().captures_response_body(),
                    )
                    .await;
                    response.headers_mut().extend(authorized.response_headers());
                    let response_context = ResponseContext::new(response.status().as_u16())
                        .with_headers(original_headers)
                        .with_latency_ms(authorized.elapsed_ms());
                    let response_context = response_body.map_or(response_context.clone(), |body| {
                        response_context.with_body(body)
                    });
                    middleware.record(authorized, response_context).await;
                    Ok(response)
                }
            }
        })
    }
}

fn request_context(request: &Request) -> RequestContext {
    let mut context = RequestContext::new(request.method().clone(), request.uri().path())
        .with_headers(request.headers().clone());
    if let Some(query) = request.uri().query() {
        context = context.with_query(query);
    }
    if let Some(ConnectInfo(address)) = request.extensions().get::<ConnectInfo<SocketAddr>>() {
        context = context.with_client_ip(address.ip().to_string());
    }
    context
}

fn denial_response(denial: &DeniedResponse) -> Response {
    let mut response = Response::new(Body::from(denial.json_body()));
    *response.status_mut() = http::StatusCode::from_u16(denial.status_code)
        .unwrap_or(http::StatusCode::INTERNAL_SERVER_ERROR);
    response.headers_mut().extend(denial.headers());
    response
}

async fn capture_response_body(response: &mut Response, enabled: bool) -> Option<String> {
    if !enabled
        || response
            .headers()
            .get(http::header::CONTENT_ENCODING)
            .is_some_and(|value| value != "identity")
    {
        return None;
    }
    let exact = usize::try_from(response.body().size_hint().exact()?).ok()?;
    if exact > MAX_CAPTURE_BYTES {
        return None;
    }
    let replacement = Response::new(Body::empty());
    let original = std::mem::replace(response, replacement);
    let (parts, body) = original.into_parts();
    match to_bytes(body, MAX_CAPTURE_BYTES).await {
        Ok(bytes) => {
            let text = String::from_utf8_lossy(&bytes)
                .chars()
                .take(crate::client::MAX_BODY_CHARACTERS)
                .collect();
            *response = Response::from_parts(parts, Body::from(bytes));
            Some(text)
        }
        Err(error) => {
            let errors = stream::once(async move { Err::<Bytes, _>(error) });
            *response = Response::from_parts(parts, Body::from_stream(errors));
            None
        }
    }
}