Skip to main content

atheneum/graph/
search.rs

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    /// Ensure the HNSW index exists. Creates it lazily on first use.
69    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    /// Add a single entity's vector to the existing HNSW index.
98    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    /// Full rebuild of the HNSW index (still useful for manual reindexing).
113    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    /// Search discoveries using a hash-projected bag-of-tokens index (HNSW).
120    ///
121    /// Finds entities that share tokens with `query`. This is **lexical similarity**,
122    /// not semantic/neural similarity — synonyms with no token overlap will not match.
123    /// For true semantic search, embeddings from a language model would be needed.
124    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}