//! 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, #[serde(skip_serializing_if = "Option::is_none")] pub include_obfuscation: Option, } /// 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, #[serde(skip_serializing_if = "Option::is_none")] pub timezone: Option, } /// 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, } /// OpenAI Chat Completions 请求体。 #[derive(Debug, Clone, Default, Serialize, Deserialize)] #[serde(rename_all = "snake_case")] pub(crate) struct OpenaiChatRequest { pub model: String, pub messages: Vec, #[serde(skip_serializing_if = "Option::is_none")] pub frequency_penalty: Option, #[serde(skip_serializing_if = "Option::is_none")] pub logit_bias: Option, #[serde(skip_serializing_if = "Option::is_none")] pub max_tokens: Option, #[serde(skip_serializing_if = "Option::is_none")] pub n: Option, #[serde(skip_serializing_if = "Option::is_none")] pub presence_penalty: Option, #[serde(skip_serializing_if = "Option::is_none")] pub response_format: Option, #[serde(skip_serializing_if = "Option::is_none")] pub seed: Option, #[serde(skip_serializing_if = "Option::is_none")] pub service_tier: Option, #[serde(skip_serializing_if = "Option::is_none")] pub stop: Option, #[serde(skip_serializing_if = "Option::is_none")] pub stream: Option, #[serde(skip_serializing_if = "Option::is_none")] pub stream_options: Option, #[serde(skip_serializing_if = "Option::is_none")] pub temperature: Option, #[serde(skip_serializing_if = "Option::is_none")] pub top_p: Option, #[serde(skip_serializing_if = "Option::is_none")] pub tools: Option>, #[serde(skip_serializing_if = "Option::is_none")] pub tool_choice: Option, #[serde(skip_serializing_if = "Option::is_none")] pub parallel_tool_calls: Option, #[serde(skip_serializing_if = "Option::is_none")] pub user: Option, #[serde(skip_serializing_if = "Option::is_none")] pub extra_headers: Option, #[serde(skip_serializing_if = "Option::is_none")] pub extra_body: Option, } // ============================================================================= // 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 { 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::().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 { // 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 = request.messages.iter().map(to_openai).collect(); let tools: Option> = 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 { 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 { 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> + 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> + 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 { self.chat_blocking(request).await } async fn chat_stream( &self, request: MessageRequest, ) -> Result> + 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 { self.0.chat(request).await } async fn chat_stream( &self, request: MessageRequest, ) -> Result> + 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> + Send>>, buffer: String, partial: PartialMessageResponse, state: BlockState, next_block_index: u32, saw_terminal: bool, } impl ChunkToEventStream { fn new(chunks: Pin> + 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 { if let Some(pos) = self.buffer.find('\n') { let line = self.buffer.drain(..=pos).collect::(); Some(line.trim().to_string()) } else { None } } /// 在状态机中处理单个 SSE 数据行(已剥去 `data: ` 前缀)。 /// /// 返回值为该 chunk 触发的事件序列(按产生顺序)。 fn handle_chunk_json(&mut self, data: &str) -> Vec { // 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 { 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; fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { 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 = 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 = 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" ); } }