use std::collections::HashMap;
use crate::error::Result;
use crate::vector::registry::declared_dimension;
use crate::vector::{EmbeddingCodec, ModelName};
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
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 async fn search_vector(
conn: &libsql::Connection,
query_vec: &[f32],
model: &ModelName,
top_k: usize,
) -> Result<Vec<VectorSearchResult>> {
if top_k == 0 {
return Ok(Vec::new());
}
let blob = encode_for_model(conn, model, query_vec).await?;
let sql = format!(
"SELECT e.concept_id, vector_distance_cos(e.embedding, ?1)
FROM vector_top_k('{index}', ?1, ?2) AS t
JOIN {table} AS e ON e.rowid = t.id
ORDER BY 2 ASC",
index = model.index(),
table = model.table(),
);
let mut rows = conn
.query(&sql, libsql::params![blob, top_k as i64])
.await?;
let mut results = Vec::new();
while let Some(row) = rows.next().await? {
results.push(VectorSearchResult {
concept_id: row.get(0)?,
score: row.get::<f64>(1)? as f32,
});
}
Ok(results)
}
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
}