From 4aea5abf29cf709ca12efc51147f9d7e00e076d0 Mon Sep 17 00:00:00 2001 From: john Date: Wed, 23 Sep 2026 14:25:12 +0800 Subject: [PATCH] feat(agents): observable loop guards and terminal Finish on panic - RepetitionInspector keeps per-session state, resets per user turn and is enabled via GOOSE_MAX_TOOL_REPETITIONS (off by default); denials are logged. - tkmind Finish reason now reports stop/max_turns/empty_response/cancelled/ error, with a 'Reply finished' log line. - RequestGuard publishes Error + Finish{aborted} when the reply task panics so clients stop waiting for the full timeout. Co-authored-by: Cursor --- .../goose/src/acp/tkmind/session_event_bus.rs | 61 +++++++++ crates/goose/src/acp/tkmind/session_events.rs | 17 ++- crates/goose/src/agents/agent.rs | 35 ++++- crates/goose/src/agents/mod.rs | 4 +- crates/goose/src/tool_monitor.rs | 110 +++++++++++---- .../goose/tests/repetition_inspector_tests.rs | 127 +++++++++++++++++- 6 files changed, 324 insertions(+), 30 deletions(-) diff --git a/crates/goose/src/acp/tkmind/session_event_bus.rs b/crates/goose/src/acp/tkmind/session_event_bus.rs index 75ab8b21b..4c14247cd 100644 --- a/crates/goose/src/acp/tkmind/session_event_bus.rs +++ b/crates/goose/src/acp/tkmind/session_event_bus.rs @@ -198,7 +198,30 @@ impl Drop for RequestGuard { if !self.disarmed { let bus = self.bus.clone(); let request_id = self.request_id.clone(); + // A panicking reply task never reaches its own Finish; without a terminal + // event here, clients wait for their full reply timeout. + let panicked = std::thread::panicking(); + if panicked { + tracing::error!(request_id = %request_id, "Reply task panicked; publishing terminal events"); + } tokio::spawn(async move { + if panicked { + bus.publish( + Some(request_id.clone()), + MessageEvent::Error { + error: "reply task aborted unexpectedly".to_string(), + }, + ) + .await; + bus.publish( + Some(request_id.clone()), + MessageEvent::Finish { + reason: "aborted".to_string(), + token_state: Default::default(), + }, + ) + .await; + } bus.cleanup_request(&request_id).await; }); } @@ -286,6 +309,44 @@ mod tests { assert!(!cancelled); } + #[tokio::test] + async fn test_request_guard_publishes_terminal_events_when_task_panics() { + let bus = std::sync::Arc::new(SessionEventBus::new()); + bus.register_request("req-p".to_string()).await; + + let task_bus = bus.clone(); + let handle = tokio::spawn(async move { + let _guard = RequestGuard::new(task_bus, "req-p".to_string()); + panic!("simulated tool panic"); + }); + assert!(handle.await.unwrap_err().is_panic()); + + let mut events = Vec::new(); + for _ in 0..50 { + let (replay, _, _rx) = bus.subscribe(None).await.unwrap(); + if replay.len() >= 2 { + events = replay; + break; + } + tokio::time::sleep(std::time::Duration::from_millis(10)).await; + } + assert!(matches!(events[0].event, MessageEvent::Error { .. })); + assert!( + matches!(&events[1].event, MessageEvent::Finish { reason, .. } if reason == "aborted") + ); + assert_eq!(events[1].request_id.as_deref(), Some("req-p")); + assert!(!bus.cancel_request("req-p").await); + } + + #[tokio::test] + async fn test_request_guard_does_not_publish_on_normal_early_return() { + let bus = std::sync::Arc::new(SessionEventBus::new()); + drop(RequestGuard::new(bus.clone(), "req-e".to_string())); + tokio::time::sleep(std::time::Duration::from_millis(20)).await; + let (replay, _, _rx) = bus.subscribe(None).await.unwrap(); + assert!(replay.is_empty()); + } + #[tokio::test] async fn test_cleanup_request() { let bus = SessionEventBus::new(); diff --git a/crates/goose/src/acp/tkmind/session_events.rs b/crates/goose/src/acp/tkmind/session_events.rs index f878acc01..c05397182 100644 --- a/crates/goose/src/acp/tkmind/session_events.rs +++ b/crates/goose/src/acp/tkmind/session_events.rs @@ -20,7 +20,7 @@ use tokio::sync::mpsc; use tokio::time::timeout; use tokio_stream::wrappers::ReceiverStream; -use crate::agents::{AgentEvent, SessionConfig}; +use crate::agents::{loop_stop_reason, AgentEvent, SessionConfig}; use crate::conversation::message::{ActionRequiredData, Message, MessageContent}; use crate::conversation::Conversation; @@ -414,10 +414,12 @@ pub async fn session_reply( } }; + let mut finish_reason = "stop"; loop { tokio::select! { _ = task_cancel.cancelled() => { tracing::info!("Agent task cancelled for request {}", task_request_id); + finish_reason = "cancelled"; break; } response = timeout(Duration::from_millis(500), stream.next()) => { @@ -426,6 +428,9 @@ pub async fn session_reply( for content in &message.content { track_tool_telemetry(content, all_messages.messages()); } + if let Some(reason) = loop_stop_reason(&message) { + finish_reason = reason; + } all_messages.push(message.clone()); let token_state = get_token_state( task_state.session_manager(), @@ -466,6 +471,7 @@ pub async fn session_reply( } Ok(Some(Err(e))) => { tracing::error!("Error processing message: {}", e); + finish_reason = "error"; publish( Some(task_request_id.clone()), MessageEvent::Error { @@ -543,10 +549,17 @@ pub async fn session_reply( let final_token_state = get_token_state(task_state.session_manager(), &task_session_id).await; + tracing::info!( + session_id = %task_session_id, + request_id = %task_request_id, + finish_reason, + duration_ms = session_duration.as_millis() as u64, + "Reply finished" + ); publish( Some(task_request_id.clone()), MessageEvent::Finish { - reason: "stop".to_string(), + reason: finish_reason.to_string(), token_state: final_token_state, }, ) diff --git a/crates/goose/src/agents/agent.rs b/crates/goose/src/agents/agent.rs index d8abdf0fe..ecee17267 100644 --- a/crates/goose/src/agents/agent.rs +++ b/crates/goose/src/agents/agent.rs @@ -94,6 +94,19 @@ const MAX_EMPTY_TURN_RETRIES: u32 = 3; const EMPTY_TURN_MESSAGE: &str = "The model returned an empty response. Please resend your message to continue."; +/// Classifies the assistant notices the loop emits when it stops on a guard, +/// so HTTP clients can report why a reply ended instead of a generic "stop". +pub fn loop_stop_reason(message: &Message) -> Option<&'static str> { + if message.role != rmcp::model::Role::Assistant { + return None; + } + match message.as_concat_text().trim() { + text if text == MAX_TURNS_MESSAGE => Some("max_turns"), + text if text == EMPTY_TURN_MESSAGE => Some("empty_response"), + _ => None, + } +} + fn provider_creation_error(error: anyhow::Error, context: impl fmt::Display) -> anyhow::Error { let message = format!("{context}: {error}"); error.context(message) @@ -801,7 +814,7 @@ impl Agent { ))); // Add repetition inspector (lower priority - basic repetition checking) - tool_inspection_manager.add_inspector(Box::new(RepetitionInspector::new(None))); + tool_inspection_manager.add_inspector(Box::new(RepetitionInspector::from_config())); tool_inspection_manager } @@ -4235,6 +4248,26 @@ mod tests { use std::sync::atomic::{AtomicUsize, Ordering}; use tempfile::TempDir; + #[test] + fn loop_stop_reason_classifies_guard_notices() { + assert_eq!( + loop_stop_reason(&Message::assistant().with_text(MAX_TURNS_MESSAGE)), + Some("max_turns") + ); + assert_eq!( + loop_stop_reason(&Message::assistant().with_text(EMPTY_TURN_MESSAGE)), + Some("empty_response") + ); + assert_eq!( + loop_stop_reason(&Message::assistant().with_text("页面已生成")), + None + ); + assert_eq!( + loop_stop_reason(&Message::user().with_text(MAX_TURNS_MESSAGE)), + None + ); + } + fn persisted_builtin(name: &str) -> ExtensionConfig { ExtensionConfig::Builtin { name: name.to_string(), diff --git a/crates/goose/src/agents/mod.rs b/crates/goose/src/agents/mod.rs index 498199212..a033c07d1 100644 --- a/crates/goose/src/agents/mod.rs +++ b/crates/goose/src/agents/mod.rs @@ -26,7 +26,9 @@ mod tool_schema_normalize; pub mod types; pub mod validate_extensions; -pub use agent::{Agent, AgentConfig, ExtensionLoadResult, GoosePlatform, MCP_PROTOCOL_VERSION}; +pub use agent::{ + loop_stop_reason, Agent, AgentConfig, ExtensionLoadResult, GoosePlatform, MCP_PROTOCOL_VERSION, +}; pub use container::Container; pub use execute_commands::{context_management_unsupported_message, COMPACT_TRIGGERS}; pub use extension::{ExtensionConfig, ExtensionError}; diff --git a/crates/goose/src/tool_monitor.rs b/crates/goose/src/tool_monitor.rs index 96a2c2779..4383fba88 100644 --- a/crates/goose/src/tool_monitor.rs +++ b/crates/goose/src/tool_monitor.rs @@ -1,11 +1,21 @@ -use crate::config::GooseMode; +use crate::config::{Config, GooseMode}; use crate::conversation::message::{Message, ToolRequest}; use crate::tool_inspection::{InspectionAction, InspectionResult, ToolInspector}; use anyhow::Result; use async_trait::async_trait; -use rmcp::model::CallToolRequestParams; +use rmcp::model::{CallToolRequestParams, Role}; use serde_json::Value; use std::collections::HashMap; +use std::sync::Mutex; + +const MAX_TRACKED_SESSIONS: usize = 4096; + +#[derive(Debug, Default)] +struct SessionRepetitionState { + turn_marker: Option, + last_call: Option, + repeat_count: u32, +} // Helper struct for internal tracking #[derive(Debug, Clone)] @@ -36,6 +46,7 @@ pub struct RepetitionInspector { last_call: Option, repeat_count: u32, call_counts: HashMap, + sessions: Mutex>, } impl RepetitionInspector { @@ -45,9 +56,26 @@ impl RepetitionInspector { last_call: None, repeat_count: 0, call_counts: HashMap::new(), + sessions: Mutex::new(HashMap::new()), } } + /// Limit from `GOOSE_MAX_TOOL_REPETITIONS`; unset or 0 disables the check. + pub fn from_config() -> Self { + let max = Config::global() + .get_param::("GOOSE_MAX_TOOL_REPETITIONS") + .ok() + .filter(|value| *value > 0); + Self::new(max) + } + + /// Index of the latest real user message; changes when a new user turn starts. + fn turn_marker(messages: &[Message]) -> Option { + messages + .iter() + .rposition(|m| m.role == Role::User && !m.is_tool_response()) + } + pub fn check_tool_call(&mut self, tool_call: CallToolRequestParams) -> bool { let internal_call = InternalToolCall::from_tool_call(&tool_call); let total_calls = self @@ -96,37 +124,69 @@ impl ToolInspector for RepetitionInspector { self } + fn is_enabled(&self) -> bool { + self.max_repetitions.is_some() + } + async fn inspect( &self, - _session_id: &str, + session_id: &str, tool_requests: &[ToolRequest], - _messages: &[Message], + messages: &[Message], _goose_mode: GooseMode, ) -> Result> { + let Some(max_repetitions) = self.max_repetitions else { + return Ok(Vec::new()); + }; + let marker = Self::turn_marker(messages); + let mut sessions = self + .sessions + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()); + if sessions.len() >= MAX_TRACKED_SESSIONS && !sessions.contains_key(session_id) { + sessions.clear(); + } + let state = sessions.entry(session_id.to_string()).or_default(); + if state.turn_marker != marker { + *state = SessionRepetitionState { + turn_marker: marker, + ..Default::default() + }; + } + let mut results = Vec::new(); - - // Check repetition limits for each tool request for tool_request in tool_requests { - if let Ok(tool_call) = &tool_request.tool_call { - // Create a temporary clone to check without modifying state - let mut temp_inspector = RepetitionInspector::new(self.max_repetitions); - temp_inspector.last_call = self.last_call.clone(); - temp_inspector.repeat_count = self.repeat_count; - temp_inspector.call_counts = self.call_counts.clone(); + let Ok(tool_call) = &tool_request.tool_call else { + continue; + }; + let call = InternalToolCall::from_tool_call(tool_call); + let repeated = state + .last_call + .as_ref() + .is_some_and(|last| last.matches(&call)); + state.repeat_count = if repeated { state.repeat_count + 1 } else { 1 }; + state.last_call = Some(call); - if !temp_inspector.check_tool_call(tool_call.clone()) { - results.push(InspectionResult { - tool_request_id: tool_request.id.clone(), - action: InspectionAction::Deny, - reason: format!( - "Tool '{}' has exceeded maximum repetitions", - tool_call.name - ), - confidence: 1.0, - inspector_name: "repetition".to_string(), - finding_id: Some("REP-001".to_string()), - }); - } + if state.repeat_count > max_repetitions { + tracing::warn!( + session_id, + tool = %tool_call.name, + repeat_count = state.repeat_count, + max_repetitions, + "Tool call denied: exceeded maximum consecutive repetitions" + ); + results.push(InspectionResult { + tool_request_id: tool_request.id.clone(), + action: InspectionAction::Deny, + reason: format!( + "Tool '{}' was called with identical arguments {} times in a row; \ + change approach instead of repeating it", + tool_call.name, state.repeat_count + ), + confidence: 1.0, + inspector_name: "repetition".to_string(), + finding_id: Some("REP-001".to_string()), + }); } } diff --git a/crates/goose/tests/repetition_inspector_tests.rs b/crates/goose/tests/repetition_inspector_tests.rs index a67377588..042b523c1 100644 --- a/crates/goose/tests/repetition_inspector_tests.rs +++ b/crates/goose/tests/repetition_inspector_tests.rs @@ -1,7 +1,132 @@ +use goose::config::GooseMode; +use goose::conversation::message::{Message, ToolRequest}; +use goose::tool_inspection::{InspectionAction, ToolInspector}; use goose::tool_monitor::RepetitionInspector; -use rmcp::model::CallToolRequestParams; +use rmcp::model::{CallToolRequestParams, CallToolResult}; use rmcp::object; +fn request(id: &str, tool: &str, args: serde_json::Map) -> ToolRequest { + ToolRequest { + id: id.into(), + tool_call: Ok(CallToolRequestParams::new(tool.to_string()).with_arguments(args)), + metadata: None, + tool_meta: None, + } +} + +async fn denied( + inspector: &RepetitionInspector, + session: &str, + req: ToolRequest, + messages: &[Message], +) -> bool { + inspector + .inspect(session, &[req], messages, GooseMode::default()) + .await + .unwrap() + .iter() + .any(|r| r.action == InspectionAction::Deny) +} + +#[tokio::test] +async fn inspect_tracks_repeats_across_calls_within_a_user_turn() { + let inspector = RepetitionInspector::new(Some(2)); + let mut history = vec![Message::user().with_text("生成页面")]; + let args = object!({"path": "public/a.html"}); + + assert!( + !denied( + &inspector, + "s1", + request("1", "read_file", args.clone()), + &history + ) + .await + ); + history.push(Message::user().with_tool_response("1", Ok(CallToolResult::success(vec![])))); + assert!( + !denied( + &inspector, + "s1", + request("2", "read_file", args.clone()), + &history + ) + .await + ); + history.push(Message::user().with_tool_response("2", Ok(CallToolResult::success(vec![])))); + assert!( + denied( + &inspector, + "s1", + request("3", "read_file", args.clone()), + &history + ) + .await + ); + + // A different call resets the consecutive counter. + assert!( + !denied( + &inspector, + "s1", + request("4", "list_dir", object!({})), + &history + ) + .await + ); + assert!(!denied(&inspector, "s1", request("5", "read_file", args), &history).await); +} + +#[tokio::test] +async fn inspect_resets_on_new_user_message_and_isolates_sessions() { + let inspector = RepetitionInspector::new(Some(1)); + let args = object!({"q": "x"}); + let mut history = vec![Message::user().with_text("第一轮")]; + + assert!( + !denied( + &inspector, + "s1", + request("1", "search", args.clone()), + &history + ) + .await + ); + assert!( + !denied( + &inspector, + "s2", + request("1", "search", args.clone()), + &history + ) + .await + ); + assert!( + denied( + &inspector, + "s1", + request("2", "search", args.clone()), + &history + ) + .await + ); + + history.push(Message::assistant().with_text("好的")); + history.push(Message::user().with_text("第二轮")); + assert!(!denied(&inspector, "s1", request("3", "search", args), &history).await); +} + +#[tokio::test] +async fn inspect_is_disabled_without_limit() { + let inspector = RepetitionInspector::new(None); + assert!(!inspector.is_enabled()); + let history = vec![Message::user().with_text("hi")]; + for i in 0..10 { + let req = request(&i.to_string(), "read_file", object!({"path": "a"})); + assert!(!denied(&inspector, "s1", req, &history).await); + } +} + // This test targets RepetitionInspector::check_tool_call // It verifies that: // - consecutive identical tool calls are allowed up to max_repetitions times