mirror of
https://github.com/aaif-goose/goose.git
synced 2026-07-03 14:10:03 +02:00
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:
@@ -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<ExtensionConfig>,
|
||||
Option<PathBuf>,
|
||||
) -> BoxFuture<'static, Result<Arc<dyn Provider>>>
|
||||
+ 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<PermissionManager> {
|
||||
Arc::clone(&self.permission_manager)
|
||||
@@ -1093,8 +1109,15 @@ impl GooseAcpAgent {
|
||||
provider_name: &str,
|
||||
model_config: crate::model::ModelConfig,
|
||||
extensions: Vec<ExtensionConfig>,
|
||||
working_dir: Option<PathBuf>,
|
||||
) -> Result<Arc<dyn Provider>> {
|
||||
(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<NewSessionResponse, agent_client_protocol::Error> {
|
||||
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<LoadSessionResponse, agent_client_protocol::Error> {
|
||||
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::<AgentSetupSignal>(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<Client>,
|
||||
args: ForkSessionRequest,
|
||||
) -> Result<ForkSessionResponse, agent_client_protocol::Error> {
|
||||
validate_absolute_cwd(&args.cwd)?;
|
||||
let source_session_id = &*args.session_id.0;
|
||||
|
||||
let new_session = self
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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<crate::config::ExtensionConfig>,
|
||||
) -> BoxFuture<'static, Result<AcpProvider>> {
|
||||
Self::from_env_with_working_dir(model, extensions, current_working_dir())
|
||||
}
|
||||
|
||||
fn from_env_with_working_dir(
|
||||
model: ModelConfig,
|
||||
extensions: Vec<crate::config::ExtensionConfig>,
|
||||
working_dir: PathBuf,
|
||||
) -> BoxFuture<'static, Result<AcpProvider>> {
|
||||
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,
|
||||
|
||||
@@ -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<ExtensionConfig>,
|
||||
_working_dir: PathBuf,
|
||||
) -> BoxFuture<'static, Result<Self::Provider>>
|
||||
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,
|
||||
|
||||
@@ -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<crate::config::ExtensionConfig>,
|
||||
) -> BoxFuture<'static, Result<AcpProvider>> {
|
||||
Self::from_env_with_working_dir(model, extensions, current_working_dir())
|
||||
}
|
||||
|
||||
fn from_env_with_working_dir(
|
||||
model: ModelConfig,
|
||||
extensions: Vec<crate::config::ExtensionConfig>,
|
||||
working_dir: PathBuf,
|
||||
) -> BoxFuture<'static, Result<AcpProvider>> {
|
||||
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,
|
||||
|
||||
@@ -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<crate::config::ExtensionConfig>,
|
||||
) -> BoxFuture<'static, Result<AcpProvider>> {
|
||||
Self::from_env_with_working_dir(model, extensions, current_working_dir())
|
||||
}
|
||||
|
||||
fn from_env_with_working_dir(
|
||||
model: ModelConfig,
|
||||
extensions: Vec<crate::config::ExtensionConfig>,
|
||||
working_dir: PathBuf,
|
||||
) -> BoxFuture<'static, Result<AcpProvider>> {
|
||||
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,
|
||||
|
||||
@@ -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<crate::config::ExtensionConfig>,
|
||||
) -> BoxFuture<'static, Result<AcpProvider>> {
|
||||
Self::from_env_with_working_dir(model, extensions, current_working_dir())
|
||||
}
|
||||
|
||||
fn from_env_with_working_dir(
|
||||
model: ModelConfig,
|
||||
extensions: Vec<crate::config::ExtensionConfig>,
|
||||
working_dir: PathBuf,
|
||||
) -> BoxFuture<'static, Result<AcpProvider>> {
|
||||
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,
|
||||
|
||||
@@ -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<ExtensionConfig>,
|
||||
working_dir: PathBuf,
|
||||
) -> Result<Arc<dyn Provider>> {
|
||||
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<str>,
|
||||
extensions: Vec<ExtensionConfig>,
|
||||
|
||||
@@ -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};
|
||||
|
||||
@@ -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<crate::config::ExtensionConfig>,
|
||||
) -> BoxFuture<'static, Result<AcpProvider>> {
|
||||
Self::from_env_with_working_dir(model, extensions, current_working_dir())
|
||||
}
|
||||
|
||||
fn from_env_with_working_dir(
|
||||
model: ModelConfig,
|
||||
extensions: Vec<crate::config::ExtensionConfig>,
|
||||
working_dir: PathBuf,
|
||||
) -> BoxFuture<'static, Result<AcpProvider>> {
|
||||
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,
|
||||
|
||||
@@ -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<ExtensionConfig>) -> BoxFuture<'static, Result<Arc<dyn Provider>>>
|
||||
dyn Fn(
|
||||
ModelConfig,
|
||||
Vec<ExtensionConfig>,
|
||||
Option<PathBuf>,
|
||||
) -> BoxFuture<'static, Result<Arc<dyn Provider>>>
|
||||
+ Send
|
||||
+ Sync,
|
||||
>;
|
||||
@@ -75,7 +80,7 @@ impl ProviderEntry {
|
||||
) -> Result<Arc<dyn Provider>> {
|
||||
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<ExtensionConfig>,
|
||||
) -> Result<Arc<dyn Provider>> {
|
||||
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<ExtensionConfig>,
|
||||
working_dir: PathBuf,
|
||||
) -> Result<Arc<dyn Provider>> {
|
||||
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<dyn Provider>)
|
||||
})
|
||||
}),
|
||||
@@ -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?;
|
||||
|
||||
@@ -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