use crate::{
neutralizer::{NeutralizeResult, ThreatNeutralizer},
scanner::Threat,
traits::{RateLimitKey, RateLimiter},
};
use anyhow::{bail, Result};
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use std::sync::Arc;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct NeutralizationRateLimitConfig {
pub enabled: bool,
pub per_minute: u32,
pub per_hour: u32,
pub max_concurrent: u32,
pub strict_for_expensive: bool,
pub expensive_threshold: usize,
pub expensive_multiplier: f32,
}
impl Default for NeutralizationRateLimitConfig {
fn default() -> Self {
Self {
enabled: true,
per_minute: 60,
per_hour: 1000,
max_concurrent: 5,
strict_for_expensive: true,
expensive_threshold: 1024 * 100, expensive_multiplier: 0.5, }
}
}
pub struct RateLimitedNeutralizer {
inner: Arc<dyn ThreatNeutralizer>,
rate_limiter: Arc<dyn RateLimiter>,
config: NeutralizationRateLimitConfig,
concurrent_operations: Arc<tokio::sync::Mutex<std::collections::HashMap<String, u32>>>,
}
impl RateLimitedNeutralizer {
pub fn new(
neutralizer: Arc<dyn ThreatNeutralizer>,
rate_limiter: Arc<dyn RateLimiter>,
config: NeutralizationRateLimitConfig,
) -> Self {
Self {
inner: neutralizer,
rate_limiter,
config,
concurrent_operations: Arc::new(tokio::sync::Mutex::new(
std::collections::HashMap::new(),
)),
}
}
fn get_client_id(&self) -> String {
"anonymous".to_string()
}
const fn is_expensive_operation(&self, content: &str) -> bool {
content.len() > self.config.expensive_threshold
}
async fn apply_expensive_penalty(&self, client_id: &str) -> Result<()> {
if self.config.strict_for_expensive {
self.rate_limiter
.apply_penalty(client_id, self.config.expensive_multiplier)
.await?;
}
Ok(())
}
async fn increment_concurrent(&self, client_id: &str) -> Result<()> {
let mut ops = self.concurrent_operations.lock().await;
let count = ops.entry(client_id.to_string()).or_insert(0);
if *count >= self.config.max_concurrent {
bail!("Maximum concurrent neutralization operations exceeded");
}
*count += 1;
Ok(())
}
}
#[async_trait]
impl ThreatNeutralizer for RateLimitedNeutralizer {
async fn neutralize(&self, threat: &Threat, content: &str) -> Result<NeutralizeResult> {
if !self.config.enabled {
return self.inner.neutralize(threat, content).await;
}
let client_id = self.get_client_id();
let key = RateLimitKey {
client_id: client_id.clone(),
method: Some("neutralize".to_string()),
};
let decision = self.rate_limiter.check_rate_limit(&key).await?;
if !decision.allowed {
bail!(
"Rate limit exceeded for neutralization. Reset in {:?}",
decision.reset_after
);
}
self.increment_concurrent(&client_id).await?;
let _guard = ConcurrentGuard {
concurrent_ops: self.concurrent_operations.clone(),
client_id: client_id.clone(),
};
if self.is_expensive_operation(content) {
self.apply_expensive_penalty(&client_id).await?;
tracing::debug!(
"Applied rate limit penalty for expensive neutralization: {} bytes",
content.len()
);
}
self.rate_limiter.record_request(&key).await?;
let result = self.inner.neutralize(threat, content).await?;
if result.processing_time_us > 10_000_000 {
self.rate_limiter.apply_penalty(&client_id, 0.75).await?;
tracing::warn!(
"Applied additional penalty for slow neutralization: {}μs",
result.processing_time_us
);
}
Ok(result)
}
fn can_neutralize(&self, threat_type: &crate::scanner::ThreatType) -> bool {
self.inner.can_neutralize(threat_type)
}
fn get_capabilities(&self) -> crate::neutralizer::NeutralizerCapabilities {
let mut capabilities = self.inner.get_capabilities();
if self.config.enabled {
capabilities.real_time = false; }
capabilities
}
async fn batch_neutralize(
&self,
threats: &[crate::scanner::Threat],
content: &str,
) -> Result<crate::neutralizer::BatchNeutralizeResult> {
let client_id = self.get_client_id();
let key = crate::traits::RateLimitKey {
client_id: client_id.clone(),
method: Some("neutralize_batch".to_string()),
};
let decision = self.rate_limiter.check_rate_limit(&key).await?;
if !decision.allowed {
return Err(anyhow::anyhow!(
"Rate limit exceeded. Try again in {:?}",
decision.reset_after
));
}
self.rate_limiter.record_request(&key).await?;
self.inner.batch_neutralize(threats, content).await
}
}
struct ConcurrentGuard {
concurrent_ops: Arc<tokio::sync::Mutex<std::collections::HashMap<String, u32>>>,
client_id: String,
}
impl Drop for ConcurrentGuard {
fn drop(&mut self) {
let concurrent_ops = self.concurrent_ops.clone();
let client_id = self.client_id.clone();
tokio::spawn(async move {
let mut ops = concurrent_ops.lock().await;
if let Some(count) = ops.get_mut(&client_id) {
*count = count.saturating_sub(1);
if *count == 0 {
ops.remove(&client_id);
}
}
});
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct NeutralizationRateLimitStats {
pub total_requests: u64,
pub rate_limited_requests: u64,
pub concurrent_limit_hits: u64,
pub expensive_operations: u64,
pub average_tokens_remaining: f64,
}
#[cfg(test)]
mod tests {
use super::*;
use crate::neutralizer::standard::StandardNeutralizer;
use crate::neutralizer::NeutralizationConfig;
struct MockRateLimiter {
allow: bool,
}
#[async_trait]
impl RateLimiter for MockRateLimiter {
async fn check_rate_limit(
&self,
_key: &crate::traits::RateLimitKey,
) -> Result<crate::traits::RateLimitDecision> {
Ok(crate::traits::RateLimitDecision {
allowed: self.allow,
tokens_remaining: 10.0,
reset_after: std::time::Duration::from_secs(60),
})
}
async fn record_request(&self, _key: &crate::traits::RateLimitKey) -> Result<()> {
Ok(())
}
async fn apply_penalty(&self, _client_id: &str, _factor: f32) -> Result<()> {
Ok(())
}
fn get_stats(&self) -> crate::traits::RateLimiterStats {
crate::traits::RateLimiterStats {
requests_allowed: 100,
requests_denied: 10,
active_buckets: 5,
}
}
}
#[tokio::test]
async fn test_rate_limited_neutralizer() {
let config = NeutralizationConfig::default();
let neutralizer = Arc::new(StandardNeutralizer::new(config));
let rate_limiter = Arc::new(MockRateLimiter { allow: true });
let rate_config = NeutralizationRateLimitConfig::default();
let limited = RateLimitedNeutralizer::new(neutralizer, rate_limiter, rate_config);
assert!(limited.can_neutralize(&crate::scanner::ThreatType::SqlInjection));
}
}