use std::sync::Arc;
use anyhow::{Context, Result};
use tokio::sync::RwLock;
use crate::cache::{CacheCategory, CacheKey, CacheManager};
use crate::scanner::Finding;
use super::config::{AiConfig, ExplanationContext};
use super::prompt::PROMPT_VERSION;
use super::provider::{create_provider, AiProvider};
use super::rate_limit::RateLimiter;
use super::response::ExplanationResponse;
pub struct ExplainEngine {
provider: Arc<dyn AiProvider>,
cache: Option<Arc<CacheManager>>,
rate_limiter: Arc<RateLimiter>,
default_context: ExplanationContext,
stats: Arc<RwLock<EngineStats>>,
}
#[derive(Debug, Default, Clone)]
pub struct EngineStats {
pub total_explanations: u64,
pub cache_hits: u64,
pub cache_misses: u64,
pub api_calls: u64,
pub tokens_used: u64,
pub total_response_time_ms: u64,
}
impl EngineStats {
pub fn cache_hit_rate(&self) -> f64 {
let total = self.cache_hits + self.cache_misses;
if total == 0 {
0.0
} else {
(self.cache_hits as f64 / total as f64) * 100.0
}
}
pub fn avg_response_time_ms(&self) -> u64 {
if self.api_calls == 0 {
0
} else {
self.total_response_time_ms / self.api_calls
}
}
}
impl ExplainEngine {
pub fn new(config: AiConfig) -> Result<Self> {
let provider = create_provider(&config)?;
let rate_limiter = Arc::new(RateLimiter::new(
config.rate_limit_rpm,
config.rate_limit_tpm,
));
Ok(Self {
provider: Arc::from(provider),
cache: None,
rate_limiter,
default_context: ExplanationContext::default(),
stats: Arc::new(RwLock::new(EngineStats::default())),
})
}
pub fn with_cache(mut self, cache: Arc<CacheManager>) -> Self {
self.cache = Some(cache);
self
}
pub fn with_default_context(mut self, context: ExplanationContext) -> Self {
self.default_context = context;
self
}
pub fn with_rate_limiter(mut self, rate_limiter: Arc<RateLimiter>) -> Self {
self.rate_limiter = rate_limiter;
self
}
pub fn provider_name(&self) -> &'static str {
self.provider.name()
}
pub fn model_name(&self) -> &str {
self.provider.model()
}
fn cache_key(&self, finding: &Finding, context: &ExplanationContext) -> CacheKey {
let key_data = format!(
"{}:{}:{}:{}:{}",
PROMPT_VERSION,
self.provider.model(),
finding.id,
finding.rule_id,
context.audience.description()
);
CacheKey::new(CacheCategory::AiResponse, &key_data)
}
async fn get_cached(
&self,
finding: &Finding,
context: &ExplanationContext,
) -> Option<ExplanationResponse> {
let cache = self.cache.as_ref()?;
let key = self.cache_key(finding, context);
match cache.get::<ExplanationResponse>(&key).await {
Ok(Some(response)) => {
let mut stats = self.stats.write().await;
stats.cache_hits += 1;
Some(response)
}
_ => {
let mut stats = self.stats.write().await;
stats.cache_misses += 1;
None
}
}
}
async fn store_cached(
&self,
finding: &Finding,
context: &ExplanationContext,
response: &ExplanationResponse,
) {
if let Some(cache) = &self.cache {
let key = self.cache_key(finding, context);
if let Err(e) = cache.set(&key, response).await {
tracing::warn!("Failed to cache explanation: {}", e);
}
}
}
pub async fn explain(&self, finding: &Finding) -> Result<ExplanationResponse> {
self.explain_with_context(finding, &self.default_context)
.await
}
pub async fn explain_with_context(
&self,
finding: &Finding,
context: &ExplanationContext,
) -> Result<ExplanationResponse> {
if let Some(cached) = self.get_cached(finding, context).await {
tracing::debug!("Cache hit for finding {}", finding.id);
return Ok(cached);
}
let estimated_tokens = estimate_tokens(finding);
self.rate_limiter
.acquire(estimated_tokens)
.await
.context("Rate limit acquisition failed")?;
tracing::debug!("Generating explanation for finding {}", finding.id);
let response = self
.provider
.explain_finding(finding, context)
.await
.context("Failed to generate explanation")?;
{
let mut stats = self.stats.write().await;
stats.total_explanations += 1;
stats.api_calls += 1;
if response.metadata.tokens_used > 0 {
stats.tokens_used += response.metadata.tokens_used as u64;
self.rate_limiter
.record_tokens(response.metadata.tokens_used, estimated_tokens)
.await;
}
if response.metadata.response_time_ms > 0 {
stats.total_response_time_ms += response.metadata.response_time_ms;
}
}
self.store_cached(finding, context, &response).await;
Ok(response)
}
pub async fn explain_batch(
&self,
findings: &[Finding],
context: Option<&ExplanationContext>,
) -> Result<Vec<ExplanationResponse>> {
let ctx = context.unwrap_or(&self.default_context);
let mut results = Vec::with_capacity(findings.len());
let mut uncached_findings = Vec::new();
let mut uncached_indices = Vec::new();
for (idx, finding) in findings.iter().enumerate() {
if let Some(cached) = self.get_cached(finding, ctx).await {
results.push((idx, cached));
} else {
uncached_findings.push(finding.clone());
uncached_indices.push(idx);
}
}
if !uncached_findings.is_empty() {
tracing::info!(
"Generating {} explanations ({} cached)",
uncached_findings.len(),
results.len()
);
for (i, finding) in uncached_findings.iter().enumerate() {
let response = self.explain_with_context(finding, ctx).await?;
results.push((uncached_indices[i], response));
}
}
results.sort_by_key(|(idx, _)| *idx);
Ok(results.into_iter().map(|(_, r)| r).collect())
}
pub async fn ask_followup(
&self,
explanation: &ExplanationResponse,
question: &str,
) -> Result<String> {
let estimated_tokens = 500 + (question.len() as u32 / 4);
self.rate_limiter
.acquire(estimated_tokens)
.await
.context("Rate limit acquisition failed")?;
let response = self
.provider
.ask_followup(explanation, question)
.await
.context("Failed to get follow-up response")?;
{
let mut stats = self.stats.write().await;
stats.api_calls += 1;
}
Ok(response)
}
pub async fn health_check(&self) -> Result<bool> {
self.provider.health_check().await
}
pub async fn stats(&self) -> EngineStats {
self.stats.read().await.clone()
}
pub async fn rate_limit_stats(&self) -> super::rate_limit::RateLimitStats {
self.rate_limiter.stats().await
}
pub async fn would_exceed_rate_limit(&self, estimated_tokens: u32) -> bool {
self.rate_limiter.would_exceed(estimated_tokens).await
}
}
fn estimate_tokens(finding: &Finding) -> u32 {
let base = 500;
let finding_tokens = (finding.title.len()
+ finding.description.len()
+ finding.evidence.iter().map(|e| e.data.len()).sum::<usize>())
/ 4;
let response_estimate = 2000;
base + finding_tokens as u32 + response_estimate
}
#[cfg(test)]
mod tests {
use super::*;
use crate::ai::config::AiProvider as AiProviderType;
use crate::ai::provider::MockProvider;
use crate::scanner::Severity;
use std::sync::Arc;
fn sample_finding() -> Finding {
Finding::new(
"MCP-TEST-001",
Severity::High,
"Test Finding",
"This is a test finding for unit tests",
)
}
fn sample_finding_with_evidence() -> Finding {
use crate::scanner::{Evidence, EvidenceKind};
sample_finding()
.with_evidence(Evidence::new(
EvidenceKind::Observation,
"let x = user_input;",
"Vulnerable code pattern",
))
.with_evidence(Evidence::new(
EvidenceKind::Observation,
"src/main.rs:42",
"Location of issue",
))
}
#[test]
fn engine_stats_calculations() {
let mut stats = EngineStats::default();
assert_eq!(stats.cache_hit_rate(), 0.0);
stats.cache_hits = 3;
stats.cache_misses = 7;
assert!((stats.cache_hit_rate() - 30.0).abs() < 0.01);
assert_eq!(stats.avg_response_time_ms(), 0);
stats.api_calls = 5;
stats.total_response_time_ms = 1000;
assert_eq!(stats.avg_response_time_ms(), 200);
}
#[test]
fn engine_stats_edge_cases() {
let stats = EngineStats::default();
assert_eq!(stats.cache_hit_rate(), 0.0);
assert_eq!(stats.avg_response_time_ms(), 0);
let stats = EngineStats {
cache_hits: 10,
cache_misses: 0,
..Default::default()
};
assert_eq!(stats.cache_hit_rate(), 100.0);
let stats = EngineStats {
cache_hits: 0,
cache_misses: 10,
..Default::default()
};
assert_eq!(stats.cache_hit_rate(), 0.0);
}
#[test]
fn token_estimation() {
let finding = sample_finding();
let tokens = estimate_tokens(&finding);
assert!(tokens > 500);
assert!(tokens < 10000);
}
#[test]
fn token_estimation_with_evidence() {
let finding = sample_finding_with_evidence();
let tokens = estimate_tokens(&finding);
let base_tokens = estimate_tokens(&sample_finding());
assert!(tokens >= base_tokens);
}
#[test]
fn token_estimation_scales_with_content() {
let small_finding = Finding::new("ID1", Severity::Low, "Short", "Brief");
let large_finding = Finding::new(
"ID2",
Severity::High,
"Very Long Title That Contains Many Words And Details",
"This is a much longer description that contains significantly more text \
and detailed information about the vulnerability, including technical details, \
impact assessment, and comprehensive analysis of the security issue.",
);
let small_tokens = estimate_tokens(&small_finding);
let large_tokens = estimate_tokens(&large_finding);
assert!(large_tokens > small_tokens);
}
#[tokio::test]
async fn engine_creation_with_ollama() {
let config = AiConfig::builder()
.provider(AiProviderType::Ollama)
.model("llama3.2")
.build();
let engine = ExplainEngine::new(config);
assert!(engine.is_ok());
let engine = engine.unwrap();
assert_eq!(engine.provider_name(), "Ollama");
assert_eq!(engine.model_name(), "llama3.2");
}
#[tokio::test]
async fn engine_builder_pattern() {
use crate::cache::{CacheBackend, CacheConfig};
let config = AiConfig::builder()
.provider(AiProviderType::Ollama)
.model("llama3.2")
.build();
let cache_config = CacheConfig {
backend: CacheBackend::Memory,
schema_ttl_secs: 3600,
result_ttl_secs: 86400,
validation_ttl_secs: 3600,
corpus_persist: false,
max_size_bytes: None,
enabled: true,
};
let cache_manager = Arc::new(CacheManager::new(cache_config).await.unwrap());
let custom_context = ExplanationContext::default();
let rate_limiter = Arc::new(RateLimiter::new(100, 10000));
let engine = ExplainEngine::new(config)
.unwrap()
.with_cache(cache_manager)
.with_default_context(custom_context)
.with_rate_limiter(rate_limiter);
assert_eq!(engine.provider_name(), "Ollama");
}
#[tokio::test]
async fn engine_stats_tracking() {
let config = AiConfig::builder()
.provider(AiProviderType::Ollama)
.model("llama3.2")
.build();
let engine = ExplainEngine::new(config).unwrap();
let stats = engine.stats().await;
assert_eq!(stats.total_explanations, 0);
assert_eq!(stats.api_calls, 0);
assert_eq!(stats.tokens_used, 0);
}
#[test]
fn cache_key_generation() {
let config = AiConfig::builder()
.provider(AiProviderType::Ollama)
.model("llama3.2")
.build();
let engine = ExplainEngine::new(config).unwrap();
let finding = sample_finding();
let context = ExplanationContext::default();
let key = engine.cache_key(&finding, &context);
assert!(key.to_string().contains("llama3.2"));
assert!(key.to_string().contains(&finding.id));
}
#[test]
fn cache_key_differs_for_different_findings() {
let config = AiConfig::builder()
.provider(AiProviderType::Ollama)
.model("llama3.2")
.build();
let engine = ExplainEngine::new(config).unwrap();
let finding1 = sample_finding();
let finding2 = Finding::new(
"MCP-TEST-002",
Severity::Medium,
"Different Finding",
"Different description",
);
let context = ExplanationContext::default();
let key1 = engine.cache_key(&finding1, &context);
let key2 = engine.cache_key(&finding2, &context);
assert_ne!(key1, key2);
}
#[test]
fn cache_key_includes_prompt_version() {
let config = AiConfig::builder()
.provider(AiProviderType::Ollama)
.model("llama3.2")
.build();
let engine = ExplainEngine::new(config).unwrap();
let finding = sample_finding();
let context = ExplanationContext::default();
let key = engine.cache_key(&finding, &context);
assert!(key.to_string().contains(PROMPT_VERSION));
}
#[tokio::test]
async fn rate_limit_stats() {
let config = AiConfig::builder()
.provider(AiProviderType::Ollama)
.model("llama3.2")
.build();
let engine = ExplainEngine::new(config).unwrap();
let stats = engine.rate_limit_stats().await;
assert!(stats.requests_limit > 0);
assert_eq!(stats.requests_used, 0);
}
#[tokio::test]
async fn would_exceed_rate_limit() {
let config = AiConfig::builder()
.provider(AiProviderType::Ollama)
.model("llama3.2")
.build();
let rate_limiter = Arc::new(RateLimiter::new(10, 1000));
let engine = ExplainEngine::new(config)
.unwrap()
.with_rate_limiter(rate_limiter);
let would_exceed = engine.would_exceed_rate_limit(100).await;
assert!(!would_exceed);
let would_exceed = engine.would_exceed_rate_limit(1_000_000).await;
assert!(would_exceed);
}
#[tokio::test]
async fn engine_getters() {
let config = AiConfig::builder()
.provider(AiProviderType::Ollama)
.model("test-model")
.build();
let engine = ExplainEngine::new(config).unwrap();
assert_eq!(engine.provider_name(), "Ollama");
assert_eq!(engine.model_name(), "test-model");
}
#[tokio::test]
async fn explain_with_mock_provider() {
let provider = Arc::new(MockProvider::new()) as Arc<dyn AiProvider>;
let rate_limiter = Arc::new(RateLimiter::new(1000, 100000));
let engine = ExplainEngine {
provider,
cache: None,
rate_limiter,
default_context: ExplanationContext::default(),
stats: Arc::new(RwLock::new(EngineStats::default())),
};
let finding = sample_finding();
let result = engine.explain(&finding).await;
assert!(result.is_ok());
let explanation = result.unwrap();
assert_eq!(explanation.finding_id, finding.id);
let stats = engine.stats().await;
assert_eq!(stats.total_explanations, 1);
assert_eq!(stats.api_calls, 1);
}
#[tokio::test]
async fn explain_batch_ordering() {
let provider = Arc::new(MockProvider::new()) as Arc<dyn AiProvider>;
let rate_limiter = Arc::new(RateLimiter::new(1000, 100000));
let engine = ExplainEngine {
provider,
cache: None,
rate_limiter,
default_context: ExplanationContext::default(),
stats: Arc::new(RwLock::new(EngineStats::default())),
};
let finding1 = Finding::new("RULE-1", Severity::High, "First", "First finding");
let finding2 = Finding::new("RULE-2", Severity::Medium, "Second", "Second finding");
let finding3 = Finding::new("RULE-3", Severity::Low, "Third", "Third finding");
let findings = vec![finding1.clone(), finding2.clone(), finding3.clone()];
let result = engine.explain_batch(&findings, None).await;
assert!(result.is_ok());
let responses = result.unwrap();
assert_eq!(responses.len(), 3);
assert_eq!(responses[0].finding_id, finding1.id);
assert_eq!(responses[1].finding_id, finding2.id);
assert_eq!(responses[2].finding_id, finding3.id);
}
#[tokio::test]
async fn explain_with_custom_context() {
let provider = Arc::new(MockProvider::new()) as Arc<dyn AiProvider>;
let rate_limiter = Arc::new(RateLimiter::new(1000, 100000));
let engine = ExplainEngine {
provider,
cache: None,
rate_limiter,
default_context: ExplanationContext::default(),
stats: Arc::new(RwLock::new(EngineStats::default())),
};
let finding = sample_finding();
let custom_context = ExplanationContext::default();
let result = engine.explain_with_context(&finding, &custom_context).await;
assert!(result.is_ok());
}
#[tokio::test]
async fn ask_followup_updates_stats() {
let provider = Arc::new(MockProvider::new()) as Arc<dyn AiProvider>;
let rate_limiter = Arc::new(RateLimiter::new(1000, 100000));
let engine = ExplainEngine {
provider,
cache: None,
rate_limiter,
default_context: ExplanationContext::default(),
stats: Arc::new(RwLock::new(EngineStats::default())),
};
let explanation = ExplanationResponse::new("test", "TEST-001");
let result = engine.ask_followup(&explanation, "How to fix?").await;
assert!(result.is_ok());
let stats = engine.stats().await;
assert_eq!(stats.api_calls, 1);
}
#[tokio::test]
async fn health_check() {
let provider = Arc::new(MockProvider::new()) as Arc<dyn AiProvider>;
let rate_limiter = Arc::new(RateLimiter::new(1000, 100000));
let engine = ExplainEngine {
provider,
cache: None,
rate_limiter,
default_context: ExplanationContext::default(),
stats: Arc::new(RwLock::new(EngineStats::default())),
};
let result = engine.health_check().await;
assert!(result.is_ok());
assert!(result.unwrap());
}
}