- README 添加 feature 组合表 + 模块级 features 清单 + 升级指南 - 18 个 example 顶部添加 Required features 注释 - roadmap.md 和 roadmap-v0.3.2.md 同步 Phase 26-27 完成状态 - cargo fmt 全量格式化(修复预存格式问题,CI format job 可通过)
703 lines
24 KiB
Rust
703 lines
24 KiB
Rust
//! 记忆检索器 -- 双通道关键词检索(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<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,
|
||
pub score: f32,
|
||
}
|
||
|
||
/// 检索结果。
|
||
#[derive(Debug, Clone)]
|
||
pub struct RetrievalResult {
|
||
/// 统一条目列表,按分数降序排列。
|
||
pub items: Vec<RetrievalItem>,
|
||
pub query: String,
|
||
/// 本次检索实际执行的策略(可能因 graph 未注入而退化),而非用户通过 `with_strategy()` 配置的值。
|
||
pub strategy: RetrievalStrategy,
|
||
}
|
||
|
||
/// 记忆检索器 -- 在 `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(保持向后兼容)。
|
||
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(),
|
||
});
|
||
}
|
||
|
||
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<Vec<RetrievalItem>, 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<RetrievalItem> = 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<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.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<String>) -> Vec<String> {
|
||
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<String> {
|
||
let chars: Vec<char> = s.chars().collect();
|
||
chars.windows(2).map(|w| w.iter().collect()).collect()
|
||
}
|
||
|
||
fn default_stop_words() -> HashSet<String> {
|
||
[
|
||
"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<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 (_, 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<dyn MemoryStore>;
|
||
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<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);
|
||
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);
|
||
}
|
||
}
|