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:
@@ -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 {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user