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            let dot: f64 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum();
42            let norm_a: f64 = a.iter().map(|x| x * x).sum::<f64>().sqrt();
43            let norm_b: f64 = b.iter().map(|x| x * x).sum::<f64>().sqrt();
44            if norm_a == 0.0 || norm_b == 0.0 {
45                return Ok(Value::Double(1.0));
46            }
47            Ok(Value::Double(dot / (norm_a * norm_b)))
48        }
49        ArrayOp::Distance => {
50            let sum_sq: f64 = a.iter().zip(b.iter()).map(|(x, y)| (x - y) * (x - y)).sum();
51            Ok(Value::Double(sum_sq.sqrt()))
52        }
53        ArrayOp::InnerProduct | ArrayOp::DotProduct => {
54            let dot: f64 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum();
55            Ok(Value::Double(dot))
56        }
57        ArrayOp::CrossProduct => {
58            if a.len() != 3 || b.len() != 3 {
59                return Err("Cross product requires 3D arrays".into());
60            }
61            let result = vec![
62                Value::Double(a[1] * b[2] - a[2] * b[1]),
63                Value::Double(a[2] * b[0] - a[0] * b[2]),
64                Value::Double(a[0] * b[1] - a[1] * b[0]),
65            ];
66            Ok(Value::List(result))
67        }
68        ArrayOp::SquaredDistance => {
69            let sum_sq: f64 = a.iter().zip(b.iter()).map(|(x, y)| (x - y) * (x - y)).sum();
70            Ok(Value::Double(sum_sq))
71        }
72        ArrayOp::Intersect => {
73            // Intersect two numeric arrays, preserving only common elements.
74            let mut set_b = std::collections::HashSet::new();
75            for item in b {
76                set_b.insert(item.to_bits());
77            }
78            let mut result = Vec::new();
79            for item in a {
80                if set_b.contains(&item.to_bits()) {
81                    result.push(Value::Double(item));
82                }
83            }
84            Ok(Value::List(result))
85        }
86    }
87}