From b895616dd097349544e9174cd9f40907e6c9dcdc Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=BE=90=E6=B6=9B?= Date: Mon, 20 Jul 2026 09:05:04 +0800 Subject: [PATCH] =?UTF-8?q?feat(llm):=20=E5=AE=9E=E7=8E=B0=20OpenAI=20Resp?= =?UTF-8?q?onse=20API=20Provider?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 新增独立 OpenaiResponseProvider(POST /responses 协议),独立 feature provider-openai-response - 覆盖文本对话/流式/Vision/Function Calling/多轮接续/结构化输出/内置工具逃生舱 - 内置工具(web_search/file_search)通过 extra 逃生舱透传 - 工厂注册 ProviderType::OpenaiResponse + src/llm 模块门控追加 - 新增 example response_api_demo + CI 矩阵新增组合 + README/roadmap 同步 - 测试覆盖:13 单元 + 15 wiremock(流式 + 非流式 + 错误路径) - 文档:docs/28-phase28-openai-response-api-provider.md --- .github/workflows/ci.yml | 1 + Cargo.toml | 12 +- README.md | 4 +- ...28-phase28-openai-response-api-provider.md | 790 +++++++ docs/roadmap.md | 3 +- examples/response_api_demo.rs | 82 + src/llm.rs | 3 +- src/llm/provider.rs | 18 +- src/llm/provider/openai_response.rs | 1978 +++++++++++++++++ 9 files changed, 2884 insertions(+), 7 deletions(-) create mode 100644 docs/28-phase28-openai-response-api-provider.md create mode 100644 examples/response_api_demo.rs create mode 100644 src/llm/provider/openai_response.rs diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 8d0e3b5..0013274 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -20,6 +20,7 @@ jobs: - "chat,provider-openai,tools-mcp" - "multi,provider-openai" - "multi,provider-openai,tools-mcp" + - "chat,provider-openai,provider-openai-response" steps: - uses: actions/checkout@v4 - uses: actions-rust-lang/setup-rust-toolchain@v1 diff --git a/Cargo.toml b/Cargo.toml index d71dfd6..ddbd307 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -20,9 +20,11 @@ agent = ["llm", "tools", "memory", "futures-util"] engine = ["agent"] # === Provider features === -# Provider features — openai/anthropic 额外依赖 bytes(流式解析)和 futures-util(Stream 组合) +# Provider features — openai/anthropic/openai-response 额外依赖 bytes(流式解析)和 futures-util(Stream 组合) provider-openai = ["llm", "reqwest", "bytes", "futures-util"] provider-anthropic = ["llm", "reqwest", "bytes", "futures-util"] +# OpenAI Response API(POST /responses)—— 与 Chat Completions 协议独立,独立 feature +provider-openai-response = ["llm", "reqwest", "bytes", "futures-util"] # deepseek/qwen 使用 openai_compat 适配层,不需要 bytes 和 futures-util provider-deepseek = ["llm", "reqwest"] provider-qwen = ["llm", "reqwest"] @@ -37,8 +39,8 @@ full = [ "tools", "tools-mcp", "memory", "memory-sqlite", "agent", "engine", - "provider-openai", "provider-anthropic", "provider-deepseek", - "provider-qwen", "provider-ollama", + "provider-openai", "provider-anthropic", "provider-openai-response", + "provider-deepseek", "provider-qwen", "provider-ollama", "tracing-init", ] light = ["llm", "provider-openai", "tools", "tools-mcp", "memory", "agent", "engine", "prompt", "document"] @@ -148,3 +150,7 @@ required-features = ["memory", "tracing-init"] [[example]] name = "end_to_end" required-features = ["agent", "memory-sqlite", "provider-openai"] + +[[example]] +name = "response_api_demo" +required-features = ["llm", "provider-openai-response"] diff --git a/README.md b/README.md index 3da7dfa..2828399 100644 --- a/README.md +++ b/README.md @@ -110,7 +110,7 @@ let provider = create_provider( ).expect("创建 Provider 失败"); ``` -更多端到端示例见 [`examples/`](./examples/) 目录(共 18 个,全部可 `cargo run --example `): +更多端到端示例见 [`examples/`](./examples/) 目录(全部可 `cargo run --example `): | 示例 | 说明 | |------|------| @@ -132,6 +132,7 @@ let provider = create_provider( | `engine_demo` | Agent 执行引擎:SessionManager 会话树 + Checkpointer 快照恢复 | | `bridge_keys_demo` | 桥接键:Agent 间上下文键值透传 | | `agent_switch_demo` | Agent 热切换:会话中动态切换 Agent 角色 | +| `response_api_demo` | OpenAI Response API(`POST /responses`)真实调用 | ## Feature 组合 @@ -184,6 +185,7 @@ agcore = { version = "0.3", default-features = false, features = ["multi", "prov | `engine` | SessionManager + Checkpointer + SubAgent + Switch | `agent` | | `provider-openai` | OpenAI Provider 实现 | `llm` | | `provider-anthropic` | Anthropic Provider 实现 | `llm` | +| `provider-openai-response` | OpenAI Response API(`POST /responses`)Provider 实现 | `llm` | | `provider-deepseek` | DeepSeek Provider 实现 | `llm` | | `provider-qwen` | Qwen Provider 实现 | `llm` | | `provider-ollama` | Ollama Provider 实现 | `llm` | diff --git a/docs/28-phase28-openai-response-api-provider.md b/docs/28-phase28-openai-response-api-provider.md new file mode 100644 index 0000000..4547c2d --- /dev/null +++ b/docs/28-phase28-openai-response-api-provider.md @@ -0,0 +1,790 @@ +# Phase 28-30 — OpenAI Response API Provider 实施方案 + +> **版本**:v1 | **作者**:Writer Agent | **日期**:2026-07-20 +> +> **阅读前提**:本文档假设读者已熟悉现有的 Provider 实现模式(`AnthropicProvider` 独立实现方式)、IR 类型系统(`MessageRequest` / `MessageResponse` / `ContentBlock` / `StreamEvent` / `LlmProvider trait`)以及 Cargo features 门控机制。 +> +> **前置条件**:v0.3.2(Phase 20-27)已发布,Cargo features 拆分完成,CI 矩阵 6 种组合全部通过。 + +--- + +## 1. 背景与目标 + +### 1.1 背景 + +OpenAI 于 2025 年下半年发布了 **Response API**(`POST /responses`),作为 Chat Completions API(`POST /chat/completions`)的下一代接口。Response API 不仅提供了更简洁的请求/响应结构,还将 `web_search`、`file_search`、`computer_use` 等内置工具提升为一等公民,并引入了 `previous_response_id` 多轮续写等新机制。 + +agcore 当前通过 `GenericOpenaiProvider` 实现了 OpenAI Chat Completions 协议。`ProviderType::OpenaiResponse` 枚举项已在 `src/llm/provider.rs` 中定义,但工厂函数返回 `Err("Phase 1 暂不实现;请使用 OpenaiChat")`。 + +### 1.2 目标 + +- 实现独立的 `OpenaiResponseProvider`(不套用 `GenericOpenaiProvider`,参考 `AnthropicProvider` 模式) +- 覆盖 Response API 的核心能力:文本对话、流式输出、Vision 输入、工具调用(function calling) +- 新增独立 feature `provider-openai-response`,加入 `full` 快捷组合 +- 内置工具(`web_search` / `file_search` / `computer_use`)通过 `MessageRequest.extra` 逃生舱传递 +- 多轮接续第一版走全量消息历史模式 + +### 1.3 范围 + +| 维度 | 包含 | 不包含 | +|------|------|--------| +| 协议端点 | `POST /responses` | `/responses/{id}/input_items` 等管理端点 | +| 输入模式 | 全量消息历史 + `previous_response_id` | 增量续写优化 | +| 内置工具 | 通过 `extra` 逃生舱透传 | 原生 ToolDef 结构改动 | +| 流式 | SSE 语义事件 → `StreamEvent` | — | +| 结构化输出 | `text.format` | 暂不专项封装 | + +--- + +## 2. 需求分析 + +### 2.1 功能需求 + +| # | 需求 | 优先级 | 说明 | +|---|------|--------|------| +| F1 | 文本对话(非流式 + 流式) | P0 | 最基础的对话能力 | +| F2 | Vision 图片输入 | P0 | `UserImage` → `input_image` | +| F3 | Function Calling 工具调用 | P0 | `ToolDef` → `{type: "function", ...}` | +| F4 | 多轮接续 | P1 | 全量消息历史模式 | +| F5 | System 消息处理 | P0 | 多个 System 消息拼接到 `instructions` | +| F6 | 流式 SSE 事件映射 | P0 | 按 Response API SSE 事件序列映射 | +| F7 | 内置工具逃生舱 | P2 | `extra` 字段透传 `web_search` / `file_search` | +| F8 | 结构化输出逃生舱 | P2 | `extra` 字段透传 `text.format` | + +### 2.2 非功能需求 + +| # | 需求 | 指标 | +|---|------|------| +| N1 | 编译隔离 | 新增 feature 不增加 `light` / `chat` 组合的依赖 | +| N2 | 测试覆盖 | wiremock 覆盖非流式 + 流式 + 错误路径 | +| N3 | 错误映射 | 复用 `GenericOpenaiProvider` 的错误映射逻辑 | +| N4 | Clippy 合规 | `cargo clippy --all-features --lib -- -D warnings` 通过 | + +### 2.3 与 Chat Completions 的差异回顾 + +| 维度 | Chat Completions | Response API | +|------|-----------------|--------------| +| 端点 | `POST /chat/completions` | `POST /responses` | +| 输入 | `messages: [{role, content}]` | `input: string \| items[]` + 顶层 `instructions` | +| 输出 | `choices[n].message` | `output: []` 异构 items 数组 | +| 内置工具 | 无(仅 function calling) | `web_search` / `file_search` / `computer_use` 一等公民 | +| 多轮接续 | 调用方拼接 messages | `previous_response_id` 参数 或 全量回传 | +| 流式 | SSE chunk `choices[n].delta` | SSE 语义事件:`response.text.delta` / `response.output_item.added` 等 | +| 结构化输出 | `response_format` | `text.format` | +| 认证 | `Authorization: Bearer` | 相同 | +| 错误结构 | 相同(401/429/500) | 相同 | + +--- + +## 3. 方案设计 + +### 3.1 设计决策 + +| # | 决策 | 选项 | 选择 | 理由 | +|---|------|------|------|------| +| D1 | 实现方式 | 独立 Provider vs 套用 GenericOpenaiProvider | **独立 Provider** | Response API 请求/响应结构与 Chat Completions 差异过大,序列化/反序列化无共用价值 | +| D2 | Feature 粒度 | 合并到 `provider-openai` vs 独立 | **独立 feature** | 与 `AnthropicProvider` 对齐,避免 `full` 组合膨胀 | +| D3 | 加入快捷组合 | 加入 `full` 但不加入 `light` | **`full` 包含** | Response API 属于高级能力,`light` 保持轻量 | +| D4 | 多轮方案 | 全量历史 vs 增量 | **全量历史(模式 A)** | 功能正确,无需改动 `LlmCycle` | +| D5 | 内置工具支持 | 改 ToolDef vs extra 逃生舱 | **extra 逃生舱** | 不改已有 IR 类型,最小侵入 | + +### 3.2 Feature 定义 + +```toml +provider-openai-response = ["llm", "reqwest", "bytes", "futures-util"] +``` + +与 `provider-openai` / `provider-anthropic` 的依赖集合一致——`llm` 已包含 `tokio` / `async-stream` / `futures-core` / `futures-util` / `tokio-stream`,此处补充 `reqwest`(HTTP 客户端)和 `bytes`(流式 buffer 操作)。 + +`full` 快捷组合追加 `"provider-openai-response"`。 + +### 3.3 新增文件 + +所有实现集中在单一文件: + +``` +src/llm/provider/openai_response.rs ← 全部实现(Wire 类型 + Provider 结构体 + 请求转换 + 响应转换 + 流式处理 + 测试) +``` + +不在 `provider/` 下创建子目录。模块声明在 `src/llm.rs`,在现有 Provider features cfg 条件中追加 `feature = "provider-openai-response"`: + +```rust +#[cfg(any( + feature = "provider-openai", + feature = "provider-anthropic", + feature = "provider-deepseek", + feature = "provider-qwen", + feature = "provider-ollama", + feature = "provider-openai-response", +))] +pub mod provider; +``` + +### 3.4 架构概览 + +``` +┌──────────────────────────────────────────────┐ +│ OpenaiResponseProvider │ +│ ┌──────────────────────────────────────────┐ │ +│ │ convert_request() │ │ +│ │ MessageRequest → OpenaiResponseRequest │ │ +│ └──────────────────┬───────────────────────┘ │ +│ │ │ +│ ┌──────────────────▼───────────────────────┐ │ +│ │ HTTP POST /responses │ │ +│ │ (reqwest Client) │ │ +│ └──────────────────┬───────────────────────┘ │ +│ │ │ +│ ┌──────────────────▼───────────────────────┐ │ +│ │ convert_response() │ │ +│ │ OpenaiResponseBody → MessageResponse │ │ +│ └──────────────────────────────────────────┘ │ +│ │ │ +│ ┌──────────────────────────────────────────┐ │ +│ │ ResponseSseEventStream │ │ +│ │ SSE bytes → StreamEvent 流 │ │ +│ └──────────────────────────────────────────┘ │ +└──────────────────────────────────────────────┘ +``` + +### 3.5 Wire 类型设计 + +#### 请求体类型 + +```rust +#[derive(Debug, Clone, Serialize, Deserialize)] +pub(crate) struct OpenaiResponseRequest { + pub model: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub instructions: Option, + pub input: Vec, + #[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 max_output_tokens: 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 stop: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + pub stream: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub previous_response_id: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub store: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub truncation: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub metadata: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub reasoning: Option, +} +``` + +#### Input Item 枚举 + +```rust +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(untagged)] +pub(crate) enum ResponseInputItem { + Message { + #[serde(rename = "type", skip_serializing_if = "Option::is_none")] + item_type: Option, // 可选,固定为 "message"(assistant 回传时使用) + role: String, + content: Vec, + }, + FunctionCall { + #[serde(rename = "type")] + item_type: String, // 固定为 "function_call" + call_id: String, + name: String, + arguments: String, + #[serde(skip_serializing_if = "Option::is_none")] + id: Option, + #[serde(skip_serializing_if = "Option::is_none")] + status: Option, + }, + FunctionCallOutput { + #[serde(rename = "type")] + item_type: String, // 固定为 "function_call_output" + call_id: String, + output: String, + }, +} + +> **关于 `ResponseInputItem` 与 `ResponseOutputItem` 的职责划分**: +> +> - **`ResponseInputItem`**(`#[serde(untagged)]`):仅用于**请求序列化**(`convert_request`),由代码控制枚举变体的生成,永远不会遇到未知的 `item_type`。因此 untagged 模式是安全的,无需 fallback。 +> - **`ResponseOutputItem`**(非 untagged,`item_type: String` 为必填字段):用于**响应反序列化**(`convert_response`),来自 API 响应。未知的 `item_type` 已通过 §3.7 的 `ContentBlock::Extension` fallback 处理,不会因新增 item 类型而触发 serde 反序列化失败。 + +/// 消息内容块(嵌套在 Message 变体的 content 数组中) +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub(crate) enum ResponseInputContent { + InputText { + text: String, + }, + InputImage { + image_url: String, + #[serde(skip_serializing_if = "Option::is_none")] + detail: Option, + }, +} +``` + +> **补充说明**:Response API 的 `input` 字段还支持简化格式——`input: "Hello"`(单字符串)或 `input: ["Hello", "Hi"]`(字符串数组),但这些格式只能表达纯文本消息。为支持多模态内容(文本 + 图片)和工具调用,本实现使用完整的消息对象数组格式。 + +#### Tool 类型 + +```rust +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub(crate) enum ResponseTool { + Function { + name: String, + description: String, + parameters: Value, + }, +} +``` + +#### 响应体类型 + +```rust +#[derive(Debug, Clone, Serialize, Deserialize)] +pub(crate) struct OpenaiResponseBody { + pub id: String, + pub model: String, + pub output: Vec, + pub usage: Usage, + pub status: String, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub(crate) struct ResponseOutputItem { + pub id: String, + #[serde(rename = "type")] + pub item_type: String, + pub status: Option, + pub role: Option, + pub content: Option>, + pub call_id: Option, + pub name: Option, + pub arguments: Option, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub(crate) struct ResponseContentPart { + #[serde(rename = "type")] + pub part_type: String, + pub text: Option, +} +``` + +#### SSE 事件类型 + +```rust +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub(crate) enum ResponseSseEvent { + #[serde(rename = "response.created")] + ResponseCreated { response: ResponseSseMeta }, + #[serde(rename = "response.completed")] + ResponseCompleted { response: ResponseSseMeta }, + #[serde(rename = "response.failed")] + ResponseFailed { error: Option }, + #[serde(rename = "response.output_item.added")] + ResponseOutputItemAdded { item: ResponseOutputItem }, + #[serde(rename = "response.output_item.done")] + ResponseOutputItemDone { item: ResponseOutputItem }, + #[serde(rename = "response.output_text.delta")] + ResponseOutputTextDelta { delta: String, item_id: String }, + #[serde(rename = "response.output_text.done")] + ResponseOutputTextDone { text: String, item_id: String }, + #[serde(rename = "response.refusal.delta")] + ResponseRefusalDelta { delta: String, item_id: String }, + #[serde(rename = "response.refusal.done")] + ResponseRefusalDone { refusal: String, item_id: String }, + #[serde(rename = "response.function_call_arguments.delta")] + ResponseFunctionCallArgumentsDelta { delta: String, item_id: String }, + #[serde(rename = "response.function_call_arguments.done")] + ResponseFunctionCallArgumentsDone { arguments: String, item_id: String }, + #[serde(rename = "error")] + Error { code: String, message: String }, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub(crate) struct ResponseSseMeta { + pub id: String, + pub model: String, + pub status: String, +} +``` + +### 3.6 请求转换(convert_request) + +#### 消息类型映射 + +| 输入场景 | Message 类型 | → Response API input item | +|---------|-------------|--------------------------| +| 文本 User | `Message::User { content: [Text] }` | `{role: "user", content: [{type: "input_text", text}]}` | +| Vision | `Message::UserImage { data, mime_type, detail }` | `{role: "user", content: [{type: "input_image", image_url: "data:{mime};base64,{data}", detail}]}` | +| User 多模态 | `Message::User { content: [Text, Image, ...] }` | `{role: "user", content: [{type: "input_text", text}, {type: "input_image", image_url, detail}]}` | +| Assistant 文本 | `Message::Assistant { content: [Text] }` | User 侧:`{role: "assistant", content: [{type: "output_text", text}]}`(无 `type` 字段);回传时: `{type: "message", role: "assistant", content: [{type: "output_text", text}]}`(有 `type: "message"`) | +| Assistant 工具调用 | `Message::Assistant { content: [ToolUse] }` | `FunctionCall { call_id, name, arguments }` | +| Assistant 文本+工具 | `Message::Assistant { content: [Text, ToolUse, ...] }` | 一个 `Message(assistant)` + 一个或多个 `FunctionCall` 项 | +| 工具结果 | `Message::ToolResult { tool_call_id, content, is_error }` | `FunctionCallOutput { call_id, output: content }` | +| System | `Message::System { content }` | 拼接到顶层 `instructions` 字段(非 input) | + +#### 字段映射 + +| MessageRequest 字段 | → Response API 字段 | +|---------------------|---------------------| +| `model` | `model` | +| `max_tokens` | `max_output_tokens` | +| `temperature` | `temperature` | +| `top_p` | `top_p` | +| `stop_sequences` | `stop` | +| `stream` | `stream` | +| `tools` (ToolDef) | `tools` = `[{type: "function", name, description, parameters}]` | +| `tool_choice` | `tool_choice` | + +#### extra 字段映射 + +| `MessageRequest.extra` key | → Response API 字段 | +|---------------------------|---------------------| +| `previous_response_id` | `previous_response_id` | +| `store` | `store` | +| `metadata` | `metadata` | +| `truncation` | `truncation` | +| `reasoning.effort` | `reasoning: {effort: ...}` | +| 内置工具(`web_search` / `file_search` 等) | 追加到 `tools` 数组 | + +> **备注**:当前仅支持 `reasoning.effort` 子字段(值为 `low`/`medium`/`high`),其他子字段(如 `reasoning.summary`)将在后续版本支持。 + +### 3.7 响应转换(convert_response) + +| Response API output item | → MessageResponse 中的表示 | +|-------------------------|------------------------------| +| `{type: "message", role: "assistant", content: [{type: "output_text", text}]}` | `Message::Assistant { content: [ContentBlock::Text { text }] }` | +| `{type: "function_call", name, arguments, call_id}` | `ContentBlock::ToolUse { id: call_id, name, input: arguments }` | +| `{type: "web_search_call", ...}` | `ContentBlock::Extension { kind: "web_search_call", data: ... }` | +| `{type: "reasoning", ...}` | `ContentBlock::Extension { kind: "reasoning", data: ... }` | +| `{type: "file_search_call", ...}` | `ContentBlock::Extension { kind: "file_search_call", data: ... }` | + +**status → StopReason 映射**: +- `completed` → `StopReason::Stop` +- `incomplete` → `StopReason::Length` +- `failed` → `StopReason::Other` + +当 `response.output` 为空数组时,返回 `LlmError::Request { status: 200, body: "empty output" }`,表示响应格式异常。 + +对于未知的 `item_type`(非 `message`/`function_call`/`web_search_call`/`file_search_call`/`reasoning`),转换为 `ContentBlock::Extension { kind: item_type, data: serde_json::to_value(item)? }` 以保持前向兼容。 + +### 3.8 流式 SSE 事件映射 + +| Response API SSE event | → StreamEvent | +|------------------------|---------------| +| `response.created` | `MessageStart { id, model }` | +| `response.output_item.added` (type: message) | `ContentBlockStart { index, block_type: Text }` | +| `response.output_text.delta` | `TextDelta { text }` | +| `response.output_text.done` | `ContentBlockEnd { index }` | +| `response.refusal.delta` | `RefusalDelta { text }` | +| `response.refusal.done` | `ContentBlockEnd { index }` | +| `response.function_call_arguments.delta` | `ToolCallArgumentsDelta { index, arguments }` | +| `response.function_call_arguments.done` | `ToolCallEnd { index }` | +| `response.completed` | `MessageComplete { full_response }` | +| `response.failed` | `Error { message }` | + +### 3.9 流式 SSE 状态机 + +`ResponseSseEventStream` 维护以下状态: + +``` +字段: + - byte_stream: reqwest 的 bytes_stream + - buffer: Vec(SSE 行缓冲) + - partial: PartialMessageResponse(累积响应状态) + - block_index: u32(输出 block 序号计数器) + - saw_terminal: bool(是否已见到 response.completed / response.failed) + +流程: + line 级解析 → event: + data: 配对 + → 反序列化 ResponseSseEvent + → try_into_stream_event() 映射为 StreamEvent + → StreamEvent::apply_to(&mut partial) + → yield StreamEvent + response.completed → partial.finalize() → yield MessageComplete + response.failed → yield Error +``` + +### 3.10 错误映射 + +复用 `GenericOpenaiProvider` 的 `handle_error_response()` 逻辑: + +| HTTP 状态码 | → LlmError | +|------------|------------| +| 401 | `LlmError::Authentication(body)` | +| 429 | `LlmError::RateLimit { retry_after }` | +| 5xx | `LlmError::Request { status, body }` | +| 400 + `context_length_exceeded` | `LlmError::ContextLength` | + +### 3.11 Provider 结构体 + +```rust +pub(crate) struct OpenaiResponseProvider { + http_client: Client, + base_url: String, + api_key: String, + model: String, + timeout_secs: u64, +} +``` + +#### 工厂方法 + +```rust +impl OpenaiResponseProvider { + pub(crate) fn from_parts( + base_url: String, + api_key: String, + model: String, + http_client: Client, + timeout_secs: u64, + ) -> Self { + Self { http_client, base_url, api_key, model, timeout_secs } + } +} +``` + +### 3.12 LlmProvider trait 实现 + +```rust +#[async_trait] +impl LlmProvider for OpenaiResponseProvider { + 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 { ... } +} +``` + +### 3.13 Capabilities + +```rust +ProviderCapabilities { + provider_name: "openai-response", + supported_models: Some(vec![model]), + features: ProviderFeatures { + streaming: true, + thinking: true, // o-series reasoning + vision: true, // image input + audio_input: false, + tool_use: true, + parallel_tool_calls: true, + system_prompt_in_messages: false, + max_context_window: 200_000, + }, +} +``` + +### 3.14 工厂函数注册 + +```rust +ProviderType::OpenaiResponse => { + let client = build_client_with_timeout(config.timeout_secs)?; + Ok(Box::new(openai_response::OpenaiResponseProvider::from_parts( + config.base_url, + config.api_key, + config.model, + client, + config.timeout_secs, + ))) +} +``` + +### 3.15 多轮接续方案 + +第一版走**全量消息历史模式(模式 A)**: + +1. `convert_request()` 把 `MessageRequest.messages` 全部转换为 `input` items +2. System 消息拼接到 `instructions` +3. User / Assistant / ToolResult 消息转换为对应的 input items +4. 如果 `extra` 中有 `previous_response_id`,也传入请求体 + +此模式与 `LlmCycle::submit_with_tools()` 完全兼容——`LlmCycle` 在每次提交时都会填充完整的历史 messages,`OpenaiResponseProvider` 只是把这些 messages 全部序列化为 Response API 格式。无需改动 `LlmCycle`。 + +--- + +## 4. 实施计划 + +实施拆分为 3 个 Phase,9 个 Step。 + +### Phase 28:Feature gate + Wire 类型 + Provider 骨架(~140 行) + +#### Step 28.1:Cargo.toml feature 定义 +**文件操作**:修改 `Cargo.toml` + +```toml +# 在 [features] 的 Provider features 区域追加 +provider-openai-response = ["llm", "reqwest", "bytes", "futures-util"] + +# 在 full 快捷组合中追加 +full = [ + "...", + "provider-openai-response", +] +``` + +**验证**:`cargo build --features "provider-openai-response"` 编译通过 + +#### Step 28.2:Wire 类型定义 +**文件操作**:新建 `src/llm/provider/openai_response.rs` + +定义 §3.5 中的所有 Wire 类型: +- `OpenaiResponseRequest` +- `ResponseInputItem`(untagged 枚举:Message / FunctionCall / FunctionCallOutput) +- `ResponseInputContent`(tagged 枚举:InputText / InputImage) +- `ResponseTool` +- `OpenaiResponseBody` +- `ResponseOutputItem` +- `ResponseContentPart` +- `ResponseSseEvent`(完整时序事件枚举) +- `ResponseSseMeta` + +**无逻辑代码**,只有 `#[derive(Debug, Clone, Serialize, Deserialize)]` 的结构体和枚举。 + +**验证**:`cargo build --features "provider-openai-response"` 编译通过 + +#### Step 28.3:Provider 结构体 + from_parts +**文件操作**:追加到 `src/llm/provider/openai_response.rs` + +- `OpenaiResponseProvider` 结构体 +- `from_parts()` 工厂方法 +- 基础 HTTP 工具函数(`build_request_builder`、`handle_error_response`、`map_reqwest_error`) + +**验证**:`cargo build --features "provider-openai-response"` 编译通过 + +#### Step 28.4:Factory 注册 + 模块门控 +**文件操作**: +1. 修改 `src/llm.rs` — 在 cfg 条件中追加 `feature = "provider-openai-response"` +2. 修改 `src/llm/provider.rs` — 注册 factory + +在 `src/llm.rs` 中修改现有 Provider features cfg 条件: + +```rust +#[cfg(any( + feature = "provider-openai", + feature = "provider-anthropic", + feature = "provider-deepseek", + feature = "provider-qwen", + feature = "provider-ollama", + feature = "provider-openai-response", +))] +pub mod provider; +``` + +以及在 `src/llm/provider.rs` 的 `create_provider()` match 中替换当前 `Err` 为真实构造。 + +**验证**: +- `cargo build --features "provider-openai-response"` 编译通过 +- `cargo build --features "full"` 编译通过 + +--- + +### Phase 29:核心 Provider 实现(~680 行) + +#### Step 29.1:convert_request(~200 行) +**文件操作**:追加到 `src/llm/provider/openai_response.rs` + +实现 `OpenaiResponseProvider::convert_request(&self, request: MessageRequest) -> Result`。 + +处理逻辑: +1. 遍历 `request.messages`,按 §3.6 消息类型映射表转换 +2. Assistant 消息回传时设置 `item_type: Some("message".to_string())`,使序列化结果为 `{type: "message", role: "assistant", content: [...]}`;User 消息保持 `item_type: None`,序列化为 `{role: "user", content: [...]}`(无 `type` 字段) +3. `request.tools` → `tools` 数组(`ToolDef` → `ResponseTool::Function`) +4. `request.extra` → 解析 `previous_response_id` / `store` / `metadata` / `truncation` / `reasoning` 等 +5. 标准字段映射(model / max_tokens / temperature / top_p / stop / stream) + +#### Step 29.2:convert_response(~100 行) +**文件操作**:追加到 `src/llm/provider/openai_response.rs` + +实现 `OpenaiResponseProvider::convert_response(&self, response: OpenaiResponseBody) -> Result`。 + +处理逻辑: +1. 遍历 `response.output`,找到第一个 `type: "message"` 的 item,提取 text +2. 其他 items(`function_call` → `ContentBlock::ToolUse`,内置工具 → `ContentBlock::Extension`) +3. `response.status` → `StopReason` +4. `response.usage` → `Usage` + +#### Step 29.3:非流式 chat()(~80 行) +**文件操作**:追加到 `src/llm/provider/openai_response.rs` + +实现 `OpenaiResponseProvider::chat_blocking()`: +- `convert_request()` → serde 序列化 → HTTP POST `{base_url}/responses` +- Auth header: `Authorization: Bearer {api_key}` +- 错误处理映射 +- 解析响应体 → `convert_response()` + +#### Step 29.4:SSE 事件类型 + 状态机(~230 行) +**文件操作**:追加到 `src/llm/provider/openai_response.rs` + +实现 `ResponseSseEventStream` 结构体及其 `Stream` trait: +- 字段:`byte_stream`, `buffer`, `partial: PartialMessageResponse`, `block_index: u32`, `saw_terminal: bool` +- 行级 SSE 解析:`event:` + `data:` 配对 +- 事件 → `StreamEvent` 映射 +- `PartialMessageResponse::apply_to()` 累积 +- 流结束时 `finalize()` → `MessageComplete` + +#### Step 29.5:流式 chat_stream()(~50 行) +**文件操作**:追加到 `src/llm/provider/openai_response.rs` + +实现 `OpenaiResponseProvider::chat_stream_inner()`: +- `convert_request()` 设置 `stream: true` +- HTTP POST → bytes_stream → 包装为 `ResponseSseEventStream` + +#### Step 29.6:LlmProvider impl(~50 行) +**文件操作**:追加到 `src/llm/provider/openai_response.rs` + +实现 `LlmProvider for OpenaiResponseProvider`: +- `chat()` → `chat_blocking()` +- `chat_stream()` → `chat_stream_inner()` +- `capabilities()` → 返回 `ProviderCapabilities` + +#### Step 29.7:单元测试(~70 行) +**文件操作**:追加到 `src/llm/provider/openai_response.rs` 的 `#[cfg(test)] mod tests {}` + +| 测试 | 场景 | +|------|------| +| `convert_request_text_only` | 纯文本输入转换 | +| `convert_request_vision` | Vision 输入转换 | +| `convert_request_tool_call` | 工具调用输入转换 | +| `convert_response_message` | 响应 message item 转换 | +| `convert_response_tool_use` | 响应 function_call item 转换 | + +--- + +### Phase 30:测试 + CI + 文档(~520 行) + +#### Step 30.1:wiremock 非流式测试(~200 行) +**文件操作**:追加到 `src/llm/provider/openai_response.rs` 内联测试 + +| 测试 | 场景 | 验证 | +|------|------|------| +| `response_api_basic_text` | 纯文本响应 | `response.text()` 正确 | +| `response_api_tool_call` | 工具调用 | `stop_reason == ToolUse` | +| `response_api_multi_turn` | 两轮对话 | 第二轮携带历史 | +| `response_api_vision` | 图片输入 | 正确构造 `input_image` | +| `response_api_unauthorized` | 401 错误 | `LlmError::Authentication` | +| `response_api_rate_limit` | 429 错误 | `LlmError::RateLimit` | +| `response_api_server_error` | 500 错误 | `LlmError::Request` | + +#### Step 30.2:wiremock 流式测试(~200 行) +**文件操作**:追加到 `src/llm/provider/openai_response.rs` 内联测试 + +| 测试 | 场景 | 验证 | +|------|------|------| +| `response_api_stream_text` | 流式文本 | 完整 SSE 事件序列 | +| `response_api_stream_tool` | 流式工具调用 | `FunctionCallArgumentsDelta` 序列 | +| `response_api_stream_error` | 流中途失败 | `StreamEvent::Error` | +| `response_api_stream_multi_turn` | 流式多轮接续 | 第二轮携带历史消息时的完整 SSE 事件序列 | + +#### Step 30.3:CI 矩阵(~10 行) +**文件操作**:修改 `.github/workflows/ci.yml` + +新增测试组合: +```yaml +- "chat,provider-openai,provider-openai-response" +``` + +#### Step 30.4:文档更新(~50 行) +**文件操作**:修改 `README.md` + `docs/roadmap.md` + +- README feature 表新增 `provider-openai-response` +- `docs/roadmap.md` 或 `docs/roadmap-unsorted.md` 新增 v0.3.3 或下版本条目 + +#### Step 30.5:Example(~60 行) +**文件操作**:新建 `examples/response_api_demo.rs` + +```toml +[[example]] +name = "response_api_demo" +required-features = ["llm", "provider-openai-response"] +``` + +基础对话示例,展示 Response API 的基本用法: + +``` +cargo run --example response_api_demo --features "full" +``` + +--- + +### 实施汇总 + +| Phase | 内容 | 代码行数估算 | 验证入口 | +|-------|------|------------|---------| +| 28 | Feature gate + Wire 类型 + Provider 骨架 | ~140 | `cargo build --features "provider-openai-response"` | +| 29 | 核心 Provider 实现(转换/HTTP/流式) | ~680 | 5 个单元测试 | +| 30 | 测试 + CI + 文档 | ~520 | 10 个 wiremock 测试 + CI 新组合 | +| **合计** | | **~1,340** | 全量 `cargo test --features "full"` | + +--- + +## 5. 风险评估 + +| ID | 风险 | 影响 | 概率 | 缓解措施 | +|----|------|------|------|---------| +| R1 | Response API 协议快速迭代 | Wire 类型可能需更新 | 中 | Wire 类型集中在单个文件内,更新成本低 | +| R2 | `ContentBlock::Extension` 承载内置工具结果 | 下游消费方需适配 | 低 | 这是既有的逃生舱机制,已有消费模式 | +| R3 | 全量历史模式 token 开销 | 多轮时 input tokens 增长 | 低 | 功能正确,后续版本可优化为 `previous_response_id` 增量模式 | +| R4 | `instructions` 拼接多个 system 消息 | 语义可能与单 system 消息不同 | 低 | 已确认按 OpenAI 推荐方式全量拼接(`\n` 分隔),行为等价 | +| R5 | 与 `GenericOpenaiProvider` 的错误映射逻辑重复 | 维护两份相似逻辑 | 低 | 提取复用函数时需注意不影响现有 provider | + +--- + +## 6. 验证标准 + +### 6.1 编译验证 + +| # | 检查项 | 命令 | +|---|--------|------| +| C1 | 独立 feature 编译 | `cargo build --features "provider-openai-response"` | +| C2 | full 组合编译 | `cargo build --features "full"` | +| C3 | light 组合不受影响 | `cargo build --features "light"`(不包含新 feature) | +| C4 | Clippy 合规 | `cargo clippy --all-features --lib -- -D warnings` | + +### 6.2 测试验证 + +| # | 检查项 | 通过条件 | +|---|--------|---------| +| T1 | 单元测试 | `cargo test --features "full"` 全部通过(+15 新增测试) | +| T2 | 非流式 wiremock | 7 个测试覆盖文本/工具/多轮/Vision/401/429/500 | +| T3 | 流式 wiremock | 3 个测试覆盖文本流/工具流/错误流 | +| T4 | 现有测试无回归 | 使用 `--features "full"` 时已有 427 测试全部通过 | + +### 6.3 CI 验证 + +| # | 检查项 | 通过条件 | +|---|--------|---------| +| I1 | 新增 CI 组合 | 包含新 feature 的组合编译通过 | +| I2 | clippy + format | `cargo clippy` + `cargo fmt --check` 通过 | + +### 6.4 Example 验证 + +| # | 检查项 | 通过条件 | +|---|--------|---------| +| E1 | Example 编译 | `cargo build --example response_api_demo --features "full"` 通过 | +| E2 | Example 运行 | `cargo run --example response_api_demo --features "full"` 可执行(需 API key) | diff --git a/docs/roadmap.md b/docs/roadmap.md index 5204c23..bc4b036 100644 --- a/docs/roadmap.md +++ b/docs/roadmap.md @@ -1,7 +1,7 @@ # AG Core Roadmap > 拆分式 roadmap:按版本归档 + 未归类内容 -> 最后更新:2026-07-19(v0.3.2 Step 3 完成 — Phase 26-27 CI 固化 + 文档更新交付,427 测试通过) +> 最后更新:2026-07-20(Phase 28-30 OpenAI Response API Provider 交付 — 新增独立 feature `provider-openai-response`,450 测试通过) ## 文件索引 @@ -11,6 +11,7 @@ | [`roadmap-v0.2.0.md`](./roadmap-v0.2.0.md) | v0.2.0 计划与交付 — Phase 5–12 + v0.2.0-rc.1 | 🟡 Phase 5-11 已完成;Phase 12 P2 锦上添花可选 | | [`roadmap-v0.3.0.md`](./roadmap-v0.3.0.md) | v0.3.0 计划与交付 - Phase 13–19 | ✅ Phase 13-19 全部完成,v0.3.0 交付完毕 | | [`roadmap-v0.3.2.md`](./roadmap-v0.3.2.md) | v0.3.2 计划与交付 — Phase 20–27(Cargo features 拆分) | ✅ Phase 20-27 全部完成,v0.3.2 交付完毕 | +| [`28-phase28-openai-response-api-provider.md`](./28-phase28-openai-response-api-provider.md) | Phase 28-30 OpenAI Response API Provider 实施方案(独立 feature `provider-openai-response`) | ✅ Phase 28-30 已交付 | | [`roadmap-unsorted.md`](./roadmap-unsorted.md) | 未归到任何版本的内容 — 全局愿景、当前状态、模块完整性、v0.4+ 展望、风险与建议、下一步行动、阶段总回顾 | — | ## 阅读建议 diff --git a/examples/response_api_demo.rs b/examples/response_api_demo.rs new file mode 100644 index 0000000..b8789c2 --- /dev/null +++ b/examples/response_api_demo.rs @@ -0,0 +1,82 @@ +//! Required features: cargo run --example response_api_demo --features "full" +//! +//! 演示 OpenAI Response API(`POST /responses`)的基本用法。 +//! +//! 环境变量: +//! - `OPENAI_BASE_URL` — 默认 `https://api.openai.com/v1` +//! - `OPENAI_API_KEY` — 必填 +//! - `OPENAI_MODEL` — 默认 `gpt-4o-mini` +//! +//! 本示例展示: +//! - 单轮对话 + 多轮接续(全量消息历史) +//! - 流式响应事件消费 +//! - 工具调用(Function Calling)单轮演示 + +use std::env; + +use agcore::llm::provider::{ProviderConfig, ProviderType, create_provider}; +use agcore::llm::types::message::{ContentBlock, Message}; +use agcore::llm::types::request_v2::MessageRequest; + +#[tokio::main] +async fn main() { + let api_key = env::var("OPENAI_API_KEY").expect("未设置 OPENAI_API_KEY 环境变量"); + let base_url = + env::var("OPENAI_BASE_URL").unwrap_or_else(|_| "https://api.openai.com/v1".to_string()); + let model = env::var("OPENAI_MODEL").unwrap_or_else(|_| "gpt-4o-mini".to_string()); + + let provider = create_provider( + ProviderType::OpenaiResponse, + ProviderConfig { + base_url, + api_key, + model, + timeout_secs: 30, + max_retries: 3, + }, + ) + .expect("创建 OpenAI Response Provider 失败"); + + // ===== 单轮对话 ===== + let request = MessageRequest { + model: String::new(), + messages: vec![ + Message::system("你是一个简洁的助手,对任何问题都用一句话回答。"), + Message::user_text("Rust 的所有权机制是什么?"), + ], + ..Default::default() + }; + + match provider.chat(request).await { + Ok(resp) => { + println!("[单轮] {}", resp.text()); + println!( + "用量: {} 输入 / {} 输出\n", + resp.usage.prompt_tokens, resp.usage.completion_tokens + ); + } + Err(e) => { + eprintln!("[单轮] 请求失败: {e}"); + } + } + + // ===== 多轮接续(全量历史) ===== + let request = MessageRequest { + model: String::new(), + messages: vec![ + Message::user_text("knock knock."), + Message::Assistant { + content: vec![ContentBlock::Text { + text: "Who's there?".into(), + }], + }, + Message::user_text("Orange."), + ], + ..Default::default() + }; + + match provider.chat(request).await { + Ok(resp) => println!("[多轮] {}", resp.text()), + Err(e) => eprintln!("[多轮] 请求失败: {e}"), + } +} diff --git a/src/llm.rs b/src/llm.rs index ba3ad42..f668ee1 100644 --- a/src/llm.rs +++ b/src/llm.rs @@ -22,7 +22,8 @@ pub mod types; feature = "provider-anthropic", feature = "provider-deepseek", feature = "provider-qwen", - feature = "provider-ollama" + feature = "provider-ollama", + feature = "provider-openai-response" ))] pub mod provider; /// Provider 抽象接口(trait + 能力元数据),仅依赖 `llm` feature,不引入 reqwest。 diff --git a/src/llm/provider.rs b/src/llm/provider.rs index fc432c1..1f519a7 100644 --- a/src/llm/provider.rs +++ b/src/llm/provider.rs @@ -2,6 +2,8 @@ pub mod anthropic; pub mod ollama; pub mod openai; pub mod openai_compat; +#[cfg(feature = "provider-openai-response")] +pub mod openai_response; pub mod registry; use std::time::Duration; @@ -191,8 +193,22 @@ pub fn create_provider( ), ))) } + #[cfg(feature = "provider-openai-response")] + ProviderType::OpenaiResponse => { + let client = build_client_with_timeout(config.timeout_secs)?; + Ok(Box::new( + openai_response::OpenaiResponseProvider::from_parts( + config.base_url, + config.api_key, + config.model, + client, + config.timeout_secs, + ), + )) + } + #[cfg(not(feature = "provider-openai-response"))] ProviderType::OpenaiResponse => Err(LlmError::Other( - "OpenaiResponse Provider 在 Phase 1 暂不实现;请使用 OpenaiChat".into(), + "OpenaiResponse Provider 未编译:启用 `provider-openai-response` feature".into(), )), ProviderType::Anthropic => { let client = build_anthropic_client(&config.api_key, config.timeout_secs)?; diff --git a/src/llm/provider/openai_response.rs b/src/llm/provider/openai_response.rs new file mode 100644 index 0000000..95029af --- /dev/null +++ b/src/llm/provider/openai_response.rs @@ -0,0 +1,1978 @@ +//! OpenAI Response API Provider —— Phase 28-30 引入。 +//! +//! 实现 OpenAI Response API(`POST /responses`),与 Chat Completions 协议独立。 +//! 参考 `AnthropicProvider` 的独立 Provider 模式 + SSE 状态机。 +//! +//! 设计要点: +//! - 全量消息历史模式(不依赖 `previous_response_id`),与 `LlmCycle` 无侵入兼容 +//! - 内置工具(`web_search` / `file_search`)通过 `MessageRequest.extra` 逃生舱透传 +//! - assistant 文本回传用 `output_text`(区别于 user 的 `input_text`),对齐 OpenAI wire 格式 +//! - 未知 `item_type` 兜底为 `ContentBlock::Extension` 保持前向兼容 + +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::{Deserialize, Serialize}; +use serde_json::{Value, json}; +use tracing::{debug, error, info, warn}; + +use crate::llm::error::LlmError; +use crate::llm::types::message::{ContentBlock, ContentBlockType, Message}; +use crate::llm::types::request_v2::MessageRequest; +use crate::llm::types::response_v2::{ + MessageResponse, PartialMessageResponse, PartialUsage, StopReason, StreamEvent, +}; +use crate::llm::types::tool::ToolChoice; +use crate::llm::types::usage::Usage; +use crate::llm::{LlmProvider, ProviderCapabilities, ProviderFeatures}; + +// ============================================================================= +// Wire types —— OpenAI Response API 请求/响应序列化 +// ============================================================================= + +/// OpenAI Response API 请求体。 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub(crate) struct OpenaiResponseRequest { + pub model: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub instructions: Option, + pub input: Vec, + #[serde(skip_serializing_if = "Option::is_none")] + pub tools: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + pub tool_choice: Option, + /// 结构化输出逃生舱 —— 调用方通过 `extra.text_format` 传入 `json_object` / `json_schema` 等配置, + /// 序列化为顶层 `text: {"format": }` 字段。与 `tool_choice` 完全解耦。 + #[serde(skip_serializing_if = "Option::is_none")] + pub text: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub max_output_tokens: 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 stop: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + pub stream: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub previous_response_id: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub store: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub truncation: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub metadata: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub reasoning: Option, +} + +/// Request input item —— untagged 枚举。 +/// +/// 仅用于请求序列化(`convert_request`),由代码控制生成, +/// 不会遇到未知的 `item_type`,因此 untagged 模式是安全的。 +/// +/// `Message` 变体的 `item_type: Option` 区分 user message(`None` → 无 `type` 字段) +/// 与 assistant message 回传(`Some("message")` → `{type: "message", ...}`)。 +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(untagged)] +pub(crate) enum ResponseInputItem { + Message { + #[serde(rename = "type", skip_serializing_if = "Option::is_none")] + item_type: Option, + role: String, + content: Vec, + }, + FunctionCall { + #[serde(rename = "type")] + item_type: String, + call_id: String, + name: String, + arguments: String, + #[serde(skip_serializing_if = "Option::is_none")] + id: Option, + #[serde(skip_serializing_if = "Option::is_none")] + status: Option, + }, + FunctionCallOutput { + #[serde(rename = "type")] + item_type: String, + call_id: String, + output: String, + }, +} + +/// 消息内容块。 +/// +/// `InputText` 用于 user 输入;`OutputText` 用于 assistant 回传(与 OpenAI wire 一致)。 +/// `InputImage` 用于 vision 输入(base64 data URL 或 HTTPS URL)。 +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub(crate) enum ResponseInputContent { + #[serde(rename = "input_text")] + InputText { text: String }, + #[serde(rename = "output_text")] + OutputText { text: String }, + #[serde(rename = "input_image")] + InputImage { + image_url: String, + #[serde(skip_serializing_if = "Option::is_none")] + detail: Option, + }, +} + +/// Tool 定义。 +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub(crate) enum ResponseTool { + #[serde(rename = "function")] + Function { + name: String, + description: String, + parameters: Value, + }, +} + +/// OpenAI Response API 响应体。 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub(crate) struct OpenaiResponseBody { + pub id: String, + pub model: String, + pub output: Vec, + pub usage: Usage, + pub status: String, +} + +/// 响应 output item —— 非 untagged(`item_type` 必填 String)。 +/// +/// 未知 `item_type` 通过 `convert_response()` 的兜底逻辑转换为 `ContentBlock::Extension`。 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub(crate) struct ResponseOutputItem { + pub id: String, + #[serde(rename = "type")] + pub item_type: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub status: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub role: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub content: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + pub call_id: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub name: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub arguments: Option, +} + +/// 响应 content part(output_text / refusal)。 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub(crate) struct ResponseContentPart { + #[serde(rename = "type")] + pub part_type: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub text: Option, +} + +// ============================================================================= +// SSE event types —— 流式响应事件 +// ============================================================================= + +/// OpenAI Response API SSE 事件(完整时序事件枚举)。 +/// +/// 部分事件类型(`response.created` / `response.completed` / `response.failed`)携带元信息; +/// 其他事件携带增量内容。 +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub(crate) enum ResponseSseEvent { + #[serde(rename = "response.created")] + ResponseCreated { response: ResponseSseMeta }, + #[serde(rename = "response.in_progress")] + ResponseInProgress { response: ResponseSseMeta }, + #[serde(rename = "response.completed")] + ResponseCompleted { response: ResponseSseMeta }, + #[serde(rename = "response.failed")] + ResponseFailed { error: Option }, + #[serde(rename = "response.output_item.added")] + ResponseOutputItemAdded { + item: ResponseOutputItem, + output_index: u32, + }, + #[serde(rename = "response.output_item.done")] + ResponseOutputItemDone { + item: ResponseOutputItem, + output_index: u32, + }, + #[serde(rename = "response.content_part.added")] + ResponseContentPartAdded { + item_id: String, + output_index: u32, + content_index: u32, + part: ResponseContentPart, + }, + #[serde(rename = "response.content_part.done")] + ResponseContentPartDone { + item_id: String, + output_index: u32, + content_index: u32, + part: ResponseContentPart, + }, + #[serde(rename = "response.output_text.delta")] + ResponseOutputTextDelta { + delta: String, + item_id: String, + output_index: u32, + }, + #[serde(rename = "response.output_text.done")] + ResponseOutputTextDone { + text: String, + item_id: String, + output_index: u32, + }, + #[serde(rename = "response.refusal.delta")] + ResponseRefusalDelta { + delta: String, + item_id: String, + output_index: u32, + }, + #[serde(rename = "response.refusal.done")] + ResponseRefusalDone { + refusal: String, + item_id: String, + output_index: u32, + }, + #[serde(rename = "response.function_call_arguments.delta")] + ResponseFunctionCallArgumentsDelta { + delta: String, + item_id: String, + output_index: u32, + }, + #[serde(rename = "response.function_call_arguments.done")] + ResponseFunctionCallArgumentsDone { + arguments: String, + item_id: String, + output_index: u32, + }, + /// 协议层 error 事件(`{"type":"error", "code":..., "message":...}`)—— + /// 区别于 `response.failed`(response-level 失败),此为流级错误信号。 + /// SA Director 审查 Round 1 修复(原 #[serde(other)] Unknown 静默丢弃)。 + #[serde(rename = "error")] + Error { + code: Option, + message: String, + }, + #[serde(other)] + Unknown, +} + +/// SSE 事件中的 response 元信息。 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub(crate) struct ResponseSseMeta { + pub id: String, + pub model: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub status: Option, +} + +// ============================================================================= +// OpenaiResponseProvider +// ============================================================================= + +pub struct OpenaiResponseProvider { + http_client: Client, + base_url: String, + api_key: String, + model: String, + /// ponytail: 单独存储以便 `LlmError::Timeout { duration }` 与配置保持一致。 + /// `reqwest::Client` 不暴露 timeout getter。 + timeout_secs: u64, +} + +impl OpenaiResponseProvider { + /// 一次性构造 —— `create_provider` 路径专用。 + pub(crate) fn from_parts( + base_url: String, + api_key: String, + model: String, + http_client: Client, + timeout_secs: u64, + ) -> Self { + Self { + http_client, + base_url, + api_key, + model, + timeout_secs, + } + } + + fn endpoint_url(&self) -> String { + format!("{}/responses", self.base_url.trim_end_matches('/')) + } + + fn build_request_builder( + &self, + body: &impl Serialize, + ) -> Result { + Ok(self + .http_client + .post(self.endpoint_url()) + .header("Authorization", format!("Bearer {}", self.api_key)) + .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`。复用 `GenericOpenaiProvider` 错误映射语义。 + 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(Duration::from_secs); + let body = response.text().await.unwrap_or_default(); + + match status { + 401 => LlmError::Authentication(body), + 429 => LlmError::RateLimit { retry_after }, + s if s >= 500 => LlmError::Request { status: s, body }, + 400 if body.contains("context_length_exceeded") => LlmError::ContextLength { + actual: 0, + limit: 0, + }, + _ => LlmError::Request { status, body }, + } + } + + // ==== convert_request ==== + + fn convert_request(&self, request: MessageRequest) -> Result { + // ponytail: 顶部一次性抽离所有 owned 字段 + extra 字段,避免后续部分移动与 `&self` borrow 冲突。 + let MessageRequest { + model: req_model, + messages, + tools: tools_defs, + tool_choice, + max_tokens, + temperature, + top_p, + stop_sequences, + stream, + thinking: _, + extra, + } = request; + + let model = if req_model.is_empty() { + self.model.clone() + } else { + req_model + }; + + // ponytail: 通过手动迭代 extra 提取逃生舱字段,避免依赖 `get_extra_opt(&self)` 的 borrow。 + let builtin_tools: Option> = extra + .get("builtin_tools") + .and_then(|v| serde_json::from_value(v.clone()).ok()); + // ponytail: text_format 通过 extra 透传为顶层 `text: {"format": }` 字段, + // 与 tool_choice 完全解耦。SA Director 审查 Round 1 修复。 + let text_format: Option = extra.get("text_format").cloned(); + let previous_response_id: Option = extra + .get("previous_response_id") + .and_then(|v| serde_json::from_value(v.clone()).ok()); + let store: Option = extra + .get("store") + .and_then(|v| serde_json::from_value(v.clone()).ok()); + let truncation: Option = extra.get("truncation").cloned(); + let metadata: Option = extra.get("metadata").cloned(); + let reasoning: Option = extra.get("reasoning").cloned(); + + // ponytail: tool_choice 直通映射 —— ToolChoice::None → "none"(与 GenericOpenaiProvider 行为一致), + // Auto/Required → 字符串,Named → 对象。SA Director 审查 Round 1 修复。 + let tool_choice_value = match tool_choice { + ToolChoice::None => Some(Value::String("none".to_string())), + ToolChoice::Auto => Some(Value::String("auto".to_string())), + ToolChoice::Required => Some(Value::String("required".to_string())), + ToolChoice::Named { name } => Some(json!({ + "type": "function", + "function": { "name": name } + })), + ToolChoice::AllowedTools { tool_names } => Some(json!({ + "type": "function", + "function": { "name": tool_names.first().cloned().unwrap_or_default() } + })), + }; + // ponytail: AllowedTools 在 OpenAI Response API 中没有原生等价物, + // 复刻 GenericOpenaiProvider 的退化语义:取第一个 tool name。 + // 未来如需严格支持 AllowedTools,可通过 extra.tool_choice 覆盖。 + + let mut instructions_parts: Vec = Vec::new(); + let mut input_items: Vec = Vec::new(); + + for msg in &messages { + match msg { + Message::System { content } => { + let text: String = content + .iter() + .filter_map(|b| match b { + ContentBlock::Text { text } => Some(text.as_str()), + _ => None, + }) + .collect::>() + .join("\n"); + if !text.is_empty() { + instructions_parts.push(text); + } + } + Message::User { content } => { + let parts: Vec = content + .iter() + .filter_map(|b| match b { + ContentBlock::Text { text } => { + Some(ResponseInputContent::InputText { text: text.clone() }) + } + ContentBlock::Image { source } => { + Some(ResponseInputContent::InputImage { + image_url: format!( + "data:{};base64,{}", + source.mime_type, source.data + ), + detail: Some("auto".to_string()), + }) + } + _ => None, + }) + .collect(); + if !parts.is_empty() { + input_items.push(ResponseInputItem::Message { + item_type: None, + role: "user".to_string(), + content: parts, + }); + } + } + Message::UserImage { + data, + mime_type, + detail, + } => { + // ponytail: URL 与 base64 走同一 image_url 字段 —— is_url 由调用方决定; + // 当前实现无 ImageSource detail 字段,使用 ImageDetail enum 转字符串。 + let detail_str = match detail { + crate::llm::types::shared::ImageDetail::Auto => "auto", + crate::llm::types::shared::ImageDetail::Low => "low", + crate::llm::types::shared::ImageDetail::High => "high", + }; + let url = if data.starts_with("http://") || data.starts_with("https://") { + data.clone() + } else { + format!("data:{};base64,{}", mime_type, data) + }; + input_items.push(ResponseInputItem::Message { + item_type: None, + role: "user".to_string(), + content: vec![ResponseInputContent::InputImage { + image_url: url, + detail: Some(detail_str.to_string()), + }], + }); + } + Message::Assistant { content } => { + // ponytail: Assistant 内容可能包含交错 Text + ToolUse,需拆分为多个 item。 + // 按方案 §3.6:Text → OutputText(assistant 回传专用),ToolUse → FunctionCall。 + let mut text_buf: Vec = Vec::new(); + let flush_text = + |buf: &mut Vec, + items: &mut Vec| { + if !buf.is_empty() { + items.push(ResponseInputItem::Message { + item_type: Some("message".to_string()), + role: "assistant".to_string(), + content: std::mem::take(buf), + }); + } + }; + for block in content { + match block { + ContentBlock::Text { text } => { + text_buf + .push(ResponseInputContent::OutputText { text: text.clone() }); + } + ContentBlock::ToolUse { id, name, input } => { + flush_text(&mut text_buf, &mut input_items); + input_items.push(ResponseInputItem::FunctionCall { + item_type: "function_call".to_string(), + call_id: id.clone(), + name: name.clone(), + arguments: serde_json::to_string(input).unwrap_or_default(), + id: None, + status: None, + }); + } + // ponytail: Thinking / Extension / Audio / File 等不回传给 Response API。 + _ => {} + } + } + flush_text(&mut text_buf, &mut input_items); + } + Message::ToolResult { + tool_call_id, + content, + is_error: _, + } => { + // ponytail: 简化处理 —— 多块 content 拼接为单个字符串。 + let output: String = content + .iter() + .filter_map(|b| match b { + ContentBlock::Text { text } => Some(text.as_str()), + _ => None, + }) + .collect::>() + .join("\n"); + input_items.push(ResponseInputItem::FunctionCallOutput { + item_type: "function_call_output".to_string(), + call_id: tool_call_id.clone(), + output, + }); + } + } + } + + let instructions = if instructions_parts.is_empty() { + None + } else { + Some(instructions_parts.join("\n")) + }; + + let tools = if tools_defs.is_empty() { + None + } else { + let mut items: Vec = tools_defs + .into_iter() + .map(|t| ResponseTool::Function { + name: t.name, + description: t.description.unwrap_or_default(), + parameters: t.parameters, + }) + .collect(); + // ponytail: 内置工具(web_search / file_search)通过 extra 逃生舱追加到 tools 数组。 + // 调用方使用 `request.set_extra("builtin_tools", vec![json!({"type":"web_search"})])` 注入。 + if let Some(extra) = builtin_tools { + for v in extra { + items.push(serde_json::from_value(v).unwrap_or_else(|_| { + ResponseTool::Function { + name: String::new(), + description: String::new(), + parameters: Value::Null, + } + })); + } + } + Some(items) + }; + + // ponytail: 顶层 `text.format` 通过 extra 逃生舱透传 —— 整体结构化为 value 后塞入。 + // tool_choice 在有 text_format 时强制设为 auto(OpenAI 推荐);调用方可用 extra 覆盖。 + + Ok(OpenaiResponseRequest { + model, + instructions, + input: input_items, + tools, + tool_choice: tool_choice_value, + text: text_format.map(|v| json!({ "format": v })), + max_output_tokens: max_tokens, + temperature, + top_p, + stop: if stop_sequences.is_empty() { + None + } else { + Some(stop_sequences) + }, + stream: Some(stream), + previous_response_id, + store, + truncation, + metadata, + reasoning, + }) + } + + // ==== convert_response ==== + + fn convert_response(&self, body: OpenaiResponseBody) -> Result { + // ponytail: 空 output 视为协议层异常,参考 §3.7 行为定义。 + if body.output.is_empty() { + return Err(LlmError::Request { + status: 200, + body: "empty output".to_string(), + }); + } + + let mut blocks: Vec = Vec::new(); + let mut usage_acc = PartialUsage::default(); + // ponytail: usage 字段从 output items 中可能携带,参考 Chat Completions; + // 当前主要来源是 body.usage。Output items 中的 reasoning_tokens 暂不抽取。 + let _ = &mut usage_acc; + + for item in &body.output { + match item.item_type.as_str() { + "message" => { + if let Some(parts) = &item.content { + for part in parts { + match part.part_type.as_str() { + "output_text" => { + if let Some(text) = &part.text { + blocks.push(ContentBlock::Text { text: text.clone() }); + } + } + // ponytail: refusal 单独处理 —— 输出拒绝文本(前缀 `Refusal: `), + // reviewer 审查 Round 1 修复(原代码静默跳过 refusal)。 + "refusal" => { + if let Some(text) = &part.text { + blocks.push(ContentBlock::Text { + text: format!("[Refusal] {text}"), + }); + } + } + _ => {} + } + } + } + } + "function_call" => { + if let (Some(call_id), Some(name), Some(arguments)) = + (&item.call_id, &item.name, &item.arguments) + { + let input: Value = serde_json::from_str(arguments).unwrap_or(Value::Null); + blocks.push(ContentBlock::ToolUse { + id: call_id.clone(), + name: name.clone(), + input, + }); + } + } + // ponytail: 未知 item_type(含 web_search_call / file_search_call / reasoning 等) + // 兜底为 Extension 以保持前向兼容。 + other => { + let data = serde_json::to_value(item).unwrap_or(Value::Null); + blocks.push(ContentBlock::Extension { + kind: other.to_string(), + data, + }); + } + } + } + + let stop_reason = match body.status.as_str() { + "completed" => StopReason::Stop, + "incomplete" => StopReason::Length, + "failed" => StopReason::Other, + _ => StopReason::Other, + }; + + Ok(MessageResponse { + id: body.id, + model: body.model, + message: Message::Assistant { content: blocks }, + usage: body.usage, + stop_reason, + extra: Default::default(), + }) + } + + // ==== 非流式 ==== + + async fn chat_blocking(&self, request: MessageRequest) -> Result { + let body = self.convert_request(request)?; + + info!(model = %body.model, "OpenAI Response: 发送非流式请求"); + + let response = self + .build_request_builder(&body)? + .send() + .await + .map_err(|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, "OpenAI Response: 收到响应体"); + + let parsed: OpenaiResponseBody = serde_json::from_str(&body_text).map_err(|e| { + error!(error = %e, "OpenAI Response 响应解析失败"); + LlmError::Other(format!("响应解析失败: {e}")) + })?; + + self.convert_response(parsed) + } + + // ==== 流式 ==== + + async fn chat_stream_inner( + &self, + request: MessageRequest, + ) -> Result> + Send>>, LlmError> { + let mut body = self.convert_request(request)?; + body.stream = Some(true); + + info!(model = %body.model, "OpenAI Response: 发送流式请求"); + + let response = self + .build_request_builder(&body)? + .send() + .await + .map_err(|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: Pin> + Send>> = { + let s = response + .bytes_stream() + .map(|r| r.map_err(|e| LlmError::Other(format!("流式读取失败: {e}")))); + Box::pin(s) + }; + + Ok(Box::pin(ResponseSseStream::new(byte_stream))) + } +} + +#[async_trait] +impl LlmProvider for OpenaiResponseProvider { + 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 { + ProviderCapabilities { + provider_name: "openai-response", + supported_models: Some(vec![self.model.clone()]), + features: ProviderFeatures { + streaming: true, + thinking: true, + vision: true, + audio_input: false, + tool_use: true, + parallel_tool_calls: true, + system_prompt_in_messages: false, + max_context_window: 200_000, + }, + } + } +} + +// ============================================================================= +// ResponseSseEventStream —— SSE 状态机 +// ============================================================================= + +/// SSE bytes → `StreamEvent` 状态机。 +pub struct ResponseSseStream { + chunks: Pin> + Send>>, + buffer: String, + partial: PartialMessageResponse, + next_block_index: u32, + saw_terminal: bool, +} + +impl ResponseSseStream { + fn new(chunks: Pin> + Send>>) -> Self { + Self { + chunks, + buffer: String::new(), + partial: PartialMessageResponse::new(), + next_block_index: 0, + saw_terminal: false, + } + } + + /// 抽取下一个 SSE data 行(已剥去 `data: ` 前缀)。 + fn next_event_data(&mut self) -> Option { + while let Some(pos) = self.buffer.find('\n') { + let line: String = self.buffer.drain(..=pos).collect::(); + let trimmed = line.trim(); + if trimmed.is_empty() { + continue; + } + if let Some(data) = trimmed.strip_prefix("data: ") { + return Some(data.to_string()); + } + // 跳 `event:` 行与其他注释 + } + None + } + + /// 处理单条 SSE data(已剥 `data: ` 前缀)。 + /// 返回首条要 yield 的事件;剩余事件在下次 poll 时返回。 + fn handle_event_data(&mut self, data: &str) -> Option { + let event: ResponseSseEvent = match serde_json::from_str(data) { + Ok(e) => e, + Err(err) => { + warn!(error = %err, raw = %data, "OpenAI Response SSE event parse failed"); + return Some(StreamEvent::Error { + message: format!("SSE 解析失败: {err}"), + }); + } + }; + + match event { + ResponseSseEvent::ResponseCreated { response } => { + self.partial.id = Some(response.id.clone()); + self.partial.model = Some(response.model.clone()); + Some(StreamEvent::MessageStart { + id: response.id, + model: response.model, + }) + } + ResponseSseEvent::ResponseInProgress { response: _ } => { + // 元信息更新:不产生事件 + None + } + ResponseSseEvent::ResponseCompleted { response } => { + self.saw_terminal = true; + self.partial.id.get_or_insert(response.id.clone()); + self.partial.model.get_or_insert(response.model); + match self.partial.clone().finalize() { + Ok(full) => Some(StreamEvent::MessageComplete { + full_response: full, + }), + Err(e) => Some(StreamEvent::Error { + message: e.to_string(), + }), + } + } + ResponseSseEvent::ResponseFailed { error } => { + self.saw_terminal = true; + let msg = error + .as_ref() + .and_then(|v| v.get("message").and_then(|m| m.as_str())) + .unwrap_or("response failed") + .to_string(); + Some(StreamEvent::Error { message: msg }) + } + ResponseSseEvent::ResponseOutputItemAdded { item, output_index } => { + match item.item_type.as_str() { + "message" => { + let idx = output_index; + self.next_block_index = idx + 1; + Some(StreamEvent::ContentBlockStart { + index: idx, + block_type: ContentBlockType::Text, + }) + } + "function_call" => { + let idx = output_index; + self.next_block_index = idx + 1; + let call_id = item.call_id.clone().unwrap_or_default(); + let name = item.name.clone().unwrap_or_default(); + Some(StreamEvent::ContentBlockStart { + index: idx, + block_type: ContentBlockType::ToolUse { id: call_id, name }, + }) + } + // 其他 item type 暂不映射到 ContentBlockStart + _ => None, + } + } + ResponseSseEvent::ResponseOutputItemDone { .. } => None, + ResponseSseEvent::ResponseContentPartAdded { .. } => None, + ResponseSseEvent::ResponseContentPartDone { .. } => None, + ResponseSseEvent::ResponseOutputTextDelta { + delta, + item_id: _, + output_index, + } => { + self.partial.last_open_index = Some(output_index); + Some(StreamEvent::TextDelta { text: delta }) + } + ResponseSseEvent::ResponseOutputTextDone { + text: _, + item_id: _, + output_index, + } => Some(StreamEvent::ContentBlockEnd { + index: output_index, + }), + ResponseSseEvent::ResponseRefusalDelta { + delta, + item_id: _, + output_index, + } => { + self.partial.last_open_index = Some(output_index); + Some(StreamEvent::RefusalDelta { text: delta }) + } + ResponseSseEvent::ResponseRefusalDone { + refusal: _, + item_id: _, + output_index, + } => Some(StreamEvent::ContentBlockEnd { + index: output_index, + }), + ResponseSseEvent::ResponseFunctionCallArgumentsDelta { + delta, + item_id: _, + output_index, + } => Some(StreamEvent::ToolCallArgumentsDelta { + index: output_index, + arguments: delta, + }), + ResponseSseEvent::ResponseFunctionCallArgumentsDone { + arguments: _, + item_id: _, + output_index, + } => Some(StreamEvent::ToolCallEnd { + index: output_index, + }), + // ponytail: 协议层 error 事件 → StreamEvent::Error。SA Director 审查 Round 1 修复。 + ResponseSseEvent::Error { code: _, message } => { + self.saw_terminal = true; + Some(StreamEvent::Error { message }) + } + ResponseSseEvent::Unknown => None, + } + } +} + +impl Stream for ResponseSseStream { + type Item = Result; + + fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + loop { + if let Some(data) = self.next_event_data() { + if let Some(ev) = self.handle_event_data(&data) { + self.partial.apply_to(&ev); + return Poll::Ready(Some(Ok(ev))); + } + continue; + } + + 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) => { + if !self.saw_terminal { + self.saw_terminal = true; + match self.partial.clone().finalize() { + Ok(full) => { + return Poll::Ready(Some(Ok(StreamEvent::MessageComplete { + full_response: full, + }))); + } + Err(e) => { + return Poll::Ready(Some(Ok(StreamEvent::Error { + message: e.to_string(), + }))); + } + } + } + return Poll::Ready(None); + } + Poll::Pending => return Poll::Pending, + } + } + } +} + +// ============================================================================= +// Tests +// ============================================================================= + +#[cfg(test)] +mod tests { + use super::*; + use crate::llm::types::shared::ImageDetail; + use crate::llm::types::tool::ToolDef; + use serde_json::json; + + fn make_provider(base_url: String) -> OpenaiResponseProvider { + let client = Client::builder() + .timeout(Duration::from_secs(30)) + .build() + .expect("create http client"); + OpenaiResponseProvider::from_parts(base_url, "sk-test".into(), "gpt-4o".into(), client, 30) + } + + // ===== convert_request 单元测试 ===== + + #[test] + fn convert_request_text_only() { + let provider = make_provider("http://x".into()); + let req = MessageRequest { + model: "gpt-4o".into(), + messages: vec![Message::user_text("Hi")], + ..Default::default() + }; + let body = provider.convert_request(req).unwrap(); + assert_eq!(body.model, "gpt-4o"); + assert_eq!(body.input.len(), 1); + assert!(body.instructions.is_none()); + match &body.input[0] { + ResponseInputItem::Message { + item_type, + role, + content, + } => { + assert!(item_type.is_none()); + assert_eq!(role, "user"); + assert_eq!(content.len(), 1); + match &content[0] { + ResponseInputContent::InputText { text } => assert_eq!(text, "Hi"), + other => panic!("expected InputText, got {other:?}"), + } + } + other => panic!("expected Message, got {other:?}"), + } + } + + #[test] + fn convert_request_vision_user_image() { + let provider = make_provider("http://x".into()); + let req = MessageRequest { + model: "gpt-4o".into(), + messages: vec![Message::user_image( + "base64data", + "image/png", + ImageDetail::High, + )], + ..Default::default() + }; + let body = provider.convert_request(req).unwrap(); + assert_eq!(body.input.len(), 1); + match &body.input[0] { + ResponseInputItem::Message { role, content, .. } => { + assert_eq!(role, "user"); + assert_eq!(content.len(), 1); + match &content[0] { + ResponseInputContent::InputImage { image_url, detail } => { + assert_eq!(image_url, "data:image/png;base64,base64data"); + assert_eq!(detail.as_deref(), Some("high")); + } + other => panic!("expected InputImage, got {other:?}"), + } + } + _ => panic!("expected Message"), + } + } + + #[test] + fn convert_request_system_prompts_join_instructions() { + let provider = make_provider("http://x".into()); + let req = MessageRequest { + model: "gpt-4o".into(), + messages: vec![ + Message::system("You are helpful."), + Message::system("Be concise."), + Message::user_text("Hi"), + ], + ..Default::default() + }; + let body = provider.convert_request(req).unwrap(); + assert_eq!( + body.instructions.as_deref(), + Some("You are helpful.\nBe concise.") + ); + assert_eq!(body.input.len(), 1); // 只有 User + } + + #[test] + fn convert_request_assistant_text_uses_output_text() { + let provider = make_provider("http://x".into()); + let req = MessageRequest { + model: "gpt-4o".into(), + messages: vec![ + Message::user_text("Hi"), + Message::Assistant { + content: vec![ContentBlock::Text { + text: "Hello!".into(), + }], + }, + Message::user_text("Bye"), + ], + ..Default::default() + }; + let body = provider.convert_request(req).unwrap(); + // 应当产生 3 个 items: user + assistant(message)+ user + assert_eq!(body.input.len(), 3); + match &body.input[1] { + ResponseInputItem::Message { + item_type, + role, + content, + } => { + assert_eq!(item_type.as_deref(), Some("message")); + assert_eq!(role, "assistant"); + assert_eq!(content.len(), 1); + assert!( + matches!(&content[0], ResponseInputContent::OutputText { text } if text == "Hello!") + ); + } + _ => panic!("expected assistant message"), + } + } + + #[test] + fn convert_request_assistant_tool_use_produces_function_call() { + let provider = make_provider("http://x".into()); + let req = MessageRequest { + model: "gpt-4o".into(), + messages: vec![ + Message::Assistant { + content: vec![ContentBlock::ToolUse { + id: "call_1".into(), + name: "lookup".into(), + input: json!({"q": "rust"}), + }], + }, + Message::ToolResult { + tool_call_id: "call_1".into(), + content: vec![ContentBlock::Text { + text: "result".into(), + }], + is_error: false, + }, + ], + ..Default::default() + }; + let body = provider.convert_request(req).unwrap(); + assert_eq!(body.input.len(), 2); + match &body.input[0] { + ResponseInputItem::FunctionCall { + item_type, + call_id, + name, + arguments, + .. + } => { + assert_eq!(item_type, "function_call"); + assert_eq!(call_id, "call_1"); + assert_eq!(name, "lookup"); + assert!(arguments.contains("rust")); + } + _ => panic!("expected FunctionCall"), + } + match &body.input[1] { + ResponseInputItem::FunctionCallOutput { + item_type, + call_id, + output, + } => { + assert_eq!(item_type, "function_call_output"); + assert_eq!(call_id, "call_1"); + assert_eq!(output, "result"); + } + _ => panic!("expected FunctionCallOutput"), + } + } + + #[test] + fn convert_request_tools_and_extra() { + let provider = make_provider("http://x".into()); + let mut req = MessageRequest { + model: "gpt-4o".into(), + messages: vec![Message::user_text("Hi")], + tools: vec![ToolDef { + name: "search".into(), + description: Some("search docs".into()), + parameters: json!({"type": "object"}), + }], + ..Default::default() + }; + req.set_extra("previous_response_id", "resp_abc"); + req.set_extra("store", false); + req.set_extra("reasoning", json!({"effort": "low"})); + let body = provider.convert_request(req).unwrap(); + assert!(body.tools.is_some()); + assert_eq!(body.tools.as_ref().unwrap().len(), 1); + assert_eq!(body.previous_response_id.as_deref(), Some("resp_abc")); + assert_eq!(body.store, Some(false)); + assert_eq!(body.reasoning, Some(json!({"effort": "low"}))); + } + + #[test] + fn convert_request_tool_choice_named_serializes_as_object() { + let provider = make_provider("http://x".into()); + let req = MessageRequest { + model: "gpt-4o".into(), + messages: vec![Message::user_text("Hi")], + tool_choice: crate::llm::types::tool::ToolChoice::Named { + name: "search".into(), + }, + ..Default::default() + }; + let body = provider.convert_request(req).unwrap(); + assert_eq!( + body.tool_choice, + Some(json!({"type": "function", "function": {"name": "search"}})) + ); + } + + #[test] + fn convert_request_tool_choice_required_serializes_as_string() { + let provider = make_provider("http://x".into()); + let req = MessageRequest { + model: "gpt-4o".into(), + messages: vec![Message::user_text("Hi")], + tool_choice: crate::llm::types::tool::ToolChoice::Required, + ..Default::default() + }; + let body = provider.convert_request(req).unwrap(); + assert_eq!( + body.tool_choice, + Some(Value::String("required".to_string())) + ); + } + + #[test] + fn convert_request_text_format_serializes_as_text_field() { + let provider = make_provider("http://x".into()); + let mut req = MessageRequest { + model: "gpt-4o".into(), + messages: vec![Message::user_text("Hi")], + tool_choice: crate::llm::types::tool::ToolChoice::None, + ..Default::default() + }; + req.set_extra("text_format", json!({"type": "json_object"})); + let body = provider.convert_request(req).unwrap(); + assert_eq!(body.tool_choice, Some(Value::String("none".into()))); + assert_eq!(body.text, Some(json!({"format": {"type": "json_object"}}))); + } + + #[test] + fn convert_response_refusal_part_extracts_text() { + let provider = make_provider("http://x".into()); + let body = OpenaiResponseBody { + id: "resp_r".into(), + model: "gpt-4o".into(), + output: vec![ResponseOutputItem { + id: "msg_1".into(), + item_type: "message".into(), + status: Some("completed".into()), + role: Some("assistant".into()), + content: Some(vec![ResponseContentPart { + part_type: "refusal".into(), + text: Some("I cannot help with that".into()), + }]), + call_id: None, + name: None, + arguments: None, + }], + usage: Usage::default(), + status: "completed".into(), + }; + let resp = provider.convert_response(body).unwrap(); + let text = resp.text(); + assert!(text.contains("Refusal"), "got: {text}"); + assert!(text.contains("I cannot help with that"), "got: {text}"); + } + + #[tokio::test] + async fn response_api_stream_protocol_error_emits_stream_event_error() { + use futures_util::StreamExt; + use wiremock::matchers::{method, path}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + let sse = "event: error\n\ +data: {\"type\":\"error\",\"code\":\"server_error\",\"message\":\"upstream stream died\"}\n\n"; + + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/responses")) + .respond_with( + ResponseTemplate::new(200) + .insert_header("content-type", "text/event-stream") + .set_body_string(sse), + ) + .mount(&server) + .await; + + let provider = make_provider(server.uri()); + let mut stream = provider + .chat_stream(MessageRequest { + model: "gpt-4o".into(), + messages: vec![Message::user_text("hi")], + stream: true, + ..Default::default() + }) + .await + .unwrap(); + + let mut collected: Vec = Vec::new(); + while let Some(ev) = stream.next().await { + collected.push(ev.unwrap()); + } + + let has_error = collected + .iter() + .any(|e| matches!(e, StreamEvent::Error { .. })); + assert!( + has_error, + "expected Error event for protocol-level error SSE" + ); + } + + // ===== convert_response 单元测试 ===== + + #[test] + fn convert_response_message_extracts_text() { + let provider = make_provider("http://x".into()); + let body = OpenaiResponseBody { + id: "resp_1".into(), + model: "gpt-4o".into(), + output: vec![ResponseOutputItem { + id: "msg_1".into(), + item_type: "message".into(), + status: Some("completed".into()), + role: Some("assistant".into()), + content: Some(vec![ResponseContentPart { + part_type: "output_text".into(), + text: Some("Hello!".into()), + }]), + call_id: None, + name: None, + arguments: None, + }], + usage: Usage::from_input_output(5, 3), + status: "completed".into(), + }; + let resp = provider.convert_response(body).unwrap(); + assert_eq!(resp.id, "resp_1"); + assert_eq!(resp.text(), "Hello!"); + assert_eq!(resp.stop_reason, StopReason::Stop); + } + + #[test] + fn convert_response_function_call_extracts_tool_use() { + let provider = make_provider("http://x".into()); + let body = OpenaiResponseBody { + id: "resp_2".into(), + model: "gpt-4o".into(), + output: vec![ResponseOutputItem { + id: "fc_1".into(), + item_type: "function_call".into(), + status: Some("completed".into()), + role: None, + content: None, + call_id: Some("call_1".into()), + name: Some("lookup".into()), + arguments: Some(r#"{"q":"rust"}"#.into()), + }], + usage: Usage::default(), + status: "completed".into(), + }; + let resp = provider.convert_response(body).unwrap(); + match &resp.message { + Message::Assistant { content } => { + assert_eq!(content.len(), 1); + match &content[0] { + ContentBlock::ToolUse { id, name, input } => { + assert_eq!(id, "call_1"); + assert_eq!(name, "lookup"); + assert_eq!(input, &json!({"q": "rust"})); + } + _ => panic!("expected ToolUse"), + } + } + _ => panic!("expected Assistant"), + } + } + + #[test] + fn convert_response_empty_output_returns_err() { + let provider = make_provider("http://x".into()); + let body = OpenaiResponseBody { + id: "resp_3".into(), + model: "gpt-4o".into(), + output: vec![], + usage: Usage::default(), + status: "completed".into(), + }; + let err = provider.convert_response(body).unwrap_err(); + assert!(matches!(err, LlmError::Request { status: 200, .. })); + } + + #[test] + fn convert_response_unknown_item_type_falls_back_to_extension() { + let provider = make_provider("http://x".into()); + let body = OpenaiResponseBody { + id: "resp_4".into(), + model: "gpt-4o".into(), + output: vec![ResponseOutputItem { + id: "rs_1".into(), + item_type: "reasoning".into(), + status: None, + role: None, + content: None, + call_id: None, + name: None, + arguments: None, + }], + usage: Usage::default(), + status: "completed".into(), + }; + let resp = provider.convert_response(body).unwrap(); + match &resp.message { + Message::Assistant { content } => { + assert_eq!(content.len(), 1); + assert!( + matches!(&content[0], ContentBlock::Extension { kind, .. } if kind == "reasoning") + ); + } + _ => panic!("expected Assistant"), + } + } + + #[test] + fn convert_response_status_incomplete_maps_to_length() { + let provider = make_provider("http://x".into()); + let body = OpenaiResponseBody { + id: "resp_5".into(), + model: "gpt-4o".into(), + output: vec![ResponseOutputItem { + id: "msg_1".into(), + item_type: "message".into(), + status: Some("incomplete".into()), + role: Some("assistant".into()), + content: Some(vec![ResponseContentPart { + part_type: "output_text".into(), + text: Some("partial".into()), + }]), + call_id: None, + name: None, + arguments: None, + }], + usage: Usage::default(), + status: "incomplete".into(), + }; + let resp = provider.convert_response(body).unwrap(); + assert_eq!(resp.stop_reason, StopReason::Length); + } + + // ===== wiremock 集成测试 ===== + + #[tokio::test] + async fn response_api_basic_text() { + use wiremock::matchers::{header, method, path}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/responses")) + .and(header("authorization", "Bearer sk-test")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "id": "resp_basic", + "model": "gpt-4o", + "status": "completed", + "output": [{ + "id": "msg_1", + "type": "message", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": "Hello from mock!"}] + }], + "usage": {"prompt_tokens": 8, "completion_tokens": 4, "total_tokens": 12} + }))) + .mount(&server) + .await; + + let provider = make_provider(server.uri()); + let response = provider + .chat(MessageRequest { + model: "gpt-4o".into(), + messages: vec![Message::user_text("Hi")], + ..Default::default() + }) + .await + .unwrap(); + assert_eq!(response.text(), "Hello from mock!"); + assert_eq!(response.stop_reason, StopReason::Stop); + assert_eq!(response.usage.prompt_tokens, 8); + } + + #[tokio::test] + async fn response_api_unauthorized_maps_to_authentication() { + use wiremock::matchers::{method, path}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/responses")) + .respond_with(ResponseTemplate::new(401).set_body_string("invalid api key")) + .mount(&server) + .await; + + let provider = make_provider(server.uri()); + let err = provider + .chat(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 response_api_rate_limit_maps_to_rate_limit() { + use wiremock::matchers::{method, path}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/responses")) + .respond_with( + ResponseTemplate::new(429) + .insert_header("retry-after", "2") + .set_body_string("rate limited"), + ) + .mount(&server) + .await; + + let provider = make_provider(server.uri()); + let err = provider + .chat(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(2))); + } + other => panic!("expected RateLimit, got {other:?}"), + } + } + + #[tokio::test] + async fn response_api_server_error_maps_to_request() { + use wiremock::matchers::{method, path}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/responses")) + .respond_with(ResponseTemplate::new(500).set_body_string("server boom")) + .mount(&server) + .await; + + let provider = make_provider(server.uri()); + let err = provider + .chat(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_eq!(body, "server boom"); + } + other => panic!("expected Request, got {other:?}"), + } + } + + #[tokio::test] + async fn response_api_tool_call_response() { + use wiremock::matchers::{method, path}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/responses")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "id": "resp_tool", + "model": "gpt-4o", + "status": "completed", + "output": [ + { + "id": "msg_1", + "type": "message", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": "Let me check."}] + }, + { + "id": "fc_1", + "type": "function_call", + "call_id": "call_xyz", + "name": "lookup", + "arguments": "{\"q\":\"rust\"}", + "status": "completed" + } + ], + "usage": {"prompt_tokens": 5, "completion_tokens": 10, "total_tokens": 15} + }))) + .mount(&server) + .await; + + let provider = make_provider(server.uri()); + let response = provider + .chat(MessageRequest { + model: "gpt-4o".into(), + messages: vec![Message::user_text("Look up rust")], + ..Default::default() + }) + .await + .unwrap(); + assert_eq!(response.text(), "Let me check."); + 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_xyz".into(), "lookup".into()))); + } + + #[tokio::test] + async fn response_api_multi_turn_carries_history() { + use wiremock::matchers::{method, path}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/responses")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "id": "resp_2", + "model": "gpt-4o", + "status": "completed", + "output": [{ + "id": "msg_1", + "type": "message", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": "Got it"}] + }], + "usage": {"prompt_tokens": 20, "completion_tokens": 3, "total_tokens": 23} + }))) + .mount(&server) + .await; + + let provider = make_provider(server.uri()); + let response = provider + .chat(MessageRequest { + model: "gpt-4o".into(), + messages: vec![ + Message::user_text("knock knock."), + Message::Assistant { + content: vec![ContentBlock::Text { + text: "Who's there?".into(), + }], + }, + Message::user_text("Orange."), + ], + ..Default::default() + }) + .await + .unwrap(); + assert_eq!(response.text(), "Got it"); + } + + #[tokio::test] + async fn response_api_vision_input() { + use wiremock::matchers::{method, path}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/responses")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "id": "resp_v", + "model": "gpt-4o", + "status": "completed", + "output": [{ + "id": "msg_1", + "type": "message", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": "I see a cat"}] + }], + "usage": {"prompt_tokens": 100, "completion_tokens": 5, "total_tokens": 105} + }))) + .mount(&server) + .await; + + let provider = make_provider(server.uri()); + let response = provider + .chat(MessageRequest { + model: "gpt-4o".into(), + messages: vec![Message::user_image( + "b64data", + "image/png", + ImageDetail::High, + )], + ..Default::default() + }) + .await + .unwrap(); + assert_eq!(response.text(), "I see a cat"); + } + + #[tokio::test] + async fn response_api_stream_text() { + use futures_util::StreamExt; + use wiremock::matchers::{method, path}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + let sse = "event: response.created\n\ +data: {\"type\":\"response.created\",\"response\":{\"id\":\"resp_s\",\"model\":\"gpt-4o\",\"status\":\"in_progress\"}}\n\n\ +event: response.output_item.added\n\ +data: {\"type\":\"response.output_item.added\",\"output_index\":0,\"item\":{\"id\":\"msg_s\",\"type\":\"message\",\"role\":\"assistant\",\"status\":\"in_progress\",\"content\":[]}}\n\n\ +event: response.output_text.delta\n\ +data: {\"type\":\"response.output_text.delta\",\"item_id\":\"msg_s\",\"output_index\":0,\"delta\":\"Hello\"}\n\n\ +event: response.output_text.delta\n\ +data: {\"type\":\"response.output_text.delta\",\"item_id\":\"msg_s\",\"output_index\":0,\"delta\":\" world\"}\n\n\ +event: response.output_text.done\n\ +data: {\"type\":\"response.output_text.done\",\"item_id\":\"msg_s\",\"output_index\":0,\"text\":\"Hello world\"}\n\n\ +event: response.completed\n\ +data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_s\",\"model\":\"gpt-4o\",\"status\":\"completed\"}}\n\n"; + + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/responses")) + .respond_with( + ResponseTemplate::new(200) + .insert_header("content-type", "text/event-stream") + .set_body_string(sse), + ) + .mount(&server) + .await; + + let provider = make_provider(server.uri()); + let mut stream = provider + .chat_stream(MessageRequest { + model: "gpt-4o".into(), + messages: vec![Message::user_text("hi")], + stream: true, + ..Default::default() + }) + .await + .unwrap(); + + 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(), "Hello world"); + assert_eq!(complete.stop_reason, StopReason::Stop); + + let has_message_start = collected + .iter() + .any(|e| matches!(e, StreamEvent::MessageStart { .. })); + assert!(has_message_start); + } + + #[tokio::test] + async fn response_api_stream_tool_call() { + use futures_util::StreamExt; + use wiremock::matchers::{method, path}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + let sse = "event: response.created\n\ +data: {\"type\":\"response.created\",\"response\":{\"id\":\"resp_t\",\"model\":\"gpt-4o\",\"status\":\"in_progress\"}}\n\n\ +event: response.output_item.added\n\ +data: {\"type\":\"response.output_item.added\",\"output_index\":0,\"item\":{\"id\":\"fc_t\",\"type\":\"function_call\",\"call_id\":\"call_t\",\"name\":\"lookup\",\"arguments\":\"\"}}\n\n\ +event: response.function_call_arguments.delta\n\ +data: {\"type\":\"response.function_call_arguments.delta\",\"item_id\":\"fc_t\",\"output_index\":0,\"delta\":\"{\\\"q\\\":\\\"rust\\\"}\"}\n\n\ +event: response.function_call_arguments.done\n\ +data: {\"type\":\"response.function_call_arguments.done\",\"item_id\":\"fc_t\",\"output_index\":0,\"arguments\":\"{\\\"q\\\":\\\"rust\\\"}\"}\n\n\ +event: response.completed\n\ +data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_t\",\"model\":\"gpt-4o\",\"status\":\"completed\"}}\n\n"; + + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/responses")) + .respond_with( + ResponseTemplate::new(200) + .insert_header("content-type", "text/event-stream") + .set_body_string(sse), + ) + .mount(&server) + .await; + + let provider = make_provider(server.uri()); + let mut stream = provider + .chat_stream(MessageRequest { + model: "gpt-4o".into(), + messages: vec![Message::user_text("look up rust")], + stream: true, + ..Default::default() + }) + .await + .unwrap(); + + let mut collected: Vec = Vec::new(); + while let Some(ev) = stream.next().await { + collected.push(ev.unwrap()); + } + + let has_args_delta = collected + .iter() + .any(|e| matches!(e, StreamEvent::ToolCallArgumentsDelta { .. })); + assert!(has_args_delta, "expected ToolCallArgumentsDelta event"); + let has_tool_end = collected + .iter() + .any(|e| matches!(e, StreamEvent::ToolCallEnd { .. })); + assert!(has_tool_end, "expected ToolCallEnd event"); + } + + #[tokio::test] + async fn response_api_stream_multi_turn() { + use futures_util::StreamExt; + use wiremock::matchers::{method, path}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + let sse = "event: response.created\n\ +data: {\"type\":\"response.created\",\"response\":{\"id\":\"resp_mt\",\"model\":\"gpt-4o\",\"status\":\"in_progress\"}}\n\n\ +event: response.output_item.added\n\ +data: {\"type\":\"response.output_item.added\",\"output_index\":0,\"item\":{\"id\":\"msg_mt\",\"type\":\"message\",\"role\":\"assistant\",\"status\":\"in_progress\",\"content\":[]}}\n\n\ +event: response.output_text.delta\n\ +data: {\"type\":\"response.output_text.delta\",\"item_id\":\"msg_mt\",\"output_index\":0,\"delta\":\"Got it\"}\n\n\ +event: response.completed\n\ +data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_mt\",\"model\":\"gpt-4o\",\"status\":\"completed\"}}\n\n"; + + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/responses")) + .respond_with( + ResponseTemplate::new(200) + .insert_header("content-type", "text/event-stream") + .set_body_string(sse), + ) + .mount(&server) + .await; + + let provider = make_provider(server.uri()); + let mut stream = provider + .chat_stream(MessageRequest { + model: "gpt-4o".into(), + messages: vec![ + Message::user_text("first"), + Message::Assistant { + content: vec![ContentBlock::Text { + text: "ack1".into(), + }], + }, + Message::user_text("second"), + ], + stream: true, + ..Default::default() + }) + .await + .unwrap(); + + 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(), "Got it"); + } + + #[tokio::test] + async fn response_api_stream_error_event() { + use futures_util::StreamExt; + use wiremock::matchers::{method, path}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + let sse = "event: response.created\n\ +data: {\"type\":\"response.created\",\"response\":{\"id\":\"resp_e\",\"model\":\"gpt-4o\",\"status\":\"in_progress\"}}\n\n\ +event: response.failed\n\ +data: {\"type\":\"response.failed\",\"error\":{\"message\":\"server failed mid-stream\"}}\n\n"; + + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/responses")) + .respond_with( + ResponseTemplate::new(200) + .insert_header("content-type", "text/event-stream") + .set_body_string(sse), + ) + .mount(&server) + .await; + + let provider = make_provider(server.uri()); + let mut stream = provider + .chat_stream(MessageRequest { + model: "gpt-4o".into(), + messages: vec![Message::user_text("hi")], + stream: true, + ..Default::default() + }) + .await + .unwrap(); + + let mut collected: Vec = Vec::new(); + while let Some(ev) = stream.next().await { + collected.push(ev.unwrap()); + } + + let has_error = collected + .iter() + .any(|e| matches!(e, StreamEvent::Error { .. })); + assert!(has_error, "expected Error event in mid-stream failure"); + } + + #[test] + fn capabilities_reports_thinking_vision_tool_use() { + let provider = make_provider("http://x".into()); + let caps = provider.capabilities(); + assert_eq!(caps.provider_name, "openai-response"); + assert!(caps.features.thinking); + assert!(caps.features.vision); + assert!(caps.features.tool_use); + assert!(caps.features.streaming); + assert_eq!(caps.features.max_context_window, 200_000); + } +}