//! 记忆检索器 —— 基于 TextOverlap (Dice 系数) 的单通道关键词检索。 use std::collections::HashSet; use crate::memory::error::MemoryError; use crate::memory::knowledge::KnowledgeStore; use crate::memory::types::KnowledgePage; /// 检索器配置。 #[derive(Debug, Clone)] pub struct RetrieverConfig { /// 最大返回条数(默认 20)。 pub max_results: usize, /// 最低分数阈值 [0.0, 1.0](默认 0.1)。 pub min_score: f32, } impl Default for RetrieverConfig { fn default() -> Self { Self { max_results: 20, min_score: 0.1, } } } /// 单条带评分的检索结果。 #[derive(Debug, Clone)] pub struct ScoredItem { pub page: KnowledgePage, /// TextOverlap 评分 [0.0, 1.0] pub score: f32, } /// 检索结果。 #[derive(Debug, Clone)] pub struct RetrievalResult { pub items: Vec, pub query: String, } /// 记忆检索器 —— 在 `KnowledgeStore` 中做关键词检索并按 TextOverlap 评分。 pub struct MemoryRetriever { knowledge_store: KnowledgeStore, config: RetrieverConfig, /// 停用词表(用于关键词提取)。 stop_words: HashSet, } impl MemoryRetriever { /// 创建一个新的 MemoryRetriever。 pub fn new(knowledge_store: KnowledgeStore, config: RetrieverConfig) -> Self { Self { knowledge_store, config, stop_words: default_stop_words(), } } /// 替换停用词表。 pub fn with_stop_words(mut self, stop_words: HashSet) -> Self { self.stop_words = stop_words; self } /// 检索相关知识页面。 pub async fn retrieve(&self, query: &str) -> Result { if query.is_empty() { return Ok(RetrievalResult { items: Vec::new(), query: query.to_string(), }); } // 1. 关键词提取 let keywords = extract_keywords(query, &self.stop_words); // 2. 用关键词在 KnowledgeStore 中搜索 let mut pages = Vec::new(); for keyword in &keywords { let found = self.knowledge_store.search(keyword).await?; for page in found { if !pages.iter().any(|p: &KnowledgePage| p.id == page.id) { pages.push(page); } } } // 3. TextOverlap 评分 let mut items: Vec = pages .into_iter() .map(|page| { let score = text_overlap_score(query, &page); ScoredItem { page, score } }) .collect(); // 4. 过滤 → 排序 → 截取 items.retain(|i| i.score >= self.config.min_score); items.sort_by(|a, b| { b.score .partial_cmp(&a.score) .unwrap_or(std::cmp::Ordering::Equal) }); items.truncate(self.config.max_results); Ok(RetrievalResult { items, query: query.to_string(), }) } } /// 从 query 中提取关键词:按非字母数字字符分割 → 转小写 → 过滤单字符和停用词。 fn extract_keywords(query: &str, stop_words: &HashSet) -> Vec { query .split(|c: char| !c.is_alphanumeric()) .filter_map(|s| { let lower = s.to_lowercase(); if lower.is_empty() || lower.chars().count() < 2 || stop_words.contains(&lower) { None } else { Some(lower) } }) .collect() } /// TextOverlap 评分(基于字符 bigram 的 Dice 系数 + 多字段加权)。 /// /// 字段权重:title 0.5 + summary 0.3 + content 0.2 /// 中文场景按字符级 bigram 处理,不依赖分词器。 pub fn text_overlap_score(query: &str, page: &KnowledgePage) -> f32 { let title = text_overlap_dice(query, &page.title); let summary = text_overlap_dice(query, &page.summary); let content = text_overlap_dice(query, &page.content); title * 0.5 + summary * 0.3 + content * 0.2 } /// Dice 系数(基于字符 bigram)。 fn text_overlap_dice(query: &str, text: &str) -> f32 { let q_bigrams = char_bigrams(query); let t_bigrams = char_bigrams(text); if q_bigrams.is_empty() || t_bigrams.is_empty() { return 0.0; } let q_set: HashSet<&String> = q_bigrams.iter().collect(); let t_set: HashSet<&String> = t_bigrams.iter().collect(); let intersect = q_set.intersection(&t_set).count(); let denom = q_set.len() + t_set.len(); if denom == 0 { 0.0 } else { (2.0 * intersect as f32) / denom as f32 } } /// 提取字符 bigrams(用于字符级 Dice 系数)。 fn char_bigrams(s: &str) -> Vec { let chars: Vec = s.chars().collect(); chars.windows(2).map(|w| w.iter().collect()).collect() } fn default_stop_words() -> HashSet { [ "the", "a", "an", "is", "are", "was", "were", "be", "been", "being", "have", "has", "had", "do", "does", "did", "will", "would", "should", "could", "may", "might", "shall", "can", "this", "that", "these", "those", "it", "its", "they", "them", "their", "what", "which", "who", "whom", "how", "when", "where", "and", "or", "but", "not", "no", "nor", "so", "if", "then", "else", "with", "without", "for", "to", "from", "in", "on", "at", "by", "of", "as", "into", "through", "during", "before", "after", "above", "below", ] .iter() .map(|s| s.to_string()) .collect() } #[cfg(test)] mod tests { use super::*; use crate::memory::knowledge::KnowledgeStore; use crate::memory::{InMemoryStore, MemoryStore}; use std::sync::Arc; use time::OffsetDateTime; fn make_page(id: &str, title: &str, summary: &str, content: &str) -> KnowledgePage { let now = OffsetDateTime::now_utc(); KnowledgePage { id: id.to_string(), title: title.to_string(), summary: summary.to_string(), content: content.to_string(), tags: Vec::new(), references: Vec::new(), created_at: now, updated_at: now, } } #[tokio::test] async fn retrieve_empty_query() { let store = Arc::new(InMemoryStore::new()) as Arc; let ks = KnowledgeStore::new(store); let retriever = MemoryRetriever::new(ks, RetrieverConfig::default()); let result = retriever.retrieve("").await.unwrap(); assert!(result.items.is_empty()); } #[tokio::test] async fn retrieve_finds_relevant_page() { let store = Arc::new(InMemoryStore::new()) as Arc; let ks = KnowledgeStore::new(store); ks.add_page(make_page( "p1", "LangGraph StateGraph", "state management", "LangGraph uses StateGraph for state machines", )) .await .unwrap(); ks.add_page(make_page("p2", "Other", "unrelated", "nothing matching")) .await .unwrap(); let retriever = MemoryRetriever::new(ks, RetrieverConfig::default()); let result = retriever.retrieve("LangGraph state").await.unwrap(); assert!(!result.items.is_empty()); assert_eq!(result.items[0].page.id, "p1"); assert!(result.items[0].score > 0.0); } #[tokio::test] async fn retrieve_respects_min_score() { let store = Arc::new(InMemoryStore::new()) as Arc; let ks = KnowledgeStore::new(store); ks.add_page(make_page("p1", "X", "Y", "Z")).await.unwrap(); let config = RetrieverConfig { max_results: 10, min_score: 0.99, }; let retriever = MemoryRetriever::new(ks, config); let result = retriever .retrieve("totally unrelated content") .await .unwrap(); assert!(result.items.is_empty()); } #[tokio::test] async fn retrieve_respects_max_results() { let store = Arc::new(InMemoryStore::new()) as Arc; let ks = KnowledgeStore::new(store); for i in 0..5 { ks.add_page(make_page( &format!("p{i}"), "LangGraph", "framework", "agent runtime", )) .await .unwrap(); } let config = RetrieverConfig { max_results: 2, min_score: 0.0, }; let retriever = MemoryRetriever::new(ks, config); let result = retriever.retrieve("LangGraph").await.unwrap(); assert_eq!(result.items.len(), 2); } #[test] fn text_overlap_dice_zero_on_empty() { assert_eq!(text_overlap_dice("hello", ""), 0.0); assert_eq!(text_overlap_dice("", "hello"), 0.0); } #[test] fn text_overlap_dice_identical() { let s = "hello world"; assert!((text_overlap_dice(s, s) - 1.0).abs() < 0.001); } #[test] fn extract_keywords_filters_stop_words() { let stop = default_stop_words(); let kws = extract_keywords("the quick brown fox is fast", &stop); assert!(!kws.contains(&"the".to_string())); assert!(!kws.contains(&"is".to_string())); assert!(kws.contains(&"quick".to_string())); assert!(kws.contains(&"brown".to_string())); } #[test] fn extract_keywords_filters_single_chars() { let stop = default_stop_words(); let kws = extract_keywords("a b c dog", &stop); assert!(!kws.contains(&"a".to_string())); assert!(!kws.contains(&"b".to_string())); } }