use std::{collections::BTreeSet, error::Error, fmt};
use serde::{Deserialize, Serialize};
use crate::{
api::HybridRetrievalResponse,
domain::{CodeRetrievalHit, RetrieverSource},
};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum EvaluationCaseKind {
ExactFact,
MultiHop,
Temporal,
NegativeRejection,
StaleIndex,
AmbiguousEntity,
CodeImpact,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct EvaluationCase {
pub id: String,
pub kind: EvaluationCaseKind,
pub query: String,
pub expected_result_ids: Vec<String>,
pub forbidden_result_ids: Vec<String>,
pub required_sources: Vec<RetrieverSource>,
pub expected_stale: Option<bool>,
}
impl EvaluationCase {
pub fn new(
id: impl Into<String>,
kind: EvaluationCaseKind,
query: impl Into<String>,
) -> Result<Self, EvaluationError> {
let id = required_text("case id", id.into())?;
let query = required_text("query", query.into())?;
Ok(Self {
id,
kind,
query,
expected_result_ids: Vec::new(),
forbidden_result_ids: Vec::new(),
required_sources: Vec::new(),
expected_stale: None,
})
}
pub fn requiring_results(mut self, ids: &[&str]) -> Result<Self, EvaluationError> {
self.expected_result_ids = normalize_ids(ids)?;
Ok(self)
}
pub fn forbidding_results(mut self, ids: &[&str]) -> Result<Self, EvaluationError> {
self.forbidden_result_ids = normalize_ids(ids)?;
Ok(self)
}
pub fn requiring_sources(mut self, sources: &[RetrieverSource]) -> Self {
self.required_sources = sources.to_vec();
self
}
pub const fn expecting_stale(mut self, stale: bool) -> Self {
self.expected_stale = Some(stale);
self
}
}
pub fn phase4_fixture_cases() -> Result<Vec<EvaluationCase>, EvaluationError> {
Ok(vec![
EvaluationCase::new(
"phase4_exact_fact",
EvaluationCaseKind::ExactFact,
"exact fact async sqlite",
)?
.requiring_results(&["ev-exact"])?
.requiring_sources(&[RetrieverSource::Bm25]),
EvaluationCase::new(
"phase4_multi_hop",
EvaluationCaseKind::MultiHop,
"GraphRAG uses vector path",
)?
.requiring_results(&["ev-path"])?
.requiring_sources(&[RetrieverSource::GraphPath]),
EvaluationCase::new(
"phase4_temporal",
EvaluationCaseKind::Temporal,
"timeline 2026 relay release",
)?
.requiring_results(&["ev-temporal"])?
.requiring_sources(&[RetrieverSource::Temporal]),
EvaluationCase::new(
"phase4_negative_rejection",
EvaluationCaseKind::NegativeRejection,
"rejected only context",
)?
.forbidding_results(&["ev-rejected"])?,
EvaluationCase::new(
"phase4_stale_index",
EvaluationCaseKind::StaleIndex,
"stale index refresh",
)?
.requiring_results(&["ev-stale"])?
.expecting_stale(true),
EvaluationCase::new(
"phase4_ambiguous_entity",
EvaluationCaseKind::AmbiguousEntity,
"rust",
)?
.requiring_results(&["ev-rust-language", "ev-rust-material"])?,
EvaluationCase::new(
"phase4_code_impact",
EvaluationCaseKind::CodeImpact,
"retry policy changed",
)?
.requiring_results(&["symbol:retry_policy"])?
.requiring_sources(&[RetrieverSource::CodeGraph]),
])
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct EvaluationObservation {
pub result_ids: Vec<String>,
pub retriever_sources: Vec<RetrieverSource>,
pub stale: bool,
}
impl EvaluationObservation {
pub fn from_retrieval(response: &HybridRetrievalResponse) -> Self {
let result_ids = response
.results
.iter()
.map(|hit| hit.evidence_id.clone())
.collect::<Vec<_>>();
let retriever_sources = response
.results
.iter()
.flat_map(|hit| hit.retriever_sources.iter().copied())
.collect::<BTreeSet<_>>()
.into_iter()
.collect::<Vec<_>>();
Self {
result_ids,
retriever_sources,
stale: response.metadata.stale,
}
}
pub fn from_code_impact(hits: &[CodeRetrievalHit]) -> Self {
let retriever_sources = (!hits.is_empty())
.then_some(RetrieverSource::CodeGraph)
.into_iter()
.collect::<Vec<_>>();
Self {
result_ids: hits
.iter()
.map(|hit| {
hit.symbol_snapshot_id
.clone()
.or_else(|| hit.file_id.clone())
.unwrap_or_else(|| hit.path.clone())
})
.collect(),
retriever_sources,
stale: hits.iter().any(|hit| hit.stale),
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct EvaluationResult {
pub case_id: String,
pub kind: EvaluationCaseKind,
pub passed: bool,
pub missing_result_ids: Vec<String>,
pub forbidden_result_ids: Vec<String>,
pub missing_sources: Vec<RetrieverSource>,
pub stale_mismatch: Option<bool>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct EvaluationReport {
pub passed: bool,
pub total: usize,
pub failed: usize,
pub results: Vec<EvaluationResult>,
}
pub fn evaluate_case(
case: &EvaluationCase,
observation: &EvaluationObservation,
) -> EvaluationResult {
let observed_ids = observation
.result_ids
.iter()
.cloned()
.collect::<BTreeSet<_>>();
let observed_sources = observation
.retriever_sources
.iter()
.copied()
.collect::<BTreeSet<_>>();
let missing_result_ids = case
.expected_result_ids
.iter()
.filter(|id| !observed_ids.contains(*id))
.cloned()
.collect::<Vec<_>>();
let forbidden_result_ids = case
.forbidden_result_ids
.iter()
.filter(|id| observed_ids.contains(*id))
.cloned()
.collect::<Vec<_>>();
let missing_sources = case
.required_sources
.iter()
.filter(|source| !observed_sources.contains(*source))
.copied()
.collect::<Vec<_>>();
let stale_mismatch = case
.expected_stale
.filter(|expected| *expected != observation.stale);
let passed = missing_result_ids.is_empty()
&& forbidden_result_ids.is_empty()
&& missing_sources.is_empty()
&& stale_mismatch.is_none();
EvaluationResult {
case_id: case.id.clone(),
kind: case.kind,
passed,
missing_result_ids,
forbidden_result_ids,
missing_sources,
stale_mismatch,
}
}
pub fn evaluate_suite(
cases: &[EvaluationCase],
observations: &[EvaluationObservation],
) -> Result<EvaluationReport, EvaluationError> {
if cases.len() != observations.len() {
return Err(EvaluationError::MismatchedObservationCount);
}
let results = cases
.iter()
.zip(observations)
.map(|(case, observation)| evaluate_case(case, observation))
.collect::<Vec<_>>();
let failed = results.iter().filter(|result| !result.passed).count();
Ok(EvaluationReport {
passed: failed == 0,
total: results.len(),
failed,
results,
})
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum EvaluationError {
EmptyField(&'static str),
MismatchedObservationCount,
}
impl fmt::Display for EvaluationError {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::EmptyField(field) => write!(formatter, "{field} must not be empty"),
Self::MismatchedObservationCount => {
write!(
formatter,
"evaluation case and observation counts must match"
)
}
}
}
}
impl Error for EvaluationError {}
fn required_text(field: &'static str, value: String) -> Result<String, EvaluationError> {
let trimmed = value.trim();
if trimmed.is_empty() {
return Err(EvaluationError::EmptyField(field));
}
Ok(trimmed.to_owned())
}
fn normalize_ids(ids: &[&str]) -> Result<Vec<String>, EvaluationError> {
ids.iter()
.map(|id| required_text("result id", (*id).to_owned()))
.collect()
}
#[cfg(test)]
#[path = "mod_tests.rs"]
mod tests;