Skip to main content

oxirs_embed/application_tasks/
search.rs

1//! Search relevance evaluation module
2//!
3//! This module provides comprehensive evaluation for search relevance using
4//! embedding models, including precision, recall, NDCG, MAP, and other
5//! information retrieval metrics.
6
7use super::ApplicationEvalConfig;
8use crate::{EmbeddingModel, Vector};
9use anyhow::{anyhow, Result};
10use serde::{Deserialize, Serialize};
11use std::collections::HashMap;
12
13/// Relevance judgment for search evaluation
14#[derive(Debug, Clone, Serialize, Deserialize)]
15pub struct RelevanceJudgment {
16    /// Query
17    pub query: String,
18    /// Document/entity identifier
19    pub document_id: String,
20    /// Relevance score (0-3: not relevant, somewhat relevant, relevant, highly relevant)
21    pub relevance_score: u8,
22    /// Annotator identifier
23    pub annotator_id: String,
24}
25
26/// Search evaluation metrics
27#[derive(Debug, Clone, Serialize, Deserialize)]
28pub enum SearchMetric {
29    /// Precision at K
30    PrecisionAtK(usize),
31    /// Recall at K
32    RecallAtK(usize),
33    /// Mean Average Precision
34    MAP,
35    /// Normalized Discounted Cumulative Gain
36    NDCG(usize),
37    /// Mean Reciprocal Rank
38    MRR,
39    /// Expected Reciprocal Rank
40    ERR,
41    /// Click-through rate simulation
42    CTR,
43}
44
45/// Per-query search results
46#[derive(Debug, Clone, Serialize, Deserialize)]
47pub struct QueryResults {
48    /// Query text
49    pub query: String,
50    /// Precision scores at different K values
51    pub precision_scores: HashMap<usize, f64>,
52    /// Recall scores at different K values
53    pub recall_scores: HashMap<usize, f64>,
54    /// NDCG scores
55    pub ndcg_scores: HashMap<usize, f64>,
56    /// Number of relevant documents
57    pub num_relevant: usize,
58    /// Query difficulty score
59    pub difficulty_score: f64,
60    /// Average precision over the full ranked result list, for aggregating
61    /// into Mean Average Precision (MAP) across queries.
62    pub average_precision: f64,
63    /// Reciprocal rank of the first relevant result (0.0 if none), for
64    /// aggregating into Mean Reciprocal Rank (MRR) across queries.
65    pub reciprocal_rank: f64,
66}
67
68/// Query performance analysis
69#[derive(Debug, Clone, Serialize, Deserialize)]
70pub struct QueryPerformanceAnalysis {
71    /// Average query length
72    pub avg_query_length: f64,
73    /// Query type distribution
74    pub query_type_distribution: HashMap<String, usize>,
75    /// Performance by query difficulty
76    pub performance_by_difficulty: HashMap<String, f64>,
77    /// Zero-result queries percentage
78    pub zero_result_queries: f64,
79}
80
81/// Search effectiveness metrics
82#[derive(Debug, Clone, Serialize, Deserialize)]
83pub struct SearchEffectivenessMetrics {
84    /// Overall search satisfaction
85    pub search_satisfaction: f64,
86    /// Result relevance distribution
87    pub relevance_distribution: HashMap<u8, usize>,
88    /// Search result diversity
89    pub result_diversity: f64,
90    /// Query success rate
91    pub query_success_rate: f64,
92}
93
94/// Search evaluation results
95#[derive(Debug, Clone, Serialize, Deserialize)]
96pub struct SearchResults {
97    /// Metric scores
98    pub metric_scores: HashMap<String, f64>,
99    /// Per-query results
100    pub per_query_results: HashMap<String, QueryResults>,
101    /// Query performance analysis
102    pub query_analysis: QueryPerformanceAnalysis,
103    /// Search effectiveness metrics
104    pub effectiveness_metrics: SearchEffectivenessMetrics,
105}
106
107/// Search relevance evaluator
108pub struct SearchEvaluator {
109    /// Search queries and their relevance judgments
110    query_relevance: HashMap<String, Vec<RelevanceJudgment>>,
111    /// Search metrics to evaluate
112    metrics: Vec<SearchMetric>,
113}
114
115impl SearchEvaluator {
116    /// Create a new search evaluator
117    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    /// Add relevance judgment
132    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    /// Evaluate search relevance
140    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        // Sample queries for evaluation
149        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        // Calculate aggregate metrics
162        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        // Analyze query performance
168        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    /// Evaluate search for a specific query
180    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        // Get search results (simplified - would use actual search system)
191        let search_results = self.perform_search(query, model).await?;
192
193        // Calculate relevance for each result
194        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        // Calculate metrics at different K values
207        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                // Calculate NDCG
227                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        // Average precision over the full ranked list (for MAP): the mean of
235        // precision@k evaluated at every rank that holds a relevant result.
236        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        // Reciprocal rank of the first relevant result (for MRR).
251        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    /// Perform search (simplified implementation)
270    async fn perform_search(
271        &self,
272        query: &str,
273        model: &dyn EmbeddingModel,
274    ) -> Result<Vec<(String, f64)>> {
275        // Create query embedding (simplified)
276        let query_words: Vec<&str> = query.split_whitespace().collect();
277        let mut query_embedding = vec![0.0f32; 100];
278
279        // Simple word-based embedding (in practice, use proper query embedding)
280        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        // Score entities (documents) against query
288        let entities = model.get_entities();
289        let mut search_results = Vec::new();
290
291        for entity in entities.iter().take(100) {
292            // Limit for efficiency
293            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        // Sort by score and return top results
300        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    /// Calculate NDCG for search results
307    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        // Calculate ideal DCG
321        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    /// Calculate query difficulty
340    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 // High difficulty
344        } 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    /// Calculate aggregate search metric
352    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    /// Analyze query performance
411    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(), // Simplified
430            performance_by_difficulty: HashMap::new(), // Simplified
431            zero_result_queries,
432        })
433    }
434
435    /// Calculate effectiveness metrics
436    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, // Simplified
449            relevance_distribution: HashMap::new(),        // Simplified
450            result_diversity: 0.6,                         // Simplified
451            query_success_rate,
452        })
453    }
454
455    /// Calculate cosine similarity
456    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    /// Regression test: MAP/MRR must be computed for real from
498    /// `average_precision`/`reciprocal_rank` instead of a hardcoded 0.5.
499    #[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    /// Metrics that genuinely cannot be computed from the tracked data must
528    /// fail loudly rather than return a fabricated placeholder score.
529    #[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}