fix(desktop): keep extensions icon visible after switching sessions (#9787)
Signed-off-by: Seydi Charyyev <seydi.charyev@gmail.com> Signed-off-by: Angie Jones <jones.angie@gmail.com> Signed-off-by: Douwe M Osinga <douwe@sidewalklabs.com> Co-authored-by: Angie Jones <jones.angie@gmail.com> Co-authored-by: Douwe M Osinga <douwe@sidewalklabs.com>
This commit is contained in:
Generated
+1
-1
@@ -1169,7 +1169,7 @@ dependencies = [
|
||||
"quote",
|
||||
"regex",
|
||||
"rustc-hash 1.1.0",
|
||||
"shlex",
|
||||
"shlex 1.3.0",
|
||||
"syn 2.0.117",
|
||||
]
|
||||
|
||||
|
||||
@@ -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<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
|
||||
@@ -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<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())
|
||||
@@ -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())
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user