#[derive(Debug, Clone, thiserror::Error)]
pub enum MathError {
#[error("vector length mismatch: {0} vs {1}")]
LengthMismatch(usize, usize),
}
pub fn cosine_similarity(a: &[f32], b: &[f32]) -> Result<f32, MathError> {
if a.len() != b.len() {
return Err(MathError::LengthMismatch(a.len(), b.len()));
}
let dot_product: f64 = a
.iter()
.zip(b.iter())
.map(|(x, y)| (*x as f64) * (*y as f64))
.sum();
let norm_a: f64 = a
.iter()
.map(|x| (*x as f64) * (*x as f64))
.sum::<f64>()
.sqrt();
let norm_b: f64 = b
.iter()
.map(|x| (*x as f64) * (*x as f64))
.sum::<f64>()
.sqrt();
if norm_a < f64::EPSILON || norm_b < f64::EPSILON {
return Ok(0.0);
}
Ok((dot_product / (norm_a * norm_b)) as f32)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_identical_vectors() {
let v = vec![1.0, 2.0, 3.0];
let sim = cosine_similarity(&v, &v).unwrap();
assert!((sim - 1.0).abs() < 1e-6);
}
#[test]
fn test_orthogonal_vectors() {
let a = vec![1.0, 0.0, 0.0];
let b = vec![0.0, 1.0, 0.0];
let sim = cosine_similarity(&a, &b).unwrap();
assert!((sim - 0.0).abs() < 1e-6);
}
#[test]
fn test_opposite_vectors() {
let a = vec![1.0, 0.0];
let b = vec![-1.0, 0.0];
let sim = cosine_similarity(&a, &b).unwrap();
assert!((sim - (-1.0)).abs() < 1e-6);
}
#[test]
fn test_different_lengths_returns_error() {
let a = vec![1.0, 2.0];
let b = vec![1.0];
assert!(cosine_similarity(&a, &b).is_err());
}
#[test]
fn test_zero_vector() {
let a = vec![0.0, 0.0];
let b = vec![1.0, 2.0];
assert_eq!(cosine_similarity(&a, &b).unwrap(), 0.0);
}
}