From 7791af162c711e5dcb027650b0c31672e139c349 Mon Sep 17 00:00:00 2001 From: Chris Beck Date: Sat, 6 Dec 2025 19:24:39 -0700 Subject: [PATCH] cargo fmt, fix limiter set evaluate semantics --- signal-gateway-bin/src/main.rs | 13 +- signal-gateway/src/gateway/log_handler.rs | 131 ++++--- signal-gateway/src/gateway/rate_limiter.rs | 330 ------------------ .../src/gateway/rate_limiter_set.rs | 29 +- signal-gateway/src/gateway/route.rs | 14 +- .../src/gateway/signal_trust_set.rs | 5 +- 6 files changed, 107 insertions(+), 415 deletions(-) delete mode 100644 signal-gateway/src/gateway/rate_limiter.rs diff --git a/signal-gateway-bin/src/main.rs b/signal-gateway-bin/src/main.rs index 1e422f8..42efb6b 100644 --- a/signal-gateway-bin/src/main.rs +++ b/signal-gateway-bin/src/main.rs @@ -230,10 +230,7 @@ limits = [ .try_parse() .expect("Failed to parse config"); - assert_eq!( - config.http_listen_addr, - "0.0.0.0:8080".parse().unwrap() - ); + assert_eq!(config.http_listen_addr, "0.0.0.0:8080".parse().unwrap()); assert_eq!(config.gateway.signal_account, "+15551234567"); assert_eq!( config.gateway.signal_cli_tcp_addr, @@ -244,7 +241,13 @@ limits = [ Duration::from_secs(10) ); assert_eq!(config.gateway.admin_signal_uuids.len(), 2); - assert!(config.gateway.admin_signal_uuids.get("abc-123-uuid").is_some()); + assert!( + config + .gateway + .admin_signal_uuids + .get("abc-123-uuid") + .is_some() + ); let syslog = config.syslog.expect("syslog should be present"); assert_eq!(syslog.listen_addr, "0.0.0.0:1514".parse().unwrap()); diff --git a/signal-gateway/src/gateway/log_handler.rs b/signal-gateway/src/gateway/log_handler.rs index 01d2c69..08e7eab 100644 --- a/signal-gateway/src/gateway/log_handler.rs +++ b/signal-gateway/src/gateway/log_handler.rs @@ -6,7 +6,7 @@ use super::{ use crate::{ concurrent_map::LazyMap, log_format::LogFormatConfig, - log_message::{LogMessage, Origin}, + log_message::{LogFilter, LogMessage, Origin}, }; use chrono::Utc; use conf::Conf; @@ -22,8 +22,8 @@ pub enum SuppressionReason { /// Suppressed by route limiters. Contains the index and result for each /// route whose filter matched but whose limiter blocked the message. Routes(Vec<(usize, LimitResult)>), - /// Suppressed by an overall limiter. Contains the limiter index and result. - Overall(usize, LimitResult), + /// Suppressed by an overall limiter. + Overall(LimitResult), } impl fmt::Display for SuppressionReason { @@ -40,8 +40,8 @@ impl fmt::Display for SuppressionReason { } write!(f, "]") } - SuppressionReason::Overall(idx, result) => { - write!(f, "overall[{idx}]:{result:?}") + SuppressionReason::Overall(result) => { + write!(f, "overall:{result:?}") } } } @@ -82,7 +82,8 @@ pub struct LogHandler { /// Routes with their associated limiter sets. routes: Vec<(Route, LimiterSet)>, /// Overall rate limiters applied after route checks pass. - overall_limits: Vec, + /// Each entry is a (filter, limiter) pair. + overall_limits: Vec<(LogFilter, Limiter)>, } impl LogHandler { @@ -118,42 +119,41 @@ impl LogHandler { /// /// If `filter` is provided, only origins matching the filter are included. pub async fn format_logs(&self, filter: Option<&str>) -> String { - self.log_buffers - .with_read_lock(|buffers| { - if buffers.is_empty() { - return "No log sources registered yet".to_string(); + self.log_buffers.with_read_lock(|buffers| { + if buffers.is_empty() { + return "No log sources registered yet".to_string(); + } + + let mut text = String::with_capacity(4096); + let now = Utc::now(); + + for (origin, buffer) in buffers.iter() { + // Apply filter if present + if filter.is_some_and(|f| !origin.matches_filter(f)) { + continue; } - let mut text = String::with_capacity(4096); - let now = Utc::now(); - - for (origin, buffer) in buffers.iter() { - // Apply filter if present - if filter.is_some_and(|f| !origin.matches_filter(f)) { - continue; + use std::fmt::Write; + writeln!(&mut text, "=== [{origin}] ===").unwrap(); + buffer.with_iter(|iter| { + writeln!(&mut text, "{} log messages (newest first):", iter.len()).unwrap(); + // Guess at how much to reserve + text.reserve(iter.len() * 128); + for log_msg in iter { + self.config + .log_format + .write_log_msg(&mut text, log_msg, now); } + }); + text.push('\n'); + } - use std::fmt::Write; - writeln!(&mut text, "=== [{origin}] ===").unwrap(); - buffer.with_iter(|iter| { - writeln!(&mut text, "{} log messages (newest first):", iter.len()).unwrap(); - // Guess at how much to reserve - text.reserve(iter.len() * 128); - for log_msg in iter { - self.config - .log_format - .write_log_msg(&mut text, log_msg, now); - } - }); - text.push('\n'); - } - - if text.is_empty() { - "No matching log sources".to_string() - } else { - text - } - }) + if text.is_empty() { + "No matching log sources".to_string() + } else { + text + } + }) } /// Consume a new log message from the given origin @@ -170,33 +170,31 @@ impl LogHandler { } // Get or create the buffer for this origin, then record the message - let formatted_text = self - .log_buffers - .get(&origin, |buffer| { - if rate_limit_result.is_err() { - buffer.push_back(log_msg); - None - } else { - // Guess at capacity, it will be faster to use too much memory than too little - // signal-cli JVM is a hog anyways. - let mut text = String::with_capacity(4096); - let mut first_msg_len = 0; - let mut is_first = true; - let now = Utc::now(); + let formatted_text = self.log_buffers.get(&origin, |buffer| { + if rate_limit_result.is_err() { + buffer.push_back(log_msg); + None + } else { + // Guess at capacity, it will be faster to use too much memory than too little + // signal-cli JVM is a hog anyways. + let mut text = String::with_capacity(4096); + let mut first_msg_len = 0; + let mut is_first = true; + let now = Utc::now(); - buffer.push_back_and_drain(log_msg, |log_msg| { - self.config - .log_format - .write_log_msg(&mut text, log_msg, now); - if is_first { - first_msg_len = text.len(); - is_first = false; - } - }); + buffer.push_back_and_drain(log_msg, |log_msg| { + self.config + .log_format + .write_log_msg(&mut text, log_msg, now); + if is_first { + first_msg_len = text.len(); + is_first = false; + } + }); - Some((text, first_msg_len)) - } - }); + Some((text, first_msg_len)) + } + }); // Send alert if we have formatted text if let Some((text, first_msg_len)) = formatted_text { @@ -269,9 +267,10 @@ impl LogHandler { }; // At least one route passed, now check overall limits - for (idx, limiter) in self.overall_limits.iter().enumerate() { - if !limiter.evaluate(log_msg, ts_sec) { - return Err(SuppressionReason::Overall(idx, LimitResult::Limiter(0))); + for (idx, (filter, limiter)) in self.overall_limits.iter().enumerate() { + // Only evaluate the limiter if the message matches the filter + if filter.matches(log_msg) && !limiter.evaluate(log_msg, ts_sec) { + return Err(SuppressionReason::Overall(LimitResult::OverallLimiter(idx))); } } diff --git a/signal-gateway/src/gateway/rate_limiter.rs b/signal-gateway/src/gateway/rate_limiter.rs deleted file mode 100644 index 8c07beb..0000000 --- a/signal-gateway/src/gateway/rate_limiter.rs +++ /dev/null @@ -1,330 +0,0 @@ -use super::route::{Limit, RateThreshold}; -use crate::log_message::{LogMessage, Origin}; -use std::{ - collections::HashMap, - sync::atomic::{AtomicI64, Ordering}, - time::Duration, -}; - -/// Maximum entries in a source-location rate limiter before triggering cleanup. -const SOURCE_LOCATION_MAX_ENTRIES: usize = 2000; - -/// A rate limiter that can be either a multi-rate limiter or a source-location limiter. -#[derive(Debug)] -pub enum Limiter { - /// Counts events regardless of source location. - Multi(MultiRateLimiter), - /// Tracks events independently per source location (file:line). - SourceLocation(SourceLocationRateLimiter), -} - -/// A set of limiters for a route, containing both per-origin and global limiters. -#[derive(Debug)] -pub struct LimiterSet { - /// Limit configurations (used to create limiters for new origins). - limits: Vec, - /// Per-origin rate limiters, keyed by origin. Lazily created. - limiters: HashMap>, - /// Global rate limiters (shared across all origins). - global_limiters: Vec, -} - -/// Result of evaluating a log message against 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), -} - -impl LimiterSet { - /// Create a new limiter set from limit configurations. - pub fn new(limits: Vec, global_limits: Vec) -> Self { - Self { - limits, - limiters: HashMap::new(), - global_limiters: global_limits.iter().map(|l| l.make_limiter()).collect(), - } - } - - /// 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.limits.iter().map(|l| l.make_limiter()).collect()); - - 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. - /// - /// Returns `true` if the event should be allowed (not rate-limited), - /// `false` if it should be suppressed. - pub fn evaluate(&mut self, log_msg: &LogMessage, ts_sec: i64) -> bool { - match self { - Limiter::Multi(limiter) => limiter.evaluate(ts_sec), - Limiter::SourceLocation(limiter) => { - let file = log_msg.file.as_deref().unwrap_or("?"); - let line = log_msg.line.as_deref().unwrap_or("?"); - limiter.evaluate(file, line, ts_sec) - } - } - } - - /// Create a multi-rate limiter from a threshold. - pub fn multi(threshold: RateThreshold) -> Self { - Limiter::Multi(MultiRateLimiter::from(threshold)) - } - - /// Create a source-location limiter from a threshold. - /// - /// Note: Only the duration is used; source-location limiters allow one event - /// per location per window. - pub fn source_location(threshold: RateThreshold) -> Self { - Limiter::SourceLocation(SourceLocationRateLimiter::new( - threshold.duration, - SOURCE_LOCATION_MAX_ENTRIES, - )) - } -} - -/// A rate limiter containing a single counter, and a minimum time window for the next event to pass -#[allow(dead_code)] -#[derive(Debug, Default)] -pub struct SimpleRateLimiter { - last_timestamp: AtomicI64, - window: i64, -} - -#[allow(dead_code)] -impl SimpleRateLimiter { - pub fn new(window: Duration) -> Self { - Self { - last_timestamp: Default::default(), - window: window.as_secs().try_into().unwrap(), - } - } - - /// Check if a particular new timestamp passes the limit. This also updates the last-known timestamp. - pub fn evaluate(&self, ts_sec: i64) -> bool { - let last_ts = self.last_timestamp.load(Ordering::SeqCst); - let rate_limited = ts_sec - last_ts < self.window; - if !rate_limited && ts_sec > last_ts { - // If this is called concurrently, guarantee that we keep going - // until the max value is stored at self.last_timestamp, - // so self.last_timestamp is "eventually" only monotonically increasing. - store_max(ts_sec, &self.last_timestamp); - } - !rate_limited - } -} - -#[allow(dead_code)] -fn store_max(val: i64, at: &AtomicI64) { - let prev = at.swap(val, Ordering::SeqCst); - if prev > val { - store_max(prev, at) - } -} - -/// A rate limiter that tracks alerts per source location (file:line). -/// -/// This allows different error locations to alert independently, preventing one noisy -/// error from suppressing alerts from completely different code paths. -#[derive(Debug)] -pub struct SourceLocationRateLimiter { - /// Maps (file, line) -> last alert timestamp - last_timestamps: HashMap<(String, String), i64>, - /// The rate limiting window in seconds - window: i64, - /// Maximum entries before triggering cleanup - max_entries: usize, -} - -impl SourceLocationRateLimiter { - pub fn new(window: Duration, max_entries: usize) -> Self { - Self { - last_timestamps: HashMap::new(), - window: window.as_secs().try_into().unwrap(), - max_entries, - } - } - - /// Check if an error from this source location should trigger an alert. - /// - /// Returns true if the alert should fire (not rate-limited), false if suppressed. - /// Updates the stored timestamp if the alert fires. - pub fn evaluate(&mut self, file: &str, line: &str, ts_sec: i64) -> bool { - let key = (file.to_owned(), line.to_owned()); - - if let Some(&last_ts) = self.last_timestamps.get(&key) - && ts_sec - last_ts < self.window - { - return false; // Rate limited - } - - // Alert should fire - update timestamp - self.last_timestamps.insert(key, ts_sec); - - // Clean up if we've exceeded max entries - if self.last_timestamps.len() > self.max_entries { - self.cleanup(ts_sec); - } - - true - } - - /// Remove entries older than the window - fn cleanup(&mut self, now: i64) { - self.last_timestamps - .retain(|_, &mut ts| now - ts < self.window); - } -} - -/// Implements rate-limiting criteria such as 'at least n in the last w seconds' -#[derive(Debug)] -pub struct MultiRateLimiter { - /// Records the last n events - timestamps: Vec, - /// Invariant: Always points to the oldest of the last n timestamps in the buffer - idx: usize, - /// The length of the window (in seconds) - window: i64, -} - -impl MultiRateLimiter { - pub fn new(num: usize, window: Duration) -> Self { - Self { - idx: 0, - timestamps: vec![Default::default(); num], - window: window.as_secs().try_into().unwrap(), - } - } - - /// Check if a particular new timestamp passes the limit. This also updates the last-known timestamp. - /// - /// Note: Assumes that new_timestamp is monotonically increasing, otherwise it might not work right. - pub fn evaluate(&mut self, new_timestamp: i64) -> bool { - let oldest = self.timestamps[self.idx]; - if oldest >= new_timestamp { - return false; - } - self.timestamps[self.idx] = new_timestamp; - self.idx += 1; - self.idx %= self.timestamps.len(); - let next_oldest = self.timestamps[self.idx]; - // If the next oldest is within 'window' of the new timestamp, - // then all of the most recent n are. Otherwise, at most n-1 of the most recent are. - next_oldest + self.window >= new_timestamp - } -} - -impl From for MultiRateLimiter { - fn from(src: RateThreshold) -> Self { - Self::new(src.times, src.duration) - } -} - -#[cfg(test)] -mod tests { - use super::*; - use std::str::FromStr; - - #[test] - fn parse_rate_threshold() { - let threshold = RateThreshold::from_str("1/10s").unwrap(); - assert_eq!(threshold.times, 1); - assert_eq!(threshold.duration, Duration::from_secs(10)); - let threshold = RateThreshold::from_str("1 / 10s").unwrap(); - assert_eq!(threshold.times, 1); - assert_eq!(threshold.duration, Duration::from_secs(10)); - - let threshold = RateThreshold::from_str("2 / 5m").unwrap(); - assert_eq!(threshold.times, 2); - assert_eq!(threshold.duration, Duration::from_secs(300)); - - let threshold = RateThreshold::from_str("> 3 / 10m").unwrap(); - assert_eq!(threshold.times, 4); - assert_eq!(threshold.duration, Duration::from_secs(600)); - - let threshold = RateThreshold::from_str(">=3/10m").unwrap(); - assert_eq!(threshold.times, 3); - assert_eq!(threshold.duration, Duration::from_secs(600)); - } - - #[test] - fn source_location_rate_limiter_basic() { - let mut limiter = SourceLocationRateLimiter::new(Duration::from_secs(600), 100); - - // First alert from location A should pass - assert!(limiter.evaluate("file_a.rs", "10", 1000)); - - // Second alert from same location within window should be rate limited - assert!(!limiter.evaluate("file_a.rs", "10", 1100)); - - // Alert from different location should pass (independent rate limiting) - assert!(limiter.evaluate("file_b.rs", "20", 1100)); - - // Same location after window passes should alert again - assert!(limiter.evaluate("file_a.rs", "10", 1700)); // 1000 + 600 + 100 - } - - #[test] - fn source_location_rate_limiter_different_lines_same_file() { - let mut limiter = SourceLocationRateLimiter::new(Duration::from_secs(600), 100); - - // Different lines in same file should be independent - assert!(limiter.evaluate("file.rs", "10", 1000)); - assert!(limiter.evaluate("file.rs", "20", 1000)); - assert!(limiter.evaluate("file.rs", "30", 1000)); - - // Each should still be rate limited individually - assert!(!limiter.evaluate("file.rs", "10", 1100)); - assert!(!limiter.evaluate("file.rs", "20", 1100)); - } - - #[test] - fn source_location_rate_limiter_cleanup() { - // Use small max_entries to trigger cleanup - let mut limiter = SourceLocationRateLimiter::new(Duration::from_secs(600), 3); - - // Fill up the limiter - assert!(limiter.evaluate("file1.rs", "1", 1000)); - assert!(limiter.evaluate("file2.rs", "2", 1000)); - assert!(limiter.evaluate("file3.rs", "3", 1000)); - assert_eq!(limiter.last_timestamps.len(), 3); - - // Add one more, triggering cleanup - but all are fresh so none removed - assert!(limiter.evaluate("file4.rs", "4", 1000)); - // Still have 4 after cleanup since none are old enough - assert_eq!(limiter.last_timestamps.len(), 4); - - // Now add with a timestamp far in the future - old entries should be cleaned - assert!(limiter.evaluate("file5.rs", "5", 2000)); - // Should have cleaned up entries from timestamp 1000 (older than 600 sec window) - assert_eq!(limiter.last_timestamps.len(), 1); - } -} diff --git a/signal-gateway/src/gateway/rate_limiter_set.rs b/signal-gateway/src/gateway/rate_limiter_set.rs index 4cff05d..6f4fe23 100644 --- a/signal-gateway/src/gateway/rate_limiter_set.rs +++ b/signal-gateway/src/gateway/rate_limiter_set.rs @@ -2,7 +2,7 @@ use crate::{ concurrent_map::LazyMap, - log_message::{LogMessage, Origin}, + log_message::{LogFilter, LogMessage, Origin}, rate_limiter::Limiter, }; @@ -15,21 +15,26 @@ pub enum LimitResult { Limiter(usize), /// The event was blocked by a global limiter at the given index. GlobalLimiter(usize), + /// The event was blocked by an overall limiter at the given index. + OverallLimiter(usize), } /// A set of limiters for a route, containing both per-origin and global limiters. +/// Each limiter is paired with a filter that must match before the limiter is evaluated. pub struct LimiterSet { /// Per-origin rate limiters, keyed by origin. Lazily created. - limiters: LazyMap>, + /// Each entry is a (filter, limiter) pair. + limiters: LazyMap>, /// Global rate limiters (shared across all origins). - global_limiters: Vec, + /// Each entry is a (filter, limiter) pair. + global_limiters: Vec<(LogFilter, Limiter)>, } 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, + make_limiters: impl Fn() -> Vec<(LogFilter, Limiter)> + Send + Sync + 'static, + global_limiters: Vec<(LogFilter, Limiter)>, ) -> Self { Self { limiters: LazyMap::new(make_limiters), @@ -39,14 +44,19 @@ impl LimiterSet { /// Evaluate whether an event should pass all rate limits. /// + /// For each limiter, first checks if the message matches the filter. + /// If it matches, evaluates the rate limiter. + /// If it doesn't match, the limiter is skipped (passes). + /// /// 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(&self, log_msg: &LogMessage, origin: &Origin, ts_sec: i64) -> LimitResult { // Check per-origin limiters let origin_result = self.limiters.get(origin, |origin_limiters| { - for (i, limiter) in origin_limiters.iter().enumerate() { - if !limiter.evaluate(log_msg, ts_sec) { + for (i, (filter, limiter)) in origin_limiters.iter().enumerate() { + // Only evaluate the limiter if the message matches the filter + if filter.matches(log_msg) && !limiter.evaluate(log_msg, ts_sec) { return Some(LimitResult::Limiter(i)); } } @@ -58,8 +68,9 @@ impl LimiterSet { } // Check global limiters - for (i, limiter) in self.global_limiters.iter().enumerate() { - if !limiter.evaluate(log_msg, ts_sec) { + for (i, (filter, limiter)) in self.global_limiters.iter().enumerate() { + // Only evaluate the limiter if the message matches the filter + if filter.matches(log_msg) && !limiter.evaluate(log_msg, ts_sec) { return LimitResult::GlobalLimiter(i); } } diff --git a/signal-gateway/src/gateway/route.rs b/signal-gateway/src/gateway/route.rs index abc9e15..22e0fb3 100644 --- a/signal-gateway/src/gateway/route.rs +++ b/signal-gateway/src/gateway/route.rs @@ -29,12 +29,14 @@ pub struct Limit { impl Limit { /// Create the appropriate limiter for this limit configuration. - pub fn make_limiter(&self) -> Limiter { - if self.by_source_location { + /// Returns a (filter, limiter) pair so the filter can be checked before rate limiting. + pub fn make_limiter(&self) -> (LogFilter, Limiter) { + let limiter = if self.by_source_location { Limiter::source_location(self.threshold) } else { Limiter::multi(self.threshold) - } + }; + (self.filter.clone(), limiter) } } @@ -81,7 +83,11 @@ impl Route { /// Create a limiter set from this route's limit configurations. pub fn make_limiter_set(&self) -> LimiterSet { let limits = self.limits.clone(); - let global_limiters = self.global_limits.iter().map(|l| l.make_limiter()).collect(); + let global_limiters = self + .global_limits + .iter() + .map(|l| l.make_limiter()) + .collect(); LimiterSet::new( move || limits.iter().map(|l| l.make_limiter()).collect(), global_limiters, diff --git a/signal-gateway/src/gateway/signal_trust_set.rs b/signal-gateway/src/gateway/signal_trust_set.rs index cd80511..1a402dc 100644 --- a/signal-gateway/src/gateway/signal_trust_set.rs +++ b/signal-gateway/src/gateway/signal_trust_set.rs @@ -170,7 +170,10 @@ impl SignalTrustSet { } } - info!("Trust reset verified for {uuid}: {} safety numbers trusted", trusted_now.len()); + info!( + "Trust reset verified for {uuid}: {} safety numbers trusted", + trusted_now.len() + ); } else { // Just add any new safety numbers that aren't already trusted let already_trusted: Vec<_> = current_identities