1use std::path::Path;
39use std::sync::OnceLock;
40
41use anyhow::{Context, Result};
42use rusqlite::Connection;
43
44pub const MAX_COSINE_DISTANCE: f64 = 0.7;
49
50const MAX_KNN_CANDIDATES: usize = 10_000;
56
57static VEC_EXT_REGISTERED: OnceLock<()> = OnceLock::new();
63
64fn ensure_extension_registered() {
68 VEC_EXT_REGISTERED.get_or_init(|| {
69 let init_fn: unsafe extern "C" fn() = sqlite_vec::sqlite3_vec_init;
77 let entry: rusqlite::auto_extension::RawAutoExtension =
78 unsafe { std::mem::transmute(init_fn as *const ()) };
79 unsafe {
82 let _ = rusqlite::auto_extension::register_auto_extension(entry);
83 }
84 });
85}
86
87const VECTOR_TABLE: &str = "vectors";
89
90pub struct VecDb {
97 conn: Connection,
98}
99
100#[derive(Clone)]
102pub struct KnnRow {
103 pub node_json: String,
105 pub distance: f64,
107}
108
109impl VecDb {
110 pub fn open(path: impl AsRef<Path>) -> Result<Self> {
112 ensure_extension_registered();
113 let conn = Connection::open(path.as_ref())
114 .context("打开向量数据库失败")?;
115 conn.pragma_update(None, "journal_mode", "WAL")
117 .context("设置 WAL 模式失败")?;
118 conn.busy_timeout(std::time::Duration::from_secs(5))
119 .context("设置 busy_timeout 失败")?;
120 Ok(Self { conn })
121 }
122
123 fn table_exists(&self) -> Result<bool> {
125 let sql = "SELECT COUNT(*) FROM sqlite_master WHERE type='table' AND name=?1";
126 let count: i64 = self
127 .conn
128 .query_row(sql, [VECTOR_TABLE], |r| r.get(0))
129 .context("查询虚表存在性失败")?;
130 Ok(count > 0)
131 }
132
133 pub fn table_dimension(&self) -> Result<Option<usize>> {
135 if !self.table_exists()? {
136 return Ok(None);
137 }
138 let sql = "SELECT sql FROM sqlite_master WHERE type='table' AND name=?1";
139 let ddl: String = self
140 .conn
141 .query_row(sql, [VECTOR_TABLE], |r| r.get(0))
142 .context("读取虚表定义失败")?;
143 let start = ddl.find("float[").map(|i| i + 6);
146 let Some(start) = start else {
147 anyhow::bail!("vec0 虚表定义缺少 float[N] 声明: {ddl}");
149 };
150 let end = ddl[start..].find(']').map(|i| start + i);
151 let Some(end) = end else {
152 anyhow::bail!("vec0 虚表定义缺少 ] 结束符: {ddl}");
153 };
154 ddl[start..end]
155 .parse::<usize>()
156 .map(Some)
157 .with_context(|| format!("解析 vec0 维度失败: {}", &ddl[start..end]))
158 }
159
160 fn create_table(&self, dim: usize) -> Result<()> {
162 let sql = format!(
163 "CREATE VIRTUAL TABLE {VECTOR_TABLE} USING vec0(\
164 embedding float[{dim}] distance_metric=cosine,\
165 file_path TEXT,\
166 node_json TEXT\
167 )"
168 );
169 self.conn
170 .execute_batch(&sql)
171 .with_context(|| format!("创建 vec0 虚表失败(dim={dim})"))
172 }
173
174 fn rebuild_table(&self, dim: usize) -> Result<()> {
176 self.conn
177 .execute_batch(&format!("DROP TABLE IF EXISTS {VECTOR_TABLE};"))
178 .context("删除旧 vec0 虚表失败")?;
179 self.create_table(dim)
180 }
181
182 fn ensure_table(&self, dim: usize) -> Result<()> {
187 match self.table_dimension()? {
188 None => self.create_table(dim),
189 Some(existing) if existing != dim => {
190 tracing::warn!(
191 "向量维度变化({} → {}),重建语义索引(embedding 模型变更需全量重新索引)",
192 existing, dim
193 );
194 self.rebuild_table(dim)
195 }
196 Some(_) => Ok(()),
197 }
198 }
199
200 pub fn insert_batch(&self, items: &[(String, String, Vec<f32>)]) -> Result<()> {
205 if items.is_empty() {
206 return Ok(());
207 }
208 let dim = items[0].2.len();
209 if dim == 0 {
210 anyhow::bail!("空向量无法入库(维度为 0)");
211 }
212 self.ensure_table(dim)?;
213 let mut stmt = self.conn.prepare(&format!(
214 "INSERT INTO {VECTOR_TABLE}(rowid, embedding, file_path, node_json) \
215 VALUES (?1, ?2, ?3, ?4)"
216 ))?;
217 for (i, (file_path, node_json, vector)) in items.iter().enumerate() {
218 let json = vector_to_json(vector);
220 stmt.execute(rusqlite::params![
221 (i + 1) as i64,
222 json,
223 crate::incremental::norm_sep(file_path),
224 node_json,
225 ])
226 .with_context(|| format!("插入向量失败: {file_path}"))?;
227 }
228 Ok(())
229 }
230
231 pub fn knn(&self, query_json: &str, limit: usize, max_distance: f64) -> Result<Vec<KnnRow>> {
249 if !self.table_exists()? {
250 return Ok(Vec::new());
251 }
252 let mut sample = limit.max(1);
253 let mut all: Vec<KnnRow> = Vec::new();
254 let mut seen: std::collections::HashSet<String> = std::collections::HashSet::new();
255 loop {
256 let sql = format!(
257 "SELECT node_json, distance FROM {VECTOR_TABLE} \
258 WHERE embedding MATCH ?1 \
259 ORDER BY distance \
260 LIMIT {sample}"
261 );
262 let mut stmt = self.conn.prepare(&sql).context("准备 KNN 查询失败")?;
263 let rows: Vec<KnnRow> = stmt
264 .query_map([query_json], |r| {
265 Ok(KnnRow {
266 node_json: r.get(0)?,
267 distance: r.get(1)?,
268 })
269 })
270 .context("执行 KNN 查询失败")?
271 .collect::<rusqlite::Result<Vec<_>>>()
272 .context("读取 KNN 结果失败")?;
273
274 for row in rows.iter().filter(|r| r.distance <= max_distance) {
276 if seen.insert(row.node_json.clone()) {
277 all.push(row.clone());
278 }
279 }
280 if rows.len() < sample || rows.last().map(|r| r.distance > max_distance).unwrap_or(true) || sample >= MAX_KNN_CANDIDATES {
286 break;
287 }
288 sample = (sample * 2).min(MAX_KNN_CANDIDATES);
289 }
290 Ok(all)
291 }
292
293 pub fn remove_by_file(&self, file_path: &str) -> Result<usize> {
295 if !self.table_exists()? {
296 return Ok(0);
297 }
298 let sql = format!("DELETE FROM {VECTOR_TABLE} WHERE file_path = ?1");
299 let count = self
300 .conn
301 .execute(&sql, rusqlite::params![crate::incremental::norm_sep(file_path)])
302 .context("删除向量失败")?;
303 Ok(count)
304 }
305
306 pub fn clear(&self) -> Result<()> {
308 if !self.table_exists()? {
309 return Ok(());
310 }
311 self.conn
312 .execute_batch(&format!("DELETE FROM {VECTOR_TABLE};"))
313 .context("清空向量失败")
314 }
315
316 pub fn entry_count(&self) -> Result<usize> {
318 if !self.table_exists()? {
319 return Ok(0);
320 }
321 let sql = format!("SELECT COUNT(*) FROM {VECTOR_TABLE}");
322 let count: i64 = self
323 .conn
324 .query_row(&sql, [], |r| r.get(0))
325 .context("查询向量数失败")?;
326 Ok(count as usize)
327 }
328}
329
330fn vector_to_json(v: &[f32]) -> String {
335 let parts: Vec<String> = v.iter().map(|f| format!("{f}")).collect();
336 format!("[{}]", parts.join(","))
337}
338
339#[cfg(test)]
340mod tests {
341 use super::*;
342
343 fn tmp_db(tag: &str) -> std::path::PathBuf {
345 use std::sync::atomic::{AtomicU64, Ordering};
346 static SEQ: AtomicU64 = AtomicU64::new(0);
347 let id = SEQ.fetch_add(1, Ordering::Relaxed);
348 let mut p = std::env::temp_dir();
349 p.push(format!("vecdb_{}_{}_{}.db", tag, std::process::id(), id));
350 let _ = std::fs::remove_file(&p);
351 p
352 }
353
354 fn make_items() -> Vec<(String, String, Vec<f32>)> {
355 vec![
356 ("src/a.rs".into(), r#"{"name":"alpha"}"#.into(), vec![1.0, 0.0, 0.0]),
357 ("src/b.rs".into(), r#"{"name":"beta"}"#.into(), vec![0.0, 1.0, 0.0]),
358 ("src/c.rs".into(), r#"{"name":"gamma"}"#.into(), vec![-1.0, 0.0, 0.0]),
359 ]
360 }
361
362 #[test]
363 fn test_insert_and_knn_cosine_ranking() {
364 let db = VecDb::open(tmp_db("knn")).unwrap();
365 db.insert_batch(&make_items()).unwrap();
366
367 let rows = db.knn("[0.9,0.1,0]", 10, MAX_COSINE_DISTANCE).unwrap();
369 assert_eq!(rows.len(), 1, "只有 alpha 在 0.3 相似度阈值内");
372 assert!(rows[0].node_json.contains("alpha"));
373 assert!(rows[0].distance < 0.01);
374 }
375
376 #[test]
377 fn test_threshold_0_3_maps_to_distance_0_7() {
378 assert_eq!(MAX_COSINE_DISTANCE, 0.7);
380 }
381
382 #[test]
383 fn test_knn_without_threshold_returns_all() {
384 let db = VecDb::open(tmp_db("knn_all")).unwrap();
385 db.insert_batch(&make_items()).unwrap();
386
387 let rows = db.knn("[0.9,0.1,0]", 10, 2.0).unwrap();
389 assert_eq!(rows.len(), 3);
390 assert!(rows[0].node_json.contains("alpha"));
391 assert!(rows[2].node_json.contains("gamma"));
392 assert!(rows[0].distance < rows[1].distance);
393 }
394
395 #[test]
396 fn test_knn_expands_sample_to_cover_threshold() {
397 let db = VecDb::open(tmp_db("expand")).unwrap();
400 let items: Vec<(String, String, Vec<f32>)> = (0..30)
401 .map(|i| (format!("src/f{i}.rs"), format!(r#"{{"name":"f{i}"}}"#), vec![1.0, 0.0, 0.0]))
402 .collect();
403 db.insert_batch(&items).unwrap();
404
405 let rows = db.knn("[1,0,0]", 5, MAX_COSINE_DISTANCE).unwrap();
406 assert_eq!(rows.len(), 30, "阈值内全部候选必须返回(扩样不能截断)");
407 }
408
409 #[test]
410 fn test_delete_by_file_and_count() {
411 let db = VecDb::open(tmp_db("del")).unwrap();
412 db.insert_batch(&make_items()).unwrap();
413 assert_eq!(db.entry_count().unwrap(), 3);
414
415 let removed = db.remove_by_file("src/b.rs").unwrap();
416 assert_eq!(removed, 1);
417 assert_eq!(db.entry_count().unwrap(), 2);
418
419 db.clear().unwrap();
420 assert_eq!(db.entry_count().unwrap(), 0);
421 }
422
423 #[test]
424 fn test_dimension_mismatch_rebuilds_table() {
425 let db = VecDb::open(tmp_db("dim")).unwrap();
426 db.insert_batch(&make_items()).unwrap(); db.insert_batch(&[(
430 "src/d.rs".into(),
431 r#"{"name":"delta"}"#.into(),
432 vec![1.0, 0.0, 0.0, 0.0],
433 )])
434 .unwrap();
435 assert_eq!(db.entry_count().unwrap(), 1, "维度变化重建后只剩新数据");
436
437 let rows = db.knn("[1,0,0,0]", 10, MAX_COSINE_DISTANCE).unwrap();
438 assert_eq!(rows.len(), 1);
439 assert!(rows[0].node_json.contains("delta"));
440 }
441
442 #[test]
443 fn test_empty_batch_and_zero_dim() {
444 let db = VecDb::open(tmp_db("empty")).unwrap();
445 db.insert_batch(&[]).unwrap(); assert!(db.insert_batch(&[("a".into(), "b".into(), vec![])]).is_err(), "零维向量应报错");
447 }
448
449 #[test]
450 fn test_empty_db_operations_are_noop() {
451 let db = VecDb::open(tmp_db("no_table")).unwrap();
452 assert_eq!(db.entry_count().unwrap(), 0);
453 assert!(db.knn("[1,0,0]", 10, MAX_COSINE_DISTANCE).unwrap().is_empty());
454 assert_eq!(db.remove_by_file("src/a.rs").unwrap(), 0);
455 db.clear().unwrap(); }
457}