diff --git a/crates/goose/src/agents/platform_extensions/summon.rs b/crates/goose/src/agents/platform_extensions/summon.rs index 4308c036a..63b0b1c86 100644 --- a/crates/goose/src/agents/platform_extensions/summon.rs +++ b/crates/goose/src/agents/platform_extensions/summon.rs @@ -1554,8 +1554,6 @@ impl SummonClient { recipe: &Recipe, session: &crate::session::Session, ) -> Result { - let (provider, model_config) = self.resolve_provider(params, recipe, session).await?; - let mut extensions = EnabledExtensionsState::extensions_or_default( Some(&session.extension_data), Config::global(), @@ -1582,6 +1580,10 @@ impl SummonClient { } } + let (provider, model_config) = self + .resolve_provider(params, recipe, session, &extensions) + .await?; + let max_turns = params .max_turns .or_else(|| recipe.settings.as_ref().and_then(|s| s.max_turns)) @@ -1676,6 +1678,7 @@ impl SummonClient { params: &DelegateParams, recipe: &Recipe, session: &crate::session::Session, + extensions: &[crate::config::ExtensionConfig], ) -> Result< ( Arc, @@ -1701,7 +1704,31 @@ impl SummonClient { .ok_or_else(|| anyhow::anyhow!("No provider configured"))?; let model_config = self.resolve_model_config(params, recipe, session, &provider_name)?; - let provider = providers::create(&provider_name, Vec::new()).await?; + let provider = match providers::get_from_registry(&provider_name).await { + Ok(entry) => entry.create(extensions.to_vec()).await?, + Err(error) => { + let parent_provider = if let Some(extension_manager) = self + .context + .extension_manager + .as_ref() + .and_then(|weak| weak.upgrade()) + { + extension_manager.get_provider().lock().await.clone() + } else { + None + }; + + match parent_provider { + Some(provider) + if provider.get_name() == provider_name + && !provider.manages_own_context() => + { + provider + } + _ => return Err(error), + } + } + }; Ok((provider, model_config)) } @@ -2611,6 +2638,78 @@ You review code."#; } } + #[tokio::test] + async fn test_resolve_provider_reuses_unregistered_parent_provider() { + let temp_dir = TempDir::new().unwrap(); + let parent_provider: Arc = Arc::new( + crate::providers::testprovider::TestProvider::new_replaying( + temp_dir.path().join("records.json").display().to_string(), + ) + .unwrap(), + ); + let extension_manager = Arc::new( + crate::agents::extension_manager::ExtensionManager::new_without_provider( + temp_dir.path().to_path_buf(), + ), + ); + *extension_manager.get_provider().lock().await = Some(Arc::clone(&parent_provider)); + let mut context = extension_manager.get_context().clone(); + context.extension_manager = Some(Arc::downgrade(&extension_manager)); + let client = SummonClient::new(context).unwrap(); + let session = crate::session::Session { + provider_name: Some(parent_provider.get_name().to_string()), + model_config: Some(goose_providers::model::ModelConfig::new("test-model")), + ..Default::default() + }; + + let params = DelegateParams { + provider: Some(parent_provider.get_name().to_string()), + model: Some("test-model".to_string()), + ..Default::default() + }; + let (resolved_provider, _) = client + .resolve_provider(¶ms, &empty_recipe(), &session, &[]) + .await + .unwrap(); + + assert!(Arc::ptr_eq(&parent_provider, &resolved_provider)); + } + + #[tokio::test] + async fn test_build_task_config_recreates_registered_parent_provider() { + let temp_dir = TempDir::new().unwrap(); + let parent_provider = providers::create("openai", Vec::new()).await.unwrap(); + let extension_manager = Arc::new( + crate::agents::extension_manager::ExtensionManager::new_without_provider( + temp_dir.path().to_path_buf(), + ), + ); + *extension_manager.get_provider().lock().await = Some(Arc::clone(&parent_provider)); + let mut context = extension_manager.get_context().clone(); + context.extension_manager = Some(Arc::downgrade(&extension_manager)); + let client = SummonClient::new(context).unwrap(); + let session = crate::session::Session { + provider_name: Some(parent_provider.get_name().to_string()), + model_config: Some(goose_providers::model::ModelConfig::new("test-model")), + working_dir: temp_dir.path().to_path_buf(), + ..Default::default() + }; + let params = DelegateParams { + extensions: Some(Vec::new()), + provider: Some(parent_provider.get_name().to_string()), + model: Some("test-model".to_string()), + ..Default::default() + }; + + let task_config = client + .build_task_config(¶ms, &empty_recipe(), &session) + .await + .unwrap(); + + assert!(!Arc::ptr_eq(&parent_provider, &task_config.provider)); + assert!(task_config.extensions.is_empty()); + } + const PARENT_MODEL: &str = "claude-3-5-sonnet-20241022"; const OVERRIDE_MODEL: &str = "claude-opus-4-6"; const PROVIDER: &str = "anthropic";