diff --git a/crates/goose-cli/src/scenario_tests/scenario_runner.rs b/crates/goose-cli/src/scenario_tests/scenario_runner.rs index 2498c78cc..7b125cae7 100644 --- a/crates/goose-cli/src/scenario_tests/scenario_runner.rs +++ b/crates/goose-cli/src/scenario_tests/scenario_runner.rs @@ -251,7 +251,7 @@ where .await?; let mut cli_session = CliSession::new( - agent, + Arc::new(agent), session.id, false, None, @@ -260,6 +260,8 @@ where None, "text".to_string(), false, + false, + None, ) .await; diff --git a/crates/goose-cli/src/session/builder.rs b/crates/goose-cli/src/session/builder.rs index eaa1766c8..60ce54b52 100644 --- a/crates/goose-cli/src/session/builder.rs +++ b/crates/goose-cli/src/session/builder.rs @@ -13,10 +13,10 @@ use goose::recipe::Recipe; use goose::session::session_manager::SessionType; use goose::session::EnabledExtensionsState; use rustyline::EditMode; -use std::collections::{BTreeSet, HashMap, HashSet}; +use std::collections::{HashMap, HashSet}; use std::process; use std::sync::Arc; -use tokio::task::JoinSet; +use tokio_util::task::AbortOnDropHandle; const EXTENSION_HINT_MAX_LEN: usize = 5; @@ -232,80 +232,33 @@ impl Default for SessionBuilderConfig { } } +pub struct ExtensionFailure { + pub label: Option, + pub error: anyhow::Error, +} + async fn load_extensions( - agent: Agent, - extensions_to_load: Vec<(String, ExtensionConfig)>, + agent: Arc, + extensions: Vec, session_id: &str, -) -> Arc { - let mut set = JoinSet::new(); - let agent_ptr = Arc::new(agent); - - let mut waiting_ids: BTreeSet = (0..extensions_to_load.len()).collect(); - for (id, (_label, extension)) in extensions_to_load.iter().enumerate() { - let agent_ptr = agent_ptr.clone(); - let cfg = extension.clone(); - let sid = session_id.to_string(); - set.spawn(async move { (id, agent_ptr.add_extension(cfg, &sid).await) }); - } - - let get_message = |waiting_ids: &BTreeSet| { - let labels: Vec = waiting_ids - .iter() - .map(|id| { - extensions_to_load - .get(*id) - .map(|e| e.0.clone()) - .unwrap_or_default() - }) - .collect(); - format!( - "starting {} extensions: {}", - waiting_ids.len(), - labels.join(", ") - ) +) -> Vec { + let results = match agent.add_extensions_bulk(extensions, session_id).await { + Ok(results) => results, + Err(error) => { + tracing::error!("failed to load extensions: {}", error); + return vec![ExtensionFailure { label: None, error }]; + } }; - let spinner = cliclack::spinner(); - spinner.start(get_message(&waiting_ids)); - - let mut failed: Vec<(usize, anyhow::Error)> = Vec::new(); - while let Some(result) = set.join_next().await { - match result { - Ok((id, Ok(_))) => { - waiting_ids.remove(&id); - spinner.set_message(get_message(&waiting_ids)); - } - Ok((id, Err(e))) => failed.push((id, e.into())), - Err(e) => tracing::error!("failed to add extension: {}", e), - } - } - - spinner.clear(); - - for (id, err) in failed { - let label = extensions_to_load - .get(id) - .map(|e| e.0.clone()) - .unwrap_or_default(); - eprintln!( - "{}", - style(format!( - "Warning: Failed to start extension '{}' ({}), continuing without it", - label, err - )) - .yellow() - ); - eprintln!( - "{}", - style(format!( - " Hint: once the session starts, ask goose to help debug the '{}' extension", - label - )) - .dim() - ); - } - - agent_ptr + results + .into_iter() + .filter_map(|r| { + r.error.map(|error| ExtensionFailure { + label: Some(r.name), + error: anyhow::anyhow!(error), + }) + }) + .collect() } struct ResolvedProviderConfig { @@ -662,33 +615,11 @@ async fn collect_extension_configs( Ok(all.into_iter().map(|(_, config)| config).collect()) } -async fn resolve_and_load_extensions( - agent: Agent, - extensions: Vec, - session_id: &str, -) -> Arc { - for warning in goose::config::get_warnings() { - eprintln!("{}", style(format!("Warning: {}", warning)).yellow()); - } - - let extensions_to_load: Vec<(String, ExtensionConfig)> = extensions - .into_iter() - .map(|cfg| (cfg.name(), cfg)) - .collect(); - - load_extensions(agent, extensions_to_load, session_id).await -} - async fn configure_session_prompts( session: &CliSession, config: &Config, session_config: &SessionBuilderConfig, - session_id: &str, ) { - if let Err(e) = session.agent.persist_extension_state(session_id).await { - tracing::warn!("Failed to save extension state: {}", e); - } - if let Some(ref additional_prompt) = session_config.additional_system_prompt { session .agent @@ -880,8 +811,26 @@ pub async fn build_session(session_config: SessionBuilderConfig) -> CliSession { } } - // Extensions are loaded after session creation because we may change directory when resuming - let agent_ptr = resolve_and_load_extensions(agent, extensions_for_provider, &session_id).await; + for warning in goose::config::get_warnings() { + eprintln!("{}", style(format!("Warning: {}", warning)).yellow()); + } + + // Extensions are loaded after session creation because we may change + // directory when resuming. + let agent_ptr = Arc::new(agent); + let loading_handle = match agent_ptr + .persist_extension_configs(&session_id, extensions_for_provider.clone()) + .await + { + Ok(()) => AbortOnDropHandle::new(tokio::spawn({ + let agent = agent_ptr.clone(); + let sid = session_id.clone(); + async move { load_extensions(agent, extensions_for_provider, &sid).await } + })), + Err(error) => AbortOnDropHandle::new(tokio::spawn(async move { + vec![ExtensionFailure { label: None, error }] + })), + }; let edit_mode = config .get_param::("EDIT_MODE") @@ -898,7 +847,7 @@ pub async fn build_session(session_config: SessionBuilderConfig) -> CliSession { let debug_mode = session_config.debug || config.get_param("GOOSE_DEBUG").unwrap_or(false); let session = CliSession::new( - Arc::try_unwrap(agent_ptr).unwrap_or_else(|_| panic!("There should be no more references")), + agent_ptr, session_id.clone(), debug_mode, session_config.scheduled_job_id.clone(), @@ -907,10 +856,12 @@ pub async fn build_session(session_config: SessionBuilderConfig) -> CliSession { recipe.and_then(|r| r.retry.clone()), session_config.output_format.clone(), session_config.stats, + session_config.interactive, + Some(loading_handle), ) .await; - configure_session_prompts(&session, config, &session_config, &session_id).await; + configure_session_prompts(&session, config, &session_config).await; if !session_config.quiet { output::display_session_info( diff --git a/crates/goose-cli/src/session/mod.rs b/crates/goose-cli/src/session/mod.rs index 790419715..2b8e73a85 100644 --- a/crates/goose-cli/src/session/mod.rs +++ b/crates/goose-cli/src/session/mod.rs @@ -19,7 +19,7 @@ use std::str::FromStr; use tokio::signal::ctrl_c; use tokio_util::task::AbortOnDropHandle; -pub use builder::{build_session, SessionBuilderConfig}; +pub use builder::{build_session, ExtensionFailure, SessionBuilderConfig}; use console::Color; use goose::agents::platform_extensions::developer::shell::{ parse_shell_output_notification, ShellOutputNotificationParams, ShellOutputStream, @@ -253,7 +253,7 @@ impl HistoryManager { } pub struct CliSession { - agent: Agent, + agent: Arc, messages: Conversation, session_id: String, completion_cache: Arc>, @@ -265,6 +265,11 @@ pub struct CliSession { retry_config: Option, output_format: String, stats: bool, + /// Background extension loader; drained exclusively by + /// [`CliSession::ensure_extensions_loaded`], the session's single loading + /// gate. + extension_loading: Option>>>, + loading_announced: bool, } #[derive(Debug, Clone, Copy, PartialEq, Eq)] @@ -351,7 +356,7 @@ fn planner_classification_text(response: &Message) -> Result { impl CliSession { #[allow(clippy::too_many_arguments)] pub async fn new( - agent: Agent, + agent: Arc, session_id: String, debug: bool, scheduled_job_id: Option, @@ -360,6 +365,8 @@ impl CliSession { retry_config: Option, output_format: String, stats: bool, + refresh_completions: bool, + extension_loading: Option>>, ) -> Self { let messages = agent .config @@ -368,12 +375,27 @@ impl CliSession { .await .map(|session| session.conversation.unwrap_or_default()) .unwrap(); + let completion_cache = Arc::new(std::sync::RwLock::new(CompletionCache::new())); + let extension_loading = extension_loading.map(|handle| { + let agent = agent.clone(); + let session_id = session_id.clone(); + let completion_cache = completion_cache.clone(); + AbortOnDropHandle::new(tokio::spawn(async move { + let failures = handle + .await + .map_err(|error| anyhow::anyhow!("Extension loading task failed: {}", error))?; + if refresh_completions { + Self::refresh_completion_cache(&agent, &session_id, &completion_cache).await?; + } + Ok(failures) + })) + }); CliSession { agent, messages, session_id, - completion_cache: Arc::new(std::sync::RwLock::new(CompletionCache::new())), + completion_cache, debug, run_mode: RunMode::Normal, scheduled_job_id, @@ -382,6 +404,8 @@ impl CliSession { retry_config, output_format, stats, + extension_loading, + loading_announced: false, } } @@ -504,6 +528,8 @@ impl CliSession { } async fn add_and_persist_extensions(&mut self, configs: Vec) -> Result<()> { + // Extension-set mutations must not race the background loader. + self.ensure_extensions_loaded(true).await?; for config in configs { self.agent .add_extension(config, &self.session_id) @@ -538,6 +564,7 @@ impl CliSession { &mut self, extension: Option, ) -> Result>> { + self.ensure_extensions_loaded(true).await?; let prompts = self.agent.list_extension_prompts(&self.session_id).await; // Early validation if filtering by extension @@ -559,6 +586,7 @@ impl CliSession { } pub async fn get_prompt_info(&mut self, name: &str) -> Result> { + self.ensure_extensions_loaded(true).await?; let prompts = self.agent.list_extension_prompts(&self.session_id).await; // Find which extension has this prompt @@ -577,6 +605,7 @@ impl CliSession { } pub async fn get_prompt(&mut self, name: &str, arguments: Value) -> Result> { + self.ensure_extensions_loaded(true).await?; Ok(self .agent .get_prompt(&self.session_id, name, arguments) @@ -591,6 +620,7 @@ impl CliSession { cancel_token: CancellationToken, interactive: bool, ) -> Result<()> { + self.ensure_extensions_loaded(interactive).await?; let cancel_token = cancel_token.clone(); self.push_message(message); self.process_agent_response(interactive, cancel_token) @@ -630,15 +660,32 @@ impl CliSession { let msg = Message::user().with_text(&prompt); self.process_message(msg, CancellationToken::default(), true) .await?; + } else if self + .extension_loading + .as_ref() + .is_some_and(|h| !h.is_finished()) + { + self.loading_announced = true; + output::show_loading_extensions_background(); } - self.update_completion_cache().await?; - + // The completion cache starts empty: populating it here would issue + // prompts/list to every connected extension, so a slow responder could + // block the prompt right after we announced loading continues in the + // background. The loader refreshes the shared cache when it finishes. let mut editor = self.create_editor()?; let history_manager = HistoryManager::new(); history_manager.load(&mut editor); loop { + if self + .extension_loading + .as_ref() + .is_some_and(|h| h.is_finished()) + { + self.ensure_extensions_loaded(true).await?; + } + self.display_context_usage().await?; let conversation_strings: Vec = self @@ -692,6 +739,9 @@ impl CliSession { editor: &mut rustyline::Editor, conversation_messages: &[String], ) -> Result<()> { + // The REPL's loading gate: every command, including future ones, waits + // here until background extension loading has finished. + self.ensure_extensions_loaded(true).await?; match input { InputResult::Message(content) => { self.handle_message_input(&content, history, editor).await?; @@ -1168,6 +1218,11 @@ impl CliSession { return Ok(()); } + // The background loader pins the session id it was spawned with, and the + // handle_input gate has already drained it, so extensions can be torn + // down and re-added under a new session id without racing an in-flight + // load. + let new_session_id = match self.prepare_successor_session().await { Ok(id) => id, Err(e) => { @@ -1864,10 +1919,45 @@ impl CliSession { Ok(()) } + /// The session's single extension-loading gate: wait for the background + /// loader and surface any failures. + /// + /// Entry points that can touch the agent or the extension set call this + /// rather than gating handlers individually, so new commands inherit the + /// gate automatically. `interactive` controls the waiting/ready + /// indicators. The loader refreshes the shared completion cache when it + /// finishes so completions become available while readline is active. + async fn ensure_extensions_loaded(&mut self, interactive: bool) -> Result<()> { + if let Some(handle) = self.extension_loading.take() { + let was_in_progress = !handle.is_finished(); + if interactive && was_in_progress { + output::show_waiting_for_extensions(); + } + let failures = handle + .await + .map_err(|e| anyhow::anyhow!("Extension loading task failed: {}", e))??; + output::show_extension_failures(&failures); + + if interactive && (was_in_progress || self.loading_announced) { + output::show_extensions_ready(); + } + self.loading_announced = false; + } + Ok(()) + } + pub async fn update_completion_cache(&mut self) -> Result<()> { - let prompts = self.agent.list_extension_prompts(&self.session_id).await; + Self::refresh_completion_cache(&self.agent, &self.session_id, &self.completion_cache).await + } + + async fn refresh_completion_cache( + agent: &Agent, + session_id: &str, + completion_cache: &Arc>, + ) -> Result<()> { + let prompts = agent.list_extension_prompts(session_id).await; let all_providers = goose::providers::providers().await; - let session_provider = self.agent.provider().await?.get_name().to_string(); + let session_provider = agent.provider().await?.get_name().to_string(); let provider_ids: Vec = all_providers.iter().map(|(m, _)| m.name.clone()).collect(); let inventory_models: HashMap> = { @@ -1896,7 +1986,7 @@ impl CliSession { }) .collect(); - let mut cache = self.completion_cache.write().unwrap(); + let mut cache = completion_cache.write().unwrap(); cache.prompts.clear(); cache.prompt_info.clear(); @@ -3308,4 +3398,187 @@ mod tests { Some(&serde_json::json!("marker")) ); } + + struct StubProvider; + + #[async_trait::async_trait] + impl Provider for StubProvider { + fn get_name(&self) -> &str { + "stub" + } + + async fn stream( + &self, + _model_config: &goose_providers::model::ModelConfig, + _system: &str, + _messages: &[Message], + _tools: &[rmcp::model::Tool], + ) -> std::result::Result< + goose::providers::base::MessageStream, + goose_providers::errors::ProviderError, + > { + Ok(goose::providers::base::stream_from_single_message( + Message::assistant().with_text("stub reply"), + ProviderUsage::new( + "stub".to_string(), + goose_providers::conversation::token_usage::Usage::default(), + ), + )) + } + } + + async fn session_with_loader( + extension_loading: Option>>, + refresh_completions: bool, + ) -> CliSession { + let temp_dir = tempfile::TempDir::new().unwrap(); + let session_manager = SessionManager::new(temp_dir.path().to_path_buf()); + let session = session_manager + .create_session( + temp_dir.path().to_path_buf(), + "Loading gate test".to_string(), + goose::session::SessionType::User, + GooseMode::default(), + ) + .await + .unwrap(); + + let agent = goose::agents::Agent::with_config(goose::agents::AgentConfig::new( + Arc::new(session_manager), + Arc::new(goose::config::PermissionManager::new( + temp_dir.path().to_path_buf(), + )), + None, + GooseMode::default(), + // Disable background session naming so the test agent starts no + // provider-dependent tasks. + true, + goose::agents::GoosePlatform::GooseCli, + )); + agent + .update_provider( + Arc::new(StubProvider), + goose_providers::model::ModelConfig::new("stub-model"), + &session.id, + ) + .await + .unwrap(); + + CliSession::new( + Arc::new(agent), + session.id, + false, + None, + None, + None, + None, + "text".to_string(), + false, + refresh_completions, + extension_loading, + ) + .await + } + + #[tokio::test] + async fn commands_wait_for_background_extension_loading() { + let (release, released) = tokio::sync::oneshot::channel::<()>(); + let loader = AbortOnDropHandle::new(tokio::spawn(async move { + let _ = released.await; + Vec::::new() + })); + + let mut session = session_with_loader(Some(loader), false).await; + let mut editor = session.create_editor().unwrap(); + + let handled = tokio::spawn(async move { + let history = HistoryManager::new(); + session + .handle_input( + InputResult::Plan(input::PlanCommandOptions { + message_text: String::new(), + }), + &history, + &mut editor, + &[], + ) + .await + .expect("handle_input failed"); + session.run_mode + }); + + tokio::time::sleep(Duration::from_millis(250)).await; + assert!( + !handled.is_finished(), + "a command was handled before background extension loading finished" + ); + + release.send(()).unwrap(); + let run_mode = handled.await.unwrap(); + assert!(matches!(run_mode, RunMode::Plan)); + } + + #[tokio::test] + async fn ensure_extensions_loaded_drains_the_loader_once() { + let loader = AbortOnDropHandle::new(tokio::spawn(async { Vec::::new() })); + let mut session = session_with_loader(Some(loader), false).await; + + session.ensure_extensions_loaded(false).await.unwrap(); + assert!(session.extension_loading.is_none()); + + // A second pass is a no-op, not an error. + session.ensure_extensions_loaded(false).await.unwrap(); + assert!(session.extension_loading.is_none()); + } + + #[tokio::test] + async fn background_loader_refreshes_completions_before_the_gate_runs() { + let (release, released) = tokio::sync::oneshot::channel::<()>(); + let loader = AbortOnDropHandle::new(tokio::spawn(async move { + let _ = released.await; + Vec::::new() + })); + let session = session_with_loader(Some(loader), true).await; + + assert!(session + .completion_cache + .read() + .unwrap() + .current_session_provider + .is_empty()); + release.send(()).unwrap(); + + tokio::time::timeout(Duration::from_secs(5), async { + loop { + if session + .completion_cache + .read() + .unwrap() + .current_session_provider + == "stub" + { + break; + } + tokio::task::yield_now().await; + } + }) + .await + .expect("completion cache was not refreshed in the background"); + assert!(session.extension_loading.is_some()); + } + + #[tokio::test] + async fn headless_loader_skips_completion_refresh() { + let loader = AbortOnDropHandle::new(tokio::spawn(async { Vec::::new() })); + let mut session = session_with_loader(Some(loader), false).await; + + session.ensure_extensions_loaded(false).await.unwrap(); + + assert!(session + .completion_cache + .read() + .unwrap() + .current_session_provider + .is_empty()); + } } diff --git a/crates/goose-cli/src/session/output.rs b/crates/goose-cli/src/session/output.rs index 725bd0fc8..48f23308c 100644 --- a/crates/goose-cli/src/session/output.rs +++ b/crates/goose-cli/src/session/output.rs @@ -1,3 +1,4 @@ +use crate::session::builder::ExtensionFailure; use anstream::{adapter::strip_str, println}; use bat::WrappingMode; use console::{measure_text_width, style, Color, StyledObject, Term}; @@ -196,6 +197,59 @@ pub fn hide_thinking() { } } +pub fn show_loading_extensions_background() { + eprintln!( + " {}", + style("⏳ loading extensions in background...").dim() + ); +} + +pub fn show_waiting_for_extensions() { + eprintln!( + " {}", + style("⏳ waiting for extensions to finish loading...").dim() + ); +} + +pub fn show_extensions_ready() { + eprintln!(" {}", style("✓ extensions ready").green()); +} + +pub fn show_extension_failures(failures: &[ExtensionFailure]) { + for failure in failures { + match failure.label.as_deref() { + None => { + eprintln!( + "{}", + style(format!( + " ⚠ Failed to start extensions ({})", + failure.error + )) + .yellow() + ); + } + Some(label) => { + eprintln!( + "{}", + style(format!( + " ⚠ Failed to start extension '{}' ({}), continuing without it", + label, failure.error + )) + .yellow() + ); + eprintln!( + "{}", + style(format!( + " Hint: ask goose to help debug the '{}' extension", + label + )) + .dim() + ); + } + } + } +} + pub fn run_status_hook(status: &str) { if let Ok(hook) = Config::global().get_param::("GOOSE_STATUS_HOOK") { let status = status.to_string(); diff --git a/crates/goose/src/agents/agent.rs b/crates/goose/src/agents/agent.rs index a366beab6..600be0fe5 100644 --- a/crates/goose/src/agents/agent.rs +++ b/crates/goose/src/agents/agent.rs @@ -1350,8 +1350,17 @@ impl Agent { /// Save current extension state to session by session_id pub async fn persist_extension_state(&self, session_id: &str) -> Result<()> { - let extensions_state = - EnabledExtensionsState::new(self.extension_configs_for_persistence().await); + self.persist_extension_configs(session_id, self.extension_configs_for_persistence().await) + .await + } + + /// Save the provided extension configuration to session metadata. + pub async fn persist_extension_configs( + &self, + session_id: &str, + extensions: Vec, + ) -> Result<()> { + let extensions_state = EnabledExtensionsState::new(extensions); let session_manager = self.config.session_manager.clone(); let session = session_manager.get_session(session_id, false).await?; @@ -1488,6 +1497,11 @@ impl Agent { /// /// Unlike `add_extension`, this avoids per-extension persistence and acquires /// the container lock once upfront to prevent serialisation of the parallel futures. + /// + /// State is persisted once every extension has settled, even when all of them + /// fail: the session's enabled list records what actually loaded, so failed + /// extensions are dropped instead of staying marked as enabled and being + /// retried on every subsequent resume. pub async fn add_extensions_bulk( self: &Arc, extensions: Vec, @@ -1511,30 +1525,41 @@ impl Agent { .into_iter() .map(|config| { let ext_manager = Arc::clone(&self.extension_manager); + let agent = Arc::clone(self); let working_dir = working_dir.clone(); let container = container.clone(); let sid = session_id.to_string(); async move { let name = config.name().to_string(); - match ext_manager - .add_extension(config, working_dir, container.as_ref(), Some(&sid)) - .await - { - Ok(_) => ExtensionLoadResult { - name, - success: true, - error: None, - }, - Err(e) => { - let error_msg = e.to_string(); - warn!("Failed to load extension {}: {}", name, error_msg); + match &config { + ExtensionConfig::Frontend { .. } => { + agent.insert_frontend_extension(config.clone()).await; ExtensionLoadResult { name, - success: false, - error: Some(error_msg), + success: true, + error: None, } } + _ => match ext_manager + .add_extension(config, working_dir, container.as_ref(), Some(&sid)) + .await + { + Ok(_) => ExtensionLoadResult { + name, + success: true, + error: None, + }, + Err(e) => { + let error = e.to_string(); + warn!("Failed to load extension {}: {}", name, error); + ExtensionLoadResult { + name, + success: false, + error: Some(error), + } + } + }, } } }) @@ -1542,9 +1567,7 @@ impl Agent { let results = futures::future::join_all(extension_futures).await; - if results.iter().any(|r| r.success) { - self.persist_extension_state(session_id).await?; - } + self.persist_extension_state(session_id).await?; Ok(results) } diff --git a/crates/goose/tests/agent.rs b/crates/goose/tests/agent.rs index a9a733190..a8265d1a9 100644 --- a/crates/goose/tests/agent.rs +++ b/crates/goose/tests/agent.rs @@ -2954,6 +2954,249 @@ mod tests { } } + mod add_extensions_bulk_tests { + use super::*; + use goose::agents::extension::Envs; + use goose::agents::{AgentConfig, ExtensionConfig}; + use goose::config::permission::PermissionManager; + use goose::config::GooseMode; + use goose::session::session_manager::SessionType; + use goose::session::{ + EnabledExtensionsState, ExtensionData, ExtensionState, SessionManager, + }; + use rmcp::model::Tool; + use rmcp::object; + use tempfile::TempDir; + + fn frontend_extension(name: &str) -> ExtensionConfig { + ExtensionConfig::Frontend { + name: name.to_string(), + description: format!("Frontend test extension {name}"), + tools: vec![Tool::new( + format!("{name}__tool"), + format!("Run a tool from {name}"), + object!({ + "type": "object", + "properties": { + "message": { "type": "string" } + }, + "required": ["message"] + }), + )], + instructions: None, + bundled: None, + available_tools: vec![], + } + } + + fn unloadable_stdio_extension(name: &str) -> ExtensionConfig { + ExtensionConfig::Stdio { + name: name.to_string(), + description: format!("Unloadable test extension {name}"), + cmd: "goose-test-definitely-missing-binary".to_string(), + args: vec![], + envs: Envs::default(), + env_keys: vec![], + timeout: Some(1), + cwd: None, + bundled: None, + available_tools: vec![], + } + } + + async fn setup_agent_and_session( + test_name: &str, + ) -> (Arc, Arc, String, TempDir) { + let temp_dir = TempDir::new().unwrap(); + let data_dir = temp_dir.path().to_path_buf(); + let session_manager = Arc::new(SessionManager::new(data_dir.clone())); + let permission_manager = Arc::new(PermissionManager::new(data_dir)); + let agent = Arc::new(Agent::with_config(AgentConfig::new( + session_manager.clone(), + permission_manager, + None, + GooseMode::default(), + false, + GoosePlatform::GooseDesktop, + ))); + + let session = session_manager + .create_session( + std::env::current_dir().unwrap(), + test_name.to_string(), + SessionType::Hidden, + GooseMode::default(), + ) + .await + .unwrap(); + + (agent, session_manager, session.id, temp_dir) + } + + async fn persisted_extension_names( + session_manager: &SessionManager, + session_id: &str, + ) -> Vec { + let session = session_manager + .get_session(session_id, false) + .await + .unwrap(); + let mut names: Vec = + EnabledExtensionsState::from_extension_data(&session.extension_data) + .expect("enabled extensions state should be persisted") + .extensions + .iter() + .map(|extension| extension.name()) + .collect(); + names.sort(); + names + } + + #[tokio::test] + async fn test_bulk_load_persists_loaded_extensions() { + let (agent, session_manager, session_id, _temp_dir) = + setup_agent_and_session("bulk-load-persist-success").await; + + let results = agent + .add_extensions_bulk( + vec![ + frontend_extension("frontend-a"), + frontend_extension("frontend-b"), + ], + &session_id, + ) + .await + .unwrap(); + + assert!(results.iter().all(|result| result.success)); + assert_eq!( + persisted_extension_names(&session_manager, &session_id).await, + vec!["frontend-a".to_string(), "frontend-b".to_string()] + ); + } + + #[tokio::test] + async fn test_bulk_load_partial_failure_persists_only_loaded_extensions() { + let (agent, session_manager, session_id, _temp_dir) = + setup_agent_and_session("bulk-load-persist-partial-failure").await; + + let results = agent + .add_extensions_bulk( + vec![ + frontend_extension("frontend-ok"), + unloadable_stdio_extension("broken"), + ], + &session_id, + ) + .await + .unwrap(); + + assert_eq!(results.len(), 2); + assert!(results + .iter() + .any(|result| result.name == "frontend-ok" && result.success)); + let broken = results + .iter() + .find(|result| result.name == "broken") + .expect("broken extension should report a result"); + assert!(!broken.success); + assert!(broken.error.is_some()); + + assert_eq!( + persisted_extension_names(&session_manager, &session_id).await, + vec!["frontend-ok".to_string()] + ); + } + + #[tokio::test] + async fn test_bulk_load_total_failure_drops_failed_extensions_from_session_state() { + let (agent, session_manager, session_id, _temp_dir) = + setup_agent_and_session("bulk-load-persist-total-failure").await; + + // Seed the session with the extensions up front, mirroring a resume + // where the enabled list is read back from session metadata. + let extensions = vec![ + unloadable_stdio_extension("broken-one"), + unloadable_stdio_extension("broken-two"), + ]; + let mut extension_data = ExtensionData::new(); + EnabledExtensionsState::new(extensions.clone()) + .to_extension_data(&mut extension_data) + .unwrap(); + session_manager + .update(&session_id) + .extension_data(extension_data) + .apply() + .await + .unwrap(); + + let results = agent + .add_extensions_bulk(extensions, &session_id) + .await + .unwrap(); + + assert_eq!(results.len(), 2); + assert!( + results.iter().all(|result| !result.success), + "expected every extension load to fail: {results:?}" + ); + + // The failed extensions must not stay marked as enabled in the + // session, otherwise every future resume retries them. + assert_eq!( + persisted_extension_names(&session_manager, &session_id).await, + Vec::::new() + ); + } + + #[tokio::test] + async fn test_bulk_load_cancellation_preserves_pending_extensions() { + let (agent, session_manager, session_id, _temp_dir) = + setup_agent_and_session("bulk-load-persist-cancellation").await; + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + let (accepted_tx, accepted_rx) = tokio::sync::oneshot::channel(); + let server = tokio::spawn(async move { + let (_connection, _) = listener.accept().await.unwrap(); + let _ = accepted_tx.send(()); + std::future::pending::<()>().await; + }); + + let extension = ExtensionConfig::streamable_http( + "pending".to_string(), + format!("http://{address}"), + "Pending test extension".to_string(), + 30_u64, + ); + agent + .persist_extension_configs(&session_id, vec![extension.clone()]) + .await + .unwrap(); + let load = tokio::spawn({ + let agent = agent.clone(); + let session_id = session_id.clone(); + async move { + agent + .add_extensions_bulk(vec![extension], &session_id) + .await + } + }); + + tokio::time::timeout(std::time::Duration::from_secs(5), accepted_rx) + .await + .expect("extension did not connect") + .expect("test server stopped before accepting a connection"); + + load.abort(); + assert!(load.await.unwrap_err().is_cancelled()); + server.abort(); + assert_eq!( + persisted_extension_names(&session_manager, &session_id).await, + vec!["pending".to_string()] + ); + } + } + mod audience_tool_result_tests { use super::*; use async_trait::async_trait;