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:
@@ -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 6(KnowledgeStore)[高]、Phase 15(VectorStore 模式参考)[低]
|
||||
- **优先级**:P0(v0.3.0 最后一个 Phase)
|
||||
- **预估规模**:约 600 行核心 + 200 行测试
|
||||
|
||||
---
|
||||
|
||||
## 需求分析
|
||||
|
||||
### 功能需求
|
||||
|
||||
| ID | 需求 | 优先级 |
|
||||
|----|------|--------|
|
||||
| F1 | `GraphEntity` / `GraphRelation` / `RelationDirection` 类型定义 | P0 |
|
||||
| F2 | `KnowledgeGraph` trait(10 个 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
@@ -1,6 +1,6 @@
|
||||
# AG Core Roadmap — v0.3.0
|
||||
|
||||
> 本文件聚焦 **v0.3.0 版本** 的规划与交付(Phase 13–19)。Phase 13-18 已完成,Phase 19 待实施。
|
||||
> 本文件聚焦 **v0.3.0 版本** 的规划与交付(Phase 13–19)。Phase 13-19 全部完成,v0.3.0 交付完毕。
|
||||
> 返回总入口:[`roadmap.md`](./roadmap.md)
|
||||
|
||||
## v0.3.0 愿景
|
||||
@@ -9,7 +9,7 @@
|
||||
|
||||
## v0.3.0 总体范围
|
||||
|
||||
**总体规模**:7 个增量 Phase(Phase 13–19),总新增代码约 2600 行,测试从 277 → 380+。已完成 6 个 Phase(M9-M14 已达成),Phase 19 待交付(M15)。
|
||||
**总体规模**:7 个增量 Phase(Phase 13–19),总新增代码约 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 个增量 Phase(Phase 13-19),总新增代码约 2600 行,测试从 277 → 380+。
|
||||
**总体规模**:7 个增量 Phase(Phase 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` trait(10 个 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
@@ -1,7 +1,7 @@
|
||||
# AG Core Roadmap
|
||||
|
||||
> 拆分式 roadmap:按版本归档 + 未归类内容
|
||||
> 最后更新:2026-07-15(v0.3.0 Phase 18 完成 + M14 里程碑达成 + 文档拆分)
|
||||
> 最后更新:2026-07-17(v0.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 0–4c + v0.1.0 Release | ✅ 已发布 2026-07-04 |
|
||||
| [`roadmap-v0.2.0.md`](./roadmap-v0.2.0.md) | v0.2.0 计划与交付 — Phase 5–12 + 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 13–19 | 🟡 Phase 13-18 已完成;Phase 19 待实施 |
|
||||
| [`roadmap-v0.3.0.md`](./roadmap-v0.3.0.md) | v0.3.0 计划与交付 - Phase 13–19 | ✅ Phase 13-19 全部完成,v0.3.0 交付完毕 |
|
||||
| [`roadmap-unsorted.md`](./roadmap-unsorted.md) | 未归到任何版本的内容 — 全局愿景、当前状态、模块完整性、v0.4+ 展望、风险与建议、下一步行动、阶段总回顾 | — |
|
||||
|
||||
## 阅读建议
|
||||
|
||||
@@ -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 完成");
|
||||
}
|
||||
@@ -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
@@ -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
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