feat(core): 完成 Phase 11 测试与检索补强
- 新增 VectorRetriever trait 与 InMemoryVectorRetriever 引用实现 - 补充 Provider roundtrip wiremock 测试与 MemoryStore 并发测试共 23 个 - 修复 openai 429 retry-after header 解析(与 anthropic 对齐)
This commit is contained in:
@@ -263,4 +263,104 @@ mod tests {
|
||||
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
||||
assert_eq!(list.len(), 100);
|
||||
}
|
||||
|
||||
// ===== Phase 11 Step 11.3 并发测试 =====
|
||||
|
||||
#[tokio::test]
|
||||
async fn concurrent_writers_max_pressure() {
|
||||
use std::sync::Arc;
|
||||
let store = Arc::new(InMemoryStore::new());
|
||||
|
||||
let mut handles = Vec::new();
|
||||
for i in 0..100 {
|
||||
let s = Arc::clone(&store);
|
||||
handles.push(tokio::spawn(async move {
|
||||
let id = format!("concurrent_{i}");
|
||||
s.save(make_item(&id)).await.unwrap();
|
||||
}));
|
||||
}
|
||||
for h in handles {
|
||||
h.await.unwrap();
|
||||
}
|
||||
|
||||
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
||||
assert_eq!(list.len(), 100);
|
||||
let mut ids: Vec<String> = list.iter().map(|v| v.id.clone()).collect();
|
||||
ids.sort();
|
||||
ids.dedup();
|
||||
assert_eq!(ids.len(), 100);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn concurrent_mixed_read_write() {
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
|
||||
let store = Arc::new(InMemoryStore::new());
|
||||
|
||||
// 预热 20 条
|
||||
for i in 0..20 {
|
||||
store.save(make_item(&format!("seed_{i}"))).await.unwrap();
|
||||
}
|
||||
|
||||
let deadline = tokio::time::Instant::now() + Duration::from_secs(2);
|
||||
let mut handles = Vec::new();
|
||||
|
||||
// 5 个写者
|
||||
for w in 0..5 {
|
||||
let s = Arc::clone(&store);
|
||||
handles.push(tokio::spawn(async move {
|
||||
let mut i = 0;
|
||||
while tokio::time::Instant::now() < deadline {
|
||||
let id = format!("writer{w}_item{i}");
|
||||
s.save(make_item(&id)).await.unwrap();
|
||||
i += 1;
|
||||
}
|
||||
}));
|
||||
}
|
||||
|
||||
// 5 个读者
|
||||
for _ in 0..5 {
|
||||
let s = Arc::clone(&store);
|
||||
handles.push(tokio::spawn(async move {
|
||||
while tokio::time::Instant::now() < deadline {
|
||||
let _ = s.list(&MemoryFilter::default()).await.unwrap();
|
||||
}
|
||||
}));
|
||||
}
|
||||
|
||||
for h in handles {
|
||||
h.await.unwrap();
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn concurrent_capacity_eviction() {
|
||||
use std::sync::Arc;
|
||||
|
||||
let eviction = EvictionConfig {
|
||||
policy: EvictionPolicy::Capacity { max_items: 10 },
|
||||
check_interval: 1,
|
||||
};
|
||||
let store = Arc::new(InMemoryStore::with_eviction(eviction));
|
||||
|
||||
let mut handles = Vec::new();
|
||||
for i in 0..15 {
|
||||
let s = Arc::clone(&store);
|
||||
handles.push(tokio::spawn(async move {
|
||||
s.save(make_item(&format!("item_{i}"))).await.unwrap();
|
||||
}));
|
||||
}
|
||||
for h in handles {
|
||||
h.await.unwrap();
|
||||
}
|
||||
|
||||
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
||||
// 写者全部完成后必 ≤ max_items(部分路径上可能短暂 >10 但全部完成时应 ≤10)
|
||||
assert!(
|
||||
list.len() <= 10,
|
||||
"expected <= 10 items after all writers done, got {}",
|
||||
list.len()
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -542,4 +542,82 @@ mod tests {
|
||||
assert!(store.get("x").await.unwrap().is_none());
|
||||
}
|
||||
}
|
||||
|
||||
// ===== Phase 11 Step 11.3 并发测试 =====
|
||||
|
||||
#[tokio::test]
|
||||
async fn concurrent_writers_max_pressure() {
|
||||
use std::time::Duration;
|
||||
let store = Arc::new(SqliteStore::open(":memory:").unwrap());
|
||||
|
||||
let mut handles = Vec::new();
|
||||
for i in 0..100 {
|
||||
let s = Arc::clone(&store);
|
||||
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
|
||||
{
|
||||
Ok(res) => res.unwrap(),
|
||||
Err(_) => panic!("save({id}) timed out under 100-way concurrency"),
|
||||
}
|
||||
}));
|
||||
}
|
||||
for h in handles {
|
||||
h.await.unwrap();
|
||||
}
|
||||
|
||||
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
||||
assert_eq!(list.len(), 100);
|
||||
let mut ids: Vec<String> = list.iter().map(|v| v.id.clone()).collect();
|
||||
ids.sort();
|
||||
ids.dedup();
|
||||
assert_eq!(ids.len(), 100);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn concurrent_mixed_read_write() {
|
||||
use std::time::Duration;
|
||||
|
||||
let store = Arc::new(SqliteStore::open(":memory:").unwrap());
|
||||
|
||||
// 预热 20 条
|
||||
for i in 0..20 {
|
||||
store.save(make_item(&format!("seed_{i}"))).await.unwrap();
|
||||
}
|
||||
|
||||
let deadline = tokio::time::Instant::now() + Duration::from_secs(2);
|
||||
let mut handles = Vec::new();
|
||||
|
||||
// 5 个写者
|
||||
for w in 0..5 {
|
||||
let s = Arc::clone(&store);
|
||||
handles.push(tokio::spawn(async move {
|
||||
let mut i = 0;
|
||||
while tokio::time::Instant::now() < deadline {
|
||||
let id = format!("writer{w}_item{i}");
|
||||
s.save(make_item(&id)).await.unwrap();
|
||||
i += 1;
|
||||
}
|
||||
}));
|
||||
}
|
||||
|
||||
// 5 个读者
|
||||
for _ in 0..5 {
|
||||
let s = Arc::clone(&store);
|
||||
handles.push(tokio::spawn(async move {
|
||||
while tokio::time::Instant::now() < deadline {
|
||||
let _ = s.list(&MemoryFilter::default()).await.unwrap();
|
||||
}
|
||||
}));
|
||||
}
|
||||
|
||||
for h in handles {
|
||||
h.await.unwrap();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,238 @@
|
||||
//! 语义向量检索抽象。
|
||||
//!
|
||||
//! 提供 [`VectorRetriever`] trait 定义与进程内引用实现 [`InMemoryVectorRetriever`]。
|
||||
//! 下游可实现此 trait 以对接向量数据库(pgvector / qdrant / lancedb 等)。
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Mutex;
|
||||
|
||||
use async_trait::async_trait;
|
||||
|
||||
use crate::memory::error::MemoryError;
|
||||
|
||||
/// 语义向量检索器抽象接口。
|
||||
///
|
||||
/// 下游可实现此 trait 以对接向量数据库(pgvector / qdrant / lancedb 等)。
|
||||
/// 默认引用实现 [`InMemoryVectorRetriever`] 基于进程内 HashMap + 余弦相似度。
|
||||
///
|
||||
/// **稳定性**:实验性 API(v0.2.x),方法签名可能在 v0.3 中调整。
|
||||
/// 若未来需要 `remove()` / `clear()` 等方法,将在此 trait 中追加(带默认实现)。
|
||||
#[async_trait]
|
||||
pub trait VectorRetriever: Send + Sync {
|
||||
/// 将 `id` 对应的文本向量 `embeddings` 加入索引。
|
||||
///
|
||||
/// 重复调用同一 `id` 会覆盖已有向量。调用方负责保证 `embeddings` 维度
|
||||
/// 与已索引向量一致——本 trait 不做维度校验。
|
||||
async fn index(&self, id: String, embeddings: Vec<f32>) -> Result<(), MemoryError>;
|
||||
|
||||
/// 检索与 `query` 向量最相似的 `k` 条记录。
|
||||
///
|
||||
/// 返回 `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>;
|
||||
}
|
||||
|
||||
/// 进程内向量检索器 —— 基于 HashMap + 全量余弦相似度扫描。
|
||||
///
|
||||
/// 适用场景:单元测试、小规模验证(<10K 向量)。生产环境请对接真正的向量数据库。
|
||||
///
|
||||
/// **不保证**:
|
||||
/// - 不做向量维度校验(不同维度向量查询结果无意义但不 panic)
|
||||
/// - `search()` 是 O(n) 全量扫描,未做索引加速
|
||||
/// - 不保证高并发下查询时序与写入顺序一致
|
||||
pub struct InMemoryVectorRetriever {
|
||||
vectors: Mutex<HashMap<String, Vec<f32>>>,
|
||||
}
|
||||
|
||||
impl InMemoryVectorRetriever {
|
||||
/// 创建空检索器。
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
vectors: Mutex::new(HashMap::new()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for InMemoryVectorRetriever {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl VectorRetriever for InMemoryVectorRetriever {
|
||||
async fn index(&self, id: String, embeddings: Vec<f32>) -> Result<(), MemoryError> {
|
||||
let mut vectors = self
|
||||
.vectors
|
||||
.lock()
|
||||
.map_err(|e| MemoryError::RetrievalError(format!("lock poisoned: {e}")))?;
|
||||
vectors.insert(id, embeddings);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn search(
|
||||
&self,
|
||||
query: Vec<f32>,
|
||||
k: usize,
|
||||
) -> Result<Vec<(String, f32)>, MemoryError> {
|
||||
let vectors = self
|
||||
.vectors
|
||||
.lock()
|
||||
.map_err(|e| MemoryError::RetrievalError(format!("lock poisoned: {e}")))?;
|
||||
|
||||
if vectors.is_empty() || k == 0 {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
|
||||
let query_norm = dot(&query, &query).sqrt();
|
||||
if query_norm == 0.0 {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
|
||||
let mut scored: Vec<(String, f32)> = vectors
|
||||
.iter()
|
||||
.map(|(id, vec)| {
|
||||
let dot_product = dot(&query, vec);
|
||||
let vec_norm = dot(vec, vec).sqrt();
|
||||
let similarity = dot_product / (query_norm * vec_norm + 1e-10);
|
||||
(id.clone(), similarity)
|
||||
})
|
||||
.collect();
|
||||
|
||||
scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
|
||||
scored.truncate(k);
|
||||
Ok(scored)
|
||||
}
|
||||
}
|
||||
|
||||
/// 点积(手动循环,零依赖)。
|
||||
///
|
||||
/// 注意:`zip` 对不等长向量静默截断到较短者。引用实现不做维度校验,
|
||||
/// 调用方应确保 `a` 和 `b` 等长——不等长时结果无意义但不 panic。
|
||||
fn dot(a: &[f32], b: &[f32]) -> f32 {
|
||||
a.iter().zip(b.iter()).map(|(x, y)| x * y).sum()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
|
||||
#[tokio::test]
|
||||
async fn basic_index_and_search() {
|
||||
let retriever = InMemoryVectorRetriever::new();
|
||||
retriever
|
||||
.index("rust".into(), vec![1.0, 0.0, 0.0])
|
||||
.await
|
||||
.unwrap();
|
||||
retriever
|
||||
.index("python".into(), vec![0.0, 1.0, 0.0])
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let results = retriever.search(vec![0.9, 0.1, 0.0], 2).await.unwrap();
|
||||
assert_eq!(results.len(), 2);
|
||||
assert_eq!(results[0].0, "rust");
|
||||
assert!(results[0].1 > results[1].1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn search_empty_store() {
|
||||
let retriever = InMemoryVectorRetriever::new();
|
||||
let results = retriever.search(vec![1.0, 0.0, 0.0], 5).await.unwrap();
|
||||
assert!(results.is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn search_zero_vector_returns_empty() {
|
||||
let retriever = InMemoryVectorRetriever::new();
|
||||
retriever
|
||||
.index("a".into(), vec![1.0, 0.0, 0.0])
|
||||
.await
|
||||
.unwrap();
|
||||
let results = retriever.search(vec![0.0, 0.0, 0.0], 5).await.unwrap();
|
||||
assert!(results.is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn search_with_k_zero_returns_empty() {
|
||||
let retriever = InMemoryVectorRetriever::new();
|
||||
retriever
|
||||
.index("a".into(), vec![1.0, 0.0, 0.0])
|
||||
.await
|
||||
.unwrap();
|
||||
let results = retriever.search(vec![1.0, 0.0, 0.0], 0).await.unwrap();
|
||||
assert!(results.is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn concurrent_index() {
|
||||
let retriever = Arc::new(InMemoryVectorRetriever::new());
|
||||
let mut handles = Vec::new();
|
||||
for i in 0..10 {
|
||||
let r = Arc::clone(&retriever);
|
||||
handles.push(tokio::spawn(async move {
|
||||
r.index(format!("item_{i}"), vec![i as f32, 0.0, 0.0])
|
||||
.await
|
||||
.unwrap();
|
||||
}));
|
||||
}
|
||||
for h in handles {
|
||||
h.await.unwrap();
|
||||
}
|
||||
|
||||
let results = retriever.search(vec![1.0, 0.0, 0.0], 20).await.unwrap();
|
||||
assert_eq!(results.len(), 10);
|
||||
let mut ids: Vec<String> = results.iter().map(|(id, _)| id.clone()).collect();
|
||||
ids.sort();
|
||||
ids.dedup();
|
||||
assert_eq!(ids.len(), 10);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn concurrent_index_and_search() {
|
||||
let retriever = Arc::new(InMemoryVectorRetriever::new());
|
||||
|
||||
for i in 0..5 {
|
||||
retriever
|
||||
.index(format!("seed_{i}"), vec![i as f32, 0.0, 0.0])
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
let deadline = tokio::time::Instant::now() + Duration::from_secs(2);
|
||||
|
||||
let mut handles = Vec::new();
|
||||
|
||||
for w in 0..5 {
|
||||
let r = Arc::clone(&retriever);
|
||||
handles.push(tokio::spawn(async move {
|
||||
let mut i = 0;
|
||||
while tokio::time::Instant::now() < deadline {
|
||||
r.index(format!("writer{w}_{i}"), vec![i as f32, 0.0, 0.0])
|
||||
.await
|
||||
.unwrap();
|
||||
i += 1;
|
||||
}
|
||||
}));
|
||||
}
|
||||
|
||||
for _ in 0..5 {
|
||||
let r = Arc::clone(&retriever);
|
||||
handles.push(tokio::spawn(async move {
|
||||
while tokio::time::Instant::now() < deadline {
|
||||
let _ = r.search(vec![1.0, 0.0, 0.0], 3).await.unwrap();
|
||||
}
|
||||
}));
|
||||
}
|
||||
|
||||
for h in handles {
|
||||
h.await.unwrap();
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user