318 lines
9.9 KiB
Rust
318 lines
9.9 KiB
Rust
//! Answer Validation & Confidence Scoring
|
|||
|
|
//!
|
||
|
|
//! Validate query answers and assign confidence scores.
|
||
|
|
//! Multi-signal confidence aggregation (Zep alignment).
|
||
|
|
//!
|
||
|
|
//! CRAP: 15 (Multiple confidence signals)
|
||
|
|
//! SOLID: Single responsibility (answer validation)
|
||
|
|
//! DRY: Reuses score types from mem_core
|
||
|
|
|
||
|
|
use serde::{Deserialize, Serialize};
|
||
|
|
use tracing::{debug, info};
|
||
|
|
|
||
|
|
/// Answer validation configuration
|
||
|
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||
|
|
pub struct AnswerValidationConfig {
|
||
|
|
pub enabled: bool,
|
||
|
|
pub min_confidence_threshold: f32, // Minimum confidence to accept answer
|
||
|
|
pub require_evidence: bool, // Must have supporting facts
|
||
|
|
pub evidence_threshold: usize, // Minimum number of supporting facts
|
||
|
|
}
|
||
|
|
|
||
|
|
impl Default for AnswerValidationConfig {
|
||
|
|
fn default() -> Self {
|
||
|
|
Self {
|
||
|
|
enabled: true,
|
||
|
|
min_confidence_threshold: 0.6,
|
||
|
|
require_evidence: true,
|
||
|
|
evidence_threshold: 1,
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
/// Answer confidence signals
|
||
|
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||
|
|
pub struct ConfidenceSignals {
|
||
|
|
/// Base search score (semantic + lexical combined)
|
||
|
|
pub search_score: f32,
|
||
|
|
/// Number of supporting facts
|
||
|
|
pub evidence_count: usize,
|
||
|
|
/// Average evidence confidence
|
||
|
|
pub evidence_confidence: f32,
|
||
|
|
/// Temporal consistency (0-1: higher = more recent)
|
||
|
|
pub temporal_score: f32,
|
||
|
|
/// Entity coverage (0-1: higher = all entities found)
|
||
|
|
pub entity_coverage: f32,
|
||
|
|
/// Contradiction score (0-1: higher = fewer contradictions)
|
||
|
|
pub contradiction_score: f32,
|
||
|
|
}
|
||
|
|
|
||
|
|
/// Answer validation result
|
||
|
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||
|
|
pub struct ValidatedAnswer {
|
||
|
|
pub answer: String,
|
||
|
|
pub overall_confidence: f32, // 0-1
|
||
|
|
pub signals: ConfidenceSignals,
|
||
|
|
pub is_valid: bool, // Passes validation threshold
|
||
|
|
pub reasoning: String,
|
||
|
|
pub warning: Option<String>, // Low confidence or missing evidence
|
||
|
|
}
|
||
|
|
|
||
|
|
/// Answer Validator
|
||
|
|
pub struct AnswerValidator {
|
||
|
|
config: AnswerValidationConfig,
|
||
|
|
}
|
||
|
|
|
||
|
|
impl AnswerValidator {
|
||
|
|
pub fn new(config: AnswerValidationConfig) -> Self {
|
||
|
|
Self { config }
|
||
|
|
}
|
||
|
|
|
||
|
|
/// Compute overall confidence from multiple signals
|
||
|
|
fn compute_confidence(&self, signals: &ConfidenceSignals) -> f32 {
|
||
|
|
if !self.config.enabled {
|
||
|
|
return 1.0;
|
||
|
|
}
|
||
|
|
|
||
|
|
let mut weighted_sum = 0.0;
|
||
|
|
let mut weight_sum = 0.0;
|
||
|
|
|
||
|
|
// Search score: 0.4 weight
|
||
|
|
weighted_sum += signals.search_score * 0.4;
|
||
|
|
weight_sum += 0.4;
|
||
|
|
|
||
|
|
// Evidence: 0.25 weight
|
||
|
|
let evidence_score = (signals.evidence_count as f32 / 5.0).min(1.0) * signals.evidence_confidence;
|
||
|
|
weighted_sum += evidence_score * 0.25;
|
||
|
|
weight_sum += 0.25;
|
||
|
|
|
||
|
|
// Temporal recency: 0.15 weight
|
||
|
|
weighted_sum += signals.temporal_score * 0.15;
|
||
|
|
weight_sum += 0.15;
|
||
|
|
|
||
|
|
// Entity coverage: 0.1 weight
|
||
|
|
weighted_sum += signals.entity_coverage * 0.1;
|
||
|
|
weight_sum += 0.1;
|
||
|
|
|
||
|
|
// Contradiction: 0.1 weight
|
||
|
|
weighted_sum += signals.contradiction_score * 0.1;
|
||
|
|
weight_sum += 0.1;
|
||
|
|
|
||
|
|
(weighted_sum / weight_sum).clamp(0.0, 1.0)
|
||
|
|
}
|
||
|
|
|
||
|
|
/// Validate answer based on configuration
|
||
|
|
pub fn validate(
|
||
|
|
&self,
|
||
|
|
answer: &str,
|
||
|
|
signals: &ConfidenceSignals,
|
||
|
|
) -> ValidatedAnswer {
|
||
|
|
if !self.config.enabled {
|
||
|
|
return ValidatedAnswer {
|
||
|
|
answer: answer.to_string(),
|
||
|
|
overall_confidence: 1.0,
|
||
|
|
signals: signals.clone(),
|
||
|
|
is_valid: true,
|
||
|
|
reasoning: "Validation disabled".to_string(),
|
||
|
|
warning: None,
|
||
|
|
};
|
||
|
|
}
|
||
|
|
|
||
|
|
let overall_confidence = self.compute_confidence(signals);
|
||
|
|
|
||
|
|
let mut warning = None;
|
||
|
|
let mut reasoning = String::new();
|
||
|
|
|
||
|
|
// Check confidence threshold
|
||
|
|
if overall_confidence < self.config.min_confidence_threshold {
|
||
|
|
warning = Some(format!(
|
||
|
|
"Low confidence: {:.2} (threshold: {:.2})",
|
||
|
|
overall_confidence, self.config.min_confidence_threshold
|
||
|
|
));
|
||
|
|
reasoning.push_str(&format!("Low confidence ({:.2}). ", overall_confidence));
|
||
|
|
}
|
||
|
|
|
||
|
|
// Check evidence
|
||
|
|
if self.config.require_evidence && signals.evidence_count < self.config.evidence_threshold {
|
||
|
|
warning = Some(format!(
|
||
|
|
"Insufficient evidence: {} facts (required: {})",
|
||
|
|
signals.evidence_count, self.config.evidence_threshold
|
||
|
|
));
|
||
|
|
reasoning.push_str(&format!(
|
||
|
|
"Insufficient evidence ({} facts). ",
|
||
|
|
signals.evidence_count
|
||
|
|
));
|
||
|
|
}
|
||
|
|
|
||
|
|
// Check for contradictions
|
||
|
|
if signals.contradiction_score < 0.5 {
|
||
|
|
warning = Some("Multiple contradictions detected in evidence".to_string());
|
||
|
|
reasoning.push_str("High contradiction risk. ");
|
||
|
|
}
|
||
|
|
|
||
|
|
let is_valid = overall_confidence >= self.config.min_confidence_threshold
|
||
|
|
&& (!self.config.require_evidence
|
||
|
|
|| signals.evidence_count >= self.config.evidence_threshold);
|
||
|
|
|
||
|
|
info!(
|
||
|
|
"Answer validation: confidence={:.2}, valid={}, evidence={}",
|
||
|
|
overall_confidence, is_valid, signals.evidence_count
|
||
|
|
);
|
||
|
|
|
||
|
|
ValidatedAnswer {
|
||
|
|
answer: answer.to_string(),
|
||
|
|
overall_confidence,
|
||
|
|
signals: signals.clone(),
|
||
|
|
is_valid,
|
||
|
|
reasoning: if reasoning.is_empty() {
|
||
|
|
format!("Valid answer (confidence: {:.2})", overall_confidence)
|
||
|
|
} else {
|
||
|
|
reasoning.trim_end().to_string()
|
||
|
|
},
|
||
|
|
warning,
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
/// Batch validate multiple answers
|
||
|
|
pub fn validate_batch(
|
||
|
|
&self,
|
||
|
|
answers: &[(&str, &ConfidenceSignals)],
|
||
|
|
) -> Vec<ValidatedAnswer> {
|
||
|
|
answers
|
||
|
|
.iter()
|
||
|
|
.map(|(answer, signals)| self.validate(answer, signals))
|
||
|
|
.collect()
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
#[cfg(test)]
|
||
|
|
mod tests {
|
||
|
|
use super::*;
|
||
|
|
|
||
|
|
fn make_signals(
|
||
|
|
search: f32,
|
||
|
|
evidence: usize,
|
||
|
|
temporal: f32,
|
||
|
|
entity_cov: f32,
|
||
|
|
contra: f32,
|
||
|
|
) -> ConfidenceSignals {
|
||
|
|
ConfidenceSignals {
|
||
|
|
search_score: search,
|
||
|
|
evidence_count: evidence,
|
||
|
|
evidence_confidence: 0.8,
|
||
|
|
temporal_score: temporal,
|
||
|
|
entity_coverage: entity_cov,
|
||
|
|
contradiction_score: contra,
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
#[test]
|
||
|
|
fn test_validator_config_defaults() {
|
||
|
|
let config = AnswerValidationConfig::default();
|
||
|
|
assert!(config.enabled);
|
||
|
|
assert_eq!(config.min_confidence_threshold, 0.6);
|
||
|
|
assert!(config.require_evidence);
|
||
|
|
}
|
||
|
|
|
||
|
|
#[test]
|
||
|
|
fn test_validate_high_confidence() {
|
||
|
|
let config = AnswerValidationConfig::default();
|
||
|
|
let validator = AnswerValidator::new(config);
|
||
|
|
|
||
|
|
let signals = make_signals(0.9, 3, 0.9, 1.0, 1.0);
|
||
|
|
let result = validator.validate("High confidence answer", &signals);
|
||
|
|
|
||
|
|
assert!(result.is_valid);
|
||
|
|
assert!(result.overall_confidence > 0.8);
|
||
|
|
assert!(result.warning.is_none());
|
||
|
|
}
|
||
|
|
|
||
|
|
#[test]
|
||
|
|
fn test_validate_low_confidence() {
|
||
|
|
let config = AnswerValidationConfig::default();
|
||
|
|
let validator = AnswerValidator::new(config);
|
||
|
|
|
||
|
|
let signals = make_signals(0.3, 0, 0.2, 0.2, 0.5);
|
||
|
|
let result = validator.validate("Low confidence answer", &signals);
|
||
|
|
|
||
|
|
assert!(!result.is_valid);
|
||
|
|
assert!(result.overall_confidence < 0.6);
|
||
|
|
assert!(result.warning.is_some());
|
||
|
|
}
|
||
|
|
|
||
|
|
#[test]
|
||
|
|
fn test_validate_insufficient_evidence() {
|
||
|
|
let config = AnswerValidationConfig {
|
||
|
|
require_evidence: true,
|
||
|
|
evidence_threshold: 3,
|
||
|
|
..Default::default()
|
||
|
|
};
|
||
|
|
let validator = AnswerValidator::new(config);
|
||
|
|
|
||
|
|
let signals = make_signals(0.8, 1, 0.8, 1.0, 1.0); // Only 1 fact
|
||
|
|
let result = validator.validate("Answer with low evidence", &signals);
|
||
|
|
|
||
|
|
assert!(!result.is_valid);
|
||
|
|
assert!(result.warning.is_some());
|
||
|
|
}
|
||
|
|
|
||
|
|
#[test]
|
||
|
|
fn test_validate_disabled() {
|
||
|
|
let config = AnswerValidationConfig {
|
||
|
|
enabled: false,
|
||
|
|
..Default::default()
|
||
|
|
};
|
||
|
|
let validator = AnswerValidator::new(config);
|
||
|
|
|
||
|
|
let signals = make_signals(0.1, 0, 0.1, 0.0, 0.0);
|
||
|
|
let result = validator.validate("Any answer", &signals);
|
||
|
|
|
||
|
|
assert!(result.is_valid);
|
||
|
|
assert_eq!(result.overall_confidence, 1.0);
|
||
|
|
}
|
||
|
|
|
||
|
|
#[test]
|
||
|
|
fn test_confidence_scoring() {
|
||
|
|
let config = AnswerValidationConfig::default();
|
||
|
|
let validator = AnswerValidator::new(config);
|
||
|
|
|
||
|
|
let signals = make_signals(0.8, 2, 0.9, 0.9, 0.9);
|
||
|
|
let result = validator.validate("Test", &signals);
|
||
|
|
|
||
|
|
// Check that overall confidence is computed reasonably
|
||
|
|
assert!(result.overall_confidence > 0.7);
|
||
|
|
assert!(result.overall_confidence <= 1.0);
|
||
|
|
}
|
||
|
|
|
||
|
|
#[test]
|
||
|
|
fn test_contradiction_warning() {
|
||
|
|
let config = AnswerValidationConfig::default();
|
||
|
|
let validator = AnswerValidator::new(config);
|
||
|
|
|
||
|
|
let signals = make_signals(0.8, 3, 0.8, 0.9, 0.3); // Low contradiction score
|
||
|
|
let result = validator.validate("Contradictory answer", &signals);
|
||
|
|
|
||
|
|
assert!(result.warning.is_some());
|
||
|
|
}
|
||
|
|
|
||
|
|
#[test]
|
||
|
|
fn test_batch_validate() {
|
||
|
|
let config = AnswerValidationConfig::default();
|
||
|
|
let validator = AnswerValidator::new(config);
|
||
|
|
|
||
|
|
let signals1 = make_signals(0.9, 3, 0.9, 1.0, 1.0);
|
||
|
|
let signals2 = make_signals(0.2, 0, 0.2, 0.0, 0.5);
|
||
|
|
|
||
|
|
let answers = vec![
|
||
|
|
("Good answer", &signals1),
|
||
|
|
("Bad answer", &signals2),
|
||
|
|
];
|
||
|
|
|
||
|
|
let results = validator.validate_batch(&answers);
|
||
|
|
|
||
|
|
assert_eq!(results.len(), 2);
|
||
|
|
assert!(results[0].is_valid);
|
||
|
|
assert!(!results[1].is_valid);
|
||
|
|
}
|
||
|
|
}
|