Skip to main content

akar_function/scalar/
array.rs

1use crate::registry::*;
2use akar_common::types::Value;
3
4// ==================== Array Math Functions ====================
5
6/// Evaluate an array math function: cosine_similarity, distance, inner_product,
7/// cross_product, squared_distance.
8pub(crate) fn evaluate_array(op: ArrayOp, args: &[Value]) -> Result<Value, String> {
9    if args.len() < 2 {
10        return Err(format!("Array function {:?} requires 2 arguments", op));
11    }
12
13    /// Extract a Vec<f64> from a Value::List or return an error.
14    fn extract_f64s(v: &Value) -> Result<Vec<f64>, String> {
15        match v {
16            Value::List(items) => items
17                .iter()
18                .map(|item| match item {
19                    Value::Int64(i) => Ok(*i as f64),
20                    Value::Double(f) => Ok(*f),
21                    Value::Float(f) => Ok(*f as f64),
22                    _ => Err(format!(
23                        "Expected numeric value in array, got {:?}",
24                        item.logical_type()
25                    )),
26                })
27                .collect(),
28            _ => Err("Expected list/array".into()),
29        }
30    }
31
32    let a = extract_f64s(&args[0])?;
33    let b = extract_f64s(&args[1])?;
34
35    if a.len() != b.len() {
36        return Err("Arrays must have the same length".into());
37    }
38
39    match op {
40        ArrayOp::CosineSimilarity => {
41            // Single-pass optimization for dot product and squared norms
42            let mut dot = 0.0;
43            let mut sq_a = 0.0;
44            let mut sq_b = 0.0;
45            for (x, y) in a.iter().zip(b.iter()) {
46                dot += x * y;
47                sq_a += x * x;
48                sq_b += y * y;
49            }
50            if sq_a == 0.0 || sq_b == 0.0 {
51                return Ok(Value::Double(1.0));
52            }
53            Ok(Value::Double(dot / (sq_a.sqrt() * sq_b.sqrt())))
54        }
55        ArrayOp::Distance => {
56            let sum_sq: f64 = a.iter().zip(b.iter()).map(|(x, y)| (x - y) * (x - y)).sum();
57            Ok(Value::Double(sum_sq.sqrt()))
58        }
59        ArrayOp::InnerProduct | ArrayOp::DotProduct => {
60            let dot: f64 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum();
61            Ok(Value::Double(dot))
62        }
63        ArrayOp::CrossProduct => {
64            if a.len() != 3 || b.len() != 3 {
65                return Err("Cross product requires 3D arrays".into());
66            }
67            let result = vec![
68                Value::Double(a[1] * b[2] - a[2] * b[1]),
69                Value::Double(a[2] * b[0] - a[0] * b[2]),
70                Value::Double(a[0] * b[1] - a[1] * b[0]),
71            ];
72            Ok(Value::List(result))
73        }
74        ArrayOp::SquaredDistance => {
75            let sum_sq: f64 = a.iter().zip(b.iter()).map(|(x, y)| (x - y) * (x - y)).sum();
76            Ok(Value::Double(sum_sq))
77        }
78        ArrayOp::Intersect => {
79            // Intersect two numeric arrays, preserving only common elements.
80            let mut set_b = std::collections::HashSet::new();
81            for item in b {
82                set_b.insert(item.to_bits());
83            }
84            let mut result = Vec::new();
85            for item in a {
86                if set_b.contains(&item.to_bits()) {
87                    result.push(Value::Double(item));
88                }
89            }
90            Ok(Value::List(result))
91        }
92    }
93}