diff --git a/docs/29-openai-response-provider-custom-headers.md b/docs/29-openai-response-provider-custom-headers.md new file mode 100644 index 0000000..6790bbe --- /dev/null +++ b/docs/29-openai-response-provider-custom-headers.md @@ -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`,通过 `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`(OpenAI API 自身的 wire 格式字段) +/// 不同——后者是 OpenAI API 参数,本字段是 reqwest 层的 HTTP 头注入。 +#[serde(skip)] +pub custom_headers: HashMap, +``` + +`#[serde(skip)]` 确保该字段不会出现在序列化后的 JSON body 中。 + +**② `convert_request()` 从 extra 提取** + +在 `parallel_tool_calls`(第 559 行)之后: + +```rust +let custom_headers: HashMap = 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 { + 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, +``` + +`#[serde(skip)]` 确保该字段不会出现在序列化后的 JSON body 中。 + +**⑤ `convert_request()` 从 extra 提取** + +第 404 行,`reasoning` 之后: + +```rust +let custom_headers: HashMap = 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`,失败时静默降级为空 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 { + 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, +} +``` + +`#[serde(skip)]` 确保该字段不会出现在序列化后的 JSON body 中。 + +**⑤ `build_request_body()` 从 extra 提取** + +```rust +let custom_headers: HashMap = 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 { + 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`(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` 枚举 +- 不新增任何平台相关代码 diff --git a/src/llm/provider.rs b/src/llm/provider.rs index 1f519a7..6979174 100644 --- a/src/llm/provider.rs +++ b/src/llm/provider.rs @@ -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 => { diff --git a/src/llm/provider/anthropic.rs b/src/llm/provider/anthropic.rs index ca15942..52fcf02 100644 --- a/src/llm/provider/anthropic.rs +++ b/src/llm/provider/anthropic.rs @@ -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 = + 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 { + 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 { 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, #[serde(skip_serializing_if = "Option::is_none")] stream: Option, + /// 请求级别自定义 HTTP 头。运行时注入,不进入 JSON 请求体。 + #[serde(skip)] + custom_headers: HashMap, } #[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:?}" + ); + } } diff --git a/src/llm/provider/openai.rs b/src/llm/provider/openai.rs index 552fef1..dd533f5 100644 --- a/src/llm/provider/openai.rs +++ b/src/llm/provider/openai.rs @@ -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, #[serde(skip_serializing_if = "Option::is_none")] pub extra_body: Option, + /// 请求级别自定义 HTTP 头。运行时注入,不进入 JSON 请求体。 + /// ⚠️ 与 struct 已有的 `extra_headers: Option`(OpenAI API 自身的 wire 格式字段) + /// 不同——后者是 OpenAI API 参数,本字段是 reqwest 层的 HTTP 头注入。 + #[serde(skip)] + pub custom_headers: HashMap, } // ============================================================================= @@ -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 { 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 = + 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:?}" + ); + } } diff --git a/src/llm/provider/openai_response.rs b/src/llm/provider/openai_response.rs index 95029af..8f8a7d5 100644 --- a/src/llm/provider/openai_response.rs +++ b/src/llm/provider/openai_response.rs @@ -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, #[serde(skip_serializing_if = "Option::is_none")] pub reasoning: Option, + /// 请求级别自定义 HTTP 头。序列化时跳过,仅运行时由 build_request_builder 消费。 + /// stream 模式的修改不影响该字段——header 由 convert_request 在请求构造时注入。 + #[serde(skip)] + pub custom_headers: HashMap, } /// 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 { - 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 = extra.get("metadata").cloned(); let reasoning: Option = extra.get("reasoning").cloned(); + // 注意:OpenaiResponseProvider 的 `convert_request` 在顶部 destructure 了 `request`, + // 因此使用 `extra.get()` 而非 `request.get_extra_opt()`。两者语义一致, + // 均反序列化为 `HashMap`,失败时静默降级为空 HashMap。 + let custom_headers: HashMap = 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:?}" + ); + } }