diff --git a/crates/goose-sdk-types/src/custom_requests.rs b/crates/goose-sdk-types/src/custom_requests.rs index df2ec19e0..36948b94b 100644 --- a/crates/goose-sdk-types/src/custom_requests.rs +++ b/crates/goose-sdk-types/src/custom_requests.rs @@ -42,7 +42,7 @@ pub struct AddSessionExtensionRequest { #[serde(rename_all = "camelCase")] pub struct RemoveSessionExtensionRequest { pub session_id: String, - pub name: String, + pub extension_key: String, } /// List all tools available in a session. @@ -405,6 +405,13 @@ pub struct GooseExtensionEntry { pub config_key: Option, } +#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema)] +#[serde(rename_all = "camelCase")] +pub struct SessionExtensionEntry { + pub extension: GooseExtension, + pub extension_key: String, +} + /// List Goose-owned extension definitions available to configure or enable. #[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcRequest)] #[request( @@ -477,7 +484,7 @@ pub struct GetSessionExtensionsRequest { #[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcResponse)] pub struct GetSessionExtensionsResponse { - pub extensions: Vec, + pub extensions: Vec, } /// Read allowlisted user preferences. Empty `keys` means all supported preferences. diff --git a/crates/goose/acp-schema.json b/crates/goose/acp-schema.json index 36b641e65..fcf4f282d 100644 --- a/crates/goose/acp-schema.json +++ b/crates/goose/acp-schema.json @@ -422,13 +422,13 @@ "sessionId": { "type": "string" }, - "name": { + "extensionKey": { "type": "string" } }, "required": [ "sessionId", - "name" + "extensionKey" ], "description": "Remove an extension from an active session.", "x-side": "agent", @@ -1603,7 +1603,7 @@ "extensions": { "type": "array", "items": { - "$ref": "#/$defs/GooseExtension" + "$ref": "#/$defs/SessionExtensionEntry" } } }, @@ -1613,6 +1613,21 @@ "x-side": "agent", "x-method": "_goose/unstable/session/extensions/list" }, + "SessionExtensionEntry": { + "type": "object", + "properties": { + "extension": { + "$ref": "#/$defs/GooseExtension" + }, + "extensionKey": { + "type": "string" + } + }, + "required": [ + "extension", + "extensionKey" + ] + }, "ListProvidersRequest_unstable": { "type": "object", "properties": { diff --git a/crates/goose/src/acp/server/extensions.rs b/crates/goose/src/acp/server/extensions.rs index 4f39f6a58..a6da99cf1 100644 --- a/crates/goose/src/acp/server/extensions.rs +++ b/crates/goose/src/acp/server/extensions.rs @@ -2,6 +2,7 @@ use super::*; use crate::agents::extension::Envs; use crate::config::extensions::ExtensionEntry; use agent_client_protocol::schema::v1::{HttpHeader, McpServer, McpServerHttp, McpServerStdio}; +use std::collections::HashSet; impl GooseAcpAgent { pub(super) async fn on_add_session_extension( @@ -24,10 +25,14 @@ impl GooseAcpAgent { ) -> Result { let session_id = &req.session_id; let agent = self.get_session_agent(&req.session_id).await?; - agent - .remove_extension(&req.name, session_id) + let removed = agent + .remove_extension_by_key(&req.extension_key, session_id) .await .internal_err()?; + if !removed { + return Err(agent_client_protocol::Error::invalid_params() + .data(format!("Extension '{}' not found", req.extension_key))); + } Ok(EmptyResponse {}) } @@ -123,18 +128,33 @@ impl GooseAcpAgent { crate::config::Config::global(), ); - let extensions = extensions - .into_iter() - .map(|config| config_to_goose_extension(&config)) - .collect::, _>>()? - .into_iter() - .flatten() - .collect::>(); - - Ok(GetSessionExtensionsResponse { extensions }) + Ok(GetSessionExtensionsResponse { + extensions: session_configs_to_entries(extensions)?, + }) } } +fn session_configs_to_entries( + configs: Vec, +) -> Result, agent_client_protocol::Error> { + let mut extension_keys = HashSet::with_capacity(configs.len()); + let mut entries = Vec::with_capacity(configs.len()); + for config in configs { + let extension_key = config.key(); + if !extension_keys.insert(extension_key.clone()) { + return Err(agent_client_protocol::Error::internal_error() + .data(format!("Duplicate session extension key '{extension_key}'"))); + } + if let Some(extension) = config_to_goose_extension(&config)? { + entries.push(SessionExtensionEntry { + extension, + extension_key, + }); + } + } + Ok(entries) +} + fn config_to_goose_extension( config: &ExtensionConfig, ) -> Result, agent_client_protocol::Error> { @@ -406,6 +426,34 @@ mod tests { use agent_client_protocol::schema::v1::{McpServer, McpServerSse}; use std::collections::HashMap; + fn builtin_config(name: &str) -> ExtensionConfig { + ExtensionConfig::Builtin { + name: name.to_string(), + description: String::new(), + display_name: None, + timeout: None, + bundled: None, + available_tools: Vec::new(), + } + } + + #[test] + fn session_entries_preserve_backend_distinct_unicode_keys() { + let entries = + session_configs_to_entries(vec![builtin_config("\u{130}"), builtin_config("i\u{307}")]) + .expect("backend-distinct keys should be listed"); + + assert_eq!(entries[0].extension_key, "_"); + assert_eq!(entries[1].extension_key, "i_"); + } + + #[test] + fn session_entries_reject_duplicate_authoritative_keys() { + let result = session_configs_to_entries(vec![builtin_config("a.b"), builtin_config("a/b")]); + + assert!(result.is_err()); + } + #[test] fn builtin_config_converts_to_goose_builtin_extension() { let config = ExtensionConfig::Builtin { diff --git a/crates/goose/src/agents/agent.rs b/crates/goose/src/agents/agent.rs index 9406aff02..a366beab6 100644 --- a/crates/goose/src/agents/agent.rs +++ b/crates/goose/src/agents/agent.rs @@ -391,6 +391,19 @@ impl Default for Agent { } } +fn has_unique_persisted_extension(configs: &[ExtensionConfig], key: &str) -> Result { + match configs + .iter() + .filter(|config| config.key() == key) + .take(2) + .count() + { + 0 => Ok(false), + 1 => Ok(true), + _ => Err(anyhow!("Duplicate session extension key '{key}'")), + } +} + impl Agent { pub fn new() -> Self { let config = Config::global(); @@ -1126,10 +1139,13 @@ impl Agent { self.rebuild_frontend_derived_state(&extensions).await; } - async fn remove_frontend_extension(&self, name: &str) { + async fn remove_frontend_extension_by_key(&self, key: &str) -> bool { let mut extensions = self.frontend_extensions.lock().await; - extensions.remove(&name_to_key(name)); - self.rebuild_frontend_derived_state(&extensions).await; + let removed = extensions.remove(key).is_some(); + if removed { + self.rebuild_frontend_derived_state(&extensions).await; + } + removed } async fn extension_configs_for_persistence(&self) -> Vec { @@ -1593,8 +1609,27 @@ impl Agent { } pub async fn remove_extension(&self, name: &str, session_id: &str) -> Result<()> { - self.extension_manager.remove_extension(name).await?; - self.remove_frontend_extension(name).await; + self.remove_extension_by_key(&name_to_key(name), session_id) + .await?; + Ok(()) + } + + pub async fn remove_extension_by_key(&self, key: &str, session_id: &str) -> Result { + let session = self + .config + .session_manager + .get_session(session_id, false) + .await?; + let persisted_extensions = EnabledExtensionsState::extensions_or_default( + Some(&session.extension_data), + Config::global(), + ); + if !has_unique_persisted_extension(&persisted_extensions, key)? { + return Ok(false); + } + + self.extension_manager.remove_extension_by_key(key).await?; + self.remove_frontend_extension_by_key(key).await; // Persist extension state after successful removal self.persist_extension_state(session_id) @@ -1604,7 +1639,7 @@ impl Agent { anyhow!("Failed to persist extension state: {}", e) })?; - Ok(()) + Ok(true) } pub async fn list_extensions(&self) -> Vec { @@ -4156,6 +4191,33 @@ mod tests { use std::sync::atomic::{AtomicUsize, Ordering}; use tempfile::TempDir; + fn persisted_builtin(name: &str) -> ExtensionConfig { + ExtensionConfig::Builtin { + name: name.to_string(), + description: String::new(), + display_name: None, + timeout: None, + bundled: None, + available_tools: Vec::new(), + } + } + + #[test] + fn persisted_extension_identity_must_be_unique_before_removal() { + let session_extension = persisted_builtin("session-only"); + assert!(has_unique_persisted_extension(&[session_extension], "session-only").unwrap()); + assert!(!has_unique_persisted_extension(&[], "missing").unwrap()); + + let duplicate_result = has_unique_persisted_extension( + &[persisted_builtin("a.b"), persisted_builtin("a/b")], + "a_b", + ); + assert_eq!( + duplicate_result.unwrap_err().to_string(), + "Duplicate session extension key 'a_b'" + ); + } + #[test] fn provider_creation_context_preserves_acp_error_code() { let source = anyhow::Error::new(agent_client_protocol::Error::auth_required()) diff --git a/crates/goose/src/agents/extension_manager.rs b/crates/goose/src/agents/extension_manager.rs index c9c6e5c28..4a53bc721 100644 --- a/crates/goose/src/agents/extension_manager.rs +++ b/crates/goose/src/agents/extension_manager.rs @@ -1761,11 +1761,18 @@ impl ExtensionManager { /// Get aggregated usage statistics pub async fn remove_extension(&self, name: &str) -> ExtensionResult<()> { let sanitized_name = name_to_key(name); - self.extensions.lock().await.remove(&sanitized_name); - self.invalidate_tools_cache_and_bump_version().await; + self.remove_extension_by_key(&sanitized_name).await?; Ok(()) } + pub async fn remove_extension_by_key(&self, key: &str) -> ExtensionResult { + let removed = self.extensions.lock().await.remove(key).is_some(); + if removed { + self.invalidate_tools_cache_and_bump_version().await; + } + Ok(removed) + } + pub async fn update_working_dir(&self, new_dir: &std::path::Path) { let extensions = self.extensions.lock().await; for (name, ext) in extensions.iter() { diff --git a/crates/goose/tests/acp_custom_requests_test.rs b/crates/goose/tests/acp_custom_requests_test.rs index 446f088b7..e6f637cc0 100644 --- a/crates/goose/tests/acp_custom_requests_test.rs +++ b/crates/goose/tests/acp_custom_requests_test.rs @@ -343,7 +343,7 @@ fn test_custom_session_extensions_add_list_remove() { .expect("extensions should be an array"); extensions .iter() - .find(|extension| extension["name"] == extension_name) + .find(|entry| entry["extension"]["name"] == extension_name) .cloned() }; @@ -369,18 +369,39 @@ fn test_custom_session_extensions_add_list_remove() { .await; assert!(add_result.is_ok(), "expected ok, got: {:?}", add_result); - let extension = list_extension() + let entry = list_extension() .await .unwrap_or_else(|| panic!("missing added session extension")); + assert_eq!(entry["extensionKey"], extension_name); + let extension = &entry["extension"]; assert_eq!(extension["type"], "platform"); assert_eq!(extension["name"], extension_name); + let unknown_remove_result = send_custom( + conn.cx(), + "_goose/unstable/session/extensions/remove", + serde_json::json!({ + "sessionId": session_id.clone(), + "extensionKey": "missing-extension", + }), + ) + .await; + assert_eq!( + unknown_remove_result.unwrap_err(), + agent_client_protocol::Error::invalid_params() + .data("Extension 'missing-extension' not found") + ); + assert!( + list_extension().await.is_some(), + "unknown key must not remove another extension" + ); + let remove_result = send_custom( conn.cx(), "_goose/unstable/session/extensions/remove", serde_json::json!({ "sessionId": session_id.clone(), - "name": extension_name, + "extensionKey": extension_name, }), ) .await; diff --git a/ui/desktop/src/acp/__tests__/session-extensions.test.ts b/ui/desktop/src/acp/__tests__/session-extensions.test.ts new file mode 100644 index 000000000..52ed7d71c --- /dev/null +++ b/ui/desktop/src/acp/__tests__/session-extensions.test.ts @@ -0,0 +1,54 @@ +import { beforeEach, describe, expect, it, vi } from 'vitest'; +import { getAcpClient } from '../acpConnection'; +import { getSessionExtensions, removeSessionExtension } from '../session-extensions'; + +vi.mock('../acpConnection', () => ({ + getAcpClient: vi.fn(), +})); + +const extension = (name: string, extensionKey: string) => ({ + extension: { + type: 'builtin' as const, + name, + }, + extensionKey, +}); + +describe('ACP session extensions', () => { + const list = vi.fn(); + const remove = vi.fn(); + + beforeEach(() => { + vi.clearAllMocks(); + vi.mocked(getAcpClient).mockResolvedValue({ + goose: { + sessionExtensionsList_unstable: list, + sessionExtensionsRemove_unstable: remove, + }, + } as unknown as Awaited>); + }); + + it('preserves the backend identity on mapped session entries', async () => { + list.mockResolvedValue({ extensions: [extension('i\u0307', 'i_')] }); + + await expect(getSessionExtensions('session')).resolves.toEqual([ + expect.objectContaining({ name: 'i\u0307', extensionKey: 'i_' }), + ]); + }); + + it('rejects duplicate authoritative identities', async () => { + list.mockResolvedValue({ + extensions: [extension('first', 'duplicate'), extension('second', 'duplicate')], + }); + + await expect(getSessionExtensions('session')).rejects.toThrow( + "Duplicate session extension key 'duplicate'" + ); + }); + + it('removes by the backend identity', async () => { + await removeSessionExtension('session', 'i_'); + + expect(remove).toHaveBeenCalledWith({ sessionId: 'session', extensionKey: 'i_' }); + }); +}); diff --git a/ui/desktop/src/acp/session-extensions.ts b/ui/desktop/src/acp/session-extensions.ts index 9a894e422..1b9f24d66 100644 --- a/ui/desktop/src/acp/session-extensions.ts +++ b/ui/desktop/src/acp/session-extensions.ts @@ -2,12 +2,27 @@ import type { ExtensionConfig } from '../types/extensions'; import { getAcpClient } from './acpConnection'; import { extensionConfigToGooseExtension, gooseExtensionToExtensionConfig } from './extensions'; -export async function getSessionExtensions(sessionId: string): Promise { +export type SessionExtension = ExtensionConfig & { extensionKey: string }; + +export async function getSessionExtensions(sessionId: string): Promise { const client = await getAcpClient(); const response = await client.goose.sessionExtensionsList_unstable({ sessionId }); - return response.extensions - .map(gooseExtensionToExtensionConfig) - .filter((config): config is ExtensionConfig => config !== null); + const extensionKeys = new Set(); + const extensions: SessionExtension[] = []; + + for (const entry of response.extensions) { + if (extensionKeys.has(entry.extensionKey)) { + throw new Error(`Duplicate session extension key '${entry.extensionKey}'`); + } + extensionKeys.add(entry.extensionKey); + + const config = gooseExtensionToExtensionConfig(entry.extension); + if (config) { + extensions.push({ ...config, extensionKey: entry.extensionKey }); + } + } + + return extensions; } export async function addSessionExtension( @@ -22,7 +37,10 @@ export async function addSessionExtension( await client.goose.sessionExtensionsAdd_unstable({ sessionId, extension }); } -export async function removeSessionExtension(sessionId: string, name: string): Promise { +export async function removeSessionExtension( + sessionId: string, + extensionKey: string +): Promise { const client = await getAcpClient(); - await client.goose.sessionExtensionsRemove_unstable({ sessionId, name }); + await client.goose.sessionExtensionsRemove_unstable({ sessionId, extensionKey }); } diff --git a/ui/desktop/src/components/ConfigContext.tsx b/ui/desktop/src/components/ConfigContext.tsx index 4cb6236b1..2cac7111f 100644 --- a/ui/desktop/src/components/ConfigContext.tsx +++ b/ui/desktop/src/components/ConfigContext.tsx @@ -18,6 +18,7 @@ export type { ExtensionConfig } from '../types/extensions'; export type FixedExtensionEntry = ExtensionConfig & { enabled: boolean; configKey?: string; + extensionKey?: string; }; type ConfigMap = Record; diff --git a/ui/desktop/src/components/bottom_menu/BottomMenuExtensionSelection.test.tsx b/ui/desktop/src/components/bottom_menu/BottomMenuExtensionSelection.test.tsx new file mode 100644 index 000000000..ff211a529 --- /dev/null +++ b/ui/desktop/src/components/bottom_menu/BottomMenuExtensionSelection.test.tsx @@ -0,0 +1,156 @@ +import { fireEvent, render, screen, waitFor } from '@testing-library/react'; +import { beforeEach, describe, expect, it, vi } from 'vitest'; +import { IntlTestWrapper } from '../../i18n/test-utils'; +import type { FixedExtensionEntry } from '../ConfigContext'; +import { BottomMenuExtensionSelection } from './BottomMenuExtensionSelection'; + +const mocks = vi.hoisted(() => ({ + addToAgent: vi.fn(), + configuredExtensions: [] as FixedExtensionEntry[], + getSessionExtensions: vi.fn(), + removeFromAgent: vi.fn(), +})); + +vi.mock('../ConfigContext', () => ({ + useConfig: () => ({ extensionsList: mocks.configuredExtensions }), +})); + +vi.mock('../../acp/session-extensions', () => ({ + getSessionExtensions: mocks.getSessionExtensions, +})); + +vi.mock('../settings/extensions/agent-api', () => ({ + addToAgent: mocks.addToAgent, + removeFromAgent: mocks.removeFromAgent, +})); + +vi.mock('./ExtensionMenu', () => ({ + ExtensionMenu: ({ + extensions, + hidden, + onToggle, + }: { + extensions: Array; + hidden: boolean; + onToggle: (extension: FixedExtensionEntry & { extensionKey?: string }) => void; + }) => ( +
+ {String(hidden)} + + {JSON.stringify( + extensions.map(({ name, enabled, extensionKey }) => ({ + name, + enabled, + extensionKey, + })) + )} + + {extensions.map((extension) => ( + + ))} +
+ ), +})); + +const configuredExtension = (name: string, configKey: string): FixedExtensionEntry => ({ + type: 'builtin', + name, + description: `${name} configured extension`, + enabled: false, + configKey, +}); + +const sessionExtension = (name: string, extensionKey: string) => ({ + type: 'stdio' as const, + name, + description: `${name} session extension`, + cmd: 'session-extension', + args: [], + extensionKey, +}); + +describe('BottomMenuExtensionSelection session identities', () => { + beforeEach(() => { + vi.clearAllMocks(); + mocks.configuredExtensions = []; + mocks.getSessionExtensions.mockResolvedValue([]); + }); + + it('keeps backend-distinct Unicode identities separate and removes by session key', async () => { + mocks.configuredExtensions = [configuredExtension('\u0130', '_')]; + mocks.getSessionExtensions.mockResolvedValue([sessionExtension('i\u0307', 'i_')]); + + render(, { + wrapper: IntlTestWrapper, + }); + + await waitFor(() => + expect(screen.getByTestId('extension-identities')).toHaveTextContent( + JSON.stringify([ + { name: '\u0130', enabled: false }, + { name: 'i\u0307', enabled: true, extensionKey: 'i_' }, + ]) + ) + ); + + fireEvent.click(screen.getByRole('button', { name: 'i\u0307' })); + + await waitFor(() => + expect(mocks.removeFromAgent).toHaveBeenCalledWith('i_', 'i\u0307', 'victim-session', true) + ); + }); + + it('merges a legitimate configured and session entry by authoritative key', async () => { + mocks.configuredExtensions = [configuredExtension('developer', 'developer')]; + mocks.getSessionExtensions.mockResolvedValue([sessionExtension('developer', 'developer')]); + + render(, { + wrapper: IntlTestWrapper, + }); + + await waitFor(() => + expect(screen.getByTestId('extension-identities')).toHaveTextContent( + JSON.stringify([{ name: 'developer', enabled: true, extensionKey: 'developer' }]) + ) + ); + }); + + it('keeps an empty authoritative key visible and removes by that exact key', async () => { + mocks.configuredExtensions = [configuredExtension('empty-key', '')]; + mocks.getSessionExtensions.mockResolvedValue([sessionExtension('empty-key', '')]); + + render(, { + wrapper: IntlTestWrapper, + }); + + await waitFor(() => expect(screen.getByTestId('hidden')).toHaveTextContent('false')); + expect(screen.getByTestId('extension-identities')).toHaveTextContent( + JSON.stringify([{ name: 'empty-key', enabled: true, extensionKey: '' }]) + ); + + fireEvent.click(screen.getByRole('button', { name: 'empty-key' })); + + await waitFor(() => + expect(mocks.removeFromAgent).toHaveBeenCalledWith('', 'empty-key', 'session', true) + ); + }); + + it('hides controls when configured entries repeat an authoritative key', async () => { + mocks.configuredExtensions = [ + configuredExtension('first', 'duplicate'), + configuredExtension('second', 'duplicate'), + ]; + + render(, { + wrapper: IntlTestWrapper, + }); + + await waitFor(() => expect(screen.getByTestId('hidden')).toHaveTextContent('true')); + expect(screen.getByTestId('extension-identities')).toHaveTextContent('[]'); + }); +}); diff --git a/ui/desktop/src/components/bottom_menu/BottomMenuExtensionSelection.tsx b/ui/desktop/src/components/bottom_menu/BottomMenuExtensionSelection.tsx index e6836cc06..b2b9892e7 100644 --- a/ui/desktop/src/components/bottom_menu/BottomMenuExtensionSelection.tsx +++ b/ui/desktop/src/components/bottom_menu/BottomMenuExtensionSelection.tsx @@ -2,9 +2,10 @@ import { useCallback, useEffect, useMemo, useRef, useState } from 'react'; import { useConfig, type FixedExtensionEntry } from '../ConfigContext'; import { toastService } from '../../toasts'; import { formatExtensionName } from '../settings/extensions/subcomponents/ExtensionList'; -import { nameToKey } from '../settings/extensions/utils'; -import type { ExtensionConfig } from '../../types/extensions'; -import { getSessionExtensions as getAcpSessionExtensions } from '../../acp/session-extensions'; +import { + getSessionExtensions as getAcpSessionExtensions, + type SessionExtension, +} from '../../acp/session-extensions'; import { addToAgent, removeFromAgent } from '../settings/extensions/agent-api'; import { defineMessages, useIntl } from '../../i18n'; import { AppEvents } from '../../constants/events'; @@ -64,6 +65,43 @@ type GetSessionExtensionsSignal = { aborted: boolean }; const EXTENSION_SORT_DELAY_MS = 800; +function mergeSessionExtensions( + configuredExtensions: FixedExtensionEntry[], + sessionExtensions: SessionExtension[] +): FixedExtensionEntry[] | null { + const sessionExtensionsByKey = new Map(); + for (const extension of sessionExtensions) { + if (sessionExtensionsByKey.has(extension.extensionKey)) { + return null; + } + sessionExtensionsByKey.set(extension.extensionKey, extension); + } + + const configuredExtensionKeys = new Set(); + const mergedExtensions: FixedExtensionEntry[] = []; + for (const extension of configuredExtensions) { + if (extension.configKey === undefined || configuredExtensionKeys.has(extension.configKey)) { + return null; + } + configuredExtensionKeys.add(extension.configKey); + + const sessionExtension = sessionExtensionsByKey.get(extension.configKey); + mergedExtensions.push({ + ...extension, + enabled: sessionExtension !== undefined, + extensionKey: sessionExtension?.extensionKey, + }); + } + + for (const sessionExtension of sessionExtensions) { + if (!configuredExtensionKeys.has(sessionExtension.extensionKey)) { + mergedExtensions.push({ ...sessionExtension, enabled: true }); + } + } + + return mergedExtensions; +} + function useExtensionMenuTransition() { const [isTransitioning, setIsTransitioning] = useState(false); const [isSortPending, setIsSortPending] = useState(false); @@ -239,7 +277,7 @@ function DraftExtensionsMenu({ function SessionExtensionsMenu({ sessionId }: { sessionId: string }) { const intl = useIntl(); - const [sessionExtensions, setSessionExtensions] = useState([]); + const [sessionExtensions, setSessionExtensions] = useState([]); const [isSessionExtensionsLoaded, setIsSessionExtensionsLoaded] = useState(false); const latestSessionIdRef = useRef(sessionId); const { extensionsList: allExtensions } = useConfig(); @@ -261,7 +299,7 @@ function SessionExtensionsMenu({ sessionId }: { sessionId: string }) { const loadSessionExtensions = useCallback( async (targetSessionId: string, signal?: GetSessionExtensionsSignal) => { - const extensions = await getAcpSessionExtensions(targetSessionId) + const extensions = await getAcpSessionExtensions(targetSessionId); if (signal?.aborted || latestSessionIdRef.current !== targetSessionId) { return; @@ -287,7 +325,7 @@ function SessionExtensionsMenu({ sessionId }: { sessionId: string }) { } console.error('Failed to fetch session extensions:', error); - setIsSessionExtensionsLoaded(true); + setIsSessionExtensionsLoaded(false); }); }; @@ -321,7 +359,15 @@ function SessionExtensionsMenu({ sessionId }: { sessionId: string }) { try { if (extensionConfig.enabled) { - await removeFromAgent(extensionConfig.name, sessionId, true); + if (extensionConfig.extensionKey === undefined) { + throw new Error('Missing session extension key'); + } + await removeFromAgent( + extensionConfig.extensionKey, + extensionConfig.name, + sessionId, + true + ); } else { await addToAgent(extensionConfig, sessionId, true); } @@ -344,45 +390,20 @@ function SessionExtensionsMenu({ sessionId }: { sessionId: string }) { [beginToggle, loadSessionExtensions, resetTransition, scheduleSort, sessionId] ); - const extensions = useMemo(() => { - const sessionExtensionKeys = new Set( - sessionExtensions.map((extension) => nameToKey(extension.name)) - ); - const configuredExtensionKeys = new Set( - allExtensions.map((extension) => nameToKey(extension.name)) - ); - - const mergedExtensions = allExtensions.map( - (extension) => - ({ - ...extension, - enabled: sessionExtensionKeys.has(nameToKey(extension.name)), - }) as FixedExtensionEntry - ); - - for (const sessionExtension of sessionExtensions) { - if (configuredExtensionKeys.has(nameToKey(sessionExtension.name))) { - continue; - } - - mergedExtensions.push({ - ...sessionExtension, - enabled: true, - }); - } - - return mergedExtensions; - }, [allExtensions, sessionExtensions]); + const extensions = useMemo( + () => mergeSessionExtensions(allExtensions, sessionExtensions), + [allExtensions, sessionExtensions] + ); return (