Fix model selector showing wrong model in tabs (#7784)

Co-authored-by: Claude Haiku 4.5 <noreply@anthropic.com>
This commit is contained in:
jh-block
2026-03-11 19:18:23 +01:00
committed by GitHub
parent b2e20c22b1
commit 6fed3e392c
10 changed files with 148 additions and 75 deletions
+6 -8
View File
@@ -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}
/>
</div>
+42 -5
View File
@@ -90,6 +90,9 @@ interface ChatInputProps {
append?: (message: Message) => void;
onWorkingDirChange?: (newDir: string) => void;
inputRef?: React.RefObject<HTMLTextAreaElement | null>;
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<HTMLDivElement>;
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<number>(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}
/>
</div>
</>
@@ -1520,6 +1553,10 @@ export default function ChatInput({
dropdownRef={dropdownRef}
setView={setView}
alerts={alerts}
sessionModel={effectiveModel}
sessionProvider={effectiveProvider}
onModelChanged={setModelOverride}
sessionLoaded={sessionLoaded}
/>
</div>
</Tooltip>
@@ -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<void>;
changeModel: (sessionId: string | null, model: Model) => Promise<boolean>;
getCurrentModelAndProvider: () => Promise<{ model: string; provider: string }>;
getFallbackModelAndProvider: () => Promise<{ model: string; provider: string }>;
getCurrentModelAndProviderForDisplay: () => Promise<{ model: string; provider: string }>;
getCurrentModelDisplayName: () => Promise<string>;
getCurrentProviderDisplayName: () => Promise<string>; // Gets provider display name from subtext
refreshCurrentModelAndProvider: () => Promise<void>;
setProviderAndModel: (provider: string, model: string) => void;
}
interface ModelAndProviderProviderProps {
@@ -47,7 +46,7 @@ export const ModelAndProviderProvider: React.FC<ModelAndProviderProviderProps> =
try {
if (sessionId) {
await updateAgentProvider({
const response = await updateAgentProvider({
body: {
session_id: sessionId,
provider: providerName,
@@ -56,24 +55,34 @@ export const ModelAndProviderProvider: React.FC<ModelAndProviderProviderProps> =
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<ModelAndProviderProviderProps> =
msg: `${error}`,
traceback: errorMessage(error),
});
return false;
}
}, []);
@@ -174,11 +184,6 @@ export const ModelAndProviderProvider: React.FC<ModelAndProviderProviderProps> =
}
}, [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<ModelAndProviderProviderProps> =
getCurrentModelDisplayName,
getCurrentProviderDisplayName,
refreshCurrentModelAndProvider,
setProviderAndModel,
}),
[
currentModel,
@@ -207,7 +211,6 @@ export const ModelAndProviderProvider: React.FC<ModelAndProviderProviderProps> =
getCurrentModelDisplayName,
getCurrentProviderDisplayName,
refreshCurrentModelAndProvider,
setProviderAndModel,
]
);
@@ -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<ModelInfoData | null>(null);
const [isLoading, setIsLoading] = useState(true);
const [showPricing, setShowPricing] = useState(true);
@@ -28,7 +28,7 @@ export const LocalInferenceSettings = () => {
const [downloads, setDownloads] = useState<Map<string, DownloadProgress>>(new Map());
const [showAllFeatured, setShowAllFeatured] = useState(false);
const [settingsOpenFor, setSettingsOpenFor] = useState<string | null>(null);
const { currentModel, currentProvider, setProviderAndModel } = useModelAndProvider();
const { currentModel, currentProvider, refreshCurrentModelAndProvider } = useModelAndProvider();
const downloadSectionRef = useRef<HTMLDivElement>(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);
}
@@ -24,6 +24,10 @@ interface ModelsBottomBarProps {
dropdownRef: React.RefObject<HTMLDivElement>;
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<string | null>(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 (
<div className="relative flex items-center" ref={dropdownRef}>
<BottomMenuAlertPopover alerts={alerts} />
@@ -146,7 +166,7 @@ export default function ModelsBottomBar({
<DropdownMenuTrigger className="flex items-center hover:cursor-pointer max-w-[180px] md:max-w-[200px] lg:max-w-[380px] min-w-0 text-text-primary/70 hover:text-text-primary transition-colors">
<div className="flex items-center truncate max-w-[130px] md:max-w-[200px] lg:max-w-[360px] min-w-0">
<Bot className="mr-1 h-4 w-4 flex-shrink-0" />
<span className="truncate text-xs">
<span className={`truncate text-xs${isModelLoading ? ' opacity-0' : ''}`}>
{displayModel}
{isLeadWorkerActive && modelMode && (
<span className="ml-1 text-[10px] opacity-60">({modelMode})</span>
@@ -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}
@@ -23,13 +23,6 @@ vi.mock('../../../ConfigContext', () => ({
}),
}));
// Minimal mock for useModelAndProvider
vi.mock('../../../ModelAndProviderContext', () => ({
useModelAndProvider: () => ({
currentModel: null,
}),
}));
describe('LeadWorkerSettings', () => {
beforeEach(() => {
vi.clearAllMocks();
@@ -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<string>('');
const [workerModel, setWorkerModel] = useState<string>('');
const [leadProvider, setLeadProvider] = useState<string>('');
@@ -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(() => {
@@ -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
+4 -2
View File
@@ -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<string | undefined>(undefined);
const prevProviderRef = useRef<string | undefined>(undefined);
// Handle model changes and accumulate costs
useEffect(() => {
if (!currentModel || !currentProvider) return;
const handleModelChange = async () => {
if (
prevModelRef.current !== undefined &&