Skip to main content

systemprompt_api/services/middleware/rate_limit/
mod.rs

1//! Router extension traits for rate limiting and authenticated route groups.
2//!
3//! `RouterExt::with_auth` attaches authentication and authorization in one
4//! call: it requires an `AuthzPolicy`, so a route group cannot be mounted
5//! authenticated-but-unauthorized — omitting the policy is a compile error.
6//!
7//! `RouterExt::with_rate_limit` mounts two throttles: the in-process governor
8//! keyed by verified identity or trusted client IP, which smooths bursts per
9//! replica, and a database-backed window keyed by verified identity only,
10//! which bounds a caller's budget across every replica of the deployment.
11//!
12//! The database-backed window applies only to routers that carry `with_auth`:
13//! the identity comes from a
14//! [`RequestContext`](systemprompt_models::RequestContext) in the request's
15//! extensions, and a router that authenticates inside its handlers never puts
16//! one there. On those routers every request is anonymous to that layer and
17//! passes straight through, so the per-address governor is the only limit in
18//! force.
19//!
20//! Copyright (c) systemprompt.io — Business Source License 1.1.
21//! See <https://systemprompt.io> for licensing details.
22
23mod global;
24mod key;
25
26use crate::services::middleware::authz::{AuthzPolicy, authz_gate};
27use crate::services::middleware::context::{
28    A2AContextMiddleware, McpContextMiddleware, PublicContextMiddleware, UserOnlyContextMiddleware,
29};
30use axum::Router;
31use axum::extract::Request;
32use axum::middleware::Next;
33use axum::response::Response;
34use ipnet::IpNet;
35use std::future::Future;
36use std::sync::Arc;
37use std::time::Duration;
38use systemprompt_extension::LoaderError;
39use systemprompt_models::Config;
40use systemprompt_models::profile::RateLimitsConfig;
41use systemprompt_runtime::AppContext;
42use systemprompt_users::UserRateLimitBucketRepository;
43
44use self::global::{GLOBAL_WINDOW_SECS, GlobalUserLimit, global_user_rate_limit};
45pub use self::key::IdentityOrTrustedIpKey;
46
47#[derive(Clone, Debug)]
48pub struct RateLimitState {
49    config: RateLimitsConfig,
50    trusted_proxies: Arc<Vec<IpNet>>,
51    buckets: Arc<UserRateLimitBucketRepository>,
52}
53
54impl RateLimitState {
55    #[must_use]
56    pub fn new(config: &Config, buckets: Arc<UserRateLimitBucketRepository>) -> Self {
57        Self {
58            config: config.rate_limits,
59            trusted_proxies: Arc::new(config.trusted_proxies.clone()),
60            buckets,
61        }
62    }
63
64    pub fn from_context(ctx: &AppContext) -> Result<Self, LoaderError> {
65        let buckets = crate::repository::user_rate_limit_buckets(ctx.db_pool()).map_err(|e| {
66            LoaderError::InitializationFailed {
67                extension: "rate_limit".to_owned(),
68                message: e.to_string(),
69            }
70        })?;
71        Ok(Self::new(ctx.config(), buckets))
72    }
73}
74
75pub trait ContextLayer: Clone + Send + Sync + 'static {
76    fn handle(self, req: Request, next: Next) -> impl Future<Output = Response> + Send;
77}
78
79impl ContextLayer for PublicContextMiddleware {
80    async fn handle(self, req: Request, next: Next) -> Response {
81        Self::handle(&self, req, next).await
82    }
83}
84
85impl ContextLayer for UserOnlyContextMiddleware {
86    async fn handle(self, req: Request, next: Next) -> Response {
87        Self::handle(&self, req, next).await
88    }
89}
90
91impl ContextLayer for A2AContextMiddleware {
92    async fn handle(self, req: Request, next: Next) -> Response {
93        Self::handle(&self, req, next).await
94    }
95}
96
97impl ContextLayer for McpContextMiddleware {
98    async fn handle(self, req: Request, next: Next) -> Response {
99        Self::handle(&self, req, next).await
100    }
101}
102
103pub trait RouterExt<S>: Sized {
104    fn with_rate_limit(
105        self,
106        limits: &RateLimitState,
107        per_second: u64,
108        scope: &'static str,
109    ) -> Result<Self, LoaderError>;
110
111    fn with_auth<L: ContextLayer>(self, auth: L, policy: AuthzPolicy) -> Self;
112}
113
114impl<S> RouterExt<S> for Router<S>
115where
116    S: Clone + Send + Sync + 'static,
117{
118    fn with_rate_limit(
119        self,
120        limits: &RateLimitState,
121        per_second: u64,
122        scope: &'static str,
123    ) -> Result<Self, LoaderError> {
124        let rate_config = &limits.config;
125        if rate_config.disabled {
126            return Ok(self);
127        }
128
129        let burst = per_second.saturating_mul(rate_config.burst_multiplier);
130        let burst_u32 = u32::try_from(burst).unwrap_or(u32::MAX).max(1);
131        let per_second_clamped = per_second.max(1);
132
133        // Why: `GovernorConfigBuilder::per_second(n)` replenishes one element
134        // every n seconds, the inverse of a rate; `period` takes 1/per_second
135        // directly.
136        let replenish = Duration::from_secs(1)
137            .checked_div(u32::try_from(per_second_clamped).unwrap_or(u32::MAX))
138            .filter(|d| !d.is_zero())
139            .unwrap_or(Duration::from_nanos(1));
140        let rate_limit = tower_governor::governor::GovernorConfigBuilder::default()
141            .period(replenish)
142            .burst_size(burst_u32)
143            .key_extractor(IdentityOrTrustedIpKey::new(Arc::clone(
144                &limits.trusted_proxies,
145            )))
146            .use_headers()
147            .finish()
148            .ok_or_else(|| LoaderError::InitializationFailed {
149                extension: "rate_limit".to_owned(),
150                message: format!(
151                    "rate limit rejected for {per_second_clamped}/s with burst {burst_u32}"
152                ),
153            })?;
154
155        let window_secs = u64::try_from(GLOBAL_WINDOW_SECS).unwrap_or(u64::MAX);
156        let budget = burst.saturating_mul(window_secs);
157        let global = GlobalUserLimit {
158            buckets: Arc::clone(&limits.buckets),
159            scope,
160            budget: i64::try_from(budget).unwrap_or(i64::MAX),
161        };
162
163        Ok(self
164            .layer(axum::middleware::from_fn_with_state(
165                global,
166                global_user_rate_limit,
167            ))
168            .layer(tower_governor::GovernorLayer::new(rate_limit)))
169    }
170
171    fn with_auth<L: ContextLayer>(self, auth: L, policy: AuthzPolicy) -> Self {
172        self.layer(axum::middleware::from_fn(move |req, next| async move {
173            authz_gate(policy, req, next).await
174        }))
175        .layer(axum::middleware::from_fn(move |req, next| {
176            let auth = auth.clone();
177            async move { auth.handle(req, next).await }
178        }))
179    }
180}