systemprompt_api/services/middleware/rate_limit/
mod.rs1mod 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 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}