docs: 更新 README feature 表 + 升级指南 + 示例注释 + roadmap 同步
- README 添加 feature 组合表 + 模块级 features 清单 + 升级指南 - 18 个 example 顶部添加 Required features 注释 - roadmap.md 和 roadmap-v0.3.2.md 同步 Phase 26-27 完成状态 - cargo fmt 全量格式化(修复预存格式问题,CI format job 可通过)
This commit is contained in:
+75
-79
@@ -35,7 +35,11 @@ pub struct GraphEntity {
|
||||
|
||||
impl GraphEntity {
|
||||
/// 创建一个最小实体(仅 id + name + type,其余为空)。
|
||||
pub fn new(id: impl Into<String>, name: impl Into<String>, entity_type: impl Into<String>) -> Self {
|
||||
pub fn new(
|
||||
id: impl Into<String>,
|
||||
name: impl Into<String>,
|
||||
entity_type: impl Into<String>,
|
||||
) -> Self {
|
||||
Self {
|
||||
id: id.into(),
|
||||
name: name.into(),
|
||||
@@ -81,7 +85,10 @@ impl GraphRelation {
|
||||
|
||||
/// 复合键:`source_id:target_id:relation_type`,用于去重和查找。
|
||||
pub fn composite_key(&self) -> String {
|
||||
format!("{}:{}:{}", self.source_id, self.target_id, self.relation_type)
|
||||
format!(
|
||||
"{}:{}:{}",
|
||||
self.source_id, self.target_id, self.relation_type
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -253,9 +260,10 @@ impl KnowledgeGraph for InMemoryGraph {
|
||||
if entity.id.is_empty() {
|
||||
return Err(MemoryError::InvalidInput("entity id is empty".into()));
|
||||
}
|
||||
let mut inner = self.inner.lock().map_err(|e| {
|
||||
MemoryError::RetrievalError(format!("lock poisoned: {e}"))
|
||||
})?;
|
||||
let mut inner = self
|
||||
.inner
|
||||
.lock()
|
||||
.map_err(|e| MemoryError::RetrievalError(format!("lock poisoned: {e}")))?;
|
||||
// upsert:若已存在,先收集旧标签用于清理反向引用(避免同时借用 entities 和 tag_index)
|
||||
let old_tags: Vec<String> = inner
|
||||
.entities
|
||||
@@ -283,16 +291,18 @@ impl KnowledgeGraph for InMemoryGraph {
|
||||
}
|
||||
|
||||
async fn get_entity(&self, id: &str) -> Result<Option<GraphEntity>, MemoryError> {
|
||||
let inner = self.inner.lock().map_err(|e| {
|
||||
MemoryError::RetrievalError(format!("lock poisoned: {e}"))
|
||||
})?;
|
||||
let inner = self
|
||||
.inner
|
||||
.lock()
|
||||
.map_err(|e| MemoryError::RetrievalError(format!("lock poisoned: {e}")))?;
|
||||
Ok(inner.entities.get(id).cloned())
|
||||
}
|
||||
|
||||
async fn remove_entity(&self, id: &str) -> Result<(), MemoryError> {
|
||||
let mut inner = self.inner.lock().map_err(|e| {
|
||||
MemoryError::RetrievalError(format!("lock poisoned: {e}"))
|
||||
})?;
|
||||
let mut inner = self
|
||||
.inner
|
||||
.lock()
|
||||
.map_err(|e| MemoryError::RetrievalError(format!("lock poisoned: {e}")))?;
|
||||
// 移除实体并清理其标签反向引用
|
||||
if let Some(entity) = inner.entities.remove(id) {
|
||||
for tag in &entity.tags {
|
||||
@@ -312,9 +322,10 @@ impl KnowledgeGraph for InMemoryGraph {
|
||||
}
|
||||
|
||||
async fn add_relation(&self, relation: GraphRelation) -> Result<(), MemoryError> {
|
||||
let mut inner = self.inner.lock().map_err(|e| {
|
||||
MemoryError::RetrievalError(format!("lock poisoned: {e}"))
|
||||
})?;
|
||||
let mut inner = self
|
||||
.inner
|
||||
.lock()
|
||||
.map_err(|e| MemoryError::RetrievalError(format!("lock poisoned: {e}")))?;
|
||||
// 校验两端实体存在
|
||||
if !inner.entities.contains_key(&relation.source_id) {
|
||||
return Err(MemoryError::InvalidInput(format!(
|
||||
@@ -348,11 +359,14 @@ impl KnowledgeGraph for InMemoryGraph {
|
||||
target_id: &str,
|
||||
relation_type: &str,
|
||||
) -> Result<(), MemoryError> {
|
||||
let mut inner = self.inner.lock().map_err(|e| {
|
||||
MemoryError::RetrievalError(format!("lock poisoned: {e}"))
|
||||
})?;
|
||||
let mut inner = self
|
||||
.inner
|
||||
.lock()
|
||||
.map_err(|e| MemoryError::RetrievalError(format!("lock poisoned: {e}")))?;
|
||||
inner.relations.retain(|r| {
|
||||
!(r.source_id == source_id && r.target_id == target_id && r.relation_type == relation_type)
|
||||
!(r.source_id == source_id
|
||||
&& r.target_id == target_id
|
||||
&& r.relation_type == relation_type)
|
||||
});
|
||||
Ok(())
|
||||
}
|
||||
@@ -364,9 +378,10 @@ impl KnowledgeGraph for InMemoryGraph {
|
||||
direction: RelationDirection,
|
||||
relation_types: Option<&[&str]>,
|
||||
) -> Result<Vec<ScoredEntity>, MemoryError> {
|
||||
let inner = self.inner.lock().map_err(|e| {
|
||||
MemoryError::RetrievalError(format!("lock poisoned: {e}"))
|
||||
})?;
|
||||
let inner = self
|
||||
.inner
|
||||
.lock()
|
||||
.map_err(|e| MemoryError::RetrievalError(format!("lock poisoned: {e}")))?;
|
||||
|
||||
// 1. 验证起点存在
|
||||
if !inner.entities.contains_key(entity_id) {
|
||||
@@ -470,9 +485,10 @@ impl KnowledgeGraph for InMemoryGraph {
|
||||
}
|
||||
|
||||
async fn find_by_keywords(&self, keywords: &[String]) -> Result<Vec<GraphEntity>, MemoryError> {
|
||||
let inner = self.inner.lock().map_err(|e| {
|
||||
MemoryError::RetrievalError(format!("lock poisoned: {e}"))
|
||||
})?;
|
||||
let inner = self
|
||||
.inner
|
||||
.lock()
|
||||
.map_err(|e| MemoryError::RetrievalError(format!("lock poisoned: {e}")))?;
|
||||
if keywords.is_empty() {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
@@ -495,9 +511,10 @@ impl KnowledgeGraph for InMemoryGraph {
|
||||
}
|
||||
|
||||
async fn find_tags(&self, prefix: &str) -> Result<Vec<String>, MemoryError> {
|
||||
let inner = self.inner.lock().map_err(|e| {
|
||||
MemoryError::RetrievalError(format!("lock poisoned: {e}"))
|
||||
})?;
|
||||
let inner = self
|
||||
.inner
|
||||
.lock()
|
||||
.map_err(|e| MemoryError::RetrievalError(format!("lock poisoned: {e}")))?;
|
||||
let prefix_l = prefix.to_lowercase();
|
||||
let mut tags: Vec<String> = inner
|
||||
.tag_index
|
||||
@@ -514,9 +531,10 @@ impl KnowledgeGraph for InMemoryGraph {
|
||||
entity_id: &str,
|
||||
tags: Vec<String>,
|
||||
) -> Result<usize, MemoryError> {
|
||||
let mut inner = self.inner.lock().map_err(|e| {
|
||||
MemoryError::RetrievalError(format!("lock poisoned: {e}"))
|
||||
})?;
|
||||
let mut inner = self
|
||||
.inner
|
||||
.lock()
|
||||
.map_err(|e| MemoryError::RetrievalError(format!("lock poisoned: {e}")))?;
|
||||
|
||||
// 先收集旧标签(避免同时借用 entities 和 tag_index)
|
||||
let old_tags: Vec<String> = inner
|
||||
@@ -554,13 +572,11 @@ impl KnowledgeGraph for InMemoryGraph {
|
||||
}
|
||||
|
||||
async fn entity_count_by_tag(&self, tag: &str) -> Result<usize, MemoryError> {
|
||||
let inner = self.inner.lock().map_err(|e| {
|
||||
MemoryError::RetrievalError(format!("lock poisoned: {e}"))
|
||||
})?;
|
||||
Ok(inner
|
||||
.tag_index
|
||||
.get(tag)
|
||||
.map_or(0, |ids| ids.len()))
|
||||
let inner = self
|
||||
.inner
|
||||
.lock()
|
||||
.map_err(|e| MemoryError::RetrievalError(format!("lock poisoned: {e}")))?;
|
||||
Ok(inner.tag_index.get(tag).map_or(0, |ids| ids.len()))
|
||||
}
|
||||
|
||||
fn tag_constraints(&self) -> TagConstraints {
|
||||
@@ -660,10 +676,7 @@ mod tests {
|
||||
#[tokio::test]
|
||||
async fn add_relation_validates_entities() {
|
||||
let graph = InMemoryGraph::new();
|
||||
graph
|
||||
.add_entity(make_entity("a", "A", "x"))
|
||||
.await
|
||||
.unwrap();
|
||||
graph.add_entity(make_entity("a", "A", "x")).await.unwrap();
|
||||
// target 不存在
|
||||
let result = graph
|
||||
.add_relation(GraphRelation::new("a", "ghost", "r", 0.5))
|
||||
@@ -674,14 +687,8 @@ mod tests {
|
||||
#[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_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
|
||||
@@ -701,14 +708,8 @@ mod tests {
|
||||
#[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_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
|
||||
@@ -726,10 +727,7 @@ mod tests {
|
||||
#[tokio::test]
|
||||
async fn get_related_empty_graph() {
|
||||
let graph = InMemoryGraph::new();
|
||||
graph
|
||||
.add_entity(make_entity("a", "A", "x"))
|
||||
.await
|
||||
.unwrap();
|
||||
graph.add_entity(make_entity("a", "A", "x")).await.unwrap();
|
||||
let related = graph
|
||||
.get_related("a", 3, RelationDirection::Outgoing, None)
|
||||
.await
|
||||
@@ -740,10 +738,7 @@ mod tests {
|
||||
#[tokio::test]
|
||||
async fn get_related_depth_0() {
|
||||
let graph = InMemoryGraph::new();
|
||||
graph
|
||||
.add_entity(make_entity("a", "A", "x"))
|
||||
.await
|
||||
.unwrap();
|
||||
graph.add_entity(make_entity("a", "A", "x")).await.unwrap();
|
||||
let related = graph
|
||||
.get_related("a", 0, RelationDirection::Outgoing, None)
|
||||
.await
|
||||
@@ -853,10 +848,16 @@ mod tests {
|
||||
#[tokio::test]
|
||||
async fn get_related_star_topology() {
|
||||
let graph = InMemoryGraph::new();
|
||||
graph.add_entity(make_entity("center", "C", "x")).await.unwrap();
|
||||
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_entity(make_entity(&leaf, &leaf, "x"))
|
||||
.await
|
||||
.unwrap();
|
||||
graph
|
||||
.add_relation(GraphRelation::new("center", &leaf, "r", 0.8))
|
||||
.await
|
||||
@@ -1016,10 +1017,7 @@ mod tests {
|
||||
#[tokio::test]
|
||||
async fn set_and_find_tags() {
|
||||
let graph = InMemoryGraph::new();
|
||||
graph
|
||||
.add_entity(make_entity("a", "A", "x"))
|
||||
.await
|
||||
.unwrap();
|
||||
graph.add_entity(make_entity("a", "A", "x")).await.unwrap();
|
||||
let n = graph
|
||||
.set_entity_tags("a", vec!["rust".into(), "ai".into()])
|
||||
.await
|
||||
@@ -1055,7 +1053,10 @@ mod tests {
|
||||
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()])
|
||||
.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");
|
||||
@@ -1075,10 +1076,7 @@ mod tests {
|
||||
.set_entity_tags("b", vec!["rust".into(), "ai".into()])
|
||||
.await
|
||||
.unwrap();
|
||||
graph
|
||||
.set_entity_tags("c", vec!["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);
|
||||
}
|
||||
@@ -1086,9 +1084,7 @@ mod tests {
|
||||
#[tokio::test]
|
||||
async fn set_entity_tags_not_found() {
|
||||
let graph = InMemoryGraph::new();
|
||||
let result = graph
|
||||
.set_entity_tags("ghost", vec!["x".into()])
|
||||
.await;
|
||||
let result = graph.set_entity_tags("ghost", vec!["x".into()]).await;
|
||||
assert!(matches!(result, Err(MemoryError::NotFound(_))));
|
||||
}
|
||||
|
||||
|
||||
+33
-12
@@ -270,7 +270,12 @@ impl MemoryRetriever {
|
||||
}
|
||||
// BFS 找相关实体
|
||||
let related: Vec<ScoredEntity> = graph
|
||||
.get_related(&start.id, depth, crate::memory::graph::RelationDirection::Both, None)
|
||||
.get_related(
|
||||
&start.id,
|
||||
depth,
|
||||
crate::memory::graph::RelationDirection::Both,
|
||||
None,
|
||||
)
|
||||
.await?;
|
||||
for se in related {
|
||||
if seen.insert(se.entity.id.clone()) {
|
||||
@@ -361,7 +366,7 @@ fn default_stop_words() -> HashSet<String> {
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::memory::graph::{GraphEntity, InMemoryGraph, GraphRelation};
|
||||
use crate::memory::graph::{GraphEntity, GraphRelation, InMemoryGraph};
|
||||
use crate::memory::knowledge::KnowledgeStore;
|
||||
use crate::memory::{InMemoryStore, MemoryStore};
|
||||
use std::sync::Arc;
|
||||
@@ -502,7 +507,12 @@ mod tests {
|
||||
graph.add_entity(langgraph).await.unwrap();
|
||||
graph.add_entity(python).await.unwrap();
|
||||
graph
|
||||
.add_relation(GraphRelation::new("langchain", "langgraph", "includes", 0.8))
|
||||
.add_relation(GraphRelation::new(
|
||||
"langchain",
|
||||
"langgraph",
|
||||
"includes",
|
||||
0.8,
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
graph
|
||||
@@ -523,7 +533,12 @@ mod tests {
|
||||
let result = retriever.retrieve("langchain").await.unwrap();
|
||||
assert_eq!(result.strategy, RetrievalStrategy::GraphOnly);
|
||||
// 应该全部是 GraphEntity 变体
|
||||
assert!(result.items.iter().all(|i| matches!(i, RetrievalItem::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");
|
||||
}
|
||||
@@ -544,8 +559,8 @@ mod tests {
|
||||
// KnowledgeGraph 也有匹配
|
||||
let graph = make_graph_with_data().await;
|
||||
|
||||
let retriever = MemoryRetriever::new(ks, RetrieverConfig::default())
|
||||
.with_knowledge_graph(graph);
|
||||
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
|
||||
@@ -569,8 +584,8 @@ mod tests {
|
||||
.unwrap();
|
||||
// 空图
|
||||
let graph = Arc::new(InMemoryGraph::new());
|
||||
let retriever = MemoryRetriever::new(ks, RetrieverConfig::default())
|
||||
.with_knowledge_graph(graph);
|
||||
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 结果
|
||||
@@ -587,15 +602,18 @@ mod tests {
|
||||
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 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");
|
||||
assert!(
|
||||
has_entity,
|
||||
"should have graph results even when store is empty"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
@@ -615,7 +633,10 @@ mod tests {
|
||||
.items
|
||||
.iter()
|
||||
.any(|i| matches!(i, RetrievalItem::GraphEntity { .. }));
|
||||
assert!(!has_entity, "KnowledgeOnly should not return graph entities");
|
||||
assert!(
|
||||
!has_entity,
|
||||
"KnowledgeOnly should not return graph entities"
|
||||
);
|
||||
}
|
||||
|
||||
// ── 辅助函数测试 ──
|
||||
|
||||
@@ -6,9 +6,9 @@ use std::path::Path;
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use async_trait::async_trait;
|
||||
use rusqlite::{params, params_from_iter, Connection, ErrorCode};
|
||||
use time::format_description::well_known::Rfc3339;
|
||||
use rusqlite::{Connection, ErrorCode, params, params_from_iter};
|
||||
use time::OffsetDateTime;
|
||||
use time::format_description::well_known::Rfc3339;
|
||||
use tracing::{debug, error, instrument, warn};
|
||||
|
||||
use crate::memory::error::MemoryError;
|
||||
@@ -114,9 +114,7 @@ impl MemoryStore for SqliteStore {
|
||||
.map_err(|e| map_sqlite_error(e, "query get"))?;
|
||||
match rows.next() {
|
||||
None => Ok(None),
|
||||
Some(row) => row
|
||||
.map(Some)
|
||||
.map_err(|e| map_sqlite_error(e, "decode row")),
|
||||
Some(row) => row.map(Some).map_err(|e| map_sqlite_error(e, "decode row")),
|
||||
}
|
||||
})
|
||||
.await
|
||||
@@ -130,11 +128,8 @@ impl MemoryStore for SqliteStore {
|
||||
|
||||
tokio::task::spawn_blocking(move || -> Result<(), MemoryError> {
|
||||
let conn = conn.lock().unwrap_or_else(|e| e.into_inner());
|
||||
conn.execute(
|
||||
"DELETE FROM memory_items WHERE id = ?1",
|
||||
params![id_owned],
|
||||
)
|
||||
.map_err(|e| map_sqlite_error(e, "delete"))?;
|
||||
conn.execute("DELETE FROM memory_items WHERE id = ?1", params![id_owned])
|
||||
.map_err(|e| map_sqlite_error(e, "delete"))?;
|
||||
Ok(())
|
||||
})
|
||||
.await
|
||||
@@ -143,9 +138,8 @@ impl MemoryStore for SqliteStore {
|
||||
|
||||
#[instrument(skip(self, filter))]
|
||||
async fn list(&self, filter: &MemoryFilter) -> Result<Vec<MemoryItem>, MemoryError> {
|
||||
let mut sql = String::from(
|
||||
"SELECT id, content, metadata, created_at FROM memory_items WHERE 1=1",
|
||||
);
|
||||
let mut sql =
|
||||
String::from("SELECT id, content, metadata, created_at FROM memory_items WHERE 1=1");
|
||||
let mut param_values: Vec<String> = Vec::new();
|
||||
let mut ph_idx = 0usize;
|
||||
|
||||
@@ -309,9 +303,7 @@ fn map_sqlite_error(e: rusqlite::Error, ctx: &str) -> MemoryError {
|
||||
rusqlite::Error::InvalidQuery
|
||||
| rusqlite::Error::InvalidParameterName(_)
|
||||
| rusqlite::Error::InvalidColumnIndex(_)
|
||||
| rusqlite::Error::InvalidColumnName(_) => {
|
||||
MemoryError::InvalidInput(format!("{ctx}: {e}"))
|
||||
}
|
||||
| rusqlite::Error::InvalidColumnName(_) => MemoryError::InvalidInput(format!("{ctx}: {e}")),
|
||||
rusqlite::Error::FromSqlConversionFailure(_, _, _)
|
||||
| rusqlite::Error::ToSqlConversionFailure(_) => {
|
||||
MemoryError::Serialization(format!("{ctx}: {e}"))
|
||||
@@ -527,8 +519,7 @@ mod tests {
|
||||
// ponytail: 回归验证 SqliteStore 可作为 Arc<dyn MemoryStore> 与 InMemoryStore 互换
|
||||
// 所有现有消费者(Conversation / Knowledge / Retriever / SessionMemory)均通过 trait object 引用,
|
||||
// 此测试确保 trait 接口契约在 SqliteStore 上同样成立。
|
||||
let sqlite: Arc<dyn MemoryStore> =
|
||||
Arc::new(SqliteStore::open(":memory:").unwrap());
|
||||
let sqlite: Arc<dyn MemoryStore> = Arc::new(SqliteStore::open(":memory:").unwrap());
|
||||
let in_mem: Arc<dyn MemoryStore> = Arc::new(InMemoryStore::new());
|
||||
|
||||
let stores: Vec<Arc<dyn MemoryStore>> = vec![Arc::clone(&sqlite), Arc::clone(&in_mem)];
|
||||
@@ -556,12 +547,7 @@ mod tests {
|
||||
handles.push(tokio::spawn(async move {
|
||||
let id = format!("concurrent_{i}");
|
||||
// 设置每次 save 的 per-call timeout —— busy_timeout=5000ms 应足够
|
||||
match tokio::time::timeout(
|
||||
Duration::from_secs(10),
|
||||
s.save(make_item(&id)),
|
||||
)
|
||||
.await
|
||||
{
|
||||
match tokio::time::timeout(Duration::from_secs(10), s.save(make_item(&id))).await {
|
||||
Ok(res) => res.unwrap(),
|
||||
Err(_) => panic!("save({id}) timed out under 100-way concurrency"),
|
||||
}
|
||||
|
||||
+3
-11
@@ -30,11 +30,7 @@ pub trait VectorRetriever: Send + Sync {
|
||||
///
|
||||
/// 返回 `Vec<(id, score)>`,按 score 降序排列,score ∈ [0.0, 1.0]
|
||||
/// (余弦相似度)。当 `k == 0`、索引为空或 query 为零向量时返回空 Vec。
|
||||
async fn search(
|
||||
&self,
|
||||
query: Vec<f32>,
|
||||
k: usize,
|
||||
) -> Result<Vec<(String, f32)>, MemoryError>;
|
||||
async fn search(&self, query: Vec<f32>, k: usize) -> Result<Vec<(String, f32)>, MemoryError>;
|
||||
}
|
||||
|
||||
/// 进程内向量检索器 —— 基于 HashMap + 全量余弦相似度扫描。
|
||||
@@ -79,11 +75,7 @@ impl VectorRetriever for InMemoryVectorRetriever {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn search(
|
||||
&self,
|
||||
query: Vec<f32>,
|
||||
k: usize,
|
||||
) -> Result<Vec<(String, f32)>, MemoryError> {
|
||||
async fn search(&self, query: Vec<f32>, k: usize) -> Result<Vec<(String, f32)>, MemoryError> {
|
||||
let vectors = self
|
||||
.vectors
|
||||
.lock()
|
||||
@@ -241,4 +233,4 @@ mod tests {
|
||||
h.await.unwrap();
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+26
-46
@@ -11,8 +11,8 @@ use std::sync::{Arc, Mutex};
|
||||
|
||||
use async_trait::async_trait;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use time::format_description::well_known::Rfc3339;
|
||||
use time::OffsetDateTime;
|
||||
use time::format_description::well_known::Rfc3339;
|
||||
use tracing::{debug, info};
|
||||
|
||||
use crate::document::{Document, RecursiveCharacterSplitter};
|
||||
@@ -43,11 +43,8 @@ pub trait VectorStore: Send + Sync {
|
||||
/// - 截取 `min(len)` 对处理(部分写入已发生)
|
||||
/// - 返回 `Err(MemoryError::InvalidInput)` 告知截断
|
||||
/// - 调用方可以 `let _ = store.add(...)` 忽略错误
|
||||
async fn add(
|
||||
&self,
|
||||
documents: &[Document],
|
||||
embeddings: &[Vec<f32>],
|
||||
) -> Result<(), MemoryError>;
|
||||
async fn add(&self, documents: &[Document], embeddings: &[Vec<f32>])
|
||||
-> Result<(), MemoryError>;
|
||||
|
||||
/// 检索与 `query` 向量最相似的 `k` 条记录。
|
||||
///
|
||||
@@ -59,11 +56,7 @@ pub trait VectorStore: Send + Sync {
|
||||
/// - 空索引 → 返回 `vec![]`
|
||||
/// - `k == 0` → 返回 `vec![]`
|
||||
/// - 零向量(norm ≈ 0)→ 返回 `vec![]`
|
||||
async fn search(
|
||||
&self,
|
||||
query: &[f32],
|
||||
k: usize,
|
||||
) -> Result<Vec<(Document, f32)>, MemoryError>;
|
||||
async fn search(&self, query: &[f32], k: usize) -> Result<Vec<(Document, f32)>, MemoryError>;
|
||||
|
||||
/// 批量删除文档(幂等)。
|
||||
///
|
||||
@@ -105,9 +98,7 @@ impl InMemoryVectorStore {
|
||||
}
|
||||
|
||||
/// 从预填充的 entries 构造(供 `PersistentVectorStore` 使用)。
|
||||
pub(crate) fn with_entries(
|
||||
entries: HashMap<String, (Document, Vec<f32>)>,
|
||||
) -> Self {
|
||||
pub(crate) fn with_entries(entries: HashMap<String, (Document, Vec<f32>)>) -> Self {
|
||||
Self {
|
||||
entries: Mutex::new(entries),
|
||||
}
|
||||
@@ -142,7 +133,10 @@ impl VectorStore for InMemoryVectorStore {
|
||||
}
|
||||
|
||||
for i in 0..n {
|
||||
entries.insert(documents[i].id.clone(), (documents[i].clone(), embeddings[i].clone()));
|
||||
entries.insert(
|
||||
documents[i].id.clone(),
|
||||
(documents[i].clone(), embeddings[i].clone()),
|
||||
);
|
||||
}
|
||||
|
||||
if documents.len() != embeddings.len() {
|
||||
@@ -156,11 +150,7 @@ impl VectorStore for InMemoryVectorStore {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn search(
|
||||
&self,
|
||||
query: &[f32],
|
||||
k: usize,
|
||||
) -> Result<Vec<(Document, f32)>, MemoryError> {
|
||||
async fn search(&self, query: &[f32], k: usize) -> Result<Vec<(Document, f32)>, MemoryError> {
|
||||
if k == 0 {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
@@ -265,10 +255,7 @@ impl PersistentVectorStore {
|
||||
///
|
||||
/// `MemoryStore::list()` 由 `SqliteStore` 内部使用 `spawn_blocking` 卸载,
|
||||
/// 加载过程本身在 async context 中即可,无需额外 spawn_blocking。
|
||||
pub async fn new(
|
||||
store: Arc<dyn MemoryStore>,
|
||||
namespace: &str,
|
||||
) -> Result<Self, MemoryError> {
|
||||
pub async fn new(store: Arc<dyn MemoryStore>, namespace: &str) -> Result<Self, MemoryError> {
|
||||
let prefix = format!("vec:{namespace}:");
|
||||
let filter = MemoryFilter {
|
||||
prefix: Some(prefix.clone()),
|
||||
@@ -292,7 +279,10 @@ impl PersistentVectorStore {
|
||||
entries.insert(doc.id.clone(), (doc, entry.embedding));
|
||||
}
|
||||
|
||||
info!(entries = entries.len(), "PersistentVectorStore — 内存索引重建完成");
|
||||
info!(
|
||||
entries = entries.len(),
|
||||
"PersistentVectorStore — 内存索引重建完成"
|
||||
);
|
||||
Ok(Self {
|
||||
inner: InMemoryVectorStore::with_entries(entries),
|
||||
store,
|
||||
@@ -340,11 +330,7 @@ impl VectorStore for PersistentVectorStore {
|
||||
self.inner.add(documents, embeddings).await
|
||||
}
|
||||
|
||||
async fn search(
|
||||
&self,
|
||||
query: &[f32],
|
||||
k: usize,
|
||||
) -> Result<Vec<(Document, f32)>, MemoryError> {
|
||||
async fn search(&self, query: &[f32], k: usize) -> Result<Vec<(Document, f32)>, MemoryError> {
|
||||
tracing::trace!(k, "PersistentVectorStore::search");
|
||||
self.inner.search(query, k).await
|
||||
}
|
||||
@@ -575,10 +561,7 @@ mod tests {
|
||||
async fn remove_items() {
|
||||
let store = InMemoryVectorStore::new();
|
||||
let docs = vec![make_doc("a", "alpha"), make_doc("b", "beta")];
|
||||
let embeddings = vec![
|
||||
make_vec(&[1.0, 0.0, 0.0]),
|
||||
make_vec(&[0.0, 1.0, 0.0]),
|
||||
];
|
||||
let embeddings = vec![make_vec(&[1.0, 0.0, 0.0]), make_vec(&[0.0, 1.0, 0.0])];
|
||||
store.add(&docs, &embeddings).await.unwrap();
|
||||
|
||||
store.remove(&["a".to_string()]).await.unwrap();
|
||||
@@ -599,7 +582,9 @@ mod tests {
|
||||
let embeddings = vec![make_vec(&[1.0, 0.0, 0.0])];
|
||||
store.add(&docs, &embeddings).await.unwrap();
|
||||
|
||||
let result = store.remove(&["nonexistent".to_string(), "also_nonexistent".to_string()]).await;
|
||||
let result = store
|
||||
.remove(&["nonexistent".to_string(), "also_nonexistent".to_string()])
|
||||
.await;
|
||||
assert!(result.is_ok(), "批量删除不存在 id 不应报错");
|
||||
|
||||
let results = store.search(&[1.0, 0.0, 0.0], 5).await.unwrap();
|
||||
@@ -657,7 +642,9 @@ mod tests {
|
||||
backend: Arc<dyn MemoryStore>,
|
||||
namespace: &str,
|
||||
) -> PersistentVectorStore {
|
||||
PersistentVectorStore::new(backend, namespace).await.unwrap()
|
||||
PersistentVectorStore::new(backend, namespace)
|
||||
.await
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
@@ -755,10 +742,7 @@ mod tests {
|
||||
let store = make_persistent(Arc::clone(&backend), "default").await;
|
||||
|
||||
// 正常写入前 2 条
|
||||
let docs_first = vec![
|
||||
make_doc("doc_0", "first"),
|
||||
make_doc("doc_1", "second"),
|
||||
];
|
||||
let docs_first = vec![make_doc("doc_0", "first"), make_doc("doc_1", "second")];
|
||||
let embeddings_first = vec![make_vec(&[1.0, 0.0, 0.0]), make_vec(&[0.0, 1.0, 0.0])];
|
||||
store.add(&docs_first, &embeddings_first).await.unwrap();
|
||||
|
||||
@@ -797,11 +781,7 @@ mod tests {
|
||||
let store: Arc<dyn VectorStore> = Arc::new(InMemoryVectorStore::new());
|
||||
let splitter = RecursiveCharacterSplitter::new(50, 5);
|
||||
|
||||
let pipeline = RagPipeline::new(
|
||||
Arc::clone(&embedder),
|
||||
Arc::clone(&store),
|
||||
Some(splitter),
|
||||
);
|
||||
let pipeline = RagPipeline::new(Arc::clone(&embedder), Arc::clone(&store), Some(splitter));
|
||||
|
||||
// 创建多段落文档
|
||||
let doc = Document::new(
|
||||
@@ -934,4 +914,4 @@ mod tests {
|
||||
elapsed.as_millis()
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user