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_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 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}