use super::super::helpers::*;
use super::super::*;
use crate::datatypes::values::Value;
struct PairSums {
dot: f64,
norm_a_sq: f64,
norm_b_sq: f64,
}
fn element(fname: &str, which: &str, index: usize, v: &Value) -> Result<f64, String> {
match v {
Value::Int64(i) => Ok(*i as f64),
Value::Float64(f) => Ok(*f),
other => Err(format!(
"{fname}(): {which} vector element {index} must be a number, got {}",
other.type_name()
)),
}
}
fn pair_sums(fname: &str, a: &[Value], b: &[Value]) -> Result<PairSums, String> {
let mut sums = PairSums {
dot: 0.0,
norm_a_sq: 0.0,
norm_b_sq: 0.0,
};
for (index, (av, bv)) in a.iter().zip(b.iter()).enumerate() {
let x = element(fname, "first", index, av)?;
let y = element(fname, "second", index, bv)?;
sums.dot += x * y;
sums.norm_a_sq += x * x;
sums.norm_b_sq += y * y;
}
Ok(sums)
}
impl<'a> CypherExecutor<'a> {
fn vector_arg(
&self,
fname: &str,
which: &str,
expr: &Expression,
row: &ResultRow,
) -> Result<Option<Vec<Value>>, String> {
match self.evaluate_expression(expr, row)? {
Value::List(items) => Ok(Some(items)),
Value::Null => Ok(None),
Value::String(s) => {
let trimmed = s.trim();
if trimmed.starts_with('[') && trimmed.ends_with(']') {
Ok(Some(parse_list_value(&Value::String(s))))
} else {
Err(format!(
"{fname}(): {which} argument must be a list of numbers, \
got a string that is not a bracketed list"
))
}
}
other => Err(format!(
"{fname}(): {which} argument must be a list of numbers, got {}",
other.type_name()
)),
}
}
pub(super) fn eval_vector_fn(
&self,
name: &str,
args: &[Expression],
row: &ResultRow,
) -> Result<Option<Value>, String> {
let result: Result<Value, String> = match name {
"dot" | "cosine" => {
if args.len() != 2 {
return Err(format!("{name}() requires 2 arguments: {name}(a, b)"));
}
let a = self.vector_arg(name, "first", &args[0], row)?;
let b = self.vector_arg(name, "second", &args[1], row)?;
let (Some(a), Some(b)) = (a, b) else {
return Ok(Some(Value::Null));
};
if a.len() != b.len() {
return Err(format!(
"{name}(): vectors must have the same length, got {} and {}",
a.len(),
b.len()
));
}
let sums = pair_sums(name, &a, &b)?;
if name == "dot" {
Ok(Value::Float64(sums.dot))
} else {
let denom = (sums.norm_a_sq * sums.norm_b_sq).sqrt();
if denom == 0.0 {
return Ok(Some(Value::Null));
}
let cos = sums.dot / denom;
if cos.is_finite() {
Ok(Value::Float64(cos))
} else {
Ok(Value::Null)
}
}
}
"norm" => {
if args.len() != 1 {
return Err("norm() requires 1 argument: norm(a)".into());
}
let Some(a) = self.vector_arg(name, "first", &args[0], row)? else {
return Ok(Some(Value::Null));
};
let mut sum = 0.0f64;
for (index, v) in a.iter().enumerate() {
let x = element(name, "first", index, v)?;
sum += x * x;
}
Ok(Value::Float64(sum.sqrt()))
}
_ => return Ok(None),
};
result.map(Some)
}
}