use anyhow::Result;
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use std::fmt;
use std::sync::Arc;
use crate::scanner::{Threat, ThreatType};
pub mod api;
#[cfg(feature = "enhanced")]
pub mod enhanced;
pub mod health;
pub mod metrics;
pub mod rate_limited;
pub mod recovery;
pub mod rollback;
pub mod security_aware;
pub mod standard;
pub mod traced;
pub mod validation;
#[cfg(test)]
mod security_tests;
#[async_trait]
pub trait ThreatNeutralizer: Send + Sync {
async fn neutralize(&self, threat: &Threat, content: &str) -> Result<NeutralizeResult>;
fn can_neutralize(&self, threat_type: &ThreatType) -> bool;
fn get_capabilities(&self) -> NeutralizerCapabilities;
async fn batch_neutralize(
&self,
threats: &[Threat],
content: &str,
) -> Result<BatchNeutralizeResult> {
let mut results = Vec::new();
let mut current_content = content.to_string();
for threat in threats {
let result = self.neutralize(threat, ¤t_content).await?;
if let Some(ref sanitized) = result.sanitized_content {
current_content = sanitized.clone();
}
results.push(result);
}
Ok(BatchNeutralizeResult {
final_content: current_content,
individual_results: results,
})
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct NeutralizeResult {
pub action_taken: NeutralizeAction,
pub sanitized_content: Option<String>,
pub confidence_score: f64,
pub processing_time_us: u64,
pub correlation_data: Option<CorrelationData>,
pub extracted_params: Option<Vec<String>>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct BatchNeutralizeResult {
pub final_content: String,
pub individual_results: Vec<NeutralizeResult>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum NeutralizeAction {
Sanitized,
Parameterized,
Normalized,
Escaped,
Removed,
Quarantined,
NoAction,
}
impl fmt::Display for NeutralizeAction {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Sanitized => write!(f, "Sanitized"),
Self::Parameterized => write!(f, "Parameterized"),
Self::Normalized => write!(f, "Normalized"),
Self::Escaped => write!(f, "Escaped"),
Self::Removed => write!(f, "Removed"),
Self::Quarantined => write!(f, "Quarantined"),
Self::NoAction => write!(f, "No Action"),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CorrelationData {
pub related_threats: Vec<String>,
pub attack_pattern: Option<AttackPattern>,
pub prediction_score: f64,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub enum AttackPattern {
CoordinatedUnicode,
SqlInjectionCampaign,
CommandEscalation,
MultiVector,
Probing,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct NeutralizerCapabilities {
pub real_time: bool,
pub batch_mode: bool,
pub predictive: bool,
pub correlation: bool,
pub rollback_depth: usize,
pub supported_threats: Vec<ThreatType>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum NeutralizationMode {
ReportOnly,
Interactive,
Automatic,
}
impl Default for NeutralizationMode {
fn default() -> Self {
Self::ReportOnly
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct NeutralizationConfig {
pub mode: NeutralizationMode,
pub backup_originals: bool,
pub audit_all_actions: bool,
pub unicode: UnicodeNeutralizationConfig,
pub injection: InjectionNeutralizationConfig,
#[serde(skip_serializing_if = "Option::is_none")]
pub recovery: Option<recovery::RecoveryConfig>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct UnicodeNeutralizationConfig {
pub bidi_replacement: BiDiReplacement,
pub zero_width_action: ZeroWidthAction,
pub homograph_action: HomographAction,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum BiDiReplacement {
Remove,
Marker,
Escape,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum ZeroWidthAction {
Remove,
Escape,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum HomographAction {
Ascii,
Warn,
Block,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct InjectionNeutralizationConfig {
pub sql_action: SqlAction,
pub command_action: CommandAction,
pub path_action: PathAction,
pub prompt_action: PromptAction,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum SqlAction {
Block,
Escape,
Parameterize,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum CommandAction {
Block,
Escape,
Sandbox,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum PathAction {
Block,
Normalize,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum PromptAction {
Block,
Escape,
Wrap,
}
impl Default for NeutralizationConfig {
fn default() -> Self {
Self {
mode: NeutralizationMode::default(),
backup_originals: true,
audit_all_actions: true,
unicode: UnicodeNeutralizationConfig {
bidi_replacement: BiDiReplacement::Marker,
zero_width_action: ZeroWidthAction::Remove,
homograph_action: HomographAction::Ascii,
},
injection: InjectionNeutralizationConfig {
sql_action: SqlAction::Parameterize,
command_action: CommandAction::Escape,
path_action: PathAction::Normalize,
prompt_action: PromptAction::Wrap,
},
recovery: Some(recovery::RecoveryConfig::default()),
}
}
}
pub fn create_neutralizer(
config: &NeutralizationConfig,
rate_limiter: Option<Arc<dyn crate::traits::RateLimiter>>,
) -> Arc<dyn ThreatNeutralizer> {
create_neutralizer_with_telemetry(config, rate_limiter, None)
}
pub fn create_neutralizer_with_telemetry(
config: &NeutralizationConfig,
rate_limiter: Option<Arc<dyn crate::traits::RateLimiter>>,
tracing_provider: Option<Arc<crate::telemetry::DistributedTracingProvider>>,
) -> Arc<dyn ThreatNeutralizer> {
let mut neutralizer: Arc<dyn ThreatNeutralizer> = {
#[cfg(feature = "enhanced")]
{
Arc::new(enhanced::EnhancedNeutralizer::new(config.clone()))
}
#[cfg(not(feature = "enhanced"))]
{
Arc::new(standard::StandardNeutralizer::new(config.clone()))
}
};
if let Some(ref recovery_config) = config.recovery {
if recovery_config.enabled {
neutralizer = Arc::new(recovery::ResilientNeutralizer::new(
neutralizer,
recovery_config.clone(),
));
}
}
if config.backup_originals {
neutralizer =
rollback::RollbackNeutralizer::new(neutralizer, rollback::RollbackConfig::default());
}
if let Some(limiter) = rate_limiter {
neutralizer = Arc::new(rate_limited::RateLimitedNeutralizer::new(
neutralizer,
limiter,
rate_limited::NeutralizationRateLimitConfig::default(),
));
}
neutralizer = health::HealthMonitoredNeutralizer::new(
neutralizer,
health::NeutralizationHealthConfig::default(),
);
if let Some(provider) = tracing_provider {
use crate::neutralizer::traced::NeutralizerTracingExt;
neutralizer = neutralizer.with_tracing(provider);
}
neutralizer
}
#[derive(Debug, thiserror::Error)]
pub enum NeutralizeError {
#[error("Threat type not supported: {0:?}")]
UnsupportedThreatType(ThreatType),
#[error("Neutralization failed: {0}")]
NeutralizationFailed(String),
#[error("Invalid configuration: {0}")]
InvalidConfig(String),
}