use std::collections::HashMap;
use std::time::Duration;
use crate::error::{DbError, Result};
use crate::vector::search::{decay_factor, rerank_depth};
use crate::vector::{reciprocal_rank_fusion, search_vector, ModelName, VectorSearchResult};
pub const RRF_K: usize = 60;
#[derive(Debug, Clone, PartialEq)]
#[non_exhaustive]
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,
as_of_valid: Option<&str>,
half_life: Option<Duration>,
) -> Result<Vec<(String, f64)>> {
if top_k == 0 || query.trim().is_empty() {
return Ok(Vec::new());
}
let reference = match (half_life, as_of_valid) {
(Some(_), None) => return Err(DbError::HalfLifeWithoutInstant),
(Some(_), Some(t)) => Some(t),
(None, _) => None,
};
let want = match half_life {
Some(_) => rerank_depth(top_k),
None => top_k,
};
let age_column = if half_life.is_some() {
", c.valid_from"
} else {
""
};
let sql = format!(
"SELECT c.id, bm25(concepts_fts) AS rank{age_column}
FROM concepts_fts
JOIN concepts c ON c.rowid_pk = concepts_fts.rowid
WHERE concepts_fts MATCH ?1
AND {visible}
ORDER BY rank ASC, c.id ASC
LIMIT ?2",
visible = crate::vector::search::visible_concept(as_of_valid.map(|_| 3)),
);
let mut params: Vec<libsql::Value> = vec![query.into(), (want as i64).into()];
if let Some(t) = as_of_valid {
params.push(t.into());
}
let mut rows = conn.query(&sql, params).await?;
let mut out: Vec<(String, f64)> = Vec::new();
while let Some(row) = rows.next().await? {
let id: String = row.get(0)?;
let rank: f64 = row.get(1)?;
let rank = match (reference, half_life) {
(Some(reference), Some(half_life)) => {
let valid_from: String = row.get(2)?;
decayed_rank(rank, decay_factor(reference, &valid_from, half_life)?)
}
_ => rank,
};
out.push((id, rank));
}
if half_life.is_some() {
out.sort_by(|a, b| {
a.1.partial_cmp(&b.1)
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| a.0.cmp(&b.0))
});
out.truncate(top_k);
}
Ok(out)
}
fn decayed_rank(rank: f64, factor: f64) -> f64 {
if rank < 0.0 {
rank * factor
} else {
rank
}
}
#[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,
as_of_valid: Option<String>,
half_life: Option<Duration>,
}
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,
as_of_valid: None,
half_life: None,
}
}
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
}
pub fn as_of_valid(mut self, ts: impl Into<String>) -> Self {
self.as_of_valid = Some(ts.into());
self
}
pub fn half_life(mut self, half_life: Duration) -> Self {
self.half_life = Some(half_life);
self
}
fn effective_depth(&self) -> usize {
self.depth.unwrap_or_else(|| rerank_depth(self.top_k))
}
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 at = self.as_of_valid.as_deref();
let vector: Vec<VectorSearchResult> = search_vector(
conn,
&self.query_vector,
&self.model,
depth,
at,
self.half_life,
)
.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, at, self.half_life).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);
fn rank_index(list: &[String]) -> HashMap<&str, usize> {
list.iter()
.enumerate()
.map(|(i, id)| (id.as_str(), i + 1))
.collect()
}
let vector_rank = rank_index(&vector_ids);
let keyword_rank = rank_index(&keyword_ids);
Ok(fused
.into_iter()
.take(self.top_k)
.map(|(concept_id, score)| HybridHit {
vector_rank: vector_rank.get(concept_id.as_str()).copied(),
keyword_rank: keyword_rank.get(concept_id.as_str()).copied(),
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""#);
}
}