tower-rate-tier 0.3.0

Tier-based rate limiting middleware for Tower
Documentation
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use std::time::SystemTime;

use bytes::Bytes;
use http::{Request, Response, StatusCode};
use http_body::Body;
use http_body_util::{BodyExt, Full, LengthLimitError, Limited};
use tower_layer::Layer;
use tower_service::Service;

use crate::check::{self, CheckOutcome};
use crate::layer::Settings;
use crate::response;
use crate::storage::StorageKey;
use crate::tier::RateTier;

/// Tower layer for tier-based rate limiting with body-based identification.
///
/// Created by calling [`buffer_body()`](crate::TierLimitLayer::buffer_body)
/// on a [`TierLimitLayer`](crate::TierLimitLayer).
///
/// This layer buffers the request body before identification, enabling
/// [`TierIdentifier::identify_with_body`](crate::TierIdentifier::identify_with_body)
/// to inspect the body contents.
/// The body is then reconstructed as `Full<Bytes>` for the downstream service.
///
/// # Body size limit
///
/// Requests exceeding [`max_body_size`](Self::max_body_size) (default: 64KB)
/// are rejected with 413 Payload Too Large. A body whose declared length
/// (e.g. `Content-Length`) is over the limit is rejected without being read;
/// otherwise reading stops at the first chunk that crosses the limit.
///
/// Requires the `buffered-body` feature.
#[derive(Clone, Debug)]
pub struct BufferedTierLimitLayer {
    pub(crate) rate_tier: Arc<RateTier>,
    pub(crate) settings: Settings,
    pub(crate) max_body_size: usize,
}

impl BufferedTierLimitLayer {
    /// Set the maximum allowed body size in bytes.
    ///
    /// Requests with bodies larger than this are rejected with 413.
    /// Default: 64KB.
    pub fn max_body_size(mut self, size: usize) -> Self {
        self.max_body_size = size;
        self
    }
}

impl<S> Layer<S> for BufferedTierLimitLayer {
    type Service = BufferedTierLimitService<S>;

    fn layer(&self, inner: S) -> Self::Service {
        BufferedTierLimitService {
            inner,
            rate_tier: self.rate_tier.clone(),
            settings: Arc::new(self.settings.clone()),
            max_body_size: self.max_body_size,
        }
    }
}

/// Tower service that buffers the request body for identification.
///
/// Created by [`BufferedTierLimitLayer`].
///
/// Requires the `buffered-body` feature.
#[derive(Clone, Debug)]
pub struct BufferedTierLimitService<S> {
    inner: S,
    rate_tier: Arc<RateTier>,
    settings: Arc<Settings>,
    max_body_size: usize,
}

impl<S, B, ResBody> Service<Request<B>> for BufferedTierLimitService<S>
where
    S: Service<Request<Full<Bytes>>, Response = Response<ResBody>> + Clone + Send + 'static,
    S::Future: Send,
    S::Error: Send,
    B: Body + Send + 'static,
    B::Data: Send,
    B::Error: Into<Box<dyn std::error::Error + Send + Sync>>,
    ResBody: From<String> + Send,
{
    type Response = Response<ResBody>;
    type Error = S::Error;
    type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;

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

    fn call(&mut self, req: Request<B>) -> Self::Future {
        let rate_tier = self.rate_tier.clone();
        let settings = self.settings.clone();
        let max_body_size = self.max_body_size;
        let mut inner = self.inner.clone();
        std::mem::swap(&mut self.inner, &mut inner);

        Box::pin(async move {
            // Split request to buffer body separately
            let (parts, body) = req.into_parts();

            // A body that already declares more than the limit is never read.
            if body.size_hint().lower() > max_body_size as u64 {
                return Ok(payload_too_large_response().map(Into::into));
            }

            // `Limited` fails on the first frame that crosses the limit, so at
            // most `max_body_size` bytes (plus that frame) are ever buffered.
            let body_bytes = match Limited::new(body, max_body_size).collect().await {
                Ok(collected) => collected.to_bytes(),
                Err(err) if err.is::<LengthLimitError>() => {
                    return Ok(payload_too_large_response().map(Into::into));
                }
                Err(_) => {
                    return Ok(response::bad_request_response().map(Into::into));
                }
            };

            // Identify using headers + body
            let identity = settings
                .identifier
                .identify_with_body(&parts.headers, &body_bytes)
                .await;

            let (user_id, tier_name) = match check::resolve_identity(identity, &rate_tier) {
                Ok(pair) => pair,
                Err(CheckOutcome::PassThrough) => {
                    let req = Request::from_parts(parts, Full::new(body_bytes));
                    return inner.call(req).await;
                }
                Err(CheckOutcome::Deny(resp)) => return Ok(resp.map(Into::into)),
                Err(CheckOutcome::Allow(_)) => unreachable!(),
            };

            let resolved = check::resolve_quota(&rate_tier, &user_id, tier_name, &settings);
            let (tier_name, quota) = match resolved {
                Ok(resolved) => resolved,
                Err(CheckOutcome::PassThrough) => {
                    let req = Request::from_parts(parts, Full::new(body_bytes));
                    return inner.call(req).await;
                }
                Err(CheckOutcome::Deny(resp)) => return Ok(resp.map(Into::into)),
                Err(CheckOutcome::Allow(_)) => unreachable!(),
            };

            let cost = check::request_cost(&parts, &settings);
            if let Some(resp) =
                check::reject_cost_over_limit(cost, quota, &user_id, &tier_name, &settings)
            {
                return Ok(resp.map(Into::into));
            }
            let now = rate_tier.clock().now();
            let key = StorageKey::new(&user_id, &tier_name);
            let result = rate_tier
                .storage()
                .check_and_update(key, quota, cost, now)
                .await;

            // Wall-clock time of the check; the durations in the result are
            // relative to it.
            let checked_at = SystemTime::now();

            match check::process_result(result, &user_id, &tier_name, &settings, checked_at) {
                CheckOutcome::Allow(info) => {
                    let req = Request::from_parts(parts, Full::new(body_bytes));
                    let mut resp = inner.call(req).await?;
                    response::inject_headers(&mut resp, &info, checked_at);
                    Ok(resp)
                }
                CheckOutcome::Deny(resp) => Ok(resp.map(Into::into)),
                CheckOutcome::PassThrough => {
                    let req = Request::from_parts(parts, Full::new(body_bytes));
                    inner.call(req).await
                }
            }
        })
    }
}

fn payload_too_large_response() -> Response<String> {
    Response::builder()
        .status(StatusCode::PAYLOAD_TOO_LARGE)
        .header("Content-Type", "application/json")
        .body(r#"{"error":"payload too large"}"#.to_string())
        .unwrap()
}