relay_knowledge/evaluation/
mod.rs1use 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)]
329#[path = "mod_tests.rs"]
330mod tests;