systemprompt_api/services/middleware/
rate_limit.rs1use crate::services::middleware::authz::{AuthzPolicy, authz_gate};
16use crate::services::middleware::client_addr::resolve_client_ip;
17use crate::services::middleware::context::{
18 A2AContextMiddleware, McpContextMiddleware, PublicContextMiddleware, UserOnlyContextMiddleware,
19};
20use axum::Router;
21use axum::extract::{ConnectInfo, Request, State};
22use axum::http::{StatusCode, header};
23use axum::middleware::Next;
24use axum::response::{IntoResponse, Response};
25use chrono::{DateTime, Utc};
26use ipnet::IpNet;
27use std::future::Future;
28use std::net::SocketAddr;
29use std::sync::Arc;
30use systemprompt_extension::LoaderError;
31use systemprompt_models::auth::UserType;
32use systemprompt_models::config::RateLimitConfig;
33use systemprompt_models::{Config, RequestContext};
34use systemprompt_runtime::AppContext;
35use systemprompt_users::UserRateLimitBucketRepository;
36
37const GLOBAL_WINDOW_SECS: i64 = 10;
38
39#[derive(Clone, Debug)]
40pub struct RateLimitState {
41 config: RateLimitConfig,
42 trusted_proxies: Arc<Vec<IpNet>>,
43 buckets: Arc<UserRateLimitBucketRepository>,
44}
45
46impl RateLimitState {
47 #[must_use]
48 pub fn new(config: &Config, buckets: Arc<UserRateLimitBucketRepository>) -> Self {
49 Self {
50 config: config.rate_limits,
51 trusted_proxies: Arc::new(config.trusted_proxies.clone()),
52 buckets,
53 }
54 }
55
56 pub fn from_context(ctx: &AppContext) -> Result<Self, LoaderError> {
57 let buckets = crate::repository::user_rate_limit_buckets(ctx.db_pool()).map_err(|e| {
58 LoaderError::InitializationFailed {
59 extension: "rate_limit".to_owned(),
60 message: e.to_string(),
61 }
62 })?;
63 Ok(Self::new(ctx.config(), buckets))
64 }
65}
66
67#[derive(Clone, Debug)]
68struct GlobalUserLimit {
69 buckets: Arc<UserRateLimitBucketRepository>,
70 scope: &'static str,
71 budget: i64,
72}
73
74fn window_start(now: DateTime<Utc>) -> DateTime<Utc> {
75 let secs = now.timestamp();
76 let start = secs - secs.rem_euclid(GLOBAL_WINDOW_SECS);
77 DateTime::from_timestamp(start, 0).unwrap_or(now)
78}
79
80fn too_many_requests(now: DateTime<Utc>, start: DateTime<Utc>) -> Response {
81 let elapsed = now.timestamp() - start.timestamp();
82 let retry_after = (GLOBAL_WINDOW_SECS - elapsed).max(1);
83 (
84 StatusCode::TOO_MANY_REQUESTS,
85 [(header::RETRY_AFTER, retry_after.to_string())],
86 "rate limit exceeded",
87 )
88 .into_response()
89}
90
91async fn global_user_rate_limit(
92 State(limit): State<GlobalUserLimit>,
93 req: Request,
94 next: Next,
95) -> Response {
96 let user_id = req
97 .extensions()
98 .get::<RequestContext>()
99 .filter(|ctx| ctx.auth.user_type != UserType::Anon)
100 .map(|ctx| ctx.user_id().clone());
101 let Some(user_id) = user_id else {
102 return next.run(req).await;
103 };
104
105 let now = Utc::now();
106 let start = window_start(now);
107 match limit.buckets.hit(&user_id, limit.scope, start).await {
108 Ok(hits) if hits > limit.budget => {
109 tracing::debug!(
110 user_id = %user_id,
111 scope = limit.scope,
112 hits,
113 budget = limit.budget,
114 "global user rate limit exceeded"
115 );
116 too_many_requests(now, start)
117 },
118 Ok(_) => next.run(req).await,
119 Err(err) => {
120 tracing::warn!(
121 user_id = %user_id,
122 scope = limit.scope,
123 error = %err,
124 "global user rate limit unavailable; admitting request"
125 );
126 next.run(req).await
127 },
128 }
129}
130
131#[derive(Clone, Debug)]
132pub struct IdentityOrTrustedIpKey {
133 trusted_proxies: Arc<Vec<IpNet>>,
134}
135
136impl IdentityOrTrustedIpKey {
137 const fn new(trusted_proxies: Arc<Vec<IpNet>>) -> Self {
138 Self { trusted_proxies }
139 }
140}
141
142impl tower_governor::key_extractor::KeyExtractor for IdentityOrTrustedIpKey {
143 type Key = String;
144
145 fn extract<T>(&self, req: &Request<T>) -> Result<Self::Key, tower_governor::GovernorError> {
146 if let Some(ctx) = req.extensions().get::<RequestContext>()
147 && ctx.auth.user_type != UserType::Anon
148 {
149 return Ok(format!("u:{}", ctx.user_id()));
150 }
151
152 resolve_client_ip(
153 req.headers(),
154 req.extensions().get::<ConnectInfo<SocketAddr>>(),
155 &self.trusted_proxies,
156 )
157 .map(|ip| format!("ip:{ip}"))
158 .ok_or(tower_governor::GovernorError::UnableToExtractKey)
159 }
160}
161
162pub trait ContextLayer: Clone + Send + Sync + 'static {
163 fn handle(self, req: Request, next: Next) -> impl Future<Output = Response> + Send;
164}
165
166impl ContextLayer for PublicContextMiddleware {
167 async fn handle(self, req: Request, next: Next) -> Response {
168 Self::handle(&self, req, next).await
169 }
170}
171
172impl ContextLayer for UserOnlyContextMiddleware {
173 async fn handle(self, req: Request, next: Next) -> Response {
174 Self::handle(&self, req, next).await
175 }
176}
177
178impl ContextLayer for A2AContextMiddleware {
179 async fn handle(self, req: Request, next: Next) -> Response {
180 Self::handle(&self, req, next).await
181 }
182}
183
184impl ContextLayer for McpContextMiddleware {
185 async fn handle(self, req: Request, next: Next) -> Response {
186 Self::handle(&self, req, next).await
187 }
188}
189
190pub trait RouterExt<S>: Sized {
191 fn with_rate_limit(
192 self,
193 limits: &RateLimitState,
194 per_second: u64,
195 scope: &'static str,
196 ) -> Result<Self, LoaderError>;
197
198 fn with_auth<L: ContextLayer>(self, auth: L, policy: AuthzPolicy) -> Self;
199}
200
201impl<S> RouterExt<S> for Router<S>
202where
203 S: Clone + Send + Sync + 'static,
204{
205 fn with_rate_limit(
206 self,
207 limits: &RateLimitState,
208 per_second: u64,
209 scope: &'static str,
210 ) -> Result<Self, LoaderError> {
211 let rate_config = &limits.config;
212 if rate_config.disabled {
213 return Ok(self);
214 }
215
216 let burst = per_second.saturating_mul(rate_config.burst_multiplier);
217 let burst_u32 = u32::try_from(burst).unwrap_or(u32::MAX).max(1);
218 let per_second_clamped = per_second.max(1);
219
220 let rate_limit = tower_governor::governor::GovernorConfigBuilder::default()
221 .per_second(per_second_clamped)
222 .burst_size(burst_u32)
223 .key_extractor(IdentityOrTrustedIpKey::new(Arc::clone(
224 &limits.trusted_proxies,
225 )))
226 .use_headers()
227 .finish()
228 .ok_or_else(|| LoaderError::InitializationFailed {
229 extension: "rate_limit".to_owned(),
230 message: format!(
231 "rate limit rejected for {per_second_clamped}/s with burst {burst_u32}"
232 ),
233 })?;
234
235 let window_secs = u64::try_from(GLOBAL_WINDOW_SECS).unwrap_or(u64::MAX);
236 let budget = burst.saturating_mul(window_secs);
237 let global = GlobalUserLimit {
238 buckets: Arc::clone(&limits.buckets),
239 scope,
240 budget: i64::try_from(budget).unwrap_or(i64::MAX),
241 };
242
243 Ok(self
244 .layer(axum::middleware::from_fn_with_state(
245 global,
246 global_user_rate_limit,
247 ))
248 .layer(tower_governor::GovernorLayer::new(rate_limit)))
249 }
250
251 fn with_auth<L: ContextLayer>(self, auth: L, policy: AuthzPolicy) -> Self {
252 self.layer(axum::middleware::from_fn(move |req, next| async move {
253 authz_gate(policy, req, next).await
254 }))
255 .layer(axum::middleware::from_fn(move |req, next| {
256 let auth = auth.clone();
257 async move { auth.handle(req, next).await }
258 }))
259 }
260}