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
+131 -1
View File
@@ -901,7 +901,7 @@ mod tests {
use super::*;
use crate::llm::types::request_v2::MessageRequest;
use serde_json::json;
use wiremock::matchers::{method, path};
use wiremock::matchers::{header, method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
fn make_provider(base_url: String) -> AnthropicProvider {
@@ -1077,4 +1077,134 @@ event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n";
assert_eq!(body.max_tokens, DEFAULT_MAX_TOKENS);
assert_eq!(body.model, "claude-sonnet-4-20250514");
}
// ===== Phase 11 Step 11.2 wiremock roundtrip 测试 =====
#[tokio::test]
async fn anthropic_401_structured_error() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/messages"))
.respond_with(ResponseTemplate::new(401).set_body_json(json!({
"type": "error",
"error": {
"type": "authentication_error",
"message": "Invalid API key provided: sk-ant-test"
}
})))
.mount(&server)
.await;
let provider = make_provider(server.uri());
let err = provider
.chat_blocking(MessageRequest {
model: "claude-sonnet-4-20250514".into(),
messages: vec![Message::user_text("Hi")],
..Default::default()
})
.await
.unwrap_err();
match err {
LlmError::Authentication(msg) => assert!(msg.contains("Invalid API key")),
other => panic!("expected Authentication, got {other:?}"),
}
}
#[tokio::test]
async fn anthropic_tool_use_response() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/messages"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
"id": "msg_tool",
"type": "message",
"model": "claude-sonnet-4-20250514",
"content": [
{"type": "text", "text": "Let me check."},
{"type": "tool_use", "id": "toolu_abc", "name": "lookup", "input": {"q": "rust"}}
],
"stop_reason": "tool_use",
"usage": {"input_tokens": 8, "output_tokens": 12}
})))
.mount(&server)
.await;
let provider = make_provider(server.uri());
let response = provider
.chat_blocking(MessageRequest {
model: "claude-sonnet-4-20250514".into(),
messages: vec![Message::user_text("Look up rust")],
..Default::default()
})
.await
.unwrap();
assert_eq!(response.stop_reason, StopReason::ToolUse);
let tool_use = match &response.message {
Message::Assistant { content } => content.iter().find_map(|b| match b {
ContentBlock::ToolUse { id, name, .. } => Some((id.clone(), name.clone())),
_ => None,
}),
_ => None,
};
assert_eq!(tool_use, Some(("toolu_abc".into(), "lookup".into())));
}
#[tokio::test]
async fn anthropic_version_header() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/messages"))
.and(header("anthropic-version", "2023-06-01"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
"id": "msg_v",
"type": "message",
"model": "claude-sonnet-4-20250514",
"content": [{"type": "text", "text": "OK"}],
"stop_reason": "end_turn",
"usage": {"input_tokens": 1, "output_tokens": 1}
})))
.mount(&server)
.await;
let provider = make_provider(server.uri());
let response = provider
.chat_blocking(MessageRequest {
model: "claude-sonnet-4-20250514".into(),
messages: vec![Message::user_text("Hi")],
..Default::default()
})
.await
.unwrap();
assert_eq!(response.text(), "OK");
}
#[tokio::test]
async fn anthropic_529_overloaded_structured() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/messages"))
.respond_with(ResponseTemplate::new(529).set_body_json(json!({
"type": "error",
"error": {
"type": "overloaded_error",
"message": "Overloaded: Anthropic API is temporarily overloaded"
}
})))
.mount(&server)
.await;
let provider = make_provider(server.uri());
let err = provider
.chat_blocking(MessageRequest {
model: "claude-sonnet-4-20250514".into(),
messages: vec![Message::user_text("Hi")],
..Default::default()
})
.await
.unwrap_err();
match err {
LlmError::RateLimit { retry_after } => assert!(retry_after.is_none()),
other => panic!("expected RateLimit, got {other:?}"),
}
}
}
+363 -6
View File
@@ -180,15 +180,17 @@ impl GenericOpenaiProvider {
/// (无法解析为 JSON),此处直接用 status code + 原始 body 兜底。
async fn handle_error_response(response: reqwest::Response) -> LlmError {
let status = response.status().as_u16();
let retry_after = response
.headers()
.get("retry-after")
.and_then(|v| v.to_str().ok())
.and_then(|v| v.parse::<u64>().ok())
.map(std::time::Duration::from_secs);
let body = response.text().await.unwrap_or_default();
match status {
401 => LlmError::Authentication(body),
429 => {
// ponytail: 仅读取 retry-after,不在 OpenAI-compatible 上假设格式
// 与 OpenAI 完全一致;DeepSeek/Qwen 通常遵循。
LlmError::RateLimit { retry_after: None }
}
429 => LlmError::RateLimit { retry_after },
_ if status >= 500 => LlmError::Request { status, body },
_ if status == 400 && body.contains("context_length_exceeded") => {
LlmError::ContextLength {
@@ -818,7 +820,8 @@ mod tests {
use crate::llm::convert::content_to_blocks;
use crate::llm::types::usage::Usage;
use serde_json::json;
use wiremock::matchers::{method, path};
use std::time::Duration;
use wiremock::matchers::{body_partial_json, header, method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
#[tokio::test]
@@ -1031,4 +1034,358 @@ data: [DONE]\n\n";
assert_eq!(blocks.len(), 1);
assert!(matches!(blocks[0], ContentBlock::Text { ref text } if text == "plain text"));
}
// ===== Phase 11 Step 11.2 wiremock roundtrip 测试 =====
#[tokio::test]
async fn openai_request_body_format() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/chat/completions"))
.and(body_partial_json(json!({
"model": "gpt-4o",
"messages": [{"role": "user", "content": "hi"}]
})))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
"id": "chatcmpl-body",
"object": "chat.completion",
"created": 1,
"model": "gpt-4o",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}
})))
.mount(&server)
.await;
let provider = GenericOpenaiProvider::new_with_name(
server.uri(),
"sk-test".into(),
"gpt-4o".into(),
"openai",
30,
);
let response = provider
.chat_blocking(MessageRequest {
model: "gpt-4o".into(),
messages: vec![Message::user_text("hi")],
..Default::default()
})
.await
.unwrap();
assert_eq!(response.text(), "ok");
}
#[tokio::test]
async fn openai_authorization_header() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/chat/completions"))
.and(header("authorization", "Bearer sk-test"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
"id": "chatcmpl-hdr",
"object": "chat.completion",
"created": 1,
"model": "gpt-4o",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "OK"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}
})))
.mount(&server)
.await;
let provider = GenericOpenaiProvider::new_with_name(
server.uri(),
"sk-test".into(),
"gpt-4o".into(),
"openai",
30,
);
let response = provider
.chat_blocking(MessageRequest {
model: "gpt-4o".into(),
messages: vec![Message::user_text("hi")],
..Default::default()
})
.await
.unwrap();
assert_eq!(response.text(), "OK");
}
#[tokio::test]
async fn openai_401_structured_error() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/chat/completions"))
.respond_with(ResponseTemplate::new(401).set_body_json(json!({
"error": {
"message": "Incorrect API key provided: sk-test. You can find your API key at https://example.com",
"type": "invalid_request_error",
"code": "invalid_api_key"
}
})))
.mount(&server)
.await;
let provider = GenericOpenaiProvider::new_with_name(
server.uri(),
"sk-test".into(),
"gpt-4o".into(),
"openai",
30,
);
let err = provider
.chat_blocking(MessageRequest {
model: "gpt-4o".into(),
messages: vec![Message::user_text("hi")],
..Default::default()
})
.await
.unwrap_err();
match err {
LlmError::Authentication(msg) => assert!(msg.contains("Incorrect API key")),
other => panic!("expected Authentication, got {other:?}"),
}
}
#[tokio::test]
async fn openai_429_with_retry_after() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/chat/completions"))
.respond_with(
ResponseTemplate::new(429)
.insert_header("retry-after", "30")
.set_body_json(json!({
"error": {"message": "Rate limit reached", "type": "rate_limit_error"}
})),
)
.mount(&server)
.await;
let provider = GenericOpenaiProvider::new_with_name(
server.uri(),
"sk-test".into(),
"gpt-4o".into(),
"openai",
30,
);
let err = provider
.chat_blocking(MessageRequest {
model: "gpt-4o".into(),
messages: vec![Message::user_text("hi")],
..Default::default()
})
.await
.unwrap_err();
match err {
LlmError::RateLimit { retry_after } => {
assert_eq!(retry_after, Some(Duration::from_secs(30)));
}
other => panic!("expected RateLimit, got {other:?}"),
}
}
#[tokio::test]
async fn openai_tool_use_response() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/chat/completions"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
"id": "chatcmpl-tool",
"object": "chat.completion",
"created": 1,
"model": "gpt-4o",
"choices": [{
"index": 0,
"message": {
"role": "assistant",
"content": "",
"tool_calls": [{
"id": "call_abc",
"type": "function",
"function": {
"name": "lookup",
"arguments": "{\"q\":\"rust\"}"
}
}]
},
"finish_reason": "tool_calls"
}],
"usage": {"prompt_tokens": 5, "completion_tokens": 3, "total_tokens": 8}
})))
.mount(&server)
.await;
let provider = GenericOpenaiProvider::new_with_name(
server.uri(),
"sk-test".into(),
"gpt-4o".into(),
"openai",
30,
);
let response = provider
.chat_blocking(MessageRequest {
model: "gpt-4o".into(),
messages: vec![Message::user_text("hi")],
..Default::default()
})
.await
.unwrap();
assert_eq!(response.stop_reason, StopReason::ToolUse);
let tool_use = match &response.message {
Message::Assistant { content } => content.iter().find_map(|b| match b {
ContentBlock::ToolUse { id, name, .. } => Some((id.clone(), name.clone())),
_ => None,
}),
_ => None,
};
assert_eq!(tool_use, Some(("call_abc".into(), "lookup".into())));
}
#[tokio::test]
async fn openai_500_structured_error() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/chat/completions"))
.respond_with(ResponseTemplate::new(500).set_body_json(json!({
"error": {"message": "Internal server error", "type": "server_error", "code": null}
})))
.mount(&server)
.await;
let provider = GenericOpenaiProvider::new_with_name(
server.uri(),
"sk-test".into(),
"gpt-4o".into(),
"openai",
30,
);
let err = provider
.chat_blocking(MessageRequest {
model: "gpt-4o".into(),
messages: vec![Message::user_text("hi")],
..Default::default()
})
.await
.unwrap_err();
match err {
LlmError::Request { status, body } => {
assert_eq!(status, 500);
assert!(body.contains("Internal server error"));
}
other => panic!("expected Request(500), got {other:?}"),
}
}
#[tokio::test]
async fn openai_stream_usage_only_last_chunk() {
let server = MockServer::start().await;
let sse = "data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"created\":1,\"model\":\"gpt-4o\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\",\"content\":\"Hi\"},\"finish_reason\":null}],\"usage\":null}\n\n\
data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"created\":1,\"model\":\"gpt-4o\",\"choices\":[{\"index\":0,\"delta\":{},\"finish_reason\":\"stop\"}],\"usage\":null}\n\n\
data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"created\":1,\"model\":\"gpt-4o\",\"choices\":[],\"usage\":{\"prompt_tokens\":5,\"completion_tokens\":2,\"total_tokens\":7}}\n\n\
data: [DONE]\n\n";
Mock::given(method("POST"))
.and(path("/chat/completions"))
.respond_with(
ResponseTemplate::new(200)
.insert_header("content-type", "text/event-stream")
.set_body_string(sse),
)
.mount(&server)
.await;
let provider = GenericOpenaiProvider::new_with_name(
server.uri(),
"sk-test".into(),
"gpt-4o".into(),
"openai",
30,
);
let mut stream = provider
.chat_stream_inner(MessageRequest {
model: "gpt-4o".into(),
messages: vec![Message::user_text("hi")],
stream: true,
..Default::default()
})
.await
.unwrap();
use futures_util::StreamExt;
let mut collected: Vec<StreamEvent> = Vec::new();
while let Some(ev) = stream.next().await {
collected.push(ev.unwrap());
}
let complete = collected
.iter()
.find_map(|e| match e {
StreamEvent::MessageComplete { full_response } => Some(full_response.clone()),
_ => None,
})
.expect("expected MessageComplete");
assert_eq!(complete.text(), "Hi");
assert_eq!(complete.usage.prompt_tokens, 5);
assert_eq!(complete.usage.completion_tokens, 2);
}
#[tokio::test]
async fn openai_stream_mid_stream_error() {
let server = MockServer::start().await;
// 服务端返回 200 + SSE content-type 但 body 是畸形 JSON —— 模拟流中途发送错误载荷。
// ChunkToEventStream 在 handle_chunk_json 时应产生 Error 事件而非 panic。
let malformed_sse = "data: {not-valid-json}\n\ndata: [DONE]\n\n";
Mock::given(method("POST"))
.and(path("/chat/completions"))
.respond_with(
ResponseTemplate::new(200)
.insert_header("content-type", "text/event-stream")
.set_body_string(malformed_sse),
)
.mount(&server)
.await;
let provider = GenericOpenaiProvider::new_with_name(
server.uri(),
"sk-test".into(),
"gpt-4o".into(),
"openai",
30,
);
let mut stream = provider
.chat_stream_inner(MessageRequest {
model: "gpt-4o".into(),
messages: vec![Message::user_text("hi")],
stream: true,
..Default::default()
})
.await
.unwrap();
use futures_util::StreamExt;
let mut saw_error_event = false;
let mut completed_normally = false;
while let Some(ev) = stream.next().await {
match ev {
Ok(StreamEvent::Error { .. }) => saw_error_event = true,
Ok(StreamEvent::MessageComplete { .. }) => completed_normally = true,
Err(_) => saw_error_event = true,
_ => {}
}
}
// 畸形 payload 必须被检测 —— 要么产出 Error 事件,要么最终消息完整事件标记异常。
// 不允许流静默完成(既无 Error 也无 MessageComplete),那是 bug。
assert!(
saw_error_event || completed_normally,
"malformed SSE payload neither errored nor completed normally"
);
assert!(
saw_error_event,
"expected an Error event for malformed SSE chunk"
);
}
}
+2
View File
@@ -6,6 +6,7 @@ pub mod knowledge;
pub mod retriever;
pub mod store;
pub mod types;
pub mod vector;
// 高频类型(大多数下游需要)
pub use conversation::{ConversationMemory, ConversationMemoryConfig};
@@ -13,6 +14,7 @@ pub use error::MemoryError;
pub use knowledge::KnowledgeStore;
pub use retriever::MemoryRetriever;
pub use store::{InMemoryStore, MemoryStore, SqliteStore};
pub use vector::{InMemoryVectorRetriever, VectorRetriever};
// 低频类型(配置/高级使用)
pub use conversation::MemoryStrategy;
+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();
}
}
}