Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion Cargo.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

95 changes: 81 additions & 14 deletions crates/goose-server/src/routes/agent.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand All @@ -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 {
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -751,6 +780,8 @@ async fn agent_remove_extension(
State(state): State<Arc<AppState>>,
Json(request): Json<RemoveExtensionRequest>,
) -> Result<StatusCode, ErrorResponse> {
ensure_extensions_loaded(&state, &request.session_id).await?;

let agent = state.get_agent(request.session_id.clone()).await?;

agent
Expand Down Expand Up @@ -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(())
}
}
}

Expand All @@ -1009,7 +1074,9 @@ async fn read_resource(
) -> Result<Json<ReadResourceResponse>, 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())
Expand Down Expand Up @@ -1091,7 +1158,7 @@ async fn call_tool(
State(state): State<Arc<AppState>>,
Json(payload): Json<CallToolRequest>,
) -> Result<Json<CallToolResponse>, 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())
Expand Down
22 changes: 17 additions & 5 deletions crates/goose-server/src/state.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<Vec<ExtensionLoadResult>> {
) -> Result<Option<Vec<ExtensionLoadResult>>, 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) {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand All @@ -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);
Comment thread
angiejones marked this conversation as resolved.
Comment thread
angiejones marked this conversation as resolved.

return () => {
controller?.abort();
window.removeEventListener(
Expand Down