use crate::error::Result;
use crate::vector::{reciprocal_rank_fusion, search_vector, ModelName, VectorSearchResult};
pub const RRF_K: usize = 60;
#[derive(Debug, Clone, PartialEq)]
pub struct HybridHit {
pub concept_id: String,
pub score: f64,
pub vector_rank: Option<usize>,
pub keyword_rank: Option<usize>,
}
pub fn escape_fts5_query(input: &str) -> String {
let mut out = String::with_capacity(input.len() + 8);
for token in input.split(|c: char| !c.is_alphanumeric()) {
if token.is_empty() {
continue;
}
if !out.is_empty() {
out.push(' ');
}
out.push('"');
out.push_str(token);
out.push('"');
}
out
}
pub async fn keyword_search(
conn: &libsql::Connection,
query: &str,
top_k: usize,
) -> Result<Vec<(String, f64)>> {
if top_k == 0 || query.trim().is_empty() {
return Ok(Vec::new());
}
let sql = "SELECT c.id, bm25(concepts_fts) AS rank
FROM concepts_fts
JOIN concepts c ON c.rowid = concepts_fts.rowid
WHERE concepts_fts MATCH ?1
AND c.retired = 0
ORDER BY rank ASC, c.id ASC
LIMIT ?2";
let mut rows = conn
.query(sql, libsql::params![query, top_k as i64])
.await?;
let mut out = Vec::new();
while let Some(row) = rows.next().await? {
out.push((row.get(0)?, row.get::<f64>(1)?));
}
Ok(out)
}
#[derive(Debug, Clone)]
pub struct HybridSearch {
model: ModelName,
query_text: String,
query_vector: Vec<f32>,
top_k: usize,
depth: Option<usize>,
rrf_k: usize,
raw_match: bool,
}
impl HybridSearch {
pub fn new(model: ModelName, query_text: impl Into<String>, query_vector: Vec<f32>) -> Self {
Self {
model,
query_text: query_text.into(),
query_vector,
top_k: 10,
depth: None,
rrf_k: RRF_K,
raw_match: false,
}
}
pub fn top_k(mut self, k: usize) -> Self {
self.top_k = k;
self
}
pub fn depth(mut self, depth: usize) -> Self {
self.depth = Some(depth);
self
}
pub fn rrf_k(mut self, k: usize) -> Self {
self.rrf_k = k;
self
}
pub fn raw_match(mut self, raw: bool) -> Self {
self.raw_match = raw;
self
}
fn effective_depth(&self) -> usize {
self.depth.unwrap_or_else(|| (self.top_k * 5).max(50))
}
pub async fn execute(&self, conn: &libsql::Connection) -> Result<Vec<HybridHit>> {
if self.top_k == 0 {
return Ok(Vec::new());
}
let depth = self.effective_depth();
let vector: Vec<VectorSearchResult> =
search_vector(conn, &self.query_vector, &self.model, depth).await?;
let match_expr = if self.raw_match {
self.query_text.clone()
} else {
escape_fts5_query(&self.query_text)
};
let keyword = keyword_search(conn, &match_expr, depth).await?;
let vector_ids: Vec<String> = vector.iter().map(|v| v.concept_id.clone()).collect();
let keyword_ids: Vec<String> = keyword.iter().map(|(id, _)| id.clone()).collect();
let fused = reciprocal_rank_fusion(&vector_ids, &keyword_ids, self.rrf_k);
let rank_of = |list: &[String], id: &str| list.iter().position(|x| x == id).map(|i| i + 1);
Ok(fused
.into_iter()
.take(self.top_k)
.map(|(concept_id, score)| HybridHit {
vector_rank: rank_of(&vector_ids, &concept_id),
keyword_rank: rank_of(&keyword_ids, &concept_id),
concept_id,
score,
})
.collect())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn escaping_turns_a_search_box_into_terms() {
assert_eq!(
escape_fts5_query("bitemporal ledger"),
r#""bitemporal" "ledger""#
);
assert_eq!(escape_fts5_query("cats NOT dogs"), r#""cats" "NOT" "dogs""#);
assert_eq!(escape_fts5_query(r#"a" OR "b"#), r#""a" "OR" "b""#);
assert_eq!(escape_fts5_query("title:macrame"), r#""title" "macrame""#);
assert_eq!(escape_fts5_query("trailing AND"), r#""trailing" "AND""#);
}
#[test]
fn a_query_with_no_terms_escapes_to_nothing() {
assert_eq!(escape_fts5_query("!!! ???"), "");
assert_eq!(escape_fts5_query(""), "");
}
#[test]
fn unicode_survives_escaping() {
assert_eq!(escape_fts5_query("Müller größe"), r#""Müller" "größe""#);
}
}