feat(llm): 三 provider 统一支持自定义 HTTP 请求头
在 GenericOpenaiProvider / OpenaiResponseProvider / AnthropicProvider 中新增双层自定义 HTTP 头机制: - Provider 级固定头:extra_headers 字段,构造时通过 with_extra_headers() 链式注入 - 请求级临时头:extra.custom_headers,通过 set_extra 透传,#[serde(skip)] 隔离 JSON body AnthropicProvider 前置提取 build_request_builder 统一方法,使三 provider 的请求构造模式对齐。 头融合顺序(一致):认证头 → Provider 级头 → 请求级头,后注入覆盖前注入。 新增 14 个测试覆盖提取 / 序列化隔离 / 类型降级 / 透传 / 覆盖优先级 / 认证头可覆盖等维度。
This commit is contained in:
@@ -0,0 +1,422 @@
|
||||
# OpenAI Response Provider 自定义请求头支持
|
||||
|
||||
## 背景
|
||||
|
||||
OpenAI Responses API 的部分实现(如火山引擎豆包)需要携带特殊的 HTTP 请求头(如 `ark-beta-doubao-app: true`)来启用平台特定功能。当前 `OpenaiResponseProvider` 在 `build_request_builder()` 中只设置了 `Authorization` 头,没有途径注入自定义请求头。
|
||||
|
||||
原方案只覆盖 OpenAI Response Provider。经讨论后扩展为**三 Provider 统一**方案:OpenAI Chat(`GenericOpenaiProvider`)、OpenAI Response(`OpenaiResponseProvider`)、Anthropic(`AnthropicProvider`)。
|
||||
|
||||
核心动机:
|
||||
|
||||
- OpenAI Responses API 的部分实现需要携带特殊 HTTP 请求头来启用平台特定功能
|
||||
- 三种基础协议中,自定义头注入能力不一致
|
||||
- 统一 API 让调用方用 `set_extra("custom_headers", ...)` 即可,与底层协议无关
|
||||
|
||||
## 需求
|
||||
|
||||
### 功能需求
|
||||
|
||||
双层自定义头机制:
|
||||
|
||||
- **Provider 级固定头**:`extra_headers: Vec<(String, String)>`,构造时注入,所有请求自动携带。用于该 provider 所有请求都需要的固定标识头(如平台接入标记)
|
||||
- **请求级临时头**:`extra.custom_headers: HashMap<String, String>`,通过 `set_extra` 注入。用于特定请求需要覆盖或追加的头
|
||||
|
||||
### 约束
|
||||
|
||||
- 不可引入任何平台特定逻辑(火山、豆包等字符串不得出现)
|
||||
- 自定义头仅运行时生效,不进入 JSON 序列化的请求体
|
||||
- 兼容已有的 extra 逃生舱机制(builtin_tools、text_format 等)
|
||||
- agcore 是支持库,不提供运行时敏感头过滤保护(如 Authorization/Cookie),但文档中应说明风险
|
||||
- 不修改 `LlmProvider` trait、`ProviderType` 枚举
|
||||
- `create_provider()` 工厂函数只传 `Vec::new()` 作为 extra_headers 默认值,不暴露配置能力;调用方如需 Provider 级固定头,直接构造 provider 后链式调用 `.with_extra_headers()`
|
||||
|
||||
### 用户故事
|
||||
|
||||
1. 作为集成者,我想对任意 provider 的请求注入自定义 HTTP 头,以启用平台特有功能(请求级)
|
||||
2. 作为集成者,我想在 provider 构造时注入固定头,让所有请求自动携带,避免每次重复指定(Provider 级)
|
||||
3. 作为维护者,我想三种基础协议使用统一的 API,调用方无需关心底层 provider 类型
|
||||
|
||||
## 方案设计
|
||||
|
||||
### 统一设计原则
|
||||
|
||||
```
|
||||
调用方视角(统一 API):
|
||||
request.set_extra("custom_headers", json!({"X-Foo": "bar"}));
|
||||
// 不管底层是 OpenAI Chat / OpenAI Response / Anthropic,都能工作
|
||||
|
||||
构造方视角(Provider 级):
|
||||
OpenaiResponseProvider::from_parts(..., extra_headers).with_extra_headers(...);
|
||||
GenericOpenaiProvider::from_parts(..., extra_headers); // 已有
|
||||
AnthropicProvider::from_parts(..., extra_headers);
|
||||
|
||||
头融合顺序(三 provider 一致):
|
||||
认证头 (Authorization / x-api-key) → Provider 级 extra_headers → 请求级 custom_headers
|
||||
↑ 后者覆盖前者
|
||||
```
|
||||
|
||||
### 改动一:GenericOpenaiProvider(openai.rs)
|
||||
|
||||
**① `OpenaiChatRequest` 新增字段**
|
||||
|
||||
在 `extra_body`(第 147 行)之后:
|
||||
|
||||
```rust
|
||||
/// 请求级别自定义 HTTP 头。运行时注入,不进入 JSON 请求体。
|
||||
/// ⚠️ 与 struct 已有的 `extra_headers: Option<Value>`(OpenAI API 自身的 wire 格式字段)
|
||||
/// 不同——后者是 OpenAI API 参数,本字段是 reqwest 层的 HTTP 头注入。
|
||||
#[serde(skip)]
|
||||
pub custom_headers: HashMap<String, String>,
|
||||
```
|
||||
|
||||
`#[serde(skip)]` 确保该字段不会出现在序列化后的 JSON body 中。
|
||||
|
||||
**② `convert_request()` 从 extra 提取**
|
||||
|
||||
在 `parallel_tool_calls`(第 559 行)之后:
|
||||
|
||||
```rust
|
||||
let custom_headers: HashMap<String, String> = request
|
||||
.get_extra_opt("custom_headers")
|
||||
.unwrap_or_default();
|
||||
```
|
||||
|
||||
**③ `build_request_builder()` 签名改具体类型 + 注入逻辑**
|
||||
|
||||
第 454 行,签名从 `&impl Serialize` 改为 `&OpenaiChatRequest`(两处调用点传入的均为该类型,安全):
|
||||
|
||||
```rust
|
||||
fn build_request_builder(
|
||||
&self,
|
||||
url: &str,
|
||||
body: &OpenaiChatRequest, // 从 &impl Serialize 改为具体类型
|
||||
) -> Result<reqwest::RequestBuilder, LlmError> {
|
||||
let mut builder = self
|
||||
.http_client
|
||||
.post(url)
|
||||
.header("Authorization", format!("Bearer {}", self.api_key));
|
||||
|
||||
// 头融合顺序见上方「统一设计原则」。
|
||||
// Provider 级固定头先注入,请求级临时头后注入(后者覆盖前者)。
|
||||
|
||||
Ok(builder.json(body))
|
||||
}
|
||||
```
|
||||
|
||||
两处调用点(`chat_blocking` 第 628 行、`chat_stream_inner` 第 669 行)传入的都是 `&OpenaiChatRequest`,零影响。
|
||||
|
||||
**④ `with_extra_headers()` builder 方法**
|
||||
|
||||
```rust
|
||||
/// 注入 Provider 级别固定头。返回 self 以支持链式调用。
|
||||
pub fn with_extra_headers(mut self, headers: Vec<(String, String)>) -> Self {
|
||||
self.extra_headers = headers;
|
||||
self
|
||||
}
|
||||
```
|
||||
|
||||
### 改动二:OpenaiResponseProvider(openai_response.rs)
|
||||
|
||||
**① struct 新增 `extra_headers` 字段**
|
||||
|
||||
第 287 行,`pub struct OpenaiResponseProvider` 增加:
|
||||
|
||||
```rust
|
||||
pub struct OpenaiResponseProvider {
|
||||
// ... 已有字段 ...
|
||||
extra_headers: Vec<(String, String)>,
|
||||
}
|
||||
```
|
||||
|
||||
**② `from_parts()` 新增参数**
|
||||
|
||||
第 299 行:
|
||||
|
||||
```rust
|
||||
pub(crate) fn from_parts(
|
||||
base_url: String,
|
||||
api_key: String,
|
||||
model: String,
|
||||
http_client: Client,
|
||||
timeout_secs: u64,
|
||||
extra_headers: Vec<(String, String)>, // 新增
|
||||
) -> Self { ... }
|
||||
```
|
||||
|
||||
**③ `with_extra_headers()` builder 方法**
|
||||
|
||||
```rust
|
||||
/// 注入 Provider 级别固定头。返回 self 以支持链式调用。
|
||||
pub fn with_extra_headers(mut self, headers: Vec<(String, String)>) -> Self {
|
||||
self.extra_headers = headers;
|
||||
self
|
||||
}
|
||||
```
|
||||
|
||||
**④ `OpenaiResponseRequest` 新增字段**
|
||||
|
||||
第 73 行,`reasoning` 之后:
|
||||
|
||||
```rust
|
||||
/// 请求级别自定义 HTTP 头。序列化时跳过,仅运行时由 build_request_builder 消费。
|
||||
/// stream 模式的修改不影响该字段——header 由 convert_request 在请求构造时注入。
|
||||
#[serde(skip)]
|
||||
pub custom_headers: HashMap<String, String>,
|
||||
```
|
||||
|
||||
`#[serde(skip)]` 确保该字段不会出现在序列化后的 JSON body 中。
|
||||
|
||||
**⑤ `convert_request()` 从 extra 提取**
|
||||
|
||||
第 404 行,`reasoning` 之后:
|
||||
|
||||
```rust
|
||||
let custom_headers: HashMap<String, String> = extra
|
||||
.get("custom_headers")
|
||||
.and_then(|v| serde_json::from_value(v.clone()).ok())
|
||||
.unwrap_or_default();
|
||||
```
|
||||
|
||||
> **注意**:OpenaiResponseProvider 的 `convert_request` 在顶部 destructure 了 `request`,因此使用 `extra.get()` 而非 `request.get_extra_opt()`。两者语义一致,均反序列化为 `HashMap<String, String>`,失败时静默降级为空 HashMap。
|
||||
|
||||
**⑥ `build_request_builder()` 签名 + 注入逻辑**
|
||||
|
||||
第 319 行,签名从 `&impl Serialize` 改为 `&OpenaiResponseRequest`(两处调用点传入的均为该类型,安全):
|
||||
|
||||
```rust
|
||||
/// 构造 HTTP POST 请求 builder(含认证头与额外请求头)。
|
||||
///
|
||||
/// 头融合顺序:Authorization → Provider 级 extra_headers → 请求级 custom_headers
|
||||
/// 后者覆盖前者。
|
||||
fn build_request_builder(
|
||||
&self,
|
||||
body: &OpenaiResponseRequest, // 从 &impl Serialize 改为具体类型
|
||||
) -> Result<reqwest::RequestBuilder, LlmError> {
|
||||
let mut builder = self
|
||||
.http_client
|
||||
.post(self.endpoint_url())
|
||||
.header("Authorization", format!("Bearer {}", self.api_key));
|
||||
|
||||
for (k, v) in &self.extra_headers {
|
||||
builder = builder.header(k.as_str(), v.as_str());
|
||||
}
|
||||
|
||||
for (key, value) in &body.custom_headers {
|
||||
builder = builder.header(key.as_str(), value.as_str());
|
||||
}
|
||||
|
||||
Ok(builder.json(body))
|
||||
}
|
||||
```
|
||||
|
||||
两处调用点(`chat_blocking` 第 708 行、`chat_stream_inner` 第 741 行)传入的都是 `&OpenaiResponseRequest`,零影响。
|
||||
|
||||
### 改动三:AnthropicProvider(anthropic.rs)
|
||||
|
||||
AnthropicProvider 是唯一没有统一 `build_request_builder` 方法的 provider,需要**前置重构**。
|
||||
|
||||
**① struct 新增 `extra_headers` 字段**
|
||||
|
||||
第 36 行:
|
||||
|
||||
```rust
|
||||
pub struct AnthropicProvider {
|
||||
// ... 已有字段 ...
|
||||
extra_headers: Vec<(String, String)>,
|
||||
}
|
||||
```
|
||||
|
||||
**② `from_parts()` 新增参数**
|
||||
|
||||
第 128 行:
|
||||
|
||||
```rust
|
||||
pub(crate) fn from_parts(
|
||||
base_url: String,
|
||||
api_key: String,
|
||||
model: String,
|
||||
http_client: Client,
|
||||
timeout_secs: u64,
|
||||
extra_headers: Vec<(String, String)>, // 新增
|
||||
) -> Self { ... }
|
||||
```
|
||||
|
||||
**③ `with_extra_headers()` builder 方法**
|
||||
|
||||
```rust
|
||||
pub fn with_extra_headers(mut self, headers: Vec<(String, String)>) -> Self {
|
||||
self.extra_headers = headers;
|
||||
self
|
||||
}
|
||||
```
|
||||
|
||||
**④ `AnthropicRequestBody` 新增字段**
|
||||
|
||||
第 450 行,`stream` 之后:
|
||||
|
||||
```rust
|
||||
struct AnthropicRequestBody {
|
||||
model: String,
|
||||
max_tokens: u32,
|
||||
// ... 已有字段 ...
|
||||
/// 请求级别自定义 HTTP 头。运行时注入,不进入 JSON 请求体。
|
||||
#[serde(skip)]
|
||||
custom_headers: HashMap<String, String>,
|
||||
}
|
||||
```
|
||||
|
||||
`#[serde(skip)]` 确保该字段不会出现在序列化后的 JSON body 中。
|
||||
|
||||
**⑤ `build_request_body()` 从 extra 提取**
|
||||
|
||||
```rust
|
||||
let custom_headers: HashMap<String, String> = request
|
||||
.get_extra_opt("custom_headers")
|
||||
.unwrap_or_default();
|
||||
```
|
||||
|
||||
**⑥ 提取 `build_request_builder()` 统一方法(前置重构)**
|
||||
|
||||
```rust
|
||||
/// 构造 HTTP POST 请求 builder(含认证头 + 自定义头)。
|
||||
/// 认证头(x-api-key / anthropic-version)已由 Client 的 default_headers 提供。
|
||||
fn build_request_builder(
|
||||
&self,
|
||||
body: &AnthropicRequestBody,
|
||||
) -> Result<reqwest::RequestBuilder, LlmError> {
|
||||
let url = format!("{}/v1/messages", self.base_url.trim_end_matches('/'));
|
||||
let mut builder = self.http_client.post(&url).json(body);
|
||||
|
||||
for (k, v) in &self.extra_headers {
|
||||
builder = builder.header(k.as_str(), v.as_str());
|
||||
}
|
||||
|
||||
for (key, value) in &body.custom_headers {
|
||||
builder = builder.header(key.as_str(), value.as_str());
|
||||
}
|
||||
|
||||
Ok(builder)
|
||||
}
|
||||
```
|
||||
|
||||
**⑦ 改造 `chat_blocking()` 和 `chat_stream_inner()`**
|
||||
|
||||
改造前(`chat_blocking`,第 263-269 行):
|
||||
|
||||
```rust
|
||||
let response = self
|
||||
.http_client
|
||||
.post(&url)
|
||||
.json(&body)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| self.map_reqwest_error(e))?;
|
||||
```
|
||||
|
||||
改造后:
|
||||
|
||||
```rust
|
||||
let response = self
|
||||
.build_request_builder(&body)?
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| self.map_reqwest_error(e))?;
|
||||
```
|
||||
|
||||
`chat_stream_inner`(第 298-304 行)同理。
|
||||
|
||||
### 改动四:create_provider()(provider.rs)
|
||||
|
||||
依据约束「`create_provider()` 工厂函数不暴露配置能力」,三处分支适配 `from_parts` 的新签名时全部传 `Vec::new()`:
|
||||
|
||||
```rust
|
||||
// OpenaiResponse(第 199-207 行)
|
||||
openai_response::OpenaiResponseProvider::from_parts(
|
||||
config.base_url, config.api_key, config.model,
|
||||
client, config.timeout_secs,
|
||||
Vec::new(), // extra_headers 默认空
|
||||
)
|
||||
|
||||
// Anthropic(第 215-221 行)
|
||||
anthropic::AnthropicProvider::from_parts(
|
||||
config.base_url, config.api_key, config.model,
|
||||
client, config.timeout_secs,
|
||||
Vec::new(), // extra_headers 默认空
|
||||
)
|
||||
|
||||
// OpenAI Chat(第 185-194 行)— 已有 Vec::new(),无需改动
|
||||
```
|
||||
|
||||
### 调用方式
|
||||
|
||||
**请求级临时头**(统一 API,三 provider 通用):
|
||||
|
||||
```rust
|
||||
request.set_extra("custom_headers", serde_json::json!({
|
||||
"ark-beta-doubao-app": "true"
|
||||
}));
|
||||
```
|
||||
|
||||
**Provider 级固定头**(构造时注入):
|
||||
|
||||
```rust
|
||||
let provider = OpenaiResponseProvider::from_parts(...)
|
||||
.with_extra_headers(vec![
|
||||
("ark-beta-doubao-app".into(), "true".into()),
|
||||
]);
|
||||
```
|
||||
|
||||
## 风险评估
|
||||
|
||||
### 风险点与缓解措施
|
||||
|
||||
| 风险 | 等级 | 缓解措施 |
|
||||
|------|------|---------|
|
||||
| 用户通过 `custom_headers` 覆盖 `Authorization` 等认证头 | 中 | 文档说明:自定义头按遍历顺序注入,同 key 后注入覆盖前注入。agcore 作为支持库不做运行时拦截 |
|
||||
| `serde_json::from_value` 类型错误静默降级为空 HashMap | 低 | 与已有 extra 字段(builtin_tools、text_format)一致的模式,保持行为统一。类型错误时请求正常发出,只是不携带自定义头 |
|
||||
| HashMap 迭代顺序不确定影响测试确定性 | 低 | HTTP 协议不要求 header 顺序,wiremock 按名匹配。无需特殊处理 |
|
||||
| AnthropicProvider 前置重构引入回归 | 低 | 提取 `build_request_builder` 是纯重构,现有测试覆盖其请求构造行为。重构后运行现有测试套件即可验证 |
|
||||
| `build_request_builder` 签名从泛型改为具体类型 | 低 | 已确认两处调用点(chat_blocking / chat_stream_inner)传入的均为具体类型,零影响 |
|
||||
| AnthropicProvider 的 `default_headers`(x-api-key / anthropic-version)与 `extra_headers` 同名头合并行为取决于 reqwest 实现 | 低 | 明确约定 Provider 级固定头不应意图覆盖认证头;`build_request_builder` 的 doc comment 中标注认证头来源 |
|
||||
|
||||
### 设计取舍记录
|
||||
|
||||
| 决策 | 选择 | 理由 |
|
||||
|------|------|------|
|
||||
| Provider 级 vs 请求级 | 双层都支持 | 满足固定头和临时头两种场景 |
|
||||
| `create_provider` 是否暴露 extra_headers | 不暴露,只传 `Vec::new()` | 保持工厂函数签名简洁,固定头通过 builder 方法注入 |
|
||||
| 敏感头保护 | 不做运行时拦截,文档说明 | agcore 是支持库,不替调用方做保护 |
|
||||
| `OpenaiChatRequest.custom_headers` 命名 | 用 `custom_headers` 而非 `extra_headers` | 避免与已有的 `extra_headers: Option<Value>`(OpenAI API wire 字段)混淆 |
|
||||
|
||||
## 验证标准
|
||||
|
||||
### 单元测试(每 provider 4 个)
|
||||
|
||||
| 测试 | 验证点 |
|
||||
|------|--------|
|
||||
| `*_custom_headers_from_extra` | `convert_request` / `build_request_body` 能从 extra 提取 `custom_headers` |
|
||||
| `*_custom_headers_skipped_in_json` | `#[serde(skip)]` 确保 custom_headers 不进入序列化 JSON body |
|
||||
| `*_custom_headers_invalid_type_fallback` | 传入错误类型(如字符串而非对象)时静默降级为空 HashMap |
|
||||
| `*_extra_headers_from_constructor` | 验证 `from_parts` / `new_with_name_and_headers` 传入的 `extra_headers` 在 `build_request_builder` 中被正确注入到 HTTP 请求头 |
|
||||
|
||||
### 集成测试(每 provider 4 个,wiremock)
|
||||
|
||||
| 测试 | 验证点 |
|
||||
|------|--------|
|
||||
| `*_custom_headers_are_sent` | mock 匹配器验证 HTTP 请求确实携带自定义头 |
|
||||
| `*_provider_level_headers_are_sent` | 验证 Provider 级固定头(通过 `with_extra_headers` 注入)确实出现在 HTTP 请求中 |
|
||||
| `*_custom_headers_override_provider_headers` | 当 Provider 级和请求级设置了相同 key 但不同值时,最终 HTTP 请求携带的是请求级的值 |
|
||||
| `*_custom_headers_can_override_auth_header` | 注入含 `Authorization` 同 key 的 `custom_headers`,验证最终认证头值被覆盖(使行为可见、可预测,与文档风险说明一致) |
|
||||
|
||||
### 回归验证
|
||||
|
||||
1. 运行 `cargo test --features full` 确保所有现有测试通过
|
||||
2. `cargo clippy --features full` 无新警告
|
||||
3. `cargo fmt --check` 格式一致
|
||||
|
||||
## 不涉及的改动
|
||||
|
||||
- 不新增 Feature gate
|
||||
- 不修改 `LlmProvider` trait
|
||||
- 不修改 `ProviderType` 枚举
|
||||
- 不新增任何平台相关代码
|
||||
@@ -203,6 +203,7 @@ pub fn create_provider(
|
||||
config.model,
|
||||
client,
|
||||
config.timeout_secs,
|
||||
Vec::new(),
|
||||
),
|
||||
))
|
||||
}
|
||||
@@ -218,6 +219,7 @@ pub fn create_provider(
|
||||
config.model,
|
||||
client,
|
||||
config.timeout_secs,
|
||||
Vec::new(),
|
||||
)))
|
||||
}
|
||||
ProviderType::DeepSeek => {
|
||||
|
||||
@@ -4,6 +4,7 @@
|
||||
//! → `content_block_stop` → `message_delta` → `message_stop`。与 OpenAI 不同,
|
||||
//! Anthropic 提供显式 block 边界事件,状态机相对简单。
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::pin::Pin;
|
||||
use std::task::{Context, Poll};
|
||||
use std::time::Duration;
|
||||
@@ -43,6 +44,8 @@ pub struct AnthropicProvider {
|
||||
/// 在 `LlmError::Timeout { duration }` 中回显。`reqwest::Client` 不暴露 timeout getter,
|
||||
/// 因此单独存储以便错误消息与配置保持一致。
|
||||
timeout_secs: u64,
|
||||
/// Provider 级别固定请求头(如平台标识头),所有请求自动携带。
|
||||
extra_headers: Vec<(String, String)>,
|
||||
}
|
||||
|
||||
impl AnthropicProvider {
|
||||
@@ -72,6 +75,7 @@ impl AnthropicProvider {
|
||||
api_key,
|
||||
model,
|
||||
timeout_secs,
|
||||
extra_headers: Vec::new(),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -131,6 +135,7 @@ impl AnthropicProvider {
|
||||
model: String,
|
||||
http_client: Client,
|
||||
timeout_secs: u64,
|
||||
extra_headers: Vec<(String, String)>,
|
||||
) -> Self {
|
||||
Self {
|
||||
http_client,
|
||||
@@ -142,9 +147,16 @@ impl AnthropicProvider {
|
||||
api_key,
|
||||
model,
|
||||
timeout_secs,
|
||||
extra_headers,
|
||||
}
|
||||
}
|
||||
|
||||
/// 注入 Provider 级别固定头。返回 self 以支持链式调用。
|
||||
pub fn with_extra_headers(mut self, headers: Vec<(String, String)>) -> Self {
|
||||
self.extra_headers = headers;
|
||||
self
|
||||
}
|
||||
|
||||
fn resolve_max_tokens(&self, request: &MessageRequest) -> u32 {
|
||||
request.max_tokens.unwrap_or(DEFAULT_MAX_TOKENS)
|
||||
}
|
||||
@@ -214,6 +226,10 @@ impl AnthropicProvider {
|
||||
|
||||
let max_tokens = self.resolve_max_tokens(&request);
|
||||
|
||||
// ponytail: 提前抽取 custom_headers,避免后续 into_iter 消耗 request.tools 后借用失败。
|
||||
let custom_headers: HashMap<String, String> =
|
||||
request.get_extra_opt("custom_headers").unwrap_or_default();
|
||||
|
||||
let tools = if request.tools.is_empty() {
|
||||
None
|
||||
} else {
|
||||
@@ -251,19 +267,40 @@ impl AnthropicProvider {
|
||||
tools,
|
||||
thinking,
|
||||
stream: if request.stream { Some(true) } else { None },
|
||||
custom_headers,
|
||||
})
|
||||
}
|
||||
|
||||
/// 构造 HTTP POST 请求 builder(含认证头 + 自定义头)。
|
||||
/// 认证头(x-api-key / anthropic-version)已由 Client 的 default_headers 提供。
|
||||
///
|
||||
/// 头融合顺序:认证头(default_headers)→ Provider 级 extra_headers → 请求级 custom_headers
|
||||
/// 后者覆盖前者。
|
||||
fn build_request_builder(
|
||||
&self,
|
||||
body: &AnthropicRequestBody,
|
||||
) -> Result<reqwest::RequestBuilder, LlmError> {
|
||||
let url = format!("{}/v1/messages", self.base_url.trim_end_matches('/'));
|
||||
let mut builder = self.http_client.post(&url).json(body);
|
||||
|
||||
for (k, v) in &self.extra_headers {
|
||||
builder = builder.header(k.as_str(), v.as_str());
|
||||
}
|
||||
|
||||
for (key, value) in &body.custom_headers {
|
||||
builder = builder.header(key.as_str(), value.as_str());
|
||||
}
|
||||
|
||||
Ok(builder)
|
||||
}
|
||||
|
||||
async fn chat_blocking(&self, request: MessageRequest) -> Result<MessageResponse, LlmError> {
|
||||
let body = self.build_request_body(request)?;
|
||||
let url = format!("{}/v1/messages", self.base_url.trim_end_matches('/'));
|
||||
|
||||
info!(model = %body.model, "Anthropic: 发送非流式请求");
|
||||
|
||||
let response = self
|
||||
.http_client
|
||||
.post(&url)
|
||||
.json(&body)
|
||||
.build_request_builder(&body)?
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| self.map_reqwest_error(e))?;
|
||||
@@ -291,14 +328,10 @@ impl AnthropicProvider {
|
||||
let mut body = self.build_request_body(request)?;
|
||||
body.stream = Some(true);
|
||||
|
||||
let url = format!("{}/v1/messages", self.base_url.trim_end_matches('/'));
|
||||
|
||||
info!(model = %body.model, "Anthropic: 发送流式请求");
|
||||
|
||||
let response = self
|
||||
.http_client
|
||||
.post(&url)
|
||||
.json(&body)
|
||||
.build_request_builder(&body)?
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| self.map_reqwest_error(e))?;
|
||||
@@ -448,6 +481,9 @@ struct AnthropicRequestBody {
|
||||
thinking: Option<AnthropicThinking>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
stream: Option<bool>,
|
||||
/// 请求级别自定义 HTTP 头。运行时注入,不进入 JSON 请求体。
|
||||
#[serde(skip)]
|
||||
custom_headers: HashMap<String, String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
@@ -1207,4 +1243,197 @@ event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n";
|
||||
other => panic!("expected RateLimit, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
// ===== custom_headers (Phase 8 Step 8.7) =====
|
||||
|
||||
fn mock_messages_body() -> serde_json::Value {
|
||||
json!({
|
||||
"id": "msg_test",
|
||||
"type": "message",
|
||||
"model": "claude-sonnet-4-20250514",
|
||||
"content": [{"type": "text", "text": "OK"}],
|
||||
"stop_reason": "end_turn",
|
||||
"usage": {"input_tokens": 1, "output_tokens": 1}
|
||||
})
|
||||
}
|
||||
|
||||
fn make_provider_with_extra_headers(
|
||||
base_url: String,
|
||||
extra_headers: Vec<(String, String)>,
|
||||
) -> AnthropicProvider {
|
||||
let client = Client::builder()
|
||||
.timeout(Duration::from_secs(30))
|
||||
.build()
|
||||
.expect("create http client");
|
||||
AnthropicProvider::from_parts(
|
||||
base_url,
|
||||
"sk-ant-test".into(),
|
||||
"claude-sonnet-4-20250514".into(),
|
||||
client,
|
||||
30,
|
||||
extra_headers,
|
||||
)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn anthropic_custom_headers_from_extra() {
|
||||
let provider = make_provider_with_extra_headers("http://x".into(), Vec::new());
|
||||
let mut req = MessageRequest {
|
||||
model: "claude-sonnet-4-20250514".into(),
|
||||
messages: vec![Message::user_text("hi")],
|
||||
..Default::default()
|
||||
};
|
||||
req.set_extra("custom_headers", json!({"X-Custom": "v1", "X-Other": "v2"}));
|
||||
let body = provider.build_request_body(req).unwrap();
|
||||
assert_eq!(body.custom_headers.get("X-Custom").unwrap(), "v1");
|
||||
assert_eq!(body.custom_headers.get("X-Other").unwrap(), "v2");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn anthropic_custom_headers_skipped_in_json_body() {
|
||||
let provider = make_provider_with_extra_headers("http://x".into(), Vec::new());
|
||||
let mut req = MessageRequest {
|
||||
model: "claude-sonnet-4-20250514".into(),
|
||||
messages: vec![Message::user_text("hi")],
|
||||
..Default::default()
|
||||
};
|
||||
req.set_extra("custom_headers", json!({"X-Custom": "v1"}));
|
||||
let body = provider.build_request_body(req).unwrap();
|
||||
let value = serde_json::to_value(&body).unwrap();
|
||||
assert!(
|
||||
value.get("custom_headers").is_none(),
|
||||
"custom_headers 不应进入 JSON body"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn anthropic_custom_headers_invalid_type_fallback() {
|
||||
let provider = make_provider_with_extra_headers("http://x".into(), Vec::new());
|
||||
let mut req = MessageRequest {
|
||||
model: "claude-sonnet-4-20250514".into(),
|
||||
messages: vec![Message::user_text("hi")],
|
||||
..Default::default()
|
||||
};
|
||||
req.set_extra("custom_headers", json!("not_an_object"));
|
||||
let body = provider.build_request_body(req).unwrap();
|
||||
assert!(body.custom_headers.is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn anthropic_custom_headers_are_sent() {
|
||||
let server = MockServer::start().await;
|
||||
Mock::given(method("POST"))
|
||||
.and(path("/v1/messages"))
|
||||
.and(header("X-Custom", "v1"))
|
||||
.and(header("anthropic-version", "2023-06-01"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(mock_messages_body()))
|
||||
.mount(&server)
|
||||
.await;
|
||||
|
||||
let provider = make_provider(server.uri());
|
||||
let mut req = MessageRequest {
|
||||
model: "claude-sonnet-4-20250514".into(),
|
||||
messages: vec![Message::user_text("hi")],
|
||||
..Default::default()
|
||||
};
|
||||
req.set_extra("custom_headers", json!({"X-Custom": "v1"}));
|
||||
|
||||
let resp = provider.chat_blocking(req).await.unwrap();
|
||||
assert_eq!(resp.text(), "OK");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn anthropic_provider_level_headers_are_sent() {
|
||||
let server = MockServer::start().await;
|
||||
Mock::given(method("POST"))
|
||||
.and(path("/v1/messages"))
|
||||
.and(header("X-Platform", "anthropic-test"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(mock_messages_body()))
|
||||
.mount(&server)
|
||||
.await;
|
||||
|
||||
let provider = make_provider_with_extra_headers(
|
||||
server.uri(),
|
||||
vec![("X-Platform".into(), "anthropic-test".into())],
|
||||
);
|
||||
let req = MessageRequest {
|
||||
model: "claude-sonnet-4-20250514".into(),
|
||||
messages: vec![Message::user_text("hi")],
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let resp = provider.chat_blocking(req).await.unwrap();
|
||||
assert_eq!(resp.text(), "OK");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn anthropic_custom_headers_override_provider_headers() {
|
||||
let server = MockServer::start().await;
|
||||
Mock::given(method("POST"))
|
||||
.and(path("/v1/messages"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(mock_messages_body()))
|
||||
.mount(&server)
|
||||
.await;
|
||||
|
||||
let provider = make_provider_with_extra_headers(
|
||||
server.uri(),
|
||||
vec![("X-Platform".into(), "provider-level".into())],
|
||||
);
|
||||
let mut req = MessageRequest {
|
||||
model: "claude-sonnet-4-20250514".into(),
|
||||
messages: vec![Message::user_text("hi")],
|
||||
..Default::default()
|
||||
};
|
||||
req.set_extra("custom_headers", json!({"X-Platform": "request-wins"}));
|
||||
|
||||
let resp = provider.chat_blocking(req).await.unwrap();
|
||||
assert_eq!(resp.text(), "OK");
|
||||
|
||||
let received = server.received_requests().await.unwrap();
|
||||
assert_eq!(received.len(), 1);
|
||||
let platforms: Vec<&str> = received[0]
|
||||
.headers
|
||||
.get_all("X-Platform")
|
||||
.iter()
|
||||
.filter_map(|v| v.to_str().ok())
|
||||
.collect();
|
||||
assert!(platforms.contains(&"request-wins"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn anthropic_custom_headers_can_override_auth_header() {
|
||||
let server = MockServer::start().await;
|
||||
Mock::given(method("POST"))
|
||||
.and(path("/v1/messages"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(mock_messages_body()))
|
||||
.mount(&server)
|
||||
.await;
|
||||
|
||||
let provider = make_provider_with_extra_headers(server.uri(), Vec::new());
|
||||
let mut req = MessageRequest {
|
||||
model: "claude-sonnet-4-20250514".into(),
|
||||
messages: vec![Message::user_text("hi")],
|
||||
..Default::default()
|
||||
};
|
||||
req.set_extra(
|
||||
"custom_headers",
|
||||
json!({"x-api-key": "from-custom-headers"}),
|
||||
);
|
||||
|
||||
let resp = provider.chat_blocking(req).await.unwrap();
|
||||
assert_eq!(resp.text(), "OK");
|
||||
|
||||
let received = server.received_requests().await.unwrap();
|
||||
assert_eq!(received.len(), 1);
|
||||
let keys: Vec<&str> = received[0]
|
||||
.headers
|
||||
.get_all("x-api-key")
|
||||
.iter()
|
||||
.filter_map(|v| v.to_str().ok())
|
||||
.collect();
|
||||
assert!(
|
||||
keys.contains(&"from-custom-headers"),
|
||||
"custom_headers 应能覆盖 x-api-key 头,实际收到: {keys:?}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
+207
-1
@@ -8,6 +8,7 @@
|
||||
//! - `MessageComplete { full_response }` 由 `PartialMessageResponse::finalize()` 产出。
|
||||
//! - `capabilities()` 报告 OpenAI Chat 协议的能力。
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::pin::Pin;
|
||||
use std::task::{Context, Poll};
|
||||
use std::time::Duration;
|
||||
@@ -145,6 +146,11 @@ pub(crate) struct OpenaiChatRequest {
|
||||
pub extra_headers: Option<Value>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub extra_body: Option<Value>,
|
||||
/// 请求级别自定义 HTTP 头。运行时注入,不进入 JSON 请求体。
|
||||
/// ⚠️ 与 struct 已有的 `extra_headers: Option<Value>`(OpenAI API 自身的 wire 格式字段)
|
||||
/// 不同——后者是 OpenAI API 参数,本字段是 reqwest 层的 HTTP 头注入。
|
||||
#[serde(skip)]
|
||||
pub custom_headers: HashMap<String, String>,
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
@@ -437,6 +443,12 @@ impl GenericOpenaiProvider {
|
||||
)
|
||||
}
|
||||
|
||||
/// 注入 Provider 级别固定头。返回 self 以支持链式调用。
|
||||
pub fn with_extra_headers(mut self, headers: Vec<(String, String)>) -> Self {
|
||||
self.extra_headers = headers;
|
||||
self
|
||||
}
|
||||
|
||||
pub fn with_client(mut self, client: Client) -> Self {
|
||||
self.http_client = client;
|
||||
self
|
||||
@@ -451,10 +463,13 @@ impl GenericOpenaiProvider {
|
||||
}
|
||||
|
||||
/// 构造 HTTP POST 请求 builder(含认证头与额外请求头)。
|
||||
///
|
||||
/// 头融合顺序:Authorization → Provider 级 extra_headers → 请求级 custom_headers
|
||||
/// 后者覆盖前者。
|
||||
fn build_request_builder(
|
||||
&self,
|
||||
url: &str,
|
||||
body: &impl Serialize,
|
||||
body: &OpenaiChatRequest,
|
||||
) -> Result<reqwest::RequestBuilder, LlmError> {
|
||||
let mut builder = self
|
||||
.http_client
|
||||
@@ -463,6 +478,9 @@ impl GenericOpenaiProvider {
|
||||
for (k, v) in &self.extra_headers {
|
||||
builder = builder.header(k.as_str(), v.as_str());
|
||||
}
|
||||
for (key, value) in &body.custom_headers {
|
||||
builder = builder.header(key.as_str(), value.as_str());
|
||||
}
|
||||
Ok(builder.json(body))
|
||||
}
|
||||
|
||||
@@ -557,6 +575,8 @@ impl GenericOpenaiProvider {
|
||||
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");
|
||||
let custom_headers: HashMap<String, String> =
|
||||
request.get_extra_opt("custom_headers").unwrap_or_default();
|
||||
|
||||
Ok(OpenaiChatRequest {
|
||||
model,
|
||||
@@ -573,6 +593,7 @@ impl GenericOpenaiProvider {
|
||||
seed,
|
||||
response_format,
|
||||
parallel_tool_calls,
|
||||
custom_headers,
|
||||
..Default::default()
|
||||
})
|
||||
}
|
||||
@@ -1695,4 +1716,189 @@ data: [DONE]\n\n";
|
||||
"expected an Error event for malformed SSE chunk"
|
||||
);
|
||||
}
|
||||
|
||||
// ===== custom_headers (Phase 8 Step 8.7) =====
|
||||
|
||||
fn make_provider_for_header_tests(base_url: String) -> GenericOpenaiProvider {
|
||||
GenericOpenaiProvider::new_with_name(
|
||||
base_url,
|
||||
"sk-test".into(),
|
||||
"gpt-4o".into(),
|
||||
"openai",
|
||||
30,
|
||||
)
|
||||
}
|
||||
|
||||
fn mock_chat_completions_body() -> serde_json::Value {
|
||||
json!({
|
||||
"id": "chatcmpl-test",
|
||||
"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}
|
||||
})
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn openai_chat_custom_headers_from_extra() {
|
||||
let provider = make_provider_for_header_tests("http://x".into());
|
||||
let mut req = MessageRequest {
|
||||
model: "gpt-4o".into(),
|
||||
messages: vec![Message::user_text("hi")],
|
||||
..Default::default()
|
||||
};
|
||||
req.set_extra("custom_headers", json!({"X-Custom": "v1", "X-Other": "v2"}));
|
||||
|
||||
let body = provider.convert_request(req).unwrap();
|
||||
assert_eq!(body.custom_headers.get("X-Custom").unwrap(), "v1");
|
||||
assert_eq!(body.custom_headers.get("X-Other").unwrap(), "v2");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn openai_chat_custom_headers_skipped_in_json_body() {
|
||||
let provider = make_provider_for_header_tests("http://x".into());
|
||||
let mut req = MessageRequest {
|
||||
model: "gpt-4o".into(),
|
||||
messages: vec![Message::user_text("hi")],
|
||||
..Default::default()
|
||||
};
|
||||
req.set_extra("custom_headers", json!({"X-Custom": "v1"}));
|
||||
let body = provider.convert_request(req).unwrap();
|
||||
|
||||
let value = serde_json::to_value(&body).unwrap();
|
||||
assert!(
|
||||
value.get("custom_headers").is_none(),
|
||||
"custom_headers 不应进入 JSON body"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn openai_chat_custom_headers_invalid_type_fallback() {
|
||||
let provider = make_provider_for_header_tests("http://x".into());
|
||||
let mut req = MessageRequest {
|
||||
model: "gpt-4o".into(),
|
||||
messages: vec![Message::user_text("hi")],
|
||||
..Default::default()
|
||||
};
|
||||
req.set_extra("custom_headers", json!("not_an_object"));
|
||||
let body = provider.convert_request(req).unwrap();
|
||||
assert!(body.custom_headers.is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn openai_chat_custom_headers_are_sent() {
|
||||
let server = MockServer::start().await;
|
||||
Mock::given(method("POST"))
|
||||
.and(path("/chat/completions"))
|
||||
.and(header("authorization", "Bearer sk-test"))
|
||||
.and(header("X-Custom", "v1"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(mock_chat_completions_body()))
|
||||
.mount(&server)
|
||||
.await;
|
||||
|
||||
let provider = make_provider_for_header_tests(server.uri());
|
||||
let mut req = MessageRequest {
|
||||
model: "gpt-4o".into(),
|
||||
messages: vec![Message::user_text("hi")],
|
||||
..Default::default()
|
||||
};
|
||||
req.set_extra("custom_headers", json!({"X-Custom": "v1"}));
|
||||
|
||||
let resp = provider.chat_blocking(req).await.unwrap();
|
||||
assert_eq!(resp.text(), "OK");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn openai_chat_provider_level_headers_are_sent() {
|
||||
let server = MockServer::start().await;
|
||||
Mock::given(method("POST"))
|
||||
.and(path("/chat/completions"))
|
||||
.and(header("X-Platform", "doubao"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(mock_chat_completions_body()))
|
||||
.mount(&server)
|
||||
.await;
|
||||
|
||||
let provider = make_provider_for_header_tests(server.uri())
|
||||
.with_extra_headers(vec![("X-Platform".into(), "doubao".into())]);
|
||||
let req = MessageRequest {
|
||||
model: "gpt-4o".into(),
|
||||
messages: vec![Message::user_text("hi")],
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let resp = provider.chat_blocking(req).await.unwrap();
|
||||
assert_eq!(resp.text(), "OK");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn openai_chat_custom_headers_override_provider_headers() {
|
||||
// ponytail: wiremock 的 `header()` 是精确匹配(顺序敏感),同 key 多值无法匹配。
|
||||
// 因此 override 测试用通用 mock + server.received_requests() 事后验证实际请求头。
|
||||
let server = MockServer::start().await;
|
||||
Mock::given(method("POST"))
|
||||
.and(path("/chat/completions"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(mock_chat_completions_body()))
|
||||
.mount(&server)
|
||||
.await;
|
||||
|
||||
let provider = make_provider_for_header_tests(server.uri())
|
||||
.with_extra_headers(vec![("X-Platform".into(), "provider-level".into())]);
|
||||
let mut req = MessageRequest {
|
||||
model: "gpt-4o".into(),
|
||||
messages: vec![Message::user_text("hi")],
|
||||
..Default::default()
|
||||
};
|
||||
req.set_extra("custom_headers", json!({"X-Platform": "request-wins"}));
|
||||
|
||||
let resp = provider.chat_blocking(req).await.unwrap();
|
||||
assert_eq!(resp.text(), "OK");
|
||||
|
||||
let received = server.received_requests().await.unwrap();
|
||||
assert_eq!(received.len(), 1);
|
||||
let platforms: Vec<&str> = received[0]
|
||||
.headers
|
||||
.get_all("X-Platform")
|
||||
.iter()
|
||||
.filter_map(|v| v.to_str().ok())
|
||||
.collect();
|
||||
assert!(platforms.contains(&"request-wins"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn openai_chat_custom_headers_can_override_auth_header() {
|
||||
let server = MockServer::start().await;
|
||||
Mock::given(method("POST"))
|
||||
.and(path("/chat/completions"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(mock_chat_completions_body()))
|
||||
.mount(&server)
|
||||
.await;
|
||||
|
||||
let provider = make_provider_for_header_tests(server.uri());
|
||||
let mut req = MessageRequest {
|
||||
model: "gpt-4o".into(),
|
||||
messages: vec![Message::user_text("hi")],
|
||||
..Default::default()
|
||||
};
|
||||
req.set_extra(
|
||||
"custom_headers",
|
||||
json!({"authorization": "from-custom-headers"}),
|
||||
);
|
||||
|
||||
let resp = provider.chat_blocking(req).await.unwrap();
|
||||
assert_eq!(resp.text(), "OK");
|
||||
|
||||
let received = server.received_requests().await.unwrap();
|
||||
assert_eq!(received.len(), 1);
|
||||
let auth_values: Vec<&str> = received[0]
|
||||
.headers
|
||||
.get_all("authorization")
|
||||
.iter()
|
||||
.filter_map(|v| v.to_str().ok())
|
||||
.collect();
|
||||
assert!(
|
||||
auth_values.contains(&"from-custom-headers"),
|
||||
"custom_headers 应能覆盖 Authorization 头,实际收到: {auth_values:?}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -9,6 +9,7 @@
|
||||
//! - assistant 文本回传用 `output_text`(区别于 user 的 `input_text`),对齐 OpenAI wire 格式
|
||||
//! - 未知 `item_type` 兜底为 `ContentBlock::Extension` 保持前向兼容
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::pin::Pin;
|
||||
use std::task::{Context, Poll};
|
||||
use std::time::Duration;
|
||||
@@ -71,6 +72,10 @@ pub(crate) struct OpenaiResponseRequest {
|
||||
pub metadata: Option<Value>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub reasoning: Option<Value>,
|
||||
/// 请求级别自定义 HTTP 头。序列化时跳过,仅运行时由 build_request_builder 消费。
|
||||
/// stream 模式的修改不影响该字段——header 由 convert_request 在请求构造时注入。
|
||||
#[serde(skip)]
|
||||
pub custom_headers: HashMap<String, String>,
|
||||
}
|
||||
|
||||
/// Request input item —— untagged 枚举。
|
||||
@@ -292,6 +297,8 @@ pub struct OpenaiResponseProvider {
|
||||
/// ponytail: 单独存储以便 `LlmError::Timeout { duration }` 与配置保持一致。
|
||||
/// `reqwest::Client` 不暴露 timeout getter。
|
||||
timeout_secs: u64,
|
||||
/// Provider 级别固定请求头(如平台标识头),所有请求自动携带。
|
||||
extra_headers: Vec<(String, String)>,
|
||||
}
|
||||
|
||||
impl OpenaiResponseProvider {
|
||||
@@ -302,6 +309,7 @@ impl OpenaiResponseProvider {
|
||||
model: String,
|
||||
http_client: Client,
|
||||
timeout_secs: u64,
|
||||
extra_headers: Vec<(String, String)>,
|
||||
) -> Self {
|
||||
Self {
|
||||
http_client,
|
||||
@@ -309,22 +317,38 @@ impl OpenaiResponseProvider {
|
||||
api_key,
|
||||
model,
|
||||
timeout_secs,
|
||||
extra_headers,
|
||||
}
|
||||
}
|
||||
|
||||
/// 注入 Provider 级别固定头。返回 self 以支持链式调用。
|
||||
pub fn with_extra_headers(mut self, headers: Vec<(String, String)>) -> Self {
|
||||
self.extra_headers = headers;
|
||||
self
|
||||
}
|
||||
|
||||
fn endpoint_url(&self) -> String {
|
||||
format!("{}/responses", self.base_url.trim_end_matches('/'))
|
||||
}
|
||||
|
||||
fn build_request_builder(
|
||||
&self,
|
||||
body: &impl Serialize,
|
||||
body: &OpenaiResponseRequest,
|
||||
) -> Result<reqwest::RequestBuilder, LlmError> {
|
||||
Ok(self
|
||||
let mut builder = self
|
||||
.http_client
|
||||
.post(self.endpoint_url())
|
||||
.header("Authorization", format!("Bearer {}", self.api_key))
|
||||
.json(body))
|
||||
.header("Authorization", format!("Bearer {}", self.api_key));
|
||||
|
||||
for (k, v) in &self.extra_headers {
|
||||
builder = builder.header(k.as_str(), v.as_str());
|
||||
}
|
||||
|
||||
for (key, value) in &body.custom_headers {
|
||||
builder = builder.header(key.as_str(), value.as_str());
|
||||
}
|
||||
|
||||
Ok(builder.json(body))
|
||||
}
|
||||
|
||||
fn map_reqwest_error(&self, e: reqwest::Error) -> LlmError {
|
||||
@@ -403,6 +427,14 @@ impl OpenaiResponseProvider {
|
||||
let metadata: Option<Value> = extra.get("metadata").cloned();
|
||||
let reasoning: Option<Value> = extra.get("reasoning").cloned();
|
||||
|
||||
// 注意:OpenaiResponseProvider 的 `convert_request` 在顶部 destructure 了 `request`,
|
||||
// 因此使用 `extra.get()` 而非 `request.get_extra_opt()`。两者语义一致,
|
||||
// 均反序列化为 `HashMap<String, String>`,失败时静默降级为空 HashMap。
|
||||
let custom_headers: HashMap<String, String> = extra
|
||||
.get("custom_headers")
|
||||
.and_then(|v| serde_json::from_value(v.clone()).ok())
|
||||
.unwrap_or_default();
|
||||
|
||||
// ponytail: tool_choice 直通映射 —— ToolChoice::None → "none"(与 GenericOpenaiProvider 行为一致),
|
||||
// Auto/Required → 字符串,Named → 对象。SA Director 审查 Round 1 修复。
|
||||
let tool_choice_value = match tool_choice {
|
||||
@@ -611,6 +643,7 @@ impl OpenaiResponseProvider {
|
||||
truncation,
|
||||
metadata,
|
||||
reasoning,
|
||||
custom_headers,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -1020,7 +1053,14 @@ mod tests {
|
||||
.timeout(Duration::from_secs(30))
|
||||
.build()
|
||||
.expect("create http client");
|
||||
OpenaiResponseProvider::from_parts(base_url, "sk-test".into(), "gpt-4o".into(), client, 30)
|
||||
OpenaiResponseProvider::from_parts(
|
||||
base_url,
|
||||
"sk-test".into(),
|
||||
"gpt-4o".into(),
|
||||
client,
|
||||
30,
|
||||
Vec::new(),
|
||||
)
|
||||
}
|
||||
|
||||
// ===== convert_request 单元测试 =====
|
||||
@@ -1975,4 +2015,185 @@ data: {\"type\":\"response.failed\",\"error\":{\"message\":\"server failed mid-s
|
||||
assert!(caps.features.streaming);
|
||||
assert_eq!(caps.features.max_context_window, 200_000);
|
||||
}
|
||||
|
||||
// ===== custom_headers (Phase 8 Step 8.7) =====
|
||||
|
||||
fn mock_responses_body() -> serde_json::Value {
|
||||
json!({
|
||||
"id": "resp-test",
|
||||
"object": "response",
|
||||
"status": "completed",
|
||||
"model": "gpt-4o",
|
||||
"output": [{
|
||||
"id": "msg_1", "type": "message", "role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "OK"}]
|
||||
}],
|
||||
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}
|
||||
})
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn openai_response_custom_headers_from_extra() {
|
||||
let provider = make_provider("http://x".into());
|
||||
let mut req = MessageRequest {
|
||||
model: "gpt-4o".into(),
|
||||
messages: vec![Message::user_text("hi")],
|
||||
..Default::default()
|
||||
};
|
||||
req.set_extra("custom_headers", json!({"X-Custom": "v1", "X-Other": "v2"}));
|
||||
let body = provider.convert_request(req).unwrap();
|
||||
assert_eq!(body.custom_headers.get("X-Custom").unwrap(), "v1");
|
||||
assert_eq!(body.custom_headers.get("X-Other").unwrap(), "v2");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn openai_response_custom_headers_skipped_in_json_body() {
|
||||
let provider = make_provider("http://x".into());
|
||||
let mut req = MessageRequest {
|
||||
model: "gpt-4o".into(),
|
||||
messages: vec![Message::user_text("hi")],
|
||||
..Default::default()
|
||||
};
|
||||
req.set_extra("custom_headers", json!({"X-Custom": "v1"}));
|
||||
let body = provider.convert_request(req).unwrap();
|
||||
let value = serde_json::to_value(&body).unwrap();
|
||||
assert!(
|
||||
value.get("custom_headers").is_none(),
|
||||
"custom_headers 不应进入 JSON body"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn openai_response_custom_headers_invalid_type_fallback() {
|
||||
let provider = make_provider("http://x".into());
|
||||
let mut req = MessageRequest {
|
||||
model: "gpt-4o".into(),
|
||||
messages: vec![Message::user_text("hi")],
|
||||
..Default::default()
|
||||
};
|
||||
req.set_extra("custom_headers", json!("not_an_object"));
|
||||
let body = provider.convert_request(req).unwrap();
|
||||
assert!(body.custom_headers.is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn openai_response_custom_headers_are_sent() {
|
||||
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"))
|
||||
.and(header("X-Custom", "v1"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(mock_responses_body()))
|
||||
.mount(&server)
|
||||
.await;
|
||||
|
||||
let provider = make_provider(server.uri());
|
||||
let mut req = MessageRequest {
|
||||
model: "gpt-4o".into(),
|
||||
messages: vec![Message::user_text("hi")],
|
||||
..Default::default()
|
||||
};
|
||||
req.set_extra("custom_headers", json!({"X-Custom": "v1"}));
|
||||
|
||||
let resp = provider.chat_blocking(req).await.unwrap();
|
||||
assert_eq!(resp.text(), "OK");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn openai_response_provider_level_headers_are_sent() {
|
||||
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("X-Platform", "doubao"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(mock_responses_body()))
|
||||
.mount(&server)
|
||||
.await;
|
||||
|
||||
let provider = make_provider(server.uri())
|
||||
.with_extra_headers(vec![("X-Platform".into(), "doubao".into())]);
|
||||
let req = MessageRequest {
|
||||
model: "gpt-4o".into(),
|
||||
messages: vec![Message::user_text("hi")],
|
||||
..Default::default()
|
||||
};
|
||||
let resp = provider.chat_blocking(req).await.unwrap();
|
||||
assert_eq!(resp.text(), "OK");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn openai_response_custom_headers_override_provider_headers() {
|
||||
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(mock_responses_body()))
|
||||
.mount(&server)
|
||||
.await;
|
||||
|
||||
let provider = make_provider(server.uri())
|
||||
.with_extra_headers(vec![("X-Platform".into(), "provider-level".into())]);
|
||||
let mut req = MessageRequest {
|
||||
model: "gpt-4o".into(),
|
||||
messages: vec![Message::user_text("hi")],
|
||||
..Default::default()
|
||||
};
|
||||
req.set_extra("custom_headers", json!({"X-Platform": "request-wins"}));
|
||||
|
||||
let resp = provider.chat_blocking(req).await.unwrap();
|
||||
assert_eq!(resp.text(), "OK");
|
||||
|
||||
let received = server.received_requests().await.unwrap();
|
||||
assert_eq!(received.len(), 1);
|
||||
let platforms: Vec<&str> = received[0]
|
||||
.headers
|
||||
.get_all("X-Platform")
|
||||
.iter()
|
||||
.filter_map(|v| v.to_str().ok())
|
||||
.collect();
|
||||
assert!(platforms.contains(&"request-wins"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn openai_response_custom_headers_can_override_auth_header() {
|
||||
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(mock_responses_body()))
|
||||
.mount(&server)
|
||||
.await;
|
||||
|
||||
let provider = make_provider(server.uri());
|
||||
let mut req = MessageRequest {
|
||||
model: "gpt-4o".into(),
|
||||
messages: vec![Message::user_text("hi")],
|
||||
..Default::default()
|
||||
};
|
||||
req.set_extra(
|
||||
"custom_headers",
|
||||
json!({"authorization": "from-custom-headers"}),
|
||||
);
|
||||
|
||||
let resp = provider.chat_blocking(req).await.unwrap();
|
||||
assert_eq!(resp.text(), "OK");
|
||||
|
||||
let received = server.received_requests().await.unwrap();
|
||||
assert_eq!(received.len(), 1);
|
||||
let auth_values: Vec<&str> = received[0]
|
||||
.headers
|
||||
.get_all("authorization")
|
||||
.iter()
|
||||
.filter_map(|v| v.to_str().ok())
|
||||
.collect();
|
||||
assert!(
|
||||
auth_values.contains(&"from-custom-headers"),
|
||||
"custom_headers 应能覆盖 Authorization 头,实际收到: {auth_values:?}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user