diff --git a/apps/app/src/hooks/thread-creation-options/persisted-selection-fields.ts b/apps/app/src/hooks/thread-creation-options/persisted-selection-fields.ts index 7d30eeab8c..6b799ff38f 100644 --- a/apps/app/src/hooks/thread-creation-options/persisted-selection-fields.ts +++ b/apps/app/src/hooks/thread-creation-options/persisted-selection-fields.ts @@ -17,6 +17,7 @@ const PERMISSION_MODE_STORAGE_KEY = "bb.promptbox.permission-mode"; const ENVIRONMENT_STORAGE_KEY = "bb.promptbox.environment"; const PROVIDER_STORAGE_KEY = "bb.promptbox.provider"; const PROVIDER_SELECTION_STORAGE_VERSION = "1"; +const HOST_SELECTION_STORAGE_VERSION = "2"; export type StoredServiceTier = "" | ServiceTier; export type StoredReasoningLevel = "" | ReasoningLevel; @@ -86,15 +87,24 @@ function isStoredPermissionMode(value: string): value is StoredPermissionMode { return value === "" || isPermissionMode(value); } -const providerIdAtom = atomWithStorage( - PROVIDER_STORAGE_KEY, - "", - rawStringLocalStorage, - { getOnInit: true }, -); const emptyModelAtom = atom(""); const emptyReasoningLevelAtom = atom(""); +function normalizeHostId(hostId?: string | null): string | null { + const normalized = hostId?.trim() ?? ""; + return normalized.length > 0 ? normalized : null; +} + +function getHostSelectionStorageKey( + storageKey: string, + hostId: string, + providerId?: string, +): string { + const scope = + providerId === undefined ? [hostId.trim()] : [hostId.trim(), providerId]; + return `${storageKey}-${encodeURIComponent(JSON.stringify(scope))}-${HOST_SELECTION_STORAGE_VERSION}`; +} + function getProviderSelectionStorageKey( storageKey: string, providerId: string, @@ -105,30 +115,63 @@ function getProviderSelectionStorageKey( function getLegacyProviderSelection( providerId: string, storageKey: string, + includeProviderScopedValue = false, ): string | null { if (typeof window === "undefined") return null; + if (includeProviderScopedValue) { + const value = window.localStorage.getItem( + getProviderSelectionStorageKey(storageKey, providerId), + ); + if (value !== null) return value; + } if (window.localStorage.getItem(PROVIDER_STORAGE_KEY) !== providerId) { return null; } return window.localStorage.getItem(storageKey); } -function createProviderModelStorage(providerId: string) { +function createHostProviderStorage() { return createLocalStorageSyncStorage({ parse: (storedValue, initialValue) => storedValue ?? - getLegacyProviderSelection(providerId, MODEL_STORAGE_KEY) ?? + (typeof window === "undefined" + ? null + : window.localStorage.getItem(PROVIDER_STORAGE_KEY)) ?? initialValue, serialize: (value) => value, }); } -function createProviderReasoningStorage(providerId: string) { +function createProviderModelStorage( + providerId: string, + includeProviderScopedValue = false, +) { + return createLocalStorageSyncStorage({ + parse: (storedValue, initialValue) => + storedValue ?? + getLegacyProviderSelection( + providerId, + MODEL_STORAGE_KEY, + includeProviderScopedValue, + ) ?? + initialValue, + serialize: (value) => value, + }); +} + +function createProviderReasoningStorage( + providerId: string, + includeProviderScopedValue = false, +) { return createLocalStorageSyncStorage({ parse: (storedValue, initialValue) => { const value = storedValue ?? - getLegacyProviderSelection(providerId, REASONING_STORAGE_KEY); + getLegacyProviderSelection( + providerId, + REASONING_STORAGE_KEY, + includeProviderScopedValue, + ); return value !== null && isStoredReasoningLevel(value) ? value : initialValue; @@ -137,28 +180,57 @@ function createProviderReasoningStorage(providerId: string) { }); } -const modelAtomFamily = atomFamily((providerId: string) => +const providerIdAtomFamily = atomFamily((hostId: string | null) => atomWithStorage( - getProviderSelectionStorageKey(MODEL_STORAGE_KEY, providerId), + hostId + ? getHostSelectionStorageKey(PROVIDER_STORAGE_KEY, hostId) + : PROVIDER_STORAGE_KEY, "", - createProviderModelStorage(providerId), + hostId ? createHostProviderStorage() : rawStringLocalStorage, { getOnInit: true }, ), ); + +type HostProviderScope = readonly [hostId: string | null, providerId: string]; + +function isSameHostProviderScope( + left: HostProviderScope, + right: HostProviderScope, +): boolean { + return left[0] === right[0] && left[1] === right[1]; +} + +const modelAtomFamily = atomFamily( + ([hostId, providerId]: HostProviderScope) => + atomWithStorage( + hostId + ? getHostSelectionStorageKey(MODEL_STORAGE_KEY, hostId, providerId) + : getProviderSelectionStorageKey(MODEL_STORAGE_KEY, providerId), + "", + createProviderModelStorage(providerId, hostId !== null), + { getOnInit: true }, + ), + isSameHostProviderScope, +); + +const reasoningLevelAtomFamily = atomFamily( + ([hostId, providerId]: HostProviderScope) => + atomWithStorage( + hostId + ? getHostSelectionStorageKey(REASONING_STORAGE_KEY, hostId, providerId) + : getProviderSelectionStorageKey(REASONING_STORAGE_KEY, providerId), + "", + createProviderReasoningStorage(providerId, hostId !== null), + { getOnInit: true }, + ), + isSameHostProviderScope, +); const serviceTierAtom = atomWithStorage( SERVICE_TIER_STORAGE_KEY, "", createLocalStorageEnumStorage(isStoredServiceTier), { getOnInit: true }, ); -const reasoningLevelAtomFamily = atomFamily((providerId: string) => - atomWithStorage( - getProviderSelectionStorageKey(REASONING_STORAGE_KEY, providerId), - "", - createProviderReasoningStorage(providerId), - { getOnInit: true }, - ), -); // Legacy preference migration: "workspace-write" maps onto the same workspace // sandbox as "accept-edits", so the user's stored intent carries forward. // Legacy "readonly" (and any other unknown value) is dropped rather than @@ -198,11 +270,18 @@ const projectEnvironmentSelectionAtomFamily = atomFamily((projectId: string) => ), ); -export function usePromptBoxProviderPreference(): PersistedStringSelectionField { - const [value, setAtomValue] = useAtom(providerIdAtom); +export function usePromptBoxProviderPreference( + hostId?: string | null, +): PersistedStringSelectionField { + const normalizedHostId = normalizeHostId(hostId); + const [value, setAtomValue] = useAtom(providerIdAtomFamily(normalizedHostId)); const setValue = useCallback( (nextValue: string) => { - if (nextValue !== value && typeof window !== "undefined") { + if ( + normalizedHostId === null && + nextValue !== value && + typeof window !== "undefined" + ) { // Once the provider changes, the legacy unscoped values no longer have // a trustworthy owner. The caller saves the current pair under its // provider-scoped keys before changing this value. @@ -211,17 +290,19 @@ export function usePromptBoxProviderPreference(): PersistedStringSelectionField } setAtomValue(nextValue); }, - [setAtomValue, value], + [normalizedHostId, setAtomValue, value], ); return { setValue, value }; } export function usePromptBoxModelPreference( providerId: string, + hostId?: string | null, ): PersistedStringSelectionField { - const selectionAtom = providerId - ? modelAtomFamily(providerId) - : emptyModelAtom; + const normalizedHostId = normalizeHostId(hostId); + const selectionAtom = !providerId + ? emptyModelAtom + : modelAtomFamily([normalizedHostId, providerId]); const [value, setAtomValue] = useAtom(selectionAtom); const setValue = useCallback( (nextValue: string) => { @@ -245,10 +326,12 @@ export function usePromptBoxServiceTierPreference(): PersistedServiceTierSelecti export function usePromptBoxReasoningLevelPreference( providerId: string, + hostId?: string | null, ): PersistedReasoningLevelSelectionField { - const selectionAtom = providerId - ? reasoningLevelAtomFamily(providerId) - : emptyReasoningLevelAtom; + const normalizedHostId = normalizeHostId(hostId); + const selectionAtom = !providerId + ? emptyReasoningLevelAtom + : reasoningLevelAtomFamily([normalizedHostId, providerId]); const [value, setAtomValue] = useAtom(selectionAtom); const setValue = useCallback( (nextValue: StoredReasoningLevel) => { @@ -259,17 +342,19 @@ export function usePromptBoxReasoningLevelPreference( return { setValue, value }; } -export function useSetPromptBoxProviderModelReasoningPreference(): ( - preference: PromptBoxProviderModelReasoningPreference, -) => void { +export function useSetPromptBoxProviderModelReasoningPreference( + hostId?: string | null, +): (preference: PromptBoxProviderModelReasoningPreference) => void { const store = useStore(); + const normalizedHostId = normalizeHostId(hostId); return useCallback( ({ providerId, model, reasoningLevel }) => { if (providerId.length === 0) return; - store.set(modelAtomFamily(providerId), model); - store.set(reasoningLevelAtomFamily(providerId), reasoningLevel); + const scope: HostProviderScope = [normalizedHostId, providerId]; + store.set(modelAtomFamily(scope), model); + store.set(reasoningLevelAtomFamily(scope), reasoningLevel); }, - [store], + [normalizedHostId, store], ); } diff --git a/apps/app/src/hooks/useThreadCreationOptions.test.tsx b/apps/app/src/hooks/useThreadCreationOptions.test.tsx index 5f0c7dbf03..5a9ca64041 100644 --- a/apps/app/src/hooks/useThreadCreationOptions.test.tsx +++ b/apps/app/src/hooks/useThreadCreationOptions.test.tsx @@ -489,6 +489,105 @@ describe("useThreadCreationOptions", () => { }); }); + it("restores provider, model, and reasoning per machine", async () => { + setProjectScopedValue("bb.promptbox.environment", "host:local-host:local"); + vi.mocked(sdk.system.executionOptions).mockImplementation(async (args) => + providerExecutionOptionsResponse(args?.providerId), + ); + const mounted = renderHook( + () => + useThreadCreationOptions({ + scope: "new-thread", + preferenceProjectId: PROJECT_ID, + }), + { wrapper: createQueryClientTestHarness().wrapper }, + ); + const selectHost = (hostId: string) => { + act(() => { + mounted.result.current.setEnvironmentSelectionValue( + `host:${hostId}:local`, + ); + }); + }; + const expectSelection = async ( + providerId: string, + model: string, + reasoningLevel: string, + ) => { + await waitFor(() => { + expect(mounted.result.current.selectedProviderId).toBe(providerId); + expect(mounted.result.current.selectedModel).toBe(model); + expect(mounted.result.current.reasoningLevel).toBe(reasoningLevel); + }); + }; + const rememberSelection = async ( + hostId: string, + providerId: string, + model: string, + reasoningLevel: "low" | "medium" | "high", + ) => { + selectHost(hostId); + await waitFor(() => { + expect(mounted.result.current.executionOptionsRouting).toEqual({ + hostId, + }); + }); + act(() => { + mounted.result.current.setProviderModelReasoning({ + providerId, + model, + reasoningLevel, + }); + }); + }; + + await rememberSelection( + "local-host", + GLOBAL_PROVIDER_ID, + "global-remembered", + "high", + ); + await rememberSelection( + "vps-host", + PROJECT_PROVIDER_ID, + "project-remembered", + "medium", + ); + selectHost("local-host"); + await expectSelection(GLOBAL_PROVIDER_ID, "global-remembered", "high"); + + selectHost("vps-host"); + await expectSelection(PROJECT_PROVIDER_ID, "project-remembered", "medium"); + + await rememberSelection( + "vps-host", + GLOBAL_PROVIDER_ID, + "global-default", + "low", + ); + selectHost("local-host"); + await expectSelection(GLOBAL_PROVIDER_ID, "global-remembered", "high"); + selectHost("vps-host"); + await expectSelection(GLOBAL_PROVIDER_ID, "global-default", "low"); + + mounted.unmount(); + const reloaded = renderHook( + () => + useThreadCreationOptions({ + scope: "new-thread", + preferenceProjectId: PROJECT_ID, + }), + { wrapper: createQueryClientTestHarness().wrapper }, + ); + await waitFor(() => { + expect(reloaded.result.current.selectedProviderId).toBe( + GLOBAL_PROVIDER_ID, + ); + expect(reloaded.result.current.selectedModel).toBe("global-default"); + expect(reloaded.result.current.reasoningLevel).toBe("low"); + }); + }); + it("keeps provider selections local in component-local composers", async () => { vi.mocked(sdk.system.executionOptions).mockImplementation(async (args) => providerExecutionOptionsResponse(args?.providerId), diff --git a/apps/app/src/hooks/useThreadCreationOptions.ts b/apps/app/src/hooks/useThreadCreationOptions.ts index 34b6afcfcd..ef6ca9a2de 100644 --- a/apps/app/src/hooks/useThreadCreationOptions.ts +++ b/apps/app/src/hooks/useThreadCreationOptions.ts @@ -202,10 +202,6 @@ export function useThreadCreationOptions( resetKey, scope = "new-thread", } = options ?? {}; - const { setValue: setStoredProviderId, value: storedProviderId } = - usePromptBoxProviderPreference(); - const setStoredProviderModelReasoning = - useSetPromptBoxProviderModelReasoningPreference(); const { setValue: setStoredServiceTier, value: storedServiceTier } = usePromptBoxServiceTierPreference(); const { setValue: setStoredPermissionMode, value: storedPermissionMode } = @@ -283,9 +279,6 @@ export function useThreadCreationOptions( usesLocalThreadSelections, ]); - const selectedProviderIdBeforeConnectedFallback = usesStoredCreateSelections - ? storedProviderId || renderedThreadSelections.selectedProviderId - : renderedThreadSelections.selectedProviderId; const rawServiceTier = usesStoredCreateSelections ? storedServiceTier || renderedThreadSelections.serviceTier : renderedThreadSelections.serviceTier; @@ -306,9 +299,19 @@ export function useThreadCreationOptions( environmentId, environmentHostId, environmentSelectionValue: rawEnvironmentSelectionValue, - providerId: selectedProviderIdBeforeConnectedFallback, + providerId: renderedThreadSelections.selectedProviderId, scope, }); + const preferenceHostId = usesStoredCreateSelections + ? executionOptionsRouting.hostId + : undefined; + const { setValue: setStoredProviderId, value: storedProviderId } = + usePromptBoxProviderPreference(preferenceHostId); + const setStoredProviderModelReasoning = + useSetPromptBoxProviderModelReasoningPreference(preferenceHostId); + const selectedProviderIdBeforeConnectedFallback = usesStoredCreateSelections + ? storedProviderId || renderedThreadSelections.selectedProviderId + : renderedThreadSelections.selectedProviderId; const shouldResolveConnectedProvider = executionOptionsQueryEnabled && scope === "new-thread" && @@ -371,9 +374,9 @@ export function useThreadCreationOptions( }, [providers, rawSelectedProviderId]); const { setValue: setStoredSelectedModel, value: storedSelectedModel } = - usePromptBoxModelPreference(effectiveProviderId); + usePromptBoxModelPreference(effectiveProviderId, preferenceHostId); const { setValue: setStoredReasoningLevel, value: storedReasoningLevel } = - usePromptBoxReasoningLevelPreference(effectiveProviderId); + usePromptBoxReasoningLevelPreference(effectiveProviderId, preferenceHostId); const effectiveProviderMatchesInitialProvider = effectiveProviderId.length > 0 && effectiveProviderId === renderedThreadSelections.selectedProviderId; @@ -840,6 +843,9 @@ export function useThreadCreationOptions( touchedThreadFieldsRef.current.add("selectedModel"); if (usesStoredCreateSelections) { setStoredSelectedModel(value); + if (effectiveProviderId.length > 0) { + setStoredProviderId(effectiveProviderId); + } return; } setLocalProvidersUsingDefaults((current) => { @@ -861,6 +867,7 @@ export function useThreadCreationOptions( [ effectiveProviderId, reasoningLevel, + setStoredProviderId, setStoredSelectedModel, usesStoredCreateSelections, ], @@ -887,6 +894,9 @@ export function useThreadCreationOptions( touchedThreadFieldsRef.current.add("reasoningLevel"); if (usesStoredCreateSelections) { setStoredReasoningLevel(value); + if (effectiveProviderId.length > 0) { + setStoredProviderId(effectiveProviderId); + } return; } setLocalProvidersUsingDefaults((current) => { @@ -910,6 +920,7 @@ export function useThreadCreationOptions( [ effectiveProviderId, selectedModel, + setStoredProviderId, setStoredReasoningLevel, usesStoredCreateSelections, ],