use saorsa_gossip_rendezvous::ProviderSummary;
use std::collections::HashMap;
use std::sync::Arc;
use std::time::{SystemTime, UNIX_EPOCH};
use tokio::sync::RwLock;
use tracing::debug;
pub struct ProviderRateLimiter {
#[allow(clippy::type_complexity)]
counts: Arc<RwLock<HashMap<[u8; 32], (u32, u64)>>>,
max_per_window: u32,
window_ms: u64,
}
impl ProviderRateLimiter {
pub fn new(max_per_window: u32, window_ms: u64) -> Self {
Self {
counts: Arc::new(RwLock::new(HashMap::new())),
max_per_window,
window_ms,
}
}
pub async fn check_and_update(&self, target_id: &[u8; 32]) -> bool {
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_else(|_| std::time::Duration::from_secs(0))
.as_millis() as u64;
let mut counts = self.counts.write().await;
let (count, window_start) = counts.entry(*target_id).or_insert((0, now));
if now - *window_start > self.window_ms {
*count = 1;
*window_start = now;
return true;
}
if *count >= self.max_per_window {
debug!(
"Rate limit exceeded for target {:?}",
hex::encode(target_id)
);
return false;
}
*count += 1;
true
}
pub async fn cleanup(&self) {
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_else(|_| std::time::Duration::from_secs(0))
.as_millis() as u64;
let mut counts = self.counts.write().await;
counts.retain(|_, (_, window_start)| {
now - *window_start < self.window_ms * 2 });
}
}
impl Default for ProviderRateLimiter {
fn default() -> Self {
Self::new(10, 1000) }
}
pub struct ProviderCollector {
target_id: [u8; 32],
providers: Arc<RwLock<Vec<ProviderSummary>>>,
rate_limiter: Arc<ProviderRateLimiter>,
stats: Arc<RwLock<CollectorStats>>,
}
#[derive(Debug, Clone, Default)]
pub struct CollectorStats {
pub total_received: u64,
pub signature_valid: u64,
pub signature_invalid: u64,
pub rate_limited: u64,
pub expired: u64,
pub accepted: u64,
}
impl ProviderCollector {
pub fn new(target_id: [u8; 32]) -> Self {
Self {
target_id,
providers: Arc::new(RwLock::new(Vec::new())),
rate_limiter: Arc::new(ProviderRateLimiter::default()),
stats: Arc::new(RwLock::new(CollectorStats::default())),
}
}
pub async fn process(&self, summary: ProviderSummary, is_valid: bool) -> bool {
let mut stats = self.stats.write().await;
stats.total_received += 1;
drop(stats);
if summary.target != self.target_id {
return false; }
if !is_valid {
let mut stats = self.stats.write().await;
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_else(|_| std::time::Duration::from_secs(0))
.as_millis() as u64;
if summary.exp <= now {
stats.expired += 1;
} else {
stats.signature_invalid += 1;
}
return false;
}
if !self.rate_limiter.check_and_update(&self.target_id).await {
let mut stats = self.stats.write().await;
stats.rate_limited += 1;
return false;
}
let mut stats = self.stats.write().await;
stats.signature_valid += 1;
drop(stats);
let mut providers = self.providers.write().await;
providers.retain(|p| p.provider != summary.provider);
providers.push(summary);
let mut stats = self.stats.write().await;
stats.accepted += 1;
true
}
pub async fn get_providers(&self) -> Vec<ProviderSummary> {
self.providers.read().await.clone()
}
pub async fn stats(&self) -> CollectorStats {
self.stats.read().await.clone()
}
}
#[cfg(test)]
mod tests {
use super::*;
use saorsa_gossip_types::PeerId;
fn create_test_summary(
target: [u8; 32],
provider: PeerId,
validity_ms: u64,
) -> ProviderSummary {
ProviderSummary::new(
target,
provider,
vec![saorsa_gossip_rendezvous::Capability::Site],
validity_ms,
)
}
#[tokio::test]
async fn test_rate_limiter_basic() {
let limiter = ProviderRateLimiter::new(3, 1000); let target = [1u8; 32];
assert!(limiter.check_and_update(&target).await);
assert!(limiter.check_and_update(&target).await);
assert!(limiter.check_and_update(&target).await);
assert!(!limiter.check_and_update(&target).await);
}
#[tokio::test]
async fn test_rate_limiter_window_reset() {
let limiter = ProviderRateLimiter::new(2, 100); let target = [2u8; 32];
assert!(limiter.check_and_update(&target).await);
assert!(limiter.check_and_update(&target).await);
assert!(!limiter.check_and_update(&target).await);
tokio::time::sleep(tokio::time::Duration::from_millis(150)).await;
assert!(limiter.check_and_update(&target).await);
}
#[tokio::test]
async fn test_collector_accepts_valid_summary() {
let target = [1u8; 32];
let collector = ProviderCollector::new(target);
let summary = create_test_summary(target, PeerId::new([2u8; 32]), 60000);
assert!(collector.process(summary, true).await);
let providers = collector.get_providers().await;
assert_eq!(providers.len(), 1);
let stats = collector.stats().await;
assert_eq!(stats.accepted, 1);
assert_eq!(stats.signature_valid, 1);
}
#[tokio::test]
async fn test_collector_rejects_invalid_summary() {
let target = [1u8; 32];
let collector = ProviderCollector::new(target);
let summary = create_test_summary(target, PeerId::new([2u8; 32]), 60000);
assert!(!collector.process(summary, false).await);
let stats = collector.stats().await;
assert_eq!(stats.accepted, 0);
assert_eq!(stats.signature_invalid, 1);
}
#[tokio::test]
async fn test_collector_enforces_rate_limit() {
let target = [1u8; 32];
let collector = ProviderCollector::new(target);
let mut accepted = 0;
for i in 0..15 {
let summary = create_test_summary(target, PeerId::new([i as u8; 32]), 60000);
if collector.process(summary, true).await {
accepted += 1;
}
}
assert_eq!(accepted, 10);
let stats = collector.stats().await;
assert_eq!(stats.rate_limited, 5);
}
#[tokio::test]
async fn test_collector_deduplicates_by_provider() {
let target = [1u8; 32];
let collector = ProviderCollector::new(target);
let provider = PeerId::new([2u8; 32]);
for _ in 0..3 {
let summary = create_test_summary(target, provider, 60000);
collector.process(summary, true).await;
tokio::time::sleep(tokio::time::Duration::from_millis(10)).await;
}
let providers = collector.get_providers().await;
assert_eq!(providers.len(), 1);
assert_eq!(providers[0].provider, provider);
}
}