fix(summon): reuse parent provider for delegates (#10754)

This commit is contained in:
Anthony
2026-07-29 06:41:01 -07:00
committed by GitHub
parent 55ad520633
commit 9fec4152a4
@@ -1554,8 +1554,6 @@ impl SummonClient {
recipe: &Recipe,
session: &crate::session::Session,
) -> Result<TaskConfig, anyhow::Error> {
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<dyn crate::providers::base::Provider>,
@@ -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<dyn crate::providers::base::Provider> = 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(&params, &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(&params, &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";