//! 记忆检索器 -- 双通道关键词检索(KnowledgeStore + KnowledgeGraph)。 //! //! 单通道模式:仅 KnowledgeStore,基于 TextOverlap (Dice 系数) 评分。 //! 双通道模式:并行检索 KnowledgeStore + KnowledgeGraph,结果合并为统一列表。 use std::collections::HashSet; use std::sync::Arc; use crate::memory::error::MemoryError; use crate::memory::graph::{GraphEntity, KnowledgeGraph, ScoredEntity}; 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, /// 图遍历的默认深度(默认 2)。 pub graph_depth: usize, } impl Default for RetrieverConfig { fn default() -> Self { Self { max_results: 20, min_score: 0.1, graph_depth: 2, } } } /// 检索策略 -- 控制双通道分流。 #[derive(Debug, Clone, Default, PartialEq, Eq)] pub enum RetrievalStrategy { /// 并行 KnowledgeStore + KnowledgeGraph,合并排序(默认)。 #[default] Hybrid, /// 仅 KnowledgeStore。 KnowledgeOnly, /// 仅 KnowledgeGraph。 GraphOnly, } /// 统一检索条目 -- enum 变体区分类别,两通道分数均在 \[0,1\] 区间。 /// /// 注意:两个通道的分数维度不同(TextOverlap vs 图距离), /// 合并排序仅用于统一返回,不代表跨通道可比性。 #[derive(Debug, Clone)] pub enum RetrievalItem { /// 知识页面(来自 KnowledgeStore)。 KnowledgePage { page: KnowledgePage, /// TextOverlap 评分 [0.0, 1.0]。 score: f32, }, /// 图谱实体(来自 KnowledgeGraph)。 GraphEntity { entity: GraphEntity, /// 图距离评分 [0.0, 1.0]。 score: f32, /// 从查询实体到当前实体的 ID 路径。 path: Vec, }, } impl RetrievalItem { /// 统一分数(用于合并排序)。 pub fn score(&self) -> f32 { match self { Self::KnowledgePage { score, .. } => *score, Self::GraphEntity { score, .. } => *score, } } } /// 旧版带评分的知识页面检索结果(已废弃)。 /// /// 请迁移到 [`RetrievalItem::KnowledgePage`]。 #[deprecated(since = "0.3.0", note = "使用 RetrievalItem::KnowledgePage 代替")] #[derive(Debug, Clone)] pub struct ScoredItem { pub page: KnowledgePage, pub score: f32, } /// 检索结果。 #[derive(Debug, Clone)] pub struct RetrievalResult { /// 统一条目列表,按分数降序排列。 pub items: Vec, pub query: String, /// 本次检索实际执行的策略(可能因 graph 未注入而退化),而非用户通过 `with_strategy()` 配置的值。 pub strategy: RetrievalStrategy, } /// 记忆检索器 -- 在 `KnowledgeStore` 中做关键词检索并按 TextOverlap 评分, /// 可选注入 `KnowledgeGraph` 启用双通道检索。 pub struct MemoryRetriever { knowledge_store: KnowledgeStore, /// 可选知识图谱(None 时退化为单通道)。 knowledge_graph: Option>, /// 检索策略(默认 Hybrid)。 strategy: RetrievalStrategy, config: RetrieverConfig, /// 停用词表(用于关键词提取)。 stop_words: HashSet, } impl MemoryRetriever { /// 创建一个新的 MemoryRetriever(保持向后兼容)。 pub fn new(knowledge_store: KnowledgeStore, config: RetrieverConfig) -> Self { Self { knowledge_store, knowledge_graph: None, strategy: RetrievalStrategy::default(), config, stop_words: default_stop_words(), } } /// 注入知识图谱,启用双通道检索。 pub fn with_knowledge_graph(mut self, graph: Arc) -> Self { self.knowledge_graph = Some(graph); self } /// 设置检索策略。 pub fn with_strategy(mut self, strategy: RetrievalStrategy) -> Self { self.strategy = strategy; self } /// 替换停用词表。 pub fn with_stop_words(mut self, stop_words: HashSet) -> Self { self.stop_words = stop_words; self } /// 检索相关记忆(双通道)。 /// /// 根据 `strategy` 和是否注入 `knowledge_graph` 分流: /// - `KnowledgeOnly` 或未注入 graph -> 仅检索 KnowledgeStore /// - `GraphOnly` -> 仅检索 KnowledgeGraph /// - `Hybrid` -> 并行检索两通道,合并排序 pub async fn retrieve(&self, query: &str) -> Result { if query.is_empty() { return Ok(RetrievalResult { items: Vec::new(), query: query.to_string(), strategy: self.strategy.clone(), }); } let keywords = extract_keywords(query, &self.stop_words); let has_graph = self.knowledge_graph.is_some(); // 按策略分流 match (&self.strategy, has_graph) { // 仅知识页面(或未注入 graph 退化为单通道) (RetrievalStrategy::KnowledgeOnly, _) | (_, false) => { let items = self.search_knowledge_store(query, &keywords).await?; Ok(RetrievalResult { items, query: query.to_string(), strategy: RetrievalStrategy::KnowledgeOnly, }) } // 仅图谱 (RetrievalStrategy::GraphOnly, true) => { let graph = self.knowledge_graph.as_ref().unwrap(); let items = self.search_graph(&keywords, graph).await?; Ok(RetrievalResult { items, query: query.to_string(), strategy: self.strategy.clone(), }) } // 混合:并行执行,合并排序 (RetrievalStrategy::Hybrid, true) => { let graph = self.knowledge_graph.as_ref().unwrap(); let (kp_items, g_items) = tokio::join!( self.search_knowledge_store(query, &keywords), self.search_graph(&keywords, graph), ); let mut items = kp_items?; items.extend(g_items?); // 过滤最低分数 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(), strategy: self.strategy.clone(), }) } } } /// 通道 1:KnowledgeStore 关键词检索 + TextOverlap 评分。 async fn search_knowledge_store( &self, query: &str, keywords: &[String], ) -> Result, MemoryError> { 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); } } } let mut items: Vec = pages .into_iter() .map(|page| { let score = text_overlap_score(query, &page); RetrievalItem::KnowledgePage { page, score } }) .collect(); // 过滤 -> 排序 -> 截取 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(items) } /// 通道 2:KnowledgeGraph 关键词检索 + BFS 图遍历。 /// /// 流程:find_by_keywords 找到起始实体 -> 对每个起始实体 BFS 遍历 -> /// 收集 ScoredEntity -> 转为 RetrievalItem::GraphEntity。 async fn search_graph( &self, keywords: &[String], graph: &Arc, ) -> Result, MemoryError> { let starts = graph.find_by_keywords(keywords).await?; if starts.is_empty() { return Ok(Vec::new()); } let depth = self.config.graph_depth; let mut seen: HashSet = HashSet::new(); let mut items: Vec = Vec::new(); for start in starts { // 起始实体自身也加入结果(score = 1.0) if seen.insert(start.id.clone()) { items.push(RetrievalItem::GraphEntity { entity: start.clone(), score: 1.0, path: vec![start.id.clone()], }); } // BFS 找相关实体 let related: Vec = graph .get_related( &start.id, depth, crate::memory::graph::RelationDirection::Both, None, ) .await?; for se in related { if seen.insert(se.entity.id.clone()) { items.push(RetrievalItem::GraphEntity { entity: se.entity, score: se.score, path: se.path, }); } } } 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(items) } } /// 从 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::graph::{GraphEntity, GraphRelation, InMemoryGraph}; 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, } } fn make_store() -> (Arc, KnowledgeStore) { let store = Arc::new(InMemoryStore::new()); let ks = KnowledgeStore::new(store.clone()); (store, ks) } fn make_retriever(ks: KnowledgeStore) -> MemoryRetriever { MemoryRetriever::new(ks, RetrieverConfig::default()) } // ── 单通道(KnowledgeStore)测试 ── #[tokio::test] async fn retrieve_empty_query() { let (_, ks) = make_store(); let retriever = make_retriever(ks); let result = retriever.retrieve("").await.unwrap(); assert!(result.items.is_empty()); // 空查询返回用户配置的策略 assert_eq!(result.strategy, RetrievalStrategy::Hybrid); } #[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()); // 第一个结果应该是 KnowledgePage 变体 match &result.items[0] { RetrievalItem::KnowledgePage { page, score } => { assert_eq!(page.id, "p1"); assert!(*score > 0.0); } _ => panic!("expected KnowledgePage variant"), } } #[tokio::test] async fn retrieve_respects_min_score() { let (_, ks) = make_store(); ks.add_page(make_page("p1", "X", "Y", "Z")).await.unwrap(); let config = RetrieverConfig { max_results: 10, min_score: 0.99, graph_depth: 2, }; 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 (_, ks) = make_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, graph_depth: 2, }; let retriever = MemoryRetriever::new(ks, config); let result = retriever.retrieve("LangGraph").await.unwrap(); assert_eq!(result.items.len(), 2); } #[tokio::test] async fn retrieve_without_graph_degrades_to_knowledge_only() { let (_, ks) = make_store(); ks.add_page(make_page("p1", "Test", "t", "t")) .await .unwrap(); // 默认 Hybrid 策略,但未注入 graph let retriever = MemoryRetriever::new(ks, RetrieverConfig::default()); let result = retriever.retrieve("Test").await.unwrap(); // 应退化为 KnowledgeOnly assert_eq!(result.strategy, RetrievalStrategy::KnowledgeOnly); } // ── 双通道(KnowledgeGraph)测试 ── async fn make_graph_with_data() -> Arc { let graph = Arc::new(InMemoryGraph::new()); // 准备一些实体 let mut langchain = GraphEntity::new("langchain", "LangChain", "framework"); langchain.description = "LLM application framework".to_string(); let mut langgraph = GraphEntity::new("langgraph", "LangGraph", "framework"); langgraph.description = "Graph-based agent runtime".to_string(); let mut python = GraphEntity::new("python", "Python", "language"); python.description = "Programming language".to_string(); graph.add_entity(langchain).await.unwrap(); graph.add_entity(langgraph).await.unwrap(); graph.add_entity(python).await.unwrap(); graph .add_relation(GraphRelation::new( "langchain", "langgraph", "includes", 0.8, )) .await .unwrap(); graph .add_relation(GraphRelation::new("langchain", "python", "built_with", 0.9)) .await .unwrap(); graph } #[tokio::test] async fn retrieve_graph_only() { let (_, ks) = make_store(); let graph = make_graph_with_data().await; let retriever = MemoryRetriever::new(ks, RetrieverConfig::default()) .with_knowledge_graph(graph) .with_strategy(RetrievalStrategy::GraphOnly); let result = retriever.retrieve("langchain").await.unwrap(); assert_eq!(result.strategy, RetrievalStrategy::GraphOnly); // 应该全部是 GraphEntity 变体 assert!( result .items .iter() .all(|i| matches!(i, RetrievalItem::GraphEntity { .. })) ); // langchain 是起始实体(score=1.0),langgraph 和 python 是 BFS 结果 assert!(!result.items.is_empty(), "should find graph entities"); } #[tokio::test] async fn retrieve_hybrid_both_channels() { let store = Arc::new(InMemoryStore::new()) as Arc; let ks = KnowledgeStore::new(store); // KnowledgeStore 有匹配 ks.add_page(make_page( "p1", "LangChain framework", "LLM application framework", "LangChain is a framework for building LLM applications", )) .await .unwrap(); // KnowledgeGraph 也有匹配 let graph = make_graph_with_data().await; let retriever = MemoryRetriever::new(ks, RetrieverConfig::default()).with_knowledge_graph(graph); let result = retriever.retrieve("langchain").await.unwrap(); assert_eq!(result.strategy, RetrievalStrategy::Hybrid); // 应该同时包含 KnowledgePage 和 GraphEntity let has_page = result .items .iter() .any(|i| matches!(i, RetrievalItem::KnowledgePage { .. })); let has_entity = result .items .iter() .any(|i| matches!(i, RetrievalItem::GraphEntity { .. })); assert!(has_page, "should have KnowledgePage results"); assert!(has_entity, "should have GraphEntity results"); } #[tokio::test] async fn retrieve_hybrid_only_store_has_results() { let (_, ks) = make_store(); ks.add_page(make_page("p1", "LangChain", "framework", "LLM app")) .await .unwrap(); // 空图 let graph = Arc::new(InMemoryGraph::new()); let retriever = MemoryRetriever::new(ks, RetrieverConfig::default()).with_knowledge_graph(graph); let result = retriever.retrieve("langchain").await.unwrap(); assert_eq!(result.strategy, RetrievalStrategy::Hybrid); // 图空,只有 Store 结果 let has_page = result .items .iter() .any(|i| matches!(i, RetrievalItem::KnowledgePage { .. })); assert!(has_page); } #[tokio::test] async fn retrieve_hybrid_only_graph_has_results() { let store = Arc::new(InMemoryStore::new()) as Arc; let ks = KnowledgeStore::new(store); // Store 空,图有数据 let graph = make_graph_with_data().await; let retriever = MemoryRetriever::new(ks, RetrieverConfig::default()).with_knowledge_graph(graph); let result = retriever.retrieve("langchain").await.unwrap(); assert_eq!(result.strategy, RetrievalStrategy::Hybrid); let has_entity = result .items .iter() .any(|i| matches!(i, RetrievalItem::GraphEntity { .. })); assert!( has_entity, "should have graph results even when store is empty" ); } #[tokio::test] async fn retrieve_strategy_knowledge_only_ignores_graph() { let (_, ks) = make_store(); ks.add_page(make_page("p1", "Test", "t", "t")) .await .unwrap(); let graph = make_graph_with_data().await; let retriever = MemoryRetriever::new(ks, RetrieverConfig::default()) .with_knowledge_graph(graph) .with_strategy(RetrievalStrategy::KnowledgeOnly); let result = retriever.retrieve("langchain").await.unwrap(); assert_eq!(result.strategy, RetrievalStrategy::KnowledgeOnly); // 即使图有数据,KnowledgeOnly 也不应该返回 GraphEntity let has_entity = result .items .iter() .any(|i| matches!(i, RetrievalItem::GraphEntity { .. })); assert!( !has_entity, "KnowledgeOnly should not return graph entities" ); } // ── 辅助函数测试 ── #[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())); } #[test] fn retrieval_item_score_accessor() { let page = KnowledgePage { id: "p1".into(), title: "T".into(), summary: "S".into(), content: "C".into(), tags: vec![], references: vec![], created_at: OffsetDateTime::now_utc(), updated_at: OffsetDateTime::now_utc(), }; let item = RetrievalItem::KnowledgePage { page, score: 0.5 }; assert!((item.score() - 0.5).abs() < 0.001); let entity = GraphEntity::new("e1", "E1", "x"); let item2 = RetrievalItem::GraphEntity { entity, score: 0.8, path: vec!["e1".into()], }; assert!((item2.score() - 0.8).abs() < 0.001); } #[test] fn default_strategy_is_hybrid() { assert_eq!(RetrievalStrategy::default(), RetrievalStrategy::Hybrid); } }