325bf396af
Co-authored-by: Alex Hancock <alexhancock@block.xyz>
591 lines
21 KiB
Rust
591 lines
21 KiB
Rust
use std::sync::Arc;
|
|
|
|
use anyhow::Result;
|
|
use futures::StreamExt;
|
|
use goose::agents::{Agent, AgentEvent, GoosePlatform};
|
|
use goose::config::extensions::{set_extension, ExtensionEntry};
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
|
|
#[cfg(test)]
|
|
mod schedule_tool_tests {
|
|
use super::*;
|
|
use async_trait::async_trait;
|
|
use chrono::{DateTime, Utc};
|
|
use goose::agents::platform_tools::PLATFORM_MANAGE_SCHEDULE_TOOL_NAME;
|
|
use goose::agents::AgentConfig;
|
|
use goose::config::permission::PermissionManager;
|
|
use goose::config::GooseMode;
|
|
use goose::scheduler::{ScheduledJob, SchedulerError};
|
|
use goose::scheduler_trait::SchedulerTrait;
|
|
use goose::session::{Session, SessionManager};
|
|
use std::path::PathBuf;
|
|
use std::sync::Arc;
|
|
use tempfile::TempDir;
|
|
|
|
struct MockScheduler {
|
|
jobs: tokio::sync::Mutex<Vec<ScheduledJob>>,
|
|
}
|
|
|
|
impl MockScheduler {
|
|
fn new() -> Self {
|
|
Self {
|
|
jobs: tokio::sync::Mutex::new(Vec::new()),
|
|
}
|
|
}
|
|
}
|
|
|
|
#[async_trait]
|
|
impl SchedulerTrait for MockScheduler {
|
|
async fn add_scheduled_job(
|
|
&self,
|
|
job: ScheduledJob,
|
|
_copy: bool,
|
|
) -> Result<(), SchedulerError> {
|
|
let mut jobs = self.jobs.lock().await;
|
|
jobs.push(job);
|
|
Ok(())
|
|
}
|
|
|
|
async fn schedule_recipe(
|
|
&self,
|
|
_recipe_path: PathBuf,
|
|
_cron_schedule: Option<String>,
|
|
) -> Result<(), SchedulerError> {
|
|
Ok(())
|
|
}
|
|
|
|
async fn list_scheduled_jobs(&self) -> Vec<ScheduledJob> {
|
|
let jobs = self.jobs.lock().await;
|
|
jobs.clone()
|
|
}
|
|
|
|
async fn remove_scheduled_job(
|
|
&self,
|
|
id: &str,
|
|
_remove: bool,
|
|
) -> Result<(), SchedulerError> {
|
|
let mut jobs = self.jobs.lock().await;
|
|
if let Some(pos) = jobs.iter().position(|job| job.id == id) {
|
|
jobs.remove(pos);
|
|
Ok(())
|
|
} else {
|
|
Err(SchedulerError::JobNotFound(id.to_string()))
|
|
}
|
|
}
|
|
|
|
async fn pause_schedule(&self, _id: &str) -> Result<(), SchedulerError> {
|
|
Ok(())
|
|
}
|
|
|
|
async fn unpause_schedule(&self, _id: &str) -> Result<(), SchedulerError> {
|
|
Ok(())
|
|
}
|
|
|
|
async fn run_now(&self, _id: &str) -> Result<String, SchedulerError> {
|
|
Ok("test_session_123".to_string())
|
|
}
|
|
|
|
async fn sessions(
|
|
&self,
|
|
_sched_id: &str,
|
|
_limit: usize,
|
|
) -> Result<Vec<(String, Session)>, SchedulerError> {
|
|
Ok(vec![])
|
|
}
|
|
|
|
async fn update_schedule(
|
|
&self,
|
|
_sched_id: &str,
|
|
_new_cron: String,
|
|
) -> Result<(), SchedulerError> {
|
|
Ok(())
|
|
}
|
|
|
|
async fn kill_running_job(&self, _sched_id: &str) -> Result<(), SchedulerError> {
|
|
Ok(())
|
|
}
|
|
|
|
async fn get_running_job_info(
|
|
&self,
|
|
_sched_id: &str,
|
|
) -> Result<Option<(String, DateTime<Utc>)>, SchedulerError> {
|
|
Ok(None)
|
|
}
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_schedule_management_tool_list() {
|
|
let temp_dir = TempDir::new().unwrap();
|
|
let data_dir = temp_dir.path().to_path_buf();
|
|
let session_manager = Arc::new(SessionManager::new(data_dir.clone()));
|
|
let permission_manager = Arc::new(PermissionManager::new(data_dir));
|
|
let mock_scheduler = Arc::new(MockScheduler::new());
|
|
let config = AgentConfig::new(
|
|
session_manager,
|
|
permission_manager,
|
|
Some(mock_scheduler),
|
|
GooseMode::Auto,
|
|
false,
|
|
GoosePlatform::GooseCli,
|
|
);
|
|
let agent = Agent::with_config(config);
|
|
|
|
let tools = agent.list_tools("test-session-id", None).await;
|
|
let schedule_tool = tools
|
|
.iter()
|
|
.find(|tool| tool.name == PLATFORM_MANAGE_SCHEDULE_TOOL_NAME);
|
|
assert!(schedule_tool.is_some());
|
|
|
|
let tool = schedule_tool.unwrap();
|
|
assert!(tool
|
|
.description
|
|
.clone()
|
|
.unwrap_or_default()
|
|
.contains("Manage goose's internal scheduled recipe execution"));
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_no_schedule_management_tool_without_scheduler() {
|
|
let agent = Agent::new();
|
|
|
|
let tools = agent.list_tools("test-session-id", None).await;
|
|
let schedule_tool = tools
|
|
.iter()
|
|
.find(|tool| tool.name == PLATFORM_MANAGE_SCHEDULE_TOOL_NAME);
|
|
assert!(schedule_tool.is_none());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_schedule_management_tool_in_platform_tools() {
|
|
let temp_dir = TempDir::new().unwrap();
|
|
let data_dir = temp_dir.path().to_path_buf();
|
|
let session_manager = Arc::new(SessionManager::new(data_dir.clone()));
|
|
let permission_manager = Arc::new(PermissionManager::new(data_dir));
|
|
let mock_scheduler = Arc::new(MockScheduler::new());
|
|
let config = AgentConfig::new(
|
|
session_manager,
|
|
permission_manager,
|
|
Some(mock_scheduler),
|
|
GooseMode::Auto,
|
|
false,
|
|
GoosePlatform::GooseCli,
|
|
);
|
|
let agent = Agent::with_config(config);
|
|
|
|
let tools = agent
|
|
.list_tools("test-session-id", Some("platform".to_string()))
|
|
.await;
|
|
|
|
// Check that the schedule management tool is included in platform tools
|
|
let schedule_tool = tools
|
|
.iter()
|
|
.find(|tool| tool.name == PLATFORM_MANAGE_SCHEDULE_TOOL_NAME);
|
|
assert!(schedule_tool.is_some());
|
|
|
|
let tool = schedule_tool.unwrap();
|
|
assert!(tool
|
|
.description
|
|
.clone()
|
|
.unwrap_or_default()
|
|
.contains("Manage goose's internal scheduled recipe execution"));
|
|
|
|
// Verify the tool has the expected actions in its schema
|
|
if let Some(properties) = tool.input_schema.get("properties") {
|
|
if let Some(action_prop) = properties.get("action") {
|
|
if let Some(enum_values) = action_prop.get("enum") {
|
|
let actions: Vec<String> = enum_values
|
|
.as_array()
|
|
.unwrap()
|
|
.iter()
|
|
.map(|v| v.as_str().unwrap().to_string())
|
|
.collect();
|
|
|
|
// Check that our session_content action is included
|
|
assert!(actions.contains(&"session_content".to_string()));
|
|
assert!(actions.contains(&"list".to_string()));
|
|
assert!(actions.contains(&"create".to_string()));
|
|
assert!(actions.contains(&"sessions".to_string()));
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_schedule_management_tool_schema_validation() {
|
|
let temp_dir = TempDir::new().unwrap();
|
|
let data_dir = temp_dir.path().to_path_buf();
|
|
let session_manager = Arc::new(SessionManager::new(data_dir.clone()));
|
|
let permission_manager = Arc::new(PermissionManager::new(data_dir));
|
|
let mock_scheduler = Arc::new(MockScheduler::new());
|
|
let config = AgentConfig::new(
|
|
session_manager,
|
|
permission_manager,
|
|
Some(mock_scheduler),
|
|
GooseMode::Auto,
|
|
false,
|
|
GoosePlatform::GooseCli,
|
|
);
|
|
let agent = Agent::with_config(config);
|
|
|
|
let tools = agent.list_tools("test-session-id", None).await;
|
|
let schedule_tool = tools
|
|
.iter()
|
|
.find(|tool| tool.name == PLATFORM_MANAGE_SCHEDULE_TOOL_NAME);
|
|
assert!(schedule_tool.is_some());
|
|
|
|
let tool = schedule_tool.unwrap();
|
|
|
|
// Verify the tool schema has the session_id parameter for session_content action
|
|
if let Some(properties) = tool.input_schema.get("properties") {
|
|
assert!(properties.get("session_id").is_some());
|
|
|
|
if let Some(session_id_prop) = properties.get("session_id") {
|
|
assert_eq!(
|
|
session_id_prop.get("type").unwrap().as_str().unwrap(),
|
|
"string"
|
|
);
|
|
assert!(session_id_prop
|
|
.get("description")
|
|
.unwrap()
|
|
.as_str()
|
|
.unwrap()
|
|
.contains("Session identifier for session_content action"));
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod retry_tests {
|
|
use super::*;
|
|
use goose::agents::types::{RetryConfig, SuccessCheck};
|
|
|
|
#[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::*;
|
|
use async_trait::async_trait;
|
|
use goose::agents::SessionConfig;
|
|
use goose::conversation::message::{Message, MessageContent};
|
|
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 rmcp::model::{CallToolRequestParams, Tool};
|
|
use rmcp::object;
|
|
use std::path::PathBuf;
|
|
|
|
struct MockToolProvider {}
|
|
|
|
impl MockToolProvider {
|
|
fn new() -> Self {
|
|
Self {}
|
|
}
|
|
}
|
|
|
|
impl ProviderDef for MockToolProvider {
|
|
type Provider = Self;
|
|
|
|
fn metadata() -> ProviderMetadata {
|
|
ProviderMetadata {
|
|
name: "mock".to_string(),
|
|
display_name: "Mock Provider".to_string(),
|
|
description: "Mock provider for testing".to_string(),
|
|
default_model: "mock-model".to_string(),
|
|
known_models: vec![],
|
|
model_doc_link: "".to_string(),
|
|
config_keys: vec![],
|
|
}
|
|
}
|
|
|
|
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 MockToolProvider {
|
|
async fn stream(
|
|
&self,
|
|
_model_config: &ModelConfig,
|
|
_session_id: &str,
|
|
_system_prompt: &str,
|
|
_messages: &[Message],
|
|
_tools: &[Tool],
|
|
) -> Result<MessageStream, ProviderError> {
|
|
let tool_call = CallToolRequestParams::new("test_tool")
|
|
.with_arguments(object!({"param": "value"}));
|
|
let message = Message::assistant().with_tool_request("call_123", Ok(tool_call));
|
|
|
|
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 {
|
|
"mock-test"
|
|
}
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_max_turns_limit() -> Result<()> {
|
|
let agent = Agent::new();
|
|
let provider = Arc::new(MockToolProvider::new());
|
|
let user_message = Message::user().with_text("Hello");
|
|
|
|
let session = agent
|
|
.config
|
|
.session_manager
|
|
.create_session(
|
|
PathBuf::default(),
|
|
"max-turn-test".to_string(),
|
|
SessionType::Hidden,
|
|
)
|
|
.await?;
|
|
|
|
agent.update_provider(provider, &session.id).await?;
|
|
|
|
let session_config = SessionConfig {
|
|
id: session.id,
|
|
schedule_id: None,
|
|
max_turns: Some(1),
|
|
retry_config: None,
|
|
};
|
|
|
|
let reply_stream = agent.reply(user_message, session_config, None).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)) => {
|
|
if let Some(MessageContent::ActionRequired(action)) =
|
|
response.content.first()
|
|
{
|
|
if let goose::conversation::message::ActionRequiredData::ToolConfirmation { id, .. } = &action.data {
|
|
agent.handle_confirmation(
|
|
id.clone(),
|
|
goose::permission::PermissionConfirmation {
|
|
principal_type: goose::permission::permission_confirmation::PrincipalType::Tool,
|
|
permission: goose::permission::Permission::AllowOnce,
|
|
}
|
|
).await;
|
|
}
|
|
}
|
|
responses.push(response);
|
|
}
|
|
Ok(AgentEvent::McpNotification(_)) => {}
|
|
Ok(AgentEvent::ModelChange { .. }) => {}
|
|
Ok(AgentEvent::HistoryReplaced(_updated_conversation)) => {
|
|
// We should update the conversation here, but we're not reading it
|
|
}
|
|
Err(e) => {
|
|
return Err(e);
|
|
}
|
|
}
|
|
}
|
|
|
|
assert!(
|
|
!responses.is_empty(),
|
|
"Expected at least 1 response, got {}",
|
|
responses.len()
|
|
);
|
|
|
|
// Look for the max turns message as the last response
|
|
let last_response = responses.last().unwrap();
|
|
let last_content = last_response.content.first().unwrap();
|
|
if let MessageContent::Text(text_content) = last_content {
|
|
assert!(text_content.text.contains(
|
|
"I've reached the maximum number of actions I can do without user input"
|
|
));
|
|
} else {
|
|
panic!("Expected text content in last message");
|
|
}
|
|
Ok(())
|
|
}
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod extension_manager_tests {
|
|
use super::*;
|
|
use goose::agents::extension::ExtensionConfig;
|
|
use goose::agents::platform_extensions::{
|
|
MANAGE_EXTENSIONS_TOOL_NAME, SEARCH_AVAILABLE_EXTENSIONS_TOOL_NAME,
|
|
};
|
|
use goose::agents::AgentConfig;
|
|
use goose::config::permission::PermissionManager;
|
|
use goose::config::GooseMode;
|
|
use goose::session::SessionManager;
|
|
|
|
async fn setup_agent_with_extension_manager() -> (Agent, String) {
|
|
use goose::session::session_manager::SessionType;
|
|
|
|
// Add the TODO extension to the config so it can be discovered by search_available_extensions
|
|
// Set it as disabled initially so tests can enable it
|
|
let todo_extension_entry = ExtensionEntry {
|
|
enabled: false,
|
|
config: ExtensionConfig::Platform {
|
|
name: "todo".to_string(),
|
|
description:
|
|
"Enable a todo list for goose so it can keep track of what it is doing"
|
|
.to_string(),
|
|
display_name: Some("Todo".to_string()),
|
|
bundled: Some(true),
|
|
available_tools: vec![],
|
|
},
|
|
};
|
|
set_extension(todo_extension_entry);
|
|
|
|
// Create agent with session_id from the start
|
|
let temp_dir = tempfile::tempdir().unwrap();
|
|
let session_manager = Arc::new(SessionManager::new(temp_dir.path().to_path_buf()));
|
|
let config = AgentConfig::new(
|
|
session_manager.clone(),
|
|
PermissionManager::instance(),
|
|
None,
|
|
GooseMode::Auto,
|
|
false,
|
|
GoosePlatform::GooseCli,
|
|
);
|
|
|
|
let agent = Agent::with_config(config);
|
|
|
|
let session = session_manager
|
|
.create_session(
|
|
std::path::PathBuf::from("."),
|
|
"Test Session".to_string(),
|
|
SessionType::Hidden,
|
|
)
|
|
.await
|
|
.expect("Failed to create session");
|
|
let session_id = session.id;
|
|
|
|
// Now add the extension manager platform extension
|
|
let ext_config = ExtensionConfig::Platform {
|
|
name: "extensionmanager".to_string(),
|
|
description: "Extension Manager".to_string(),
|
|
display_name: Some("Extension Manager".to_string()),
|
|
bundled: Some(true),
|
|
available_tools: vec![],
|
|
};
|
|
|
|
agent
|
|
.add_extension(ext_config, &session_id)
|
|
.await
|
|
.expect("Failed to add extension manager");
|
|
(agent, session_id)
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_extension_manager_tools_available() {
|
|
let (agent, session_id) = setup_agent_with_extension_manager().await;
|
|
let tools = agent.list_tools(&session_id, None).await;
|
|
|
|
// Note: Tool names are prefixed with the normalized extension name "extensionmanager"
|
|
// not the display name "Extension Manager"
|
|
let search_tool = tools.iter().find(|tool| {
|
|
tool.name == format!("extensionmanager__{SEARCH_AVAILABLE_EXTENSIONS_TOOL_NAME}")
|
|
});
|
|
assert!(
|
|
search_tool.is_some(),
|
|
"search_available_extensions tool should be available"
|
|
);
|
|
|
|
let manage_tool = tools.iter().find(|tool| {
|
|
tool.name == format!("extensionmanager__{MANAGE_EXTENSIONS_TOOL_NAME}")
|
|
});
|
|
assert!(
|
|
manage_tool.is_some(),
|
|
"manage_extensions tool should be available"
|
|
);
|
|
}
|
|
}
|
|
}
|