Skip to main content

code_repo_wiki/search/
vecdb.rs

1//! 语义向量存储层:基于 sqlite-vec 0.1.9 的 vec0 虚表封装
2//!
3//! ## 技术选型(v6 决策 1 修正)
4//!
5//! 原定 sqlite-vector-rs(HNSW via usearch),实测其依赖链在 Windows/MSVC
6//! 三重阻断不可编译:
7//! 1. sqlite3_ext 0.2.1 的 `cfg!(unix)` 误用(运行时宏门控编译期 use
8//!    `std::os::unix`)→ E0433(纯 Rust 层问题,gcc 工具链同样存在,
9//!    windows-gnu target 也没有 std::os::unix);
10//! 2. numkong(usearch 的 C 依赖)的 C99 混合声明 + cc 传 `-std:c99`
11//!    → MSVC C 编译器不支持声明后语句 → C2059;
12//! 3. 切 windows-gnu 工具链可解 C 层,但 E0433 无解且全项目 C 依赖
13//!    (rusqlite bundled/git2/leiden)需重编译,风险不可接受。
14//!
15//! 改用 **sqlite-vec 0.1.9**(纯 C 静态嵌入,cc 编译,无 usearch/numkong/
16//! sqlite3_ext 依赖链):官方 Rust 用法 = `sqlite3_vec_init` + rusqlite
17//! `sqlite3_auto_extension` 注册,Windows 实测编译通过 + 全链路探针通过。
18//!
19//! ## 语义对齐(阈值换算依据)
20//!
21//! vec0 的 `distance_metric=cosine` 返回 `1 - cosine_similarity`
22//! (sqlite-vec.c:479,与 usearch Cos 同语义)。原实现的相似度阈值 0.3
23//! (semantic.rs 硬编码)等价换算为 **distance ≤ 0.7**(1 - 0.3)。
24//! 排序方向也一致:distance 升序 = 相似度降序。
25//!
26//! ## 维度契约
27//!
28//! vec0 建表时维度固定(`float[N]`)。首次插入时探测向量维度建表;
29//! 换 embedding 模型(EmbedSection.model)导致维度变化时自动重建表
30//! (丢弃旧索引并告警——与旧实现"删旧索引全量重建"行为一致)。
31//!
32//! ## 表结构
33//!
34//! 单表 vec0:`embedding float[N] distance_metric=cosine` + metadata 列
35//! `file_path TEXT`(删除/过滤键,与 FTS5 entities 表同归一化规则)、
36//! `node_json TEXT`(CodeNode 序列化,KNN 结果直接反序列化,无 join)。
37
38use std::path::Path;
39use std::sync::OnceLock;
40
41use anyhow::{Context, Result};
42use rusqlite::Connection;
43
44/// 相似度阈值 0.3 换算后的余弦距离上限(1 - 0.3 = 0.7)
45///
46/// 换算依据见模块头:vec0 cosine distance = 1 - cosine_similarity。
47/// 保持与旧实现(cosine_similarity > 0.3)逐位一致的行为。
48pub const MAX_COSINE_DISTANCE: f64 = 0.7;
49
50/// KNN 超采样扩大的上限:单次查询最多取这么多候选行
51///
52/// 循环扩样策略:从 `limit` 起步,若返回行数等于采样数且最后一行
53/// 仍 ≤ 阈值(可能有更多候选被截断)则翻倍重查,直到表尽或达到
54/// 本上限。上限是防御真实路径(万级实体全相似)退化为全表扫描。
55const MAX_KNN_CANDIDATES: usize = 10_000;
56
57/// sqlite-vec 扩展的进程级注册(OnceLock 保证只注册一次)
58///
59/// `sqlite3_auto_extension` 是全局注册:注册后所有新建的 SQLite 连接
60/// 自动加载扩展。重复注册会重复执行 init(无害但冗余),且并发注册
61/// 存在竞态风险——用 OnceLock 收敛为一次。
62static VEC_EXT_REGISTERED: OnceLock<()> = OnceLock::new();
63
64/// 注册 sqlite-vec 扩展(进程级,幂等)
65///
66/// 必须在任何 vec0 虚表操作之前调用;所有连接共享该注册。
67fn ensure_extension_registered() {
68    VEC_EXT_REGISTERED.get_or_init(|| {
69        // 类型注解:`transmute` 需显式标注目标类型(clippy::missing_transmute_annotations)。
70        // RawAutoExtension = unsafe extern "C" fn(*mut sqlite3, *mut *mut c_char,
71        // *const sqlite3_api_routines) -> c_int(SQLite 扩展入口约定);
72        // sqlite3_vec_init 是 extern "C" fn()——ABI 兼容(C 扩展入口),
73        // 官方示例即用 transmute(sqlite-vec crate 自带测试同款写法)。
74        // register_auto_extension 是 rusqlite 对 sqlite3_auto_extension 的安全封装,
75        // 只做条目注册不打开连接,返回值校验由封装完成。
76        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        // register_auto_extension 内部调用 ffi::sqlite3_auto_extension(unsafe),
80        // 封装自身也标注 unsafe:扩展注册是进程级全局状态变更
81        unsafe {
82            let _ = rusqlite::auto_extension::register_auto_extension(entry);
83        }
84    });
85}
86
87/// 虚表名(与 FTS5 entities 表同库共存的向量存储)
88const VECTOR_TABLE: &str = "vectors";
89
90/// vec0 虚表封装:向量持久化 + 带阈值过滤的 KNN 查询
91///
92/// 职责边界:只做 vec0 虚表的建表/增删查,不含 embedding 生成
93/// (EmbeddingEngine 在 SemanticEngine 层调用)与 CodeNode 业务逻辑
94/// (node_json 的序列化由调用方负责)。单库单连接(与 SearchStore 同
95/// 进程契约:WAL + busy_timeout)。
96pub struct VecDb {
97    conn: Connection,
98}
99
100/// KNN 查询结果行:node_json + 余弦距离
101#[derive(Clone)]
102pub struct KnnRow {
103    /// CodeNode 的 JSON 序列化(调用方反序列化)
104    pub node_json: String,
105    /// 余弦距离(1 - similarity,升序 = 相似度降序)
106    pub distance: f64,
107}
108
109impl VecDb {
110    /// 打开或创建向量数据库(注册扩展 + 建表延迟到首次插入)
111    pub fn open(path: impl AsRef<Path>) -> Result<Self> {
112        ensure_extension_registered();
113        let conn = Connection::open(path.as_ref())
114            .context("打开向量数据库失败")?;
115        // WAL 模式:与 FTS5 存储同并发契约(多读单写)
116        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    /// 虚表是否存在
124    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    /// 当前虚表的向量维度(表不存在返回 None)
134    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        // vec0 建表 SQL 形如 "CREATE VIRTUAL TABLE vectors USING vec0(embedding float[3] ..., ...)"
144        // 提取 float[N] 中的 N
145        let start = ddl.find("float[").map(|i| i + 6);
146        let Some(start) = start else {
147            // 定义异常(理论不可达:表由本模块创建)→ 显式报错而非猜测维度
148            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    /// 以指定维度创建 vec0 虚表
161    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    /// 丢弃并重建虚表(维度变化时调用,旧索引一并丢弃)
175    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    /// 确保虚表存在且维度匹配;不匹配时重建
183    ///
184    /// `dim` 为本次插入向量的维度(EmbeddingEngine 首次产出后才知道)。
185    /// 表不存在 → 按 dim 建表;表存在但维度不同(换模型)→ 重建并告警。
186    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    /// 批量插入向量(node_json 由调用方序列化)
201    ///
202    /// 首次插入时按首个向量维度建表。file_path 归一化与 FTS5 表同规则
203    /// (跨平台路径键统一,删除/过滤同基准)。
204    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            // vec0 的 embedding 列接受 JSON 数组字符串(探针验证:'[1,0,0]' 直接可用)
219            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    /// 带阈值过滤的 KNN 查询:返回余弦距离 ≤ `max_distance` 的全部行
232    ///
233    /// vec0 的 KNN 是 LIMIT 截断,无法在 SQL 侧表达距离谓词
234    /// (knn_match 是唯一 distance 约束,见 sqlite-vec best_index 实现),
235    /// 因此用**循环扩样**保证阈值语义正确:
236    /// - 从 `limit` 起步采样,若返回行数 == 采样数且最后一行仍 ≤ 阈值,
237    ///   说明可能有更多候选被截断 → 翻倍重查;
238    /// - 直到返回行数 < 采样数(表尽)或最后一行 > 阈值(边界内已全)
239    ///   或达到 MAX_KNN_CANDIDATES 上限。
240    ///
241    /// 该策略与旧实现(全量加载 + 逐条余弦 + 过滤)在阈值语义上
242    /// 完全等价,且正常场景(阈值过滤后结果远小于 limit)只查一次。
243    /// 表不存在时返回空(等价于空索引)。
244    ///
245    /// 扩样重查的去重:翻倍 LIMIT 后,前一次已取过的行会再次返回
246    /// (KNN 结果按距离稳定排序),按 node_json 去重合并,保证
247    /// 返回行数 = 阈值内候选数(不因扩样重复)。
248    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            // 本次采样内全部 ≤ 阈值 → 并入结果(按 node_json 去重,扩样重查防重复)
275            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            // 提前终止条件:
281            // 1. 表尽(返回行数 < 采样数)——已取完所有候选
282            // 2. 采样内最后一行 > 阈值——阈值边界之后的(更远)行必然也 > 阈值
283            //    (distance 升序),无更多可并入结果
284            // 3. 达到扩样上限——防御全表相似退化
285            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    /// 删除指定文件路径关联的所有向量(file_path 键与插入同基准)
294    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    /// 清空所有向量
307    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    /// 当前向量条数(表不存在视为 0)
317    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
330/// f32 向量 → vec0 可用的 JSON 数组字符串(`[0.1,0.2,...]`)
331///
332/// vec0 的 embedding 列接受 JSON 数组文本(探针验证);不使用二进制
333/// blob 传输,避免引入额外序列化依赖。
334fn 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    /// 临时数据库路径(进程内自增序号防并行冲突)
344    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        // 查询 alpha 附近:alpha 距离最小(≈1-cos),gamma 最远
368        let rows = db.knn("[0.9,0.1,0]", 10, MAX_COSINE_DISTANCE).unwrap();
369        // 阈值 0.7:alpha(0.006) 与 beta(0.89>0.7? 否——beta 距离 0.89 > 0.7 被过滤)
370        // 实际:alpha d≈0.006 ≤0.7 ✓;beta d≈0.89 >0.7 ✗;gamma d≈1.99 >0.7 ✗
371        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        // 阈值换算锚定:0.3 相似度 ↔ 0.7 距离(常量即契约)
379        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        // max_distance=2.0(放行全部):3 行按距离升序返回
388        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        // 扩样正确性:插入 30 个与查询高度相似(distance ≤ 0.7)的向量,
398        // limit=5。若不做扩样只取 5 条会漏掉;循环扩样必须返回全部 30 条。
399        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(); // 3 维建表
427
428        // 换"模型"(4 维)→ 自动重建,旧数据丢弃
429        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(); // 空批次静默成功
446        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(); // 表不存在时静默成功
456    }
457}