#![allow(dead_code)]
use super::engine::{EvaluationContext, PolicyDecision, PolicyEngine};
use super::model::*;
use super::parser::validate_policy;
use crate::core::api::types::RunRequest;
use anyhow::Result;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct PolicyTestScenario {
pub name: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub description: Option<String>,
pub request: TestRequest,
pub url: String,
#[serde(default)]
pub context: TestContext,
pub expected: ExpectedOutcome,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TestRequest {
pub operation_id: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub auth_profile: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub env: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub parameters: Option<HashMap<String, serde_json::Value>>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct TestContext {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub method: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub tags: Option<Vec<String>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub source_ip: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ExpectedOutcome {
pub decision: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub rule: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub reason_contains: Option<String>,
}
#[derive(Debug)]
pub struct TestResult {
pub scenario_name: String,
pub passed: bool,
pub message: String,
pub actual_decision: PolicyDecision,
}
pub struct PolicyTestRunner {
engine: PolicyEngine,
}
impl PolicyTestRunner {
pub fn new(policy: PolicySet) -> Result<Self> {
validate_policy(&policy)?;
let engine = PolicyEngine::new(policy)?;
Ok(Self { engine })
}
pub fn run_scenario(&self, scenario: &PolicyTestScenario) -> TestResult {
let request = RunRequest {
operation_id: scenario.request.operation_id.clone(),
auth_profile: scenario.request.auth_profile.clone(),
env: scenario.request.env.clone(),
parameters: scenario.request.parameters.clone(),
body: None,
spec_path: None,
};
let context = EvaluationContext {
method: scenario.context.method.clone(),
tags: scenario.context.tags.clone(),
source_ip: scenario.context.source_ip.clone(),
timestamp: chrono::Utc::now(),
};
let decision = self.engine.evaluate(&request, &scenario.url, &context);
let (passed, message) = self.check_outcome(&decision, &scenario.expected);
TestResult {
scenario_name: scenario.name.clone(),
passed,
message,
actual_decision: decision,
}
}
pub fn run_scenarios(&self, scenarios: &[PolicyTestScenario]) -> Vec<TestResult> {
scenarios
.iter()
.map(|scenario| self.run_scenario(scenario))
.collect()
}
fn check_outcome(
&self,
decision: &PolicyDecision,
expected: &ExpectedOutcome,
) -> (bool, String) {
let actual_decision = if decision.is_allowed() {
"allow"
} else {
"deny"
};
if actual_decision != expected.decision.to_lowercase() {
return (
false,
format!("Expected {} but got {}", expected.decision, actual_decision),
);
}
if let Some(expected_rule) = &expected.rule {
if decision.rule_name() != expected_rule {
return (
false,
format!(
"Expected rule '{}' but got '{}'",
expected_rule,
decision.rule_name()
),
);
}
}
if let PolicyDecision::Deny { reason, .. } = decision {
if let Some(expected_pattern) = &expected.reason_contains {
if !reason.contains(expected_pattern) {
return (
false,
format!(
"Expected reason to contain '{}' but got '{}'",
expected_pattern, reason
),
);
}
}
}
(true, "Test passed".to_string())
}
}
pub fn load_test_scenarios(content: &str) -> Result<Vec<PolicyTestScenario>> {
let scenarios: Vec<PolicyTestScenario> = serde_yaml::from_str(content)?;
Ok(scenarios)
}
pub fn generate_test_report(results: &[TestResult]) -> String {
let mut report = String::new();
let total = results.len();
let passed = results.iter().filter(|r| r.passed).count();
let failed = total - passed;
report.push_str(&format!("Policy Test Results\n"));
report.push_str(&format!("==================\n\n"));
report.push_str(&format!(
"Total: {} | Passed: {} | Failed: {}\n\n",
total, passed, failed
));
if failed > 0 {
report.push_str("FAILED TESTS:\n");
for result in results.iter().filter(|r| !r.passed) {
report.push_str(&format!(
" ❌ {} - {}\n",
result.scenario_name, result.message
));
}
report.push_str("\n");
}
if passed > 0 {
report.push_str("PASSED TESTS:\n");
for result in results.iter().filter(|r| r.passed) {
report.push_str(&format!(" ✅ {}\n", result.scenario_name));
}
}
report
}
#[cfg(test)]
mod tests {
use super::*;
fn create_test_policy() -> PolicySet {
PolicySet {
version: "1.0".to_string(),
metadata: None,
defaults: PolicyDefaults {
allow_methods: vec!["GET".to_string()],
deny_external_refs: true,
require_auth: true,
audit_level: "basic".to_string(),
default_classification: DataClassification::Internal,
read_only: false,
max_calls_per_session: 0,
max_calls_per_minute: 0,
warn_on_confidential_to_llm: false,
block_regulated_to_llm: false,
},
classifications: std::collections::BTreeMap::new(),
credential_scopes: vec![],
rules: vec![PolicyRule {
name: "public-read".to_string(),
description: None,
pattern: "api.public.com/*".to_string(),
conditions: None,
allow: Some(PolicyAction {
methods: Some(vec!["GET".to_string()]),
operations: None,
all: None,
tags: None,
}),
deny: None,
audit: None,
explain: None,
}],
}
}
#[test]
fn test_scenario_pass() {
let policy = create_test_policy();
let runner = PolicyTestRunner::new(policy).unwrap();
let scenario = PolicyTestScenario {
name: "Allow public GET".to_string(),
description: None,
request: TestRequest {
operation_id: "getResource".to_string(),
auth_profile: Some("default".to_string()),
env: None,
parameters: None,
},
url: "api.public.com/resource".to_string(),
context: TestContext {
method: Some("GET".to_string()),
tags: None,
source_ip: None,
},
expected: ExpectedOutcome {
decision: "allow".to_string(),
rule: Some("public-read".to_string()),
reason_contains: None,
},
};
let result = runner.run_scenario(&scenario);
assert!(result.passed);
}
#[test]
fn test_scenario_fail_wrong_method() {
let policy = create_test_policy();
let runner = PolicyTestRunner::new(policy).unwrap();
let scenario = PolicyTestScenario {
name: "Deny POST to public".to_string(),
description: None,
request: TestRequest {
operation_id: "createResource".to_string(),
auth_profile: Some("default".to_string()),
env: None,
parameters: None,
},
url: "api.public.com/resource".to_string(),
context: TestContext {
method: Some("POST".to_string()),
tags: None,
source_ip: None,
},
expected: ExpectedOutcome {
decision: "deny".to_string(),
rule: None,
reason_contains: Some("Method POST not allowed".to_string()),
},
};
let result = runner.run_scenario(&scenario);
assert!(result.passed);
}
#[test]
fn test_load_scenarios_from_yaml() {
let yaml = r#"
- name: "Test scenario 1"
request:
operation_id: "getUser"
auth_profile: "default"
url: "https://api.example.com/users/123"
context:
method: "GET"
expected:
decision: "allow"
"#;
let scenarios = load_test_scenarios(yaml).unwrap();
assert_eq!(scenarios.len(), 1);
assert_eq!(scenarios[0].name, "Test scenario 1");
}
}