1use std::{collections::BTreeSet, sync::Arc, time::Instant};
4
5use runifold_core::Usage;
6use runifold_retrieval::{DocumentId, RetrievalContext, RetrievalError, RetrievalQuery, Retriever};
7use thiserror::Error;
8
9#[derive(Debug, Error)]
11#[non_exhaustive]
12pub enum RetrievalEvaluationError {
13 #[error("retrieval evaluation case id cannot be empty")]
15 EmptyCaseId,
16 #[error("retrieval evaluation case must contain at least one relevant document")]
18 EmptyRelevantDocuments,
19 #[error("retrieval evaluation cutoff must be greater than zero")]
21 ZeroCutoff,
22 #[error("retrieval evaluation collection is too large for metric calculation")]
24 CountOutOfRange,
25 #[error("retrieval evaluation failed: {0}")]
27 Retrieval(#[from] RetrievalError),
28}
29
30#[derive(Clone, Debug, Eq, PartialEq)]
32pub struct RetrievalEvaluationCase {
33 pub id: String,
35 pub query: String,
37 pub relevant: BTreeSet<DocumentId>,
39 pub cutoff: usize,
41}
42
43impl RetrievalEvaluationCase {
44 pub fn new(
51 id: impl Into<String>,
52 query: impl Into<String>,
53 relevant: impl IntoIterator<Item = DocumentId>,
54 cutoff: usize,
55 ) -> Result<Self, RetrievalEvaluationError> {
56 let id = id.into();
57 if id.trim().is_empty() {
58 return Err(RetrievalEvaluationError::EmptyCaseId);
59 }
60 if cutoff == 0 {
61 return Err(RetrievalEvaluationError::ZeroCutoff);
62 }
63 let query = query.into();
64 RetrievalQuery::new(query.clone(), cutoff)?;
65 let relevant = relevant.into_iter().collect::<BTreeSet<_>>();
66 if relevant.is_empty() {
67 return Err(RetrievalEvaluationError::EmptyRelevantDocuments);
68 }
69 Ok(Self {
70 id,
71 query,
72 relevant,
73 cutoff,
74 })
75 }
76}
77
78#[derive(Clone, Debug, PartialEq)]
80pub struct RetrievalCaseMetrics {
81 pub id: String,
83 pub precision_at_k: f64,
85 pub recall_at_k: f64,
87 pub reciprocal_rank: f64,
89 pub ndcg_at_k: f64,
91 pub usage: Usage,
93 pub elapsed_micros: u64,
95}
96
97#[derive(Clone, Debug, PartialEq)]
99pub struct RetrievalEvaluationReport {
100 pub cases: Vec<RetrievalCaseMetrics>,
102 pub mean_precision_at_k: f64,
104 pub mean_recall_at_k: f64,
106 pub mean_reciprocal_rank: f64,
108 pub mean_ndcg_at_k: f64,
110 pub mean_elapsed_micros: u64,
112}
113
114#[derive(Clone)]
116pub struct RetrievalEvaluationRunner {
117 retriever: Arc<dyn Retriever>,
118}
119
120impl RetrievalEvaluationRunner {
121 pub fn new(retriever: Arc<dyn Retriever>) -> Self {
123 Self { retriever }
124 }
125
126 pub async fn run(
132 &self,
133 cases: &[RetrievalEvaluationCase],
134 ) -> Result<RetrievalEvaluationReport, RetrievalEvaluationError> {
135 let mut metrics = Vec::with_capacity(cases.len());
136 for case in cases {
137 let started = Instant::now();
138 let response = self
139 .retriever
140 .retrieve(
141 RetrievalQuery::new(case.query.clone(), case.cutoff)?,
142 RetrievalContext::new(),
143 )
144 .await?;
145 let elapsed_micros = u64::try_from(started.elapsed().as_micros()).unwrap_or(u64::MAX);
146 let ranked = response
147 .documents
148 .iter()
149 .take(case.cutoff)
150 .map(|result| &result.document.id)
151 .collect::<Vec<_>>();
152 let hits = ranked
153 .iter()
154 .filter(|id| case.relevant.contains(*id))
155 .count();
156 let reciprocal_rank = ranked
157 .iter()
158 .position(|id| case.relevant.contains(*id))
159 .map(|index| count_as_f64(index + 1).map(|rank| 1.0 / rank))
160 .transpose()?
161 .unwrap_or(0.0);
162 let dcg = ranked
163 .iter()
164 .enumerate()
165 .filter(|(_, id)| case.relevant.contains(**id))
166 .map(|(index, _)| count_as_f64(index + 2).map(|rank| 1.0 / rank.log2()))
167 .collect::<Result<Vec<_>, _>>()?
168 .into_iter()
169 .sum::<f64>();
170 let ideal = case.relevant.len().min(case.cutoff);
171 let idcg = (0..ideal)
172 .map(|index| count_as_f64(index + 2).map(|rank| 1.0 / rank.log2()))
173 .collect::<Result<Vec<_>, _>>()?
174 .into_iter()
175 .sum::<f64>();
176 metrics.push(RetrievalCaseMetrics {
177 id: case.id.clone(),
178 precision_at_k: count_as_f64(hits)? / count_as_f64(case.cutoff)?,
179 recall_at_k: count_as_f64(hits)? / count_as_f64(case.relevant.len())?,
180 reciprocal_rank,
181 ndcg_at_k: dcg / idcg,
182 usage: response.usage,
183 elapsed_micros,
184 });
185 }
186 aggregate(metrics)
187 }
188}
189
190impl std::fmt::Debug for RetrievalEvaluationRunner {
191 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
192 formatter
193 .debug_struct("RetrievalEvaluationRunner")
194 .field("retriever", self.retriever.descriptor())
195 .finish()
196 }
197}
198
199fn aggregate(
200 cases: Vec<RetrievalCaseMetrics>,
201) -> Result<RetrievalEvaluationReport, RetrievalEvaluationError> {
202 if cases.is_empty() {
203 return Ok(RetrievalEvaluationReport {
204 cases,
205 mean_precision_at_k: 0.0,
206 mean_recall_at_k: 0.0,
207 mean_reciprocal_rank: 0.0,
208 mean_ndcg_at_k: 0.0,
209 mean_elapsed_micros: 0,
210 });
211 }
212 let count = count_as_f64(cases.len())?;
213 let elapsed_total = cases
214 .iter()
215 .map(|case| u128::from(case.elapsed_micros))
216 .sum::<u128>();
217 let elapsed_mean = elapsed_total
218 / u128::try_from(cases.len()).map_err(|_| RetrievalEvaluationError::CountOutOfRange)?;
219 Ok(RetrievalEvaluationReport {
220 mean_precision_at_k: cases.iter().map(|case| case.precision_at_k).sum::<f64>() / count,
221 mean_recall_at_k: cases.iter().map(|case| case.recall_at_k).sum::<f64>() / count,
222 mean_reciprocal_rank: cases.iter().map(|case| case.reciprocal_rank).sum::<f64>() / count,
223 mean_ndcg_at_k: cases.iter().map(|case| case.ndcg_at_k).sum::<f64>() / count,
224 mean_elapsed_micros: u64::try_from(elapsed_mean)
225 .map_err(|_| RetrievalEvaluationError::CountOutOfRange)?,
226 cases,
227 })
228}
229
230fn count_as_f64(value: usize) -> Result<f64, RetrievalEvaluationError> {
231 u32::try_from(value)
232 .map(f64::from)
233 .map_err(|_| RetrievalEvaluationError::CountOutOfRange)
234}
235
236#[cfg(test)]
237mod tests {
238 use std::collections::BTreeMap;
239
240 use runifold_retrieval::{
241 Document, RetrievalFuture, RetrievalResponse, RetrievedDocument, RetrieverDescriptor,
242 };
243
244 use super::*;
245
246 struct RankedRetriever {
247 descriptor: RetrieverDescriptor,
248 }
249
250 impl Retriever for RankedRetriever {
251 fn descriptor(&self) -> &RetrieverDescriptor {
252 &self.descriptor
253 }
254
255 fn retrieve(
256 &self,
257 _query: RetrievalQuery,
258 _context: RetrievalContext,
259 ) -> RetrievalFuture<'_, Result<RetrievalResponse, RetrievalError>> {
260 Box::pin(async {
261 Ok(RetrievalResponse {
262 documents: vec![
263 RetrievedDocument {
264 document: Document::new("irrelevant", "noise").unwrap(),
265 score: 1.0,
266 },
267 RetrievedDocument {
268 document: Document::new("relevant", "answer").unwrap(),
269 score: 0.9,
270 },
271 ],
272 usage: Usage::default(),
273 })
274 })
275 }
276 }
277
278 #[test]
279 fn computes_rank_sensitive_metrics_from_stable_evidence() {
280 let case = RetrievalEvaluationCase::new(
281 "case",
282 "query",
283 [DocumentId::new("relevant").unwrap()],
284 2,
285 )
286 .unwrap();
287 let runner = RetrievalEvaluationRunner::new(Arc::new(RankedRetriever {
288 descriptor: RetrieverDescriptor {
289 metadata: BTreeMap::new(),
290 ..RetrieverDescriptor::read_only("ranked")
291 },
292 }));
293
294 let report = futures_executor::block_on(runner.run(&[case])).unwrap();
295
296 assert!((report.mean_precision_at_k - 0.5).abs() < f64::EPSILON);
297 assert!((report.mean_recall_at_k - 1.0).abs() < f64::EPSILON);
298 assert!((report.mean_reciprocal_rank - 0.5).abs() < f64::EPSILON);
299 assert!(report.mean_ndcg_at_k > 0.6 && report.mean_ndcg_at_k < 0.7);
300 }
301}