760de46623
- 删除 types/request.rs(187 行) - 所有 OpenAI wire-format 类型迁入 provider/openai.rs,可见性 pub(crate): StreamOptions / OpenaiTool / AudioParam / PredictionContent / UserLocation / Approximate / WebSearchOptions / OpenaiChatRequest - types/mod.rs 删除 pub mod request; 与对应 re-export - convert_request 同步降级为 pub(crate) 以匹配 OpenaiChatRequest 可见性 - 公共 re-export 路径 agcore::llm::types::OpenaiChatRequest 等已删除(Breaking Change,见 CHANGELOG)
1508 lines
54 KiB
Rust
1508 lines
54 KiB
Rust
//! 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::{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_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, ResponseFormat, ServiceTier, StopSequence};
|
||
use crate::llm::types::tool::{OpenaiToolCall, OpenaiToolDefinition, ToolChoice};
|
||
use serde::Deserialize;
|
||
|
||
// =============================================================================
|
||
// 0. OpenAI wire-format 类型(Phase 13 从 types::request 迁入)
|
||
// =============================================================================
|
||
|
||
/// 流式响应选项。
|
||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||
pub(crate) struct StreamOptions {
|
||
#[serde(skip_serializing_if = "Option::is_none")]
|
||
pub include_usage: Option<bool>,
|
||
#[serde(skip_serializing_if = "Option::is_none")]
|
||
pub include_obfuscation: Option<bool>,
|
||
}
|
||
|
||
/// OpenAI wire-format 工具定义。
|
||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||
#[serde(rename_all = "snake_case", tag = "type")]
|
||
pub(crate) enum OpenaiTool {
|
||
Function { function: OpenaiToolDefinition },
|
||
}
|
||
|
||
/// 音频输出参数。
|
||
#[allow(dead_code)]
|
||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||
pub(crate) struct AudioParam {
|
||
pub format: String,
|
||
pub voice: String,
|
||
}
|
||
|
||
/// 预测内容(OpenAI `prediction` 字段)。
|
||
#[allow(dead_code)]
|
||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||
pub(crate) struct PredictionContent {
|
||
#[serde(rename = "type")]
|
||
pub pred_type: String,
|
||
pub content: String,
|
||
}
|
||
|
||
/// 用户位置(web search 用)。
|
||
#[allow(dead_code)]
|
||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||
pub(crate) struct UserLocation {
|
||
#[serde(rename = "type")]
|
||
pub loc_type: String,
|
||
pub approximate: Approximate,
|
||
}
|
||
|
||
/// 近似位置。
|
||
#[allow(dead_code)]
|
||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||
pub(crate) struct Approximate {
|
||
pub city: String,
|
||
pub country: String,
|
||
#[serde(skip_serializing_if = "Option::is_none")]
|
||
pub region: Option<String>,
|
||
#[serde(skip_serializing_if = "Option::is_none")]
|
||
pub timezone: Option<String>,
|
||
}
|
||
|
||
/// Web search 选项。
|
||
#[allow(dead_code)]
|
||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||
pub(crate) struct WebSearchOptions {
|
||
pub search_context_size: String,
|
||
#[serde(skip_serializing_if = "Option::is_none")]
|
||
pub user_location: Option<UserLocation>,
|
||
}
|
||
|
||
/// OpenAI Chat Completions 请求体。
|
||
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
|
||
#[serde(rename_all = "snake_case")]
|
||
pub(crate) struct OpenaiChatRequest {
|
||
pub model: String,
|
||
pub messages: Vec<OpenaiChatMessage>,
|
||
#[serde(skip_serializing_if = "Option::is_none")]
|
||
pub frequency_penalty: Option<f32>,
|
||
#[serde(skip_serializing_if = "Option::is_none")]
|
||
pub logit_bias: Option<Value>,
|
||
#[serde(skip_serializing_if = "Option::is_none")]
|
||
pub max_tokens: Option<u32>,
|
||
#[serde(skip_serializing_if = "Option::is_none")]
|
||
pub n: Option<u32>,
|
||
#[serde(skip_serializing_if = "Option::is_none")]
|
||
pub presence_penalty: Option<f32>,
|
||
#[serde(skip_serializing_if = "Option::is_none")]
|
||
pub response_format: Option<ResponseFormat>,
|
||
#[serde(skip_serializing_if = "Option::is_none")]
|
||
pub seed: Option<i64>,
|
||
#[serde(skip_serializing_if = "Option::is_none")]
|
||
pub service_tier: Option<ServiceTier>,
|
||
#[serde(skip_serializing_if = "Option::is_none")]
|
||
pub stop: Option<StopSequence>,
|
||
#[serde(skip_serializing_if = "Option::is_none")]
|
||
pub stream: Option<bool>,
|
||
#[serde(skip_serializing_if = "Option::is_none")]
|
||
pub stream_options: Option<StreamOptions>,
|
||
#[serde(skip_serializing_if = "Option::is_none")]
|
||
pub temperature: Option<f32>,
|
||
#[serde(skip_serializing_if = "Option::is_none")]
|
||
pub top_p: Option<f32>,
|
||
#[serde(skip_serializing_if = "Option::is_none")]
|
||
pub tools: Option<Vec<OpenaiTool>>,
|
||
#[serde(skip_serializing_if = "Option::is_none")]
|
||
pub tool_choice: Option<ToolChoice>,
|
||
#[serde(skip_serializing_if = "Option::is_none")]
|
||
pub parallel_tool_calls: Option<bool>,
|
||
#[serde(skip_serializing_if = "Option::is_none")]
|
||
pub user: Option<String>,
|
||
#[serde(skip_serializing_if = "Option::is_none")]
|
||
pub extra_headers: Option<Value>,
|
||
#[serde(skip_serializing_if = "Option::is_none")]
|
||
pub extra_body: Option<Value>,
|
||
}
|
||
|
||
// =============================================================================
|
||
// 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)>,
|
||
/// HTTP 请求超时秒数。由 `ProviderConfig::timeout_secs` 传入,
|
||
/// 在 `LlmError::Timeout { duration }` 中回显。`reqwest::Client` 不暴露 timeout getter,
|
||
/// 因此单独存储以便错误消息与配置保持一致。
|
||
timeout_secs: u64,
|
||
}
|
||
|
||
impl GenericOpenaiProvider {
|
||
/// 一次性构造 —— `create_provider` 路径专用,避免 `new_with_name` + `with_client` 的双重 client 构造。
|
||
///
|
||
/// 调用方负责预先构造好带正确 timeout 的 `http_client`。`extra_headers` 与 `timeout_secs`
|
||
/// 一并设置字段,避免后续修改。
|
||
pub(crate) fn from_parts(
|
||
base_url: String,
|
||
api_key: String,
|
||
model: String,
|
||
provider_name: &'static str,
|
||
http_client: Client,
|
||
extra_headers: Vec<(String, String)>,
|
||
timeout_secs: u64,
|
||
) -> Self {
|
||
Self {
|
||
http_client,
|
||
base_url,
|
||
api_key,
|
||
model,
|
||
provider_name,
|
||
extra_headers,
|
||
timeout_secs,
|
||
}
|
||
}
|
||
|
||
/// 基础构造器。
|
||
///
|
||
/// `timeout_secs` 应用于 `reqwest::Client` 的请求超时配置。
|
||
/// 应由 `ProviderConfig::timeout_secs` 传入(调用方如不知道,可传 30)。
|
||
pub fn new_with_name(
|
||
base_url: String,
|
||
api_key: String,
|
||
model: String,
|
||
provider_name: &'static str,
|
||
timeout_secs: u64,
|
||
) -> Self {
|
||
let http_client = Client::builder()
|
||
.timeout(Duration::from_secs(timeout_secs))
|
||
.build()
|
||
.expect("创建 HTTP 客户端失败");
|
||
Self::from_parts(
|
||
base_url,
|
||
api_key,
|
||
model,
|
||
provider_name,
|
||
http_client,
|
||
Vec::new(),
|
||
timeout_secs,
|
||
)
|
||
}
|
||
|
||
/// 带额外请求头的构造器(如 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)>,
|
||
timeout_secs: u64,
|
||
) -> Self {
|
||
let http_client = Client::builder()
|
||
.timeout(Duration::from_secs(timeout_secs))
|
||
.build()
|
||
.expect("创建 HTTP 客户端失败");
|
||
Self::from_parts(
|
||
base_url,
|
||
api_key,
|
||
model,
|
||
provider_name,
|
||
http_client,
|
||
extra_headers,
|
||
timeout_secs,
|
||
)
|
||
}
|
||
|
||
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(&self, e: reqwest::Error) -> LlmError {
|
||
if e.is_timeout() {
|
||
LlmError::Timeout {
|
||
duration: Duration::from_secs(self.timeout_secs),
|
||
}
|
||
} 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 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 => LlmError::RateLimit { retry_after },
|
||
_ 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(crate) 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.into() })
|
||
.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, timeout_secs: u64) -> Self {
|
||
Self(GenericOpenaiProvider::new_with_name(
|
||
base_url,
|
||
api_key,
|
||
model,
|
||
"openai",
|
||
timeout_secs,
|
||
))
|
||
}
|
||
|
||
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 let Some(p) = trimmed.strip_prefix("data:") {
|
||
p
|
||
} 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 std::time::Duration;
|
||
use wiremock::matchers::{body_partial_json, header, 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",
|
||
30,
|
||
);
|
||
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",
|
||
30,
|
||
);
|
||
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",
|
||
30,
|
||
);
|
||
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",
|
||
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());
|
||
}
|
||
|
||
// 应当包含 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",
|
||
30,
|
||
);
|
||
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"));
|
||
}
|
||
|
||
// ===== 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"
|
||
);
|
||
}
|
||
}
|