Skip to main content

dhive_core/
search.rs

1//! BM25 搜索 — 分词、索引构建、搜索排序
2//!
3//! BM25 是信息检索领域的标准算法,比简单的关键词匹配更准确。
4//! 集成 effectivenessBoost — 频繁使用的高质量记忆搜索排名更高。
5
6use std::collections::{HashMap, HashSet};
7use unicode_segmentation::UnicodeSegmentation;
8use crate::memory::Mem;
9
10/// BM25 参数
11const K1: f64 = 1.5;   // term frequency saturation
12const B: f64 = 0.75;   // length normalization
13
14/// 分词 — Unicode 感知的分词,CJK 回退到字符 bigrams
15///
16/// 优先使用 Unicode 单词边界分割。如果 unicode_words 没有产出 (纯 CJK 文本),则回退到字符 bigrams。
17pub 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    // 回退: 如果 unicode_words 没有产出 (CJK 文本不被视为 alphabetic),用字符 bigrams
23    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            // 过滤单字符 (非 CJK) 和纯数字
37            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)   // CJK Unified
50        || ('\u{3400}'..='\u{4DBF}').contains(&c)  // CJK Ext-A
51        || ('\u{F900}'..='\u{FAFF}').contains(&c)  // CJK Compat
52        || ('\u{3040}'..='\u{309F}').contains(&c)  // Hiragana
53        || ('\u{30A0}'..='\u{30FF}').contains(&c)  // Katakana
54        || ('\u{AC00}'..='\u{D7AF}').contains(&c)  // Hangul
55}
56
57/// BM25 文档 — 构建索引时的中间表示
58#[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/// BM25 索引
67#[derive(Debug, Clone)]
68pub struct Bm25Index {
69    pub docs: Vec<Bm25Doc>,
70    /// 平均文档长度
71    pub avgdl: f64,
72    /// 文档频率: term → 包含该 term 的文档数
73    pub df: HashMap<String, u32>,
74}
75
76impl Bm25Index {
77    /// 从记忆列表构建 BM25 索引
78    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    /// 文档总数
116    pub fn doc_count(&self) -> u32 {
117        self.docs.len() as u32
118    }
119}
120
121/// BM25 单文档得分
122pub 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        // IDF
135        let idf = ((n - df_t + 0.5) / (df_t + 0.5)).ln_1p();
136        // TF saturation
137        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
145/// 效果追踪加权 — 频繁使用的高质量记忆搜索排名更高
146pub 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
152/// BM25 搜索 — 返回评分排序的结果
153pub 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/// 搜索结果
175#[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        // CJK text should produce tokens (either via unicode_words or bigram fallback)
222        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); // (1+1)*(1+1) = 4.0
253    }
254}