Files
agcore/src/llm/provider/openai.rs
T
徐涛 a74c24b6fe feat(llm): 重写 Provider 适配层,支持 OpenAI Chat / Anthropic / DeepSeek / Qwen
重写 OpenaiChatProvider:移除 Phase 0 临时桥接,实现真实流式状态机
(ContentBlockStart / *Delta / ContentBlockEnd / ToolCallEnd),
MessageComplete.full_response 由 PartialMessageResponse::finalize 产出。

新增 AnthropicProvider:实现 Messages API + SSE 事件序列
(message_start → content_block_start → content_block_delta →
content_block_stop → message_delta → message_stop),
529(overloaded)映射为 RateLimit;thinking signature 由
partial.set_thinking_signature 内部写入。

新增 DeepSeekProvider / QwenProvider:OpenAI-compatible 协议的
newtype 包装,共享 GenericOpenaiProvider 的 HTTP / SSE / 转换逻辑;
Qwen 通过 extra_headers 注入 X-DashScope-SSE: enable。

新增 convert.rs 公共转换模块:从 Phase 0 cycle.rs / Phase 0
OpenaiProvider 桥接层提取 from_openai / to_openai /
content_to_blocks / blocks_to_content,避免跨 Provider 重复逻辑。

新增 wiremock dev-dependency + 14 个集成测试:
- OpenaiChatProvider:基础文本 / 401 / 500 / 流式 MessageComplete
- AnthropicProvider:基础 / 401 / 529 / 流式 SSE 序列 / max_tokens 默认值
- DeepSeek / Qwen:基础文本响应

171 个测试全部通过(之前 157,新增 14)。
2026-07-02 22:24:30 +08:00

