use astraea_core::error::{AstraeaError, Result};
use astraea_core::types::DistanceMetric;
pub fn cosine_distance(a: &[f32], b: &[f32]) -> Result<f32> {
validate_dimensions(a, b)?;
let mut dot = 0.0_f32;
let mut norm_a = 0.0_f32;
let mut norm_b = 0.0_f32;
for i in 0..a.len() {
dot += a[i] * b[i];
norm_a += a[i] * a[i];
norm_b += b[i] * b[i];
}
let denom = norm_a.sqrt() * norm_b.sqrt();
if denom == 0.0 {
return Ok(1.0);
}
let cosine_sim = (dot / denom).clamp(-1.0, 1.0);
Ok(1.0 - cosine_sim)
}
pub fn euclidean_distance(a: &[f32], b: &[f32]) -> Result<f32> {
validate_dimensions(a, b)?;
let mut sum_sq = 0.0_f32;
for i in 0..a.len() {
let diff = a[i] - b[i];
sum_sq += diff * diff;
}
Ok(sum_sq.sqrt())
}
pub fn dot_product_distance(a: &[f32], b: &[f32]) -> Result<f32> {
validate_dimensions(a, b)?;
let mut dot = 0.0_f32;
for i in 0..a.len() {
dot += a[i] * b[i];
}
Ok(-dot)
}
pub fn compute_distance(metric: DistanceMetric, a: &[f32], b: &[f32]) -> Result<f32> {
match metric {
DistanceMetric::Cosine => cosine_distance(a, b),
DistanceMetric::Euclidean => euclidean_distance(a, b),
DistanceMetric::DotProduct => dot_product_distance(a, b),
}
}
fn validate_dimensions(a: &[f32], b: &[f32]) -> Result<()> {
if a.len() != b.len() {
return Err(AstraeaError::DimensionMismatch {
expected: a.len(),
got: b.len(),
});
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_cosine_distance_identical_vectors() {
let v = vec![1.0, 2.0, 3.0];
let d = cosine_distance(&v, &v).unwrap();
assert!(
d.abs() < 1e-6,
"identical vectors should have cosine distance ~0, got {d}"
);
}
#[test]
fn test_cosine_distance_orthogonal_vectors() {
let a = vec![1.0, 0.0];
let b = vec![0.0, 1.0];
let d = cosine_distance(&a, &b).unwrap();
assert!(
(d - 1.0).abs() < 1e-6,
"orthogonal vectors should have cosine distance ~1.0, got {d}"
);
}
#[test]
fn test_cosine_distance_opposite_vectors() {
let a = vec![1.0, 0.0];
let b = vec![-1.0, 0.0];
let d = cosine_distance(&a, &b).unwrap();
assert!(
(d - 2.0).abs() < 1e-6,
"opposite vectors should have cosine distance ~2.0, got {d}"
);
}
#[test]
fn test_cosine_distance_zero_vector() {
let a = vec![0.0, 0.0, 0.0];
let b = vec![1.0, 2.0, 3.0];
let d = cosine_distance(&a, &b).unwrap();
assert!(
(d - 1.0).abs() < 1e-6,
"zero vector should yield cosine distance 1.0, got {d}"
);
}
#[test]
fn test_cosine_distance_known_value() {
let a = vec![1.0, 0.0];
let b = vec![1.0, 1.0];
let d = cosine_distance(&a, &b).unwrap();
let expected = 1.0 - (1.0 / 2.0_f32.sqrt());
assert!(
(d - expected).abs() < 1e-5,
"expected cosine distance {expected}, got {d}"
);
}
#[test]
fn test_euclidean_distance_identical() {
let v = vec![1.0, 2.0, 3.0];
let d = euclidean_distance(&v, &v).unwrap();
assert!(
d.abs() < 1e-6,
"identical vectors should have L2 distance 0, got {d}"
);
}
#[test]
fn test_euclidean_distance_known_value() {
let a = vec![0.0, 0.0];
let b = vec![3.0, 4.0];
let d = euclidean_distance(&a, &b).unwrap();
assert!(
(d - 5.0).abs() < 1e-6,
"expected euclidean distance 5.0 (3-4-5 triangle), got {d}"
);
}
#[test]
fn test_dot_product_distance_known_value() {
let a = vec![1.0, 2.0, 3.0];
let b = vec![4.0, 5.0, 6.0];
let d = dot_product_distance(&a, &b).unwrap();
assert!(
(d - (-32.0)).abs() < 1e-6,
"expected dot product distance -32.0, got {d}"
);
}
#[test]
fn test_dot_product_distance_identical() {
let v = vec![1.0, 1.0];
let d = dot_product_distance(&v, &v).unwrap();
assert!(
(d - (-2.0)).abs() < 1e-6,
"expected dot product distance -2.0, got {d}"
);
}
#[test]
fn test_dimension_mismatch() {
let a = vec![1.0, 2.0];
let b = vec![1.0, 2.0, 3.0];
assert!(cosine_distance(&a, &b).is_err());
assert!(euclidean_distance(&a, &b).is_err());
assert!(dot_product_distance(&a, &b).is_err());
}
#[test]
fn test_compute_distance_dispatch() {
let a = vec![1.0, 0.0];
let b = vec![0.0, 1.0];
let cos = compute_distance(DistanceMetric::Cosine, &a, &b).unwrap();
let euc = compute_distance(DistanceMetric::Euclidean, &a, &b).unwrap();
let dot = compute_distance(DistanceMetric::DotProduct, &a, &b).unwrap();
assert!((cos - 1.0).abs() < 1e-6);
assert!((euc - 2.0_f32.sqrt()).abs() < 1e-6);
assert!(dot.abs() < 1e-6); }
}