fix(summon): reuse parent provider for delegates (#10754)
This commit is contained in:
@@ -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(¶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";
|
||||
|
||||
Reference in New Issue
Block a user