Skip to main content

oxirs_vec/
similarity.rs

1//! Advanced similarity algorithms and semantic matching for vectors
2
3use crate::Vector;
4use anyhow::{anyhow, Result};
5use oxirs_core::simd::SimdOps;
6use serde::{Deserialize, Serialize};
7use std::collections::hash_map::DefaultHasher;
8use std::collections::HashMap;
9use std::hash::{Hash, Hasher};
10use std::time::{SystemTime, UNIX_EPOCH};
11
12/// Similarity measurement configuration
13#[derive(Debug, Clone, Serialize, Deserialize, oxicode::Encode, oxicode::Decode)]
14pub struct SimilarityConfig {
15    /// Primary similarity metric
16    pub primary_metric: SimilarityMetric,
17    /// Secondary metrics for ensemble scoring
18    pub ensemble_metrics: Vec<SimilarityMetric>,
19    /// Weights for ensemble metrics
20    pub ensemble_weights: Vec<f32>,
21    /// Threshold for considering vectors similar
22    pub similarity_threshold: f32,
23    /// Enable semantic boosting
24    pub semantic_boost: bool,
25    /// Enable temporal decay
26    pub temporal_decay: bool,
27}
28
29impl Default for SimilarityConfig {
30    fn default() -> Self {
31        Self {
32            primary_metric: SimilarityMetric::Cosine,
33            ensemble_metrics: vec![
34                SimilarityMetric::Cosine,
35                SimilarityMetric::Pearson,
36                SimilarityMetric::Jaccard,
37            ],
38            ensemble_weights: vec![0.5, 0.3, 0.2],
39            similarity_threshold: 0.7,
40            semantic_boost: true,
41            temporal_decay: false,
42        }
43    }
44}
45
46/// Available similarity metrics
47#[derive(
48    Debug, Clone, Copy, Serialize, Deserialize, PartialEq, oxicode::Encode, oxicode::Decode,
49)]
50pub enum SimilarityMetric {
51    /// Cosine similarity
52    Cosine,
53    /// Euclidean distance (converted to similarity)
54    Euclidean,
55    /// Manhattan distance (converted to similarity)
56    Manhattan,
57    /// Minkowski distance (general Lp norm)
58    Minkowski(f32),
59    /// Pearson correlation coefficient
60    Pearson,
61    /// Spearman rank correlation
62    Spearman,
63    /// Jaccard similarity (for sparse vectors)
64    Jaccard,
65    /// Dice coefficient
66    Dice,
67    /// Jensen-Shannon divergence
68    JensenShannon,
69    /// Bhattacharyya distance
70    Bhattacharyya,
71    /// Mahalanobis distance (requires covariance matrix)
72    Mahalanobis,
73    /// Hamming distance (for binary vectors)
74    Hamming,
75    /// Canberra distance
76    Canberra,
77    /// Angular distance
78    Angular,
79    /// Chebyshev distance (L∞ norm)
80    Chebyshev,
81    /// Dot product (inner product)
82    DotProduct,
83}
84
85impl SimilarityMetric {
86    /// Calculate similarity between two vectors
87    pub fn similarity(&self, a: &[f32], b: &[f32]) -> Result<f32> {
88        if a.len() != b.len() {
89            return Err(anyhow!("Vector dimensions must match"));
90        }
91
92        let similarity = match self {
93            SimilarityMetric::Cosine => cosine_similarity(a, b),
94            SimilarityMetric::Euclidean => euclidean_similarity(a, b),
95            SimilarityMetric::Manhattan => manhattan_similarity(a, b),
96            SimilarityMetric::Minkowski(p) => minkowski_similarity(a, b, *p),
97            SimilarityMetric::Pearson => pearson_correlation(a, b)?,
98            SimilarityMetric::Spearman => spearman_correlation(a, b)?,
99            SimilarityMetric::Jaccard => jaccard_similarity(a, b),
100            SimilarityMetric::Dice => dice_coefficient(a, b),
101            SimilarityMetric::JensenShannon => jensen_shannon_similarity(a, b)?,
102            SimilarityMetric::Bhattacharyya => bhattacharyya_similarity(a, b)?,
103            SimilarityMetric::Mahalanobis => {
104                // Mahalanobis requires a covariance matrix, which this stateless
105                // metric API does not carry. Do not silently return Euclidean
106                // (Mahalanobis with Σ = I) mislabeled as Mahalanobis — fail loud.
107                // Use `SemanticSimilarity::set_covariance_matrix` +
108                // `SemanticSimilarity::mahalanobis_similarity` for the real metric.
109                return Err(anyhow!(
110                    "Mahalanobis similarity requires a covariance matrix; use \
111                     SemanticSimilarity::set_covariance_matrix() rather than the \
112                     stateless SimilarityMetric::Mahalanobis"
113                ));
114            }
115            SimilarityMetric::Hamming => hamming_similarity(a, b),
116            SimilarityMetric::Canberra => canberra_similarity(a, b),
117            SimilarityMetric::Angular => angular_similarity(a, b),
118            SimilarityMetric::Chebyshev => chebyshev_similarity(a, b),
119            SimilarityMetric::DotProduct => dot_product_similarity(a, b),
120        };
121
122        Ok(similarity.clamp(0.0, 1.0))
123    }
124
125    /// Calculate distance between two vectors (lower is more similar)
126    pub fn distance(&self, a: &Vector, b: &Vector) -> Result<f32> {
127        let a_f32 = a.as_f32();
128        let b_f32 = b.as_f32();
129        self.distance_slices(&a_f32, &b_f32)
130    }
131
132    /// Calculate distance between two already-materialised `f32` slices.
133    ///
134    /// This is the slice-based core of [`Self::distance`]. Prefer this method
135    /// over `distance()` whenever the caller already has cached `f32` data on
136    /// hand (e.g. [`crate::hnsw::Node::vector_data_f32`]) — `distance()` must
137    /// call [`Vector::as_f32`] on each argument, which allocates and clones a
138    /// fresh `Vec<f32>` for every single call. In hot paths that compare many
139    /// node pairs (HNSW neighbor selection / pruning), that redundant
140    /// allocate-and-copy dominates runtime; calling this method directly with
141    /// pre-materialised slices avoids it entirely while computing the exact
142    /// same result.
143    pub fn distance_slices(&self, a_f32: &[f32], b_f32: &[f32]) -> Result<f32> {
144        if a_f32.len() != b_f32.len() {
145            return Err(anyhow!("Vector dimensions must match"));
146        }
147
148        let distance = match self {
149            // Distance metrics - use direct calculation
150            SimilarityMetric::Euclidean => euclidean_distance(a_f32, b_f32),
151            SimilarityMetric::Manhattan => manhattan_distance(a_f32, b_f32),
152            SimilarityMetric::Minkowski(p) => minkowski_distance(a_f32, b_f32, *p),
153            SimilarityMetric::Hamming => hamming_distance(a_f32, b_f32),
154            SimilarityMetric::Canberra => canberra_distance(a_f32, b_f32),
155            SimilarityMetric::Chebyshev => chebyshev_distance(a_f32, b_f32),
156
157            // Similarity metrics - convert to distance (1 - similarity)
158            _ => {
159                let similarity = self.similarity(a_f32, b_f32)?;
160                1.0 - similarity
161            }
162        };
163
164        Ok(distance.max(0.0))
165    }
166
167    /// Compute similarity between two vectors (alias for similarity method)
168    pub fn compute(&self, a: &Vector, b: &Vector) -> Result<f32> {
169        let a_f32 = a.as_f32();
170        let b_f32 = b.as_f32();
171        self.similarity(&a_f32, &b_f32)
172    }
173}
174
175/// Semantic similarity computer with multiple algorithms
176pub struct SemanticSimilarity {
177    config: SimilarityConfig,
178    feature_weights: Option<Vec<f32>>,
179    covariance_matrix: Option<Vec<Vec<f32>>>,
180}
181
182impl SemanticSimilarity {
183    pub fn new(config: SimilarityConfig) -> Self {
184        Self {
185            config,
186            feature_weights: None,
187            covariance_matrix: None,
188        }
189    }
190
191    /// Set feature importance weights
192    pub fn set_feature_weights(&mut self, weights: Vec<f32>) {
193        self.feature_weights = Some(weights);
194    }
195
196    /// Set covariance matrix for Mahalanobis distance
197    pub fn set_covariance_matrix(&mut self, matrix: Vec<Vec<f32>>) {
198        self.covariance_matrix = Some(matrix);
199    }
200
201    /// Compute the true Mahalanobis distance `sqrt((a-b)ᵀ · Σ⁻¹ · (a-b))` using
202    /// the configured covariance matrix Σ (set via
203    /// [`Self::set_covariance_matrix`]).
204    ///
205    /// Fails loudly if no covariance matrix has been configured or if it is not
206    /// a square matrix matching the vector dimensionality / is singular — never
207    /// silently degrades to Euclidean distance.
208    pub fn mahalanobis_distance(&self, a: &[f32], b: &[f32]) -> Result<f32> {
209        if a.len() != b.len() {
210            return Err(anyhow!("Vector dimensions must match"));
211        }
212        let cov = self.covariance_matrix.as_ref().ok_or_else(|| {
213            anyhow!(
214                "Mahalanobis distance requires a covariance matrix; call \
215                 SemanticSimilarity::set_covariance_matrix() first"
216            )
217        })?;
218        let n = a.len();
219        if cov.len() != n || cov.iter().any(|row| row.len() != n) {
220            return Err(anyhow!(
221                "Covariance matrix must be {n}x{n} to match the vector dimensionality"
222            ));
223        }
224        let inv = invert_matrix(cov)
225            .ok_or_else(|| anyhow!("Covariance matrix is singular and cannot be inverted"))?;
226
227        // diff = a - b
228        let diff: Vec<f32> = a.iter().zip(b).map(|(x, y)| x - y).collect();
229        // tmp = Σ⁻¹ · diff
230        let mut quadratic = 0.0f32;
231        for (i, &di) in diff.iter().enumerate() {
232            let mut row_sum = 0.0f32;
233            for (j, &dj) in diff.iter().enumerate() {
234                row_sum += inv[i][j] * dj;
235            }
236            quadratic += di * row_sum;
237        }
238        // Numerical guard: a valid inverse-covariance is PSD so quadratic >= 0,
239        // but clamp tiny negative round-off to 0 before sqrt.
240        Ok(quadratic.max(0.0).sqrt())
241    }
242
243    /// Mahalanobis *similarity* in `[0, 1]` derived from the distance as
244    /// `1 / (1 + d)` (larger = more similar), consistent with the crate's
245    /// distance→similarity convention.
246    pub fn mahalanobis_similarity(&self, a: &[f32], b: &[f32]) -> Result<f32> {
247        let d = self.mahalanobis_distance(a, b)?;
248        Ok(1.0 / (1.0 + d))
249    }
250
251    /// Calculate similarity using primary metric
252    pub fn similarity(&self, a: &Vector, b: &Vector) -> Result<f32> {
253        let a_f32 = a.as_f32();
254        let b_f32 = b.as_f32();
255
256        // Mahalanobis is stateful (needs the covariance matrix), so it is
257        // handled here rather than via the stateless `SimilarityMetric` path,
258        // which deliberately fails loudly for this variant.
259        let mut similarity = if matches!(self.config.primary_metric, SimilarityMetric::Mahalanobis)
260        {
261            self.mahalanobis_similarity(&a_f32, &b_f32)?
262        } else {
263            self.config.primary_metric.similarity(&a_f32, &b_f32)?
264        };
265
266        // Apply feature weighting if available
267        if let Some(ref weights) = self.feature_weights {
268            similarity = self.apply_feature_weights(&a_f32, &b_f32, weights);
269        }
270
271        // Apply semantic boosting
272        if self.config.semantic_boost {
273            similarity = self.apply_semantic_boost(similarity, a, b);
274        }
275
276        Ok(similarity)
277    }
278
279    /// Calculate ensemble similarity using multiple metrics
280    pub fn ensemble_similarity(&self, a: &Vector, b: &Vector) -> Result<f32> {
281        if self.config.ensemble_metrics.len() != self.config.ensemble_weights.len() {
282            return Err(anyhow!("Ensemble metrics and weights length mismatch"));
283        }
284
285        let a_f32 = a.as_f32();
286        let b_f32 = b.as_f32();
287
288        let mut weighted_sum = 0.0;
289        let mut total_weight = 0.0;
290
291        for (metric, weight) in self
292            .config
293            .ensemble_metrics
294            .iter()
295            .zip(&self.config.ensemble_weights)
296        {
297            let similarity = metric.similarity(&a_f32, &b_f32)?;
298            weighted_sum += similarity * weight;
299            total_weight += weight;
300        }
301
302        if total_weight == 0.0 {
303            return Ok(0.0);
304        }
305
306        let ensemble_score = weighted_sum / total_weight;
307
308        // Apply semantic boosting
309        if self.config.semantic_boost {
310            Ok(self.apply_semantic_boost(ensemble_score, a, b))
311        } else {
312            Ok(ensemble_score)
313        }
314    }
315
316    /// Calculate similarity matrix for a set of vectors
317    pub fn similarity_matrix(&self, vectors: &[Vector]) -> Result<Vec<Vec<f32>>> {
318        let n = vectors.len();
319        let mut matrix = vec![vec![0.0; n]; n];
320
321        for i in 0..n {
322            for j in i..n {
323                let similarity = if i == j {
324                    1.0
325                } else {
326                    self.similarity(&vectors[i], &vectors[j])?
327                };
328
329                matrix[i][j] = similarity;
330                matrix[j][i] = similarity;
331            }
332        }
333
334        Ok(matrix)
335    }
336
337    /// Find most similar vectors to a query
338    pub fn find_similar(
339        &self,
340        query: &Vector,
341        candidates: &[(String, Vector)],
342        k: usize,
343    ) -> Result<Vec<(String, f32)>> {
344        let mut similarities: Vec<(String, f32)> = candidates
345            .iter()
346            .map(|(uri, vector)| {
347                let sim = self.similarity(query, vector).unwrap_or(0.0);
348                (uri.clone(), sim)
349            })
350            .collect();
351
352        similarities.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
353        similarities.truncate(k);
354
355        Ok(similarities)
356    }
357
358    /// Calculate semantic clusters based on similarity
359    pub fn cluster_by_similarity(
360        &self,
361        vectors: &[(String, Vector)],
362        threshold: f32,
363    ) -> Result<Vec<Vec<String>>> {
364        let mut clusters: Vec<Vec<String>> = Vec::new();
365        let mut assigned: Vec<bool> = vec![false; vectors.len()];
366
367        for i in 0..vectors.len() {
368            if assigned[i] {
369                continue;
370            }
371
372            let mut cluster = vec![vectors[i].0.clone()];
373            assigned[i] = true;
374
375            for j in (i + 1)..vectors.len() {
376                if assigned[j] {
377                    continue;
378                }
379
380                let similarity = self.similarity(&vectors[i].1, &vectors[j].1)?;
381                if similarity >= threshold {
382                    cluster.push(vectors[j].0.clone());
383                    assigned[j] = true;
384                }
385            }
386
387            clusters.push(cluster);
388        }
389
390        Ok(clusters)
391    }
392
393    fn apply_feature_weights(&self, a: &[f32], b: &[f32], weights: &[f32]) -> f32 {
394        let weighted_a: Vec<f32> = a.iter().zip(weights).map(|(x, w)| x * w).collect();
395        let weighted_b: Vec<f32> = b.iter().zip(weights).map(|(x, w)| x * w).collect();
396
397        cosine_similarity(&weighted_a, &weighted_b)
398    }
399
400    fn apply_semantic_boost(&self, similarity: f32, a: &Vector, b: &Vector) -> f32 {
401        // Simple semantic boosting based on vector magnitude similarity
402        let a_f32 = a.as_f32();
403        let b_f32 = b.as_f32();
404        let mag_a = vector_magnitude(&a_f32);
405        let mag_b = vector_magnitude(&b_f32);
406        let magnitude_similarity = 1.0 - (mag_a - mag_b).abs() / (mag_a + mag_b + f32::EPSILON);
407
408        // Weighted combination
409        0.8 * similarity + 0.2 * magnitude_similarity
410    }
411}
412
413/// Adaptive similarity that learns from user feedback
414pub struct AdaptiveSimilarity {
415    base_similarity: SemanticSimilarity,
416    feedback_weights: HashMap<String, f32>,
417    learning_rate: f32,
418}
419
420impl AdaptiveSimilarity {
421    pub fn new(config: SimilarityConfig, learning_rate: f32) -> Self {
422        Self {
423            base_similarity: SemanticSimilarity::new(config),
424            feedback_weights: HashMap::new(),
425            learning_rate,
426        }
427    }
428
429    /// Provide feedback on similarity result
430    pub fn add_feedback(&mut self, uri: &str, expected_similarity: f32, actual_similarity: f32) {
431        let error = expected_similarity - actual_similarity;
432        let adjustment = self.learning_rate * error;
433
434        *self.feedback_weights.entry(uri.to_string()).or_insert(0.0) += adjustment;
435    }
436
437    /// Calculate similarity with learned adjustments
438    pub fn adaptive_similarity(
439        &self,
440        a: &Vector,
441        b: &Vector,
442        uri_a: &str,
443        uri_b: &str,
444    ) -> Result<f32> {
445        let base_sim = self.base_similarity.similarity(a, b)?;
446
447        let weight_a = self.feedback_weights.get(uri_a).unwrap_or(&0.0);
448        let weight_b = self.feedback_weights.get(uri_b).unwrap_or(&0.0);
449        let adjustment = (weight_a + weight_b) / 2.0;
450
451        Ok((base_sim + adjustment).clamp(0.0, 1.0))
452    }
453
454    /// Get learned weights for analysis
455    pub fn get_feedback_weights(&self) -> &HashMap<String, f32> {
456        &self.feedback_weights
457    }
458}
459
460/// Temporal similarity that considers time decay
461pub struct TemporalSimilarity {
462    base_similarity: SemanticSimilarity,
463    decay_rate: f32,
464    time_weights: HashMap<String, f32>,
465}
466
467impl TemporalSimilarity {
468    pub fn new(config: SimilarityConfig, decay_rate: f32) -> Self {
469        Self {
470            base_similarity: SemanticSimilarity::new(config),
471            decay_rate,
472            time_weights: HashMap::new(),
473        }
474    }
475
476    /// Set time weight for a URI (higher = more recent)
477    pub fn set_time_weight(&mut self, uri: &str, time_weight: f32) {
478        self.time_weights.insert(uri.to_string(), time_weight);
479    }
480
481    /// Calculate similarity with temporal decay
482    pub fn temporal_similarity(
483        &self,
484        a: &Vector,
485        b: &Vector,
486        uri_a: &str,
487        uri_b: &str,
488    ) -> Result<f32> {
489        let base_sim = self.base_similarity.similarity(a, b)?;
490
491        let time_a = self.time_weights.get(uri_a).unwrap_or(&1.0);
492        let time_b = self.time_weights.get(uri_b).unwrap_or(&1.0);
493
494        let time_factor = (time_a + time_b) / 2.0;
495        let decay = (-self.decay_rate * (1.0 - time_factor)).exp();
496
497        Ok(base_sim * decay)
498    }
499}
500
501// Individual similarity function implementations
502
503/// Compute similarity between two vectors using the specified metric
504pub fn compute_similarity(a: &[f32], b: &[f32], metric: SimilarityMetric) -> Result<f32> {
505    metric.similarity(a, b)
506}
507
508/// Normalize a vector to unit length (in-place)
509pub fn normalize_vector(vector: &mut [f32]) -> Result<()> {
510    let magnitude: f32 = vector.iter().map(|x| x * x).sum::<f32>().sqrt();
511    if magnitude > 0.0 {
512        for value in vector.iter_mut() {
513            *value /= magnitude;
514        }
515    }
516    Ok(())
517}
518
519pub fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
520    // Use oxirs-core SIMD operations
521    1.0 - f32::cosine_distance(a, b)
522}
523
524fn euclidean_similarity(a: &[f32], b: &[f32]) -> f32 {
525    // Use oxirs-core SIMD operations
526    let distance = f32::euclidean_distance(a, b);
527    1.0 / (1.0 + distance)
528}
529
530fn manhattan_similarity(a: &[f32], b: &[f32]) -> f32 {
531    // Use oxirs-core SIMD operations
532    let distance = f32::manhattan_distance(a, b);
533    1.0 / (1.0 + distance)
534}
535
536fn minkowski_similarity(a: &[f32], b: &[f32], p: f32) -> f32 {
537    if p <= 0.0 {
538        // Handle edge case
539        return euclidean_similarity(a, b);
540    }
541
542    let distance: f32 = a
543        .iter()
544        .zip(b)
545        .map(|(x, y)| (x - y).abs().powf(p))
546        .sum::<f32>()
547        .powf(1.0 / p);
548    1.0 / (1.0 + distance)
549}
550
551fn chebyshev_similarity(a: &[f32], b: &[f32]) -> f32 {
552    let distance: f32 = a
553        .iter()
554        .zip(b)
555        .map(|(x, y)| (x - y).abs())
556        .fold(0.0, |acc, diff| acc.max(diff));
557    1.0 / (1.0 + distance)
558}
559
560fn pearson_correlation(a: &[f32], b: &[f32]) -> Result<f32> {
561    let n = a.len() as f32;
562    if n == 0.0 {
563        return Ok(0.0);
564    }
565
566    // Use oxirs-core SIMD operations for mean calculation
567    let mean_a = f32::mean(a);
568    let mean_b = f32::mean(b);
569
570    let numerator: f32 = a
571        .iter()
572        .zip(b)
573        .map(|(x, y)| (x - mean_a) * (y - mean_b))
574        .sum();
575    let sum_sq_a: f32 = a.iter().map(|x| (x - mean_a).powi(2)).sum();
576    let sum_sq_b: f32 = b.iter().map(|x| (x - mean_b).powi(2)).sum();
577
578    let denominator = (sum_sq_a * sum_sq_b).sqrt();
579
580    if denominator == 0.0 {
581        Ok(0.0)
582    } else {
583        Ok(numerator / denominator)
584    }
585}
586
587fn spearman_correlation(a: &[f32], b: &[f32]) -> Result<f32> {
588    let ranks_a = compute_ranks(a);
589    let ranks_b = compute_ranks(b);
590    pearson_correlation(&ranks_a, &ranks_b)
591}
592
593fn compute_ranks(values: &[f32]) -> Vec<f32> {
594    let mut indexed: Vec<(usize, f32)> = values.iter().enumerate().map(|(i, &v)| (i, v)).collect();
595    indexed.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal));
596
597    let mut ranks = vec![0.0; values.len()];
598    for (rank, (original_index, _)) in indexed.iter().enumerate() {
599        ranks[*original_index] = rank as f32 + 1.0;
600    }
601
602    ranks
603}
604
605fn jaccard_similarity(a: &[f32], b: &[f32]) -> f32 {
606    let threshold = 0.01; // Consider values above this as "present"
607    let set_a: Vec<bool> = a.iter().map(|&x| x > threshold).collect();
608    let set_b: Vec<bool> = b.iter().map(|&x| x > threshold).collect();
609
610    let intersection: usize = set_a
611        .iter()
612        .zip(&set_b)
613        .map(|(x, y)| (*x && *y) as usize)
614        .sum();
615    let union: usize = set_a
616        .iter()
617        .zip(&set_b)
618        .map(|(x, y)| (*x || *y) as usize)
619        .sum();
620
621    if union == 0 {
622        1.0 // Both empty sets
623    } else {
624        intersection as f32 / union as f32
625    }
626}
627
628fn dice_coefficient(a: &[f32], b: &[f32]) -> f32 {
629    let threshold = 0.01;
630    let set_a: Vec<bool> = a.iter().map(|&x| x > threshold).collect();
631    let set_b: Vec<bool> = b.iter().map(|&x| x > threshold).collect();
632
633    let intersection: usize = set_a
634        .iter()
635        .zip(&set_b)
636        .map(|(x, y)| (*x && *y) as usize)
637        .sum();
638    let size_a: usize = set_a.iter().map(|&x| x as usize).sum();
639    let size_b: usize = set_b.iter().map(|&x| x as usize).sum();
640
641    if size_a + size_b == 0 {
642        1.0
643    } else {
644        2.0 * intersection as f32 / (size_a + size_b) as f32
645    }
646}
647
648fn jensen_shannon_similarity(a: &[f32], b: &[f32]) -> Result<f32> {
649    // Normalize to probability distributions
650    let sum_a: f32 = a.iter().sum();
651    let sum_b: f32 = b.iter().sum();
652
653    if sum_a == 0.0 || sum_b == 0.0 {
654        return Ok(0.0);
655    }
656
657    let p: Vec<f32> = a.iter().map(|x| x / sum_a).collect();
658    let q: Vec<f32> = b.iter().map(|x| x / sum_b).collect();
659
660    // Compute average distribution
661    let m: Vec<f32> = p.iter().zip(&q).map(|(x, y)| (x + y) / 2.0).collect();
662
663    // Compute KL divergences
664    let kl_pm = kl_divergence(&p, &m);
665    let kl_qm = kl_divergence(&q, &m);
666
667    let js_distance = (kl_pm + kl_qm) / 2.0;
668    Ok(1.0 - js_distance.sqrt()) // Convert distance to similarity
669}
670
671fn kl_divergence(p: &[f32], q: &[f32]) -> f32 {
672    p.iter()
673        .zip(q)
674        .map(|(pi, qi)| {
675            if *pi > 0.0 && *qi > 0.0 {
676                pi * (pi / qi).ln()
677            } else {
678                0.0
679            }
680        })
681        .sum()
682}
683
684fn bhattacharyya_similarity(a: &[f32], b: &[f32]) -> Result<f32> {
685    let sum_a: f32 = a.iter().sum();
686    let sum_b: f32 = b.iter().sum();
687
688    if sum_a == 0.0 || sum_b == 0.0 {
689        return Ok(0.0);
690    }
691
692    let p: Vec<f32> = a.iter().map(|x| x / sum_a).collect();
693    let q: Vec<f32> = b.iter().map(|x| x / sum_b).collect();
694
695    let bc: f32 = p.iter().zip(&q).map(|(x, y)| (x * y).sqrt()).sum();
696    Ok(bc)
697}
698
699fn hamming_similarity(a: &[f32], b: &[f32]) -> f32 {
700    let threshold = 0.5;
701    let matches = a
702        .iter()
703        .zip(b)
704        .filter(|(x, y)| (**x > threshold) == (**y > threshold))
705        .count();
706
707    matches as f32 / a.len() as f32
708}
709
710fn canberra_similarity(a: &[f32], b: &[f32]) -> f32 {
711    let distance: f32 = a
712        .iter()
713        .zip(b)
714        .map(|(x, y)| {
715            let numerator = (x - y).abs();
716            let denominator = x.abs() + y.abs();
717            if denominator > 0.0 {
718                numerator / denominator
719            } else {
720                0.0
721            }
722        })
723        .sum();
724
725    1.0 / (1.0 + distance)
726}
727
728fn angular_similarity(a: &[f32], b: &[f32]) -> f32 {
729    let cosine_sim = cosine_similarity(a, b);
730    let angle = cosine_sim.acos();
731    1.0 - (angle / std::f32::consts::PI)
732}
733
734fn dot_product_similarity(a: &[f32], b: &[f32]) -> f32 {
735    // Simple dot product implementation
736    a.iter().zip(b.iter()).map(|(x, y)| x * y).sum()
737}
738
739fn vector_magnitude(vector: &[f32]) -> f32 {
740    // Calculate vector magnitude (L2 norm)
741    vector.iter().map(|x| x * x).sum::<f32>().sqrt()
742}
743
744/// Invert a square matrix via Gauss-Jordan elimination with partial pivoting.
745///
746/// Returns `None` if the matrix is singular (or numerically indistinguishable
747/// from singular). Used to obtain Σ⁻¹ for the Mahalanobis distance. Pure Rust,
748/// no external linear-algebra dependency.
749fn invert_matrix(matrix: &[Vec<f32>]) -> Option<Vec<Vec<f32>>> {
750    let n = matrix.len();
751    if n == 0 || matrix.iter().any(|row| row.len() != n) {
752        return None;
753    }
754
755    // Build the augmented matrix [A | I] in f64 for numerical stability.
756    let mut aug: Vec<Vec<f64>> = matrix
757        .iter()
758        .enumerate()
759        .map(|(i, row)| {
760            let mut r: Vec<f64> = row.iter().map(|&v| v as f64).collect();
761            r.extend((0..n).map(|j| if i == j { 1.0 } else { 0.0 }));
762            r
763        })
764        .collect();
765
766    for col in 0..n {
767        // Partial pivot: pick the row with the largest absolute value in `col`.
768        let mut pivot = col;
769        let mut best = aug[col][col].abs();
770        for (row, aug_row) in aug.iter().enumerate().skip(col + 1) {
771            let v = aug_row[col].abs();
772            if v > best {
773                best = v;
774                pivot = row;
775            }
776        }
777        if best < 1e-12 {
778            return None; // Singular.
779        }
780        aug.swap(col, pivot);
781
782        // Normalize the pivot row.
783        let pivot_val = aug[col][col];
784        for x in aug[col].iter_mut() {
785            *x /= pivot_val;
786        }
787
788        // Eliminate `col` from every other row. Clone the (already normalized)
789        // pivot row so we can iterate the target row mutably without a second
790        // index into `aug`.
791        let pivot_row = aug[col].clone();
792        for (row, target_row) in aug.iter_mut().enumerate() {
793            if row == col {
794                continue;
795            }
796            let factor = target_row[col];
797            if factor != 0.0 {
798                for (target, &pv) in target_row.iter_mut().zip(pivot_row.iter()) {
799                    *target -= factor * pv;
800                }
801            }
802        }
803    }
804
805    // Extract the right half (the inverse), back into f32.
806    Some(
807        aug.into_iter()
808            .map(|row| row[n..].iter().map(|&v| v as f32).collect())
809            .collect(),
810    )
811}
812
813// Distance function implementations (lower values mean more similar)
814
815fn euclidean_distance(a: &[f32], b: &[f32]) -> f32 {
816    // Use oxirs-core SIMD operations
817    f32::euclidean_distance(a, b)
818}
819
820fn manhattan_distance(a: &[f32], b: &[f32]) -> f32 {
821    // Use oxirs-core SIMD operations
822    f32::manhattan_distance(a, b)
823}
824
825fn minkowski_distance(a: &[f32], b: &[f32], p: f32) -> f32 {
826    if p <= 0.0 {
827        return euclidean_distance(a, b);
828    }
829
830    a.iter()
831        .zip(b)
832        .map(|(x, y)| (x - y).abs().powf(p))
833        .sum::<f32>()
834        .powf(1.0 / p)
835}
836
837fn chebyshev_distance(a: &[f32], b: &[f32]) -> f32 {
838    a.iter()
839        .zip(b)
840        .map(|(x, y)| (x - y).abs())
841        .fold(0.0, |acc, diff| acc.max(diff))
842}
843
844fn hamming_distance(a: &[f32], b: &[f32]) -> f32 {
845    let threshold = 0.5;
846    let mismatches = a
847        .iter()
848        .zip(b)
849        .filter(|(x, y)| (**x > threshold) != (**y > threshold))
850        .count();
851
852    mismatches as f32 / a.len() as f32
853}
854
855fn canberra_distance(a: &[f32], b: &[f32]) -> f32 {
856    a.iter()
857        .zip(b)
858        .map(|(x, y)| {
859            let numerator = (x - y).abs();
860            let denominator = x.abs() + y.abs();
861            if denominator > 0.0 {
862                numerator / denominator
863            } else {
864                0.0
865            }
866        })
867        .sum()
868}
869
870/// Similarity search result with metadata
871#[derive(Debug, Clone, Serialize, Deserialize)]
872pub struct SimilarityResult {
873    pub id: String,
874    pub uri: String,
875    pub similarity: f32,
876    pub metrics: HashMap<String, f32>,
877    pub metadata: Option<HashMap<String, String>>,
878}
879
880/// Batch similarity processor for efficient computation
881pub struct BatchSimilarityProcessor {
882    similarity: SemanticSimilarity,
883    cache: HashMap<(String, String), f32>,
884    max_cache_size: usize,
885}
886
887impl BatchSimilarityProcessor {
888    pub fn new(config: SimilarityConfig, max_cache_size: usize) -> Self {
889        Self {
890            similarity: SemanticSimilarity::new(config),
891            cache: HashMap::new(),
892            max_cache_size,
893        }
894    }
895
896    /// Process batch of similarity computations with caching
897    pub fn process_batch(
898        &mut self,
899        queries: &[(String, Vector)],
900        candidates: &[(String, Vector)],
901    ) -> Result<Vec<Vec<SimilarityResult>>> {
902        let mut results = Vec::new();
903
904        for (query_uri, query_vec) in queries {
905            let mut query_results = Vec::new();
906
907            for (candidate_uri, candidate_vec) in candidates {
908                let cache_key = if query_uri < candidate_uri {
909                    (query_uri.clone(), candidate_uri.clone())
910                } else {
911                    (candidate_uri.clone(), query_uri.clone())
912                };
913
914                let similarity = if let Some(&cached_sim) = self.cache.get(&cache_key) {
915                    cached_sim
916                } else {
917                    let sim = self.similarity.similarity(query_vec, candidate_vec)?;
918
919                    // Cache management
920                    if self.cache.len() >= self.max_cache_size {
921                        // Simple eviction: remove first entry
922                        if let Some(key) = self.cache.keys().next().cloned() {
923                            self.cache.remove(&key);
924                        }
925                    }
926
927                    self.cache.insert(cache_key, sim);
928                    sim
929                };
930
931                query_results.push(SimilarityResult {
932                    id: generate_similarity_id(candidate_uri, similarity),
933                    uri: candidate_uri.clone(),
934                    similarity,
935                    metrics: HashMap::new(),
936                    metadata: None,
937                });
938            }
939
940            // Sort by similarity (descending)
941            query_results.sort_by(|a, b| {
942                b.similarity
943                    .partial_cmp(&a.similarity)
944                    .unwrap_or(std::cmp::Ordering::Equal)
945            });
946            results.push(query_results);
947        }
948
949        Ok(results)
950    }
951
952    pub fn cache_stats(&self) -> (usize, usize) {
953        (self.cache.len(), self.max_cache_size)
954    }
955
956    pub fn clear_cache(&mut self) {
957        self.cache.clear();
958    }
959}
960
961/// Generate a unique ID for similarity results
962fn generate_similarity_id(uri: &str, similarity: f32) -> String {
963    let mut hasher = DefaultHasher::new();
964    uri.hash(&mut hasher);
965    similarity.to_bits().hash(&mut hasher);
966
967    let timestamp = SystemTime::now()
968        .duration_since(UNIX_EPOCH)
969        .unwrap_or_default()
970        .as_millis();
971
972    timestamp.hash(&mut hasher);
973
974    format!("sim_{:x}", hasher.finish())
975}
976
977#[cfg(test)]
978mod mahalanobis_tests {
979    use super::*;
980    use crate::distance_metrics::ExtendedDistanceMetric;
981
982    #[test]
983    fn regression_stateless_mahalanobis_fails_loud() {
984        // The stateless metric APIs must NOT silently return Euclidean disguised
985        // as Mahalanobis — they must error.
986        let a = [1.0f32, 2.0, 3.0];
987        let b = [4.0f32, 5.0, 6.0];
988        assert!(SimilarityMetric::Mahalanobis.similarity(&a, &b).is_err());
989        let va = Vector::new(a.to_vec());
990        let vb = Vector::new(b.to_vec());
991        assert!(SimilarityMetric::Mahalanobis.distance(&va, &vb).is_err());
992        assert!(ExtendedDistanceMetric::Mahalanobis
993            .distance(&va, &vb)
994            .is_err());
995    }
996
997    #[test]
998    fn regression_mahalanobis_identity_equals_euclidean() -> Result<()> {
999        // With Σ = I, Mahalanobis distance reduces to Euclidean distance.
1000        let mut sem = SemanticSimilarity::new(SimilarityConfig {
1001            primary_metric: SimilarityMetric::Mahalanobis,
1002            ..Default::default()
1003        });
1004        sem.set_covariance_matrix(vec![
1005            vec![1.0, 0.0, 0.0],
1006            vec![0.0, 1.0, 0.0],
1007            vec![0.0, 0.0, 1.0],
1008        ]);
1009        let a = [1.0f32, 2.0, 3.0];
1010        let b = [4.0f32, 6.0, 3.0];
1011        let maha = sem.mahalanobis_distance(&a, &b)?;
1012        let eucl = ((3.0f32).powi(2) + (4.0f32).powi(2)).sqrt(); // = 5.0
1013        assert!((maha - eucl).abs() < 1e-4, "maha={maha}, eucl={eucl}");
1014        Ok(())
1015    }
1016
1017    #[test]
1018    fn regression_mahalanobis_uses_covariance() -> Result<()> {
1019        // A diagonal covariance with a large variance on axis 0 should shrink
1020        // that axis's contribution, so it must differ from plain Euclidean.
1021        let mut sem = SemanticSimilarity::new(SimilarityConfig {
1022            primary_metric: SimilarityMetric::Mahalanobis,
1023            ..Default::default()
1024        });
1025        sem.set_covariance_matrix(vec![vec![100.0, 0.0], vec![0.0, 1.0]]);
1026        let a = [0.0f32, 0.0];
1027        let b = [10.0f32, 0.0];
1028        // d = sqrt((10)^2 / 100) = 1.0, NOT the Euclidean 10.0.
1029        let maha = sem.mahalanobis_distance(&a, &b)?;
1030        assert!((maha - 1.0).abs() < 1e-4, "maha={maha}");
1031        Ok(())
1032    }
1033
1034    #[test]
1035    fn regression_mahalanobis_requires_covariance() {
1036        let sem = SemanticSimilarity::new(SimilarityConfig {
1037            primary_metric: SimilarityMetric::Mahalanobis,
1038            ..Default::default()
1039        });
1040        assert!(sem.mahalanobis_distance(&[1.0, 2.0], &[3.0, 4.0]).is_err());
1041    }
1042
1043    #[test]
1044    fn regression_mahalanobis_singular_covariance_errs() {
1045        let mut sem = SemanticSimilarity::new(SimilarityConfig::default());
1046        // A zero matrix is singular -> must error, not silently proceed.
1047        sem.set_covariance_matrix(vec![vec![0.0, 0.0], vec![0.0, 0.0]]);
1048        assert!(sem.mahalanobis_distance(&[1.0, 2.0], &[3.0, 4.0]).is_err());
1049    }
1050}