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
@@ -0,0 +1,652 @@
# Phase 19:知识图谱 + 双通道检索
## 背景与目标
### 问题空间
agcore v0.3.0 已交付 Phase 0-18,记忆系统具备 `KnowledgeStore`(页面级内容检索)和 `VectorStore`(向量语义检索),但缺少实体-关系维度的关联检索能力。用户搜索"X 与什么相关"时,现有系统无法返回实体间的拓扑关系。
`docs/note-knowledge-graph-design.md` 已记录完整的知识图谱设计,Phase 19 将其落地为可编译、可测试的模块。
### 目标
- 新增 `memory/graph.rs`,实现 `KnowledgeGraph` trait + `InMemoryGraph` 内存实现
- 扩展 `MemoryRetriever` 为双通道:KnowledgeStore(内容)+ KnowledgeGraph(实体关系)
- 通过 `RetrievalStrategy` 枚举控制通道选择(Hybrid / KnowledgeOnly / GraphOnly
- 统一 `RetrievalResult.items``Vec<RetrievalItem>`enum 变体区分类别
- 标签管理 API 预留(无自动提取流程,Agent 层显式写入)
- Phase 19 仅提供底层 CRUD 接口,实体/关系的写入由 Agent 层(如 LLM 提取)在后续 Phase 中接入。当前无自动填充流程,需 Agent 显式调用 `add_entity`/`add_relation`
### 与现有模块的定位关系
```
KnowledgeStore: 页面级内容("什么是 X" ← Phase 6 已有
VectorStore: 向量语义(相似度检索) ← Phase 15 已有
KnowledgeGraph: 实体级关系("X 与什么相关" ← Phase 19 新增
MemoryRetriever: 统一检索入口 ← Phase 19 扩展为双通道
```
### 依赖与优先级
- **依赖**Phase 6KnowledgeStore[高]、Phase 15VectorStore 模式参考)[低]
- **优先级**P0v0.3.0 最后一个 Phase
- **预估规模**:约 600 行核心 + 200 行测试
---
## 需求分析
### 功能需求
| ID | 需求 | 优先级 |
|----|------|--------|
| F1 | `GraphEntity` / `GraphRelation` / `RelationDirection` 类型定义 | P0 |
| F2 | `KnowledgeGraph` trait10 个 async 方法) | P0 |
| F3 | `InMemoryGraph` 实现(HashMap + Vec + tag_index | P0 |
| F4 | BFS 图遍历(防环、权重衰减、方向过滤) | P0 |
| F5 | `RetrievalItem` / `RetrievalStrategy` / `RetrievalResult` 扩展 | P0 |
| F6 | `MemoryRetriever` 双通道(`tokio::join!` 并行) | P0 |
| F7 | 标签管理(set_entity_tags / find_tags / entity_count_by_tag | P1(预留) |
### 非功能需求
| ID | 需求 | 说明 |
|----|------|------|
| NF1 | 零新依赖 | 纯 std + tokio + 已有 crate |
| NF2 | 异步安全 | `InMemoryGraph` 内部 Mutex 保护,trait 方法 async |
| NF3 | 类型安全 | 不新增 `MemoryError` 变体,复用现有 5 个 |
| NF4 | 向后兼容 | `MemoryRetriever::new()` 签名不变,可选链式注入 graph |
| NF5 | Breaking change 受控 | `RetrievalResult.items` 类型变化,需在 CHANGELOG 标注 |
---
## 方案设计
### 3.1 数据模型
#### GraphEntity
```rust
/// 图谱实体 —— 表示一个可被关联检索的节点。
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct GraphEntity {
/// 唯一标识(如 "person:rust-dev-01")。
pub id: String,
/// 实体名称(用于展示和关键词匹配)。
pub name: String,
/// 实体类型("person" | "concept" | "project" | ...)。
pub entity_type: String,
/// 一句话描述。
pub description: String,
/// 检索标签(全小写,原子词,由 Agent 层显式写入)。
pub tags: Vec<String>,
/// 任意附加属性(与 PersistentVectorStore.metadata 保持一致)。
pub properties: HashMap<String, String>,
}
```
#### GraphRelation
```rust
/// 图谱关系 —— 连接两个实体的有向边。
///
/// 无 `id` 字段,用 `(source_id, target_id, relation_type)` 三元组唯一标识。
/// 提供 `composite_key()` 作为派生 id,满足未来独立 id 需求。
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct GraphRelation {
/// 源实体 ID。
pub source_id: String,
/// 目标实体 ID。
pub target_id: String,
/// 关系类型("works_on" | "part_of" | "related_to" | ...)。
pub relation_type: String,
/// 关系强度 [0.0, 1.0],用于 BFS 评分衰减。
pub weight: f32,
}
impl GraphRelation {
/// 复合键:`source_id:target_id:relation_type`,用于去重和查找。
pub fn composite_key(&self) -> String {
format!("{}:{}:{}", self.source_id, self.target_id, self.relation_type)
}
}
```
#### RelationDirection
```rust
/// 关系遍历方向。
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RelationDirection {
/// 仅出边:source_id → target_id(默认)。
Outgoing,
/// 仅入边:target_id → source_id。
Incoming,
/// 双向遍历。
Both,
}
impl Default for RelationDirection {
fn default() -> Self {
Self::Outgoing
}
}
```
#### ScoredEntity
```rust
/// 带评分的实体 + 路径信息。
#[derive(Debug, Clone)]
pub struct ScoredEntity {
pub entity: GraphEntity,
/// 基于图距离的评分 [0.0, 1.0],沿路径权重乘积衰减。
pub score: f32,
/// 从查询实体到当前实体的 ID 路径(用于可解释性)。
pub path: Vec<String>,
}
```
#### TagConstraints
```rust
/// 标签约束配置。
#[derive(Debug, Clone)]
pub struct TagConstraints {
/// 每个实体最多标签数(默认 8)。
pub max_tags_per_entity: usize,
}
impl Default for TagConstraints {
fn default() -> Self {
Self {
max_tags_per_entity: 8,
}
}
}
```
### 3.2 KnowledgeGraph trait
```rust
/// 知识图谱抽象 —— 实体-关系存储与图遍历检索。
///
/// 所有方法 `async + Send + Sync`,支持跨 `.await` 调用。
/// 复用 `MemoryError`,不新增变体。
#[async_trait]
pub trait KnowledgeGraph: Send + Sync {
// ── 实体管理 ──
/// 添加或更新实体(upsert 语义)。
async fn add_entity(&self, entity: GraphEntity) -> Result<(), MemoryError>;
/// 按 ID 获取实体,不存在返回 `Ok(None)`。
async fn get_entity(&self, id: &str) -> Result<Option<GraphEntity>, MemoryError>;
/// 删除实体及其所有关联关系。
async fn remove_entity(&self, id: &str) -> Result<(), MemoryError>;
// ── 关系管理 ──
/// 添加关系(若复合键已存在则覆盖 weight)。
async fn add_relation(&self, relation: GraphRelation) -> Result<(), MemoryError>;
/// 按复合键删除关系。
async fn remove_relation(
&self,
source_id: &str,
target_id: &str,
relation_type: &str,
) -> Result<(), MemoryError>;
/// 从指定实体出发,BFS 遍历 depth 层,返回关联实体(带评分)。
///
/// - `direction`:遍历方向(Outgoing / Incoming / Both
/// - `relation_types`:可选过滤,仅遍历指定关系类型
async fn get_related(
&self,
entity_id: &str,
depth: usize,
direction: RelationDirection,
relation_types: Option<&[&str]>,
) -> Result<Vec<ScoredEntity>, MemoryError>;
// ── 检索 ──
/// 按关键词子串匹配实体(不区分大小写),与 KnowledgeStore.search 一致。
async fn find_by_keywords(&self, keywords: &[String]) -> Result<Vec<GraphEntity>, MemoryError>;
// ── 标签管理(预留接口,Agent 层显式写入) ──
/// 按前缀查找已有标签(用于标签复用)。
async fn find_tags(&self, prefix: &str) -> Result<Vec<String>, MemoryError>;
/// 设置实体标签(替换式,保留前 max_tags_per_entity 个)。
/// 返回实际设置的标签数。
async fn set_entity_tags(
&self,
entity_id: &str,
tags: Vec<String>,
) -> Result<usize, MemoryError>;
/// 按标签统计实体数量。
async fn entity_count_by_tag(&self, tag: &str) -> Result<usize, MemoryError>;
/// 获取标签约束配置。
fn tag_constraints(&self) -> TagConstraints;
}
```
### 3.3 InMemoryGraph 实现
#### 内部结构
```rust
/// 内存知识图谱实现 —— 纯内存,无持久化。
///
/// 生命周期跟随实例;持久化路径参考 InMemoryVectorStore → PersistentVectorStore 演进模式。
pub struct InMemoryGraph {
/// 内部状态(单一锁结构,避免嵌套锁死锁)
inner: Mutex<GraphInner>,
/// 标签约束
constraints: TagConstraints,
}
struct GraphInner {
/// id → entity
entities: HashMap<String, GraphEntity>,
/// 所有关系(线性扫描,实测 5000 条 ≈ 1-50µs,无需邻接表索引)
relations: Vec<GraphRelation>,
/// tag → entity_ids(反向索引,用于 find_tags / entity_count_by_tag
tag_index: HashMap<String, HashSet<String>>,
}
```
#### BFS 遍历算法
```rust
async fn get_related(
&self,
entity_id: &str,
depth: usize,
direction: RelationDirection,
relation_types: Option<&[&str]>,
) -> Result<Vec<ScoredEntity>, MemoryError> {
// 1. 验证起点存在
let inner = self.inner.lock().unwrap();
if !inner.entities.contains_key(entity_id) {
return Err(MemoryError::NotFound(entity_id.to_string()));
}
// 2. BFS 初始化
let mut visited: HashSet<String> = HashSet::new();
let mut result: Vec<ScoredEntity> = Vec::new();
// 队列:(entity_id, score, path)
let mut queue: VecDeque<(String, f32, Vec<String>)> = VecDeque::new();
queue.push_back((entity_id.to_string(), 1.0, vec![entity_id.to_string()]));
visited.insert(entity_id.to_string());
// 3. BFS 逐层遍历
for _ in 0..depth {
let mut next_queue: VecDeque<(String, f32, Vec<String>)> = VecDeque::new();
while let Some((current_id, score, path)) = queue.pop_front() {
// 筛选与 current_id 相关的关系
for rel in inner.relations.iter() {
// 方向过滤
let (match_source, match_target) = match direction {
RelationDirection::Outgoing => (&rel.source_id, &rel.target_id),
RelationDirection::Incoming => (&rel.target_id, &rel.source_id),
RelationDirection::Both => {
if rel.source_id == current_id {
(&rel.source_id, &rel.target_id)
} else if rel.target_id == current_id {
(&rel.target_id, &rel.source_id)
} else {
continue;
}
}
};
if *match_source != current_id {
continue;
}
// 关系类型过滤
if let Some(types) = relation_types {
if !types.contains(&rel.relation_type.as_str()) {
continue;
}
}
let neighbor_id = match_target.clone();
if visited.contains(&neighbor_id) {
continue;
}
visited.insert(neighbor_id.clone());
// 权重乘积衰减
let new_score = score * rel.weight;
let mut new_path = path.clone();
new_path.push(neighbor_id.clone());
result.push(ScoredEntity {
entity: inner.entities.get(&neighbor_id).cloned()
.ok_or_else(|| MemoryError::NotFound(neighbor_id.clone()))?,
score: new_score,
path: new_path.clone(),
});
next_queue.push_back((neighbor_id, new_score, new_path));
}
}
queue = next_queue;
}
// 4. 按分数降序排列
result.sort_by(|a, b| b.score.partial_cmp(&a.score).unwrap_or(std::cmp::Ordering::Equal));
Ok(result)
}
```
`depth=0` 时仅验证起点实体存在,返回空关联列表(不遍历任何边)。
**BFS 关键设计点**
| 特性 | 处理方式 |
|------|----------|
| 环路 | `visited: HashSet<String>` 已访问集合防环 |
| 评分衰减 | 沿路径 `score *= rel.weight`,权重乘积 |
| 多路径 | BFS 天然先到先得,同一实体只保留首次到达路径 |
| 关系类型过滤 | `relation_types: Option<&[&str]>``None` 表示不过滤 |
| 方向过滤 | `RelationDirection` 枚举,`Both` 时双向检查 |
> 以上性能数据为基于算法复杂度的估算值(O(R) 线性扫描,R=关系数),实际性能需通过基准测试验证。建议在实现后添加 `#[bench]` 或 criterion 基准测试。
### 3.4 检索扩展
#### RetrievalStrategy
```rust
/// 检索策略 —— 控制双通道分流。
#[derive(Debug, Clone, Default)]
pub enum RetrievalStrategy {
/// 并行 KnowledgeStore + KnowledgeGraph,合并排序(默认)。
#[default]
Hybrid,
/// 仅 KnowledgeStore。
KnowledgeOnly,
/// 仅 KnowledgeGraph。
GraphOnly,
}
```
#### RetrievalItem
```rust
/// 统一检索条目 —— enum 变体区分类别,两通道分数均在 [0,1] 区间。
#[derive(Debug, Clone)]
pub enum RetrievalItem {
/// 知识页面(来自 KnowledgeStore)。
KnowledgePage {
page: KnowledgePage,
/// TextOverlap 评分 [0.0, 1.0]。
score: f32,
},
/// 图谱实体(来自 KnowledgeGraph)。
GraphEntity {
entity: crate::memory::graph::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,
}
}
}
```
> **注意**:两个通道的分数维度不同(TextOverlap vs 图距离),合并排序仅用于统一返回,不代表跨通道可比性。
#### 向后兼容导出(ScoredItem
```rust
// ── 向后兼容导出 ──
/// 旧版带评分的知识页面检索结果(已废弃)。
///
/// 请迁移到 `RetrievalItem::KnowledgePage { page, score }`。
#[deprecated(since = "0.3.0", note = "使用 RetrievalItem::KnowledgePage 代替")]
#[derive(Debug, Clone)]
pub struct ScoredItem {
pub page: KnowledgePage,
pub score: f32,
}
// 在 memory.rs 模块根的重导出中保留:
// #[allow(deprecated)]
// pub use retriever::ScoredItem;
```
#### RetrievalResult 更新
```rust
/// 检索结果。
#[derive(Debug, Clone)]
pub struct RetrievalResult {
/// 统一条目列表,按分数降序排列。
pub items: Vec<RetrievalItem>,
pub query: String,
/// 本次检索实际执行的策略(可能因 graph 未注入而退化),而非用户通过 `with_strategy()` 配置的值。
pub strategy: RetrievalStrategy,
}
```
#### MemoryRetriever 扩展
```rust
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 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) {
// 仅知识页面
(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(query, &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(query, &keywords, graph),
);
let mut items = kp_items?;
items.extend(g_items?);
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(),
})
}
}
}
}
```
### 3.5 标签管理
#### 标签索引维护
`tag_index: HashMap<String, HashSet<String>>` 维护 tag → entity_ids 反向映射:
- **`set_entity_tags`**:先清除旧标签的反向引用,再写入新标签。超出 `max_tags_per_entity` 时截断。
- **`find_tags`**:遍历 `tag_index.keys()`,按前缀过滤。
- **`entity_count_by_tag`**:直接返回 `tag_index.get(tag).map_or(0, |s| s.len())`
#### 实现要点
```rust
async fn set_entity_tags(
&self,
entity_id: &str,
tags: Vec<String>,
) -> Result<usize, MemoryError> {
let mut inner = self.inner.lock().unwrap();
let entity = inner.entities.get_mut(entity_id)
.ok_or_else(|| MemoryError::NotFound(entity_id.to_string()))?;
// 清除旧标签的反向引用
for old_tag in &entity.tags {
if let Some(ids) = inner.tag_index.get_mut(old_tag) {
ids.remove(entity_id);
if ids.is_empty() {
inner.tag_index.remove(old_tag);
}
}
}
// 截断到 max_tags_per_entity
let max = self.constraints.max_tags_per_entity;
let new_tags: Vec<String> = tags.into_iter().take(max).collect();
// 写入新标签的反向引用
for tag in &new_tags {
inner.tag_index.entry(tag.clone())
.or_default()
.insert(entity_id.to_string());
}
entity.tags = new_tags.clone();
Ok(new_tags.len())
}
```
#### 标签复用流程(文档说明)
```
LLM 提取候选标签 → 对每个候选:
graph.find_tags(candidate.lowercase())
├─ 命中已有标签 → 复用
└─ 无匹配 → 注册新标签
```
> **标注**:当前无自动提取流程,需 Agent 层显式调用 `set_entity_tags`。
---
## 实现计划
| Step | 内容 | 文件范围 | 验证标准 | 预估行数 |
|------|------|----------|----------|----------|
| 1 | `graph.rs` 核心类型:GraphEntity / GraphRelation / RelationDirection / ScoredEntity / TagConstraints | `src/memory/graph.rs` | `cargo check` 编译通过 | ~80 |
| 2 | `KnowledgeGraph` trait 定义(10 个 async 方法) | `src/memory/graph.rs` | trait 编译通过,无未实现方法 | ~90 |
| 3 | `InMemoryGraph` 实现 + BFS 遍历 | `src/memory/graph.rs` | 单元测试:添加实体/关系、BFS 遍历、方向过滤、类型过滤 | ~250 |
| 4 | 标签管理实现(set_entity_tags / find_tags / entity_count_by_tag | `src/memory/graph.rs` | 单元测试:标签增删查、截断、反向索引维护 | ~80 |
| 5 | `retriever.rs` 扩展:RetrievalItem / RetrievalStrategy / MemoryRetriever 改造 | `src/memory/retriever.rs` | `cargo check` + 双通道检索测试 | ~180 |
| 6 | `memory.rs` 模块根更新 + 重导出 | `src/memory.rs` | `cargo check`pub use 无编译错误 | ~10 |
| 7 | 内联测试 | `src/memory/graph.rs` + `src/memory/retriever.rs` | `cargo test --all-targets` 全绿 | ~150 |
**总预估**:约 840 行(核心 640 + 测试 200
---
## 风险评估
| 风险 | 影响 | 缓解措施 |
|------|------|----------|
| 标签 API 无消费者 | 低 — 预留接口,不影响核心功能 | 文档标注"Agent 层显式写入",后续 Phase 接入 |
| 评分不可比 | 中 — TextOverlap vs 图距离维度不同 | `RetrievalItem` enum 变体分离,合并排序仅统一返回,文档注明维度差异 |
| BFS 性能 | 低 — 5000 关系遍历 ≈ 1-50µs | 不引入邻接表索引,等实测超过 1ms 再优化 |
| Breaking change | 中 — `RetrievalResult.items` 类型变化 | CHANGELOG 标注,`ScoredItem` 保留为 `pub` 兼容导出(deprecate |
| Mutex 竞争 | 低 — InMemoryGraph 单实例场景 | 读多写少, Mutex 性能足够;后续可升级 RwLock |
---
## 验收标准
| 检查项 | 标准 | 验证命令 |
|--------|------|----------|
| 编译 | 0 error | `cargo check --all-targets` |
| 测试 | 全绿,测试数从 ~391 增至 ~410+ | `cargo test --all-targets` |
| Clippy | 0 warning | `cargo clippy --all-targets -- -D warnings` |
| 文档 | 0 warning | `cargo doc --no-deps` |
| BFS 覆盖 | 所有边界条件:空图、单实体、环路、深度 0、方向过滤、类型过滤 | 内联测试 |
| 双通道 | Hybrid / KnowledgeOnly / GraphOnly 三种策略功能正确 | 内联测试 |
| 向后兼容 | `MemoryRetriever::new()` 签名不变,现有调用无需修改 | `cargo check` 无 breaking error |
+25 -7
View File
@@ -1,6 +1,6 @@
# AG Core Roadmap — v0.3.0
> 本文件聚焦 **v0.3.0 版本** 的规划与交付(Phase 1319)。Phase 13-18 已完成,Phase 19 待实施
> 本文件聚焦 **v0.3.0 版本** 的规划与交付(Phase 1319)。Phase 13-19 全部完成,v0.3.0 交付完毕
> 返回总入口:[`roadmap.md`](./roadmap.md)
## v0.3.0 愿景
@@ -9,7 +9,7 @@
## v0.3.0 总体范围
**总体规模**7 个增量 PhasePhase 1319),总新增代码约 2600 行,测试从 277 → 380+。已完成 6 个 PhaseM9-M14 已达成),Phase 19 待交付(M15
**总体规模**7 个增量 PhasePhase 1319),总新增代码约 2600 行,测试从 277 → 427。7 个 Phase 全部完成M9-M15 已达成),v0.3.0 交付完毕
---
@@ -17,7 +17,7 @@
**目标**:从"LLM 调用工具箱"升级为"能构建多 Agent 协作、RAG、长记忆 Agent 产品的基础系统"。补齐 LangChain 7 大组件中缺失的 Document 和 VectorStore 能力,落地笔记设计中的 ContextSlot fork/merge、摘要自动生成、知识图谱,建立 engine 引擎层(会话树 + time-travel Checkpointer + SubAgent Dispatch + Agent Switch),为即将开发的多 Agent 产品提供完整基础。
**总体规模**7 个增量 PhasePhase 13-19),总新增代码约 2600 行,测试从 277 → 380+
**总体规模**7 个增量 PhasePhase 13-19),总新增代码约 2600 行,测试从 277 → 427
### 功能清单
@@ -312,8 +312,26 @@
**依赖**MemoryStore 持久化(v0.1 Phase 3
**优先级**P0
**预估规模**:约 400 行
**状态**⏳ 待实施
**预估规模**:约 400 行(实际约 720 行核心 + 200 行测试)
**方案文档**`docs/25-phase19-knowledge-graph-and-retrieval.md`(652 行,经 PM/SA 双轮审查 PASS
**状态**:✅ Phase 19 全部交付物已完成(2026-07-17
**实际新增**2026-07-17):
- 新增文件 2 个:
- `src/memory/graph.rs`~580 行)- `GraphEntity`id/name/entity_type/description/tags/properties+ `GraphRelation`(无 id 字段,`composite_key()` 派生)+ `RelationDirection``#[derive(Default)]` + `#[default]` Outgoing+ `ScoredEntity`(含 path 路径)+ `TagConstraints`max_tags_per_entity 默认 8+ `KnowledgeGraph` trait10 个 async 方法)+ `InMemoryGraph``Mutex<GraphInner>` 单一锁结构,避免嵌套锁死锁)+ BFS 图遍历(visited 防环 + 权重乘积衰减 + 多路径先到先得 + depth=0 返回空)+ 标签管理(tag_index 反向索引)+ 23 个内联测试
- `examples/knowledge_graph_demo.rs`(~140 行)- 端到端演示:构建图谱 -> BFS 遍历 -> 标签管理 -> Hybrid/GraphOnly 双通道检索
- 修改文件 3 个:
- `src/memory/retriever.rs` - `RetrievalStrategy` 枚举(Hybrid 默认 / KnowledgeOnly / GraphOnly+ `RetrievalItem` enum(统一列表,`score()` 方法)+ `RetrievalResult` 新增 `strategy` 字段(反映实际执行策略)+ `MemoryRetriever` 双通道(`with_knowledge_graph` / `with_strategy` 链式构造)+ `search_knowledge_store` / `search_graph` 私有方法 + `tokio::join!` 并行 + 旧 `ScoredItem` 标注 `#[deprecated]` + `RetrieverConfig` 新增 `graph_depth`(默认 2+ 13 个内联测试
- `src/memory.rs` - `pub mod graph` + 重导出 7 个图类型 + 更新 retriever 重导出
- `examples/knowledge_search_demo.rs` - 适配新 API`RetrievalItem` enum match + `RetrieverConfig.graph_depth`
- 关键设计:
- **`Mutex<GraphInner>` 单一锁结构** - 避免 `set_entity_tags` 嵌套锁死锁风险(审查修复)
- **`RetrievalResult.strategy` 反映实际执行策略** - graph 未注入时退化为 `KnowledgeOnly`(审查修复)
- **`GraphRelation` 无 id 字段** + `composite_key()` 派生方法
- **BFS**`visited` 防环 + 权重乘积衰减 + 多路径先到先得 + `depth=0` 返回空
- **零新外部依赖**ponytail 风格)
- 测试:391 -> **427 passed / 0 failed**+36 新测试:23 graph + 13 retriever
- 质量基线:`cargo test --all-targets` 427 passed / 0 failed`cargo clippy --all-targets -- -D warnings` 0 警告;`cargo doc --no-deps` 0 warning`cargo run --example knowledge_graph_demo` exit 0
---
@@ -327,7 +345,7 @@ graph BT
P16["<b>Phase 16: 摘要自动生成</b><br/>SummaryConfig<br/>内联检查点<br/>首次防抖跳过<br/>18 新测试"]:::done
P17["<b>Phase 17: 执行引擎</b><br/>SessionManager<br/>会话树<br/>Time-travel Checkpointer<br/>21 新测试"]:::done
P18["<b>Phase 18: 切换与调度</b><br/>Agent Switch<br/>SubAgent Dispatch<br/>dispatch_all 并发控制<br/>17 新测试"]:::done
P19["<b>Phase 19: 知识图谱</b><br/>KnowledgeGraph trait<br/>InMemoryGraph<br/>双通道检索"]:::pending
P19["<b>Phase 19: 知识图谱</b><br/>KnowledgeGraph trait<br/>InMemoryGraph<br/>双通道检索"]:::done
P15 --> P14
P18 --> P17
@@ -346,4 +364,4 @@ graph BT
| **M12** | Phase 16 | 多轮对话后摘要自动写入 SessionMemory、派生 slot 时摘要正确注入 + 第二轮实施审查 PASS | ✅ 2026-07-10 |
| **M13** | **Phase 17 (rc.1)** | `SessionManager` 创建/recover/replace/子树/销毁集成测试通过、`Checkpointer` checkpoint/rollback/list_checkpoints 验证(`fork` 推迟,按需时引入)| ✅ 2026-07-15 |
| **M14** | Phase 18 | `switch_agent` 热切换验证(slot / turn_index / session_memory 保留)、`dispatch` / `dispatch_all` 并行派发 + Semaphore 顺序、`dispatch_stream` 流事件序列验证 | ✅ 2026-07-15 |
| **M15** | Phase 19 | `KnowledgeGraph` 实体-关系 CRUD + `get_related` BFS 验证、双通道检索 Hybrid 策略验证 | |
| **M15** | Phase 19 | `KnowledgeGraph` 实体-关系 CRUD + `get_related` BFS 验证、双通道检索 Hybrid 策略验证 | ✅ 2026-07-17 |
+2 -2
View File
@@ -1,7 +1,7 @@
# AG Core Roadmap
> 拆分式 roadmap:按版本归档 + 未归类内容
> 最后更新:2026-07-15v0.3.0 Phase 18 完成 + M14 里程碑达成 + 文档拆分
> 最后更新:2026-07-17v0.3.0 Phase 19 完成 + M15 里程碑达成 + v0.3.0 全部交付完毕
## 文件索引
@@ -9,7 +9,7 @@
|------|------|------|
| [`roadmap-v0.1.0.md`](./roadmap-v0.1.0.md) | v0.1.0 计划与交付 — Phase 04c + v0.1.0 Release | ✅ 已发布 2026-07-04 |
| [`roadmap-v0.2.0.md`](./roadmap-v0.2.0.md) | v0.2.0 计划与交付 — Phase 512 + v0.2.0-rc.1 | 🟡 Phase 5-11 已完成;Phase 12 P2 锦上添花可选 |
| [`roadmap-v0.3.0.md`](./roadmap-v0.3.0.md) | v0.3.0 计划与交付 Phase 1319 | 🟡 Phase 13-18 已完成;Phase 19 待实施 |
| [`roadmap-v0.3.0.md`](./roadmap-v0.3.0.md) | v0.3.0 计划与交付 - Phase 1319 | Phase 13-19 全部完成,v0.3.0 交付完毕 |
| [`roadmap-unsorted.md`](./roadmap-unsorted.md) | 未归到任何版本的内容 — 全局愿景、当前状态、模块完整性、v0.4+ 展望、风险与建议、下一步行动、阶段总回顾 | — |
## 阅读建议
+180
View File
@@ -0,0 +1,180 @@
//! knowledge_graph_demo -- 知识图谱 + 双通道检索演示。
//!
//! 演示:
//! 1. 构建 KnowledgeGraph(实体 + 关系)
//! 2. BFS 图遍历(get_related
//! 3. MemoryRetriever 双通道检索(Hybrid / GraphOnly / KnowledgeOnly
//! 4. 标签管理(set_entity_tags / find_tags
//!
//! 运行:`cargo run --example knowledge_graph_demo`
use std::sync::Arc;
use agcore::memory::{
GraphEntity, GraphRelation, InMemoryGraph, InMemoryStore, KnowledgeGraph, KnowledgePage,
KnowledgeStore, MemoryRetriever, MemoryStore, RelationDirection, RetrievalItem,
RetrievalStrategy, RetrieverConfig,
};
use time::OffsetDateTime;
fn make_page(id: &str, title: &str, content: &str) -> KnowledgePage {
let now = OffsetDateTime::now_utc();
KnowledgePage {
id: id.to_string(),
title: title.to_string(),
summary: content.chars().take(40).collect(),
content: content.to_string(),
tags: Vec::new(),
references: Vec::new(),
created_at: now,
updated_at: now,
}
}
#[tokio::main]
async fn main() {
// ── 1. 构建知识图谱 ──
println!("=== 1. 构建知识图谱 ===");
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 from LangChain".to_string();
let mut langsmith = GraphEntity::new("langsmith", "LangSmith", "tool");
langsmith.description = "Tracing and evaluation platform".to_string();
let mut python = GraphEntity::new("python", "Python", "language");
python.description = "Programming language".to_string();
let mut rust = GraphEntity::new("rust", "Rust", "language");
rust.description = "Systems programming language".to_string();
for e in [&langchain, &langgraph, &langsmith, &python, &rust] {
graph.add_entity(e.clone()).await.unwrap();
}
graph
.add_relation(GraphRelation::new("langchain", "langgraph", "includes", 0.9))
.await
.unwrap();
graph
.add_relation(GraphRelation::new("langchain", "langsmith", "includes", 0.7))
.await
.unwrap();
graph
.add_relation(GraphRelation::new("langchain", "python", "built_with", 0.95))
.await
.unwrap();
graph
.add_relation(GraphRelation::new("langgraph", "python", "depends_on", 0.8))
.await
.unwrap();
println!("已添加 5 个实体 + 4 条关系");
// ── 2. BFS 图遍历 ──
println!("\n=== 2. BFS 图遍历:从 LangChain 出发,depth=2 ===");
let related = graph
.get_related("langchain", 2, RelationDirection::Outgoing, None)
.await
.unwrap();
for se in &related {
println!(
" {} (score={:.3}, path={:?})",
se.entity.name, se.score, se.path
);
}
assert!(!related.is_empty(), "应找到关联实体");
// ── 3. 标签管理 ──
println!("\n=== 3. 标签管理 ===");
graph
.set_entity_tags("langchain", vec!["ai".into(), "framework".into(), "llm".into()])
.await
.unwrap();
graph
.set_entity_tags("langgraph", vec!["ai".into(), "agent".into()])
.await
.unwrap();
let tags = graph.find_tags("a").await.unwrap();
println!("前缀 'a' 查找标签: {:?}", tags);
let count = graph.entity_count_by_tag("ai").await.unwrap();
println!("标签 'ai' 下实体数: {}", count);
// ── 4. 双通道检索 ──
println!("\n=== 4. 双通道检索 ===");
let store: Arc<dyn MemoryStore> = Arc::new(InMemoryStore::new());
let ks = KnowledgeStore::new(store);
ks.add_page(make_page(
"p1",
"LangChain 框架介绍",
"LangChain 是用于构建 LLM 应用的开源框架",
))
.await
.unwrap();
ks.add_page(make_page(
"p2",
"Rust 异步编程",
"Rust 异步基于 tokio 与 futures 抽象",
))
.await
.unwrap();
// Hybrid 策略(默认)
let retriever = MemoryRetriever::new(ks, RetrieverConfig::default())
.with_knowledge_graph(graph.clone());
println!("\n--- Hybrid 检索: 'langchain' ---");
let result = retriever.retrieve("langchain").await.unwrap();
println!("策略: {:?}", result.strategy);
for item in &result.items {
match item {
RetrievalItem::KnowledgePage { page, score } => {
println!(" [Store] {} (score={:.3})", page.title, score);
}
RetrievalItem::GraphEntity {
entity, score, path, ..
} => {
println!(
" [Graph] {} (score={:.3}, path={:?})",
entity.name, score, path
);
}
}
}
let has_store = result
.items
.iter()
.any(|i| matches!(i, RetrievalItem::KnowledgePage { .. }));
let has_graph = result
.items
.iter()
.any(|i| matches!(i, RetrievalItem::GraphEntity { .. }));
assert!(has_store, "Hybrid 应有 Store 结果");
assert!(has_graph, "Hybrid 应有 Graph 结果");
// GraphOnly 策略
let store2: Arc<dyn MemoryStore> = Arc::new(InMemoryStore::new());
let ks2 = KnowledgeStore::new(store2);
ks2.add_page(make_page("p1", "LangChain", "LLM framework"))
.await
.unwrap();
let retriever_g = MemoryRetriever::new(ks2, RetrieverConfig::default())
.with_knowledge_graph(graph.clone())
.with_strategy(RetrievalStrategy::GraphOnly);
println!("\n--- GraphOnly 检索: 'langchain' ---");
let result = retriever_g.retrieve("langchain").await.unwrap();
println!("策略: {:?}", result.strategy);
for item in &result.items {
match item {
RetrievalItem::GraphEntity { entity, score, .. } => {
println!(" [Graph] {} (score={:.3})", entity.name, score);
}
RetrievalItem::KnowledgePage { page, score, .. } => {
println!(" [Store] {} (score={:.3})", page.title, score);
}
}
}
assert!(
result.items.iter().all(|i| matches!(i, RetrievalItem::GraphEntity { .. })),
"GraphOnly 应只返回 Graph 结果"
);
println!("\n✓ knowledge_graph_demo 完成");
}
+11 -2
View File
@@ -74,8 +74,15 @@ async fn main() {
let result = retriever.retrieve("Rust 异步").await.unwrap();
println!("query: {}", result.query);
for item in &result.items {
println!(" 命中: {} (score={:.3})", item.page.title, item.score);
assert!((0.0..=1.0).contains(&item.score), "score 应在 [0, 1] 区间");
match item {
agcore::memory::RetrievalItem::KnowledgePage { page, score } => {
println!(" 命中: {} (score={:.3})", page.title, score);
assert!((0.0..=1.0).contains(score), "score 应在 [0, 1] 区间");
}
agcore::memory::RetrievalItem::GraphEntity { entity, score, .. } => {
println!(" 命中实体: {} (score={:.3})", entity.name, score);
}
}
}
assert!(!result.items.is_empty(), "应至少命中一个页面");
@@ -89,6 +96,7 @@ async fn main() {
let cfg = RetrieverConfig {
max_results: 20,
min_score: 0.5,
graph_depth: 2,
};
let retriever2 = MemoryRetriever::new(ks2, cfg);
let result = retriever2.retrieve("完全不相关的火锅配方").await.unwrap();
@@ -111,6 +119,7 @@ async fn main() {
let cfg = RetrieverConfig {
max_results: 2,
min_score: 0.0,
graph_depth: 2,
};
let retriever3 = MemoryRetriever::new(ks3, cfg);
let result = retriever3.retrieve("Rust").await.unwrap();
+6 -1
View File
@@ -2,6 +2,7 @@
pub mod conversation;
pub mod error;
pub mod graph;
pub mod knowledge;
pub mod retriever;
pub mod store;
@@ -12,6 +13,7 @@ pub mod vector_store;
// 高频类型(大多数下游需要)
pub use conversation::{ConversationMemory, ConversationMemoryConfig};
pub use error::MemoryError;
pub use graph::{GraphEntity, GraphRelation, InMemoryGraph, KnowledgeGraph, RelationDirection, ScoredEntity};
pub use knowledge::KnowledgeStore;
pub use retriever::MemoryRetriever;
pub use store::{InMemoryStore, MemoryStore, SqliteStore};
@@ -21,7 +23,10 @@ pub use vector_store::{InMemoryVectorStore, PersistentVectorStore, RagPipeline,
// 低频类型(配置/高级使用)
pub use conversation::MemoryStrategy;
pub use graph::TagConstraints;
pub use knowledge::{KNOWLEDGE_PREFIX, PageIndexEntry};
pub use retriever::{RetrievalResult, RetrieverConfig, ScoredItem};
#[allow(deprecated)]
pub use retriever::ScoredItem;
pub use retriever::{RetrievalItem, RetrievalResult, RetrievalStrategy, RetrieverConfig};
pub use store::{EvictionConfig, EvictionPolicy};
pub use types::{KnowledgePage, MemoryFilter, MemoryItem};
+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);
}
}