1use anyhow::Result;
2use serde::{Deserialize, Serialize};
3use std::collections::HashMap;
4
5use crate::db::Database;
6use crate::embedding::{self, Embedder};
7
8#[derive(Debug, Clone, Copy, PartialEq)]
9pub enum SearchMode {
10 Fts,
11 Vector,
12 Hybrid,
13}
14
15impl std::str::FromStr for SearchMode {
16 type Err = anyhow::Error;
17 fn from_str(s: &str) -> Result<Self> {
18 match s.to_lowercase().as_str() {
19 "fts" => Ok(Self::Fts),
20 "vector" => Ok(Self::Vector),
21 "hybrid" => Ok(Self::Hybrid),
22 _ => anyhow::bail!("Invalid search mode: {}. Use fts, vector, or hybrid", s),
23 }
24 }
25}
26
27#[derive(Debug, Clone, Serialize, Deserialize)]
28pub struct SearchResult {
29 pub issue_id: String,
30 pub identifier: String,
31 pub title: String,
32 pub state_name: String,
33 pub priority: i32,
34 pub score: f64,
35 #[serde(skip_serializing_if = "Option::is_none")]
36 pub fts_rank: Option<usize>,
37 #[serde(skip_serializing_if = "Option::is_none")]
38 pub vector_rank: Option<usize>,
39 #[serde(skip_serializing_if = "Option::is_none")]
40 pub similarity: Option<f32>,
41}
42
43pub struct SearchParams<'a> {
44 pub query: &'a str,
45 pub mode: SearchMode,
46 pub team_key: Option<&'a str>,
47 pub state_filter: Option<&'a str>,
48 pub label_ids: Option<&'a [String]>,
49 pub limit: usize,
50 pub embedder: Option<&'a Embedder>,
51 pub rrf_k: u32,
52 pub workspace_id: &'a str,
53}
54
55pub async fn search(db: &Database, params: SearchParams<'_>) -> Result<Vec<SearchResult>> {
57 let SearchParams {
58 query,
59 mode,
60 team_key,
61 state_filter,
62 label_ids,
63 limit,
64 embedder,
65 rrf_k,
66 workspace_id,
67 } = params;
68 let results = match mode {
69 SearchMode::Fts => fts_search(db, query, limit * 2, workspace_id, label_ids)?,
70 SearchMode::Vector => {
71 let embedder =
72 embedder.ok_or_else(|| anyhow::anyhow!("Embedder required for vector search"))?;
73 vector_search(
74 db,
75 query,
76 team_key,
77 limit * 2,
78 embedder,
79 workspace_id,
80 label_ids,
81 )
82 .await?
83 }
84 SearchMode::Hybrid => {
85 let fts_results = fts_search(db, query, limit * 3, workspace_id, label_ids)?;
86
87 if let Some(embedder) = embedder {
88 let vec_results = vector_search(
89 db,
90 query,
91 team_key,
92 limit * 3,
93 embedder,
94 workspace_id,
95 label_ids,
96 )
97 .await?;
98 reciprocal_rank_fusion(fts_results, vec_results, rrf_k, 0.3, 0.7)
99 } else {
100 fts_results
102 }
103 }
104 };
105
106 let results: Vec<_> = results
108 .into_iter()
109 .filter(|_r| {
110 if let Some(_team) = team_key {
111 true
114 } else {
115 true
116 }
117 })
118 .filter(|r| {
119 if let Some(state) = state_filter {
120 r.state_name.to_lowercase().contains(&state.to_lowercase())
121 } else {
122 true
123 }
124 })
125 .take(limit)
126 .collect();
127
128 Ok(results)
129}
130
131fn fts_search(
132 db: &Database,
133 query: &str,
134 limit: usize,
135 workspace_id: &str,
136 label_ids: Option<&[String]>,
137) -> Result<Vec<SearchResult>> {
138 let fts_query = build_fts_query(query);
139 let fts_results = db.fts_search_filtered(&fts_query, limit, workspace_id, label_ids)?;
140
141 Ok(fts_results
142 .into_iter()
143 .enumerate()
144 .map(|(rank, r)| SearchResult {
145 issue_id: r.issue_id,
146 identifier: r.identifier,
147 title: r.title,
148 state_name: r.state_name,
149 priority: r.priority,
150 score: -r.bm25_score, fts_rank: Some(rank + 1),
152 vector_rank: None,
153 similarity: None,
154 })
155 .collect())
156}
157
158async fn vector_search(
159 db: &Database,
160 query: &str,
161 team_key: Option<&str>,
162 limit: usize,
163 embedder: &Embedder,
164 workspace_id: &str,
165 label_ids: Option<&[String]>,
166) -> Result<Vec<SearchResult>> {
167 let query_embedding = embedder.embed_single(query).await?;
168
169 let chunks = if let Some(team) = team_key {
170 db.get_chunks_for_team(team, workspace_id)?
171 } else {
172 db.get_all_chunks(workspace_id)?
173 };
174
175 let mut issue_max_sim: HashMap<String, (f32, String)> = HashMap::new(); for chunk in &chunks {
179 let chunk_embedding = embedding::bytes_to_embedding(&chunk.embedding);
180 let sim = embedding::cosine_similarity(&query_embedding, &chunk_embedding);
181
182 let entry = issue_max_sim
183 .entry(chunk.issue_id.clone())
184 .or_insert((0.0, chunk.identifier.clone()));
185 if sim > entry.0 {
186 entry.0 = sim;
187 }
188 }
189
190 let mut results: Vec<_> = issue_max_sim.into_iter().collect();
192 results.sort_by(|a, b| b.1 .0.partial_cmp(&a.1 .0).unwrap());
193
194 let results: Vec<_> = results
196 .into_iter()
197 .take(limit)
198 .enumerate()
199 .filter_map(|(rank, (issue_id, (sim, _identifier)))| {
200 let issue = db.get_issue(&issue_id).ok()??;
201 Some(SearchResult {
202 issue_id,
203 identifier: issue.identifier,
204 title: issue.title,
205 state_name: issue.state_name,
206 priority: issue.priority,
207 score: sim as f64,
208 fts_rank: None,
209 vector_rank: Some(rank + 1),
210 similarity: Some(sim),
211 })
212 })
213 .collect();
214
215 let results = if let Some(required_ids) = label_ids.filter(|ids| !ids.is_empty()) {
217 results
218 .into_iter()
219 .filter(|r| {
220 let issue_labels = db.get_issue_label_ids(&r.issue_id).unwrap_or_default();
221 required_ids.iter().all(|req| issue_labels.contains(req))
222 })
223 .collect()
224 } else {
225 results
226 };
227
228 Ok(results)
229}
230
231fn reciprocal_rank_fusion(
233 fts_results: Vec<SearchResult>,
234 vec_results: Vec<SearchResult>,
235 k: u32,
236 fts_weight: f64,
237 vec_weight: f64,
238) -> Vec<SearchResult> {
239 let mut scores: HashMap<String, (f64, SearchResult)> = HashMap::new();
240 let k = k as f64;
241
242 for (rank, result) in fts_results.into_iter().enumerate() {
243 let rrf_score = fts_weight / (k + (rank + 1) as f64);
244 let entry = scores
245 .entry(result.issue_id.clone())
246 .or_insert((0.0, result.clone()));
247 entry.0 += rrf_score;
248 entry.1.fts_rank = result.fts_rank;
249 }
250
251 for (rank, result) in vec_results.into_iter().enumerate() {
252 let rrf_score = vec_weight / (k + (rank + 1) as f64);
253 let entry = scores
254 .entry(result.issue_id.clone())
255 .or_insert((0.0, result.clone()));
256 entry.0 += rrf_score;
257 entry.1.vector_rank = result.vector_rank;
258 entry.1.similarity = result.similarity;
259 }
260
261 let mut results: Vec<_> = scores
262 .into_values()
263 .map(|(score, mut result)| {
264 result.score = score;
265 result
266 })
267 .collect();
268
269 results.sort_by(|a, b| b.score.partial_cmp(&a.score).unwrap());
270 results
271}
272
273pub async fn find_duplicates(
275 db: &Database,
276 text: &str,
277 team_key: Option<&str>,
278 threshold: f32,
279 limit: usize,
280 embedder: &Embedder,
281 rrf_k: u32,
282 workspace_id: &str,
283) -> Result<Vec<SearchResult>> {
284 let mut results = search(
285 db,
286 SearchParams {
287 query: text,
288 mode: SearchMode::Hybrid,
289 team_key,
290 state_filter: None,
291 label_ids: None,
292 limit,
293 embedder: Some(embedder),
294 rrf_k,
295 workspace_id,
296 },
297 )
298 .await?;
299
300 let _vec_results =
302 vector_search(db, text, team_key, limit, embedder, workspace_id, None).await?;
303
304 results.retain(|r| r.similarity.unwrap_or(0.0) >= threshold || r.score > 0.01);
306
307 Ok(results)
308}
309
310fn build_fts_query(input: &str) -> String {
312 let words: Vec<_> = input
314 .split_whitespace()
315 .filter_map(|w| {
316 let clean: String = w
318 .chars()
319 .filter(|c| c.is_alphanumeric() || *c == '_' || *c == '-')
320 .collect();
321 if clean.is_empty() {
322 None
323 } else {
324 Some(format!("\"{}\"", clean))
325 }
326 })
327 .collect();
328
329 if words.is_empty() {
330 "\"\"".to_string()
331 } else {
332 words.join(" OR ")
333 }
334}