diff --git a/docs/16-phase9-streaming-experience.md b/docs/16-phase9-streaming-experience.md index 507f9ef..7472f0e 100644 --- a/docs/16-phase9-streaming-experience.md +++ b/docs/16-phase9-streaming-experience.md @@ -241,6 +241,12 @@ pub async fn submit_turn_stream( pub async fn finalize_turn(&mut self, response: &MessageResponse) ``` +> **实施偏差(Phase 10 适配)**:实际签名扩展为 +> `pub async fn finalize_turn(&mut self, response: &MessageResponse, new_messages_from_cycle: Vec) -> Result<(), AgentError>`。 +> - `new_messages_from_cycle`:本轮新增消息(`[user_input, ...tool_results, final_response]`),由消费者在流消费完毕后从 `cycle.messages()[input_len..]` 提取并传入;`finalize_turn` 增量追加到当前 slot(不覆盖已有消息)。 +> - 返回 `Result<(), AgentError>`:错误传播更清晰,与 `submit_turn` 的 slot 边界错误(`SlotReadonly` / `SlotNotFound`)对齐。 +> - Phase 10 ContextSlot 实施时扩展。Phase 9 消费者若不接入 slot 持久化,可传 `vec![response.message.clone()]` 兜底。 + ### 3.5 消费者使用模式 ```rust @@ -410,10 +416,10 @@ if let Some(response) = final_response { **目标**:端到端验证 `submit_turn_stream` + `finalize_turn` 的完整链路,确保零回归。 -**新增**(`agent/session.rs` 内联测试): +**新增**(`agent/session.rs` 内联测试,2026-07-08 实施审查补全): -- **场景**:`submit_turn_stream` 跑通 mock provider → 消费流(验证各事件到达) → `finalize_turn` 后 cost 更新正确 -- **场景**:verify `OnTurnStart` hook 在 `submit_turn_stream` 返回流之前已触发 +- `submit_turn_stream_end_to_end` — `submit_turn_stream` 跑通 mock provider → 消费流(验证收到 TextDelta + MessageComplete) → `finalize_turn` 后 `cost_so_far` 正确更新(`prompt_tokens=10, completion_tokens=5`) + `turn_index=1` + default slot 包含 user/assistant 消息 +- `submit_turn_stream_triggers_turn_hooks` — 验证 `OnTurnStart` 在 `submit_turn_stream` 返回流之前已触发(计数=1)+ `OnTurnEnd` 在 `finalize_turn` 之前**不**触发(计数=0)+ `finalize_turn` 后 `OnTurnEnd` 触发(计数=1) **验证**: diff --git a/src/agent/session.rs b/src/agent/session.rs index 2f714d6..6a0e422 100644 --- a/src/agent/session.rs +++ b/src/agent/session.rs @@ -418,9 +418,6 @@ impl AgentSession { // 6. turn_index 递增 —— 配合 finalize_turn 用 (turn_index - 1) 传递正确的 OnTurnEnd 序号 self.turn_index += 1; - // 注:hook_executor 不显式 drop,生命周期由 Arc 自动管理 - let _ = hook_executor; - Ok(stream) } @@ -475,10 +472,12 @@ mod tests { use crate::agent::builder::AgentBuilder; use crate::llm::hooks::{Hook, HookContext, HookExecutor, HookResult}; use crate::llm::mock::MockProvider; + use crate::llm::stream::StreamEvent; use crate::llm::types::message::ContentBlock; use crate::llm::types::response_v2::{MessageResponse, StopReason}; use crate::tools::ToolRegistry; use async_trait::async_trait; + use futures_util::StreamExt; use std::sync::atomic::{AtomicU32, Ordering}; /// 计数 hook —— 每被调用一次 +1。 @@ -948,4 +947,107 @@ mod tests { .unwrap_err(); assert!(matches!(err, AgentError::SlotReadonly(_))); } + + // ====== Phase 9 Step 5: 集成测试 ====== + + /// Phase 9 Step 5.1 — `submit_turn_stream` 端到端链路。 + /// + /// 验证:mock provider → `submit_turn_stream` 消费流 → 收到 TextDelta + MessageComplete + /// → `finalize_turn` 后 `cost_so_far` 正确更新,turn_index 递增。 + #[tokio::test(flavor = "multi_thread")] + async fn submit_turn_stream_end_to_end() { + let (mut session, _, _) = build_session(vec![assistant_text("hi back")]); + + let mut stream = session + .submit_turn_stream("user msg") + .await + .expect("submit_turn_stream 应成功"); + + // 消费流并提取 MessageComplete + let mut final_response: Option = None; + while let Some(ev) = stream.next().await { + if let StreamEvent::MessageComplete { full_response } = &ev { + final_response = Some(full_response.clone()); + } + } + + let response = final_response.expect("流中应有 MessageComplete"); + // consumer 负责构造本轮新增消息列表(user_input + assistant_response)。 + // submit_turn_stream 不会自动写入 self.slots(流是延迟的), + // 消费者需在 finalize_turn 时把 [user_input, ...tool_results, final_response] 一并传入。 + let new_messages = vec![Message::user_text("user msg"), response.message.clone()]; + session + .finalize_turn(&response, new_messages) + .await + .expect("finalize_turn 应成功"); + + // cost_so_far 已累计(assistant_text 的 usage 是 from_input_output(10, 5)) + assert_eq!(session.usage().total().prompt_tokens, 10); + assert_eq!(session.usage().total().completion_tokens, 5); + // turn_index 已递增 + assert_eq!(session.turn_index(), 1); + + // default slot 应包含 user 输入和 assistant 响应 + let slot = session.slots.get("default").expect("default slot"); + let has_user = slot + .messages + .iter() + .any(|m| matches!(m, Message::User { .. })); + let has_resp = slot.messages.iter().any(|m| extract_text(m) == "hi back"); + assert!(has_user && has_resp, "default slot 应包含 user 和 assistant 消息"); + } + + /// Phase 9 Step 5.2 — `submit_turn_stream` 触发 OnTurnStart / OnTurnEnd hook。 + /// + /// 验证:OnTurnStart 在 `submit_turn_stream` 返回流之前已触发; + /// OnTurnEnd 在 `finalize_turn` 调用后才触发。 + #[tokio::test(flavor = "multi_thread")] + async fn submit_turn_stream_triggers_turn_hooks() { + let (mut session, start_count, end_count) = build_session(vec![assistant_text("ok")]); + + // 初始状态:两个 hook 计数都是 0 + assert_eq!(start_count.0.load(Ordering::SeqCst), 0); + assert_eq!(end_count.0.load(Ordering::SeqCst), 0); + + // 调 submit_turn_stream + let mut stream = session + .submit_turn_stream("user msg") + .await + .expect("submit_turn_stream 应成功"); + + // OnTurnStart 应在流返回前已触发 + assert_eq!( + start_count.0.load(Ordering::SeqCst), + 1, + "OnTurnStart 应在 submit_turn_stream 返回流之前触发" + ); + // OnTurnEnd 此时尚未触发 + assert_eq!( + end_count.0.load(Ordering::SeqCst), + 0, + "OnTurnEnd 不应在 submit_turn_stream 阶段触发" + ); + + // 消费流 + let mut final_response: Option = None; + while let Some(ev) = stream.next().await { + if let StreamEvent::MessageComplete { full_response } = &ev { + final_response = Some(full_response.clone()); + } + } + + // finalize_turn + let response = final_response.expect("流中应有 MessageComplete"); + session + .finalize_turn(&response, vec![]) + .await + .expect("finalize_turn 应成功"); + + // OnTurnEnd 已触发 + assert_eq!( + end_count.0.load(Ordering::SeqCst), + 1, + "OnTurnEnd 应在 finalize_turn 后触发" + ); + } } \ No newline at end of file diff --git a/src/llm/cycle.rs b/src/llm/cycle.rs index 02d1059..e4e126b 100644 --- a/src/llm/cycle.rs +++ b/src/llm/cycle.rs @@ -18,7 +18,7 @@ use tokio_stream::wrappers::UnboundedReceiverStream; use crate::llm::compact::{CompactConfig, CompactState, microcompact, should_compact}; use crate::llm::cycle::retry::should_retry; use crate::llm::error::LlmError; -use crate::llm::hooks::{HookContext, HookExecutor}; +use crate::llm::hooks::{HookContext, HookEvent, HookExecutor}; use crate::llm::provider::LlmProvider; use crate::llm::stream::StreamEvent; use crate::llm::types::message::{ContentBlock, Message}; @@ -774,8 +774,21 @@ async fn run_tool_loop( ..Default::default() }; - // ② PreRequest hook(fire-and-forget,仅占位保留以保持接口对称) - let _ = hook_executor.as_ref(); + // ② PreRequest hook —— 与 `submit_with_tools` / `submit_stream` 行为对齐: + // 触发 hook → 检查 should_block → 阻断则事件化 Error + return 结束 task。 + // 阻断原因透传,让消费者看到完整的拒绝原因。 + if let Some(ref executor) = hook_executor { + let ctx = HookContext::new(HookEvent::PreRequest).with_request(&request); + let results = executor.execute(HookEvent::PreRequest, &ctx).await; + if let Some(blocking) = results.iter().find(|r| r.should_block) { + let reason = blocking + .reason + .clone() + .unwrap_or_else(|| "Blocked by pre-request hook".to_string()); + let _ = tx.send(StreamEvent::Error { message: reason }); + return; + } + } // ③ chat_stream —— 第一层错误 let mut stream = match provider.chat_stream(request).await {