From 61cd6f636626ccf86f73a16e8c380b6867b5fa68 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Stadler=20Gell=C3=A9rt?= Date: Thu, 25 Jun 2026 08:52:39 +0200 Subject: [PATCH 1/5] Improve Bedrock model discovery and validation --- Cargo.lock | 26 ++ crates/goose-server/src/routes/agent.rs | 6 + .../src/routes/config_management.rs | 17 +- crates/goose/Cargo.toml | 2 + crates/goose/src/providers/bedrock.rs | 397 ++++++++++++++++-- crates/goose/src/providers/init.rs | 5 +- .../src/providers/inventory/registrations.rs | 16 + crates/goose/src/providers/provider_test.rs | 26 +- .../models/subcomponents/SwitchModelModal.tsx | 69 +-- 9 files changed, 493 insertions(+), 71 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index d4a76d6a8ffc..d483367b3b14 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -580,6 +580,31 @@ dependencies = [ "uuid", ] +[[package]] +name = "aws-sdk-bedrock" +version = "1.146.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bd3d870793928e6c18de8273dcdceed08af8cc9083be9d7de3e52efc7ac01e9e" +dependencies = [ + "arc-swap", + "aws-credential-types", + "aws-runtime", + "aws-smithy-async", + "aws-smithy-http", + "aws-smithy-json", + "aws-smithy-observability", + "aws-smithy-runtime", + "aws-smithy-runtime-api", + "aws-smithy-types", + "aws-types", + "bytes", + "fastrand", + "http 0.2.12", + "http 1.4.2", + "regex-lite", + "tracing", +] + [[package]] name = "aws-sdk-bedrockruntime" version = "1.133.0" @@ -4761,6 +4786,7 @@ dependencies = [ "async-stream", "async-trait", "aws-config", + "aws-sdk-bedrock", "aws-sdk-bedrockruntime", "aws-sdk-sagemakerruntime", "aws-smithy-types", diff --git a/crates/goose-server/src/routes/agent.rs b/crates/goose-server/src/routes/agent.rs index c93f6b1596f5..1908afb8ffb4 100644 --- a/crates/goose-server/src/routes/agent.rs +++ b/crates/goose-server/src/routes/agent.rs @@ -20,6 +20,7 @@ use goose::agents::ExtensionConfig; use goose::config::resolve_extensions_for_new_session; use goose::config::{Config, GooseMode}; use goose::providers::create; +use goose::providers::provider_test::test_provider_model; use goose::recipe::Recipe; use goose::recipe_deeplink; use goose::session::session_manager::SessionType; @@ -651,6 +652,11 @@ async fn update_agent_provider( if let Some(request_params) = payload.request_params { model_config = model_config.with_merged_request_params(request_params); } + + test_provider_model(&payload.provider, &model) + .await + .map_err(|e| (StatusCode::BAD_REQUEST, e.to_string()))?; + let model_info = resolve_provider_model_info(&payload.provider, &model) .await .map_err(|e| (e.status, e.message))?; diff --git a/crates/goose-server/src/routes/config_management.rs b/crates/goose-server/src/routes/config_management.rs index 4ad7f189fa25..e3c7899f9b37 100644 --- a/crates/goose-server/src/routes/config_management.rs +++ b/crates/goose-server/src/routes/config_management.rs @@ -21,6 +21,7 @@ use goose::providers::catalog::{ }; use goose::providers::create_with_default_model; use goose::providers::huggingface_auth; +use goose::providers::provider_test::test_provider_model; use goose::providers::providers as get_providers; use goose::{ agents::execute_commands, agents::ExtensionConfig, config::permission::PermissionLevel, @@ -1341,20 +1342,22 @@ pub async fn check_provider( pub async fn set_config_provider( Json(SetProviderRequest { provider, model }): Json, ) -> Result<(), ErrorResponse> { - // Provider validation does not use extensions. - create_with_default_model(&provider, Vec::new()) + test_provider_model(&provider, &model) .await - .and_then(|_| { - let config = Config::global(); - goose::config::set_active_provider(config, &provider, &model) - .map_err(|e| anyhow::anyhow!(e)) - }) .map_err(|err| { ErrorResponse::bad_request(format!( "Failed to set provider to '{}' with model '{}': {}", provider, model, err )) })?; + + let config = Config::global(); + goose::config::set_active_provider(config, &provider, &model).map_err(|err| { + ErrorResponse::bad_request(format!( + "Failed to set provider to '{}' with model '{}': {}", + provider, model, err + )) + })?; Ok(()) } diff --git a/crates/goose/Cargo.toml b/crates/goose/Cargo.toml index 9a52ef60f3e3..4a7487105d95 100644 --- a/crates/goose/Cargo.toml +++ b/crates/goose/Cargo.toml @@ -35,6 +35,7 @@ local-inference = [ aws-providers = [ "dep:aws-config", "dep:aws-smithy-types", + "dep:aws-sdk-bedrock", "dep:aws-sdk-bedrockruntime", "dep:aws-sdk-sagemakerruntime", "dep:smithy-transport-reqwest", @@ -218,6 +219,7 @@ llama-cpp-sys-2 = { workspace = true, optional = true } image = { version = "0.24.9", default-features = false, features = ["png", "jpeg", "gif", "webp"] } subtle = { version = "2.5", default-features = false, features = ["std"] } gethostname = "1.1.0" +aws-sdk-bedrock = { version = "1.132", default-features = false, features = ["rt-tokio"], optional = true } [target.'cfg(target_os = "windows")'.dependencies] winapi = { version = "0.3.9", default-features = false, features = ["wincred", "std", "jobapi2", "winbase", "winnt", "processthreadsapi", "handleapi", "minwindef"] } diff --git a/crates/goose/src/providers/bedrock.rs b/crates/goose/src/providers/bedrock.rs index b23bb89b90d3..9c9699631e7b 100644 --- a/crates/goose/src/providers/bedrock.rs +++ b/crates/goose/src/providers/bedrock.rs @@ -1,4 +1,4 @@ -use std::collections::HashMap; +use std::collections::{BTreeSet, HashMap}; use super::base::{ConfigKey, MessageStream, Provider, ProviderDef, ProviderMetadata}; use super::openai_compatible::{handle_status, stream_responses_compat}; @@ -8,11 +8,16 @@ use crate::session_context::SESSION_ID_HEADER; use anyhow::Result; use async_stream::try_stream; use async_trait::async_trait; +use aws_sdk_bedrock::types::{ + FoundationModelLifecycleStatus, FoundationModelSummary, InferenceProfileStatus, + InferenceProfileType, InferenceType, ModelModality, +}; +use aws_sdk_bedrock::Client as BedrockControlClient; use aws_sdk_bedrockruntime::config::ProvideCredentials; use aws_sdk_bedrockruntime::operation::converse::ConverseError; use aws_sdk_bedrockruntime::operation::converse_stream::ConverseStreamError; use aws_sdk_bedrockruntime::types::error::ConverseStreamOutputError; -use aws_sdk_bedrockruntime::{types as bedrock, Client}; +use aws_sdk_bedrockruntime::{types as bedrock, Client as BedrockRuntimeClient}; use base64::Engine; use futures::future::BoxFuture; use goose_providers::conversation::token_usage::{ProviderUsage, Usage}; @@ -36,15 +41,8 @@ pub const BEDROCK_DOC_LINK: &str = "https://docs.aws.amazon.com/bedrock/latest/userguide/models-supported.html"; pub const BEDROCK_DEFAULT_MODEL: &str = "us.anthropic.claude-sonnet-4-5-20250929-v1:0"; -pub const BEDROCK_KNOWN_MODELS: &[&str] = &[ - "us.anthropic.claude-sonnet-4-5-20250929-v1:0", - "us.anthropic.claude-sonnet-4-20250514-v1:0", - "us.anthropic.claude-3-7-sonnet-20250219-v1:0", - "us.anthropic.claude-opus-4-20250514-v1:0", - "us.anthropic.claude-opus-4-1-20250805-v1:0", - "openai.gpt-5.5", - "openai.gpt-5.4", -]; +pub const BEDROCK_MANTLE_MODELS: &[&str] = &["openai.gpt-5.5", "openai.gpt-5.4"]; +pub const BEDROCK_BOOTSTRAP_MODELS: &[&str] = &[BEDROCK_DEFAULT_MODEL]; pub const BEDROCK_DEFAULT_MAX_RETRIES: usize = 6; pub const BEDROCK_DEFAULT_INITIAL_RETRY_INTERVAL_MS: u64 = 2000; @@ -54,7 +52,9 @@ pub const BEDROCK_DEFAULT_MAX_RETRY_INTERVAL_MS: u64 = 120_000; #[derive(Debug, serde::Serialize)] pub struct BedrockProvider { #[serde(skip)] - client: Client, + client: BedrockRuntimeClient, + #[serde(skip)] + control_plane_client: BedrockControlClient, #[serde(skip)] retry_config: RetryConfig, #[serde(skip)] @@ -162,15 +162,26 @@ impl BedrockProvider { )) .build(); - Client::from_conf(bedrock_config) + BedrockRuntimeClient::from_conf(bedrock_config) + } else { + Self::create_runtime_client_with_credentials(&sdk_config).await? + }; + + let control_plane_client = if let Some(ref token) = bearer_token { + let bedrock_config = aws_sdk_bedrock::Config::new(&sdk_config) + .to_builder() + .bearer_token(aws_sdk_bedrock::config::Token::new(token.clone(), None)) + .build(); + BedrockControlClient::from_conf(bedrock_config) } else { - Self::create_client_with_credentials(&sdk_config).await? + BedrockControlClient::new(&sdk_config) }; let retry_config = Self::load_retry_config(config); Ok(Self { client, + control_plane_client, retry_config, name: BEDROCK_PROVIDER_NAME.to_string(), region: resolved_region, @@ -180,7 +191,9 @@ impl BedrockProvider { }) } - async fn create_client_with_credentials(sdk_config: &aws_config::SdkConfig) -> Result { + async fn create_runtime_client_with_credentials( + sdk_config: &aws_config::SdkConfig, + ) -> Result { sdk_config .credentials_provider() .ok_or_else(|| anyhow::anyhow!("No AWS credentials provider configured"))? @@ -193,7 +206,162 @@ impl BedrockProvider { ) })?; - Ok(Client::new(sdk_config)) + Ok(BedrockRuntimeClient::new(sdk_config)) + } + + fn bootstrap_models(include_mantle: bool) -> Vec { + let mut models: Vec = BEDROCK_BOOTSTRAP_MODELS + .iter() + .map(|model| model.to_string()) + .collect(); + if include_mantle { + models.extend(BEDROCK_MANTLE_MODELS.iter().map(|model| model.to_string())); + } + models + } + + fn merge_discovered_model_ids( + inference_profiles: impl IntoIterator, + foundation_models: impl IntoIterator, + extra_models: impl IntoIterator, + ) -> Vec { + let mut models = BTreeSet::new(); + models.extend(inference_profiles); + models.extend(foundation_models); + models.extend(extra_models); + models.into_iter().collect() + } + + fn is_excluded_non_chat_model_id(model_id: &str) -> bool { + let id = model_id.to_lowercase(); + id.contains(".embed") + || id.contains("-embed-") + || id.ends_with("-embed") + || id.contains(".rerank") + || id.contains("-rerank-") + || id.ends_with("-rerank") + } + + fn is_chat_capable_model_id(model_id: &str) -> bool { + !Self::is_excluded_non_chat_model_id(model_id) + } + + fn is_chat_capable_foundation_model(summary: &FoundationModelSummary) -> bool { + if !Self::is_chat_capable_model_id(summary.model_id()) { + return false; + } + + let has_text_input = summary + .input_modalities() + .iter() + .any(|modality| *modality == ModelModality::Text); + let has_text_output = summary + .output_modalities() + .iter() + .any(|modality| *modality == ModelModality::Text); + let has_embedding_output = summary + .output_modalities() + .iter() + .any(|modality| *modality == ModelModality::Embedding); + let streaming = summary.response_streaming_supported().unwrap_or(false); + let not_legacy = summary + .model_lifecycle() + .map(|lifecycle| lifecycle.status() != &FoundationModelLifecycleStatus::Legacy) + .unwrap_or(true); + + has_text_input && has_text_output && !has_embedding_output && streaming && not_legacy + } + + fn is_mantle_model_name(model_name: &str) -> bool { + model_name + .strip_prefix("openai.") + .unwrap_or(model_name) + .starts_with("gpt-") + } + + async fn fetch_inference_profile_ids(&self) -> Result, ProviderError> { + let mut ids = BTreeSet::new(); + + for profile_type in [ + InferenceProfileType::SystemDefined, + InferenceProfileType::Application, + ] { + let mut next_token = None; + loop { + let mut request = self + .control_plane_client + .list_inference_profiles() + .type_equals(profile_type.clone()); + if let Some(token) = &next_token { + request = request.next_token(token); + } + + let response = request.send().await.map_err(|err| { + ProviderError::ExecutionError(format!( + "Failed to list Bedrock inference profiles: {}", + err + )) + })?; + + for summary in response.inference_profile_summaries() { + if summary.status() == &InferenceProfileStatus::Active { + let id = summary.inference_profile_id(); + if Self::is_chat_capable_model_id(id) { + ids.insert(id.to_string()); + } + } + } + + next_token = response.next_token().map(|token| token.to_string()); + if next_token.is_none() { + break; + } + } + } + + Ok(ids.into_iter().collect()) + } + + async fn fetch_foundation_model_ids(&self) -> Result, ProviderError> { + let response = self + .control_plane_client + .list_foundation_models() + .by_inference_type(InferenceType::OnDemand) + .by_output_modality(ModelModality::Text) + .send() + .await + .map_err(|err| { + ProviderError::ExecutionError(format!( + "Failed to list Bedrock foundation models: {}", + err + )) + })?; + + Ok(response + .model_summaries() + .iter() + .filter(|summary| Self::is_chat_capable_foundation_model(summary)) + .map(|summary| summary.model_id().to_string()) + .collect()) + } + + async fn fetch_models_from_aws(&self) -> Result, ProviderError> { + let inference_profiles = self.fetch_inference_profile_ids().await?; + let foundation_models = self.fetch_foundation_model_ids().await?; + let extra_models = if self.bearer_token.is_some() { + BEDROCK_MANTLE_MODELS + .iter() + .map(|model| model.to_string()) + .collect() + } else { + Vec::new() + }; + + Ok(Self::merge_discovered_model_ids( + inference_profiles, + foundation_models, + extra_models, + )) } fn load_retry_config(config: &crate::config::Config) -> RetryConfig { @@ -399,6 +567,12 @@ impl BedrockProvider { "Bedrock validation error: {}", err.message().unwrap_or("unknown validation error") )), + ConverseError::ResourceNotFoundException(err) => { + ProviderError::ExecutionError(format!( + "Bedrock model not found or not accessible: {}", + err.message().unwrap_or("unknown resource error") + )) + } ConverseError::ModelErrorException(err) => { ProviderError::ExecutionError(format!("Failed to call Bedrock: {:?}", err)) } @@ -490,6 +664,18 @@ impl BedrockProvider { err )) } + ConverseStreamError::ValidationException(err) => { + ProviderError::ExecutionError(format!( + "Bedrock validation error: {}", + err.message().unwrap_or("unknown validation error") + )) + } + ConverseStreamError::ResourceNotFoundException(err) => { + ProviderError::ExecutionError(format!( + "Bedrock model not found or not accessible: {}", + err.message().unwrap_or("unknown resource error") + )) + } ConverseStreamError::ModelErrorException(err) => { ProviderError::ExecutionError(format!("Failed to call Bedrock: {:?}", err)) } @@ -685,9 +871,9 @@ impl goose_providers::base::ProviderDescriptor for BedrockProvider { ProviderMetadata::new( BEDROCK_PROVIDER_NAME, "Amazon Bedrock", - "Run models through Amazon Bedrock. Supports AWS SSO profiles - run 'aws sso login --profile ' before using. Configure with AWS_PROFILE and AWS_REGION, use environment variables/credentials, or use AWS_BEARER_TOKEN_BEDROCK for bearer token authentication. Region is required for bearer token auth (can be set via AWS_REGION, AWS_DEFAULT_REGION, or AWS profile). Prompt caching can be enabled for Anthropic Claude models by setting BEDROCK_ENABLE_CACHING=true. Responses stream via the ConverseStream API; set BEDROCK_DISABLE_STREAMING=true to fall back to blocking Converse calls.", + "Run models through Amazon Bedrock. Supports AWS SSO profiles - run 'aws sso login --profile ' before using. Configure with AWS_PROFILE and AWS_REGION, use environment variables/credentials, or use AWS_BEARER_TOKEN_BEDROCK for bearer token authentication. Region is required for bearer token auth (can be set via AWS_REGION, AWS_DEFAULT_REGION, or AWS profile). Model discovery requires bedrock:ListFoundationModels and bedrock:ListInferenceProfiles permissions. Prompt caching can be enabled for Anthropic Claude models by setting BEDROCK_ENABLE_CACHING=true. Responses stream via the ConverseStream API; set BEDROCK_DISABLE_STREAMING=true to fall back to blocking Converse calls.", BEDROCK_DEFAULT_MODEL, - BEDROCK_KNOWN_MODELS.to_vec(), + BEDROCK_BOOTSTRAP_MODELS.to_vec(), BEDROCK_DOC_LINK, vec![ ConfigKey::new("AWS_PROFILE", false, false, Some("default"), true), @@ -727,8 +913,25 @@ impl Provider for BedrockProvider { self.retry_config.clone() } + fn skip_canonical_filtering(&self) -> bool { + true + } + async fn fetch_supported_models(&self) -> Result, ProviderError> { - Ok(BEDROCK_KNOWN_MODELS.iter().map(|s| s.to_string()).collect()) + match self.fetch_models_from_aws().await { + Ok(models) if !models.is_empty() => Ok(models), + Ok(_) => { + tracing::debug!("Bedrock model discovery returned no models, using bootstrap list"); + Ok(Self::bootstrap_models(self.bearer_token.is_some())) + } + Err(err) => { + tracing::warn!( + "Bedrock model discovery failed ({}), using bootstrap list", + err + ); + Ok(Self::bootstrap_models(self.bearer_token.is_some())) + } + } } async fn stream( @@ -752,7 +955,7 @@ impl Provider for BedrockProvider { let (base_name, effort) = extract_reasoning_effort(without_prefix); let bedrock_model_id = format!("openai.{}", base_name); - let is_mantle_model = BEDROCK_KNOWN_MODELS.contains(&bedrock_model_id.as_str()); + let is_mantle_model = Self::is_mantle_model_name(&model_config.model_name); if is_mantle_model { let mut normalized_config = ModelConfig { @@ -915,11 +1118,23 @@ mod tests { .behavior_version(aws_config::BehaviorVersion::latest()) .region(aws_config::Region::new("us-east-1")) .build(); - let client = Client::new(&sdk_config); + let client = BedrockRuntimeClient::new(&sdk_config); + let control_plane_client = BedrockControlClient::new(&sdk_config); + let model = ModelConfig { + model_name: model_name.to_string(), + context_limit: None, + temperature: None, + max_tokens: None, + toolshim: false, + toolshim_model: None, + request_params: None, + reasoning: None, + }; ( BedrockProvider { client, + control_plane_client, retry_config: RetryConfig::default(), name: "aws_bedrock".to_string(), region: None, @@ -927,16 +1142,7 @@ mod tests { http_client: reqwest::Client::new(), mantle_base_url: None, }, - ModelConfig { - model_name: model_name.to_string(), - context_limit: None, - temperature: None, - max_tokens: None, - toolshim: false, - toolshim_model: None, - request_params: None, - reasoning: None, - }, + model, ) } @@ -1002,6 +1208,21 @@ mod tests { ); } + #[test] + #[serial] + fn test_caching_disabled_by_default() { + std::env::set_var("BEDROCK_ENABLE_CACHING", "false"); + + let (provider, model) = + create_mock_provider_and_model("us.anthropic.claude-sonnet-4-5-20250929-v1:0"); + assert!( + !provider.should_enable_caching(&model), + "Caching should be disabled by default" + ); + + std::env::remove_var("BEDROCK_ENABLE_CACHING"); + } + #[test] fn test_caching_disabled_for_non_claude_models() { let (provider, model) = create_mock_provider_and_model("amazon.titan-text-express-v1"); @@ -1089,9 +1310,10 @@ mod tests { .region(aws_config::Region::new("us-east-1")) .build(); - let model = ModelConfig::new("openai.gpt-5.5"); + let model = ModelConfig::new("openai.gpt-5.5").unwrap(); let provider = BedrockProvider { - client: Client::new(&sdk_config), + client: BedrockRuntimeClient::new(&sdk_config), + control_plane_client: BedrockControlClient::new(&sdk_config), retry_config: RetryConfig::default(), name: "aws_bedrock".to_string(), region: Some("us-east-1".to_string()), @@ -1493,4 +1715,113 @@ mod tests { other => panic!("expected RedactedThinking, got {:?}", other), } } + + #[test] + fn test_merge_discovered_model_ids_dedupes_and_sorts() { + let models = BedrockProvider::merge_discovered_model_ids( + [ + "eu.anthropic.claude-sonnet-4-5-20250929-v1:0".to_string(), + "us.anthropic.claude-sonnet-4-5-20250929-v1:0".to_string(), + ], + [ + "anthropic.claude-3-5-sonnet-20240620-v1:0".to_string(), + "eu.anthropic.claude-sonnet-4-5-20250929-v1:0".to_string(), + ], + ["openai.gpt-5.5".to_string()], + ); + + assert_eq!( + models, + vec![ + "anthropic.claude-3-5-sonnet-20240620-v1:0".to_string(), + "eu.anthropic.claude-sonnet-4-5-20250929-v1:0".to_string(), + "openai.gpt-5.5".to_string(), + "us.anthropic.claude-sonnet-4-5-20250929-v1:0".to_string(), + ] + ); + } + + #[test] + fn test_bootstrap_models_includes_mantle_with_bearer_token() { + let without_mantle = BedrockProvider::bootstrap_models(false); + assert_eq!(without_mantle, vec![BEDROCK_DEFAULT_MODEL.to_string()]); + + let with_mantle = BedrockProvider::bootstrap_models(true); + assert!(with_mantle.contains(&"openai.gpt-5.5".to_string())); + assert!(with_mantle.contains(&BEDROCK_DEFAULT_MODEL.to_string())); + } + + #[test] + fn test_is_mantle_model_name() { + assert!(BedrockProvider::is_mantle_model_name("openai.gpt-5.5")); + assert!(BedrockProvider::is_mantle_model_name("gpt-5.4")); + assert!(!BedrockProvider::is_mantle_model_name( + "us.anthropic.claude-sonnet-4-5-20250929-v1:0" + )); + } + + #[test] + fn test_skip_canonical_filtering_enabled() { + let provider = create_mock_provider("test"); + assert!(provider.skip_canonical_filtering()); + } + + #[test] + fn test_is_chat_capable_foundation_model_filters_embeddings() { + let chat_model = FoundationModelSummary::builder() + .model_arn("arn:aws:bedrock:us-east-1::foundation-model/anthropic.claude-3-5-sonnet-20240620-v1:0") + .model_id("anthropic.claude-3-5-sonnet-20240620-v1:0") + .input_modalities(ModelModality::Text) + .output_modalities(ModelModality::Text) + .response_streaming_supported(true) + .build() + .unwrap(); + assert!(BedrockProvider::is_chat_capable_foundation_model( + &chat_model + )); + + let embedding_model = FoundationModelSummary::builder() + .model_arn("arn:aws:bedrock:us-east-1::foundation-model/amazon.titan-embed-text-v2:0") + .model_id("amazon.titan-embed-text-v2:0") + .input_modalities(ModelModality::Text) + .output_modalities(ModelModality::Embedding) + .response_streaming_supported(false) + .build() + .unwrap(); + assert!(!BedrockProvider::is_chat_capable_foundation_model( + &embedding_model + )); + + let cohere_embed_v4 = FoundationModelSummary::builder() + .model_arn("arn:aws:bedrock:us-east-1::foundation-model/cohere.embed-v4:0") + .model_id("cohere.embed-v4:0") + .input_modalities(ModelModality::Text) + .output_modalities(ModelModality::Text) + .output_modalities(ModelModality::Embedding) + .response_streaming_supported(true) + .build() + .unwrap(); + assert!(!BedrockProvider::is_chat_capable_foundation_model( + &cohere_embed_v4 + )); + } + + #[test] + fn test_is_chat_capable_model_id_excludes_embed_and_rerank_profiles() { + assert!(!BedrockProvider::is_chat_capable_model_id( + "cohere.embed-v4:0" + )); + assert!(!BedrockProvider::is_chat_capable_model_id( + "us.cohere.embed-v4:0" + )); + assert!(!BedrockProvider::is_chat_capable_model_id( + "amazon.titan-embed-text-v2:0" + )); + assert!(!BedrockProvider::is_chat_capable_model_id( + "cohere.rerank-v3-5:0" + )); + assert!(BedrockProvider::is_chat_capable_model_id( + "us.anthropic.claude-sonnet-4-5-20250929-v1:0" + )); + } } diff --git a/crates/goose/src/providers/init.rs b/crates/goose/src/providers/init.rs index 4861fe7a932b..826cba30aec7 100644 --- a/crates/goose/src/providers/init.rs +++ b/crates/goose/src/providers/init.rs @@ -70,7 +70,10 @@ async fn init_registry() -> RwLock { registry.register::(false); registry.register::(false); #[cfg(feature = "aws-providers")] - registry.register::(false); + registry.register_with_inventory::( + false, + Some(registrations::bedrock_inventory()), + ); #[cfg(feature = "local-inference")] registry.register::(false); registry.register_with_inventory::( diff --git a/crates/goose/src/providers/inventory/registrations.rs b/crates/goose/src/providers/inventory/registrations.rs index d51cc63bb5e9..b57a05d146ff 100644 --- a/crates/goose/src/providers/inventory/registrations.rs +++ b/crates/goose/src/providers/inventory/registrations.rs @@ -6,6 +6,8 @@ use crate::config::Config; use crate::providers::acp_tooling::{acp_adapter_installed, resolved_acp_command}; use crate::providers::amp_acp::{AMP_ACP_BINARY, AMP_ACP_PROVIDER_NAME}; use crate::providers::base::ProviderDescriptor; +#[cfg(feature = "aws-providers")] +use crate::providers::bedrock::{BedrockProvider, BEDROCK_PROVIDER_NAME}; use crate::providers::chatgpt_codex::TokenCache as ChatGptCodexTokenCache; use crate::providers::claude_acp::{CLAUDE_ACP_BINARY, CLAUDE_ACP_PROVIDER_NAME}; use crate::providers::codex_acp::CODEX_ACP_PROVIDER_NAME; @@ -121,6 +123,20 @@ pub fn ollama_inventory() -> InventoryRegistration { .with_configured(|| ollama_host_configured(Config::global())) } +#[cfg(feature = "aws-providers")] +pub fn bedrock_inventory() -> InventoryRegistration { + InventoryRegistration::new(true, || { + let config = Config::global(); + let metadata = BedrockProvider::metadata(); + Ok(default_inventory_identity( + BEDROCK_PROVIDER_NAME, + BEDROCK_PROVIDER_NAME, + &metadata.config_keys, + config, + )) + }) +} + pub fn huggingface_inventory() -> InventoryRegistration { InventoryRegistration::new(false, || { let metadata = HuggingFaceProvider::metadata(); diff --git a/crates/goose/src/providers/provider_test.rs b/crates/goose/src/providers/provider_test.rs index 5284f2a5d3e5..ffc16aa3655a 100644 --- a/crates/goose/src/providers/provider_test.rs +++ b/crates/goose/src/providers/provider_test.rs @@ -3,6 +3,23 @@ use anyhow::Result; use futures::StreamExt; use rmcp::model::ToolAnnotations; use rmcp::{model::Tool, object}; +use std::time::Duration; +use tokio::time::timeout; + +const PROVIDER_TEST_TIMEOUT: Duration = Duration::from_secs(60); + +pub fn toolshim_settings_from_env() -> (bool, Option) { + let toolshim_enabled = std::env::var("GOOSE_TOOLSHIM") + .map(|val| val == "1" || val.to_lowercase() == "true") + .unwrap_or(false); + let toolshim_model = std::env::var("GOOSE_TOOLSHIM_OLLAMA_MODEL").ok(); + (toolshim_enabled, toolshim_model) +} + +pub async fn test_provider_model(provider_name: &str, model: &str) -> Result<()> { + let (toolshim_enabled, toolshim_model) = toolshim_settings_from_env(); + test_provider_configuration(provider_name, model, toolshim_enabled, toolshim_model).await +} pub async fn test_provider_configuration( provider_name: &str, @@ -36,9 +53,14 @@ pub async fn test_provider_configuration( ) .await?; - let first_chunk = stream - .next() + let first_chunk = timeout(PROVIDER_TEST_TIMEOUT, stream.next()) .await + .map_err(|_| { + anyhow::anyhow!( + "Provider configuration test timed out after {}s", + PROVIDER_TEST_TIMEOUT.as_secs() + ) + })? .ok_or_else(|| anyhow::anyhow!("Provider test stream returned no events"))?; first_chunk?; diff --git a/ui/desktop/src/components/settings/models/subcomponents/SwitchModelModal.tsx b/ui/desktop/src/components/settings/models/subcomponents/SwitchModelModal.tsx index a7e734d62dd2..b424599e824d 100644 --- a/ui/desktop/src/components/settings/models/subcomponents/SwitchModelModal.tsx +++ b/ui/desktop/src/components/settings/models/subcomponents/SwitchModelModal.tsx @@ -127,6 +127,10 @@ const i18n = defineMessages({ id: 'switchModelModal.loadingModels', defaultMessage: 'Loading models…', }, + checkingModel: { + id: 'switchModelModal.checkingModel', + defaultMessage: 'Checking model…', + }, selectModelPlaceholder: { id: 'switchModelModal.selectModelPlaceholder', defaultMessage: 'Select a model, type to search', @@ -298,6 +302,7 @@ export const SwitchModelModal = ({ const [selectedPredefinedModel, setSelectedPredefinedModel] = useState(null); const [predefinedModels, setPredefinedModels] = useState([]); const [loadingModels, setLoadingModels] = useState(false); + const [isSubmitting, setIsSubmitting] = useState(false); const [userClearedModel, setUserClearedModel] = useState(false); const [providerErrors, setProviderErrors] = useState>({}); const [providerWarnings, setProviderWarnings] = useState>({}); @@ -393,41 +398,47 @@ export const SwitchModelModal = ({ setAttemptedSubmit(true); const isFormValid = validateForm(); - if (isFormValid) { - let modelObj: Model; + if (!isFormValid) { + return; + } - if (usePredefinedModels && selectedPredefinedModel) { - modelObj = selectedPredefinedModel; - } else { - const providerMetaData = await getProviderMetadata(provider || '', getProviders); - const providerDisplayName = providerMetaData.display_name; - modelObj = { - name: model, - provider: provider, - subtext: providerDisplayName, - } as Model; - } + let modelObj: Model; + + if (usePredefinedModels && selectedPredefinedModel) { + modelObj = selectedPredefinedModel; + } else { + const providerMetaData = await getProviderMetadata(provider || '', getProviders); + const providerDisplayName = providerMetaData.display_name; + modelObj = { + name: model, + provider: provider, + subtext: providerDisplayName, + } as Model; + } + modelObj = { + ...modelObj, + reasoning: selectedModelReasoning ?? modelObj.reasoning, + }; + + if (showThinkingControl) { + const effort = thinkingEffort ?? modelObj.request_params?.thinking_effort ?? 'off'; modelObj = { ...modelObj, - reasoning: selectedModelReasoning ?? modelObj.reasoning, + request_params: { ...modelObj.request_params, thinking_effort: effort }, }; + upsert('GOOSE_THINKING_EFFORT', effort, false).catch(console.warn); + } - if (showThinkingControl) { - const effort = thinkingEffort ?? modelObj.request_params?.thinking_effort ?? 'off'; - modelObj = { - ...modelObj, - request_params: { ...modelObj.request_params, thinking_effort: effort }, - }; - upsert('GOOSE_THINKING_EFFORT', effort, false).catch(console.warn); - } - + setIsSubmitting(true); + try { const success = await changeModel(sessionId, modelObj); if (success) { onModelSelected?.(modelObj.name, modelObj.provider || ''); trackModelChanged(modelObj.provider || '', modelObj.name); + onClose(); } - - onClose(); + } finally { + setIsSubmitting(false); } }; @@ -963,11 +974,13 @@ export const SwitchModelModal = ({ {intl.formatMessage(i18n.quickStartGuide)}
- -
From 00272832ac96251c5a45687e7fd8aa8577f950a6 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Stadler=20Gell=C3=A9rt?= Date: Thu, 25 Jun 2026 09:15:14 +0200 Subject: [PATCH 2/5] Fix model switch error display --- ui/desktop/src/components/ModelAndProviderContext.tsx | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/ui/desktop/src/components/ModelAndProviderContext.tsx b/ui/desktop/src/components/ModelAndProviderContext.tsx index 8b5f19d6249c..fa23957a7c11 100644 --- a/ui/desktop/src/components/ModelAndProviderContext.tsx +++ b/ui/desktop/src/components/ModelAndProviderContext.tsx @@ -81,7 +81,7 @@ export const ModelAndProviderProvider: React.FC = }, }); if (response.error) { - throw new Error(`Failed to update agent provider: ${response.error}`); + throw new Error(`Failed to update agent provider: ${errorMessage(response.error)}`); } } @@ -113,13 +113,14 @@ export const ModelAndProviderProvider: React.FC = return true; } catch (error) { console.error(`Failed to change model at ${phase} step -- ${modelName} ${providerName}`); + const message = errorMessage(error); toastError({ title: intl.formatMessage(i18n.modelChangeFailed, { provider: providerName, model: modelName, }), - msg: `${error}`, - traceback: errorMessage(error), + msg: message, + traceback: message, }); return false; } From f29f58e7096229805176d33ca4a4ffad728a3ef7 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Stadler=20Gell=C3=A9rt?= Date: Thu, 25 Jun 2026 10:30:09 +0200 Subject: [PATCH 3/5] Fix Bedrock CI and i18n validation failures --- crates/goose-cli/src/commands/configure.rs | 6 ++--- crates/goose/src/providers/bedrock.rs | 26 +++++++++++---------- crates/goose/src/providers/provider_test.rs | 25 ++++++++++++-------- ui/desktop/src/i18n/messages/es.json | 3 +++ ui/desktop/src/i18n/messages/hi.json | 3 +++ ui/desktop/src/i18n/messages/ja.json | 3 +++ ui/desktop/src/i18n/messages/ko.json | 3 +++ ui/desktop/src/i18n/messages/ru.json | 3 +++ ui/desktop/src/i18n/messages/tr.json | 3 +++ ui/desktop/src/i18n/messages/zh-CN.json | 3 +++ 10 files changed, 52 insertions(+), 26 deletions(-) diff --git a/crates/goose-cli/src/commands/configure.rs b/crates/goose-cli/src/commands/configure.rs index d94b1d6e5e23..3f21e036251b 100644 --- a/crates/goose-cli/src/commands/configure.rs +++ b/crates/goose-cli/src/commands/configure.rs @@ -811,10 +811,8 @@ pub async fn configure_provider_dialog() -> anyhow::Result { let spin = spinner(); spin.start("Checking your configuration..."); - let toolshim_enabled = std::env::var("GOOSE_TOOLSHIM") - .map(|val| val == "1" || val.to_lowercase() == "true") - .unwrap_or(false); - let toolshim_model = std::env::var("GOOSE_TOOLSHIM_OLLAMA_MODEL").ok(); + let (toolshim_enabled, toolshim_model) = + goose::providers::provider_test::toolshim_settings_from_env(); match test_provider_configuration(provider_name, &model, toolshim_enabled, toolshim_model).await { diff --git a/crates/goose/src/providers/bedrock.rs b/crates/goose/src/providers/bedrock.rs index 9c9699631e7b..6495747dfac2 100644 --- a/crates/goose/src/providers/bedrock.rs +++ b/crates/goose/src/providers/bedrock.rs @@ -251,18 +251,11 @@ impl BedrockProvider { return false; } - let has_text_input = summary - .input_modalities() - .iter() - .any(|modality| *modality == ModelModality::Text); - let has_text_output = summary - .output_modalities() - .iter() - .any(|modality| *modality == ModelModality::Text); + let has_text_input = summary.input_modalities().contains(&ModelModality::Text); + let has_text_output = summary.output_modalities().contains(&ModelModality::Text); let has_embedding_output = summary .output_modalities() - .iter() - .any(|modality| *modality == ModelModality::Embedding); + .contains(&ModelModality::Embedding); let streaming = summary.response_streaming_supported().unwrap_or(false); let not_legacy = summary .model_lifecycle() @@ -1310,7 +1303,16 @@ mod tests { .region(aws_config::Region::new("us-east-1")) .build(); - let model = ModelConfig::new("openai.gpt-5.5").unwrap(); + let model = ModelConfig { + model_name: "openai.gpt-5.5".to_string(), + context_limit: None, + temperature: None, + max_tokens: None, + toolshim: false, + toolshim_model: None, + request_params: None, + reasoning: None, + }; let provider = BedrockProvider { client: BedrockRuntimeClient::new(&sdk_config), control_plane_client: BedrockControlClient::new(&sdk_config), @@ -1762,7 +1764,7 @@ mod tests { #[test] fn test_skip_canonical_filtering_enabled() { - let provider = create_mock_provider("test"); + let (provider, _) = create_mock_provider_and_model("test"); assert!(provider.skip_canonical_filtering()); } diff --git a/crates/goose/src/providers/provider_test.rs b/crates/goose/src/providers/provider_test.rs index ffc16aa3655a..98fd82c83e3f 100644 --- a/crates/goose/src/providers/provider_test.rs +++ b/crates/goose/src/providers/provider_test.rs @@ -8,36 +8,41 @@ use tokio::time::timeout; const PROVIDER_TEST_TIMEOUT: Duration = Duration::from_secs(60); -pub fn toolshim_settings_from_env() -> (bool, Option) { +pub fn toolshim_settings_from_env() -> (Option, Option) { let toolshim_enabled = std::env::var("GOOSE_TOOLSHIM") .map(|val| val == "1" || val.to_lowercase() == "true") - .unwrap_or(false); + .ok(); let toolshim_model = std::env::var("GOOSE_TOOLSHIM_OLLAMA_MODEL").ok(); (toolshim_enabled, toolshim_model) } pub async fn test_provider_model(provider_name: &str, model: &str) -> Result<()> { - let (toolshim_enabled, toolshim_model) = toolshim_settings_from_env(); - test_provider_configuration(provider_name, model, toolshim_enabled, toolshim_model).await + test_provider_configuration(provider_name, model, None, None).await } pub async fn test_provider_configuration( provider_name: &str, model: &str, - toolshim_enabled: bool, + toolshim_enabled: Option, toolshim_model: Option, ) -> Result<()> { - let model_config = crate::model_config::model_config_from_user_config(provider_name, model)? - .with_max_tokens(Some(50)) - .with_toolshim(toolshim_enabled) - .with_toolshim_model(toolshim_model); + let mut model_config = + crate::model_config::model_config_from_user_config(provider_name, model)? + .with_max_tokens(Some(50)); + + if let Some(toolshim_enabled) = toolshim_enabled { + model_config = model_config.with_toolshim(toolshim_enabled); + } + if toolshim_model.is_some() { + model_config = model_config.with_toolshim_model(toolshim_model); + } let provider = create(provider_name, Vec::new()).await?; let messages = vec![Message::user().with_text("What is the weather like in San Francisco today?")]; - let tools = if !toolshim_enabled { + let tools = if !model_config.toolshim { vec![create_sample_weather_tool()] } else { vec![] diff --git a/ui/desktop/src/i18n/messages/es.json b/ui/desktop/src/i18n/messages/es.json index dda634bf1ab5..b8746ab07e16 100644 --- a/ui/desktop/src/i18n/messages/es.json +++ b/ui/desktop/src/i18n/messages/es.json @@ -4247,6 +4247,9 @@ "switchModelModal.loadingModels": { "defaultMessage": "Cargando modelos…" }, + "switchModelModal.checkingModel": { + "defaultMessage": "Comprobando modelo…" + }, "switchModelModal.localModelsDescription": { "defaultMessage": "Para usar inferencia local, primero debes descargar un modelo a tu computadora. Ve a Ajustes → Modelos para gestionar los modelos locales." }, diff --git a/ui/desktop/src/i18n/messages/hi.json b/ui/desktop/src/i18n/messages/hi.json index d2bdbde46a60..8bcc2007faca 100644 --- a/ui/desktop/src/i18n/messages/hi.json +++ b/ui/desktop/src/i18n/messages/hi.json @@ -4247,6 +4247,9 @@ "switchModelModal.loadingModels": { "defaultMessage": "मॉडल लोड हो रहे हैं…" }, + "switchModelModal.checkingModel": { + "defaultMessage": "मॉडल की जाँच की जा रही है…" + }, "switchModelModal.localModelsDescription": { "defaultMessage": "स्थानीय अनुमान का उपयोग करने के लिए, आपको पहले अपने कंप्यूटर पर एक मॉडल डाउनलोड करना होगा। स्थानीय मॉडल प्रबंधित करने के लिए Settings → मॉडल पर जाएं।" }, diff --git a/ui/desktop/src/i18n/messages/ja.json b/ui/desktop/src/i18n/messages/ja.json index 3feeb5b0e6e4..87e042608320 100644 --- a/ui/desktop/src/i18n/messages/ja.json +++ b/ui/desktop/src/i18n/messages/ja.json @@ -4247,6 +4247,9 @@ "switchModelModal.loadingModels": { "defaultMessage": "モデルを読み込み中…" }, + "switchModelModal.checkingModel": { + "defaultMessage": "モデルを確認中…" + }, "switchModelModal.localModelsDescription": { "defaultMessage": "ローカル推論を使用するには、先にモデルをコンピューターにダウンロードする必要があります。設定 → モデルでローカルモデルを管理できます。" }, diff --git a/ui/desktop/src/i18n/messages/ko.json b/ui/desktop/src/i18n/messages/ko.json index 6de27b4df4ff..2727b7a3d9c7 100644 --- a/ui/desktop/src/i18n/messages/ko.json +++ b/ui/desktop/src/i18n/messages/ko.json @@ -4247,6 +4247,9 @@ "switchModelModal.loadingModels": { "defaultMessage": "모델 로드 중…" }, + "switchModelModal.checkingModel": { + "defaultMessage": "모델 확인 중…" + }, "switchModelModal.localModelsDescription": { "defaultMessage": "로컬 추론을 사용하려면 먼저 모델을 컴퓨터에 다운로드해야 합니다. 로컬 모델을 관리하려면 설정 → 모델로 이동하세요." }, diff --git a/ui/desktop/src/i18n/messages/ru.json b/ui/desktop/src/i18n/messages/ru.json index 3390fbf1573d..0e92d31c75aa 100644 --- a/ui/desktop/src/i18n/messages/ru.json +++ b/ui/desktop/src/i18n/messages/ru.json @@ -4247,6 +4247,9 @@ "switchModelModal.loadingModels": { "defaultMessage": "Загрузка моделей…" }, + "switchModelModal.checkingModel": { + "defaultMessage": "Проверка модели…" + }, "switchModelModal.localModelsDescription": { "defaultMessage": "Чтобы использовать локальный инференс, сначала скачайте модель на компьютер. Перейдите в Настройки → Модели для управления локальными моделями." }, diff --git a/ui/desktop/src/i18n/messages/tr.json b/ui/desktop/src/i18n/messages/tr.json index 46481e610f8f..839459b3376f 100644 --- a/ui/desktop/src/i18n/messages/tr.json +++ b/ui/desktop/src/i18n/messages/tr.json @@ -4247,6 +4247,9 @@ "switchModelModal.loadingModels": { "defaultMessage": "Modeller yükleniyor…" }, + "switchModelModal.checkingModel": { + "defaultMessage": "Model kontrol ediliyor…" + }, "switchModelModal.localModelsDescription": { "defaultMessage": "Yerel çıkarımı kullanmak için öncelikle bilgisayarınıza bir model indirmeniz gerekir. Yerel modelleri yönetmek için Ayarlar → Modeller'e gidin." }, diff --git a/ui/desktop/src/i18n/messages/zh-CN.json b/ui/desktop/src/i18n/messages/zh-CN.json index 7c88961e813d..4c20c9c1b859 100644 --- a/ui/desktop/src/i18n/messages/zh-CN.json +++ b/ui/desktop/src/i18n/messages/zh-CN.json @@ -4247,6 +4247,9 @@ "switchModelModal.loadingModels": { "defaultMessage": "正在加载模型…" }, + "switchModelModal.checkingModel": { + "defaultMessage": "正在检查模型…" + }, "switchModelModal.localModelsDescription": { "defaultMessage": "要使用本地推理,你需要先下载一个模型到电脑上。前往 设置 → 模型 管理本地模型。" }, From 0475cc0d301c132e182ec676f05b270995ec16e8 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Stadler=20Gell=C3=A9rt?= Date: Thu, 25 Jun 2026 10:55:57 +0200 Subject: [PATCH 4/5] Update English desktop i18n messages --- ui/desktop/src/i18n/messages/en.json | 3 +++ 1 file changed, 3 insertions(+) diff --git a/ui/desktop/src/i18n/messages/en.json b/ui/desktop/src/i18n/messages/en.json index 8933bae98df6..e69c9e42ad5c 100644 --- a/ui/desktop/src/i18n/messages/en.json +++ b/ui/desktop/src/i18n/messages/en.json @@ -4199,6 +4199,9 @@ "switchModelModal.checkProviderConfig": { "defaultMessage": "Check your provider configuration in Settings → Providers" }, + "switchModelModal.checkingModel": { + "defaultMessage": "Checking model…" + }, "switchModelModal.chooseModel": { "defaultMessage": "Choose a model:" }, From db7072c93f3bfc2b18f3b76d0a786d03b7e29b68 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Stadler=20Gell=C3=A9rt?= Date: Thu, 25 Jun 2026 11:37:17 +0200 Subject: [PATCH 5/5] Implemented Codex Bot Suggestions: 1)Keep GPT-OSS Bedrock models on the Converse path 2)Apply the probe timeout to stream creation too --- crates/goose/src/providers/bedrock.rs | 20 ++++++------- crates/goose/src/providers/provider_test.rs | 32 ++++++++++++--------- 2 files changed, 29 insertions(+), 23 deletions(-) diff --git a/crates/goose/src/providers/bedrock.rs b/crates/goose/src/providers/bedrock.rs index 6495747dfac2..b8dfae197ab6 100644 --- a/crates/goose/src/providers/bedrock.rs +++ b/crates/goose/src/providers/bedrock.rs @@ -265,11 +265,8 @@ impl BedrockProvider { has_text_input && has_text_output && !has_embedding_output && streaming && not_legacy } - fn is_mantle_model_name(model_name: &str) -> bool { - model_name - .strip_prefix("openai.") - .unwrap_or(model_name) - .starts_with("gpt-") + fn is_mantle_model_id(model_id: &str) -> bool { + BEDROCK_MANTLE_MODELS.contains(&model_id) } async fn fetch_inference_profile_ids(&self) -> Result, ProviderError> { @@ -948,7 +945,7 @@ impl Provider for BedrockProvider { let (base_name, effort) = extract_reasoning_effort(without_prefix); let bedrock_model_id = format!("openai.{}", base_name); - let is_mantle_model = Self::is_mantle_model_name(&model_config.model_name); + let is_mantle_model = Self::is_mantle_model_id(&bedrock_model_id); if is_mantle_model { let mut normalized_config = ModelConfig { @@ -1754,10 +1751,13 @@ mod tests { } #[test] - fn test_is_mantle_model_name() { - assert!(BedrockProvider::is_mantle_model_name("openai.gpt-5.5")); - assert!(BedrockProvider::is_mantle_model_name("gpt-5.4")); - assert!(!BedrockProvider::is_mantle_model_name( + fn test_is_mantle_model_id() { + assert!(BedrockProvider::is_mantle_model_id("openai.gpt-5.5")); + assert!(BedrockProvider::is_mantle_model_id("openai.gpt-5.4")); + assert!(!BedrockProvider::is_mantle_model_id( + "openai.gpt-oss-120b-1:0" + )); + assert!(!BedrockProvider::is_mantle_model_id( "us.anthropic.claude-sonnet-4-5-20250929-v1:0" )); } diff --git a/crates/goose/src/providers/provider_test.rs b/crates/goose/src/providers/provider_test.rs index 98fd82c83e3f..f534a5e2f827 100644 --- a/crates/goose/src/providers/provider_test.rs +++ b/crates/goose/src/providers/provider_test.rs @@ -48,26 +48,32 @@ pub async fn test_provider_configuration( vec![] }; - let mut stream = provider - .stream( - &model_config, - "test-session-id", - "You are an AI agent called goose. You use tools of connected extensions to solve problems.", - &messages, - &tools.into_iter().collect::>(), - ) - .await?; + timeout(PROVIDER_TEST_TIMEOUT, async { + let mut stream = provider + .stream( + &model_config, + "test-session-id", + "You are an AI agent called goose. You use tools of connected extensions to solve problems.", + &messages, + &tools.into_iter().collect::>(), + ) + .await?; + + let first_chunk = stream + .next() + .await + .ok_or_else(|| anyhow::anyhow!("Provider test stream returned no events"))?; + first_chunk?; - let first_chunk = timeout(PROVIDER_TEST_TIMEOUT, stream.next()) + Ok::<(), anyhow::Error>(()) + }) .await .map_err(|_| { anyhow::anyhow!( "Provider configuration test timed out after {}s", PROVIDER_TEST_TIMEOUT.as_secs() ) - })? - .ok_or_else(|| anyhow::anyhow!("Provider test stream returned no events"))?; - first_chunk?; + })??; Ok(()) }