1use std::{collections::BTreeSet, error::Error, fmt};
9
10use serde::{Deserialize, Serialize};
11
12use crate::{
13 api::HybridRetrievalResponse,
14 domain::{CodeRetrievalHit, RetrieverSource},
15};
16
17#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
19#[serde(rename_all = "snake_case")]
20pub enum EvaluationCaseKind {
21 ExactFact,
22 MultiHop,
23 Temporal,
24 NegativeRejection,
25 StaleIndex,
26 AmbiguousEntity,
27 CodeImpact,
28}
29
30#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
32pub struct EvaluationCase {
33 pub id: String,
34 pub kind: EvaluationCaseKind,
35 pub query: String,
36 pub expected_result_ids: Vec<String>,
37 pub forbidden_result_ids: Vec<String>,
38 pub required_sources: Vec<RetrieverSource>,
39 pub expected_stale: Option<bool>,
40}
41
42impl EvaluationCase {
43 pub fn new(
45 id: impl Into<String>,
46 kind: EvaluationCaseKind,
47 query: impl Into<String>,
48 ) -> Result<Self, EvaluationError> {
49 let id = required_text("case id", id.into())?;
50 let query = required_text("query", query.into())?;
51
52 Ok(Self {
53 id,
54 kind,
55 query,
56 expected_result_ids: Vec::new(),
57 forbidden_result_ids: Vec::new(),
58 required_sources: Vec::new(),
59 expected_stale: None,
60 })
61 }
62
63 pub fn requiring_results(mut self, ids: &[&str]) -> Result<Self, EvaluationError> {
65 self.expected_result_ids = normalize_ids(ids)?;
66 Ok(self)
67 }
68
69 pub fn forbidding_results(mut self, ids: &[&str]) -> Result<Self, EvaluationError> {
71 self.forbidden_result_ids = normalize_ids(ids)?;
72 Ok(self)
73 }
74
75 pub fn requiring_sources(mut self, sources: &[RetrieverSource]) -> Self {
77 self.required_sources = sources.to_vec();
78 self
79 }
80
81 pub const fn expecting_stale(mut self, stale: bool) -> Self {
83 self.expected_stale = Some(stale);
84 self
85 }
86}
87
88pub fn phase4_fixture_cases() -> Result<Vec<EvaluationCase>, EvaluationError> {
90 Ok(vec![
91 EvaluationCase::new(
92 "phase4_exact_fact",
93 EvaluationCaseKind::ExactFact,
94 "exact fact async sqlite",
95 )?
96 .requiring_results(&["ev-exact"])?
97 .requiring_sources(&[RetrieverSource::Bm25]),
98 EvaluationCase::new(
99 "phase4_multi_hop",
100 EvaluationCaseKind::MultiHop,
101 "GraphRAG uses vector path",
102 )?
103 .requiring_results(&["ev-path"])?
104 .requiring_sources(&[RetrieverSource::GraphPath]),
105 EvaluationCase::new(
106 "phase4_temporal",
107 EvaluationCaseKind::Temporal,
108 "timeline 2026 relay release",
109 )?
110 .requiring_results(&["ev-temporal"])?
111 .requiring_sources(&[RetrieverSource::Temporal]),
112 EvaluationCase::new(
113 "phase4_negative_rejection",
114 EvaluationCaseKind::NegativeRejection,
115 "rejected only context",
116 )?
117 .forbidding_results(&["ev-rejected"])?,
118 EvaluationCase::new(
119 "phase4_stale_index",
120 EvaluationCaseKind::StaleIndex,
121 "stale index refresh",
122 )?
123 .requiring_results(&["ev-stale"])?
124 .expecting_stale(true),
125 EvaluationCase::new(
126 "phase4_ambiguous_entity",
127 EvaluationCaseKind::AmbiguousEntity,
128 "rust",
129 )?
130 .requiring_results(&["ev-rust-language", "ev-rust-material"])?,
131 EvaluationCase::new(
132 "phase4_code_impact",
133 EvaluationCaseKind::CodeImpact,
134 "retry policy changed",
135 )?
136 .requiring_results(&["symbol:retry_policy"])?
137 .requiring_sources(&[RetrieverSource::CodeGraph]),
138 ])
139}
140
141#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
143pub struct EvaluationObservation {
144 pub result_ids: Vec<String>,
145 pub retriever_sources: Vec<RetrieverSource>,
146 pub stale: bool,
147}
148
149impl EvaluationObservation {
150 pub fn from_retrieval(response: &HybridRetrievalResponse) -> Self {
152 let result_ids = response
153 .results
154 .iter()
155 .map(|hit| hit.evidence_id.clone())
156 .collect::<Vec<_>>();
157 let retriever_sources = response
158 .results
159 .iter()
160 .flat_map(|hit| hit.retriever_sources.iter().copied())
161 .collect::<BTreeSet<_>>()
162 .into_iter()
163 .collect::<Vec<_>>();
164
165 Self {
166 result_ids,
167 retriever_sources,
168 stale: response.metadata.stale,
169 }
170 }
171
172 pub fn from_code_impact(hits: &[CodeRetrievalHit]) -> Self {
174 let retriever_sources = (!hits.is_empty())
175 .then_some(RetrieverSource::CodeGraph)
176 .into_iter()
177 .collect::<Vec<_>>();
178 Self {
179 result_ids: hits
180 .iter()
181 .map(|hit| {
182 hit.symbol_snapshot_id
183 .clone()
184 .or_else(|| hit.file_id.clone())
185 .unwrap_or_else(|| hit.path.clone())
186 })
187 .collect(),
188 retriever_sources,
189 stale: hits.iter().any(|hit| hit.stale),
190 }
191 }
192}
193
194#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
196pub struct EvaluationResult {
197 pub case_id: String,
198 pub kind: EvaluationCaseKind,
199 pub passed: bool,
200 pub missing_result_ids: Vec<String>,
201 pub forbidden_result_ids: Vec<String>,
202 pub missing_sources: Vec<RetrieverSource>,
203 pub stale_mismatch: Option<bool>,
204}
205
206#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
208pub struct EvaluationReport {
209 pub passed: bool,
210 pub total: usize,
211 pub failed: usize,
212 pub results: Vec<EvaluationResult>,
213}
214
215pub fn evaluate_case(
217 case: &EvaluationCase,
218 observation: &EvaluationObservation,
219) -> EvaluationResult {
220 let observed_ids = observation
221 .result_ids
222 .iter()
223 .cloned()
224 .collect::<BTreeSet<_>>();
225 let observed_sources = observation
226 .retriever_sources
227 .iter()
228 .copied()
229 .collect::<BTreeSet<_>>();
230 let missing_result_ids = case
231 .expected_result_ids
232 .iter()
233 .filter(|id| !observed_ids.contains(*id))
234 .cloned()
235 .collect::<Vec<_>>();
236 let forbidden_result_ids = case
237 .forbidden_result_ids
238 .iter()
239 .filter(|id| observed_ids.contains(*id))
240 .cloned()
241 .collect::<Vec<_>>();
242 let missing_sources = case
243 .required_sources
244 .iter()
245 .filter(|source| !observed_sources.contains(*source))
246 .copied()
247 .collect::<Vec<_>>();
248 let stale_mismatch = case
249 .expected_stale
250 .filter(|expected| *expected != observation.stale);
251 let passed = missing_result_ids.is_empty()
252 && forbidden_result_ids.is_empty()
253 && missing_sources.is_empty()
254 && stale_mismatch.is_none();
255
256 EvaluationResult {
257 case_id: case.id.clone(),
258 kind: case.kind,
259 passed,
260 missing_result_ids,
261 forbidden_result_ids,
262 missing_sources,
263 stale_mismatch,
264 }
265}
266
267pub fn evaluate_suite(
269 cases: &[EvaluationCase],
270 observations: &[EvaluationObservation],
271) -> Result<EvaluationReport, EvaluationError> {
272 if cases.len() != observations.len() {
273 return Err(EvaluationError::MismatchedObservationCount);
274 }
275 let results = cases
276 .iter()
277 .zip(observations)
278 .map(|(case, observation)| evaluate_case(case, observation))
279 .collect::<Vec<_>>();
280 let failed = results.iter().filter(|result| !result.passed).count();
281
282 Ok(EvaluationReport {
283 passed: failed == 0,
284 total: results.len(),
285 failed,
286 results,
287 })
288}
289
290#[derive(Debug, Clone, PartialEq, Eq)]
292pub enum EvaluationError {
293 EmptyField(&'static str),
294 MismatchedObservationCount,
295}
296
297impl fmt::Display for EvaluationError {
298 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
299 match self {
300 Self::EmptyField(field) => write!(formatter, "{field} must not be empty"),
301 Self::MismatchedObservationCount => {
302 write!(
303 formatter,
304 "evaluation case and observation counts must match"
305 )
306 }
307 }
308 }
309}
310
311impl Error for EvaluationError {}
312
313fn required_text(field: &'static str, value: String) -> Result<String, EvaluationError> {
314 let trimmed = value.trim();
315 if trimmed.is_empty() {
316 return Err(EvaluationError::EmptyField(field));
317 }
318
319 Ok(trimmed.to_owned())
320}
321
322fn normalize_ids(ids: &[&str]) -> Result<Vec<String>, EvaluationError> {
323 ids.iter()
324 .map(|id| required_text("result id", (*id).to_owned()))
325 .collect()
326}
327
328#[cfg(test)]
329mod tests {
330 use super::*;
331 use crate::domain::{CodeRetrievalLayer, RepositoryCodeRange};
332
333 #[test]
334 fn evaluates_phase4_case_families() {
335 let cases = phase4_fixture_cases().expect("fixture cases should validate");
336 let observations = vec![
337 observation(&["ev-exact"], &[RetrieverSource::Bm25], false),
338 observation(&["ev-path"], &[RetrieverSource::GraphPath], false),
339 observation(&["ev-temporal"], &[RetrieverSource::Temporal], false),
340 observation(&[], &[], false),
341 observation(&["ev-stale"], &[], true),
342 observation(&["ev-rust-language", "ev-rust-material"], &[], false),
343 observation(
344 &["symbol:retry_policy"],
345 &[RetrieverSource::CodeGraph],
346 false,
347 ),
348 ];
349
350 let report = evaluate_suite(&cases, &observations).expect("suite should score");
351
352 assert!(report.passed);
353 assert_eq!(report.total, 7);
354 }
355
356 #[test]
357 fn reports_missing_forbidden_and_stale_failures() {
358 let case = EvaluationCase::new("case", EvaluationCaseKind::ExactFact, "query")
359 .unwrap()
360 .requiring_results(&["wanted"])
361 .unwrap()
362 .forbidding_results(&["forbidden"])
363 .unwrap()
364 .requiring_sources(&[RetrieverSource::Vector])
365 .expecting_stale(false);
366 let result = evaluate_case(
367 &case,
368 &observation(&["forbidden"], &[RetrieverSource::Bm25], true),
369 );
370
371 assert!(!result.passed);
372 assert_eq!(result.missing_result_ids, ["wanted"]);
373 assert_eq!(result.forbidden_result_ids, ["forbidden"]);
374 assert_eq!(result.missing_sources, [RetrieverSource::Vector]);
375 assert_eq!(result.stale_mismatch, Some(false));
376 }
377
378 #[test]
379 fn code_impact_observation_preserves_sources_and_stale_state() {
380 let observation = EvaluationObservation::from_code_impact(&[code_hit(true)]);
381
382 assert_eq!(observation.result_ids, ["symbol:retry_policy"]);
383 assert_eq!(observation.retriever_sources, [RetrieverSource::CodeGraph]);
384 assert!(observation.stale);
385 }
386
387 fn observation(
388 ids: &[&str],
389 sources: &[RetrieverSource],
390 stale: bool,
391 ) -> EvaluationObservation {
392 EvaluationObservation {
393 result_ids: ids.iter().map(|id| (*id).to_owned()).collect(),
394 retriever_sources: sources.to_vec(),
395 stale,
396 }
397 }
398
399 fn code_hit(stale: bool) -> CodeRetrievalHit {
400 CodeRetrievalHit {
401 repository_id: "repo".to_owned(),
402 scope_id: "main".to_owned(),
403 resolved_commit_sha: "abc".to_owned(),
404 tree_hash: "tree".to_owned(),
405 path: "src/lib.rs".to_owned(),
406 language_id: "rust".to_owned(),
407 byte_range: RepositoryCodeRange { start: 0, end: 10 },
408 line_range: RepositoryCodeRange { start: 1, end: 1 },
409 symbol_snapshot_id: Some("symbol:retry_policy".to_owned()),
410 canonical_symbol_id: Some("repo://repo/src::lib::retry_policy".to_owned()),
411 file_id: Some("file:src/lib.rs".to_owned()),
412 retrieval_layers: vec![CodeRetrievalLayer::Impact],
413 index_versions: vec!["code_graph:1".to_owned()],
414 stale,
415 degraded_reason: None,
416 edge_kind: None,
417 edge_resolution_state: None,
418 edge_target_hint: None,
419 edge_confidence_basis_points: None,
420 edge_confidence_tier: None,
421 score: 1.0,
422 excerpt: "fn retry_policy() {}".to_owned(),
423 }
424 }
425}