feat(core): 新增 Phase 9 流式体验增强 submit_turn_stream
- 新增 StreamEvent::ToolExecutionStarted/Completed 变体(apply_to 元事件) - 新增 LlmCycle::submit_with_tools_stream + run_tool_loop(spawn + mpsc 状态机) - 新增 AgentSession::submit_turn_stream + finalize_turn(手动同步状态) - CycleConfig 加 Clone derive - 新增 9 个单元测试 + 2 个集成测试(211 passed) - clippy 0 警告,存量 0 回归 docs: 追加 Phase 9 实施方案(docs/16-phase9-streaming-experience.md)
This commit is contained in:
+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 后触发");
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user