Skip to main content

runifold_testkit/
retrieval_evaluation.rs

1//! Deterministic information-retrieval quality metrics.
2
3use 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/// Invalid dataset or failed retrieval evaluation.
10#[derive(Debug, Error)]
11#[non_exhaustive]
12pub enum RetrievalEvaluationError {
13    /// A case identity was blank.
14    #[error("retrieval evaluation case id cannot be empty")]
15    EmptyCaseId,
16    /// A case had no relevant documents.
17    #[error("retrieval evaluation case must contain at least one relevant document")]
18    EmptyRelevantDocuments,
19    /// A case requested no ranked results.
20    #[error("retrieval evaluation cutoff must be greater than zero")]
21    ZeroCutoff,
22    /// A collection size could not be represented by the metric format.
23    #[error("retrieval evaluation collection is too large for metric calculation")]
24    CountOutOfRange,
25    /// The retriever failed.
26    #[error("retrieval evaluation failed: {0}")]
27    Retrieval(#[from] RetrievalError),
28}
29
30/// One versionable retrieval-quality case.
31#[derive(Clone, Debug, Eq, PartialEq)]
32pub struct RetrievalEvaluationCase {
33    /// Stable case identity.
34    pub id: String,
35    /// Search query.
36    pub query: String,
37    /// Ground-truth relevant document identities.
38    pub relevant: BTreeSet<DocumentId>,
39    /// Ranking cutoff.
40    pub cutoff: usize,
41}
42
43impl RetrievalEvaluationCase {
44    /// Validates one retrieval evaluation case.
45    ///
46    /// # Errors
47    ///
48    /// Rejects blank identities or queries, empty relevance sets, and zero
49    /// cutoffs.
50    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/// Metrics for one ranked retrieval case.
79#[derive(Clone, Debug, PartialEq)]
80pub struct RetrievalCaseMetrics {
81    /// Case identity.
82    pub id: String,
83    /// Precision at the configured cutoff.
84    pub precision_at_k: f64,
85    /// Recall at the configured cutoff.
86    pub recall_at_k: f64,
87    /// Reciprocal rank of the first relevant result.
88    pub reciprocal_rank: f64,
89    /// Normalized discounted cumulative gain.
90    pub ndcg_at_k: f64,
91    /// Retriever-reported usage.
92    pub usage: Usage,
93    /// Host-observed end-to-end duration.
94    pub elapsed_micros: u64,
95}
96
97/// Macro-averaged retrieval-quality report.
98#[derive(Clone, Debug, PartialEq)]
99pub struct RetrievalEvaluationReport {
100    /// Per-case evidence in input order.
101    pub cases: Vec<RetrievalCaseMetrics>,
102    /// Mean precision at K.
103    pub mean_precision_at_k: f64,
104    /// Mean recall at K.
105    pub mean_recall_at_k: f64,
106    /// Mean reciprocal rank.
107    pub mean_reciprocal_rank: f64,
108    /// Mean normalized discounted cumulative gain.
109    pub mean_ndcg_at_k: f64,
110    /// Mean host-observed latency.
111    pub mean_elapsed_micros: u64,
112}
113
114/// Deterministic evaluator for any provider-neutral retriever.
115#[derive(Clone)]
116pub struct RetrievalEvaluationRunner {
117    retriever: Arc<dyn Retriever>,
118}
119
120impl RetrievalEvaluationRunner {
121    /// Creates an evaluator around one retriever.
122    pub fn new(retriever: Arc<dyn Retriever>) -> Self {
123        Self { retriever }
124    }
125
126    /// Runs cases sequentially to preserve deterministic evidence ordering.
127    ///
128    /// # Errors
129    ///
130    /// Fails when any retriever invocation fails.
131    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}