From e501e3b93c44c66f93408e71f349d348b06699d4 Mon Sep 17 00:00:00 2001 From: David Katz Date: Wed, 20 Aug 2025 13:48:49 -0400 Subject: [PATCH 01/31] my hints for goose --- crates/goose/src/context_mgmt/summarize.rs | 2 +- crates/goose/src/model.rs | 1 + crates/goose/src/providers/base.rs | 36 ++++++++++++++++++++++ 3 files changed, 38 insertions(+), 1 deletion(-) diff --git a/crates/goose/src/context_mgmt/summarize.rs b/crates/goose/src/context_mgmt/summarize.rs index 68947cfa98f7..3d53fcc4ed85 100644 --- a/crates/goose/src/context_mgmt/summarize.rs +++ b/crates/goose/src/context_mgmt/summarize.rs @@ -44,7 +44,7 @@ pub async fn summarize_messages( // Send the request to the provider and fetch the response let (mut response, mut provider_usage) = provider - .complete(&system_prompt, &summarization_request, &[]) + .complete_fast(&system_prompt, &summarization_request, &[]) .await?; // Set role to user as it will be used in following conversation as user content diff --git a/crates/goose/src/model.rs b/crates/goose/src/model.rs index b606eaa3b043..ef38382953aa 100644 --- a/crates/goose/src/model.rs +++ b/crates/goose/src/model.rs @@ -72,6 +72,7 @@ pub struct ModelConfig { pub max_tokens: Option, pub toolshim: bool, pub toolshim_model: Option, + pub fast_model: Option, } #[derive(Debug, Clone, Serialize, Deserialize)] diff --git a/crates/goose/src/providers/base.rs b/crates/goose/src/providers/base.rs index 60623abb3a3e..65156143dcd7 100644 --- a/crates/goose/src/providers/base.rs +++ b/crates/goose/src/providers/base.rs @@ -337,6 +337,42 @@ pub trait Provider: Send + Sync { tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError>; + /// Generate the next message using the configured model and other parameters + /// + /// # Arguments + /// * `system` - The system prompt that guides the model's behavior + /// * `messages` - The conversation history as a sequence of messages + /// * `tools` - Optional list of tools the model can use + /// + /// # Returns + /// A tuple containing the model's response message and provider usage statistics + /// + /// # Errors + /// ProviderError + /// - It's important to raise ContextLengthExceeded correctly since agent handles it + async fn complete_fast( + &self, + system: &str, + messages: &[Message], + tools: &[Tool], + ) -> Result<(Message, ProviderUsage), ProviderError>; + + /// This shouldn't be exposed externally, but each provider should implement it + /// so we can swap out the model for a fast model if configured + /// + /// # Arguments + /// * `system` - The system prompt that guides the model's behavior + /// * `messages` - The conversation history as a sequence of messages + /// * `tools` - Optional list of tools the model can use + async fn _complete_with_model( + &self, + system: &str, + messages: &[Message], + tools: &[Tool], + model: &str, + ) -> Result<(Message, ProviderUsage), ProviderError>; + + /// Get the model config from the provider fn get_model_config(&self) -> ModelConfig; From 060bbdd9c75bf444e1ab7c7be57aedacea064b62 Mon Sep 17 00:00:00 2001 From: David Katz Date: Wed, 20 Aug 2025 13:59:24 -0400 Subject: [PATCH 02/31] base impl --- crates/goose/src/model.rs | 1 + crates/goose/src/providers/base.rs | 28 +++------- .../goose/src/providers/formats/databricks.rs | 3 + crates/goose/src/providers/formats/openai.rs | 3 + crates/goose/src/providers/openai.rs | 55 ++++++++++++++----- 5 files changed, 58 insertions(+), 32 deletions(-) diff --git a/crates/goose/src/model.rs b/crates/goose/src/model.rs index ef38382953aa..89bf1167a9e9 100644 --- a/crates/goose/src/model.rs +++ b/crates/goose/src/model.rs @@ -102,6 +102,7 @@ impl ModelConfig { max_tokens: None, toolshim, toolshim_model, + fast_model: None, }) } diff --git a/crates/goose/src/providers/base.rs b/crates/goose/src/providers/base.rs index 65156143dcd7..ca159faf9553 100644 --- a/crates/goose/src/providers/base.rs +++ b/crates/goose/src/providers/base.rs @@ -337,7 +337,9 @@ pub trait Provider: Send + Sync { tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError>; - /// Generate the next message using the configured model and other parameters + /// Generate the next message using a fast/cheaper model when available + /// + /// Default implementation just calls regular complete() for providers that don't support fast models /// /// # Arguments /// * `system` - The system prompt that guides the model's behavior @@ -355,23 +357,11 @@ pub trait Provider: Send + Sync { system: &str, messages: &[Message], tools: &[Tool], - ) -> Result<(Message, ProviderUsage), ProviderError>; - - /// This shouldn't be exposed externally, but each provider should implement it - /// so we can swap out the model for a fast model if configured - /// - /// # Arguments - /// * `system` - The system prompt that guides the model's behavior - /// * `messages` - The conversation history as a sequence of messages - /// * `tools` - Optional list of tools the model can use - async fn _complete_with_model( - &self, - system: &str, - messages: &[Message], - tools: &[Tool], - model: &str, - ) -> Result<(Message, ProviderUsage), ProviderError>; - + ) -> Result<(Message, ProviderUsage), ProviderError> { + // Default implementation: just call regular complete + // Providers that support fast models should override this + self.complete(system, messages, tools).await + } /// Get the model config from the provider fn get_model_config(&self) -> ModelConfig; @@ -454,7 +444,7 @@ pub trait Provider: Send + Sync { let prompt = self.create_session_name_prompt(&context); let message = Message::user().with_text(&prompt); let result = self - .complete( + .complete_fast( "Reply with only a description in four words or less", &[message], &[], diff --git a/crates/goose/src/providers/formats/databricks.rs b/crates/goose/src/providers/formats/databricks.rs index 06f052b59c2a..b7472ca23464 100644 --- a/crates/goose/src/providers/formats/databricks.rs +++ b/crates/goose/src/providers/formats/databricks.rs @@ -1045,6 +1045,7 @@ mod tests { max_tokens: Some(1024), toolshim: false, toolshim_model: None, + fast_model: None, }; let request = create_request(&model_config, "system", &[], &[], &ImageFormat::OpenAi)?; let obj = request.as_object().unwrap(); @@ -1076,6 +1077,7 @@ mod tests { max_tokens: Some(1024), toolshim: false, toolshim_model: None, + fast_model: None, }; let request = create_request(&model_config, "system", &[], &[], &ImageFormat::OpenAi)?; let obj = request.as_object().unwrap(); @@ -1108,6 +1110,7 @@ mod tests { max_tokens: Some(1024), toolshim: false, toolshim_model: None, + fast_model: None, }; let request = create_request(&model_config, "system", &[], &[], &ImageFormat::OpenAi)?; let obj = request.as_object().unwrap(); diff --git a/crates/goose/src/providers/formats/openai.rs b/crates/goose/src/providers/formats/openai.rs index e5bdfdb08dfc..3ff4b712a867 100644 --- a/crates/goose/src/providers/formats/openai.rs +++ b/crates/goose/src/providers/formats/openai.rs @@ -1077,6 +1077,7 @@ mod tests { max_tokens: Some(1024), toolshim: false, toolshim_model: None, + fast_model: None, }; let request = create_request(&model_config, "system", &[], &[], &ImageFormat::OpenAi)?; let obj = request.as_object().unwrap(); @@ -1108,6 +1109,7 @@ mod tests { max_tokens: Some(1024), toolshim: false, toolshim_model: None, + fast_model: None, }; let request = create_request(&model_config, "system", &[], &[], &ImageFormat::OpenAi)?; let obj = request.as_object().unwrap(); @@ -1140,6 +1142,7 @@ mod tests { max_tokens: Some(1024), toolshim: false, toolshim_model: None, + fast_model: None, }; let request = create_request(&model_config, "system", &[], &[], &ImageFormat::OpenAi)?; let obj = request.as_object().unwrap(); diff --git a/crates/goose/src/providers/openai.rs b/crates/goose/src/providers/openai.rs index d49dcfeffbdf..5c5890338388 100644 --- a/crates/goose/src/providers/openai.rs +++ b/crates/goose/src/providers/openai.rs @@ -28,6 +28,8 @@ use crate::providers::base::MessageStream; use crate::providers::formats::openai::response_to_streaming_message; use rmcp::model::Tool; +const OPEN_AI_FAST_MODEL: &str = "gpt-4o-mini"; + pub const OPEN_AI_DEFAULT_MODEL: &str = "gpt-4o"; pub const OPEN_AI_KNOWN_MODELS: &[(&str, usize)] = &[ ("gpt-4o", 128_000), @@ -160,6 +162,31 @@ impl OpenAiProvider { .await?; handle_response_openai_compat(response).await } + + // Core completion logic that takes a model config + async fn complete_with_model( + &self, + model_config: &ModelConfig, + system: &str, + messages: &[Message], + tools: &[Tool], + ) -> Result<(Message, ProviderUsage), ProviderError> { + let payload = create_request(model_config, system, messages, tools, &ImageFormat::OpenAi)?; + + let json_response = self.post(&payload).await?; + + let message = response_to_message(&json_response)?; + let usage = json_response + .get("usage") + .map(get_usage) + .unwrap_or_else(|| { + tracing::debug!("Failed to get usage data"); + Usage::default() + }); + let model = get_model(&json_response); + emit_debug_trace(model_config, &payload, &json_response, &usage); + Ok((message, ProviderUsage::new(model, usage))) + } } #[async_trait] @@ -202,21 +229,23 @@ impl Provider for OpenAiProvider { messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - let payload = create_request(&self.model, system, messages, tools, &ImageFormat::OpenAi)?; + // Thin wrapper that calls complete_with_model with the configured model + self.complete_with_model(&self.model, system, messages, tools) + .await + } - let json_response = self.post(&payload).await?; + async fn complete_fast( + &self, + system: &str, + messages: &[Message], + tools: &[Tool], + ) -> Result<(Message, ProviderUsage), ProviderError> { + // Use the fast model (gpt-4o-mini) for fast completions + let mut fast_config = self.model.clone(); + fast_config.model_name = OPEN_AI_FAST_MODEL.to_string(); - let message = response_to_message(&json_response)?; - let usage = json_response - .get("usage") - .map(get_usage) - .unwrap_or_else(|| { - tracing::debug!("Failed to get usage data"); - Usage::default() - }); - let model = get_model(&json_response); - emit_debug_trace(&self.model, &payload, &json_response, &usage); - Ok((message, ProviderUsage::new(model, usage))) + self.complete_with_model(&fast_config, system, messages, tools) + .await } async fn fetch_supported_models(&self) -> Result>, ProviderError> { From 8b34e0296cacccc14af195ecc7ea76924a4292e0 Mon Sep 17 00:00:00 2001 From: David Katz Date: Wed, 20 Aug 2025 15:28:15 -0400 Subject: [PATCH 03/31] update new providers with api --- crates/goose/src/providers/anthropic.rs | 17 +++--- crates/goose/src/providers/azure.rs | 17 +++--- crates/goose/src/providers/base.rs | 42 +++++++++++++-- crates/goose/src/providers/bedrock.rs | 5 +- crates/goose/src/providers/claude_code.rs | 15 ++++-- crates/goose/src/providers/cursor_agent.rs | 15 ++++-- crates/goose/src/providers/databricks.rs | 24 ++++++--- crates/goose/src/providers/factory.rs | 2 +- crates/goose/src/providers/gcpvertexai.rs | 13 +++-- crates/goose/src/providers/gemini_cli.rs | 11 ++-- crates/goose/src/providers/githubcopilot.rs | 19 ++++--- crates/goose/src/providers/google.rs | 21 +++++--- crates/goose/src/providers/groq.rs | 17 +++--- crates/goose/src/providers/lead_worker.rs | 7 +-- crates/goose/src/providers/litellm.rs | 15 ++++-- crates/goose/src/providers/ollama.rs | 15 ++++-- crates/goose/src/providers/openai.rs | 57 +++++++++------------ crates/goose/src/providers/openrouter.rs | 15 ++++-- crates/goose/src/providers/sagemaker_tgi.rs | 9 +++- crates/goose/src/providers/snowflake.rs | 19 ++++--- crates/goose/src/providers/testprovider.rs | 5 +- crates/goose/src/providers/venice.rs | 15 ++++-- crates/goose/src/providers/xai.rs | 17 +++--- 23 files changed, 257 insertions(+), 135 deletions(-) diff --git a/crates/goose/src/providers/anthropic.rs b/crates/goose/src/providers/anthropic.rs index 952006ee25de..747a28ac662f 100644 --- a/crates/goose/src/providers/anthropic.rs +++ b/crates/goose/src/providers/anthropic.rs @@ -179,16 +179,21 @@ impl Provider for AnthropicProvider { } #[tracing::instrument( - skip(self, system, messages, tools), + skip(self, model, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] - async fn complete( + async fn complete_with_model( &self, + model: &str, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - let payload = create_request(&self.model, system, messages, tools)?; + // Create a temporary model config with the specified model + let mut model_config = self.model.clone(); + model_config.model_name = model.to_string(); + + let payload = create_request(&model_config, system, messages, tools)?; let response = self .with_retry(|| async { self.post(&payload).await }) @@ -201,9 +206,9 @@ impl Provider for AnthropicProvider { tracing::debug!("🔍 Anthropic non-streaming parsed usage: input_tokens={:?}, output_tokens={:?}, total_tokens={:?}", usage.input_tokens, usage.output_tokens, usage.total_tokens); - let model = get_model(&json_response); - emit_debug_trace(&self.model, &payload, &json_response, &usage); - let provider_usage = ProviderUsage::new(model, usage); + let response_model = get_model(&json_response); + emit_debug_trace(&model_config, &payload, &json_response, &usage); + let provider_usage = ProviderUsage::new(response_model, usage); tracing::debug!( "🔍 Anthropic non-streaming returning ProviderUsage: {:?}", provider_usage diff --git a/crates/goose/src/providers/azure.rs b/crates/goose/src/providers/azure.rs index f40993d67657..1255ca38dbc4 100644 --- a/crates/goose/src/providers/azure.rs +++ b/crates/goose/src/providers/azure.rs @@ -135,16 +135,21 @@ impl Provider for AzureProvider { } #[tracing::instrument( - skip(self, system, messages, tools), + skip(self, model, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] - async fn complete( + async fn complete_with_model( &self, + model: &str, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - let payload = create_request(&self.model, system, messages, tools, &ImageFormat::OpenAi)?; + // Create a temporary model config with the specified model + let mut model_config = self.model.clone(); + model_config.model_name = model.to_string(); + + let payload = create_request(&model_config, system, messages, tools, &ImageFormat::OpenAi)?; let response = self .with_retry(|| async { let payload_clone = payload.clone(); @@ -157,8 +162,8 @@ impl Provider for AzureProvider { tracing::debug!("Failed to get usage data"); Usage::default() }); - let model = get_model(&response); - emit_debug_trace(&self.model, &payload, &response, &usage); - Ok((message, ProviderUsage::new(model, usage))) + let response_model = get_model(&response); + emit_debug_trace(&model_config, &payload, &response, &usage); + Ok((message, ProviderUsage::new(response_model, usage))) } } diff --git a/crates/goose/src/providers/base.rs b/crates/goose/src/providers/base.rs index ca159faf9553..33eefaea1fc6 100644 --- a/crates/goose/src/providers/base.rs +++ b/crates/goose/src/providers/base.rs @@ -317,6 +317,29 @@ pub trait Provider: Send + Sync { where Self: Sized; + /// Internal method that performs completion with a specific model + /// This is where providers implement their actual completion logic + /// + /// # Arguments + /// * `model` - The model name to use + /// * `system` - The system prompt that guides the model's behavior + /// * `messages` - The conversation history as a sequence of messages + /// * `tools` - Optional list of tools the model can use + /// + /// # Returns + /// A tuple containing the model's response message and provider usage statistics + /// + /// # Errors + /// ProviderError + /// - It's important to raise ContextLengthExceeded correctly since agent handles it + async fn complete_with_model( + &self, + model: &str, + system: &str, + messages: &[Message], + tools: &[Tool], + ) -> Result<(Message, ProviderUsage), ProviderError>; + /// Generate the next message using the configured model and other parameters /// /// # Arguments @@ -335,7 +358,12 @@ pub trait Provider: Send + Sync { system: &str, messages: &[Message], tools: &[Tool], - ) -> Result<(Message, ProviderUsage), ProviderError>; + ) -> Result<(Message, ProviderUsage), ProviderError> { + // Default implementation: use the provider's configured model + let model_config = self.get_model_config(); + self.complete_with_model(&model_config.model_name, system, messages, tools) + .await + } /// Generate the next message using a fast/cheaper model when available /// @@ -358,9 +386,15 @@ pub trait Provider: Send + Sync { messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - // Default implementation: just call regular complete - // Providers that support fast models should override this - self.complete(system, messages, tools).await + // Check if a fast model is configured, otherwise fall back to regular model + let model_config = self.get_model_config(); + let model = model_config + .fast_model + .as_deref() + .unwrap_or(&model_config.model_name); + + self.complete_with_model(model, system, messages, tools) + .await } /// Get the model config from the provider diff --git a/crates/goose/src/providers/bedrock.rs b/crates/goose/src/providers/bedrock.rs index 7579a7de7141..a5864c7d577d 100644 --- a/crates/goose/src/providers/bedrock.rs +++ b/crates/goose/src/providers/bedrock.rs @@ -152,11 +152,12 @@ impl Provider for BedrockProvider { } #[tracing::instrument( - skip(self, system, messages, tools), + skip(self, model, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] - async fn complete( + async fn complete_with_model( &self, + model: &str, system: &str, messages: &[Message], tools: &[Tool], diff --git a/crates/goose/src/providers/claude_code.rs b/crates/goose/src/providers/claude_code.rs index 3185a7961bc2..b2b024dcc9a9 100644 --- a/crates/goose/src/providers/claude_code.rs +++ b/crates/goose/src/providers/claude_code.rs @@ -474,15 +474,20 @@ impl Provider for ClaudeCodeProvider { } #[tracing::instrument( - skip(self, system, messages, tools), + skip(self, model, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] - async fn complete( + async fn complete_with_model( &self, + model: &str, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { + // Create a temporary model config with the specified model + let mut model_config = self.model.clone(); + model_config.model_name = model.to_string(); + // Check if this is a session description request (short system prompt asking for 4 words or less) if system.contains("four words or less") || system.contains("4 words or less") { return self.generate_simple_session_description(messages); @@ -495,7 +500,7 @@ impl Provider for ClaudeCodeProvider { // Create a dummy payload for debug tracing let payload = json!({ "command": self.command, - "model": self.model.model_name, + "model": model_config.model_name, "system": system, "messages": messages.len() }); @@ -505,11 +510,11 @@ impl Provider for ClaudeCodeProvider { "usage": usage }); - emit_debug_trace(&self.model, &payload, &response, &usage); + emit_debug_trace(&model_config, &payload, &response, &usage); Ok(( message, - ProviderUsage::new(self.model.model_name.clone(), usage), + ProviderUsage::new(model_config.model_name.clone(), usage), )) } } diff --git a/crates/goose/src/providers/cursor_agent.rs b/crates/goose/src/providers/cursor_agent.rs index 432093df0b51..dc393c24a6c6 100644 --- a/crates/goose/src/providers/cursor_agent.rs +++ b/crates/goose/src/providers/cursor_agent.rs @@ -407,15 +407,20 @@ impl Provider for CursorAgentProvider { } #[tracing::instrument( - skip(self, system, messages, tools), + skip(self, model, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] - async fn complete( + async fn complete_with_model( &self, + model: &str, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { + // Create a temporary model config with the specified model + let mut model_config = self.model.clone(); + model_config.model_name = model.to_string(); + // Check if this is a session description request (short system prompt asking for 4 words or less) if system.contains("four words or less") || system.contains("4 words or less") { return self.generate_simple_session_description(messages); @@ -428,7 +433,7 @@ impl Provider for CursorAgentProvider { // Create a dummy payload for debug tracing let payload = json!({ "command": self.command, - "model": self.model.model_name, + "model": model_config.model_name, "system": system, "messages": messages.len() }); @@ -438,11 +443,11 @@ impl Provider for CursorAgentProvider { "usage": usage }); - emit_debug_trace(&self.model, &payload, &response, &usage); + emit_debug_trace(&model_config, &payload, &response, &usage); Ok(( message, - ProviderUsage::new(self.model.model_name.clone(), usage), + ProviderUsage::new(model_config.model_name.clone(), usage), )) } } diff --git a/crates/goose/src/providers/databricks.rs b/crates/goose/src/providers/databricks.rs index c635fe589470..5e33d3582374 100644 --- a/crates/goose/src/providers/databricks.rs +++ b/crates/goose/src/providers/databricks.rs @@ -238,16 +238,22 @@ impl Provider for DatabricksProvider { } #[tracing::instrument( - skip(self, system, messages, tools), + skip(self, model, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] - async fn complete( + async fn complete_with_model( &self, + model: &str, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - let mut payload = create_request(&self.model, system, messages, tools, &self.image_format)?; + // Create a temporary model config with the specified model + let mut model_config = self.model.clone(); + model_config.model_name = model.to_string(); + + let mut payload = + create_request(&model_config, system, messages, tools, &self.image_format)?; payload .as_object_mut() .expect("payload should have model key") @@ -260,10 +266,10 @@ impl Provider for DatabricksProvider { tracing::debug!("Failed to get usage data"); Usage::default() }); - let model = get_model(&response); - super::utils::emit_debug_trace(&self.model, &payload, &response, &usage); + let response_model = get_model(&response); + super::utils::emit_debug_trace(&model_config, &payload, &response, &usage); - Ok((message, ProviderUsage::new(model, usage))) + Ok((message, ProviderUsage::new(response_model, usage))) } async fn stream( @@ -272,7 +278,11 @@ impl Provider for DatabricksProvider { messages: &[Message], tools: &[Tool], ) -> Result { - let mut payload = create_request(&self.model, system, messages, tools, &self.image_format)?; + // Create a temporary model config with the specified model + let model_config = self.model.clone(); + + let mut payload = + create_request(&model_config, system, messages, tools, &self.image_format)?; payload .as_object_mut() .expect("payload should have model key") diff --git a/crates/goose/src/providers/factory.rs b/crates/goose/src/providers/factory.rs index fbcab89b8329..7b43806b1ed0 100644 --- a/crates/goose/src/providers/factory.rs +++ b/crates/goose/src/providers/factory.rs @@ -200,7 +200,7 @@ mod tests { self.model_config.clone() } - async fn complete( + async fn complete_with_model( &self, _system: &str, _messages: &[Message], diff --git a/crates/goose/src/providers/gcpvertexai.rs b/crates/goose/src/providers/gcpvertexai.rs index 969d7146d7e2..23925a65d70f 100644 --- a/crates/goose/src/providers/gcpvertexai.rs +++ b/crates/goose/src/providers/gcpvertexai.rs @@ -512,23 +512,28 @@ impl Provider for GcpVertexAIProvider { /// * `messages` - Array of previous messages in the conversation /// * `tools` - Array of available tools for the model #[tracing::instrument( - skip(self, system, messages, tools), + skip(self, model, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] - async fn complete( + async fn complete_with_model( &self, + model: &str, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { + // Create a temporary model config with the specified model + let mut model_config = self.model.clone(); + model_config.model_name = model.to_string(); + // Create request and context - let (request, context) = create_request(&self.model, system, messages, tools)?; + let (request, context) = create_request(&model_config, system, messages, tools)?; // Send request and process response let response = self.post(&request, &context).await?; let usage = get_usage(&response, &context)?; - emit_debug_trace(&self.model, &request, &response, &usage); + emit_debug_trace(&model_config, &request, &response, &usage); // Convert response to message let message = response_to_message(response, context)?; diff --git a/crates/goose/src/providers/gemini_cli.rs b/crates/goose/src/providers/gemini_cli.rs index fcfc0f75c369..824f0f7fa5cc 100644 --- a/crates/goose/src/providers/gemini_cli.rs +++ b/crates/goose/src/providers/gemini_cli.rs @@ -319,15 +319,20 @@ impl Provider for GeminiCliProvider { } #[tracing::instrument( - skip(self, system, messages, tools), + skip(self, model, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] - async fn complete( + async fn complete_with_model( &self, + model: &str, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { + // Create a temporary model config with the specified model + let mut model_config = self.model.clone(); + model_config.model_name = model.to_string(); + // Check if this is a session description request (short system prompt asking for 4 words or less) if system.contains("four words or less") || system.contains("4 words or less") { return self.generate_simple_session_description(messages); @@ -350,7 +355,7 @@ impl Provider for GeminiCliProvider { "usage": usage }); - emit_debug_trace(&self.model, &payload, &response, &usage); + emit_debug_trace(&model_config, &payload, &response, &usage); Ok(( message, diff --git a/crates/goose/src/providers/githubcopilot.rs b/crates/goose/src/providers/githubcopilot.rs index fc1e7fb640dd..91a32fd647d7 100644 --- a/crates/goose/src/providers/githubcopilot.rs +++ b/crates/goose/src/providers/githubcopilot.rs @@ -401,16 +401,23 @@ impl Provider for GithubCopilotProvider { } #[tracing::instrument( - skip(self, system, messages, tools), + skip(self, model, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] - async fn complete( + async fn complete_with_model( &self, + model: &str, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - let payload = create_request(&self.model, system, messages, tools, &ImageFormat::OpenAi)?; + // Create a temporary model config with the specified model + + let mut model_config = self.model.clone(); + + model_config.model_name = model.to_string(); + + let payload = create_request(&model_config, system, messages, tools, &ImageFormat::OpenAi)?; // Make request with retry let response = self @@ -426,9 +433,9 @@ impl Provider for GithubCopilotProvider { tracing::debug!("Failed to get usage data"); Usage::default() }); - let model = get_model(&response); - emit_debug_trace(&self.model, &payload, &response, &usage); - Ok((message, ProviderUsage::new(model, usage))) + let response_model = get_model(&response); + emit_debug_trace(&model_config, &payload, &response, &usage); + Ok((message, ProviderUsage::new(response_model, usage))) } /// Fetch supported models from GitHub Copliot; returns Err on failure, Ok(None) if not present diff --git a/crates/goose/src/providers/google.rs b/crates/goose/src/providers/google.rs index fa262f403c3a..8b6f92e6f946 100644 --- a/crates/goose/src/providers/google.rs +++ b/crates/goose/src/providers/google.rs @@ -101,16 +101,23 @@ impl Provider for GoogleProvider { } #[tracing::instrument( - skip(self, system, messages, tools), + skip(self, model, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] - async fn complete( + async fn complete_with_model( &self, + model: &str, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - let payload = create_request(&self.model, system, messages, tools)?; + // Create a temporary model config with the specified model + + let mut model_config = self.model.clone(); + + model_config.model_name = model.to_string(); + + let payload = create_request(&model_config, system, messages, tools)?; // Make request let response = self @@ -123,12 +130,12 @@ impl Provider for GoogleProvider { // Parse response let message = response_to_message(unescape_json_values(&response))?; let usage = get_usage(&response)?; - let model = match response.get("modelVersion") { + let response_model = match response.get("modelVersion") { Some(model_version) => model_version.as_str().unwrap_or_default().to_string(), - None => self.model.model_name.clone(), + None => model_config.model_name.clone(), }; - emit_debug_trace(&self.model, &payload, &response, &usage); - let provider_usage = ProviderUsage::new(model, usage); + emit_debug_trace(&model_config, &payload, &response, &usage); + let provider_usage = ProviderUsage::new(response_model, usage); Ok((message, provider_usage)) } diff --git a/crates/goose/src/providers/groq.rs b/crates/goose/src/providers/groq.rs index acb51d1fc75d..55f69715311c 100644 --- a/crates/goose/src/providers/groq.rs +++ b/crates/goose/src/providers/groq.rs @@ -77,17 +77,22 @@ impl Provider for GroqProvider { } #[tracing::instrument( - skip(self, system, messages, tools), + skip(self, model, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] - async fn complete( + async fn complete_with_model( &self, + model: &str, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { + // Create a temporary model config with the specified model + let mut model_config = self.model.clone(); + model_config.model_name = model.to_string(); + let payload = create_request( - &self.model, + &model_config, system, messages, tools, @@ -101,9 +106,9 @@ impl Provider for GroqProvider { tracing::debug!("Failed to get usage data"); Usage::default() }); - let model = get_model(&response); - super::utils::emit_debug_trace(&self.model, &payload, &response, &usage); - Ok((message, ProviderUsage::new(model, usage))) + let response_model = get_model(&response); + super::utils::emit_debug_trace(&model_config, &payload, &response, &usage); + Ok((message, ProviderUsage::new(response_model, usage))) } /// Fetch supported models from Groq; returns Err on failure, Ok(None) if no models found diff --git a/crates/goose/src/providers/lead_worker.rs b/crates/goose/src/providers/lead_worker.rs index 18564a2d1261..93d39593a2e7 100644 --- a/crates/goose/src/providers/lead_worker.rs +++ b/crates/goose/src/providers/lead_worker.rs @@ -326,8 +326,9 @@ impl Provider for LeadWorkerProvider { self.lead_provider.get_model_config() } - async fn complete( + async fn complete_with_model( &self, + model: &str, system: &str, messages: &[Message], tools: &[Tool], @@ -475,7 +476,7 @@ mod tests { self.model_config.clone() } - async fn complete( + async fn complete_with_model( &self, _system: &str, _messages: &[Message], @@ -635,7 +636,7 @@ mod tests { self.model_config.clone() } - async fn complete( + async fn complete_with_model( &self, _system: &str, _messages: &[Message], diff --git a/crates/goose/src/providers/litellm.rs b/crates/goose/src/providers/litellm.rs index 8911341b4d42..500ab4994543 100644 --- a/crates/goose/src/providers/litellm.rs +++ b/crates/goose/src/providers/litellm.rs @@ -161,14 +161,19 @@ impl Provider for LiteLLMProvider { } #[tracing::instrument(skip_all, name = "provider_complete")] - async fn complete( + async fn complete_with_model( &self, + model: &str, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { + // Create a temporary model config with the specified model + let mut model_config = self.model.clone(); + model_config.model_name = model.to_string(); + let mut payload = super::formats::openai::create_request( - &self.model, + &model_config, system, messages, tools, @@ -188,9 +193,9 @@ impl Provider for LiteLLMProvider { let message = super::formats::openai::response_to_message(&response)?; let usage = super::formats::openai::get_usage(&response); - let model = get_model(&response); - emit_debug_trace(&self.model, &payload, &response, &usage); - Ok((message, ProviderUsage::new(model, usage))) + let response_model = get_model(&response); + emit_debug_trace(&model_config, &payload, &response, &usage); + Ok((message, ProviderUsage::new(response_model, usage))) } fn supports_embeddings(&self) -> bool { diff --git a/crates/goose/src/providers/ollama.rs b/crates/goose/src/providers/ollama.rs index 84b6a5e5e184..2fc80f9daaa3 100644 --- a/crates/goose/src/providers/ollama.rs +++ b/crates/goose/src/providers/ollama.rs @@ -165,15 +165,20 @@ impl Provider for OllamaProvider { } #[tracing::instrument( - skip(self, system, messages, tools), + skip(self, model, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] - async fn complete( + async fn complete_with_model( &self, + model: &str, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { + // Create a temporary model config with the specified model + let mut model_config = self.model.clone(); + model_config.model_name = model.to_string(); + let config = crate::config::Config::global(); let goose_mode = config.get_param("GOOSE_MODE").unwrap_or("auto".to_string()); let filtered_tools = if goose_mode == "chat" { &[] } else { tools }; @@ -197,9 +202,9 @@ impl Provider for OllamaProvider { tracing::debug!("Failed to get usage data"); Usage::default() }); - let model = get_model(&response); - super::utils::emit_debug_trace(&self.model, &payload, &response, &usage); - Ok((message, ProviderUsage::new(model, usage))) + let response_model = get_model(&response); + super::utils::emit_debug_trace(&model_config, &payload, &response, &usage); + Ok((message, ProviderUsage::new(response_model, usage))) } /// Generate a session name based on the conversation history diff --git a/crates/goose/src/providers/openai.rs b/crates/goose/src/providers/openai.rs index 5c5890338388..01f8447402b0 100644 --- a/crates/goose/src/providers/openai.rs +++ b/crates/goose/src/providers/openai.rs @@ -162,31 +162,6 @@ impl OpenAiProvider { .await?; handle_response_openai_compat(response).await } - - // Core completion logic that takes a model config - async fn complete_with_model( - &self, - model_config: &ModelConfig, - system: &str, - messages: &[Message], - tools: &[Tool], - ) -> Result<(Message, ProviderUsage), ProviderError> { - let payload = create_request(model_config, system, messages, tools, &ImageFormat::OpenAi)?; - - let json_response = self.post(&payload).await?; - - let message = response_to_message(&json_response)?; - let usage = json_response - .get("usage") - .map(get_usage) - .unwrap_or_else(|| { - tracing::debug!("Failed to get usage data"); - Usage::default() - }); - let model = get_model(&json_response); - emit_debug_trace(model_config, &payload, &json_response, &usage); - Ok((message, ProviderUsage::new(model, usage))) - } } #[async_trait] @@ -220,18 +195,35 @@ impl Provider for OpenAiProvider { } #[tracing::instrument( - skip(self, system, messages, tools), + skip(self, model, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] - async fn complete( + async fn complete_with_model( &self, + model: &str, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - // Thin wrapper that calls complete_with_model with the configured model - self.complete_with_model(&self.model, system, messages, tools) - .await + // Create a temporary model config with the specified model + let mut model_config = self.model.clone(); + model_config.model_name = model.to_string(); + + let payload = create_request(&model_config, system, messages, tools, &ImageFormat::OpenAi)?; + + let json_response = self.post(&payload).await?; + + let message = response_to_message(&json_response)?; + let usage = json_response + .get("usage") + .map(get_usage) + .unwrap_or_else(|| { + tracing::debug!("Failed to get usage data"); + Usage::default() + }); + let response_model = get_model(&json_response); + emit_debug_trace(&model_config, &payload, &json_response, &usage); + Ok((message, ProviderUsage::new(response_model, usage))) } async fn complete_fast( @@ -241,10 +233,7 @@ impl Provider for OpenAiProvider { tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { // Use the fast model (gpt-4o-mini) for fast completions - let mut fast_config = self.model.clone(); - fast_config.model_name = OPEN_AI_FAST_MODEL.to_string(); - - self.complete_with_model(&fast_config, system, messages, tools) + self.complete_with_model(OPEN_AI_FAST_MODEL, system, messages, tools) .await } diff --git a/crates/goose/src/providers/openrouter.rs b/crates/goose/src/providers/openrouter.rs index 00fa77bdea7a..c8ec0258d96b 100644 --- a/crates/goose/src/providers/openrouter.rs +++ b/crates/goose/src/providers/openrouter.rs @@ -238,15 +238,20 @@ impl Provider for OpenRouterProvider { } #[tracing::instrument( - skip(self, system, messages, tools), + skip(self, model, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] - async fn complete( + async fn complete_with_model( &self, + model: &str, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { + // Create a temporary model config with the specified model + let mut model_config = self.model.clone(); + model_config.model_name = model.to_string(); + // Create the base payload let payload = create_request_based_on_model(self, system, messages, tools)?; @@ -264,9 +269,9 @@ impl Provider for OpenRouterProvider { tracing::debug!("Failed to get usage data"); Usage::default() }); - let model = get_model(&response); - emit_debug_trace(&self.model, &payload, &response, &usage); - Ok((message, ProviderUsage::new(model, usage))) + let response_model = get_model(&response); + emit_debug_trace(&model_config, &payload, &response, &usage); + Ok((message, ProviderUsage::new(response_model, usage))) } /// Fetch supported models from OpenRouter API (only models with tool support) diff --git a/crates/goose/src/providers/sagemaker_tgi.rs b/crates/goose/src/providers/sagemaker_tgi.rs index 90a2498d0977..8d0180fbc760 100644 --- a/crates/goose/src/providers/sagemaker_tgi.rs +++ b/crates/goose/src/providers/sagemaker_tgi.rs @@ -280,15 +280,20 @@ impl Provider for SageMakerTgiProvider { } #[tracing::instrument( - skip(self, system, messages, tools), + skip(self, model, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] - async fn complete( + async fn complete_with_model( &self, + model: &str, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { + // Create a temporary model config with the specified model + let mut model_config = self.model.clone(); + model_config.model_name = model.to_string(); + let model_name = &self.model.model_name; let request_payload = self.create_tgi_request(system, messages).map_err(|e| { diff --git a/crates/goose/src/providers/snowflake.rs b/crates/goose/src/providers/snowflake.rs index 8e8ea663b5ae..6089ea744149 100644 --- a/crates/goose/src/providers/snowflake.rs +++ b/crates/goose/src/providers/snowflake.rs @@ -299,16 +299,23 @@ impl Provider for SnowflakeProvider { } #[tracing::instrument( - skip(self, system, messages, tools), + skip(self, model, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] - async fn complete( + async fn complete_with_model( &self, + model: &str, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - let payload = create_request(&self.model, system, messages, tools)?; + // Create a temporary model config with the specified model + + let mut model_config = self.model.clone(); + + model_config.model_name = model.to_string(); + + let payload = create_request(&model_config, system, messages, tools)?; let response = self .with_retry(|| async { @@ -320,9 +327,9 @@ impl Provider for SnowflakeProvider { // Parse response let message = response_to_message(&response)?; let usage = get_usage(&response)?; - let model = get_model(&response); - super::utils::emit_debug_trace(&self.model, &payload, &response, &usage); + let response_model = get_model(&response); + super::utils::emit_debug_trace(&model_config, &payload, &response, &usage); - Ok((message, ProviderUsage::new(model, usage))) + Ok((message, ProviderUsage::new(response_model, usage))) } } diff --git a/crates/goose/src/providers/testprovider.rs b/crates/goose/src/providers/testprovider.rs index eca2c87627b1..fb225e061cc7 100644 --- a/crates/goose/src/providers/testprovider.rs +++ b/crates/goose/src/providers/testprovider.rs @@ -112,8 +112,9 @@ impl Provider for TestProvider { ) } - async fn complete( + async fn complete_with_model( &self, + model: &str, system: &str, messages: &[Message], tools: &[Tool], @@ -188,7 +189,7 @@ mod tests { ) } - async fn complete( + async fn complete_with_model( &self, _system: &str, _messages: &[Message], diff --git a/crates/goose/src/providers/venice.rs b/crates/goose/src/providers/venice.rs index 185587c6df6c..3b76ebaaffb0 100644 --- a/crates/goose/src/providers/venice.rs +++ b/crates/goose/src/providers/venice.rs @@ -246,23 +246,28 @@ impl Provider for VeniceProvider { } #[tracing::instrument( - skip(_system, messages, tools), + skip(self, model, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] - async fn complete( + async fn complete_with_model( &self, - _system: &str, + model: &str, + system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { + // Create a temporary model config with the specified model + let mut model_config = self.model.clone(); + model_config.model_name = model.to_string(); + // Create properly formatted messages for Venice API let mut formatted_messages = Vec::new(); // Add the system message if present - if !_system.is_empty() { + if !system.is_empty() { formatted_messages.push(json!({ "role": "system", - "content": _system + "content": system })); } diff --git a/crates/goose/src/providers/xai.rs b/crates/goose/src/providers/xai.rs index 7b2aed5f15c8..b8adcb567d76 100644 --- a/crates/goose/src/providers/xai.rs +++ b/crates/goose/src/providers/xai.rs @@ -93,17 +93,22 @@ impl Provider for XaiProvider { } #[tracing::instrument( - skip(self, system, messages, tools), + skip(self, model, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] - async fn complete( + async fn complete_with_model( &self, + model: &str, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { + // Create a temporary model config with the specified model + let mut model_config = self.model.clone(); + model_config.model_name = model.to_string(); + let payload = create_request( - &self.model, + &model_config, system, messages, tools, @@ -117,8 +122,8 @@ impl Provider for XaiProvider { tracing::debug!("Failed to get usage data"); Usage::default() }); - let model = get_model(&response); - super::utils::emit_debug_trace(&self.model, &payload, &response, &usage); - Ok((message, ProviderUsage::new(model, usage))) + let response_model = get_model(&response); + super::utils::emit_debug_trace(&model_config, &payload, &response, &usage); + Ok((message, ProviderUsage::new(response_model, usage))) } } From 966e3919475e1145d81f5c8740157fd9c7331e07 Mon Sep 17 00:00:00 2001 From: David Katz Date: Wed, 20 Aug 2025 15:47:36 -0400 Subject: [PATCH 04/31] add provider defaults --- crates/goose/src/providers/anthropic.rs | 13 +++++++++-- crates/goose/src/providers/bedrock.rs | 4 ++-- crates/goose/src/providers/databricks.rs | 10 +++++++-- crates/goose/src/providers/google.rs | 11 +++++---- crates/goose/src/providers/groq.rs | 5 ++++- crates/goose/src/providers/lead_worker.rs | 2 +- crates/goose/src/providers/openai.rs | 26 +++++++++------------- crates/goose/src/providers/testprovider.rs | 2 +- 8 files changed, 45 insertions(+), 28 deletions(-) diff --git a/crates/goose/src/providers/anthropic.rs b/crates/goose/src/providers/anthropic.rs index 747a28ac662f..57f1ea33c805 100644 --- a/crates/goose/src/providers/anthropic.rs +++ b/crates/goose/src/providers/anthropic.rs @@ -49,7 +49,10 @@ pub struct AnthropicProvider { impl_provider_default!(AnthropicProvider); impl AnthropicProvider { - pub fn from_env(model: ModelConfig) -> Result { + pub fn from_env(mut model: ModelConfig) -> Result { + // Set the default fast model for Anthropic + model.fast_model = Some("claude-3-5-haiku-latest".to_string()); + let config = crate::config::Config::global(); let api_key: String = config.get_secret("ANTHROPIC_API_KEY")?; let host: String = config @@ -71,7 +74,13 @@ impl AnthropicProvider { }) } - pub fn from_custom_config(model: ModelConfig, config: CustomProviderConfig) -> Result { + pub fn from_custom_config( + mut model: ModelConfig, + config: CustomProviderConfig, + ) -> Result { + // Set the default fast model for Anthropic + model.fast_model = Some("claude-3-5-haiku-latest".to_string()); + let global_config = crate::config::Config::global(); let api_key: String = global_config .get_secret(&config.api_key_env) diff --git a/crates/goose/src/providers/bedrock.rs b/crates/goose/src/providers/bedrock.rs index a5864c7d577d..57961c789885 100644 --- a/crates/goose/src/providers/bedrock.rs +++ b/crates/goose/src/providers/bedrock.rs @@ -152,12 +152,12 @@ impl Provider for BedrockProvider { } #[tracing::instrument( - skip(self, model, system, messages, tools), + skip(self, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] async fn complete_with_model( &self, - model: &str, + _model: &str, system: &str, messages: &[Message], tools: &[Tool], diff --git a/crates/goose/src/providers/databricks.rs b/crates/goose/src/providers/databricks.rs index 5e33d3582374..de0326ec25a7 100644 --- a/crates/goose/src/providers/databricks.rs +++ b/crates/goose/src/providers/databricks.rs @@ -107,7 +107,10 @@ pub struct DatabricksProvider { impl_provider_default!(DatabricksProvider); impl DatabricksProvider { - pub fn from_env(model: ModelConfig) -> Result { + pub fn from_env(mut model: ModelConfig) -> Result { + // Set the default fast model for Databricks - using a smaller/faster model + model.fast_model = Some("databricks-mixtral-8x7b-instruct".to_string()); + let config = crate::config::Config::global(); let mut host: Result = config.get_param("DATABRICKS_HOST"); @@ -179,7 +182,10 @@ impl DatabricksProvider { } } - pub fn from_params(host: String, api_key: String, model: ModelConfig) -> Result { + pub fn from_params(host: String, api_key: String, mut model: ModelConfig) -> Result { + // Set the default fast model for Databricks + model.fast_model = Some("databricks-mixtral-8x7b-instruct".to_string()); + let auth = DatabricksAuth::token(api_key); let auth_method = AuthMethod::Custom(Box::new(DatabricksAuthProvider { auth: auth.clone() })); diff --git a/crates/goose/src/providers/google.rs b/crates/goose/src/providers/google.rs index 8b6f92e6f946..64547e87c39d 100644 --- a/crates/goose/src/providers/google.rs +++ b/crates/goose/src/providers/google.rs @@ -54,7 +54,10 @@ pub struct GoogleProvider { impl_provider_default!(GoogleProvider); impl GoogleProvider { - pub fn from_env(model: ModelConfig) -> Result { + pub fn from_env(mut model: ModelConfig) -> Result { + // Set the default fast model for Google - using Gemini Flash + model.fast_model = Some("gemini-1.5-flash".to_string()); + let config = crate::config::Config::global(); let api_key: String = config.get_secret("GOOGLE_API_KEY")?; let host: String = config @@ -72,8 +75,8 @@ impl GoogleProvider { Ok(Self { api_client, model }) } - async fn post(&self, payload: &Value) -> Result { - let path = format!("v1beta/models/{}:generateContent", self.model.model_name); + async fn post(&self, model_name: &str, payload: &Value) -> Result { + let path = format!("v1beta/models/{}:generateContent", model_name); let response = self.api_client.response_post(&path, payload).await?; handle_response_google_compat(response).await } @@ -123,7 +126,7 @@ impl Provider for GoogleProvider { let response = self .with_retry(|| async { let payload_clone = payload.clone(); - self.post(&payload_clone).await + self.post(model, &payload_clone).await }) .await?; diff --git a/crates/goose/src/providers/groq.rs b/crates/goose/src/providers/groq.rs index 55f69715311c..7de887b813f6 100644 --- a/crates/goose/src/providers/groq.rs +++ b/crates/goose/src/providers/groq.rs @@ -33,7 +33,10 @@ pub struct GroqProvider { impl_provider_default!(GroqProvider); impl GroqProvider { - pub fn from_env(model: ModelConfig) -> Result { + pub fn from_env(mut model: ModelConfig) -> Result { + // Set the default fast model for Groq - using a smaller/faster model + model.fast_model = Some("gemma2-9b-it".to_string()); + let config = crate::config::Config::global(); let api_key: String = config.get_secret("GROQ_API_KEY")?; let host: String = config diff --git a/crates/goose/src/providers/lead_worker.rs b/crates/goose/src/providers/lead_worker.rs index 93d39593a2e7..c15360cfc076 100644 --- a/crates/goose/src/providers/lead_worker.rs +++ b/crates/goose/src/providers/lead_worker.rs @@ -328,7 +328,7 @@ impl Provider for LeadWorkerProvider { async fn complete_with_model( &self, - model: &str, + _model: &str, system: &str, messages: &[Message], tools: &[Tool], diff --git a/crates/goose/src/providers/openai.rs b/crates/goose/src/providers/openai.rs index 01f8447402b0..fe665eb20ecd 100644 --- a/crates/goose/src/providers/openai.rs +++ b/crates/goose/src/providers/openai.rs @@ -28,8 +28,6 @@ use crate::providers::base::MessageStream; use crate::providers::formats::openai::response_to_streaming_message; use rmcp::model::Tool; -const OPEN_AI_FAST_MODEL: &str = "gpt-4o-mini"; - pub const OPEN_AI_DEFAULT_MODEL: &str = "gpt-4o"; pub const OPEN_AI_KNOWN_MODELS: &[(&str, usize)] = &[ ("gpt-4o", 128_000), @@ -60,7 +58,10 @@ pub struct OpenAiProvider { impl_provider_default!(OpenAiProvider); impl OpenAiProvider { - pub fn from_env(model: ModelConfig) -> Result { + pub fn from_env(mut model: ModelConfig) -> Result { + // Set the default fast model for OpenAI + model.fast_model = Some("gpt-4o-mini".to_string()); + let config = crate::config::Config::global(); let api_key: String = config.get_secret("OPENAI_API_KEY")?; let host: String = config @@ -111,7 +112,13 @@ impl OpenAiProvider { }) } - pub fn from_custom_config(model: ModelConfig, config: CustomProviderConfig) -> Result { + pub fn from_custom_config( + mut model: ModelConfig, + config: CustomProviderConfig, + ) -> Result { + // Set the default fast model for OpenAI + model.fast_model = Some("gpt-4o-mini".to_string()); + let global_config = crate::config::Config::global(); let api_key: String = global_config .get_secret(&config.api_key_env) @@ -226,17 +233,6 @@ impl Provider for OpenAiProvider { Ok((message, ProviderUsage::new(response_model, usage))) } - async fn complete_fast( - &self, - system: &str, - messages: &[Message], - tools: &[Tool], - ) -> Result<(Message, ProviderUsage), ProviderError> { - // Use the fast model (gpt-4o-mini) for fast completions - self.complete_with_model(OPEN_AI_FAST_MODEL, system, messages, tools) - .await - } - async fn fetch_supported_models(&self) -> Result>, ProviderError> { let models_path = self.base_path.replace("v1/chat/completions", "v1/models"); let response = self.api_client.response_get(&models_path).await?; diff --git a/crates/goose/src/providers/testprovider.rs b/crates/goose/src/providers/testprovider.rs index fb225e061cc7..ccef229ea0ef 100644 --- a/crates/goose/src/providers/testprovider.rs +++ b/crates/goose/src/providers/testprovider.rs @@ -114,7 +114,7 @@ impl Provider for TestProvider { async fn complete_with_model( &self, - model: &str, + _model: &str, system: &str, messages: &[Message], tools: &[Tool], From f95b13d64a42c370a0ad0d0e2ae2bdd522581370 Mon Sep 17 00:00:00 2001 From: David Katz Date: Wed, 20 Aug 2025 15:51:31 -0400 Subject: [PATCH 05/31] cleanup --- crates/goose/src/providers/bedrock.rs | 2 +- crates/goose/src/providers/databricks.rs | 10 ++-------- 2 files changed, 3 insertions(+), 9 deletions(-) diff --git a/crates/goose/src/providers/bedrock.rs b/crates/goose/src/providers/bedrock.rs index 57961c789885..00e1befa0176 100644 --- a/crates/goose/src/providers/bedrock.rs +++ b/crates/goose/src/providers/bedrock.rs @@ -152,7 +152,7 @@ impl Provider for BedrockProvider { } #[tracing::instrument( - skip(self, system, messages, tools), + skip(self, _model, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] async fn complete_with_model( diff --git a/crates/goose/src/providers/databricks.rs b/crates/goose/src/providers/databricks.rs index de0326ec25a7..5e33d3582374 100644 --- a/crates/goose/src/providers/databricks.rs +++ b/crates/goose/src/providers/databricks.rs @@ -107,10 +107,7 @@ pub struct DatabricksProvider { impl_provider_default!(DatabricksProvider); impl DatabricksProvider { - pub fn from_env(mut model: ModelConfig) -> Result { - // Set the default fast model for Databricks - using a smaller/faster model - model.fast_model = Some("databricks-mixtral-8x7b-instruct".to_string()); - + pub fn from_env(model: ModelConfig) -> Result { let config = crate::config::Config::global(); let mut host: Result = config.get_param("DATABRICKS_HOST"); @@ -182,10 +179,7 @@ impl DatabricksProvider { } } - pub fn from_params(host: String, api_key: String, mut model: ModelConfig) -> Result { - // Set the default fast model for Databricks - model.fast_model = Some("databricks-mixtral-8x7b-instruct".to_string()); - + pub fn from_params(host: String, api_key: String, model: ModelConfig) -> Result { let auth = DatabricksAuth::token(api_key); let auth_method = AuthMethod::Custom(Box::new(DatabricksAuthProvider { auth: auth.clone() })); From 32256b0711d4dc76722484bd33c9a1ad7e8ab76a Mon Sep 17 00:00:00 2001 From: David Katz Date: Wed, 20 Aug 2025 15:56:35 -0400 Subject: [PATCH 06/31] clean comments --- crates/goose/src/providers/anthropic.rs | 1 - crates/goose/src/providers/azure.rs | 1 - crates/goose/src/providers/claude_code.rs | 1 - crates/goose/src/providers/cursor_agent.rs | 1 - crates/goose/src/providers/databricks.rs | 2 -- crates/goose/src/providers/gcpvertexai.rs | 1 - crates/goose/src/providers/gemini_cli.rs | 1 - crates/goose/src/providers/githubcopilot.rs | 1 - crates/goose/src/providers/google.rs | 1 - crates/goose/src/providers/groq.rs | 1 - crates/goose/src/providers/litellm.rs | 1 - crates/goose/src/providers/ollama.rs | 1 - crates/goose/src/providers/openai.rs | 1 - crates/goose/src/providers/openrouter.rs | 1 - crates/goose/src/providers/sagemaker_tgi.rs | 1 - crates/goose/src/providers/snowflake.rs | 1 - crates/goose/src/providers/venice.rs | 1 - crates/goose/src/providers/xai.rs | 1 - 18 files changed, 19 deletions(-) diff --git a/crates/goose/src/providers/anthropic.rs b/crates/goose/src/providers/anthropic.rs index 57f1ea33c805..85bd2c7e13f2 100644 --- a/crates/goose/src/providers/anthropic.rs +++ b/crates/goose/src/providers/anthropic.rs @@ -198,7 +198,6 @@ impl Provider for AnthropicProvider { messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - // Create a temporary model config with the specified model let mut model_config = self.model.clone(); model_config.model_name = model.to_string(); diff --git a/crates/goose/src/providers/azure.rs b/crates/goose/src/providers/azure.rs index 1255ca38dbc4..5be63357deaa 100644 --- a/crates/goose/src/providers/azure.rs +++ b/crates/goose/src/providers/azure.rs @@ -145,7 +145,6 @@ impl Provider for AzureProvider { messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - // Create a temporary model config with the specified model let mut model_config = self.model.clone(); model_config.model_name = model.to_string(); diff --git a/crates/goose/src/providers/claude_code.rs b/crates/goose/src/providers/claude_code.rs index b2b024dcc9a9..22b4b22adc1c 100644 --- a/crates/goose/src/providers/claude_code.rs +++ b/crates/goose/src/providers/claude_code.rs @@ -484,7 +484,6 @@ impl Provider for ClaudeCodeProvider { messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - // Create a temporary model config with the specified model let mut model_config = self.model.clone(); model_config.model_name = model.to_string(); diff --git a/crates/goose/src/providers/cursor_agent.rs b/crates/goose/src/providers/cursor_agent.rs index dc393c24a6c6..2e677ffce101 100644 --- a/crates/goose/src/providers/cursor_agent.rs +++ b/crates/goose/src/providers/cursor_agent.rs @@ -417,7 +417,6 @@ impl Provider for CursorAgentProvider { messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - // Create a temporary model config with the specified model let mut model_config = self.model.clone(); model_config.model_name = model.to_string(); diff --git a/crates/goose/src/providers/databricks.rs b/crates/goose/src/providers/databricks.rs index 5e33d3582374..d89913ff08f7 100644 --- a/crates/goose/src/providers/databricks.rs +++ b/crates/goose/src/providers/databricks.rs @@ -248,7 +248,6 @@ impl Provider for DatabricksProvider { messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - // Create a temporary model config with the specified model let mut model_config = self.model.clone(); model_config.model_name = model.to_string(); @@ -278,7 +277,6 @@ impl Provider for DatabricksProvider { messages: &[Message], tools: &[Tool], ) -> Result { - // Create a temporary model config with the specified model let model_config = self.model.clone(); let mut payload = diff --git a/crates/goose/src/providers/gcpvertexai.rs b/crates/goose/src/providers/gcpvertexai.rs index 23925a65d70f..d23fa62a4e9a 100644 --- a/crates/goose/src/providers/gcpvertexai.rs +++ b/crates/goose/src/providers/gcpvertexai.rs @@ -522,7 +522,6 @@ impl Provider for GcpVertexAIProvider { messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - // Create a temporary model config with the specified model let mut model_config = self.model.clone(); model_config.model_name = model.to_string(); diff --git a/crates/goose/src/providers/gemini_cli.rs b/crates/goose/src/providers/gemini_cli.rs index 824f0f7fa5cc..d25525e074c9 100644 --- a/crates/goose/src/providers/gemini_cli.rs +++ b/crates/goose/src/providers/gemini_cli.rs @@ -329,7 +329,6 @@ impl Provider for GeminiCliProvider { messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - // Create a temporary model config with the specified model let mut model_config = self.model.clone(); model_config.model_name = model.to_string(); diff --git a/crates/goose/src/providers/githubcopilot.rs b/crates/goose/src/providers/githubcopilot.rs index 91a32fd647d7..48b2d3df29c5 100644 --- a/crates/goose/src/providers/githubcopilot.rs +++ b/crates/goose/src/providers/githubcopilot.rs @@ -411,7 +411,6 @@ impl Provider for GithubCopilotProvider { messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - // Create a temporary model config with the specified model let mut model_config = self.model.clone(); diff --git a/crates/goose/src/providers/google.rs b/crates/goose/src/providers/google.rs index 64547e87c39d..138acf27ad22 100644 --- a/crates/goose/src/providers/google.rs +++ b/crates/goose/src/providers/google.rs @@ -114,7 +114,6 @@ impl Provider for GoogleProvider { messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - // Create a temporary model config with the specified model let mut model_config = self.model.clone(); diff --git a/crates/goose/src/providers/groq.rs b/crates/goose/src/providers/groq.rs index 7de887b813f6..68c30a2440f5 100644 --- a/crates/goose/src/providers/groq.rs +++ b/crates/goose/src/providers/groq.rs @@ -90,7 +90,6 @@ impl Provider for GroqProvider { messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - // Create a temporary model config with the specified model let mut model_config = self.model.clone(); model_config.model_name = model.to_string(); diff --git a/crates/goose/src/providers/litellm.rs b/crates/goose/src/providers/litellm.rs index 500ab4994543..f198dde99d15 100644 --- a/crates/goose/src/providers/litellm.rs +++ b/crates/goose/src/providers/litellm.rs @@ -168,7 +168,6 @@ impl Provider for LiteLLMProvider { messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - // Create a temporary model config with the specified model let mut model_config = self.model.clone(); model_config.model_name = model.to_string(); diff --git a/crates/goose/src/providers/ollama.rs b/crates/goose/src/providers/ollama.rs index 2fc80f9daaa3..9afea70651a5 100644 --- a/crates/goose/src/providers/ollama.rs +++ b/crates/goose/src/providers/ollama.rs @@ -175,7 +175,6 @@ impl Provider for OllamaProvider { messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - // Create a temporary model config with the specified model let mut model_config = self.model.clone(); model_config.model_name = model.to_string(); diff --git a/crates/goose/src/providers/openai.rs b/crates/goose/src/providers/openai.rs index fe665eb20ecd..954bd35a094b 100644 --- a/crates/goose/src/providers/openai.rs +++ b/crates/goose/src/providers/openai.rs @@ -212,7 +212,6 @@ impl Provider for OpenAiProvider { messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - // Create a temporary model config with the specified model let mut model_config = self.model.clone(); model_config.model_name = model.to_string(); diff --git a/crates/goose/src/providers/openrouter.rs b/crates/goose/src/providers/openrouter.rs index c8ec0258d96b..f2e607b60276 100644 --- a/crates/goose/src/providers/openrouter.rs +++ b/crates/goose/src/providers/openrouter.rs @@ -248,7 +248,6 @@ impl Provider for OpenRouterProvider { messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - // Create a temporary model config with the specified model let mut model_config = self.model.clone(); model_config.model_name = model.to_string(); diff --git a/crates/goose/src/providers/sagemaker_tgi.rs b/crates/goose/src/providers/sagemaker_tgi.rs index 8d0180fbc760..551bb464ee80 100644 --- a/crates/goose/src/providers/sagemaker_tgi.rs +++ b/crates/goose/src/providers/sagemaker_tgi.rs @@ -290,7 +290,6 @@ impl Provider for SageMakerTgiProvider { messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - // Create a temporary model config with the specified model let mut model_config = self.model.clone(); model_config.model_name = model.to_string(); diff --git a/crates/goose/src/providers/snowflake.rs b/crates/goose/src/providers/snowflake.rs index 6089ea744149..390849e8ba55 100644 --- a/crates/goose/src/providers/snowflake.rs +++ b/crates/goose/src/providers/snowflake.rs @@ -309,7 +309,6 @@ impl Provider for SnowflakeProvider { messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - // Create a temporary model config with the specified model let mut model_config = self.model.clone(); diff --git a/crates/goose/src/providers/venice.rs b/crates/goose/src/providers/venice.rs index 3b76ebaaffb0..6a84d9c244ba 100644 --- a/crates/goose/src/providers/venice.rs +++ b/crates/goose/src/providers/venice.rs @@ -256,7 +256,6 @@ impl Provider for VeniceProvider { messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - // Create a temporary model config with the specified model let mut model_config = self.model.clone(); model_config.model_name = model.to_string(); diff --git a/crates/goose/src/providers/xai.rs b/crates/goose/src/providers/xai.rs index b8adcb567d76..30d7a347e4c4 100644 --- a/crates/goose/src/providers/xai.rs +++ b/crates/goose/src/providers/xai.rs @@ -103,7 +103,6 @@ impl Provider for XaiProvider { messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - // Create a temporary model config with the specified model let mut model_config = self.model.clone(); model_config.model_name = model.to_string(); From e8cd8fcd73dc95eb34bed2c89a8e1d05f59278cf Mon Sep 17 00:00:00 2001 From: David Katz Date: Wed, 20 Aug 2025 15:56:43 -0400 Subject: [PATCH 07/31] Fmt --- examples/frontend_tools.py | 52 ++-- examples/mcp-wiki/src/mcp_wiki/__init__.py | 3 +- examples/mcp-wiki/src/mcp_wiki/__main__.py | 2 +- examples/mcp-wiki/src/mcp_wiki/server.py | 7 +- .../generate_leaderboard.py | 134 ++++++----- .../calculate_final_scores_vibes.py | 33 ++- .../llm-judges/llm_judge.py | 136 ++++++----- .../prepare_aggregate_metrics.py | 224 ++++++++++-------- 8 files changed, 331 insertions(+), 260 deletions(-) diff --git a/examples/frontend_tools.py b/examples/frontend_tools.py index d8824f3a9828..ce25db3b60a5 100644 --- a/examples/frontend_tools.py +++ b/examples/frontend_tools.py @@ -117,6 +117,7 @@ def execute_calculator(args: Dict[str, Any]) -> List[Dict[str, Any]]: } ] + def get_tools() -> Dict[str, Any]: with httpx.Client() as client: response = client.get( @@ -154,17 +155,21 @@ def execute_enable_extension(args: Dict[str, Any]) -> List[Dict[str, Any]]: ) if add_response.status_code != 200: error_text = add_response.text - return [{ - "type": "text", - "text": f"Error: Failed to enable extension: {error_text}", - "annotations": None, - }] - - return [{ - "type": "text", - "text": f"Successfully enabled extension: {extension_name}", - "annotations": None, - }] + return [ + { + "type": "text", + "text": f"Error: Failed to enable extension: {error_text}", + "annotations": None, + } + ] + + return [ + { + "type": "text", + "text": f"Successfully enabled extension: {extension_name}", + "annotations": None, + } + ] def submit_tool_result(tool_id: str, result: List[Dict[str, Any]]) -> None: @@ -217,7 +222,7 @@ async def chat_loop() -> None: # Process the stream of responses async with client.stream( "POST", - f"{GOOSE_URL}/reply", # lock + f"{GOOSE_URL}/reply", # lock json=payload, headers={ "X-Secret-Key": SECRET_KEY, @@ -252,25 +257,26 @@ async def chat_loop() -> None: tool_call = content["toolCall"]["value"] print(f"\nTool Request: {tool_call}") - if tool_call['name'] == "calculator": + if tool_call["name"] == "calculator": print(f"Calculator: {tool_call}") # Execute the tool result = execute_calculator(tool_call["arguments"]) - elif tool_call['name'] == "enable_extension": + elif tool_call["name"] == "enable_extension": # to trigger this tool, use the instruction "use enable_extension tool with "fetch" extension name" print(f"Enabling fetch extension") - result = execute_enable_extension(args={ - "type": "stdio", - "name": "fetch", - "cmd": "uvx", - "args": ["mcp-server-fetch"], - "timeout": 300, - "bundled": False - }) + result = execute_enable_extension( + args={ + "type": "stdio", + "name": "fetch", + "cmd": "uvx", + "args": ["mcp-server-fetch"], + "timeout": 300, + "bundled": False, + } + ) listed_tools = get_tools() print(f"\nTools after enabling extension: {listed_tools}") - # Submit the result submit_tool_result(content["id"], result) diff --git a/examples/mcp-wiki/src/mcp_wiki/__init__.py b/examples/mcp-wiki/src/mcp_wiki/__init__.py index 20c873b5068f..1d60c32ebec6 100644 --- a/examples/mcp-wiki/src/mcp_wiki/__init__.py +++ b/examples/mcp-wiki/src/mcp_wiki/__init__.py @@ -1,6 +1,7 @@ import argparse from .server import mcp + def main(): """MCP Wiki: read Wikipedia articles and convert them to Markdown.""" parser = argparse.ArgumentParser( @@ -12,4 +13,4 @@ def main(): if __name__ == "__main__": - main() \ No newline at end of file + main() diff --git a/examples/mcp-wiki/src/mcp_wiki/__main__.py b/examples/mcp-wiki/src/mcp_wiki/__main__.py index 8579afa979d7..cdb4dc1b54c5 100644 --- a/examples/mcp-wiki/src/mcp_wiki/__main__.py +++ b/examples/mcp-wiki/src/mcp_wiki/__main__.py @@ -1,3 +1,3 @@ from mcp_wiki import main -main() \ No newline at end of file +main() diff --git a/examples/mcp-wiki/src/mcp_wiki/server.py b/examples/mcp-wiki/src/mcp_wiki/server.py index 33329ebb00e7..1936e2271bd7 100644 --- a/examples/mcp-wiki/src/mcp_wiki/server.py +++ b/examples/mcp-wiki/src/mcp_wiki/server.py @@ -11,6 +11,7 @@ mcp = FastMCP("wiki") + @mcp.tool() def read_wikipedia_article(url: str) -> str: """ @@ -31,7 +32,7 @@ def read_wikipedia_article(url: str) -> str: raise McpError( ErrorData( INTERNAL_ERROR, - f"Failed to retrieve the article. HTTP status code: {response.status_code}" + f"Failed to retrieve the article. HTTP status code: {response.status_code}", ) ) @@ -42,7 +43,7 @@ def read_wikipedia_article(url: str) -> str: raise McpError( ErrorData( INVALID_PARAMS, - "Could not find the main content on the provided Wikipedia URL." + "Could not find the main content on the provided Wikipedia URL.", ) ) @@ -59,5 +60,3 @@ def read_wikipedia_article(url: str) -> str: except Exception as e: # Catch-all for any other unexpected errors raise McpError(ErrorData(INTERNAL_ERROR, f"Unexpected error: {str(e)}")) from e - - diff --git a/scripts/bench-postprocess-scripts/generate_leaderboard.py b/scripts/bench-postprocess-scripts/generate_leaderboard.py index 16292c727c45..45227dadc59a 100755 --- a/scripts/bench-postprocess-scripts/generate_leaderboard.py +++ b/scripts/bench-postprocess-scripts/generate_leaderboard.py @@ -23,7 +23,7 @@ def find_aggregate_metrics_files(benchmark_dir: Path) -> list: """Find all aggregate_metrics.csv files in model subdirectories.""" csv_files = [] - + # Look for model directories in the benchmark directory for model_dir in benchmark_dir.iterdir(): if model_dir.is_dir(): @@ -33,7 +33,7 @@ def find_aggregate_metrics_files(benchmark_dir: Path) -> list: csv_path = eval_results_dir / "aggregate_metrics.csv" if csv_path.exists(): csv_files.append(csv_path) - + return csv_files @@ -44,67 +44,73 @@ def process_csv_files(csv_files: list) -> tuple: 2. A leaderboard grouping by provider and model_name with averaged metrics """ selected_columns = [ - 'provider', - 'model_name', - 'eval_suite', - 'eval_name', - 'total_tool_calls_mean', - 'prompt_execution_time_mean', - 'total_tokens_mean', - 'score_mean', - 'prompt_error_mean', - 'server_error_mean' + "provider", + "model_name", + "eval_suite", + "eval_name", + "total_tool_calls_mean", + "prompt_execution_time_mean", + "total_tokens_mean", + "score_mean", + "prompt_error_mean", + "server_error_mean", ] - + all_data = [] - + for csv_file in csv_files: try: df = pd.read_csv(csv_file) - + # Check which selected columns are available missing_columns = [col for col in selected_columns if col not in df.columns] if missing_columns: print(f"Warning: {csv_file} is missing columns: {missing_columns}") - + # For missing columns, add them with NaN values for col in missing_columns: - df[col] = float('nan') - + df[col] = float("nan") + # Select only the columns we care about - df_subset = df[selected_columns].copy() # Create a copy to avoid SettingWithCopyWarning - + df_subset = df[ + selected_columns + ].copy() # Create a copy to avoid SettingWithCopyWarning + # Add model folder name as additional context model_folder = csv_file.parent.parent.name - df_subset['model_folder'] = model_folder - + df_subset["model_folder"] = model_folder + all_data.append(df_subset) - + except Exception as e: print(f"Error processing {csv_file}: {str(e)}") - + if not all_data: raise ValueError("No valid CSV files found with required columns") - + # Concatenate all dataframes to create a union union_df = pd.concat(all_data, ignore_index=True) - + # Create leaderboard by grouping and averaging numerical columns numeric_columns = [ - 'total_tool_calls_mean', - 'prompt_execution_time_mean', - 'total_tokens_mean', - 'score_mean', - 'prompt_error_mean', - 'server_error_mean' + "total_tool_calls_mean", + "prompt_execution_time_mean", + "total_tokens_mean", + "score_mean", + "prompt_error_mean", + "server_error_mean", ] - + # Group by provider and model_name, then calculate averages for numeric columns - leaderboard_df = union_df.groupby(['provider', 'model_name'])[numeric_columns].mean().reset_index() - + leaderboard_df = ( + union_df.groupby(["provider", "model_name"])[numeric_columns] + .mean() + .reset_index() + ) + # Sort by score_mean in descending order (highest scores first) - leaderboard_df = leaderboard_df.sort_values('score_mean', ascending=False) - + leaderboard_df = leaderboard_df.sort_values("score_mean", ascending=False) + return union_df, leaderboard_df @@ -116,65 +122,75 @@ def main(): "--benchmark-dir", type=str, required=True, - help="Path to the benchmark directory containing model subdirectories" + help="Path to the benchmark directory containing model subdirectories", ) parser.add_argument( "--union-output", type=str, default="all_metrics.csv", - help="Output filename for the union of all CSVs (default: all_metrics.csv)" + help="Output filename for the union of all CSVs (default: all_metrics.csv)", ) parser.add_argument( "--leaderboard-output", type=str, default="leaderboard.csv", - help="Output filename for the leaderboard (default: leaderboard.csv)" + help="Output filename for the leaderboard (default: leaderboard.csv)", ) - + args = parser.parse_args() - + benchmark_dir = Path(args.benchmark_dir) if not benchmark_dir.exists() or not benchmark_dir.is_dir(): - print(f"Error: Benchmark directory {benchmark_dir} does not exist or is not a directory") + print( + f"Error: Benchmark directory {benchmark_dir} does not exist or is not a directory" + ) sys.exit(1) - + try: # Find all aggregate_metrics.csv files in model subdirectories csv_files = find_aggregate_metrics_files(benchmark_dir) - + if not csv_files: - print(f"No aggregate_metrics.csv files found in any model directory under {benchmark_dir}") + print( + f"No aggregate_metrics.csv files found in any model directory under {benchmark_dir}" + ) sys.exit(1) - - print(f"Found {len(csv_files)} aggregate_metrics.csv files in model directories") - + + print( + f"Found {len(csv_files)} aggregate_metrics.csv files in model directories" + ) + # Process and create the union and leaderboard dataframes union_df, leaderboard_df = process_csv_files(csv_files) - + # Save the union CSV to the benchmark directory union_output_path = benchmark_dir / args.union_output union_df.to_csv(union_output_path, index=False) print(f"Union CSV with all metrics saved to: {union_output_path}") - + # Save the leaderboard CSV to the benchmark directory leaderboard_output_path = benchmark_dir / args.leaderboard_output leaderboard_df.to_csv(leaderboard_output_path, index=False) - print(f"Leaderboard CSV with averaged metrics saved to: {leaderboard_output_path}") - + print( + f"Leaderboard CSV with averaged metrics saved to: {leaderboard_output_path}" + ) + # Print a summary of the leaderboard print("\nLeaderboard Summary:") - pd.set_option('display.max_columns', None) # Show all columns + pd.set_option("display.max_columns", None) # Show all columns print(leaderboard_df.to_string(index=False)) - + # Highlight models with server errors - if 'server_error_mean' in leaderboard_df.columns: - models_with_errors = leaderboard_df[leaderboard_df['server_error_mean'] > 0] + if "server_error_mean" in leaderboard_df.columns: + models_with_errors = leaderboard_df[leaderboard_df["server_error_mean"] > 0] if not models_with_errors.empty: print("\nWARNING - Models with server errors detected:") for _, row in models_with_errors.iterrows(): - print(f" * {row['provider']} {row['model_name']} - {row['server_error_mean']*100:.1f}% of evaluations had server errors") + print( + f" * {row['provider']} {row['model_name']} - {row['server_error_mean'] * 100:.1f}% of evaluations had server errors" + ) print("\nThese models may need to be re-run to get accurate results.") - + except Exception as e: print(f"Error: {str(e)}") sys.exit(1) diff --git a/scripts/bench-postprocess-scripts/llm-judges/calculate_final_scores_vibes.py b/scripts/bench-postprocess-scripts/llm-judges/calculate_final_scores_vibes.py index 261fbe52832d..17c4259775d7 100755 --- a/scripts/bench-postprocess-scripts/llm-judges/calculate_final_scores_vibes.py +++ b/scripts/bench-postprocess-scripts/llm-judges/calculate_final_scores_vibes.py @@ -28,14 +28,14 @@ def calculate_score(eval_name, metrics): llm_judge_score = get_metric_value(metrics, "llm_judge_score") used_fetch_tool = get_metric_value(metrics, "used_fetch_tool") valid_markdown_format = get_metric_value(metrics, "valid_markdown_format") - + if llm_judge_score is None: raise ValueError("llm_judge_score not found in metrics") - + # Convert boolean metrics to 0/1 if needed used_fetch_tool = 1.0 if used_fetch_tool else 0.0 valid_markdown_format = 1.0 if valid_markdown_format else 0.0 - + if eval_name == "blog_summary": # max score is 4.0 as llm_judge_score is between [0,2] and used_fetch_tool/valid_markedown_format have values [0,1] score = (llm_judge_score + used_fetch_tool + valid_markdown_format) / 4.0 @@ -43,7 +43,7 @@ def calculate_score(eval_name, metrics): score = (llm_judge_score + valid_markdown_format + used_fetch_tool) / 4.0 else: raise ValueError(f"Unknown evaluation type: {eval_name}") - + return score @@ -51,34 +51,31 @@ def main(): if len(sys.argv) != 2: print("Usage: calculate_final_score.py ") sys.exit(1) - + eval_name = sys.argv[1] - + # Load eval results from current directory eval_results_path = Path("eval-results.json") if not eval_results_path.exists(): print(f"Error: eval-results.json not found in current directory") sys.exit(1) - - with open(eval_results_path, 'r') as f: + + with open(eval_results_path, "r") as f: eval_results = json.load(f) - + try: # Calculate the final score score = calculate_score(eval_name, eval_results["metrics"]) - + # Add the score metric - eval_results["metrics"].append([ - "score", - {"Float": score} - ]) - + eval_results["metrics"].append(["score", {"Float": score}]) + # Save updated results - with open(eval_results_path, 'w') as f: + with open(eval_results_path, "w") as f: json.dump(eval_results, f, indent=2) - + print(f"Successfully added final score: {score}") - + except Exception as e: print(f"Error calculating final score: {str(e)}") sys.exit(1) diff --git a/scripts/bench-postprocess-scripts/llm-judges/llm_judge.py b/scripts/bench-postprocess-scripts/llm-judges/llm_judge.py index 2b22dc24489c..e32bcf7744f7 100755 --- a/scripts/bench-postprocess-scripts/llm-judges/llm_judge.py +++ b/scripts/bench-postprocess-scripts/llm-judges/llm_judge.py @@ -8,7 +8,7 @@ Usage: python llm_judge.py [--rubric-max-score N] [--prompt-file PATH] - + Arguments: output_file: Name of the file containing the output to evaluate (e.g., blog_summary_output.txt) --rubric-max-score: Maximum score for the rubric (default: 2) @@ -33,15 +33,15 @@ def evaluate_with_openai(prompt: str, text: str, rubric_max_score: int = 2) -> float: """Evaluate response using OpenAI's API. - + Args: prompt: System prompt for evaluation text: Text to evaluate rubric_max_score: Maximum score for the rubric (default: 2.0) - + Returns: float: Evaluation score (0 to rubric_max_score) - + Raises: ValueError: If OPENAI_API_KEY environment variable is not set """ @@ -49,11 +49,13 @@ def evaluate_with_openai(prompt: str, text: str, rubric_max_score: int = 2) -> f api_key = os.getenv("OPENAI_API_KEY") if not api_key: print("No OpenAI API key found!") - raise ValueError("OPENAI_API_KEY environment variable is not set, but is needed to run this evaluation.") - + raise ValueError( + "OPENAI_API_KEY environment variable is not set, but is needed to run this evaluation." + ) + try: client = OpenAI(api_key=api_key) - + # Append output instructions to system prompt output_instructions = f""" Output Instructions: @@ -68,25 +70,23 @@ def evaluate_with_openai(prompt: str, text: str, rubric_max_score: int = 2) -> f - Do not include any additional text before or after the JSON - Return only the raw JSON object - The score must be an integer between 0 and {rubric_max_score}""" - + input_prompt = f"{prompt} {output_instructions}\nResponse to evaluate: {text}" - + # Run the chat completion 3 times and collect scores scores = [] for i in range(3): max_retries = 5 retry_count = 0 - + while retry_count < max_retries: try: response = client.chat.completions.create( model="gpt-4o", - messages=[ - {"role": "user", "content": input_prompt} - ], - temperature=0.9 + messages=[{"role": "user", "content": input_prompt}], + temperature=0.9, ) - + # Extract and parse JSON from response response_text = response.choices[0].message.content.strip() try: @@ -94,14 +94,18 @@ def evaluate_with_openai(prompt: str, text: str, rubric_max_score: int = 2) -> f score = float(evaluation.get("score", 0.0)) score = max(0.0, min(score, rubric_max_score)) scores.append(score) - print(f"Run {i+1} score: {score}") + print(f"Run {i + 1} score: {score}") break # Successfully parsed, exit retry loop except (json.JSONDecodeError, ValueError) as e: retry_count += 1 - print(f"Error parsing OpenAI response as JSON (attempt {retry_count}/{max_retries}): {str(e)}") + print( + f"Error parsing OpenAI response as JSON (attempt {retry_count}/{max_retries}): {str(e)}" + ) print(f"Response text: {response_text}") if retry_count == max_retries: - raise ValueError(f"Failed to parse OpenAI evaluation response after {max_retries} attempts: {str(e)}") + raise ValueError( + f"Failed to parse OpenAI evaluation response after {max_retries} attempts: {str(e)}" + ) print("Retrying...") time.sleep(1) # Wait 1 second before retrying continue @@ -109,26 +113,24 @@ def evaluate_with_openai(prompt: str, text: str, rubric_max_score: int = 2) -> f # For other exceptions (API errors, etc.), raise immediately print(f"API error: {str(e)}") raise - + # Count occurrences of each score score_counts = Counter(scores) - + # If there's no single most common score (all scores are different), run one more time if len(scores) == 3 and max(score_counts.values()) == 1: print("No majority score found. Running tie-breaker...") max_retries = 5 retry_count = 0 - + while retry_count < max_retries: try: response = client.chat.completions.create( model="gpt-4o", - messages=[ - {"role": "user", "content": input_prompt} - ], - temperature=0.9 + messages=[{"role": "user", "content": input_prompt}], + temperature=0.9, ) - + response_text = response.choices[0].message.content.strip() try: evaluation = json.loads(response_text) @@ -140,10 +142,14 @@ def evaluate_with_openai(prompt: str, text: str, rubric_max_score: int = 2) -> f break # Successfully parsed, exit retry loop except (json.JSONDecodeError, ValueError) as e: retry_count += 1 - print(f"Error parsing tie-breaker response as JSON (attempt {retry_count}/{max_retries}): {str(e)}") + print( + f"Error parsing tie-breaker response as JSON (attempt {retry_count}/{max_retries}): {str(e)}" + ) print(f"Response text: {response_text}") if retry_count == max_retries: - raise ValueError(f"Failed to parse tie-breaker response after {max_retries} attempts: {str(e)}") + raise ValueError( + f"Failed to parse tie-breaker response after {max_retries} attempts: {str(e)}" + ) print("Retrying tie-breaker...") time.sleep(1) # Wait 1 second before retrying continue @@ -151,12 +157,14 @@ def evaluate_with_openai(prompt: str, text: str, rubric_max_score: int = 2) -> f # For other exceptions (API errors, etc.), raise immediately print(f"API error in tie-breaker: {str(e)}") raise - + # Get the most common score most_common_score = score_counts.most_common(1)[0][0] - print(f"Most common score: {most_common_score} (occurred {score_counts[most_common_score]} times)") + print( + f"Most common score: {most_common_score} (occurred {score_counts[most_common_score]} times)" + ) return most_common_score - + except Exception as e: if "OPENAI_API_KEY" in str(e): raise # Re-raise API key errors @@ -169,8 +177,8 @@ def load_eval_results(working_dir: Path) -> Dict[str, Any]: eval_results_path = working_dir / "eval-results.json" if not eval_results_path.exists(): raise FileNotFoundError(f"eval-results.json not found in {working_dir}") - - with open(eval_results_path, 'r') as f: + + with open(eval_results_path, "r") as f: return json.load(f) @@ -179,22 +187,22 @@ def load_output_file(working_dir: Path, output_file: str) -> str: output_path = working_dir / output_file if not output_path.exists(): raise FileNotFoundError(f"Output file not found: {output_path}") - - with open(output_path, 'r') as f: + + with open(output_path, "r") as f: return f.read().strip() def load_evaluation_prompt(working_dir: Path) -> str: """Load the evaluation prompt from a file or use a default. - + This function looks for a prompt.txt file in the working directory. If not found, it returns a default evaluation prompt. """ prompt_file = working_dir / "prompt.txt" if prompt_file.exists(): - with open(prompt_file, 'r') as f: + with open(prompt_file, "r") as f: return f.read().strip() - + # Default evaluation prompt return """You are an expert evaluator assessing the quality of AI responses. Evaluate the response based on the following criteria: @@ -210,46 +218,58 @@ def load_evaluation_prompt(working_dir: Path) -> str: def main(): - parser = argparse.ArgumentParser(description="LLM Judge post-processing script for Goose benchmarks") - parser.add_argument("output_file", type=str, help="Name of the output file to evaluate (e.g., blog_summary_output.txt)") - parser.add_argument("--rubric-max-score", type=int, default=2, help="Maximum score for the rubric (default: 2)") - parser.add_argument("--prompt-file", type=str, help="Path to custom evaluation prompt file") - + parser = argparse.ArgumentParser( + description="LLM Judge post-processing script for Goose benchmarks" + ) + parser.add_argument( + "output_file", + type=str, + help="Name of the output file to evaluate (e.g., blog_summary_output.txt)", + ) + parser.add_argument( + "--rubric-max-score", + type=int, + default=2, + help="Maximum score for the rubric (default: 2)", + ) + parser.add_argument( + "--prompt-file", type=str, help="Path to custom evaluation prompt file" + ) + args = parser.parse_args() - + # Use current working directory working_dir = Path.cwd() - + try: # Load eval results eval_results = load_eval_results(working_dir) - + # Load the output file to evaluate response_text = load_output_file(working_dir, args.output_file) - + # Load evaluation prompt if args.prompt_file: - with open(args.prompt_file, 'r') as f: + with open(args.prompt_file, "r") as f: evaluation_prompt = f.read().strip() else: evaluation_prompt = load_evaluation_prompt(working_dir) - + # Evaluate with OpenAI - score = evaluate_with_openai(evaluation_prompt, response_text, args.rubric_max_score) - + score = evaluate_with_openai( + evaluation_prompt, response_text, args.rubric_max_score + ) + # Update eval results with the score - eval_results["metrics"].append([ - "llm_judge_score", - {"Float": score} - ]) + eval_results["metrics"].append(["llm_judge_score", {"Float": score}]) # Save updated results eval_results_path = working_dir / "eval-results.json" - with open(eval_results_path, 'w') as f: + with open(eval_results_path, "w") as f: json.dump(eval_results, f, indent=2) - + print(f"Successfully updated eval-results.json with LLM judge score: {score}") - + except Exception as e: print(f"Error: {str(e)}") sys.exit(1) diff --git a/scripts/bench-postprocess-scripts/prepare_aggregate_metrics.py b/scripts/bench-postprocess-scripts/prepare_aggregate_metrics.py index 399172918968..ffdd5e09ebd5 100755 --- a/scripts/bench-postprocess-scripts/prepare_aggregate_metrics.py +++ b/scripts/bench-postprocess-scripts/prepare_aggregate_metrics.py @@ -21,60 +21,74 @@ from pathlib import Path import sys + def extract_provider_model(model_dir): """Extract provider and model name from directory name.""" dir_name = model_dir.name - parts = dir_name.split('-') - + parts = dir_name.split("-") + if len(parts) > 1: model_name = parts[-1] # Last part is the model name - provider = '-'.join(parts[:-1]) # Everything else is the provider + provider = "-".join(parts[:-1]) # Everything else is the provider else: model_name = dir_name provider = "unknown" - + return provider, model_name + def find_eval_results_files(model_dir): """Find all eval-results.json files in a model directory.""" return list(model_dir.glob("**/eval-results.json")) + def find_session_files(model_dir): """Find all session jsonl files in a model directory.""" return list(model_dir.glob("**/*.jsonl")) + def check_for_errors_in_session(session_file): """Check if a session file contains server errors.""" try: error_found = False error_messages = [] - - with open(session_file, 'r') as f: + + with open(session_file, "r") as f: for line in f: try: message_obj = json.loads(line.strip()) # Check for error messages in the content - if 'content' in message_obj and isinstance(message_obj['content'], list): - for content_item in message_obj['content']: - if isinstance(content_item, dict) and 'text' in content_item: - text = content_item['text'] - if 'Server error' in text or 'error_code' in text or 'TEMPORARILY_UNAVAILABLE' in text: + if "content" in message_obj and isinstance( + message_obj["content"], list + ): + for content_item in message_obj["content"]: + if ( + isinstance(content_item, dict) + and "text" in content_item + ): + text = content_item["text"] + if ( + "Server error" in text + or "error_code" in text + or "TEMPORARILY_UNAVAILABLE" in text + ): error_found = True error_messages.append(text) except json.JSONDecodeError: continue - + return error_found, error_messages except Exception as e: print(f"Error checking session file {session_file}: {str(e)}") return False, [] + def extract_metrics_from_eval_file(eval_file, provider, model_name, session_files): """Extract metrics from an eval-results.json file.""" try: - with open(eval_file, 'r') as f: + with open(eval_file, "r") as f: data = json.load(f) - + # Extract directory structure to determine eval suite and name path_parts = eval_file.parts run_index = -1 @@ -82,145 +96,154 @@ def extract_metrics_from_eval_file(eval_file, provider, model_name, session_file if part.startswith("run-"): run_index = i break - + if run_index == -1 or run_index + 2 >= len(path_parts): print(f"Warning: Could not determine eval suite and name from {eval_file}") return None - - run_number = path_parts[run_index].split('-')[1] # Extract "0" from "run-0" + + run_number = path_parts[run_index].split("-")[1] # Extract "0" from "run-0" eval_suite = path_parts[run_index + 1] # Directory after run-N eval_name = path_parts[run_index + 2] # Directory after eval_suite - + # Create a row with basic identification row = { - 'provider': provider, - 'model_name': model_name, - 'eval_suite': eval_suite, - 'eval_name': eval_name, - 'run': run_number + "provider": provider, + "model_name": model_name, + "eval_suite": eval_suite, + "eval_name": eval_name, + "run": run_number, } - + # Check for server errors in session files for this evaluation eval_dir = eval_file.parent related_session_files = [sf for sf in session_files if eval_dir in sf.parents] - + server_error_found = False for session_file in related_session_files: error_found, _ = check_for_errors_in_session(session_file) if error_found: server_error_found = True break - + # Add server error flag - row['server_error'] = 1 if server_error_found else 0 - + row["server_error"] = 1 if server_error_found else 0 + # Extract all metrics (flatten the JSON structure) if isinstance(data, dict): metrics = {} - + # Extract top-level metrics for key, value in data.items(): if isinstance(value, (int, float)) and not isinstance(value, bool): metrics[key] = value - + # Look for nested metrics structure (list of [name, value] pairs) - if 'metrics' in data and isinstance(data['metrics'], list): - for metric_item in data['metrics']: + if "metrics" in data and isinstance(data["metrics"], list): + for metric_item in data["metrics"]: if isinstance(metric_item, list) and len(metric_item) == 2: metric_name = metric_item[0] metric_value = metric_item[1] - + # Handle different value formats if isinstance(metric_value, dict): - if 'Integer' in metric_value: - metrics[metric_name] = int(metric_value['Integer']) - elif 'Float' in metric_value: - metrics[metric_name] = float(metric_value['Float']) - elif 'Bool' in metric_value: - metrics[metric_name] = 1 if metric_value['Bool'] else 0 + if "Integer" in metric_value: + metrics[metric_name] = int(metric_value["Integer"]) + elif "Float" in metric_value: + metrics[metric_name] = float(metric_value["Float"]) + elif "Bool" in metric_value: + metrics[metric_name] = 1 if metric_value["Bool"] else 0 # Skip string values for aggregation - elif isinstance(metric_value, (int, float)) and not isinstance(metric_value, bool): + elif isinstance(metric_value, (int, float)) and not isinstance( + metric_value, bool + ): metrics[metric_name] = metric_value elif isinstance(metric_value, bool): metrics[metric_name] = 1 if metric_value else 0 - + # Look for metrics in other common locations - for metric_location in ['metrics', 'result', 'evaluation']: + for metric_location in ["metrics", "result", "evaluation"]: if metric_location in data and isinstance(data[metric_location], dict): for key, value in data[metric_location].items(): - if isinstance(value, (int, float)) and not isinstance(value, bool): + if isinstance(value, (int, float)) and not isinstance( + value, bool + ): metrics[key] = value elif isinstance(value, bool): metrics[key] = 1 if value else 0 - + # Add all metrics to the row row.update(metrics) - + # Ensure a score is present (if not, add a placeholder) - if 'score' not in row: + if "score" not in row: # Try to use existing fields to calculate a score if server_error_found: - row['score'] = 0 # Failed runs get a zero score + row["score"] = 0 # Failed runs get a zero score else: # Set a default based on presence of "success" fields for key in row: - if 'success' in key.lower() and isinstance(row[key], (int, float)): - row['score'] = row[key] + if "success" in key.lower() and isinstance( + row[key], (int, float) + ): + row["score"] = row[key] break else: # No success field found, mark as NaN - row['score'] = float('nan') - + row["score"] = float("nan") + return row else: print(f"Warning: Unexpected format in {eval_file}") return None - + except Exception as e: print(f"Error processing {eval_file}: {str(e)}") return None + def process_model_directory(model_dir): """Process a model directory to create aggregate_metrics.csv.""" provider, model_name = extract_provider_model(model_dir) - + # Find all eval results files eval_files = find_eval_results_files(model_dir) if not eval_files: print(f"No eval-results.json files found in {model_dir}") return False - + # Find all session files for error checking session_files = find_session_files(model_dir) - + # Extract metrics from each eval file rows = [] for eval_file in eval_files: - row = extract_metrics_from_eval_file(eval_file, provider, model_name, session_files) + row = extract_metrics_from_eval_file( + eval_file, provider, model_name, session_files + ) if row is not None: rows.append(row) - + if not rows: print(f"No valid metrics extracted from {model_dir}") return False - + # Create a dataframe from all rows combined_df = pd.DataFrame(rows) - + # Calculate aggregates for numeric columns, grouped by eval_suite, eval_name - numeric_cols = combined_df.select_dtypes(include=['number']).columns.tolist() + numeric_cols = combined_df.select_dtypes(include=["number"]).columns.tolist() # Exclude the run column from aggregation - if 'run' in numeric_cols: - numeric_cols.remove('run') - + if "run" in numeric_cols: + numeric_cols.remove("run") + # Group by provider, model_name, eval_suite, eval_name and calculate mean for numeric columns - group_by_cols = ['provider', 'model_name', 'eval_suite', 'eval_name'] - agg_dict = {col: 'mean' for col in numeric_cols} - + group_by_cols = ["provider", "model_name", "eval_suite", "eval_name"] + agg_dict = {col: "mean" for col in numeric_cols} + # Only perform aggregation if we have numeric columns if numeric_cols: aggregate_df = combined_df.groupby(group_by_cols).agg(agg_dict).reset_index() - + # Rename columns to add _mean suffix for the averaged metrics for col in numeric_cols: aggregate_df = aggregate_df.rename(columns={col: f"{col}_mean"}) @@ -228,38 +251,41 @@ def process_model_directory(model_dir): print(f"Warning: No numeric metrics found in {model_dir}") # Create a minimal dataframe with just the grouping columns aggregate_df = combined_df[group_by_cols].drop_duplicates() - + # Make sure we have prompt_execution_time_mean and prompt_error_mean columns # These are expected by the generate_leaderboard.py script - if 'prompt_execution_time_mean' not in aggregate_df.columns: - aggregate_df['prompt_execution_time_mean'] = float('nan') - - if 'prompt_error_mean' not in aggregate_df.columns: - aggregate_df['prompt_error_mean'] = float('nan') - + if "prompt_execution_time_mean" not in aggregate_df.columns: + aggregate_df["prompt_execution_time_mean"] = float("nan") + + if "prompt_error_mean" not in aggregate_df.columns: + aggregate_df["prompt_error_mean"] = float("nan") + # Add server_error_mean column if not present - if 'server_error_mean' not in aggregate_df.columns: - aggregate_df['server_error_mean'] = 0.0 - + if "server_error_mean" not in aggregate_df.columns: + aggregate_df["server_error_mean"] = 0.0 + # Create eval-results directory eval_results_dir = model_dir / "eval-results" eval_results_dir.mkdir(exist_ok=True) - + # Save to CSV csv_path = eval_results_dir / "aggregate_metrics.csv" aggregate_df.to_csv(csv_path, index=False) - + # Count number of evaluations that had server errors - if 'server_error_mean' in aggregate_df.columns: - error_count = len(aggregate_df[aggregate_df['server_error_mean'] > 0]) + if "server_error_mean" in aggregate_df.columns: + error_count = len(aggregate_df[aggregate_df["server_error_mean"] > 0]) total_count = len(aggregate_df) - print(f"Saved aggregate metrics to {csv_path} with {len(aggregate_df)} rows " + - f"({error_count}/{total_count} evals had server errors)") + print( + f"Saved aggregate metrics to {csv_path} with {len(aggregate_df)} rows " + + f"({error_count}/{total_count} evals had server errors)" + ) else: print(f"Saved aggregate metrics to {csv_path} with {len(aggregate_df)} rows") - + return True + def main(): parser = argparse.ArgumentParser( description="Prepare aggregate_metrics.csv files from eval-results.json files with error detection" @@ -268,33 +294,39 @@ def main(): "--benchmark-dir", type=str, required=True, - help="Path to the benchmark directory containing model subdirectories" + help="Path to the benchmark directory containing model subdirectories", ) - + args = parser.parse_args() - + # Convert path to Path object and validate it exists benchmark_dir = Path(args.benchmark_dir) if not benchmark_dir.exists() or not benchmark_dir.is_dir(): - print(f"Error: Benchmark directory {benchmark_dir} does not exist or is not a directory") + print( + f"Error: Benchmark directory {benchmark_dir} does not exist or is not a directory" + ) sys.exit(1) - + success_count = 0 - + # Process each model directory for model_dir in benchmark_dir.iterdir(): - if model_dir.is_dir() and not model_dir.name.startswith('.'): + if model_dir.is_dir() and not model_dir.name.startswith("."): if process_model_directory(model_dir): success_count += 1 - + if success_count == 0: print("No aggregate_metrics.csv files were created") sys.exit(1) - - print(f"Successfully created aggregate_metrics.csv files for {success_count} model directories") + + print( + f"Successfully created aggregate_metrics.csv files for {success_count} model directories" + ) print("You can now run generate_leaderboard.py to create the final leaderboard.") - print("Note: The server_error_mean column indicates the average rate of server errors across evaluations.") + print( + "Note: The server_error_mean column indicates the average rate of server errors across evaluations." + ) + if __name__ == "__main__": main() - From 3c3c3768aa282c6077e779c5be076ed57d0e2dbd Mon Sep 17 00:00:00 2001 From: David Katz Date: Wed, 20 Aug 2025 16:24:43 -0400 Subject: [PATCH 08/31] parse context limit --- crates/goose/src/model.rs | 46 +++++++++++++++++++-- crates/goose/src/providers/githubcopilot.rs | 1 - crates/goose/src/providers/google.rs | 1 - crates/goose/src/providers/snowflake.rs | 1 - 4 files changed, 43 insertions(+), 6 deletions(-) diff --git a/crates/goose/src/model.rs b/crates/goose/src/model.rs index 89bf1167a9e9..9dfad80f171f 100644 --- a/crates/goose/src/model.rs +++ b/crates/goose/src/model.rs @@ -21,13 +21,16 @@ static MODEL_SPECIFIC_LIMITS: Lazy> = Lazy::new(|| { ("gpt-4-turbo", 128_000), ("gpt-4.1", 1_000_000), ("gpt-4-1", 1_000_000), + ("gpt-4o-mini", 128_000), // Fast model for OpenAI ("gpt-4o", 128_000), ("o4-mini", 200_000), ("o3-mini", 200_000), ("o3", 200_000), // anthropic - all 200k + ("claude-3-5-haiku", 200_000), // Fast model for Anthropic ("claude", 200_000), // google + ("gemini-1.5-flash", 1_000_000), // Fast model for Google ("gemini-1", 128_000), ("gemini-2", 1_000_000), ("gemma-3-27b", 128_000), @@ -41,6 +44,7 @@ static MODEL_SPECIFIC_LIMITS: Lazy> = Lazy::new(|| { ("gemma-2-27b", 8_192), ("gemma-2-9b", 8_192), ("gemma-2-2b", 8_192), + ("gemma2-9b", 8_192), // Fast model for Groq ("gemma2-", 8_192), ("gemma-7b", 8_192), ("gemma-2b", 8_192), @@ -90,7 +94,7 @@ impl ModelConfig { model_name: String, context_env_var: Option<&str>, ) -> Result { - let context_limit = Self::parse_context_limit(&model_name, context_env_var)?; + let context_limit = Self::parse_context_limit(&model_name, None, context_env_var)?; let temperature = Self::parse_temperature()?; let toolshim = Self::parse_toolshim()?; let toolshim_model = Self::parse_toolshim_model()?; @@ -108,8 +112,10 @@ impl ModelConfig { fn parse_context_limit( model_name: &str, + fast_model: Option<&str>, custom_env_var: Option<&str>, ) -> Result, ConfigError> { + // First check if there's an explicit environment variable override if let Some(env_var) = custom_env_var { if let Ok(val) = std::env::var(env_var) { return Self::validate_context_limit(&val, env_var).map(Some); @@ -118,7 +124,24 @@ impl ModelConfig { if let Ok(val) = std::env::var("GOOSE_CONTEXT_LIMIT") { return Self::validate_context_limit(&val, "GOOSE_CONTEXT_LIMIT").map(Some); } - Ok(Self::get_model_specific_limit(model_name)) + + // Get the model's limit + let model_limit = Self::get_model_specific_limit(model_name); + + // If there's a fast_model, get its limit and use the minimum + if let Some(fast_model_name) = fast_model { + let fast_model_limit = Self::get_model_specific_limit(fast_model_name); + + // Return the minimum of both limits (if both exist) + match (model_limit, fast_model_limit) { + (Some(m), Some(f)) => Ok(Some(m.min(f))), + (Some(m), None) => Ok(Some(m)), + (None, Some(f)) => Ok(Some(f)), + (None, None) => Ok(None), + } + } else { + Ok(model_limit) + } } fn validate_context_limit(val: &str, env_var: &str) -> Result { @@ -234,7 +257,24 @@ impl ModelConfig { } pub fn context_limit(&self) -> usize { - self.context_limit.unwrap_or(DEFAULT_CONTEXT_LIMIT) + // If we have a fast_model, use the minimum of both limits + if let Some(fast_model) = &self.fast_model { + let main_limit = self + .context_limit + .or_else(|| Self::get_model_specific_limit(&self.model_name)); + let fast_limit = Self::get_model_specific_limit(fast_model); + + match (main_limit, fast_limit) { + (Some(m), Some(f)) => m.min(f), + (Some(m), None) => m, + (None, Some(f)) => f, + (None, None) => DEFAULT_CONTEXT_LIMIT, + } + } else { + self.context_limit.unwrap_or_else(|| { + Self::get_model_specific_limit(&self.model_name).unwrap_or(DEFAULT_CONTEXT_LIMIT) + }) + } } pub fn new_or_fail(model_name: &str) -> ModelConfig { diff --git a/crates/goose/src/providers/githubcopilot.rs b/crates/goose/src/providers/githubcopilot.rs index 48b2d3df29c5..ba1d16de6f16 100644 --- a/crates/goose/src/providers/githubcopilot.rs +++ b/crates/goose/src/providers/githubcopilot.rs @@ -411,7 +411,6 @@ impl Provider for GithubCopilotProvider { messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - let mut model_config = self.model.clone(); model_config.model_name = model.to_string(); diff --git a/crates/goose/src/providers/google.rs b/crates/goose/src/providers/google.rs index 138acf27ad22..1816cf5728de 100644 --- a/crates/goose/src/providers/google.rs +++ b/crates/goose/src/providers/google.rs @@ -114,7 +114,6 @@ impl Provider for GoogleProvider { messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - let mut model_config = self.model.clone(); model_config.model_name = model.to_string(); diff --git a/crates/goose/src/providers/snowflake.rs b/crates/goose/src/providers/snowflake.rs index 390849e8ba55..8df85e5f4a99 100644 --- a/crates/goose/src/providers/snowflake.rs +++ b/crates/goose/src/providers/snowflake.rs @@ -309,7 +309,6 @@ impl Provider for SnowflakeProvider { messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - let mut model_config = self.model.clone(); model_config.model_name = model.to_string(); From bdb5583ac2bae25378f5e92b4ad54db2d435f606 Mon Sep 17 00:00:00 2001 From: David Katz Date: Wed, 20 Aug 2025 16:34:10 -0400 Subject: [PATCH 09/31] more model config --- crates/goose/src/model.rs | 5 +---- crates/goose/src/providers/groq.rs | 4 ++-- 2 files changed, 3 insertions(+), 6 deletions(-) diff --git a/crates/goose/src/model.rs b/crates/goose/src/model.rs index 9dfad80f171f..d81ef6a81732 100644 --- a/crates/goose/src/model.rs +++ b/crates/goose/src/model.rs @@ -21,16 +21,14 @@ static MODEL_SPECIFIC_LIMITS: Lazy> = Lazy::new(|| { ("gpt-4-turbo", 128_000), ("gpt-4.1", 1_000_000), ("gpt-4-1", 1_000_000), - ("gpt-4o-mini", 128_000), // Fast model for OpenAI ("gpt-4o", 128_000), ("o4-mini", 200_000), ("o3-mini", 200_000), ("o3", 200_000), // anthropic - all 200k - ("claude-3-5-haiku", 200_000), // Fast model for Anthropic ("claude", 200_000), // google - ("gemini-1.5-flash", 1_000_000), // Fast model for Google + ("gemini-1.5-flash", 1_000_000), ("gemini-1", 128_000), ("gemini-2", 1_000_000), ("gemma-3-27b", 128_000), @@ -44,7 +42,6 @@ static MODEL_SPECIFIC_LIMITS: Lazy> = Lazy::new(|| { ("gemma-2-27b", 8_192), ("gemma-2-9b", 8_192), ("gemma-2-2b", 8_192), - ("gemma2-9b", 8_192), // Fast model for Groq ("gemma2-", 8_192), ("gemma-7b", 8_192), ("gemma-2b", 8_192), diff --git a/crates/goose/src/providers/groq.rs b/crates/goose/src/providers/groq.rs index 68c30a2440f5..2bff2086eeec 100644 --- a/crates/goose/src/providers/groq.rs +++ b/crates/goose/src/providers/groq.rs @@ -34,8 +34,8 @@ impl_provider_default!(GroqProvider); impl GroqProvider { pub fn from_env(mut model: ModelConfig) -> Result { - // Set the default fast model for Groq - using a smaller/faster model - model.fast_model = Some("gemma2-9b-it".to_string()); + // Set the default fast model for Groq - using llama 70b for better context + model.fast_model = Some("llama-3.3-70b-versatile".to_string()); let config = crate::config::Config::global(); let api_key: String = config.get_secret("GROQ_API_KEY")?; From 352e0c5d73e3e69336c77b75e301ce74fd0ffe7f Mon Sep 17 00:00:00 2001 From: David Katz Date: Wed, 20 Aug 2025 16:34:59 -0400 Subject: [PATCH 10/31] rm groq model --- crates/goose/src/providers/groq.rs | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/crates/goose/src/providers/groq.rs b/crates/goose/src/providers/groq.rs index 2bff2086eeec..ebd1578f77a3 100644 --- a/crates/goose/src/providers/groq.rs +++ b/crates/goose/src/providers/groq.rs @@ -33,9 +33,7 @@ pub struct GroqProvider { impl_provider_default!(GroqProvider); impl GroqProvider { - pub fn from_env(mut model: ModelConfig) -> Result { - // Set the default fast model for Groq - using llama 70b for better context - model.fast_model = Some("llama-3.3-70b-versatile".to_string()); + pub fn from_env(model: ModelConfig) -> Result { let config = crate::config::Config::global(); let api_key: String = config.get_secret("GROQ_API_KEY")?; From ee5f40b362bf072b2907e27763b3048e17f3dd91 Mon Sep 17 00:00:00 2001 From: David Katz Date: Wed, 20 Aug 2025 16:38:14 -0400 Subject: [PATCH 11/31] Reset python files --- examples/frontend_tools.py | 52 ++-- examples/mcp-wiki/src/mcp_wiki/__init__.py | 3 +- examples/mcp-wiki/src/mcp_wiki/__main__.py | 2 +- examples/mcp-wiki/src/mcp_wiki/server.py | 7 +- .../generate_leaderboard.py | 134 +++++------ .../calculate_final_scores_vibes.py | 33 +-- .../llm-judges/llm_judge.py | 136 +++++------ .../prepare_aggregate_metrics.py | 224 ++++++++---------- 8 files changed, 260 insertions(+), 331 deletions(-) diff --git a/examples/frontend_tools.py b/examples/frontend_tools.py index ce25db3b60a5..d8824f3a9828 100644 --- a/examples/frontend_tools.py +++ b/examples/frontend_tools.py @@ -117,7 +117,6 @@ def execute_calculator(args: Dict[str, Any]) -> List[Dict[str, Any]]: } ] - def get_tools() -> Dict[str, Any]: with httpx.Client() as client: response = client.get( @@ -155,21 +154,17 @@ def execute_enable_extension(args: Dict[str, Any]) -> List[Dict[str, Any]]: ) if add_response.status_code != 200: error_text = add_response.text - return [ - { - "type": "text", - "text": f"Error: Failed to enable extension: {error_text}", - "annotations": None, - } - ] - - return [ - { - "type": "text", - "text": f"Successfully enabled extension: {extension_name}", - "annotations": None, - } - ] + return [{ + "type": "text", + "text": f"Error: Failed to enable extension: {error_text}", + "annotations": None, + }] + + return [{ + "type": "text", + "text": f"Successfully enabled extension: {extension_name}", + "annotations": None, + }] def submit_tool_result(tool_id: str, result: List[Dict[str, Any]]) -> None: @@ -222,7 +217,7 @@ async def chat_loop() -> None: # Process the stream of responses async with client.stream( "POST", - f"{GOOSE_URL}/reply", # lock + f"{GOOSE_URL}/reply", # lock json=payload, headers={ "X-Secret-Key": SECRET_KEY, @@ -257,26 +252,25 @@ async def chat_loop() -> None: tool_call = content["toolCall"]["value"] print(f"\nTool Request: {tool_call}") - if tool_call["name"] == "calculator": + if tool_call['name'] == "calculator": print(f"Calculator: {tool_call}") # Execute the tool result = execute_calculator(tool_call["arguments"]) - elif tool_call["name"] == "enable_extension": + elif tool_call['name'] == "enable_extension": # to trigger this tool, use the instruction "use enable_extension tool with "fetch" extension name" print(f"Enabling fetch extension") - result = execute_enable_extension( - args={ - "type": "stdio", - "name": "fetch", - "cmd": "uvx", - "args": ["mcp-server-fetch"], - "timeout": 300, - "bundled": False, - } - ) + result = execute_enable_extension(args={ + "type": "stdio", + "name": "fetch", + "cmd": "uvx", + "args": ["mcp-server-fetch"], + "timeout": 300, + "bundled": False + }) listed_tools = get_tools() print(f"\nTools after enabling extension: {listed_tools}") + # Submit the result submit_tool_result(content["id"], result) diff --git a/examples/mcp-wiki/src/mcp_wiki/__init__.py b/examples/mcp-wiki/src/mcp_wiki/__init__.py index 1d60c32ebec6..20c873b5068f 100644 --- a/examples/mcp-wiki/src/mcp_wiki/__init__.py +++ b/examples/mcp-wiki/src/mcp_wiki/__init__.py @@ -1,7 +1,6 @@ import argparse from .server import mcp - def main(): """MCP Wiki: read Wikipedia articles and convert them to Markdown.""" parser = argparse.ArgumentParser( @@ -13,4 +12,4 @@ def main(): if __name__ == "__main__": - main() + main() \ No newline at end of file diff --git a/examples/mcp-wiki/src/mcp_wiki/__main__.py b/examples/mcp-wiki/src/mcp_wiki/__main__.py index cdb4dc1b54c5..8579afa979d7 100644 --- a/examples/mcp-wiki/src/mcp_wiki/__main__.py +++ b/examples/mcp-wiki/src/mcp_wiki/__main__.py @@ -1,3 +1,3 @@ from mcp_wiki import main -main() +main() \ No newline at end of file diff --git a/examples/mcp-wiki/src/mcp_wiki/server.py b/examples/mcp-wiki/src/mcp_wiki/server.py index 1936e2271bd7..33329ebb00e7 100644 --- a/examples/mcp-wiki/src/mcp_wiki/server.py +++ b/examples/mcp-wiki/src/mcp_wiki/server.py @@ -11,7 +11,6 @@ mcp = FastMCP("wiki") - @mcp.tool() def read_wikipedia_article(url: str) -> str: """ @@ -32,7 +31,7 @@ def read_wikipedia_article(url: str) -> str: raise McpError( ErrorData( INTERNAL_ERROR, - f"Failed to retrieve the article. HTTP status code: {response.status_code}", + f"Failed to retrieve the article. HTTP status code: {response.status_code}" ) ) @@ -43,7 +42,7 @@ def read_wikipedia_article(url: str) -> str: raise McpError( ErrorData( INVALID_PARAMS, - "Could not find the main content on the provided Wikipedia URL.", + "Could not find the main content on the provided Wikipedia URL." ) ) @@ -60,3 +59,5 @@ def read_wikipedia_article(url: str) -> str: except Exception as e: # Catch-all for any other unexpected errors raise McpError(ErrorData(INTERNAL_ERROR, f"Unexpected error: {str(e)}")) from e + + diff --git a/scripts/bench-postprocess-scripts/generate_leaderboard.py b/scripts/bench-postprocess-scripts/generate_leaderboard.py index 45227dadc59a..16292c727c45 100755 --- a/scripts/bench-postprocess-scripts/generate_leaderboard.py +++ b/scripts/bench-postprocess-scripts/generate_leaderboard.py @@ -23,7 +23,7 @@ def find_aggregate_metrics_files(benchmark_dir: Path) -> list: """Find all aggregate_metrics.csv files in model subdirectories.""" csv_files = [] - + # Look for model directories in the benchmark directory for model_dir in benchmark_dir.iterdir(): if model_dir.is_dir(): @@ -33,7 +33,7 @@ def find_aggregate_metrics_files(benchmark_dir: Path) -> list: csv_path = eval_results_dir / "aggregate_metrics.csv" if csv_path.exists(): csv_files.append(csv_path) - + return csv_files @@ -44,73 +44,67 @@ def process_csv_files(csv_files: list) -> tuple: 2. A leaderboard grouping by provider and model_name with averaged metrics """ selected_columns = [ - "provider", - "model_name", - "eval_suite", - "eval_name", - "total_tool_calls_mean", - "prompt_execution_time_mean", - "total_tokens_mean", - "score_mean", - "prompt_error_mean", - "server_error_mean", + 'provider', + 'model_name', + 'eval_suite', + 'eval_name', + 'total_tool_calls_mean', + 'prompt_execution_time_mean', + 'total_tokens_mean', + 'score_mean', + 'prompt_error_mean', + 'server_error_mean' ] - + all_data = [] - + for csv_file in csv_files: try: df = pd.read_csv(csv_file) - + # Check which selected columns are available missing_columns = [col for col in selected_columns if col not in df.columns] if missing_columns: print(f"Warning: {csv_file} is missing columns: {missing_columns}") - + # For missing columns, add them with NaN values for col in missing_columns: - df[col] = float("nan") - + df[col] = float('nan') + # Select only the columns we care about - df_subset = df[ - selected_columns - ].copy() # Create a copy to avoid SettingWithCopyWarning - + df_subset = df[selected_columns].copy() # Create a copy to avoid SettingWithCopyWarning + # Add model folder name as additional context model_folder = csv_file.parent.parent.name - df_subset["model_folder"] = model_folder - + df_subset['model_folder'] = model_folder + all_data.append(df_subset) - + except Exception as e: print(f"Error processing {csv_file}: {str(e)}") - + if not all_data: raise ValueError("No valid CSV files found with required columns") - + # Concatenate all dataframes to create a union union_df = pd.concat(all_data, ignore_index=True) - + # Create leaderboard by grouping and averaging numerical columns numeric_columns = [ - "total_tool_calls_mean", - "prompt_execution_time_mean", - "total_tokens_mean", - "score_mean", - "prompt_error_mean", - "server_error_mean", + 'total_tool_calls_mean', + 'prompt_execution_time_mean', + 'total_tokens_mean', + 'score_mean', + 'prompt_error_mean', + 'server_error_mean' ] - + # Group by provider and model_name, then calculate averages for numeric columns - leaderboard_df = ( - union_df.groupby(["provider", "model_name"])[numeric_columns] - .mean() - .reset_index() - ) - + leaderboard_df = union_df.groupby(['provider', 'model_name'])[numeric_columns].mean().reset_index() + # Sort by score_mean in descending order (highest scores first) - leaderboard_df = leaderboard_df.sort_values("score_mean", ascending=False) - + leaderboard_df = leaderboard_df.sort_values('score_mean', ascending=False) + return union_df, leaderboard_df @@ -122,75 +116,65 @@ def main(): "--benchmark-dir", type=str, required=True, - help="Path to the benchmark directory containing model subdirectories", + help="Path to the benchmark directory containing model subdirectories" ) parser.add_argument( "--union-output", type=str, default="all_metrics.csv", - help="Output filename for the union of all CSVs (default: all_metrics.csv)", + help="Output filename for the union of all CSVs (default: all_metrics.csv)" ) parser.add_argument( "--leaderboard-output", type=str, default="leaderboard.csv", - help="Output filename for the leaderboard (default: leaderboard.csv)", + help="Output filename for the leaderboard (default: leaderboard.csv)" ) - + args = parser.parse_args() - + benchmark_dir = Path(args.benchmark_dir) if not benchmark_dir.exists() or not benchmark_dir.is_dir(): - print( - f"Error: Benchmark directory {benchmark_dir} does not exist or is not a directory" - ) + print(f"Error: Benchmark directory {benchmark_dir} does not exist or is not a directory") sys.exit(1) - + try: # Find all aggregate_metrics.csv files in model subdirectories csv_files = find_aggregate_metrics_files(benchmark_dir) - + if not csv_files: - print( - f"No aggregate_metrics.csv files found in any model directory under {benchmark_dir}" - ) + print(f"No aggregate_metrics.csv files found in any model directory under {benchmark_dir}") sys.exit(1) - - print( - f"Found {len(csv_files)} aggregate_metrics.csv files in model directories" - ) - + + print(f"Found {len(csv_files)} aggregate_metrics.csv files in model directories") + # Process and create the union and leaderboard dataframes union_df, leaderboard_df = process_csv_files(csv_files) - + # Save the union CSV to the benchmark directory union_output_path = benchmark_dir / args.union_output union_df.to_csv(union_output_path, index=False) print(f"Union CSV with all metrics saved to: {union_output_path}") - + # Save the leaderboard CSV to the benchmark directory leaderboard_output_path = benchmark_dir / args.leaderboard_output leaderboard_df.to_csv(leaderboard_output_path, index=False) - print( - f"Leaderboard CSV with averaged metrics saved to: {leaderboard_output_path}" - ) - + print(f"Leaderboard CSV with averaged metrics saved to: {leaderboard_output_path}") + # Print a summary of the leaderboard print("\nLeaderboard Summary:") - pd.set_option("display.max_columns", None) # Show all columns + pd.set_option('display.max_columns', None) # Show all columns print(leaderboard_df.to_string(index=False)) - + # Highlight models with server errors - if "server_error_mean" in leaderboard_df.columns: - models_with_errors = leaderboard_df[leaderboard_df["server_error_mean"] > 0] + if 'server_error_mean' in leaderboard_df.columns: + models_with_errors = leaderboard_df[leaderboard_df['server_error_mean'] > 0] if not models_with_errors.empty: print("\nWARNING - Models with server errors detected:") for _, row in models_with_errors.iterrows(): - print( - f" * {row['provider']} {row['model_name']} - {row['server_error_mean'] * 100:.1f}% of evaluations had server errors" - ) + print(f" * {row['provider']} {row['model_name']} - {row['server_error_mean']*100:.1f}% of evaluations had server errors") print("\nThese models may need to be re-run to get accurate results.") - + except Exception as e: print(f"Error: {str(e)}") sys.exit(1) diff --git a/scripts/bench-postprocess-scripts/llm-judges/calculate_final_scores_vibes.py b/scripts/bench-postprocess-scripts/llm-judges/calculate_final_scores_vibes.py index 17c4259775d7..261fbe52832d 100755 --- a/scripts/bench-postprocess-scripts/llm-judges/calculate_final_scores_vibes.py +++ b/scripts/bench-postprocess-scripts/llm-judges/calculate_final_scores_vibes.py @@ -28,14 +28,14 @@ def calculate_score(eval_name, metrics): llm_judge_score = get_metric_value(metrics, "llm_judge_score") used_fetch_tool = get_metric_value(metrics, "used_fetch_tool") valid_markdown_format = get_metric_value(metrics, "valid_markdown_format") - + if llm_judge_score is None: raise ValueError("llm_judge_score not found in metrics") - + # Convert boolean metrics to 0/1 if needed used_fetch_tool = 1.0 if used_fetch_tool else 0.0 valid_markdown_format = 1.0 if valid_markdown_format else 0.0 - + if eval_name == "blog_summary": # max score is 4.0 as llm_judge_score is between [0,2] and used_fetch_tool/valid_markedown_format have values [0,1] score = (llm_judge_score + used_fetch_tool + valid_markdown_format) / 4.0 @@ -43,7 +43,7 @@ def calculate_score(eval_name, metrics): score = (llm_judge_score + valid_markdown_format + used_fetch_tool) / 4.0 else: raise ValueError(f"Unknown evaluation type: {eval_name}") - + return score @@ -51,31 +51,34 @@ def main(): if len(sys.argv) != 2: print("Usage: calculate_final_score.py ") sys.exit(1) - + eval_name = sys.argv[1] - + # Load eval results from current directory eval_results_path = Path("eval-results.json") if not eval_results_path.exists(): print(f"Error: eval-results.json not found in current directory") sys.exit(1) - - with open(eval_results_path, "r") as f: + + with open(eval_results_path, 'r') as f: eval_results = json.load(f) - + try: # Calculate the final score score = calculate_score(eval_name, eval_results["metrics"]) - + # Add the score metric - eval_results["metrics"].append(["score", {"Float": score}]) - + eval_results["metrics"].append([ + "score", + {"Float": score} + ]) + # Save updated results - with open(eval_results_path, "w") as f: + with open(eval_results_path, 'w') as f: json.dump(eval_results, f, indent=2) - + print(f"Successfully added final score: {score}") - + except Exception as e: print(f"Error calculating final score: {str(e)}") sys.exit(1) diff --git a/scripts/bench-postprocess-scripts/llm-judges/llm_judge.py b/scripts/bench-postprocess-scripts/llm-judges/llm_judge.py index e32bcf7744f7..2b22dc24489c 100755 --- a/scripts/bench-postprocess-scripts/llm-judges/llm_judge.py +++ b/scripts/bench-postprocess-scripts/llm-judges/llm_judge.py @@ -8,7 +8,7 @@ Usage: python llm_judge.py [--rubric-max-score N] [--prompt-file PATH] - + Arguments: output_file: Name of the file containing the output to evaluate (e.g., blog_summary_output.txt) --rubric-max-score: Maximum score for the rubric (default: 2) @@ -33,15 +33,15 @@ def evaluate_with_openai(prompt: str, text: str, rubric_max_score: int = 2) -> float: """Evaluate response using OpenAI's API. - + Args: prompt: System prompt for evaluation text: Text to evaluate rubric_max_score: Maximum score for the rubric (default: 2.0) - + Returns: float: Evaluation score (0 to rubric_max_score) - + Raises: ValueError: If OPENAI_API_KEY environment variable is not set """ @@ -49,13 +49,11 @@ def evaluate_with_openai(prompt: str, text: str, rubric_max_score: int = 2) -> f api_key = os.getenv("OPENAI_API_KEY") if not api_key: print("No OpenAI API key found!") - raise ValueError( - "OPENAI_API_KEY environment variable is not set, but is needed to run this evaluation." - ) - + raise ValueError("OPENAI_API_KEY environment variable is not set, but is needed to run this evaluation.") + try: client = OpenAI(api_key=api_key) - + # Append output instructions to system prompt output_instructions = f""" Output Instructions: @@ -70,23 +68,25 @@ def evaluate_with_openai(prompt: str, text: str, rubric_max_score: int = 2) -> f - Do not include any additional text before or after the JSON - Return only the raw JSON object - The score must be an integer between 0 and {rubric_max_score}""" - + input_prompt = f"{prompt} {output_instructions}\nResponse to evaluate: {text}" - + # Run the chat completion 3 times and collect scores scores = [] for i in range(3): max_retries = 5 retry_count = 0 - + while retry_count < max_retries: try: response = client.chat.completions.create( model="gpt-4o", - messages=[{"role": "user", "content": input_prompt}], - temperature=0.9, + messages=[ + {"role": "user", "content": input_prompt} + ], + temperature=0.9 ) - + # Extract and parse JSON from response response_text = response.choices[0].message.content.strip() try: @@ -94,18 +94,14 @@ def evaluate_with_openai(prompt: str, text: str, rubric_max_score: int = 2) -> f score = float(evaluation.get("score", 0.0)) score = max(0.0, min(score, rubric_max_score)) scores.append(score) - print(f"Run {i + 1} score: {score}") + print(f"Run {i+1} score: {score}") break # Successfully parsed, exit retry loop except (json.JSONDecodeError, ValueError) as e: retry_count += 1 - print( - f"Error parsing OpenAI response as JSON (attempt {retry_count}/{max_retries}): {str(e)}" - ) + print(f"Error parsing OpenAI response as JSON (attempt {retry_count}/{max_retries}): {str(e)}") print(f"Response text: {response_text}") if retry_count == max_retries: - raise ValueError( - f"Failed to parse OpenAI evaluation response after {max_retries} attempts: {str(e)}" - ) + raise ValueError(f"Failed to parse OpenAI evaluation response after {max_retries} attempts: {str(e)}") print("Retrying...") time.sleep(1) # Wait 1 second before retrying continue @@ -113,24 +109,26 @@ def evaluate_with_openai(prompt: str, text: str, rubric_max_score: int = 2) -> f # For other exceptions (API errors, etc.), raise immediately print(f"API error: {str(e)}") raise - + # Count occurrences of each score score_counts = Counter(scores) - + # If there's no single most common score (all scores are different), run one more time if len(scores) == 3 and max(score_counts.values()) == 1: print("No majority score found. Running tie-breaker...") max_retries = 5 retry_count = 0 - + while retry_count < max_retries: try: response = client.chat.completions.create( model="gpt-4o", - messages=[{"role": "user", "content": input_prompt}], - temperature=0.9, + messages=[ + {"role": "user", "content": input_prompt} + ], + temperature=0.9 ) - + response_text = response.choices[0].message.content.strip() try: evaluation = json.loads(response_text) @@ -142,14 +140,10 @@ def evaluate_with_openai(prompt: str, text: str, rubric_max_score: int = 2) -> f break # Successfully parsed, exit retry loop except (json.JSONDecodeError, ValueError) as e: retry_count += 1 - print( - f"Error parsing tie-breaker response as JSON (attempt {retry_count}/{max_retries}): {str(e)}" - ) + print(f"Error parsing tie-breaker response as JSON (attempt {retry_count}/{max_retries}): {str(e)}") print(f"Response text: {response_text}") if retry_count == max_retries: - raise ValueError( - f"Failed to parse tie-breaker response after {max_retries} attempts: {str(e)}" - ) + raise ValueError(f"Failed to parse tie-breaker response after {max_retries} attempts: {str(e)}") print("Retrying tie-breaker...") time.sleep(1) # Wait 1 second before retrying continue @@ -157,14 +151,12 @@ def evaluate_with_openai(prompt: str, text: str, rubric_max_score: int = 2) -> f # For other exceptions (API errors, etc.), raise immediately print(f"API error in tie-breaker: {str(e)}") raise - + # Get the most common score most_common_score = score_counts.most_common(1)[0][0] - print( - f"Most common score: {most_common_score} (occurred {score_counts[most_common_score]} times)" - ) + print(f"Most common score: {most_common_score} (occurred {score_counts[most_common_score]} times)") return most_common_score - + except Exception as e: if "OPENAI_API_KEY" in str(e): raise # Re-raise API key errors @@ -177,8 +169,8 @@ def load_eval_results(working_dir: Path) -> Dict[str, Any]: eval_results_path = working_dir / "eval-results.json" if not eval_results_path.exists(): raise FileNotFoundError(f"eval-results.json not found in {working_dir}") - - with open(eval_results_path, "r") as f: + + with open(eval_results_path, 'r') as f: return json.load(f) @@ -187,22 +179,22 @@ def load_output_file(working_dir: Path, output_file: str) -> str: output_path = working_dir / output_file if not output_path.exists(): raise FileNotFoundError(f"Output file not found: {output_path}") - - with open(output_path, "r") as f: + + with open(output_path, 'r') as f: return f.read().strip() def load_evaluation_prompt(working_dir: Path) -> str: """Load the evaluation prompt from a file or use a default. - + This function looks for a prompt.txt file in the working directory. If not found, it returns a default evaluation prompt. """ prompt_file = working_dir / "prompt.txt" if prompt_file.exists(): - with open(prompt_file, "r") as f: + with open(prompt_file, 'r') as f: return f.read().strip() - + # Default evaluation prompt return """You are an expert evaluator assessing the quality of AI responses. Evaluate the response based on the following criteria: @@ -218,58 +210,46 @@ def load_evaluation_prompt(working_dir: Path) -> str: def main(): - parser = argparse.ArgumentParser( - description="LLM Judge post-processing script for Goose benchmarks" - ) - parser.add_argument( - "output_file", - type=str, - help="Name of the output file to evaluate (e.g., blog_summary_output.txt)", - ) - parser.add_argument( - "--rubric-max-score", - type=int, - default=2, - help="Maximum score for the rubric (default: 2)", - ) - parser.add_argument( - "--prompt-file", type=str, help="Path to custom evaluation prompt file" - ) - + parser = argparse.ArgumentParser(description="LLM Judge post-processing script for Goose benchmarks") + parser.add_argument("output_file", type=str, help="Name of the output file to evaluate (e.g., blog_summary_output.txt)") + parser.add_argument("--rubric-max-score", type=int, default=2, help="Maximum score for the rubric (default: 2)") + parser.add_argument("--prompt-file", type=str, help="Path to custom evaluation prompt file") + args = parser.parse_args() - + # Use current working directory working_dir = Path.cwd() - + try: # Load eval results eval_results = load_eval_results(working_dir) - + # Load the output file to evaluate response_text = load_output_file(working_dir, args.output_file) - + # Load evaluation prompt if args.prompt_file: - with open(args.prompt_file, "r") as f: + with open(args.prompt_file, 'r') as f: evaluation_prompt = f.read().strip() else: evaluation_prompt = load_evaluation_prompt(working_dir) - + # Evaluate with OpenAI - score = evaluate_with_openai( - evaluation_prompt, response_text, args.rubric_max_score - ) - + score = evaluate_with_openai(evaluation_prompt, response_text, args.rubric_max_score) + # Update eval results with the score - eval_results["metrics"].append(["llm_judge_score", {"Float": score}]) + eval_results["metrics"].append([ + "llm_judge_score", + {"Float": score} + ]) # Save updated results eval_results_path = working_dir / "eval-results.json" - with open(eval_results_path, "w") as f: + with open(eval_results_path, 'w') as f: json.dump(eval_results, f, indent=2) - + print(f"Successfully updated eval-results.json with LLM judge score: {score}") - + except Exception as e: print(f"Error: {str(e)}") sys.exit(1) diff --git a/scripts/bench-postprocess-scripts/prepare_aggregate_metrics.py b/scripts/bench-postprocess-scripts/prepare_aggregate_metrics.py index ffdd5e09ebd5..399172918968 100755 --- a/scripts/bench-postprocess-scripts/prepare_aggregate_metrics.py +++ b/scripts/bench-postprocess-scripts/prepare_aggregate_metrics.py @@ -21,74 +21,60 @@ from pathlib import Path import sys - def extract_provider_model(model_dir): """Extract provider and model name from directory name.""" dir_name = model_dir.name - parts = dir_name.split("-") - + parts = dir_name.split('-') + if len(parts) > 1: model_name = parts[-1] # Last part is the model name - provider = "-".join(parts[:-1]) # Everything else is the provider + provider = '-'.join(parts[:-1]) # Everything else is the provider else: model_name = dir_name provider = "unknown" - + return provider, model_name - def find_eval_results_files(model_dir): """Find all eval-results.json files in a model directory.""" return list(model_dir.glob("**/eval-results.json")) - def find_session_files(model_dir): """Find all session jsonl files in a model directory.""" return list(model_dir.glob("**/*.jsonl")) - def check_for_errors_in_session(session_file): """Check if a session file contains server errors.""" try: error_found = False error_messages = [] - - with open(session_file, "r") as f: + + with open(session_file, 'r') as f: for line in f: try: message_obj = json.loads(line.strip()) # Check for error messages in the content - if "content" in message_obj and isinstance( - message_obj["content"], list - ): - for content_item in message_obj["content"]: - if ( - isinstance(content_item, dict) - and "text" in content_item - ): - text = content_item["text"] - if ( - "Server error" in text - or "error_code" in text - or "TEMPORARILY_UNAVAILABLE" in text - ): + if 'content' in message_obj and isinstance(message_obj['content'], list): + for content_item in message_obj['content']: + if isinstance(content_item, dict) and 'text' in content_item: + text = content_item['text'] + if 'Server error' in text or 'error_code' in text or 'TEMPORARILY_UNAVAILABLE' in text: error_found = True error_messages.append(text) except json.JSONDecodeError: continue - + return error_found, error_messages except Exception as e: print(f"Error checking session file {session_file}: {str(e)}") return False, [] - def extract_metrics_from_eval_file(eval_file, provider, model_name, session_files): """Extract metrics from an eval-results.json file.""" try: - with open(eval_file, "r") as f: + with open(eval_file, 'r') as f: data = json.load(f) - + # Extract directory structure to determine eval suite and name path_parts = eval_file.parts run_index = -1 @@ -96,154 +82,145 @@ def extract_metrics_from_eval_file(eval_file, provider, model_name, session_file if part.startswith("run-"): run_index = i break - + if run_index == -1 or run_index + 2 >= len(path_parts): print(f"Warning: Could not determine eval suite and name from {eval_file}") return None - - run_number = path_parts[run_index].split("-")[1] # Extract "0" from "run-0" + + run_number = path_parts[run_index].split('-')[1] # Extract "0" from "run-0" eval_suite = path_parts[run_index + 1] # Directory after run-N eval_name = path_parts[run_index + 2] # Directory after eval_suite - + # Create a row with basic identification row = { - "provider": provider, - "model_name": model_name, - "eval_suite": eval_suite, - "eval_name": eval_name, - "run": run_number, + 'provider': provider, + 'model_name': model_name, + 'eval_suite': eval_suite, + 'eval_name': eval_name, + 'run': run_number } - + # Check for server errors in session files for this evaluation eval_dir = eval_file.parent related_session_files = [sf for sf in session_files if eval_dir in sf.parents] - + server_error_found = False for session_file in related_session_files: error_found, _ = check_for_errors_in_session(session_file) if error_found: server_error_found = True break - + # Add server error flag - row["server_error"] = 1 if server_error_found else 0 - + row['server_error'] = 1 if server_error_found else 0 + # Extract all metrics (flatten the JSON structure) if isinstance(data, dict): metrics = {} - + # Extract top-level metrics for key, value in data.items(): if isinstance(value, (int, float)) and not isinstance(value, bool): metrics[key] = value - + # Look for nested metrics structure (list of [name, value] pairs) - if "metrics" in data and isinstance(data["metrics"], list): - for metric_item in data["metrics"]: + if 'metrics' in data and isinstance(data['metrics'], list): + for metric_item in data['metrics']: if isinstance(metric_item, list) and len(metric_item) == 2: metric_name = metric_item[0] metric_value = metric_item[1] - + # Handle different value formats if isinstance(metric_value, dict): - if "Integer" in metric_value: - metrics[metric_name] = int(metric_value["Integer"]) - elif "Float" in metric_value: - metrics[metric_name] = float(metric_value["Float"]) - elif "Bool" in metric_value: - metrics[metric_name] = 1 if metric_value["Bool"] else 0 + if 'Integer' in metric_value: + metrics[metric_name] = int(metric_value['Integer']) + elif 'Float' in metric_value: + metrics[metric_name] = float(metric_value['Float']) + elif 'Bool' in metric_value: + metrics[metric_name] = 1 if metric_value['Bool'] else 0 # Skip string values for aggregation - elif isinstance(metric_value, (int, float)) and not isinstance( - metric_value, bool - ): + elif isinstance(metric_value, (int, float)) and not isinstance(metric_value, bool): metrics[metric_name] = metric_value elif isinstance(metric_value, bool): metrics[metric_name] = 1 if metric_value else 0 - + # Look for metrics in other common locations - for metric_location in ["metrics", "result", "evaluation"]: + for metric_location in ['metrics', 'result', 'evaluation']: if metric_location in data and isinstance(data[metric_location], dict): for key, value in data[metric_location].items(): - if isinstance(value, (int, float)) and not isinstance( - value, bool - ): + if isinstance(value, (int, float)) and not isinstance(value, bool): metrics[key] = value elif isinstance(value, bool): metrics[key] = 1 if value else 0 - + # Add all metrics to the row row.update(metrics) - + # Ensure a score is present (if not, add a placeholder) - if "score" not in row: + if 'score' not in row: # Try to use existing fields to calculate a score if server_error_found: - row["score"] = 0 # Failed runs get a zero score + row['score'] = 0 # Failed runs get a zero score else: # Set a default based on presence of "success" fields for key in row: - if "success" in key.lower() and isinstance( - row[key], (int, float) - ): - row["score"] = row[key] + if 'success' in key.lower() and isinstance(row[key], (int, float)): + row['score'] = row[key] break else: # No success field found, mark as NaN - row["score"] = float("nan") - + row['score'] = float('nan') + return row else: print(f"Warning: Unexpected format in {eval_file}") return None - + except Exception as e: print(f"Error processing {eval_file}: {str(e)}") return None - def process_model_directory(model_dir): """Process a model directory to create aggregate_metrics.csv.""" provider, model_name = extract_provider_model(model_dir) - + # Find all eval results files eval_files = find_eval_results_files(model_dir) if not eval_files: print(f"No eval-results.json files found in {model_dir}") return False - + # Find all session files for error checking session_files = find_session_files(model_dir) - + # Extract metrics from each eval file rows = [] for eval_file in eval_files: - row = extract_metrics_from_eval_file( - eval_file, provider, model_name, session_files - ) + row = extract_metrics_from_eval_file(eval_file, provider, model_name, session_files) if row is not None: rows.append(row) - + if not rows: print(f"No valid metrics extracted from {model_dir}") return False - + # Create a dataframe from all rows combined_df = pd.DataFrame(rows) - + # Calculate aggregates for numeric columns, grouped by eval_suite, eval_name - numeric_cols = combined_df.select_dtypes(include=["number"]).columns.tolist() + numeric_cols = combined_df.select_dtypes(include=['number']).columns.tolist() # Exclude the run column from aggregation - if "run" in numeric_cols: - numeric_cols.remove("run") - + if 'run' in numeric_cols: + numeric_cols.remove('run') + # Group by provider, model_name, eval_suite, eval_name and calculate mean for numeric columns - group_by_cols = ["provider", "model_name", "eval_suite", "eval_name"] - agg_dict = {col: "mean" for col in numeric_cols} - + group_by_cols = ['provider', 'model_name', 'eval_suite', 'eval_name'] + agg_dict = {col: 'mean' for col in numeric_cols} + # Only perform aggregation if we have numeric columns if numeric_cols: aggregate_df = combined_df.groupby(group_by_cols).agg(agg_dict).reset_index() - + # Rename columns to add _mean suffix for the averaged metrics for col in numeric_cols: aggregate_df = aggregate_df.rename(columns={col: f"{col}_mean"}) @@ -251,41 +228,38 @@ def process_model_directory(model_dir): print(f"Warning: No numeric metrics found in {model_dir}") # Create a minimal dataframe with just the grouping columns aggregate_df = combined_df[group_by_cols].drop_duplicates() - + # Make sure we have prompt_execution_time_mean and prompt_error_mean columns # These are expected by the generate_leaderboard.py script - if "prompt_execution_time_mean" not in aggregate_df.columns: - aggregate_df["prompt_execution_time_mean"] = float("nan") - - if "prompt_error_mean" not in aggregate_df.columns: - aggregate_df["prompt_error_mean"] = float("nan") - + if 'prompt_execution_time_mean' not in aggregate_df.columns: + aggregate_df['prompt_execution_time_mean'] = float('nan') + + if 'prompt_error_mean' not in aggregate_df.columns: + aggregate_df['prompt_error_mean'] = float('nan') + # Add server_error_mean column if not present - if "server_error_mean" not in aggregate_df.columns: - aggregate_df["server_error_mean"] = 0.0 - + if 'server_error_mean' not in aggregate_df.columns: + aggregate_df['server_error_mean'] = 0.0 + # Create eval-results directory eval_results_dir = model_dir / "eval-results" eval_results_dir.mkdir(exist_ok=True) - + # Save to CSV csv_path = eval_results_dir / "aggregate_metrics.csv" aggregate_df.to_csv(csv_path, index=False) - + # Count number of evaluations that had server errors - if "server_error_mean" in aggregate_df.columns: - error_count = len(aggregate_df[aggregate_df["server_error_mean"] > 0]) + if 'server_error_mean' in aggregate_df.columns: + error_count = len(aggregate_df[aggregate_df['server_error_mean'] > 0]) total_count = len(aggregate_df) - print( - f"Saved aggregate metrics to {csv_path} with {len(aggregate_df)} rows " - + f"({error_count}/{total_count} evals had server errors)" - ) + print(f"Saved aggregate metrics to {csv_path} with {len(aggregate_df)} rows " + + f"({error_count}/{total_count} evals had server errors)") else: print(f"Saved aggregate metrics to {csv_path} with {len(aggregate_df)} rows") - + return True - def main(): parser = argparse.ArgumentParser( description="Prepare aggregate_metrics.csv files from eval-results.json files with error detection" @@ -294,39 +268,33 @@ def main(): "--benchmark-dir", type=str, required=True, - help="Path to the benchmark directory containing model subdirectories", + help="Path to the benchmark directory containing model subdirectories" ) - + args = parser.parse_args() - + # Convert path to Path object and validate it exists benchmark_dir = Path(args.benchmark_dir) if not benchmark_dir.exists() or not benchmark_dir.is_dir(): - print( - f"Error: Benchmark directory {benchmark_dir} does not exist or is not a directory" - ) + print(f"Error: Benchmark directory {benchmark_dir} does not exist or is not a directory") sys.exit(1) - + success_count = 0 - + # Process each model directory for model_dir in benchmark_dir.iterdir(): - if model_dir.is_dir() and not model_dir.name.startswith("."): + if model_dir.is_dir() and not model_dir.name.startswith('.'): if process_model_directory(model_dir): success_count += 1 - + if success_count == 0: print("No aggregate_metrics.csv files were created") sys.exit(1) - - print( - f"Successfully created aggregate_metrics.csv files for {success_count} model directories" - ) + + print(f"Successfully created aggregate_metrics.csv files for {success_count} model directories") print("You can now run generate_leaderboard.py to create the final leaderboard.") - print( - "Note: The server_error_mean column indicates the average rate of server errors across evaluations." - ) - + print("Note: The server_error_mean column indicates the average rate of server errors across evaluations.") if __name__ == "__main__": main() + From 0c85d4d31bc2ce97ee42aca4118d9e8704893b06 Mon Sep 17 00:00:00 2001 From: David Katz Date: Wed, 20 Aug 2025 17:00:50 -0400 Subject: [PATCH 12/31] should build now --- crates/goose/src/context_mgmt/auto_compact.rs | 3 +- crates/goose/src/context_mgmt/summarize.rs | 3 +- crates/goose/src/model.rs | 29 +++++++++---------- .../goose/src/permission/permission_judge.rs | 3 +- crates/goose/src/providers/factory.rs | 1 + crates/goose/src/providers/lead_worker.rs | 2 ++ crates/goose/src/providers/testprovider.rs | 1 + crates/goose/src/scheduler.rs | 3 +- 8 files changed, 26 insertions(+), 19 deletions(-) diff --git a/crates/goose/src/context_mgmt/auto_compact.rs b/crates/goose/src/context_mgmt/auto_compact.rs index 062004650b69..dd8aaa04b478 100644 --- a/crates/goose/src/context_mgmt/auto_compact.rs +++ b/crates/goose/src/context_mgmt/auto_compact.rs @@ -221,8 +221,9 @@ mod tests { self.model_config.clone() } - async fn complete( + async fn complete_with_model( &self, + _model: &str, _system: &str, _messages: &[Message], _tools: &[Tool], diff --git a/crates/goose/src/context_mgmt/summarize.rs b/crates/goose/src/context_mgmt/summarize.rs index 3d53fcc4ed85..50513d4b5b1a 100644 --- a/crates/goose/src/context_mgmt/summarize.rs +++ b/crates/goose/src/context_mgmt/summarize.rs @@ -87,8 +87,9 @@ mod tests { self.model_config.clone() } - async fn complete( + async fn complete_with_model( &self, + _model: &str, _system: &str, _messages: &[Message], _tools: &[Tool], diff --git a/crates/goose/src/model.rs b/crates/goose/src/model.rs index d81ef6a81732..cff2239ff067 100644 --- a/crates/goose/src/model.rs +++ b/crates/goose/src/model.rs @@ -254,23 +254,22 @@ impl ModelConfig { } pub fn context_limit(&self) -> usize { - // If we have a fast_model, use the minimum of both limits + // If we have an explicit context limit set, use it + if let Some(limit) = self.context_limit { + return limit; + } + + // Otherwise, get the model's default limit + let main_limit = Self::get_model_specific_limit(&self.model_name) + .unwrap_or(DEFAULT_CONTEXT_LIMIT); + + // If we have a fast_model, also check its limit and use the minimum if let Some(fast_model) = &self.fast_model { - let main_limit = self - .context_limit - .or_else(|| Self::get_model_specific_limit(&self.model_name)); - let fast_limit = Self::get_model_specific_limit(fast_model); - - match (main_limit, fast_limit) { - (Some(m), Some(f)) => m.min(f), - (Some(m), None) => m, - (None, Some(f)) => f, - (None, None) => DEFAULT_CONTEXT_LIMIT, - } + let fast_limit = Self::get_model_specific_limit(fast_model) + .unwrap_or(DEFAULT_CONTEXT_LIMIT); + main_limit.min(fast_limit) } else { - self.context_limit.unwrap_or_else(|| { - Self::get_model_specific_limit(&self.model_name).unwrap_or(DEFAULT_CONTEXT_LIMIT) - }) + main_limit } } diff --git a/crates/goose/src/permission/permission_judge.rs b/crates/goose/src/permission/permission_judge.rs index 71438d38cb4c..c438272b5b0f 100644 --- a/crates/goose/src/permission/permission_judge.rs +++ b/crates/goose/src/permission/permission_judge.rs @@ -292,8 +292,9 @@ mod tests { self.model_config.clone() } - async fn complete( + async fn complete_with_model( &self, + _model: &str, _system: &str, _messages: &[Message], _tools: &[Tool], diff --git a/crates/goose/src/providers/factory.rs b/crates/goose/src/providers/factory.rs index 7b43806b1ed0..bd0013a3c372 100644 --- a/crates/goose/src/providers/factory.rs +++ b/crates/goose/src/providers/factory.rs @@ -202,6 +202,7 @@ mod tests { async fn complete_with_model( &self, + _model: &str, _system: &str, _messages: &[Message], _tools: &[Tool], diff --git a/crates/goose/src/providers/lead_worker.rs b/crates/goose/src/providers/lead_worker.rs index c15360cfc076..d0e57e4c43c0 100644 --- a/crates/goose/src/providers/lead_worker.rs +++ b/crates/goose/src/providers/lead_worker.rs @@ -478,6 +478,7 @@ mod tests { async fn complete_with_model( &self, + _model: &str, _system: &str, _messages: &[Message], _tools: &[Tool], @@ -638,6 +639,7 @@ mod tests { async fn complete_with_model( &self, + _model: &str, _system: &str, _messages: &[Message], _tools: &[Tool], diff --git a/crates/goose/src/providers/testprovider.rs b/crates/goose/src/providers/testprovider.rs index ccef229ea0ef..8c9a0cbe2cb3 100644 --- a/crates/goose/src/providers/testprovider.rs +++ b/crates/goose/src/providers/testprovider.rs @@ -191,6 +191,7 @@ mod tests { async fn complete_with_model( &self, + _model: &str, _system: &str, _messages: &[Message], _tools: &[Tool], diff --git a/crates/goose/src/scheduler.rs b/crates/goose/src/scheduler.rs index 0116e203fb30..e6c74cb0899f 100644 --- a/crates/goose/src/scheduler.rs +++ b/crates/goose/src/scheduler.rs @@ -1390,8 +1390,9 @@ mod tests { self.model_config.clone() } - async fn complete( + async fn complete_with_model( &self, + _model: &str, _system: &str, _messages: &[Message], _tools: &[Tool], From 3637a1a3483e74d1b8969dfdd63123ed7cea7ca8 Mon Sep 17 00:00:00 2001 From: David Katz Date: Thu, 21 Aug 2025 00:00:13 -0400 Subject: [PATCH 13/31] with fast abstraction + env set --- crates/goose/src/model.rs | 12 ++++++++++++ crates/goose/src/providers/anthropic.rs | 6 +++--- crates/goose/src/providers/google.rs | 4 ++-- crates/goose/src/providers/openai.rs | 6 +++--- 4 files changed, 20 insertions(+), 8 deletions(-) diff --git a/crates/goose/src/model.rs b/crates/goose/src/model.rs index cff2239ff067..0af02a21390b 100644 --- a/crates/goose/src/model.rs +++ b/crates/goose/src/model.rs @@ -253,6 +253,18 @@ impl ModelConfig { self } + pub fn with_fast(mut self, fast_model: String) -> Self { + self.fast_model = Some(fast_model); + self + } + + pub fn use_fast_model(mut self) -> Self { + if let Some(fast_model) = self.fast_model.clone() { + self.model_name = fast_model; + } + self + } + pub fn context_limit(&self) -> usize { // If we have an explicit context limit set, use it if let Some(limit) = self.context_limit { diff --git a/crates/goose/src/providers/anthropic.rs b/crates/goose/src/providers/anthropic.rs index 85bd2c7e13f2..dc16764a7d48 100644 --- a/crates/goose/src/providers/anthropic.rs +++ b/crates/goose/src/providers/anthropic.rs @@ -23,6 +23,7 @@ use crate::providers::retry::ProviderRetry; use rmcp::model::Tool; const ANTHROPIC_DEFAULT_MODEL: &str = "claude-sonnet-4-0"; +const ANTHROPIC_DEFAULT_FAST_MODEL: &str = "claude-3-5-haiku-latest"; const ANTHROPIC_KNOWN_MODELS: &[&str] = &[ "claude-sonnet-4-0", "claude-sonnet-4-20250514", @@ -49,9 +50,8 @@ pub struct AnthropicProvider { impl_provider_default!(AnthropicProvider); impl AnthropicProvider { - pub fn from_env(mut model: ModelConfig) -> Result { - // Set the default fast model for Anthropic - model.fast_model = Some("claude-3-5-haiku-latest".to_string()); + pub fn from_env(model: ModelConfig) -> Result { + let model = model.with_fast(ANTHROPIC_DEFAULT_FAST_MODEL.to_string()); let config = crate::config::Config::global(); let api_key: String = config.get_secret("ANTHROPIC_API_KEY")?; diff --git a/crates/goose/src/providers/google.rs b/crates/goose/src/providers/google.rs index 1816cf5728de..d6e2bcbbc53a 100644 --- a/crates/goose/src/providers/google.rs +++ b/crates/goose/src/providers/google.rs @@ -14,6 +14,7 @@ use serde_json::Value; pub const GOOGLE_API_HOST: &str = "https://generativelanguage.googleapis.com"; pub const GOOGLE_DEFAULT_MODEL: &str = "gemini-2.5-flash"; +pub const GOOGLE_DEFAULT_FAST_MODEL: &str = "gemini-1.5-flash"; pub const GOOGLE_KNOWN_MODELS: &[&str] = &[ // Gemini 2.5 models (latest generation) "gemini-2.5-pro", @@ -55,8 +56,7 @@ impl_provider_default!(GoogleProvider); impl GoogleProvider { pub fn from_env(mut model: ModelConfig) -> Result { - // Set the default fast model for Google - using Gemini Flash - model.fast_model = Some("gemini-1.5-flash".to_string()); + let model = model.with_fast(GOOGLE_DEFAULT_FAST_MODEL.to_string()); let config = crate::config::Config::global(); let api_key: String = config.get_secret("GOOGLE_API_KEY")?; diff --git a/crates/goose/src/providers/openai.rs b/crates/goose/src/providers/openai.rs index 954bd35a094b..fff09573e55b 100644 --- a/crates/goose/src/providers/openai.rs +++ b/crates/goose/src/providers/openai.rs @@ -29,6 +29,7 @@ use crate::providers::formats::openai::response_to_streaming_message; use rmcp::model::Tool; pub const OPEN_AI_DEFAULT_MODEL: &str = "gpt-4o"; +pub const OPEN_AI_DEFAULT_FAST_MODEL: &str = "gpt-4o-mini"; pub const OPEN_AI_KNOWN_MODELS: &[(&str, usize)] = &[ ("gpt-4o", 128_000), ("gpt-4o-mini", 128_000), @@ -58,9 +59,8 @@ pub struct OpenAiProvider { impl_provider_default!(OpenAiProvider); impl OpenAiProvider { - pub fn from_env(mut model: ModelConfig) -> Result { - // Set the default fast model for OpenAI - model.fast_model = Some("gpt-4o-mini".to_string()); + pub fn from_env(model: ModelConfig) -> Result { + let model = model.with_fast(OPEN_AI_DEFAULT_FAST_MODEL.to_string()); let config = crate::config::Config::global(); let api_key: String = config.get_secret("OPENAI_API_KEY")?; From 5528a29d1958ff8b9dfea1f25df10b6e0d457fb6 Mon Sep 17 00:00:00 2001 From: David Katz Date: Thu, 21 Aug 2025 00:02:10 -0400 Subject: [PATCH 14/31] no fast model on custom config --- crates/goose/src/providers/anthropic.rs | 5 +---- crates/goose/src/providers/openai.rs | 5 +---- 2 files changed, 2 insertions(+), 8 deletions(-) diff --git a/crates/goose/src/providers/anthropic.rs b/crates/goose/src/providers/anthropic.rs index dc16764a7d48..1003114efe37 100644 --- a/crates/goose/src/providers/anthropic.rs +++ b/crates/goose/src/providers/anthropic.rs @@ -75,12 +75,9 @@ impl AnthropicProvider { } pub fn from_custom_config( - mut model: ModelConfig, + model: ModelConfig, config: CustomProviderConfig, ) -> Result { - // Set the default fast model for Anthropic - model.fast_model = Some("claude-3-5-haiku-latest".to_string()); - let global_config = crate::config::Config::global(); let api_key: String = global_config .get_secret(&config.api_key_env) diff --git a/crates/goose/src/providers/openai.rs b/crates/goose/src/providers/openai.rs index fff09573e55b..b0fb0f7fb585 100644 --- a/crates/goose/src/providers/openai.rs +++ b/crates/goose/src/providers/openai.rs @@ -113,12 +113,9 @@ impl OpenAiProvider { } pub fn from_custom_config( - mut model: ModelConfig, + model: ModelConfig, config: CustomProviderConfig, ) -> Result { - // Set the default fast model for OpenAI - model.fast_model = Some("gpt-4o-mini".to_string()); - let global_config = crate::config::Config::global(); let api_key: String = global_config .get_secret(&config.api_key_env) From 04475871e0f807fbe8e68666e664c3433899883b Mon Sep 17 00:00:00 2001 From: David Katz Date: Thu, 21 Aug 2025 00:05:13 -0400 Subject: [PATCH 15/31] fn comments --- crates/goose/src/providers/base.rs | 52 ++++-------------------------- 1 file changed, 7 insertions(+), 45 deletions(-) diff --git a/crates/goose/src/providers/base.rs b/crates/goose/src/providers/base.rs index 33eefaea1fc6..b97c314b0b7c 100644 --- a/crates/goose/src/providers/base.rs +++ b/crates/goose/src/providers/base.rs @@ -317,21 +317,8 @@ pub trait Provider: Send + Sync { where Self: Sized; - /// Internal method that performs completion with a specific model - /// This is where providers implement their actual completion logic - /// - /// # Arguments - /// * `model` - The model name to use - /// * `system` - The system prompt that guides the model's behavior - /// * `messages` - The conversation history as a sequence of messages - /// * `tools` - Optional list of tools the model can use - /// - /// # Returns - /// A tuple containing the model's response message and provider usage statistics - /// - /// # Errors - /// ProviderError - /// - It's important to raise ContextLengthExceeded correctly since agent handles it + // Internal implementation of complete, used by complete_fast and complete + // Providers should override this to implement their actual completion logic async fn complete_with_model( &self, model: &str, @@ -340,53 +327,28 @@ pub trait Provider: Send + Sync { tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError>; - /// Generate the next message using the configured model and other parameters - /// - /// # Arguments - /// * `system` - The system prompt that guides the model's behavior - /// * `messages` - The conversation history as a sequence of messages - /// * `tools` - Optional list of tools the model can use - /// - /// # Returns - /// A tuple containing the model's response message and provider usage statistics - /// - /// # Errors - /// ProviderError - /// - It's important to raise ContextLengthExceeded correctly since agent handles it + + // Default implementation: use the provider's configured model async fn complete( &self, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - // Default implementation: use the provider's configured model + let model_config = self.get_model_config(); self.complete_with_model(&model_config.model_name, system, messages, tools) .await } - /// Generate the next message using a fast/cheaper model when available - /// - /// Default implementation just calls regular complete() for providers that don't support fast models - /// - /// # Arguments - /// * `system` - The system prompt that guides the model's behavior - /// * `messages` - The conversation history as a sequence of messages - /// * `tools` - Optional list of tools the model can use - /// - /// # Returns - /// A tuple containing the model's response message and provider usage statistics - /// - /// # Errors - /// ProviderError - /// - It's important to raise ContextLengthExceeded correctly since agent handles it + // Check if a fast model is configured, otherwise fall back to regular model async fn complete_fast( &self, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - // Check if a fast model is configured, otherwise fall back to regular model + let model_config = self.get_model_config(); let model = model_config .fast_model From 51bfa48beeb5e1a17f6755f9ae7da46f0d360fe1 Mon Sep 17 00:00:00 2001 From: David Katz Date: Thu, 21 Aug 2025 00:26:11 -0400 Subject: [PATCH 16/31] Swap model to modelconfig --- crates/goose/src/context_mgmt/auto_compact.rs | 2 +- crates/goose/src/context_mgmt/summarize.rs | 2 +- crates/goose/src/model.rs | 19 +++++---- .../goose/src/permission/permission_judge.rs | 2 +- crates/goose/src/providers/anthropic.rs | 20 ++++----- crates/goose/src/providers/azure.rs | 11 ++--- crates/goose/src/providers/base.rs | 15 ++----- crates/goose/src/providers/bedrock.rs | 6 +-- crates/goose/src/providers/claude_code.rs | 9 ++-- crates/goose/src/providers/cursor_agent.rs | 9 ++-- crates/goose/src/providers/databricks.rs | 15 +++---- crates/goose/src/providers/factory.rs | 2 +- crates/goose/src/providers/gcpvertexai.rs | 11 ++--- crates/goose/src/providers/gemini_cli.rs | 9 ++-- crates/goose/src/providers/githubcopilot.rs | 12 ++---- crates/goose/src/providers/google.rs | 16 +++---- crates/goose/src/providers/groq.rs | 10 ++--- crates/goose/src/providers/lead_worker.rs | 6 +-- crates/goose/src/providers/litellm.rs | 7 +--- crates/goose/src/providers/ollama.rs | 9 ++-- crates/goose/src/providers/openai.rs | 16 +++---- crates/goose/src/providers/openrouter.rs | 9 ++-- crates/goose/src/providers/sagemaker_tgi.rs | 7 +--- crates/goose/src/providers/snowflake.rs | 12 ++---- crates/goose/src/providers/testprovider.rs | 4 +- crates/goose/src/providers/venice.rs | 7 +--- crates/goose/src/providers/xai.rs | 9 ++-- crates/goose/src/scheduler.rs | 2 +- fix_model_usage.sh | 42 +++++++++++++++++++ fix_references.sh | 20 +++++++++ fix_skip.sh | 34 +++++++++++++++ update_providers.sh | 42 +++++++++++++++++++ 32 files changed, 233 insertions(+), 163 deletions(-) create mode 100755 fix_model_usage.sh create mode 100755 fix_references.sh create mode 100755 fix_skip.sh create mode 100755 update_providers.sh diff --git a/crates/goose/src/context_mgmt/auto_compact.rs b/crates/goose/src/context_mgmt/auto_compact.rs index dd8aaa04b478..4fce6d18be3d 100644 --- a/crates/goose/src/context_mgmt/auto_compact.rs +++ b/crates/goose/src/context_mgmt/auto_compact.rs @@ -223,7 +223,7 @@ mod tests { async fn complete_with_model( &self, - _model: &str, + _model_config: &ModelConfig, _system: &str, _messages: &[Message], _tools: &[Tool], diff --git a/crates/goose/src/context_mgmt/summarize.rs b/crates/goose/src/context_mgmt/summarize.rs index 50513d4b5b1a..360ef2a08d9c 100644 --- a/crates/goose/src/context_mgmt/summarize.rs +++ b/crates/goose/src/context_mgmt/summarize.rs @@ -89,7 +89,7 @@ mod tests { async fn complete_with_model( &self, - _model: &str, + _model_config: &ModelConfig, _system: &str, _messages: &[Message], _tools: &[Tool], diff --git a/crates/goose/src/model.rs b/crates/goose/src/model.rs index 0af02a21390b..228013f0b469 100644 --- a/crates/goose/src/model.rs +++ b/crates/goose/src/model.rs @@ -258,11 +258,14 @@ impl ModelConfig { self } - pub fn use_fast_model(mut self) -> Self { - if let Some(fast_model) = self.fast_model.clone() { - self.model_name = fast_model; + pub fn use_fast_model(&self) -> Self { + if let Some(fast_model) = &self.fast_model { + let mut config = self.clone(); + config.model_name = fast_model.clone(); + config + } else { + self.clone() } - self } pub fn context_limit(&self) -> usize { @@ -272,13 +275,13 @@ impl ModelConfig { } // Otherwise, get the model's default limit - let main_limit = Self::get_model_specific_limit(&self.model_name) - .unwrap_or(DEFAULT_CONTEXT_LIMIT); + let main_limit = + Self::get_model_specific_limit(&self.model_name).unwrap_or(DEFAULT_CONTEXT_LIMIT); // If we have a fast_model, also check its limit and use the minimum if let Some(fast_model) = &self.fast_model { - let fast_limit = Self::get_model_specific_limit(fast_model) - .unwrap_or(DEFAULT_CONTEXT_LIMIT); + let fast_limit = + Self::get_model_specific_limit(fast_model).unwrap_or(DEFAULT_CONTEXT_LIMIT); main_limit.min(fast_limit) } else { main_limit diff --git a/crates/goose/src/permission/permission_judge.rs b/crates/goose/src/permission/permission_judge.rs index c438272b5b0f..922b8e21ea44 100644 --- a/crates/goose/src/permission/permission_judge.rs +++ b/crates/goose/src/permission/permission_judge.rs @@ -294,7 +294,7 @@ mod tests { async fn complete_with_model( &self, - _model: &str, + _model_config: &ModelConfig, _system: &str, _messages: &[Message], _tools: &[Tool], diff --git a/crates/goose/src/providers/anthropic.rs b/crates/goose/src/providers/anthropic.rs index 1003114efe37..2b51aab7e7cb 100644 --- a/crates/goose/src/providers/anthropic.rs +++ b/crates/goose/src/providers/anthropic.rs @@ -74,10 +74,7 @@ impl AnthropicProvider { }) } - pub fn from_custom_config( - model: ModelConfig, - config: CustomProviderConfig, - ) -> Result { + pub fn from_custom_config(model: ModelConfig, config: CustomProviderConfig) -> Result { let global_config = crate::config::Config::global(); let api_key: String = global_config .get_secret(&config.api_key_env) @@ -185,20 +182,17 @@ impl Provider for AnthropicProvider { } #[tracing::instrument( - skip(self, model, system, messages, tools), + skip(self, model_config, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] async fn complete_with_model( &self, - model: &str, + model_config: &ModelConfig, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - let mut model_config = self.model.clone(); - model_config.model_name = model.to_string(); - - let payload = create_request(&model_config, system, messages, tools)?; + let payload = create_request(model_config, system, messages, tools)?; let response = self .with_retry(|| async { self.post(&payload).await }) @@ -212,7 +206,7 @@ impl Provider for AnthropicProvider { usage.input_tokens, usage.output_tokens, usage.total_tokens); let response_model = get_model(&json_response); - emit_debug_trace(&model_config, &payload, &json_response, &usage); + emit_debug_trace(&self.model, &payload, &json_response, &usage); let provider_usage = ProviderUsage::new(response_model, usage); tracing::debug!( "🔍 Anthropic non-streaming returning ProviderUsage: {:?}", @@ -281,7 +275,7 @@ impl Provider for AnthropicProvider { let stream = response.bytes_stream().map_err(io::Error::other); - let model_config = self.model.clone(); + let model = self.model.clone(); Ok(Box::pin(try_stream! { let stream_reader = StreamReader::new(stream); let framed = tokio_util::codec::FramedRead::new(stream_reader, tokio_util::codec::LinesCodec::new()).map_err(anyhow::Error::from); @@ -290,7 +284,7 @@ impl Provider for AnthropicProvider { pin!(message_stream); while let Some(message) = futures::StreamExt::next(&mut message_stream).await { let (message, usage) = message.map_err(|e| ProviderError::RequestFailed(format!("Stream decode error: {}", e)))?; - emit_debug_trace(&model_config, &payload, &message, &usage.as_ref().map(|f| f.usage).unwrap_or_default()); + emit_debug_trace(&model, &payload, &message, &usage.as_ref().map(|f| f.usage).unwrap_or_default()); yield (message, usage); } })) diff --git a/crates/goose/src/providers/azure.rs b/crates/goose/src/providers/azure.rs index 5be63357deaa..12bf93f20e3d 100644 --- a/crates/goose/src/providers/azure.rs +++ b/crates/goose/src/providers/azure.rs @@ -135,20 +135,17 @@ impl Provider for AzureProvider { } #[tracing::instrument( - skip(self, model, system, messages, tools), + skip(self, model_config, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] async fn complete_with_model( &self, - model: &str, + model_config: &ModelConfig, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - let mut model_config = self.model.clone(); - model_config.model_name = model.to_string(); - - let payload = create_request(&model_config, system, messages, tools, &ImageFormat::OpenAi)?; + let payload = create_request(model_config, system, messages, tools, &ImageFormat::OpenAi)?; let response = self .with_retry(|| async { let payload_clone = payload.clone(); @@ -162,7 +159,7 @@ impl Provider for AzureProvider { Usage::default() }); let response_model = get_model(&response); - emit_debug_trace(&model_config, &payload, &response, &usage); + emit_debug_trace(model_config, &payload, &response, &usage); Ok((message, ProviderUsage::new(response_model, usage))) } } diff --git a/crates/goose/src/providers/base.rs b/crates/goose/src/providers/base.rs index b97c314b0b7c..bed31d59909c 100644 --- a/crates/goose/src/providers/base.rs +++ b/crates/goose/src/providers/base.rs @@ -321,13 +321,12 @@ pub trait Provider: Send + Sync { // Providers should override this to implement their actual completion logic async fn complete_with_model( &self, - model: &str, + model_config: &ModelConfig, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError>; - // Default implementation: use the provider's configured model async fn complete( &self, @@ -335,9 +334,8 @@ pub trait Provider: Send + Sync { messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - let model_config = self.get_model_config(); - self.complete_with_model(&model_config.model_name, system, messages, tools) + self.complete_with_model(&model_config, system, messages, tools) .await } @@ -348,14 +346,9 @@ pub trait Provider: Send + Sync { messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - let model_config = self.get_model_config(); - let model = model_config - .fast_model - .as_deref() - .unwrap_or(&model_config.model_name); - - self.complete_with_model(model, system, messages, tools) + let fast_config = model_config.use_fast_model(); + self.complete_with_model(&fast_config, system, messages, tools) .await } diff --git a/crates/goose/src/providers/bedrock.rs b/crates/goose/src/providers/bedrock.rs index 00e1befa0176..a2e04dbbf539 100644 --- a/crates/goose/src/providers/bedrock.rs +++ b/crates/goose/src/providers/bedrock.rs @@ -152,17 +152,17 @@ impl Provider for BedrockProvider { } #[tracing::instrument( - skip(self, _model, system, messages, tools), + skip(self, model_config, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] async fn complete_with_model( &self, - _model: &str, + model_config: &ModelConfig, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - let model_name = &self.model.model_name; + let model_name = model_config.model_name.clone(); let (bedrock_message, bedrock_usage) = self .with_retry(|| self.converse(system, messages, tools)) diff --git a/crates/goose/src/providers/claude_code.rs b/crates/goose/src/providers/claude_code.rs index 22b4b22adc1c..9ef2950c272c 100644 --- a/crates/goose/src/providers/claude_code.rs +++ b/crates/goose/src/providers/claude_code.rs @@ -474,19 +474,16 @@ impl Provider for ClaudeCodeProvider { } #[tracing::instrument( - skip(self, model, system, messages, tools), + skip(self, model_config, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] async fn complete_with_model( &self, - model: &str, + model_config: &ModelConfig, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - let mut model_config = self.model.clone(); - model_config.model_name = model.to_string(); - // Check if this is a session description request (short system prompt asking for 4 words or less) if system.contains("four words or less") || system.contains("4 words or less") { return self.generate_simple_session_description(messages); @@ -509,7 +506,7 @@ impl Provider for ClaudeCodeProvider { "usage": usage }); - emit_debug_trace(&model_config, &payload, &response, &usage); + emit_debug_trace(model_config, &payload, &response, &usage); Ok(( message, diff --git a/crates/goose/src/providers/cursor_agent.rs b/crates/goose/src/providers/cursor_agent.rs index 2e677ffce101..bbce315f648b 100644 --- a/crates/goose/src/providers/cursor_agent.rs +++ b/crates/goose/src/providers/cursor_agent.rs @@ -407,19 +407,16 @@ impl Provider for CursorAgentProvider { } #[tracing::instrument( - skip(self, model, system, messages, tools), + skip(self, model_config, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] async fn complete_with_model( &self, - model: &str, + model_config: &ModelConfig, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - let mut model_config = self.model.clone(); - model_config.model_name = model.to_string(); - // Check if this is a session description request (short system prompt asking for 4 words or less) if system.contains("four words or less") || system.contains("4 words or less") { return self.generate_simple_session_description(messages); @@ -442,7 +439,7 @@ impl Provider for CursorAgentProvider { "usage": usage }); - emit_debug_trace(&model_config, &payload, &response, &usage); + emit_debug_trace(model_config, &payload, &response, &usage); Ok(( message, diff --git a/crates/goose/src/providers/databricks.rs b/crates/goose/src/providers/databricks.rs index d89913ff08f7..0bc80ed1d4e5 100644 --- a/crates/goose/src/providers/databricks.rs +++ b/crates/goose/src/providers/databricks.rs @@ -238,21 +238,18 @@ impl Provider for DatabricksProvider { } #[tracing::instrument( - skip(self, model, system, messages, tools), + skip(self, model_config, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] async fn complete_with_model( &self, - model: &str, + model_config: &ModelConfig, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - let mut model_config = self.model.clone(); - model_config.model_name = model.to_string(); - let mut payload = - create_request(&model_config, system, messages, tools, &self.image_format)?; + create_request(model_config, system, messages, tools, &self.image_format)?; payload .as_object_mut() .expect("payload should have model key") @@ -266,7 +263,7 @@ impl Provider for DatabricksProvider { Usage::default() }); let response_model = get_model(&response); - super::utils::emit_debug_trace(&model_config, &payload, &response, &usage); + super::utils::emit_debug_trace(&self.model, &payload, &response, &usage); Ok((message, ProviderUsage::new(response_model, usage))) } @@ -307,8 +304,8 @@ impl Provider for DatabricksProvider { .await?; let stream = response.bytes_stream().map_err(io::Error::other); - let model_config = self.model.clone(); + let model = self.model.clone(); Ok(Box::pin(try_stream! { let stream_reader = StreamReader::new(stream); let framed = FramedRead::new(stream_reader, LinesCodec::new()).map_err(anyhow::Error::from); @@ -317,7 +314,7 @@ impl Provider for DatabricksProvider { pin!(message_stream); while let Some(message) = message_stream.next().await { let (message, usage) = message.map_err(|e| ProviderError::RequestFailed(format!("Stream decode error: {}", e)))?; - super::utils::emit_debug_trace(&model_config, &payload, &message, &usage.as_ref().map(|f| f.usage).unwrap_or_default()); + super::utils::emit_debug_trace(&model, &payload, &message, &usage.as_ref().map(|f| f.usage).unwrap_or_default()); yield (message, usage); } })) diff --git a/crates/goose/src/providers/factory.rs b/crates/goose/src/providers/factory.rs index bd0013a3c372..80b9abb2aa7a 100644 --- a/crates/goose/src/providers/factory.rs +++ b/crates/goose/src/providers/factory.rs @@ -202,7 +202,7 @@ mod tests { async fn complete_with_model( &self, - _model: &str, + _model_config: &ModelConfig, _system: &str, _messages: &[Message], _tools: &[Tool], diff --git a/crates/goose/src/providers/gcpvertexai.rs b/crates/goose/src/providers/gcpvertexai.rs index d23fa62a4e9a..609d77ab7eb6 100644 --- a/crates/goose/src/providers/gcpvertexai.rs +++ b/crates/goose/src/providers/gcpvertexai.rs @@ -512,27 +512,24 @@ impl Provider for GcpVertexAIProvider { /// * `messages` - Array of previous messages in the conversation /// * `tools` - Array of available tools for the model #[tracing::instrument( - skip(self, model, system, messages, tools), + skip(self, model_config, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] async fn complete_with_model( &self, - model: &str, + model_config: &ModelConfig, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - let mut model_config = self.model.clone(); - model_config.model_name = model.to_string(); - // Create request and context - let (request, context) = create_request(&model_config, system, messages, tools)?; + let (request, context) = create_request(model_config, system, messages, tools)?; // Send request and process response let response = self.post(&request, &context).await?; let usage = get_usage(&response, &context)?; - emit_debug_trace(&model_config, &request, &response, &usage); + emit_debug_trace(model_config, &request, &response, &usage); // Convert response to message let message = response_to_message(response, context)?; diff --git a/crates/goose/src/providers/gemini_cli.rs b/crates/goose/src/providers/gemini_cli.rs index d25525e074c9..4b14270244fd 100644 --- a/crates/goose/src/providers/gemini_cli.rs +++ b/crates/goose/src/providers/gemini_cli.rs @@ -319,19 +319,16 @@ impl Provider for GeminiCliProvider { } #[tracing::instrument( - skip(self, model, system, messages, tools), + skip(self, model_config, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] async fn complete_with_model( &self, - model: &str, + model_config: &ModelConfig, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - let mut model_config = self.model.clone(); - model_config.model_name = model.to_string(); - // Check if this is a session description request (short system prompt asking for 4 words or less) if system.contains("four words or less") || system.contains("4 words or less") { return self.generate_simple_session_description(messages); @@ -354,7 +351,7 @@ impl Provider for GeminiCliProvider { "usage": usage }); - emit_debug_trace(&model_config, &payload, &response, &usage); + emit_debug_trace(model_config, &payload, &response, &usage); Ok(( message, diff --git a/crates/goose/src/providers/githubcopilot.rs b/crates/goose/src/providers/githubcopilot.rs index ba1d16de6f16..720d8089b717 100644 --- a/crates/goose/src/providers/githubcopilot.rs +++ b/crates/goose/src/providers/githubcopilot.rs @@ -401,21 +401,17 @@ impl Provider for GithubCopilotProvider { } #[tracing::instrument( - skip(self, model, system, messages, tools), + skip(self, model_config, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] async fn complete_with_model( &self, - model: &str, + model_config: &ModelConfig, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - let mut model_config = self.model.clone(); - - model_config.model_name = model.to_string(); - - let payload = create_request(&model_config, system, messages, tools, &ImageFormat::OpenAi)?; + let payload = create_request(model_config, system, messages, tools, &ImageFormat::OpenAi)?; // Make request with retry let response = self @@ -432,7 +428,7 @@ impl Provider for GithubCopilotProvider { Usage::default() }); let response_model = get_model(&response); - emit_debug_trace(&model_config, &payload, &response, &usage); + emit_debug_trace(model_config, &payload, &response, &usage); Ok((message, ProviderUsage::new(response_model, usage))) } diff --git a/crates/goose/src/providers/google.rs b/crates/goose/src/providers/google.rs index d6e2bcbbc53a..4bf843088e34 100644 --- a/crates/goose/src/providers/google.rs +++ b/crates/goose/src/providers/google.rs @@ -55,7 +55,7 @@ pub struct GoogleProvider { impl_provider_default!(GoogleProvider); impl GoogleProvider { - pub fn from_env(mut model: ModelConfig) -> Result { + pub fn from_env(model: ModelConfig) -> Result { let model = model.with_fast(GOOGLE_DEFAULT_FAST_MODEL.to_string()); let config = crate::config::Config::global(); @@ -104,27 +104,23 @@ impl Provider for GoogleProvider { } #[tracing::instrument( - skip(self, model, system, messages, tools), + skip(self, model_config, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] async fn complete_with_model( &self, - model: &str, + model_config: &ModelConfig, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - let mut model_config = self.model.clone(); - - model_config.model_name = model.to_string(); - - let payload = create_request(&model_config, system, messages, tools)?; + let payload = create_request(model_config, system, messages, tools)?; // Make request let response = self .with_retry(|| async { let payload_clone = payload.clone(); - self.post(model, &payload_clone).await + self.post(&model_config.model_name, &payload_clone).await }) .await?; @@ -135,7 +131,7 @@ impl Provider for GoogleProvider { Some(model_version) => model_version.as_str().unwrap_or_default().to_string(), None => model_config.model_name.clone(), }; - emit_debug_trace(&model_config, &payload, &response, &usage); + emit_debug_trace(model_config, &payload, &response, &usage); let provider_usage = ProviderUsage::new(response_model, usage); Ok((message, provider_usage)) } diff --git a/crates/goose/src/providers/groq.rs b/crates/goose/src/providers/groq.rs index ebd1578f77a3..91825f3d41ab 100644 --- a/crates/goose/src/providers/groq.rs +++ b/crates/goose/src/providers/groq.rs @@ -34,7 +34,6 @@ impl_provider_default!(GroqProvider); impl GroqProvider { pub fn from_env(model: ModelConfig) -> Result { - let config = crate::config::Config::global(); let api_key: String = config.get_secret("GROQ_API_KEY")?; let host: String = config @@ -78,19 +77,16 @@ impl Provider for GroqProvider { } #[tracing::instrument( - skip(self, model, system, messages, tools), + skip(self, model_config, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] async fn complete_with_model( &self, - model: &str, + model_config: &ModelConfig, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - let mut model_config = self.model.clone(); - model_config.model_name = model.to_string(); - let payload = create_request( &model_config, system, @@ -107,7 +103,7 @@ impl Provider for GroqProvider { Usage::default() }); let response_model = get_model(&response); - super::utils::emit_debug_trace(&model_config, &payload, &response, &usage); + super::utils::emit_debug_trace(model_config, &payload, &response, &usage); Ok((message, ProviderUsage::new(response_model, usage))) } diff --git a/crates/goose/src/providers/lead_worker.rs b/crates/goose/src/providers/lead_worker.rs index d0e57e4c43c0..07638b2198bc 100644 --- a/crates/goose/src/providers/lead_worker.rs +++ b/crates/goose/src/providers/lead_worker.rs @@ -328,7 +328,7 @@ impl Provider for LeadWorkerProvider { async fn complete_with_model( &self, - _model: &str, + _model_config: &ModelConfig, system: &str, messages: &[Message], tools: &[Tool], @@ -478,7 +478,7 @@ mod tests { async fn complete_with_model( &self, - _model: &str, + _model_config: &ModelConfig, _system: &str, _messages: &[Message], _tools: &[Tool], @@ -639,7 +639,7 @@ mod tests { async fn complete_with_model( &self, - _model: &str, + _model_config: &ModelConfig, _system: &str, _messages: &[Message], _tools: &[Tool], diff --git a/crates/goose/src/providers/litellm.rs b/crates/goose/src/providers/litellm.rs index f198dde99d15..41a3a89f1618 100644 --- a/crates/goose/src/providers/litellm.rs +++ b/crates/goose/src/providers/litellm.rs @@ -163,14 +163,11 @@ impl Provider for LiteLLMProvider { #[tracing::instrument(skip_all, name = "provider_complete")] async fn complete_with_model( &self, - model: &str, + model_config: &ModelConfig, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - let mut model_config = self.model.clone(); - model_config.model_name = model.to_string(); - let mut payload = super::formats::openai::create_request( &model_config, system, @@ -193,7 +190,7 @@ impl Provider for LiteLLMProvider { let message = super::formats::openai::response_to_message(&response)?; let usage = super::formats::openai::get_usage(&response); let response_model = get_model(&response); - emit_debug_trace(&model_config, &payload, &response, &usage); + emit_debug_trace(model_config, &payload, &response, &usage); Ok((message, ProviderUsage::new(response_model, usage))) } diff --git a/crates/goose/src/providers/ollama.rs b/crates/goose/src/providers/ollama.rs index 9afea70651a5..0cf35cc408c8 100644 --- a/crates/goose/src/providers/ollama.rs +++ b/crates/goose/src/providers/ollama.rs @@ -165,19 +165,16 @@ impl Provider for OllamaProvider { } #[tracing::instrument( - skip(self, model, system, messages, tools), + skip(self, model_config, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] async fn complete_with_model( &self, - model: &str, + model_config: &ModelConfig, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - let mut model_config = self.model.clone(); - model_config.model_name = model.to_string(); - let config = crate::config::Config::global(); let goose_mode = config.get_param("GOOSE_MODE").unwrap_or("auto".to_string()); let filtered_tools = if goose_mode == "chat" { &[] } else { tools }; @@ -202,7 +199,7 @@ impl Provider for OllamaProvider { Usage::default() }); let response_model = get_model(&response); - super::utils::emit_debug_trace(&model_config, &payload, &response, &usage); + super::utils::emit_debug_trace(model_config, &payload, &response, &usage); Ok((message, ProviderUsage::new(response_model, usage))) } diff --git a/crates/goose/src/providers/openai.rs b/crates/goose/src/providers/openai.rs index b0fb0f7fb585..1c85a804e179 100644 --- a/crates/goose/src/providers/openai.rs +++ b/crates/goose/src/providers/openai.rs @@ -112,10 +112,7 @@ impl OpenAiProvider { }) } - pub fn from_custom_config( - model: ModelConfig, - config: CustomProviderConfig, - ) -> Result { + pub fn from_custom_config(model: ModelConfig, config: CustomProviderConfig) -> Result { let global_config = crate::config::Config::global(); let api_key: String = global_config .get_secret(&config.api_key_env) @@ -199,20 +196,17 @@ impl Provider for OpenAiProvider { } #[tracing::instrument( - skip(self, model, system, messages, tools), + skip(self, model_config, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] async fn complete_with_model( &self, - model: &str, + model_config: &ModelConfig, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - let mut model_config = self.model.clone(); - model_config.model_name = model.to_string(); - - let payload = create_request(&model_config, system, messages, tools, &ImageFormat::OpenAi)?; + let payload = create_request(model_config, system, messages, tools, &ImageFormat::OpenAi)?; let json_response = self.post(&payload).await?; @@ -225,7 +219,7 @@ impl Provider for OpenAiProvider { Usage::default() }); let response_model = get_model(&json_response); - emit_debug_trace(&model_config, &payload, &json_response, &usage); + emit_debug_trace(model_config, &payload, &json_response, &usage); Ok((message, ProviderUsage::new(response_model, usage))) } diff --git a/crates/goose/src/providers/openrouter.rs b/crates/goose/src/providers/openrouter.rs index f2e607b60276..49b9cea164e2 100644 --- a/crates/goose/src/providers/openrouter.rs +++ b/crates/goose/src/providers/openrouter.rs @@ -238,19 +238,16 @@ impl Provider for OpenRouterProvider { } #[tracing::instrument( - skip(self, model, system, messages, tools), + skip(self, model_config, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] async fn complete_with_model( &self, - model: &str, + model_config: &ModelConfig, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - let mut model_config = self.model.clone(); - model_config.model_name = model.to_string(); - // Create the base payload let payload = create_request_based_on_model(self, system, messages, tools)?; @@ -269,7 +266,7 @@ impl Provider for OpenRouterProvider { Usage::default() }); let response_model = get_model(&response); - emit_debug_trace(&model_config, &payload, &response, &usage); + emit_debug_trace(model_config, &payload, &response, &usage); Ok((message, ProviderUsage::new(response_model, usage))) } diff --git a/crates/goose/src/providers/sagemaker_tgi.rs b/crates/goose/src/providers/sagemaker_tgi.rs index 551bb464ee80..ea12beba5c15 100644 --- a/crates/goose/src/providers/sagemaker_tgi.rs +++ b/crates/goose/src/providers/sagemaker_tgi.rs @@ -280,19 +280,16 @@ impl Provider for SageMakerTgiProvider { } #[tracing::instrument( - skip(self, model, system, messages, tools), + skip(self, model_config, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] async fn complete_with_model( &self, - model: &str, + model_config: &ModelConfig, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - let mut model_config = self.model.clone(); - model_config.model_name = model.to_string(); - let model_name = &self.model.model_name; let request_payload = self.create_tgi_request(system, messages).map_err(|e| { diff --git a/crates/goose/src/providers/snowflake.rs b/crates/goose/src/providers/snowflake.rs index 8df85e5f4a99..5b9344a29d9e 100644 --- a/crates/goose/src/providers/snowflake.rs +++ b/crates/goose/src/providers/snowflake.rs @@ -299,21 +299,17 @@ impl Provider for SnowflakeProvider { } #[tracing::instrument( - skip(self, model, system, messages, tools), + skip(self, model_config, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] async fn complete_with_model( &self, - model: &str, + model_config: &ModelConfig, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - let mut model_config = self.model.clone(); - - model_config.model_name = model.to_string(); - - let payload = create_request(&model_config, system, messages, tools)?; + let payload = create_request(model_config, system, messages, tools)?; let response = self .with_retry(|| async { @@ -326,7 +322,7 @@ impl Provider for SnowflakeProvider { let message = response_to_message(&response)?; let usage = get_usage(&response)?; let response_model = get_model(&response); - super::utils::emit_debug_trace(&model_config, &payload, &response, &usage); + super::utils::emit_debug_trace(model_config, &payload, &response, &usage); Ok((message, ProviderUsage::new(response_model, usage))) } diff --git a/crates/goose/src/providers/testprovider.rs b/crates/goose/src/providers/testprovider.rs index 8c9a0cbe2cb3..98c4f22739d9 100644 --- a/crates/goose/src/providers/testprovider.rs +++ b/crates/goose/src/providers/testprovider.rs @@ -114,7 +114,7 @@ impl Provider for TestProvider { async fn complete_with_model( &self, - _model: &str, + _model_config: &ModelConfig, system: &str, messages: &[Message], tools: &[Tool], @@ -191,7 +191,7 @@ mod tests { async fn complete_with_model( &self, - _model: &str, + _model_config: &ModelConfig, _system: &str, _messages: &[Message], _tools: &[Tool], diff --git a/crates/goose/src/providers/venice.rs b/crates/goose/src/providers/venice.rs index 6a84d9c244ba..6f30320cbe8b 100644 --- a/crates/goose/src/providers/venice.rs +++ b/crates/goose/src/providers/venice.rs @@ -246,19 +246,16 @@ impl Provider for VeniceProvider { } #[tracing::instrument( - skip(self, model, system, messages, tools), + skip(self, model_config, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] async fn complete_with_model( &self, - model: &str, + model_config: &ModelConfig, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - let mut model_config = self.model.clone(); - model_config.model_name = model.to_string(); - // Create properly formatted messages for Venice API let mut formatted_messages = Vec::new(); diff --git a/crates/goose/src/providers/xai.rs b/crates/goose/src/providers/xai.rs index 30d7a347e4c4..9b862eeaddf4 100644 --- a/crates/goose/src/providers/xai.rs +++ b/crates/goose/src/providers/xai.rs @@ -93,19 +93,16 @@ impl Provider for XaiProvider { } #[tracing::instrument( - skip(self, model, system, messages, tools), + skip(self, model_config, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] async fn complete_with_model( &self, - model: &str, + model_config: &ModelConfig, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - let mut model_config = self.model.clone(); - model_config.model_name = model.to_string(); - let payload = create_request( &model_config, system, @@ -122,7 +119,7 @@ impl Provider for XaiProvider { Usage::default() }); let response_model = get_model(&response); - super::utils::emit_debug_trace(&model_config, &payload, &response, &usage); + super::utils::emit_debug_trace(model_config, &payload, &response, &usage); Ok((message, ProviderUsage::new(response_model, usage))) } } diff --git a/crates/goose/src/scheduler.rs b/crates/goose/src/scheduler.rs index e6c74cb0899f..c6dd933cb7ac 100644 --- a/crates/goose/src/scheduler.rs +++ b/crates/goose/src/scheduler.rs @@ -1392,7 +1392,7 @@ mod tests { async fn complete_with_model( &self, - _model: &str, + _model_config: &ModelConfig, _system: &str, _messages: &[Message], _tools: &[Tool], diff --git a/fix_model_usage.sh b/fix_model_usage.sh new file mode 100755 index 000000000000..0178f3b509fe --- /dev/null +++ b/fix_model_usage.sh @@ -0,0 +1,42 @@ +#!/bin/bash + +# For providers that were creating a temporary model_config and setting model_name +# We need to remove those lines since model_config is now passed directly + +providers=( + "anthropic.rs" + "azure.rs" + "claude_code.rs" + "cursor_agent.rs" + "databricks.rs" + "gcpvertexai.rs" + "gemini_cli.rs" + "githubcopilot.rs" + "google.rs" + "groq.rs" + "litellm.rs" + "ollama.rs" + "openrouter.rs" + "sagemaker_tgi.rs" + "snowflake.rs" + "venice.rs" + "xai.rs" + "together.rs" + "fireworks.rs" + "octoai.rs" +) + +for provider in "${providers[@]}"; do + file="crates/goose/src/providers/$provider" + if [ -f "$file" ]; then + echo "Fixing $file..." + # Remove the lines that create temporary model_config + sed -i '' '/let mut model_config = self\.model\.clone();/d' "$file" + sed -i '' '/model_config\.model_name = model\.to_string();/d' "$file" + # Change &model_config to model_config in create_request calls + sed -i '' 's/create_request(&model_config,/create_request(model_config,/g' "$file" + sed -i '' 's/emit_debug_trace(&model_config,/emit_debug_trace(model_config,/g' "$file" + fi +done + +echo "Done fixing model usage" diff --git a/fix_references.sh b/fix_references.sh new file mode 100755 index 000000000000..b6839b07d700 --- /dev/null +++ b/fix_references.sh @@ -0,0 +1,20 @@ +#!/bin/bash + +# Fix references to model_config in streaming code +providers=( + "anthropic.rs" + "databricks.rs" +) + +for provider in "${providers[@]}"; do + file="crates/goose/src/providers/$provider" + if [ -f "$file" ]; then + echo "Fixing references in $file..." + # In streaming code, use self.model instead of model_config + sed -i '' 's/let model_config = self\.model\.clone();//' "$file" + sed -i '' 's/emit_debug_trace(model_config,/emit_debug_trace(\&self.model,/g' "$file" + sed -i '' 's/emit_debug_trace(&model_config,/emit_debug_trace(\&self.model,/g' "$file" + fi +done + +echo "Done fixing references" diff --git a/fix_skip.sh b/fix_skip.sh new file mode 100755 index 000000000000..64ec0237e9f2 --- /dev/null +++ b/fix_skip.sh @@ -0,0 +1,34 @@ +#!/bin/bash + +# List of provider files that need updating +providers=( + "anthropic.rs" + "azure.rs" + "bedrock.rs" + "claude_code.rs" + "cursor_agent.rs" + "databricks.rs" + "gcpvertexai.rs" + "gemini_cli.rs" + "githubcopilot.rs" + "google.rs" + "groq.rs" + "ollama.rs" + "openrouter.rs" + "sagemaker_tgi.rs" + "snowflake.rs" + "venice.rs" + "xai.rs" +) + +for provider in "${providers[@]}"; do + file="crates/goose/src/providers/$provider" + if [ -f "$file" ]; then + echo "Fixing skip in $file..." + # Update skip(self, model, ...) to skip(self, model_config, ...) + sed -i '' 's/skip(self, model,/skip(self, model_config,/g' "$file" + sed -i '' 's/skip(self, _model,/skip(self, _model_config,/g' "$file" + fi +done + +echo "Done fixing skip parameters" diff --git a/update_providers.sh b/update_providers.sh new file mode 100755 index 000000000000..14b7ebe7db1d --- /dev/null +++ b/update_providers.sh @@ -0,0 +1,42 @@ +#!/bin/bash + +# List of provider files that need updating +providers=( + "anthropic.rs" + "azure.rs" + "bedrock.rs" + "claude_code.rs" + "cursor_agent.rs" + "databricks.rs" + "factory.rs" + "gcpvertexai.rs" + "gemini_cli.rs" + "githubcopilot.rs" + "google.rs" + "groq.rs" + "lead_worker.rs" + "litellm.rs" + "ollama.rs" + "openrouter.rs" + "sagemaker_tgi.rs" + "snowflake.rs" + "testprovider.rs" + "together.rs" + "venice.rs" + "xai.rs" + "fireworks.rs" + "octoai.rs" +) + +for provider in "${providers[@]}"; do + file="crates/goose/src/providers/$provider" + if [ -f "$file" ]; then + echo "Updating $file..." + # Update the parameter from model: &str to model_config: &ModelConfig + sed -i '' 's/async fn complete_with_model(/async fn complete_with_model(/g' "$file" + sed -i '' 's/model: &str,/model_config: \&ModelConfig,/g' "$file" + sed -i '' 's/_model: &str,/_model_config: \&ModelConfig,/g' "$file" + fi +done + +echo "Done updating provider files" From d3a5307297044b77edfa75f30ebb1be877998a89 Mon Sep 17 00:00:00 2001 From: David Katz Date: Thu, 21 Aug 2025 00:26:37 -0400 Subject: [PATCH 17/31] rm extra scripts --- fix_model_usage.sh | 42 ------------------------------------------ fix_references.sh | 20 -------------------- fix_skip.sh | 34 ---------------------------------- update_providers.sh | 42 ------------------------------------------ 4 files changed, 138 deletions(-) delete mode 100755 fix_model_usage.sh delete mode 100755 fix_references.sh delete mode 100755 fix_skip.sh delete mode 100755 update_providers.sh diff --git a/fix_model_usage.sh b/fix_model_usage.sh deleted file mode 100755 index 0178f3b509fe..000000000000 --- a/fix_model_usage.sh +++ /dev/null @@ -1,42 +0,0 @@ -#!/bin/bash - -# For providers that were creating a temporary model_config and setting model_name -# We need to remove those lines since model_config is now passed directly - -providers=( - "anthropic.rs" - "azure.rs" - "claude_code.rs" - "cursor_agent.rs" - "databricks.rs" - "gcpvertexai.rs" - "gemini_cli.rs" - "githubcopilot.rs" - "google.rs" - "groq.rs" - "litellm.rs" - "ollama.rs" - "openrouter.rs" - "sagemaker_tgi.rs" - "snowflake.rs" - "venice.rs" - "xai.rs" - "together.rs" - "fireworks.rs" - "octoai.rs" -) - -for provider in "${providers[@]}"; do - file="crates/goose/src/providers/$provider" - if [ -f "$file" ]; then - echo "Fixing $file..." - # Remove the lines that create temporary model_config - sed -i '' '/let mut model_config = self\.model\.clone();/d' "$file" - sed -i '' '/model_config\.model_name = model\.to_string();/d' "$file" - # Change &model_config to model_config in create_request calls - sed -i '' 's/create_request(&model_config,/create_request(model_config,/g' "$file" - sed -i '' 's/emit_debug_trace(&model_config,/emit_debug_trace(model_config,/g' "$file" - fi -done - -echo "Done fixing model usage" diff --git a/fix_references.sh b/fix_references.sh deleted file mode 100755 index b6839b07d700..000000000000 --- a/fix_references.sh +++ /dev/null @@ -1,20 +0,0 @@ -#!/bin/bash - -# Fix references to model_config in streaming code -providers=( - "anthropic.rs" - "databricks.rs" -) - -for provider in "${providers[@]}"; do - file="crates/goose/src/providers/$provider" - if [ -f "$file" ]; then - echo "Fixing references in $file..." - # In streaming code, use self.model instead of model_config - sed -i '' 's/let model_config = self\.model\.clone();//' "$file" - sed -i '' 's/emit_debug_trace(model_config,/emit_debug_trace(\&self.model,/g' "$file" - sed -i '' 's/emit_debug_trace(&model_config,/emit_debug_trace(\&self.model,/g' "$file" - fi -done - -echo "Done fixing references" diff --git a/fix_skip.sh b/fix_skip.sh deleted file mode 100755 index 64ec0237e9f2..000000000000 --- a/fix_skip.sh +++ /dev/null @@ -1,34 +0,0 @@ -#!/bin/bash - -# List of provider files that need updating -providers=( - "anthropic.rs" - "azure.rs" - "bedrock.rs" - "claude_code.rs" - "cursor_agent.rs" - "databricks.rs" - "gcpvertexai.rs" - "gemini_cli.rs" - "githubcopilot.rs" - "google.rs" - "groq.rs" - "ollama.rs" - "openrouter.rs" - "sagemaker_tgi.rs" - "snowflake.rs" - "venice.rs" - "xai.rs" -) - -for provider in "${providers[@]}"; do - file="crates/goose/src/providers/$provider" - if [ -f "$file" ]; then - echo "Fixing skip in $file..." - # Update skip(self, model, ...) to skip(self, model_config, ...) - sed -i '' 's/skip(self, model,/skip(self, model_config,/g' "$file" - sed -i '' 's/skip(self, _model,/skip(self, _model_config,/g' "$file" - fi -done - -echo "Done fixing skip parameters" diff --git a/update_providers.sh b/update_providers.sh deleted file mode 100755 index 14b7ebe7db1d..000000000000 --- a/update_providers.sh +++ /dev/null @@ -1,42 +0,0 @@ -#!/bin/bash - -# List of provider files that need updating -providers=( - "anthropic.rs" - "azure.rs" - "bedrock.rs" - "claude_code.rs" - "cursor_agent.rs" - "databricks.rs" - "factory.rs" - "gcpvertexai.rs" - "gemini_cli.rs" - "githubcopilot.rs" - "google.rs" - "groq.rs" - "lead_worker.rs" - "litellm.rs" - "ollama.rs" - "openrouter.rs" - "sagemaker_tgi.rs" - "snowflake.rs" - "testprovider.rs" - "together.rs" - "venice.rs" - "xai.rs" - "fireworks.rs" - "octoai.rs" -) - -for provider in "${providers[@]}"; do - file="crates/goose/src/providers/$provider" - if [ -f "$file" ]; then - echo "Updating $file..." - # Update the parameter from model: &str to model_config: &ModelConfig - sed -i '' 's/async fn complete_with_model(/async fn complete_with_model(/g' "$file" - sed -i '' 's/model: &str,/model_config: \&ModelConfig,/g' "$file" - sed -i '' 's/_model: &str,/_model_config: \&ModelConfig,/g' "$file" - fi -done - -echo "Done updating provider files" From 55aeedf3aa65482889752c191af3000dad6cd94c Mon Sep 17 00:00:00 2001 From: David Katz Date: Thu, 21 Aug 2025 00:34:25 -0400 Subject: [PATCH 18/31] openai output model --- crates/goose/src/providers/openai.rs | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/crates/goose/src/providers/openai.rs b/crates/goose/src/providers/openai.rs index 1c85a804e179..5a8a3c16413b 100644 --- a/crates/goose/src/providers/openai.rs +++ b/crates/goose/src/providers/openai.rs @@ -218,9 +218,8 @@ impl Provider for OpenAiProvider { tracing::debug!("Failed to get usage data"); Usage::default() }); - let response_model = get_model(&json_response); emit_debug_trace(model_config, &payload, &json_response, &usage); - Ok((message, ProviderUsage::new(response_model, usage))) + Ok((message, ProviderUsage::new(model_config.model_name.clone(), usage))) } async fn fetch_supported_models(&self) -> Result>, ProviderError> { From b666551ec927a786a81a2c32d1d5249ba2443642 Mon Sep 17 00:00:00 2001 From: David Katz Date: Thu, 21 Aug 2025 00:40:07 -0400 Subject: [PATCH 19/31] fix warnings --- crates/goose/src/providers/openai.rs | 2 +- crates/goose/src/providers/sagemaker_tgi.rs | 2 +- crates/goose/src/providers/venice.rs | 4 ++-- 3 files changed, 4 insertions(+), 4 deletions(-) diff --git a/crates/goose/src/providers/openai.rs b/crates/goose/src/providers/openai.rs index 5a8a3c16413b..fe86e5c31d5d 100644 --- a/crates/goose/src/providers/openai.rs +++ b/crates/goose/src/providers/openai.rs @@ -17,7 +17,7 @@ use super::embedding::{EmbeddingCapable, EmbeddingRequest, EmbeddingResponse}; use super::errors::ProviderError; use super::formats::openai::{create_request, get_usage, response_to_message}; use super::utils::{ - emit_debug_trace, get_model, handle_response_openai_compat, handle_status_openai_compat, + emit_debug_trace, handle_response_openai_compat, handle_status_openai_compat, ImageFormat, }; use crate::config::custom_providers::CustomProviderConfig; diff --git a/crates/goose/src/providers/sagemaker_tgi.rs b/crates/goose/src/providers/sagemaker_tgi.rs index ea12beba5c15..0c65e32d75cc 100644 --- a/crates/goose/src/providers/sagemaker_tgi.rs +++ b/crates/goose/src/providers/sagemaker_tgi.rs @@ -290,7 +290,7 @@ impl Provider for SageMakerTgiProvider { messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - let model_name = &self.model.model_name; + let model_name = &model_config.model_name; let request_payload = self.create_tgi_request(system, messages).map_err(|e| { ProviderError::RequestFailed(format!("Failed to create request: {}", e)) diff --git a/crates/goose/src/providers/venice.rs b/crates/goose/src/providers/venice.rs index 6f30320cbe8b..9c075bec9c06 100644 --- a/crates/goose/src/providers/venice.rs +++ b/crates/goose/src/providers/venice.rs @@ -392,7 +392,7 @@ impl Provider for VeniceProvider { // Build Venice-specific payload let mut payload = json!({ - "model": strip_flags(&self.model.model_name), + "model": strip_flags(&model_config.model_name), "messages": formatted_messages, "stream": false, "temperature": 0.7, @@ -471,7 +471,7 @@ impl Provider for VeniceProvider { return Ok(( message, ProviderUsage::new( - strip_flags(&self.model.model_name).to_string(), + strip_flags(&model_config.model_name).to_string(), Usage::default(), ), )); From f6c2550d74f3d973cdc20498651fe1504df67dee Mon Sep 17 00:00:00 2001 From: David Katz Date: Thu, 21 Aug 2025 00:40:19 -0400 Subject: [PATCH 20/31] fmt --- crates/goose/src/providers/openai.rs | 8 +++++--- 1 file changed, 5 insertions(+), 3 deletions(-) diff --git a/crates/goose/src/providers/openai.rs b/crates/goose/src/providers/openai.rs index fe86e5c31d5d..1afb88ac3025 100644 --- a/crates/goose/src/providers/openai.rs +++ b/crates/goose/src/providers/openai.rs @@ -17,8 +17,7 @@ use super::embedding::{EmbeddingCapable, EmbeddingRequest, EmbeddingResponse}; use super::errors::ProviderError; use super::formats::openai::{create_request, get_usage, response_to_message}; use super::utils::{ - emit_debug_trace, handle_response_openai_compat, handle_status_openai_compat, - ImageFormat, + emit_debug_trace, handle_response_openai_compat, handle_status_openai_compat, ImageFormat, }; use crate::config::custom_providers::CustomProviderConfig; use crate::conversation::message::Message; @@ -219,7 +218,10 @@ impl Provider for OpenAiProvider { Usage::default() }); emit_debug_trace(model_config, &payload, &json_response, &usage); - Ok((message, ProviderUsage::new(model_config.model_name.clone(), usage))) + Ok(( + message, + ProviderUsage::new(model_config.model_name.clone(), usage), + )) } async fn fetch_supported_models(&self) -> Result>, ProviderError> { From 8cc66f1ac647338c55fe7e352933d6ecc6378bee Mon Sep 17 00:00:00 2001 From: David Katz Date: Thu, 21 Aug 2025 00:49:47 -0400 Subject: [PATCH 21/31] support databricks --- crates/goose/src/providers/databricks.rs | 2 ++ 1 file changed, 2 insertions(+) diff --git a/crates/goose/src/providers/databricks.rs b/crates/goose/src/providers/databricks.rs index 0bc80ed1d4e5..c3d9e8087e62 100644 --- a/crates/goose/src/providers/databricks.rs +++ b/crates/goose/src/providers/databricks.rs @@ -37,6 +37,7 @@ const DEFAULT_SCOPES: &[&str] = &["all-apis", "offline_access"]; const DEFAULT_TIMEOUT_SECS: u64 = 600; pub const DATABRICKS_DEFAULT_MODEL: &str = "databricks-claude-3-7-sonnet"; +const DATABRICKS_DEFAULT_FAST_MODEL: &str = "claude-3-5-haiku"; pub const DATABRICKS_KNOWN_MODELS: &[&str] = &[ "databricks-meta-llama-3-3-70b-instruct", "databricks-meta-llama-3-1-405b-instruct", @@ -109,6 +110,7 @@ impl_provider_default!(DatabricksProvider); impl DatabricksProvider { pub fn from_env(model: ModelConfig) -> Result { let config = crate::config::Config::global(); + let model = model.with_fast(DATABRICKS_DEFAULT_FAST_MODEL.to_string()); let mut host: Result = config.get_param("DATABRICKS_HOST"); if host.is_err() { From 4ec0d9b6e98165ae3cfc594d07da7ecbae200b36 Mon Sep 17 00:00:00 2001 From: David Katz Date: Thu, 21 Aug 2025 01:10:35 -0400 Subject: [PATCH 22/31] bring back output --- crates/goose/src/providers/openai.rs | 11 +++++------ 1 file changed, 5 insertions(+), 6 deletions(-) diff --git a/crates/goose/src/providers/openai.rs b/crates/goose/src/providers/openai.rs index 1afb88ac3025..7a1398cfbfa0 100644 --- a/crates/goose/src/providers/openai.rs +++ b/crates/goose/src/providers/openai.rs @@ -17,7 +17,8 @@ use super::embedding::{EmbeddingCapable, EmbeddingRequest, EmbeddingResponse}; use super::errors::ProviderError; use super::formats::openai::{create_request, get_usage, response_to_message}; use super::utils::{ - emit_debug_trace, handle_response_openai_compat, handle_status_openai_compat, ImageFormat, + emit_debug_trace, get_model, handle_response_openai_compat, handle_status_openai_compat, + ImageFormat, }; use crate::config::custom_providers::CustomProviderConfig; use crate::conversation::message::Message; @@ -217,11 +218,9 @@ impl Provider for OpenAiProvider { tracing::debug!("Failed to get usage data"); Usage::default() }); - emit_debug_trace(model_config, &payload, &json_response, &usage); - Ok(( - message, - ProviderUsage::new(model_config.model_name.clone(), usage), - )) + let model = get_model(&json_response); + emit_debug_trace(&self.model, &payload, &json_response, &usage); + Ok((message, ProviderUsage::new(model, usage))) } async fn fetch_supported_models(&self) -> Result>, ProviderError> { From bdbfcf4a1d5f4b7532009dae8b6abdce15e6de72 Mon Sep 17 00:00:00 2001 From: David Katz Date: Thu, 21 Aug 2025 01:20:58 -0400 Subject: [PATCH 23/31] fix databricks --- crates/goose/src/providers/databricks.rs | 17 ++++++++++------- 1 file changed, 10 insertions(+), 7 deletions(-) diff --git a/crates/goose/src/providers/databricks.rs b/crates/goose/src/providers/databricks.rs index c3d9e8087e62..7bca4b72892f 100644 --- a/crates/goose/src/providers/databricks.rs +++ b/crates/goose/src/providers/databricks.rs @@ -197,17 +197,18 @@ impl DatabricksProvider { }) } - fn get_endpoint_path(&self, is_embedding: bool) -> String { + fn get_endpoint_path(&self, model_name: &str, is_embedding: bool) -> String { if is_embedding { "serving-endpoints/text-embedding-3-small/invocations".to_string() } else { - format!("serving-endpoints/{}/invocations", self.model.model_name) + format!("serving-endpoints/{}/invocations", model_name) } } - async fn post(&self, payload: Value) -> Result { + async fn post(&self, payload: Value, model_name: Option<&str>) -> Result { let is_embedding = payload.get("input").is_some() && payload.get("messages").is_none(); - let path = self.get_endpoint_path(is_embedding); + let model_to_use = model_name.unwrap_or(&self.model.model_name); + let path = self.get_endpoint_path(model_to_use, is_embedding); let response = self.api_client.response_post(&path, &payload).await?; handle_response_openai_compat(response).await @@ -257,7 +258,9 @@ impl Provider for DatabricksProvider { .expect("payload should have model key") .remove("model"); - let response = self.with_retry(|| self.post(payload.clone())).await?; + let response = self + .with_retry(|| self.post(payload.clone(), Some(&model_config.model_name))) + .await?; let message = response_to_message(&response)?; let usage = response.get("usage").map(get_usage).unwrap_or_else(|| { @@ -290,7 +293,7 @@ impl Provider for DatabricksProvider { .unwrap() .insert("stream".to_string(), Value::Bool(true)); - let path = self.get_endpoint_path(false); + let path = self.get_endpoint_path(&model_config.model_name, false); let response = self .with_retry(|| async { let resp = self.api_client.response_post(&path, &payload).await?; @@ -415,7 +418,7 @@ impl EmbeddingCapable for DatabricksProvider { "input": texts, }); - let response = self.with_retry(|| self.post(request.clone())).await?; + let response = self.with_retry(|| self.post(request.clone(), None)).await?; let embeddings = response["data"] .as_array() From 7a5f390bc2b5d0552c195762104943ac0a15973b Mon Sep 17 00:00:00 2001 From: David Katz Date: Thu, 21 Aug 2025 10:47:31 -0400 Subject: [PATCH 24/31] databricks to 3.7sonnet --- crates/goose/src/providers/databricks.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/crates/goose/src/providers/databricks.rs b/crates/goose/src/providers/databricks.rs index 7bca4b72892f..4de326123e06 100644 --- a/crates/goose/src/providers/databricks.rs +++ b/crates/goose/src/providers/databricks.rs @@ -37,7 +37,7 @@ const DEFAULT_SCOPES: &[&str] = &["all-apis", "offline_access"]; const DEFAULT_TIMEOUT_SECS: u64 = 600; pub const DATABRICKS_DEFAULT_MODEL: &str = "databricks-claude-3-7-sonnet"; -const DATABRICKS_DEFAULT_FAST_MODEL: &str = "claude-3-5-haiku"; +const DATABRICKS_DEFAULT_FAST_MODEL: &str = "claude-3-7-sonnet"; pub const DATABRICKS_KNOWN_MODELS: &[&str] = &[ "databricks-meta-llama-3-3-70b-instruct", "databricks-meta-llama-3-1-405b-instruct", From 1e6f164c36f25542e0f0a8113847af5cb874191f Mon Sep 17 00:00:00 2001 From: David Katz Date: Thu, 21 Aug 2025 10:50:39 -0400 Subject: [PATCH 25/31] fix clippy --- crates/goose/src/providers/groq.rs | 2 +- crates/goose/src/providers/litellm.rs | 2 +- crates/goose/src/providers/xai.rs | 2 +- 3 files changed, 3 insertions(+), 3 deletions(-) diff --git a/crates/goose/src/providers/groq.rs b/crates/goose/src/providers/groq.rs index 91825f3d41ab..e70a9d36aa58 100644 --- a/crates/goose/src/providers/groq.rs +++ b/crates/goose/src/providers/groq.rs @@ -88,7 +88,7 @@ impl Provider for GroqProvider { tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { let payload = create_request( - &model_config, + model_config, system, messages, tools, diff --git a/crates/goose/src/providers/litellm.rs b/crates/goose/src/providers/litellm.rs index 41a3a89f1618..d08e6b5c286b 100644 --- a/crates/goose/src/providers/litellm.rs +++ b/crates/goose/src/providers/litellm.rs @@ -169,7 +169,7 @@ impl Provider for LiteLLMProvider { tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { let mut payload = super::formats::openai::create_request( - &model_config, + model_config, system, messages, tools, diff --git a/crates/goose/src/providers/xai.rs b/crates/goose/src/providers/xai.rs index 9b862eeaddf4..65ecccbd7f36 100644 --- a/crates/goose/src/providers/xai.rs +++ b/crates/goose/src/providers/xai.rs @@ -104,7 +104,7 @@ impl Provider for XaiProvider { tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { let payload = create_request( - &model_config, + model_config, system, messages, tools, From 395f4343b48ad245b33d6123595b84f36e19c7e7 Mon Sep 17 00:00:00 2001 From: David Katz Date: Thu, 21 Aug 2025 11:09:51 -0400 Subject: [PATCH 26/31] summary model -> 1.5flash --- crates/goose/src/providers/databricks.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/crates/goose/src/providers/databricks.rs b/crates/goose/src/providers/databricks.rs index 4de326123e06..6611d335665f 100644 --- a/crates/goose/src/providers/databricks.rs +++ b/crates/goose/src/providers/databricks.rs @@ -37,7 +37,7 @@ const DEFAULT_SCOPES: &[&str] = &["all-apis", "offline_access"]; const DEFAULT_TIMEOUT_SECS: u64 = 600; pub const DATABRICKS_DEFAULT_MODEL: &str = "databricks-claude-3-7-sonnet"; -const DATABRICKS_DEFAULT_FAST_MODEL: &str = "claude-3-7-sonnet"; +const DATABRICKS_DEFAULT_FAST_MODEL: &str = "gemini-1-5-flash"; pub const DATABRICKS_KNOWN_MODELS: &[&str] = &[ "databricks-meta-llama-3-3-70b-instruct", "databricks-meta-llama-3-1-405b-instruct", From f99626224acfea9df5078da896e78388f40084a9 Mon Sep 17 00:00:00 2001 From: David Katz Date: Thu, 21 Aug 2025 11:30:54 -0400 Subject: [PATCH 27/31] fix titrate --- crates/goose/src/providers/tetrate.rs | 36 ++++++++++----------------- 1 file changed, 13 insertions(+), 23 deletions(-) diff --git a/crates/goose/src/providers/tetrate.rs b/crates/goose/src/providers/tetrate.rs index ce865f4cc519..2aacfc9e0b9c 100644 --- a/crates/goose/src/providers/tetrate.rs +++ b/crates/goose/src/providers/tetrate.rs @@ -1,4 +1,4 @@ -use anyhow::{Error, Result}; +use anyhow::Result; use async_trait::async_trait; use serde_json::Value; @@ -113,23 +113,6 @@ impl TetrateProvider { } } -fn create_request_based_on_model( - provider: &TetrateProvider, - system: &str, - messages: &[Message], - tools: &[Tool], -) -> anyhow::Result { - let payload = create_request( - &provider.model, - system, - messages, - tools, - &super::utils::ImageFormat::OpenAi, - )?; - - Ok(payload) -} - #[async_trait] impl Provider for TetrateProvider { fn metadata() -> ProviderMetadata { @@ -157,17 +140,24 @@ impl Provider for TetrateProvider { } #[tracing::instrument( - skip(self, system, messages, tools), + skip(self, model_config, system, messages, tools), fields(model_config, input, output, input_tokens, output_tokens, total_tokens) )] - async fn complete( + async fn complete_with_model( &self, + model_config: &ModelConfig, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(Message, ProviderUsage), ProviderError> { - // Create the base payload - let payload = create_request_based_on_model(self, system, messages, tools)?; + // Create the base payload using the provided model_config + let payload = create_request( + model_config, + system, + messages, + tools, + &super::utils::ImageFormat::OpenAi, + )?; // Make request let response = self @@ -184,7 +174,7 @@ impl Provider for TetrateProvider { Usage::default() }); let model = get_model(&response); - emit_debug_trace(&self.model, &payload, &response, &usage); + emit_debug_trace(model_config, &payload, &response, &usage); Ok((message, ProviderUsage::new(model, usage))) } From 9e43de98ddeeb6b7ac93a1139789d95d639916c1 Mon Sep 17 00:00:00 2001 From: David Katz Date: Thu, 21 Aug 2025 11:51:16 -0400 Subject: [PATCH 28/31] one more test fix --- crates/goose-server/src/routes/reply.rs | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/crates/goose-server/src/routes/reply.rs b/crates/goose-server/src/routes/reply.rs index 66480c6fb070..dd45ff598903 100644 --- a/crates/goose-server/src/routes/reply.rs +++ b/crates/goose-server/src/routes/reply.rs @@ -537,8 +537,9 @@ mod tests { goose::providers::base::ProviderMetadata::empty() } - async fn complete( + async fn complete_with_model( &self, + _model_config: &ModelConfig, _system: &str, _messages: &[Message], _tools: &[rmcp::model::Tool], From 2f232f8553f7aba1cf5c2a709aacf094eb4041f2 Mon Sep 17 00:00:00 2001 From: David Katz Date: Thu, 21 Aug 2025 12:47:12 -0400 Subject: [PATCH 29/31] fix agent tests --- crates/goose/tests/agent.rs | 40 +++++++++++++++++++++++++++++++++++++ 1 file changed, 40 insertions(+) diff --git a/crates/goose/tests/agent.rs b/crates/goose/tests/agent.rs index b4659332122b..9c7e84838c0a 100644 --- a/crates/goose/tests/agent.rs +++ b/crates/goose/tests/agent.rs @@ -592,6 +592,16 @@ mod final_output_tool_tests { ProviderUsage::new("mock".to_string(), Usage::default()), )) } + + async fn complete_with_model( + &self, + _model_config: &ModelConfig, + system: &str, + messages: &[Message], + tools: &[Tool], + ) -> anyhow::Result<(Message, ProviderUsage), ProviderError> { + self.complete(system, messages, tools).await + } } let agent = Agent::new(); @@ -713,6 +723,16 @@ mod final_output_tool_tests { ) -> Result<(Message, ProviderUsage), ProviderError> { Err(ProviderError::NotImplemented("Not implemented".to_string())) } + + async fn complete_with_model( + &self, + _model_config: &ModelConfig, + system: &str, + messages: &[Message], + tools: &[Tool], + ) -> anyhow::Result<(Message, ProviderUsage), ProviderError> { + self.complete(system, messages, tools).await + } } let agent = Agent::new(); @@ -829,6 +849,16 @@ mod retry_tests { )) } } + + async fn complete_with_model( + &self, + _model_config: &ModelConfig, + system: &str, + messages: &[Message], + tools: &[Tool], + ) -> anyhow::Result<(Message, ProviderUsage), ProviderError> { + self.complete(system, messages, tools).await + } } #[tokio::test] @@ -1002,6 +1032,16 @@ mod max_turns_tests { Ok((message, usage)) } + async fn complete_with_model( + &self, + _model_config: &ModelConfig, + system_prompt: &str, + messages: &[Message], + tools: &[Tool], + ) -> anyhow::Result<(Message, ProviderUsage), ProviderError> { + self.complete(system_prompt, messages, tools).await + } + fn get_model_config(&self) -> ModelConfig { ModelConfig::new("mock-model").unwrap() } From 0cd418979ce0be60b49d91e23ac2ad12d1ed6d5f Mon Sep 17 00:00:00 2001 From: David Katz Date: Thu, 21 Aug 2025 13:38:55 -0400 Subject: [PATCH 30/31] add fast model exists check --- crates/goose/src/providers/databricks.rs | 35 +++++++++++++++++++++--- 1 file changed, 31 insertions(+), 4 deletions(-) diff --git a/crates/goose/src/providers/databricks.rs b/crates/goose/src/providers/databricks.rs index 6611d335665f..0172eba58df8 100644 --- a/crates/goose/src/providers/databricks.rs +++ b/crates/goose/src/providers/databricks.rs @@ -110,7 +110,6 @@ impl_provider_default!(DatabricksProvider); impl DatabricksProvider { pub fn from_env(model: ModelConfig) -> Result { let config = crate::config::Config::global(); - let model = model.with_fast(DATABRICKS_DEFAULT_FAST_MODEL.to_string()); let mut host: Result = config.get_param("DATABRICKS_HOST"); if host.is_err() { @@ -139,13 +138,41 @@ impl DatabricksProvider { let api_client = ApiClient::with_timeout(host, auth_method, Duration::from_secs(DEFAULT_TIMEOUT_SECS))?; - Ok(Self { + // Create the provider without the fast model first + let mut provider = Self { api_client, auth, - model, + model: model.clone(), image_format: ImageFormat::OpenAi, retry_config, - }) + }; + + // Check if the default fast model exists in the workspace + let model_with_fast = tokio::task::block_in_place(|| { + tokio::runtime::Handle::current().block_on(async { + if let Ok(Some(models)) = provider.fetch_supported_models().await { + if models.contains(&DATABRICKS_DEFAULT_FAST_MODEL.to_string()) { + tracing::debug!( + "Found {} in Databricks workspace, setting as fast model", + DATABRICKS_DEFAULT_FAST_MODEL + ); + model.with_fast(DATABRICKS_DEFAULT_FAST_MODEL.to_string()) + } else { + tracing::debug!( + "{} not found in Databricks workspace, not setting fast model", + DATABRICKS_DEFAULT_FAST_MODEL + ); + model + } + } else { + tracing::debug!("Could not fetch Databricks models, not setting fast model"); + model + } + }) + }); + + provider.model = model_with_fast; + Ok(provider) } fn load_retry_config(config: &crate::config::Config) -> RetryConfig { From 74a83594bae3290d5d451e5ae8cf012dbb2fb652 Mon Sep 17 00:00:00 2001 From: David Katz Date: Thu, 21 Aug 2025 14:51:51 -0400 Subject: [PATCH 31/31] bump sonnet model --- crates/goose/src/providers/anthropic.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/crates/goose/src/providers/anthropic.rs b/crates/goose/src/providers/anthropic.rs index 2b51aab7e7cb..9db8d048e851 100644 --- a/crates/goose/src/providers/anthropic.rs +++ b/crates/goose/src/providers/anthropic.rs @@ -23,7 +23,7 @@ use crate::providers::retry::ProviderRetry; use rmcp::model::Tool; const ANTHROPIC_DEFAULT_MODEL: &str = "claude-sonnet-4-0"; -const ANTHROPIC_DEFAULT_FAST_MODEL: &str = "claude-3-5-haiku-latest"; +const ANTHROPIC_DEFAULT_FAST_MODEL: &str = "claude-3-7-sonnet-latest"; const ANTHROPIC_KNOWN_MODELS: &[&str] = &[ "claude-sonnet-4-0", "claude-sonnet-4-20250514",