use std::collections::{BTreeMap, BTreeSet};
use lora_store::{
GraphStorage, IndexConfigValue, LoraVector, RawCoordinate, StoredIndexEntity, StoredIndexKind,
VectorCoordinateType,
};
use super::super::errors::set_eval_error;
use super::super::expr::EvalContext;
use crate::value::LoraValue;
pub(super) fn dispatch<S: GraphStorage>(
op: &str,
args: &[LoraValue],
ctx: &EvalContext<'_, S>,
) -> Option<LoraValue> {
let result = match op {
"fulltext_nodes" => fulltext(args, ctx, StoredIndexEntity::Node),
"fulltext_relationships" => fulltext(args, ctx, StoredIndexEntity::Relationship),
"vector_nodes" => vector(args, ctx, StoredIndexEntity::Node),
"vector_relationships" => vector(args, ctx, StoredIndexEntity::Relationship),
_ => return None,
};
Some(result.unwrap_or_else(|msg| {
set_eval_error(msg);
LoraValue::Null
}))
}
#[cfg(test)]
pub(super) fn known(op: &str) -> Option<()> {
matches!(
op,
"fulltext_nodes" | "fulltext_relationships" | "vector_nodes" | "vector_relationships"
)
.then_some(())
}
fn fulltext<S: GraphStorage>(
args: &[LoraValue],
ctx: &EvalContext<'_, S>,
entity: StoredIndexEntity,
) -> Result<LoraValue, String> {
let name = string_arg(args, 0, "indexName")?;
let query = string_arg(args, 1, "queryString")?;
check_index(ctx, &name, StoredIndexKind::Fulltext, entity)?;
Ok(hits(
ctx.storage.fulltext_search(&name, &query),
None,
entity,
))
}
fn vector<S: GraphStorage>(
args: &[LoraValue],
ctx: &EvalContext<'_, S>,
entity: StoredIndexEntity,
) -> Result<LoraValue, String> {
let name = string_arg(args, 0, "indexName")?;
let k = match args.get(1) {
Some(LoraValue::Int(n)) if *n > 0 => *n as usize,
other => return Err(format!("k must be a positive integer, got {other:?}")),
};
let query = match args.get(2) {
Some(LoraValue::Vector(v)) => v.clone(),
Some(LoraValue::List(items)) => float32_vector(items)?,
other => {
return Err(format!(
"query must be a VECTOR or LIST<NUMBER>, got {other:?}"
))
}
};
let restrict_to = match args.get(3) {
None | Some(LoraValue::Null) => None,
Some(LoraValue::Map(options)) => restrict_to(options)?,
Some(other) => return Err(format!("vector options must be a MAP, got {other:?}")),
};
let def = check_index(ctx, &name, StoredIndexKind::Vector, entity)?;
if let Some(IndexConfigValue::Integer(dim)) = def.options.get("vector.dimensions") {
if query.dimension as i64 != *dim {
return Err(format!(
"query vector has dimension {} but index `{name}` expects {dim}",
query.dimension
));
}
}
let scored = ctx
.storage
.vector_search(&name, &query, k, restrict_to.as_ref());
Ok(hits(scored, Some(k), entity))
}
fn check_index<S: GraphStorage>(
ctx: &EvalContext<'_, S>,
name: &str,
kind: StoredIndexKind,
entity: StoredIndexEntity,
) -> Result<lora_store::IndexDefinition, String> {
let def = ctx
.storage
.get_index(name)
.ok_or_else(|| format!("no {} index named `{name}`", kind.as_str()))?;
if def.kind != kind {
return Err(format!(
"index `{name}` is not a {} index (kind={})",
kind.as_str(),
def.kind.as_str()
));
}
if def.entity != entity {
return Err(format!(
"index `{name}` is on {} entities; procedure expects {}",
def.entity.as_str(),
entity.as_str()
));
}
Ok(def)
}
fn hits(mut scored: Vec<(u64, f64)>, limit: Option<usize>, entity: StoredIndexEntity) -> LoraValue {
scored.sort_by(|a, b| {
b.1.partial_cmp(&a.1)
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| a.0.cmp(&b.0))
});
if let Some(limit) = limit {
scored.truncate(limit);
}
let (key, wrap): (&str, fn(u64) -> LoraValue) = match entity {
StoredIndexEntity::Node => ("node", LoraValue::Node),
StoredIndexEntity::Relationship => ("relationship", LoraValue::Relationship),
};
LoraValue::List(
scored
.into_iter()
.map(|(id, score)| {
LoraValue::Map(BTreeMap::from([
(key.to_string(), wrap(id)),
("score".to_string(), LoraValue::Float(score)),
]))
})
.collect(),
)
}
fn string_arg(args: &[LoraValue], idx: usize, label: &str) -> Result<String, String> {
match args.get(idx) {
Some(LoraValue::String(s)) => Ok(s.clone()),
other => Err(format!("{label} must be a string, got {other:?}")),
}
}
fn float32_vector(items: &[LoraValue]) -> Result<LoraVector, String> {
let raw = items
.iter()
.map(|item| match item {
LoraValue::Int(n) => Ok(RawCoordinate::Int(*n)),
LoraValue::Float(f) => Ok(RawCoordinate::Float(*f)),
other => Err(format!(
"query list elements must be numbers, got {other:?}"
)),
})
.collect::<Result<Vec<_>, _>>()?;
let dim = raw.len() as i64;
LoraVector::try_new(raw, dim, VectorCoordinateType::Float32)
.map_err(|e| format!("invalid query vector: {e}"))
}
fn restrict_to(options: &BTreeMap<String, LoraValue>) -> Result<Option<BTreeSet<u64>>, String> {
let mut out = None;
for (key, value) in options {
match key.as_str() {
"restrictTo" => {
let LoraValue::List(items) = value else {
return Err(format!("`restrictTo` must be a LIST of ids, got {value:?}"));
};
let mut ids = BTreeSet::new();
for item in items {
match item {
LoraValue::Int(n) if *n >= 0 => {
ids.insert(*n as u64);
}
LoraValue::Node(id) | LoraValue::Relationship(id) => {
ids.insert(*id);
}
other => {
return Err(format!("`restrictTo` entries must be ids, got {other:?}"))
}
}
}
out = Some(ids);
}
other => {
return Err(format!(
"unknown vector option `{other}` (known: `restrictTo`)"
))
}
}
}
Ok(out)
}