diff --git a/signal-gateway/src/gateway/mod.rs b/signal-gateway/src/gateway/mod.rs index 83c882f..5b15310 100644 --- a/signal-gateway/src/gateway/mod.rs +++ b/signal-gateway/src/gateway/mod.rs @@ -37,10 +37,13 @@ mod log_buffer; mod log_handler; use log_handler::{LogHandler, LogHandlerConfig}; +mod rate_limiter_set; +pub use rate_limiter_set::{LimitResult, LimiterSet}; + mod route; pub use route::{Destination, Limit, Route}; -pub use crate::rate_limiter::{LimitResult, Limiter, LimiterSet, RateThreshold}; +pub use crate::rate_limiter::{Limiter, RateThreshold}; /// Configuration for the gateway. #[derive(Conf, Debug)] diff --git a/signal-gateway/src/gateway/rate_limiter_set.rs b/signal-gateway/src/gateway/rate_limiter_set.rs new file mode 100644 index 0000000..81fe387 --- /dev/null +++ b/signal-gateway/src/gateway/rate_limiter_set.rs @@ -0,0 +1,76 @@ +//! Rate limiter set for managing per-route rate limiting. + +use crate::log_message::{LogMessage, Origin}; +use crate::rate_limiter::Limiter; +use std::collections::HashMap; + +/// Result of evaluating a limiter set. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum LimitResult { + /// The event passed all limits (not rate-limited). + Passed, + /// The event was blocked by a per-origin limiter at the given index. + Limiter(usize), + /// The event was blocked by a global limiter at the given index. + GlobalLimiter(usize), +} + +/// A set of limiters for a route, containing both per-origin and global limiters. +pub struct LimiterSet { + /// Factory to create limiters for new origins. + make_limiters: Box Vec + Send + Sync>, + /// Per-origin rate limiters, keyed by origin. Lazily created. + limiters: HashMap>, + /// Global rate limiters (shared across all origins). + global_limiters: Vec, +} + +impl std::fmt::Debug for LimiterSet { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("LimiterSet") + .field("limiters", &self.limiters) + .field("global_limiters", &self.global_limiters) + .finish_non_exhaustive() + } +} + +impl LimiterSet { + /// Create a new limiter set with factories for per-origin and global limiters. + pub fn new( + make_limiters: impl Fn() -> Vec + Send + Sync + 'static, + global_limiters: Vec, + ) -> Self { + Self { + make_limiters: Box::new(make_limiters), + limiters: HashMap::new(), + global_limiters, + } + } + + /// Evaluate whether an event should pass all rate limits. + /// + /// Returns [`LimitResult::Passed`] if the event passes all limits. + /// Returns [`LimitResult::Limiter(i)`] if blocked by per-origin limiter at index `i`. + /// Returns [`LimitResult::GlobalLimiter(i)`] if blocked by global limiter at index `i`. + pub fn evaluate(&mut self, log_msg: &LogMessage, origin: &Origin, ts_sec: i64) -> LimitResult { + // Get or create limiters for this origin + let origin_limiters = self + .limiters + .entry(origin.clone()) + .or_insert_with(&self.make_limiters); + + for (i, limiter) in origin_limiters.iter_mut().enumerate() { + if !limiter.evaluate(log_msg, ts_sec) { + return LimitResult::Limiter(i); + } + } + + for (i, limiter) in self.global_limiters.iter_mut().enumerate() { + if !limiter.evaluate(log_msg, ts_sec) { + return LimitResult::GlobalLimiter(i); + } + } + + LimitResult::Passed + } +} diff --git a/signal-gateway/src/gateway/route.rs b/signal-gateway/src/gateway/route.rs index be97b6c..d9c097d 100644 --- a/signal-gateway/src/gateway/route.rs +++ b/signal-gateway/src/gateway/route.rs @@ -3,8 +3,9 @@ //! Routes define how log messages are processed based on filters, severity levels, //! and destination overrides. +use super::rate_limiter_set::LimiterSet; use crate::log_message::{Level, LogFilter}; -use crate::rate_limiter::{Limiter, LimiterSet, RateThreshold}; +use crate::rate_limiter::{Limiter, RateThreshold}; use serde::Deserialize; /// A rate limit rule for suppressing repeated alerts. diff --git a/signal-gateway/src/rate_limiter.rs b/signal-gateway/src/rate_limiter.rs index ef7d3b2..daa7141 100644 --- a/signal-gateway/src/rate_limiter.rs +++ b/signal-gateway/src/rate_limiter.rs @@ -1,6 +1,6 @@ //! Rate limiting for log alerts. -use crate::log_message::{LogMessage, Origin}; +use crate::log_message::LogMessage; use serde::Deserialize; use std::{ collections::HashMap, @@ -87,77 +87,6 @@ pub enum Limiter { SourceLocation(SourceLocationRateLimiter), } -/// Result of evaluating a limiter set. -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum LimitResult { - /// The event passed all limits (not rate-limited). - Passed, - /// The event was blocked by a per-origin limiter at the given index. - Limiter(usize), - /// The event was blocked by a global limiter at the given index. - GlobalLimiter(usize), -} - -/// A set of limiters for a route, containing both per-origin and global limiters. -pub struct LimiterSet { - /// Factory to create limiters for new origins. - make_limiters: Box Vec + Send + Sync>, - /// Per-origin rate limiters, keyed by origin. Lazily created. - limiters: HashMap>, - /// Global rate limiters (shared across all origins). - global_limiters: Vec, -} - -impl std::fmt::Debug for LimiterSet { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - f.debug_struct("LimiterSet") - .field("limiters", &self.limiters) - .field("global_limiters", &self.global_limiters) - .finish_non_exhaustive() - } -} - -impl LimiterSet { - /// Create a new limiter set with factories for per-origin and global limiters. - pub fn new( - make_limiters: impl Fn() -> Vec + Send + Sync + 'static, - global_limiters: Vec, - ) -> Self { - Self { - make_limiters: Box::new(make_limiters), - limiters: HashMap::new(), - global_limiters, - } - } - - /// Evaluate whether an event should pass all rate limits. - /// - /// Returns [`LimitResult::Passed`] if the event passes all limits. - /// Returns [`LimitResult::Limiter(i)`] if blocked by per-origin limiter at index `i`. - /// Returns [`LimitResult::GlobalLimiter(i)`] if blocked by global limiter at index `i`. - pub fn evaluate(&mut self, log_msg: &LogMessage, origin: &Origin, ts_sec: i64) -> LimitResult { - // Get or create limiters for this origin - let origin_limiters = self - .limiters - .entry(origin.clone()) - .or_insert_with(&self.make_limiters); - - for (i, limiter) in origin_limiters.iter_mut().enumerate() { - if !limiter.evaluate(log_msg, ts_sec) { - return LimitResult::Limiter(i); - } - } - - for (i, limiter) in self.global_limiters.iter_mut().enumerate() { - if !limiter.evaluate(log_msg, ts_sec) { - return LimitResult::GlobalLimiter(i); - } - } - - LimitResult::Passed - } -} - impl Limiter { /// Evaluate whether an event should pass the rate limit. ///