diff --git a/docs/25-phase19-knowledge-graph-and-retrieval.md b/docs/25-phase19-knowledge-graph-and-retrieval.md new file mode 100644 index 0000000..242b468 --- /dev/null +++ b/docs/25-phase19-knowledge-graph-and-retrieval.md @@ -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`,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, + /// 任意附加属性(与 PersistentVectorStore.metadata 保持一致)。 + pub properties: HashMap, +} +``` + +#### 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, +} +``` + +#### 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, 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, MemoryError>; + + // ── 检索 ── + + /// 按关键词子串匹配实体(不区分大小写),与 KnowledgeStore.search 一致。 + async fn find_by_keywords(&self, keywords: &[String]) -> Result, MemoryError>; + + // ── 标签管理(预留接口,Agent 层显式写入) ── + + /// 按前缀查找已有标签(用于标签复用)。 + async fn find_tags(&self, prefix: &str) -> Result, MemoryError>; + + /// 设置实体标签(替换式,保留前 max_tags_per_entity 个)。 + /// 返回实际设置的标签数。 + async fn set_entity_tags( + &self, + entity_id: &str, + tags: Vec, + ) -> Result; + + /// 按标签统计实体数量。 + async fn entity_count_by_tag(&self, tag: &str) -> Result; + + /// 获取标签约束配置。 + fn tag_constraints(&self) -> TagConstraints; +} +``` + +### 3.3 InMemoryGraph 实现 + +#### 内部结构 + +```rust +/// 内存知识图谱实现 —— 纯内存,无持久化。 +/// +/// 生命周期跟随实例;持久化路径参考 InMemoryVectorStore → PersistentVectorStore 演进模式。 +pub struct InMemoryGraph { + /// 内部状态(单一锁结构,避免嵌套锁死锁) + inner: Mutex, + /// 标签约束 + constraints: TagConstraints, +} + +struct GraphInner { + /// id → entity + entities: HashMap, + /// 所有关系(线性扫描,实测 5000 条 ≈ 1-50µs,无需邻接表索引) + relations: Vec, + /// tag → entity_ids(反向索引,用于 find_tags / entity_count_by_tag) + tag_index: HashMap>, +} +``` + +#### BFS 遍历算法 + +```rust +async fn get_related( + &self, + entity_id: &str, + depth: usize, + direction: RelationDirection, + relation_types: Option<&[&str]>, +) -> Result, 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 = HashSet::new(); + let mut result: Vec = Vec::new(); + // 队列:(entity_id, score, path) + let mut queue: VecDeque<(String, f32, Vec)> = 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)> = 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` 已访问集合防环 | +| 评分衰减 | 沿路径 `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, + }, +} + +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, + pub query: String, + /// 本次检索实际执行的策略(可能因 graph 未注入而退化),而非用户通过 `with_strategy()` 配置的值。 + pub strategy: RetrievalStrategy, +} +``` + +#### MemoryRetriever 扩展 + +```rust +pub struct MemoryRetriever { + knowledge_store: KnowledgeStore, + /// 可选知识图谱(None 时退化为单通道)。 + knowledge_graph: Option>, + /// 检索策略(默认 Hybrid)。 + strategy: RetrievalStrategy, + config: RetrieverConfig, + stop_words: HashSet, +} + +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) -> 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 { + 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>` 维护 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, +) -> Result { + 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 = 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 | diff --git a/docs/roadmap-v0.3.0.md b/docs/roadmap-v0.3.0.md index f1b500a..e08318c 100644 --- a/docs/roadmap-v0.3.0.md +++ b/docs/roadmap-v0.3.0.md @@ -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` 单一锁结构,避免嵌套锁死锁)+ 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` 单一锁结构** - 避免 `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["Phase 16: 摘要自动生成
SummaryConfig
内联检查点
首次防抖跳过
18 新测试"]:::done P17["Phase 17: 执行引擎
SessionManager
会话树
Time-travel Checkpointer
21 新测试"]:::done P18["Phase 18: 切换与调度
Agent Switch
SubAgent Dispatch
dispatch_all 并发控制
17 新测试"]:::done - P19["Phase 19: 知识图谱
KnowledgeGraph trait
InMemoryGraph
双通道检索"]:::pending + P19["Phase 19: 知识图谱
KnowledgeGraph trait
InMemoryGraph
双通道检索"]:::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 | diff --git a/docs/roadmap.md b/docs/roadmap.md index f1f91a9..e9ba643 100644 --- a/docs/roadmap.md +++ b/docs/roadmap.md @@ -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+ 展望、风险与建议、下一步行动、阶段总回顾 | — | ## 阅读建议 diff --git a/examples/knowledge_graph_demo.rs b/examples/knowledge_graph_demo.rs new file mode 100644 index 0000000..2b6dfe1 --- /dev/null +++ b/examples/knowledge_graph_demo.rs @@ -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 = 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 = 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 完成"); +} diff --git a/examples/knowledge_search_demo.rs b/examples/knowledge_search_demo.rs index d101c9f..c79f860 100644 --- a/examples/knowledge_search_demo.rs +++ b/examples/knowledge_search_demo.rs @@ -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(); diff --git a/src/memory.rs b/src/memory.rs index 177801f..f7c3564 100644 --- a/src/memory.rs +++ b/src/memory.rs @@ -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}; diff --git a/src/memory/graph.rs b/src/memory/graph.rs new file mode 100644 index 0000000..4ac71bf --- /dev/null +++ b/src/memory/graph.rs @@ -0,0 +1,1080 @@ +//! 知识图谱 -- 实体-关系存储与图遍历检索。 +//! +//! 提供 [`KnowledgeGraph`] trait 定义和内存实现 [`InMemoryGraph`], +//! 支持 BFS 图遍历、关键词检索和标签管理。 +//! +//! 与 [`KnowledgeStore`](crate::memory::knowledge::KnowledgeStore)(页面级内容)互补, +//! 提供实体级 + 关系维度的检索能力。 + +use std::collections::{HashMap, HashSet, VecDeque}; +use std::sync::Mutex; + +use async_trait::async_trait; +use serde::{Deserialize, Serialize}; + +use crate::memory::error::MemoryError; + +// ───────────────────────── 核心数据类型 ───────────────────────── + +/// 图谱实体 -- 表示一个可被关联检索的节点。 +#[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, + /// 任意附加属性(与 PersistentVectorStore.metadata 保持一致)。 + pub properties: HashMap, +} + +impl GraphEntity { + /// 创建一个最小实体(仅 id + name + type,其余为空)。 + pub fn new(id: impl Into, name: impl Into, entity_type: impl Into) -> Self { + Self { + id: id.into(), + name: name.into(), + entity_type: entity_type.into(), + description: String::new(), + tags: Vec::new(), + properties: HashMap::new(), + } + } +} + +/// 图谱关系 -- 连接两个实体的有向边。 +/// +/// 无 `id` 字段,用 `(source_id, target_id, relation_type)` 三元组唯一标识。 +/// 提供 [`composite_key`](Self::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 { + /// 创建一条新关系。 + pub fn new( + source_id: impl Into, + target_id: impl Into, + relation_type: impl Into, + weight: f32, + ) -> Self { + Self { + source_id: source_id.into(), + target_id: target_id.into(), + relation_type: relation_type.into(), + weight, + } + } + + /// 复合键:`source_id:target_id:relation_type`,用于去重和查找。 + pub fn composite_key(&self) -> String { + format!("{}:{}:{}", self.source_id, self.target_id, self.relation_type) + } +} + +/// 关系遍历方向。 +#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)] +pub enum RelationDirection { + /// 仅出边:source_id -> target_id(默认)。 + #[default] + Outgoing, + /// 仅入边:target_id -> source_id。 + Incoming, + /// 双向遍历。 + Both, +} + +/// 带评分的实体 + 路径信息。 +#[derive(Debug, Clone)] +pub struct ScoredEntity { + pub entity: GraphEntity, + /// 基于图距离的评分 [0.0, 1.0],沿路径权重乘积衰减。 + pub score: f32, + /// 从查询实体到当前实体的 ID 路径(用于可解释性)。 + pub path: Vec, +} + +/// 标签约束配置。 +#[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, + } + } +} + +// ───────────────────────── KnowledgeGraph trait ───────────────────────── + +/// 知识图谱抽象 -- 实体-关系存储与图遍历检索。 +/// +/// 所有方法 `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, 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`:可选过滤,仅遍历指定关系类型 + /// + /// `depth=0` 时仅验证起点实体存在,返回空关联列表(不遍历任何边)。 + async fn get_related( + &self, + entity_id: &str, + depth: usize, + direction: RelationDirection, + relation_types: Option<&[&str]>, + ) -> Result, MemoryError>; + + // ── 检索 ── + + /// 按关键词子串匹配实体(不区分大小写),与 KnowledgeStore.search 一致。 + async fn find_by_keywords(&self, keywords: &[String]) -> Result, MemoryError>; + + // ── 标签管理(预留接口,Agent 层显式写入) ── + + /// 按前缀查找已有标签(用于标签复用)。 + async fn find_tags(&self, prefix: &str) -> Result, MemoryError>; + + /// 设置实体标签(替换式,保留前 max_tags_per_entity 个)。 + /// 返回实际设置的标签数。 + async fn set_entity_tags( + &self, + entity_id: &str, + tags: Vec, + ) -> Result; + + /// 按标签统计实体数量。 + async fn entity_count_by_tag(&self, tag: &str) -> Result; + + /// 获取标签约束配置。 + fn tag_constraints(&self) -> TagConstraints; +} + +// ───────────────────────── InMemoryGraph 实现 ───────────────────────── + +/// 内部状态(单一锁结构,避免嵌套锁死锁)。 +struct GraphInner { + /// id -> entity + entities: HashMap, + /// 所有关系(线性扫描,5000 条估算约 1-50µs,无需邻接表索引) + relations: Vec, + /// tag -> entity_ids(反向索引,用于 find_tags / entity_count_by_tag) + tag_index: HashMap>, +} + +/// 内存知识图谱实现 -- 纯内存,无持久化。 +/// +/// 生命周期跟随实例;持久化路径参考 InMemoryVectorStore -> PersistentVectorStore 演进模式。 +pub struct InMemoryGraph { + /// 内部状态(单一锁结构,避免嵌套锁死锁) + inner: Mutex, + /// 标签约束 + constraints: TagConstraints, +} + +impl Default for InMemoryGraph { + fn default() -> Self { + Self::new() + } +} + +impl InMemoryGraph { + /// 创建一个空的内存知识图谱。 + pub fn new() -> Self { + Self { + inner: Mutex::new(GraphInner { + entities: HashMap::new(), + relations: Vec::new(), + tag_index: HashMap::new(), + }), + constraints: TagConstraints::default(), + } + } + + /// 使用自定义标签约束创建图谱。 + pub fn with_constraints(constraints: TagConstraints) -> Self { + Self { + inner: Mutex::new(GraphInner { + entities: HashMap::new(), + relations: Vec::new(), + tag_index: HashMap::new(), + }), + constraints, + } + } +} + +#[async_trait] +impl KnowledgeGraph for InMemoryGraph { + async fn add_entity(&self, entity: GraphEntity) -> Result<(), MemoryError> { + if entity.id.is_empty() { + return Err(MemoryError::InvalidInput("entity id is empty".into())); + } + let mut inner = self.inner.lock().unwrap(); + // upsert:若已存在,先收集旧标签用于清理反向引用(避免同时借用 entities 和 tag_index) + let old_tags: Vec = inner + .entities + .get(&entity.id) + .map(|old| old.tags.clone()) + .unwrap_or_default(); + for old_tag in &old_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); + } + } + } + // 写入新标签的反向引用 + for tag in &entity.tags { + inner + .tag_index + .entry(tag.clone()) + .or_default() + .insert(entity.id.clone()); + } + inner.entities.insert(entity.id.clone(), entity); + Ok(()) + } + + async fn get_entity(&self, id: &str) -> Result, MemoryError> { + let inner = self.inner.lock().unwrap(); + Ok(inner.entities.get(id).cloned()) + } + + async fn remove_entity(&self, id: &str) -> Result<(), MemoryError> { + let mut inner = self.inner.lock().unwrap(); + // 移除实体并清理其标签反向引用 + if let Some(entity) = inner.entities.remove(id) { + for tag in &entity.tags { + if let Some(ids) = inner.tag_index.get_mut(tag) { + ids.remove(id); + if ids.is_empty() { + inner.tag_index.remove(tag); + } + } + } + } + // 移除所有涉及该实体的关系 + inner + .relations + .retain(|r| r.source_id != id && r.target_id != id); + Ok(()) + } + + async fn add_relation(&self, relation: GraphRelation) -> Result<(), MemoryError> { + let mut inner = self.inner.lock().unwrap(); + // 校验两端实体存在 + if !inner.entities.contains_key(&relation.source_id) { + return Err(MemoryError::InvalidInput(format!( + "source entity '{}' not found", + relation.source_id + ))); + } + if !inner.entities.contains_key(&relation.target_id) { + return Err(MemoryError::InvalidInput(format!( + "target entity '{}' not found", + relation.target_id + ))); + } + // upsert:若复合键已存在则覆盖 + let key = relation.composite_key(); + if let Some(existing) = inner + .relations + .iter_mut() + .find(|r| r.composite_key() == key) + { + existing.weight = relation.weight; + } else { + inner.relations.push(relation); + } + Ok(()) + } + + async fn remove_relation( + &self, + source_id: &str, + target_id: &str, + relation_type: &str, + ) -> Result<(), MemoryError> { + let mut inner = self.inner.lock().unwrap(); + inner.relations.retain(|r| { + !(r.source_id == source_id && r.target_id == target_id && r.relation_type == relation_type) + }); + Ok(()) + } + + async fn get_related( + &self, + entity_id: &str, + depth: usize, + direction: RelationDirection, + relation_types: Option<&[&str]>, + ) -> Result, MemoryError> { + let inner = self.inner.lock().unwrap(); + + // 1. 验证起点存在 + if !inner.entities.contains_key(entity_id) { + return Err(MemoryError::NotFound(entity_id.to_string())); + } + + // depth=0 时仅验证起点,返回空列表 + if depth == 0 { + return Ok(Vec::new()); + } + + // 2. BFS 初始化 + let mut visited: HashSet = HashSet::new(); + let mut result: Vec = Vec::new(); + // 队列:(entity_id, score, path) + let mut queue: VecDeque<(String, f32, Vec)> = 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)> = VecDeque::new(); + + while let Some((current_id, score, path)) = queue.pop_front() { + for rel in inner.relations.iter() { + // 方向过滤 + 邻居解析 + let neighbor_id = match direction { + RelationDirection::Outgoing => { + if rel.source_id == current_id { + Some(rel.target_id.clone()) + } else { + None + } + } + RelationDirection::Incoming => { + if rel.target_id == current_id { + Some(rel.source_id.clone()) + } else { + None + } + } + RelationDirection::Both => { + if rel.source_id == current_id { + Some(rel.target_id.clone()) + } else if rel.target_id == current_id { + Some(rel.source_id.clone()) + } else { + None + } + } + }; + + let Some(neighbor_id) = neighbor_id else { + continue; + }; + + // 关系类型过滤 + if let Some(types) = relation_types + && !types.contains(&rel.relation_type.as_str()) + { + continue; + } + + 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()); + + let entity = inner + .entities + .get(&neighbor_id) + .cloned() + .ok_or_else(|| MemoryError::NotFound(neighbor_id.clone()))?; + + result.push(ScoredEntity { + entity, + 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) + } + + async fn find_by_keywords(&self, keywords: &[String]) -> Result, MemoryError> { + let inner = self.inner.lock().unwrap(); + if keywords.is_empty() { + return Ok(Vec::new()); + } + let lowered: Vec = keywords.iter().map(|k| k.to_lowercase()).collect(); + let mut seen: HashSet = HashSet::new(); + let mut results: Vec = Vec::new(); + for entity in inner.entities.values() { + let name_l = entity.name.to_lowercase(); + let desc_l = entity.description.to_lowercase(); + let tags_l: Vec = entity.tags.iter().map(|t| t.to_lowercase()).collect(); + // 任一关键词子串命中即返回(与 KnowledgeStore.search 一致) + let hit = lowered.iter().any(|k| { + name_l.contains(k) || desc_l.contains(k) || tags_l.iter().any(|t| t.contains(k)) + }); + if hit && seen.insert(entity.id.clone()) { + results.push(entity.clone()); + } + } + Ok(results) + } + + async fn find_tags(&self, prefix: &str) -> Result, MemoryError> { + let inner = self.inner.lock().unwrap(); + let prefix_l = prefix.to_lowercase(); + let mut tags: Vec = inner + .tag_index + .keys() + .filter(|t| t.to_lowercase().starts_with(&prefix_l)) + .cloned() + .collect(); + tags.sort(); + Ok(tags) + } + + async fn set_entity_tags( + &self, + entity_id: &str, + tags: Vec, + ) -> Result { + let mut inner = self.inner.lock().unwrap(); + + // 先收集旧标签(避免同时借用 entities 和 tag_index) + let old_tags: Vec = inner + .entities + .get(entity_id) + .map(|e| e.tags.clone()) + .ok_or_else(|| MemoryError::NotFound(entity_id.to_string()))?; + + // 清除旧标签的反向引用 + for old_tag in &old_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 = 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 + inner.entities.get_mut(entity_id).unwrap().tags = new_tags.clone(); + Ok(new_tags.len()) + } + + async fn entity_count_by_tag(&self, tag: &str) -> Result { + let inner = self.inner.lock().unwrap(); + Ok(inner + .tag_index + .get(tag) + .map_or(0, |ids| ids.len())) + } + + fn tag_constraints(&self) -> TagConstraints { + self.constraints.clone() + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn make_entity(id: &str, name: &str, entity_type: &str) -> GraphEntity { + GraphEntity::new(id, name, entity_type) + } + + fn make_entity_with_desc(id: &str, name: &str, desc: &str) -> GraphEntity { + let mut e = GraphEntity::new(id, name, "concept"); + e.description = desc.to_string(); + e + } + + // ── 实体管理 ── + + #[tokio::test] + async fn add_get_entity() { + let graph = InMemoryGraph::new(); + graph + .add_entity(make_entity("p1", "Person One", "person")) + .await + .unwrap(); + let got = graph.get_entity("p1").await.unwrap(); + assert!(got.is_some()); + assert_eq!(got.unwrap().name, "Person One"); + } + + #[tokio::test] + async fn add_entity_rejects_empty_id() { + let graph = InMemoryGraph::new(); + let result = graph.add_entity(make_entity("", "NoId", "x")).await; + assert!(matches!(result, Err(MemoryError::InvalidInput(_)))); + } + + #[tokio::test] + async fn get_entity_not_found() { + let graph = InMemoryGraph::new(); + let got = graph.get_entity("ghost").await.unwrap(); + assert!(got.is_none()); + } + + #[tokio::test] + async fn remove_entity_cleans_relations() { + let graph = InMemoryGraph::new(); + graph + .add_entity(make_entity("a", "A", "concept")) + .await + .unwrap(); + graph + .add_entity(make_entity("b", "B", "concept")) + .await + .unwrap(); + graph + .add_relation(GraphRelation::new("a", "b", "related_to", 0.5)) + .await + .unwrap(); + + graph.remove_entity("a").await.unwrap(); + assert!(graph.get_entity("a").await.unwrap().is_none()); + + // 删除实体后,相关关系也应被清理 + let related = graph + .get_related("b", 1, RelationDirection::Both, None) + .await + .unwrap(); + assert!(related.is_empty(), "relation should be removed with entity"); + } + + #[tokio::test] + async fn add_entity_upsert_cleans_old_tags() { + let graph = InMemoryGraph::new(); + let mut e = make_entity("p1", "P1", "person"); + e.tags = vec!["old_tag".to_string()]; + graph.add_entity(e).await.unwrap(); + assert_eq!(graph.entity_count_by_tag("old_tag").await.unwrap(), 1); + + // upsert:新版本不带 old_tag + let e2 = make_entity("p1", "P1-updated", "person"); + graph.add_entity(e2).await.unwrap(); + assert_eq!( + graph.entity_count_by_tag("old_tag").await.unwrap(), + 0, + "old tag index should be cleaned on upsert" + ); + } + + // ── 关系管理 ── + + #[tokio::test] + async fn add_relation_validates_entities() { + let graph = InMemoryGraph::new(); + graph + .add_entity(make_entity("a", "A", "x")) + .await + .unwrap(); + // target 不存在 + let result = graph + .add_relation(GraphRelation::new("a", "ghost", "r", 0.5)) + .await; + assert!(matches!(result, Err(MemoryError::InvalidInput(_)))); + } + + #[tokio::test] + async fn add_relation_upsert_overrides_weight() { + let graph = InMemoryGraph::new(); + graph + .add_entity(make_entity("a", "A", "x")) + .await + .unwrap(); + graph + .add_entity(make_entity("b", "B", "x")) + .await + .unwrap(); + graph + .add_relation(GraphRelation::new("a", "b", "r", 0.5)) + .await + .unwrap(); + graph + .add_relation(GraphRelation::new("a", "b", "r", 0.9)) + .await + .unwrap(); + let related = graph + .get_related("a", 1, RelationDirection::Outgoing, None) + .await + .unwrap(); + assert_eq!(related.len(), 1); + assert!((related[0].score - 0.9).abs() < 0.001); + } + + #[tokio::test] + async fn remove_relation() { + let graph = InMemoryGraph::new(); + graph + .add_entity(make_entity("a", "A", "x")) + .await + .unwrap(); + graph + .add_entity(make_entity("b", "B", "x")) + .await + .unwrap(); + graph + .add_relation(GraphRelation::new("a", "b", "r", 0.5)) + .await + .unwrap(); + graph.remove_relation("a", "b", "r").await.unwrap(); + let related = graph + .get_related("a", 1, RelationDirection::Outgoing, None) + .await + .unwrap(); + assert!(related.is_empty()); + } + + // ── BFS 遍历 ── + + #[tokio::test] + async fn get_related_empty_graph() { + let graph = InMemoryGraph::new(); + graph + .add_entity(make_entity("a", "A", "x")) + .await + .unwrap(); + let related = graph + .get_related("a", 3, RelationDirection::Outgoing, None) + .await + .unwrap(); + assert!(related.is_empty()); + } + + #[tokio::test] + async fn get_related_depth_0() { + let graph = InMemoryGraph::new(); + graph + .add_entity(make_entity("a", "A", "x")) + .await + .unwrap(); + let related = graph + .get_related("a", 0, RelationDirection::Outgoing, None) + .await + .unwrap(); + assert!(related.is_empty(), "depth=0 should return empty list"); + } + + #[tokio::test] + async fn get_related_not_found() { + let graph = InMemoryGraph::new(); + let result = graph + .get_related("ghost", 1, RelationDirection::Outgoing, None) + .await; + assert!(matches!(result, Err(MemoryError::NotFound(_)))); + } + + #[tokio::test] + async fn get_related_linear_chain_depth_1() { + let graph = InMemoryGraph::new(); + for id in ["a", "b", "c", "d"] { + graph.add_entity(make_entity(id, id, "x")).await.unwrap(); + } + graph + .add_relation(GraphRelation::new("a", "b", "r", 0.5)) + .await + .unwrap(); + graph + .add_relation(GraphRelation::new("b", "c", "r", 0.5)) + .await + .unwrap(); + graph + .add_relation(GraphRelation::new("c", "d", "r", 0.5)) + .await + .unwrap(); + + let related = graph + .get_related("a", 1, RelationDirection::Outgoing, None) + .await + .unwrap(); + assert_eq!(related.len(), 1); + assert_eq!(related[0].entity.id, "b"); + assert!((related[0].score - 0.5).abs() < 0.001); + } + + #[tokio::test] + async fn get_related_linear_chain_depth_3() { + let graph = InMemoryGraph::new(); + for id in ["a", "b", "c", "d"] { + graph.add_entity(make_entity(id, id, "x")).await.unwrap(); + } + graph + .add_relation(GraphRelation::new("a", "b", "r", 0.5)) + .await + .unwrap(); + graph + .add_relation(GraphRelation::new("b", "c", "r", 0.5)) + .await + .unwrap(); + graph + .add_relation(GraphRelation::new("c", "d", "r", 0.5)) + .await + .unwrap(); + + let related = graph + .get_related("a", 3, RelationDirection::Outgoing, None) + .await + .unwrap(); + assert_eq!(related.len(), 3); + // 降序排列:b (0.5) > c (0.25) > d (0.125) + assert_eq!(related[0].entity.id, "b"); + assert_eq!(related[1].entity.id, "c"); + assert_eq!(related[2].entity.id, "d"); + } + + #[tokio::test] + async fn get_related_cycle() { + let graph = InMemoryGraph::new(); + for id in ["a", "b", "c"] { + graph.add_entity(make_entity(id, id, "x")).await.unwrap(); + } + // 环 a -> b -> c -> a + graph + .add_relation(GraphRelation::new("a", "b", "r", 0.5)) + .await + .unwrap(); + graph + .add_relation(GraphRelation::new("b", "c", "r", 0.5)) + .await + .unwrap(); + graph + .add_relation(GraphRelation::new("c", "a", "r", 0.5)) + .await + .unwrap(); + + let related = graph + .get_related("a", 5, RelationDirection::Outgoing, None) + .await + .unwrap(); + // 应该返回 b 和 c,不会重复返回 a + assert_eq!(related.len(), 2); + let ids: Vec<&str> = related.iter().map(|s| s.entity.id.as_str()).collect(); + assert!(ids.contains(&"b")); + assert!(ids.contains(&"c")); + assert!(!ids.contains(&"a"), "start node should not be revisited"); + } + + #[tokio::test] + async fn get_related_star_topology() { + let graph = InMemoryGraph::new(); + graph.add_entity(make_entity("center", "C", "x")).await.unwrap(); + for i in 0..5 { + let leaf = format!("leaf{i}"); + graph.add_entity(make_entity(&leaf, &leaf, "x")).await.unwrap(); + graph + .add_relation(GraphRelation::new("center", &leaf, "r", 0.8)) + .await + .unwrap(); + } + let related = graph + .get_related("center", 1, RelationDirection::Outgoing, None) + .await + .unwrap(); + assert_eq!(related.len(), 5); + // 所有叶子节点的 score 都应该是 0.8 + for s in &related { + assert!((s.score - 0.8).abs() < 0.001); + } + } + + #[tokio::test] + async fn get_related_isolated_components() { + let graph = InMemoryGraph::new(); + // 分量 1: a -> b + graph.add_entity(make_entity("a", "A", "x")).await.unwrap(); + graph.add_entity(make_entity("b", "B", "x")).await.unwrap(); + graph + .add_relation(GraphRelation::new("a", "b", "r", 0.5)) + .await + .unwrap(); + // 分量 2: c -> d + graph.add_entity(make_entity("c", "C", "x")).await.unwrap(); + graph.add_entity(make_entity("d", "D", "x")).await.unwrap(); + graph + .add_relation(GraphRelation::new("c", "d", "r", 0.5)) + .await + .unwrap(); + + let related = graph + .get_related("a", 3, RelationDirection::Both, None) + .await + .unwrap(); + assert_eq!(related.len(), 1); + assert_eq!(related[0].entity.id, "b"); + } + + #[tokio::test] + async fn get_related_direction_incoming() { + let graph = InMemoryGraph::new(); + graph.add_entity(make_entity("a", "A", "x")).await.unwrap(); + graph.add_entity(make_entity("b", "B", "x")).await.unwrap(); + // a -> b + graph + .add_relation(GraphRelation::new("a", "b", "r", 0.5)) + .await + .unwrap(); + // 从 b 看 Incoming 方向应该找到 a + let related = graph + .get_related("b", 1, RelationDirection::Incoming, None) + .await + .unwrap(); + assert_eq!(related.len(), 1); + assert_eq!(related[0].entity.id, "a"); + } + + #[tokio::test] + async fn get_related_direction_both() { + let graph = InMemoryGraph::new(); + graph.add_entity(make_entity("a", "A", "x")).await.unwrap(); + graph.add_entity(make_entity("b", "B", "x")).await.unwrap(); + graph.add_entity(make_entity("c", "C", "x")).await.unwrap(); + // a -> b(出边) + graph + .add_relation(GraphRelation::new("a", "b", "r", 0.5)) + .await + .unwrap(); + // c -> a(入边) + graph + .add_relation(GraphRelation::new("c", "a", "r", 0.7)) + .await + .unwrap(); + + let related = graph + .get_related("a", 1, RelationDirection::Both, None) + .await + .unwrap(); + assert_eq!(related.len(), 2); + // 降序:c (0.7) > b (0.5) + assert_eq!(related[0].entity.id, "c"); + assert_eq!(related[1].entity.id, "b"); + } + + #[tokio::test] + async fn get_related_filter_by_type() { + let graph = InMemoryGraph::new(); + graph.add_entity(make_entity("a", "A", "x")).await.unwrap(); + graph.add_entity(make_entity("b", "B", "x")).await.unwrap(); + graph.add_entity(make_entity("c", "C", "x")).await.unwrap(); + graph + .add_relation(GraphRelation::new("a", "b", "knows", 0.5)) + .await + .unwrap(); + graph + .add_relation(GraphRelation::new("a", "c", "works_with", 0.5)) + .await + .unwrap(); + + let related = graph + .get_related("a", 1, RelationDirection::Outgoing, Some(&["knows"])) + .await + .unwrap(); + assert_eq!(related.len(), 1); + assert_eq!(related[0].entity.id, "b"); + } + + // ── 关键词检索 ── + + #[tokio::test] + async fn find_by_keywords_matches_name() { + let graph = InMemoryGraph::new(); + graph + .add_entity(make_entity("a", "LangGraph", "framework")) + .await + .unwrap(); + graph + .add_entity(make_entity("b", "Other", "tool")) + .await + .unwrap(); + let results = graph + .find_by_keywords(&["langgraph".to_string()]) + .await + .unwrap(); + assert_eq!(results.len(), 1); + assert_eq!(results[0].id, "a"); + } + + #[tokio::test] + async fn find_by_keywords_matches_description_case_insensitive() { + let graph = InMemoryGraph::new(); + graph + .add_entity(make_entity_with_desc("a", "A", "Framework for graphs")) + .await + .unwrap(); + let results = graph + .find_by_keywords(&["FRAMEWORK".to_string()]) + .await + .unwrap(); + assert_eq!(results.len(), 1); + } + + #[tokio::test] + async fn find_by_keywords_empty_input() { + let graph = InMemoryGraph::new(); + graph.add_entity(make_entity("a", "A", "x")).await.unwrap(); + let results = graph.find_by_keywords(&[]).await.unwrap(); + assert!(results.is_empty()); + } + + // ── 标签管理 ── + + #[tokio::test] + async fn set_and_find_tags() { + let graph = InMemoryGraph::new(); + graph + .add_entity(make_entity("a", "A", "x")) + .await + .unwrap(); + let n = graph + .set_entity_tags("a", vec!["rust".into(), "ai".into()]) + .await + .unwrap(); + assert_eq!(n, 2); + + let tags = graph.find_tags("ru").await.unwrap(); + assert_eq!(tags, vec!["rust".to_string()]); + } + + #[tokio::test] + async fn set_entity_tags_replaces_old() { + let graph = InMemoryGraph::new(); + graph.add_entity(make_entity("a", "A", "x")).await.unwrap(); + graph + .set_entity_tags("a", vec!["old1".into(), "old2".into()]) + .await + .unwrap(); + graph + .set_entity_tags("a", vec!["new1".into()]) + .await + .unwrap(); + assert_eq!(graph.entity_count_by_tag("old1").await.unwrap(), 0); + assert_eq!(graph.entity_count_by_tag("old2").await.unwrap(), 0); + assert_eq!(graph.entity_count_by_tag("new1").await.unwrap(), 1); + } + + #[tokio::test] + async fn set_entity_tags_truncates_to_max() { + let constraints = TagConstraints { + max_tags_per_entity: 2, + }; + let graph = InMemoryGraph::with_constraints(constraints); + graph.add_entity(make_entity("a", "A", "x")).await.unwrap(); + let n = graph + .set_entity_tags("a", vec!["t1".into(), "t2".into(), "t3".into(), "t4".into()]) + .await + .unwrap(); + assert_eq!(n, 2, "should truncate to max_tags_per_entity"); + } + + #[tokio::test] + async fn entity_count_by_tag_multi_entities() { + let graph = InMemoryGraph::new(); + graph.add_entity(make_entity("a", "A", "x")).await.unwrap(); + graph.add_entity(make_entity("b", "B", "x")).await.unwrap(); + graph.add_entity(make_entity("c", "C", "x")).await.unwrap(); + graph + .set_entity_tags("a", vec!["rust".into()]) + .await + .unwrap(); + graph + .set_entity_tags("b", vec!["rust".into(), "ai".into()]) + .await + .unwrap(); + graph + .set_entity_tags("c", vec!["ai".into()]) + .await + .unwrap(); + assert_eq!(graph.entity_count_by_tag("rust").await.unwrap(), 2); + assert_eq!(graph.entity_count_by_tag("ai").await.unwrap(), 2); + } + + #[tokio::test] + async fn set_entity_tags_not_found() { + let graph = InMemoryGraph::new(); + let result = graph + .set_entity_tags("ghost", vec!["x".into()]) + .await; + assert!(matches!(result, Err(MemoryError::NotFound(_)))); + } + + #[tokio::test] + async fn composite_key_format() { + let r = GraphRelation::new("a", "b", "knows", 0.5); + assert_eq!(r.composite_key(), "a:b:knows"); + } +} diff --git a/src/memory/retriever.rs b/src/memory/retriever.rs index 5abec8f..6f69130 100644 --- a/src/memory/retriever.rs +++ b/src/memory/retriever.rs @@ -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, + }, +} + +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, + /// 统一条目列表,按分数降序排列。 + pub items: Vec, pub query: String, + /// 本次检索实际执行的策略(可能因 graph 未注入而退化),而非用户通过 `with_strategy()` 配置的值。 + pub strategy: RetrievalStrategy, } -/// 记忆检索器 —— 在 `KnowledgeStore` 中做关键词检索并按 TextOverlap 评分。 +/// 记忆检索器 -- 在 `KnowledgeStore` 中做关键词检索并按 TextOverlap 评分, +/// 可选注入 `KnowledgeGraph` 启用双通道检索。 pub struct MemoryRetriever { knowledge_store: KnowledgeStore, + /// 可选知识图谱(None 时退化为单通道)。 + knowledge_graph: Option>, + /// 检索策略(默认 Hybrid)。 + strategy: RetrievalStrategy, config: RetrieverConfig, /// 停用词表(用于关键词提取)。 stop_words: HashSet, } 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) -> 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) -> Self { self.stop_words = stop_words; self } - /// 检索相关知识页面。 + /// 检索相关记忆(双通道)。 + /// + /// 根据 `strategy` 和是否注入 `knowledge_graph` 分流: + /// - `KnowledgeOnly` 或未注入 graph -> 仅检索 KnowledgeStore + /// - `GraphOnly` -> 仅检索 KnowledgeGraph + /// - `Hybrid` -> 并行检索两通道,合并排序 pub async fn retrieve(&self, query: &str) -> Result { 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, 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 = pages + let mut items: Vec = 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, + ) -> Result, 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 = HashSet::new(); + let mut items: Vec = 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 = 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) -> Vec { query .split(|c: char| !c.is_alphanumeric()) @@ -178,6 +360,7 @@ fn default_stop_words() -> HashSet { #[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, 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; - 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; - 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; - 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 { + 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; + 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; + 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); + } }