diff --git a/Cargo.toml b/Cargo.toml
index aa71fd2..b4a0de0 100644
--- a/Cargo.toml
+++ b/Cargo.toml
@@ -40,7 +40,7 @@ nats-transport = ["dep:async-nats", "dep:tokio"]
rabbitmq-transport = ["dep:lapin", "dep:tokio"]
redis-transport = ["dep:redis", "dep:tokio"]
schedule = ["dep:chrono", "dep:cron", "dep:tokio"]
-security = []
+security = ["dep:sha2"]
session = ["dep:getrandom"]
shutdown-hooks = ["dep:tokio", "tokio/macros", "tokio/signal"]
static = ["dep:tokio"]
@@ -73,6 +73,7 @@ schemars = { version = "1", optional = true }
serde = { version = "1", features = ["derive"] }
serde_json = "1"
serde_urlencoded = "0.7"
+sha2 = { version = "0.10", optional = true }
thiserror = "2"
tokio = { version = "1", features = ["fs", "io-util", "net", "rt", "sync", "time"], optional = true }
tonic = { version = "0.14.6", default-features = false, features = ["codegen", "transport"], optional = true }
diff --git a/README.md b/README.md
index fcccd8a..e08fff5 100644
--- a/README.md
+++ b/README.md
@@ -113,7 +113,7 @@ are opt-in.
| Lifecycle | `shutdown-hooks` | SIGINT and SIGTERM shutdown handling |
| Configuration | `config` | ACL-backed typed configuration parsing |
| Authentication | `auth` | Strategy-backed authentication guards |
-| Security | `security` | CORS, CSRF, rate limiting, and security headers |
+| Security | `security` | CORS, CSRF, local or provider-backed rate limiting, and security headers |
| Sessions | `session` | Session middleware and replaceable stores |
| Cache | `cache` | Cache abstraction, interceptor, and in-memory store |
| Database | `database` | Replaceable database facade and in-memory backend |
@@ -284,6 +284,23 @@ from `schemars::JsonSchema` types.
Optional HTTP modules add multipart uploads, static files, gzip compression,
views, sessions, security policies, request context, and an outbound HTTP client.
+### Provider-backed rate limiting
+
+The `security` feature keeps `use_global_rate_limit` process-local by default.
+Applications that need one budget across multiple processes can implement the
+public `RateLimitProvider` contract and register it with
+`use_global_rate_limit_provider`. Each atomic acquisition receives a stable
+policy identifier, a policy-scoped SHA-256 subject digest, and the configured
+request limit and window. Selected header values and bearer credentials do not
+cross the provider boundary in plaintext.
+
+Every process using the same policy identifier must use identical limits and
+windows. Provider errors reject guarded requests instead of bypassing the
+limit. Boot deliberately does not select or bundle a distributed backend; the
+built-in `InMemoryRateLimitProvider` does not share state between processes.
+This boundary does not cover the separate streaming-disconnect, backpressure,
+or graceful-drain work.
+
## Protocols
### WebSocket
diff --git a/ROADMAP.md b/ROADMAP.md
index 27a6c82..6a05d2c 100644
--- a/ROADMAP.md
+++ b/ROADMAP.md
@@ -956,7 +956,8 @@ Nest equivalent areas:
- streamable file and download responses (implemented)
- file upload (implemented)
- static assets and SPA shells (implemented)
-- security helpers such as CORS, CSRF, helmet-like headers, and rate limiting
+- security helpers such as CORS, CSRF, helmet-like headers, process-local rate
+ limiting, and a provider-neutral atomic boundary for shared rate limits
(implemented)
- sessions (implemented)
@@ -1042,7 +1043,9 @@ Acceptance:
hidden dotfile defaults, and root traversal protection. (Covered)
- Security helpers can handle CORS preflight and actual response headers, add
helmet-like response headers, reject invalid CSRF tokens on unsafe methods,
- and enforce in-memory fixed-window rate limits. (Covered)
+ enforce in-memory fixed-window rate limits, and delegate atomic acquisitions
+ to an application-supplied shared provider without exposing plaintext
+ credentials. Boot does not select a distributed backend. (Covered)
- Sessions can register a provider-backed `SessionManager`, expose
request-bound `Session` handles through `BootRequest::session()` and
Nest-style `#[session]` arguments, bind session ids before handlers, persist
diff --git a/src/app/builder.rs b/src/app/builder.rs
index 3cbe2cc..b057f41 100644
--- a/src/app/builder.rs
+++ b/src/app/builder.rs
@@ -18,7 +18,7 @@ use crate::{CompressionInterceptor, CompressionOptions};
#[cfg(feature = "security")]
use crate::{
CorsMiddleware, CorsOptions, CorsPreflightRoute, CorsResponseInterceptor, CsrfGuard,
- CsrfOptions, RateLimitGuard, RateLimitOptions, SecurityHeadersInterceptor,
+ CsrfOptions, RateLimitGuard, RateLimitOptions, RateLimitProvider, SecurityHeadersInterceptor,
SecurityHeadersOptions,
};
#[cfg(feature = "session")]
@@ -315,6 +315,21 @@ impl BootApplicationBuilder {
self
}
+ /// Add an application-wide rate limit guard backed by a shared provider.
+ #[cfg(feature = "security")]
+ pub fn use_global_rate_limit_provider
(
+ mut self,
+ options: RateLimitOptions,
+ provider: P,
+ ) -> Self
+ where
+ P: RateLimitProvider,
+ {
+ self.global_pipeline
+ .push_guard(RateLimitGuard::with_provider(options, provider));
+ self
+ }
+
/// Add global session middleware and cookie persistence.
#[cfg(feature = "session")]
pub fn use_global_sessions(mut self, manager: SessionManager) -> Self {
diff --git a/src/lib.rs b/src/lib.rs
index 0626f7c..a3877e7 100644
--- a/src/lib.rs
+++ b/src/lib.rs
@@ -198,8 +198,8 @@ pub use schedule::{
#[cfg(feature = "security")]
pub use security::{
CorsMiddleware, CorsOptions, CorsPreflightRoute, CorsResponseInterceptor, CsrfGuard,
- CsrfOptions, RateLimitGuard, RateLimitOptions, SecurityHeadersInterceptor,
- SecurityHeadersOptions,
+ CsrfOptions, InMemoryRateLimitProvider, RateLimitDecision, RateLimitGuard, RateLimitOptions,
+ RateLimitProvider, RateLimitRequest, SecurityHeadersInterceptor, SecurityHeadersOptions,
};
pub use serialization::{SerializationInterceptor, SerializationOptions};
#[cfg(feature = "session")]
diff --git a/src/security/mod.rs b/src/security/mod.rs
index f8d6e9a..0bd82f1 100644
--- a/src/security/mod.rs
+++ b/src/security/mod.rs
@@ -7,4 +7,7 @@ mod rate_limit;
pub use cors::{CorsMiddleware, CorsOptions, CorsPreflightRoute, CorsResponseInterceptor};
pub use csrf::{CsrfGuard, CsrfOptions};
pub use headers::{SecurityHeadersInterceptor, SecurityHeadersOptions};
-pub use rate_limit::{RateLimitGuard, RateLimitOptions};
+pub use rate_limit::{
+ InMemoryRateLimitProvider, RateLimitDecision, RateLimitGuard, RateLimitOptions,
+ RateLimitProvider, RateLimitRequest,
+};
diff --git a/src/security/rate_limit.rs b/src/security/rate_limit.rs
index b6a30d4..70d9737 100644
--- a/src/security/rate_limit.rs
+++ b/src/security/rate_limit.rs
@@ -1,11 +1,17 @@
use crate::{BootError, BootRequest, BoxFuture, ExecutionContext, Guard, Result};
+use sha2::{Digest, Sha256};
use std::collections::BTreeMap;
+use std::fmt;
use std::sync::{Arc, Mutex};
use std::time::{Duration, Instant};
-/// In-memory rate limit settings.
+const DEFAULT_POLICY_ID: &str = "global";
+const MAX_POLICY_ID_BYTES: usize = 128;
+
+/// Application-wide rate limit settings.
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct RateLimitOptions {
+ policy_id: String,
max_requests: u32,
window: Duration,
key_headers: Vec,
@@ -16,6 +22,7 @@ pub struct RateLimitOptions {
impl Default for RateLimitOptions {
fn default() -> Self {
Self {
+ policy_id: DEFAULT_POLICY_ID.to_string(),
max_requests: 60,
window: Duration::from_secs(60),
key_headers: vec!["x-forwarded-for".to_string(), "x-real-ip".to_string()],
@@ -30,6 +37,12 @@ impl RateLimitOptions {
Self::default()
}
+ /// Set a stable policy identifier shared by every process using the same provider policy.
+ pub fn with_policy_id(mut self, policy_id: impl Into) -> Self {
+ self.policy_id = policy_id.into();
+ self
+ }
+
pub fn with_max_requests(mut self, max_requests: u32) -> Self {
self.max_requests = max_requests;
self
@@ -64,6 +77,10 @@ impl RateLimitOptions {
self
}
+ pub fn policy_id(&self) -> &str {
+ &self.policy_id
+ }
+
pub fn max_requests(&self) -> u32 {
self.max_requests
}
@@ -71,13 +88,108 @@ impl RateLimitOptions {
pub fn window(&self) -> Duration {
self.window
}
+
+ fn validate(&self) -> Result<()> {
+ if self.max_requests == 0 {
+ return Err(BootError::Internal(
+ "rate limit max_requests must be greater than zero".to_string(),
+ ));
+ }
+ if self.window.is_zero() {
+ return Err(BootError::Internal(
+ "rate limit window must be greater than zero".to_string(),
+ ));
+ }
+ if !valid_policy_id(&self.policy_id) {
+ return Err(BootError::Internal(format!(
+ "rate limit policy_id must be 1 to {MAX_POLICY_ID_BYTES} ASCII identifier bytes"
+ )));
+ }
+ Ok(())
+ }
}
-/// Guard that enforces an in-memory fixed-window rate limit.
-#[derive(Debug, Clone)]
-pub struct RateLimitGuard {
- options: RateLimitOptions,
- buckets: Arc>>,
+/// One atomic request to a [`RateLimitProvider`].
+///
+/// The subject is a policy-scoped, domain-separated SHA-256 digest. Header values and bearer
+/// credentials never cross the provider boundary in plaintext.
+#[derive(Debug, Clone, PartialEq, Eq)]
+pub struct RateLimitRequest {
+ policy_id: String,
+ subject_hash: String,
+ max_requests: u32,
+ window: Duration,
+}
+
+impl RateLimitRequest {
+ pub fn policy_id(&self) -> &str {
+ &self.policy_id
+ }
+
+ pub fn subject_hash(&self) -> &str {
+ &self.subject_hash
+ }
+
+ pub fn max_requests(&self) -> u32 {
+ self.max_requests
+ }
+
+ pub fn window(&self) -> Duration {
+ self.window
+ }
+}
+
+/// Result of an atomic rate limit acquisition.
+#[derive(Debug, Clone, Copy, PartialEq, Eq)]
+pub enum RateLimitDecision {
+ /// The provider consumed one request from the policy budget.
+ Allowed,
+ /// The provider rejected the request because its policy budget is exhausted.
+ Limited,
+}
+
+impl RateLimitDecision {
+ pub fn is_allowed(self) -> bool {
+ matches!(self, Self::Allowed)
+ }
+}
+
+/// Provider-neutral atomic rate limit boundary.
+///
+/// A distributed implementation can use Redis, PostgreSQL, or another shared service without
+/// exposing that backend to Boot. Implementations must atomically consume one request for the
+/// `(policy_id, subject_hash)` pair. Every client of a stable policy identifier must use the same
+/// request limit and window. Providers must reject conflicting settings, and returning any error
+/// rejects the guarded request.
+pub trait RateLimitProvider: Send + Sync + 'static {
+ fn acquire(&self, request: RateLimitRequest) -> BoxFuture<'static, Result>;
+}
+
+impl RateLimitProvider for Arc
+where
+ T: RateLimitProvider + ?Sized,
+{
+ fn acquire(&self, request: RateLimitRequest) -> BoxFuture<'static, Result> {
+ self.as_ref().acquire(request)
+ }
+}
+
+/// Process-local fixed-window provider used by [`RateLimitGuard::with_options`].
+#[derive(Debug, Clone, Default)]
+pub struct InMemoryRateLimitProvider {
+ state: Arc>,
+}
+
+#[derive(Debug, Default)]
+struct InMemoryRateLimitState {
+ policies: BTreeMap,
+ buckets: BTreeMap<(String, String), RateLimitBucket>,
+}
+
+#[derive(Debug, Clone, Copy, PartialEq, Eq)]
+struct RateLimitPolicy {
+ max_requests: u32,
+ window: Duration,
}
#[derive(Debug, Clone)]
@@ -86,12 +198,82 @@ struct RateLimitBucket {
count: u32,
}
+impl InMemoryRateLimitProvider {
+ pub fn new() -> Self {
+ Self::default()
+ }
+}
+
+impl RateLimitProvider for InMemoryRateLimitProvider {
+ fn acquire(&self, request: RateLimitRequest) -> BoxFuture<'static, Result> {
+ let state = Arc::clone(&self.state);
+ Box::pin(async move {
+ let now = Instant::now();
+ let mut state = state.lock().map_err(|_| {
+ BootError::Internal("rate limit state lock is poisoned".to_string())
+ })?;
+ let requested_policy = RateLimitPolicy {
+ max_requests: request.max_requests,
+ window: request.window,
+ };
+ if let Some(policy) = state.policies.get(&request.policy_id) {
+ if *policy != requested_policy {
+ return Err(BootError::Internal(
+ "rate limit policy settings conflict across provider clients".to_string(),
+ ));
+ }
+ } else {
+ state
+ .policies
+ .insert(request.policy_id.clone(), requested_policy);
+ }
+
+ let InMemoryRateLimitState { policies, buckets } = &mut *state;
+ buckets.retain(|(policy_id, _), bucket| {
+ policies.get(policy_id).is_some_and(|policy| {
+ now.duration_since(bucket.window_started_at) < policy.window
+ })
+ });
+
+ let key = (request.policy_id.clone(), request.subject_hash.clone());
+ let bucket = buckets.entry(key).or_insert_with(|| RateLimitBucket {
+ window_started_at: now,
+ count: 0,
+ });
+ let elapsed = now.duration_since(bucket.window_started_at);
+ if elapsed >= requested_policy.window {
+ bucket.window_started_at = now;
+ bucket.count = 0;
+ }
+ if bucket.count >= requested_policy.max_requests {
+ return Ok(RateLimitDecision::Limited);
+ }
+
+ bucket.count += 1;
+ Ok(RateLimitDecision::Allowed)
+ })
+ }
+}
+
+/// Guard that enforces a fixed-window policy through a [`RateLimitProvider`].
+#[derive(Clone)]
+pub struct RateLimitGuard {
+ options: RateLimitOptions,
+ provider: Arc,
+}
+
+impl fmt::Debug for RateLimitGuard {
+ fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
+ formatter
+ .debug_struct("RateLimitGuard")
+ .field("options", &self.options)
+ .finish_non_exhaustive()
+ }
+}
+
impl Default for RateLimitGuard {
fn default() -> Self {
- Self {
- options: RateLimitOptions::default(),
- buckets: Arc::new(Mutex::new(BTreeMap::new())),
- }
+ Self::with_options(RateLimitOptions::default())
}
}
@@ -100,10 +282,19 @@ impl RateLimitGuard {
Self::default()
}
+ /// Construct a guard with the process-local provider.
pub fn with_options(options: RateLimitOptions) -> Self {
+ Self::with_provider(options, InMemoryRateLimitProvider::new())
+ }
+
+ /// Construct a guard with an application-supplied provider.
+ pub fn with_provider(options: RateLimitOptions, provider: P) -> Self
+ where
+ P: RateLimitProvider,
+ {
Self {
options,
- buckets: Arc::new(Mutex::new(BTreeMap::new())),
+ provider: Arc::new(provider),
}
}
@@ -115,39 +306,32 @@ impl RateLimitGuard {
impl Guard for RateLimitGuard {
fn can_activate(&self, context: ExecutionContext) -> BoxFuture<'static, Result> {
let options = self.options.clone();
- let buckets = Arc::clone(&self.buckets);
+ let provider = Arc::clone(&self.provider);
Box::pin(async move {
- let key = rate_limit_key(&context.request, &options);
- let now = Instant::now();
- let mut buckets = buckets.lock().map_err(|_| {
- BootError::Internal("rate limit state lock is poisoned".to_string())
- })?;
- buckets
- .retain(|_, bucket| now.duration_since(bucket.window_started_at) < options.window);
-
- let bucket = buckets.entry(key).or_insert_with(|| RateLimitBucket {
- window_started_at: now,
- count: 0,
- });
-
- if now.duration_since(bucket.window_started_at) >= options.window {
- bucket.window_started_at = now;
- bucket.count = 0;
- }
-
- if bucket.count >= options.max_requests {
- return Err(BootError::TooManyRequests(
+ options.validate()?;
+ let request = rate_limit_request(&context.request, &options);
+ let decision = provider.acquire(request).await?;
+ if decision.is_allowed() {
+ Ok(true)
+ } else {
+ Err(BootError::TooManyRequests(
"rate limit exceeded".to_string(),
- ));
+ ))
}
-
- bucket.count += 1;
- Ok(true)
})
}
}
-fn rate_limit_key(request: &BootRequest, options: &RateLimitOptions) -> String {
+fn rate_limit_request(request: &BootRequest, options: &RateLimitOptions) -> RateLimitRequest {
+ RateLimitRequest {
+ policy_id: options.policy_id.clone(),
+ subject_hash: rate_limit_subject_hash(request, options),
+ max_requests: options.max_requests,
+ window: options.window,
+ }
+}
+
+fn rate_limit_subject_hash(request: &BootRequest, options: &RateLimitOptions) -> String {
for header in &options.key_headers {
if let Some(value) = request
.header(header)
@@ -155,15 +339,88 @@ fn rate_limit_key(request: &BootRequest, options: &RateLimitOptions) -> String {
.map(str::trim)
.filter(|value| !value.is_empty())
{
- return format!("header:{header}:{value}");
+ return hash_subject(&[
+ options.policy_id(),
+ "header",
+ &header.to_ascii_lowercase(),
+ value,
+ ]);
}
}
if options.use_bearer_token {
if let Some(token) = request.bearer_token() {
- return format!("bearer:{token}");
+ return hash_subject(&[options.policy_id(), "bearer", token]);
}
}
- options.anonymous_key.clone()
+ hash_subject(&[options.policy_id(), "anonymous", &options.anonymous_key])
+}
+
+fn hash_subject(parts: &[&str]) -> String {
+ let mut digest = Sha256::new();
+ digest.update(b"a3s-boot-rate-limit-subject-v1\0");
+ for part in parts {
+ digest.update((part.len() as u64).to_be_bytes());
+ digest.update(part.as_bytes());
+ }
+ format!("{:x}", digest.finalize())
+}
+
+fn valid_policy_id(value: &str) -> bool {
+ !value.is_empty()
+ && value.len() <= MAX_POLICY_ID_BYTES
+ && value.bytes().all(|byte| {
+ byte.is_ascii_alphanumeric() || matches!(byte, b'.' | b'_' | b'-' | b':' | b'/')
+ })
+}
+
+#[cfg(test)]
+mod tests {
+ use super::*;
+
+ #[test]
+ fn subject_hash_is_stable_and_does_not_embed_input() {
+ let options = RateLimitOptions::new().with_key_header("x-user-id");
+ let request = BootRequest::new(crate::HttpMethod::Get, "/")
+ .with_header("x-user-id", "tenant-user-secret-material");
+ let first = rate_limit_subject_hash(&request, &options);
+ let second = rate_limit_subject_hash(&request, &options);
+
+ assert_eq!(first, second);
+ assert_eq!(first.len(), 64);
+ assert!(!first.contains("tenant-user-secret-material"));
+ }
+
+ #[test]
+ fn subject_hash_is_scoped_to_the_policy() {
+ let request = BootRequest::new(crate::HttpMethod::Get, "/")
+ .with_header("authorization", "Bearer tenant-secret-token");
+ let first = rate_limit_subject_hash(
+ &request,
+ &RateLimitOptions::new().with_policy_id("public-api"),
+ );
+ let second = rate_limit_subject_hash(
+ &request,
+ &RateLimitOptions::new().with_policy_id("admin-api"),
+ );
+
+ assert_ne!(first, second);
+ }
+
+ #[test]
+ fn invalid_policy_settings_fail_before_provider_use() {
+ assert!(RateLimitOptions::new()
+ .with_policy_id("invalid policy")
+ .validate()
+ .is_err());
+ assert!(RateLimitOptions::new()
+ .with_max_requests(0)
+ .validate()
+ .is_err());
+ assert!(RateLimitOptions::new()
+ .with_window(Duration::ZERO)
+ .validate()
+ .is_err());
+ }
}
diff --git a/tests/security.rs b/tests/security.rs
index 021c7d4..3526d0b 100644
--- a/tests/security.rs
+++ b/tests/security.rs
@@ -1,10 +1,14 @@
#![cfg(feature = "security")]
use a3s_boot::{
- BootApplication, BootRequest, BootResponse, CorsOptions, CsrfOptions, HttpMethod,
- RateLimitOptions, RouteDefinition, SecurityHeadersOptions,
+ BootApplication, BootError, BootRequest, BootResponse, BoxFuture, CorsOptions, CsrfOptions,
+ HttpMethod, InMemoryRateLimitProvider, RateLimitDecision, RateLimitOptions, RateLimitProvider,
+ RateLimitRequest, Result, RouteDefinition, SecurityHeadersOptions,
};
use serde_json::json;
+use std::collections::BTreeMap;
+use std::sync::atomic::{AtomicUsize, Ordering};
+use std::sync::{Arc, Mutex};
use std::time::Duration;
#[tokio::test]
@@ -213,3 +217,140 @@ async fn rate_limit_guard_rejects_requests_after_the_window_limit() {
);
assert_eq!(separate_key.status(), 200);
}
+
+#[derive(Clone, Default)]
+struct SharedRateLimitProvider {
+ counts: Arc>>,
+ requests: Arc>>,
+}
+
+impl RateLimitProvider for SharedRateLimitProvider {
+ fn acquire(&self, request: RateLimitRequest) -> BoxFuture<'static, Result> {
+ let counts = Arc::clone(&self.counts);
+ let requests = Arc::clone(&self.requests);
+ Box::pin(async move {
+ requests.lock().unwrap().push(request.clone());
+ let key = (
+ request.policy_id().to_string(),
+ request.subject_hash().to_string(),
+ );
+ let mut counts = counts.lock().unwrap();
+ let count = counts.entry(key).or_default();
+ if *count >= request.max_requests() {
+ return Ok(RateLimitDecision::Limited);
+ }
+ *count += 1;
+ Ok(RateLimitDecision::Allowed)
+ })
+ }
+}
+
+fn limited_app(provider: P, policy_id: &str, max_requests: u32) -> BootApplication
+where
+ P: RateLimitProvider,
+{
+ BootApplication::builder()
+ .use_global_rate_limit_provider(
+ RateLimitOptions::new()
+ .with_policy_id(policy_id)
+ .with_max_requests(max_requests)
+ .with_window(Duration::from_secs(60)),
+ provider,
+ )
+ .route(
+ RouteDefinition::get("/limited", |_| async { Ok(BootResponse::text("ok")) }).unwrap(),
+ )
+ .build()
+ .unwrap()
+}
+
+#[tokio::test]
+async fn public_rate_limit_provider_shares_state_without_receiving_credentials() {
+ let provider = SharedRateLimitProvider::default();
+ let first_process = limited_app(provider.clone(), "cloud-api", 1);
+ let second_process = limited_app(provider.clone(), "cloud-api", 1);
+ let request = || {
+ BootRequest::new(HttpMethod::Get, "/limited")
+ .with_header("authorization", "Bearer tenant-secret-token")
+ };
+
+ assert_eq!(first_process.call(request()).await.unwrap().status(), 200);
+ assert_eq!(second_process.handle(request()).await.status(), 429);
+
+ let requests = provider.requests.lock().unwrap();
+ assert_eq!(requests.len(), 2);
+ assert!(requests
+ .iter()
+ .all(|request| request.policy_id() == "cloud-api"));
+ assert!(requests
+ .iter()
+ .all(|request| request.subject_hash().len() == 64));
+ assert!(requests
+ .iter()
+ .all(|request| !request.subject_hash().contains("tenant-secret-token")));
+ assert_eq!(requests[0].subject_hash(), requests[1].subject_hash());
+}
+
+#[tokio::test]
+async fn in_memory_provider_rejects_conflicting_clients_for_one_policy() {
+ let provider = InMemoryRateLimitProvider::new();
+ let first_client = limited_app(provider.clone(), "cloud-api", 1);
+ let conflicting_client = limited_app(provider, "cloud-api", 2);
+
+ assert_eq!(
+ first_client
+ .call(BootRequest::new(HttpMethod::Get, "/limited"))
+ .await
+ .unwrap()
+ .status(),
+ 200
+ );
+
+ let error = conflicting_client
+ .call(BootRequest::new(HttpMethod::Get, "/limited"))
+ .await
+ .unwrap_err();
+ assert!(matches!(error, BootError::Internal(_)));
+}
+
+#[derive(Clone)]
+struct UnavailableRateLimitProvider;
+
+impl RateLimitProvider for UnavailableRateLimitProvider {
+ fn acquire(&self, _request: RateLimitRequest) -> BoxFuture<'static, Result> {
+ Box::pin(async {
+ Err(BootError::ServiceUnavailable(
+ "rate limit provider unavailable".to_string(),
+ ))
+ })
+ }
+}
+
+#[tokio::test]
+async fn provider_failure_rejects_work_instead_of_bypassing_the_limit() {
+ let handler_calls = Arc::new(AtomicUsize::new(0));
+ let handler_calls_for_route = Arc::clone(&handler_calls);
+ let app = BootApplication::builder()
+ .use_global_rate_limit_provider(
+ RateLimitOptions::new().with_policy_id("cloud-api"),
+ UnavailableRateLimitProvider,
+ )
+ .route(
+ RouteDefinition::get("/limited", move |_| {
+ let handler_calls = Arc::clone(&handler_calls_for_route);
+ async move {
+ handler_calls.fetch_add(1, Ordering::SeqCst);
+ Ok(BootResponse::text("ok"))
+ }
+ })
+ .unwrap(),
+ )
+ .build()
+ .unwrap();
+
+ let response = app
+ .handle(BootRequest::new(HttpMethod::Get, "/limited"))
+ .await;
+ assert_eq!(response.status(), 503);
+ assert_eq!(handler_calls.load(Ordering::SeqCst), 0);
+}