akar_function/scalar/
array.rs1use crate::registry::*;
2use akar_common::types::Value;
3
4pub(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 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 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}