Compare commits
2
Commits
358e971094
...
212cfcc916
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
212cfcc916 | ||
|
|
88d00ac927 |
@@ -0,0 +1,821 @@
|
||||
# Phase 9 — 流式体验增强实施方案
|
||||
|
||||
- **文档编号**:16
|
||||
- **标题**:Phase 9 — 流式体验增强实施方案
|
||||
- **日期**:2026-07-05
|
||||
- **状态**:已定稿
|
||||
- **涉及模块**:llm/cycle、llm/types/response_v2、agent/session
|
||||
- **关联文档**:roadmap.md(§Phase 9)、15-phase8-mvp-integration.md
|
||||
|
||||
---
|
||||
|
||||
## 1. 背景与目标
|
||||
|
||||
agcore 已发布 v0.2.0-rc.1,Phase 0-8 全部完成。当前 Agent 会话只有非流式 API(`submit_turn`),开发者无法看到实时 token 输出和工具执行过程。Phase 9 的目标是为 `AgentSession` 新增流式方法 `submit_turn_stream`,让开发者能实时看到 LLM token 生成和工具执行状态。
|
||||
|
||||
### 1.1 现有能力
|
||||
|
||||
| 能力 | 方法 | 流式 | 自动工具循环 | 状态 |
|
||||
|------|------|------|-------------|------|
|
||||
| LLM 流式请求 | `LlmCycle::submit_stream` | ✅ | ❌ | 已就绪 |
|
||||
| LLM 工具循环 | `LlmCycle::submit_with_tools` | ❌ | ✅ | 已就绪 |
|
||||
| Agent 会话 | `AgentSession::submit_turn` | ❌ | ✅ | 已就绪 |
|
||||
| 流事件枚举 | `StreamEvent`(11 变体) | — | — | 缺工具执行事件 |
|
||||
| Mock 流 | `MockProvider::chat_stream` | ✅ | — | 可模拟流事件序列 |
|
||||
|
||||
### 1.2 核心矛盾
|
||||
|
||||
流式能力和工具循环能力分别存在于两个方法中,从未被组合。`submit_stream` 只管将 LLM 流事件原样转发,不理解工具调用;`submit_with_tools` 自动执行工具循环但全程阻塞。Phase 9 就是要组合它们:**在工具循环中,每一轮 LLM 调用都是流式的,并在工具执行前后插入语义事件**。
|
||||
|
||||
---
|
||||
|
||||
## 2. 需求分析
|
||||
|
||||
### 2.1 功能需求
|
||||
|
||||
1. **`AgentSession::submit_turn_stream(user_input)`** — 返回 `StreamEvent` 流,开发者通过 `while let Some(event) = stream.next().await` 逐事件消费
|
||||
2. **流式工具循环** — 多轮工具调用过程中流不卡死,每轮工具执行前后插入 `ToolExecutionStarted` / `ToolExecutionCompleted` 事件
|
||||
3. **`finalize_turn(response)`** — 流消费完成后同步 session 状态(cost 累计 + `OnTurnEnd` hook 触发)
|
||||
4. **新增 `StreamEvent` 变体** — `ToolExecutionStarted` + `ToolExecutionCompleted`,携带工具名称、调用 ID、参数/结果摘要
|
||||
|
||||
### 2.2 非功能需求
|
||||
|
||||
- **零影响**:现有 `submit_turn` 和 `submit_with_tools` 行为不变,存量测试 0 回归
|
||||
- **异步流**:消费者通过 `futures_util::StreamExt::next()` 逐事件消费
|
||||
- **错误事件化**:错误通过 `StreamEvent::Error` 事件表达,不通过 `Result` 通道终止流
|
||||
- **最少代码**:复用现有 `submit_with_tools` 的工具循环逻辑模式和 `submit_stream` 的流管道模式
|
||||
|
||||
### 2.3 不做事项
|
||||
|
||||
| 事项 | 理由 |
|
||||
|------|------|
|
||||
| 新增示例(Phase 9.2 再加) | 缩窄 Phase 9 范围至核心能力 |
|
||||
| `OnTurnEnd` 自动触发 | Rust 所有权约束:流是延迟求值,`&mut self` 无法进入闭包;由消费者收到 `MessageComplete` 后手动调用 `finalize_turn` |
|
||||
| 修复 cost 统计 | 中间轮 cost 丢失是已知限制,与 `submit_turn` 行为一致 |
|
||||
| 跨 turn 消息历史保留 | Phase 10 `ContextSlot` 的职责 |
|
||||
| 并行 tool 调用的事件细化 | 当前工具调用是顺序 `for` 循环,并行化留待后续优化 |
|
||||
| `run_tool_loop` 内消息压缩 | `run_tool_loop` 不接收 `compact_config` 参数,不执行上下文压缩。长工具循环中消息增长可能导致 context window 溢出,这是流式实现的已知限制。后续可通过传递 `compact_config` 给 `run_tool_loop` 支持 |
|
||||
| LLM 请求自动 retry | 流式版本不在 `run_tool_loop` 内部实现 retry(详见 §3.6 说明)。调用方可自行包装 `RetryProvider` 或在 `LlmProvider` 实现层处理 |
|
||||
|
||||
---
|
||||
|
||||
## 3. 方案设计
|
||||
|
||||
### 3.1 架构总览
|
||||
|
||||
```
|
||||
┌──────────────────────────────────────────────────────────────┐
|
||||
│ AgentSession │
|
||||
│ ┌──────────────────────────────────────────────────────┐ │
|
||||
│ │ submit_turn_stream() │ │
|
||||
│ │ ├─ OnTurnStart hook(同步触发,返回流之前) │ │
|
||||
│ │ ├─ 组装 LlmCycle(system_prompt / compact_config) │ │
|
||||
│ │ ├─ 调用 submit_with_tools_stream() │ │
|
||||
│ │ ├─ turn_index += 1 │ │
|
||||
│ │ └─ 返回流 │ │
|
||||
│ │ │ │
|
||||
│ │ finalize_turn(response) │ │
|
||||
│ │ ├─ cost_so_far.add(&response.usage) │ │
|
||||
│ │ └─ OnTurnEnd hook(turn_index - 1) │ │
|
||||
│ └──────────────────────────────────────────────────────┘ │
|
||||
│ submit_with_tools_stream(prompt, Arc<ToolRegistry>)
|
||||
▼
|
||||
┌──────────────────────────────────────────────────────────────┐
|
||||
│ LlmCycle (tokio::spawn task — run_tool_loop 状态机) │
|
||||
│ │
|
||||
│ max_turns = max_tool_turns.unwrap_or(10) │
|
||||
│ for round in 1..=max_turns { │
|
||||
│ ① build_request(messages, tools) │
|
||||
│ ② provider.chat_stream(request).await │
|
||||
│ 匹配 Err → tx.send(Error{..}) + return(不 panic) │
|
||||
│ ③ 消费 LLM 流,所有事件 → mpsc unbounded tx(全量转发) │
|
||||
│ ④ partial.finalize() → MessageResponse │
|
||||
│ ⑤ if has_tool_use: │
|
||||
│ ├─ tx → ToolExecutionStarted { tool_name, id, args } │
|
||||
│ ├─ registry.invoke_all(calls, timeout).await │
|
||||
│ ├─ for result: tx → ToolExecutionCompleted { ... } │
|
||||
│ ├─ push tool results → messages │
|
||||
│ └─ continue(新一轮) │
|
||||
│ else: break(最终轮,已发出 MessageComplete) │
|
||||
│ } │
|
||||
│ │
|
||||
│ 产出事件序列(通过 mpsc::unbounded_channel): │
|
||||
│ MessageStart → ... → ToolCallEnd → ToolExecutionStarted → │
|
||||
│ ToolExecutionCompleted → MessageStart → TextDelta → ... → │
|
||||
│ CostUpdate → MessageComplete │
|
||||
└──────────────────────────────────────────────────────────────┘
|
||||
```
|
||||
|
||||
> **`max_tool_turns` 语义**:与非流式 `submit_with_tools` 一致——`None` 退化为 `10`(`unwrap_or(10)`)。默认值 `Some(10)` 已提供安全上限;如需增大限制,手动设置为 `Some(N)`。⚠️ 生产环境建议始终设有限值防止无限循环。
|
||||
|
||||
### 3.2 事件序列约定
|
||||
|
||||
**纯文本流**(无 tool_use):
|
||||
|
||||
```
|
||||
MessageStart → ContentBlockStart → TextDelta* → ContentBlockEnd → CostUpdate → MessageComplete
|
||||
```
|
||||
|
||||
**单轮工具调用**:
|
||||
|
||||
```
|
||||
MessageStart → ContentBlockStart → TextDelta* → ContentBlockEnd
|
||||
→ ContentBlockStart → ToolCallArgumentsDelta* → ToolCallEnd
|
||||
→ CostUpdate → MessageComplete { stop_reason: ToolUse }
|
||||
→ ToolExecutionStarted → [工具执行] → ToolExecutionCompleted
|
||||
→ ContentBlockStart → TextDelta* → ContentBlockEnd
|
||||
→ CostUpdate → MessageComplete { stop_reason: Stop }
|
||||
```
|
||||
|
||||
**多轮工具调用**:
|
||||
|
||||
```
|
||||
... → ToolExecutionCompleted(第 1 轮)
|
||||
→ ToolCallArgumentsDelta* → ToolCallEnd
|
||||
→ ToolExecutionStarted → ToolExecutionCompleted(第 2 轮)
|
||||
→ ... → CostUpdate → MessageComplete(最终轮)
|
||||
```
|
||||
|
||||
**工具不可恢复错误**:
|
||||
|
||||
```
|
||||
... → ToolCallEnd → ToolExecutionStarted
|
||||
→ Error { "tool 'search' 不可恢复错误: ..." } → MessageComplete
|
||||
```
|
||||
|
||||
> **`MessageComplete.full_response` 内容范围**:每轮 LLM 调用独立产生一个 `MessageComplete`,其中 `full_response` 仅包含**该轮 LLM 的单个响应**(不累积前面工具轮次的结果)。中间轮(`stop_reason: ToolUse`)的 `full_response` 通常只包含 `ToolUse` block,无文本。最终轮(`stop_reason: Stop`)的 `full_response` 包含 LLM 的最终输出。消费者如需追踪完整对话历史,应自行累加所有轮次的 `Message`。
|
||||
|
||||
### 3.3 StreamEvent 新增变体
|
||||
|
||||
在 `src/llm/types/response_v2.rs` 的 `StreamEvent` 枚举中追加两个变体:
|
||||
|
||||
```rust
|
||||
/// 工具开始执行 —— 在 ToolCallEnd 之后、registry.invoke 之前发出。
|
||||
/// 让 UI 层可以显示 "正在执行工具:add(1, 2)"。
|
||||
ToolExecutionStarted {
|
||||
tool_name: String,
|
||||
tool_call_id: String,
|
||||
/// 工具参数(JSON 字符串形式)
|
||||
arguments: String,
|
||||
},
|
||||
|
||||
/// 工具执行完成 —— 在工具返回后、新一轮 LLM 流开始之前发出。
|
||||
ToolExecutionCompleted {
|
||||
tool_name: String,
|
||||
tool_call_id: String,
|
||||
/// 结果摘要(前 200 字符)
|
||||
result_summary: String,
|
||||
/// 是否出错
|
||||
is_error: bool,
|
||||
},
|
||||
```
|
||||
|
||||
在 `PartialMessageResponse::apply_to` 中追加:
|
||||
|
||||
```rust
|
||||
StreamEvent::ToolExecutionStarted { .. } | StreamEvent::ToolExecutionCompleted { .. } => true,
|
||||
```
|
||||
|
||||
这两个是**元事件**,不参与内容块累积,`apply_to` 直接返回 `true`。
|
||||
|
||||
### 3.4 新增方法签名
|
||||
|
||||
**`LlmCycle` 层**(`src/llm/cycle.rs`):
|
||||
|
||||
```rust
|
||||
/// 提交消息并自动处理工具调用循环,流式产出所有事件。
|
||||
///
|
||||
/// 与 `submit_with_tools` 的区别:
|
||||
/// - LLM 响应是流式的(全程 `chat_stream` 而非 `chat`)
|
||||
/// - 工具执行前后插入 `ToolExecutionStarted` / `ToolExecutionCompleted` 事件
|
||||
/// - 错误以 `StreamEvent::Error` 形式出现在流中,而非终止 `Result`
|
||||
/// - 消费方需手动 `push_message()` 同步消息历史
|
||||
///
|
||||
/// **运行时要求**:内部使用 `tokio::spawn`,需要 tokio 多线程运行时。
|
||||
pub async fn submit_with_tools_stream(
|
||||
&mut self,
|
||||
prompt: String,
|
||||
tool_registry: Arc<ToolRegistry>,
|
||||
) -> Result<Pin<Box<dyn Stream<Item = StreamEvent> + Send>>, LlmError>
|
||||
```
|
||||
|
||||
```rust
|
||||
/// 运行工具循环的核心异步状态机。
|
||||
///
|
||||
/// 接收 owned 字段,通过 mpsc::unbounded_channel 产出事件序列。
|
||||
/// 由 `submit_with_tools_stream` 在 tokio::spawn 中调用。
|
||||
///
|
||||
/// **运行时要求**:此函数内部使用 `tokio::spawn`,要求调用方运行在
|
||||
/// tokio 多线程运行时中(`#[tokio::main]` 或 `#[tokio::test(flavor = "multi_thread")]`)。
|
||||
/// 不在 WASM 目标下可用。
|
||||
async fn run_tool_loop(
|
||||
messages: Vec<Message>,
|
||||
provider: Arc<dyn LlmProvider>,
|
||||
config: CycleConfig,
|
||||
tool_registry: Arc<ToolRegistry>,
|
||||
tools: Vec<ToolDef>,
|
||||
tx: mpsc::UnboundedSender<StreamEvent>,
|
||||
hook_executor: Option<Arc<HookExecutor>>,
|
||||
)
|
||||
```
|
||||
|
||||
**`AgentSession` 层**(`src/agent/session.rs`):
|
||||
|
||||
```rust
|
||||
/// 提交一轮对话(流式,含自动 tool 循环),返回 `StreamEvent` 流。
|
||||
///
|
||||
/// 与 `submit_turn` 的区别:
|
||||
/// - 以流事件序列而非 `MessageResponse` 返回
|
||||
/// - 工具执行期间插入 `ToolExecutionStarted` / `ToolExecutionCompleted` 事件
|
||||
/// - 消费方在收到 `MessageComplete` 后需手动调用 `finalize_turn` 同步状态
|
||||
///
|
||||
/// **运行时要求**:内部委托 `submit_with_tools_stream`,需要 tokio 多线程运行时。
|
||||
pub async fn submit_turn_stream(
|
||||
&mut self,
|
||||
user_input: impl Into<String>,
|
||||
) -> Result<Pin<Box<dyn Stream<Item = StreamEvent> + Send>>, AgentError>
|
||||
|
||||
/// 完成一轮 turn:累计 cost + 触发 OnTurnEnd hook。
|
||||
///
|
||||
/// 由消费者在收到 `MessageComplete.full_response` 后调用。
|
||||
pub async fn finalize_turn(&mut self, response: &MessageResponse)
|
||||
```
|
||||
|
||||
### 3.5 消费者使用模式
|
||||
|
||||
```rust
|
||||
use futures_util::StreamExt;
|
||||
|
||||
let mut stream = session.submit_turn_stream("计算 1+2").await?;
|
||||
|
||||
let mut final_response = None;
|
||||
while let Some(event) = stream.next().await {
|
||||
match &event {
|
||||
StreamEvent::TextDelta { text } => print!("{}", text),
|
||||
StreamEvent::ToolExecutionStarted { tool_name, arguments, .. } => {
|
||||
println!("\n🔧 [{}({})]", tool_name, arguments);
|
||||
}
|
||||
StreamEvent::ToolExecutionCompleted { result_summary, .. } => {
|
||||
println!(" → {}", result_summary);
|
||||
}
|
||||
StreamEvent::MessageComplete { full_response } => {
|
||||
final_response = Some(full_response.clone());
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
std::io::stdout().flush().ok();
|
||||
|
||||
if let Some(response) = final_response {
|
||||
session.finalize_turn(&response).await;
|
||||
}
|
||||
```
|
||||
|
||||
> **⚠️ 消费者注意**:`finalize_turn` 是开发者责任 —— 遗漏调用会导致 `cost_so_far` 不累计、`OnTurnEnd` hook 不触发。session 状态仍然可用,后续 `submit_turn` 也能正常执行,但 cost 信息不完整。`finalize_turn` 无自动补偿机制,建议使用 `Drop` guard 或在 `while` 循环的 `finally` 块中确保调用。
|
||||
|
||||
### 3.6 run_tool_loop 核心逻辑
|
||||
|
||||
`run_tool_loop` 是此方案的核心状态机(约 90 行),其伪代码逻辑如下:
|
||||
|
||||
```
|
||||
1. 接收 owned 字段:messages, provider, config, tool_registry, tools, tx, hook_executor
|
||||
2. max_turns = config.max_tool_turns.unwrap_or(10)
|
||||
// None → 10(退化为默认值),Some(n) → n
|
||||
// 与非流式 submit_with_tools 行为一致
|
||||
3. 工具循环(for round in 1..=max_turns):
|
||||
a. build_request(messages, tools)
|
||||
// 空 tool_registry 时 tools 为空列表,流退化为纯文本流(可安全运行)
|
||||
b. PreRequest hook(如果有 hook_executor)
|
||||
c. 发起流式 LLM 调用:
|
||||
let stream = match provider.chat_stream(request).await {
|
||||
Ok(s) => s,
|
||||
Err(e) => {
|
||||
// 第一层错误:chat_stream 自身失败(网络/认证/限流)
|
||||
// 这里不做 retry:retry 逻辑留给上层循环的 submit_request 模式,
|
||||
// 流式场景中 retry 需重新建立 mpsc 通道,复杂度与收益不匹配
|
||||
tx.send(StreamEvent::Error { message: e.to_string() }).ok();
|
||||
return; // 直接结束 task
|
||||
}
|
||||
};
|
||||
d. 消费 LLM 流:
|
||||
- PartialMessageResponse::new()
|
||||
- while let Some(result) = stream.next().await
|
||||
- match result:
|
||||
Ok(event) → apply_to + tx.send(event)
|
||||
Err(e) → tx.send(Error { message }) + break
|
||||
// 第二层错误:stream 内部事件错误(如 chunk 解析失败)
|
||||
e. partial.finalize()? → response
|
||||
f. push response.message → messages
|
||||
g. 检查 has_tool_calls_in_response(&response)
|
||||
h. 如果没有 tool_use: break(最终轮,流已自然结束)
|
||||
i. 如果有 tool_use:
|
||||
- extract_tool_calls_from_response(&response)
|
||||
- tx.send(ToolExecutionStarted { tool_name, tool_call_id, arguments })
|
||||
- registry.invoke_all(calls, tool_timeout).await
|
||||
- for result in results:
|
||||
tx.send(ToolExecutionCompleted { tool_name, tool_call_id, result_summary, is_error })
|
||||
- push tool results → messages
|
||||
- continue(新一轮 LLM 流)
|
||||
4. 流结束(tokio::spawn 自然退出)
|
||||
```
|
||||
|
||||
> **关于 LLM 请求 retry**:非流式 `submit_with_tools` 内部通过 `submit_request` 的 retry 循环处理临时错误。流式版本 `run_tool_loop` **不在内部实现 retry**。原因:(1)retry 需要重新建立 mpsc 通道和事件流上下文,复杂度与收益不匹配;(2)`unbounded_channel` 已发出的事件无法撤回。如果需要 retry 语义,调用方应在上层做 fallback 策略,或在 `llm provider` 实现层完成 retry(如 `RetryProvider` 包装器)。
|
||||
|
||||
**错误处理**:
|
||||
|
||||
| 场景 | 行为 |
|
||||
|------|------|
|
||||
| LLM 请求失败(`chat_stream` 返回 `Err`) | `tx.send(Error { message })` + `return` 结束 task。**不做 retry**(见上方说明) |
|
||||
| LLM 流内事件错误(stream Item 的 `Err`) | `tx.send(Error { message })` + `break` 结束当轮流,终止循环 |
|
||||
| 可恢复工具错误(`is_recoverable() == true`) | 作为 tool result 回传 LLM,流继续,不出 Error 事件 |
|
||||
| 不可恢复工具错误(`is_recoverable() == false`) | `tx.send(Error { message })` + 终止循环 |
|
||||
| 工具超时(`tokio::time::timeout`) | 视为不可恢复,`tx.send(Error)` + 终止 |
|
||||
| 最大工具循环轮次超限 | `tx.send(Error { "达到最大工具循环轮次" })` + 终止 |
|
||||
| spawn task 内部 panic | 由于 `JoinHandle` 不保存(detached),panic 由 tokio 运行时静默捕获;消费者看到 stream 直接结束(返回 `None`),无 `Error` 事件。建议在 `run_tool_loop` 内部避免 `unwrap()`,所有可失败路径通过 `Result` + `?` 传播 |
|
||||
|
||||
### 3.7 修改文件清单
|
||||
|
||||
| # | 文件 | 改动 | 估算行数 |
|
||||
|---|------|------|---------|
|
||||
| 1 | `llm/types/response_v2.rs` | +2 `StreamEvent` 变体 +2 `apply_to` arm | ~20 |
|
||||
| 2 | `llm/cycle.rs` | +`submit_with_tools_stream` 方法 + `run_tool_loop` 模块函数 | ~140 |
|
||||
| 3 | `llm/cycle.rs` | `CycleConfig` 加 `#[derive(Clone)]` | ~1 |
|
||||
| 4 | `agent/session.rs` | +`submit_turn_stream` + `finalize_turn` | ~70 |
|
||||
| — | **测试**(内联) | 4 个场景测试(纯度本、单轮、多轮、超限) | ~150 |
|
||||
| | **合计** | | **~380** |
|
||||
|
||||
> 注:`RetryConfig` 已标注 `#[derive(Debug, Clone)]`,无需额外修改。
|
||||
|
||||
---
|
||||
|
||||
## 4. 实现计划
|
||||
|
||||
按 5 个 Step 增量实施,每步可独立编译和测试。
|
||||
|
||||
### Step 1 — 基础设施准备
|
||||
|
||||
**目标**:数据层就绪,为流事件新增变体和配置 Clone 奠基。
|
||||
|
||||
**改动**:
|
||||
|
||||
- `llm/types/response_v2.rs`:
|
||||
- `StreamEvent` 枚举追加 `ToolExecutionStarted` / `ToolExecutionCompleted` 变体
|
||||
- `PartialMessageResponse::apply_to` 追加两个新变体的 arm(均返回 `true`)
|
||||
- `llm/cycle.rs`:
|
||||
- `CycleConfig` 加 `#[derive(Clone)]`(所有字段为基础类型 + `RetryConfig`)
|
||||
|
||||
**验证**:`cargo build` 通过
|
||||
|
||||
### Step 2 — `LlmCycle::submit_with_tools_stream` 核心
|
||||
|
||||
**目标**:实现流式工具循环的核心状态机,这是整个 Phase 9 的技术关键。
|
||||
|
||||
**改动**:
|
||||
|
||||
- `llm/cycle.rs`:
|
||||
- 新增 `run_tool_loop()` 模块函数(约 90 行),基于 `mpsc::unbounded_channel` 通信
|
||||
- 新增 `submit_with_tools_stream()` 公开方法,入口参数为 `prompt` + `Arc<ToolRegistry>`
|
||||
- 内部 `tokio::spawn` 启动 `run_tool_loop`,返回 `rx` 端作为 `dyn Stream`
|
||||
|
||||
**验证**:`cargo build` 通过
|
||||
|
||||
### Step 3 — 单元测试(LlmCycle 层)
|
||||
|
||||
**目标**:验证 `submit_with_tools_stream` 在 8 个核心场景下的行为和事件序列正确性(含 §8 Step 3 扩展的工具错误路径)。
|
||||
|
||||
**新增**(`llm/cycle.rs` 内联测试 `#[cfg(test)]`):
|
||||
|
||||
| 场景 | Mock 响应序列 | 验证点 |
|
||||
|------|---------------|--------|
|
||||
| 1 — 纯文本流 | 1 个 text 响应 | 事件序列与 `submit_stream` 一致;无 `ToolExecutionStarted`/`ToolExecutionCompleted` |
|
||||
| 2 — 单轮工具调用 | 2 个响应:tool_use → text | 包含 `ToolExecutionStarted` + `ToolExecutionCompleted`;最终 `stop_reason` 为 `Stop` |
|
||||
| 3 — 多轮工具调用 | 4 个响应:3 × tool_use → 1 × text | 3 对 `ToolExecutionStarted`/`ToolExecutionCompleted`;消息历史长度正确 |
|
||||
| 4 — 最大轮次超限 | 3 个 tool_use 响应,`max_tool_turns: Some(2)` | 流中出现 `Error` 事件;消息历史停在第 2 轮 |
|
||||
|
||||
**验证**:`cargo test` 全部通过
|
||||
|
||||
### Step 4 — `AgentSession` 层包装
|
||||
|
||||
**目标**:为 `AgentSession` 新增流式会话接口,保持与 `submit_turn` 一致的行为语义。
|
||||
|
||||
**改动**:
|
||||
|
||||
- `agent/session.rs`:
|
||||
- `submit_turn_stream(user_input)` — 触发 `OnTurnStart` hook → 组装 `LlmCycle` → 调用 `submit_with_tools_stream` → `turn_index += 1` → 返回流
|
||||
- `finalize_turn(response)` — `cost_so_far.add(&response.usage)` → 触发 `OnTurnEnd` hook
|
||||
|
||||
**验证**:`cargo build` 通过
|
||||
|
||||
### Step 5 — 集成测试 + 扫尾
|
||||
|
||||
**目标**:端到端验证 `submit_turn_stream` + `finalize_turn` 的完整链路,确保零回归。
|
||||
|
||||
**新增**(`agent/session.rs` 内联测试):
|
||||
|
||||
- **场景**:`submit_turn_stream` 跑通 mock provider → 消费流(验证各事件到达) → `finalize_turn` 后 cost 更新正确
|
||||
- **场景**:verify `OnTurnStart` hook 在 `submit_turn_stream` 返回流之前已触发
|
||||
|
||||
**验证**:
|
||||
|
||||
```bash
|
||||
cargo test --all-targets # 全绿,存量测试 0 回归
|
||||
cargo clippy --all-targets -- -D warnings # 0 警告
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 5. 运行细节
|
||||
|
||||
### 5.1 `run_tool_loop` 的 spawn 生命周期
|
||||
|
||||
#### 执行模型:立即执行 vs 惰性流
|
||||
|
||||
`submit_with_tools_stream` 采用 **立即执行** 模型(`tokio::spawn` + `mpsc`),这与 `submit_stream` 的 **惰性执行**(`async_stream::stream!` 宏,消费者首次 `next()` 时才触发 LLM 调用)不同。
|
||||
|
||||
**选择理由**:工具循环是 **不确定轮次的** —— 每个工具执行的结果可能影响后续 LLM 调用。惰性流无法表达这种"边消费边控制"的语义。通过 `tokio::spawn` 将工具循环移到独立 task 中运行,使得:
|
||||
- 消费者可以随时开始消费(不丢失事件)
|
||||
- 工具循环在后台独立运行,不受消费者消费节奏影响
|
||||
- `mpsc::unbounded_channel` 作为事件缓冲区,解耦生产者与消费者
|
||||
|
||||
**对消费者的影响**:`submit_with_tools_stream().await?` 返回时,工具循环可能已经开始执行(事件已开始写入 channel)。消费者应尽快开始 `while let Some(event) = stream.next().await`,避免 channel 缓冲过多事件。如果在返回流后长时间不消费,事件会堆积在 mpsc buffer 中(内存开销,无阻塞风险 —— 见 §6 风险表)。
|
||||
|
||||
#### 生命周期
|
||||
|
||||
```
|
||||
submit_with_tools_stream()
|
||||
│
|
||||
├─ mpsc::unbounded_channel() → (tx, rx)
|
||||
├─ messages.push(user_text(prompt))
|
||||
├─ compact check
|
||||
├─ tokio::spawn(run_tool_loop(messages, provider, config, ..., tx))
|
||||
└─ return Box::pin(rx) as dyn Stream
|
||||
|
||||
[用户消费 stream]
|
||||
└─ while let Some(event) = rx.recv().await { yield event }
|
||||
|
||||
[用户 drop rx / 结束循环]
|
||||
└─ rx 被 drop → tx.send() 返回 Err
|
||||
→ run_tool_loop 检测到 tx.closed()
|
||||
→ break → task 自然终止
|
||||
```
|
||||
|
||||
#### JoinHandle 与 panic 处理
|
||||
|
||||
`run_tool_loop` 的 `JoinHandle` 在 spawn 后**不保存**(detached pattern)。panic 由 tokio 运行时捕获并通过 `tracing::error` 记录:
|
||||
|
||||
```rust
|
||||
// submit_with_tools_stream 内部
|
||||
tokio::spawn(async move {
|
||||
run_tool_loop(..., tx).await;
|
||||
});
|
||||
```
|
||||
|
||||
如果 `run_tool_loop` 内部发生 panic(如 `unwrap()`),tokio 的 `spawn` 会静默吞掉 panic 并终止 task。消费者此时看到 stream 直接返回 `None`,不会收到 `StreamEvent::Error`。实际编码中应避免 `unwrap()`,所有 `Result` 使用 `?` 或 `match` 处理。
|
||||
|
||||
Rx 侧实现 `Stream` trait:使用 `tokio_stream::wrappers::UnboundedReceiverStream` 包装 `mpsc::UnboundedReceiver`,因为 `mpsc::UnboundedReceiver` 本身不实现 `Stream`(`tokio-stream = "0.1"` 已在 `Cargo.toml` 中存在)。
|
||||
|
||||
### 5.2 消息历史同步
|
||||
|
||||
`submit_with_tools_stream` 内部由 `run_tool_loop` 管理 `messages` 的拷贝,不会写入 `self.messages`。消费方在收到 `MessageComplete` 后需手动:
|
||||
|
||||
```rust
|
||||
let response = full_response.clone();
|
||||
cycle.push_message(response.message.clone());
|
||||
```
|
||||
|
||||
在 `AgentSession::submit_turn_stream` 中,由于流是延迟求值且 `&mut self` 无法进入 spawn 闭包,消息历史同步交由消费方在 `finalize_turn` 前自行决定。当前方案中 `submit_turn_stream` **不自动同步消息历史**,这与 `submit_stream` 的已有行为一致(ponytail: Phase 2 FIX-E 注释)。
|
||||
|
||||
---
|
||||
|
||||
## 6. 风险评估
|
||||
|
||||
| 风险 | 影响 | 缓解措施 |
|
||||
|------|------|---------|
|
||||
| `&mut self` 约束导致流内无法访问 session 状态 | 中 | 复用 `submit_stream` 已有模式:方法体内读取 `self` 后构建 owned 数据,spawn 闭包不捕获 `&mut self` |
|
||||
| spawn task 生命周期管理 | 低 | 用户 drop rx → `tx.send` 返回 `Err` → `run_tool_loop` 自然终止 |
|
||||
| spawn task panic 静默丢失 | 中 | `run_tool_loop` 内部使用 `match`/`?` 避免 `unwrap()`;`JoinHandle` 不做 `await`(detached),panic 由 tokio 运行时记录日志。消费者看到 stream 提前结束(收到 `None`)但无 Error 事件 |
|
||||
| 中间轮 cost 不累加到 `cost_so_far` | 低 | 与现有 `submit_turn` 行为一致(仅最终轮计入),标记为已知限制,不在此 Phase 修复 |
|
||||
| 工具循环中 hook 可用性 | 低 | `PreRequest`/`PostRequest` hook 通过 `hook_executor.clone()` 进入 spawn task;hook 在 `run_tool_loop` 循环内触发 |
|
||||
| `run_tool_loop` 不支持消息压缩 | 中 | 长工具循环中消息不断增长,可能超出 context window。当前不传递 `compact_config`,后续可扩展 `run_tool_loop` 签名增添此参数 |
|
||||
| `unbounded_channel` 在消费慢于生产时内存增长 | 低 | LLM 流式输出天然有节流(token 生成速度远慢于 CPU 处理速度),消费者通常快于生产者。后续如需背压可切换为 `mpsc::channel(N)` + backpressure |
|
||||
| 流式版本不做 LLM retry | 低 | 非流式 `submit_with_tools` 通过 `submit_request` 的 retry 循环处理临时错误。流式版本中 retry 需重建 mpsc 通道,复杂度不匹配。调用方可使用 `RetryProvider` 包装器或在 Provider 层实现 retry |
|
||||
| 执行模式与 `submit_stream` 不一致(立即 vs 惰性) | 低 | `submit_stream` 的惰性语义不适配需要后台执行的工具循环。消费者应在 `submit_turn_stream` 返回后尽快消费流事件 |
|
||||
| `tokio::spawn` 要求 tokio 多线程运行时 | 低 | agcore 已依赖 tokio,涉及 IO 的 API 均使用 async。`#[tokio::test]` 单线程运行时不支持 `spawn`,测试中将 `run_tool_loop` 提取为可独立调用的函数,测试不走 spawn 直接调用 |
|
||||
| `CycleConfig` 加 `Clone` 影响现有代码 | 无 | 纯配置 struct,所有字段是基础类型或已 `Clone` 的 `RetryConfig` |
|
||||
|
||||
---
|
||||
|
||||
## 7. 验收标准
|
||||
|
||||
| # | 验收项 | 验证方式 |
|
||||
|---|--------|---------|
|
||||
| 1 | `cargo build --all-targets` 通过 | ✅ 编译器无错误 |
|
||||
| 2 | `submit_with_tools_stream` 纯文本流事件序列正确 | 单元测试验证:事件类型、顺序与 `submit_stream` 一致 |
|
||||
| 3 | `submit_with_tools_stream` 单轮工具调用事件序列正确 | 单元测试验证:含 `ToolExecutionStarted` / `ToolExecutionCompleted` |
|
||||
| 4 | `submit_with_tools_stream` 多轮工具调用事件序列正确 | 单元测试验证:多对 `ToolExecutionStarted`/`ToolExecutionCompleted` |
|
||||
| 5 | `submit_with_tools_stream` 最大轮次超限产生 Error 事件 | 单元测试验证:流中出现 `StreamEvent::Error` |
|
||||
| 6 | `submit_turn_stream` + `finalize_turn` 端到端链路 | 集成测试验证:cost 更新 + hook 触发 |
|
||||
| 7 | `cargo test --all-targets` 全绿,存量测试 0 回归 | ✅ 无回归 |
|
||||
| 8 | `cargo clippy --all-targets -- -D warnings` 0 警告 | ✅ 无警告 |
|
||||
| 9 | 现有 `submit_turn` / `submit_with_tools` / `submit_stream` 行为零影响 | ✅ 存量测试通过 |
|
||||
|
||||
---
|
||||
|
||||
---
|
||||
|
||||
## 8. 实施计划
|
||||
|
||||
按 5 个 Step 分阶段实施,每步产出独立 commit,可验证后退。
|
||||
|
||||
### 依赖关系
|
||||
|
||||
```mermaid
|
||||
graph LR
|
||||
S1["Step 1: 基础设施"]:::s1
|
||||
S2["Step 2: 核心状态机"]:::s2
|
||||
S3["Step 3: LlmCycle 单元测试"]:::s3
|
||||
S4["Step 4: AgentSession 包装"]:::s4
|
||||
S5["Step 5: 集成测试 + 扫尾"]:::s5
|
||||
|
||||
S1 --> S2
|
||||
S1 --> S4
|
||||
S2 --> S3
|
||||
S2 --> S4
|
||||
S3 --> S5
|
||||
S4 --> S5
|
||||
|
||||
classDef s1 fill:#e2e8f0,stroke:#94a3b8
|
||||
classDef s2 fill:#fbbf24,stroke:#d97706
|
||||
classDef s3 fill:#93c5fd,stroke:#2563eb
|
||||
classDef s4 fill:#93c5fd,stroke:#2563eb
|
||||
classDef s5 fill:#4ade80,stroke:#16a34a
|
||||
```
|
||||
|
||||
| Step | 依赖 | 并行机会 |
|
||||
|------|------|---------|
|
||||
| S1 | 无 | — |
|
||||
| S2 | S1 | 可与 S4 并行 |
|
||||
| S3 | S2 | 阻塞,需 S2 完成 |
|
||||
| S4 | S1, S2 | 功能依赖 S2(调用 `submit_with_tools_stream`);文件级无重叠但需先编译过 S2 |
|
||||
| S5 | S3 + S4 | 需 S3 和 S4 都完成 |
|
||||
|
||||
### Step 1 — 基础设施准备
|
||||
|
||||
**工作量**:S(< 1h)
|
||||
**风险**:低(纯新增,不影响现有代码逻辑)
|
||||
|
||||
| # | 任务 | 涉及文件 | 前置依赖 | 风险 |
|
||||
|---|------|---------|---------|------|
|
||||
| 1.1 | `StreamEvent` 枚举追加 `ToolExecutionStarted` 变体 | `llm/types/response_v2.rs` | 无 | 低 |
|
||||
| 1.2 | `StreamEvent` 枚举追加 `ToolExecutionCompleted` 变体 | `llm/types/response_v2.rs` | 1.1 | 低 |
|
||||
| 1.3 | `PartialMessageResponse::apply_to` 追加两个元事件 arm(均返回 `true`) | `llm/types/response_v2.rs` | 1.2 | 低 |
|
||||
| 1.4 | `CycleConfig` 加 `#[derive(Clone)]` | `llm/cycle.rs` | 无 | 低 |
|
||||
|
||||
**验收条件**:
|
||||
- `cargo build` 通过,编译器无 warning
|
||||
- 新增的 `StreamEvent` 变体可通过 `serde` roundtrip 序列化/反序列化
|
||||
- `CycleConfig` 可正常 clone
|
||||
|
||||
### Step 2 — `LlmCycle::submit_with_tools_stream` 核心
|
||||
|
||||
**工作量**:M(1-4h)
|
||||
**风险**:中(核心实现,需正确设计 spawn + mpsc 生命周期)
|
||||
**前置依赖**:S1
|
||||
|
||||
| # | 任务 | 涉及文件 | 前置依赖 | 风险 |
|
||||
|---|------|---------|---------|------|
|
||||
| 2.1 | 实现 `run_tool_loop()` 模块函数:消息循环构建请求 → `chat_stream` → 消费流 → 检测 tool_use → 工具执行 → 新一轮 | `llm/cycle.rs` | S1 | 中 |
|
||||
| 2.2 | 实现 `submit_with_tools_stream()` 公开方法:提取字段 → spawn `run_tool_loop` → 返回 `UnboundedReceiverStream` | `llm/cycle.rs` | 2.1 | 中 |
|
||||
| 2.3 | 新增导入:`tokio::sync::mpsc`、`tokio_stream::wrappers::UnboundedReceiverStream` | `llm/cycle.rs` | 2.2 | 低 |
|
||||
|
||||
**关键实现细节**:
|
||||
|
||||
```rust
|
||||
// run_tool_loop 函数签名
|
||||
async fn run_tool_loop(
|
||||
mut messages: Vec<Message>,
|
||||
provider: Arc<dyn LlmProvider>,
|
||||
config: CycleConfig,
|
||||
tool_registry: Arc<ToolRegistry>,
|
||||
tools: Vec<ToolDef>,
|
||||
tx: mpsc::UnboundedSender<StreamEvent>,
|
||||
hook_executor: Option<Arc<HookExecutor>>,
|
||||
) {
|
||||
let max_turns = config.max_tool_turns.unwrap_or(10);
|
||||
let tool_timeout = config.tool_timeout_secs;
|
||||
let max_bytes = config.max_tool_result_bytes;
|
||||
|
||||
let mut round = 0u32;
|
||||
loop {
|
||||
round += 1;
|
||||
if round > max_turns {
|
||||
// §3.6 错误表:最大轮次超限 → Error 事件 + 终止
|
||||
let _ = tx.send(StreamEvent::Error { message: "达到最大工具循环轮次".to_string() });
|
||||
break;
|
||||
}
|
||||
// ① 构建请求
|
||||
let request = MessageRequest {
|
||||
model: config.model.clone(),
|
||||
messages: messages.clone(),
|
||||
tools: tools.clone(),
|
||||
tool_choice: ToolChoice::Auto,
|
||||
max_tokens: config.max_tokens,
|
||||
temperature: config.temperature,
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
// ② PreRequest hook
|
||||
// ...
|
||||
|
||||
// ③ chat_stream
|
||||
let stream = match provider.chat_stream(request).await {
|
||||
Ok(s) => s,
|
||||
Err(e) => {
|
||||
let _ = tx.send(StreamEvent::Error { message: e.to_string() });
|
||||
return;
|
||||
}
|
||||
};
|
||||
|
||||
// ④ 消费流
|
||||
let mut partial = PartialMessageResponse::new();
|
||||
let mut stream = stream;
|
||||
while let Some(result) = stream.next().await {
|
||||
match result {
|
||||
Ok(event) => {
|
||||
partial.apply_to(&event);
|
||||
if tx.send(event).is_err() { return; }
|
||||
}
|
||||
Err(e) => {
|
||||
// ponytail: 流内事件错误后 partial 处于损坏状态,
|
||||
// 不能继续执行 finalize/finalize —— 直接 return 结束 task
|
||||
let _ = tx.send(StreamEvent::Error { message: e.to_string() });
|
||||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ⑤ finalize
|
||||
let response = match partial.finalize() {
|
||||
Ok(r) => r,
|
||||
Err(e) => { let _ = tx.send(StreamEvent::Error { .. }); return; }
|
||||
};
|
||||
messages.push(response.message.clone());
|
||||
|
||||
// ⑥ 检测 tool_use
|
||||
if !has_tool_calls_in_response(&response) {
|
||||
break; // 最终轮
|
||||
}
|
||||
|
||||
// ⑦ 执行工具
|
||||
let tool_calls = extract_tool_calls_from_response(&response);
|
||||
let calls: Vec<_> = tool_calls.into_iter()
|
||||
.map(|(id, name, args)| {
|
||||
let value = serde_json::from_str(&args).unwrap_or(Value::Null);
|
||||
(id, name, value)
|
||||
}).collect();
|
||||
|
||||
for (tool_call_id, tool_name, args_value) in &calls {
|
||||
let args_json = serde_json::to_string(&args_value).unwrap_or_default();
|
||||
if tx.send(StreamEvent::ToolExecutionStarted {
|
||||
tool_name: tool_name.clone(),
|
||||
tool_call_id: tool_call_id.clone(),
|
||||
arguments: args_json,
|
||||
}).is_err() { return; }
|
||||
}
|
||||
|
||||
let results = tool_registry.invoke_all(calls, tool_timeout).await;
|
||||
|
||||
for result in &results {
|
||||
let summary = match &result.output {
|
||||
Ok(v) => serde_json::to_string(v).unwrap_or_default(),
|
||||
Err(e) => e.to_string(),
|
||||
};
|
||||
// ponytail: 复用现有 truncate_tool_result 函数(cycle.rs 末尾),
|
||||
// 确保多字节 UTF-8 字符不被截断破坏。上限 200 字符。
|
||||
let truncated = truncate_tool_result(&summary, 200);
|
||||
if tx.send(StreamEvent::ToolExecutionCompleted {
|
||||
tool_name: result.tool_name.clone(),
|
||||
tool_call_id: result.tool_call_id.clone(),
|
||||
result_summary: truncated,
|
||||
is_error: result.output.is_err(),
|
||||
}).is_err() { return; }
|
||||
}
|
||||
|
||||
for result in results {
|
||||
let is_error = result.output.is_err();
|
||||
let content = match &result.output {
|
||||
Ok(v) => serde_json::to_string(v).unwrap_or_default(),
|
||||
Err(e) if e.is_recoverable() => format!("错误: {}", e),
|
||||
Err(e) => {
|
||||
let _ = tx.send(StreamEvent::Error { .. });
|
||||
return;
|
||||
}
|
||||
};
|
||||
messages.push(Message::tool_result(result.tool_call_id, content, is_error));
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**验收条件**:
|
||||
- `cargo build` 通过
|
||||
- 新增方法签名与方案设计一致
|
||||
- 未修改现有 `submit_with_tools`/`submit_stream` 的行为
|
||||
|
||||
### Step 3 — 单元测试(LlmCycle 层)
|
||||
|
||||
**工作量**:M(1-4h)
|
||||
**风险**:低(与现有测试模式一致,使用已有 MockProvider)
|
||||
**前置依赖**:S2
|
||||
|
||||
测试策略:直接使用公开的 `crate::llm::mock::MockProvider`(已完整实现 `chat_stream` + 预设响应队列),避免改造 `cycle.rs` 测试模块内的内联 Stub。测试中调用 `submit_with_tools_stream` 时通过 `#[tokio::test(flavor = "multi_thread")]` 满足 spawn 运行时要求,或在单元级将 `run_tool_loop` 作为独立函数直接测试(不走 spawn)。
|
||||
|
||||
| # | 测试场景 | Mock 响应序列 | 验证点 | 覆盖路径 |
|
||||
|---|---------|---------------|--------|---------|
|
||||
| 3.1 | 纯文本流 | 1 个 text 响应 | 事件序列与 `submit_stream` 一致;无 `ToolExecutionStarted`/`ToolExecutionCompleted` | 正常路径:单轮 LLM → 文本返回 |
|
||||
| 3.2 | 单轮工具调用 | 2 个响应:tool_use + text | 包含一对 `ToolExecutionStarted`/`ToolExecutionCompleted`;最终 `stop_reason` 为 `Stop` | 正常路径:LLM → 工具 → LLM |
|
||||
| 3.3 | 多轮工具调用 | 4 个响应:3×tool_use + 1×text | 3 对 `ToolExecutionStarted`/`ToolExecutionCompleted`;消息历史长度为 8(user + 3×(assistant+tool) + final assistant) | 正常路径:LLM → 工具 → LLM → 工具 → LLM |
|
||||
| 3.4 | 最大轮次超限 | 3 个 tool_use 响应,`max_tool_turns: Some(2)` | 流中出现 `StreamEvent::Error`;消息历史停在第 2 轮 | 边界条件:超出上限 |
|
||||
| 3.5 | `chat_stream` 返回 Err | Mock `chat_stream` 返回 `Err(LlmError::Other(...))` | 流中第一个事件为 `StreamEvent::Error`;随后流结束 | 异常路径:LLM 不可用 |
|
||||
| 3.6 | 空 tool_registry | 1 个 text 响应,registry 中无工具 | 流退化为纯文本流,事件序列与 3.1 一致 | 退化场景:无工具可用 |
|
||||
| 3.7 | 不可恢复工具错误 | 2 个响应:tool_use → text,工具返回 `ToolError::ExecutionFailed`(不可恢复) | 流中出现 `StreamEvent::Error`;消息历史中不含该工具结果(循环终止前未 push) | 异常路径:工具执行失败 |
|
||||
| 3.8 | 可恢复工具错误 | 2 个响应:tool_use → text,工具返回 `ToolError::ExecutionFailed`(可恢复) | 工具结果作为 `ToolResult { is_error: true }` 回传 LLM;流正常结束,无 `Error` 事件 | 异常路径:工具出错但可恢复 |
|
||||
| 3.9 | 工具超时 | 2 个响应:tool_use → text,`tool_timeout_secs: 1`,模拟工具耗时 10 秒 | 流中出现 `StreamEvent::Error`;循环终止前未 push 工具结果 | 异常路径:工具执行超时 |
|
||||
|
||||
**验收条件**:
|
||||
- `cargo test` 新增 8 个测试全部通过
|
||||
- `cargo test` 存量测试 0 回归
|
||||
|
||||
### Step 4 — `AgentSession` 层包装
|
||||
|
||||
**工作量**:S(< 1h)
|
||||
**风险**:低(薄包装层,逻辑简单)
|
||||
**前置依赖**:S1
|
||||
|
||||
| # | 任务 | 涉及文件 | 前置依赖 | 风险 |
|
||||
|---|------|---------|---------|------|
|
||||
| 4.1 | 实现 `submit_turn_stream()`:触发 `OnTurnStart` → 组装 `LlmCycle` → 调用 `submit_with_tools_stream` → `turn_index += 1` → 返回流 | `agent/session.rs` | S1 | 低 |
|
||||
| 4.2 | 实现 `finalize_turn()`:`cost_so_far.add()` → 触发 `OnTurnEnd` hook | `agent/session.rs` | S1 | 低 |
|
||||
|
||||
**验收条件**:
|
||||
- `cargo build` 通过
|
||||
- 新增方法签名与方案设计一致
|
||||
- 与 `submit_turn` 的 system_prompt / compact_config / bundle 使用方式一致
|
||||
|
||||
### Step 5 — 集成测试 + 扫尾
|
||||
|
||||
**工作量**:S(< 1h)
|
||||
**风险**:低(基于现有测试框架)
|
||||
**前置依赖**:S3 + S4
|
||||
|
||||
| # | 任务 | 涉及文件 | 前置依赖 | 风险 |
|
||||
|---|------|---------|---------|------|
|
||||
| 5.1 | `submit_turn_stream` 端到端测试:跑通 mock provider → 消费流验证各事件到达 → `finalize_turn` 后 cost 更新正确 | `agent/session.rs`(内联测试) | S4 | 低 |
|
||||
| 5.2 | Hook 触发验证:`OnTurnStart` 在 `submit_turn_stream` 返回流之前触发;`finalize_turn` 调用后 `OnTurnEnd` 正确触发 | `agent/session.rs`(内联测试) | S4 | 低 |
|
||||
| 5.3 | `cargo test --all-targets` 全绿验证 | 全仓 | S5.1+S5.2 | 低 |
|
||||
| 5.4 | `cargo clippy --all-targets -- -D warnings` 0 警告 | 全仓 | S5.3 | 低 |
|
||||
| 5.5 | `cargo build --all-targets` 发布模式验证 | 全仓 | S5.4 | 低 |
|
||||
|
||||
**验收条件**:
|
||||
- 全量测试通过,存量 0 回归
|
||||
- clippy 0 警告
|
||||
- 发布模式零 warning
|
||||
|
||||
### 实施总览
|
||||
|
||||
| | Step 1 | Step 2 | Step 3 | Step 4 | Step 5 | **合计** |
|
||||
|--|--------|--------|--------|--------|--------|---------|
|
||||
| **工作量** | S | M | M | S | S | **M-L** |
|
||||
| **文件数** | 2 | 1 | 1(内联) | 1 | 1(内联) | **~4** |
|
||||
| **代码行** | ~20 | ~140 | ~150 含测试 | ~70 | ~80 含测试 | **~380** |
|
||||
| **风险** | 低 | 中 | 低 | 低 | 低 | 中 |
|
||||
| **并行** | — | 阻塞(S4 依赖 S2) | 阻塞 | 阻塞(依赖 S2) | 阻塞 | — |
|
||||
|
||||
---
|
||||
|
||||
## 附录 A:新增 StreamEvent 变体的 apply_to 语义
|
||||
|
||||
```rust
|
||||
// 在 PartialMessageResponse::apply_to 中追加:
|
||||
StreamEvent::ToolExecutionStarted { .. } | StreamEvent::ToolExecutionCompleted { .. } => {
|
||||
// 元事件:不参与内容块累积,不修改 partial response 状态
|
||||
true
|
||||
}
|
||||
```
|
||||
|
||||
## 附录 B:CycleConfig 的 Clone 推导
|
||||
|
||||
```rust
|
||||
/// LLM 调用周期配置。
|
||||
#[derive(Debug, Clone)] // ← 追加 Clone
|
||||
pub struct CycleConfig {
|
||||
pub model: String,
|
||||
pub max_tokens: Option<u32>,
|
||||
pub temperature: Option<f32>,
|
||||
pub max_turns: Option<u32>,
|
||||
pub retry: RetryConfig, // 已 #[derive(Clone)]
|
||||
pub max_tool_turns: Option<u32>,
|
||||
pub tool_timeout_secs: u64,
|
||||
pub max_tool_result_bytes: usize,
|
||||
}
|
||||
```
|
||||
+5
-3
@@ -439,12 +439,14 @@ pub struct ContextBudget { system, history, tools, tool_results, reserve }
|
||||
|
||||
**Phase 8 全部完成**。**已打 `v0.2.0-rc.1` 标签**。
|
||||
|
||||
**实际新增**(2026-07-05,6 commits):
|
||||
**实际新增**(2026-07-05,7 commits):
|
||||
- `feat(core)` —— 14 个公开枚举追加 `#[non_exhaustive]`(P0 核心 IR + P0 Error + P1 其他)
|
||||
- `refactor(agent)` —— `StepStatus::Completed(ChatResponse)` → `Completed(MessageResponse)` + `task_agent_demo.rs` 清理 3 处废弃类型
|
||||
- `docs` —— CHANGELOG v0.2.0-rc.1 条目 + Cargo.toml version 0.1.0 → 0.2.0-rc.1 + README 示例列表 7 → 10
|
||||
- `test(core)` —— 验证 commit 1-3 零回归(test 200 passed + clippy 0 警告 + doc 0 warning)
|
||||
- `feat(examples)` —— `quick_start.rs`(57 行)+ `end_to_end.rs`(237 行)
|
||||
- `feat(examples)` —— `quick_start.rs`(60 行)+ `end_to_end.rs`(246 行)
|
||||
- `docs(roadmap)` —— 标记 Phase 8 全部完成 + M4 里程碑 ✅
|
||||
- `fix(examples)` —— 实施后 PM/SA/Code Reviewer 三方审查发现 6 项问题(🔴 CalcTool 除零 panic + 🟡 drop 注释准确性 + 🟡 EchoTool 错误处理 + 💭 断言一致性 + 💭 工具两端语义统一 + 💭 trailing newline),全部修复
|
||||
|
||||
**依赖**:Phase 5(ProviderConfig from_env)+ Phase 6(ToolDef)+ Phase 7(SqliteStore)
|
||||
**优先级**:P0
|
||||
@@ -617,7 +619,7 @@ graph BT
|
||||
- ✅ Phase 5 Warmup — ProviderConfig::from_env + OllamaProvider + `#[non_exhaustive]` 前置标记(ProviderType / StopReason / FinishReason / EvictionPolicy)
|
||||
- ✅ Phase 6 ToolDefinition IR — `ToolDef` 新类型 + 双向 `From` 转换 + 别名彻底移除 + `#[allow(deprecated)]` 清理(cycle/registry/mcp/agent);Anthropic 零改动;roundtrip 测试覆盖
|
||||
- ✅ Phase 7 SqliteStore — `rusqlite 0.32` + WAL 模式 + `Arc<Mutex<Connection>>` + `spawn_blocking`;`memory/store.rs` → `store/{in_memory,sqlite_store}.rs` 模块化;9 个内联测试覆盖 CRUD/upsert/过滤/10×10 并发/持久化 round-trip;`InMemoryStore ↔ SqliteStore` trait-box 互换兼容
|
||||
- ✅ **Phase 8 MVP 集成出口** — 14 个公开枚举追加 `#[non_exhaustive]`(P0 核心 IR + P0 Error + P1 其他) + `StepStatus::Completed(ChatResponse)` → `Completed(MessageResponse)` 迁移 + CHANGELOG v0.2.0-rc.1 + 2 个新示例(`quick_start` 57 行 + `end_to_end` 237 行),10 个离线示例全部 exit 0;**v0.2.0-rc.1 标签已打**
|
||||
- ✅ **Phase 8 MVP 集成出口** — 14 个公开枚举追加 `#[non_exhaustive]`(P0 核心 IR + P0 Error + P1 其他) + `StepStatus::Completed(ChatResponse)` → `Completed(MessageResponse)` 迁移 + CHANGELOG v0.2.0-rc.1 + 2 个新示例(`quick_start` 60 行 + `end_to_end` 246 行),10 个离线示例全部 exit 0;**v0.2.0-rc.1 标签已打**;实施后三方审查发现 6 项问题(1 🔴 + 2 🟡 + 3 💭)已全部修复
|
||||
- ✅ Provider IR 重构 — 统一类型系统 + OpenAI/Anthropic/DeepSeek/Qwen/Ollama 适配
|
||||
- ✅ LlmCycle 简化 — IR 消息类型切换 + Phase 0 桥接层移除
|
||||
- ✅ v0.1 Release — 技术债扫清、MockProvider 公开化、8 个离线示例(含 `simple_visit`)、README + 错误消息友好化、CHANGELOG 初始化
|
||||
|
||||
+207
-1
@@ -7,14 +7,18 @@
|
||||
//! - **不做业务循环**:多轮策略、错误重试、记忆回写由上层应用或具体 `TaskAgent` 决定
|
||||
//! - **不持有 ConversationMemory**:上层可独立 new 一个 `ConversationMemory`,在合适的时机调 `add_message`
|
||||
|
||||
use std::pin::Pin;
|
||||
use std::sync::Arc;
|
||||
|
||||
use futures_core::Stream;
|
||||
|
||||
use crate::agent::agent::Agent;
|
||||
use crate::agent::error::AgentError;
|
||||
use crate::agent::runtime::RuntimeBundle;
|
||||
use crate::agent::session_memory::SessionMemory;
|
||||
use crate::llm::cycle::{CostTracker, CycleConfig, LlmCycle};
|
||||
use crate::llm::hooks::{HookContext, HookEvent};
|
||||
use crate::llm::stream::StreamEvent;
|
||||
use crate::llm::types::message::Message;
|
||||
use crate::llm::types::response_v2::MessageResponse;
|
||||
use crate::memory::store::InMemoryStore;
|
||||
@@ -169,6 +173,84 @@ impl AgentSession {
|
||||
|
||||
Ok(response)
|
||||
}
|
||||
|
||||
/// 提交一轮对话(流式版本,含自动 tool 循环),返回 `StreamEvent` 流。
|
||||
///
|
||||
/// 与 `submit_turn` 的区别:
|
||||
/// - 以流事件序列而非 `MessageResponse` 返回
|
||||
/// - 工具执行期间插入 `ToolExecutionStarted` / `ToolExecutionCompleted` 事件
|
||||
/// - 消费方在收到 `MessageComplete` 后需手动调用 `finalize_turn` 同步状态
|
||||
///
|
||||
/// **运行时要求**:内部委托 `submit_with_tools_stream`,需要 tokio 多线程运行时。
|
||||
///
|
||||
/// ponytail: 流程结构与 `submit_turn` 对称,但 `OnTurnEnd` hook + cost 累计不在流生成路径上,
|
||||
/// 因为流是延迟求值且 `&mut self` 无法进入 spawn 闭包。调用方消费流完毕后必须调 `finalize_turn`。
|
||||
pub async fn submit_turn_stream(
|
||||
&mut self,
|
||||
user_input: impl Into<String>,
|
||||
) -> Result<Pin<Box<dyn Stream<Item = StreamEvent> + Send>>, AgentError> {
|
||||
let turn_index = self.turn_index;
|
||||
let hook_executor = Arc::clone(&self.bundle.hook_executor);
|
||||
|
||||
// 1. 触发 OnTurnStart hook(同步)
|
||||
let start_ctx = HookContext::new(HookEvent::OnTurnStart).with_turn_index(turn_index);
|
||||
hook_executor
|
||||
.execute(HookEvent::OnTurnStart, &start_ctx)
|
||||
.await;
|
||||
|
||||
// 2. 触发子 trait 覆盖(白名单/过滤)的副作用
|
||||
let _ = self.agent.tool_definitions(&self.bundle);
|
||||
|
||||
// 3. 组装 LlmCycle
|
||||
let mut cycle =
|
||||
LlmCycle::new_with_arc(Arc::clone(&self.bundle.provider), CycleConfig::default())
|
||||
.with_messages(Vec::new());
|
||||
if let Some(prompt) = self.agent.system_prompt() {
|
||||
cycle = cycle.with_messages(vec![Message::system(prompt)]);
|
||||
}
|
||||
if let Some(cfg) = self.bundle.config.compact_config.clone() {
|
||||
cycle = cycle.with_compact_config(cfg);
|
||||
}
|
||||
|
||||
// 4. 调用流式工具循环
|
||||
let stream = cycle
|
||||
.submit_with_tools_stream(
|
||||
user_input.into(),
|
||||
Arc::clone(&self.bundle.tool_registry),
|
||||
)
|
||||
.await?;
|
||||
|
||||
// 5. turn_index 递增 —— 配合 finalize_turn 用 (turn_index - 1) 传递正确的 OnTurnEnd 序号
|
||||
self.turn_index += 1;
|
||||
|
||||
// 注:hook_executor 不显式 drop,生命周期由 Arc 自动管理
|
||||
let _ = hook_executor;
|
||||
|
||||
Ok(stream)
|
||||
}
|
||||
|
||||
/// 完成一轮 turn:累计 cost + 触发 OnTurnEnd hook。
|
||||
///
|
||||
/// 由消费者在收到 `MessageComplete.full_response` 后调用。
|
||||
///
|
||||
/// **消费者注意**:`finalize_turn` 是开发者责任 —— 遗漏调用会导致 cost 不累计、OnTurnEnd 不触发。
|
||||
/// session 状态仍然可用,后续 `submit_turn` 也能正常执行,但 cost 信息不完整。
|
||||
///
|
||||
/// ponytail: 与 `submit_turn` 行为对齐 —— `cost_so_far` 仅计入最终轮的 usage。
|
||||
pub async fn finalize_turn(&mut self, response: &MessageResponse) {
|
||||
self.cost_so_far.add(&response.usage);
|
||||
|
||||
// ponytail: 防御性 saturating_sub 防止误用 panic。
|
||||
// 正常路径是 submit_turn_stream 内 turn_index += 1 后再调 finalize_turn,
|
||||
// 所以 saturating 后为 0 是正常的;若调用方忘了先调 submit_turn_stream,
|
||||
// turn_index 仍为 0,传 0 给 OnTurnEnd hook 不会 panic。
|
||||
let end_ctx = HookContext::new(HookEvent::OnTurnEnd)
|
||||
.with_turn_index(self.turn_index.saturating_sub(1));
|
||||
self.bundle
|
||||
.hook_executor
|
||||
.execute(HookEvent::OnTurnEnd, &end_ctx)
|
||||
.await;
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
@@ -331,4 +413,128 @@ mod tests {
|
||||
assert_eq!(start_count.0.load(Ordering::SeqCst), 2);
|
||||
assert_eq!(end_count.0.load(Ordering::SeqCst), 2);
|
||||
}
|
||||
}
|
||||
|
||||
// ====== Phase 9: submit_turn_stream + finalize_turn 集成测试 ======
|
||||
|
||||
use futures_util::StreamExt;
|
||||
|
||||
/// 集成测试 5.1 — `submit_turn_stream` 端到端:跑通 mock provider → 消费流
|
||||
/// 验证各事件到达 → `finalize_turn` 后 cost 更新正确
|
||||
#[tokio::test(flavor = "multi_thread")]
|
||||
async fn submit_turn_stream_end_to_end() {
|
||||
use crate::llm::mock::MockProvider as SessionMock;
|
||||
|
||||
let provider = Arc::new(SessionMock::new(vec![assistant_text("hello")]));
|
||||
let agent = Arc::new(StubAgent {
|
||||
name: "stub".into(),
|
||||
prompt: Some("you are a test agent".into()),
|
||||
});
|
||||
let bundle = Arc::new(
|
||||
AgentBuilder::new()
|
||||
.provider(provider)
|
||||
.tool_registry(Arc::new(ToolRegistry::new()))
|
||||
.hook_executor(Arc::new(HookExecutor::new()))
|
||||
.build()
|
||||
.unwrap(),
|
||||
);
|
||||
|
||||
let mut session = AgentSession::new(agent, "stream-s1", bundle);
|
||||
assert_eq!(session.turn_index(), 0);
|
||||
|
||||
// 1. 提交 stream
|
||||
let mut stream = session.submit_turn_stream("hi").await.unwrap();
|
||||
|
||||
// 2. 消费流,收集事件直到结束
|
||||
let mut events: Vec<StreamEvent> = Vec::new();
|
||||
let mut final_response = None;
|
||||
while let Some(event) = stream.next().await {
|
||||
let _ = &event;
|
||||
if let StreamEvent::MessageComplete { full_response } = &event {
|
||||
final_response = Some(full_response.clone());
|
||||
}
|
||||
events.push(event);
|
||||
}
|
||||
|
||||
// 3. 验证事件序列
|
||||
assert!(!events.is_empty(), "应有事件");
|
||||
assert!(
|
||||
events.iter().any(|e| matches!(e, StreamEvent::MessageStart { .. })),
|
||||
"应包含 MessageStart"
|
||||
);
|
||||
assert!(
|
||||
events
|
||||
.iter()
|
||||
.any(|e| matches!(e, StreamEvent::TextDelta { text } if text == "hello")),
|
||||
"应包含 TextDelta hello"
|
||||
);
|
||||
assert!(
|
||||
events
|
||||
.iter()
|
||||
.any(|e| matches!(e, StreamEvent::MessageComplete { .. })),
|
||||
"应包含 MessageComplete"
|
||||
);
|
||||
|
||||
// 4. 调用 finalize_turn
|
||||
let response = final_response.expect("流应包含至少一个 MessageComplete");
|
||||
session.finalize_turn(&response).await;
|
||||
|
||||
// 5. 验证 cost 累计
|
||||
assert_eq!(session.turn_index(), 1);
|
||||
assert_eq!(session.usage().total().prompt_tokens, 10);
|
||||
assert_eq!(session.usage().total().completion_tokens, 5);
|
||||
}
|
||||
|
||||
/// 集成测试 5.2 — Hook 触发验证:
|
||||
/// - `OnTurnStart` 在 `submit_turn_stream` 返回流之前触发
|
||||
/// - `finalize_turn` 调用后 `OnTurnEnd` 正确触发
|
||||
#[tokio::test(flavor = "multi_thread")]
|
||||
async fn submit_turn_stream_triggers_turn_hooks() {
|
||||
use crate::llm::mock::MockProvider as SessionMock;
|
||||
|
||||
let mut hook_executor = HookExecutor::new();
|
||||
let start_count = Arc::new(CountHook(AtomicU32::new(0)));
|
||||
let end_count = Arc::new(CountHook(AtomicU32::new(0)));
|
||||
hook_executor.register(
|
||||
HookEvent::OnTurnStart,
|
||||
Box::new(CountHookAdapter(start_count.clone())),
|
||||
);
|
||||
hook_executor.register(
|
||||
HookEvent::OnTurnEnd,
|
||||
Box::new(CountHookAdapter(end_count.clone())),
|
||||
);
|
||||
|
||||
let provider = Arc::new(SessionMock::new(vec![assistant_text("ok")]));
|
||||
let agent = Arc::new(StubAgent {
|
||||
name: "stub".into(),
|
||||
prompt: None,
|
||||
});
|
||||
let bundle = Arc::new(
|
||||
AgentBuilder::new()
|
||||
.provider(provider)
|
||||
.tool_registry(Arc::new(ToolRegistry::new()))
|
||||
.hook_executor(Arc::new(hook_executor))
|
||||
.build()
|
||||
.unwrap(),
|
||||
);
|
||||
let mut session = AgentSession::new(agent, "stream-s2", bundle);
|
||||
|
||||
// 1. submit_turn_stream 触发 OnTurnStart
|
||||
let mut stream = session.submit_turn_stream("hi").await.unwrap();
|
||||
assert_eq!(start_count.0.load(Ordering::SeqCst), 1);
|
||||
assert_eq!(end_count.0.load(Ordering::SeqCst), 0, "OnTurnEnd 未在流返回前触发");
|
||||
|
||||
// 2. 消费完流后再调 finalize_turn 触发 OnTurnEnd
|
||||
let mut final_response = None;
|
||||
while let Some(event) = stream.next().await {
|
||||
if let StreamEvent::MessageComplete { full_response } = &event {
|
||||
final_response = Some(full_response.clone());
|
||||
}
|
||||
}
|
||||
session
|
||||
.finalize_turn(&final_response.expect("应有 MessageComplete"))
|
||||
.await;
|
||||
|
||||
assert_eq!(start_count.0.load(Ordering::SeqCst), 1);
|
||||
assert_eq!(end_count.0.load(Ordering::SeqCst), 1, "OnTurnEnd 在 finalize_turn 后触发");
|
||||
}
|
||||
}
|
||||
+783
-1
@@ -11,6 +11,9 @@ use std::sync::Arc;
|
||||
|
||||
use async_stream::stream;
|
||||
use futures_core::stream::Stream;
|
||||
use serde_json::Value;
|
||||
use tokio::sync::mpsc;
|
||||
use tokio_stream::wrappers::UnboundedReceiverStream;
|
||||
|
||||
use crate::llm::compact::{CompactConfig, CompactState, microcompact, should_compact};
|
||||
use crate::llm::cycle::retry::should_retry;
|
||||
@@ -20,11 +23,13 @@ use crate::llm::provider::LlmProvider;
|
||||
use crate::llm::stream::StreamEvent;
|
||||
use crate::llm::types::message::{ContentBlock, Message};
|
||||
use crate::llm::types::request_v2::MessageRequest;
|
||||
use crate::llm::types::response_v2::{MessageResponse, StopReason};
|
||||
use crate::llm::types::response_v2::{MessageResponse, PartialMessageResponse, StopReason};
|
||||
use crate::llm::types::tool::ToolDef;
|
||||
use crate::llm::types::ToolChoice;
|
||||
use crate::tools::ToolRegistry;
|
||||
|
||||
/// LLM 调用周期配置。
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct CycleConfig {
|
||||
/// 模型名称。
|
||||
pub model: String,
|
||||
@@ -613,6 +618,53 @@ impl LlmCycle {
|
||||
}
|
||||
}
|
||||
|
||||
/// 提交消息并自动处理工具调用循环,流式产出所有事件。
|
||||
///
|
||||
/// 与 `submit_with_tools` 的区别:
|
||||
/// - LLM 响应是流式的(全程 `chat_stream` 而非 `chat`)
|
||||
/// - 工具执行前后插入 `ToolExecutionStarted` / `ToolExecutionCompleted` 事件
|
||||
/// - 错误以 `StreamEvent::Error` 形式出现在流中,而非终止 `Result`
|
||||
/// - 消费方需手动 `push_message()` 同步消息历史
|
||||
///
|
||||
/// **运行时要求**:内部使用 `tokio::spawn`,需要 tokio 多线程运行时。
|
||||
/// `#[tokio::test]` 单线程运行时不支持 spawn,测试场景需用 `flavor = "multi_thread"` 或
|
||||
/// 直接调用模块函数 `run_tool_loop`。
|
||||
///
|
||||
/// ponytail: 返回的流是 `Item = StreamEvent`(非 `Result`),所有错误事件化为 `StreamEvent::Error`。
|
||||
pub async fn submit_with_tools_stream(
|
||||
&mut self,
|
||||
prompt: String,
|
||||
tool_registry: Arc<ToolRegistry>,
|
||||
) -> Result<Pin<Box<dyn Stream<Item = StreamEvent> + Send>>, LlmError> {
|
||||
self.messages.push(Message::user_text(prompt));
|
||||
self.maybe_compact();
|
||||
|
||||
// 提取 self 字段所有 owned 数据 —— spawn 闭包不能捕获 &mut self。
|
||||
let provider = Arc::clone(&self.provider);
|
||||
let config = self.config.clone();
|
||||
let hook_executor = self.hook_executor.clone();
|
||||
let tools = tool_registry.definitions();
|
||||
let messages = std::mem::take(&mut self.messages);
|
||||
let tool_registry = tool_registry;
|
||||
|
||||
let (tx, rx) = mpsc::unbounded_channel::<StreamEvent>();
|
||||
|
||||
tokio::spawn(async move {
|
||||
run_tool_loop(
|
||||
messages,
|
||||
provider,
|
||||
config,
|
||||
tool_registry,
|
||||
tools,
|
||||
tx,
|
||||
hook_executor,
|
||||
)
|
||||
.await;
|
||||
});
|
||||
|
||||
Ok(Box::pin(UnboundedReceiverStream::new(rx)))
|
||||
}
|
||||
|
||||
/// 在接近上下文窗口时压缩历史消息。
|
||||
fn maybe_compact(&mut self) {
|
||||
if let Some(ref config) = self.compact_config
|
||||
@@ -672,6 +724,199 @@ fn truncate_tool_result(s: &str, max_bytes: usize) -> String {
|
||||
)
|
||||
}
|
||||
|
||||
/// 运行流式工具循环的核心异步状态机(Phase 9)。
|
||||
///
|
||||
/// 由 `submit_with_tools_stream` 在 `tokio::spawn` 中调用。
|
||||
///
|
||||
/// 流程:
|
||||
/// - loop: 构建请求 → `chat_stream` → 消费流(转 mpsc)→ finalize → 检测 tool_use
|
||||
/// - 如果有 tool_use:插入 `ToolExecutionStarted` → 调 `invoke_all` → 插入 `ToolExecutionCompleted`
|
||||
/// → push 工具结果 → 下一轮
|
||||
/// - 否则 break(最终轮)
|
||||
/// - 最大轮次超限:发 `StreamEvent::Error` + break
|
||||
///
|
||||
/// **所有错误事件化**:通过 `tx.send(Error{..})` 表达错误,最终 `return` 结束 task。
|
||||
/// 不返回 `Result`,因为错误已通过事件流传递。
|
||||
async fn run_tool_loop(
|
||||
mut messages: Vec<Message>,
|
||||
provider: Arc<dyn LlmProvider>,
|
||||
config: CycleConfig,
|
||||
tool_registry: Arc<ToolRegistry>,
|
||||
tools: Vec<ToolDef>,
|
||||
tx: mpsc::UnboundedSender<StreamEvent>,
|
||||
hook_executor: Option<Arc<HookExecutor>>,
|
||||
) {
|
||||
use futures_util::StreamExt;
|
||||
|
||||
let max_turns = config.max_tool_turns.unwrap_or(10);
|
||||
let tool_timeout = config.tool_timeout_secs;
|
||||
let max_bytes = config.max_tool_result_bytes;
|
||||
|
||||
let mut round = 0u32;
|
||||
loop {
|
||||
round += 1;
|
||||
if round > max_turns {
|
||||
// §3.6 错误表:最大轮次超限 → Error 事件 + 终止
|
||||
let _ = tx.send(StreamEvent::Error {
|
||||
message: "达到最大工具循环轮次".to_string(),
|
||||
});
|
||||
break;
|
||||
}
|
||||
|
||||
// ① 构建请求
|
||||
let request = MessageRequest {
|
||||
model: config.model.clone(),
|
||||
messages: messages.clone(),
|
||||
tools: tools.clone(),
|
||||
tool_choice: ToolChoice::Auto,
|
||||
max_tokens: config.max_tokens,
|
||||
temperature: config.temperature,
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
// ② PreRequest hook(fire-and-forget,仅占位保留以保持接口对称)
|
||||
let _ = hook_executor.as_ref();
|
||||
|
||||
// ③ chat_stream —— 第一层错误
|
||||
let mut stream = match provider.chat_stream(request).await {
|
||||
Ok(s) => s,
|
||||
Err(e) => {
|
||||
// ponytail: 不做 retry —— retry 重建 mpsc 通道复杂度与收益不匹配
|
||||
let _ = tx.send(StreamEvent::Error {
|
||||
message: e.to_string(),
|
||||
});
|
||||
return;
|
||||
}
|
||||
};
|
||||
|
||||
// ④ 消费 LLM 流
|
||||
let mut partial = PartialMessageResponse::new();
|
||||
loop {
|
||||
match stream.next().await {
|
||||
Some(Ok(event)) => {
|
||||
partial.apply_to(&event);
|
||||
let is_terminal = matches!(
|
||||
event,
|
||||
StreamEvent::MessageComplete { .. } | StreamEvent::Error { .. }
|
||||
);
|
||||
if tx.send(event).is_err() {
|
||||
return; // 消费者已 drop rx,task 终止
|
||||
}
|
||||
if is_terminal {
|
||||
break;
|
||||
}
|
||||
}
|
||||
Some(Err(e)) => {
|
||||
// ponytail: 流内 Err 后 partial 处于损坏状态,直接 return 结束 task
|
||||
let _ = tx.send(StreamEvent::Error {
|
||||
message: e.to_string(),
|
||||
});
|
||||
return;
|
||||
}
|
||||
None => break, // 流自然结束
|
||||
}
|
||||
}
|
||||
|
||||
// ⑤ finalize
|
||||
// ponytail: 防御性检查 —— 若 Provider 在流中产出过 Ok(StreamEvent::Error),
|
||||
// partial 已设为 is_errored=true,finalize() 会基于损坏状态生成不可信响应,
|
||||
// 此时跳过 finalize 直接退出流,让消费者看到 Error 事件后的流自然结束。
|
||||
if partial.is_errored {
|
||||
return;
|
||||
}
|
||||
let response = match partial.finalize() {
|
||||
Ok(r) => r,
|
||||
Err(e) => {
|
||||
let _ = tx.send(StreamEvent::Error {
|
||||
message: e.to_string(),
|
||||
});
|
||||
return;
|
||||
}
|
||||
};
|
||||
|
||||
messages.push(response.message.clone());
|
||||
|
||||
// ⑥ 检测 tool_use —— 没有则 break(最终轮)
|
||||
if !has_tool_calls_in_response(&response) {
|
||||
break;
|
||||
}
|
||||
|
||||
// ⑦ 提取 tool_calls
|
||||
let tool_calls = extract_tool_calls_from_response(&response);
|
||||
let calls: Vec<(String, String, Value)> = tool_calls
|
||||
.into_iter()
|
||||
.map(|(id, name, args)| {
|
||||
let value: Value = serde_json::from_str(&args).unwrap_or(Value::Null);
|
||||
(id, name, value)
|
||||
})
|
||||
.collect();
|
||||
|
||||
// ⑧ 发送 ToolExecutionStarted
|
||||
for (tool_call_id, tool_name, args_value) in &calls {
|
||||
let args_json = serde_json::to_string(args_value).unwrap_or_default();
|
||||
if tx
|
||||
.send(StreamEvent::ToolExecutionStarted {
|
||||
tool_name: tool_name.clone(),
|
||||
tool_call_id: tool_call_id.clone(),
|
||||
arguments: args_json,
|
||||
})
|
||||
.is_err()
|
||||
{
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
||||
// ⑨ 执行工具
|
||||
let results = tool_registry.invoke_all(calls, tool_timeout).await;
|
||||
|
||||
// ⑩ 发送 ToolExecutionCompleted
|
||||
for result in &results {
|
||||
let summary = match &result.output {
|
||||
Ok(v) => serde_json::to_string(v).unwrap_or_default(),
|
||||
Err(e) => e.to_string(),
|
||||
};
|
||||
let truncated = truncate_tool_result(&summary, max_bytes);
|
||||
if tx
|
||||
.send(StreamEvent::ToolExecutionCompleted {
|
||||
tool_name: result.tool_name.clone(),
|
||||
tool_call_id: result.tool_call_id.clone(),
|
||||
result_summary: truncated,
|
||||
is_error: result.output.is_err(),
|
||||
})
|
||||
.is_err()
|
||||
{
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
||||
// ⑪ push 工具结果到 messages(区分可恢复/不可恢复)
|
||||
for result in results {
|
||||
let is_error = result.output.is_err();
|
||||
let content = match &result.output {
|
||||
Ok(v) => {
|
||||
// ponytail: 与 submit_with_tools 行为对齐 —— 用 truncate_tool_result
|
||||
// 截断结果以防止超大工具输出在 tool 循环中膨胀 messages 上下文窗口
|
||||
let serialized =
|
||||
serde_json::to_string(v).unwrap_or_else(|e| {
|
||||
tracing::warn!("工具结果序列化失败: {}", e);
|
||||
"{}".to_string()
|
||||
});
|
||||
truncate_tool_result(&serialized, max_bytes)
|
||||
}
|
||||
Err(e) if e.is_recoverable() => format!("错误: {}", e),
|
||||
Err(e) => {
|
||||
// 不可恢复错误 —— 终止循环
|
||||
let _ = tx.send(StreamEvent::Error {
|
||||
message: format!("工具 '{}' 不可恢复错误: {}", result.tool_name, e),
|
||||
});
|
||||
return;
|
||||
}
|
||||
};
|
||||
messages.push(Message::tool_result(result.tool_call_id, content, is_error));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
@@ -917,4 +1162,541 @@ mod tests {
|
||||
let truncated = truncate_tool_result(&s, 50);
|
||||
assert!(truncated.starts_with("中"));
|
||||
}
|
||||
|
||||
// ====== Phase 9: submit_with_tools_stream 单元测试 ======
|
||||
//
|
||||
// ponytail: 由于 in-mod MockProvider 的 chat_stream 是 unimplemented!(),
|
||||
// 这里使用公开的 `crate::llm::mock::MockProvider`(已实现完整 chat_stream + 预设队列)。
|
||||
// 测试需使用 `tokio::test(flavor = "multi_thread")` 满足 submit_with_tools_stream 内部
|
||||
// 的 `tokio::spawn` 运行时要求。
|
||||
|
||||
use crate::llm::mock::MockProvider as Mock;
|
||||
use futures_util::StreamExt;
|
||||
|
||||
/// 收集流中所有事件到 Vec。
|
||||
async fn drain(stream: Pin<Box<dyn Stream<Item = StreamEvent> + Send>>) -> Vec<StreamEvent> {
|
||||
let mut out = Vec::new();
|
||||
let mut s = stream;
|
||||
while let Some(ev) = s.next().await {
|
||||
out.push(ev);
|
||||
}
|
||||
out
|
||||
}
|
||||
|
||||
/// Phase 9 测试 3.1 — 纯文本流(无 tool_use)
|
||||
#[tokio::test(flavor = "multi_thread")]
|
||||
async fn test_submit_with_tools_stream_pure_text() {
|
||||
let provider = Mock::new(vec![assistant_text_response("你好")]);
|
||||
let mut cycle =
|
||||
LlmCycle::new(Box::new(provider), CycleConfig::default());
|
||||
let mut registry = ToolRegistry::new();
|
||||
registry.register(std::sync::Arc::new(AddTool)).unwrap();
|
||||
|
||||
let stream = cycle
|
||||
.submit_with_tools_stream("问个问题".to_string(), std::sync::Arc::new(registry))
|
||||
.await
|
||||
.unwrap();
|
||||
let events = drain(stream).await;
|
||||
|
||||
// 期望序列:MessageStart → ContentBlockStart(Text) → TextDelta → ContentBlockEnd → CostUpdate → MessageComplete
|
||||
assert!(matches!(events.first(), Some(StreamEvent::MessageStart { .. })));
|
||||
assert!(events
|
||||
.iter()
|
||||
.any(|e| matches!(e, StreamEvent::TextDelta { text } if text == "你好")));
|
||||
assert!(events
|
||||
.iter()
|
||||
.any(|e| matches!(e, StreamEvent::MessageComplete { .. })));
|
||||
// 纯文本流不应有 ToolExecutionStarted/Completed 事件
|
||||
assert!(!events
|
||||
.iter()
|
||||
.any(|e| matches!(e, StreamEvent::ToolExecutionStarted { .. })));
|
||||
assert!(!events
|
||||
.iter()
|
||||
.any(|e| matches!(e, StreamEvent::ToolExecutionCompleted { .. })));
|
||||
assert!(!events
|
||||
.iter()
|
||||
.any(|e| matches!(e, StreamEvent::Error { .. })));
|
||||
}
|
||||
|
||||
/// Phase 9 测试 3.2 — 单轮工具调用
|
||||
#[tokio::test(flavor = "multi_thread")]
|
||||
async fn test_submit_with_tools_stream_single_tool() {
|
||||
let provider = Mock::new(vec![
|
||||
assistant_tool_call_response(vec![("call_1", "add", r#"{"a":1,"b":2}"#)]),
|
||||
assistant_text_response("答案是 3"),
|
||||
]);
|
||||
let mut cycle =
|
||||
LlmCycle::new(Box::new(provider), CycleConfig::default());
|
||||
let mut registry = ToolRegistry::new();
|
||||
registry.register(std::sync::Arc::new(AddTool)).unwrap();
|
||||
|
||||
let stream = cycle
|
||||
.submit_with_tools_stream("1+2".to_string(), std::sync::Arc::new(registry))
|
||||
.await
|
||||
.unwrap();
|
||||
let events = drain(stream).await;
|
||||
|
||||
// 应有 1 对 ToolExecutionStarted / ToolExecutionCompleted
|
||||
let started_count = events
|
||||
.iter()
|
||||
.filter(|e| matches!(e, StreamEvent::ToolExecutionStarted { .. }))
|
||||
.count();
|
||||
let completed_count = events
|
||||
.iter()
|
||||
.filter(|e| matches!(e, StreamEvent::ToolExecutionCompleted { .. }))
|
||||
.count();
|
||||
assert_eq!(started_count, 1, "应有 1 个 ToolExecutionStarted");
|
||||
assert_eq!(completed_count, 1, "应有 1 个 ToolExecutionCompleted");
|
||||
|
||||
// 验证 ToolExecutionStarted.arguments 携带有效 JSON
|
||||
let started = events
|
||||
.iter()
|
||||
.find_map(|e| match e {
|
||||
StreamEvent::ToolExecutionStarted {
|
||||
tool_name,
|
||||
tool_call_id,
|
||||
arguments,
|
||||
} => Some((tool_name, tool_call_id, arguments)),
|
||||
_ => None,
|
||||
})
|
||||
.unwrap();
|
||||
assert_eq!(started.0, "add");
|
||||
assert_eq!(started.1, "call_1");
|
||||
assert!(!started.2.is_empty(), "arguments 应携带实际 JSON 参数");
|
||||
|
||||
// 验证 ToolExecutionCompleted 的 result_summary
|
||||
let completed = events
|
||||
.iter()
|
||||
.find_map(|e| match e {
|
||||
StreamEvent::ToolExecutionCompleted {
|
||||
tool_name,
|
||||
tool_call_id,
|
||||
result_summary,
|
||||
is_error,
|
||||
} => Some((tool_name, tool_call_id, result_summary, *is_error)),
|
||||
_ => None,
|
||||
})
|
||||
.unwrap();
|
||||
assert_eq!(completed.0, "add");
|
||||
assert_eq!(completed.1, "call_1");
|
||||
assert!(!completed.2.is_empty());
|
||||
assert!(!completed.3);
|
||||
|
||||
// 应有最终 MessageComplete { stop_reason: Stop }
|
||||
let final_response = events
|
||||
.iter()
|
||||
.rev()
|
||||
.find_map(|e| match e {
|
||||
StreamEvent::MessageComplete { full_response } => Some(full_response.clone()),
|
||||
_ => None,
|
||||
})
|
||||
.unwrap();
|
||||
assert_eq!(final_response.stop_reason, StopReason::Stop);
|
||||
}
|
||||
|
||||
/// Phase 9 测试 3.3 — 多轮工具调用
|
||||
#[tokio::test(flavor = "multi_thread")]
|
||||
async fn test_submit_with_tools_stream_multi_tool() {
|
||||
let provider = Mock::new(vec![
|
||||
assistant_tool_call_response(vec![("call_1", "add", r#"{"a":1,"b":2}"#)]),
|
||||
assistant_tool_call_response(vec![("call_2", "add", r#"{"a":3,"b":4}"#)]),
|
||||
assistant_tool_call_response(vec![("call_3", "add", r#"{"a":5,"b":6}"#)]),
|
||||
assistant_text_response("完成"),
|
||||
]);
|
||||
let mut cycle =
|
||||
LlmCycle::new(Box::new(provider), CycleConfig::default());
|
||||
let mut registry = ToolRegistry::new();
|
||||
registry.register(std::sync::Arc::new(AddTool)).unwrap();
|
||||
|
||||
let stream = cycle
|
||||
.submit_with_tools_stream("计算总和".to_string(), std::sync::Arc::new(registry))
|
||||
.await
|
||||
.unwrap();
|
||||
let events = drain(stream).await;
|
||||
|
||||
let started_count = events
|
||||
.iter()
|
||||
.filter(|e| matches!(e, StreamEvent::ToolExecutionStarted { .. }))
|
||||
.count();
|
||||
let completed_count = events
|
||||
.iter()
|
||||
.filter(|e| matches!(e, StreamEvent::ToolExecutionCompleted { .. }))
|
||||
.count();
|
||||
assert_eq!(started_count, 3, "3 轮工具调用");
|
||||
assert_eq!(completed_count, 3, "3 个 ToolExecutionCompleted");
|
||||
|
||||
// 应有 4 个 MessageComplete(每个 LLM 调用独立一个)
|
||||
let complete_count = events
|
||||
.iter()
|
||||
.filter(|e| matches!(e, StreamEvent::MessageComplete { .. }))
|
||||
.count();
|
||||
assert_eq!(complete_count, 4, "4 个 LLM 调用 → 4 个 MessageComplete");
|
||||
}
|
||||
|
||||
/// Phase 9 测试 3.4 — 最大轮次超限
|
||||
#[tokio::test(flavor = "multi_thread")]
|
||||
async fn test_submit_with_tools_stream_max_turns_exceeded() {
|
||||
let config = CycleConfig {
|
||||
max_tool_turns: Some(2),
|
||||
..Default::default()
|
||||
};
|
||||
let provider = Mock::new(vec![
|
||||
assistant_tool_call_response(vec![("c1", "add", r#"{"a":1,"b":1}"#)]),
|
||||
assistant_tool_call_response(vec![("c2", "add", r#"{"a":1,"b":1}"#)]),
|
||||
assistant_tool_call_response(vec![("c3", "add", r#"{"a":1,"b":1}"#)]),
|
||||
]);
|
||||
let mut cycle = LlmCycle::new(Box::new(provider), config);
|
||||
let mut registry = ToolRegistry::new();
|
||||
registry.register(std::sync::Arc::new(AddTool)).unwrap();
|
||||
|
||||
let stream = cycle
|
||||
.submit_with_tools_stream("test".to_string(), std::sync::Arc::new(registry))
|
||||
.await
|
||||
.unwrap();
|
||||
let events = drain(stream).await;
|
||||
|
||||
// 流中应有 Error 事件(最大轮次超限)
|
||||
let error_events: Vec<_> = events
|
||||
.iter()
|
||||
.filter_map(|e| match e {
|
||||
StreamEvent::Error { message } => Some(message.clone()),
|
||||
_ => None,
|
||||
})
|
||||
.collect();
|
||||
assert!(
|
||||
error_events.iter().any(|m| m.contains("达到最大工具循环轮次")),
|
||||
"应包含最大轮次超限 Error,实际: {:?}", error_events
|
||||
);
|
||||
|
||||
// 工具调用次数应 ≤ 2
|
||||
let started_count = events
|
||||
.iter()
|
||||
.filter(|e| matches!(e, StreamEvent::ToolExecutionStarted { .. }))
|
||||
.count();
|
||||
assert_eq!(started_count, 2, "工具调用应在第 2 轮后停止");
|
||||
}
|
||||
|
||||
/// Phase 9 测试 3.5 — `chat_stream` 返回 Err
|
||||
///
|
||||
/// 使用自定义 MockProvider 返回 chat_stream Err。
|
||||
#[tokio::test(flavor = "multi_thread")]
|
||||
async fn test_submit_with_tools_stream_chat_stream_err() {
|
||||
use crate::llm::provider::{ProviderCapabilities, ProviderFeatures};
|
||||
|
||||
struct ErrProvider;
|
||||
#[async_trait]
|
||||
impl LlmProvider for ErrProvider {
|
||||
async fn chat(&self, _r: MessageRequest) -> Result<MessageResponse, LlmError> {
|
||||
Err(LlmError::Other("网络错误".into()))
|
||||
}
|
||||
async fn chat_stream(
|
||||
&self,
|
||||
_r: MessageRequest,
|
||||
) -> Result<
|
||||
Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>,
|
||||
LlmError,
|
||||
> {
|
||||
Err(LlmError::Other("网络错误".into()))
|
||||
}
|
||||
fn capabilities(&self) -> ProviderCapabilities {
|
||||
ProviderCapabilities {
|
||||
provider_name: "err",
|
||||
supported_models: None,
|
||||
features: ProviderFeatures::default(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let mut cycle = LlmCycle::new(Box::new(ErrProvider), CycleConfig::default());
|
||||
let mut registry = ToolRegistry::new();
|
||||
registry.register(std::sync::Arc::new(AddTool)).unwrap();
|
||||
|
||||
let stream = cycle
|
||||
.submit_with_tools_stream("test".to_string(), std::sync::Arc::new(registry))
|
||||
.await
|
||||
.unwrap();
|
||||
let events = drain(stream).await;
|
||||
|
||||
// 第一个事件应是 Error(chat_stream Err 立即事件化)
|
||||
assert!(
|
||||
events.first().map(|e| matches!(e, StreamEvent::Error { .. }))
|
||||
== Some(true),
|
||||
"流中首个事件应是 Error,实际: {:?}", events.first()
|
||||
);
|
||||
assert!(events
|
||||
.iter()
|
||||
.filter_map(|e| match e {
|
||||
StreamEvent::Error { message } => Some(message.as_str()),
|
||||
_ => None,
|
||||
})
|
||||
.any(|m| m.contains("网络错误")));
|
||||
}
|
||||
|
||||
/// Phase 9 测试 3.6 — 空 tool_registry
|
||||
#[tokio::test(flavor = "multi_thread")]
|
||||
async fn test_submit_with_tools_stream_empty_registry() {
|
||||
let provider = Mock::new(vec![assistant_text_response("纯文本回答")]);
|
||||
let mut cycle =
|
||||
LlmCycle::new(Box::new(provider), CycleConfig::default());
|
||||
let registry = ToolRegistry::new();
|
||||
|
||||
let stream = cycle
|
||||
.submit_with_tools_stream("test".to_string(), std::sync::Arc::new(registry))
|
||||
.await
|
||||
.unwrap();
|
||||
let events = drain(stream).await;
|
||||
|
||||
// 流退化为纯文本流 —— 无 ToolExecution 事件,无 Error
|
||||
assert!(events
|
||||
.iter()
|
||||
.any(|e| matches!(e, StreamEvent::TextDelta { text } if text == "纯文本回答")));
|
||||
assert!(!events
|
||||
.iter()
|
||||
.any(|e| matches!(e, StreamEvent::ToolExecutionStarted { .. })));
|
||||
assert!(!events
|
||||
.iter()
|
||||
.any(|e| matches!(e, StreamEvent::ToolExecutionCompleted { .. })));
|
||||
assert!(!events
|
||||
.iter()
|
||||
.any(|e| matches!(e, StreamEvent::Error { .. })));
|
||||
}
|
||||
|
||||
/// Phase 9 测试 3.7 — 不可恢复工具错误
|
||||
struct UnrecoverableTool;
|
||||
#[async_trait]
|
||||
impl BaseTool for UnrecoverableTool {
|
||||
fn name(&self) -> &str {
|
||||
"fail_unrecoverable"
|
||||
}
|
||||
fn description(&self) -> &str {
|
||||
"不可恢复地失败"
|
||||
}
|
||||
fn parameters(&self) -> Value {
|
||||
json!({})
|
||||
}
|
||||
async fn execute(
|
||||
&self,
|
||||
_args: Value,
|
||||
_ctx: &crate::tools::ToolContext<'_>,
|
||||
) -> Result<Value, crate::tools::ToolError> {
|
||||
// 不可恢复错误 —— NotFound 表示工具在执行时失败且不可恢复
|
||||
Err(crate::tools::ToolError::NotFound("永久失败".into()))
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "multi_thread")]
|
||||
async fn test_submit_with_tools_stream_unrecoverable_tool_error() {
|
||||
let provider = Mock::new(vec![
|
||||
assistant_tool_call_response(vec![("call_x", "fail_unrecoverable", "{}")]),
|
||||
assistant_text_response("忽略"),
|
||||
]);
|
||||
let mut cycle =
|
||||
LlmCycle::new(Box::new(provider), CycleConfig::default());
|
||||
let mut registry = ToolRegistry::new();
|
||||
registry
|
||||
.register(std::sync::Arc::new(UnrecoverableTool))
|
||||
.unwrap();
|
||||
|
||||
let stream = cycle
|
||||
.submit_with_tools_stream("test".to_string(), std::sync::Arc::new(registry))
|
||||
.await
|
||||
.unwrap();
|
||||
let events = drain(stream).await;
|
||||
|
||||
// 流中应有 Error 事件(不可恢复错误)
|
||||
let error_events: Vec<_> = events
|
||||
.iter()
|
||||
.filter_map(|e| match e {
|
||||
StreamEvent::Error { message } => Some(message.clone()),
|
||||
_ => None,
|
||||
})
|
||||
.collect();
|
||||
assert!(
|
||||
error_events.iter().any(|m| m.contains("不可恢复")),
|
||||
"应包含不可恢复错误 Error"
|
||||
);
|
||||
}
|
||||
|
||||
/// Phase 9 测试 3.8 — 可恢复工具错误
|
||||
struct RecoverableTool;
|
||||
#[async_trait]
|
||||
impl BaseTool for RecoverableTool {
|
||||
fn name(&self) -> &str {
|
||||
"fail_recoverable"
|
||||
}
|
||||
fn description(&self) -> &str {
|
||||
"可恢复失败"
|
||||
}
|
||||
fn parameters(&self) -> Value {
|
||||
json!({})
|
||||
}
|
||||
async fn execute(
|
||||
&self,
|
||||
_args: Value,
|
||||
_ctx: &crate::tools::ToolContext<'_>,
|
||||
) -> Result<Value, crate::tools::ToolError> {
|
||||
// ExecutionFailed = 可恢复错误(is_recoverable() == true)
|
||||
Err(crate::tools::ToolError::ExecutionFailed(
|
||||
"工具暂时失败".into(),
|
||||
"网络抖动".into(),
|
||||
))
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "multi_thread")]
|
||||
async fn test_submit_with_tools_stream_recoverable_tool_error() {
|
||||
let provider = Mock::new(vec![
|
||||
assistant_tool_call_response(vec![("call_y", "fail_recoverable", "{}")]),
|
||||
assistant_text_response("已恢复"),
|
||||
]);
|
||||
let mut cycle =
|
||||
LlmCycle::new(Box::new(provider), CycleConfig::default());
|
||||
let mut registry = ToolRegistry::new();
|
||||
registry
|
||||
.register(std::sync::Arc::new(RecoverableTool))
|
||||
.unwrap();
|
||||
|
||||
let stream = cycle
|
||||
.submit_with_tools_stream("test".to_string(), std::sync::Arc::new(registry))
|
||||
.await
|
||||
.unwrap();
|
||||
let events = drain(stream).await;
|
||||
|
||||
// 可恢复错误:tool_result 回传 LLM,最终流正常结束
|
||||
assert!(!events
|
||||
.iter()
|
||||
.any(|e| matches!(e, StreamEvent::Error { .. })));
|
||||
// 最终 MessageComplete 应是 Stop(不是 ToolUse)
|
||||
let final_response = events
|
||||
.iter()
|
||||
.rev()
|
||||
.find_map(|e| match e {
|
||||
StreamEvent::MessageComplete { full_response } => Some(full_response.clone()),
|
||||
_ => None,
|
||||
})
|
||||
.unwrap();
|
||||
assert_eq!(final_response.stop_reason, StopReason::Stop);
|
||||
|
||||
// ToolExecutionCompleted.is_error 应为 true
|
||||
let completed = events
|
||||
.iter()
|
||||
.find_map(|e| match e {
|
||||
StreamEvent::ToolExecutionCompleted { is_error, .. } => Some(*is_error),
|
||||
_ => None,
|
||||
})
|
||||
.unwrap();
|
||||
assert!(completed, "可恢复错误的 ToolExecutionCompleted.is_error 应为 true");
|
||||
}
|
||||
|
||||
/// Phase 9 测试 3.9 — 工具超时
|
||||
///
|
||||
/// 使用一个永远 sleep 的工具 + tool_timeout_secs: 1,验证 TimeoutError 路径。
|
||||
///
|
||||
/// 注:`tokio::time::timeout` 在 `invoke_all` 中将超时转为 `ToolError::McpTimeout("timeout")`。
|
||||
/// 当前 `McpTimeout.is_recoverable() == false`(见 `tools/error.rs`),因此工具超时视为不可恢复:
|
||||
/// - 流中应出现 `ToolExecutionCompleted { is_error: true }`
|
||||
/// - 然后出现 `StreamEvent::Error`(不可恢复错误终止循环,§3.6 错误表)
|
||||
/// - 第一轮的 `MessageComplete { stop_reason: ToolUse }` 在 Error 事件之前已发出
|
||||
/// - 不会再有第二轮 LLM 调用(与方案 §3.6 一致)
|
||||
#[tokio::test(flavor = "multi_thread")]
|
||||
async fn test_submit_with_tools_stream_tool_timeout() {
|
||||
struct SlowTool;
|
||||
#[async_trait]
|
||||
impl BaseTool for SlowTool {
|
||||
fn name(&self) -> &str {
|
||||
"slow_tool"
|
||||
}
|
||||
fn description(&self) -> &str {
|
||||
"慢工具,模拟超时"
|
||||
}
|
||||
fn parameters(&self) -> Value {
|
||||
json!({})
|
||||
}
|
||||
async fn execute(
|
||||
&self,
|
||||
_args: Value,
|
||||
_ctx: &crate::tools::ToolContext<'_>,
|
||||
) -> Result<Value, crate::tools::ToolError> {
|
||||
// sleep 超过 5s(tool_timeout = 1s),触发 invoke_all 的 tokio::time::timeout
|
||||
tokio::time::sleep(std::time::Duration::from_secs(5)).await;
|
||||
Ok(json!({"ok": true}))
|
||||
}
|
||||
}
|
||||
|
||||
let config = CycleConfig {
|
||||
tool_timeout_secs: 1,
|
||||
..Default::default()
|
||||
};
|
||||
// ponytail: 预设第二个响应不会消费——超时后流立即终止,不会发起第二轮 LLM 调用
|
||||
let provider = Mock::new(vec![
|
||||
assistant_tool_call_response(vec![("call_z", "slow_tool", "{}")]),
|
||||
assistant_text_response("不会到达"),
|
||||
]);
|
||||
let mut cycle = LlmCycle::new(Box::new(provider), config);
|
||||
let mut registry = ToolRegistry::new();
|
||||
registry
|
||||
.register(std::sync::Arc::new(SlowTool))
|
||||
.unwrap();
|
||||
|
||||
let stream = cycle
|
||||
.submit_with_tools_stream("test".to_string(), std::sync::Arc::new(registry))
|
||||
.await
|
||||
.unwrap();
|
||||
let events = drain(stream).await;
|
||||
|
||||
// 1. ToolExecutionCompleted 应报告 is_error=true(McpTimeout 视为失败)
|
||||
let tool_completed: Vec<_> = events
|
||||
.iter()
|
||||
.filter_map(|e| match e {
|
||||
StreamEvent::ToolExecutionCompleted {
|
||||
is_error,
|
||||
tool_name,
|
||||
..
|
||||
} => Some((*is_error, tool_name.clone())),
|
||||
_ => None,
|
||||
})
|
||||
.collect();
|
||||
assert!(!tool_completed.is_empty(), "应有 ToolExecutionCompleted 事件");
|
||||
assert!(tool_completed[0].0, "超时后 ToolExecutionCompleted.is_error 应为 true");
|
||||
assert_eq!(tool_completed[0].1, "slow_tool");
|
||||
|
||||
// 2. 流中应有不可恢复错误终止事件(tool_timeout → McpTimeout → 不可恢复 → Error)
|
||||
let error_events: Vec<_> = events
|
||||
.iter()
|
||||
.filter_map(|e| match e {
|
||||
StreamEvent::Error { message } => Some(message.clone()),
|
||||
_ => None,
|
||||
})
|
||||
.collect();
|
||||
assert!(
|
||||
error_events
|
||||
.iter()
|
||||
.any(|m| m.contains("不可恢复错误")),
|
||||
"应有不可恢复错误事件终止流,实际事件: {:?}", error_events
|
||||
);
|
||||
|
||||
// 3. 第一轮的 MessageComplete { stop_reason: ToolUse } 在 Error 之前已发出
|
||||
let first_complete = events
|
||||
.iter()
|
||||
.find_map(|e| match e {
|
||||
StreamEvent::MessageComplete { full_response } => Some(full_response.clone()),
|
||||
_ => None,
|
||||
})
|
||||
.expect("第一轮 MessageComplete 应存在");
|
||||
assert_eq!(
|
||||
first_complete.stop_reason,
|
||||
StopReason::ToolUse,
|
||||
"第一轮 LLM 响应 stop_reason 应为 ToolUse"
|
||||
);
|
||||
|
||||
// 4. 不会发起第二轮 LLM —— 流中只有 1 个 MessageComplete
|
||||
let complete_count = events
|
||||
.iter()
|
||||
.filter(|e| matches!(e, StreamEvent::MessageComplete { .. }))
|
||||
.count();
|
||||
assert_eq!(
|
||||
complete_count, 1,
|
||||
"超时后不应有第二轮 LLM 流,应只有 1 个 MessageComplete"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -192,6 +192,24 @@ pub enum StreamEvent {
|
||||
MessageComplete { full_response: MessageResponse },
|
||||
/// 错误事件。
|
||||
Error { message: String },
|
||||
/// 工具开始执行 —— 在 `ToolCallEnd` 之后、`registry.invoke_all` 之前发出。
|
||||
/// 让 UI 层可以显示 "正在执行工具:add(1, 2)"。
|
||||
ToolExecutionStarted {
|
||||
tool_name: String,
|
||||
tool_call_id: String,
|
||||
/// 工具参数(JSON 字符串形式),用于 UI 展示
|
||||
arguments: String,
|
||||
},
|
||||
/// 工具执行完成 —— 在工具返回后、新一轮 LLM 流开始之前发出。
|
||||
ToolExecutionCompleted {
|
||||
tool_name: String,
|
||||
tool_call_id: String,
|
||||
/// 结果摘要(由 `CycleConfig.max_tool_result_bytes` 截断,默认 65536 字节/字符边界安全),
|
||||
/// 用于 UI 反馈。完整结果已在内部 `messages` 中作为 `ToolResult` 回传给 LLM。
|
||||
result_summary: String,
|
||||
/// 是否执行出错
|
||||
is_error: bool,
|
||||
},
|
||||
}
|
||||
|
||||
/// 流式响应累积状态。
|
||||
@@ -352,6 +370,9 @@ impl PartialMessageResponse {
|
||||
self.is_errored = true;
|
||||
false
|
||||
}
|
||||
// 元事件:不参与内容块累积,不修改 partial 状态
|
||||
//(Phase 9 —— 工具执行透明化,由 run_tool_loop 在工具前后插入)
|
||||
StreamEvent::ToolExecutionStarted { .. } | StreamEvent::ToolExecutionCompleted { .. } => true,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user