feat(memory): 实现知识图谱与双通道检索

- 新增 KnowledgeGraph trait + InMemoryGraph(BFS 图遍历 + 标签管理)
- 扩展 MemoryRetriever 为双通道检索(Hybrid/KnowledgeOnly/GraphOnly)
- 统一 RetrievalResult 为 RetrievalItem enum 变体
- GraphRelation 使用复合键 + composite_key() 派生 id
- 旧 ScoredItem 标注 #[deprecated]
- 新增 knowledge_graph_demo 示例
- 全量测试 427 passed,clippy/doc 0 警告
This commit is contained in:
徐涛
2026-07-17 14:37:49 +08:00
parent 5bb349d177
commit 7e72e102a2
8 changed files with 2365 additions and 43 deletions
+180
View File
@@ -0,0 +1,180 @@
//! knowledge_graph_demo -- 知识图谱 + 双通道检索演示。
//!
//! 演示:
//! 1. 构建 KnowledgeGraph(实体 + 关系)
//! 2. BFS 图遍历(get_related
//! 3. MemoryRetriever 双通道检索(Hybrid / GraphOnly / KnowledgeOnly
//! 4. 标签管理(set_entity_tags / find_tags
//!
//! 运行:`cargo run --example knowledge_graph_demo`
use std::sync::Arc;
use agcore::memory::{
GraphEntity, GraphRelation, InMemoryGraph, InMemoryStore, KnowledgeGraph, KnowledgePage,
KnowledgeStore, MemoryRetriever, MemoryStore, RelationDirection, RetrievalItem,
RetrievalStrategy, RetrieverConfig,
};
use time::OffsetDateTime;
fn make_page(id: &str, title: &str, content: &str) -> KnowledgePage {
let now = OffsetDateTime::now_utc();
KnowledgePage {
id: id.to_string(),
title: title.to_string(),
summary: content.chars().take(40).collect(),
content: content.to_string(),
tags: Vec::new(),
references: Vec::new(),
created_at: now,
updated_at: now,
}
}
#[tokio::main]
async fn main() {
// ── 1. 构建知识图谱 ──
println!("=== 1. 构建知识图谱 ===");
let graph = Arc::new(InMemoryGraph::new());
let mut langchain = GraphEntity::new("langchain", "LangChain", "framework");
langchain.description = "LLM application framework".to_string();
let mut langgraph = GraphEntity::new("langgraph", "LangGraph", "framework");
langgraph.description = "Graph-based agent runtime from LangChain".to_string();
let mut langsmith = GraphEntity::new("langsmith", "LangSmith", "tool");
langsmith.description = "Tracing and evaluation platform".to_string();
let mut python = GraphEntity::new("python", "Python", "language");
python.description = "Programming language".to_string();
let mut rust = GraphEntity::new("rust", "Rust", "language");
rust.description = "Systems programming language".to_string();
for e in [&langchain, &langgraph, &langsmith, &python, &rust] {
graph.add_entity(e.clone()).await.unwrap();
}
graph
.add_relation(GraphRelation::new("langchain", "langgraph", "includes", 0.9))
.await
.unwrap();
graph
.add_relation(GraphRelation::new("langchain", "langsmith", "includes", 0.7))
.await
.unwrap();
graph
.add_relation(GraphRelation::new("langchain", "python", "built_with", 0.95))
.await
.unwrap();
graph
.add_relation(GraphRelation::new("langgraph", "python", "depends_on", 0.8))
.await
.unwrap();
println!("已添加 5 个实体 + 4 条关系");
// ── 2. BFS 图遍历 ──
println!("\n=== 2. BFS 图遍历:从 LangChain 出发,depth=2 ===");
let related = graph
.get_related("langchain", 2, RelationDirection::Outgoing, None)
.await
.unwrap();
for se in &related {
println!(
" {} (score={:.3}, path={:?})",
se.entity.name, se.score, se.path
);
}
assert!(!related.is_empty(), "应找到关联实体");
// ── 3. 标签管理 ──
println!("\n=== 3. 标签管理 ===");
graph
.set_entity_tags("langchain", vec!["ai".into(), "framework".into(), "llm".into()])
.await
.unwrap();
graph
.set_entity_tags("langgraph", vec!["ai".into(), "agent".into()])
.await
.unwrap();
let tags = graph.find_tags("a").await.unwrap();
println!("前缀 'a' 查找标签: {:?}", tags);
let count = graph.entity_count_by_tag("ai").await.unwrap();
println!("标签 'ai' 下实体数: {}", count);
// ── 4. 双通道检索 ──
println!("\n=== 4. 双通道检索 ===");
let store: Arc<dyn MemoryStore> = Arc::new(InMemoryStore::new());
let ks = KnowledgeStore::new(store);
ks.add_page(make_page(
"p1",
"LangChain 框架介绍",
"LangChain 是用于构建 LLM 应用的开源框架",
))
.await
.unwrap();
ks.add_page(make_page(
"p2",
"Rust 异步编程",
"Rust 异步基于 tokio 与 futures 抽象",
))
.await
.unwrap();
// Hybrid 策略(默认)
let retriever = MemoryRetriever::new(ks, RetrieverConfig::default())
.with_knowledge_graph(graph.clone());
println!("\n--- Hybrid 检索: 'langchain' ---");
let result = retriever.retrieve("langchain").await.unwrap();
println!("策略: {:?}", result.strategy);
for item in &result.items {
match item {
RetrievalItem::KnowledgePage { page, score } => {
println!(" [Store] {} (score={:.3})", page.title, score);
}
RetrievalItem::GraphEntity {
entity, score, path, ..
} => {
println!(
" [Graph] {} (score={:.3}, path={:?})",
entity.name, score, path
);
}
}
}
let has_store = result
.items
.iter()
.any(|i| matches!(i, RetrievalItem::KnowledgePage { .. }));
let has_graph = result
.items
.iter()
.any(|i| matches!(i, RetrievalItem::GraphEntity { .. }));
assert!(has_store, "Hybrid 应有 Store 结果");
assert!(has_graph, "Hybrid 应有 Graph 结果");
// GraphOnly 策略
let store2: Arc<dyn MemoryStore> = Arc::new(InMemoryStore::new());
let ks2 = KnowledgeStore::new(store2);
ks2.add_page(make_page("p1", "LangChain", "LLM framework"))
.await
.unwrap();
let retriever_g = MemoryRetriever::new(ks2, RetrieverConfig::default())
.with_knowledge_graph(graph.clone())
.with_strategy(RetrievalStrategy::GraphOnly);
println!("\n--- GraphOnly 检索: 'langchain' ---");
let result = retriever_g.retrieve("langchain").await.unwrap();
println!("策略: {:?}", result.strategy);
for item in &result.items {
match item {
RetrievalItem::GraphEntity { entity, score, .. } => {
println!(" [Graph] {} (score={:.3})", entity.name, score);
}
RetrievalItem::KnowledgePage { page, score, .. } => {
println!(" [Store] {} (score={:.3})", page.title, score);
}
}
}
assert!(
result.items.iter().all(|i| matches!(i, RetrievalItem::GraphEntity { .. })),
"GraphOnly 应只返回 Graph 结果"
);
println!("\n✓ knowledge_graph_demo 完成");
}
+11 -2
View File
@@ -74,8 +74,15 @@ async fn main() {
let result = retriever.retrieve("Rust 异步").await.unwrap();
println!("query: {}", result.query);
for item in &result.items {
println!(" 命中: {} (score={:.3})", item.page.title, item.score);
assert!((0.0..=1.0).contains(&item.score), "score 应在 [0, 1] 区间");
match item {
agcore::memory::RetrievalItem::KnowledgePage { page, score } => {
println!(" 命中: {} (score={:.3})", page.title, score);
assert!((0.0..=1.0).contains(score), "score 应在 [0, 1] 区间");
}
agcore::memory::RetrievalItem::GraphEntity { entity, score, .. } => {
println!(" 命中实体: {} (score={:.3})", entity.name, score);
}
}
}
assert!(!result.items.is_empty(), "应至少命中一个页面");
@@ -89,6 +96,7 @@ async fn main() {
let cfg = RetrieverConfig {
max_results: 20,
min_score: 0.5,
graph_depth: 2,
};
let retriever2 = MemoryRetriever::new(ks2, cfg);
let result = retriever2.retrieve("完全不相关的火锅配方").await.unwrap();
@@ -111,6 +119,7 @@ async fn main() {
let cfg = RetrieverConfig {
max_results: 2,
min_score: 0.0,
graph_depth: 2,
};
let retriever3 = MemoryRetriever::new(ks3, cfg);
let result = retriever3.retrieve("Rust").await.unwrap();