feat(core): 完成 Phase 11 测试与检索补强

- 新增 VectorRetriever trait 与 InMemoryVectorRetriever 引用实现
- 补充 Provider roundtrip wiremock 测试与 MemoryStore 并发测试共 23 个
- 修复 openai 429 retry-after header 解析(与 anthropic 对齐)
This commit is contained in:
徐涛
2026-07-06 14:52:49 +08:00
parent 2af92cd554
commit 71abe881ed
6 changed files with 912 additions and 7 deletions
+100
View File
@@ -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()
);
}
}
+78
View File
@@ -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();
}
}
}
+238
View File
@@ -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 + 余弦相似度。
///
/// **稳定性**:实验性 APIv0.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();
}
}
}