Skip subagents for gemini (#5257)

Co-authored-by: Douwe Osinga <douwe@squareup.com>
This commit is contained in:
Douwe Osinga
2025-10-18 17:35:29 -04:00
committed by GitHub
parent 890393bb68
commit 64b37339e0
12 changed files with 156 additions and 298 deletions
+2 -7
View File
@@ -1522,13 +1522,9 @@ impl Agent {
self.extension_manager self.extension_manager
.suggest_disable_extensions_prompt() .suggest_disable_extensions_prompt()
.await, .await,
Some(model_name), model_name,
false, false,
); );
tracing::debug!(
"Built system prompt with {} characters",
system_prompt.len()
);
let recipe_prompt = prompt_manager.get_recipe_prompt().await; let recipe_prompt = prompt_manager.get_recipe_prompt().await;
let tools = self let tools = self
@@ -1539,7 +1535,6 @@ impl Agent {
tracing::error!("Failed to get tools for recipe creation: {}", e); tracing::error!("Failed to get tools for recipe creation: {}", e);
e e
})?; })?;
tracing::debug!("Retrieved {} tools for recipe creation", tools.len());
messages.push(Message::user().with_text(recipe_prompt)); messages.push(Message::user().with_text(recipe_prompt));
@@ -1754,7 +1749,7 @@ mod tests {
let prompt_manager = agent.prompt_manager.lock().await; let prompt_manager = agent.prompt_manager.lock().await;
let system_prompt = let system_prompt =
prompt_manager.build_system_prompt(vec![], None, Value::Null, None, false); prompt_manager.build_system_prompt(vec![], None, Value::Null, "gpt-4o", false);
let final_output_tool_ref = agent.final_output_tool.lock().await; let final_output_tool_ref = agent.final_output_tool.lock().await;
let final_output_tool_system_prompt = let final_output_tool_system_prompt =
+41 -111
View File
@@ -3,8 +3,8 @@ use serde_json::Value;
use std::collections::HashMap; use std::collections::HashMap;
use crate::agents::extension::ExtensionInfo; use crate::agents::extension::ExtensionInfo;
use crate::agents::recipe_tools::dynamic_task_tools::should_enabled_subagents;
use crate::agents::router_tools::llm_search_tool_prompt; use crate::agents::router_tools::llm_search_tool_prompt;
use crate::providers::base::get_current_model;
use crate::{config::Config, prompt_template, utils::sanitize_unicode_tags}; use crate::{config::Config, prompt_template, utils::sanitize_unicode_tags};
pub struct PromptManager { pub struct PromptManager {
@@ -39,35 +39,12 @@ impl PromptManager {
self.system_prompt_override = Some(template); self.system_prompt_override = Some(template);
} }
/// Normalize a model name (replace - and / with _, lower case)
fn normalize_model_name(name: &str) -> String {
name.replace(['-', '/', '.'], "_").to_lowercase()
}
/// Map model (normalized) to prompt filenames; returns filename if a key is contained in the normalized model
fn model_prompt_map(model: &str) -> &'static str {
let mut map = HashMap::new();
map.insert("gpt_4_1", "system_gpt_4.1.md");
// Add more mappings as needed
let norm_model = Self::normalize_model_name(model);
for (key, val) in &map {
if norm_model.contains(key) {
return val;
}
}
"system.md"
}
/// Build the final system prompt
///
/// * `extensions_info` extension information for each extension/MCP
/// * `frontend_instructions` instructions for the "frontend" tool
pub fn build_system_prompt( pub fn build_system_prompt(
&self, &self,
extensions_info: Vec<ExtensionInfo>, extensions_info: Vec<ExtensionInfo>,
frontend_instructions: Option<String>, frontend_instructions: Option<String>,
suggest_disable_extensions_prompt: Value, suggest_disable_extensions_prompt: Value,
model_name: Option<&str>, model_name: &str,
router_enabled: bool, router_enabled: bool,
) -> String { ) -> String {
let mut context: HashMap<&str, Value> = HashMap::new(); let mut context: HashMap<&str, Value> = HashMap::new();
@@ -113,36 +90,23 @@ impl PromptManager {
Value::String(suggest_disable_extensions_prompt.to_string()), Value::String(suggest_disable_extensions_prompt.to_string()),
); );
// Add the mode to the context for conditional rendering
let config = Config::global(); let config = Config::global();
let goose_mode = config.get_param("GOOSE_MODE").unwrap_or("auto".to_string()); let goose_mode = config.get_param("GOOSE_MODE").unwrap_or("auto".to_string());
context.insert("goose_mode", Value::String(goose_mode.clone())); context.insert("goose_mode", Value::String(goose_mode.clone()));
context.insert("is_autonomous", Value::Bool(goose_mode == "auto")); context.insert(
"enable_subagents",
Value::Bool(should_enabled_subagents(model_name)),
);
// First check the global store, and only if it's not available, fall back to the provided model_name
let model_to_use: Option<String> =
get_current_model().or_else(|| model_name.map(|s| s.to_string()));
// Conditionally load the override prompt or the global system prompt
let base_prompt = if let Some(override_prompt) = &self.system_prompt_override { let base_prompt = if let Some(override_prompt) = &self.system_prompt_override {
let sanitized_override_prompt = sanitize_unicode_tags(override_prompt); let sanitized_override_prompt = sanitize_unicode_tags(override_prompt);
prompt_template::render_inline_once(&sanitized_override_prompt, &context) prompt_template::render_inline_once(&sanitized_override_prompt, &context)
.expect("Prompt should render")
} else if let Some(model) = &model_to_use {
// Use the fuzzy mapping to determine the prompt file, or fall back to legacy logic
let prompt_file = Self::model_prompt_map(model);
match prompt_template::render_global_file(prompt_file, &context) {
Ok(prompt) => prompt,
Err(_) => {
// Fall back to the standard system.md if model-specific one doesn't exist
prompt_template::render_global_file("system.md", &context)
.expect("Prompt should render")
}
}
} else { } else {
prompt_template::render_global_file("system.md", &context) prompt_template::render_global_file("system.md", &context)
.expect("Prompt should render") }
}; .unwrap_or_else(|_| {
"You are a general-purpose AI agent called goose, created by Block".to_string()
});
let mut system_prompt_extras = self.system_prompt_extras.clone(); let mut system_prompt_extras = self.system_prompt_extras.clone();
if goose_mode == "chat" { if goose_mode == "chat" {
@@ -150,9 +114,6 @@ impl PromptManager {
"Right now you are in the chat only mode, no access to any tool use and system." "Right now you are in the chat only mode, no access to any tool use and system."
.to_string(), .to_string(),
); );
} else {
system_prompt_extras
.push("Right now you are *NOT* in the chat only mode and have access to tool use and system.".to_string());
} }
let sanitized_system_prompt_extras: Vec<String> = system_prompt_extras let sanitized_system_prompt_extras: Vec<String> = system_prompt_extras
@@ -173,7 +134,8 @@ impl PromptManager {
pub async fn get_recipe_prompt(&self) -> String { pub async fn get_recipe_prompt(&self) -> String {
let context: HashMap<&str, Value> = HashMap::new(); let context: HashMap<&str, Value> = HashMap::new();
prompt_template::render_global_file("recipe.md", &context).expect("Prompt should render") prompt_template::render_global_file("recipe.md", &context)
.unwrap_or_else(|_| "The recipe prompt is busted. Tell the user.".to_string())
} }
} }
@@ -181,66 +143,19 @@ impl PromptManager {
mod tests { mod tests {
use super::*; use super::*;
#[test]
fn test_normalize_model_name() {
assert_eq!(PromptManager::normalize_model_name("gpt-4.1"), "gpt_4_1");
assert_eq!(PromptManager::normalize_model_name("gpt/3.5"), "gpt_3_5");
assert_eq!(
PromptManager::normalize_model_name("GPT-3.5/PLUS"),
"gpt_3_5_plus"
);
}
#[test]
fn test_model_prompt_map_matches() {
// should match prompts based on contained normalized keys
assert_eq!(
PromptManager::model_prompt_map("gpt-4.1"),
"system_gpt_4.1.md"
);
assert_eq!(
PromptManager::model_prompt_map("gpt-4.1-2025-04-14"),
"system_gpt_4.1.md"
);
assert_eq!(
PromptManager::model_prompt_map("openai/gpt-4.1"),
"system_gpt_4.1.md"
);
assert_eq!(
PromptManager::model_prompt_map("goose-gpt-4-1"),
"system_gpt_4.1.md"
);
assert_eq!(
PromptManager::model_prompt_map("gpt-4-1-huge"),
"system_gpt_4.1.md"
);
}
#[test]
fn test_model_prompt_map_none() {
// should return system.md for unrecognized/unsupported model names
assert_eq!(PromptManager::model_prompt_map("llama-3-70b"), "system.md");
assert_eq!(PromptManager::model_prompt_map("goose"), "system.md");
assert_eq!(
PromptManager::model_prompt_map("claude-3.7-sonnet"),
"system.md"
);
assert_eq!(
PromptManager::model_prompt_map("xxx-unknown-model"),
"system.md"
);
}
#[test] #[test]
fn test_build_system_prompt_sanitizes_override() { fn test_build_system_prompt_sanitizes_override() {
let mut manager = PromptManager::new(); let mut manager = PromptManager::new();
let malicious_override = "System prompt\u{E0041}\u{E0042}\u{E0043}with hidden text"; let malicious_override = "System prompt\u{E0041}\u{E0042}\u{E0043}with hidden text";
manager.set_system_prompt_override(malicious_override.to_string()); manager.set_system_prompt_override(malicious_override.to_string());
let result = let result = manager.build_system_prompt(
manager.build_system_prompt(vec![], None, Value::String("".to_string()), None, false); vec![],
None,
Value::String("".to_string()),
"gpt-4o",
false,
);
assert!(!result.contains('\u{E0041}')); assert!(!result.contains('\u{E0041}'));
assert!(!result.contains('\u{E0042}')); assert!(!result.contains('\u{E0042}'));
@@ -255,8 +170,13 @@ mod tests {
let malicious_extra = "Extra instruction\u{E0041}\u{E0042}\u{E0043}hidden"; let malicious_extra = "Extra instruction\u{E0041}\u{E0042}\u{E0043}hidden";
manager.add_system_prompt_extra(malicious_extra.to_string()); manager.add_system_prompt_extra(malicious_extra.to_string());
let result = let result = manager.build_system_prompt(
manager.build_system_prompt(vec![], None, Value::String("".to_string()), None, false); vec![],
None,
Value::String("".to_string()),
"gpt-4o",
false,
);
assert!(!result.contains('\u{E0041}')); assert!(!result.contains('\u{E0041}'));
assert!(!result.contains('\u{E0042}')); assert!(!result.contains('\u{E0042}'));
@@ -272,8 +192,13 @@ mod tests {
manager.add_system_prompt_extra("Second\u{E0042}instruction".to_string()); manager.add_system_prompt_extra("Second\u{E0042}instruction".to_string());
manager.add_system_prompt_extra("Third\u{E0043}instruction".to_string()); manager.add_system_prompt_extra("Third\u{E0043}instruction".to_string());
let result = let result = manager.build_system_prompt(
manager.build_system_prompt(vec![], None, Value::String("".to_string()), None, false); vec![],
None,
Value::String("".to_string()),
"gpt-4o",
false,
);
assert!(!result.contains('\u{E0041}')); assert!(!result.contains('\u{E0041}'));
assert!(!result.contains('\u{E0042}')); assert!(!result.contains('\u{E0042}'));
@@ -289,8 +214,13 @@ mod tests {
let legitimate_unicode = "Instruction with 世界 and 🌍 emojis"; let legitimate_unicode = "Instruction with 世界 and 🌍 emojis";
manager.add_system_prompt_extra(legitimate_unicode.to_string()); manager.add_system_prompt_extra(legitimate_unicode.to_string());
let result = let result = manager.build_system_prompt(
manager.build_system_prompt(vec![], None, Value::String("".to_string()), None, false); vec![],
None,
Value::String("".to_string()),
"gpt-4o",
false,
);
assert!(result.contains("世界")); assert!(result.contains("世界"));
assert!(result.contains("🌍")); assert!(result.contains("🌍"));
@@ -311,7 +241,7 @@ mod tests {
vec![malicious_extension_info], vec![malicious_extension_info],
None, None,
Value::String("".to_string()), Value::String("".to_string()),
None, "gpt-4o",
false, false,
); );
@@ -27,6 +27,7 @@ pub struct CreateDynamicTaskParams {
/// How to execute multiple tasks (default: parallel for multiple tasks, sequential for single task) /// How to execute multiple tasks (default: parallel for multiple tasks, sequential for single task)
#[serde(skip_serializing_if = "Option::is_none")] #[serde(skip_serializing_if = "Option::is_none")]
#[schemars(with = "Option<String>")]
pub execution_mode: Option<ExecutionModeParam>, pub execution_mode: Option<ExecutionModeParam>,
} }
@@ -90,6 +91,18 @@ pub struct TaskParameter {
pub return_last_only: Option<bool>, pub return_last_only: Option<bool>,
} }
pub fn should_enabled_subagents(model_name: &str) -> bool {
let config = crate::config::Config::global();
let is_autonomous = config.get_param("GOOSE_MODE").unwrap_or("auto".to_string()) == "auto";
if !is_autonomous {
return false;
}
if model_name.starts_with("gemini") {
return false;
}
true
}
pub fn create_dynamic_task_tool() -> Tool { pub fn create_dynamic_task_tool() -> Tool {
let schema = schema_for!(CreateDynamicTaskParams); let schema = schema_for!(CreateDynamicTaskParams);
let schema_value = let schema_value =
+7 -9
View File
@@ -15,6 +15,7 @@ use crate::providers::toolshim::{
modify_system_prompt_for_tool_json, OllamaInterpreter, modify_system_prompt_for_tool_json, OllamaInterpreter,
}; };
use crate::agents::recipe_tools::dynamic_task_tools::should_enabled_subagents;
use crate::session::SessionManager; use crate::session::SessionManager;
use rmcp::model::Tool; use rmcp::model::Tool;
@@ -32,23 +33,20 @@ async fn toolshim_postprocess(
} }
impl Agent { impl Agent {
/// Prepares tools and system prompt for a provider request pub async fn prepare_tools_and_prompt(&self) -> Result<(Vec<Tool>, Vec<Tool>, String)> {
pub async fn prepare_tools_and_prompt(&self) -> anyhow::Result<(Vec<Tool>, Vec<Tool>, String)> {
// Get router enabled status // Get router enabled status
let router_enabled = self.tool_route_manager.is_router_enabled().await; let router_enabled = self.tool_route_manager.is_router_enabled().await;
// Get tools from extension manager // Get tools from extension manager
let mut tools = self.list_tools_for_router().await; let mut tools = self.list_tools_for_router().await;
let config = crate::config::Config::global();
let is_autonomous = config.get_param("GOOSE_MODE").unwrap_or("auto".to_string()) == "auto";
// If router is disabled and no tools were returned, fall back to regular tools // If router is disabled and no tools were returned, fall back to regular tools
if !router_enabled && tools.is_empty() { if !router_enabled && tools.is_empty() {
// Get all tools but filter out subagent tools if not in autonomous mode
tools = self.list_tools(None).await; tools = self.list_tools(None).await;
if !is_autonomous { let provider = self.provider().await?;
// Filter out subagent-related tools let model_name = provider.get_model_config().model_name;
if !should_enabled_subagents(&model_name) {
tools.retain(|tool| { tools.retain(|tool| {
tool.name != crate::agents::subagent_execution_tool::subagent_execute_task_tool::SUBAGENT_EXECUTE_TASK_TOOL_NAME tool.name != crate::agents::subagent_execution_tool::subagent_execute_task_tool::SUBAGENT_EXECUTE_TASK_TOOL_NAME
&& tool.name != crate::agents::recipe_tools::dynamic_task_tools::DYNAMIC_TASK_TOOL_NAME_PREFIX && tool.name != crate::agents::recipe_tools::dynamic_task_tools::DYNAMIC_TASK_TOOL_NAME_PREFIX
@@ -77,7 +75,7 @@ impl Agent {
self.extension_manager self.extension_manager
.suggest_disable_extensions_prompt() .suggest_disable_extensions_prompt()
.await, .await,
Some(model_name), model_name,
router_enabled, router_enabled,
); );
@@ -473,56 +473,4 @@ mod tests {
assert!(result.needs_approval.iter().any(|req| req.id == "tool_3")); assert!(result.needs_approval.iter().any(|req| req.id == "tool_3"));
assert!(enable_extension_request_ids.iter().any(|id| id == "tool_3")); assert!(enable_extension_request_ids.iter().any(|id| id == "tool_3"));
} }
#[tokio::test]
async fn test_check_tool_permissions_auto() {
// Setup mocks
let temp_file = NamedTempFile::new().unwrap();
let temp_path = temp_file.path();
let mut permission_manager = PermissionManager::new(temp_path);
let provider = create_mock_provider();
let tools_with_readonly_annotation: HashSet<String> =
vec!["file_reader".to_string()].into_iter().collect();
let tools_without_annotation: HashSet<String> =
vec!["data_fetcher".to_string()].into_iter().collect();
permission_manager.update_user_permission("file_reader", PermissionLevel::AlwaysAllow);
permission_manager
.update_smart_approve_permission("data_fetcher", PermissionLevel::AskBefore);
let tool_request_1 = ToolRequest {
id: "tool_1".to_string(),
tool_call: Ok(CallToolRequestParam {
name: "file_reader".into(),
arguments: Some(object!({"path": "/path/to/file"})),
}),
};
let tool_request_2 = ToolRequest {
id: "tool_2".to_string(),
tool_call: Ok(CallToolRequestParam {
name: "data_fetcher".into(),
arguments: Some(object!({"url": "http://example.com"})),
}),
};
let candidate_requests: Vec<ToolRequest> = vec![tool_request_1, tool_request_2];
// Call the function under test
let (result, _) = check_tool_permissions(
&candidate_requests,
"auto",
tools_with_readonly_annotation,
tools_without_annotation,
&mut permission_manager,
provider,
)
.await;
// Validate the result
assert_eq!(result.approved.len(), 2); // file_reader should be approved
assert_eq!(result.needs_approval.len(), 0); // data_fetcher should need approval
assert_eq!(result.denied.len(), 0); // No tool should be denied in this test
}
} }
+3 -2
View File
@@ -48,7 +48,8 @@ No extensions are defined. You should let the user know that they should add ext
{% endif %} {% endif %}
{{tool_selection_strategy}} {{tool_selection_strategy}}
{% if is_autonomous %} {% if enable_subagents %}
# sub agents # sub agents
Execute self contained tasks where step-by-step visibility is not important through subagents. Execute self contained tasks where step-by-step visibility is not important through subagents.
@@ -60,7 +61,7 @@ Execute self contained tasks where step-by-step visibility is not important thro
- Provide all needed context — subagents cannot see your context - Provide all needed context — subagents cannot see your context
- Use extension filters to limit resource access - Use extension filters to limit resource access
- Use return_last_only when only a summary or simple answer is required — inform subagent of this choice. - Use return_last_only when only a summary or simple answer is required — inform subagent of this choice.
{% endif %} {% endif %}
# Response Guidelines # Response Guidelines
-1
View File
@@ -46,7 +46,6 @@ pub const DATABRICKS_KNOWN_MODELS: &[&str] = &[
"databricks-meta-llama-3-3-70b-instruct", "databricks-meta-llama-3-3-70b-instruct",
"databricks-meta-llama-3-1-405b-instruct", "databricks-meta-llama-3-1-405b-instruct",
"databricks-dbrx-instruct", "databricks-dbrx-instruct",
"databricks-mixtral-8x7b-instruct",
]; ];
pub const DATABRICKS_DOC_URL: &str = pub const DATABRICKS_DOC_URL: &str =
@@ -264,23 +264,31 @@ fn format_messages(messages: &[Message], image_format: &ImageFormat) -> Vec<Data
result result
} }
/// Convert internal Tool format to OpenAI's API tool specification pub fn format_tools(tools: &[Tool], model_name: &str) -> anyhow::Result<Vec<Value>> {
pub fn format_tools(tools: &[Tool]) -> anyhow::Result<Vec<Value>> {
let mut tool_names = std::collections::HashSet::new(); let mut tool_names = std::collections::HashSet::new();
let mut result = Vec::new(); let mut result = Vec::new();
let is_gemini = model_name.starts_with("gemini");
for tool in tools { for tool in tools {
if !tool_names.insert(&tool.name) { if !tool_names.insert(&tool.name) {
return Err(anyhow!("Duplicate tool name: {}", tool.name)); return Err(anyhow!("Duplicate tool name: {}", tool.name));
} }
let parameters = if is_gemini {
let mut cleaned_schema = tool.input_schema.as_ref().clone();
cleaned_schema.remove("$schema");
json!(cleaned_schema)
} else {
json!(tool.input_schema)
};
result.push(json!({ result.push(json!({
"type": "function", "type": "function",
"function": { "function": {
"name": tool.name, "name": tool.name,
// do not silently truncate description
"description": tool.description, "description": tool.description,
"parameters": tool.input_schema, "parameters": parameters,
} }
})); }));
} }
@@ -544,7 +552,7 @@ pub fn create_request(
let messages_spec = format_messages(messages, image_format); let messages_spec = format_messages(messages, image_format);
let mut tools_spec = if !tools.is_empty() { let mut tools_spec = if !tools.is_empty() {
format_tools(tools)? format_tools(tools, &model_config.model_name)?
} else { } else {
vec![] vec![]
}; };
@@ -639,82 +647,6 @@ mod tests {
use rmcp::object; use rmcp::object;
use serde_json::json; use serde_json::json;
#[test]
fn test_validate_tool_schemas() {
// Test case 1: Empty parameters object
// Input JSON with an incomplete parameters object
let mut actual = vec![json!({
"type": "function",
"function": {
"name": "test_func",
"description": "test description",
"parameters": {
"type": "object"
}
}
})];
// Run the function to validate and update schemas
validate_tool_schemas(&mut actual);
// Expected JSON after validation
let expected = vec![json!({
"type": "function",
"function": {
"name": "test_func",
"description": "test description",
"parameters": {
"type": "object",
"properties": {},
"required": []
}
}
})];
// Compare entire JSON structures instead of individual fields
assert_eq!(actual, expected);
// Test case 2: Missing type field
let mut tools = vec![json!({
"type": "function",
"function": {
"name": "test_func",
"description": "test description",
"parameters": {
"properties": {}
}
}
})];
validate_tool_schemas(&mut tools);
let params = tools[0]["function"]["parameters"].as_object().unwrap();
assert_eq!(params["type"], "object");
// Test case 3: Complete valid schema should remain unchanged
let original_schema = json!({
"type": "function",
"function": {
"name": "test_func",
"description": "test description",
"parameters": {
"type": "object",
"properties": {
"location": {
"type": "string",
"description": "City and country"
}
},
"required": ["location"]
}
}
});
let mut tools = vec![original_schema.clone()];
validate_tool_schemas(&mut tools);
assert_eq!(tools[0], original_schema);
}
const OPENAI_TOOL_USE_RESPONSE: &str = r#"{ const OPENAI_TOOL_USE_RESPONSE: &str = r#"{
"choices": [{ "choices": [{
"role": "assistant", "role": "assistant",
@@ -752,6 +684,7 @@ mod tests {
"test_tool", "test_tool",
"A test tool", "A test tool",
object!({ object!({
"$schema": "http://json-schema.org/draft-07/schema#",
"type": "object", "type": "object",
"properties": { "properties": {
"input": { "input": {
@@ -763,11 +696,16 @@ mod tests {
}), }),
); );
let spec = format_tools(&[tool])?; let spec = format_tools(&[tool.clone()], "gpt-4o")?;
assert_eq!(
spec[0]["function"]["parameters"]["$schema"],
"http://json-schema.org/draft-07/schema#"
);
let spec = format_tools(&[tool], "gemini-2-5-flash")?;
assert!(spec[0]["function"]["parameters"].get("$schema").is_none());
assert_eq!(spec[0]["function"]["parameters"]["type"], "object");
assert_eq!(spec.len(), 1);
assert_eq!(spec[0]["type"], "function");
assert_eq!(spec[0]["function"]["name"], "test_tool");
Ok(()) Ok(())
} }
@@ -785,7 +723,6 @@ mod tests {
), ),
]; ];
// Get the ID from the tool request to use in the response
let tool_id = if let MessageContent::ToolRequest(request) = &messages[2].content[0] { let tool_id = if let MessageContent::ToolRequest(request) = &messages[2].content[0] {
&request.id &request.id
} else { } else {
@@ -823,7 +760,6 @@ mod tests {
}), }),
)]; )];
// Get the ID from the tool request to use in the response
let tool_id = if let MessageContent::ToolRequest(request) = &messages[0].content[0] { let tool_id = if let MessageContent::ToolRequest(request) = &messages[0].content[0] {
&request.id &request.id
} else { } else {
@@ -879,7 +815,7 @@ mod tests {
}), }),
); );
let result = format_tools(&[tool1, tool2]); let result = format_tools(&[tool1, tool2], "gpt-4o");
assert!(result.is_err()); assert!(result.is_err());
assert!(result assert!(result
.unwrap_err() .unwrap_err()
@@ -889,16 +825,8 @@ mod tests {
Ok(()) Ok(())
} }
#[test]
fn test_format_tools_empty() -> anyhow::Result<()> {
let spec = format_tools(&[])?;
assert!(spec.is_empty());
Ok(())
}
#[test] #[test]
fn test_format_messages_with_image_path() -> anyhow::Result<()> { fn test_format_messages_with_image_path() -> anyhow::Result<()> {
// Create a temporary PNG file with valid PNG magic numbers
let temp_dir = tempfile::tempdir()?; let temp_dir = tempfile::tempdir()?;
let png_path = temp_dir.path().join("test.png"); let png_path = temp_dir.path().join("test.png");
let png_data = [ let png_data = [
+1 -2
View File
@@ -130,7 +130,6 @@ pub fn format_messages(messages: &[Message]) -> Vec<Value> {
.collect() .collect()
} }
/// Convert internal Tool format to Google's API tool specification
pub fn format_tools(tools: &[Tool]) -> Vec<Value> { pub fn format_tools(tools: &[Tool]) -> Vec<Value> {
tools tools
.iter() .iter()
@@ -139,7 +138,7 @@ pub fn format_tools(tools: &[Tool]) -> Vec<Value> {
parameters.insert("name".to_string(), json!(tool.name)); parameters.insert("name".to_string(), json!(tool.name));
parameters.insert("description".to_string(), json!(tool.description)); parameters.insert("description".to_string(), json!(tool.description));
let tool_input_schema = &tool.input_schema; let tool_input_schema = &tool.input_schema;
// Only add the parameters key if the tool schema has non-empty properties.
if tool_input_schema if tool_input_schema
.get("properties") .get("properties")
.and_then(|v| v.as_object()) .and_then(|v| v.as_object())
+2 -16
View File
@@ -13,10 +13,9 @@ use rmcp::model::Tool;
use serde_json::Value; use serde_json::Value;
pub const GOOGLE_API_HOST: &str = "https://generativelanguage.googleapis.com"; pub const GOOGLE_API_HOST: &str = "https://generativelanguage.googleapis.com";
pub const GOOGLE_DEFAULT_MODEL: &str = "gemini-2.5-flash"; pub const GOOGLE_DEFAULT_MODEL: &str = "gemini-2.5-pro";
pub const GOOGLE_DEFAULT_FAST_MODEL: &str = "gemini-1.5-flash"; pub const GOOGLE_DEFAULT_FAST_MODEL: &str = "gemini-2.5-flash";
pub const GOOGLE_KNOWN_MODELS: &[&str] = &[ pub const GOOGLE_KNOWN_MODELS: &[&str] = &[
// Gemini 2.5 models (latest generation)
"gemini-2.5-pro", "gemini-2.5-pro",
"gemini-2.5-pro-preview-06-05", "gemini-2.5-pro-preview-06-05",
"gemini-2.5-pro-preview-05-06", "gemini-2.5-pro-preview-05-06",
@@ -27,20 +26,10 @@ pub const GOOGLE_KNOWN_MODELS: &[&str] = &[
"gemini-2.5-flash-exp-native-audio-thinking-dialog", "gemini-2.5-flash-exp-native-audio-thinking-dialog",
"gemini-2.5-flash-preview-tts", "gemini-2.5-flash-preview-tts",
"gemini-2.5-pro-preview-tts", "gemini-2.5-pro-preview-tts",
// Gemini 2.0 models
"gemini-2.0-flash", "gemini-2.0-flash",
"gemini-2.0-flash-exp", "gemini-2.0-flash-exp",
"gemini-2.0-flash-preview-image-generation", "gemini-2.0-flash-preview-image-generation",
"gemini-2.0-flash-lite", "gemini-2.0-flash-lite",
// Gemini 1.5 models
"gemini-1.5-flash",
"gemini-1.5-flash-latest",
"gemini-1.5-flash-002",
"gemini-1.5-flash-8b",
"gemini-1.5-flash-8b-latest",
"gemini-1.5-pro",
"gemini-1.5-pro-latest",
"gemini-1.5-pro-002",
]; ];
pub const GOOGLE_DOC_URL: &str = "https://ai.google.dev/gemini-api/docs/models"; pub const GOOGLE_DOC_URL: &str = "https://ai.google.dev/gemini-api/docs/models";
@@ -115,7 +104,6 @@ impl Provider for GoogleProvider {
let payload = create_request(model_config, system, messages, tools)?; let payload = create_request(model_config, system, messages, tools)?;
let mut log = RequestLog::start(model_config, &payload)?; let mut log = RequestLog::start(model_config, &payload)?;
// Make request
let response = self let response = self
.with_retry(|| async { .with_retry(|| async {
let payload_clone = payload.clone(); let payload_clone = payload.clone();
@@ -123,7 +111,6 @@ impl Provider for GoogleProvider {
}) })
.await?; .await?;
// Parse response
let message = response_to_message(unescape_json_values(&response))?; let message = response_to_message(unescape_json_values(&response))?;
let usage = get_usage(&response)?; let usage = get_usage(&response)?;
let response_model = match response.get("modelVersion") { let response_model = match response.get("modelVersion") {
@@ -135,7 +122,6 @@ impl Provider for GoogleProvider {
Ok((message, provider_usage)) Ok((message, provider_usage))
} }
/// Fetch supported models from Google Generative Language API; returns Err on failure, Ok(None) if not present
async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> { async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
let response = self.api_client.response_get("v1beta/models").await?; let response = self.api_client.response_get("v1beta/models").await?;
let json: serde_json::Value = response.json().await?; let json: serde_json::Value = response.json().await?;
+2 -2
View File
@@ -17,7 +17,7 @@ use crate::providers::formats::openai::{create_request, get_usage, response_to_m
use rmcp::model::Tool; use rmcp::model::Tool;
pub const OPENROUTER_DEFAULT_MODEL: &str = "anthropic/claude-sonnet-4"; pub const OPENROUTER_DEFAULT_MODEL: &str = "anthropic/claude-sonnet-4";
pub const OPENROUTER_DEFAULT_FAST_MODEL: &str = "google/gemini-flash-1.5"; pub const OPENROUTER_DEFAULT_FAST_MODEL: &str = "google/gemini-flash-2.5";
pub const OPENROUTER_MODEL_PREFIX_ANTHROPIC: &str = "anthropic"; pub const OPENROUTER_MODEL_PREFIX_ANTHROPIC: &str = "anthropic";
// OpenRouter can run many models, we suggest the default // OpenRouter can run many models, we suggest the default
@@ -28,7 +28,7 @@ pub const OPENROUTER_KNOWN_MODELS: &[&str] = &[
"anthropic/claude-opus-4", "anthropic/claude-opus-4",
"anthropic/claude-3.7-sonnet", "anthropic/claude-3.7-sonnet",
"google/gemini-2.5-pro", "google/gemini-2.5-pro",
"google/gemini-flash-1.5", "google/gemini-flash-2.5",
"deepseek/deepseek-r1-0528", "deepseek/deepseek-r1-0528",
"qwen/qwen3-coder", "qwen/qwen3-coder",
"moonshotai/kimi-k2", "moonshotai/kimi-k2",
+61
View File
@@ -0,0 +1,61 @@
#!/bin/bash
if [ -f .env ]; then
export $(grep -v '^#' .env | xargs)
fi
echo "Building goose..."
cargo build --release --bin goose
echo ""
SCRIPT_DIR=$(pwd)
PROVIDERS=(
"openrouter:anthropic/claude-sonnet-4.5:google/gemini-flash-2.5:qwen/qwen3-coder"
"openai:gpt-4o:gpt-4o-mini:gpt-3.5-turbo"
"anthropic:claude-sonnet-4-0:claude-3-7-sonnet-latest"
"google:gemini-2.5-pro:gemini-2.5-pro:gemini-2.5-flash"
"databricks:databricks-claude-sonnet-4:gemini-2-5-flash:gpt-4o"
)
RESULTS=()
for provider_config in "${PROVIDERS[@]}"; do
IFS=':' read -ra PARTS <<< "$provider_config"
PROVIDER="${PARTS[0]}"
for i in $(seq 1 $((${#PARTS[@]} - 1))); do
MODEL="${PARTS[$i]}"
export GOOSE_PROVIDER="$PROVIDER"
export GOOSE_MODEL="$MODEL"
TESTDIR=$(mktemp -d)
echo "hello" > "$TESTDIR/hello.txt"
echo "Provider: ${PROVIDER}"
echo "Model: ${MODEL}"
echo ""
TMPFILE=$(mktemp)
(cd "$TESTDIR" && "$SCRIPT_DIR/target/release/goose" run --text "please list files in the current directory" --with-builtin developer 2>&1) | tee "$TMPFILE"
echo ""
if grep -q "shell | developer" "$TMPFILE"; then
echo "✓ SUCCESS: Test passed - developer tool called"
RESULTS+=("${PROVIDER}/${MODEL}")
else
echo "✗ FAILED: Test failed - no developer tools called"
RESULTS+=("${PROVIDER}/${MODEL}")
fi
rm "$TMPFILE"
rm -rf "$TESTDIR"
echo "---"
done
done
echo ""
echo "=== Test Summary ==="
for result in "${RESULTS[@]}"; do
echo "$result"
done
if echo "${RESULTS[@]}" | grep -q "✗"; then
echo ""
echo "Some tests failed!"
exit 1
else
echo ""
echo "All tests passed!"
fi