feat(cli): load extensions in the background so the prompt is immediately usable (#10403)

Co-authored-by: Jasper Hugo <jasper@jasperhugo.com>
This commit is contained in:
SynthLuvr
2026-08-26 11:19:05 +00:00
committed by GitHub
parent 0ad009d153
commit 9a05f0207c
6 changed files with 673 additions and 127 deletions
@@ -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;
+49 -98
View File
@@ -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<String>,
pub error: anyhow::Error,
}
async fn load_extensions(
agent: Agent,
extensions_to_load: Vec<(String, ExtensionConfig)>,
agent: Arc<Agent>,
extensions: Vec<ExtensionConfig>,
session_id: &str,
) -> Arc<Agent> {
let mut set = JoinSet::new();
let agent_ptr = Arc::new(agent);
let mut waiting_ids: BTreeSet<usize> = (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<usize>| {
let labels: Vec<String> = 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<ExtensionFailure> {
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<ExtensionConfig>,
session_id: &str,
) -> Arc<Agent> {
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::<String>("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(
+282 -9
View File
@@ -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<Agent>,
messages: Conversation,
session_id: String,
completion_cache: Arc<std::sync::RwLock<CompletionCache>>,
@@ -265,6 +265,11 @@ pub struct CliSession {
retry_config: Option<RetryConfig>,
output_format: String,
stats: bool,
/// Background extension loader; drained exclusively by
/// [`CliSession::ensure_extensions_loaded`], the session's single loading
/// gate.
extension_loading: Option<AbortOnDropHandle<Result<Vec<ExtensionFailure>>>>,
loading_announced: bool,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
@@ -351,7 +356,7 @@ fn planner_classification_text(response: &Message) -> Result<String> {
impl CliSession {
#[allow(clippy::too_many_arguments)]
pub async fn new(
agent: Agent,
agent: Arc<Agent>,
session_id: String,
debug: bool,
scheduled_job_id: Option<String>,
@@ -360,6 +365,8 @@ impl CliSession {
retry_config: Option<RetryConfig>,
output_format: String,
stats: bool,
refresh_completions: bool,
extension_loading: Option<AbortOnDropHandle<Vec<ExtensionFailure>>>,
) -> 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<ExtensionConfig>) -> 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<String>,
) -> Result<HashMap<String, Vec<String>>> {
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<Option<output::PromptInfo>> {
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<Vec<PromptMessage>> {
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<String> = self
@@ -692,6 +739,9 @@ impl CliSession {
editor: &mut rustyline::Editor<GooseCompleter, rustyline::history::DefaultHistory>,
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<std::sync::RwLock<CompletionCache>>,
) -> 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<String> = all_providers.iter().map(|(m, _)| m.name.clone()).collect();
let inventory_models: HashMap<String, Vec<String>> = {
@@ -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<AbortOnDropHandle<Vec<ExtensionFailure>>>,
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::<ExtensionFailure>::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::<ExtensionFailure>::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::<ExtensionFailure>::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::<ExtensionFailure>::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());
}
}
+54
View File
@@ -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::<String>("GOOSE_STATUS_HOOK") {
let status = status.to_string();
+42 -19
View File
@@ -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<ExtensionConfig>,
) -> 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<Self>,
extensions: Vec<ExtensionConfig>,
@@ -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)
}
+243
View File
@@ -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<Agent>, Arc<SessionManager>, 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<String> {
let session = session_manager
.get_session(session_id, false)
.await
.unwrap();
let mut names: Vec<String> =
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::<String>::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;