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