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_identifiers::ExtensionId;
40use systemprompt_manifest::Config;
41use systemprompt_manifest::profile::RateLimitsConfig;
42use systemprompt_runtime::AppContext;
43use systemprompt_users::UserRateLimitBucketRepository;
44
45use self::global::{GLOBAL_WINDOW_SECS, GlobalUserLimit, global_user_rate_limit};
46pub use self::key::IdentityOrTrustedIpKey;
47
48#[derive(Clone, Debug)]
49pub struct RateLimitState {
50    config: RateLimitsConfig,
51    trusted_proxies: Arc<Vec<IpNet>>,
52    buckets: Arc<UserRateLimitBucketRepository>,
53}
54
55impl RateLimitState {
56    #[must_use]
57    pub fn new(config: &Config, buckets: Arc<UserRateLimitBucketRepository>) -> Self {
58        Self {
59            config: config.rate_limits,
60            trusted_proxies: Arc::new(config.trusted_proxies.clone()),
61            buckets,
62        }
63    }
64
65    pub fn from_context(ctx: &AppContext) -> Self {
66        let buckets = crate::repository::user_rate_limit_buckets(ctx.db_pool());
67        Self::new(ctx.config(), buckets)
68    }
69}
70
71pub trait ContextLayer: Clone + Send + Sync + 'static {
72    fn handle(self, req: Request, next: Next) -> impl Future<Output = Response> + Send;
73}
74
75impl ContextLayer for PublicContextMiddleware {
76    async fn handle(self, req: Request, next: Next) -> Response {
77        Self::handle(&self, req, next).await
78    }
79}
80
81impl ContextLayer for UserOnlyContextMiddleware {
82    async fn handle(self, req: Request, next: Next) -> Response {
83        Self::handle(&self, req, next).await
84    }
85}
86
87impl ContextLayer for A2AContextMiddleware {
88    async fn handle(self, req: Request, next: Next) -> Response {
89        Self::handle(&self, req, next).await
90    }
91}
92
93impl ContextLayer for McpContextMiddleware {
94    async fn handle(self, req: Request, next: Next) -> Response {
95        Self::handle(&self, req, next).await
96    }
97}
98
99pub trait RouterExt<S>: Sized {
100    fn with_rate_limit(
101        self,
102        limits: &RateLimitState,
103        per_second: u64,
104        scope: &'static str,
105    ) -> Result<Self, LoaderError>;
106
107    fn with_auth<L: ContextLayer>(self, auth: L, policy: AuthzPolicy) -> Self;
108}
109
110impl<S> RouterExt<S> for Router<S>
111where
112    S: Clone + Send + Sync + 'static,
113{
114    fn with_rate_limit(
115        self,
116        limits: &RateLimitState,
117        per_second: u64,
118        scope: &'static str,
119    ) -> Result<Self, LoaderError> {
120        let rate_config = &limits.config;
121        if rate_config.disabled {
122            return Ok(self);
123        }
124
125        let burst = per_second.saturating_mul(rate_config.burst_multiplier);
126        let burst_u32 = u32::try_from(burst).unwrap_or(u32::MAX).max(1);
127        let per_second_clamped = per_second.max(1);
128
129        // Why: `GovernorConfigBuilder::per_second(n)` replenishes one element
130        // every n seconds, the inverse of a rate; `period` takes 1/per_second
131        // directly.
132        let replenish = Duration::from_secs(1)
133            .checked_div(u32::try_from(per_second_clamped).unwrap_or(u32::MAX))
134            .filter(|d| !d.is_zero())
135            .unwrap_or(Duration::from_nanos(1));
136        let rate_limit = tower_governor::governor::GovernorConfigBuilder::default()
137            .period(replenish)
138            .burst_size(burst_u32)
139            .key_extractor(IdentityOrTrustedIpKey::new(Arc::clone(
140                &limits.trusted_proxies,
141            )))
142            .use_headers()
143            .finish()
144            .ok_or_else(|| LoaderError::InitializationFailed {
145                extension: ExtensionId::new("rate_limit"),
146                message: format!(
147                    "rate limit rejected for {per_second_clamped}/s with burst {burst_u32}"
148                ),
149            })?;
150
151        let window_secs = u64::try_from(GLOBAL_WINDOW_SECS).unwrap_or(u64::MAX);
152        let budget = burst.saturating_mul(window_secs);
153        let global = GlobalUserLimit {
154            buckets: Arc::clone(&limits.buckets),
155            scope,
156            budget: i64::try_from(budget).unwrap_or(i64::MAX),
157        };
158
159        Ok(self
160            .layer(axum::middleware::from_fn_with_state(
161                global,
162                global_user_rate_limit,
163            ))
164            .layer(tower_governor::GovernorLayer::new(rate_limit)))
165    }
166
167    fn with_auth<L: ContextLayer>(self, auth: L, policy: AuthzPolicy) -> Self {
168        self.layer(axum::middleware::from_fn(move |req, next| async move {
169            authz_gate(policy, req, next).await
170        }))
171        .layer(axum::middleware::from_fn(move |req, next| {
172            let auth = auth.clone();
173            async move { auth.handle(req, next).await }
174        }))
175    }
176}