Skip to main content

rectilinear_core/search/
mod.rs

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
55/// Perform a search using the specified mode
56pub 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                // Fall back to FTS-only if no embedder
101                fts_results
102            }
103        }
104    };
105
106    // Post-filter
107    let results: Vec<_> = results
108        .into_iter()
109        .filter(|_r| {
110            if let Some(_team) = team_key {
111                // We'd need team info - for now skip team filter in post-filter
112                // since FTS doesn't return team info directly
113                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, // BM25 returns negative scores, lower = better
151            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    // Compute similarity for each chunk, take max per issue
176    let mut issue_max_sim: HashMap<String, (f32, String)> = HashMap::new(); // issue_id -> (max_sim, identifier)
177
178    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    // Sort by similarity descending
191    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    // Get issue details for top results
195    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    // Post-filter by label_ids: keep only issues that have ALL required labels
216    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
231/// Reciprocal Rank Fusion combining FTS and vector results
232fn 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
273/// Find duplicates for a given text using vector similarity
274pub 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    // For duplicate finding, also do a pure vector search and merge
301    let _vec_results =
302        vector_search(db, text, team_key, limit, embedder, workspace_id, None).await?;
303
304    // Keep results above threshold
305    results.retain(|r| r.similarity.unwrap_or(0.0) >= threshold || r.score > 0.01);
306
307    Ok(results)
308}
309
310/// Build an FTS5 query from free-text input
311fn build_fts_query(input: &str) -> String {
312    // Split into words, wrap each in quotes to handle special chars
313    let words: Vec<_> = input
314        .split_whitespace()
315        .filter_map(|w| {
316            // Remove FTS5 special characters
317            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}