1use 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
26pub trait SemanticSearch {
33 fn index(&mut self, node: &CodeNode, source_code: &str) -> Result<()>;
35 fn index_batch(&mut self, items: &[(CodeNode, String)]) -> Result<()>;
37 fn search(&self, query: &str, limit: usize) -> Result<Vec<(CodeNode, f32)>>;
39 fn remove_by_file(&mut self, file_path: &str) -> Result<usize>;
41 fn clear(&mut self) -> Result<()>;
43 fn entry_count(&self) -> usize;
45}
46
47pub struct SemanticEngine {
52 db: VecDb,
53 embedder: Arc<EmbeddingEngine>,
54 rt: Arc<Runtime>,
55}
56
57impl SemanticEngine {
58 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 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 pub fn table_dimension(&self) -> Result<Option<usize>> {
101 self.db.table_dimension()
102 }
103
104 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 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 let rows: Vec<(String, String, Vec<f32>)> = items
142 .iter()
143 .zip(vectors)
144 .map(|((node, _), vector)| {
145 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 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 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 if let Ok(node) = serde_json::from_str::<CodeNode>(&row.node_json) {
176 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 match self.db.entry_count() {
195 Ok(n) => n,
196 Err(e) => {
197 tracing::warn!("语义索引条目计数失败,按空库处理: {}", e);
198 0
199 }
200 }
201 }
202}
203
204fn 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 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 fn header_complete(buf: &[u8]) -> bool {
278 buf.windows(4).any(|w| w == b"\r\n\r\n")
279 }
280
281 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 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 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 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 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 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 #[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}