diff --git a/crates/goose/src/acp/server.rs b/crates/goose/src/acp/server.rs index 5597d52f8e..2a703b590c 100644 --- a/crates/goose/src/acp/server.rs +++ b/crates/goose/src/acp/server.rs @@ -61,6 +61,7 @@ use rmcp::model::{ use serde::Deserialize; use std::collections::{HashMap, HashSet}; use std::panic::AssertUnwindSafe; +use std::path::{Path, PathBuf}; use std::sync::Arc; use strum::{EnumMessage, VariantNames}; use tokio::sync::{Mutex, OnceCell}; @@ -86,6 +87,7 @@ pub type AcpProviderFactory = Arc< String, crate::model::ModelConfig, Vec, + Option, ) -> BoxFuture<'static, Result>> + Send + Sync, @@ -1044,6 +1046,20 @@ fn build_usage_update(session: &Session, context_limit: usize) -> UsageUpdate { UsageUpdate::new(used, context_limit as u64) } +fn validate_absolute_cwd(cwd: &Path) -> Result<(), agent_client_protocol::Error> { + if !cwd.is_absolute() { + return Err( + agent_client_protocol::Error::invalid_params().data("cwd must be an absolute path") + ); + } + + if !cwd.exists() || !cwd.is_dir() { + return Err(agent_client_protocol::Error::invalid_params().data("invalid directory path")); + } + + Ok(()) +} + impl GooseAcpAgent { pub fn permission_manager(&self) -> Arc { Arc::clone(&self.permission_manager) @@ -1093,8 +1109,15 @@ impl GooseAcpAgent { provider_name: &str, model_config: crate::model::ModelConfig, extensions: Vec, + working_dir: Option, ) -> Result> { - (self.provider_factory)(provider_name.to_string(), model_config, extensions).await + (self.provider_factory)( + provider_name.to_string(), + model_config, + extensions, + working_dir, + ) + .await } async fn prepare_session_init_config( @@ -1131,7 +1154,12 @@ impl GooseAcpAgent { ); Config::global().invalidate_secrets_cache(); match self - .create_provider(provider_name, model_config.clone(), ext_state) + .create_provider( + provider_name, + model_config.clone(), + ext_state, + Some(goose_session.working_dir.clone()), + ) .await { Ok(provider) => { @@ -1348,9 +1376,14 @@ impl GooseAcpAgent { ); let provider = match prebuilt_provider { Some(provider) => provider, - None => provider_factory(provider_name.to_string(), model_config, ext_state) - .await - .map_err(|e| e.to_string())?, + None => provider_factory( + provider_name.to_string(), + model_config, + ext_state, + Some(goose_session.working_dir.clone()), + ) + .await + .map_err(|e| e.to_string())?, }; agent .update_provider(provider.clone(), &goose_session.id) @@ -1440,15 +1473,17 @@ impl GooseAcpAgent { } let ext_manager = &agent.extension_manager; + let working_dir = goose_session.working_dir.clone(); let extension_futures = extensions .into_iter() .map(|ext| { let ext_manager = Arc::clone(ext_manager); let sid_inner = sid_str.clone(); + let working_dir = working_dir.clone(); async move { let name = ext.name().to_string(); if let Err(e) = ext_manager - .add_extension(ext, None, None, sid_inner.as_deref()) + .add_extension(ext, Some(working_dir), None, sid_inner.as_deref()) .await { warn!(extension = %name, error = %e, "extension load failed"); @@ -2412,6 +2447,7 @@ impl GooseAcpAgent { ) -> Result { debug!(?args, "new session request"); let t_start = std::time::Instant::now(); + validate_absolute_cwd(&args.cwd)?; let requested_provider = args .meta @@ -2664,6 +2700,7 @@ impl GooseAcpAgent { args: LoadSessionRequest, ) -> Result { debug!(?args, "load session request"); + validate_absolute_cwd(&args.cwd)?; let session_id = args.session_id.0.to_string(); let sid = sid_short(&session_id); @@ -2835,6 +2872,11 @@ impl GooseAcpAgent { .apply() .await .internal_err_ctx("Failed to update session working directory")?; + let goose_session = self + .session_manager + .get_session(&session_id, false) + .await + .internal_err_ctx("Failed to reload session")?; // Register the session with a Loading handle. let (agent_tx, agent_rx) = tokio::sync::watch::channel::(None); @@ -3137,8 +3179,18 @@ impl GooseAcpAgent { let model_config = crate::model::ModelConfig::new(model_id) .invalid_params_err_ctx("Invalid model config")? .with_canonical_limits(&provider_name); + let session = self + .session_manager + .get_session(session_id, false) + .await + .internal_err_ctx("Failed to get session")?; let provider = self - .create_provider(&provider_name, model_config, extensions) + .create_provider( + &provider_name, + model_config, + extensions, + Some(session.working_dir), + ) .await .internal_err_ctx("Failed to create provider")?; agent @@ -3264,8 +3316,18 @@ impl GooseAcpAgent { let extensions = EnabledExtensionsState::for_session(&self.session_manager, session_id, &config).await; + let session = self + .session_manager + .get_session(session_id, false) + .await + .internal_err_ctx("Failed to get session")?; let new_provider = self - .create_provider(&resolved_provider_name, model_config, extensions) + .create_provider( + &resolved_provider_name, + model_config, + extensions, + Some(session.working_dir), + ) .await .internal_err_ctx("Failed to create provider")?; agent @@ -3332,6 +3394,7 @@ impl GooseAcpAgent { cx: &ConnectionTo, args: ForkSessionRequest, ) -> Result { + validate_absolute_cwd(&args.cwd)?; let source_session_id = &*args.session_id.0; let new_session = self diff --git a/crates/goose/src/acp/server/providers.rs b/crates/goose/src/acp/server/providers.rs index c4729fcbcc..964b16d464 100644 --- a/crates/goose/src/acp/server/providers.rs +++ b/crates/goose/src/acp/server/providers.rs @@ -703,7 +703,7 @@ impl GooseAcpAgent { let model_config = crate::model::ModelConfig::new(&metadata.metadata().default_model)? .with_canonical_limits(&provider_id); - provider_factory(provider_id.clone(), model_config, Vec::new()).await + provider_factory(provider_id.clone(), model_config, Vec::new(), None).await }) .catch_unwind() .await; diff --git a/crates/goose/src/acp/server/sessions.rs b/crates/goose/src/acp/server/sessions.rs index 517c78f2b1..02cb742ede 100644 --- a/crates/goose/src/acp/server/sessions.rs +++ b/crates/goose/src/acp/server/sessions.rs @@ -11,11 +11,7 @@ impl GooseAcpAgent { .data("working directory cannot be empty")); } let path = std::path::PathBuf::from(&working_dir); - if !path.exists() || !path.is_dir() { - return Err( - agent_client_protocol::Error::invalid_params().data("invalid directory path") - ); - } + validate_absolute_cwd(&path)?; let session_id = &req.session_id; self.session_manager .update(session_id) diff --git a/crates/goose/src/acp/server_factory.rs b/crates/goose/src/acp/server_factory.rs index 05c745c39c..079287d967 100644 --- a/crates/goose/src/acp/server_factory.rs +++ b/crates/goose/src/acp/server_factory.rs @@ -34,12 +34,26 @@ impl AcpServer { .unwrap_or(crate::config::GooseMode::Auto); let disable_session_naming = config.get_goose_disable_session_naming().unwrap_or(false); - let provider_factory: AcpProviderFactory = - Arc::new(move |provider_name, model_config, extensions| { + let provider_factory: AcpProviderFactory = Arc::new( + move |provider_name, model_config, extensions, working_dir| { Box::pin(async move { - crate::providers::create(&provider_name, model_config, extensions).await + match working_dir { + Some(working_dir) => { + crate::providers::create_with_working_dir( + &provider_name, + model_config, + extensions, + working_dir, + ) + .await + } + None => { + crate::providers::create(&provider_name, model_config, extensions).await + } + } }) - }); + }, + ); let agent = GooseAcpAgent::new(GooseAcpAgentOptions { provider_factory, diff --git a/crates/goose/src/providers/amp_acp.rs b/crates/goose/src/providers/amp_acp.rs index 5d6339b13f..8640c446f8 100644 --- a/crates/goose/src/providers/amp_acp.rs +++ b/crates/goose/src/providers/amp_acp.rs @@ -10,7 +10,7 @@ use crate::config::search_path::SearchPaths; use crate::config::{Config, GooseMode}; use crate::model::ModelConfig; use crate::providers::acp_tooling::{acp_adapter_installed, acp_inventory_identity}; -use crate::providers::base::{ProviderDef, ProviderMetadata}; +use crate::providers::base::{current_working_dir, ProviderDef, ProviderMetadata}; use crate::providers::inventory::InventoryIdentityInput; const AMP_ACP_PROVIDER_NAME: &str = "amp-acp"; @@ -45,6 +45,14 @@ impl ProviderDef for AmpAcpProvider { fn from_env( model: ModelConfig, extensions: Vec, + ) -> BoxFuture<'static, Result> { + Self::from_env_with_working_dir(model, extensions, current_working_dir()) + } + + fn from_env_with_working_dir( + model: ModelConfig, + extensions: Vec, + working_dir: PathBuf, ) -> BoxFuture<'static, Result> { Box::pin(async move { let config = Config::global(); @@ -65,7 +73,7 @@ impl ProviderDef for AmpAcpProvider { args: vec![], env: vec![], env_remove: vec![], - work_dir: std::env::current_dir().unwrap_or_else(|_| PathBuf::from(".")), + work_dir: working_dir, mcp_servers: extension_configs_to_mcp_servers(&extensions), session_mode_id: Some(mode_mapping[&goose_mode].clone()), mode_mapping, diff --git a/crates/goose/src/providers/base.rs b/crates/goose/src/providers/base.rs index f7a61090f1..cfa0418721 100644 --- a/crates/goose/src/providers/base.rs +++ b/crates/goose/src/providers/base.rs @@ -26,6 +26,7 @@ use utoipa::ToSchema; use once_cell::sync::Lazy; use regex::Regex; use std::ops::{Add, AddAssign}; +use std::path::PathBuf; use std::pin::Pin; use std::sync::LazyLock; use std::sync::Mutex; @@ -766,6 +767,10 @@ impl Usage { } } +pub(crate) fn current_working_dir() -> PathBuf { + std::env::current_dir().unwrap_or_else(|_| PathBuf::from(".")) +} + pub trait ProviderDef: Send + Sync { type Provider: Provider + 'static; @@ -780,6 +785,19 @@ pub trait ProviderDef: Send + Sync { where Self: Sized; + fn from_env_with_working_dir( + model: ModelConfig, + extensions: Vec, + _working_dir: PathBuf, + ) -> BoxFuture<'static, Result> + where + Self: Sized, + { + // ACP subprocess providers must override this so session cwd is preserved. + // Non-subprocess providers can rely on the default because cwd is irrelevant. + Self::from_env(model, extensions) + } + fn supports_inventory_refresh() -> bool where Self: Sized, diff --git a/crates/goose/src/providers/claude_acp.rs b/crates/goose/src/providers/claude_acp.rs index 282b7e0dd9..415889af63 100644 --- a/crates/goose/src/providers/claude_acp.rs +++ b/crates/goose/src/providers/claude_acp.rs @@ -10,7 +10,7 @@ use crate::config::search_path::SearchPaths; use crate::config::{Config, GooseMode}; use crate::model::ModelConfig; use crate::providers::acp_tooling::{acp_adapter_installed, acp_inventory_identity}; -use crate::providers::base::{ProviderDef, ProviderMetadata}; +use crate::providers::base::{current_working_dir, ProviderDef, ProviderMetadata}; use crate::providers::inventory::InventoryIdentityInput; const CLAUDE_ACP_PROVIDER_NAME: &str = "claude-acp"; @@ -43,6 +43,14 @@ impl ProviderDef for ClaudeAcpProvider { fn from_env( model: ModelConfig, extensions: Vec, + ) -> BoxFuture<'static, Result> { + Self::from_env_with_working_dir(model, extensions, current_working_dir()) + } + + fn from_env_with_working_dir( + model: ModelConfig, + extensions: Vec, + working_dir: PathBuf, ) -> BoxFuture<'static, Result> { Box::pin(async move { let config = Config::global(); @@ -69,7 +77,7 @@ impl ProviderDef for ClaudeAcpProvider { env: vec![], // Prevent nested-session detection in claude-agent-acp (wraps Claude Code) env_remove: vec!["CLAUDECODE".to_string()], - work_dir: std::env::current_dir().unwrap_or_else(|_| PathBuf::from(".")), + work_dir: working_dir, mcp_servers: extension_configs_to_mcp_servers(&extensions), session_mode_id: Some(mode_mapping[&goose_mode].clone()), mode_mapping, diff --git a/crates/goose/src/providers/codex_acp.rs b/crates/goose/src/providers/codex_acp.rs index 4ef6ffef97..5c3f46b363 100644 --- a/crates/goose/src/providers/codex_acp.rs +++ b/crates/goose/src/providers/codex_acp.rs @@ -10,7 +10,7 @@ use crate::config::search_path::SearchPaths; use crate::config::{Config, GooseMode}; use crate::model::ModelConfig; use crate::providers::acp_tooling::{acp_adapter_installed, acp_inventory_identity}; -use crate::providers::base::{ProviderDef, ProviderMetadata}; +use crate::providers::base::{current_working_dir, ProviderDef, ProviderMetadata}; use crate::providers::inventory::InventoryIdentityInput; const CODEX_ACP_PROVIDER_NAME: &str = "codex-acp"; @@ -42,6 +42,14 @@ impl ProviderDef for CodexAcpProvider { fn from_env( model: ModelConfig, extensions: Vec, + ) -> BoxFuture<'static, Result> { + Self::from_env_with_working_dir(model, extensions, current_working_dir()) + } + + fn from_env_with_working_dir( + model: ModelConfig, + extensions: Vec, + working_dir: PathBuf, ) -> BoxFuture<'static, Result> { Box::pin(async move { let config = Config::global(); @@ -49,7 +57,6 @@ impl ProviderDef for CodexAcpProvider { let resolved_command = SearchPaths::builder() .with_npm() .resolve(CODEX_ACP_PROVIDER_NAME)?; - let work_dir = std::env::current_dir().unwrap_or_else(|_| PathBuf::from(".")); let env = vec![]; let goose_mode = config.get_goose_mode().unwrap_or(GooseMode::Auto); let mcp_servers = extension_configs_to_mcp_servers(&extensions); @@ -88,7 +95,7 @@ impl ProviderDef for CodexAcpProvider { args, env, env_remove: vec![], - work_dir, + work_dir: working_dir, mcp_servers, // Disabled until https://github.com/zed-industries/codex-acp/issues/179 is fixed. session_mode_id: None, diff --git a/crates/goose/src/providers/copilot_acp.rs b/crates/goose/src/providers/copilot_acp.rs index 35bed3dec0..7f02c7b81d 100644 --- a/crates/goose/src/providers/copilot_acp.rs +++ b/crates/goose/src/providers/copilot_acp.rs @@ -10,7 +10,7 @@ use crate::config::search_path::SearchPaths; use crate::config::{Config, GooseMode}; use crate::model::ModelConfig; use crate::providers::acp_tooling::{acp_adapter_installed, acp_inventory_identity}; -use crate::providers::base::{ProviderDef, ProviderMetadata}; +use crate::providers::base::{current_working_dir, ProviderDef, ProviderMetadata}; use crate::providers::inventory::InventoryIdentityInput; const COPILOT_ACP_PROVIDER_NAME: &str = "copilot-acp"; @@ -46,6 +46,14 @@ impl ProviderDef for CopilotAcpProvider { fn from_env( model: ModelConfig, extensions: Vec, + ) -> BoxFuture<'static, Result> { + Self::from_env_with_working_dir(model, extensions, current_working_dir()) + } + + fn from_env_with_working_dir( + model: ModelConfig, + extensions: Vec, + working_dir: PathBuf, ) -> BoxFuture<'static, Result> { Box::pin(async move { let config = Config::global(); @@ -75,7 +83,7 @@ impl ProviderDef for CopilotAcpProvider { args, env: vec![], env_remove: vec![], - work_dir: std::env::current_dir().unwrap_or_else(|_| PathBuf::from(".")), + work_dir: working_dir, mcp_servers: extension_configs_to_mcp_servers(&extensions), session_mode_id: Some(mode_mapping[&goose_mode].clone()), mode_mapping, diff --git a/crates/goose/src/providers/init.rs b/crates/goose/src/providers/init.rs index b8ea286080..2061c5ebed 100644 --- a/crates/goose/src/providers/init.rs +++ b/crates/goose/src/providers/init.rs @@ -1,3 +1,4 @@ +use std::path::PathBuf; use std::sync::{Arc, RwLock}; #[cfg(feature = "aws-providers")] @@ -160,6 +161,18 @@ pub async fn create( entry.create(model, extensions).await } +pub async fn create_with_working_dir( + name: &str, + model: ModelConfig, + extensions: Vec, + working_dir: PathBuf, +) -> Result> { + let entry = get_from_registry(name).await?; + entry + .create_with_working_dir(model, extensions, working_dir) + .await +} + pub async fn create_with_default_model( name: impl AsRef, extensions: Vec, diff --git a/crates/goose/src/providers/mod.rs b/crates/goose/src/providers/mod.rs index 23e993a8d9..44f7dd0a96 100644 --- a/crates/goose/src/providers/mod.rs +++ b/crates/goose/src/providers/mod.rs @@ -59,6 +59,7 @@ pub mod xai; pub use init::{ cleanup_provider, create, create_with_default_model, create_with_named_model, - get_from_registry, inventory_identity, providers, refresh_custom_providers, + create_with_working_dir, get_from_registry, inventory_identity, providers, + refresh_custom_providers, }; pub use retry::{retry_operation, RetryConfig}; diff --git a/crates/goose/src/providers/pi_acp.rs b/crates/goose/src/providers/pi_acp.rs index b02278f97d..52215d56fc 100644 --- a/crates/goose/src/providers/pi_acp.rs +++ b/crates/goose/src/providers/pi_acp.rs @@ -10,7 +10,7 @@ use crate::config::search_path::SearchPaths; use crate::config::{Config, GooseMode}; use crate::model::ModelConfig; use crate::providers::acp_tooling::{acp_adapter_installed, acp_inventory_identity}; -use crate::providers::base::{ProviderDef, ProviderMetadata}; +use crate::providers::base::{current_working_dir, ProviderDef, ProviderMetadata}; use crate::providers::inventory::InventoryIdentityInput; const PI_ACP_PROVIDER_NAME: &str = "pi-acp"; @@ -44,6 +44,14 @@ impl ProviderDef for PiAcpProvider { fn from_env( model: ModelConfig, extensions: Vec, + ) -> BoxFuture<'static, Result> { + Self::from_env_with_working_dir(model, extensions, current_working_dir()) + } + + fn from_env_with_working_dir( + model: ModelConfig, + extensions: Vec, + working_dir: PathBuf, ) -> BoxFuture<'static, Result> { Box::pin(async move { let config = Config::global(); @@ -62,7 +70,7 @@ impl ProviderDef for PiAcpProvider { args: vec![], env: vec![], env_remove: vec![], - work_dir: std::env::current_dir().unwrap_or_else(|_| PathBuf::from(".")), + work_dir: working_dir, mcp_servers: extension_configs_to_mcp_servers(&extensions), session_mode_id: Some(mode_mapping[&goose_mode].clone()), mode_mapping, diff --git a/crates/goose/src/providers/provider_registry.rs b/crates/goose/src/providers/provider_registry.rs index 2684bfd5f0..a534057e2f 100644 --- a/crates/goose/src/providers/provider_registry.rs +++ b/crates/goose/src/providers/provider_registry.rs @@ -5,10 +5,15 @@ use crate::model::ModelConfig; use anyhow::Result; use futures::future::BoxFuture; use std::collections::HashMap; +use std::path::PathBuf; use std::sync::Arc; pub type ProviderConstructor = Arc< - dyn Fn(ModelConfig, Vec) -> BoxFuture<'static, Result>> + dyn Fn( + ModelConfig, + Vec, + Option, + ) -> BoxFuture<'static, Result>> + Send + Sync, >; @@ -75,7 +80,7 @@ impl ProviderEntry { ) -> Result> { let default_model = &self.metadata.default_model; let model_config = self.normalize_model_config(ModelConfig::new(default_model.as_str())?); - (self.constructor)(model_config, extensions).await + (self.constructor)(model_config, extensions, None).await } pub async fn create( @@ -84,7 +89,17 @@ impl ProviderEntry { extensions: Vec, ) -> Result> { let model = self.normalize_model_config(model); - (self.constructor)(model, extensions).await + (self.constructor)(model, extensions, None).await + } + + pub async fn create_with_working_dir( + &self, + model: ModelConfig, + extensions: Vec, + working_dir: PathBuf, + ) -> Result> { + let model = self.normalize_model_config(model); + (self.constructor)(model, extensions, Some(working_dir)).await } } @@ -111,9 +126,14 @@ impl ProviderRegistry { name, ProviderEntry { metadata, - constructor: Arc::new(|model, extensions| { + constructor: Arc::new(|model, extensions, working_dir| { Box::pin(async move { - let provider = F::from_env(model, extensions).await?; + let provider = match working_dir { + Some(working_dir) => { + F::from_env_with_working_dir(model, extensions, working_dir).await? + } + None => F::from_env(model, extensions).await?, + }; Ok(Arc::new(provider) as Arc) }) }), @@ -220,7 +240,7 @@ impl ProviderRegistry { config.name.clone(), ProviderEntry { metadata: custom_metadata, - constructor: Arc::new(move |model, _extensions| { + constructor: Arc::new(move |model, _extensions, _working_dir| { let result = constructor(model); Box::pin(async move { let provider = result?; diff --git a/crates/goose/tests/acp_common_tests/mod.rs b/crates/goose/tests/acp_common_tests/mod.rs index cad6f9336e..90c225d834 100644 --- a/crates/goose/tests/acp_common_tests/mod.rs +++ b/crates/goose/tests/acp_common_tests/mod.rs @@ -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) }) }) } @@ -448,7 +448,7 @@ pub async fn run_fs_write_text_file_true() { pub async fn run_initialize_doesnt_hit_provider() { 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 { diff --git a/crates/goose/tests/acp_custom_requests_test.rs b/crates/goose/tests/acp_custom_requests_test.rs index de0117a210..e0ec073f9a 100644 --- a/crates/goose/tests/acp_custom_requests_test.rs +++ b/crates/goose/tests/acp_custom_requests_test.rs @@ -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::>::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) + }) + }, + ); + + 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::>::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) + }) + }, + ); + + 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 { diff --git a/crates/goose/tests/acp_fixtures/mod.rs b/crates/goose/tests/acp_fixtures/mod.rs index 01023fdf13..17a307d27d 100644 --- a/crates/goose/tests/acp_fixtures/mod.rs +++ b/crates/goose/tests/acp_fixtures/mod.rs @@ -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 = - 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 = + Arc::new(OpenAiProvider::new(api_client, model_config)); + Ok(provider) + }) + }, + ) }); let agent = GooseAcpAgent::new(GooseAcpAgentOptions { diff --git a/crates/goose/tests/acp_secret_cache_invalidation_test.rs b/crates/goose/tests/acp_secret_cache_invalidation_test.rs index 0c7f5e5e05..f2cf3377da 100644 --- a/crates/goose/tests/acp_secret_cache_invalidation_test.rs +++ b/crates/goose/tests/acp_secret_cache_invalidation_test.rs @@ -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,