8 Commits
Author SHA1 Message Date
徐涛 528a17f5fa docs(roadmap): 更新 v0.4.0 规划状态,归档已完成 Phase 2026-07-21 22:44:01 +08:00
徐涛 d48286a942 docs(v0.4.0): 新增 v0.4.0 版本路线图规划文档
涵盖多 Agent 编排、Human-in-the-loop、语义压缩与自动校正等 Phase A-E 共 5 个增量的实施计划、依赖关系与里程碑定义
2026-07-21 22:43:55 +08:00
徐涛 eeae943727 fix(core): 升级 Cargo.toml 版本号到 v0.3.5
CI / test (chat,provider-openai) (push) Has been cancelled
CI / test (chat,provider-openai,provider-openai-response) (push) Has been cancelled
CI / test (chat,provider-openai,tools-mcp) (push) Has been cancelled
CI / test (full) (push) Has been cancelled
CI / test (light) (push) Has been cancelled
CI / test (multi,provider-openai) (push) Has been cancelled
CI / test (multi,provider-openai,tools-mcp) (push) Has been cancelled
CI / clippy (push) Has been cancelled
CI / fmt (push) Has been cancelled
CI / examples (push) Has been cancelled
2026-07-20 16:47:38 +08:00
徐涛 40e4b3d8fe fix(llm): 修复 OpenaiResponseProvider builtin_tools 注入逻辑缺陷
修复当 tools_defs 为空时 builtin_tools 完全不生效的问题:
- 修复 convert_request 分支错位,解耦 tools_defs 与 builtin_tools 处理
- 扩展 ResponseTool 枚举,添加 Builtin(Value) 变体 + 自定义 Serialize/Deserialize
- 改进错误处理:match + warn! 替代 unwrap_or_else 兜底
- 添加 6 个测试用例覆盖纯 builtin / 混用 / 回归 / wire format / roundtrip 场景
- 补充方案文档 docs/30-builtin-tools-injection-fix.md
2026-07-20 16:45:47 +08:00
徐涛 8ea01d373e fix(core): 升级 Cargo.toml 版本号到 v0.3.4
CI / test (multi,provider-openai) (push) Has been cancelled
CI / test (chat,provider-openai) (push) Has been cancelled
CI / test (chat,provider-openai,provider-openai-response) (push) Has been cancelled
CI / test (chat,provider-openai,tools-mcp) (push) Has been cancelled
CI / test (full) (push) Has been cancelled
CI / test (light) (push) Has been cancelled
CI / test (multi,provider-openai,tools-mcp) (push) Has been cancelled
CI / clippy (push) Has been cancelled
CI / fmt (push) Has been cancelled
CI / examples (push) Has been cancelled
版本号 PATCH 升级(0.3.3 → 0.3.4),向后兼容,包含三 Provider 自定义 HTTP 请求头支持(feat 77321db + 安全修复 76bbeed + doc 清理 5e475e1)。
2026-07-20 15:11:18 +08:00
徐涛 5e475e1303 docs(llm): 清理 build_request_builder 重复 doc comment
- anthropic.rs:删除 build_request_builder doc comment 中重复的 4 行(Round 1 安全性修复时追加内容未清理原段落)
- openai_response.rs:在 build_request_builder doc comment 补充「构造 HTTP POST 请求 builder(含认证头与额外请求头)」描述句,与另两 provider 对齐
2026-07-20 15:09:21 +08:00
徐涛 76bbeed596 fix(llm): 自定义头注入安全性修复 + 测试覆盖补全
三 Provider 自定义 HTTP 头机制的审查后修复:

- 安全性:build_request_builder 改用 HeaderName::from_bytes / HeaderValue::from_str
  安全转换;非法 header 名/值(如控制字符)静默跳过 + warn,避免 reqwest panic。

- doc comment:with_extra_headers() 的"注入"措辞与"替换"语义不符
  (self.extra_headers = headers),改为"设置 Provider 级别固定头,替换已有的"。

- 测试覆盖:补全方案验证标准缺失的 *_extra_headers_from_constructor 单元测试
  (三 provider 各 2 个,共 6 新增),用 RequestBuilder::build() 直检 headers。
