feat(agents): observable loop guards and terminal Finish on panic
Unused Dependencies / machete (push) Has been cancelled
Unused Dependencies / machete (push) Has been cancelled
- 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 <cursoragent@cursor.com>
This commit is contained in:
@@ -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();
|
||||
|
||||
@@ -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,
|
||||
},
|
||||
)
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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};
|
||||
|
||||
@@ -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<usize>,
|
||||
last_call: Option<InternalToolCall>,
|
||||
repeat_count: u32,
|
||||
}
|
||||
|
||||
// Helper struct for internal tracking
|
||||
#[derive(Debug, Clone)]
|
||||
@@ -36,6 +46,7 @@ pub struct RepetitionInspector {
|
||||
last_call: Option<InternalToolCall>,
|
||||
repeat_count: u32,
|
||||
call_counts: HashMap<String, u32>,
|
||||
sessions: Mutex<HashMap<String, SessionRepetitionState>>,
|
||||
}
|
||||
|
||||
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::<u32>("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<usize> {
|
||||
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<Vec<InspectionResult>> {
|
||||
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()),
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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<String, serde_json::Value>) -> 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
|
||||
|
||||
Reference in New Issue
Block a user