feat(acp): pass session cwd param to acp providers (#9229)

Signed-off-by: Matt Toohey <contact@matttoohey.com>
Signed-off-by: Kalvin Chau <kalvin@block.xyz>
This commit is contained in:
Kalvin C
2026-05-17 22:47:23 -07:00
committed by GitHub
parent 06e6e2e850
commit bf54314b67
17 changed files with 320 additions and 51 deletions
+2 -2
View File
@@ -93,7 +93,7 @@ impl Provider for NamingProvider {
}
fn naming_provider_factory() -> AcpProviderFactory {
Arc::new(|_provider_name, model_config, _extensions| {
Arc::new(|_provider_name, model_config, _extensions, _working_dir| {
Box::pin(async move { Ok(Arc::new(NamingProvider { model_config }) as Arc<dyn Provider>) })
})
}
@@ -448,7 +448,7 @@ pub async fn run_fs_write_text_file_true<C: Connection>() {
pub async fn run_initialize_doesnt_hit_provider<C: Connection>() {
let provider_factory: AcpProviderFactory =
Arc::new(|_, _, _| Box::pin(async { Err(anyhow::anyhow!("no provider configured")) }));
Arc::new(|_, _, _, _| Box::pin(async { Err(anyhow::anyhow!("no provider configured")) }));
let openai = OpenAiFixture::new(vec![], C::expected_session_id()).await;
let config = TestConnectionConfig {
+102 -1
View File
@@ -12,6 +12,7 @@ use goose::model::ModelConfig;
use goose::providers::base::{MessageStream, Provider};
use goose::providers::errors::ProviderError;
use goose_test_support::{EnforceSessionId, IgnoreSessionId};
use std::path::PathBuf;
use std::sync::{Arc, Mutex};
use common_tests::fixtures::OpenAiFixture;
@@ -49,7 +50,7 @@ impl Provider for MockProvider {
}
fn mock_provider_factory() -> AcpProviderFactory {
Arc::new(|provider_name, model_config, _extensions| {
Arc::new(|provider_name, model_config, _extensions, _working_dir| {
Box::pin(async move {
let recommended_models = match provider_name.as_str() {
"anthropic" => vec![
@@ -112,6 +113,106 @@ fn test_custom_get_extensions() {
});
}
#[test]
fn test_new_session_passes_cwd_to_provider_factory() {
run_test(async move {
let openai = OpenAiFixture::new(vec![], Arc::new(EnforceSessionId::default())).await;
let cwd = tempfile::tempdir().unwrap();
let expected_cwd = cwd.path().to_path_buf();
let captured_cwds = Arc::new(Mutex::new(Vec::<Option<PathBuf>>::new()));
let factory_cwds = Arc::clone(&captured_cwds);
let provider_factory: AcpProviderFactory = Arc::new(
move |provider_name, model_config, _extensions, working_dir| {
factory_cwds.lock().unwrap().push(working_dir);
Box::pin(async move {
Ok(Arc::new(MockProvider {
name: provider_name,
model_config,
recommended_models: Vec::new(),
}) as Arc<dyn Provider>)
})
},
);
let mut conn = AcpServerConnection::new(
TestConnectionConfig {
cwd: Some(cwd),
provider_factory: Some(provider_factory),
..Default::default()
},
openai,
)
.await;
conn.new_session().await.unwrap();
let captured_cwd = tokio::time::timeout(std::time::Duration::from_secs(1), async {
loop {
if let Some(cwd) = captured_cwds.lock().unwrap().first().cloned() {
break cwd;
}
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
}
})
.await
.expect("provider factory was not called");
assert_eq!(captured_cwd, Some(expected_cwd));
});
}
#[test]
fn test_load_session_passes_load_cwd_to_provider_factory() {
run_test(async move {
let openai = OpenAiFixture::new(vec![], Arc::new(EnforceSessionId::default())).await;
let initial_cwd = tempfile::tempdir().unwrap();
let captured_cwds = Arc::new(Mutex::new(Vec::<Option<PathBuf>>::new()));
let factory_cwds = Arc::clone(&captured_cwds);
let provider_factory: AcpProviderFactory = Arc::new(
move |provider_name, model_config, _extensions, working_dir| {
factory_cwds.lock().unwrap().push(working_dir);
Box::pin(async move {
Ok(Arc::new(MockProvider {
name: provider_name,
model_config,
recommended_models: Vec::new(),
}) as Arc<dyn Provider>)
})
},
);
let mut conn = AcpServerConnection::new(
TestConnectionConfig {
cwd: Some(initial_cwd),
provider_factory: Some(provider_factory),
..Default::default()
},
openai,
)
.await;
let SessionData { session, .. } = conn.new_session().await.unwrap();
let session_id = session.session_id().0.to_string();
let SessionData {
session: loaded, ..
} = conn.load_session(&session_id, vec![]).await.unwrap();
let expected_cwd = loaded.work_dir();
let captured_cwd = tokio::time::timeout(std::time::Duration::from_secs(1), async {
loop {
if let Some(cwd) = captured_cwds.lock().unwrap().get(1).cloned() {
break cwd;
}
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
}
})
.await
.expect("provider factory was not called for load session");
assert_eq!(captured_cwd, Some(expected_cwd));
});
}
#[test]
fn test_custom_list_builtin_skill_sources() {
run_test(async move {
+15 -11
View File
@@ -173,17 +173,21 @@ pub async fn spawn_acp_server_in_process(
}
let provider_factory = provider_factory.unwrap_or_else(|| {
let base_url = openai_base_url.to_string();
Arc::new(move |_provider_name, model_config, _extensions| {
let base_url = base_url.clone();
Box::pin(async move {
let api_client =
ApiClient::new(base_url, ApiAuthMethod::BearerToken("test-key".to_string()))
.unwrap();
let provider: Arc<dyn Provider> =
Arc::new(OpenAiProvider::new(api_client, model_config));
Ok(provider)
})
})
Arc::new(
move |_provider_name, model_config, _extensions, _working_dir| {
let base_url = base_url.clone();
Box::pin(async move {
let api_client = ApiClient::new(
base_url,
ApiAuthMethod::BearerToken("test-key".to_string()),
)
.unwrap();
let provider: Arc<dyn Provider> =
Arc::new(OpenAiProvider::new(api_client, model_config));
Ok(provider)
})
},
)
});
let agent = GooseAcpAgent::new(GooseAcpAgentOptions {
@@ -47,7 +47,7 @@ impl Provider for MockProvider {
}
fn mock_provider_factory() -> goose::acp::server::AcpProviderFactory {
Arc::new(|provider_name, model_config, _extensions| {
Arc::new(|provider_name, model_config, _extensions, _working_dir| {
Box::pin(async move {
Ok(Arc::new(MockProvider {
name: provider_name,