dhive-core 0.1.0

D-HIVE Trust Protocol — Rust core: BM25 search, LCS similarity, canonicalHash, pack format, curation scoring
Documentation
//! BM25 搜索 — 分词、索引构建、搜索排序
//!
//! BM25 是信息检索领域的标准算法,比简单的关键词匹配更准确。
//! 集成 effectivenessBoost — 频繁使用的高质量记忆搜索排名更高。

use std::collections::{HashMap, HashSet};
use unicode_segmentation::UnicodeSegmentation;
use crate::memory::Mem;

/// BM25 参数
const K1: f64 = 1.5;   // term frequency saturation
const B: f64 = 0.75;   // length normalization

/// 分词 — Unicode 感知的分词,CJK 回退到字符 bigrams
///
/// 优先使用 Unicode 单词边界分割。如果 unicode_words 没有产出 (纯 CJK 文本),则回退到字符 bigrams。
pub fn tokenize(text: &str) -> Vec<String> {
    let mut tokens: Vec<String> = text.unicode_words()
        .map(|w| w.to_lowercase())
        .collect();

    // 回退: 如果 unicode_words 没有产出 (CJK 文本不被视为 alphabetic),用字符 bigrams
    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| {
            // 过滤单字符 (非 CJK) 和纯数字
            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)   // CJK Unified
        || ('\u{3400}'..='\u{4DBF}').contains(&c)  // CJK Ext-A
        || ('\u{F900}'..='\u{FAFF}').contains(&c)  // CJK Compat
        || ('\u{3040}'..='\u{309F}').contains(&c)  // Hiragana
        || ('\u{30A0}'..='\u{30FF}').contains(&c)  // Katakana
        || ('\u{AC00}'..='\u{D7AF}').contains(&c)  // Hangul
}

/// BM25 文档 — 构建索引时的中间表示
#[derive(Debug, Clone)]
pub struct Bm25Doc {
    pub id: String,
    pub content: String,
    pub term_freqs: HashMap<String, u32>,
    pub doc_len: usize,
}

/// BM25 索引
#[derive(Debug, Clone)]
pub struct Bm25Index {
    pub docs: Vec<Bm25Doc>,
    /// 平均文档长度
    pub avgdl: f64,
    /// 文档频率: term → 包含该 term 的文档数
    pub df: HashMap<String, u32>,
}

impl Bm25Index {
    /// 从记忆列表构建 BM25 索引
    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
    }
}

/// BM25 单文档得分
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;
        // IDF
        let idf = ((n - df_t + 0.5) / (df_t + 0.5)).ln_1p();
        // TF saturation
        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)
}

/// BM25 搜索 — 返回评分排序的结果
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("基石原则 零静默修改");
        // CJK text should produce tokens (either via unicode_words or bigram fallback)
        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); // (1+1)*(1+1) = 4.0
    }
}