2026-07-20 14:58:32 +08:00
徐涛 77321db8f6 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 个测试覆盖提取 / 序列化隔离 / 类型降级 / 透传 / 覆盖优先级 / 认证头可覆盖等维度。
2026-07-20 14:47:38 +08:00
10 changed files with 2203 additions and 61 deletions
+1 -1
View File
@@ -1,6 +1,6 @@
[package]
name = "agcore"
version = "0.3.3"
version = "0.3.5"
edition = "2024"
[features]
@@ -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
↑ 后者覆盖前者
```
### 改动一:GenericOpenaiProvideropenai.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
}
```
### 改动二:OpenaiResponseProvideropenai_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`,零影响。
### 改动三:AnthropicProvideranthropic.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` 枚举
- 不新增任何平台相关代码
+268
View File
@@ -0,0 +1,268 @@
# Builtin Tools 注入修复方案
## 背景与目标
### 问题描述
agcore 的 OpenaiResponseProvider 在通过 extra 逃生舱注入内置工具(`web_search` / `file_search`)时存在两层缺陷,导致 builtin_tools 完全不生效:
1. **分支逻辑错位**`convert_request`,第 615-640 行):builtin_tools 注入代码被嵌套在 `tools_defs` 非空的 `else` 分支内。当调用方只提供 builtin_tools 而不提供 tools_defs 时,分支走 `if tools_defs.is_empty() { None }`,注入代码完全不执行。
2. **枚举不完整**`ResponseTool`,第 136-146 行):枚举只有 `Function` 一种变体,非 `function` 类型的 builtin 工具(如 `type: "web_search"`)反序列化失败,退化为 `name: ""` 的空函数定义,API 层面被拒绝。
### 目标
- 修复 builtin_tools 注入逻辑,使纯内置工具、混用场景均正常工作
- 不破坏现有 tools_defs 功能
- 添加回归测试
## 需求分析
### 功能需求
| # | 需求 | 优先级 |
|---|------|--------|
| F1 | `tools_defs` 为空、`builtin_tools` 非空时,正确注入内置工具 | P0 |
| F2 | `tools_defs``builtin_tools` 同时非空时,合并注入 | P0 |
| F3 | 两端均为空时,tools 字段为 None(回归保底) | P0 |
| F4 | 无效的 builtin_tools 值不导致崩溃,跳过并告警 | P1 |
### 非功能需求
- 不做底层架构改造(extra 逃生舱机制不变)
- 不改 `openai.rs`Chat Completions API 不支持内置工具)
## 方案设计
### 总体架构
修复分三步,对应三层独立但不相互依赖的改动:
```
┌─────────────────────────────────────────────────┐
│ convert_request() │
│ │
│ [改动二] 重构分支逻辑 │
│ ┌─────────────────────────────────────────┐ │
│ │ tools_defs ──→ 生成 Vec<ResponseTool> │ │
│ │ builtin_tools ──→ 追加到同一 Vec │ │
│ │ 两者都空 ──→ None; 否则 ──→ Some(items) │ │
│ └─────────────────────────────────────────┘ │
│ │
│ [改动三] 错误处理 │
│ unwrap_or_else ──→ match + warn! │
└─────────────────────────────────────────────────┘
┌─────────────────────────────────────────────────┐
│ ResponseTool 枚举 │
│ │
│ [改动一] 添加 Builtin(Value) 变体 │
│ 自定义 Serialize/Deserialize 避免信息丢失 │
└─────────────────────────────────────────────────┘
```
### 改动一:扩展 `ResponseTool` 枚举
**位置**:第 136-146 行
**现状**
```rust
#[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,
},
}
```
**改后**(需要自定义 Serialize/Deserialize):
```rust
/// NOTE: 仅在请求序列化路径使用(convert_request → build_request_builder → HTTP body)。
/// 响应反序列化走 ResponseOutputItem,不经过此类型。
/// 自定义 Deserialize 服务于 convert_request 内 extra 字段反序列化。
#[derive(Debug, Clone)]
pub(crate) enum ResponseTool {
Function {
name: String,
description: String,
parameters: Value,
},
/// 非 function 类型的工具(如 web_search / file_search / code_interpreter)。
/// 直接透传原始 JSON Value,不做结构化解析,避免信息丢失。
Builtin(Value),
}
impl Serialize for ResponseTool {
fn serialize<S: serde::Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
match self {
ResponseTool::Function { name, description, parameters } => {
let mut map = serde_json::Map::new();
map.insert("type".into(), Value::String("function".into()));
map.insert("name".into(), Value::String(name.clone()));
map.insert("description".into(), Value::String(description.clone()));
map.insert("parameters".into(), parameters.clone());
map.serialize(serializer)
}
ResponseTool::Builtin(value) => value.serialize(serializer),
}
}
}
impl<'de> Deserialize<'de> for ResponseTool {
fn deserialize<D: serde::Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
let value = Value::deserialize(deserializer)?;
match value.get("type").and_then(|t| t.as_str()) {
Some("function") => {
let name = value.get("name").and_then(|n| n.as_str()).unwrap_or_default().to_string();
let description = value.get("description").and_then(|d| d.as_str()).unwrap_or_default().to_string();
let parameters = value.get("parameters").cloned().unwrap_or(Value::Null);
Ok(ResponseTool::Function { name, description, parameters })
}
_ => Ok(ResponseTool::Builtin(value)),
}
}
}
```
关键点:
- 自定义 `Serialize``Builtin` 变体直接输出原始 Value(不包裹额外标记)
- 自定义 `Deserialize`:非 `"function"` 类型自动走 `Builtin(Value)` 分支
- 保留原始 JSON 结构,避免信息丢失(如 `search_context_size``user_location` 等字段)
### 改动二:修复 `convert_request` 分支逻辑
**位置**:第 615-640 行
**现状**(伪代码):
```
if tools_defs.is_empty() {
None // ← builtin_tools 被完全跳过
} else {
从 tools_defs 生成 Vec<ResponseTool>
if let Some(builtin_tools) {
for v in extra {
items.push(from_value(v)) // ← 只有进了 else 才执行
}
}
Some(items)
}
```
**改后**(伪代码):
```
let mut items: Vec<ResponseTool> = Vec::new();
// 1. 始终处理 tools_defs
items.extend(tools_defs.into_iter().map(|t| ResponseTool::Function { ... }));
// 2. 始终处理 builtin_tools(与 tools_defs 解耦)
if let Some(builtin_tools) = builtin_tools {
for v in extra {
// 见改动三
items.push(serde_json::from_value(v).unwrap_or_else(|_| { ... }));
}
}
// 3. 两者都空 → None;否则 → Some
if items.is_empty() { None } else { Some(items) }
```
### 改动三:改进错误处理
**位置**:第 630-636 行(`unwrap_or_else` 部分)
**现状**
```rust
items.push(serde_json::from_value(v).unwrap_or_else(|_| {
ResponseTool::Function {
name: String::new(),
description: String::new(),
parameters: Value::Null,
}
}));
```
**改后**
```rust
match serde_json::from_value(v.clone()) {
Ok(tool) => items.push(tool),
Err(e) => {
let raw = serde_json::to_string(&v).unwrap_or_default();
warn!(tool = %raw, error = %e, "skipped invalid builtin_tool");
}
}
```
`warn!` 输出被反序列化的 Value 摘要,便于生产排障时定位问题(无需复现调用方输入)。
`unwrap_or_else` 的错误值(空函数定义)在 OpenAI API 层会被拒绝,无实际价值。替换为 `match` + `warn!` 可明确跳过并记录原因。
### 改动四:添加测试
在文件末尾 `#[cfg(test)]` 区域或 `tests/` 目录新增 6 个测试用例:
| # | 用例名 | 场景 | 验证点 |
|---|--------|------|--------|
| T1 | `test_builtin_only` | 仅提供 builtin_tools | tools 为 Some,含正确 type |
| T2 | `test_mixed_tools` | 同时提供 tools_defs + builtin_tools | 合并后 items 顺序/数量正确 |
| T3 | `test_no_tools` | 两端均为空 | tools 为 None |
| T4 | `test_invalid_builtin` | builtin_tools 含无效 JSON | 不崩溃,有效项保留 |
| T5 | `test_function_wire_format` | ResponseTool::Function 序列化 | JSON 结构与改动前一致(AC7) |
| T6 | `test_builtin_roundtrip` | ResponseTool::Builtin 反序列化+序列化 | 原始 JSON 结构保留 |
## 实现计划
### 步骤
| 步骤 | 改动 | 文件 | 估算 |
|------|------|------|------|
| 1 | 扩展 `ResponseTool` 枚举,添加 `Builtin(Value)` + 自定义 Serialize/Deserialize | `openai_response.rs:136-146` | 40 行 |
| 2 | 重构 `convert_request` 分支逻辑,解耦 tools_defs 与 builtin_tools | `openai_response.rs:615-640` | 15 行 |
| 3 | 替换 `unwrap_or_else``match` + `warn!` | `openai_response.rs:630-636` | 5 行 |
| 4 | 添加 6 个测试用例 | `openai_response.rs` 末尾 | 70 行 |
| 5 | `cargo test` 验证全部通过 | - | - |
### 优先级
**P0(核心修复)**:步骤 1 + 2,修复分支逻辑和枚举不完整问题。
**P1(健壮性)**:步骤 3,改进错误处理。
**P1(质量保障)**:步骤 4 + 5,测试覆盖。
### 依赖关系
无外部依赖。全部改动限定在 `openai_response.rs` 一个文件内。
## 风险评估
| 风险 | 概率 | 影响 | 缓解措施 |
|------|------|------|----------|
| 自定义 Serialize/Deserialize 实现遗漏边界情况 | 低 | 中 | 测试覆盖所有分支:Function / Builtin / 无效值 |
| 现有 `Function` 序列化格式变化 | 低 | 高 | 自定义 Serialize 保持与原 derive 行为一致,测试覆盖 wire 格式 |
| `warn!` 日志在生产环境未配置 logger 导致 panic | 低 | 中 | 使用 `tracing::warn!`(已导入),项目已初始化 tracing logger |
### 回滚方案
单文件改动,回滚只需 `git checkout -- src/llm/provider/openai_response.rs`
## 验收标准
| # | 验收条件 | 验证方式 |
|---|----------|----------|
| AC1 | `tools_defs` 空 + `builtin_tools``{"type":"web_search"}` → 请求体 `tools` 包含 `{"type":"web_search"}` | 单测 T1 |
| AC2 | 混用场景 → tools 数组同时包含 function 和非 function 工具 | 单测 T2 |
| AC3 | 两端空 → tools 字段为 null/None | 单测 T3 |
| AC4 | 无效 builtin_tools → 不 panic,有效项不受影响 | 单测 T4 |
| AC5 | 全部现有测试通过 | `cargo test` |
| AC6 | `cargo clippy` 无新增警告 | `cargo clippy` |
| AC7 | `ResponseTool::Function` 序列化后的 JSON 结构与改动前一致 | 单测:验证字段顺序和值 |
---
**编写人**Writer Agent
**编写日期**2026-07-20
**基于**agcore builtin_tools 注入问题分析结论
+110 -19
View File
@@ -5,7 +5,8 @@
> **已分版本的内容**:请查阅
> - [`roadmap-v0.1.0.md`](./roadmap-v0.1.0.md) — Phase 04c + v0.1.0 Release
> - [`roadmap-v0.2.0.md`](./roadmap-v0.2.0.md) — Phase 512 + v0.2.0-rc.1
> - [`roadmap-v0.3.0.md`](./roadmap-v0.3.0.md) — Phase 131913-18 已完成,19 待实施
> - [`roadmap-v0.3.0.md`](./roadmap-v0.3.0.md) — Phase 1319全部完成
> - [`roadmap-v0.4.0.md`](./roadmap-v0.4.0.md) — Phase A-E 多 Agent 编排路线图
>
> 返回总入口:[`roadmap.md`](./roadmap.md)
@@ -15,7 +16,7 @@
AG Core 定位为构建 AI 智能体的底层工具箱,通过模块化、可插拔的架构,提供大模型调用、提示词工程、工具系统、记忆检索四大核心能力,支持快速组合出符合业务需求的智能体应用。
**当前状态**v0.2.0-rc.1 已打标签。Phase 0-18 全部完成。v0.3.0 实施中,Phase 19 共 1 个增量 Phase 待交付。目标是从"LLM 调用工具箱"升级为"能构建多 Agent 协作、RAG、长记忆 Agent 产品的基础系统"。
**当前状态**v0.3.5。Phase 0-30 全部完成。v0.4.0 规划已确定,覆盖 5 个增量 Phase(A-E):Swarm 编排抽象、结果聚合、Human-in-the-loop + 用户 Steering、TokenJuice 语义压缩、自动校正。目标是从"多 Agent 基础系统"升级为"多 Agent 多职责编排系统"。
---
@@ -35,21 +36,30 @@ AG Core 定位为构建 AI 智能体的底层工具箱,通过模块化、可
---
## v0.4+ 展望
## v0.4.0 规划
### 已规划的功能
v0.4.0 的完整规划已移入独立的 [`roadmap-v0.4.0.md`](./roadmap-v0.4.0.md),包含 5 个增量 Phase
| 功能 | 说明 | 预计版本 |
|------|------|---------|
| Multi-Agent Swarm 编排 | Supervisor/Subgraph 模式,基于 v0.3 dispatch 构建 | v0.4 |
| Human-in-the-loop 审批 | `interrupt()` + `Command(resume=...)` 异步审批回调 | v0.4 |
| Agent 自动创生 | LLM 自主决定何时派发子 agent、派发什么角色 | v0.4 |
| 分布式 session 共享 | SessionManager Redis 后端支持跨进程 | v0.4 |
| 精确 tokenizer 计数 | 引入 `tiktoken-rs`,绑定模型具体 tokenizer,替换字符估算 | v0.4+ |
| TokenJuice 语义压缩 | 对工具结果做语义压缩而非字节截断 | v0.4+ |
| Markdown 技能按需加载 | 技能注册表 + 按 prompt 上下文动态加载 | v0.4+ |
| 增量 checkpoint | 仅存储变化部分,替换当前全量 JSON 模式 | v0.4+ |
| RL 轨迹导出 | ShareGPT 格式轨迹、Atropos 集成 | v0.4+ |
| Phase | 内容 | 状态 |
|-------|------|------|
| **Phase A** | Swarm 编排(Star/Sequential/Hierarchical + Subgraph | 📋 待实施 |
| **Phase B** | 结果聚合 + 编排模式完善 | 📋 待实施 |
| **Phase C** | Human-in-the-loop + 用户 Steering | 📋 待实施 |
| **Phase D** | TokenJuice 语义压缩(工具结果/历史/跨 Agent) | 📋 待实施 |
| **Phase E** | 自动校正 / Reflection | 📋 待实施 |
### 未来版本(v0.5+
以下功能已从 v0.4 范围移出:
| 功能 | 说明 |
|------|------|
| Agent 自动创生 | LLM 自主决定何时派发子 agent — 设计复杂,v0.4 专注显式声明式编排 |
| 分布式 session 共享(Redis 后端) | 与编排正交,多数用户单进程即可 |
| 精确 tokenizer 计数(tiktoken-rs | 依赖引入,不在 v0.4 核心范围内 |
| 增量 Checkpoint | 存储优化,当前全量 JSON 够用 |
| 路线 BStateGraph 通用图引擎) | 预留为路线 A 的未来升级路径 |
| RL 轨迹导出 | 专项需求 |
### 明确不做(agcore 范围外)
@@ -74,10 +84,10 @@ AG Core 定位为构建 AI 智能体的底层工具箱,通过模块化、可
## 下一步行动
1. **v0.3.0 Phase 19 启动**KnowledgeGraph + 双通道检索,落地 `docs/note-knowledge-graph-design.md` 中记录的知识图谱设计
2. **Phase 19 收尾**完成 v0.3.0 最后一个 Phase 后准备 rc.1 标签 + CHANGELOG
3. **示例先行**完成 Phase 19 后立即创建对应的 knowledge_graph_demo 示例,确保 `cargo run --example` 可验证
4. **里程碑追踪**:以 M13Phase 17+ M14Phase 18)为已达成里程碑,逐 Phase 推进 M15
1. **v0.4.0 启动**:按 [`roadmap-v0.4.0.md`](./roadmap-v0.4.0.md) 规划,从 Phase A(Swarm 编排)开始实施
2. **Phase A 实施**engine/supervisor.rs + tools/builtin.rs + Swarm::star/sequential/hierarchical
3. **示例先行**每个 Phase 交付时同步提交对应的示例程序
4. **里程碑追踪**:以 M16-M20 为目标里程碑,逐 Phase 推进
---
@@ -107,3 +117,84 @@ AG Core 定位为构建 AI 智能体的底层工具箱,通过模块化、可
-**v0.3.0 Phase 16 完成**`SummaryConfig` 配置结构体(6 个字段:`trigger_token_ratio=0.75` / `max_context_tokens=32_000` / `summary_prompt` / `debounce_turns=3` / `summary_model=None` / `max_tool_result_chars=500`,默认 `None` 沿用主模型避断裂非 OpenAI 用户)+ `AgentBuilder::summary_config(cfg)` 链式方法 + `AgentConfig.summary_config: Option<SummaryConfig>` 字段;`AgentSession` 新增 `last_summary_turn: Option<u32>` 字段(首次不受防抖约束,`should_summarize``Option` 哨兵实现)+ `maybe_summarize(current_turn)` 内联检查点(OnTurnEnd 之后 / `turn_index` 之前,对称 `submit_turn` / `finalize_turn` 两个入口,流式路径 `saturating_sub(1)` 修正)+ 关联函数 `generate_summary`(构造独立 `LlmCycle``submit_messages``vec![Message::user_text(prompt)]``max_tokens=1024`,空消息守卫直接返回空串)+ 公开 API `get_conversation_summary()``src/agent/summary.rs`~240 行,含 8 个 SummaryConfig/`format_messages_as_text` 内联测试——默认值/空输入/系统用户助理/ToolResult(含 `tool_call_id`/工具调用/Unicode 安全截断/整体 30K 截断保留最新;有效字符数截断多字节安全,droptest 验证保留尾部消息)+ `src/agent/session.rs` 注入 10 个摘要集成测试(默认值不触发 / 超阈值触发 / 防抖阻止重复 / SessionMemory 写入 / Full 模式不注入 / 失败不阻断主流程 / 流式路径触发 / 默认配置零影响 / **Focused `summary_override` 写入正向验证** / **空消息不调用 LLM** / **巨型 `max_context_tokens` 永不触发**);`format_messages_as_text` 简洁版消息格式化(`[Tool: name]` + `Tool Result [id]:` + ToolResult 字符级 `chars().take(max_tool_result_chars)` 截断 + 整段 30K 总长度截断从头部保留最新);所有错误静默(失败用 `tracing::error!`,成功用 `tracing::info!(turn, summary_len)`);`MergeStrategy` 注释中过时 "Summarize 指向"与 `context.rs:78` "v0.3 将支持 Hook 驱动" 过时注释在实施时同步移除/更新;方案文档 `docs/22-phase16-summary-auto-generation.md`(471 行),实施后**两轮审查 PASS**:第一轮 PM/SA 审查 11 项问题修复 + 第二轮实施审查 9 项问题修复(🔴 `generate_summary` 空消息 bug + 🟡 W4 流式路径防抖 + 🟡 W2 模型硬编码 + 🟡 W5 Full 模式无谓 save + 🟡 W3 30K 截断 + 🟡 W6 成功无日志 + 🟡 W1/W7 测试补全 + 💭 注释同步);零新外部依赖;全量 335 → **353**(+18 新测试,含二次审查增补 4 个),clippy 0 警告,doc 0 warning`quick_start` 示例正常 exit 0;**M12 里程碑达成** + 第二轮审查门禁 PASS
-**v0.3.0 Phase 17 完成** — 新建 `src/engine/` 模块(5 文件:`mod.rs`/`error.rs`/`snapshot.rs`/`checkpointer.rs`/`session_manager.rs`),实现 **SessionManager**10 个公开方法:`create`/`create_child`/`get`/`recover`/`replace`/`children`/`parent`/`destroy`/`submit_turn`/`submit_turn_stream`/`finalize_turn_stream`,内部 `RwLock<HashMap>` + `Arc<tokio::sync::Mutex<AgentSession>>` + `Checkpointer` 组合)和 **Checkpointer**5 个公开方法:`checkpoint`/`rollback_load`/`list_checkpoints`/`delete_all`/`latest_snapshot`);`SessionSnapshot` 独立 struct 避开 `Arc<dyn Agent>` 不可序列化,配套 `SessionMemoryEntry` 保留 metadata/created_at`AgentSession` 扩展三段式快照(`to_snapshot` async 读 MemoryStore + `from_snapshot` 纯同步构造 + `restore_memory` &mut self async 写回持久层);`SessionMemory` 新增 `list_entries()``set_with_meta()` 方法(恢复时保留完整 entry 数据);存储 key 风格统一为 `session:{id}:meta` / `ckpt:{id}:{ckpt_id}`(与 `slot_data:` 风格一致);`EngineError` 6 个变体(含 `Memory(#[from] MemoryError)` 透传 + `Agent(#[from] AgentError)`);`CkptMeta``created_at_nanos` 字段确保同秒内精确降序排序;ckpt_id 用纳秒+单调计数器生成(零外部依赖,ponytail);session_id 用纳秒+计数器自动生成(统一策略,UUID v4 备选);自动 checkpoint 失败 `tracing::error!` 不阻断主流程(不提供强持久化保证);流式 checkpoint 仅在 `finalize_turn_stream` 创建(不留半成品污染);孤儿策略:`destroy()` 不递归删除子 session,父被销毁后 `parent()` 返回 `Ok(None)`3 处 derive 改动(`CostTracker` + `ContextSlot` + `MergeStrategy` 加 serde`CostTracker` 额外加 `Clone`);`SessionManager::recover` + `replace` 内部自动 `restore_memory` 写回持久层;零新外部依赖;方案文档 `docs/23-phase17-agent-execution-engine.md`(775 行,经两轮 PM+SA 审查 + 实施后第三轮 PM+SA+Code Reviewer 三方联合审查),实施后**两轮审查门禁 PASS**:第一轮修复 6 🔴 + 第二轮修复 2 🔴(to_snapshot 同步→async + Roadmap 同步)+ 实施后修复 8 个 🟡(restore_memory metadata/created_at 完整恢复 + &mut self 签名 + 死代码清理 + 3 个边界测试 + tracing 补全 + 文档语义统一 + 示例 rollback 一致性 assert);15 个 `SessionManager` 内联测试(CRUD/recover/replace/树形/孤儿/auto_checkpoint on-off+ 6 个 `Checkpointer` 内联测试(roundtrip/不存在的 ckpt/同秒降序/delete_all 幂等/latest/隔离)+ 1 个 `snapshot_deserialize_with_minimal_fields` 序列化兼容测试;全量 353 → **374**+21 新测试),clippy 0 警告,doc 0 warning`engine_demo` 示例端到端演示 create→submit_turn→checkpoint→rollback→replace→destroy 全链路并验证 rollback 一致性;**M13 里程碑达成** + 两轮审查门禁 PASS
-**v0.3.0 Phase 18 完成** — 新增 `src/engine/switch.rs`222 行)实现 `SessionManager::switch_agent()` 热切换(替换 `Arc<dyn Agent>`slot 历史 / `turn_index` / `session_memory` / `cost_so_far` 全部保留,同步更新 `SessionMeta.agent_name` 到持久层,`created_at` / `parent_id` 保持原始不可变)+ 新增 `src/engine/sub_agent.rs`1071 行)实现 4 个公开方法(`dispatch` / `dispatch_all` / `dispatch_stream` 与前述 `switch_agent` 共 4 个 Phase 18 核心 API+ 3 个公开类型(`DispatchConfig` / `SubTaskResult` / `SubTaskStreamEvent`);`DispatchConfig` 4 字段(`max_concurrency=10` / `inherit_session_memory=true` / `bridge_keys=None` / `shared_namespace=None`+ 三态 `bridge_keys` 语义(`None` = 不继承 / `Some(vec![])` = 全部 / `Some(keys)` = 指定 keys+ 约定式 `shared_namespace` 子↔子共享(`shared:{prefix}:{key}`)不触发自动注入;`dispatch` 流程:`create_child``inherit_session_memory`(快照语义)→ `submit_turn` → 返回 `SubTaskResult``dispatch_all` `tokio::sync::Semaphore` 并发控制 + `Vec<Result<...>>` 部分成功语义按输入顺序 indexed 收集;`dispatch_stream` `unbounded_channel` + spawn task 消息重建 + `finalize_turn` 后台落库(明确不参与 `auto_checkpoint` 防重复);`SubTaskStreamEvent` 事件序列:`ChildCreated``Stream(StreamEvent) × N``Completed(SubTaskResult)``Error { child_id, error }``EngineError` 新增 `DispatchFailed(#[source] String)` 变体 + `CostTracker``From<Usage>` 转换;`save_session_meta` / `load_session_meta``pub(crate)``switch.rs` 调用;4 个端到端示例:`agent_switch_demo`115 行)+ `sub_agent_dispatch_demo`141 行)+ `bridge_keys_demo`197 行)+ `dispatch_stream_demo`121 行)全部 exit 017 个内联测试(4 switch + 5 dispatch + 4 dispatch_all + 4 dispatch_stream);零新外部依赖;方案文档 `docs/24-phase18-agent-switch-and-dispatch.md`700 行);全量 374 → **391**(+17 新测试,0 失败),clippy 0 警告,doc 0 warning**M14 里程碑达成**
-**v0.3.0 Phase 19 完成** — 知识图谱 + 双通道检索,详见 `docs/25-phase19-knowledge-graph-and-retrieval.md`;全量 391 → **427 passed / 0 failed**(+36 新测试);**M15 里程碑达成**
-**v0.3.2 Phase 20-27 全部完成** — Cargo features 拆分(16 模块级 + 5 provider + 4 快捷组合),详见 `docs/roadmap-v0.3.2.md`;全量 427 → **427 passed**(不变,门控验证)
-**Phase 28-30 OpenAI Response API Provider 完成** — 独立 feature `provider-openai-response`,全量约 450 passed
- 📋 **v0.4.0 规划完成** — 5 个增量 PhaseA-E)覆盖多 Agent 编排、HITL + Steering、TokenJuice、自动校正。详见 [`roadmap-v0.4.0.md`](./roadmap-v0.4.0.md)
---
## 设计笔记
### Checkpointer 分层存储模型
> 来源:v0.4.0 规划讨论中涉及增量 Checkpoint 的技术推演。当前全量 JSON checkpoint 够用,但为未来优化预留设计方案。
#### 分层叠加模型(OverlayFS 模式)
受容器分层文件系统启发,增量 Checkpoint 可以借鉴 overlayfs 的"底层只读 + 上层可写叠加"设计:
**全量基座(只读)**
```rust
pub struct SnapshotBase {
pub checkpoint_id: String,
pub session_id: String,
pub snapshot: SessionSnapshot, // 完整 JSON 化状态
}
```
**增量层(叠加 diff**
```rust
pub struct SnapshotLayer {
pub base_checkpoint_id: String,
pub applies_to_id: String, // 在哪个 checkpoint 上叠加
pub diff: Vec<DiffOp>, // JSON Patch 操作集合
}
pub enum DiffOp {
MessageAppended { message: Message },
SlotChanged { slot_id: String, diff: serde_json::Value },
TurnIndexIncremented { from: u32, to: u32 },
CostUpdated { diff: CostTracker },
}
```
**重建路径**
```
rollback_load("session_x", 6)
→ 读取 "ckpt:{session_x}:base"(全量)
→ 读取 "ckpt:{session_x}:layer:1" ~ "ckpt:{session_x}:layer:6"
→ 依次应用 layer.1 → layer.2 → ... → layer.6
→ 得到 session_6 的状态
```
**层折叠(类似 docker squash**
```
layer.1 → layer.2 → layer.3 → layer.4 → layer.5
↓ 合并
base.ckpt'(包含 layer.1-3)→ layer.4 → layer.5
```
#### Shadow FS 模型(运行中保护)
与分层模型互补,shadow 模型适用于运行中的 session 保护而非长期存储:
```rust
// submit_turn 在 shadow session 上执行,commit 时才原子切换
let shadow = current_session.fork(); // 复用 ContextSlot::fork
let result = shadow.submit_turn(input).await;
if result.is_ok() {
current_session.commit(shadow); // 原子替换
} else {
drop(shadow); // 丢弃,当前 session 完好无损
}
```
#### 适用场景对比
| 模型 | 适合场景 | 不适合场景 |
|------|---------|-----------|
| **分层叠加(OverlayFS** | Checkpoint 链长期存储、time-travel、多版本回退 | session 较小(< 10KB/轮)时复杂度不值得 |
| **Shadow FSCoW** | 运行中 session 保护、防止 submit_turn 失败污染 | 不能替代 checkpoint 链、不支持多时间点回退 |
**触发条件**:当单 session checkpoint 超过 500KB 且频繁保存导致性能瓶颈时,考虑实现分层模型。
+233
View File
@@ -0,0 +1,233 @@
# AG Core Roadmap — v0.4.0
> 本文件聚焦 **v0.4.0 版本** 的规划。Phase A-E 计划中,覆盖多 Agent 编排、Human-in-the-loop 与 Steering、语义压缩、自动校正。
> 返回总入口:[`roadmap.md`](./roadmap.md)
## v0.4.0 愿景
从 v0.3 的"多 Agent 基础系统"升级为"多 Agent 多职责编排系统"。补齐高层编排抽象(Swarm/Supervisor/Subgraph)、生产级人工干预能力(HITL + Steering)、工具与消息的语义压缩(TokenJuice),以及自动质量校正(Reflection)。为即将开发的多 Agent 协作产品提供完整的编排、干预与质量保证层。
## v0.4.0 总体范围
**总体规模**5 个增量 PhasePhase A-E),总新增代码约 1,950 行,零强制新外部依赖,零破坏性变更。
### 架构决策
**路线选择**:采用轻量编排模式(路线 A),不引入通用有向图引擎。通过 `Swarm::star()` / `Swarm::sequential()` / `Swarm::hierarchical()` 等具名模式提供编排能力,底层复用现有 `dispatch` / `create_child` / `SessionManager` 基础设施。预留路线 B(StateGraph 抽象)作为未来版本的升级路径。
**模块位置**
- 编排逻辑 → `src/engine/supervisor.rs`(新增)
- 内建工具 → `src/tools/builtin.rs`(新增)
- TokenJuice 压缩 → `src/llm/compress.rs`(新增)
- Steering 机制 → `src/engine/steer.rs`(新增,或并入 supervisor.rs
### 功能清单
#### P0 — 必须交付
| # | 功能 | 模块 | 方案要点 |
|---|------|------|---------|
| 1 | Swarm 编排(Star/Sequential/Hierarchical + Subgraph | `engine/supervisor` | `Swarm::star().supervisor(A).worker(B)` 声明式 API`Swarm::sequential().link(A).link(B)` 串联;`Swarm::hierarchical().supervisor(root).group("sub", ...)` 层次嵌套 |
| 2 | 结果聚合 | `engine/supervisor` | `aggregation_prompt` 模板将子 Agent 结果合并到 Supervisor 上下文;`DispatchConfig` 扩展 `result_key` 字段 |
| 3 | Human-in-the-loop 审批 | `engine/steer` | `interrupt()` 暂停执行 + `Command(resume=bool)` 恢复;`HookEvent::OnInterrupt` 新变体 |
| 4 | 用户 Steering(运行中校正) | `engine/steer` | `Command(resume=Correction{...})` 结构化校正;Steer 消息在工具批处理边界注入 |
| 5 | TokenJuice 语义压缩 | `llm/compress` | `Compressor` trait 统一抽象;覆盖工具结果、对话历史、跨 Agent 消息三层;LLM 摘要压缩 + 确定性兜底 |
#### P1 — 推荐交付
| # | 功能 | 模块 | 方案要点 |
|---|------|------|---------|
| 6 | 自动校正 / Reflection | `engine/reflect` | Evaluator-Optimizer 循环;Producer-Critic 角色分离;上限 2-3 轮迭代 |
### 实施计划 — 5 个增量 Phase
> **编号说明**Phase A-E 为 v0.4.0 专属编号,接续已完成的 Phase 30。
---
#### Phase A: Swarm 编排抽象(Star / Sequential / Hierarchical + Subgraph
**目标**:在现有 `dispatch` 原语基础上,提供声明式多 Agent 编排 API。Supervisor 作为 `Arc<dyn Agent>`,通过内建工具 `dispatch_sub_agent` 驱动子 Agent 执行。
**交付物**
1. `src/engine/supervisor.rs` 新文件:
- `Swarm` 枚举/结构体:`Swarm::star()`(星型,一个 Supervisor + N 个 Worker)、`Swarm::sequential()`(顺序链 A→B→C)、`Swarm::hierarchical()`(层次嵌套,Supervisor 下的 Sub-Supervisor
- 各模式的 `build()``run(input)` 方法
- 底层通过 `SessionManager::dispatch()` / `dispatch_all()` 实现
2. `src/tools/builtin.rs` 新文件:
- `dispatch_sub_agent(name, task, config)` 内建工具 — 从 Agent 注册表查找 Agent 工厂 → `SessionManager::dispatch()`
3. `AgentRegistry``HashMap<String, Box<dyn Fn() -> Arc<dyn Agent>>>` 轻量工厂注册表(约 50 行)
4. Subgraph 嵌套:`Swarm::hierarchical()` 支持 `group(name, inner_swarm)`,内层 Swarm 作为子节点编译后嵌入
**设计要点**
- Supervisor 就是 `Arc<dyn Agent>`,不新增 `SupervisorAgent` trait
- 路由逻辑写在 Supervisor 的 system prompt 中(LLM 决定的动态路由)
- 三种模式覆盖常见编排拓扑,不引入通用图引擎(路线 B 留作未来)
- Subgraph 编译为独立的 `SessionManager` 子树(复用 `create_child` 的父子关系)
**依赖**Phase 18SubAgent dispatch / SessionManager
**优先级**P0
**预估规模**:约 500 行
**状态**:📋 待实施
---
#### Phase B: 结果聚合 + 编排模式完善
**目标**:让 Supervisor 能智能地合并 Worker 结果。完善三种编排模式的容错性和易用性。
**交付物**
1. `aggregation_prompt` 模板系统 — 内建 `DEFAULT_AGGREGATION_PROMPT`,用户可自定义聚合逻辑
2. `DispatchConfig` 扩展:
- `result_key: Option<String>` — 将子结果存入 `session_memory` 的指定 key,供后续阶段使用
- `aggregate_strategy: AggregateStrategy``Concatenate` / `Summarize` / `Custom(Value)`
3. 编排模式增强:
- `Swarm::sequential()` 支持失败时停止 / 跳过 / 重试策略
- `Swarm::star()` 支持 Worker 超时
4. 端到端示例 3 个:
- `swarm_star_demo.rs` — 星型编排 + 并发派发 + 结果聚合
- `swarm_sequential_demo.rs` — 串联流水线
- `swarm_hierarchical_demo.rs` — 层次嵌套(Supervisor → Sub-Supervisor → Worker
**依赖**Phase A
**优先级**P0
**预估规模**:约 200 行
**状态**:📋 待实施
---
#### Phase C: Human-in-the-loop + 用户 Steering
**目标**:生产级多 Agent 系统的关键门禁。提供执行中暂停-审批-恢复机制,以及用户运行中校正方向的能力。
**交付物**
1. `src/engine/steer.rs` 新文件:
- `interrupt(value)` 函数 — 在工具循环中插入暂停点,持久化当前状态后返回控制权
- `Command` 枚举:
- `Command::Resume(bool)` — 二元审批(批准/拒绝)
- `Command::ResumeWith(Correction)` — 结构化校正(修改工具参数 / 调整方向)
2. `LlmCycle` 扩展:可中断工具循环模式
- `submit_with_tools_interruptible()` — 支持在工具批处理边界检查中断信号
- 中断时保存当前 `LlmCycle` 状态到 checkpoint
3. `HookEvent::OnInterrupt` / `OnSteer` 新变体 — 监听中断和校正事件
4. `SessionManager::resume_turn(session_id, resume_data)` — 从 checkpoint 恢复并注入审批结果
5. `tools/builtin.rs` 扩展:
- `request_approval(question, context)` — 请求用户审批
- `emit_steer(correction)` — 用户校正
6. Steering 生命周期:
- `interrupt` → 用户收到提示 → 用户决定方向 → `Command::ResumeWith(correction)` → Agent 在新方向上继续
**设计要点**
- User Steering 不是简单的"批准/拒绝",而是 `Correction { action, reason, amended_params }` 结构化指令
- Steering 消息在工具批处理边界(Worker 返回后、Supervisor 决策前)注入,不中断正在执行的工具
- 继承 `ContextSlot::fork/merge` 模式,steer 前 fork 快照,允许用户回退到 steer 前状态
**依赖**Phase ASwarm 编排)
**优先级**P0
**预估规模**:约 500 行
**状态**:📋 待实施
---
#### Phase D: TokenJuice 语义压缩
**目标**:替代当前字节级截断(`microcompact``[pruned]`),提供语义级别的压缩。在三层管道中接入:工具结果压缩、对话历史压缩、跨 Agent 消息压缩。
**交付物**
1. `src/llm/compress.rs` 新文件:
- `Compressor` trait`async fn compress(&self, input: &str, ctx: &CompressionContext) -> Result<String>`
- `CompressionContext``target_tokens` / `preserve_keys` / `strategy`
- `CompressionStrategy` 枚举:`Semantic { model }`LLM 摘要)、`Extractive { ratio }`(抽取式)、`Hybrid { semantic_first }`(混合)
- `SemanticCompressor` 实现(复用已有 provider 做 LLM 摘要压缩)
- `ExtractiveCompressor` 实现(确定性关键句提取,零 LLM 调用)
2. 三层接入点:
- **工具结果压缩**:在 `run_tool_loop` 中,`tool.execute()` 后插入 `compress_result()`,压缩结果再 `push ToolResult`
- **对话历史压缩**:在 `load_messages()` 后插入 `compress_history()`,替代/补充 `microcompact`
- **跨 Agent 消息压缩**:在 `inherit_session_memory` 的子 memory 写入前压缩(减少子 Agent 的 context 水位)
3. `CycleConfig` / `CompactConfig` 扩展:
- `token_compression: Option<CompressionConfig>` — 可选语义压缩配置
- `fallback_to_microcompact: bool`(默认 `true`)— LLM 压缩失败时退化为字节截断
4. TokenJuice 与现有 `microcompact` 的关系:
- `microcompact` 保留为最轻量级兜底(零 LLM 调用)
- TokenJuice 是可选增强层(默认关闭,用户 opt-in)
**设计要点**
- 零新外部依赖:LLM 摘要压缩复用已有 provider,抽取式压缩纯 Rust 实现
- 与现有 `CompactState` 断路器模式兼容(LLM 压缩失败 3 次后自动降级到 `microcompact`
- `preserve_keys` 确保关键数据(数字、ID、SQL、代码片段)不被压缩掉
**依赖**Phase 14Embedding trait 可选参考)
**优先级**P0
**预估规模**:约 400 行
**状态**:📋 待实施
---
#### Phase E: 自动校正 / Reflection
**目标**:实现 Agent 输出后的自我质量评估与自动修正循环。基于 `interrupt/resume` 基础设施,构建 Producer-Critic 闭环。
**交付物**
1. `src/engine/reflect.rs` 新文件:
- `ReflectionConfig``max_cycles`(默认 2/ `critic_agent`(可选不同模型)/ `criteria: Vec<String>`(评估标准)
- `Reflectable` trait`fn reflection_criteria(&self) -> Vec<String>` + `fn needs_refinement(&self, critique: &Critique) -> bool`
- `ReflectionLoop``evaluate(output) → Critique``should_refine? → yes: refine(output, critique) → 循环 / no: 返回`
2. Swarm 内建 Reflection 模式:
- `Swarm::reflect(producer_agent, critic_agent)` — 专用 Reflection Swarm
- 可在 Supervisor 流程中嵌入 `reflect_on(worker_result)` — 对 Worker 结果自动过一遍质量检查
3. `Critique` 结构体:`issues: Vec<Issue>` / `score: f32` / `should_refine: bool` / `suggestions: Vec<String>`
4. `tools/builtin.rs` 扩展:`verify_output(claim, evidence)` 工具 — 让 Agent 自行验证输出真实性
**设计要点**
- Producer 和 Critic 使用**不同模型**(避免同一模型的自我审查盲区 bias)
- 上限 2-3 轮(第一轮修正捕获 70–80% 改善空间,第 4+ 轮收益递减)
- 基于已有 `HookEvent::OnTurnEnd` 或扩展 `HookEvent::OnOutputGenerated` 触发反思
- 失败静默:Reflection 失败不阻断主流程(`tracing::warn!` 后继续交付原始输出)
**依赖**Phase Cinterrupt/resume 基础设施)
**优先级**P0
**预估规模**:约 350 行
**状态**:📋 待实施
---
### v0.4.0 Phase 依赖关系图
```mermaid
graph BT
PA["<b>Phase A: Swarm 编排</b><br/>Swarm::star/sequential/hierarchical<br/>Subgraph 嵌套<br/>内建 dispatch_sub_agent 工具<br/>~500 行"]:::pending
PB["<b>Phase B: 结果聚合</b><br/>aggregation_prompt 模板<br/>DispatchConfig result_key<br/>编排模式完善<br/>3 个端到端示例<br/>~200 行"]:::pending
PC["<b>Phase C: HITL + Steering</b><br/>interrupt/resume<br/>Command(ResumeWith Correction)<br/>HookEvent::OnInterrupt<br/>~500 行"]:::pending
PD["<b>Phase D: TokenJuice</b><br/>Compressor trait<br/>工具结果/历史/跨 Agent 压缩<br/>Semantic + Extractive 策略<br/>~400 行"]:::pending
PE["<b>Phase E: 自动校正</b><br/>ReflectionLoop<br/>Producer-Critic<br/>上限 2-3 轮<br/>~350 行"]:::pending
PB --> PA
PC --> PA
PE --> PC
classDef done fill:#4ade80,stroke:#16a34a,color:#1a1a1a
classDef pending fill:#fbbf24,stroke:#d97706,color:#1a1a1a
classDef future fill:#94a3b8,stroke:#64748b,color:#1a1a1a
```
### 关键里程碑
| 里程碑 | Phase 完成条件 | 可验证指标 | 状态 |
|--------|---------------|-----------|------|
| **M16** | Phase A | `Swarm::star().supervisor(A).worker(B).run(input)` 端到端验证;`dispatch_sub_agent` 内建工具注册并可用;2 个示例 exit 0 | 📋 待启动 |
| **M17** | Phase B | `Swarm::sequential()` 串联执行验证;`Swarm::hierarchical()` 层次嵌套验证;结果聚合正确合并;3 个新示例 exit 0 | 📋 待启动 |
| **M18** | Phase C | `interrupt()` 暂停 + `Command::Resume(bool)` 恢复全链路验证;`Command::ResumeWith(Correction)` 结构化校正验证;HookEvent 触发验证 | 📋 待启动 |
| **M19** | Phase D | 工具结果经语义压缩后保留关键信息(验证压缩比 ≥ 3:1);`microcompact` 降级路径验证;对话历史压缩验证 | 📋 待启动 |
| **M20** | Phase E | ReflectionLoop 正确性验证:已知缺陷的输出被修复、无缺陷的输出不被修改(不变性保证);2 轮迭代上限验证;Critic 不同模型配置验证 | 📋 待启动 |
### 不做(v0.5+
| 功能 | 原因 |
|------|------|
| Agent 自动创生(LLM 驱动动态分派) | 设计复杂且不确定性高,v0.4 专注显式声明式编排 |
| 分布式 Session 共享(Redis 后端) | 与编排正交,大多数用户单进程即可 |
| 精确 tokenizer 计数(tiktoken-rs | 依赖引入,v0.4 专注编排与压缩能力本身 |
| 增量 Checkpoint | 存储优化,当前全量 JSON 够用 |
| 路线 BStateGraph 通用图引擎) | 当前编排需求在路线 A 范围内,图引擎留给未来版本 |
| RL 轨迹导出 | 专项需求,非通用 |
| Markdown 技能按需加载 | 独立功能 |
+3 -2
View File
@@ -1,7 +1,7 @@
# AG Core Roadmap
> 拆分式 roadmap:按版本归档 + 未归类内容
> 最后更新:2026-07-20Phase 28-30 OpenAI Response API Provider 交付 — 新增独立 feature `provider-openai-response`450 测试通过
> 最后更新:2026-07-21v0.4.0 规划完成 — Phase A-E 多 Agent 编排路线图制定
## 文件索引
@@ -12,11 +12,12 @@
| [`roadmap-v0.3.0.md`](./roadmap-v0.3.0.md) | v0.3.0 计划与交付 - Phase 1319 | ✅ Phase 13-19 全部完成,v0.3.0 交付完毕 |
| [`roadmap-v0.3.2.md`](./roadmap-v0.3.2.md) | v0.3.2 计划与交付 — Phase 2027Cargo 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-v0.4.0.md`](./roadmap-v0.4.0.md) | v0.4.0 计划 — Phase A-ESwarm 编排、HITL + Steering、TokenJuice 语义压缩、自动校正) | 📋 计划中 |
| [`roadmap-unsorted.md`](./roadmap-unsorted.md) | 未归到任何版本的内容 — 全局愿景、当前状态、模块完整性、v0.4+ 展望、风险与建议、下一步行动、阶段总回顾 | — |
## 阅读建议
- **按版本顺序追溯历史**v0.1.0 → v0.2.0 → v0.3.0
- **按版本顺序追溯历史**v0.1.0 → v0.2.0 → v0.3.0 → v0.4.0
- **了解产品演进全貌**:从 `roadmap-unsorted.md` 顶部开始读
- **查找特定 Phase**:每个版本文件内按 Phase 编号顺序排列
- **了解项目当前关注点**:从 `roadmap-unsorted.md` 的「下一步行动」开始
+2
View File
@@ -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 => {
+328 -10
View File
@@ -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;
@@ -13,7 +14,7 @@ use bytes::Bytes;
use futures_core::Stream;
use futures_util::StreamExt;
use reqwest::Client;
use reqwest::header::{HeaderMap, HeaderValue};
use reqwest::header::{HeaderMap, HeaderName, HeaderValue};
use serde::{Deserialize, Serialize};
use serde_json::{Value, json};
use tracing::{debug, error, info, warn};
@@ -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,17 @@ impl AnthropicProvider {
api_key,
model,
timeout_secs,
extra_headers,
}
}
/// 设置 Provider 级别固定头,替换已有的 extra_headers(如有)。
/// 返回 self 以支持链式调用。如需追加语义,在外部自行 `extend`。
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 +227,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 +268,54 @@ 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
/// 后者覆盖前者。非法 header 名/值(如控制字符)静默跳过 + warn,避免 reqwest panic。
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 {
if let (Ok(name), Ok(value)) = (
HeaderName::from_bytes(k.as_bytes()),
HeaderValue::from_str(v),
) {
builder = builder.header(name, value);
} else {
warn!(header = %k, "skipping invalid extra_header (key or value contains illegal characters)");
}
}
for (key, value) in &body.custom_headers {
if let (Ok(name), Ok(value)) = (
HeaderName::from_bytes(key.as_bytes()),
HeaderValue::from_str(value),
) {
builder = builder.header(name, value);
} else {
warn!(header = %key, "skipping invalid custom_header (key or value contains illegal characters)");
}
}
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 +343,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 +496,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 +1258,271 @@ 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:?}"
);
}
#[test]
fn anthropic_extra_headers_from_constructor_unit() {
// ponytail: 用 RequestBuilder::build() 直检 headers,无需 wiremock。
let client = Client::builder()
.timeout(Duration::from_secs(30))
.build()
.expect("create http client");
let provider = AnthropicProvider::from_parts(
"http://x".into(),
"sk-ant-test".into(),
"claude-sonnet-4-20250514".into(),
client,
30,
vec![("X-Platform".into(), "anthropic-test".into())],
);
let body = AnthropicRequestBody {
model: "claude-sonnet-4-20250514".into(),
max_tokens: 4096,
system: None,
messages: Vec::new(),
tools: None,
thinking: None,
stream: None,
custom_headers: HashMap::new(),
};
let req = provider
.build_request_builder(&body)
.unwrap()
.build()
.unwrap();
assert_eq!(req.headers().get("X-Platform").unwrap(), "anthropic-test");
}
#[test]
fn anthropic_invalid_header_name_is_skipped() {
// ponytail: extra_headers 含非法 key 应被跳过,不应让 reqwest panic。
let client = Client::builder()
.timeout(Duration::from_secs(30))
.build()
.expect("create http client");
let provider = AnthropicProvider::from_parts(
"http://x".into(),
"sk-ant-test".into(),
"claude-sonnet-4-20250514".into(),
client,
30,
Vec::new(),
)
.with_extra_headers(vec![
("X-Valid".into(), "v1".into()),
("bad\nname".into(), "v2".into()),
]);
let body = AnthropicRequestBody {
model: "claude-sonnet-4-20250514".into(),
max_tokens: 4096,
system: None,
messages: Vec::new(),
tools: None,
thinking: None,
stream: None,
custom_headers: HashMap::new(),
};
let req = provider
.build_request_builder(&body)
.unwrap()
.build()
.unwrap();
assert_eq!(req.headers().get("X-Valid").unwrap(), "v1");
assert!(
req.headers().get("bad\nname").is_none(),
"非法 header 名应被静默跳过"
);
}
}
+276 -3
View File
@@ -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;
@@ -17,9 +18,10 @@ use bytes::Bytes;
use futures_core::Stream;
use futures_util::StreamExt;
use reqwest::Client;
use reqwest::header::{HeaderName, HeaderValue};
use serde::Serialize;
use serde_json::Value;
use tracing::{debug, error, info};
use tracing::{debug, error, info, warn};
use crate::llm::convert::{from_openai, to_openai};
use crate::llm::error::LlmError;
@@ -145,6 +147,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 +444,13 @@ impl GenericOpenaiProvider {
)
}
/// 设置 Provider 级别固定头,替换已有的 extra_headers(如有)。
/// 返回 self 以支持链式调用。如需追加语义,在外部自行 `extend`。
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,17 +465,37 @@ impl GenericOpenaiProvider {
}
/// 构造 HTTP POST 请求 builder(含认证头与额外请求头)。
///
/// 头融合顺序:Authorization → Provider 级 extra_headers → 请求级 custom_headers
/// 后者覆盖前者。非法 header 名/值(如控制字符)静默跳过 + warn,避免 reqwest panic。
fn build_request_builder(
&self,
url: &str,
body: &impl Serialize,
body: &OpenaiChatRequest,
) -> Result<reqwest::RequestBuilder, LlmError> {
let mut builder = self
.http_client
.post(url)
.header("Authorization", format!("Bearer {}", self.api_key));
for (k, v) in &self.extra_headers {
builder = builder.header(k.as_str(), v.as_str());
if let (Ok(name), Ok(value)) = (
HeaderName::from_bytes(k.as_bytes()),
HeaderValue::from_str(v),
) {
builder = builder.header(name, value);
} else {
warn!(header = %k, "skipping invalid extra_header (key or value contains illegal characters)");
}
}
for (key, value) in &body.custom_headers {
if let (Ok(name), Ok(value)) = (
HeaderName::from_bytes(key.as_bytes()),
HeaderValue::from_str(value),
) {
builder = builder.header(name, value);
} else {
warn!(header = %key, "skipping invalid custom_header (key or value contains illegal characters)");
}
}
Ok(builder.json(body))
}
@@ -557,6 +591,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 +609,7 @@ impl GenericOpenaiProvider {
seed,
response_format,
parallel_tool_calls,
custom_headers,
..Default::default()
})
}
@@ -1695,4 +1732,240 @@ 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:?}"
);
}
#[test]
fn openai_chat_extra_headers_from_constructor_unit() {
// 验证 `from_parts` 传入的 extra_headers 在 build_request_builder 中被注入。
// ponytail: 用 RequestBuilder::build() 直检 headers,无需 wiremock。
let provider = GenericOpenaiProvider::from_parts(
"http://x".into(),
"sk-test".into(),
"gpt-4o".into(),
"openai",
Client::builder()
.timeout(Duration::from_secs(30))
.build()
.expect("create http client"),
vec![("X-Platform".into(), "doubao".into())],
30,
);
let body = OpenaiChatRequest {
model: "gpt-4o".into(),
..Default::default()
};
let req = provider
.build_request_builder("http://x/chat/completions", &body)
.unwrap()
.build()
.unwrap();
assert_eq!(req.headers().get("X-Platform").unwrap(), "doubao");
}
#[test]
fn openai_chat_invalid_header_name_is_skipped() {
// ponytail: extra_headers 含非法 key(如换行符)应被跳过,不应让 reqwest panic。
let provider = make_provider_for_header_tests("http://x".into()).with_extra_headers(vec![
("X-Valid".into(), "v1".into()),
("bad\nname".into(), "v2".into()),
]);
let body = OpenaiChatRequest {
model: "gpt-4o".into(),
..Default::default()
};
let req = provider
.build_request_builder("http://x/chat/completions", &body)
.unwrap()
.build()
.unwrap();
assert_eq!(req.headers().get("X-Valid").unwrap(), "v1");
assert!(
req.headers().get("bad\nname").is_none(),
"非法 header 名应被静默跳过"
);
}
}
+560 -26
View File
@@ -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;
@@ -18,6 +19,7 @@ use bytes::Bytes;
use futures_core::Stream;
use futures_util::StreamExt;
use reqwest::Client;
use reqwest::header::{HeaderName, HeaderValue};
use serde::{Deserialize, Serialize};
use serde_json::{Value, json};
use tracing::{debug, error, info, warn};
@@ -71,6 +73,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 枚举。
@@ -128,15 +134,68 @@ pub(crate) enum ResponseInputContent {
}
/// Tool 定义。
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
///
/// NOTE: 仅在请求序列化路径使用(convert_request → build_request_builder → HTTP body)。
/// 响应反序列化走 ResponseOutputItem,不经过此类型。
/// 自定义 Deserialize 服务于 convert_request 内 extra 字段反序列化。
#[derive(Debug, Clone)]
pub(crate) enum ResponseTool {
#[serde(rename = "function")]
Function {
name: String,
description: String,
parameters: Value,
},
/// 非 function 类型的工具(如 web_search / file_search / code_interpreter)。
/// 直接透传原始 JSON Value,不做结构化解析,避免信息丢失。
Builtin(Value),
}
impl Serialize for ResponseTool {
fn serialize<S: serde::Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
match self {
ResponseTool::Function {
name,
description,
parameters,
} => {
use serde::ser::SerializeMap;
let mut map = serializer.serialize_map(Some(4))?;
map.serialize_entry("type", "function")?;
map.serialize_entry("name", name)?;
map.serialize_entry("description", description)?;
map.serialize_entry("parameters", parameters)?;
map.end()
}
ResponseTool::Builtin(value) => value.serialize(serializer),
}
}
}
impl<'de> Deserialize<'de> for ResponseTool {
fn deserialize<D: serde::Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
let value = Value::deserialize(deserializer)?;
match value.get("type").and_then(|t| t.as_str()) {
Some("function") => {
let name = value
.get("name")
.and_then(|n| n.as_str())
.unwrap_or_default()
.to_string();
let description = value
.get("description")
.and_then(|d| d.as_str())
.unwrap_or_default()
.to_string();
let parameters = value.get("parameters").cloned().unwrap_or(Value::Null);
Ok(ResponseTool::Function {
name,
description,
parameters,
})
}
_ => Ok(ResponseTool::Builtin(value)),
}
}
}
/// OpenAI Response API 响应体。
@@ -292,6 +351,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 +363,7 @@ impl OpenaiResponseProvider {
model: String,
http_client: Client,
timeout_secs: u64,
extra_headers: Vec<(String, String)>,
) -> Self {
Self {
http_client,
@@ -309,22 +371,57 @@ impl OpenaiResponseProvider {
api_key,
model,
timeout_secs,
extra_headers,
}
}
/// 设置 Provider 级别固定头,替换已有的 extra_headers(如有)。
/// 返回 self 以支持链式调用。如需追加语义,在外部自行 `extend`。
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('/'))
}
/// 构造 HTTP POST 请求 builder(含认证头与额外请求头)。
///
/// 头融合顺序:Authorization → Provider 级 extra_headers → 请求级 custom_headers
/// 后者覆盖前者。非法 header 名/值(如控制字符)静默跳过 + warn,避免 reqwest panic。
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 {
if let (Ok(name), Ok(value)) = (
HeaderName::from_bytes(k.as_bytes()),
HeaderValue::from_str(v),
) {
builder = builder.header(name, value);
} else {
warn!(header = %k, "skipping invalid extra_header (key or value contains illegal characters)");
}
}
for (key, value) in &body.custom_headers {
if let (Ok(name), Ok(value)) = (
HeaderName::from_bytes(key.as_bytes()),
HeaderValue::from_str(value),
) {
builder = builder.header(name, value);
} else {
warn!(header = %key, "skipping invalid custom_header (key or value contains illegal characters)");
}
}
Ok(builder.json(body))
}
fn map_reqwest_error(&self, e: reqwest::Error) -> LlmError {
@@ -403,6 +500,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 {
@@ -560,31 +665,41 @@ impl OpenaiResponseProvider {
Some(instructions_parts.join("\n"))
};
let tools = if tools_defs.is_empty() {
None
} else {
let mut items: Vec<ResponseTool> = tools_defs
.into_iter()
.map(|t| ResponseTool::Function {
let tools = {
let mut items: Vec<ResponseTool> = Vec::new();
// 1. 自定义工具 → Function 变体
for t in tools_defs {
items.push(ResponseTool::Function {
name: t.name,
description: t.description.unwrap_or_default(),
parameters: t.parameters,
})
.collect();
// ponytail: 内置工具(web_search / file_search)通过 extra 逃生舱追加到 tools 数组。
});
}
// 2. 内置工具 → Builtin 变体(与 tools_defs 解耦)
// 调用方使用 `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,
if let Some(extra_tools) = builtin_tools {
for v in extra_tools {
match serde_json::from_value(v.clone()) {
Ok(tool) => items.push(tool),
// ponytail: ResponseTool::Deserialize 对任何 Value 都返回 Ok
// "function" 分支抽字段,其余走 Builtin),因此 Err 分支实际不可达。
// 防御性保留:未来若 ResponseTool 增加严格校验逻辑,此分支即可激活。
Err(e) => {
let raw = serde_json::to_string(&v).unwrap_or_default();
warn!(tool = %raw, error = %e, "skipped invalid builtin_tool");
}
}));
}
}
}
Some(items)
// 3. 两者都空 → None;否则 → Some
if items.is_empty() {
None
} else {
Some(items)
}
};
// ponytail: 顶层 `text.format` 通过 extra 逃生舱透传 —— 整体结构化为 value 后塞入。
@@ -611,6 +726,7 @@ impl OpenaiResponseProvider {
truncation,
metadata,
reasoning,
custom_headers,
})
}
@@ -1020,7 +1136,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 +2098,415 @@ 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:?}"
);
}
#[test]
fn openai_response_extra_headers_from_constructor_unit() {
// ponytail: 用 RequestBuilder::build() 直检 headers,无需 wiremock。
let client = Client::builder()
.timeout(Duration::from_secs(30))
.build()
.expect("create http client");
let provider = OpenaiResponseProvider::from_parts(
"http://x".into(),
"sk-test".into(),
"gpt-4o".into(),
client,
30,
vec![("X-Platform".into(), "doubao".into())],
);
let body = OpenaiResponseRequest {
model: "gpt-4o".into(),
instructions: None,
input: Vec::new(),
tools: None,
tool_choice: None,
text: None,
max_output_tokens: None,
temperature: None,
top_p: None,
stop: None,
stream: None,
previous_response_id: None,
store: None,
truncation: None,
metadata: None,
reasoning: None,
custom_headers: HashMap::new(),
};
let req = provider
.build_request_builder(&body)
.unwrap()
.build()
.unwrap();
assert_eq!(req.headers().get("X-Platform").unwrap(), "doubao");
}
#[test]
fn openai_response_invalid_header_name_is_skipped() {
// ponytail: extra_headers 含非法 key 应被跳过,不应让 reqwest panic。
let client = Client::builder()
.timeout(Duration::from_secs(30))
.build()
.expect("create http client");
let provider = OpenaiResponseProvider::from_parts(
"http://x".into(),
"sk-test".into(),
"gpt-4o".into(),
client,
30,
Vec::new(),
)
.with_extra_headers(vec![
("X-Valid".into(), "v1".into()),
("bad\nname".into(), "v2".into()),
]);
let body = OpenaiResponseRequest {
model: "gpt-4o".into(),
instructions: None,
input: Vec::new(),
tools: None,
tool_choice: None,
text: None,
max_output_tokens: None,
temperature: None,
top_p: None,
stop: None,
stream: None,
previous_response_id: None,
store: None,
truncation: None,
metadata: None,
reasoning: None,
custom_headers: HashMap::new(),
};
let req = provider
.build_request_builder(&body)
.unwrap()
.build()
.unwrap();
assert_eq!(req.headers().get("X-Valid").unwrap(), "v1");
assert!(
req.headers().get("bad\nname").is_none(),
"非法 header 名应被静默跳过"
);
}
// ===== builtin_tools 注入测试(Phase 30 =====
/// T1: 纯内置工具场景 —— tools_defs 为空、builtin_tools 非空。
#[test]
fn convert_request_builtin_tools_only() {
let provider = make_provider("http://x".into());
let mut req = MessageRequest {
model: "gpt-4o".into(),
messages: vec![Message::user_text("search web")],
tools: vec![],
..Default::default()
};
req.set_extra(
"builtin_tools",
json!([{"type": "web_search", "search_context_size": "medium"}]),
);
let body = provider.convert_request(req).unwrap();
assert!(body.tools.is_some(), "tools should be Some with builtin tools");
let tools = body.tools.as_ref().unwrap();
assert_eq!(tools.len(), 1);
// 验证 wire 序列化后 type 字段正确
let wire = serde_json::to_value(tools).unwrap();
assert_eq!(wire[0]["type"], "web_search");
assert_eq!(wire[0]["search_context_size"], "medium");
}
/// T2: 混用场景 —— tools_defs + builtin_tools 同时非空。
#[test]
fn convert_request_mixed_tools() {
let provider = make_provider("http://x".into());
let mut req = MessageRequest {
model: "gpt-4o".into(),
messages: vec![Message::user_text("compute and search")],
tools: vec![ToolDef {
name: "my_func".into(),
description: Some("a custom function".into()),
parameters: json!({"type": "object"}),
}],
..Default::default()
};
req.set_extra("builtin_tools", json!([{"type": "code_interpreter"}]));
let body = provider.convert_request(req).unwrap();
let tools = body.tools.as_ref().unwrap();
assert_eq!(tools.len(), 2);
let wire = serde_json::to_value(tools).unwrap();
assert_eq!(wire[0]["type"], "function");
assert_eq!(wire[0]["name"], "my_func");
assert_eq!(wire[1]["type"], "code_interpreter");
}
/// T3: 两端空 —— 回归测试。
#[test]
fn convert_request_no_tools_at_all() {
let provider = make_provider("http://x".into());
let req = MessageRequest {
model: "gpt-4o".into(),
messages: vec![Message::user_text("hi")],
tools: vec![],
..Default::default()
};
let body = provider.convert_request(req).unwrap();
assert!(body.tools.is_none(), "tools should be None when no tools at all");
}
/// T4: 多个有效 builtin_tools 都被正确保留。
///
/// 设计说明:自定义 Deserialize 对任何 Value 都返回 OkFunction 或 Builtin),
/// 因此 `from_value::<ResponseTool>` 实际不会失败。本测试验证多个有效 builtin_tools
/// 都能被正确处理并保留。
#[test]
fn convert_request_valid_builtin_tools_retained() {
let provider = make_provider("http://x".into());
let mut req = MessageRequest {
model: "gpt-4o".into(),
messages: vec![Message::user_text("hi")],
tools: vec![],
..Default::default()
};
// 传入多个有效 builtin tool 值(含一个看似"非 function"的字符串值,
// 会被反序列化为 Builtin 变体 —— 这是设计预期)
req.set_extra(
"builtin_tools",
json!([
{"type": "web_search"},
{"type": "file_search", "max_num_results": 5}
]),
);
let body = provider.convert_request(req).unwrap();
let tools = body.tools.as_ref().unwrap();
assert_eq!(tools.len(), 2, "both builtin_tools should be retained");
}
/// T5: ResponseTool::Function 序列化后的 JSON 结构(AC7 wire format 回归)。
///
/// 注:serde_json::Map 默认是 BTreeMap,序列化后字段按字母序排列;
/// JSON wire 格式不要求字段顺序,OpenAI API 不会拒绝任意字段顺序。
#[test]
fn test_function_wire_format() {
let tool = ResponseTool::Function {
name: "search".into(),
description: "search docs".into(),
parameters: json!({"type": "object"}),
};
let wire = serde_json::to_value(&tool).unwrap();
// 验证字段值正确(顺序无关)
assert_eq!(wire["type"], "function");
assert_eq!(wire["name"], "search");
assert_eq!(wire["description"], "search docs");
assert_eq!(wire["parameters"], json!({"type": "object"}));
// 验证字段集合完整(无缺失、无多余)
let obj = wire.as_object().unwrap();
let mut keys: Vec<&str> = obj.keys().map(|s| s.as_str()).collect();
keys.sort();
assert_eq!(keys, vec!["description", "name", "parameters", "type"]);
}
/// T6: ResponseTool::Builtin 反序列化 + 序列化 roundtrip —— 原始 JSON 结构保留。
#[test]
fn test_builtin_roundtrip() {
let original = json!({
"type": "web_search",
"search_context_size": "high",
"user_location": {"country": "US"}
});
let tool: ResponseTool = serde_json::from_value(original.clone()).unwrap();
// 验证反序列化为 Builtin 变体
match &tool {
ResponseTool::Builtin(value) => {
assert_eq!(value["type"], "web_search");
assert_eq!(value["user_location"]["country"], "US");
}
_ => panic!("expected Builtin variant"),
}
// 验证序列化后原始结构完整保留
let wire = serde_json::to_value(&tool).unwrap();
assert_eq!(wire, original);
}
}