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:
@@ -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;
|
||||
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
|
||||
Reference in New Issue
Block a user