use std::collections::HashMap;
use std::time::Duration;
use crate::error::{DbError, Result};
use crate::vector::registry::declared_dimension;
use crate::vector::{EmbeddingCodec, ModelName};
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
#[non_exhaustive]
pub struct VectorSearchResult {
pub concept_id: String,
pub score: f32,
}
#[doc(hidden)]
pub async fn upsert_embedding(
conn: &libsql::Connection,
model: &ModelName,
concept_id: &str,
vector: &[f32],
) -> Result<()> {
let blob = encode_for_model(conn, model, vector).await?;
conn.execute(
&format!(
"INSERT INTO {table} (concept_id, embedding) VALUES (?1, ?2)
ON CONFLICT(concept_id) DO UPDATE SET embedding = excluded.embedding",
table = model.table()
),
libsql::params![concept_id, blob],
)
.await?;
Ok(())
}
pub(crate) async fn upsert_embedding_chunk(
conn: &libsql::Connection,
model: &ModelName,
rows: &[(String, Vec<f32>)],
) -> Result<usize> {
if rows.is_empty() {
return Ok(0);
}
let dim = declared_dimension(conn, model).await?;
let sql = format!(
"INSERT INTO {table} (concept_id, embedding) VALUES (?1, ?2)
ON CONFLICT(concept_id) DO UPDATE SET embedding = excluded.embedding",
table = model.table()
);
let tx = conn
.transaction_with_behavior(libsql::TransactionBehavior::Immediate)
.await?;
let stmt = tx.prepare(&sql).await?;
let res: Result<()> = async {
for (concept_id, vector) in rows {
let blob = EmbeddingCodec::encode(vector, dim, model.as_str())?;
stmt.reset();
stmt.execute(libsql::params![concept_id.as_str(), blob])
.await?;
}
Ok(())
}
.await;
drop(stmt);
match res {
Ok(()) => {
tx.commit().await?;
Ok(rows.len())
}
Err(e) => {
let _ = tx.rollback().await;
Err(e)
}
}
}
pub(crate) const VISIBLE_CONCEPT: &str = "c.retired = 0";
pub(crate) fn visible_concept(at_param: Option<usize>) -> String {
match at_param {
None => VISIBLE_CONCEPT.to_string(),
Some(p) => {
format!("{VISIBLE_CONCEPT} AND c.valid_from <= ?{p} AND ?{p} < c.valid_to")
}
}
}
pub(crate) fn rerank_depth(top_k: usize) -> usize {
(top_k * 5).max(50)
}
pub(crate) fn decay_factor(reference: &str, valid_from: &str, half_life: Duration) -> Result<f64> {
let age = crate::util::timestamp::parse(reference)?
.duration_since(crate::util::timestamp::parse(valid_from)?)
.unwrap_or(Duration::ZERO);
if half_life.is_zero() {
return Ok(if age.is_zero() { 1.0 } else { 0.0 });
}
Ok(0.5f64.powf(age.as_secs_f64() / half_life.as_secs_f64()))
}
pub(crate) fn decayed_distance(distance: f32, factor: f64) -> f32 {
let similarity = ((2.0 - distance as f64) / 2.0).clamp(0.0, 1.0);
(2.0 - 2.0 * similarity * factor) as f32
}
pub async fn search_vector(
conn: &libsql::Connection,
query_vec: &[f32],
model: &ModelName,
top_k: usize,
as_of_valid: Option<&str>,
half_life: Option<Duration>,
) -> Result<Vec<VectorSearchResult>> {
if top_k == 0 {
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 blob = encode_for_model(conn, model, query_vec).await?;
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 e.concept_id, vector_distance_cos(e.embedding, ?1){age_column}
FROM vector_top_k('{index}', ?1, ?2) AS t
JOIN {table} AS e ON e.rowid = t.id
JOIN concepts AS c ON c.id = e.concept_id
WHERE {visible}
ORDER BY 2 ASC
LIMIT ?3",
index = model.index(),
table = model.table(),
visible = visible_concept(as_of_valid.map(|_| 4)),
);
let mut k_prime = want;
let mut indexed: Option<usize> = None;
loop {
let mut params: Vec<libsql::Value> = vec![
blob.clone().into(),
(k_prime as i64).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 results = Vec::new();
while let Some(row) = rows.next().await? {
let hit = VectorSearchResult {
concept_id: row.get(0)?,
score: row.get::<f64>(1)? as f32,
};
let valid_from: Option<String> = match reference {
Some(_) => Some(row.get(2)?),
None => None,
};
results.push((hit, valid_from));
}
if results.len() >= want {
return rank_by_age(results, reference, half_life, top_k);
}
let n = match indexed {
Some(n) => n,
None => {
let n = indexed_rows(conn, model).await?;
indexed = Some(n);
n
}
};
if k_prime >= n {
return rank_by_age(results, reference, half_life, top_k);
}
k_prime = k_prime.saturating_mul(2).min(n);
}
}
fn rank_by_age(
results: Vec<(VectorSearchResult, Option<String>)>,
reference: Option<&str>,
half_life: Option<Duration>,
top_k: usize,
) -> Result<Vec<VectorSearchResult>> {
let (Some(reference), Some(half_life)) = (reference, half_life) else {
return Ok(results.into_iter().map(|(hit, _)| hit).collect());
};
let mut out = Vec::with_capacity(results.len());
for (mut hit, valid_from) in results {
let valid_from = valid_from.unwrap_or_default();
let factor = decay_factor(reference, &valid_from, half_life)?;
hit.score = decayed_distance(hit.score, factor);
out.push(hit);
}
out.sort_by(|a, b| {
a.score
.partial_cmp(&b.score)
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| a.concept_id.cmp(&b.concept_id))
});
out.truncate(top_k);
Ok(out)
}
async fn indexed_rows(conn: &libsql::Connection, model: &ModelName) -> Result<usize> {
let n: i64 = conn
.query(&format!("SELECT COUNT(*) FROM {}", model.table()), ())
.await?
.next()
.await?
.map(|row| row.get(0))
.transpose()?
.unwrap_or(0);
Ok(n.max(0) as usize)
}
async fn encode_for_model(
conn: &libsql::Connection,
model: &ModelName,
vector: &[f32],
) -> Result<Vec<u8>> {
let dim = declared_dimension(conn, model).await?;
EmbeddingCodec::encode(vector, dim, model.as_str())
}
pub fn reciprocal_rank_fusion(
vector_ranks: &[String],
keyword_ranks: &[String],
k: usize,
) -> Vec<(String, f64)> {
let mut scores = HashMap::new();
for (rank, id) in vector_ranks.iter().enumerate() {
let score = 1.0 / ((k + rank + 1) as f64);
*scores.entry(id.clone()).or_insert(0.0) += score;
}
for (rank, id) in keyword_ranks.iter().enumerate() {
let score = 1.0 / ((k + rank + 1) as f64);
*scores.entry(id.clone()).or_insert(0.0) += score;
}
let mut sorted: Vec<_> = scores.into_iter().collect();
sorted.sort_by(|a, b| {
b.1.partial_cmp(&a.1)
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| a.0.cmp(&b.0))
});
sorted
}
#[cfg(test)]
mod decay_tests {
use super::*;
const HOUR: Duration = Duration::from_secs(3600);
const T0: &str = "2026-01-01T00:00:00.000000Z";
const T1: &str = "2026-01-01T01:00:00.000000Z";
const T2: &str = "2026-01-01T02:00:00.000000Z";
#[test]
fn a_half_life_halves_at_a_half_life() {
assert_eq!(decay_factor(T0, T0, HOUR).unwrap(), 1.0);
assert_eq!(decay_factor(T1, T0, HOUR).unwrap(), 0.5);
assert_eq!(decay_factor(T2, T0, HOUR).unwrap(), 0.25);
}
#[test]
fn a_future_validity_is_zero_age_rather_than_negative() {
assert_eq!(decay_factor(T0, T1, HOUR).unwrap(), 1.0);
}
#[test]
fn a_zero_half_life_is_the_limit_and_not_a_nan() {
assert_eq!(decay_factor(T0, T0, Duration::ZERO).unwrap(), 1.0);
assert_eq!(decay_factor(T1, T0, Duration::ZERO).unwrap(), 0.0);
}
#[test]
fn decay_moves_a_hit_away_and_never_toward() {
let near = 0.2_f32;
assert_eq!(decayed_distance(near, 1.0), near);
assert!(decayed_distance(near, 0.5) > near);
assert!(decayed_distance(near, 0.01) > decayed_distance(near, 0.5));
assert!(decayed_distance(near, 0.0) <= 2.0);
}
#[test]
fn one_factor_preserves_the_distance_order() {
assert!(decayed_distance(0.1, 0.7) < decayed_distance(0.9, 0.7));
}
}