feat(memory): 实现知识图谱与双通道检索

- 新增 KnowledgeGraph trait + InMemoryGraph(BFS 图遍历 + 标签管理)
- 扩展 MemoryRetriever 为双通道检索(Hybrid/KnowledgeOnly/GraphOnly)
- 统一 RetrievalResult 为 RetrievalItem enum 变体
- GraphRelation 使用复合键 + composite_key() 派生 id
- 旧 ScoredItem 标注 #[deprecated]
- 新增 knowledge_graph_demo 示例
- 全量测试 427 passed,clippy/doc 0 警告
This commit is contained in:
徐涛
2026-07-17 14:37:49 +08:00
parent 5bb349d177
commit 7e72e102a2
8 changed files with 2365 additions and 43 deletions
+1080
View File
File diff suppressed because it is too large Load Diff
+409 -31
View File
@@ -1,8 +1,13 @@
//! 记忆检索器 —— 基于 TextOverlap (Dice 系数) 的单通道关键词检索
//! 记忆检索器 -- 双通道关键词检索(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;
@@ -13,6 +18,8 @@ pub struct RetrieverConfig {
pub max_results: usize,
/// 最低分数阈值 [0.0, 1.0](默认 0.1)。
pub min_score: f32,
/// 图遍历的默认深度(默认 2)。
pub graph_depth: usize,
}
impl Default for RetrieverConfig {
@@ -20,64 +27,193 @@ impl Default for RetrieverConfig {
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<String>,
},
}
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,
/// TextOverlap 评分 [0.0, 1.0]
pub score: f32,
}
/// 检索结果。
#[derive(Debug, Clone)]
pub struct RetrievalResult {
pub items: Vec<ScoredItem>,
/// 统一条目列表,按分数降序排列。
pub items: Vec<RetrievalItem>,
pub query: String,
/// 本次检索实际执行的策略(可能因 graph 未注入而退化),而非用户通过 `with_strategy()` 配置的值。
pub strategy: RetrievalStrategy,
}
/// 记忆检索器 —— 在 `KnowledgeStore` 中做关键词检索并按 TextOverlap 评分
/// 记忆检索器 -- 在 `KnowledgeStore` 中做关键词检索并按 TextOverlap 评分
/// 可选注入 `KnowledgeGraph` 启用双通道检索。
pub struct MemoryRetriever {
knowledge_store: KnowledgeStore,
/// 可选知识图谱(None 时退化为单通道)。
knowledge_graph: Option<Arc<dyn KnowledgeGraph>>,
/// 检索策略(默认 Hybrid)。
strategy: RetrievalStrategy,
config: RetrieverConfig,
/// 停用词表(用于关键词提取)。
stop_words: HashSet<String>,
}
impl MemoryRetriever {
/// 创建一个新的 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<dyn KnowledgeGraph>) -> 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<String>) -> Self {
self.stop_words = stop_words;
self
}
/// 检索相关知识页面
/// 检索相关记忆(双通道)
///
/// 根据 `strategy` 和是否注入 `knowledge_graph` 分流:
/// - `KnowledgeOnly` 或未注入 graph -> 仅检索 KnowledgeStore
/// - `GraphOnly` -> 仅检索 KnowledgeGraph
/// - `Hybrid` -> 并行检索两通道,合并排序
pub async fn retrieve(&self, query: &str) -> Result<RetrievalResult, MemoryError> {
if query.is_empty() {
return Ok(RetrievalResult {
items: Vec::new(),
query: query.to_string(),
strategy: self.strategy.clone(),
});
}
// 1. 关键词提取
let keywords = extract_keywords(query, &self.stop_words);
let has_graph = self.knowledge_graph.is_some();
// 2. 用关键词在 KnowledgeStore 中搜索
// 按策略分流
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(),
})
}
}
}
/// 通道 1KnowledgeStore 关键词检索 + TextOverlap 评分。
async fn search_knowledge_store(
&self,
query: &str,
keywords: &[String],
) -> Result<Vec<RetrievalItem>, MemoryError> {
let mut pages = Vec::new();
for keyword in &keywords {
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) {
@@ -86,32 +222,78 @@ impl MemoryRetriever {
}
}
// 3. TextOverlap 评分
let mut items: Vec<ScoredItem> = pages
let mut items: Vec<RetrievalItem> = pages
.into_iter()
.map(|page| {
let score = text_overlap_score(query, &page);
ScoredItem { page, score }
RetrievalItem::KnowledgePage { page, score }
})
.collect();
// 4. 过滤 排序 截取
items.retain(|i| i.score >= self.config.min_score);
// 过滤 -> 排序 -> 截取
items.retain(|i| i.score() >= self.config.min_score);
items.sort_by(|a, b| {
b.score
.partial_cmp(&a.score)
b.score()
.partial_cmp(&a.score())
.unwrap_or(std::cmp::Ordering::Equal)
});
items.truncate(self.config.max_results);
Ok(items)
}
Ok(RetrievalResult {
items,
query: query.to_string(),
})
/// 通道 2KnowledgeGraph 关键词检索 + BFS 图遍历。
///
/// 流程:find_by_keywords 找到起始实体 -> 对每个起始实体 BFS 遍历 ->
/// 收集 ScoredEntity -> 转为 RetrievalItem::GraphEntity。
async fn search_graph(
&self,
keywords: &[String],
graph: &Arc<dyn KnowledgeGraph>,
) -> Result<Vec<RetrievalItem>, 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<String> = HashSet::new();
let mut items: Vec<RetrievalItem> = 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<ScoredEntity> = 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.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 中提取关键词:按非字母数字字符分割 转小写 过滤单字符和停用词。
/// 从 query 中提取关键词:按非字母数字字符分割 -> 转小写 -> 过滤单字符和停用词。
fn extract_keywords(query: &str, stop_words: &HashSet<String>) -> Vec<String> {
query
.split(|c: char| !c.is_alphanumeric())
@@ -178,6 +360,7 @@ fn default_stop_words() -> HashSet<String> {
#[cfg(test)]
mod tests {
use super::*;
use crate::memory::graph::{GraphEntity, InMemoryGraph, GraphRelation};
use crate::memory::knowledge::KnowledgeStore;
use crate::memory::{InMemoryStore, MemoryStore};
use std::sync::Arc;
@@ -197,13 +380,26 @@ mod tests {
}
}
fn make_store() -> (Arc<InMemoryStore>, 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 store = Arc::new(InMemoryStore::new()) as Arc<dyn MemoryStore>;
let ks = KnowledgeStore::new(store);
let retriever = MemoryRetriever::new(ks, RetrieverConfig::default());
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]
@@ -225,19 +421,25 @@ mod tests {
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);
// 第一个结果应该是 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 store = Arc::new(InMemoryStore::new()) as Arc<dyn MemoryStore>;
let ks = KnowledgeStore::new(store);
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
@@ -249,8 +451,7 @@ mod tests {
#[tokio::test]
async fn retrieve_respects_max_results() {
let store = Arc::new(InMemoryStore::new()) as Arc<dyn MemoryStore>;
let ks = KnowledgeStore::new(store);
let (_, ks) = make_store();
for i in 0..5 {
ks.add_page(make_page(
&format!("p{i}"),
@@ -264,12 +465,160 @@ mod tests {
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<InMemoryGraph> {
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<dyn MemoryStore>;
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<dyn MemoryStore>;
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);
@@ -299,4 +648,33 @@ mod tests {
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);
}
}