use rusqlite::Connection;
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum SymbolKind {
Function,
Struct,
Enum,
Trait,
Interface,
Class,
Method,
Constant,
Module,
TypeAlias,
Variable,
}
impl SymbolKind {
pub fn as_str(&self) -> &'static str {
match self {
Self::Function => "function",
Self::Struct => "struct",
Self::Enum => "enum",
Self::Trait => "trait",
Self::Interface => "interface",
Self::Class => "class",
Self::Method => "method",
Self::Constant => "constant",
Self::Module => "module",
Self::TypeAlias => "type_alias",
Self::Variable => "variable",
}
}
pub fn from_str(s: &str) -> Self {
match s {
"function" => Self::Function,
"struct" => Self::Struct,
"enum" => Self::Enum,
"trait" => Self::Trait,
"interface" => Self::Interface,
"class" => Self::Class,
"method" => Self::Method,
"constant" => Self::Constant,
"module" => Self::Module,
"type_alias" => Self::TypeAlias,
"variable" => Self::Variable,
_ => Self::Variable,
}
}
}
impl std::fmt::Display for SymbolKind {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}", self.as_str())
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct IndexedSymbol {
pub id: i64,
pub name: String,
pub kind: SymbolKind,
pub file: String,
pub line: u32,
pub signature: String,
pub language: String,
}
#[derive(Debug, Clone, Default)]
pub struct SymbolQuery {
pub text: Option<String>,
pub kind: Option<SymbolKind>,
pub file_prefix: Option<String>,
pub language: Option<String>,
pub limit: usize,
}
impl SymbolQuery {
#[allow(dead_code)]
pub fn text(text: &str) -> Self {
Self {
text: Some(text.to_string()),
limit: 50,
..Default::default()
}
}
#[allow(dead_code)]
pub fn kind(kind: SymbolKind) -> Self {
Self {
kind: Some(kind),
limit: 50,
..Default::default()
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SearchResult {
pub symbol: IndexedSymbol,
pub score: f64,
}
pub fn search(
conn: &Connection,
project_id: i64,
query: &SymbolQuery,
) -> anyhow::Result<Vec<SearchResult>> {
let limit = if query.limit > 0 {
query.limit as i64
} else {
50
};
if let Some(text) = &query.text {
let fts_query = sanitize_fts_query(text);
let mut sql = String::from(
"SELECT s.id, s.name, s.kind, s.file, s.line, s.signature, s.language,
bm25(symbols_fts) as score
FROM symbols_fts
JOIN symbols s ON s.id = symbols_fts.rowid
WHERE symbols_fts MATCH ?1 AND s.project_id = ?2",
);
let mut params: Vec<Box<dyn rusqlite::ToSql>> =
vec![Box::new(fts_query.clone()), Box::new(project_id)];
let mut param_idx = 3;
if let Some(kind) = &query.kind {
sql.push_str(&format!(" AND s.kind = ?{param_idx}"));
params.push(Box::new(kind.as_str().to_string()));
param_idx += 1;
}
if let Some(lang) = &query.language {
sql.push_str(&format!(" AND s.language = ?{param_idx}"));
params.push(Box::new(lang.clone()));
param_idx += 1;
}
if let Some(prefix) = &query.file_prefix {
sql.push_str(&format!(" AND s.file LIKE ?{param_idx}"));
params.push(Box::new(format!("{prefix}%")));
}
sql.push_str(&format!(" ORDER BY score ASC LIMIT {limit}"));
let mut stmt = conn.prepare(&sql)?;
let param_refs: Vec<&dyn rusqlite::ToSql> = params.iter().map(|p| p.as_ref()).collect();
let rows = stmt.query_map(param_refs.as_slice(), |row| {
Ok(SearchResult {
symbol: IndexedSymbol {
id: row.get(0)?,
name: row.get(1)?,
kind: SymbolKind::from_str(&row.get::<_, String>(2)?),
file: row.get(3)?,
line: row.get::<_, i64>(4)? as u32,
signature: row.get(5)?,
language: row.get(6)?,
},
score: {
let raw: f64 = row.get(7)?;
if raw < 0.0 { -raw } else { raw }
},
})
})?;
let mut results: Vec<SearchResult> = rows.filter_map(|r| r.ok()).collect();
if results.is_empty() {
results = like_search(conn, project_id, text, query, limit)?;
}
Ok(results)
} else {
filter_search(conn, project_id, query, limit)
}
}
#[allow(unused_assignments)]
fn like_search(
conn: &Connection,
project_id: i64,
text: &str,
query: &SymbolQuery,
limit: i64,
) -> anyhow::Result<Vec<SearchResult>> {
let pattern = format!("%{text}%");
let mut sql = String::from(
"SELECT id, name, kind, file, line, signature, language
FROM symbols WHERE name LIKE ?1 AND project_id = ?2",
);
let mut params: Vec<Box<dyn rusqlite::ToSql>> = vec![Box::new(pattern), Box::new(project_id)];
let mut idx = 3;
if let Some(kind) = &query.kind {
sql.push_str(&format!(" AND kind = ?{idx}"));
params.push(Box::new(kind.as_str().to_string()));
idx += 1;
}
sql.push_str(&format!(" LIMIT {limit}"));
let mut stmt = conn.prepare(&sql)?;
let param_refs: Vec<&dyn rusqlite::ToSql> = params.iter().map(|p| p.as_ref()).collect();
let rows = stmt.query_map(param_refs.as_slice(), |row| {
Ok(SearchResult {
symbol: IndexedSymbol {
id: row.get(0)?,
name: row.get(1)?,
kind: SymbolKind::from_str(&row.get::<_, String>(2)?),
file: row.get(3)?,
line: row.get::<_, i64>(4)? as u32,
signature: row.get(5)?,
language: row.get(6)?,
},
score: 1.0,
})
})?;
Ok(rows.filter_map(|r| r.ok()).collect())
}
#[allow(unused_assignments)]
fn filter_search(
conn: &Connection,
project_id: i64,
query: &SymbolQuery,
limit: i64,
) -> anyhow::Result<Vec<SearchResult>> {
let mut sql = String::from(
"SELECT id, name, kind, file, line, signature, language FROM symbols WHERE project_id = ?1",
);
let mut params: Vec<Box<dyn rusqlite::ToSql>> = vec![Box::new(project_id)];
let mut idx = 2;
if let Some(kind) = &query.kind {
sql.push_str(&format!(" AND kind = ?{idx}"));
params.push(Box::new(kind.as_str().to_string()));
idx += 1;
}
if let Some(lang) = &query.language {
sql.push_str(&format!(" AND language = ?{idx}"));
params.push(Box::new(lang.clone()));
idx += 1;
}
if let Some(prefix) = &query.file_prefix {
sql.push_str(&format!(" AND file LIKE ?{idx}"));
params.push(Box::new(format!("{prefix}%")));
idx += 1;
}
sql.push_str(&format!(" LIMIT {limit}"));
let mut stmt = conn.prepare(&sql)?;
let param_refs: Vec<&dyn rusqlite::ToSql> = params.iter().map(|p| p.as_ref()).collect();
let rows = stmt.query_map(param_refs.as_slice(), |row| {
Ok(SearchResult {
symbol: IndexedSymbol {
id: row.get(0)?,
name: row.get(1)?,
kind: SymbolKind::from_str(&row.get::<_, String>(2)?),
file: row.get(3)?,
line: row.get::<_, i64>(4)? as u32,
signature: row.get(5)?,
language: row.get(6)?,
},
score: 0.0,
})
})?;
Ok(rows.filter_map(|r| r.ok()).collect())
}
fn sanitize_fts_query(text: &str) -> String {
text.split_whitespace()
.map(|token| {
let clean: String = token
.chars()
.filter(|c| c.is_alphanumeric() || *c == '_' || *c == ':')
.collect();
if clean.is_empty() {
String::new()
} else {
format!("\"{clean}\"")
}
})
.filter(|s| !s.is_empty())
.collect::<Vec<_>>()
.join(" ")
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_symbol_kind_round_trip() {
let kinds = [
SymbolKind::Function,
SymbolKind::Struct,
SymbolKind::Enum,
SymbolKind::Trait,
SymbolKind::Method,
SymbolKind::Constant,
];
for kind in &kinds {
let s = kind.as_str();
assert_eq!(SymbolKind::from_str(s), *kind);
}
}
#[test]
fn test_symbol_kind_display() {
assert_eq!(SymbolKind::Function.to_string(), "function");
assert_eq!(SymbolKind::Struct.to_string(), "struct");
}
#[test]
fn test_sanitize_fts_query() {
assert_eq!(sanitize_fts_query("auth"), "\"auth\"");
assert_eq!(sanitize_fts_query("auth login"), "\"auth\" \"login\"");
assert_eq!(
sanitize_fts_query("auth; DROP TABLE"),
"\"auth\" \"DROP\" \"TABLE\""
);
assert_eq!(sanitize_fts_query(""), "");
}
#[test]
fn test_symbol_query_text() {
let q = SymbolQuery::text("authenticate");
assert_eq!(q.text.as_deref(), Some("authenticate"));
assert_eq!(q.limit, 50);
}
#[test]
fn test_symbol_query_kind() {
let q = SymbolQuery::kind(SymbolKind::Function);
assert_eq!(q.kind, Some(SymbolKind::Function));
assert_eq!(q.limit, 50);
}
}