From 32d49601a2b15e77a05e2a35731fd1d44d59d7e1 Mon Sep 17 00:00:00 2001 From: oudecheng <13802883547@139.com> Date: Thu, 2 Jul 2026 17:17:03 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E6=9B=B4=E6=96=B0=E5=8F=96=E6=B6=88?= =?UTF-8?q?=E4=BF=A1=E5=8F=B7=E5=A4=84=E7=90=86=EF=BC=8C=E6=94=AF=E6=8C=81?= =?UTF-8?q?=20interior=20mutability=EF=BC=9B=E4=BC=98=E5=8C=96=20LLM=20?= =?UTF-8?q?=E8=B0=83=E7=94=A8=E4=B8=8E=E5=B7=A5=E5=85=B7=E6=89=A7=E8=A1=8C?= =?UTF-8?q?=E7=9A=84=E5=8F=96=E6=B6=88=E9=80=BB=E8=BE=91?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/agent/agent_loop.rs | 134 ++++++++++++++++++++++++++++++++-------- 1 file changed, 108 insertions(+), 26 deletions(-) diff --git a/src/agent/agent_loop.rs b/src/agent/agent_loop.rs index 5c0ee7f..9960c73 100644 --- a/src/agent/agent_loop.rs +++ b/src/agent/agent_loop.rs @@ -644,8 +644,10 @@ pub struct AgentLoop { observer: Option>, emitted_message_handler: Option>, max_iterations: usize, - /// 取消信号接收端:Agent 在每次迭代开始时检查是否被取消 - cancel_token: Option>, + /// 取消信号接收端:Agent 在每次迭代开始时检查是否被取消。 + /// 包装在 Mutex 中以支持 interior mutability —— + /// watch::Receiver::changed() 需要 &mut self,但 process() 持有 &self。 + cancel_token: Option>>, } #[derive(Debug, Clone)] @@ -862,8 +864,9 @@ impl AgentLoop { /// /// Agent 在每次迭代开始时检查 `cancel_token.has_changed()`, /// 如果已收到取消信号则提前返回。 + /// 同时,LLM 调用和工具执行期间通过 tokio::select! 与 cancel_signal() 竞速。 pub fn with_cancel_token(mut self, token: tokio::sync::watch::Receiver<()>) -> Self { - self.cancel_token = Some(token); + self.cancel_token = Some(tokio::sync::Mutex::new(token)); self } @@ -907,21 +910,14 @@ impl AgentLoop { tracing::debug!(iteration, "Agent iteration started"); // 检查取消信号 - if let Some(ref token) = self.cancel_token { - if token.has_changed().unwrap_or(false) { + // 使用 unwrap_or(true):即使 watch channel 因异常关闭(sender drop 但未 send), + // 也视为取消信号。defense-in-depth —— fail safe 而非 fail silent。 + if let Some(ref mutex) = self.cancel_token { + if mutex.lock().await.has_changed().unwrap_or(true) { tracing::info!(iteration, "Agent execution cancelled by user"); - let cancel_message = format!( - "\n\n[用户已取消执行。已迭代 {} 次,取消前共生成了 {} 条消息。]", - iteration, - emitted_messages.len() - ); - let assistant_message = ChatMessage::assistant(cancel_message); - emitted_messages.push(assistant_message.clone()); - self.emit_live_tool_call_message(assistant_message.clone()).await; - return Ok(AgentProcessResult { - final_response: assistant_message, - emitted_messages, - }); + let cancel = Self::build_cancel_result(iteration, emitted_messages.len()); + self.emit_live_tool_call_message(cancel.final_response.clone()).await; + return Ok(cancel); } } @@ -1016,11 +1012,40 @@ impl AgentLoop { let _ = delta_tx.try_send(delta); }); - let response = match (*self.provider).chat_with_streaming(request, stream_callback).await { + // LLM 调用与取消信号竞速:若取消信号到达,drop LLM future 以 abort HTTP 请求。 + // stream_callback 是 Arc<...>,LLM future 持有其 clone。 + // 取消时需显式 drop 外部 stream_callback 以关闭 mpsc channel, + // 让 consumer_task 自然退出。 + let llm_result: Result< + crate::providers::ChatCompletionResponse, + Box, + >; + if self.cancel_token.is_some() { + tokio::select! { + _ = self.cancel_signal() => { + // LLM future 已被 select! drop → stream_callback clone 已释放。 + // 显式 drop 外部 stream_callback → delta_tx 释放 → channel 关闭。 + drop(stream_callback); + let _ = consumer_task.await; + let cancel = Self::build_cancel_result(iteration, emitted_messages.len()); + self.emit_live_tool_call_message(cancel.final_response.clone()).await; + return Ok(cancel); + } + result = self.provider.chat_with_streaming(request, stream_callback.clone()) => { + llm_result = result; + } + } + } else { + llm_result = self.provider.chat_with_streaming(request, stream_callback).await; + } + + // Close delta channel and wait for consumer to finish processing + // (delta_tx is dropped when the callback closure is dropped) + let _ = consumer_task.await; + + let response = match llm_result { Ok(response) => response, Err(e) => { - // delta_tx is dropped with the callback; await consumer to finish - let _ = consumer_task.await; tracing::error!( provider = %self.provider.name(), model = %self.provider.model_id(), @@ -1039,10 +1064,6 @@ impl AgentLoop { } }; - // Close delta channel and wait for consumer to finish processing - // (delta_tx is dropped when the callback closure is dropped) - let _ = consumer_task.await; - // Signal stream end if handler exists let had_streaming = self.emitted_message_handler.is_some(); if had_streaming { @@ -1121,7 +1142,22 @@ impl AgentLoop { .await; // Execute tools and add results to messages - let tool_results = self.execute_tools(&response.tool_calls).await; + // 工具执行与取消信号竞速:取消时 drop join_all 或 sequential future, + // 未完成的工具调用被丢弃。 + let tool_results = if self.cancel_token.is_some() { + tokio::select! { + _ = self.cancel_signal() => { + let cancel = Self::build_cancel_result(iteration, emitted_messages.len()); + self.emit_live_tool_call_message(cancel.final_response.clone()).await; + return Ok(cancel); + } + results = self.execute_tools(&response.tool_calls) => { + results + } + } + } else { + self.execute_tools(&response.tool_calls).await + }; for (tool_call, result) in response.tool_calls.iter().zip(tool_results.iter()) { // Truncate tool result if too large @@ -1238,7 +1274,27 @@ impl AgentLoop { tools: None, // No tools in final summary call }; - match (*self.provider).chat(request).await { + // 最终 summary 调用也与取消信号竞速 + let final_result: Result< + crate::providers::ChatCompletionResponse, + Box, + >; + if self.cancel_token.is_some() { + tokio::select! { + _ = self.cancel_signal() => { + let cancel = Self::build_cancel_result(self.max_iterations, emitted_messages.len()); + self.emit_live_tool_call_message(cancel.final_response.clone()).await; + return Ok(cancel); + } + result = self.provider.chat(request) => { + final_result = result; + } + } + } else { + final_result = self.provider.chat(request).await; + } + + match final_result { Ok(response) => { let assistant_message = if let Some(reasoning_content) = response.reasoning_content { @@ -1272,6 +1328,32 @@ impl AgentLoop { } } + /// 等待取消信号。若未配置 cancel_token,永远不返回。 + /// + /// 封装了与 watch channel 的交互:changed() 返回 Ok 表示收到信号, + /// 返回 Err(Closed) 表示 sender 已 drop,两种情况都视为取消。 + async fn cancel_signal(&self) { + if let Some(ref mutex) = self.cancel_token { + let mut token = mutex.lock().await; + let _ = token.changed().await; + } else { + std::future::pending::<()>().await; + } + } + + /// 构建取消响应,包含已完成的迭代次数和已生成的消息数量。 + fn build_cancel_result(iteration: usize, emitted_count: usize) -> AgentProcessResult { + let cancel_message = format!( + "\n\n[用户已取消执行。已迭代 {} 次,取消前共生成了 {} 条消息。]", + iteration, emitted_count + ); + let assistant_message = ChatMessage::assistant(cancel_message); + AgentProcessResult { + final_response: assistant_message.clone(), + emitted_messages: vec![assistant_message], + } + } + async fn emit_live_tool_call_message(&self, message: ChatMessage) { if let Some(handler) = &self.emitted_message_handler { handler.handle(message).await;