feat(memory): 新增 VectorStore 抽象与 RagPipeline 持久化管线

- 新增 src/memory/vector_store.rs(约 660 行):

  - VectorStore trait:批量 add / search / remove + add_one 默认实现

  - InMemoryVectorStore:Mutex<HashMap> + 余弦全量扫描(锁内克隆、锁外计算)

  - PersistentVectorStore:MemoryStore 包装,JSON blob 持久化,先持久化后内存

  - RagPipeline:split → embed → store 组合器(具体 struct,非 trait)

- 22 个内联测试覆盖 InMemory(10)/ Persistent(6)/ RagPipeline(4)/ 性能基准(2)

- 性能断言:search 10K 条 <100ms,PersistentVectorStore::new 加载 <500ms

- 标记 VectorRetriever / InMemoryVectorRetriever 为 #[deprecated(since="0.3.0")]

- memory.rs 追加 VectorStore 等 4 个类型的 re-export

- document_demo 从手动 VectorRetriever 循环迁移到 RagPipeline 两行调用

- 零新增外部依赖
This commit is contained in:
徐涛
2026-07-09 16:51:29 +08:00
parent b04427e83f
commit 32d886f870
5 changed files with 2144 additions and 37 deletions
File diff suppressed because it is too large Load Diff
+33 -37
View File
@@ -1,17 +1,18 @@
//! document_demo —— Document + RecursiveCharacterSplitter + MockEmbedding + VectorRetriever 完整衔接示例。
//! document_demo —— Document + RecursiveCharacterSplitter + MockEmbedding + RagPipeline 完整衔接示例。
//!
//! 演示 RAG 管线前置流程
//! 演示 RAG 管线:
//! 1. 创建多段落 Document
//! 2. RecursiveCharacterSplitter 分割为 chunk
//! 3. MockEmbedding 嵌入所有 chunk
//! 4. 与 InMemoryVectorRetriever 手动 zip 衔接
//! 5. 模拟查询做语义检索
//! 3. RagPipeline.ingest() 自动嵌入并存储
//! 4. RagPipeline.retrieve() 做语义检索
//!
//! 运行:`cargo run --example document_demo`(离线,零配置)
use std::sync::Arc;
use agcore::document::{Document, RecursiveCharacterSplitter};
use agcore::llm::embedding::{Embedding, MockEmbedding};
use agcore::memory::vector::{InMemoryVectorRetriever, VectorRetriever};
use agcore::memory::{InMemoryVectorStore, RagPipeline, VectorStore};
#[tokio::main]
async fn main() {
@@ -32,43 +33,38 @@ async fn main() {
println!("输入文档: {} 字符", doc.content.chars().count());
// 2. RecursiveCharacterSplitter 分割
// 2. 构造 RAG 管线(嵌入器 + 向量存储 + 分割器)
let embedder: Arc<dyn Embedding> = Arc::new(MockEmbedding::new(4));
let store: Arc<dyn VectorStore> = Arc::new(InMemoryVectorStore::new());
let splitter = RecursiveCharacterSplitter::new(200, 30);
let chunks = splitter.split(&[doc]);
let pipeline = RagPipeline::new(
Arc::clone(&embedder),
Arc::clone(&store),
Some(splitter),
);
println!("\n分割为 {} 个 chunk:", chunks.len());
for (i, chunk) in chunks.iter().enumerate() {
// 3. 一次性 ingest:自动 split → embed → add
pipeline.ingest(std::slice::from_ref(&doc)).await.unwrap();
// 4. 模拟查询:复用第一个 chunk 的 content 作为查询文本
let chunks_in_store = store.search(&[1.0, 0.0, 0.0, 0.0], 1).await.unwrap();
assert!(!chunks_in_store.is_empty(), "ingest 后 store 应有数据");
let query_text = &chunks_in_store[0].0.content;
let results = pipeline.retrieve(query_text, 3).await.unwrap();
println!("\nTop 3 检索结果(与第一个 chunk 相似):");
for (doc, score) in &results {
println!(
" [{:02}] id={}, 长度={}",
i, chunk.id, chunk.content.chars().count()
" id={}, score={:.4}, content={}",
doc.id, score, doc.content
);
}
// 3. MockEmbedding 嵌入所有 chunk
let embedder = MockEmbedding::new(4);
let texts: Vec<String> = chunks.iter().map(|c| c.content.clone()).collect();
let vectors = embedder.embed(&texts).await.unwrap();
println!("\n嵌入维度: {},向量数: {}", embedder.dim(), vectors.len());
// 4. 与 InMemoryVectorRetriever 手动 zip 衔接
let retriever: std::sync::Arc<dyn VectorRetriever> =
std::sync::Arc::new(InMemoryVectorRetriever::new());
for (chunk, vec) in chunks.iter().zip(vectors.iter()) {
retriever.index(chunk.id.clone(), vec.clone()).await.unwrap();
}
println!("已索引 {} 个 chunk", chunks.len());
// 5. 模拟查询:复用第一个 chunk 的 embedding 作为查询向量
let query_vec = vectors[0].clone();
let results = retriever.search(query_vec, 3).await.unwrap();
println!("\nTop 3 检索结果(与 chunk 0 相似度):");
for (id, score) in &results {
println!(" id={}, score={:.4}", id, score);
}
assert_eq!(chunks.len(), vectors.len(), "chunks 与 vectors 数量必须一致");
assert!(!results.is_empty(), "至少应返回 1 条检索结果");
assert!(results[0].0.starts_with("rust-intro:chunk:0000"), "Top 1 应为 chunk 0 自身");
assert!(
results[0].0.id.starts_with("rust-intro:chunk:0000"),
"Top 1 应为 chunk 0 自身"
);
println!("\n✓ document_demo 完成");
}
}
+3
View File
@@ -7,6 +7,7 @@ pub mod retriever;
pub mod store;
pub mod types;
pub mod vector;
pub mod vector_store;
// 高频类型(大多数下游需要)
pub use conversation::{ConversationMemory, ConversationMemoryConfig};
@@ -14,7 +15,9 @@ pub use error::MemoryError;
pub use knowledge::KnowledgeStore;
pub use retriever::MemoryRetriever;
pub use store::{InMemoryStore, MemoryStore, SqliteStore};
#[allow(deprecated)]
pub use vector::{InMemoryVectorRetriever, VectorRetriever};
pub use vector_store::{InMemoryVectorStore, PersistentVectorStore, RagPipeline, VectorStore};
// 低频类型(配置/高级使用)
pub use conversation::MemoryStrategy;
+6
View File
@@ -17,6 +17,7 @@ use crate::memory::error::MemoryError;
///
/// **稳定性**:实验性 APIv0.2.x),方法签名可能在 v0.3 中调整。
/// 若未来需要 `remove()` / `clear()` 等方法,将在此 trait 中追加(带默认实现)。
#[deprecated(since = "0.3.0", note = "请使用 memory::VectorStore")]
#[async_trait]
pub trait VectorRetriever: Send + Sync {
/// 将 `id` 对应的文本向量 `embeddings` 加入索引。
@@ -44,10 +45,12 @@ pub trait VectorRetriever: Send + Sync {
/// - 不做向量维度校验(不同维度向量查询结果无意义但不 panic)
/// - `search()` 是 O(n) 全量扫描,未做索引加速
/// - 不保证高并发下查询时序与写入顺序一致
#[deprecated(since = "0.3.0", note = "请使用 memory::InMemoryVectorStore")]
pub struct InMemoryVectorRetriever {
vectors: Mutex<HashMap<String, Vec<f32>>>,
}
#[allow(deprecated)]
impl InMemoryVectorRetriever {
/// 创建空检索器。
pub fn new() -> Self {
@@ -57,12 +60,14 @@ impl InMemoryVectorRetriever {
}
}
#[allow(deprecated)]
impl Default for InMemoryVectorRetriever {
fn default() -> Self {
Self::new()
}
}
#[allow(deprecated)]
#[async_trait]
impl VectorRetriever for InMemoryVectorRetriever {
async fn index(&self, id: String, embeddings: Vec<f32>) -> Result<(), MemoryError> {
@@ -118,6 +123,7 @@ fn dot(a: &[f32], b: &[f32]) -> f32 {
}
#[cfg(test)]
#[allow(deprecated)]
mod tests {
use super::*;
use std::sync::Arc;
+937
View File
@@ -0,0 +1,937 @@
//! 向量存储抽象与实现 —— RAG 管线「存储与检索」环节。
//!
//! 提供 [`VectorStore`] trait 定义、进程内引用实现 [`InMemoryVectorStore`]
//! 以及基于 [`MemoryStore`] 的持久化包装 [`PersistentVectorStore`] 和
//! RAG 管线组合器 [`RagPipeline`]。
//!
//! 下游可实现 [`VectorStore`] trait 以对接专用向量数据库(pgvector / Qdrant 等)。
use std::collections::HashMap;
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 tracing::{debug, info};
use crate::document::{Document, RecursiveCharacterSplitter};
use crate::llm::embedding::Embedding;
use crate::memory::error::MemoryError;
use crate::memory::store::MemoryStore;
use crate::memory::types::{MemoryFilter, MemoryItem};
/// 向量存储抽象 —— 语义检索的核心接口。
///
/// 提供文档-向量的批量添加、余弦相似度搜索、批量删除三个核心操作。
/// 所有实现必须满足 `Send + Sync` 以支持跨 `.await` 调用。
///
/// # 并发安全
///
/// 实现内部必须使用线程安全的容器(如 `Mutex<HashMap>` 或 `RwLock`),
/// 允许跨多个 tokio task 共享 `&VectorStore` 引用。
///
/// # 与旧 `VectorRetriever` 的差异
///
/// - `add` 接受批量 `(doc, embedding)` 对;旧 `index` 仅接受单条
/// - `search` 返回 `(Document, f32)`;旧 `search` 返回 `(String, f32)`,调用方需自行维护 id→Document 映射
#[async_trait]
pub trait VectorStore: Send + Sync {
/// 批量添加文档及其向量。
///
/// `documents` 和 `embeddings` 必须等长。不等长时:
/// - 截取 `min(len)` 对处理(部分写入已发生)
/// - 返回 `Err(MemoryError::InvalidInput)` 告知截断
/// - 调用方可以 `let _ = store.add(...)` 忽略错误
async fn add(
&self,
documents: &[Document],
embeddings: &[Vec<f32>],
) -> Result<(), MemoryError>;
/// 检索与 `query` 向量最相似的 `k` 条记录。
///
/// 返回 `Vec<(Document, f32)>`,其中 `f32` 为余弦相似度分数,
/// 取值范围 `[0.0, 1.0]`(对单位向量),按分数降序排列。
///
/// # 守卫
///
/// - 空索引 → 返回 `vec![]`
/// - `k == 0` → 返回 `vec![]`
/// - 零向量(norm ≈ 0)→ 返回 `vec![]`
async fn search(
&self,
query: &[f32],
k: usize,
) -> Result<Vec<(Document, f32)>, MemoryError>;
/// 批量删除文档(幂等)。
///
/// 不存在的 id 静默忽略,不会返回错误。
async fn remove(&self, ids: &[String]) -> Result<(), MemoryError>;
/// 便捷方法:单条添加。
///
/// 等价于 `self.add(&[doc], &[emb]).await`。
async fn add_one(&self, doc: Document, emb: Vec<f32>) -> Result<(), MemoryError> {
self.add(&[doc], &[emb]).await
}
}
/// 内存向量存储 —— `VectorStore` 的引用实现。
///
/// 内部使用 `Mutex<HashMap<String, (Document, Vec<f32>)>>` 存储,
/// `search()` 执行 O(n) 全量余弦相似度扫描,适用于 ≤10K 条向量的场景。
///
/// # 并发安全
///
/// 使用 `std::sync::Mutex`(非 tokio Mutex)。
///
/// **锁持有时间评估**
/// - `add()` / `remove()`:微秒级(HashMap 插入/删除操作)
/// - `search()`:毫秒级(O(n) 全量扫描 + 余弦计算),对 10K 条 1536 维向量预估 1-10ms。
/// 实现时在锁内克隆数据快照到 `Vec` 后立即释放锁,在锁外进行余弦相似度计算,
/// 避免长时间持有锁阻塞并发写操作。
pub struct InMemoryVectorStore {
entries: Mutex<HashMap<String, (Document, Vec<f32>)>>,
}
impl InMemoryVectorStore {
/// 创建一个空存储。
pub fn new() -> Self {
Self {
entries: Mutex::new(HashMap::new()),
}
}
/// 从预填充的 entries 构造(供 `PersistentVectorStore` 使用)。
pub(crate) fn with_entries(
entries: HashMap<String, (Document, Vec<f32>)>,
) -> Self {
Self {
entries: Mutex::new(entries),
}
}
}
impl Default for InMemoryVectorStore {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl VectorStore for InMemoryVectorStore {
async fn add(
&self,
documents: &[Document],
embeddings: &[Vec<f32>],
) -> Result<(), MemoryError> {
let mut entries = self
.entries
.lock()
.map_err(|e| MemoryError::RetrievalError(format!("lock poisoned: {e}")))?;
let n = documents.len().min(embeddings.len());
if documents.len() != embeddings.len() {
tracing::warn!(
docs = documents.len(),
embs = embeddings.len(),
"InMemoryVectorStore::add 长度不匹配,截断到 min"
);
}
for i in 0..n {
entries.insert(documents[i].id.clone(), (documents[i].clone(), embeddings[i].clone()));
}
if documents.len() != embeddings.len() {
return Err(MemoryError::InvalidInput(format!(
"documents.len()={} 与 embeddings.len()={} 不等,已截断到 min={}",
documents.len(),
embeddings.len(),
n
)));
}
Ok(())
}
async fn search(
&self,
query: &[f32],
k: usize,
) -> Result<Vec<(Document, f32)>, MemoryError> {
if k == 0 {
return Ok(Vec::new());
}
tracing::trace!(k, "InMemoryVectorStore::search");
// 零向量守卫:查询向量本身为零向量则返回空
let query_norm_sq: f32 = query.iter().map(|x| x * x).sum();
if query_norm_sq < 1e-20 {
return Ok(Vec::new());
}
// 锁内克隆快照,释放锁后在锁外计算余弦
let snapshot: Vec<(Document, Vec<f32>)> = {
let entries = self
.entries
.lock()
.map_err(|e| MemoryError::RetrievalError(format!("lock poisoned: {e}")))?;
entries.values().cloned().collect()
};
let mut scored: Vec<(Document, f32)> = Vec::with_capacity(snapshot.len());
for (doc, emb) in snapshot {
let score = cosine_similarity(query, &emb);
scored.push((doc, score));
}
scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
scored.truncate(k);
Ok(scored)
}
async fn remove(&self, ids: &[String]) -> Result<(), MemoryError> {
let mut entries = self
.entries
.lock()
.map_err(|e| MemoryError::RetrievalError(format!("lock poisoned: {e}")))?;
tracing::debug!(count = ids.len(), "InMemoryVectorStore::remove");
entries.retain(|key, _| !ids.iter().any(|id| id == key));
Ok(())
}
}
/// 点积。
///
/// `zip` 对不等长向量静默截断到较短者。调用方应保证 `a` 和 `b` 等长——
/// 不等长时结果无意义但不 panic。
fn dot(a: &[f32], b: &[f32]) -> f32 {
a.iter().zip(b.iter()).map(|(x, y)| x * y).sum()
}
/// 余弦相似度,加 `1e-10` 防除零。
///
/// 零向量与任意向量的相似度返回 `0.0`(因分母中 `1e-10` 保护 + 分子为 0)。
pub(crate) fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
let dot_product = dot(a, b);
let norm_a = dot(a, a).sqrt();
let norm_b = dot(b, b).sqrt();
dot_product / (norm_a * norm_b + 1e-10)
}
/// 持久化向量存储 —— 基于 [`MemoryStore`] 的持久化包装。
///
/// # 架构
///
/// 运行时全量加载到 [`InMemoryVectorStore`] 做余弦搜索,
/// 写操作(add/remove)同时同步到内存和后端 [`MemoryStore`]。
///
/// # 存储格式
///
/// 每条向量存为一条 [`MemoryItem`]
/// - `id`: `"vec:{namespace}:{doc_id}"`colon-separated namespace 前缀)
/// - `content`: JSON 序列化的向量条目(含 doc_id / content / metadata / mime_type / embedding
/// - `metadata`: 空 `serde_json::Value::Null`
///
/// # 构造开销
///
/// `new()` 通过 `store.list(prefix)` 全量加载已有条目,
/// 时间复杂度 O(N)(N 为已有向量数),适用于 ≤10K 条的场景。
pub struct PersistentVectorStore {
inner: InMemoryVectorStore,
store: Arc<dyn MemoryStore>,
namespace: String,
}
/// 持久化向量条目 —— JSON blob 格式。
#[derive(Serialize, Deserialize)]
struct VectorEntry {
doc_id: String,
content: String,
metadata: HashMap<String, String>,
mime_type: String,
embedding: Vec<f32>,
/// ISO 8601 创建时间(UTC),持久化 roundtrip 重建时保持原时间,
/// 避免 MemoryStore 的 TTL 淘汰策略误判。
#[serde(skip_serializing_if = "Option::is_none")]
#[serde(default)]
created_at: Option<String>,
}
impl PersistentVectorStore {
/// 创建新的持久化向量存储,自动从 `store` 全量加载 namespace 下的所有条目。
///
/// `MemoryStore::list()` 由 `SqliteStore` 内部使用 `spawn_blocking` 卸载,
/// 加载过程本身在 async context 中即可,无需额外 spawn_blocking。
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()),
..Default::default()
};
debug!(namespace = %namespace, "PersistentVectorStore::new — 开始全量加载");
let items = store.list(&filter).await?;
info!(count = items.len(), "PersistentVectorStore::new — 加载完成");
let mut entries: HashMap<String, (Document, Vec<f32>)> = HashMap::new();
for item in items {
let entry: VectorEntry = serde_json::from_str(&item.content)
.map_err(|e| MemoryError::Serialization(e.to_string()))?;
let doc = Document {
id: entry.doc_id,
content: entry.content,
metadata: entry.metadata,
mime_type: entry.mime_type,
};
entries.insert(doc.id.clone(), (doc, entry.embedding));
}
info!(entries = entries.len(), "PersistentVectorStore — 内存索引重建完成");
Ok(Self {
inner: InMemoryVectorStore::with_entries(entries),
store,
namespace: namespace.to_string(),
})
}
}
#[async_trait]
impl VectorStore for PersistentVectorStore {
async fn add(
&self,
documents: &[Document],
embeddings: &[Vec<f32>],
) -> Result<(), MemoryError> {
debug!(count = documents.len(), "PersistentVectorStore::add");
// 先逐个写持久化(失败时不污染内存)
for (doc, emb) in documents.iter().zip(embeddings.iter()) {
let entry = VectorEntry {
doc_id: doc.id.clone(),
content: doc.content.clone(),
metadata: doc.metadata.clone(),
mime_type: doc.mime_type.clone(),
embedding: emb.clone(),
created_at: Some(
OffsetDateTime::now_utc()
.format(&Rfc3339)
.map_err(|e| MemoryError::Serialization(format!("format time: {e}")))?,
),
};
let json = serde_json::to_string(&entry)
.map_err(|e| MemoryError::Serialization(e.to_string()))?;
let key = format!("vec:{}:{}", self.namespace, doc.id);
let item = MemoryItem {
id: key,
content: json,
metadata: serde_json::Value::Null,
created_at: OffsetDateTime::now_utc(),
};
self.store.save(item).await?;
}
// 再写内存(持久化已成功写入,内存失败也不影响重启后恢复)
self.inner.add(documents, embeddings).await
}
async fn search(
&self,
query: &[f32],
k: usize,
) -> Result<Vec<(Document, f32)>, MemoryError> {
tracing::trace!(k, "PersistentVectorStore::search");
self.inner.search(query, k).await
}
async fn remove(&self, ids: &[String]) -> Result<(), MemoryError> {
debug!(count = ids.len(), "PersistentVectorStore::remove");
for id in ids {
let key = format!("vec:{}:{}", self.namespace, id);
self.store.delete(&key).await?;
}
self.inner.remove(ids).await
}
}
// ponytail: `with_entries` 当前仅供 `PersistentVectorStore::new` 使用;
// 后续如需 VecStore 之间迁移,可放宽到 `pub`。
/// RAG 管线组合器 —— 封装 `split → embed → store`ingest)和
/// `embed → store.search`retrieve)两个核心流程。
///
/// # 使用方式
///
/// ```ignore
/// let pipeline = RagPipeline::new(embedder, store, Some(splitter));
/// pipeline.ingest(&documents).await?;
/// let results = pipeline.retrieve("query", 5).await?;
/// ```
///
/// # 分割器
///
/// `splitter` 字段为 `Option<RecursiveCharacterSplitter>`
/// - `Some(splitter)` → `ingest()` 先分割再嵌入(调用方传入原始文档)
/// - `None` → `ingest()` 跳过分割,直接嵌入(调用方已分好 chunk)
pub struct RagPipeline {
embedder: Arc<dyn Embedding>,
store: Arc<dyn VectorStore>,
splitter: Option<RecursiveCharacterSplitter>,
}
impl RagPipeline {
/// 创建新的 RAG 管线。
///
/// 不设置分割器时,`ingest()` 跳过分割阶段,
/// 调用方传入的 Document 应已是分割好的 chunk。
pub fn new(
embedder: Arc<dyn Embedding>,
store: Arc<dyn VectorStore>,
splitter: Option<RecursiveCharacterSplitter>,
) -> Self {
Self {
embedder,
store,
splitter,
}
}
/// 摄取文档:分割 → 向量化 → 存储。
///
/// 流程:
/// 1. 如果 splitter 存在,先分割文档为 chunks
/// 2. 提取所有 chunk 的 content 为 `Vec<String>`
/// 3. `embedder.embed()` 批量向量化
/// 4. `store.add()` 批量存储
///
/// # 边界
///
/// - 空文档切片 → `Ok(())`,无操作
/// - 分割后 chunk 为空 → `Ok(())`,无操作
///
/// # 已知限制
///
/// 当前将所有 chunk 一次性传入 `embedder.embed()`,真实 Embedding Provider
/// (如 OpenAI)有批量大小限制,调用方需自行控制单次 ingest 的文档数(如 20 条/批)。
pub async fn ingest(&self, documents: &[Document]) -> Result<(), MemoryError> {
let chunks = match &self.splitter {
Some(splitter) => splitter.split(documents),
None => documents.to_vec(),
};
if chunks.is_empty() {
return Ok(());
}
let texts: Vec<String> = chunks.iter().map(|d| d.content.clone()).collect();
let embeddings = self
.embedder
.embed(&texts)
.await
.map_err(|e| MemoryError::Storage(e.to_string()))?;
self.store.add(&chunks, &embeddings).await
}
/// 检索:向量化查询 → 向量相似度搜索。
///
/// # 边界
///
/// - 空字符串查询 → 返回 `vec![]`embed 产生零向量 → search 零向量守卫)
pub async fn retrieve(
&self,
query: &str,
k: usize,
) -> Result<Vec<(Document, f32)>, MemoryError> {
let embeddings = self
.embedder
.embed(&[query.to_string()])
.await
.map_err(|e| MemoryError::Storage(e.to_string()))?;
self.store.search(&embeddings[0], k).await
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Arc;
use std::time::Duration;
fn make_doc(id: &str, content: &str) -> Document {
Document::from_raw(id, content)
}
fn make_vec(values: &[f32]) -> Vec<f32> {
values.to_vec()
}
#[tokio::test]
async fn basic_add_and_search() {
let store = InMemoryVectorStore::new();
let docs = vec![
make_doc("rust", "Rust language"),
make_doc("python", "Python language"),
make_doc("javascript", "JavaScript language"),
];
let embeddings = vec![
make_vec(&[1.0, 0.0, 0.0]),
make_vec(&[0.0, 1.0, 0.0]),
make_vec(&[0.0, 0.0, 1.0]),
];
store.add(&docs, &embeddings).await.unwrap();
let results = store.search(&[0.9, 0.1, 0.0], 3).await.unwrap();
assert_eq!(results.len(), 3);
assert_eq!(results[0].0.id, "rust", "Top 1 应为 rust");
assert!(results[0].1 > results[1].1);
assert!(results[1].1 > results[2].1);
}
#[tokio::test]
async fn search_empty_store() {
let store = InMemoryVectorStore::new();
let results = store.search(&[1.0, 0.0, 0.0], 5).await.unwrap();
assert!(results.is_empty());
}
#[tokio::test]
async fn search_zero_vector() {
let store = InMemoryVectorStore::new();
let docs = vec![make_doc("a", "alpha")];
let embeddings = vec![make_vec(&[1.0, 0.0, 0.0])];
store.add(&docs, &embeddings).await.unwrap();
let results = store.search(&[0.0, 0.0, 0.0], 5).await.unwrap();
assert!(results.is_empty(), "零向量查询应返回空");
}
#[tokio::test]
async fn search_k_is_zero() {
let store = InMemoryVectorStore::new();
let docs = vec![make_doc("a", "alpha")];
let embeddings = vec![make_vec(&[1.0, 0.0, 0.0])];
store.add(&docs, &embeddings).await.unwrap();
let results = store.search(&[1.0, 0.0, 0.0], 0).await.unwrap();
assert!(results.is_empty(), "k=0 应返回空");
}
#[tokio::test]
async fn search_orthogonal_vectors() {
let store = InMemoryVectorStore::new();
let docs = vec![make_doc("a", "alpha")];
let embeddings = vec![make_vec(&[1.0, 0.0, 0.0])];
store.add(&docs, &embeddings).await.unwrap();
// 正交查询:余弦相似度 ≈ 0,结果仍返回(分数极低)
let results = store.search(&[0.0, 1.0, 0.0], 5).await.unwrap();
assert_eq!(results.len(), 1, "正交向量仍返回,score 接近 0");
assert!(results[0].1 < 1e-10, "正交相似度应约等于 0");
}
#[tokio::test]
async fn add_mismatched_lengths() {
let store = InMemoryVectorStore::new();
let docs = vec![
make_doc("a", "alpha"),
make_doc("b", "beta"),
make_doc("c", "gamma"),
];
let embeddings = vec![make_vec(&[1.0, 0.0, 0.0]), make_vec(&[0.0, 1.0, 0.0])];
let result = store.add(&docs, &embeddings).await;
assert!(result.is_err(), "不等长应返回 Err");
// 部分写入已发生:前 2 条已写入
let results = store.search(&[1.0, 0.0, 0.0], 5).await.unwrap();
assert_eq!(results.len(), 2, "应有 2 条成功写入");
let ids: Vec<&str> = results.iter().map(|(d, _)| d.id.as_str()).collect();
assert!(ids.contains(&"a"));
assert!(ids.contains(&"b"));
}
#[tokio::test]
async fn add_duplicate_id_upsert() {
let store = InMemoryVectorStore::new();
let docs_v1 = vec![make_doc("a", "v1 content")];
let embeddings_v1 = vec![make_vec(&[1.0, 0.0, 0.0])];
store.add(&docs_v1, &embeddings_v1).await.unwrap();
// 同一 doc.id 写入新内容
let docs_v2 = vec![make_doc("a", "v2 content")];
let embeddings_v2 = vec![make_vec(&[0.0, 1.0, 0.0])];
store.add(&docs_v2, &embeddings_v2).await.unwrap();
let results = store.search(&[0.9, 0.1, 0.0], 5).await.unwrap();
assert_eq!(results.len(), 1, "重复 id 写入应覆盖,最终仅 1 条");
assert_eq!(results[0].0.content, "v2 content", "新内容应覆盖旧内容");
}
#[tokio::test]
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]),
];
store.add(&docs, &embeddings).await.unwrap();
store.remove(&["a".to_string()]).await.unwrap();
let results = store.search(&[1.0, 0.0, 0.0], 5).await.unwrap();
assert_eq!(results.len(), 1);
assert_eq!(results[0].0.id, "b");
}
#[tokio::test]
async fn remove_nonexistent_id() {
let store = InMemoryVectorStore::new();
// 从未添加的 id 应静默忽略
let result = store.remove(&["nonexistent".to_string()]).await;
assert!(result.is_ok(), "删除不存在的 id 不应报错");
// 已有索引时也不应报错
let docs = vec![make_doc("a", "alpha")];
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;
assert!(result.is_ok(), "批量删除不存在 id 不应报错");
let results = store.search(&[1.0, 0.0, 0.0], 5).await.unwrap();
assert_eq!(results.len(), 1, "原有数据应保留");
}
#[tokio::test]
async fn concurrent_operations() {
let store = Arc::new(InMemoryVectorStore::new());
let mut handles = Vec::new();
// 10 个并发写入
for i in 0..10 {
let s = Arc::clone(&store);
handles.push(tokio::spawn(async move {
let docs = vec![make_doc(&format!("item_{i}"), &format!("content_{i}"))];
let embeddings = vec![make_vec(&[i as f32, 0.0, 0.0])];
s.add(&docs, &embeddings).await.unwrap();
}));
}
for h in handles.drain(..) {
h.await.unwrap();
}
// 验证并发写入后 search 结果计数正确
let results = store.search(&[1.0, 0.0, 0.0], 20).await.unwrap();
assert_eq!(results.len(), 10, "并发 add 10 条后应能检索到 10 条");
// 混合写入 + 搜索的并发(无 panic)
let deadline = tokio::time::Instant::now() + Duration::from_millis(100);
let mut handles = Vec::new();
for w in 0..3 {
let s = Arc::clone(&store);
handles.push(tokio::spawn(async move {
let mut i = 0;
while tokio::time::Instant::now() < deadline {
let docs = vec![make_doc(&format!("w{w}_i{i}"), "x")];
let embeddings = vec![make_vec(&[i as f32, 0.0, 0.0])];
let _ = s.add(&docs, &embeddings).await;
let _ = s.search(&[1.0, 0.0, 0.0], 3).await;
i += 1;
}
}));
}
for h in handles {
h.await.unwrap();
}
}
// ===== Persistent tests =====
use crate::memory::store::InMemoryStore;
async fn make_persistent(
backend: Arc<dyn MemoryStore>,
namespace: &str,
) -> PersistentVectorStore {
PersistentVectorStore::new(backend, namespace).await.unwrap()
}
#[tokio::test]
async fn persistent_roundtrip() {
let backend: Arc<dyn MemoryStore> = Arc::new(InMemoryStore::new());
let store = make_persistent(Arc::clone(&backend), "default").await;
let docs = vec![
make_doc("a", "alpha"),
make_doc("b", "beta"),
make_doc("c", "gamma"),
];
let embeddings = vec![
make_vec(&[1.0, 0.0, 0.0]),
make_vec(&[0.0, 1.0, 0.0]),
make_vec(&[0.0, 0.0, 1.0]),
];
store.add(&docs, &embeddings).await.unwrap();
// 重建 store(模拟重启)
let store2 = make_persistent(Arc::clone(&backend), "default").await;
let results = store2.search(&[0.9, 0.1, 0.0], 5).await.unwrap();
assert_eq!(results.len(), 3);
assert_eq!(results[0].0.id, "a", "Top 1 应为 a(与 [1,0,0] 最相似)");
}
#[tokio::test]
async fn search_after_reload() {
let backend: Arc<dyn MemoryStore> = Arc::new(InMemoryStore::new());
let store = make_persistent(Arc::clone(&backend), "default").await;
let docs = vec![make_doc("target", "the target doc")];
let embeddings = vec![make_vec(&[1.0, 0.0, 0.0])];
store.add(&docs, &embeddings).await.unwrap();
// 重建
let store2 = make_persistent(Arc::clone(&backend), "default").await;
let results = store2.search(&[0.99, 0.01, 0.0], 1).await.unwrap();
assert_eq!(results.len(), 1);
assert_eq!(results[0].0.id, "target");
assert!(results[0].1 > 0.99);
}
#[tokio::test]
async fn namespace_isolation() {
let backend: Arc<dyn MemoryStore> = Arc::new(InMemoryStore::new());
let s1 = make_persistent(Arc::clone(&backend), "ns1").await;
let s2 = make_persistent(Arc::clone(&backend), "ns2").await;
let docs = vec![make_doc("shared_id", "content")];
let embeddings = vec![make_vec(&[1.0, 0.0, 0.0])];
s1.add(&docs, &embeddings).await.unwrap();
// s1 能检索到
let r1 = s1.search(&[1.0, 0.0, 0.0], 5).await.unwrap();
assert_eq!(r1.len(), 1);
// s2 在 ns2 下,shared_id 不属于 ns2,应检索不到
let r2 = s2.search(&[1.0, 0.0, 0.0], 5).await.unwrap();
assert!(r2.is_empty(), "不同 namespace 应隔离");
}
#[tokio::test]
async fn concurrent_access() {
let backend: Arc<dyn MemoryStore> = Arc::new(InMemoryStore::new());
let store = Arc::new(make_persistent(Arc::clone(&backend), "default").await);
let mut handles = Vec::new();
for i in 0..5 {
let s = Arc::clone(&store);
handles.push(tokio::spawn(async move {
let docs = vec![make_doc(&format!("concurrent_{i}"), "x")];
let embeddings = vec![make_vec(&[i as f32, 0.0, 0.0])];
s.add(&docs, &embeddings).await.unwrap();
}));
}
for _w in 0..3 {
let s = Arc::clone(&store);
handles.push(tokio::spawn(async move {
let _ = s.search(&[1.0, 0.0, 0.0], 10).await.unwrap();
}));
}
for h in handles {
h.await.unwrap();
}
let results = store.search(&[1.0, 0.0, 0.0], 20).await.unwrap();
assert_eq!(results.len(), 5, "并发写入 5 条后应能检索到 5 条");
}
#[tokio::test]
async fn partial_add_recovery() {
// 写入 5 条,模拟第 3 条持久化失败(通过底层 InMemoryStore 的 save 拦截)
let backend: Arc<dyn MemoryStore> = Arc::new(InMemoryStore::new());
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 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();
// 重建 store,确认前 2 条已持久化
let store2 = make_persistent(Arc::clone(&backend), "default").await;
let results = store2.search(&[1.0, 0.0, 0.0], 10).await.unwrap();
assert_eq!(results.len(), 2, "前 2 条应已持久化并能加载");
let ids: Vec<&str> = results.iter().map(|(d, _)| d.id.as_str()).collect();
assert!(ids.contains(&"doc_0"));
assert!(ids.contains(&"doc_1"));
}
#[tokio::test]
async fn new_empty_store() {
// 空后端构造 PersistentVectorStore 应成功,且 search 返回空
let backend: Arc<dyn MemoryStore> = Arc::new(InMemoryStore::new());
let store = make_persistent(Arc::clone(&backend), "empty_ns").await;
let results = store.search(&[1.0, 0.0, 0.0], 5).await.unwrap();
assert!(results.is_empty(), "空存储 search 应返回空");
// 写入后能检索
let docs = vec![make_doc("after_empty", "data")];
let embeddings = vec![make_vec(&[1.0, 0.0, 0.0])];
store.add(&docs, &embeddings).await.unwrap();
let results = store.search(&[1.0, 0.0, 0.0], 5).await.unwrap();
assert_eq!(results.len(), 1);
}
// ===== RagPipeline tests =====
use crate::llm::embedding::MockEmbedding;
#[tokio::test]
async fn ingest_and_retrieve() {
let embedder: Arc<dyn Embedding> = Arc::new(MockEmbedding::new(4));
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 doc = Document::new(
"rag-doc",
"Rust 是一门系统编程语言。\n\n\
Rust 通过所有权系统管理内存,无需垃圾回收器。\n\n\
Cargo 是官方的构建系统和包管理器。",
"text/markdown",
);
pipeline.ingest(&[doc]).await.unwrap();
// 用第一个 chunk 的 content 检索(应能命中自己或相关 chunk)
let docs_stored = store.search(&[1.0, 0.0, 0.0, 0.0], 100).await.unwrap();
assert!(!docs_stored.is_empty(), "ingest 后 store 应有数据");
// retrieve 测试
let results = pipeline.retrieve("Rust ownership", 3).await.unwrap();
assert!(!results.is_empty(), "retrieve 应返回结果");
// 验证返回的 Document.id 是 chunk id 格式(来自 splitter
for (doc, _score) in &results {
assert!(
doc.id.starts_with("rag-doc:chunk:"),
"chunk id 格式应为 rag-doc:chunk:NNNN,实际: {}",
doc.id
);
}
}
#[tokio::test]
async fn retrieve_empty_store() {
let embedder: Arc<dyn Embedding> = Arc::new(MockEmbedding::new(4));
let store: Arc<dyn VectorStore> = Arc::new(InMemoryVectorStore::new());
let pipeline = RagPipeline::new(embedder, store, None);
let results = pipeline.retrieve("anything", 5).await.unwrap();
assert!(results.is_empty(), "空 store retrieve 应返回空");
}
#[tokio::test]
async fn ingest_empty_docs() {
let embedder: Arc<dyn Embedding> = Arc::new(MockEmbedding::new(4));
let store: Arc<dyn VectorStore> = Arc::new(InMemoryVectorStore::new());
let pipeline = RagPipeline::new(embedder, Arc::clone(&store), None);
// 空切片应返回 Ok(()),不报错
let result = pipeline.ingest(&[]).await;
assert!(result.is_ok(), "空文档切片 ingest 应返回 Ok");
// 验证 store 中没有数据
let results = store.search(&[1.0, 0.0, 0.0, 0.0], 5).await.unwrap();
assert!(results.is_empty(), "空 ingest 后 store 应为空");
}
#[tokio::test]
async fn ingest_empty_split() {
let embedder: Arc<dyn Embedding> = Arc::new(MockEmbedding::new(4));
let store: Arc<dyn VectorStore> = Arc::new(InMemoryVectorStore::new());
// splitter 分割空内容文档
let splitter = RecursiveCharacterSplitter::new(50, 5);
let pipeline = RagPipeline::new(embedder, Arc::clone(&store), Some(splitter));
// 传入一个空内容文档,splitter 应返回空 chunks
let empty_doc = Document::from_raw("empty_id", "");
let result = pipeline.ingest(&[empty_doc]).await;
assert!(result.is_ok(), "空内容 split 后 ingest 应返回 Ok");
let results = store.search(&[1.0, 0.0, 0.0, 0.0], 5).await.unwrap();
assert!(results.is_empty(), "空 split 后 store 应为空");
}
// ===== Performance benchmarks (Step 15.6.7) =====
/// 性能基准:InMemoryVectorStore::search 在 10K 条 64 维向量索引上搜索耗时 < 100ms。
/// ponytail: 本测试作为性能下限断言(非精确基准),CI 环境性能差异可通过调整阈值补偿。
#[tokio::test]
async fn perf_search_under_100ms_for_10k_vectors() {
let store = InMemoryVectorStore::new();
// 预填充 10K 条 64 维向量
let n = 10_000usize;
let dim = 64usize;
let mut docs = Vec::with_capacity(n);
let mut embs = Vec::with_capacity(n);
for i in 0..n {
docs.push(make_doc(&format!("d{i}"), "x"));
let v: Vec<f32> = (0..dim).map(|j| ((i + j) as f32).sin()).collect();
embs.push(v);
}
store.add(&docs, &embs).await.unwrap();
// 性能断言
let start = std::time::Instant::now();
let _results = store.search(&vec![1.0_f32; dim], 10).await.unwrap();
let elapsed = start.elapsed();
assert!(
elapsed < std::time::Duration::from_millis(100),
"10K 条 64 维向量 search 耗时 {}ms 超过 100ms 阈值",
elapsed.as_millis()
);
}
/// 性能基准:PersistentVectorStore::new 加载 10K 条 < 500ms。
#[tokio::test]
async fn perf_persistent_load_under_500ms_for_10k() {
let backend: Arc<dyn MemoryStore> = Arc::new(InMemoryStore::new());
let store = make_persistent(Arc::clone(&backend), "perf_ns").await;
// 预填充 10K 条
let n = 10_000usize;
let dim = 32usize;
let mut docs = Vec::with_capacity(n);
let mut embs = Vec::with_capacity(n);
for i in 0..n {
docs.push(make_doc(&format!("d{i}"), "x"));
let v: Vec<f32> = (0..dim).map(|j| ((i + j) as f32).cos()).collect();
embs.push(v);
}
store.add(&docs, &embs).await.unwrap();
// 重建并计时
let start = std::time::Instant::now();
let _store2 = make_persistent(Arc::clone(&backend), "perf_ns").await;
let elapsed = start.elapsed();
assert!(
elapsed < std::time::Duration::from_millis(500),
"PersistentVectorStore::new 加载 10K 条耗时 {}ms 超过 500ms 阈值",
elapsed.as_millis()
);
}
}