From 6fed3e392ce6a1b8c52b8260032fe65fcbbd9da6 Mon Sep 17 00:00:00 2001 From: jh-block Date: Wed, 11 Mar 2026 19:18:23 +0100 Subject: [PATCH] Fix model selector showing wrong model in tabs (#7784) Co-authored-by: Claude Haiku 4.5 --- ui/desktop/src/components/BaseChat.tsx | 14 ++--- ui/desktop/src/components/ChatInput.tsx | 47 ++++++++++++-- .../components/ModelAndProviderContext.tsx | 43 +++++++------ .../components/bottom_menu/CostTracker.tsx | 6 +- .../localInference/LocalInferenceSettings.tsx | 4 +- .../models/bottom_bar/ModelsBottomBar.tsx | 27 +++++++- .../subcomponents/LeadWorkerSettings.test.tsx | 7 --- .../subcomponents/LeadWorkerSettings.tsx | 8 +-- .../models/subcomponents/SwitchModelModal.tsx | 61 +++++++++++++------ ui/desktop/src/hooks/useCostTracking.ts | 6 +- 10 files changed, 148 insertions(+), 75 deletions(-) diff --git a/ui/desktop/src/components/BaseChat.tsx b/ui/desktop/src/components/BaseChat.tsx index cee981fb..b88c9cc7 100644 --- a/ui/desktop/src/components/BaseChat.tsx +++ b/ui/desktop/src/components/BaseChat.tsx @@ -35,7 +35,6 @@ import { useToolCount } from './alerts/useToolCount'; import { getThinkingMessage, getTextAndImageContent } from '../types/message'; import ParameterInputModal from './ParameterInputModal'; import { substituteParameters } from '../utils/parameterSubstitution'; -import { useModelAndProvider } from './ModelAndProviderContext'; import CreateRecipeFromSessionModal from './recipes/CreateRecipeFromSessionModal'; import { toastSuccess } from '../toasts'; import { Recipe } from '../recipe'; @@ -182,13 +181,9 @@ export default function BaseChat({ session, }); - const { setProviderAndModel } = useModelAndProvider(); - - useEffect(() => { - if (session?.provider_name && session?.model_config?.model_name) { - setProviderAndModel(session.provider_name, session.model_config.model_name); - } - }, [session?.provider_name, session?.model_config?.model_name, setProviderAndModel]); + const sessionModel = session?.model_config?.model_name ?? null; + const sessionProvider = session?.provider_name ?? null; + const sessionLoaded = session !== undefined; useEffect(() => { if (!recipe) return; @@ -502,6 +497,9 @@ export default function BaseChat({ recipeAccepted={!hasNotAcceptedRecipe} initialPrompt={initialPrompt} toolCount={toolCount || 0} + sessionModel={sessionModel} + sessionProvider={sessionProvider} + sessionLoaded={sessionLoaded} {...customChatInputProps} /> diff --git a/ui/desktop/src/components/ChatInput.tsx b/ui/desktop/src/components/ChatInput.tsx index e68f8bc6..c01c1e0e 100644 --- a/ui/desktop/src/components/ChatInput.tsx +++ b/ui/desktop/src/components/ChatInput.tsx @@ -90,6 +90,9 @@ interface ChatInputProps { append?: (message: Message) => void; onWorkingDirChange?: (newDir: string) => void; inputRef?: React.RefObject; + sessionModel?: string | null; + sessionProvider?: string | null; + sessionLoaded?: boolean; } export default function ChatInput({ @@ -117,6 +120,9 @@ export default function ChatInput({ append: _append, onWorkingDirChange, inputRef, + sessionModel, + sessionProvider, + sessionLoaded, }: ChatInputProps) { const [_value, setValue] = useState(initialValue); const [displayValue, setDisplayValue] = useState(initialValue); // For immediate visual feedback @@ -139,7 +145,24 @@ export default function ChatInput({ null ) as React.RefObject; const { getProviders } = useConfig(); - const { getCurrentModelAndProvider, currentModel, currentProvider } = useModelAndProvider(); + const { getCurrentModelAndProvider, currentModel: configModel, currentProvider: configProvider } = useModelAndProvider(); + + // Local override for when the user changes the model in the modal, + // before the session object is re-fetched from the backend. + const [modelOverride, setModelOverride] = useState<{ model: string; provider: string } | null>(null); + const effectiveModel = modelOverride?.model ?? sessionModel ?? configModel; + const effectiveProvider = modelOverride?.provider ?? sessionProvider ?? configProvider; + + // Clear override when the underlying data catches up (session props for + // active chats, config defaults for Hub / no-session contexts). + useEffect(() => { + if (!modelOverride) return; + const sessionCaughtUp = sessionModel === modelOverride.model && sessionProvider === modelOverride.provider; + const configCaughtUp = !sessionId && configModel === modelOverride.model && configProvider === modelOverride.provider; + if (sessionCaughtUp || configCaughtUp) { + setModelOverride(null); + } + }, [sessionModel, sessionProvider, configModel, configProvider, sessionId, modelOverride]); const [tokenLimit, setTokenLimit] = useState(TOKEN_LIMIT_DEFAULT); const [isTokenLimitLoaded, setIsTokenLimitLoaded] = useState(false); const [diagnosticsOpen, setDiagnosticsOpen] = useState(false); @@ -369,8 +392,15 @@ export default function ChatInput({ // Reset token limit loaded state setIsTokenLimitLoaded(false); - // Get current model and provider first to avoid unnecessary provider fetches - const { model, provider } = await getCurrentModelAndProvider(); + // Use effective model/provider (includes overrides from in-session model changes), + // fall back to config defaults + let model = effectiveModel; + let provider = effectiveProvider; + if (!model || !provider) { + const configModelAndProvider = await getCurrentModelAndProvider(); + model = configModelAndProvider.model; + provider = configModelAndProvider.provider; + } if (!model || !provider) { console.log('No model or provider found'); setIsTokenLimitLoaded(true); @@ -417,11 +447,12 @@ export default function ChatInput({ } }; - // Initial load and refresh when model changes + // Initial load and refresh when model changes (effective model includes overrides, + // config model is the fallback for Hub/no-session contexts) useEffect(() => { loadProviderDetails(); // eslint-disable-next-line react-hooks/exhaustive-deps - }, [currentModel, currentProvider]); + }, [effectiveModel, effectiveProvider, configModel, configProvider]); // Handle tool count alerts and token usage useEffect(() => { @@ -1509,6 +1540,8 @@ export default function ChatInput({ inputTokens={accumulatedInputTokens} outputTokens={accumulatedOutputTokens} sessionCosts={sessionCosts} + model={effectiveModel} + provider={effectiveProvider} /> @@ -1520,6 +1553,10 @@ export default function ChatInput({ dropdownRef={dropdownRef} setView={setView} alerts={alerts} + sessionModel={effectiveModel} + sessionProvider={effectiveProvider} + onModelChanged={setModelOverride} + sessionLoaded={sessionLoaded} /> diff --git a/ui/desktop/src/components/ModelAndProviderContext.tsx b/ui/desktop/src/components/ModelAndProviderContext.tsx index 9b5e2fab..6a851696 100644 --- a/ui/desktop/src/components/ModelAndProviderContext.tsx +++ b/ui/desktop/src/components/ModelAndProviderContext.tsx @@ -19,14 +19,13 @@ const SWITCH_MODEL_SUCCESS_MSG = 'Successfully switched models'; interface ModelAndProviderContextType { currentModel: string | null; currentProvider: string | null; - changeModel: (sessionId: string | null, model: Model) => Promise; + changeModel: (sessionId: string | null, model: Model) => Promise; getCurrentModelAndProvider: () => Promise<{ model: string; provider: string }>; getFallbackModelAndProvider: () => Promise<{ model: string; provider: string }>; getCurrentModelAndProviderForDisplay: () => Promise<{ model: string; provider: string }>; getCurrentModelDisplayName: () => Promise; getCurrentProviderDisplayName: () => Promise; // Gets provider display name from subtext refreshCurrentModelAndProvider: () => Promise; - setProviderAndModel: (provider: string, model: string) => void; } interface ModelAndProviderProviderProps { @@ -47,7 +46,7 @@ export const ModelAndProviderProvider: React.FC = try { if (sessionId) { - await updateAgentProvider({ + const response = await updateAgentProvider({ body: { session_id: sessionId, provider: providerName, @@ -56,24 +55,34 @@ export const ModelAndProviderProvider: React.FC = request_params: model.request_params, }, }); + if (response.error) { + throw new Error(`Failed to update agent provider: ${response.error}`); + } } - phase = 'config'; - await setConfigProvider({ - body: { - provider: providerName, - model: modelName, - }, - throwOnError: true, - }); + // Only update the global config default when there's no session + // (i.e. changing from settings, not from within an existing chat) + if (!sessionId) { + phase = 'config'; + await setConfigProvider({ + body: { + provider: providerName, + model: modelName, + }, + throwOnError: true, + }); + } - setCurrentProvider(providerName); - setCurrentModel(modelName); + if (!sessionId) { + setCurrentProvider(providerName); + setCurrentModel(modelName); + } toastSuccess({ title: CHANGE_MODEL_TOAST_TITLE, msg: `${SWITCH_MODEL_SUCCESS_MSG} -- using ${model.alias ?? modelName} from ${model.subtext ?? providerName}`, }); + return true; } catch (error) { console.error(`Failed to change model at ${phase} step -- ${modelName} ${providerName}`); toastError({ @@ -81,6 +90,7 @@ export const ModelAndProviderProvider: React.FC = msg: `${error}`, traceback: errorMessage(error), }); + return false; } }, []); @@ -174,11 +184,6 @@ export const ModelAndProviderProvider: React.FC = } }, [getCurrentModelAndProvider]); - const setProviderAndModel = useCallback((provider: string, model: string) => { - setCurrentProvider(provider); - setCurrentModel(model); - }, []); - // Load initial model and provider on mount useEffect(() => { refreshCurrentModelAndProvider(); @@ -195,7 +200,6 @@ export const ModelAndProviderProvider: React.FC = getCurrentModelDisplayName, getCurrentProviderDisplayName, refreshCurrentModelAndProvider, - setProviderAndModel, }), [ currentModel, @@ -207,7 +211,6 @@ export const ModelAndProviderProvider: React.FC = getCurrentModelDisplayName, getCurrentProviderDisplayName, refreshCurrentModelAndProvider, - setProviderAndModel, ] ); diff --git a/ui/desktop/src/components/bottom_menu/CostTracker.tsx b/ui/desktop/src/components/bottom_menu/CostTracker.tsx index acb039ba..990fddd2 100644 --- a/ui/desktop/src/components/bottom_menu/CostTracker.tsx +++ b/ui/desktop/src/components/bottom_menu/CostTracker.tsx @@ -1,5 +1,4 @@ import { useState, useEffect } from 'react'; -import { useModelAndProvider } from '../ModelAndProviderContext'; import { CoinIcon } from '../icons'; import { Tooltip, TooltipContent, TooltipTrigger } from '../ui/Tooltip'; import { fetchCanonicalModelInfo } from '../../utils/canonical'; @@ -15,10 +14,11 @@ interface CostTrackerProps { totalCost: number; }; }; + model: string | null; + provider: string | null; } -export function CostTracker({ inputTokens = 0, outputTokens = 0, sessionCosts }: CostTrackerProps) { - const { currentModel, currentProvider } = useModelAndProvider(); +export function CostTracker({ inputTokens = 0, outputTokens = 0, sessionCosts, model: currentModel, provider: currentProvider }: CostTrackerProps) { const [costInfo, setCostInfo] = useState(null); const [isLoading, setIsLoading] = useState(true); const [showPricing, setShowPricing] = useState(true); diff --git a/ui/desktop/src/components/settings/localInference/LocalInferenceSettings.tsx b/ui/desktop/src/components/settings/localInference/LocalInferenceSettings.tsx index a7dc7e79..b2191d32 100644 --- a/ui/desktop/src/components/settings/localInference/LocalInferenceSettings.tsx +++ b/ui/desktop/src/components/settings/localInference/LocalInferenceSettings.tsx @@ -28,7 +28,7 @@ export const LocalInferenceSettings = () => { const [downloads, setDownloads] = useState>(new Map()); const [showAllFeatured, setShowAllFeatured] = useState(false); const [settingsOpenFor, setSettingsOpenFor] = useState(null); - const { currentModel, currentProvider, setProviderAndModel } = useModelAndProvider(); + const { currentModel, currentProvider, refreshCurrentModelAndProvider } = useModelAndProvider(); const downloadSectionRef = useRef(null); const selectedModelId = currentProvider === 'local' ? currentModel : null; @@ -67,12 +67,12 @@ export const LocalInferenceSettings = () => { }, [models]); const selectModel = async (modelId: string) => { - setProviderAndModel('local', modelId); try { await setConfigProvider({ body: { provider: 'local', model: modelId }, throwOnError: true, }); + await refreshCurrentModelAndProvider(); } catch (error) { console.error('Failed to select model:', error); } diff --git a/ui/desktop/src/components/settings/models/bottom_bar/ModelsBottomBar.tsx b/ui/desktop/src/components/settings/models/bottom_bar/ModelsBottomBar.tsx index 2de9a95e..e3e1a51b 100644 --- a/ui/desktop/src/components/settings/models/bottom_bar/ModelsBottomBar.tsx +++ b/ui/desktop/src/components/settings/models/bottom_bar/ModelsBottomBar.tsx @@ -24,6 +24,10 @@ interface ModelsBottomBarProps { dropdownRef: React.RefObject; setView: (view: View) => void; alerts: Alert[]; + sessionModel?: string | null; + sessionProvider?: string | null; + onModelChanged: (override: { model: string; provider: string }) => void; + sessionLoaded?: boolean; } export default function ModelsBottomBar({ @@ -31,8 +35,17 @@ export default function ModelsBottomBar({ dropdownRef, setView, alerts, + sessionModel, + sessionProvider, + onModelChanged, + sessionLoaded, }: ModelsBottomBarProps) { - const { currentModel, currentProvider } = useModelAndProvider(); + // ChatInput owns the override state and passes effective model/provider as sessionModel/sessionProvider. + // Fall back to config defaults when no session-specific model is available. + const { currentModel: configModel, currentProvider: configProvider } = useModelAndProvider(); + const currentModel = sessionModel ?? configModel; + const currentProvider = sessionProvider ?? configProvider; + const currentModelInfo = useCurrentModelInfo(); const { read, getProviders } = useConfig(); const [displayProvider, setDisplayProvider] = useState(null); @@ -101,6 +114,9 @@ export default function ModelsBottomBar({ : undefined; // Determine which model to display - activeModel takes priority when lead/worker is active + // Hide label while session data is still being fetched (avoids flashing + // the config default before the session's actual model arrives). + const isModelLoading = sessionId && !sessionLoaded; const displayModel = isLeadWorkerActive && currentModelInfo?.model ? currentModelInfo.model @@ -139,6 +155,10 @@ export default function ModelsBottomBar({ setDisplayModelName(getModelDisplayName(currentModel)); }, [currentModel]); + const handleModelSelected = (model: string, provider: string) => { + onModelChanged({ model, provider }); + }; + return (
@@ -146,7 +166,7 @@ export default function ModelsBottomBar({
- + {displayModel} {isLeadWorkerActive && modelMode && ( ({modelMode}) @@ -182,6 +202,9 @@ export default function ModelsBottomBar({ sessionId={sessionId} setView={setView} onClose={() => setIsAddModelModalOpen(false)} + sessionModel={currentModel} + sessionProvider={currentProvider} + onModelSelected={(model, provider) => handleModelSelected(model, provider)} /> ) : null} diff --git a/ui/desktop/src/components/settings/models/subcomponents/LeadWorkerSettings.test.tsx b/ui/desktop/src/components/settings/models/subcomponents/LeadWorkerSettings.test.tsx index d4cea641..5dcebda3 100644 --- a/ui/desktop/src/components/settings/models/subcomponents/LeadWorkerSettings.test.tsx +++ b/ui/desktop/src/components/settings/models/subcomponents/LeadWorkerSettings.test.tsx @@ -23,13 +23,6 @@ vi.mock('../../../ConfigContext', () => ({ }), })); -// Minimal mock for useModelAndProvider -vi.mock('../../../ModelAndProviderContext', () => ({ - useModelAndProvider: () => ({ - currentModel: null, - }), -})); - describe('LeadWorkerSettings', () => { beforeEach(() => { vi.clearAllMocks(); diff --git a/ui/desktop/src/components/settings/models/subcomponents/LeadWorkerSettings.tsx b/ui/desktop/src/components/settings/models/subcomponents/LeadWorkerSettings.tsx index 2fea7c46..8acd5648 100644 --- a/ui/desktop/src/components/settings/models/subcomponents/LeadWorkerSettings.tsx +++ b/ui/desktop/src/components/settings/models/subcomponents/LeadWorkerSettings.tsx @@ -1,6 +1,5 @@ import { useState, useEffect } from 'react'; import { useConfig } from '../../../ConfigContext'; -import { useModelAndProvider } from '../../../ModelAndProviderContext'; import { Button } from '../../../ui/button'; import { Select } from '../../../ui/Select'; import { Input } from '../../../ui/input'; @@ -15,7 +14,6 @@ interface LeadWorkerSettingsProps { export function LeadWorkerSettings({ isOpen, onClose }: LeadWorkerSettingsProps) { const { read, upsert, getProviders, remove } = useConfig(); - const { currentModel } = useModelAndProvider(); const [leadModel, setLeadModel] = useState(''); const [workerModel, setWorkerModel] = useState(''); const [leadProvider, setLeadProvider] = useState(''); @@ -69,12 +67,10 @@ export function LeadWorkerSettings({ isOpen, onClose }: LeadWorkerSettingsProps) if (fallbackTurnsConfig) setFallbackTurns(Number(fallbackTurnsConfig)); else setFallbackTurns(2); - // Set worker model to current model or from config + // Set worker model from config const workerModelConfig = await read('GOOSE_MODEL', false); if (workerModelConfig) { setWorkerModel(workerModelConfig as string); - } else if (currentModel) { - setWorkerModel(currentModel as string); } else { setWorkerModel(''); } @@ -140,7 +136,7 @@ export function LeadWorkerSettings({ isOpen, onClose }: LeadWorkerSettingsProps) }; loadConfig(); - }, [read, getProviders, currentModel, isOpen]); + }, [read, getProviders, isOpen]); // If current models are not in the list (e.g., previously set to custom), switch to custom mode useEffect(() => { diff --git a/ui/desktop/src/components/settings/models/subcomponents/SwitchModelModal.tsx b/ui/desktop/src/components/settings/models/subcomponents/SwitchModelModal.tsx index 3bd5bc17..05785e5e 100644 --- a/ui/desktop/src/components/settings/models/subcomponents/SwitchModelModal.tsx +++ b/ui/desktop/src/components/settings/models/subcomponents/SwitchModelModal.tsx @@ -1,4 +1,4 @@ -import { useEffect, useState, useCallback } from 'react'; +import { useEffect, useState, useCallback, useRef } from 'react'; import { Bot, ExternalLink } from 'lucide-react'; import { @@ -84,9 +84,11 @@ type SwitchModelModalProps = { sessionId: string | null; onClose: () => void; setView: (view: View) => void; - onModelSelected?: (model: string) => void; + onModelSelected?: (model: string, provider: string) => void; initialProvider?: string | null; titleOverride?: string; + sessionModel?: string | null; + sessionProvider?: string | null; }; export const SwitchModelModal = ({ sessionId, @@ -95,9 +97,14 @@ export const SwitchModelModal = ({ onModelSelected, initialProvider, titleOverride, + sessionModel, + sessionProvider, }: SwitchModelModalProps) => { const { getProviders, read, upsert } = useConfig(); - const { changeModel, currentModel, currentProvider } = useModelAndProvider(); + const { changeModel, currentModel: configModel, currentProvider: configProvider } = useModelAndProvider(); + // Use session-specific model/provider if available, otherwise fall back to config defaults + const currentModel = sessionModel ?? configModel; + const currentProvider = sessionProvider ?? configProvider; const [providerOptions, setProviderOptions] = useState<{ value: string; label: string }[]>([]); type ModelOption = { value: string; label: string; provider: string; isDisabled?: boolean }; const [modelOptions, setModelOptions] = useState<{ options: ModelOption[] }[]>([]); @@ -242,10 +249,11 @@ export const SwitchModelModal = ({ } } - await changeModel(sessionId, modelObj); - onModelSelected?.(modelObj.name); - - trackModelChanged(modelObj.provider || '', modelObj.name); + const success = await changeModel(sessionId, modelObj); + if (success) { + onModelSelected?.(modelObj.name, modelObj.provider || ''); + trackModelChanged(modelObj.provider || '', modelObj.name); + } onClose(); } @@ -258,24 +266,37 @@ export const SwitchModelModal = ({ } }, [attemptedSubmit, validateForm]); + // Initialize predefined model selection from session/config model. + // Separate effect so it re-runs when currentModel loads asynchronously. + useEffect(() => { + if (!usePredefinedModels || !currentModel) return; + const models = getPredefinedModelsFromEnv(); + const matchingModel = models.find((m) => m.name === currentModel); + if (matchingModel) { + setSelectedPredefinedModel(matchingModel); + } + }, [usePredefinedModels, currentModel]); + + // For manual mode: one-time sync of provider/model when session data + // arrives after the modal has already mounted. Uses a ref so it only + // fires once and doesn't interfere with user-driven changes (e.g. + // switching provider clears model intentionally). + const manualSyncDone = useRef(false); + useEffect(() => { + if (usePredefinedModels || manualSyncDone.current) return; + if (initialProvider && initialProvider !== currentProvider) return; + if (currentModel && currentProvider) { + if (!provider) setProvider(currentProvider); + if (!model) setModel(currentModel); + manualSyncDone.current = true; + } + }, [currentModel, currentProvider, usePredefinedModels, provider, model, initialProvider]); + useEffect(() => { // Load predefined models if enabled if (usePredefinedModels) { const models = getPredefinedModelsFromEnv(); setPredefinedModels(models); - - // Initialize selected predefined model with current model - (async () => { - try { - const currentModelName = (await read('GOOSE_MODEL', false)) as string; - const matchingModel = models.find((model) => model.name === currentModelName); - if (matchingModel) { - setSelectedPredefinedModel(matchingModel); - } - } catch (error) { - console.error('Failed to get current model for selection:', error); - } - })(); } // Load providers for manual model selection diff --git a/ui/desktop/src/hooks/useCostTracking.ts b/ui/desktop/src/hooks/useCostTracking.ts index 11773adb..7db12dc0 100644 --- a/ui/desktop/src/hooks/useCostTracking.ts +++ b/ui/desktop/src/hooks/useCostTracking.ts @@ -1,5 +1,4 @@ import { useEffect, useRef, useState } from 'react'; -import { useModelAndProvider } from '../components/ModelAndProviderContext'; import { fetchCanonicalModelInfo } from '../utils/canonical'; import { Session } from '../api'; @@ -26,12 +25,15 @@ export const useCostTracking = ({ }; }>({}); - const { currentModel, currentProvider } = useModelAndProvider(); + const currentModel = session?.model_config?.model_name ?? undefined; + const currentProvider = session?.provider_name ?? undefined; const prevModelRef = useRef(undefined); const prevProviderRef = useRef(undefined); // Handle model changes and accumulate costs useEffect(() => { + if (!currentModel || !currentProvider) return; + const handleModelChange = async () => { if ( prevModelRef.current !== undefined &&