feat: recipes can retry with success criteria (#3474)

This commit is contained in:
Prem Pillai
2025-07-22 10:49:21 +10:00
committed by GitHub
parent 5f3c7d339c
commit 99cc0a9c81
17 changed files with 1078 additions and 82 deletions
+179
View File
@@ -761,6 +761,184 @@ mod final_output_tool_tests {
}
}
#[cfg(test)]
mod retry_tests {
use super::*;
use async_trait::async_trait;
use goose::agents::types::{RetryConfig, SessionConfig, SuccessCheck};
use goose::model::ModelConfig;
use goose::providers::base::{Provider, ProviderUsage, Usage};
use goose::providers::errors::ProviderError;
use mcp_core::tool::Tool;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
#[derive(Clone)]
struct MockRetryProvider {
model_config: ModelConfig,
call_count: Arc<AtomicUsize>,
fail_until: usize,
}
#[async_trait]
impl Provider for MockRetryProvider {
fn metadata() -> goose::providers::base::ProviderMetadata {
goose::providers::base::ProviderMetadata::empty()
}
fn get_model_config(&self) -> ModelConfig {
self.model_config.clone()
}
async fn complete(
&self,
_system: &str,
_messages: &[Message],
_tools: &[Tool],
) -> anyhow::Result<(Message, ProviderUsage), ProviderError> {
let count = self.call_count.fetch_add(1, Ordering::SeqCst);
if count < self.fail_until {
Ok((
Message::assistant().with_text("Task failed - will retry."),
ProviderUsage::new("mock".to_string(), Usage::default()),
))
} else {
Ok((
Message::assistant().with_text("Task completed successfully."),
ProviderUsage::new("mock".to_string(), Usage::default()),
))
}
}
}
#[tokio::test]
async fn test_retry_config_validation_integration() -> Result<()> {
let agent = Agent::new();
let model_config = ModelConfig::new("test-model".to_string());
let mock_provider = Arc::new(MockRetryProvider {
model_config,
call_count: Arc::new(AtomicUsize::new(0)),
fail_until: 0,
});
agent.update_provider(mock_provider.clone()).await?;
let retry_config = RetryConfig {
max_retries: 3,
checks: vec![SuccessCheck::Shell {
command: "echo 'success check'".to_string(),
}],
on_failure: Some("echo 'cleanup executed'".to_string()),
timeout_seconds: Some(30),
on_failure_timeout_seconds: Some(60),
};
assert!(
retry_config.validate().is_ok(),
"Valid config should pass validation"
);
let session_config = SessionConfig {
id: goose::session::Identifier::Name("test-retry".to_string()),
working_dir: std::env::current_dir()?,
schedule_id: None,
execution_mode: None,
max_turns: None,
retry_config: Some(retry_config),
};
let initial_messages = vec![Message::user().with_text("Complete this task")];
let reply_stream = agent.reply(&initial_messages, Some(session_config)).await?;
tokio::pin!(reply_stream);
let mut responses = Vec::new();
while let Some(response_result) = reply_stream.next().await {
match response_result {
Ok(AgentEvent::Message(response)) => responses.push(response),
Ok(_) => {}
Err(e) => return Err(e),
}
}
assert!(!responses.is_empty(), "Should have received responses");
Ok(())
}
#[tokio::test]
async fn test_retry_success_check_execution() -> Result<()> {
use goose::agents::retry::execute_success_checks;
let retry_config = RetryConfig {
max_retries: 3,
checks: vec![],
on_failure: None,
timeout_seconds: Some(30),
on_failure_timeout_seconds: Some(60),
};
let success_checks = vec![SuccessCheck::Shell {
command: "echo 'test'".to_string(),
}];
let result = execute_success_checks(&success_checks, &retry_config).await;
assert!(result.is_ok(), "Success check should pass");
assert!(result.unwrap(), "Command should succeed");
let fail_checks = vec![SuccessCheck::Shell {
command: "false".to_string(),
}];
let result = execute_success_checks(&fail_checks, &retry_config).await;
assert!(result.is_ok(), "Success check execution should not error");
assert!(!result.unwrap(), "Command should fail");
Ok(())
}
#[tokio::test]
async fn test_retry_logic_with_validation_errors() -> Result<()> {
let invalid_retry_config = RetryConfig {
max_retries: 0,
checks: vec![],
on_failure: None,
timeout_seconds: Some(0),
on_failure_timeout_seconds: None,
};
let validation_result = invalid_retry_config.validate();
assert!(
validation_result.is_err(),
"Should validate max_retries > 0"
);
assert!(validation_result
.unwrap_err()
.contains("max_retries must be greater than 0"));
Ok(())
}
#[tokio::test]
async fn test_retry_attempts_counter_reset() -> Result<()> {
let agent = Agent::new();
agent.reset_retry_attempts().await;
let initial_attempts = agent.get_retry_attempts().await;
assert_eq!(initial_attempts, 0);
let new_attempts = agent.increment_retry_attempts().await;
assert_eq!(new_attempts, 1);
agent.reset_retry_attempts().await;
let reset_attempts = agent.get_retry_attempts().await;
assert_eq!(reset_attempts, 0);
Ok(())
}
}
#[cfg(test)]
mod max_turns_tests {
use super::*;
@@ -831,6 +1009,7 @@ mod max_turns_tests {
schedule_id: None,
execution_mode: None,
max_turns: Some(1),
retry_config: None,
};
let messages = vec![Message::user().with_text("Hello")];