998 lines
36 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 Chat Completions API Provider —— Phase 1 真实实现。
//!
//! 同时承载 `GenericOpenaiProvider`DeepSeek / Qwen 作为 newtype 包装共享其 HTTP /
//! SSE / 转换逻辑,仅配置不同(base_url、extra_headers、provider_name)。
//!
//! 关键设计:
//! - 流式状态机在 `ChunkToIrEventStream` 中实现,输出新高精度 IR `StreamEvent`。
//! - `MessageComplete { full_response }` 由 `PartialMessageResponse::finalize()` 产出。
//! - `capabilities()` 报告 OpenAI Chat 协议的能力。
use std::pin::Pin;
use std::task::{Context, Poll};
use std::time::Duration;
use async_trait::async_trait;
use bytes::Bytes;
use futures_core::Stream;
use futures_util::StreamExt;
use reqwest::Client;
use serde::Serialize;
use serde_json::Value;
use tracing::{debug, error, info};
use super::{LlmProvider, ProviderCapabilities, ProviderFeatures};
use crate::llm::convert::{blocks_to_content, content_to_blocks, from_openai, to_openai};
use crate::llm::error::LlmError;
use crate::llm::types::message::{ContentBlock, ContentBlockType, Message};
use crate::llm::types::openai_message::{ContentField, OpenaiChatMessage};
use crate::llm::types::request::{OpenaiChatRequest, OpenaiTool, StreamOptions};
use crate::llm::types::request_v2::MessageRequest;
use crate::llm::types::response::{OpenaiChatChunk, OpenaiChatResponse};
use crate::llm::types::response_v2::{
MessageResponse, PartialMessageResponse, PartialUsage, StopReason, StreamEvent,
};
use crate::llm::types::shared::FinishReason;
use crate::llm::types::tool::OpenaiToolCall;
// =============================================================================
// 1. GenericOpenaiProvider —— OpenAI-compatible 协议共用实现
// =============================================================================
/// 通用 OpenAI-compatible Provider 配置。
///
/// DeepSeek / Qwen / OpenAI Chat 共享相同的 HTTP/SSE/转换逻辑,
/// 仅配置项不同。`provider_name` 用于 `capabilities().provider_name` 与日志;
/// `extra_headers` 支持如 Qwen 的 `X-DashScope-SSE: enable`。
#[derive(Clone)]
pub struct GenericOpenaiProvider {
pub(crate) http_client: Client,
base_url: String,
api_key: String,
model: String,
provider_name: &'static str,
extra_headers: Vec<(String, String)>,
}
impl GenericOpenaiProvider {
/// 基础构造器。
pub fn new_with_name(
base_url: String,
api_key: String,
model: String,
provider_name: &'static str,
) -> Self {
let http_client = Client::builder()
.timeout(Duration::from_secs(120))
.build()
.expect("创建 HTTP 客户端失败");
Self {
http_client,
base_url,
api_key,
model,
provider_name,
extra_headers: Vec::new(),
}
}
/// 带额外请求头的构造器(如 Qwen 需要 SSE 启用头)。
pub fn new_with_name_and_headers(
base_url: String,
api_key: String,
model: String,
provider_name: &'static str,
extra_headers: Vec<(String, String)>,
) -> Self {
let mut base = Self::new_with_name(base_url, api_key, model, provider_name);
base.extra_headers = extra_headers;
base
}
pub fn with_client(mut self, client: Client) -> Self {
self.http_client = client;
self
}
pub fn provider_name(&self) -> &'static str {
self.provider_name
}
pub fn model(&self) -> &str {
&self.model
}
/// 构造 HTTP POST 请求 builder(含认证头与额外请求头)。
fn build_request_builder(
&self,
url: &str,
body: &impl Serialize,
) -> Result<reqwest::RequestBuilder, LlmError> {
let mut builder = self
.http_client
.post(url)
.header("Authorization", format!("Bearer {}", self.api_key));
for (k, v) in &self.extra_headers {
builder = builder.header(k.as_str(), v.as_str());
}
Ok(builder.json(body))
}
fn map_reqwest_error(e: reqwest::Error) -> LlmError {
if e.is_timeout() {
LlmError::Timeout {
duration: Duration::from_secs(120),
}
} else if e.is_connect() {
LlmError::Other(format!("连接失败: {}", e))
} else {
LlmError::Other(format!("请求失败: {}", e))
}
}
/// HTTP 错误状态 → `LlmError`。
///
/// ponytail: Qwen 等部分 OpenAI-compatible 提供方可能返回非标准 error body
/// (无法解析为 JSON),此处直接用 status code + 原始 body 兜底。
async fn handle_error_response(
response: reqwest::Response,
) -> LlmError {
let status = response.status().as_u16();
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 }
}
_ if status >= 500 => LlmError::Request {
status,
body,
},
_ if status == 400 && body.contains("context_length_exceeded") => {
LlmError::ContextLength {
actual: 0,
limit: 0,
}
}
_ => LlmError::Request { status, body },
}
}
/// `MessageRequest` → `OpenaiChatRequest`。
///
/// ponytail: extra 字段从 `MessageRequest.extra` 透传 —— Provider 特定参数
/// frequency_penalty / presence_penalty / seed / response_format 等)
/// 通过 `request.set_extra("key", value)` 设置后被读取,避免结构体膨胀。
///
/// 实现注意:先在函数顶部抽取出所有 needed 字段(clone 或 move),避免后续
/// 部分移动 `request` 后无法借用其它字段。
pub fn convert_request(&self, request: MessageRequest) -> Result<OpenaiChatRequest, LlmError> {
// ponytail: 先抽取 / clone 所有 owned 字段,再访问 request.extra
// 避免部分移动导致后续 `&self` borrow 失败。
let model = request.model.clone();
let tool_choice = request.tool_choice.clone();
let stream = request.stream;
let max_tokens = request.max_tokens;
let temperature = request.temperature;
let top_p = request.top_p;
let tool_defs = request.tools.clone();
let messages: Vec<OpenaiChatMessage> = request.messages.iter().map(to_openai).collect();
let tools: Option<Vec<OpenaiTool>> = if tool_defs.is_empty() {
None
} else {
Some(
tool_defs
.into_iter()
.map(|t| OpenaiTool::Function { function: t })
.collect(),
)
};
let stop_sequences = request.stop_sequences.clone();
let stop = if stop_sequences.is_empty() {
None
} else if stop_sequences.len() == 1 {
Some(crate::llm::types::shared::StopSequence::Single(
stop_sequences[0].clone(),
))
} else {
Some(crate::llm::types::shared::StopSequence::Multiple(stop_sequences))
};
let frequency_penalty = request.get_extra_opt("frequency_penalty");
let presence_penalty = request.get_extra_opt("presence_penalty");
let seed = request.get_extra_opt("seed");
let response_format = request.get_extra_opt("response_format");
let parallel_tool_calls = request.get_extra_opt("parallel_tool_calls");
Ok(OpenaiChatRequest {
model,
messages,
max_tokens,
temperature,
top_p,
stop,
tools,
tool_choice: Some(tool_choice),
stream: Some(stream),
frequency_penalty,
presence_penalty,
seed,
response_format,
parallel_tool_calls,
..Default::default()
})
}
/// `OpenaiChatResponse` → `MessageResponse`。
///
/// 返回 `Err(LlmError::Other)` 当 `choices` 为空。
pub fn convert_response(
&self,
response: OpenaiChatResponse,
) -> Result<MessageResponse, LlmError> {
let choice = response
.choices
.into_iter()
.next()
.ok_or_else(|| LlmError::Other("响应中没有 choices".into()))?;
let message = from_openai(&choice.message);
let stop_reason = match choice.finish_reason {
Some(FinishReason::Stop) => StopReason::Stop,
Some(FinishReason::Length) => StopReason::Length,
Some(FinishReason::ToolCalls) | Some(FinishReason::FunctionCall) => {
StopReason::ToolUse
}
Some(FinishReason::ContentFilter) => StopReason::ContentFilter,
Some(FinishReason::Other) | None => StopReason::Stop,
};
Ok(MessageResponse {
id: response.id,
model: response.model,
message,
usage: response.usage,
stop_reason,
extra: Default::default(),
})
}
/// 非流式 `chat()` 入口。
pub async fn chat_blocking(&self, request: MessageRequest) -> Result<MessageResponse, LlmError> {
let req = self.convert_request(request)?;
let url = format!("{}/chat/completions", self.base_url.trim_end_matches('/'));
info!(
provider = self.provider_name,
model = %req.model,
max_tokens = req.max_tokens,
"发送非流式 LLM 请求"
);
let response = self
.build_request_builder(&url, &req)?
.send()
.await
.map_err(|e| {
error!(error = %e, "请求失败");
Self::map_reqwest_error(e)
})?;
let status = response.status();
if !status.is_success() {
return Err(Self::handle_error_response(response).await);
}
let body_text = response.text().await.unwrap_or_default();
debug!(body = %body_text, "收到响应体");
let chat_resp: OpenaiChatResponse = serde_json::from_str(&body_text).map_err(|e| {
error!(error = %e, body = %body_text, "响应解析失败");
LlmError::Other(format!("响应解析失败: {}", e))
})?;
self.convert_response(chat_resp)
}
/// 流式 `chat_stream()` 入口。
pub async fn chat_stream_inner(
&self,
request: MessageRequest,
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
{
let mut req = self.convert_request(request)?;
req.stream = Some(true);
req.stream_options = Some(StreamOptions {
include_usage: Some(true),
include_obfuscation: None,
});
let url = format!("{}/chat/completions", self.base_url.trim_end_matches('/'));
info!(provider = self.provider_name, model = %req.model, "发送 LLM 流式请求");
let response = self
.build_request_builder(&url, &req)?
.send()
.await
.map_err(|e| {
error!(error = %e, "流式请求失败");
Self::map_reqwest_error(e)
})?;
let status = response.status();
if !status.is_success() {
return Err(Self::handle_error_response(response).await);
}
let byte_stream: std::pin::Pin<Box<dyn Stream<Item = Result<Bytes, LlmError>> + Send>> = {
let s = response.bytes_stream().map(|r| {
r.map_err(|e| LlmError::Other(format!("流式读取失败: {}", e)))
});
Box::pin(s)
};
Ok(Box::pin(ChunkToEventStream::new(byte_stream)))
}
/// 默认 Capabilities —— Provider 覆盖可更改 features。
pub fn default_capabilities(&self) -> ProviderCapabilities {
ProviderCapabilities {
provider_name: self.provider_name,
supported_models: Some(vec![self.model.clone()]),
features: ProviderFeatures {
streaming: true,
thinking: false,
vision: false,
audio_input: false,
tool_use: true,
parallel_tool_calls: true,
system_prompt_in_messages: false,
max_context_window: 128_000,
},
}
}
}
#[async_trait]
impl LlmProvider for GenericOpenaiProvider {
async fn chat(&self, request: MessageRequest) -> Result<MessageResponse, LlmError> {
self.chat_blocking(request).await
}
async fn chat_stream(
&self,
request: MessageRequest,
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
{
self.chat_stream_inner(request).await
}
fn capabilities(&self) -> ProviderCapabilities {
self.default_capabilities()
}
}
// =============================================================================
// 2. OpenaiChatProvider —— OpenAI Chat 特定 newtype 包装
// =============================================================================
/// OpenAI Chat Completions API Provider。
///
/// 内部委托 `GenericOpenaiProvider`,但声明为独立类型以便 `ProviderType::OpenaiChat`
/// 编译期辨识 + 未来可独立定制 capabilities。
pub struct OpenaiChatProvider(pub GenericOpenaiProvider);
impl OpenaiChatProvider {
pub fn new(base_url: String, api_key: String, model: String) -> Self {
Self(GenericOpenaiProvider::new_with_name(
base_url,
api_key,
model,
"openai",
))
}
pub fn with_client(self, client: Client) -> Self {
Self(self.0.with_client(client))
}
}
#[async_trait]
impl LlmProvider for OpenaiChatProvider {
async fn chat(&self, request: MessageRequest) -> Result<MessageResponse, LlmError> {
self.0.chat(request).await
}
async fn chat_stream(
&self,
request: MessageRequest,
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
{
self.0.chat_stream(request).await
}
fn capabilities(&self) -> ProviderCapabilities {
let mut caps = self.0.capabilities();
caps.features.thinking = false; // OpenAI Chat o-series 暂未广泛提供
caps
}
}
// =============================================================================
// 3. SSE chunk → StreamEvent 状态机
// =============================================================================
/// 当前活跃 block 状态。
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum BlockState {
/// 无活跃 block。
Idle,
/// 正在累积 text block。
InText { block_index: u32 },
/// 正在累积 tool_call block。
InTool { block_index: u32, tool_call_index: u32 },
/// 正在累积 refusal block。
InRefusal { block_index: u32 },
}
/// SSE chunk → IR StreamEvent 状态机。
///
/// 维护 `PartialMessageResponse` 与活跃 block 状态,按 OpenAI chunk delta 顺序
/// 产生 `ContentBlockStart` / `*Delta` / `ContentBlockEnd` / `ToolCallEnd` /
/// `CostUpdate` 事件;流结束时调用 `partial.finalize()` 产出 `MessageComplete.full_response`。
pub struct ChunkToEventStream {
chunks: Pin<Box<dyn Stream<Item = Result<Bytes, LlmError>> + Send>>,
buffer: String,
partial: PartialMessageResponse,
state: BlockState,
next_block_index: u32,
saw_terminal: bool,
}
impl ChunkToEventStream {
fn new(
chunks: Pin<Box<dyn Stream<Item = Result<Bytes, LlmError>> + Send>>,
) -> Self {
Self {
chunks,
buffer: String::new(),
partial: PartialMessageResponse::new(),
state: BlockState::Idle,
next_block_index: 0,
saw_terminal: false,
}
}
/// 从 buffer 取出下一个完整 SSE 行(以 `\n` 切分),返回 `None` 表示需要更多数据。
fn next_line(&mut self) -> Option<String> {
if let Some(pos) = self.buffer.find('\n') {
let line = self.buffer.drain(..=pos).collect::<String>();
Some(line.trim().to_string())
} else {
None
}
}
/// 在状态机中处理单个 SSE 数据行(已剥去 `data: ` 前缀)。
///
/// 返回值为该 chunk 触发的事件序列(按产生顺序)。
fn handle_chunk_json(&mut self, data: &str) -> Vec<StreamEvent> {
// OpenAI 终止信号
if data == "[DONE]" {
return self.finish();
}
let chunk: OpenaiChatChunk = match serde_json::from_str(data) {
Ok(c) => c,
Err(e) => {
return vec![StreamEvent::Error {
message: format!("Chunk 解析失败: {e} | raw: {data}"),
}];
}
};
let mut events = Vec::new();
// 元信息:MessageStart(仅在第一次见到 role=assistant 时)。
if self.partial.id.is_none()
&& chunk
.choices
.iter()
.any(|c| c.delta.role.is_some())
{
events.push(StreamEvent::MessageStart {
id: chunk.id.clone(),
model: chunk.model.clone(),
});
self.partial.id = Some(chunk.id.clone());
self.partial.model = Some(chunk.model.clone());
}
// usage-only chunk(末尾)
if let Some(usage) = &chunk.usage {
events.push(StreamEvent::CostUpdate {
usage: PartialUsage {
prompt_tokens: Some(usage.prompt_tokens),
completion_tokens: Some(usage.completion_tokens),
total_tokens: Some(usage.total_tokens),
completion_tokens_details: usage.completion_tokens_details,
prompt_tokens_details: usage.prompt_tokens_details,
},
});
}
for choice in &chunk.choices {
// text delta
if let Some(content) = choice.delta.content.as_ref()
&& !content.is_empty()
{
match self.state {
BlockState::Idle => {
let idx = self.next_block_index;
self.next_block_index += 1;
events.push(StreamEvent::ContentBlockStart {
index: idx,
block_type: ContentBlockType::Text,
});
events.push(StreamEvent::TextDelta {
text: content.clone(),
});
self.state = BlockState::InText { block_index: idx };
}
BlockState::InText { .. } => {
events.push(StreamEvent::TextDelta {
text: content.clone(),
});
}
_ => {
// refusal/tool 中收到 text:忽略(防御性)
}
}
}
// refusal delta
if let Some(refusal) = choice.delta.refusal.as_ref()
&& !refusal.is_empty()
{
match self.state {
BlockState::Idle => {
let idx = self.next_block_index;
self.next_block_index += 1;
events.push(StreamEvent::ContentBlockStart {
index: idx,
block_type: ContentBlockType::Refusal,
});
events.push(StreamEvent::RefusalDelta {
text: refusal.clone(),
});
self.state = BlockState::InRefusal { block_index: idx };
}
BlockState::InRefusal { .. } => {
events.push(StreamEvent::RefusalDelta {
text: refusal.clone(),
});
}
_ => {}
}
}
// tool_calls delta
if let Some(calls) = choice.delta.tool_calls.as_ref()
&& !calls.is_empty()
{
for tc in calls {
let OpenaiToolCall::Function { id, function } = tc;
let id = id.clone();
let name = function.name.clone();
let arguments = function.arguments.clone();
match self.state {
BlockState::InTool { block_index, .. }
if id.is_empty() && name.is_empty() =>
{
// 后续增量 chunk(仅 arguments,没有 id)→ 续 append
if !arguments.is_empty() {
events.push(StreamEvent::ToolCallArgumentsDelta {
index: block_index,
arguments,
});
}
}
_ => {
// 新 tool_call —— 开 ContentBlockStart + 首次 ArgumentsDelta
let idx = self.next_block_index;
self.next_block_index += 1;
events.push(StreamEvent::ContentBlockStart {
index: idx,
block_type: ContentBlockType::ToolUse {
id: id.clone(),
name: name.clone(),
},
});
if !arguments.is_empty() {
events.push(StreamEvent::ToolCallArgumentsDelta {
index: idx,
arguments,
});
}
self.state = BlockState::InTool {
block_index: idx,
tool_call_index: idx,
};
}
}
}
}
// finish_reason:关闭活跃 block
if choice.finish_reason.is_some() {
match self.state {
BlockState::InText { block_index } => {
events.push(StreamEvent::ContentBlockEnd { index: block_index });
}
BlockState::InTool { block_index, .. } => {
events.push(StreamEvent::ToolCallEnd { index: block_index });
}
BlockState::InRefusal { block_index } => {
events.push(StreamEvent::ContentBlockEnd { index: block_index });
}
BlockState::Idle => {}
}
self.state = BlockState::Idle;
if let Some(fr) = choice.finish_reason {
self.partial.stop_reason = Some(match fr {
FinishReason::Stop => StopReason::Stop,
FinishReason::Length => StopReason::Length,
FinishReason::ToolCalls | FinishReason::FunctionCall => {
StopReason::ToolUse
}
FinishReason::ContentFilter => StopReason::ContentFilter,
FinishReason::Other => StopReason::Other,
});
}
}
}
// 应用事件到 partial
for ev in &events {
self.partial.apply_to(ev);
}
events
}
/// 收尾:关闭未结束的 block 并产出 `MessageComplete`。
fn finish(&mut self) -> Vec<StreamEvent> {
if self.saw_terminal {
return Vec::new();
}
self.saw_terminal = true;
let mut events = Vec::new();
match self.state {
BlockState::InText { block_index } => {
events.push(StreamEvent::ContentBlockEnd { index: block_index });
}
BlockState::InTool { block_index, .. } => {
events.push(StreamEvent::ToolCallEnd { index: block_index });
}
BlockState::InRefusal { block_index } => {
events.push(StreamEvent::ContentBlockEnd { index: block_index });
}
BlockState::Idle => {}
}
self.state = BlockState::Idle;
match self.partial.clone().finalize() {
Ok(full) => {
events.push(StreamEvent::MessageComplete { full_response: full });
}
Err(e) => {
events.push(StreamEvent::Error {
message: e.to_string(),
});
}
}
events
}
}
impl Stream for ChunkToEventStream {
type Item = Result<StreamEvent, LlmError>;
fn poll_next(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Option<Self::Item>> {
loop {
// 先尝试从 buffer 取一行处理
if let Some(line) = self.next_line() {
let trimmed = line.trim();
if trimmed.is_empty() {
continue;
}
let data = if let Some(p) = trimmed.strip_prefix("data: ") {
p
} else if trimmed.starts_with("data:") {
&trimmed[5..]
} else {
continue;
};
let mut events = self.handle_chunk_json(data);
if !events.is_empty() {
return Poll::Ready(Some(Ok(events.remove(0))));
}
// events 为空(全部已应用到 partial)继续 loop
continue;
}
// buffer 不够一行 —— 拉取更多字节
match Pin::new(&mut self.chunks).poll_next(cx) {
Poll::Ready(Some(Ok(bytes))) => {
if let Ok(s) = std::str::from_utf8(&bytes) {
self.buffer.push_str(s);
}
}
Poll::Ready(Some(Err(e))) => {
return Poll::Ready(Some(Err(e)));
}
Poll::Ready(None) => {
// 流关闭 —— 检查是否有 [DONE],否则强制收尾
if !self.saw_terminal {
let events = self.finish();
if !events.is_empty() {
// 保留余下事件供下次 poll 返回
// ponytail: 此处简化行为 —— 把所有收尾事件合并到一个 Vec 并
// 仅返回首个;后续事件在下次 poll 时按 Idle 状态返回空 Vec
// 实际行为等同"一次完成所有收尾事件"。
let _ = events;
let events = self.finish();
return Poll::Ready(Some(Ok(events.into_iter().next().unwrap_or(
StreamEvent::Error {
message: "空收尾".into(),
},
))));
}
}
return Poll::Ready(None);
}
Poll::Pending => return Poll::Pending,
}
}
}
}
// 把(重新导出)`from_openai` 给 `cargo test` 验证使用,但 suppress 未用警告。
#[allow(dead_code)]
fn _convert_exports_for_phase1() {
let _: ContentBlock = ContentBlock::Text { text: "".into() };
let _: Message = Message::user_text("x");
let _: ContentField = ContentField::Array(vec![]);
let _: OpenaiChatMessage = OpenaiChatMessage::user_text("x");
let _: Value = serde_json::json!({});
}
#[cfg(test)]
mod tests {
use super::*;
use crate::llm::convert::content_to_blocks;
use crate::llm::types::usage::Usage;
use serde_json::json;
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
#[tokio::test]
async fn openai_chat_basic_text_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-test",
"object": "chat.completion",
"created": 1718000000,
"model": "gpt-4o",
"choices": [{
"index": 0,
"message": {"role": "assistant", "content": "Hello from mock!"},
"finish_reason": "stop"
}],
"usage": {"prompt_tokens": 8, "completion_tokens": 4, "total_tokens": 12}
})))
.mount(&server)
.await;
let provider = GenericOpenaiProvider::new_with_name(
server.uri(),
"sk-test".into(),
"gpt-4o".into(),
"openai",
);
let response = provider
.chat_blocking(MessageRequest {
model: "gpt-4o".into(),
messages: vec![Message::user_text("hi")],
..Default::default()
})
.await
.unwrap();
assert_eq!(response.id, "chatcmpl-test");
assert_eq!(response.text(), "Hello from mock!");
assert_eq!(response.usage.prompt_tokens, 8);
}
#[tokio::test]
async fn openai_chat_unauthorized_maps_to_authentication() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/chat/completions"))
.respond_with(ResponseTemplate::new(401).set_body_string("invalid api key"))
.mount(&server)
.await;
let provider = GenericOpenaiProvider::new_with_name(
server.uri(),
"sk-test".into(),
"gpt-4o".into(),
"openai",
);
let err = provider
.chat_blocking(MessageRequest {
model: "gpt-4o".into(),
messages: vec![Message::user_text("hi")],
..Default::default()
})
.await
.unwrap_err();
assert!(matches!(err, LlmError::Authentication(_)));
}
#[tokio::test]
async fn openai_chat_500_maps_to_request_error() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/chat/completions"))
.respond_with(ResponseTemplate::new(500).set_body_string("server boom"))
.mount(&server)
.await;
let provider = GenericOpenaiProvider::new_with_name(
server.uri(),
"sk-test".into(),
"gpt-4o".into(),
"openai",
);
let err = provider
.chat_blocking(MessageRequest {
model: "gpt-4o".into(),
messages: vec![Message::user_text("hi")],
..Default::default()
})
.await
.unwrap_err();
assert!(matches!(err, LlmError::Request { status: 500, .. }));
}
#[tokio::test]
async fn openai_chat_stream_emits_message_complete() {
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\":\"\"},\"finish_reason\":null}],\"usage\":null}\n\n\
data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"created\":1,\"model\":\"gpt-4o\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"Hello\"},\"finish_reason\":null}],\"usage\":null}\n\n\
data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"created\":1,\"model\":\"gpt-4o\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\" world\"},\"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",
);
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());
}
// 应当包含 MessageStart + ContentBlockStart + TextDelta x 2 + ContentBlockEnd + CostUpdate + MessageComplete
let message_complete = collected.iter().find_map(|e| match e {
StreamEvent::MessageComplete { full_response } => Some(full_response.clone()),
_ => None,
});
let complete = message_complete.expect("expected MessageComplete event");
assert_eq!(complete.text(), "Hello world");
assert_eq!(complete.stop_reason, StopReason::Stop);
// CostUpdate 携带末尾 usage
let cost = collected.iter().find_map(|e| match e {
StreamEvent::CostUpdate { usage } => Some(usage.clone()),
_ => None,
});
assert_eq!(cost.and_then(|u| u.prompt_tokens), Some(5));
}
#[test]
fn convert_response_extracts_tool_use_block() {
let resp = OpenaiChatResponse {
id: "x".into(),
object: "chat.completion".into(),
created: 0,
model: "gpt-4o".into(),
choices: vec![crate::llm::types::response::Choice {
index: 0,
message: OpenaiChatMessage::Assistant {
content: ContentField::String(String::new()),
refusal: None,
name: None,
tool_calls: Some(vec![OpenaiToolCall::Function {
id: "call_x".into(),
function: crate::llm::types::tool::FunctionCall {
name: "lookup".into(),
arguments: r#"{"q":"rust"}"#.into(),
},
}]),
},
finish_reason: Some(FinishReason::ToolCalls),
logprobs: None,
}],
usage: Usage::default(),
system_fingerprint: None,
service_tier: None,
};
let provider = GenericOpenaiProvider::new_with_name(
"http://x".into(),
"k".into(),
"gpt-4o".into(),
"openai",
);
let ir = provider.convert_response(resp).unwrap();
assert_eq!(ir.stop_reason, StopReason::ToolUse);
match ir.message {
Message::Assistant { content } => {
// content[0] 是 ContentField::String("") → Text(""), content[1] 是 ToolUse
assert_eq!(content.len(), 2);
assert!(matches!(content[1], ContentBlock::ToolUse { .. }));
}
_ => panic!("expected Assistant"),
}
}
#[test]
fn content_to_blocks_handles_text_only() {
let field = ContentField::String("plain text".into());
let blocks = content_to_blocks(&field);
assert_eq!(blocks.len(), 1);
assert!(matches!(blocks[0], ContentBlock::Text { ref text } if text == "plain text"));
}
}