Files
agcore/src/memory/retriever.rs
T
徐涛 5baa170508 docs: 更新 README feature 表 + 升级指南 + 示例注释 + roadmap 同步
- README 添加 feature 组合表 + 模块级 features 清单 + 升级指南
- 18 个 example 顶部添加 Required features 注释
- roadmap.md 和 roadmap-v0.3.2.md 同步 Phase 26-27 完成状态
- cargo fmt 全量格式化(修复预存格式问题,CI format job 可通过)
2026-07-19 08:18:04 +08:00

703 lines
24 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! 记忆检索器 -- 双通道关键词检索(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(),
})
}
}
}
/// 通道 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 {
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)
}
/// 通道 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.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);
}
}