use crate::search::code_tokenizer::expand_identifiers;
use crate::types::{Chunk, ScoredChunkId};
use regex::Regex;
use std::collections::HashMap;
use std::path::Path;
use std::sync::LazyLock;
pub const STRONG_PENALTY: f32 = 0.30; pub const MODERATE_PENALTY: f32 = 0.50; pub const MILD_PENALTY: f32 = 0.70;
pub const STEM_EXACT_BOOST_FRAC: f32 = 0.40; pub const STEM_PREFIX_BOOST_FRAC: f32 = 0.20; pub const STEM_PREFIX_MIN_LEN: usize = 3;
pub const DEFINITION_BOOST_FRAC: f32 = 0.20;
pub const FILE_COHERENCE_BOOST_FRAC: f32 = 0.20;
const PENALTY_DISABLE_TOKENS: &[&str] = &[
"test",
"tests",
"spec",
"specs",
"benchmark",
"benchmarks",
"bench",
"example",
"examples",
"demo",
];
const TEST_DIR_PATTERNS: &[&str] = &["/tests/", "/test/", "/__tests__/", "/spec/"];
const COMPAT_DIR_PATTERNS: &[&str] = &["/compat/", "/legacy/", "/deprecated/", "/polyfill/"];
const EXAMPLE_DIR_PATTERNS: &[&str] = &[
"/examples/",
"/example/",
"/demos/",
"/demo/",
"/samples/",
"/sample/",
];
const BARREL_FILENAMES: &[&str] = &[
"__init__.py",
"mod.rs",
"index.ts",
"index.js",
"package-info.java",
];
static TEST_FILENAME_RE: LazyLock<Regex> = LazyLock::new(|| {
Regex::new(r"(?i)([_./-])(test|tests|spec|specs|_test|\.spec|\.test)\.[A-Za-z0-9]+$")
.expect("path_signals: test filename regex MUST compile")
});
static TEST_FILENAME_PREFIX_RE: LazyLock<Regex> = LazyLock::new(|| {
Regex::new(r"(?i)^(test_|spec_)[A-Za-z0-9_]+\.[A-Za-z0-9]+$")
.expect("path_signals: test prefix regex MUST compile")
});
pub fn should_apply_path_penalty(query: &str) -> bool {
let lower = query.to_lowercase();
for tok in PENALTY_DISABLE_TOKENS {
if has_word_boundary_match(&lower, tok) {
return false;
}
}
true
}
fn has_word_boundary_match(haystack: &str, needle: &str) -> bool {
let nlen = needle.len();
let bytes = haystack.as_bytes();
let nbytes = needle.as_bytes();
if nlen == 0 || nlen > bytes.len() {
return false;
}
let mut i = 0;
while i + nlen <= bytes.len() {
if &bytes[i..i + nlen] == nbytes {
let before_ok = i == 0 || !(bytes[i - 1] as char).is_ascii_alphanumeric();
let after_idx = i + nlen;
let after_ok =
after_idx == bytes.len() || !(bytes[after_idx] as char).is_ascii_alphanumeric();
if before_ok && after_ok {
return true;
}
}
i += 1;
}
false
}
pub fn file_path_penalty(path: &Path) -> f32 {
let raw = path.to_string_lossy();
let lower = raw.to_lowercase();
let normalized = lower.replace('\\', "/");
let padded = format!("/{}/", normalized.trim_matches('/'));
let mut factor: f32 = 1.0;
if TEST_DIR_PATTERNS.iter().any(|p| padded.contains(p)) {
factor *= STRONG_PENALTY;
}
let filename = path
.file_name()
.map(|n| n.to_string_lossy().into_owned())
.unwrap_or_default();
if !filename.is_empty()
&& (TEST_FILENAME_RE.is_match(&filename) || TEST_FILENAME_PREFIX_RE.is_match(&filename))
{
factor *= STRONG_PENALTY;
}
if COMPAT_DIR_PATTERNS.iter().any(|p| padded.contains(p)) {
factor *= STRONG_PENALTY;
}
if EXAMPLE_DIR_PATTERNS.iter().any(|p| padded.contains(p)) {
factor *= STRONG_PENALTY;
}
if filename.to_lowercase().ends_with(".d.ts") {
factor *= MILD_PENALTY;
}
if BARREL_FILENAMES
.iter()
.any(|b| filename.eq_ignore_ascii_case(b))
{
factor *= MODERATE_PENALTY;
}
factor
}
const QUERY_STOPWORDS: &[&str] = &[
"how", "the", "for", "of", "in", "on", "at", "to", "is", "are", "and", "or", "what", "where",
"when", "why",
];
pub fn query_tokens_for_path_signals(text: &str) -> Vec<String> {
let mut out: Vec<String> = Vec::new();
let mut seen: std::collections::HashSet<String> = std::collections::HashSet::new();
let expanded = expand_identifiers(text);
for tok in expanded.split_whitespace() {
if tok.contains('_') {
continue;
}
push_token(tok, &mut out, &mut seen);
}
for raw in split_identifier_spans(text) {
for sub in split_into_subtokens(&raw) {
push_token(&sub, &mut out, &mut seen);
}
}
out
}
fn push_token(tok: &str, out: &mut Vec<String>, seen: &mut std::collections::HashSet<String>) {
if tok.is_empty() {
return;
}
let lower = tok.to_lowercase();
if QUERY_STOPWORDS.contains(&lower.as_str()) {
return;
}
if seen.insert(lower.clone()) {
out.push(lower);
}
}
fn split_identifier_spans(text: &str) -> Vec<String> {
let mut spans: Vec<String> = Vec::new();
let mut current = String::new();
for c in text.chars() {
if c.is_alphanumeric() || c == '_' {
current.push(c);
} else if !current.is_empty() {
spans.push(std::mem::take(&mut current));
}
}
if !current.is_empty() {
spans.push(current);
}
spans
}
fn split_into_subtokens(span: &str) -> Vec<String> {
let mut out: Vec<String> = Vec::new();
for piece in span.split('_').filter(|s| !s.is_empty()) {
for cp in split_camel_case_local(piece) {
if !cp.is_empty() {
out.push(cp.to_lowercase());
}
}
}
out
}
fn split_camel_case_local(word: &str) -> Vec<&str> {
let bytes = word.as_bytes();
if bytes.is_empty() {
return vec![];
}
let mut parts = Vec::new();
let mut start = 0;
for i in 1..bytes.len() {
let prev = bytes[i - 1];
let cur = bytes[i];
if (prev.is_ascii_lowercase() || prev.is_ascii_digit()) && cur.is_ascii_uppercase() {
parts.push(&word[start..i]);
start = i;
continue;
}
if i >= 2
&& bytes[i - 2].is_ascii_uppercase()
&& prev.is_ascii_uppercase()
&& cur.is_ascii_lowercase()
&& (i - 1 - start) >= 2
{
parts.push(&word[start..i - 1]);
start = i - 1;
}
}
parts.push(&word[start..]);
parts
}
pub fn path_stem_boost_factor(path: &Path, query_tokens: &[String]) -> f32 {
if query_tokens.is_empty() {
return 0.0;
}
let mut stem_str = match path.file_stem().and_then(|s| s.to_str()) {
Some(s) => s.to_string(),
None => return 0.0,
};
while let Some(inner) = Path::new(&stem_str).file_stem().and_then(|s| s.to_str()) {
if inner == stem_str {
break;
}
stem_str = inner.to_string();
}
let mut stem_tokens: Vec<String> = Vec::new();
for span in split_identifier_spans(&stem_str) {
for sub in split_into_subtokens(&span) {
if !sub.is_empty() {
stem_tokens.push(sub);
}
}
}
if stem_tokens.is_empty() {
return 0.0;
}
let mut best: f32 = 0.0;
for qt in query_tokens {
for st in &stem_tokens {
if qt == st {
return STEM_EXACT_BOOST_FRAC;
}
let shared = qt.len().min(st.len());
if shared >= STEM_PREFIX_MIN_LEN
&& (qt.starts_with(st.as_str()) || st.starts_with(qt.as_str()))
{
best = best.max(STEM_PREFIX_BOOST_FRAC);
}
}
}
best
}
pub fn definition_boost_factor(symbol_name: &str, query_tokens: &[String]) -> f32 {
if symbol_name.is_empty() || query_tokens.is_empty() {
return 0.0;
}
let mut sym_tokens: Vec<String> = Vec::new();
for span in split_identifier_spans(symbol_name) {
for sub in split_into_subtokens(&span) {
if !sub.is_empty() {
sym_tokens.push(sub);
}
}
}
if sym_tokens.is_empty() {
return 0.0;
}
for qt in query_tokens {
let lower = qt.to_lowercase();
if sym_tokens.iter().any(|s| s == &lower) {
return DEFINITION_BOOST_FRAC;
}
}
0.0
}
pub fn file_coherence_boosts<S: std::hash::BuildHasher>(
fused: &[ScoredChunkId],
chunk_map: &HashMap<u64, Chunk, S>,
) -> HashMap<u64, f32> {
#[derive(Default)]
struct Group {
sum: f32,
top_id: u64,
top_score: f32,
count: usize,
}
let mut groups: HashMap<String, Group> = HashMap::new();
for s in fused {
let Some(chunk) = chunk_map.get(&s.chunk_id) else {
continue;
};
let key = chunk.file_path.to_string_lossy().into_owned();
let entry = groups.entry(key).or_default();
entry.sum += s.score;
entry.count += 1;
if entry.count == 1 || s.score > entry.top_score {
entry.top_score = s.score;
entry.top_id = s.chunk_id;
}
}
groups.retain(|_, g| g.count >= 2);
if groups.is_empty() {
return HashMap::new();
}
let max_file_sum = groups.values().map(|g| g.sum).fold(f32::MIN, f32::max);
if !matches!(
max_file_sum.partial_cmp(&0.0),
Some(std::cmp::Ordering::Greater)
) {
return HashMap::new();
}
let mut out: HashMap<u64, f32> = HashMap::with_capacity(groups.len());
for g in groups.values() {
let frac = FILE_COHERENCE_BOOST_FRAC * (g.sum / max_file_sum);
if frac > 0.0 {
out.insert(g.top_id, frac);
}
}
out
}
#[cfg(test)]
mod tests {
use super::*;
use crate::types::{Chunk, ChunkType, ScoredChunkId};
use std::path::PathBuf;
fn pb(s: &str) -> PathBuf {
PathBuf::from(s)
}
fn approx(a: f32, b: f32) -> bool {
(a - b).abs() < 1e-5
}
fn make_chunk(id: u64, file_path: &str) -> Chunk {
Chunk {
id,
file_path: pb(file_path),
start_line: 1,
end_line: 10,
content: String::new(),
chunk_type: ChunkType::TextWindow { window_index: 0 },
}
}
#[test]
fn split_camel_case_local_matches_code_tokenizer() {
let cases: &[(&str, &[&str])] = &[
("OAuth", &["OAuth"]),
("OAuthClient", &["OAuth", "Client"]),
("XMLParser", &["XML", "Parser"]),
("HTTPServer", &["HTTP", "Server"]),
("getUserById", &["get", "User", "By", "Id"]),
("MyHTTPHandler", &["My", "HTTP", "Handler"]),
("id3Tag", &["id3", "Tag"]),
("HTML5Parser", &["HTML5", "Parser"]),
];
for (input, expected) in cases {
let got = split_camel_case_local(input);
assert_eq!(
got, *expected,
"split_camel_case_local({input:?}) = {got:?}, expected {expected:?}"
);
}
}
#[test]
fn item6_pure_test_file_path_gets_strong_penalty() {
let p = pb("src/auth/auth_test.go");
let f = file_path_penalty(&p);
assert!(approx(f, STRONG_PENALTY), "got {f}");
}
#[test]
fn item6_python_pytest_test_prefix_strong_penalty() {
let p = pb("tests/unit/test_auth.py");
let f = file_path_penalty(&p);
assert!(approx(f, STRONG_PENALTY * STRONG_PENALTY), "got {f}");
let p2 = pb("src/auth/test_login.py");
let f2 = file_path_penalty(&p2);
assert!(approx(f2, STRONG_PENALTY), "got {f2}");
let p3 = pb("spec_helper.rb");
let f3 = file_path_penalty(&p3);
assert!(approx(f3, STRONG_PENALTY), "got {f3}");
}
#[test]
fn item6_pure_test_dir_strong_penalty() {
let p = pb("crates/foo/tests/integration_smoke.rs");
let f = file_path_penalty(&p);
assert!(approx(f, STRONG_PENALTY), "got {f}");
}
#[test]
fn item6_compat_plus_test_file_compounds() {
let p = pb("vendor/compat/legacy_auth_test.py");
let f = file_path_penalty(&p);
assert!(approx(f, STRONG_PENALTY * STRONG_PENALTY), "got {f}");
}
#[test]
fn item6_d_ts_only_mild_penalty() {
let p = pb("types/foo.d.ts");
let f = file_path_penalty(&p);
assert!(approx(f, MILD_PENALTY), "got {f}");
}
#[test]
fn item6_mod_rs_barrel_moderate_penalty() {
let p = pb("crates/semantex-core/src/search/mod.rs");
let f = file_path_penalty(&p);
assert!(approx(f, MODERATE_PENALTY), "got {f}");
}
#[test]
fn item6_index_ts_barrel_moderate_penalty() {
let p = pb("packages/foo/src/index.ts");
let f = file_path_penalty(&p);
assert!(approx(f, MODERATE_PENALTY), "got {f}");
}
#[test]
fn item6_no_match_returns_one() {
let p = pb("crates/semantex-core/src/search/hybrid.rs");
let f = file_path_penalty(&p);
assert!(approx(f, 1.0), "got {f}");
}
#[test]
fn item6_examples_dir_strong_penalty() {
let p = pb("examples/quickstart/main.rs");
let f = file_path_penalty(&p);
assert!(approx(f, STRONG_PENALTY), "got {f}");
}
#[test]
#[allow(non_snake_case)] fn item6_bare_tests_rs_at_root_does_NOT_match() {
let p = pb("crates/foo/src/tests.rs");
let f = file_path_penalty(&p);
assert!(
approx(f, 1.0),
"tests.rs at filename root must not be penalised, got {f}"
);
for name in &["crates/foo/src/spec.dart", "src/test.go", "lib/specs.rb"] {
let f = file_path_penalty(&pb(name));
assert!(
approx(f, 1.0),
"{name} (no left boundary) must not be penalised, got {f}"
);
}
}
#[test]
fn item6_underscore_separated_tests_still_matches() {
let p = pb("src/foo_tests.rs");
let f = file_path_penalty(&p);
assert!(
approx(f, STRONG_PENALTY),
"foo_tests.rs must still be penalised, got {f}"
);
let f2 = file_path_penalty(&pb("src/foo.test.ts"));
assert!(
approx(f2, STRONG_PENALTY),
"foo.test.ts must still be penalised, got {f2}"
);
let f3 = file_path_penalty(&pb("src/foo.spec.vue"));
assert!(
approx(f3, STRONG_PENALTY),
"foo.spec.vue must still be penalised, got {f3}"
);
}
#[test]
fn item6_penalty_disabled_when_query_mentions_test() {
assert!(!should_apply_path_penalty("auth middleware test"));
assert!(!should_apply_path_penalty("Test integration flow"));
assert!(!should_apply_path_penalty("how do I write unit tests"));
assert!(!should_apply_path_penalty("show me an example"));
assert!(!should_apply_path_penalty("benchmark harness"));
}
#[test]
fn item6_penalty_enabled_for_normal_queries() {
assert!(should_apply_path_penalty("authentication middleware"));
assert!(should_apply_path_penalty("code tokenizer"));
assert!(should_apply_path_penalty("how does graph propagation work"));
}
#[test]
fn item6_penalty_not_disabled_by_substring_inside_word() {
assert!(should_apply_path_penalty("attestation flow"));
assert!(should_apply_path_penalty("specification document"));
}
#[test]
fn item7_exact_match_returns_exact_frac() {
let tokens = query_tokens_for_path_signals("parse request");
let f = path_stem_boost_factor(&pb("src/parse_request.py"), &tokens);
assert!(
approx(f, STEM_EXACT_BOOST_FRAC),
"got {f}, tokens={tokens:?}"
);
}
#[test]
fn item7_prefix_match_returns_prefix_frac() {
let tokens = query_tokens_for_path_signals("toke");
let f = path_stem_boost_factor(&pb("src/code_tokenizer.rs"), &tokens);
assert!(
approx(f, STEM_PREFIX_BOOST_FRAC),
"got {f}, tokens={tokens:?}"
);
}
#[test]
fn item7_no_match_returns_zero() {
let tokens = query_tokens_for_path_signals("authentication middleware");
let f = path_stem_boost_factor(&pb("src/unrelated.rs"), &tokens);
assert!(approx(f, 0.0), "got {f}");
}
#[test]
fn item7_stopwords_filtered() {
let tokens = query_tokens_for_path_signals("how the why of for");
assert!(tokens.is_empty(), "got {tokens:?}");
}
#[test]
fn item7_short_prefix_below_min_len_does_not_match() {
let tokens = vec!["co".to_string()];
let f = path_stem_boost_factor(&pb("src/code_tokenizer.rs"), &tokens);
assert!(approx(f, 0.0), "got {f}");
}
#[test]
fn item7_acceptance_code_tokenizer_query() {
let tokens = query_tokens_for_path_signals("code tokenizer");
let f = path_stem_boost_factor(
&pb("crates/semantex-core/src/search/code_tokenizer.rs"),
&tokens,
);
assert!(
approx(f, STEM_EXACT_BOOST_FRAC),
"got {f}, tokens={tokens:?}"
);
}
#[test]
fn item7_camel_case_query_extracts_subtokens() {
let tokens = query_tokens_for_path_signals("getUserById");
let f = path_stem_boost_factor(&pb("src/user_by_id.rs"), &tokens);
assert!(
approx(f, STEM_EXACT_BOOST_FRAC),
"got {f}, tokens={tokens:?}"
);
}
#[test]
fn query_tokens_excludes_bigrams_for_camelcase_input() {
let tokens = query_tokens_for_path_signals("getUserById");
for expected in ["get", "user", "by", "id"] {
assert!(
tokens.iter().any(|t| t == expected),
"missing expected token {expected:?}; got {tokens:?}"
);
}
for tok in &tokens {
assert!(
!tok.contains('_'),
"bigram token leaked into query tokens: {tok:?} (full: {tokens:?})"
);
}
}
#[test]
fn query_tokens_excludes_bigrams_for_snake_case_input() {
let tokens = query_tokens_for_path_signals("get_user_by_id");
for tok in &tokens {
assert!(
!tok.contains('_'),
"bigram token leaked from snake_case input: {tok:?} (full: {tokens:?})"
);
}
}
#[test]
fn item8_symbol_match_returns_definition_frac() {
let tokens = query_tokens_for_path_signals("expand identifiers");
let f = definition_boost_factor("expand_identifiers", &tokens);
assert!(
approx(f, DEFINITION_BOOST_FRAC),
"got {f}, tokens={tokens:?}"
);
}
#[test]
fn item8_camel_case_symbol_match() {
let tokens = query_tokens_for_path_signals("retry handler");
let f = definition_boost_factor("RetryHandler", &tokens);
assert!(approx(f, DEFINITION_BOOST_FRAC), "got {f}");
}
#[test]
fn item8_no_match_returns_zero() {
let tokens = query_tokens_for_path_signals("authentication flow");
let f = definition_boost_factor("expand_identifiers", &tokens);
assert!(approx(f, 0.0), "got {f}");
}
#[test]
fn item8_empty_symbol_returns_zero() {
let tokens = vec!["foo".to_string()];
let f = definition_boost_factor("", &tokens);
assert!(approx(f, 0.0), "got {f}");
}
#[test]
fn item8_empty_tokens_returns_zero() {
let f = definition_boost_factor("RetryHandler", &[]);
assert!(approx(f, 0.0), "got {f}");
}
#[test]
fn item9_multi_chunk_file_boosts_only_top_chunk() {
let fused = vec![
ScoredChunkId::new(1, 1.0),
ScoredChunkId::new(2, 0.8),
ScoredChunkId::new(3, 0.6),
];
let mut chunk_map: HashMap<u64, Chunk> = HashMap::new();
chunk_map.insert(1, make_chunk(1, "src/foo.rs"));
chunk_map.insert(2, make_chunk(2, "src/foo.rs"));
chunk_map.insert(3, make_chunk(3, "src/foo.rs"));
let boosts = file_coherence_boosts(&fused, &chunk_map);
assert_eq!(boosts.len(), 1);
let frac = *boosts.get(&1).expect("top chunk should be present");
assert!(approx(frac, FILE_COHERENCE_BOOST_FRAC), "got {frac}");
assert!(!boosts.contains_key(&2));
assert!(!boosts.contains_key(&3));
}
#[test]
fn item9_single_chunk_files_get_no_boost() {
let fused = vec![
ScoredChunkId::new(1, 1.0),
ScoredChunkId::new(2, 0.9),
ScoredChunkId::new(3, 0.8),
];
let mut chunk_map: HashMap<u64, Chunk> = HashMap::new();
chunk_map.insert(1, make_chunk(1, "src/a.rs"));
chunk_map.insert(2, make_chunk(2, "src/b.rs"));
chunk_map.insert(3, make_chunk(3, "src/c.rs"));
let boosts = file_coherence_boosts(&fused, &chunk_map);
assert!(boosts.is_empty(), "got {boosts:?}");
}
#[test]
fn item9_multiple_files_normalized_by_max_sum() {
let fused = vec![
ScoredChunkId::new(1, 1.0),
ScoredChunkId::new(2, 0.8),
ScoredChunkId::new(3, 0.5),
ScoredChunkId::new(4, 0.4),
];
let mut chunk_map: HashMap<u64, Chunk> = HashMap::new();
chunk_map.insert(1, make_chunk(1, "src/foo.rs"));
chunk_map.insert(2, make_chunk(2, "src/foo.rs"));
chunk_map.insert(3, make_chunk(3, "src/bar.rs"));
chunk_map.insert(4, make_chunk(4, "src/bar.rs"));
let boosts = file_coherence_boosts(&fused, &chunk_map);
assert_eq!(boosts.len(), 2);
let foo = *boosts.get(&1).expect("foo top");
let bar = *boosts.get(&3).expect("bar top");
assert!(approx(foo, FILE_COHERENCE_BOOST_FRAC), "foo: {foo}");
assert!(approx(bar, FILE_COHERENCE_BOOST_FRAC * 0.5), "bar: {bar}");
}
}