use crate::vector::metric::{DistanceMetric, VectorElementType};
#[derive(Debug, Clone)]
pub struct VectorSearchBuilder {
main_table: String,
main_id_column: String,
vector_table: String,
rowid_column: String,
vector_column: String,
element_type: VectorElementType,
metric: DistanceMetric,
query_json: Option<String>,
limit: Option<u32>,
}
impl VectorSearchBuilder {
pub fn new(main_table: impl Into<String>, vector_column: impl Into<String>) -> Self {
let main = main_table.into();
let vector_table = format!("{}_vectors", main);
let rowid = format!("{}_id", singularize(&main));
Self {
main_table: main,
main_id_column: "id".to_string(),
vector_table,
rowid_column: rowid,
vector_column: vector_column.into(),
element_type: VectorElementType::Float4,
metric: DistanceMetric::Cosine,
query_json: None,
limit: None,
}
}
pub fn vector_table(mut self, name: impl Into<String>) -> Self {
self.vector_table = name.into();
self
}
pub fn rowid_column(mut self, name: impl Into<String>) -> Self {
self.rowid_column = name.into();
self
}
pub fn main_id_column(mut self, name: impl Into<String>) -> Self {
self.main_id_column = name.into();
self
}
pub fn element_type(mut self, t: VectorElementType) -> Self {
self.element_type = t;
self
}
pub fn metric(mut self, m: DistanceMetric) -> Self {
self.metric = m;
self
}
pub fn query_json(mut self, json: impl Into<String>) -> Self {
self.query_json = Some(json.into());
self
}
pub fn query_embedding(mut self, embedding: &crate::vector::types::Embedding) -> Self {
self.element_type = VectorElementType::Float4;
self.query_json = Some(embedding.to_json());
self
}
pub fn limit(mut self, n: u32) -> Self {
self.limit = Some(n);
self
}
pub fn to_sql(&self) -> crate::vector::error::VectorResult<String> {
use crate::vector::error::VectorError;
use crate::vector::{escape_sql_literal, quote_ident};
let q = self
.query_json
.as_deref()
.ok_or(VectorError::BuilderIncomplete {
field: "query_json",
})?;
let limit_clause = match self.limit {
Some(n) => format!("\nLIMIT {}", n),
None => String::new(),
};
Ok(format!(
"SELECT {main}.*, \
vector_distance(v.{vector_column}, vector_from_json('{q}', '{et}'), '{metric}', '{et}') AS distance\n\
FROM {vtable} v\n\
JOIN {main} ON {main}.{main_id} = v.{rowid}\n\
ORDER BY distance ASC{limit}",
main = quote_ident(&self.main_table),
main_id = quote_ident(&self.main_id_column),
vector_column = quote_ident(&self.vector_column),
q = escape_sql_literal(q),
et = self.element_type.as_sql(),
metric = self.metric.as_sql(),
vtable = quote_ident(&self.vector_table),
rowid = quote_ident(&self.rowid_column),
limit = limit_clause,
))
}
}
fn singularize(name: &str) -> String {
if name.ends_with('s') && !name.ends_with("ss") {
name[..name.len() - 1].to_string()
} else {
name.to_string()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::vector::types::Embedding;
#[test]
fn test_default_table_and_rowid_naming() {
let sql = VectorSearchBuilder::new("documents", "embedding")
.query_json("[0.1,0.2,0.3]")
.limit(10)
.to_sql()
.unwrap();
assert!(sql.contains("FROM \"documents_vectors\" v"));
assert!(sql.contains("JOIN \"documents\" ON \"documents\".\"id\" = v.\"document_id\""));
assert!(sql.contains("LIMIT 10"));
assert!(sql.contains("ORDER BY distance"));
}
#[test]
fn test_custom_vector_table_and_rowid() {
let sql = VectorSearchBuilder::new("docs", "emb")
.vector_table("docs_vec_tbl")
.rowid_column("doc_ref")
.query_json("[0.1]")
.to_sql()
.unwrap();
assert!(sql.contains("FROM \"docs_vec_tbl\" v"));
assert!(sql.contains("v.\"doc_ref\""));
}
#[test]
fn test_metric_and_element_type_in_sql() {
let sql = VectorSearchBuilder::new("docs", "e")
.metric(DistanceMetric::L2)
.element_type(VectorElementType::Float8)
.query_json("[1.0]")
.to_sql()
.unwrap();
assert!(sql.contains("'l2'"));
assert!(sql.contains("'float8'"));
}
#[test]
fn test_query_embedding_uses_float4() {
let emb = Embedding::new(vec![0.5, 1.5]).unwrap();
let sql = VectorSearchBuilder::new("items", "embedding")
.query_embedding(&emb)
.to_sql()
.unwrap();
assert!(sql.contains("'float4'"));
assert!(sql.contains("[0.5,1.5]"));
}
#[test]
fn test_missing_query_returns_error() {
let result = VectorSearchBuilder::new("documents", "embedding").to_sql();
match result {
Err(crate::vector::error::VectorError::BuilderIncomplete { field }) => {
assert_eq!(field, "query_json");
}
other => panic!("expected BuilderIncomplete error, got {:?}", other),
}
}
#[test]
fn test_identifier_with_embedded_double_quote_is_escaped() {
let sql = VectorSearchBuilder::new("tbl\"evil", "emb")
.query_json("[0.1]")
.to_sql()
.unwrap();
assert!(sql.contains("\"tbl\"\"evil\""));
}
#[test]
fn test_main_id_column_override() {
let sql = VectorSearchBuilder::new("documents", "embedding")
.main_id_column("doc_uid")
.query_json("[0.1]")
.to_sql()
.unwrap();
assert!(sql.contains("ON \"documents\".\"doc_uid\" = v.\"document_id\""));
assert!(!sql.contains("\"documents\".\"id\""));
}
}