194 lines
9.3 KiB
Rust
194 lines
9.3 KiB
Rust
//! OpenAI provider(Responses + Chat Completions)的录制/回放集成测试。
|
||
//! Recorded/replayed integration tests for the OpenAI provider (both
|
||
//! Responses and Chat Completions).
|
||
|
||
mod common;
|
||
|
||
use focus_core::model::*;
|
||
use focus_core::provider::{ProviderRequest, StreamEvent};
|
||
use focus_providers::config::ProviderConfig;
|
||
use focus_providers::openai::{OpenAiProtocol, OpenAiProvider};
|
||
use std::sync::Arc;
|
||
|
||
fn provider(mock: common::MockTransport, protocol: OpenAiProtocol) -> OpenAiProvider {
|
||
OpenAiProvider::with_transport(ProviderConfig::new("sk-test"), protocol, Arc::new(mock))
|
||
}
|
||
|
||
fn request() -> ProviderRequest {
|
||
ProviderRequest {
|
||
model: "gpt-4o".into(),
|
||
system_prompt: "Be concise.".into(),
|
||
messages: vec![Message::user_text("hi")],
|
||
tools: focus_json::JsonValue::arr(),
|
||
max_tokens: Some(512),
|
||
temperature: None,
|
||
}
|
||
}
|
||
|
||
fn final_message(events: &[StreamEvent]) -> &AssistantMessage {
|
||
events
|
||
.iter()
|
||
.find_map(|e| match e {
|
||
StreamEvent::Done { message } => Some(message),
|
||
_ => None,
|
||
})
|
||
.expect("done event")
|
||
}
|
||
|
||
// ---- Chat Completions ----
|
||
|
||
#[test]
|
||
fn chat_streams_text_and_usage() {
|
||
let mut mock = common::MockTransport::new();
|
||
mock.push_body(
|
||
"data: {\"id\":\"c1\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\",\"content\":\"\"},\"finish_reason\":null}]}\n\n\
|
||
data: {\"id\":\"c1\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"Hello\"},\"finish_reason\":null}]}\n\n\
|
||
data: {\"id\":\"c1\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\" world\"},\"finish_reason\":null}]}\n\n\
|
||
data: {\"id\":\"c1\",\"choices\":[{\"index\":0,\"delta\":{},\"finish_reason\":\"stop\"}]}\n\n\
|
||
data: {\"id\":\"c1\",\"choices\":[],\"usage\":{\"prompt_tokens\":9,\"completion_tokens\":5}}\n\n\
|
||
data: [DONE]\n\n",
|
||
);
|
||
let provider = provider(mock, OpenAiProtocol::ChatCompletions);
|
||
let events = common::collect(&provider, &request());
|
||
|
||
// 事件序列:start → text_start → 两个 text_delta → text_end → done。
|
||
// Sequence: start → text_start → two text_deltas → text_end → done.
|
||
assert!(matches!(events[0], StreamEvent::Start { .. }));
|
||
assert!(matches!(events[1], StreamEvent::TextStart { .. }));
|
||
assert!(matches!(events[2], StreamEvent::TextDelta { ref delta, .. } if delta == "Hello"));
|
||
assert!(matches!(events[3], StreamEvent::TextDelta { ref delta, .. } if delta == " world"));
|
||
assert!(matches!(events[4], StreamEvent::TextEnd { .. }));
|
||
|
||
let msg = final_message(&events);
|
||
let text = msg.content[0].as_text().expect("text block");
|
||
assert_eq!(text.text, "Hello world");
|
||
assert_eq!(msg.stop_reason, StopReason::Stop);
|
||
assert_eq!(msg.usage.input_tokens, 9);
|
||
assert_eq!(msg.usage.output_tokens, 5);
|
||
}
|
||
|
||
#[test]
|
||
fn chat_streams_interleaved_tool_calls() {
|
||
let mut mock = common::MockTransport::new();
|
||
mock.push_body(
|
||
// 第一个 chunk 同时打开两个调用;随后参数交错到达。
|
||
// First chunk opens both calls; arguments then arrive interleaved.
|
||
"data: {\"choices\":[{\"index\":0,\"delta\":{\"tool_calls\":[{\"index\":0,\"id\":\"call_a\",\"type\":\"function\",\"function\":{\"name\":\"read\",\"arguments\":\"\"}},{\"index\":1,\"id\":\"call_b\",\"type\":\"function\",\"function\":{\"name\":\"write\",\"arguments\":\"\"}}]},\"finish_reason\":null}]}\n\n\
|
||
data: {\"choices\":[{\"index\":0,\"delta\":{\"tool_calls\":[{\"index\":0,\"function\":{\"arguments\":\"{\\\"path\\\":\\\"a\\\"}\"}}]},\"finish_reason\":null}]}\n\n\
|
||
data: {\"choices\":[{\"index\":0,\"delta\":{\"tool_calls\":[{\"index\":1,\"function\":{\"arguments\":\"{\\\"path\\\":\\\"b\\\"}\"}}]},\"finish_reason\":null}]}\n\n\
|
||
data: {\"choices\":[{\"index\":0,\"delta\":{},\"finish_reason\":\"tool_calls\"}]}\n\n\
|
||
data: [DONE]\n\n",
|
||
);
|
||
let provider = provider(mock, OpenAiProtocol::ChatCompletions);
|
||
let events = common::collect(&provider, &request());
|
||
|
||
let msg = final_message(&events);
|
||
assert_eq!(msg.stop_reason, StopReason::ToolUse);
|
||
assert_eq!(msg.content.len(), 2, "two tool calls");
|
||
let tc0 = msg.content[0].as_tool_call().expect("call 0");
|
||
let tc1 = msg.content[1].as_tool_call().expect("call 1");
|
||
assert_eq!(tc0.name, "read");
|
||
assert_eq!(tc0.arguments.get_str("path"), Some("a"));
|
||
assert_eq!(tc1.name, "write");
|
||
assert_eq!(tc1.arguments.get_str("path"), Some("b"));
|
||
// 交错参数必须归位(回归:见 focus-core reducer 修复)。
|
||
// Interleaved arguments must land in the right slots (regression: see the
|
||
// focus-core reducer fix).
|
||
assert!(tc0.arguments.get_str("path") == Some("a"));
|
||
}
|
||
|
||
#[test]
|
||
fn chat_encodes_error_chunk() {
|
||
let mut mock = common::MockTransport::new();
|
||
mock.push_body(
|
||
"data: {\"error\":{\"message\":\"invalid api key\",\"type\":\"authentication_error\"}}\n\n",
|
||
);
|
||
let provider = provider(mock, OpenAiProtocol::ChatCompletions);
|
||
let events = common::collect(&provider, &request());
|
||
|
||
match events.last().unwrap() {
|
||
StreamEvent::Error { error } => {
|
||
let msg = error.error_message.clone().unwrap_or_default();
|
||
assert!(msg.contains("invalid api key"), "got: {}", msg);
|
||
}
|
||
other => panic!("expected error event, got {:?}", other),
|
||
}
|
||
}
|
||
|
||
// ---- Responses API ----
|
||
|
||
#[test]
|
||
fn responses_streams_text() {
|
||
let mut mock = common::MockTransport::new();
|
||
mock.push_body(
|
||
"event: response.created\ndata: {\"type\":\"response.created\",\"response\":{\"id\":\"r1\",\"status\":\"in_progress\"}}\n\n\
|
||
event: response.output_item.added\ndata: {\"type\":\"response.output_item.added\",\"output_index\":0,\"item\":{\"id\":\"it1\",\"type\":\"message\",\"role\":\"assistant\",\"content\":[{\"type\":\"output_text\",\"text\":\"\",\"annotations\":[]}]}}\n\n\
|
||
event: response.output_text.delta\ndata: {\"type\":\"response.output_text.delta\",\"item_id\":\"it1\",\"output_index\":0,\"delta\":\"Hi\"}\n\n\
|
||
event: response.output_text.delta\ndata: {\"type\":\"response.output_text.delta\",\"item_id\":\"it1\",\"output_index\":0,\"delta\":\" there\"}\n\n\
|
||
event: response.completed\ndata: {\"type\":\"response.completed\",\"response\":{\"id\":\"r1\",\"status\":\"completed\",\"usage\":{\"input_tokens\":12,\"output_tokens\":4,\"total_tokens\":16,\"input_tokens_details\":{\"cached_tokens\":3}}}}\n\n",
|
||
);
|
||
let provider = provider(mock, OpenAiProtocol::Responses);
|
||
let events = common::collect(&provider, &request());
|
||
|
||
let msg = final_message(&events);
|
||
let text = msg.content[0].as_text().expect("text block");
|
||
assert_eq!(text.text, "Hi there");
|
||
assert_eq!(msg.stop_reason, StopReason::Stop);
|
||
assert_eq!(msg.usage.input_tokens, 12);
|
||
assert_eq!(msg.usage.output_tokens, 4);
|
||
assert_eq!(msg.usage.cache_read_tokens, 3);
|
||
}
|
||
|
||
#[test]
|
||
fn responses_streams_function_call() {
|
||
let mut mock = common::MockTransport::new();
|
||
mock.push_body(
|
||
"event: response.output_item.added\ndata: {\"type\":\"response.output_item.added\",\"output_index\":0,\"item\":{\"id\":\"fc1\",\"type\":\"function_call\",\"call_id\":\"call_x\",\"name\":\"read\",\"arguments\":\"\",\"status\":\"in_progress\"}}\n\n\
|
||
event: response.function_call_arguments.delta\ndata: {\"type\":\"response.function_call_arguments.delta\",\"item_id\":\"fc1\",\"output_index\":0,\"delta\":\"{\\\"path\\\":\\\"x\"}\n\n\
|
||
event: response.function_call_arguments.delta\ndata: {\"type\":\"response.function_call_arguments.delta\",\"item_id\":\"fc1\",\"output_index\":0,\"delta\":\"y.txt\\\"}\"}\n\n\
|
||
event: response.completed\ndata: {\"type\":\"response.completed\",\"response\":{\"status\":\"completed\",\"usage\":{\"input_tokens\":5,\"output_tokens\":6}}}\n\n",
|
||
);
|
||
let provider = provider(mock, OpenAiProtocol::Responses);
|
||
let events = common::collect(&provider, &request());
|
||
|
||
let msg = final_message(&events);
|
||
assert_eq!(msg.content.len(), 1);
|
||
let tc = msg.content[0].as_tool_call().expect("tool call");
|
||
assert_eq!(tc.name, "read");
|
||
assert_eq!(tc.arguments.get_str("path"), Some("xy.txt"));
|
||
}
|
||
|
||
#[test]
|
||
fn responses_encodes_error() {
|
||
let mut mock = common::MockTransport::new();
|
||
mock.push_body(
|
||
"event: error\ndata: {\"type\":\"error\",\"code\":\"invalid_request_error\",\"message\":\"bad model\"}\n\n",
|
||
);
|
||
let provider = provider(mock, OpenAiProtocol::Responses);
|
||
let events = common::collect(&provider, &request());
|
||
|
||
match events.last().unwrap() {
|
||
StreamEvent::Error { error } => {
|
||
let msg = error.error_message.clone().unwrap_or_default();
|
||
assert!(msg.contains("bad model"), "got: {}", msg);
|
||
}
|
||
other => panic!("expected error event, got {:?}", other),
|
||
}
|
||
}
|
||
|
||
#[test]
|
||
fn request_headers_and_path() {
|
||
let mut mock = common::MockTransport::new();
|
||
mock.push_body("data: [DONE]\n\n");
|
||
let provider = provider(mock.clone(), OpenAiProtocol::ChatCompletions);
|
||
let _ = common::collect(&provider, &request());
|
||
|
||
let req = mock.last_request();
|
||
assert_eq!(req.host, "api.openai.com");
|
||
assert_eq!(req.path, "/v1/chat/completions");
|
||
let headers: Vec<(String, String)> = req.headers.clone();
|
||
assert!(headers
|
||
.iter()
|
||
.any(|(k, v)| { k.eq_ignore_ascii_case("authorization") && v == "Bearer sk-test" }));
|
||
}
|