use std::collections::{HashMap, HashSet};
use unicode_segmentation::UnicodeSegmentation;
use crate::memory::Mem;
const K1: f64 = 1.5; const B: f64 = 0.75;
pub fn tokenize(text: &str) -> Vec<String> {
let mut tokens: Vec<String> = text.unicode_words()
.map(|w| w.to_lowercase())
.collect();
if tokens.is_empty() {
let chars: Vec<char> = text.chars().filter(|c| !c.is_whitespace()).collect();
if chars.len() == 1 {
tokens.push(chars[0].to_string().to_lowercase());
} else {
for w in chars.windows(2) {
tokens.push(w.iter().collect::<String>().to_lowercase());
}
}
}
tokens.into_iter()
.filter(|w| {
if w.len() == 1 && !is_cjk(w.chars().next().unwrap()) {
return false;
}
if w.chars().all(|c| c.is_ascii_digit()) {
return false;
}
true
})
.collect()
}
fn is_cjk(c: char) -> bool {
('\u{4E00}'..='\u{9FFF}').contains(&c) || ('\u{3400}'..='\u{4DBF}').contains(&c) || ('\u{F900}'..='\u{FAFF}').contains(&c) || ('\u{3040}'..='\u{309F}').contains(&c) || ('\u{30A0}'..='\u{30FF}').contains(&c) || ('\u{AC00}'..='\u{D7AF}').contains(&c) }
#[derive(Debug, Clone)]
pub struct Bm25Doc {
pub id: String,
pub content: String,
pub term_freqs: HashMap<String, u32>,
pub doc_len: usize,
}
#[derive(Debug, Clone)]
pub struct Bm25Index {
pub docs: Vec<Bm25Doc>,
pub avgdl: f64,
pub df: HashMap<String, u32>,
}
impl Bm25Index {
pub fn build(memories: &[&Mem]) -> Self {
let docs: Vec<Bm25Doc> = memories
.iter()
.map(|m| {
let tokens = tokenize(&m.content);
let doc_len = tokens.len();
let mut term_freqs: HashMap<String, u32> = HashMap::new();
for t in &tokens {
*term_freqs.entry(t.clone()).or_insert(0) += 1;
}
Bm25Doc {
id: m.id.clone(),
content: m.content.clone(),
term_freqs,
doc_len,
}
})
.collect();
let n = docs.len() as f64;
let avgdl = if n > 0.0 {
docs.iter().map(|d| d.doc_len).sum::<usize>() as f64 / n
} else {
0.0
};
let mut df: HashMap<String, u32> = HashMap::new();
for doc in &docs {
let seen: HashSet<&String> = doc.term_freqs.keys().collect();
for term in seen {
*df.entry(term.clone()).or_insert(0) += 1;
}
}
Bm25Index { docs, avgdl, df }
}
pub fn doc_count(&self) -> u32 {
self.docs.len() as u32
}
}
pub fn bm25_score(query: &str, doc: &Bm25Doc, avgdl: f64, df: &HashMap<String, u32>, doc_count: u32) -> f64 {
let query_terms = tokenize(query);
let n = doc_count as f64;
let mut score = 0.0;
for term in &query_terms {
let tf = *doc.term_freqs.get(term).unwrap_or(&0) as f64;
if tf == 0.0 {
continue;
}
let df_t = *df.get(term).unwrap_or(&0) as f64;
let idf = ((n - df_t + 0.5) / (df_t + 0.5)).ln_1p();
let tf_part = (tf * (K1 + 1.0)) / (tf + K1 * (1.0 - B + B * doc.doc_len as f64 / avgdl));
score += idf * tf_part;
}
score
}
pub fn effectiveness_boost(mem: &Mem) -> f64 {
let loaded = mem.loaded_count as f64;
let referenced = mem.referenced_count as f64;
(1.0 + loaded * 0.01) * (1.0 + referenced * 0.02)
}
pub fn bm25_search(query: &str, index: &Bm25Index, memories: &[&Mem]) -> Vec<SearchResult> {
let doc_count = index.doc_count();
let mut results: Vec<SearchResult> = index
.docs
.iter()
.enumerate()
.map(|(i, doc)| {
let bm = bm25_score(query, doc, index.avgdl, &index.df, doc_count);
let eff = memories.get(i).map(|m| effectiveness_boost(m)).unwrap_or(1.0);
SearchResult {
id: doc.id.clone(),
score: bm * eff,
}
})
.filter(|r| r.score > 0.0)
.collect();
results.sort_by(|a, b| b.score.partial_cmp(&a.score).unwrap_or(std::cmp::Ordering::Equal));
results
}
#[derive(Debug, Clone)]
pub struct SearchResult {
pub id: String,
pub score: f64,
}
#[cfg(test)]
mod tests {
use super::*;
fn make_mem(id: &str, content: &str, loaded: u32, referenced: u32) -> Mem {
Mem {
id: id.into(),
level: crate::memory::MemLevel::Fact,
zone: crate::memory::MemoryZone::General,
content: content.into(),
weight: 0.5,
tags: vec![],
source: None, metadata: None,
last_matched: None, match_count: 0,
loaded_count: loaded, referenced_count: referenced,
last_loaded: None, last_referenced: None,
supersedes: None, superseded_at: None, superseded_by: None,
trust: None,
created_at: "2026-01-01T00:00:00Z".into(),
updated_at: "2026-01-01T00:00:00Z".into(),
}
}
#[test]
fn test_tokenize_basic() {
let tokens = tokenize("Hello World Rust");
assert!(tokens.contains(&"hello".to_string()));
assert!(tokens.contains(&"world".to_string()));
assert!(tokens.contains(&"rust".to_string()));
}
#[test]
fn test_tokenize_filters_numbers() {
let tokens = tokenize("test 1234");
assert!(!tokens.contains(&"1234".to_string()));
}
#[test]
fn test_tokenize_chinese() {
let tokens = tokenize("基石原则 零静默修改");
assert!(!tokens.is_empty(), "expected non-empty tokens for Chinese text, got: {:?}", tokens);
}
#[test]
fn test_bm25_build_and_search() {
let m1 = make_mem("m1", "基石原则:最少依赖 新增依赖前思考是否真的需要", 10, 5);
let m2 = make_mem("m2", "测试驱动的开发方法提高代码质量", 2, 1);
let m3 = make_mem("m3", "使用 pnpm 作为包管理器", 3, 2);
let mems = vec![&m1, &m2, &m3];
let index = Bm25Index::build(&mems);
let results = bm25_search("基石", &index, &mems);
assert!(!results.is_empty());
assert_eq!(results[0].id, "m1");
}
#[test]
fn test_bm25_no_match_returns_empty() {
let m1 = make_mem("m1", "基石原则", 1, 1);
let mems = vec![&m1];
let index = Bm25Index::build(&mems);
let results = bm25_search("xyz不存在的关键词abc", &index, &mems);
assert!(results.is_empty());
}
#[test]
fn test_effectiveness_boost() {
let mem = make_mem("m1", "test", 100, 50);
let boost = effectiveness_boost(&mem);
assert!(boost > 2.0, "boost = {}", boost); }
}