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(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                // Fall back to FTS-only if no embedder
84                fts_results
85            }
86        }
87    };
88
89    // Post-filter
90    let results: Vec<_> = results
91        .into_iter()
92        .filter(|_r| {
93            if let Some(_team) = team_key {
94                // We'd need team info - for now skip team filter in post-filter
95                // since FTS doesn't return team info directly
96                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, // BM25 returns negative scores, lower = better
134            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    // Compute similarity for each chunk, take max per issue
159    let mut issue_max_sim: HashMap<String, (f32, String)> = HashMap::new(); // issue_id -> (max_sim, identifier)
160
161    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    // Sort by similarity descending
174    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    // Get issue details for top results
178    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    // Post-filter by label_ids: keep only issues that have ALL required labels
199    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
214/// Reciprocal Rank Fusion combining FTS and vector results
215fn 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
256/// Find duplicates for a given text using vector similarity
257pub 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    // For duplicate finding, also do a pure vector search and merge
284    let _vec_results = vector_search(db, text, team_key, limit, embedder, workspace_id, None).await?;
285
286    // Keep results above threshold
287    results.retain(|r| r.similarity.unwrap_or(0.0) >= threshold || r.score > 0.01);
288
289    Ok(results)
290}
291
292/// Build an FTS5 query from free-text input
293fn build_fts_query(input: &str) -> String {
294    // Split into words, wrap each in quotes to handle special chars
295    let words: Vec<_> = input
296        .split_whitespace()
297        .filter_map(|w| {
298            // Remove FTS5 special characters
299            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}