1use std::collections::{HashMap, HashSet};
7use unicode_segmentation::UnicodeSegmentation;
8use crate::memory::Mem;
9
10const K1: f64 = 1.5; const B: f64 = 0.75; pub fn tokenize(text: &str) -> Vec<String> {
18 let mut tokens: Vec<String> = text.unicode_words()
19 .map(|w| w.to_lowercase())
20 .collect();
21
22 if tokens.is_empty() {
24 let chars: Vec<char> = text.chars().filter(|c| !c.is_whitespace()).collect();
25 if chars.len() == 1 {
26 tokens.push(chars[0].to_string().to_lowercase());
27 } else {
28 for w in chars.windows(2) {
29 tokens.push(w.iter().collect::<String>().to_lowercase());
30 }
31 }
32 }
33
34 tokens.into_iter()
35 .filter(|w| {
36 if w.len() == 1 && !is_cjk(w.chars().next().unwrap()) {
38 return false;
39 }
40 if w.chars().all(|c| c.is_ascii_digit()) {
41 return false;
42 }
43 true
44 })
45 .collect()
46}
47
48fn is_cjk(c: char) -> bool {
49 ('\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) }
56
57#[derive(Debug, Clone)]
59pub struct Bm25Doc {
60 pub id: String,
61 pub content: String,
62 pub term_freqs: HashMap<String, u32>,
63 pub doc_len: usize,
64}
65
66#[derive(Debug, Clone)]
68pub struct Bm25Index {
69 pub docs: Vec<Bm25Doc>,
70 pub avgdl: f64,
72 pub df: HashMap<String, u32>,
74}
75
76impl Bm25Index {
77 pub fn build(memories: &[&Mem]) -> Self {
79 let docs: Vec<Bm25Doc> = memories
80 .iter()
81 .map(|m| {
82 let tokens = tokenize(&m.content);
83 let doc_len = tokens.len();
84 let mut term_freqs: HashMap<String, u32> = HashMap::new();
85 for t in &tokens {
86 *term_freqs.entry(t.clone()).or_insert(0) += 1;
87 }
88 Bm25Doc {
89 id: m.id.clone(),
90 content: m.content.clone(),
91 term_freqs,
92 doc_len,
93 }
94 })
95 .collect();
96
97 let n = docs.len() as f64;
98 let avgdl = if n > 0.0 {
99 docs.iter().map(|d| d.doc_len).sum::<usize>() as f64 / n
100 } else {
101 0.0
102 };
103
104 let mut df: HashMap<String, u32> = HashMap::new();
105 for doc in &docs {
106 let seen: HashSet<&String> = doc.term_freqs.keys().collect();
107 for term in seen {
108 *df.entry(term.clone()).or_insert(0) += 1;
109 }
110 }
111
112 Bm25Index { docs, avgdl, df }
113 }
114
115 pub fn doc_count(&self) -> u32 {
117 self.docs.len() as u32
118 }
119}
120
121pub fn bm25_score(query: &str, doc: &Bm25Doc, avgdl: f64, df: &HashMap<String, u32>, doc_count: u32) -> f64 {
123 let query_terms = tokenize(query);
124 let n = doc_count as f64;
125 let mut score = 0.0;
126
127 for term in &query_terms {
128 let tf = *doc.term_freqs.get(term).unwrap_or(&0) as f64;
129 if tf == 0.0 {
130 continue;
131 }
132
133 let df_t = *df.get(term).unwrap_or(&0) as f64;
134 let idf = ((n - df_t + 0.5) / (df_t + 0.5)).ln_1p();
136 let tf_part = (tf * (K1 + 1.0)) / (tf + K1 * (1.0 - B + B * doc.doc_len as f64 / avgdl));
138
139 score += idf * tf_part;
140 }
141
142 score
143}
144
145pub fn effectiveness_boost(mem: &Mem) -> f64 {
147 let loaded = mem.loaded_count as f64;
148 let referenced = mem.referenced_count as f64;
149 (1.0 + loaded * 0.01) * (1.0 + referenced * 0.02)
150}
151
152pub fn bm25_search(query: &str, index: &Bm25Index, memories: &[&Mem]) -> Vec<SearchResult> {
154 let doc_count = index.doc_count();
155 let mut results: Vec<SearchResult> = index
156 .docs
157 .iter()
158 .enumerate()
159 .map(|(i, doc)| {
160 let bm = bm25_score(query, doc, index.avgdl, &index.df, doc_count);
161 let eff = memories.get(i).map(|m| effectiveness_boost(m)).unwrap_or(1.0);
162 SearchResult {
163 id: doc.id.clone(),
164 score: bm * eff,
165 }
166 })
167 .filter(|r| r.score > 0.0)
168 .collect();
169
170 results.sort_by(|a, b| b.score.partial_cmp(&a.score).unwrap_or(std::cmp::Ordering::Equal));
171 results
172}
173
174#[derive(Debug, Clone)]
176pub struct SearchResult {
177 pub id: String,
178 pub score: f64,
179}
180
181#[cfg(test)]
182mod tests {
183 use super::*;
184
185 fn make_mem(id: &str, content: &str, loaded: u32, referenced: u32) -> Mem {
186 Mem {
187 id: id.into(),
188 level: crate::memory::MemLevel::Fact,
189 zone: crate::memory::MemoryZone::General,
190 content: content.into(),
191 weight: 0.5,
192 tags: vec![],
193 source: None, metadata: None,
194 last_matched: None, match_count: 0,
195 loaded_count: loaded, referenced_count: referenced,
196 last_loaded: None, last_referenced: None,
197 supersedes: None, superseded_at: None, superseded_by: None,
198 trust: None,
199 created_at: "2026-01-01T00:00:00Z".into(),
200 updated_at: "2026-01-01T00:00:00Z".into(),
201 }
202 }
203
204 #[test]
205 fn test_tokenize_basic() {
206 let tokens = tokenize("Hello World Rust");
207 assert!(tokens.contains(&"hello".to_string()));
208 assert!(tokens.contains(&"world".to_string()));
209 assert!(tokens.contains(&"rust".to_string()));
210 }
211
212 #[test]
213 fn test_tokenize_filters_numbers() {
214 let tokens = tokenize("test 1234");
215 assert!(!tokens.contains(&"1234".to_string()));
216 }
217
218 #[test]
219 fn test_tokenize_chinese() {
220 let tokens = tokenize("基石原则 零静默修改");
221 assert!(!tokens.is_empty(), "expected non-empty tokens for Chinese text, got: {:?}", tokens);
223 }
224
225 #[test]
226 fn test_bm25_build_and_search() {
227 let m1 = make_mem("m1", "基石原则:最少依赖 新增依赖前思考是否真的需要", 10, 5);
228 let m2 = make_mem("m2", "测试驱动的开发方法提高代码质量", 2, 1);
229 let m3 = make_mem("m3", "使用 pnpm 作为包管理器", 3, 2);
230
231 let mems = vec![&m1, &m2, &m3];
232 let index = Bm25Index::build(&mems);
233
234 let results = bm25_search("基石", &index, &mems);
235 assert!(!results.is_empty());
236 assert_eq!(results[0].id, "m1");
237 }
238
239 #[test]
240 fn test_bm25_no_match_returns_empty() {
241 let m1 = make_mem("m1", "基石原则", 1, 1);
242 let mems = vec![&m1];
243 let index = Bm25Index::build(&mems);
244 let results = bm25_search("xyz不存在的关键词abc", &index, &mems);
245 assert!(results.is_empty());
246 }
247
248 #[test]
249 fn test_effectiveness_boost() {
250 let mem = make_mem("m1", "test", 100, 50);
251 let boost = effectiveness_boost(&mem);
252 assert!(boost > 2.0, "boost = {}", boost); }
254}