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
+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"
);
}
}