use std::path::Path;
use std::sync::Arc;
use anyhow::{Context, Result};
use tokio::runtime::Runtime;
use crate::generate::embed::EmbeddingEngine;
use crate::model::CodeNode;
use crate::search::vecdb::VecDb;
pub trait SemanticSearch {
fn index(&mut self, node: &CodeNode, source_code: &str) -> Result<()>;
fn index_batch(&mut self, items: &[(CodeNode, String)]) -> Result<()>;
fn search(&self, query: &str, limit: usize) -> Result<Vec<(CodeNode, f32)>>;
fn remove_by_file(&mut self, file_path: &str) -> Result<usize>;
fn clear(&mut self) -> Result<()>;
fn entry_count(&self) -> usize;
}
pub struct SemanticEngine {
db: VecDb,
embedder: Arc<EmbeddingEngine>,
rt: Arc<Runtime>,
}
impl SemanticEngine {
pub fn open(path: impl AsRef<Path>, embedder: Arc<EmbeddingEngine>, rt: Arc<Runtime>) -> Result<Self> {
let db = VecDb::open(path)?;
Ok(Self { db, embedder, rt })
}
pub fn index(&mut self, node: &CodeNode, source_code: &str) -> Result<()> {
SemanticSearch::index(self, node, source_code)
}
pub fn index_batch(&mut self, items: &[(CodeNode, String)]) -> Result<()> {
SemanticSearch::index_batch(self, items)
}
pub fn search(&self, query: &str, limit: usize) -> Result<Vec<(CodeNode, f32)>> {
SemanticSearch::search(self, query, limit)
}
pub fn remove_by_file(&mut self, file_path: &str) -> Result<usize> {
SemanticSearch::remove_by_file(self, file_path)
}
pub fn clear(&mut self) -> Result<()> {
SemanticSearch::clear(self)
}
pub fn entry_count(&self) -> usize {
SemanticSearch::entry_count(self)
}
pub fn table_dimension(&self) -> Result<Option<usize>> {
self.db.table_dimension()
}
fn index_text(node: &CodeNode, source_code: &str) -> String {
format!(
"{} {:?} {} {}",
node.name, node.kind,
node.signature.as_deref().unwrap_or(""), source_code
)
}
}
impl SemanticSearch for SemanticEngine {
fn index(&mut self, node: &CodeNode, source_code: &str) -> Result<()> {
let text = Self::index_text(node, source_code);
let vector = self
.rt
.block_on(self.embedder.embed(&text))
.context("生成 embedding 失败")?;
let node_json = serde_json::to_string(node).context("序列化 CodeNode 失败")?;
let file = node.file_path.as_deref().unwrap_or("").to_string();
self.db.insert_batch(&[(file, node_json, vector)])
}
fn index_batch(&mut self, items: &[(CodeNode, String)]) -> Result<()> {
if items.is_empty() {
return Ok(());
}
let texts: Vec<String> = items
.iter()
.map(|(node, source)| Self::index_text(node, source))
.collect();
let vectors = self
.rt
.block_on(self.embedder.embed_batch(&texts))
.context("批量生成 embedding 失败")?;
let rows: Vec<(String, String, Vec<f32>)> = items
.iter()
.zip(vectors)
.map(|((node, _), vector)| {
let node_json = serde_json::to_string(node).unwrap_or_default();
let file = node.file_path.as_deref().unwrap_or("").to_string();
(file, node_json, vector)
})
.collect();
self.db.insert_batch(&rows)
}
fn search(&self, query: &str, limit: usize) -> Result<Vec<(CodeNode, f32)>> {
if self.db.entry_count()? == 0 {
return Ok(Vec::new());
}
let q_vec = self.rt.block_on(self.embedder.embed(query))?;
let query_json = vec_to_json(&q_vec);
let rows = self
.db
.knn(&query_json, limit, crate::search::vecdb::MAX_COSINE_DISTANCE)?;
let mut results = Vec::with_capacity(rows.len());
for row in rows {
if let Ok(node) = serde_json::from_str::<CodeNode>(&row.node_json) {
results.push((node, (1.0 - row.distance) as f32));
}
}
Ok(results)
}
fn remove_by_file(&mut self, file_path: &str) -> Result<usize> {
self.db.remove_by_file(file_path)
}
fn clear(&mut self) -> Result<()> {
self.db.clear()
}
fn entry_count(&self) -> usize {
match self.db.entry_count() {
Ok(n) => n,
Err(e) => {
tracing::warn!("语义索引条目计数失败,按空库处理: {}", e);
0
}
}
}
}
fn vec_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::*;
use crate::config::schema::EmbedSection;
use crate::model::{NodeId, NodeKind};
use std::collections::HashMap;
use std::io::{Read, Write};
use std::net::{TcpListener, TcpStream};
use std::sync::Mutex;
use tokio::runtime::Runtime;
fn test_runtime() -> Arc<Runtime> {
Arc::new(Runtime::new().unwrap())
}
fn mock_embedder() -> Arc<EmbeddingEngine> {
let config = EmbedSection {
model: "text-embedding-3-small".into(),
api_key: Some("test-key".into()),
api_key_env: "OPENAI_API_KEY".into(),
base_url: Some("http://localhost:9999/v1".into()),
};
Arc::new(EmbeddingEngine::new(&config, test_runtime().handle().clone()).unwrap())
}
fn embedder_with_server(base_url: &str, rt: &Arc<Runtime>) -> Arc<EmbeddingEngine> {
let config = EmbedSection {
model: "text-embedding-3-small".into(),
api_key: Some("test-key".into()),
api_key_env: "OPENAI_API_KEY".into(),
base_url: Some(format!("{}/v1", base_url)),
};
Arc::new(EmbeddingEngine::new(&config, rt.handle().clone()).unwrap())
}
fn tmp_path(label: &str) -> std::path::PathBuf {
use std::sync::atomic::{AtomicU64, Ordering};
static SEM_COUNTER: AtomicU64 = AtomicU64::new(0);
let id = SEM_COUNTER.fetch_add(1, Ordering::Relaxed);
let mut p = std::env::temp_dir();
p.push(format!("semantic_fts_{}_{}.db", label, id));
let _ = std::fs::remove_file(&p);
p
}
fn make_node(name: &str, file: &str) -> CodeNode {
CodeNode {
id: NodeId::new(0),
kind: NodeKind::Function,
name: name.into(),
file_path: Some(file.into()),
line_range: Some((1, 10)),
doc_comment: None,
signature: Some(format!("fn {}()", name)), visibility: None,
module_path: vec![],
}
}
fn header_complete(buf: &[u8]) -> bool {
buf.windows(4).any(|w| w == b"\r\n\r\n")
}
fn read_request_body(stream: &mut TcpStream) -> String {
let mut buf = Vec::new();
let mut tmp = [0u8; 4096];
while !header_complete(&buf) {
match stream.read(&mut tmp) {
Ok(0) | Err(_) => break,
Ok(n) => buf.extend_from_slice(&tmp[..n]),
}
}
let head_end = buf.windows(4).position(|w| w == b"\r\n\r\n").unwrap_or(buf.len());
let head = String::from_utf8_lossy(&buf[..head_end]).to_string();
let content_length = head
.split("\r\n")
.filter_map(|l| l.split_once(':'))
.find(|(k, _)| k.trim().eq_ignore_ascii_case("content-length"))
.and_then(|(_, v)| v.trim().parse::<usize>().ok())
.unwrap_or(0);
const HEADER_SEP: usize = 4;
while buf.len() < head_end + HEADER_SEP + content_length {
match stream.read(&mut tmp) {
Ok(0) | Err(_) => break,
Ok(n) => buf.extend_from_slice(&tmp[..n]),
}
}
String::from_utf8_lossy(&buf[head_end + HEADER_SEP..head_end + HEADER_SEP + content_length]).to_string()
}
fn pseudo_vector(keyword: &str, seen: &mut HashMap<String, usize>) -> Vec<f32> {
match keyword {
"alpha" => vec![1.0, 0.0, 0.0],
"beta" => vec![std::f32::consts::FRAC_1_SQRT_2, std::f32::consts::FRAC_1_SQRT_2, 0.0],
"gamma" => vec![-1.0, 0.0, 0.0],
_ => {
let next = seen.len();
let idx = *seen.entry(keyword.to_string()).or_insert(next);
let mut v = vec![0.0f32; 3];
v[idx % 3] = 1.0;
v
}
}
}
fn spawn_pseudo_embed_server() -> String {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let base_url = format!("http://{}", listener.local_addr().unwrap());
let seen: Arc<Mutex<HashMap<String, usize>>> = Arc::new(Mutex::new(HashMap::new()));
std::thread::spawn(move || {
for stream in listener.incoming() {
let Ok(mut stream) = stream else { break };
let seen = seen.clone();
std::thread::spawn(move || {
let body = read_request_body(&mut stream);
let inputs: Vec<String> = serde_json::from_str::<serde_json::Value>(&body)
.ok()
.and_then(|v| v["input"].as_array().map(|a| {
a.iter().filter_map(|x| x.as_str().map(String::from)).collect()
}))
.unwrap_or_default();
let mut guard = seen.lock().unwrap();
let vectors: Vec<Vec<f32>> = inputs.iter()
.map(|t| pseudo_vector(t.split_whitespace().next().unwrap_or(""), &mut guard))
.collect();
drop(guard);
let payload = serde_json::json!({
"data": vectors.iter().map(|v| serde_json::json!({"embedding": v})).collect::<Vec<_>>()
}).to_string();
let raw = format!(
"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{}",
payload.len(), payload
);
let _ = stream.write_all(raw.as_bytes());
});
}
});
base_url
}
#[test]
fn test_semantic_new() {
let engine = SemanticEngine::open(tmp_path("new"), mock_embedder(), test_runtime()).unwrap();
assert_eq!(engine.entry_count(), 0);
}
#[test]
fn test_search_empty() {
let engine = SemanticEngine::open(tmp_path("empty"), mock_embedder(), test_runtime()).unwrap();
assert!(engine.search("test", 10).unwrap().is_empty());
}
#[test]
fn test_semantic_search_ranks_by_similarity() {
let base_url = spawn_pseudo_embed_server();
let rt = test_runtime();
let embedder = embedder_with_server(&base_url, &rt);
let mut engine = SemanticEngine::open(tmp_path("rank"), embedder, rt.clone()).unwrap();
let items = vec![
(make_node("alpha", "src/a.rs"), "fn alpha()".to_string()),
(make_node("beta", "src/b.rs"), "fn beta()".to_string()),
(make_node("gamma", "src/c.rs"), "fn gamma()".to_string()),
];
engine.index_batch(&items).unwrap();
let results = engine.search("alpha", 10).unwrap();
assert_eq!(results.len(), 2);
assert_eq!(results[0].0.name, "alpha");
assert_eq!(results[1].0.name, "beta");
assert!(results[0].1 > results[1].1);
}
#[test]
fn test_semantic_search_filters_below_threshold() {
let base_url = spawn_pseudo_embed_server();
let rt = test_runtime();
let embedder = embedder_with_server(&base_url, &rt);
let mut engine = SemanticEngine::open(tmp_path("thr"), embedder, rt.clone()).unwrap();
let items = vec![
(make_node("x1", "src/x1.rs"), "x1 unrelated".to_string()),
(make_node("x2", "src/x2.rs"), "x2 unrelated".to_string()),
];
engine.index_batch(&items).unwrap();
let results = engine.search("q", 10).unwrap();
assert!(results.is_empty());
}
#[test]
fn test_semantic_remove_by_file() {
let base_url = spawn_pseudo_embed_server();
let rt = test_runtime();
let embedder = embedder_with_server(&base_url, &rt);
let mut engine = SemanticEngine::open(tmp_path("rm"), embedder, rt.clone()).unwrap();
let items = vec![
(make_node("a1", "src/a.rs"), "fn a1()".to_string()),
(make_node("b1", "src/b.rs"), "fn b1()".to_string()),
];
engine.index_batch(&items).unwrap();
assert_eq!(engine.entry_count(), 2);
let removed = engine.remove_by_file("src/a.rs").unwrap();
assert_eq!(removed, 1);
assert_eq!(engine.entry_count(), 1);
}
#[test]
fn test_semantic_clear() {
let base_url = spawn_pseudo_embed_server();
let rt = test_runtime();
let embedder = embedder_with_server(&base_url, &rt);
let mut engine = SemanticEngine::open(tmp_path("clr"), embedder, rt.clone()).unwrap();
engine.index_batch(&[(make_node("a", "src/a.rs"), "fn a()".to_string())]).unwrap();
assert_eq!(engine.entry_count(), 1);
engine.clear().unwrap();
assert_eq!(engine.entry_count(), 0);
}
#[test]
fn test_semantic_table_dimension() {
let base_url = spawn_pseudo_embed_server();
let rt = test_runtime();
let embedder = embedder_with_server(&base_url, &rt);
let mut engine = SemanticEngine::open(tmp_path("dim"), embedder, rt.clone()).unwrap();
assert_eq!(engine.table_dimension().unwrap(), None, "空库(表未创建)应返回 None");
engine
.index_batch(&[(make_node("a1", "src/a.rs"), "fn a1()".to_string())])
.unwrap();
assert_eq!(engine.table_dimension().unwrap(), Some(3), "伪向量统一 3 维");
}
}