use std::any::Any;
use std::collections::{BTreeMap, BTreeSet};
use anyhow::Result;
use lora_ast::{Expr, ProcedureInvocationKind, StandaloneCall, YieldItem};
use lora_executor::{LoraValue, Row};
use lora_store::{
GraphStorage, GraphStorageMut, IndexDefinition, LoraVector, StoredIndexEntity, StoredIndexKind,
VectorCoordinateType,
};
use crate::database::{
row_projection::{project_yield_items, row_from_columns, ColumnLookupContext, NamedColumn},
Database,
};
use crate::error::{DatabaseOperationError, LoraError, LoraErrorCode};
fn invalid_params_error(message: impl Into<String>) -> anyhow::Error {
DatabaseOperationError::invalid_params(message).into()
}
fn invalid_vector_error(message: impl Into<String>) -> anyhow::Error {
DatabaseOperationError::invalid_vector(message).into()
}
fn validation_error(message: impl Into<String>) -> anyhow::Error {
DatabaseOperationError::validation(message).into()
}
fn not_found_error(message: impl Into<String>) -> anyhow::Error {
DatabaseOperationError::not_found(message).into()
}
fn internal_error(message: impl Into<String>) -> anyhow::Error {
LoraError::new(LoraErrorCode::Internal, message).into()
}
pub(crate) fn is_procedure_call_text(query: &str) -> bool {
let trimmed = query.trim_start();
if !trimmed
.get(..4)
.map(|s| s.eq_ignore_ascii_case("CALL"))
.unwrap_or(false)
{
return false;
}
let rest = trimmed[4..].trim_start();
rest.starts_with("db.index.vector.queryNodes")
|| rest.starts_with("db.index.vector.queryRelationships")
|| rest.starts_with("db.index.fulltext.queryNodes")
|| rest.starts_with("db.index.fulltext.queryRelationships")
}
impl<S> Database<S>
where
S: GraphStorage + GraphStorageMut + Any + Clone + Send + Sync + 'static,
{
pub(crate) fn execute_procedure_call(
&self,
call: &StandaloneCall,
params: BTreeMap<String, LoraValue>,
) -> Result<Vec<Row>> {
let (name, args) = invocation_parts(&call.procedure)?;
let qualified = name.parts.join(".");
let procedure = Procedure::parse(&qualified)?;
let raw_rows = match procedure.kind {
ProcedureKind::Vector => self.vector_query(args, ¶ms, procedure.entity)?,
ProcedureKind::Fulltext => self.fulltext_query(args, ¶ms, procedure.entity)?,
};
project_yield(raw_rows, &call.yield_items, call.yield_all)
}
fn vector_query(
&self,
args: &[Expr],
params: &BTreeMap<String, LoraValue>,
entity: StoredIndexEntity,
) -> Result<Vec<Row>> {
let (index_name, k, query_vec, restrict_to) = parse_vector_args(args, params)?;
let snapshot = self.read_store();
let def = snapshot
.get_index(&index_name)
.ok_or_else(|| not_found_error(format!("no vector index named `{index_name}`")))?;
validate_procedure_index(&def, StoredIndexKind::Vector, entity, "vector")?;
if def.label.is_none() {
return Err(internal_error(format!(
"vector index `{index_name}` has no label/type"
)));
}
if def.properties.is_empty() {
return Err(internal_error(format!(
"vector index `{index_name}` has no property column"
)));
}
let expected_dim = expected_dimension(&def);
if let Some(dim) = expected_dim {
if query_vec.dimension != dim {
return Err(invalid_vector_error(format!(
"query vector has dimension {} but index `{index_name}` expects {}",
query_vec.dimension, dim
)));
}
}
let scored = snapshot.vector_search(&index_name, &query_vec, k, restrict_to.as_ref());
Ok(hydrated(scored_rows(scored, Some(k), entity), &*snapshot))
}
fn fulltext_query(
&self,
args: &[Expr],
params: &BTreeMap<String, LoraValue>,
entity: StoredIndexEntity,
) -> Result<Vec<Row>> {
let (index_name, query_text) = parse_fulltext_args(args, params)?;
let snapshot = self.read_store();
let def = snapshot
.get_index(&index_name)
.ok_or_else(|| not_found_error(format!("no fulltext index named `{index_name}`")))?;
validate_procedure_index(&def, StoredIndexKind::Fulltext, entity, "fulltext")?;
let scored = snapshot.fulltext_search(&index_name, &query_text);
Ok(hydrated(scored_rows(scored, None, entity), &*snapshot))
}
}
#[derive(Debug, Clone, Copy)]
struct Procedure {
kind: ProcedureKind,
entity: StoredIndexEntity,
}
impl Procedure {
fn parse(qualified: &str) -> Result<Self> {
match qualified {
"db.index.vector.queryNodes" => Ok(Self {
kind: ProcedureKind::Vector,
entity: StoredIndexEntity::Node,
}),
"db.index.vector.queryRelationships" => Ok(Self {
kind: ProcedureKind::Vector,
entity: StoredIndexEntity::Relationship,
}),
"db.index.fulltext.queryNodes" => Ok(Self {
kind: ProcedureKind::Fulltext,
entity: StoredIndexEntity::Node,
}),
"db.index.fulltext.queryRelationships" => Ok(Self {
kind: ProcedureKind::Fulltext,
entity: StoredIndexEntity::Relationship,
}),
other => Err(validation_error(format!("unknown procedure: {other}"))),
}
}
}
#[derive(Debug, Clone, Copy)]
enum ProcedureKind {
Vector,
Fulltext,
}
fn expected_dimension(def: &IndexDefinition) -> Option<usize> {
def.options.get("vector.dimensions").and_then(|v| match v {
lora_store::IndexConfigValue::Integer(n) if *n > 0 => Some(*n as usize),
_ => None,
})
}
fn validate_procedure_index(
def: &IndexDefinition,
expected_kind: StoredIndexKind,
expected_entity: StoredIndexEntity,
procedure_kind: &str,
) -> Result<()> {
if def.kind != expected_kind {
return Err(validation_error(format!(
"index `{}` is not a {} index (kind={})",
def.name,
expected_kind.as_str(),
def.kind.as_str()
)));
}
if def.entity != expected_entity {
return Err(validation_error(format!(
"{procedure_kind} index `{}` is on {} entities; procedure expects {}",
def.name,
def.entity.as_str(),
expected_entity.as_str()
)));
}
Ok(())
}
fn scored_rows(
mut scored: Vec<(u64, f64)>,
limit: Option<usize>,
entity: StoredIndexEntity,
) -> Vec<Row> {
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);
}
scored
.into_iter()
.map(|(id, score)| {
let (name, value) = match entity {
StoredIndexEntity::Node => ("node", LoraValue::Node(id)),
StoredIndexEntity::Relationship => ("relationship", LoraValue::Relationship(id)),
};
row_from_columns([
NamedColumn::new(name, value),
NamedColumn::new("score", LoraValue::Float(score)),
])
})
.collect()
}
fn hydrated<G: GraphStorage>(rows: Vec<Row>, storage: &G) -> Vec<Row> {
rows.into_iter()
.map(|row| lora_executor::hydrate_row(row, storage))
.collect()
}
fn invocation_parts(
invocation: &ProcedureInvocationKind,
) -> Result<(&lora_ast::ProcedureName, &[Expr])> {
match invocation {
ProcedureInvocationKind::Explicit(call) => Ok((&call.name, call.args.as_slice())),
ProcedureInvocationKind::Implicit(name) => Ok((name, &[])),
}
}
fn parse_fulltext_args(
args: &[Expr],
params: &BTreeMap<String, LoraValue>,
) -> Result<(String, String)> {
if args.len() < 2 || args.len() > 3 {
return Err(invalid_params_error(format!(
"fulltext procedure expects 2 or 3 arguments (indexName, queryString, options? ); got {}",
args.len()
)));
}
let name = eval_string_arg(&args[0], params, "indexName")?;
let query = eval_string_arg(&args[1], params, "queryString")?;
Ok((name, query))
}
fn parse_vector_args(
args: &[Expr],
params: &BTreeMap<String, LoraValue>,
) -> Result<(String, usize, LoraVector, Option<BTreeSet<u64>>)> {
if args.len() < 3 || args.len() > 4 {
return Err(invalid_params_error(format!(
"vector procedure expects 3 or 4 arguments (indexName, k, query, options?); got {}",
args.len()
)));
}
let name = eval_string_arg(&args[0], params, "indexName")?;
let k = eval_usize_arg(&args[1], params, "k")?;
if k == 0 {
return Err(invalid_params_error("k must be positive"));
}
let query = eval_vector_arg(&args[2], params)?;
let restrict_to = if args.len() == 4 {
parse_vector_options(&args[3], params)?
} else {
None
};
Ok((name, k, query, restrict_to))
}
fn parse_vector_options(
expr: &Expr,
params: &BTreeMap<String, LoraValue>,
) -> Result<Option<BTreeSet<u64>>> {
let entries = match expr {
Expr::Map(items, _) => items.clone(),
other => {
return Err(invalid_params_error(format!(
"vector procedure options must be a MAP literal like {{restrictTo: [...]}}, got {other:?}"
)));
}
};
let mut restrict_to: Option<BTreeSet<u64>> = None;
for (key, value) in entries {
match key.as_str() {
"restrictTo" => {
let resolved = resolve_literal(&value, params)?;
restrict_to = Some(coerce_id_set(&resolved)?);
}
other => {
return Err(invalid_params_error(format!(
"unknown option `{other}` for db.index.vector.queryNodes (known: `restrictTo`)"
)));
}
}
}
Ok(restrict_to)
}
fn coerce_id_set(value: &LoraValue) -> Result<BTreeSet<u64>> {
let items = match value {
LoraValue::List(xs) => xs,
other => {
return Err(invalid_params_error(format!(
"`restrictTo` must be a LIST of node ids, got {other:?}"
)));
}
};
let mut out = BTreeSet::new();
for item in items {
match item {
LoraValue::Int(n) if *n >= 0 => {
out.insert(*n as u64);
}
LoraValue::Node(id) => {
out.insert(*id);
}
LoraValue::Relationship(id) => {
out.insert(*id);
}
other => {
return Err(invalid_params_error(format!(
"`restrictTo` entries must be non-negative integers or node/relationship references, got {other:?}"
)));
}
}
}
Ok(out)
}
fn eval_string_arg(
expr: &Expr,
params: &BTreeMap<String, LoraValue>,
label: &str,
) -> Result<String> {
match resolve_literal(expr, params)? {
LoraValue::String(s) => Ok(s),
other => Err(invalid_params_error(format!(
"{label} must be a string, got {other:?}"
))),
}
}
fn eval_usize_arg(expr: &Expr, params: &BTreeMap<String, LoraValue>, label: &str) -> Result<usize> {
match resolve_literal(expr, params)? {
LoraValue::Int(n) if n >= 0 => usize::try_from(n).map_err(|_| {
invalid_params_error(format!("{label} is too large for this platform: {n}"))
}),
other => Err(invalid_params_error(format!(
"{label} must be a non-negative integer, got {other:?}"
))),
}
}
fn eval_vector_arg(expr: &Expr, params: &BTreeMap<String, LoraValue>) -> Result<LoraVector> {
let value = resolve_literal(expr, params)?;
match value {
LoraValue::Vector(v) => Ok(v),
LoraValue::List(items) => list_to_float32_vector(&items),
other => Err(invalid_vector_error(format!(
"query must be a VECTOR or LIST<NUMBER>; got {other:?}"
))),
}
}
fn list_to_float32_vector(items: &[LoraValue]) -> Result<LoraVector> {
let mut raw = Vec::with_capacity(items.len());
for item in items {
let coord = match item {
LoraValue::Int(n) => lora_store::RawCoordinate::Int(*n),
LoraValue::Float(f) => lora_store::RawCoordinate::Float(*f),
other => {
return Err(invalid_vector_error(format!(
"query list elements must be INTEGER or FLOAT; got {other:?}"
)))
}
};
raw.push(coord);
}
let dim = items.len() as i64;
LoraVector::try_new(raw, dim, VectorCoordinateType::Float32)
.map_err(|e| invalid_vector_error(format!("invalid query vector: {e}")))
}
fn resolve_literal(expr: &Expr, params: &BTreeMap<String, LoraValue>) -> Result<LoraValue> {
match expr {
Expr::Integer(n, _) => Ok(LoraValue::Int(*n)),
Expr::Float(f, _) => Ok(LoraValue::Float(*f)),
Expr::String(s, _) => Ok(LoraValue::String(s.clone())),
Expr::Bool(b, _) => Ok(LoraValue::Bool(*b)),
Expr::Null(_) => Ok(LoraValue::Null),
Expr::Parameter(name, _) => params
.get(name)
.cloned()
.ok_or_else(|| invalid_params_error(format!("parameter `${name}` not supplied"))),
Expr::List(items, _) => {
let mut out = Vec::with_capacity(items.len());
for item in items {
out.push(resolve_literal(item, params)?);
}
Ok(LoraValue::List(out))
}
Expr::Unary {
op: lora_ast::UnaryOp::Neg,
expr: inner,
..
} => match resolve_literal(inner, params)? {
LoraValue::Int(n) => Ok(LoraValue::Int(-n)),
LoraValue::Float(f) => Ok(LoraValue::Float(-f)),
other => Err(invalid_params_error(format!("cannot negate {other:?}"))),
},
Expr::FunctionCall { name, args, .. } if matches_vector_ctor(name) => {
resolve_vector_ctor(args, params)
}
other => Err(invalid_params_error(format!(
"procedure arguments must be literals or $params; got {other:?}"
))),
}
}
fn matches_vector_ctor(name: &[String]) -> bool {
match name {
[n] if n.eq_ignore_ascii_case("vector") => true,
[ns, op] if ns.eq_ignore_ascii_case("vector") && op.eq_ignore_ascii_case("new") => true,
_ => false,
}
}
fn resolve_vector_ctor(args: &[Expr], params: &BTreeMap<String, LoraValue>) -> Result<LoraValue> {
if args.len() != 3 {
return Err(invalid_params_error(format!(
"vector() expects 3 arguments, got {}",
args.len()
)));
}
let LoraValue::List(values) = resolve_literal(&args[0], params)? else {
return Err(invalid_vector_error(
"vector() first argument must be a list",
));
};
let dim_value = resolve_literal(&args[1], params)?;
let LoraValue::Int(dim) = dim_value else {
return Err(invalid_vector_error(format!(
"vector() dimension must be integer, got {dim_value:?}"
)));
};
let coord = coord_type_from_expr(&args[2], params)?;
let mut raw = Vec::with_capacity(values.len());
for v in values {
match v {
LoraValue::Int(n) => raw.push(lora_store::RawCoordinate::Int(n)),
LoraValue::Float(f) => raw.push(lora_store::RawCoordinate::Float(f)),
other => {
return Err(invalid_vector_error(format!(
"vector() list elements must be INTEGER or FLOAT; got {other:?}"
)))
}
}
}
let v = LoraVector::try_new(raw, dim, coord)
.map_err(|e| invalid_vector_error(format!("vector(): {e}")))?;
Ok(LoraValue::Vector(v))
}
fn coord_type_from_expr(
expr: &Expr,
params: &BTreeMap<String, LoraValue>,
) -> Result<VectorCoordinateType> {
let name = match expr {
Expr::Variable(v) => v.name.clone(),
Expr::String(s, _) => s.clone(),
Expr::Parameter(_, _) => {
let value = resolve_literal(expr, params)?;
let LoraValue::String(s) = value else {
return Err(invalid_params_error(format!(
"coordinate type parameter must be a string; got {value:?}"
)));
};
s
}
other => {
return Err(invalid_params_error(format!(
"invalid coordinate type expression: {other:?}"
)));
}
};
VectorCoordinateType::parse(&name)
.ok_or_else(|| invalid_params_error(format!("unknown coordinate type `{name}`")))
}
fn project_yield(rows: Vec<Row>, items: &[YieldItem], yield_all: bool) -> Result<Vec<Row>> {
project_yield_items(rows, items, yield_all, ColumnLookupContext::ProcedureYield)
}