Files
agcore/src/llm/provider/openai.rs
T
徐涛 760de46623 refactor(types): request.rs 类型移入 provider/openai.rs
- 删除 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)
2026-07-08 22:55:17 +08:00

1508 lines
54 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::{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"
);
}
}