diff --git a/Cargo.lock b/Cargo.lock index 9f014e3bd..b5af4a3fe 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1169,7 +1169,7 @@ dependencies = [ "quote", "regex", "rustc-hash 1.1.0", - "shlex", + "shlex 1.3.0", "syn 2.0.117", ] diff --git a/crates/goose-server/src/routes/agent.rs b/crates/goose-server/src/routes/agent.rs index 25b27eecf..081a11654 100644 --- a/crates/goose-server/src/routes/agent.rs +++ b/crates/goose-server/src/routes/agent.rs @@ -400,6 +400,20 @@ async fn resume_agent( status: code, })?; + if !state.has_extension_loading_task(&payload.session_id).await { + let session_for_task = session.clone(); + let agent_for_task = agent.clone(); + let session_id_for_task = payload.session_id.clone(); + let task = tokio::spawn(async move { + agent_for_task + .load_extensions_from_session(&session_for_task) + .await + }); + state + .set_extension_loading_task(session_id_for_task, task) + .await; + } + let provider_changed = agent .restore_provider_from_session(&session) .await @@ -421,8 +435,8 @@ async fn resume_agent( session }; - let extension_results = - if let Some(results) = state.take_extension_loading_task(&payload.session_id).await { + let extension_results = match state.take_extension_loading_task(&payload.session_id).await { + Ok(Some(results)) => { tracing::debug!( "Using background extension loading results for session {}", payload.session_id @@ -431,13 +445,26 @@ async fn resume_agent( .remove_extension_loading_task(&payload.session_id) .await; results - } else { + } + Ok(None) => { tracing::debug!( - "No background task found, loading extensions for session {}", + "Extension loading task for session {} was already consumed", payload.session_id ); + vec![] + } + Err(e) => { + state + .remove_extension_loading_task(&payload.session_id) + .await; + tracing::warn!( + "Background extension loading failed for session {}, retrying synchronously: {}", + payload.session_id, + e + ); agent.load_extensions_from_session(&session).await - }; + } + }; (Some(extension_results), session) } else { @@ -719,6 +746,8 @@ async fn agent_add_extension( #[cfg(feature = "telemetry")] let extension_name = request.config.name(); + ensure_extensions_loaded(&state, &request.session_id).await?; + let agent = state.get_agent(request.session_id.clone()).await?; agent @@ -751,6 +780,8 @@ async fn agent_remove_extension( State(state): State>, Json(request): Json, ) -> Result { + ensure_extensions_loaded(&state, &request.session_id).await?; + let agent = state.get_agent(request.session_id.clone()).await?; agent @@ -981,13 +1012,47 @@ async fn update_working_dir( Ok(StatusCode::OK) } -async fn ensure_extensions_loaded(state: &AppState, session_id: &str) { - if let Some(_results) = state.take_extension_loading_task(session_id).await { - tracing::debug!( - "Awaited background extension loading for session {} before serving request", - session_id - ); - state.remove_extension_loading_task(session_id).await; +async fn ensure_extensions_loaded(state: &AppState, session_id: &str) -> Result<(), ErrorResponse> { + match state.take_extension_loading_task(session_id).await { + Ok(Some(_)) => { + tracing::debug!( + "Awaited background extension loading for session {} before serving request", + session_id + ); + state.remove_extension_loading_task(session_id).await; + Ok(()) + } + Ok(None) => Ok(()), + Err(e) => { + state.remove_extension_loading_task(session_id).await; + tracing::warn!( + "Background extension loading failed for session {}, retrying synchronously: {}", + session_id, + e + ); + let session = state + .session_manager() + .get_session(session_id, false) + .await + .map_err(|err| ErrorResponse { + message: format!( + "Failed to get session after extension loading failed: {}", + err + ), + status: StatusCode::NOT_FOUND, + })?; + let agent = state + .get_agent(session_id.to_string()) + .await + .map_err(|err| { + ErrorResponse::internal(format!( + "Failed to get agent after extension loading failed: {}", + err + )) + })?; + agent.load_extensions_from_session(&session).await; + Ok(()) + } } } @@ -1009,7 +1074,9 @@ async fn read_resource( ) -> Result, StatusCode> { use rmcp::model::ResourceContents; - ensure_extensions_loaded(&state, &payload.session_id).await; + ensure_extensions_loaded(&state, &payload.session_id) + .await + .map_err(|err| err.status)?; let agent = state .get_agent_for_route(payload.session_id.clone()) @@ -1091,7 +1158,7 @@ async fn call_tool( State(state): State>, Json(payload): Json, ) -> Result, ErrorResponse> { - ensure_extensions_loaded(&state, &payload.session_id).await; + ensure_extensions_loaded(&state, &payload.session_id).await?; let agent = state .get_agent_for_route(payload.session_id.clone()) diff --git a/crates/goose-server/src/state.rs b/crates/goose-server/src/state.rs index 6d109c51e..3543d051e 100644 --- a/crates/goose-server/src/state.rs +++ b/crates/goose-server/src/state.rs @@ -84,27 +84,39 @@ impl AppState { tasks.insert(session_id, Arc::new(Mutex::new(Some(task)))); } + pub async fn has_extension_loading_task(&self, session_id: &str) -> bool { + let tasks = self.extension_loading_tasks.lock().await; + tasks.contains_key(session_id) + } + pub async fn take_extension_loading_task( &self, session_id: &str, - ) -> Option> { + ) -> Result>, tokio::task::JoinError> { let task_holder = { let tasks = self.extension_loading_tasks.lock().await; tasks.get(session_id).cloned() }; if let Some(holder) = task_holder { - let task = holder.lock().await.take(); - if let Some(handle) = task { + let mut task = holder.lock().await; + if let Some(handle) = task.as_mut() { + // Keep the per-session task locked and discoverable while awaiting so + // concurrent routes cannot mutate extensions before background loading finishes. match handle.await { - Ok(results) => return Some(results), + Ok(results) => { + task.take(); + return Ok(Some(results)); + } Err(e) => { + task.take(); tracing::warn!("Background extension loading task failed: {}", e); + return Err(e); } } } } - None + Ok(None) } pub async fn remove_extension_loading_task(&self, session_id: &str) { diff --git a/ui/desktop/src/components/bottom_menu/BottomMenuExtensionSelection.tsx b/ui/desktop/src/components/bottom_menu/BottomMenuExtensionSelection.tsx index 7ac91a9de..1c7ca464a 100644 --- a/ui/desktop/src/components/bottom_menu/BottomMenuExtensionSelection.tsx +++ b/ui/desktop/src/components/bottom_menu/BottomMenuExtensionSelection.tsx @@ -133,13 +133,7 @@ export const BottomMenuExtensionSelection = ({ sessionId }: BottomMenuExtensionS let controller: AbortController | null = null; - const loadExtensionsForCurrentSession = (event: Event) => { - const targetSessionId = (event as CustomEvent<{ sessionId?: string }>).detail?.sessionId; - - if (targetSessionId !== sessionId) { - return; - } - + const loadForSession = (targetSessionId: string) => { controller?.abort(); const currentController = new AbortController(); controller = currentController; @@ -154,8 +148,21 @@ export const BottomMenuExtensionSelection = ({ sessionId }: BottomMenuExtensionS }); }; + const loadExtensionsForCurrentSession = (event: Event) => { + const targetSessionId = (event as CustomEvent<{ sessionId?: string }>).detail?.sessionId; + + if (targetSessionId !== sessionId) { + return; + } + + loadForSession(targetSessionId); + }; + window.addEventListener(AppEvents.SESSION_EXTENSIONS_LOADED, loadExtensionsForCurrentSession); + // Load immediately in case no SESSION_EXTENSIONS_LOADED event fires for this session. + loadForSession(sessionId); + return () => { controller?.abort(); window.removeEventListener(