1use anyhow::Result;
2use serde_json::json;
3use sqlitegraph::hnsw::{DistanceMetric, HnswConfigBuilder};
4use sqlitegraph::GraphEntity;
5use std::collections::HashSet;
6
7use super::{AtheneumGraph, SearchResult};
8
9const SEARCH_INDEX_NAME: &str = "discoveries";
10
11fn embed_text_for_entity(entity: &GraphEntity) -> String {
12 let mut parts = vec![entity.kind.clone(), entity.name.clone()];
13 for key in [
14 "target",
15 "agent",
16 "discovery_type",
17 "file",
18 "file_path",
19 "summary",
20 "signature",
21 "title",
22 "path",
23 "body",
24 "kind",
25 ] {
26 if let Some(value) = entity.data.get(key).and_then(|v| v.as_str()) {
27 parts.push(value.to_string());
28 }
29 }
30 if let Some(items) = entity.data.get("wikilinks").and_then(|v| v.as_array()) {
31 for item in items {
32 if let Some(value) = item.as_str() {
33 parts.push(value.to_string());
34 }
35 }
36 }
37 parts.join(" ")
38}
39
40fn query_tokens(query: &str) -> Vec<String> {
41 let mut seen = HashSet::new();
42 query
43 .split(|c: char| !c.is_alphanumeric())
44 .filter(|t| !t.is_empty())
45 .map(|t| t.to_ascii_lowercase())
46 .filter(|t| seen.insert(t.clone()))
47 .collect()
48}
49
50fn lexical_token_score(entity: &GraphEntity, tokens: &[String]) -> f32 {
51 if tokens.is_empty() {
52 return 0.0;
53 }
54 let text = embed_text_for_entity(entity).to_ascii_lowercase();
55 let matched = tokens.iter().filter(|token| text.contains(*token)).count();
56 matched as f32 / tokens.len() as f32
57}
58
59fn search_config(dim: usize) -> Result<sqlitegraph::hnsw::HnswConfig> {
60 HnswConfigBuilder::new()
61 .dimension(dim)
62 .distance_metric(DistanceMetric::Cosine)
63 .build()
64 .map_err(|e| anyhow::anyhow!("HNSW config build failed: {}", e))
65}
66
67impl AtheneumGraph {
68 fn ensure_search_index(&self) -> Result<()> {
70 let existing = self
71 .inner
72 .list_hnsw_indexes()
73 .map_err(|e| anyhow::anyhow!("list_hnsw_indexes failed: {}", e))?;
74 if existing.iter().any(|n| n == SEARCH_INDEX_NAME) {
75 return Ok(());
76 }
77 let config = search_config(self.embedder.dimension())?;
78 {
79 let _guard = self
80 .inner
81 .hnsw_index_persistent(SEARCH_INDEX_NAME, config)
82 .map_err(|e| anyhow::anyhow!("hnsw_index_persistent create failed: {}", e))?;
83 }
84 for entity in self.all_entities()? {
85 let text = embed_text_for_entity(&entity);
86 let vector = self.embedder.embed(&text)?;
87 let entity_id = entity.id;
88 let _ = self
89 .inner
90 .get_hnsw_index_mut(SEARCH_INDEX_NAME, move |idx| {
91 idx.insert_vector(&vector, Some(json!({"entity_id": entity_id})))
92 });
93 }
94 Ok(())
95 }
96
97 pub(super) fn add_entity_to_search_index(&self, entity: &GraphEntity) -> Result<()> {
99 self.ensure_search_index()?;
100 let text = embed_text_for_entity(entity);
101 let vector = self.embedder.embed(&text)?;
102 let entity_id = entity.id;
103 self.inner
104 .get_hnsw_index_mut(SEARCH_INDEX_NAME, move |idx| {
105 idx.insert_vector(&vector, Some(json!({"entity_id": entity_id})))
106 })
107 .map_err(|e| anyhow::anyhow!("get_hnsw_index_mut failed: {}", e))?
108 .map_err(|e| anyhow::anyhow!("insert_vector failed: {}", e))?;
109 Ok(())
110 }
111
112 pub fn build_search_index(&self) -> Result<()> {
114 let _ = self.inner.delete_hnsw_index(SEARCH_INDEX_NAME);
115 self.ensure_search_index()?;
116 Ok(())
117 }
118
119 pub fn lexical_search(
125 &self,
126 query: &str,
127 k: usize,
128 project_id: Option<&str>,
129 entity_kind: Option<&str>,
130 ) -> Result<Vec<SearchResult>> {
131 self.ensure_search_index()?;
132 let query_vec = self.embedder.embed(query)?;
133 let fetch_k = if project_id.is_some() || entity_kind.is_some() {
134 k * 4
135 } else {
136 k
137 };
138
139 let hits = self
140 .inner
141 .get_hnsw_index_ref(SEARCH_INDEX_NAME, |idx| idx.search(&query_vec, fetch_k))
142 .map_err(|e| anyhow::anyhow!("search index lookup failed: {}", e))?
143 .map_err(|e| anyhow::anyhow!("hnsw search failed: {}", e))?;
144
145 let mut results = Vec::with_capacity(hits.len());
146 let mut seen_entities = HashSet::new();
147 for (vector_id, score) in hits {
148 let metadata = self
149 .inner
150 .get_hnsw_index_ref(SEARCH_INDEX_NAME, |idx| {
151 idx.get_vector(vector_id).ok().flatten()
152 })
153 .map_err(|e| anyhow::anyhow!("get_vector failed: {}", e))?;
154 let Some((_vec, meta)) = metadata else {
155 continue;
156 };
157 let Some(entity_id) = meta.get("entity_id").and_then(|v| v.as_i64()) else {
158 continue;
159 };
160
161 let entity = match self.get_entity(entity_id) {
162 Ok(e) => e,
163 Err(_) => continue,
164 };
165 if !seen_entities.insert(entity.id) {
166 continue;
167 }
168
169 if let Some(pid) = project_id {
170 let entity_project = entity
171 .data
172 .get("project_id")
173 .and_then(|v| v.as_str())
174 .unwrap_or("");
175 if entity_project != pid {
176 continue;
177 }
178 }
179
180 if let Some(kind) = entity_kind {
181 if entity.kind != kind {
182 continue;
183 }
184 }
185
186 results.push(SearchResult {
187 id: entity.id,
188 name: entity.name,
189 kind: entity.kind,
190 score,
191 data: entity.data,
192 });
193
194 if results.len() >= k {
195 break;
196 }
197 }
198
199 if results.len() < k {
200 let tokens = query_tokens(query);
201 let mut fallback = Vec::new();
202 for entity in self.all_entities()? {
203 if seen_entities.contains(&entity.id) {
204 continue;
205 }
206 if let Some(pid) = project_id {
207 let entity_project = entity
208 .data
209 .get("project_id")
210 .and_then(|v| v.as_str())
211 .unwrap_or("");
212 if entity_project != pid {
213 continue;
214 }
215 }
216 if let Some(kind) = entity_kind {
217 if entity.kind != kind {
218 continue;
219 }
220 }
221 let score = lexical_token_score(&entity, &tokens);
222 if score > 0.0 {
223 fallback.push((entity, score));
224 }
225 }
226 fallback.sort_by(|(left, left_score), (right, right_score)| {
227 right_score
228 .partial_cmp(left_score)
229 .unwrap_or(std::cmp::Ordering::Equal)
230 .then_with(|| left.name.cmp(&right.name))
231 });
232 for (entity, score) in fallback {
233 if !seen_entities.insert(entity.id) {
234 continue;
235 }
236 results.push(SearchResult {
237 id: entity.id,
238 name: entity.name,
239 kind: entity.kind,
240 score,
241 data: entity.data,
242 });
243 if results.len() >= k {
244 break;
245 }
246 }
247 }
248 Ok(results)
249 }
250}