feat(core): 完成 Phase 11 测试与检索补强
- 新增 VectorRetriever trait 与 InMemoryVectorRetriever 引用实现 - 补充 Provider roundtrip wiremock 测试与 MemoryStore 并发测试共 23 个 - 修复 openai 429 retry-after header 解析(与 anthropic 对齐)
This commit is contained in:
@@ -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
@@ -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"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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