1use super::ApplicationEvalConfig;
8use crate::{EmbeddingModel, Vector};
9use anyhow::{anyhow, Result};
10use serde::{Deserialize, Serialize};
11use std::collections::HashMap;
12
13#[derive(Debug, Clone, Serialize, Deserialize)]
15pub struct RelevanceJudgment {
16 pub query: String,
18 pub document_id: String,
20 pub relevance_score: u8,
22 pub annotator_id: String,
24}
25
26#[derive(Debug, Clone, Serialize, Deserialize)]
28pub enum SearchMetric {
29 PrecisionAtK(usize),
31 RecallAtK(usize),
33 MAP,
35 NDCG(usize),
37 MRR,
39 ERR,
41 CTR,
43}
44
45#[derive(Debug, Clone, Serialize, Deserialize)]
47pub struct QueryResults {
48 pub query: String,
50 pub precision_scores: HashMap<usize, f64>,
52 pub recall_scores: HashMap<usize, f64>,
54 pub ndcg_scores: HashMap<usize, f64>,
56 pub num_relevant: usize,
58 pub difficulty_score: f64,
60 pub average_precision: f64,
63 pub reciprocal_rank: f64,
66}
67
68#[derive(Debug, Clone, Serialize, Deserialize)]
70pub struct QueryPerformanceAnalysis {
71 pub avg_query_length: f64,
73 pub query_type_distribution: HashMap<String, usize>,
75 pub performance_by_difficulty: HashMap<String, f64>,
77 pub zero_result_queries: f64,
79}
80
81#[derive(Debug, Clone, Serialize, Deserialize)]
83pub struct SearchEffectivenessMetrics {
84 pub search_satisfaction: f64,
86 pub relevance_distribution: HashMap<u8, usize>,
88 pub result_diversity: f64,
90 pub query_success_rate: f64,
92}
93
94#[derive(Debug, Clone, Serialize, Deserialize)]
96pub struct SearchResults {
97 pub metric_scores: HashMap<String, f64>,
99 pub per_query_results: HashMap<String, QueryResults>,
101 pub query_analysis: QueryPerformanceAnalysis,
103 pub effectiveness_metrics: SearchEffectivenessMetrics,
105}
106
107pub struct SearchEvaluator {
109 query_relevance: HashMap<String, Vec<RelevanceJudgment>>,
111 metrics: Vec<SearchMetric>,
113}
114
115impl SearchEvaluator {
116 pub fn new() -> Self {
118 Self {
119 query_relevance: HashMap::new(),
120 metrics: vec![
121 SearchMetric::PrecisionAtK(1),
122 SearchMetric::PrecisionAtK(5),
123 SearchMetric::PrecisionAtK(10),
124 SearchMetric::NDCG(10),
125 SearchMetric::MAP,
126 SearchMetric::MRR,
127 ],
128 }
129 }
130
131 pub fn add_relevance_judgment(&mut self, judgment: RelevanceJudgment) {
133 self.query_relevance
134 .entry(judgment.query.clone())
135 .or_default()
136 .push(judgment);
137 }
138
139 pub async fn evaluate(
141 &self,
142 model: &dyn EmbeddingModel,
143 config: &ApplicationEvalConfig,
144 ) -> Result<SearchResults> {
145 let mut metric_scores = HashMap::new();
146 let mut per_query_results = HashMap::new();
147
148 let queries_to_evaluate: Vec<_> = self
150 .query_relevance
151 .keys()
152 .take(config.sample_size)
153 .cloned()
154 .collect();
155
156 for query in &queries_to_evaluate {
157 let query_results = self.evaluate_query_search(query, model).await?;
158 per_query_results.insert(query.clone(), query_results);
159 }
160
161 for metric in &self.metrics {
163 let score = self.calculate_search_metric(metric, &per_query_results)?;
164 metric_scores.insert(format!("{metric:?}"), score);
165 }
166
167 let query_analysis = self.analyze_query_performance(&per_query_results)?;
169 let effectiveness_metrics = self.calculate_effectiveness_metrics(&per_query_results)?;
170
171 Ok(SearchResults {
172 metric_scores,
173 per_query_results,
174 query_analysis,
175 effectiveness_metrics,
176 })
177 }
178
179 async fn evaluate_query_search(
181 &self,
182 query: &str,
183 model: &dyn EmbeddingModel,
184 ) -> Result<QueryResults> {
185 let judgments = self
186 .query_relevance
187 .get(query)
188 .expect("query should exist in query_relevance");
189
190 let search_results = self.perform_search(query, model).await?;
192
193 let mut relevance_scores = Vec::new();
195 for (doc_id, _score) in &search_results {
196 let relevance = judgments
197 .iter()
198 .find(|j| &j.document_id == doc_id)
199 .map(|j| j.relevance_score)
200 .unwrap_or(0);
201 relevance_scores.push(relevance);
202 }
203
204 let num_relevant = judgments.iter().filter(|j| j.relevance_score > 0).count();
205
206 let mut precision_scores = HashMap::new();
208 let mut recall_scores = HashMap::new();
209 let mut ndcg_scores = HashMap::new();
210
211 for &k in &[1, 3, 5, 10] {
212 if k <= search_results.len() {
213 let relevant_at_k =
214 relevance_scores.iter().take(k).filter(|&&r| r > 0).count() as f64;
215
216 let precision = relevant_at_k / k as f64;
217 let recall = if num_relevant > 0 {
218 relevant_at_k / num_relevant as f64
219 } else {
220 0.0
221 };
222
223 precision_scores.insert(k, precision);
224 recall_scores.insert(k, recall);
225
226 let ndcg = self.calculate_search_ndcg(&relevance_scores, k)?;
228 ndcg_scores.insert(k, ndcg);
229 }
230 }
231
232 let difficulty_score = self.calculate_query_difficulty(query, num_relevant);
233
234 let average_precision = if num_relevant > 0 {
237 let mut hits = 0usize;
238 let mut precision_sum = 0.0;
239 for (rank, &relevance) in relevance_scores.iter().enumerate() {
240 if relevance > 0 {
241 hits += 1;
242 precision_sum += hits as f64 / (rank + 1) as f64;
243 }
244 }
245 precision_sum / num_relevant as f64
246 } else {
247 0.0
248 };
249
250 let reciprocal_rank = relevance_scores
252 .iter()
253 .position(|&relevance| relevance > 0)
254 .map(|rank| 1.0 / (rank + 1) as f64)
255 .unwrap_or(0.0);
256
257 Ok(QueryResults {
258 query: query.to_string(),
259 precision_scores,
260 recall_scores,
261 ndcg_scores,
262 num_relevant,
263 difficulty_score,
264 average_precision,
265 reciprocal_rank,
266 })
267 }
268
269 async fn perform_search(
271 &self,
272 query: &str,
273 model: &dyn EmbeddingModel,
274 ) -> Result<Vec<(String, f64)>> {
275 let query_words: Vec<&str> = query.split_whitespace().collect();
277 let mut query_embedding = vec![0.0f32; 100];
278
279 for (i, word) in query_words.iter().enumerate() {
281 if i < query_embedding.len() {
282 query_embedding[i] = word.len() as f32 / 10.0;
283 }
284 }
285 let query_vector = Vector::new(query_embedding);
286
287 let entities = model.get_entities();
289 let mut search_results = Vec::new();
290
291 for entity in entities.iter().take(100) {
292 if let Ok(entity_embedding) = model.get_entity_embedding(entity) {
294 let score = self.cosine_similarity(&query_vector, &entity_embedding);
295 search_results.push((entity.clone(), score));
296 }
297 }
298
299 search_results.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
301 search_results.truncate(20);
302
303 Ok(search_results)
304 }
305
306 fn calculate_search_ndcg(&self, relevance_scores: &[u8], k: usize) -> Result<f64> {
308 if k == 0 || relevance_scores.is_empty() {
309 return Ok(0.0);
310 }
311
312 let mut dcg = 0.0;
313 for (i, &relevance) in relevance_scores.iter().take(k).enumerate() {
314 if relevance > 0 {
315 let gain = (2_u32.pow(relevance as u32) - 1) as f64;
316 dcg += gain / (i as f64 + 2.0).log2();
317 }
318 }
319
320 let mut ideal_relevance: Vec<u8> = relevance_scores.to_vec();
322 ideal_relevance.sort_by(|a, b| b.cmp(a));
323
324 let mut idcg = 0.0;
325 for (i, &relevance) in ideal_relevance.iter().take(k).enumerate() {
326 if relevance > 0 {
327 let gain = (2_u32.pow(relevance as u32) - 1) as f64;
328 idcg += gain / (i as f64 + 2.0).log2();
329 }
330 }
331
332 if idcg > 0.0 {
333 Ok(dcg / idcg)
334 } else {
335 Ok(0.0)
336 }
337 }
338
339 fn calculate_query_difficulty(&self, query: &str, num_relevant: usize) -> f64 {
341 let query_length = query.split_whitespace().count() as f64;
342 let relevance_factor = if num_relevant == 0 {
343 1.0 } else {
345 1.0 / (num_relevant as f64).log2()
346 };
347
348 (query_length * 0.1 + relevance_factor * 0.9).min(1.0)
349 }
350
351 fn calculate_search_metric(
353 &self,
354 metric: &SearchMetric,
355 per_query_results: &HashMap<String, QueryResults>,
356 ) -> Result<f64> {
357 if per_query_results.is_empty() {
358 return Ok(0.0);
359 }
360
361 match metric {
362 SearchMetric::PrecisionAtK(k) => {
363 let scores: Vec<f64> = per_query_results
364 .values()
365 .filter_map(|r| r.precision_scores.get(k))
366 .cloned()
367 .collect();
368 Ok(scores.iter().sum::<f64>() / scores.len() as f64)
369 }
370 SearchMetric::RecallAtK(k) => {
371 let scores: Vec<f64> = per_query_results
372 .values()
373 .filter_map(|r| r.recall_scores.get(k))
374 .cloned()
375 .collect();
376 if scores.is_empty() {
377 return Err(anyhow!("No recall@{k} scores available across queries"));
378 }
379 Ok(scores.iter().sum::<f64>() / scores.len() as f64)
380 }
381 SearchMetric::NDCG(k) => {
382 let scores: Vec<f64> = per_query_results
383 .values()
384 .filter_map(|r| r.ndcg_scores.get(k))
385 .cloned()
386 .collect();
387 Ok(scores.iter().sum::<f64>() / scores.len() as f64)
388 }
389 SearchMetric::MAP => {
390 let scores: Vec<f64> = per_query_results
391 .values()
392 .map(|r| r.average_precision)
393 .collect();
394 Ok(scores.iter().sum::<f64>() / scores.len() as f64)
395 }
396 SearchMetric::MRR => {
397 let scores: Vec<f64> = per_query_results
398 .values()
399 .map(|r| r.reciprocal_rank)
400 .collect();
401 Ok(scores.iter().sum::<f64>() / scores.len() as f64)
402 }
403 SearchMetric::ERR | SearchMetric::CTR => Err(anyhow!(
404 "Search metric {metric:?} is not yet implemented: computing it requires a \
405 graded-relevance cascade/click model that QueryResults does not currently track"
406 )),
407 }
408 }
409
410 fn analyze_query_performance(
412 &self,
413 per_query_results: &HashMap<String, QueryResults>,
414 ) -> Result<QueryPerformanceAnalysis> {
415 let avg_query_length = per_query_results
416 .keys()
417 .map(|q| q.split_whitespace().count() as f64)
418 .sum::<f64>()
419 / per_query_results.len() as f64;
420
421 let zero_result_queries = per_query_results
422 .values()
423 .filter(|r| r.num_relevant == 0)
424 .count() as f64
425 / per_query_results.len() as f64;
426
427 Ok(QueryPerformanceAnalysis {
428 avg_query_length,
429 query_type_distribution: HashMap::new(), performance_by_difficulty: HashMap::new(), zero_result_queries,
432 })
433 }
434
435 fn calculate_effectiveness_metrics(
437 &self,
438 per_query_results: &HashMap<String, QueryResults>,
439 ) -> Result<SearchEffectivenessMetrics> {
440 let successful_queries = per_query_results
441 .values()
442 .filter(|r| r.precision_scores.get(&1).unwrap_or(&0.0) > &0.0)
443 .count() as f64;
444
445 let query_success_rate = successful_queries / per_query_results.len() as f64;
446
447 Ok(SearchEffectivenessMetrics {
448 search_satisfaction: query_success_rate * 0.8, relevance_distribution: HashMap::new(), result_diversity: 0.6, query_success_rate,
452 })
453 }
454
455 fn cosine_similarity(&self, v1: &Vector, v2: &Vector) -> f64 {
457 let dot_product: f32 = v1
458 .values
459 .iter()
460 .zip(v2.values.iter())
461 .map(|(a, b)| a * b)
462 .sum();
463 let norm_a: f32 = v1.values.iter().map(|x| x * x).sum::<f32>().sqrt();
464 let norm_b: f32 = v2.values.iter().map(|x| x * x).sum::<f32>().sqrt();
465
466 if norm_a > 0.0 && norm_b > 0.0 {
467 (dot_product / (norm_a * norm_b)) as f64
468 } else {
469 0.0
470 }
471 }
472}
473
474impl Default for SearchEvaluator {
475 fn default() -> Self {
476 Self::new()
477 }
478}
479
480#[cfg(test)]
481mod tests {
482 use super::*;
483
484 fn sample_query_results(average_precision: f64, reciprocal_rank: f64) -> QueryResults {
485 QueryResults {
486 query: "test query".to_string(),
487 precision_scores: HashMap::new(),
488 recall_scores: [(5usize, 0.5)].into_iter().collect(),
489 ndcg_scores: HashMap::new(),
490 num_relevant: 2,
491 difficulty_score: 0.0,
492 average_precision,
493 reciprocal_rank,
494 }
495 }
496
497 #[test]
500 fn test_calculate_search_metric_map_and_mrr_are_real() -> Result<()> {
501 let evaluator = SearchEvaluator::new();
502 let mut per_query = HashMap::new();
503 per_query.insert("q1".to_string(), sample_query_results(1.0, 1.0));
504 per_query.insert("q2".to_string(), sample_query_results(0.5, 0.5));
505
506 let map = evaluator.calculate_search_metric(&SearchMetric::MAP, &per_query)?;
507 assert!((map - 0.75).abs() < 1e-9, "map = {map}");
508
509 let mrr = evaluator.calculate_search_metric(&SearchMetric::MRR, &per_query)?;
510 assert!((mrr - 0.75).abs() < 1e-9, "mrr = {mrr}");
511
512 Ok(())
513 }
514
515 #[test]
516 fn test_calculate_search_metric_recall_at_k_is_real() -> Result<()> {
517 let evaluator = SearchEvaluator::new();
518 let mut per_query = HashMap::new();
519 per_query.insert("q1".to_string(), sample_query_results(0.0, 0.0));
520
521 let recall = evaluator.calculate_search_metric(&SearchMetric::RecallAtK(5), &per_query)?;
522 assert!((recall - 0.5).abs() < 1e-9, "recall = {recall}");
523
524 Ok(())
525 }
526
527 #[test]
530 fn test_calculate_search_metric_err_and_ctr_are_explicit_errors() {
531 let evaluator = SearchEvaluator::new();
532 let mut per_query = HashMap::new();
533 per_query.insert("q1".to_string(), sample_query_results(1.0, 1.0));
534
535 assert!(evaluator
536 .calculate_search_metric(&SearchMetric::ERR, &per_query)
537 .is_err());
538 assert!(evaluator
539 .calculate_search_metric(&SearchMetric::CTR, &per_query)
540 .is_err());
541 }
542}