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(db, query, team_key, limit * 2, embedder, workspace_id, label_ids).await?
74 }
75 SearchMode::Hybrid => {
76 let fts_results = fts_search(db, query, limit * 3, workspace_id, label_ids)?;
77
78 if let Some(embedder) = embedder {
79 let vec_results =
80 vector_search(db, query, team_key, limit * 3, embedder, workspace_id, label_ids).await?;
81 reciprocal_rank_fusion(fts_results, vec_results, rrf_k, 0.3, 0.7)
82 } else {
83 fts_results
85 }
86 }
87 };
88
89 let results: Vec<_> = results
91 .into_iter()
92 .filter(|_r| {
93 if let Some(_team) = team_key {
94 true
97 } else {
98 true
99 }
100 })
101 .filter(|r| {
102 if let Some(state) = state_filter {
103 r.state_name.to_lowercase().contains(&state.to_lowercase())
104 } else {
105 true
106 }
107 })
108 .take(limit)
109 .collect();
110
111 Ok(results)
112}
113
114fn fts_search(
115 db: &Database,
116 query: &str,
117 limit: usize,
118 workspace_id: &str,
119 label_ids: Option<&[String]>,
120) -> Result<Vec<SearchResult>> {
121 let fts_query = build_fts_query(query);
122 let fts_results = db.fts_search_filtered(&fts_query, limit, workspace_id, label_ids)?;
123
124 Ok(fts_results
125 .into_iter()
126 .enumerate()
127 .map(|(rank, r)| SearchResult {
128 issue_id: r.issue_id,
129 identifier: r.identifier,
130 title: r.title,
131 state_name: r.state_name,
132 priority: r.priority,
133 score: -r.bm25_score, fts_rank: Some(rank + 1),
135 vector_rank: None,
136 similarity: None,
137 })
138 .collect())
139}
140
141async fn vector_search(
142 db: &Database,
143 query: &str,
144 team_key: Option<&str>,
145 limit: usize,
146 embedder: &Embedder,
147 workspace_id: &str,
148 label_ids: Option<&[String]>,
149) -> Result<Vec<SearchResult>> {
150 let query_embedding = embedder.embed_single(query).await?;
151
152 let chunks = if let Some(team) = team_key {
153 db.get_chunks_for_team(team, workspace_id)?
154 } else {
155 db.get_all_chunks(workspace_id)?
156 };
157
158 let mut issue_max_sim: HashMap<String, (f32, String)> = HashMap::new(); for chunk in &chunks {
162 let chunk_embedding = embedding::bytes_to_embedding(&chunk.embedding);
163 let sim = embedding::cosine_similarity(&query_embedding, &chunk_embedding);
164
165 let entry = issue_max_sim
166 .entry(chunk.issue_id.clone())
167 .or_insert((0.0, chunk.identifier.clone()));
168 if sim > entry.0 {
169 entry.0 = sim;
170 }
171 }
172
173 let mut results: Vec<_> = issue_max_sim.into_iter().collect();
175 results.sort_by(|a, b| b.1 .0.partial_cmp(&a.1 .0).unwrap());
176
177 let results: Vec<_> = results
179 .into_iter()
180 .take(limit)
181 .enumerate()
182 .filter_map(|(rank, (issue_id, (sim, _identifier)))| {
183 let issue = db.get_issue(&issue_id).ok()??;
184 Some(SearchResult {
185 issue_id,
186 identifier: issue.identifier,
187 title: issue.title,
188 state_name: issue.state_name,
189 priority: issue.priority,
190 score: sim as f64,
191 fts_rank: None,
192 vector_rank: Some(rank + 1),
193 similarity: Some(sim),
194 })
195 })
196 .collect();
197
198 let results = if let Some(required_ids) = label_ids.filter(|ids| !ids.is_empty()) {
200 results
201 .into_iter()
202 .filter(|r| {
203 let issue_labels = db.get_issue_label_ids(&r.issue_id).unwrap_or_default();
204 required_ids.iter().all(|req| issue_labels.contains(req))
205 })
206 .collect()
207 } else {
208 results
209 };
210
211 Ok(results)
212}
213
214fn reciprocal_rank_fusion(
216 fts_results: Vec<SearchResult>,
217 vec_results: Vec<SearchResult>,
218 k: u32,
219 fts_weight: f64,
220 vec_weight: f64,
221) -> Vec<SearchResult> {
222 let mut scores: HashMap<String, (f64, SearchResult)> = HashMap::new();
223 let k = k as f64;
224
225 for (rank, result) in fts_results.into_iter().enumerate() {
226 let rrf_score = fts_weight / (k + (rank + 1) as f64);
227 let entry = scores
228 .entry(result.issue_id.clone())
229 .or_insert((0.0, result.clone()));
230 entry.0 += rrf_score;
231 entry.1.fts_rank = result.fts_rank;
232 }
233
234 for (rank, result) in vec_results.into_iter().enumerate() {
235 let rrf_score = vec_weight / (k + (rank + 1) as f64);
236 let entry = scores
237 .entry(result.issue_id.clone())
238 .or_insert((0.0, result.clone()));
239 entry.0 += rrf_score;
240 entry.1.vector_rank = result.vector_rank;
241 entry.1.similarity = result.similarity;
242 }
243
244 let mut results: Vec<_> = scores
245 .into_values()
246 .map(|(score, mut result)| {
247 result.score = score;
248 result
249 })
250 .collect();
251
252 results.sort_by(|a, b| b.score.partial_cmp(&a.score).unwrap());
253 results
254}
255
256pub async fn find_duplicates(
258 db: &Database,
259 text: &str,
260 team_key: Option<&str>,
261 threshold: f32,
262 limit: usize,
263 embedder: &Embedder,
264 rrf_k: u32,
265 workspace_id: &str,
266) -> Result<Vec<SearchResult>> {
267 let mut results = search(
268 db,
269 SearchParams {
270 query: text,
271 mode: SearchMode::Hybrid,
272 team_key,
273 state_filter: None,
274 label_ids: None,
275 limit,
276 embedder: Some(embedder),
277 rrf_k,
278 workspace_id,
279 },
280 )
281 .await?;
282
283 let _vec_results = vector_search(db, text, team_key, limit, embedder, workspace_id, None).await?;
285
286 results.retain(|r| r.similarity.unwrap_or(0.0) >= threshold || r.score > 0.01);
288
289 Ok(results)
290}
291
292fn build_fts_query(input: &str) -> String {
294 let words: Vec<_> = input
296 .split_whitespace()
297 .filter_map(|w| {
298 let clean: String = w
300 .chars()
301 .filter(|c| c.is_alphanumeric() || *c == '_' || *c == '-')
302 .collect();
303 if clean.is_empty() {
304 None
305 } else {
306 Some(format!("\"{}\"", clean))
307 }
308 })
309 .collect();
310
311 if words.is_empty() {
312 "\"\"".to_string()
313 } else {
314 words.join(" OR ")
315 }
316}