feat: recipes can retry with success criteria (#3474)
This commit is contained in:
@@ -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")];
|
||||
|
||||
|
||||
Reference in New Issue
Block a user