feat: add /goal command for agent self-evaluation before finishing (#9069)
Signed-off-by: Michael Neale <michael.neale@gmail.com> Co-authored-by: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
@@ -213,6 +213,8 @@ pub struct Agent {
|
|||||||
pub(super) tool_inspection_manager: ToolInspectionManager,
|
pub(super) tool_inspection_manager: ToolInspectionManager,
|
||||||
pub(super) hook_manager: crate::hooks::HookManager,
|
pub(super) hook_manager: crate::hooks::HookManager,
|
||||||
container: Mutex<Option<Container>>,
|
container: Mutex<Option<Container>>,
|
||||||
|
goal: Mutex<Option<String>>,
|
||||||
|
grind: Mutex<Option<String>>,
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(Clone, Debug)]
|
#[derive(Clone, Debug)]
|
||||||
@@ -330,6 +332,8 @@ impl Agent {
|
|||||||
),
|
),
|
||||||
hook_manager: crate::hooks::HookManager::load(std::env::current_dir().ok().as_deref()),
|
hook_manager: crate::hooks::HookManager::load(std::env::current_dir().ok().as_deref()),
|
||||||
container: Mutex::new(None),
|
container: Mutex::new(None),
|
||||||
|
goal: Mutex::new(None),
|
||||||
|
grind: Mutex::new(None),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1636,6 +1640,7 @@ impl Agent {
|
|||||||
});
|
});
|
||||||
let mut compaction_attempts = 0;
|
let mut compaction_attempts = 0;
|
||||||
let mut last_assistant_text = String::new();
|
let mut last_assistant_text = String::new();
|
||||||
|
let mut goal_check_pending = false;
|
||||||
let mut tool_pair_summarization_done = false;
|
let mut tool_pair_summarization_done = false;
|
||||||
|
|
||||||
loop {
|
loop {
|
||||||
@@ -1986,6 +1991,8 @@ impl Agent {
|
|||||||
}
|
}
|
||||||
|
|
||||||
no_tools_called = false;
|
no_tools_called = false;
|
||||||
|
// Agent is actively working — re-check goal when it next finishes
|
||||||
|
goal_check_pending = false;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
#[allow(unused_variables)]
|
#[allow(unused_variables)]
|
||||||
@@ -2136,7 +2143,46 @@ impl Agent {
|
|||||||
None if did_recovery_compact_this_iteration => {
|
None if did_recovery_compact_this_iteration => {
|
||||||
// continue from last user message after recovery compact
|
// continue from last user message after recovery compact
|
||||||
}
|
}
|
||||||
|
None if self.goal.lock().await.is_some() && !goal_check_pending => {
|
||||||
|
goal_check_pending = true;
|
||||||
|
let goal = self.goal.lock().await.clone().unwrap();
|
||||||
|
let nudge = format!(
|
||||||
|
"Before finishing, check whether the following goal has been fully met:\n\n\
|
||||||
|
**Goal:** {goal}\n\n\
|
||||||
|
If not, continue working toward it."
|
||||||
|
);
|
||||||
|
let message = Message::user().with_text(&nudge)
|
||||||
|
.with_visibility(false, true);
|
||||||
|
messages_to_add.push(message);
|
||||||
|
yield AgentEvent::Message(
|
||||||
|
Message::assistant().with_system_notification(
|
||||||
|
SystemNotificationType::InlineMessage,
|
||||||
|
format!("Goal: {goal}"),
|
||||||
|
)
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
None if self.grind.lock().await.is_some() => {
|
||||||
|
let grind = self.grind.lock().await.clone().unwrap();
|
||||||
|
let nudge = format!(
|
||||||
|
"Keep working. The grind goal is not yet complete:\n\n\
|
||||||
|
**Goal:** {grind}\n\n\
|
||||||
|
Continue until it is fully done."
|
||||||
|
);
|
||||||
|
let message = Message::user().with_text(&nudge)
|
||||||
|
.with_visibility(false, true);
|
||||||
|
messages_to_add.push(message);
|
||||||
|
yield AgentEvent::Message(
|
||||||
|
Message::assistant().with_system_notification(
|
||||||
|
SystemNotificationType::InlineMessage,
|
||||||
|
format!("Grind: {grind}"),
|
||||||
|
)
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
None => {
|
None => {
|
||||||
|
self.set_goal(None).await;
|
||||||
|
self.set_grind(None).await;
|
||||||
match self.handle_retry_logic(&mut conversation, &session_config, &initial_messages).await {
|
match self.handle_retry_logic(&mut conversation, &session_config, &initial_messages).await {
|
||||||
Ok(should_retry) => {
|
Ok(should_retry) => {
|
||||||
if should_retry {
|
if should_retry {
|
||||||
@@ -2224,6 +2270,22 @@ impl Agent {
|
|||||||
prompt_manager.add_system_prompt_extra(key, instruction);
|
prompt_manager.add_system_prompt_extra(key, instruction);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pub async fn set_goal(&self, goal: Option<String>) {
|
||||||
|
*self.goal.lock().await = goal;
|
||||||
|
}
|
||||||
|
|
||||||
|
pub async fn get_goal(&self) -> Option<String> {
|
||||||
|
self.goal.lock().await.clone()
|
||||||
|
}
|
||||||
|
|
||||||
|
pub async fn set_grind(&self, goal: Option<String>) {
|
||||||
|
*self.grind.lock().await = goal;
|
||||||
|
}
|
||||||
|
|
||||||
|
pub async fn get_grind(&self) -> Option<String> {
|
||||||
|
self.grind.lock().await.clone()
|
||||||
|
}
|
||||||
|
|
||||||
pub async fn update_provider(
|
pub async fn update_provider(
|
||||||
&self,
|
&self,
|
||||||
provider: Arc<dyn Provider>,
|
provider: Arc<dyn Provider>,
|
||||||
|
|||||||
@@ -41,6 +41,15 @@ static COMMANDS: &[CommandDef] = &[
|
|||||||
name: "doctor",
|
name: "doctor",
|
||||||
description: "Check that your Goose setup is working",
|
description: "Check that your Goose setup is working",
|
||||||
},
|
},
|
||||||
|
CommandDef {
|
||||||
|
name: "goal",
|
||||||
|
description: "Set a goal the agent must satisfy before finishing, or clear with /goal off",
|
||||||
|
},
|
||||||
|
CommandDef {
|
||||||
|
name: "grind",
|
||||||
|
description:
|
||||||
|
"Set a goal the agent pursues relentlessly until max_turns, or clear with /grind off",
|
||||||
|
},
|
||||||
];
|
];
|
||||||
|
|
||||||
pub struct ParsedSlashCommand<'a> {
|
pub struct ParsedSlashCommand<'a> {
|
||||||
@@ -101,6 +110,8 @@ impl Agent {
|
|||||||
"clear" => self.handle_clear_command(session_id).await,
|
"clear" => self.handle_clear_command(session_id).await,
|
||||||
"skills" => self.handle_skills_command(session_id).await,
|
"skills" => self.handle_skills_command(session_id).await,
|
||||||
"doctor" => Ok(Some(crate::doctor::run(self, session_id).await?)),
|
"doctor" => Ok(Some(crate::doctor::run(self, session_id).await?)),
|
||||||
|
"goal" => self.handle_goal_command(params_str).await,
|
||||||
|
"grind" => self.handle_grind_command(params_str).await,
|
||||||
_ => {
|
_ => {
|
||||||
self.handle_recipe_command(command, params_str, session_id)
|
self.handle_recipe_command(command, params_str, session_id)
|
||||||
.await
|
.await
|
||||||
@@ -460,6 +471,54 @@ impl Agent {
|
|||||||
|
|
||||||
Ok(Some(Message::user().with_text(prompt)))
|
Ok(Some(Message::user().with_text(prompt)))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
async fn handle_goal_command(&self, params_str: &str) -> Result<Option<Message>> {
|
||||||
|
if params_str.is_empty() {
|
||||||
|
let current = self.get_goal().await;
|
||||||
|
let text = match current {
|
||||||
|
Some(goal) => format!("Current goal: {goal}"),
|
||||||
|
None => "No goal set. Use `/goal <description>` to set one.".to_string(),
|
||||||
|
};
|
||||||
|
return Ok(Some(Message::assistant().with_text(text)));
|
||||||
|
}
|
||||||
|
|
||||||
|
if params_str == "off" || params_str == "clear" || params_str == "none" {
|
||||||
|
self.set_goal(None).await;
|
||||||
|
return Ok(Some(
|
||||||
|
Message::assistant().with_text("Goal cleared. The agent will finish normally."),
|
||||||
|
));
|
||||||
|
}
|
||||||
|
|
||||||
|
let goal = params_str.to_string();
|
||||||
|
self.set_goal(Some(goal.clone())).await;
|
||||||
|
Ok(Some(Message::assistant().with_text(format!(
|
||||||
|
"Goal set. The agent will verify this goal is met before finishing:\n\n> {goal}"
|
||||||
|
))))
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn handle_grind_command(&self, params_str: &str) -> Result<Option<Message>> {
|
||||||
|
if params_str.is_empty() {
|
||||||
|
let current = self.get_grind().await;
|
||||||
|
let text = match current {
|
||||||
|
Some(goal) => format!("Current grind goal: {goal}"),
|
||||||
|
None => "No grind goal set. Use `/grind <description>` to set one.".to_string(),
|
||||||
|
};
|
||||||
|
return Ok(Some(Message::assistant().with_text(text)));
|
||||||
|
}
|
||||||
|
|
||||||
|
if params_str == "off" || params_str == "clear" {
|
||||||
|
self.set_grind(None).await;
|
||||||
|
return Ok(Some(
|
||||||
|
Message::assistant().with_text("Grind cleared. The agent will finish normally."),
|
||||||
|
));
|
||||||
|
}
|
||||||
|
|
||||||
|
let goal = params_str.to_string();
|
||||||
|
self.set_grind(Some(goal.clone())).await;
|
||||||
|
Ok(Some(Message::assistant().with_text(format!(
|
||||||
|
"Grind goal set. The agent will keep working until max_turns is reached:\n\n> {goal}"
|
||||||
|
))))
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
|
|||||||
@@ -1119,6 +1119,276 @@ mod tests {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod goal_checking_tests {
|
||||||
|
use super::*;
|
||||||
|
use async_trait::async_trait;
|
||||||
|
use goose::agents::AgentConfig;
|
||||||
|
use goose::agents::SessionConfig;
|
||||||
|
use goose::config::permission::PermissionManager;
|
||||||
|
use goose::config::GooseMode;
|
||||||
|
use goose::conversation::message::Message;
|
||||||
|
use goose::model::ModelConfig;
|
||||||
|
use goose::providers::base::{
|
||||||
|
stream_from_single_message, MessageStream, Provider, ProviderDef, ProviderMetadata,
|
||||||
|
ProviderUsage, Usage,
|
||||||
|
};
|
||||||
|
use goose::providers::errors::ProviderError;
|
||||||
|
use goose::session::session_manager::SessionType;
|
||||||
|
use goose::session::SessionManager;
|
||||||
|
use rmcp::model::Tool;
|
||||||
|
use std::path::PathBuf;
|
||||||
|
use std::sync::atomic::{AtomicU32, Ordering};
|
||||||
|
use tempfile::TempDir;
|
||||||
|
|
||||||
|
struct GoalTextProvider {
|
||||||
|
call_count: AtomicU32,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl GoalTextProvider {
|
||||||
|
fn new() -> Self {
|
||||||
|
Self {
|
||||||
|
call_count: AtomicU32::new(0),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl ProviderDef for GoalTextProvider {
|
||||||
|
type Provider = Self;
|
||||||
|
|
||||||
|
fn metadata() -> ProviderMetadata {
|
||||||
|
ProviderMetadata {
|
||||||
|
name: "goal-mock".to_string(),
|
||||||
|
display_name: "Goal Mock Provider".to_string(),
|
||||||
|
description: "Mock provider for goal testing".to_string(),
|
||||||
|
default_model: "mock-model".to_string(),
|
||||||
|
known_models: vec![],
|
||||||
|
model_doc_link: "".to_string(),
|
||||||
|
config_keys: vec![],
|
||||||
|
setup_steps: vec![],
|
||||||
|
model_selection_hint: None,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn from_env(
|
||||||
|
_model: ModelConfig,
|
||||||
|
_extensions: Vec<goose::config::ExtensionConfig>,
|
||||||
|
) -> futures::future::BoxFuture<'static, anyhow::Result<Self>> {
|
||||||
|
Box::pin(async { Ok(Self::new()) })
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[async_trait]
|
||||||
|
impl Provider for GoalTextProvider {
|
||||||
|
async fn stream(
|
||||||
|
&self,
|
||||||
|
_model_config: &ModelConfig,
|
||||||
|
_session_id: &str,
|
||||||
|
_system_prompt: &str,
|
||||||
|
_messages: &[Message],
|
||||||
|
_tools: &[Tool],
|
||||||
|
) -> Result<MessageStream, ProviderError> {
|
||||||
|
let count = self.call_count.fetch_add(1, Ordering::SeqCst);
|
||||||
|
let text = format!("Response number {count}");
|
||||||
|
let message = Message::assistant().with_text(&text);
|
||||||
|
let usage = ProviderUsage::new(
|
||||||
|
"mock-model".to_string(),
|
||||||
|
Usage::new(Some(10), Some(5), Some(15)),
|
||||||
|
);
|
||||||
|
Ok(stream_from_single_message(message, usage))
|
||||||
|
}
|
||||||
|
|
||||||
|
fn get_model_config(&self) -> ModelConfig {
|
||||||
|
ModelConfig::new("mock-model").unwrap()
|
||||||
|
}
|
||||||
|
|
||||||
|
fn get_name(&self) -> &str {
|
||||||
|
"goal-mock"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn create_agent_with_session_naming_disabled(
|
||||||
|
session_manager: Arc<SessionManager>,
|
||||||
|
) -> Agent {
|
||||||
|
let config = AgentConfig::new(
|
||||||
|
session_manager,
|
||||||
|
PermissionManager::instance(),
|
||||||
|
None,
|
||||||
|
GooseMode::Auto,
|
||||||
|
true,
|
||||||
|
GoosePlatform::GooseCli,
|
||||||
|
);
|
||||||
|
Agent::with_config(config)
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn test_goal_nudges_agent_before_exit() -> Result<()> {
|
||||||
|
let temp_dir = TempDir::new()?;
|
||||||
|
let session_manager = Arc::new(SessionManager::new(temp_dir.path().to_path_buf()));
|
||||||
|
let agent = create_agent_with_session_naming_disabled(session_manager.clone());
|
||||||
|
let provider = Arc::new(GoalTextProvider::new());
|
||||||
|
|
||||||
|
let session = session_manager
|
||||||
|
.create_session(
|
||||||
|
PathBuf::default(),
|
||||||
|
"goal-test".to_string(),
|
||||||
|
SessionType::Hidden,
|
||||||
|
GooseMode::default(),
|
||||||
|
)
|
||||||
|
.await?;
|
||||||
|
|
||||||
|
agent.update_provider(provider.clone(), &session.id).await?;
|
||||||
|
agent
|
||||||
|
.set_goal(Some("Ensure the sky is blue".to_string()))
|
||||||
|
.await;
|
||||||
|
|
||||||
|
let session_config = SessionConfig {
|
||||||
|
id: session.id.clone(),
|
||||||
|
schedule_id: None,
|
||||||
|
max_turns: Some(10),
|
||||||
|
retry_config: None,
|
||||||
|
};
|
||||||
|
|
||||||
|
let reply_stream = agent
|
||||||
|
.reply(Message::user().with_text("Hello"), session_config, None)
|
||||||
|
.await?;
|
||||||
|
tokio::pin!(reply_stream);
|
||||||
|
|
||||||
|
let mut messages = Vec::new();
|
||||||
|
while let Some(event) = reply_stream.next().await {
|
||||||
|
match event {
|
||||||
|
Ok(AgentEvent::Message(msg)) => messages.push(msg),
|
||||||
|
Ok(_) => {}
|
||||||
|
Err(e) => return Err(e),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
let call_count = provider.call_count.load(Ordering::SeqCst);
|
||||||
|
assert!(
|
||||||
|
call_count > 1,
|
||||||
|
"Expected provider to be called more than once due to goal checking, got {call_count}"
|
||||||
|
);
|
||||||
|
assert!(
|
||||||
|
call_count <= 3,
|
||||||
|
"Expected at most 3 provider calls (1 initial + 1 goal check + 1 exit), got {call_count}"
|
||||||
|
);
|
||||||
|
|
||||||
|
// The goal nudge should NOT appear in yielded events (it's internal)
|
||||||
|
let nudge_messages: Vec<_> = messages
|
||||||
|
.iter()
|
||||||
|
.filter(|m| {
|
||||||
|
m.as_concat_text()
|
||||||
|
.contains("check whether the following goal")
|
||||||
|
})
|
||||||
|
.collect();
|
||||||
|
assert!(
|
||||||
|
nudge_messages.is_empty(),
|
||||||
|
"Goal nudge should be hidden from user, but found {} in events",
|
||||||
|
nudge_messages.len()
|
||||||
|
);
|
||||||
|
|
||||||
|
// Goal should be cleared after being met
|
||||||
|
assert_eq!(
|
||||||
|
agent.get_goal().await,
|
||||||
|
None,
|
||||||
|
"Goal should be cleared after the agent finishes with it met"
|
||||||
|
);
|
||||||
|
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn test_no_goal_exits_immediately() -> Result<()> {
|
||||||
|
let temp_dir = TempDir::new()?;
|
||||||
|
let session_manager = Arc::new(SessionManager::new(temp_dir.path().to_path_buf()));
|
||||||
|
let agent = create_agent_with_session_naming_disabled(session_manager.clone());
|
||||||
|
let provider = Arc::new(GoalTextProvider::new());
|
||||||
|
|
||||||
|
let session = session_manager
|
||||||
|
.create_session(
|
||||||
|
PathBuf::default(),
|
||||||
|
"no-goal-test".to_string(),
|
||||||
|
SessionType::Hidden,
|
||||||
|
GooseMode::default(),
|
||||||
|
)
|
||||||
|
.await?;
|
||||||
|
|
||||||
|
agent.update_provider(provider.clone(), &session.id).await?;
|
||||||
|
|
||||||
|
let session_config = SessionConfig {
|
||||||
|
id: session.id.clone(),
|
||||||
|
schedule_id: None,
|
||||||
|
max_turns: Some(10),
|
||||||
|
retry_config: None,
|
||||||
|
};
|
||||||
|
|
||||||
|
let reply_stream = agent
|
||||||
|
.reply(Message::user().with_text("Hello"), session_config, None)
|
||||||
|
.await?;
|
||||||
|
tokio::pin!(reply_stream);
|
||||||
|
|
||||||
|
while let Some(event) = reply_stream.next().await {
|
||||||
|
match event {
|
||||||
|
Ok(_) => {}
|
||||||
|
Err(e) => return Err(e),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
let call_count = provider.call_count.load(Ordering::SeqCst);
|
||||||
|
assert_eq!(
|
||||||
|
call_count, 1,
|
||||||
|
"Without a goal, provider should be called exactly once, got {call_count}"
|
||||||
|
);
|
||||||
|
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn test_goal_command_set_and_clear() -> Result<()> {
|
||||||
|
let temp_dir = TempDir::new()?;
|
||||||
|
let session_manager = Arc::new(SessionManager::new(temp_dir.path().to_path_buf()));
|
||||||
|
let agent = create_agent_with_session_naming_disabled(session_manager.clone());
|
||||||
|
|
||||||
|
let session = session_manager
|
||||||
|
.create_session(
|
||||||
|
PathBuf::default(),
|
||||||
|
"goal-cmd-test".to_string(),
|
||||||
|
SessionType::Hidden,
|
||||||
|
GooseMode::default(),
|
||||||
|
)
|
||||||
|
.await?;
|
||||||
|
|
||||||
|
// No goal initially
|
||||||
|
let result = agent.execute_command("/goal", &session.id).await?.unwrap();
|
||||||
|
assert!(result.as_concat_text().contains("No goal set"));
|
||||||
|
|
||||||
|
// Set a goal
|
||||||
|
let result = agent
|
||||||
|
.execute_command("/goal make all tests pass", &session.id)
|
||||||
|
.await?
|
||||||
|
.unwrap();
|
||||||
|
assert!(result.as_concat_text().contains("Goal set"));
|
||||||
|
assert_eq!(
|
||||||
|
agent.get_goal().await,
|
||||||
|
Some("make all tests pass".to_string())
|
||||||
|
);
|
||||||
|
|
||||||
|
// Query it
|
||||||
|
let result = agent.execute_command("/goal", &session.id).await?.unwrap();
|
||||||
|
assert!(result.as_concat_text().contains("make all tests pass"));
|
||||||
|
|
||||||
|
// Clear it
|
||||||
|
let result = agent
|
||||||
|
.execute_command("/goal off", &session.id)
|
||||||
|
.await?
|
||||||
|
.unwrap();
|
||||||
|
assert!(result.as_concat_text().contains("cleared"));
|
||||||
|
assert_eq!(agent.get_goal().await, None);
|
||||||
|
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
mod cumulative_token_tests {
|
mod cumulative_token_tests {
|
||||||
use super::*;
|
use super::*;
|
||||||
use async_trait::async_trait;
|
use async_trait::async_trait;
|
||||||
|
|||||||
Reference in New Issue
Block a user