use serde::{Deserialize, Serialize};
use std::time::Instant;
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)]
pub enum Severity {
Low,
Medium,
High,
Critical,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum VulnerabilityType {
SqlInjection,
Xss,
AuthBypass,
RateLimitBypass,
CommandInjection,
PathTraversal,
Csrf,
SensitiveDataExposure,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct VulnerabilityFinding {
pub vuln_type: VulnerabilityType,
pub severity: Severity,
pub description: String,
pub affected_component: String,
pub poc: Option<String>,
pub remediation: String,
}
impl VulnerabilityFinding {
pub fn new(
vuln_type: VulnerabilityType,
severity: Severity,
description: &str,
component: &str,
) -> Self {
Self {
vuln_type,
severity,
description: description.to_string(),
affected_component: component.to_string(),
poc: None,
remediation: String::new(),
}
}
pub fn with_poc(mut self, poc: &str) -> Self {
self.poc = Some(poc.to_string());
self
}
pub fn with_remediation(mut self, remediation: &str) -> Self {
self.remediation = remediation.to_string();
self
}
}
const SQL_INJECTION_PATTERNS: &[&str] = &[
"' OR '1'='1",
"1' OR '1'='1",
"\" OR \"1\"=\"1",
"1 OR 1=1",
"'; DROP TABLE users--",
"1'; DROP TABLE users--",
"1 UNION SELECT * FROM users",
"' UNION SELECT NULL--",
];
const XSS_PATTERNS: &[&str] = &[
"<script>alert('XSS')</script>",
"<img src=x onerror=alert('XSS')>",
"<svg onload=alert('XSS')>",
"javascript:alert('XSS')",
"<iframe src='javascript:alert(\"XSS\")'></iframe>",
"<body onload=alert('XSS')>",
];
const PATH_TRAVERSAL_PATTERNS: &[&str] = &[
"../../../etc/passwd",
"..\\..\\..\\windows\\system32\\config\\sam",
"....//....//....//etc/passwd",
"%2e%2e%2f%2e%2e%2f%2e%2e%2fetc%2fpasswd",
];
pub struct VulnerabilityScanner {
findings: Vec<VulnerabilityFinding>,
tests_run: usize,
vulnerabilities_found: usize,
}
impl VulnerabilityScanner {
pub fn new() -> Self {
Self {
findings: Vec::new(),
tests_run: 0,
vulnerabilities_found: 0,
}
}
pub fn scan_sql_injection(&mut self, query: &str, user_input: &str) -> bool {
self.tests_run += 1;
let mut vulnerable = false;
for pattern in SQL_INJECTION_PATTERNS {
let test_input = format!("{}{}", user_input, pattern);
if query.contains(&test_input) || self.is_sql_injection_pattern(&test_input) {
let finding = VulnerabilityFinding::new(
VulnerabilityType::SqlInjection,
Severity::Critical,
"Potential SQL injection vulnerability detected",
query,
)
.with_poc(pattern)
.with_remediation("Use parameterized queries or prepared statements");
self.findings.push(finding);
self.vulnerabilities_found += 1;
vulnerable = true;
break;
}
}
vulnerable
}
pub fn scan_xss(&mut self, output: &str, user_input: &str) -> bool {
self.tests_run += 1;
let mut vulnerable = false;
for pattern in XSS_PATTERNS {
let test_input = format!("{}{}", user_input, pattern);
if output.contains(pattern) || self.is_xss_pattern(&test_input) {
let finding = VulnerabilityFinding::new(
VulnerabilityType::Xss,
Severity::High,
"Cross-site scripting vulnerability detected",
"Output",
)
.with_poc(pattern)
.with_remediation("Sanitize and escape user input before rendering");
self.findings.push(finding);
self.vulnerabilities_found += 1;
vulnerable = true;
break;
}
}
vulnerable
}
pub fn test_auth_bypass(&mut self, endpoint: &str, bypass_attempts: &[(&str, &str)]) -> bool {
self.tests_run += 1;
let mut vulnerable = false;
for (username, password) in bypass_attempts {
let bypass_patterns = [
("admin' --", ""),
("admin'/*", ""),
("' OR 1=1--", ""),
("admin", "' OR '1'='1"),
];
for (user_pattern, pass_pattern) in &bypass_patterns {
if username.contains(user_pattern) || password.contains(pass_pattern) {
let finding = VulnerabilityFinding::new(
VulnerabilityType::AuthBypass,
Severity::Critical,
"Authentication bypass vulnerability detected",
endpoint,
)
.with_poc(&format!(
"username: {}, password: {}",
user_pattern, pass_pattern
))
.with_remediation("Use parameterized queries and proper password hashing");
self.findings.push(finding);
self.vulnerabilities_found += 1;
vulnerable = true;
break;
}
}
}
vulnerable
}
pub fn test_rate_limit(&mut self, endpoint: &str, requests_per_second: usize) -> bool {
self.tests_run += 1;
let mut vulnerable = false;
let threshold = 100;
if requests_per_second > threshold {
let finding = VulnerabilityFinding::new(
VulnerabilityType::RateLimitBypass,
Severity::Medium,
&format!(
"Rate limit bypass detected: {} requests/sec exceeds threshold of {}",
requests_per_second, threshold
),
endpoint,
)
.with_poc(&format!(
"Sent {} requests in 1 second",
requests_per_second
))
.with_remediation("Implement stricter rate limiting with token bucket or leaky bucket");
self.findings.push(finding);
self.vulnerabilities_found += 1;
vulnerable = true;
}
vulnerable
}
pub fn test_path_traversal(&mut self, file_path: &str, user_input: &str) -> bool {
self.tests_run += 1;
let mut vulnerable = false;
for pattern in PATH_TRAVERSAL_PATTERNS {
let test_path = format!("{}{}", user_input, pattern);
if file_path.contains("..") || test_path.contains("..") {
let finding = VulnerabilityFinding::new(
VulnerabilityType::PathTraversal,
Severity::High,
"Path traversal vulnerability detected",
file_path,
)
.with_poc(pattern)
.with_remediation("Validate and sanitize file paths, use whitelisting");
self.findings.push(finding);
self.vulnerabilities_found += 1;
vulnerable = true;
break;
}
}
vulnerable
}
pub fn get_findings(&self) -> &[VulnerabilityFinding] {
&self.findings
}
pub fn get_findings_by_severity(&self, severity: Severity) -> Vec<&VulnerabilityFinding> {
self.findings
.iter()
.filter(|f| f.severity == severity)
.collect()
}
pub fn get_summary(&self) -> ScanSummary {
let critical = self.get_findings_by_severity(Severity::Critical).len();
let high = self.get_findings_by_severity(Severity::High).len();
let medium = self.get_findings_by_severity(Severity::Medium).len();
let low = self.get_findings_by_severity(Severity::Low).len();
ScanSummary {
tests_run: self.tests_run,
vulnerabilities_found: self.vulnerabilities_found,
critical_findings: critical,
high_findings: high,
medium_findings: medium,
low_findings: low,
}
}
fn is_sql_injection_pattern(&self, input: &str) -> bool {
let input_lower = input.to_lowercase();
input_lower.contains("union")
|| input_lower.contains("drop")
|| input_lower.contains("delete")
|| input_lower.contains("insert")
|| input_lower.contains("update")
|| input_lower.contains("exec")
|| input.contains("--")
|| input.contains("/*")
}
fn is_xss_pattern(&self, input: &str) -> bool {
let input_lower = input.to_lowercase();
input_lower.contains("<script")
|| input_lower.contains("javascript:")
|| input_lower.contains("onerror")
|| input_lower.contains("onload")
|| input_lower.contains("<iframe")
}
}
impl Default for VulnerabilityScanner {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ScanSummary {
pub tests_run: usize,
pub vulnerabilities_found: usize,
pub critical_findings: usize,
pub high_findings: usize,
pub medium_findings: usize,
pub low_findings: usize,
}
impl ScanSummary {
pub fn is_clean(&self) -> bool {
self.vulnerabilities_found == 0
}
pub fn risk_score(&self) -> u32 {
(self.critical_findings * 25)
.min(100)
.saturating_add((self.high_findings * 10).min(100))
.saturating_add((self.medium_findings * 5).min(100))
.saturating_add(self.low_findings.min(100))
.min(100) as u32
}
}
pub struct PentestSuite {
pub name: String,
scanners: Vec<VulnerabilityScanner>,
start_time: Option<Instant>,
}
impl PentestSuite {
pub fn new(name: &str) -> Self {
Self {
name: name.to_string(),
scanners: Vec::new(),
start_time: None,
}
}
pub fn add_scanner(&mut self, scanner: VulnerabilityScanner) {
self.scanners.push(scanner);
}
pub fn run_all(&mut self) -> PentestReport {
self.start_time = Some(Instant::now());
let mut all_findings = Vec::new();
let mut total_tests = 0;
let mut total_vulns = 0;
for scanner in &self.scanners {
all_findings.extend(scanner.get_findings().iter().cloned());
total_tests += scanner.tests_run;
total_vulns += scanner.vulnerabilities_found;
}
let duration_ms = self
.start_time
.map(|s| s.elapsed().as_millis() as u64)
.unwrap_or(0);
PentestReport {
suite_name: self.name.clone(),
total_tests,
total_vulnerabilities: total_vulns,
findings: all_findings,
duration_ms,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct PentestReport {
pub suite_name: String,
pub total_tests: usize,
pub total_vulnerabilities: usize,
pub findings: Vec<VulnerabilityFinding>,
pub duration_ms: u64,
}
impl PentestReport {
pub fn summary(&self) -> ScanSummary {
let critical = self
.findings
.iter()
.filter(|f| f.severity == Severity::Critical)
.count();
let high = self
.findings
.iter()
.filter(|f| f.severity == Severity::High)
.count();
let medium = self
.findings
.iter()
.filter(|f| f.severity == Severity::Medium)
.count();
let low = self
.findings
.iter()
.filter(|f| f.severity == Severity::Low)
.count();
ScanSummary {
tests_run: self.total_tests,
vulnerabilities_found: self.total_vulnerabilities,
critical_findings: critical,
high_findings: high,
medium_findings: medium,
low_findings: low,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_vulnerability_scanner_creation() {
let scanner = VulnerabilityScanner::new();
assert_eq!(scanner.tests_run, 0);
assert_eq!(scanner.vulnerabilities_found, 0);
}
#[test]
fn test_sql_injection_detection() {
let mut scanner = VulnerabilityScanner::new();
let query = "SELECT * FROM users WHERE id = ";
let vulnerable = scanner.scan_sql_injection(query, "1' OR '1'='1");
assert!(vulnerable);
assert_eq!(scanner.vulnerabilities_found, 1);
}
#[test]
fn test_xss_detection() {
let mut scanner = VulnerabilityScanner::new();
let output = "<div><script>alert('XSS')</script></div>";
let vulnerable = scanner.scan_xss(output, "");
assert!(vulnerable);
assert_eq!(scanner.vulnerabilities_found, 1);
}
#[test]
fn test_auth_bypass_detection() {
let mut scanner = VulnerabilityScanner::new();
let attempts = vec![("admin' --", "password")];
let vulnerable = scanner.test_auth_bypass("/login", &attempts);
assert!(vulnerable);
assert_eq!(scanner.vulnerabilities_found, 1);
}
#[test]
fn test_rate_limit_bypass_detection() {
let mut scanner = VulnerabilityScanner::new();
let vulnerable = scanner.test_rate_limit("/api/login", 500);
assert!(vulnerable);
assert_eq!(scanner.vulnerabilities_found, 1);
}
#[test]
fn test_path_traversal_detection() {
let mut scanner = VulnerabilityScanner::new();
let vulnerable = scanner.test_path_traversal("/files/", "../../../etc/passwd");
assert!(vulnerable);
assert_eq!(scanner.vulnerabilities_found, 1);
}
#[test]
fn test_get_findings_by_severity() {
let mut scanner = VulnerabilityScanner::new();
scanner.scan_sql_injection("SELECT * FROM users WHERE id = ", "1' OR '1'='1");
let critical_findings = scanner.get_findings_by_severity(Severity::Critical);
assert_eq!(critical_findings.len(), 1);
}
#[test]
fn test_scan_summary() {
let mut scanner = VulnerabilityScanner::new();
scanner.scan_sql_injection("SELECT * FROM users WHERE id = ", "1' OR '1'='1");
let summary = scanner.get_summary();
assert_eq!(summary.vulnerabilities_found, 1);
assert_eq!(summary.critical_findings, 1);
}
#[test]
fn test_risk_score() {
let summary = ScanSummary {
tests_run: 10,
vulnerabilities_found: 2,
critical_findings: 1,
high_findings: 1,
medium_findings: 0,
low_findings: 0,
};
let risk_score = summary.risk_score();
assert!(risk_score > 0);
}
#[test]
fn test_pentest_suite() {
let mut suite = PentestSuite::new("Security Test Suite");
let mut scanner = VulnerabilityScanner::new();
scanner.scan_sql_injection("SELECT * FROM users WHERE id = ", "1' OR '1'='1");
suite.add_scanner(scanner);
let report = suite.run_all();
assert_eq!(report.total_vulnerabilities, 1);
}
}