use crate::{Error, Result, ScanResult, Vault};
use async_trait::async_trait;
use std::sync::Arc;
#[async_trait]
pub trait Scanner: Send + Sync {
fn name(&self) -> &str;
async fn scan(&self, input: &str, vault: &Vault) -> Result<ScanResult>;
fn scanner_type(&self) -> ScannerType {
ScannerType::Input
}
fn version(&self) -> &str {
"1.0.0"
}
fn description(&self) -> &str {
"No description provided"
}
fn requires_async(&self) -> bool {
false
}
fn validate_config(&self) -> Result<()> {
Ok(())
}
}
#[async_trait]
pub trait InputScanner: Scanner {
async fn scan_prompt(&self, prompt: &str, vault: &Vault) -> Result<ScanResult> {
self.scan(prompt, vault).await
}
}
#[async_trait]
pub trait OutputScanner: Scanner {
async fn scan_output(
&self,
prompt: &str,
output: &str,
vault: &Vault,
) -> Result<ScanResult>;
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ScannerType {
Input,
Output,
Bidirectional,
}
pub struct ScannerPipeline {
scanners: Vec<Arc<dyn Scanner>>,
short_circuit: bool,
short_circuit_threshold: f32,
}
impl ScannerPipeline {
pub fn new() -> Self {
Self {
scanners: Vec::new(),
short_circuit: false,
short_circuit_threshold: 0.9,
}
}
pub fn add(mut self, scanner: Arc<dyn Scanner>) -> Self {
self.scanners.push(scanner);
self
}
pub fn with_short_circuit(mut self, threshold: f32) -> Self {
self.short_circuit = true;
self.short_circuit_threshold = threshold;
self
}
pub async fn execute(&self, input: &str, vault: &Vault) -> Result<Vec<ScanResult>> {
let mut results = Vec::new();
for scanner in &self.scanners {
let result = scanner.scan(input, vault).await?;
if self.short_circuit && result.risk_score >= self.short_circuit_threshold {
results.push(result);
break;
}
results.push(result);
}
Ok(results)
}
pub async fn execute_parallel(&self, input: &str, vault: &Vault) -> Result<Vec<ScanResult>> {
use futures::future::join_all;
let futures: Vec<_> = self
.scanners
.iter()
.map(|scanner| {
let input = input.to_string();
let vault = vault.clone();
let scanner = Arc::clone(scanner);
async move { scanner.scan(&input, &vault).await }
})
.collect();
let results: Vec<Result<ScanResult>> = join_all(futures).await;
results.into_iter().collect()
}
pub async fn execute_aggregated(&self, input: &str, vault: &Vault) -> Result<ScanResult> {
let results = self.execute(input, vault).await?;
Ok(ScanResult::combine(results))
}
}
impl Default for ScannerPipeline {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
struct MockScanner {
name: String,
risk_score: f32,
}
#[async_trait]
impl Scanner for MockScanner {
fn name(&self) -> &str {
&self.name
}
async fn scan(&self, input: &str, _vault: &Vault) -> Result<ScanResult> {
Ok(ScanResult::new(
input.to_string(),
self.risk_score < 0.5,
self.risk_score,
))
}
fn scanner_type(&self) -> ScannerType {
ScannerType::Input
}
}
#[tokio::test]
async fn test_scanner_pipeline_sequential() {
let vault = Vault::new();
let scanner1 = Arc::new(MockScanner {
name: "test1".to_string(),
risk_score: 0.3,
});
let scanner2 = Arc::new(MockScanner {
name: "test2".to_string(),
risk_score: 0.5,
});
let pipeline = ScannerPipeline::new().add(scanner1).add(scanner2);
let results = pipeline.execute("test input", &vault).await.unwrap();
assert_eq!(results.len(), 2);
assert_eq!(results[0].risk_score, 0.3);
assert_eq!(results[1].risk_score, 0.5);
}
#[tokio::test]
async fn test_scanner_pipeline_short_circuit() {
let vault = Vault::new();
let scanner1 = Arc::new(MockScanner {
name: "test1".to_string(),
risk_score: 0.95,
});
let scanner2 = Arc::new(MockScanner {
name: "test2".to_string(),
risk_score: 0.2,
});
let pipeline = ScannerPipeline::new()
.add(scanner1)
.add(scanner2)
.with_short_circuit(0.9);
let results = pipeline.execute("test input", &vault).await.unwrap();
assert_eq!(results.len(), 1);
assert_eq!(results[0].risk_score, 0.95);
}
#[tokio::test]
async fn test_scanner_pipeline_aggregated() {
let vault = Vault::new();
let scanner1 = Arc::new(MockScanner {
name: "test1".to_string(),
risk_score: 0.3,
});
let scanner2 = Arc::new(MockScanner {
name: "test2".to_string(),
risk_score: 0.7,
});
let pipeline = ScannerPipeline::new().add(scanner1).add(scanner2);
let result = pipeline
.execute_aggregated("test input", &vault)
.await
.unwrap();
assert_eq!(result.risk_score, 0.7);
}
}