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:
+1080
File diff suppressed because it is too large
Load Diff
+409
-31
@@ -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(),
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 通道 1:KnowledgeStore 关键词检索 + 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(),
|
||||
})
|
||||
/// 通道 2:KnowledgeGraph 关键词检索 + 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);
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user