use std::collections::HashSet;
use gitcortex_core::{error::Result, graph::Node, schema::NodeKind, store::GraphStore};
use serde::Serialize;
#[derive(Debug, Clone, Serialize)]
pub struct SearchHit {
pub id: String,
pub name: String,
pub qualified_name: String,
pub kind: String,
pub file: String,
pub start_line: u32,
pub score: i32,
}
#[derive(Debug, Clone, Serialize)]
pub struct FileGroup {
pub file: String,
pub symbol_count: usize,
pub top_symbols: Vec<String>,
}
pub fn group_by_file(hits: &[SearchHit]) -> Vec<FileGroup> {
let mut order: Vec<String> = Vec::new();
let mut map: std::collections::HashMap<String, (usize, Vec<String>)> =
std::collections::HashMap::new();
for h in hits {
let e = map.entry(h.file.clone()).or_insert_with(|| {
order.push(h.file.clone());
(0, Vec::new())
});
e.0 += 1;
if e.1.len() < 3 {
e.1.push(h.name.clone());
}
}
let mut groups: Vec<FileGroup> = order
.into_iter()
.map(|file| {
let (count, top) = map.remove(&file).unwrap();
FileGroup {
file,
symbol_count: count,
top_symbols: top,
}
})
.collect();
groups.sort_by(|a, b| {
b.symbol_count
.cmp(&a.symbol_count)
.then_with(|| a.file.cmp(&b.file))
});
groups
}
const DEFAULT_LIMIT: usize = 10;
const MAX_LIMIT: usize = 200;
const MIN_TOKEN_LEN: usize = 3;
pub(crate) fn tokenize(s: &str) -> Vec<String> {
let mut tokens = Vec::new();
let mut current = String::new();
let chars: Vec<char> = s.chars().collect();
for (i, &ch) in chars.iter().enumerate() {
if ch == '_' || ch == '-' || ch == '.' || ch == ':' || ch == '/' || ch == ' ' {
if !current.is_empty() {
tokens.push(current.to_ascii_lowercase());
current = String::new();
}
} else if ch.is_uppercase() {
let next_is_lower = chars.get(i + 1).map(|c| c.is_lowercase()).unwrap_or(false);
let prev_is_upper = i > 0 && chars[i - 1].is_uppercase();
if !current.is_empty() && (!prev_is_upper || next_is_lower) {
tokens.push(current.to_ascii_lowercase());
current = String::new();
}
current.push(ch.to_ascii_lowercase());
} else {
current.push(ch);
}
}
if !current.is_empty() {
tokens.push(current.to_ascii_lowercase());
}
tokens
}
fn edit_distance(a: &str, b: &str) -> usize {
let a: Vec<char> = a.chars().collect();
let b: Vec<char> = b.chars().collect();
let m = a.len();
let n = b.len();
if m.abs_diff(n) > 3 {
return usize::MAX;
}
let mut prev: Vec<usize> = (0..=n).collect();
let mut curr = vec![0usize; n + 1];
for i in 1..=m {
curr[0] = i;
for j in 1..=n {
curr[j] = if a[i - 1] == b[j - 1] {
prev[j - 1]
} else {
1 + prev[j - 1].min(prev[j]).min(curr[j - 1])
};
}
std::mem::swap(&mut prev, &mut curr);
}
prev[n]
}
const EXACT_CASE_BONUS: i32 = 8;
const TEST_FILE_PENALTY: i32 = -45;
fn score(n: &Node, q_lower: &str, q_tokens: &[String], query: &str) -> Option<i32> {
if n.kind == NodeKind::Section {
return None;
}
let name_lower = n.name.to_ascii_lowercase();
let qname_lower = n.qualified_name.to_ascii_lowercase();
let name_tokens = tokenize(&n.name);
let base = if name_lower == q_lower {
100
} else if name_lower.starts_with(q_lower) {
60
} else if !q_tokens.is_empty() && q_tokens.iter().all(|t| name_tokens.contains(t)) {
50
} else if name_lower.contains(q_lower) {
30
} else {
let overlap = q_tokens
.iter()
.filter(|qt| qt.len() >= MIN_TOKEN_LEN && name_tokens.contains(*qt))
.count();
if overlap > 0 {
10 + (overlap as i32 * 5).min(15)
} else if qname_lower.contains(q_lower) {
10
} else if q_lower.len() >= 4 && q_lower.len() <= 15 && name_lower.len() <= 25 {
let dist = edit_distance(q_lower, &name_lower);
if dist <= 1 {
20
} else if dist <= 2 {
10
} else {
return None;
}
} else {
return None;
}
};
let exact_case = if n.name == query { EXACT_CASE_BONUS } else { 0 };
let test_penalty = if super::helpers::is_test_file(&n.file) {
TEST_FILE_PENALTY
} else {
0
};
Some(base + kind_boost(&n.kind) + exact_case + test_penalty)
}
fn kind_boost(k: &NodeKind) -> i32 {
match k {
NodeKind::Function | NodeKind::Method => 5,
NodeKind::Struct | NodeKind::Trait | NodeKind::Interface => 4,
NodeKind::Enum | NodeKind::TypeAlias => 3,
NodeKind::Constant | NodeKind::Macro | NodeKind::Annotation => 2,
NodeKind::File => -25,
NodeKind::Folder => -45,
_ => 0,
}
}
fn to_hit(n: Node, score: i32) -> SearchHit {
SearchHit {
id: n.id.as_str().to_owned(),
name: n.name,
qualified_name: n.qualified_name,
kind: n.kind.to_string(),
file: n.file.display().to_string(),
start_line: n.span.start_line,
score,
}
}
pub fn search<S: GraphStore + ?Sized>(
store: &S,
branch: &str,
query: &str,
limit: Option<usize>,
) -> Result<Vec<SearchHit>> {
let limit = limit.unwrap_or(DEFAULT_LIMIT).min(MAX_LIMIT);
let q = query.trim();
if q.is_empty() {
return Ok(Vec::new());
}
let q_lower = q.to_ascii_lowercase();
let q_tokens = tokenize(q);
let candidate_limit = (limit * 50).max(500);
let mut seen: HashSet<String> = HashSet::new();
let mut nodes: Vec<Node> = Vec::new();
let push = |nodes: &mut Vec<Node>, seen: &mut HashSet<String>, batch: Vec<Node>| {
for n in batch {
let id = n.id.as_str();
if seen.insert(id) {
nodes.push(n);
}
}
};
push(
&mut nodes,
&mut seen,
store.search_nodes(branch, q, candidate_limit)?,
);
for token in &q_tokens {
if token.len() < MIN_TOKEN_LEN {
continue;
}
if token.as_str() == q_lower {
continue;
}
push(
&mut nodes,
&mut seen,
store.search_nodes(branch, token, candidate_limit)?,
);
}
if nodes.is_empty() && q_lower.len() >= 4 && q_lower.len() <= 20 {
push(&mut nodes, &mut seen, store.list_all_nodes(branch)?);
}
let mut hits: Vec<SearchHit> = nodes
.into_iter()
.filter_map(|n| score(&n, &q_lower, &q_tokens, query).map(|s| to_hit(n, s)))
.collect();
hits.sort_by(|a, b| {
b.score
.cmp(&a.score)
.then_with(|| a.name.len().cmp(&b.name.len()))
.then_with(|| a.qualified_name.cmp(&b.qualified_name))
});
hits.truncate(limit);
Ok(hits)
}
#[cfg(test)]
mod tests {
use super::*;
use gitcortex_core::graph::{NodeId, NodeMetadata, Span};
use std::path::PathBuf;
fn node_of(kind: NodeKind, name: &str) -> Node {
Node {
id: NodeId::new(),
kind,
name: name.to_owned(),
qualified_name: name.to_owned(),
file: PathBuf::from("src/lib.rs"),
span: Span {
start_line: 1,
end_line: 2,
},
metadata: NodeMetadata::default(),
}
}
fn score_of(kind: NodeKind, name: &str, query: &str) -> Option<i32> {
let q_lower = query.to_ascii_lowercase();
let q_tokens = tokenize(query);
score(&node_of(kind, name), &q_lower, &q_tokens, query)
}
fn score_at(kind: NodeKind, name: &str, file: &str, query: &str) -> Option<i32> {
let q_lower = query.to_ascii_lowercase();
let q_tokens = tokenize(query);
let mut n = node_of(kind, name);
n.file = PathBuf::from(file);
score(&n, &q_lower, &q_tokens, query)
}
#[test]
fn case_exact_definition_outranks_case_insensitive_method() {
let struct_hit = score_of(NodeKind::Struct, "Searcher", "Searcher").unwrap();
let method_hit = score_of(NodeKind::Method, "searcher", "Searcher").unwrap();
assert!(
struct_hit > method_hit,
"case-exact struct {struct_hit} must outrank case-insensitive method {method_hit}"
);
}
#[test]
fn case_insensitive_match_still_scores() {
assert!(score_of(NodeKind::Method, "searcher", "Searcher").is_some());
}
#[test]
fn folder_ranks_below_a_weaker_code_match() {
let folder = score_of(NodeKind::Folder, "searcher", "Searcher").unwrap();
let prefix_struct = score_of(NodeKind::Struct, "SearcherBuilder", "Searcher").unwrap();
assert!(
folder < prefix_struct,
"folder {folder} must rank below prefix-matched struct {prefix_struct}"
);
}
#[test]
fn file_ranks_below_a_definition_of_the_same_name() {
let file = score_of(NodeKind::File, "Searcher", "Searcher").unwrap();
let definition = score_of(NodeKind::Struct, "Searcher", "Searcher").unwrap();
assert!(
file < definition,
"file {file} must rank below struct {definition}"
);
}
#[test]
fn file_is_still_findable_by_its_own_name() {
assert!(score_of(NodeKind::File, "sessions.py", "sessions.py").is_some());
}
#[test]
fn test_symbols_rank_below_production_symbols() {
let test_exact = score_at(
NodeKind::Method,
"jsonReader",
"src/test/java/NumberLimitsTest.java",
"JsonReader",
)
.unwrap();
let prod_prefix = score_at(
NodeKind::Struct,
"JsonReaderInternal",
"src/main/java/JsonReaderInternal.java",
"JsonReader",
)
.unwrap();
assert!(
test_exact < prod_prefix,
"exact-match test symbol {test_exact} must rank below production prefix match {prod_prefix}"
);
}
#[test]
fn test_support_modules_rank_below_production() {
let helper = score_at(
NodeKind::Struct,
"SearcherTester",
"crates/searcher/src/testutil.rs",
"Searcher",
)
.unwrap();
let production = score_at(
NodeKind::Struct,
"SearcherBuilder",
"crates/searcher/src/searcher/mod.rs",
"Searcher",
)
.unwrap();
assert!(
helper < production,
"test-support struct {helper} must rank below production struct {production}"
);
}
#[test]
fn test_symbols_are_still_findable() {
assert!(score_at(
NodeKind::Struct,
"JsonReaderTest",
"src/test/java/JsonReaderTest.java",
"JsonReaderTest"
)
.is_some());
}
#[test]
fn markdown_sections_are_excluded_from_code_search() {
assert_eq!(score_of(NodeKind::Section, "Searcher", "Searcher"), None);
}
#[test]
fn tokenize_camel_case() {
assert_eq!(tokenize("AuthConfig"), vec!["auth", "config"]);
assert_eq!(tokenize("validateToken"), vec!["validate", "token"]);
assert_eq!(tokenize("HTTPClient"), vec!["http", "client"]);
}
#[test]
fn tokenize_snake_case() {
assert_eq!(tokenize("validate_token"), vec!["validate", "token"]);
assert_eq!(tokenize("auth_middleware"), vec!["auth", "middleware"]);
}
#[test]
fn tokenize_pascal_case() {
assert_eq!(tokenize("KuzuGraphStore"), vec!["kuzu", "graph", "store"]);
}
#[test]
fn edit_distance_exact() {
assert_eq!(edit_distance("validate", "validate"), 0);
}
#[test]
fn edit_distance_typo() {
assert_eq!(edit_distance("vlidate", "validate"), 1);
assert_eq!(edit_distance("authnticate", "authenticate"), 1);
}
#[test]
fn edit_distance_length_short_circuit() {
assert_eq!(edit_distance("a", "abcde"), usize::MAX);
}
fn hit(name: &str, file: &str, score: i32) -> SearchHit {
SearchHit {
id: String::new(),
name: name.to_owned(),
qualified_name: name.to_owned(),
kind: "Function".to_owned(),
file: file.to_owned(),
start_line: 1,
score,
}
}
#[test]
fn group_by_file_empty_returns_empty() {
assert!(group_by_file(&[]).is_empty());
}
#[test]
fn group_by_file_single_file() {
let hits = vec![hit("parse_args", "src/parser.rs", 100)];
let groups = group_by_file(&hits);
assert_eq!(groups.len(), 1);
assert_eq!(groups[0].file, "src/parser.rs");
assert_eq!(groups[0].symbol_count, 1);
assert_eq!(groups[0].top_symbols, vec!["parse_args"]);
}
#[test]
fn group_by_file_sorted_by_count_desc() {
let hits = vec![
hit("a", "src/big.rs", 80),
hit("b", "src/big.rs", 70),
hit("c", "src/big.rs", 60),
hit("x", "src/small.rs", 50),
];
let groups = group_by_file(&hits);
assert_eq!(groups[0].file, "src/big.rs");
assert_eq!(groups[0].symbol_count, 3);
assert_eq!(groups[1].file, "src/small.rs");
assert_eq!(groups[1].symbol_count, 1);
}
#[test]
fn group_by_file_top_symbols_capped_at_three() {
let hits: Vec<SearchHit> = (0..8)
.map(|i| hit(&format!("fn{i}"), "src/big.rs", 100 - i))
.collect();
let groups = group_by_file(&hits);
assert_eq!(groups[0].symbol_count, 8);
assert_eq!(groups[0].top_symbols.len(), 3);
assert_eq!(groups[0].top_symbols[0], "fn0");
}
#[test]
fn group_by_file_alphabetical_tie_break() {
let hits = vec![hit("a", "src/z.rs", 50), hit("b", "src/a.rs", 50)];
let groups = group_by_file(&hits);
assert_eq!(groups[0].file, "src/a.rs");
}
}