use std::path::Path;
use std::sync::OnceLock;
use anyhow::{Context, Result};
use rusqlite::Connection;
pub const MAX_COSINE_DISTANCE: f64 = 0.7;
const MAX_KNN_CANDIDATES: usize = 10_000;
static VEC_EXT_REGISTERED: OnceLock<()> = OnceLock::new();
fn ensure_extension_registered() {
VEC_EXT_REGISTERED.get_or_init(|| {
let init_fn: unsafe extern "C" fn() = sqlite_vec::sqlite3_vec_init;
let entry: rusqlite::auto_extension::RawAutoExtension =
unsafe { std::mem::transmute(init_fn as *const ()) };
unsafe {
let _ = rusqlite::auto_extension::register_auto_extension(entry);
}
});
}
const VECTOR_TABLE: &str = "vectors";
pub struct VecDb {
conn: Connection,
}
#[derive(Clone)]
pub struct KnnRow {
pub node_json: String,
pub distance: f64,
}
impl VecDb {
pub fn open(path: impl AsRef<Path>) -> Result<Self> {
ensure_extension_registered();
let conn = Connection::open(path.as_ref())
.context("打开向量数据库失败")?;
conn.pragma_update(None, "journal_mode", "WAL")
.context("设置 WAL 模式失败")?;
conn.busy_timeout(std::time::Duration::from_secs(5))
.context("设置 busy_timeout 失败")?;
Ok(Self { conn })
}
fn table_exists(&self) -> Result<bool> {
let sql = "SELECT COUNT(*) FROM sqlite_master WHERE type='table' AND name=?1";
let count: i64 = self
.conn
.query_row(sql, [VECTOR_TABLE], |r| r.get(0))
.context("查询虚表存在性失败")?;
Ok(count > 0)
}
pub fn table_dimension(&self) -> Result<Option<usize>> {
if !self.table_exists()? {
return Ok(None);
}
let sql = "SELECT sql FROM sqlite_master WHERE type='table' AND name=?1";
let ddl: String = self
.conn
.query_row(sql, [VECTOR_TABLE], |r| r.get(0))
.context("读取虚表定义失败")?;
let start = ddl.find("float[").map(|i| i + 6);
let Some(start) = start else {
anyhow::bail!("vec0 虚表定义缺少 float[N] 声明: {ddl}");
};
let end = ddl[start..].find(']').map(|i| start + i);
let Some(end) = end else {
anyhow::bail!("vec0 虚表定义缺少 ] 结束符: {ddl}");
};
ddl[start..end]
.parse::<usize>()
.map(Some)
.with_context(|| format!("解析 vec0 维度失败: {}", &ddl[start..end]))
}
fn create_table(&self, dim: usize) -> Result<()> {
let sql = format!(
"CREATE VIRTUAL TABLE {VECTOR_TABLE} USING vec0(\
embedding float[{dim}] distance_metric=cosine,\
file_path TEXT,\
node_json TEXT\
)"
);
self.conn
.execute_batch(&sql)
.with_context(|| format!("创建 vec0 虚表失败(dim={dim})"))
}
fn rebuild_table(&self, dim: usize) -> Result<()> {
self.conn
.execute_batch(&format!("DROP TABLE IF EXISTS {VECTOR_TABLE};"))
.context("删除旧 vec0 虚表失败")?;
self.create_table(dim)
}
fn ensure_table(&self, dim: usize) -> Result<()> {
match self.table_dimension()? {
None => self.create_table(dim),
Some(existing) if existing != dim => {
tracing::warn!(
"向量维度变化({} → {}),重建语义索引(embedding 模型变更需全量重新索引)",
existing, dim
);
self.rebuild_table(dim)
}
Some(_) => Ok(()),
}
}
pub fn insert_batch(&self, items: &[(String, String, Vec<f32>)]) -> Result<()> {
if items.is_empty() {
return Ok(());
}
let dim = items[0].2.len();
if dim == 0 {
anyhow::bail!("空向量无法入库(维度为 0)");
}
self.ensure_table(dim)?;
let mut stmt = self.conn.prepare(&format!(
"INSERT INTO {VECTOR_TABLE}(rowid, embedding, file_path, node_json) \
VALUES (?1, ?2, ?3, ?4)"
))?;
for (i, (file_path, node_json, vector)) in items.iter().enumerate() {
let json = vector_to_json(vector);
stmt.execute(rusqlite::params![
(i + 1) as i64,
json,
crate::incremental::norm_sep(file_path),
node_json,
])
.with_context(|| format!("插入向量失败: {file_path}"))?;
}
Ok(())
}
pub fn knn(&self, query_json: &str, limit: usize, max_distance: f64) -> Result<Vec<KnnRow>> {
if !self.table_exists()? {
return Ok(Vec::new());
}
let mut sample = limit.max(1);
let mut all: Vec<KnnRow> = Vec::new();
let mut seen: std::collections::HashSet<String> = std::collections::HashSet::new();
loop {
let sql = format!(
"SELECT node_json, distance FROM {VECTOR_TABLE} \
WHERE embedding MATCH ?1 \
ORDER BY distance \
LIMIT {sample}"
);
let mut stmt = self.conn.prepare(&sql).context("准备 KNN 查询失败")?;
let rows: Vec<KnnRow> = stmt
.query_map([query_json], |r| {
Ok(KnnRow {
node_json: r.get(0)?,
distance: r.get(1)?,
})
})
.context("执行 KNN 查询失败")?
.collect::<rusqlite::Result<Vec<_>>>()
.context("读取 KNN 结果失败")?;
for row in rows.iter().filter(|r| r.distance <= max_distance) {
if seen.insert(row.node_json.clone()) {
all.push(row.clone());
}
}
if rows.len() < sample || rows.last().map(|r| r.distance > max_distance).unwrap_or(true) || sample >= MAX_KNN_CANDIDATES {
break;
}
sample = (sample * 2).min(MAX_KNN_CANDIDATES);
}
Ok(all)
}
pub fn remove_by_file(&self, file_path: &str) -> Result<usize> {
if !self.table_exists()? {
return Ok(0);
}
let sql = format!("DELETE FROM {VECTOR_TABLE} WHERE file_path = ?1");
let count = self
.conn
.execute(&sql, rusqlite::params![crate::incremental::norm_sep(file_path)])
.context("删除向量失败")?;
Ok(count)
}
pub fn clear(&self) -> Result<()> {
if !self.table_exists()? {
return Ok(());
}
self.conn
.execute_batch(&format!("DELETE FROM {VECTOR_TABLE};"))
.context("清空向量失败")
}
pub fn entry_count(&self) -> Result<usize> {
if !self.table_exists()? {
return Ok(0);
}
let sql = format!("SELECT COUNT(*) FROM {VECTOR_TABLE}");
let count: i64 = self
.conn
.query_row(&sql, [], |r| r.get(0))
.context("查询向量数失败")?;
Ok(count as usize)
}
}
fn vector_to_json(v: &[f32]) -> String {
let parts: Vec<String> = v.iter().map(|f| format!("{f}")).collect();
format!("[{}]", parts.join(","))
}
#[cfg(test)]
mod tests {
use super::*;
fn tmp_db(tag: &str) -> std::path::PathBuf {
use std::sync::atomic::{AtomicU64, Ordering};
static SEQ: AtomicU64 = AtomicU64::new(0);
let id = SEQ.fetch_add(1, Ordering::Relaxed);
let mut p = std::env::temp_dir();
p.push(format!("vecdb_{}_{}_{}.db", tag, std::process::id(), id));
let _ = std::fs::remove_file(&p);
p
}
fn make_items() -> Vec<(String, String, Vec<f32>)> {
vec![
("src/a.rs".into(), r#"{"name":"alpha"}"#.into(), vec![1.0, 0.0, 0.0]),
("src/b.rs".into(), r#"{"name":"beta"}"#.into(), vec![0.0, 1.0, 0.0]),
("src/c.rs".into(), r#"{"name":"gamma"}"#.into(), vec![-1.0, 0.0, 0.0]),
]
}
#[test]
fn test_insert_and_knn_cosine_ranking() {
let db = VecDb::open(tmp_db("knn")).unwrap();
db.insert_batch(&make_items()).unwrap();
let rows = db.knn("[0.9,0.1,0]", 10, MAX_COSINE_DISTANCE).unwrap();
assert_eq!(rows.len(), 1, "只有 alpha 在 0.3 相似度阈值内");
assert!(rows[0].node_json.contains("alpha"));
assert!(rows[0].distance < 0.01);
}
#[test]
fn test_threshold_0_3_maps_to_distance_0_7() {
assert_eq!(MAX_COSINE_DISTANCE, 0.7);
}
#[test]
fn test_knn_without_threshold_returns_all() {
let db = VecDb::open(tmp_db("knn_all")).unwrap();
db.insert_batch(&make_items()).unwrap();
let rows = db.knn("[0.9,0.1,0]", 10, 2.0).unwrap();
assert_eq!(rows.len(), 3);
assert!(rows[0].node_json.contains("alpha"));
assert!(rows[2].node_json.contains("gamma"));
assert!(rows[0].distance < rows[1].distance);
}
#[test]
fn test_knn_expands_sample_to_cover_threshold() {
let db = VecDb::open(tmp_db("expand")).unwrap();
let items: Vec<(String, String, Vec<f32>)> = (0..30)
.map(|i| (format!("src/f{i}.rs"), format!(r#"{{"name":"f{i}"}}"#), vec![1.0, 0.0, 0.0]))
.collect();
db.insert_batch(&items).unwrap();
let rows = db.knn("[1,0,0]", 5, MAX_COSINE_DISTANCE).unwrap();
assert_eq!(rows.len(), 30, "阈值内全部候选必须返回(扩样不能截断)");
}
#[test]
fn test_delete_by_file_and_count() {
let db = VecDb::open(tmp_db("del")).unwrap();
db.insert_batch(&make_items()).unwrap();
assert_eq!(db.entry_count().unwrap(), 3);
let removed = db.remove_by_file("src/b.rs").unwrap();
assert_eq!(removed, 1);
assert_eq!(db.entry_count().unwrap(), 2);
db.clear().unwrap();
assert_eq!(db.entry_count().unwrap(), 0);
}
#[test]
fn test_dimension_mismatch_rebuilds_table() {
let db = VecDb::open(tmp_db("dim")).unwrap();
db.insert_batch(&make_items()).unwrap();
db.insert_batch(&[(
"src/d.rs".into(),
r#"{"name":"delta"}"#.into(),
vec![1.0, 0.0, 0.0, 0.0],
)])
.unwrap();
assert_eq!(db.entry_count().unwrap(), 1, "维度变化重建后只剩新数据");
let rows = db.knn("[1,0,0,0]", 10, MAX_COSINE_DISTANCE).unwrap();
assert_eq!(rows.len(), 1);
assert!(rows[0].node_json.contains("delta"));
}
#[test]
fn test_empty_batch_and_zero_dim() {
let db = VecDb::open(tmp_db("empty")).unwrap();
db.insert_batch(&[]).unwrap(); assert!(db.insert_batch(&[("a".into(), "b".into(), vec![])]).is_err(), "零维向量应报错");
}
#[test]
fn test_empty_db_operations_are_noop() {
let db = VecDb::open(tmp_db("no_table")).unwrap();
assert_eq!(db.entry_count().unwrap(), 0);
assert!(db.knn("[1,0,0]", 10, MAX_COSINE_DISTANCE).unwrap().is_empty());
assert_eq!(db.remove_by_file("src/a.rs").unwrap(), 0);
db.clear().unwrap(); }
}