use crate::scanner::ThreatRules;
use serde_json::Value;
fn render(entry: &Value) -> String {
let kind = entry.get("kind").and_then(Value::as_str).unwrap_or("tool");
let name = entry.get("name").and_then(Value::as_str).unwrap_or("");
let description = entry.get("description").and_then(Value::as_str);
let mut text = format!("{}: {}\n", kind.to_uppercase(), name);
if let Some(description) = description {
text.push_str(&format!("DESCRIPTION: {description}\n"));
}
if let Some(uri) = entry.get("uri").and_then(Value::as_str) {
text.push_str(&format!("URI: {uri}\n"));
}
if let Some(mime) = entry.get("mime_type").and_then(Value::as_str) {
text.push_str(&format!("MIME_TYPE: {mime}\n"));
}
if let Some(schema) = entry.get("input_schema") {
if !schema.is_null() {
text.push_str(&format!("INPUT_SCHEMA: {schema}\n"));
}
}
if let Some(arguments) = entry.get("arguments").and_then(Value::as_array) {
for argument in arguments {
let arg_name = argument.get("name").and_then(Value::as_str).unwrap_or("");
let arg_desc = argument
.get("description")
.and_then(Value::as_str)
.unwrap_or("");
text.push_str(&format!("ARGUMENT: {arg_name} — {arg_desc}\n"));
}
}
text
}
fn scan_all_views(engine: &ThreatRules, text: &str) -> Vec<String> {
let mut rules: Vec<String> = engine
.pre_scan(text, "corpus")
.into_iter()
.map(|h| h.rule_name)
.collect();
for view in crate::normalize::additional_scan_views(text) {
for h in engine.pre_scan(&view, "corpus") {
rules.push(h.rule_name);
}
}
rules.sort();
rules.dedup();
rules
}
fn load(path: &str) -> Vec<Value> {
let raw = std::fs::read_to_string(path)
.unwrap_or_else(|e| panic!("corpus {path} must be readable: {e}"));
serde_json::from_str(&raw).unwrap_or_else(|e| panic!("corpus {path} must be valid JSON: {e}"))
}
struct Outcome {
benign_total: usize,
benign_flagged: Vec<(String, String, Vec<String>)>,
malicious_total: usize,
malicious_missed: Vec<(String, String)>,
}
fn evaluate() -> Outcome {
let engine = ThreatRules::new(
&crate::scanner::resolve_rules_dir()
.expect("rules directory must resolve when running the evaluation")
.to_string_lossy(),
)
.expect("rules must compile");
let benign = load("tests/corpus/benign.json");
let malicious = load("tests/corpus/malicious.json");
let mut benign_flagged = Vec::new();
for entry in &benign {
let rules = scan_all_views(&engine, &render(entry));
if !rules.is_empty() {
benign_flagged.push((
entry["origin"].as_str().unwrap_or("?").to_string(),
entry["name"].as_str().unwrap_or("?").to_string(),
rules,
));
}
}
let mut malicious_missed = Vec::new();
for entry in &malicious {
let rules = scan_all_views(&engine, &render(entry));
if rules.is_empty() {
malicious_missed.push((
entry["id"].as_str().unwrap_or("?").to_string(),
entry["attack"].as_str().unwrap_or("?").to_string(),
));
}
}
Outcome {
benign_total: benign.len(),
benign_flagged,
malicious_total: malicious.len(),
malicious_missed,
}
}
#[test]
fn report_rule_quality() {
let outcome = evaluate();
let fp = outcome.benign_flagged.len();
let tp = outcome.malicious_total - outcome.malicious_missed.len();
println!("\n=== RULE QUALITY ===");
println!(
"false positives : {fp}/{} benign items ({:.0}%)",
outcome.benign_total,
100.0 * fp as f64 / outcome.benign_total as f64
);
println!(
"true positives : {tp}/{} malicious items ({:.0}%)",
outcome.malicious_total,
100.0 * tp as f64 / outcome.malicious_total as f64
);
if !outcome.benign_flagged.is_empty() {
println!("\nfalse positives:");
for (origin, name, rules) in &outcome.benign_flagged {
println!(" [{origin}] {name} -> {rules:?}");
}
}
if !outcome.malicious_missed.is_empty() {
println!("\nmissed attacks:");
for (id, attack) in &outcome.malicious_missed {
println!(" {id} ({attack})");
}
}
println!();
}
#[test]
fn benign_servers_produce_no_findings() {
let outcome = evaluate();
assert!(
outcome.benign_flagged.is_empty(),
"{} of {} benign items flagged; every one is a false positive:\n{}",
outcome.benign_flagged.len(),
outcome.benign_total,
outcome
.benign_flagged
.iter()
.map(|(o, n, r)| format!(" [{o}] {n} -> {r:?}"))
.collect::<Vec<_>>()
.join("\n")
);
}
#[test]
fn known_attacks_are_all_detected() {
let outcome = evaluate();
assert!(
outcome.malicious_missed.is_empty(),
"{} of {} attacks missed:\n{}",
outcome.malicious_missed.len(),
outcome.malicious_total,
outcome
.malicious_missed
.iter()
.map(|(id, attack)| format!(" {id} ({attack})"))
.collect::<Vec<_>>()
.join("\n")
);
}