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); +}