focus/crates/focus-providers/tests/openai_tests.rs

194 lines
9.3 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

//! OpenAI providerResponses + 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" }));
}