use super::*;
#[test]
#[serial_test::serial]
fn bm25_scores_relevant_doc_higher() {
let mut idx = BM25Index::new();
idx.add_document(0, "authentication login password secure");
idx.add_document(1, "rendering ui components svelte");
let s0 = idx.score("authentication", 0);
let s1 = idx.score("authentication", 1);
assert!(s0 > s1, "relevant doc should score higher: {s0} vs {s1}");
}
#[test]
fn tokenize_splits_code() {
let tokens = tokenize("fn search_hybrid(query: &str) -> Vec<Hit>");
assert!(tokens.contains(&"search".to_string()));
assert!(tokens.contains(&"hybrid".to_string()));
assert!(tokens.contains(&"query".to_string()));
}
#[test]
fn tokenize_camel_case_pascal() {
let tokens = tokenize("CodeIndexer");
assert!(tokens.contains(&"code".to_string()), "got {tokens:?}");
assert!(tokens.contains(&"indexer".to_string()), "got {tokens:?}");
assert!(
tokens.contains(&"codeindexer".to_string()),
"got {tokens:?}"
);
}
#[test]
fn tokenize_pascal_two_words() {
let tokens = tokenize("UsearchStore");
assert!(tokens.contains(&"usearch".to_string()), "got {tokens:?}");
assert!(tokens.contains(&"store".to_string()), "got {tokens:?}");
}
#[test]
fn tokenize_snake_case() {
let tokens = tokenize("use_kg_first");
assert!(tokens.contains(&"use".to_string()), "got {tokens:?}");
assert!(tokens.contains(&"kg".to_string()), "got {tokens:?}");
assert!(tokens.contains(&"first".to_string()), "got {tokens:?}");
}
#[test]
fn tokenize_alpha_digit_split() {
let tokens = tokenize("HTTP2Client");
assert!(tokens.contains(&"http".to_string()), "got {tokens:?}");
assert!(tokens.contains(&"2".to_string()), "got {tokens:?}");
assert!(tokens.contains(&"client".to_string()), "got {tokens:?}");
}
#[test]
fn tokenize_acronym_then_word() {
let tokens = tokenize("HTTPSClient");
assert!(tokens.contains(&"https".to_string()), "got {tokens:?}");
assert!(tokens.contains(&"client".to_string()), "got {tokens:?}");
}
#[test]
#[serial_test::serial]
fn bm25_incremental_upsert_and_remove() {
let mut idx = BM25Index::new();
idx.upsert_document("a", "authentication login password");
idx.upsert_document("b", "rendering ui components svelte");
idx.upsert_document("c", "database connection pool postgres");
assert_eq!(idx.len(), 3);
let hits = idx.score_query_all("authentication", 10);
assert!(hits.iter().any(|(id, _)| id == "a"));
assert!(!hits.iter().any(|(id, _)| id == "b"));
idx.remove_document("a");
assert_eq!(idx.len(), 2);
let hits_after = idx.score_query_all("authentication", 10);
assert!(!hits_after.iter().any(|(id, _)| id == "a"));
let svelte_hits = idx.score_query_all("svelte", 10);
assert!(svelte_hits.iter().any(|(id, _)| id == "b"));
}
#[test]
#[serial_test::serial]
fn bm25_upsert_replaces_existing_doc() {
let mut idx = BM25Index::new();
idx.upsert_document("a", "alpha beta gamma");
idx.upsert_document("a", "delta epsilon");
assert_eq!(idx.len(), 1);
assert!(idx.score_query_all("alpha", 10).is_empty());
assert!(!idx.score_query_all("delta", 10).is_empty());
}
#[test]
#[serial_test::serial]
fn score_query_all_returns_sorted_unique_results() {
let mut idx = BM25Index::new();
idx.upsert_document("a", "search rust async tokio");
idx.upsert_document("b", "search rust");
idx.upsert_document("c", "unrelated content");
let hits = idx.score_query_all("rust async", 10);
for w in hits.windows(2) {
assert!(w[0].1 >= w[1].1, "results must be sorted desc: {hits:?}");
}
let mut ids: Vec<&str> = hits.iter().map(|(id, _)| id.as_str()).collect();
ids.sort();
let unique = ids.len();
ids.dedup();
assert_eq!(unique, ids.len());
}
#[test]
#[serial_test::serial]
fn score_query_all_with_filter_recovers_match_beyond_top_k() {
let mut idx = BM25Index::new();
for i in 0..5 {
idx.upsert_document(&format!("filler{i}"), "rust rust rust async");
}
idx.upsert_document("target", "rust async");
let unfiltered = idx.score_query_all("rust async", 3);
assert_eq!(unfiltered.len(), 3);
assert!(
unfiltered.iter().all(|(id, _)| id != "target"),
"precondition failed: target must rank below top_k=3 unfiltered; \
got {unfiltered:?}"
);
let filtered = idx.score_query_all_with_filter("rust async", 3, &|id: &str| id == "target");
assert_eq!(
filtered.len(),
1,
"the filter admits only \"target\" — it must be the one result: {filtered:?}"
);
assert_eq!(filtered[0].0, "target");
}
#[test]
fn tokenize_dedups_and_sorts() {
let tokens = tokenize("foo foo bar");
let foos: Vec<&String> = tokens.iter().filter(|t| t.as_str() == "foo").collect();
assert_eq!(foos.len(), 1, "duplicates must collapse: {tokens:?}");
let mut sorted = tokens.clone();
sorted.sort();
assert_eq!(tokens, sorted, "tokens must be sorted: {tokens:?}");
}
#[test]
#[serial_test::serial]
fn bm25_corpus_cap_env_override() {
let prev = std::env::var("TRUSTY_BM25_CORPUS_CAP").ok();
unsafe {
std::env::set_var("TRUSTY_BM25_CORPUS_CAP", "0");
}
assert_eq!(
bm25_corpus_cap(),
DEFAULT_BM25_CORPUS_CAP,
"zero must fall back to default"
);
unsafe {
std::env::set_var("TRUSTY_BM25_CORPUS_CAP", "123");
}
assert_eq!(bm25_corpus_cap(), 123, "positive value must be honoured");
match prev {
Some(v) => unsafe { std::env::set_var("TRUSTY_BM25_CORPUS_CAP", v) },
None => unsafe { std::env::remove_var("TRUSTY_BM25_CORPUS_CAP") },
}
}
#[test]
#[serial_test::serial]
fn upsert_document_reporting_returns_false_when_capped() {
let prev = std::env::var("TRUSTY_BM25_CORPUS_CAP").ok();
unsafe { std::env::set_var("TRUSTY_BM25_CORPUS_CAP", "2") };
let mut idx = BM25Index::new();
assert!(idx.upsert_document_reporting("a", "alpha one"));
assert!(idx.upsert_document_reporting("b", "beta two"));
assert!(
!idx.upsert_document_reporting("c", "gamma three"),
"a brand-new doc_id past the cap must be reported as dropped"
);
assert_eq!(idx.len(), 2, "the dropped doc must not grow the corpus");
assert!(
idx.upsert_document_reporting("a", "alpha one updated"),
"updates to an existing doc_id must never be dropped by the cap"
);
assert_eq!(idx.len(), 2);
match prev {
Some(v) => unsafe { std::env::set_var("TRUSTY_BM25_CORPUS_CAP", v) },
None => unsafe { std::env::remove_var("TRUSTY_BM25_CORPUS_CAP") },
}
}
#[test]
#[serial_test::serial]
fn upsert_document_delegates_to_reporting_and_still_logs_once() {
let prev = std::env::var("TRUSTY_BM25_CORPUS_CAP").ok();
unsafe { std::env::set_var("TRUSTY_BM25_CORPUS_CAP", "1") };
let mut idx = BM25Index::new();
idx.upsert_document("only", "the one doc that fits");
assert_eq!(idx.len(), 1);
idx.upsert_document("dropped", "this one does not fit");
assert_eq!(
idx.len(),
1,
"a brand-new doc past the cap must still be silently dropped via upsert_document"
);
match prev {
Some(v) => unsafe { std::env::set_var("TRUSTY_BM25_CORPUS_CAP", v) },
None => unsafe { std::env::remove_var("TRUSTY_BM25_CORPUS_CAP") },
}
}