Skip to main content

code_repo_wiki/search/
semantic.rs

1//! 语义搜索引擎——sqlite-vec vec0 向量存储 + 余弦距离 KNN
2//!
3//! ## 职责边界(高内聚低耦合)
4//!
5//! - `SemanticEngine`:对外语义搜索门面。负责 embedding 生成
6//!   (EmbeddingEngine 调用)与 CodeNode 序列化(node_json),
7//!   向量存储细节全部委托 `VecDb`(src/search/vecdb.rs)。
8//! - `SemanticSearch` trait:语义引擎的抽象接口,供 SearchAgent
9//!   依赖抽象(可注入 mock 测试混合检索路径)。
10//!
11//! ## 阈值语义(v6 决策 4:保持硬编码 0.3)
12//!
13//! 相似度阈值 0.3 硬编码(OpenAI 官方 cosine 参考线),换算为余弦
14//! 距离 `MAX_COSINE_DISTANCE = 0.7`(vecdb 常量)下推到存储层过滤。
15
16use std::path::Path;
17use std::sync::Arc;
18
19use anyhow::{Context, Result};
20use tokio::runtime::Runtime;
21
22use crate::generate::embed::EmbeddingEngine;
23use crate::model::CodeNode;
24use crate::search::vecdb::VecDb;
25
26/// 语义搜索抽象接口(SearchAgent 依赖抽象,可注入 mock)
27///
28/// 方法集与 SemanticEngine 公开面一致。**不要求 Send/Sync**:rusqlite
29/// Connection 非 Sync(RefCell 内部),而 SearchAgent 是单线程调用
30/// (lib.rs execute_search 同步执行);若未来需要跨线程共享语义引擎,
31/// 由调用方用 Mutex 包装(trait 不应为此牺牲可测试性)。
32pub trait SemanticSearch {
33    /// 索引单个实体(生成 embedding 并持久化)
34    fn index(&mut self, node: &CodeNode, source_code: &str) -> Result<()>;
35    /// 批量索引多个实体
36    fn index_batch(&mut self, items: &[(CodeNode, String)]) -> Result<()>;
37    /// 搜索最相似的 k 个实体(0.3 相似度阈值过滤)
38    fn search(&self, query: &str, limit: usize) -> Result<Vec<(CodeNode, f32)>>;
39    /// 删除指定文件路径关联的所有向量条目
40    fn remove_by_file(&mut self, file_path: &str) -> Result<usize>;
41    /// 清空所有向量数据
42    fn clear(&mut self) -> Result<()>;
43    /// 当前向量条目数
44    fn entry_count(&self) -> usize;
45}
46
47/// 语义搜索引擎
48///
49/// 内部委托 VecDb(sqlite-vec vec0 虚表)完成向量持久化与 KNN,
50/// 自身只做 embedding 生成与 CodeNode 序列化。
51pub struct SemanticEngine {
52    db: VecDb,
53    embedder: Arc<EmbeddingEngine>,
54    rt: Arc<Runtime>,
55}
56
57impl SemanticEngine {
58    /// 打开或创建语义搜索数据库
59    ///
60    /// vec0 虚表延迟到首次插入时创建(维度首次探测)。
61    pub fn open(path: impl AsRef<Path>, embedder: Arc<EmbeddingEngine>, rt: Arc<Runtime>) -> Result<Self> {
62        let db = VecDb::open(path)?;
63        Ok(Self { db, embedder, rt })
64    }
65
66    // ============ 固有方法(薄封装,委托 trait 实现) ============
67    // 库内调用点(lib.rs build_search_index/update_search_index_incremental)
68    // 以具体类型调用,不走 trait object;这里直接转发到 trait impl,
69    // 避免调用点改用 Box<dyn> 语法,同时保证两处行为一致。
70
71    pub fn index(&mut self, node: &CodeNode, source_code: &str) -> Result<()> {
72        SemanticSearch::index(self, node, source_code)
73    }
74
75    pub fn index_batch(&mut self, items: &[(CodeNode, String)]) -> Result<()> {
76        SemanticSearch::index_batch(self, items)
77    }
78
79    pub fn search(&self, query: &str, limit: usize) -> Result<Vec<(CodeNode, f32)>> {
80        SemanticSearch::search(self, query, limit)
81    }
82
83    pub fn remove_by_file(&mut self, file_path: &str) -> Result<usize> {
84        SemanticSearch::remove_by_file(self, file_path)
85    }
86
87    pub fn clear(&mut self) -> Result<()> {
88        SemanticSearch::clear(self)
89    }
90
91    pub fn entry_count(&self) -> usize {
92        SemanticSearch::entry_count(self)
93    }
94
95    /// 当前 vec0 表的向量维度(表不存在返回 None)——U04/D2 维度探测用:
96    /// 增量路径在回填前比对 embedding 产出维度,变化则回退全量重建。
97    ///
98    /// 返回 Result:数据库读取错误(损坏/权限)向上传播,由调用方决定
99    /// 处理(lib.rs 维度探测失败 warn + 视为维度未知,不静默吞掉)。
100    pub fn table_dimension(&self) -> Result<Option<usize>> {
101        self.db.table_dimension()
102    }
103
104    /// 组装实体索引文本(与旧实现一致,保持索引兼容性)
105    fn index_text(node: &CodeNode, source_code: &str) -> String {
106        format!(
107            "{} {:?} {} {}",
108            node.name, node.kind,
109            node.signature.as_deref().unwrap_or(""), source_code
110        )
111    }
112}
113
114impl SemanticSearch for SemanticEngine {
115    fn index(&mut self, node: &CodeNode, source_code: &str) -> Result<()> {
116        let text = Self::index_text(node, source_code);
117        let vector = self
118            .rt
119            .block_on(self.embedder.embed(&text))
120            .context("生成 embedding 失败")?;
121        let node_json = serde_json::to_string(node).context("序列化 CodeNode 失败")?;
122        let file = node.file_path.as_deref().unwrap_or("").to_string();
123        self.db.insert_batch(&[(file, node_json, vector)])
124    }
125
126    fn index_batch(&mut self, items: &[(CodeNode, String)]) -> Result<()> {
127        if items.is_empty() {
128            return Ok(());
129        }
130        // 组装批量嵌入文本(一次 API 调用,避免逐条创建 tokio Runtime 开销)
131        let texts: Vec<String> = items
132            .iter()
133            .map(|(node, source)| Self::index_text(node, source))
134            .collect();
135        let vectors = self
136            .rt
137            .block_on(self.embedder.embed_batch(&texts))
138            .context("批量生成 embedding 失败")?;
139
140        // 组装 (file_path, node_json, vector) 三元组一次性入库
141        let rows: Vec<(String, String, Vec<f32>)> = items
142            .iter()
143            .zip(vectors)
144            .map(|((node, _), vector)| {
145                // CodeNode 是纯数据模型(无自定义 serde 错误路径),
146                // 序列化失败在类型层面不可达;unwrap_or_default 只为
147                // 满足 map 闭包签名,空串行由搜索侧反序列化失败自然跳过
148                let node_json = serde_json::to_string(node).unwrap_or_default();
149                let file = node.file_path.as_deref().unwrap_or("").to_string();
150                (file, node_json, vector)
151            })
152            .collect();
153        self.db.insert_batch(&rows)
154    }
155
156    fn search(&self, query: &str, limit: usize) -> Result<Vec<(CodeNode, f32)>> {
157        // 空库路径无命中时直接返回空,免对空库发 embedding 请求
158        //(embedding API 有成本与延迟,空库查询无意义);实际实现
159        // load_all_vectors 收集时语义与此一致。
160        // 计数失败(数据库损坏)向上传播,不静默当作空库跳过查询。
161        if self.db.entry_count()? == 0 {
162            return Ok(Vec::new());
163        }
164        let q_vec = self.rt.block_on(self.embedder.embed(query))?;
165        let query_json = vec_to_json(&q_vec);
166        // 阈值换算:相似度 0.3 ↔ 距离 0.7(vecdb 常量,见模块头)
167        let rows = self
168            .db
169            .knn(&query_json, limit, crate::search::vecdb::MAX_COSINE_DISTANCE)?;
170        let mut results = Vec::with_capacity(rows.len());
171        for row in rows {
172            // 反序列化失败 = 索引数据损坏(外部篡改/旧版本写入的异构格式),
173            // 单条跳过不中断整个搜索(坏行对结果质量影响有限,搜索是
174            // 只读尽力而为路径);索引重建由维度探测/全量重建机制覆盖
175            if let Ok(node) = serde_json::from_str::<CodeNode>(&row.node_json) {
176                // 距离 → 相似度(1 - distance),与旧实现返回语义一致
177                results.push((node, (1.0 - row.distance) as f32));
178            }
179        }
180        Ok(results)
181    }
182
183    fn remove_by_file(&mut self, file_path: &str) -> Result<usize> {
184        self.db.remove_by_file(file_path)
185    }
186
187    fn clear(&mut self) -> Result<()> {
188        self.db.clear()
189    }
190
191    fn entry_count(&self) -> usize {
192        // trait 签名不含 Result(计数是幂等只读操作,调用方无错误上下文);
193        // 数据库损坏时显式告警 + 按空库处理(计数 0),不静默吞错
194        match self.db.entry_count() {
195            Ok(n) => n,
196            Err(e) => {
197                tracing::warn!("语义索引条目计数失败,按空库处理: {}", e);
198                0
199            }
200        }
201    }
202}
203
204/// f32 向量 → vec0 查询向量 JSON(与 VecDb 内部序列化同格式)
205fn vec_to_json(v: &[f32]) -> String {
206    let parts: Vec<String> = v.iter().map(|f| format!("{f}")).collect();
207    format!("[{}]", parts.join(","))
208}
209
210#[cfg(test)]
211mod tests {
212    use super::*;
213    use crate::config::schema::EmbedSection;
214    use crate::model::{NodeId, NodeKind};
215    use std::collections::HashMap;
216    use std::io::{Read, Write};
217    use std::net::{TcpListener, TcpStream};
218    use std::sync::Mutex;
219    use tokio::runtime::Runtime;
220
221    fn test_runtime() -> Arc<Runtime> {
222        Arc::new(Runtime::new().unwrap())
223    }
224
225    fn mock_embedder() -> Arc<EmbeddingEngine> {
226        let config = EmbedSection {
227
228            model: "text-embedding-3-small".into(),
229            api_key: Some("test-key".into()),
230            api_key_env: "OPENAI_API_KEY".into(),
231            base_url: Some("http://localhost:9999/v1".into()),
232        };
233        Arc::new(EmbeddingEngine::new(&config, test_runtime().handle().clone()).unwrap())
234    }
235
236    /// 构造指向本地 mock 的 Embedding 引擎(base_url 带 /v1 前缀)
237    fn embedder_with_server(base_url: &str, rt: &Arc<Runtime>) -> Arc<EmbeddingEngine> {
238        let config = EmbedSection {
239
240            model: "text-embedding-3-small".into(),
241            api_key: Some("test-key".into()),
242            api_key_env: "OPENAI_API_KEY".into(),
243            base_url: Some(format!("{}/v1", base_url)),
244        };
245        Arc::new(EmbeddingEngine::new(&config, rt.handle().clone()).unwrap())
246    }
247
248    fn tmp_path(label: &str) -> std::path::PathBuf {
249        use std::sync::atomic::{AtomicU64, Ordering};
250        static SEM_COUNTER: AtomicU64 = AtomicU64::new(0);
251        let id = SEM_COUNTER.fetch_add(1, Ordering::Relaxed);
252        let mut p = std::env::temp_dir();
253        p.push(format!("semantic_fts_{}_{}.db", label, id));
254        let _ = std::fs::remove_file(&p);
255        p
256    }
257
258    fn make_node(name: &str, file: &str) -> CodeNode {
259        CodeNode {
260            id: NodeId::new(0),
261            kind: NodeKind::Function,
262            name: name.into(),
263            file_path: Some(file.into()),
264            line_range: Some((1, 10)),
265            doc_comment: None,
266            signature: Some(format!("fn {}()", name)), visibility: None,
267            module_path: vec![],
268        }
269    }
270
271    // ============ 伪 Embedding mock server ============
272    // 语义引擎的 Embedder 是具体类型 EmbeddingEngine(非 trait object),
273    // 无法注入 FakeEmbedder,因此用本地 mock HTTP 返回确定性伪向量。
274    // 维度统一为 3(vec0 虚表固定维度契约,v6 决策 1)。
275
276    /// 缓冲区中是否已出现完整请求头(含 \r\n\r\n 分隔符,其后可能还有请求体)
277    fn header_complete(buf: &[u8]) -> bool {
278        buf.windows(4).any(|w| w == b"\r\n\r\n")
279    }
280
281    /// 读取一个 HTTP 请求体(按 Content-Length 跨包缓冲读取)
282    fn read_request_body(stream: &mut TcpStream) -> String {
283        let mut buf = Vec::new();
284        let mut tmp = [0u8; 4096];
285        while !header_complete(&buf) {
286            match stream.read(&mut tmp) {
287                Ok(0) | Err(_) => break,
288                Ok(n) => buf.extend_from_slice(&tmp[..n]),
289            }
290        }
291        let head_end = buf.windows(4).position(|w| w == b"\r\n\r\n").unwrap_or(buf.len());
292        let head = String::from_utf8_lossy(&buf[..head_end]).to_string();
293        let content_length = head
294            .split("\r\n")
295            .filter_map(|l| l.split_once(':'))
296            .find(|(k, _)| k.trim().eq_ignore_ascii_case("content-length"))
297            .and_then(|(_, v)| v.trim().parse::<usize>().ok())
298            .unwrap_or(0);
299        const HEADER_SEP: usize = 4;
300        while buf.len() < head_end + HEADER_SEP + content_length {
301            match stream.read(&mut tmp) {
302                Ok(0) | Err(_) => break,
303                Ok(n) => buf.extend_from_slice(&tmp[..n]),
304            }
305        }
306        String::from_utf8_lossy(&buf[head_end + HEADER_SEP..head_end + HEADER_SEP + content_length]).to_string()
307    }
308
309    /// 关键词 → 确定性伪向量(统一 3 维):
310    /// - 已知关键词(alpha/beta/gamma):固定向量,相似度受控
311    ///   (alpha↔beta≈0.707、alpha↔gamma=-1.0),用于验证排序正确性;
312    /// - 未知关键词:按首次出现顺序分配 3 维单位基向量(同词同向量、
313    ///   异词正交),用于确定性验证 0.3 阈值过滤。
314    fn pseudo_vector(keyword: &str, seen: &mut HashMap<String, usize>) -> Vec<f32> {
315        match keyword {
316            "alpha" => vec![1.0, 0.0, 0.0],
317            "beta" => vec![std::f32::consts::FRAC_1_SQRT_2, std::f32::consts::FRAC_1_SQRT_2, 0.0],
318            "gamma" => vec![-1.0, 0.0, 0.0],
319            _ => {
320                let next = seen.len();
321                let idx = *seen.entry(keyword.to_string()).or_insert(next);
322                let mut v = vec![0.0f32; 3];
323                v[idx % 3] = 1.0;
324                v
325            }
326        }
327    }
328
329    /// 启动本地伪 Embedding mock server(std 线程 + std::net,无 tokio net 依赖):
330    /// 解析请求体中的 input 数组,按首个单词分配伪向量,返回同序 embedding 列表。
331    fn spawn_pseudo_embed_server() -> String {
332        let listener = TcpListener::bind("127.0.0.1:0").unwrap();
333        let base_url = format!("http://{}", listener.local_addr().unwrap());
334        let seen: Arc<Mutex<HashMap<String, usize>>> = Arc::new(Mutex::new(HashMap::new()));
335
336        std::thread::spawn(move || {
337            for stream in listener.incoming() {
338                let Ok(mut stream) = stream else { break };
339                let seen = seen.clone();
340                std::thread::spawn(move || {
341                    let body = read_request_body(&mut stream);
342                    let inputs: Vec<String> = serde_json::from_str::<serde_json::Value>(&body)
343                        .ok()
344                        .and_then(|v| v["input"].as_array().map(|a| {
345                            a.iter().filter_map(|x| x.as_str().map(String::from)).collect()
346                        }))
347                        .unwrap_or_default();
348                    let mut guard = seen.lock().unwrap();
349                    let vectors: Vec<Vec<f32>> = inputs.iter()
350                        .map(|t| pseudo_vector(t.split_whitespace().next().unwrap_or(""), &mut guard))
351                        .collect();
352                    drop(guard);
353
354                    let payload = serde_json::json!({
355                        "data": vectors.iter().map(|v| serde_json::json!({"embedding": v})).collect::<Vec<_>>()
356                    }).to_string();
357                    let raw = format!(
358                        "HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{}",
359                        payload.len(), payload
360                    );
361                    let _ = stream.write_all(raw.as_bytes());
362                });
363            }
364        });
365
366        base_url
367    }
368
369    #[test]
370    fn test_semantic_new() {
371        let engine = SemanticEngine::open(tmp_path("new"), mock_embedder(), test_runtime()).unwrap();
372        assert_eq!(engine.entry_count(), 0);
373    }
374
375    #[test]
376    fn test_search_empty() {
377        let engine = SemanticEngine::open(tmp_path("empty"), mock_embedder(), test_runtime()).unwrap();
378        assert!(engine.search("test", 10).unwrap().is_empty());
379    }
380
381    #[test]
382    fn test_semantic_search_ranks_by_similarity() {
383        // 索引 3 个实体,查询 alpha:相似度 alpha=1.0 > beta≈0.707 > gamma=-1.0,
384        // gamma 被 0.3 阈值过滤,剩余结果按相似度降序
385        let base_url = spawn_pseudo_embed_server();
386        let rt = test_runtime();
387        let embedder = embedder_with_server(&base_url, &rt);
388        let mut engine = SemanticEngine::open(tmp_path("rank"), embedder, rt.clone()).unwrap();
389
390        let items = vec![
391            (make_node("alpha", "src/a.rs"), "fn alpha()".to_string()),
392            (make_node("beta", "src/b.rs"), "fn beta()".to_string()),
393            (make_node("gamma", "src/c.rs"), "fn gamma()".to_string()),
394        ];
395        engine.index_batch(&items).unwrap();
396
397        let results = engine.search("alpha", 10).unwrap();
398        // 最相似实体排第一,次相似排第二
399        assert_eq!(results.len(), 2);
400        assert_eq!(results[0].0.name, "alpha");
401        assert_eq!(results[1].0.name, "beta");
402        assert!(results[0].1 > results[1].1);
403    }
404
405    #[test]
406    fn test_semantic_search_filters_below_threshold() {
407        // 索引与查询完全无关的实体:查询 "q" 与已索引关键词正交(相似度 0.0),
408        // 0.3 阈值过滤后结果为空
409        let base_url = spawn_pseudo_embed_server();
410        let rt = test_runtime();
411        let embedder = embedder_with_server(&base_url, &rt);
412        let mut engine = SemanticEngine::open(tmp_path("thr"), embedder, rt.clone()).unwrap();
413
414        let items = vec![
415            (make_node("x1", "src/x1.rs"), "x1 unrelated".to_string()),
416            (make_node("x2", "src/x2.rs"), "x2 unrelated".to_string()),
417        ];
418        engine.index_batch(&items).unwrap();
419
420        let results = engine.search("q", 10).unwrap();
421        assert!(results.is_empty());
422    }
423
424    #[test]
425    fn test_semantic_remove_by_file() {
426        let base_url = spawn_pseudo_embed_server();
427        let rt = test_runtime();
428        let embedder = embedder_with_server(&base_url, &rt);
429        let mut engine = SemanticEngine::open(tmp_path("rm"), embedder, rt.clone()).unwrap();
430
431        let items = vec![
432            (make_node("a1", "src/a.rs"), "fn a1()".to_string()),
433            (make_node("b1", "src/b.rs"), "fn b1()".to_string()),
434        ];
435        engine.index_batch(&items).unwrap();
436        assert_eq!(engine.entry_count(), 2);
437
438        let removed = engine.remove_by_file("src/a.rs").unwrap();
439        assert_eq!(removed, 1);
440        assert_eq!(engine.entry_count(), 1);
441    }
442
443    #[test]
444    fn test_semantic_clear() {
445        let base_url = spawn_pseudo_embed_server();
446        let rt = test_runtime();
447        let embedder = embedder_with_server(&base_url, &rt);
448        let mut engine = SemanticEngine::open(tmp_path("clr"), embedder, rt.clone()).unwrap();
449        engine.index_batch(&[(make_node("a", "src/a.rs"), "fn a()".to_string())]).unwrap();
450        assert_eq!(engine.entry_count(), 1);
451
452        engine.clear().unwrap();
453        assert_eq!(engine.entry_count(), 0);
454    }
455
456    /// U04/D2:维度探测——空库 None,入库后返回实际维度
457    #[test]
458    fn test_semantic_table_dimension() {
459        let base_url = spawn_pseudo_embed_server();
460        let rt = test_runtime();
461        let embedder = embedder_with_server(&base_url, &rt);
462        let mut engine = SemanticEngine::open(tmp_path("dim"), embedder, rt.clone()).unwrap();
463        assert_eq!(engine.table_dimension().unwrap(), None, "空库(表未创建)应返回 None");
464
465        engine
466            .index_batch(&[(make_node("a1", "src/a.rs"), "fn a1()".to_string())])
467            .unwrap();
468        assert_eq!(engine.table_dimension().unwrap(), Some(3), "伪向量统一 3 维");
469    }
470}