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 { getThinkingMessage, getTextAndImageContent } from '../types/message';
import ParameterInputModal from './ParameterInputModal'; import ParameterInputModal from './ParameterInputModal';
import { substituteParameters } from '../utils/parameterSubstitution'; import { substituteParameters } from '../utils/parameterSubstitution';
import { useModelAndProvider } from './ModelAndProviderContext';
import CreateRecipeFromSessionModal from './recipes/CreateRecipeFromSessionModal'; import CreateRecipeFromSessionModal from './recipes/CreateRecipeFromSessionModal';
import { toastSuccess } from '../toasts'; import { toastSuccess } from '../toasts';
import { Recipe } from '../recipe'; import { Recipe } from '../recipe';
@@ -182,13 +181,9 @@ export default function BaseChat({
session, session,
}); });
const { setProviderAndModel } = useModelAndProvider(); const sessionModel = session?.model_config?.model_name ?? null;
const sessionProvider = session?.provider_name ?? null;
useEffect(() => { const sessionLoaded = session !== undefined;
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]);
useEffect(() => { useEffect(() => {
if (!recipe) return; if (!recipe) return;
@@ -502,6 +497,9 @@ export default function BaseChat({
recipeAccepted={!hasNotAcceptedRecipe} recipeAccepted={!hasNotAcceptedRecipe}
initialPrompt={initialPrompt} initialPrompt={initialPrompt}
toolCount={toolCount || 0} toolCount={toolCount || 0}
sessionModel={sessionModel}
sessionProvider={sessionProvider}
sessionLoaded={sessionLoaded}
{...customChatInputProps} {...customChatInputProps}
/> />
</div> </div>
+42 -5
View File
@@ -90,6 +90,9 @@ interface ChatInputProps {
append?: (message: Message) => void; append?: (message: Message) => void;
onWorkingDirChange?: (newDir: string) => void; onWorkingDirChange?: (newDir: string) => void;
inputRef?: React.RefObject<HTMLTextAreaElement | null>; inputRef?: React.RefObject<HTMLTextAreaElement | null>;
sessionModel?: string | null;
sessionProvider?: string | null;
sessionLoaded?: boolean;
} }
export default function ChatInput({ export default function ChatInput({
@@ -117,6 +120,9 @@ export default function ChatInput({
append: _append, append: _append,
onWorkingDirChange, onWorkingDirChange,
inputRef, inputRef,
sessionModel,
sessionProvider,
sessionLoaded,
}: ChatInputProps) { }: ChatInputProps) {
const [_value, setValue] = useState(initialValue); const [_value, setValue] = useState(initialValue);
const [displayValue, setDisplayValue] = useState(initialValue); // For immediate visual feedback const [displayValue, setDisplayValue] = useState(initialValue); // For immediate visual feedback
@@ -139,7 +145,24 @@ export default function ChatInput({
null null
) as React.RefObject<HTMLDivElement>; ) as React.RefObject<HTMLDivElement>;
const { getProviders } = useConfig(); 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 [tokenLimit, setTokenLimit] = useState<number>(TOKEN_LIMIT_DEFAULT);
const [isTokenLimitLoaded, setIsTokenLimitLoaded] = useState(false); const [isTokenLimitLoaded, setIsTokenLimitLoaded] = useState(false);
const [diagnosticsOpen, setDiagnosticsOpen] = useState(false); const [diagnosticsOpen, setDiagnosticsOpen] = useState(false);
@@ -369,8 +392,15 @@ export default function ChatInput({
// Reset token limit loaded state // Reset token limit loaded state
setIsTokenLimitLoaded(false); setIsTokenLimitLoaded(false);
// Get current model and provider first to avoid unnecessary provider fetches // Use effective model/provider (includes overrides from in-session model changes),
const { model, provider } = await getCurrentModelAndProvider(); // 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) { if (!model || !provider) {
console.log('No model or provider found'); console.log('No model or provider found');
setIsTokenLimitLoaded(true); 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(() => { useEffect(() => {
loadProviderDetails(); loadProviderDetails();
// eslint-disable-next-line react-hooks/exhaustive-deps // eslint-disable-next-line react-hooks/exhaustive-deps
}, [currentModel, currentProvider]); }, [effectiveModel, effectiveProvider, configModel, configProvider]);
// Handle tool count alerts and token usage // Handle tool count alerts and token usage
useEffect(() => { useEffect(() => {
@@ -1509,6 +1540,8 @@ export default function ChatInput({
inputTokens={accumulatedInputTokens} inputTokens={accumulatedInputTokens}
outputTokens={accumulatedOutputTokens} outputTokens={accumulatedOutputTokens}
sessionCosts={sessionCosts} sessionCosts={sessionCosts}
model={effectiveModel}
provider={effectiveProvider}
/> />
</div> </div>
</> </>
@@ -1520,6 +1553,10 @@ export default function ChatInput({
dropdownRef={dropdownRef} dropdownRef={dropdownRef}
setView={setView} setView={setView}
alerts={alerts} alerts={alerts}
sessionModel={effectiveModel}
sessionProvider={effectiveProvider}
onModelChanged={setModelOverride}
sessionLoaded={sessionLoaded}
/> />
</div> </div>
</Tooltip> </Tooltip>
@@ -19,14 +19,13 @@ const SWITCH_MODEL_SUCCESS_MSG = 'Successfully switched models';
interface ModelAndProviderContextType { interface ModelAndProviderContextType {
currentModel: string | null; currentModel: string | null;
currentProvider: 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 }>; getCurrentModelAndProvider: () => Promise<{ model: string; provider: string }>;
getFallbackModelAndProvider: () => Promise<{ model: string; provider: string }>; getFallbackModelAndProvider: () => Promise<{ model: string; provider: string }>;
getCurrentModelAndProviderForDisplay: () => Promise<{ model: string; provider: string }>; getCurrentModelAndProviderForDisplay: () => Promise<{ model: string; provider: string }>;
getCurrentModelDisplayName: () => Promise<string>; getCurrentModelDisplayName: () => Promise<string>;
getCurrentProviderDisplayName: () => Promise<string>; // Gets provider display name from subtext getCurrentProviderDisplayName: () => Promise<string>; // Gets provider display name from subtext
refreshCurrentModelAndProvider: () => Promise<void>; refreshCurrentModelAndProvider: () => Promise<void>;
setProviderAndModel: (provider: string, model: string) => void;
} }
interface ModelAndProviderProviderProps { interface ModelAndProviderProviderProps {
@@ -47,7 +46,7 @@ export const ModelAndProviderProvider: React.FC<ModelAndProviderProviderProps> =
try { try {
if (sessionId) { if (sessionId) {
await updateAgentProvider({ const response = await updateAgentProvider({
body: { body: {
session_id: sessionId, session_id: sessionId,
provider: providerName, provider: providerName,
@@ -56,24 +55,34 @@ export const ModelAndProviderProvider: React.FC<ModelAndProviderProviderProps> =
request_params: model.request_params, request_params: model.request_params,
}, },
}); });
if (response.error) {
throw new Error(`Failed to update agent provider: ${response.error}`);
}
} }
phase = 'config'; // Only update the global config default when there's no session
await setConfigProvider({ // (i.e. changing from settings, not from within an existing chat)
body: { if (!sessionId) {
provider: providerName, phase = 'config';
model: modelName, await setConfigProvider({
}, body: {
throwOnError: true, provider: providerName,
}); model: modelName,
},
throwOnError: true,
});
}
setCurrentProvider(providerName); if (!sessionId) {
setCurrentModel(modelName); setCurrentProvider(providerName);
setCurrentModel(modelName);
}
toastSuccess({ toastSuccess({
title: CHANGE_MODEL_TOAST_TITLE, title: CHANGE_MODEL_TOAST_TITLE,
msg: `${SWITCH_MODEL_SUCCESS_MSG} -- using ${model.alias ?? modelName} from ${model.subtext ?? providerName}`, msg: `${SWITCH_MODEL_SUCCESS_MSG} -- using ${model.alias ?? modelName} from ${model.subtext ?? providerName}`,
}); });
return true;
} catch (error) { } catch (error) {
console.error(`Failed to change model at ${phase} step -- ${modelName} ${providerName}`); console.error(`Failed to change model at ${phase} step -- ${modelName} ${providerName}`);
toastError({ toastError({
@@ -81,6 +90,7 @@ export const ModelAndProviderProvider: React.FC<ModelAndProviderProviderProps> =
msg: `${error}`, msg: `${error}`,
traceback: errorMessage(error), traceback: errorMessage(error),
}); });
return false;
} }
}, []); }, []);
@@ -174,11 +184,6 @@ export const ModelAndProviderProvider: React.FC<ModelAndProviderProviderProps> =
} }
}, [getCurrentModelAndProvider]); }, [getCurrentModelAndProvider]);
const setProviderAndModel = useCallback((provider: string, model: string) => {
setCurrentProvider(provider);
setCurrentModel(model);
}, []);
// Load initial model and provider on mount // Load initial model and provider on mount
useEffect(() => { useEffect(() => {
refreshCurrentModelAndProvider(); refreshCurrentModelAndProvider();
@@ -195,7 +200,6 @@ export const ModelAndProviderProvider: React.FC<ModelAndProviderProviderProps> =
getCurrentModelDisplayName, getCurrentModelDisplayName,
getCurrentProviderDisplayName, getCurrentProviderDisplayName,
refreshCurrentModelAndProvider, refreshCurrentModelAndProvider,
setProviderAndModel,
}), }),
[ [
currentModel, currentModel,
@@ -207,7 +211,6 @@ export const ModelAndProviderProvider: React.FC<ModelAndProviderProviderProps> =
getCurrentModelDisplayName, getCurrentModelDisplayName,
getCurrentProviderDisplayName, getCurrentProviderDisplayName,
refreshCurrentModelAndProvider, refreshCurrentModelAndProvider,
setProviderAndModel,
] ]
); );
@@ -1,5 +1,4 @@
import { useState, useEffect } from 'react'; import { useState, useEffect } from 'react';
import { useModelAndProvider } from '../ModelAndProviderContext';
import { CoinIcon } from '../icons'; import { CoinIcon } from '../icons';
import { Tooltip, TooltipContent, TooltipTrigger } from '../ui/Tooltip'; import { Tooltip, TooltipContent, TooltipTrigger } from '../ui/Tooltip';
import { fetchCanonicalModelInfo } from '../../utils/canonical'; import { fetchCanonicalModelInfo } from '../../utils/canonical';
@@ -15,10 +14,11 @@ interface CostTrackerProps {
totalCost: number; totalCost: number;
}; };
}; };
model: string | null;
provider: string | null;
} }
export function CostTracker({ inputTokens = 0, outputTokens = 0, sessionCosts }: CostTrackerProps) { export function CostTracker({ inputTokens = 0, outputTokens = 0, sessionCosts, model: currentModel, provider: currentProvider }: CostTrackerProps) {
const { currentModel, currentProvider } = useModelAndProvider();
const [costInfo, setCostInfo] = useState<ModelInfoData | null>(null); const [costInfo, setCostInfo] = useState<ModelInfoData | null>(null);
const [isLoading, setIsLoading] = useState(true); const [isLoading, setIsLoading] = useState(true);
const [showPricing, setShowPricing] = useState(true); const [showPricing, setShowPricing] = useState(true);
@@ -28,7 +28,7 @@ export const LocalInferenceSettings = () => {
const [downloads, setDownloads] = useState<Map<string, DownloadProgress>>(new Map()); const [downloads, setDownloads] = useState<Map<string, DownloadProgress>>(new Map());
const [showAllFeatured, setShowAllFeatured] = useState(false); const [showAllFeatured, setShowAllFeatured] = useState(false);
const [settingsOpenFor, setSettingsOpenFor] = useState<string | null>(null); const [settingsOpenFor, setSettingsOpenFor] = useState<string | null>(null);
const { currentModel, currentProvider, setProviderAndModel } = useModelAndProvider(); const { currentModel, currentProvider, refreshCurrentModelAndProvider } = useModelAndProvider();
const downloadSectionRef = useRef<HTMLDivElement>(null); const downloadSectionRef = useRef<HTMLDivElement>(null);
const selectedModelId = currentProvider === 'local' ? currentModel : null; const selectedModelId = currentProvider === 'local' ? currentModel : null;
@@ -67,12 +67,12 @@ export const LocalInferenceSettings = () => {
}, [models]); }, [models]);
const selectModel = async (modelId: string) => { const selectModel = async (modelId: string) => {
setProviderAndModel('local', modelId);
try { try {
await setConfigProvider({ await setConfigProvider({
body: { provider: 'local', model: modelId }, body: { provider: 'local', model: modelId },
throwOnError: true, throwOnError: true,
}); });
await refreshCurrentModelAndProvider();
} catch (error) { } catch (error) {
console.error('Failed to select model:', error); console.error('Failed to select model:', error);
} }
@@ -24,6 +24,10 @@ interface ModelsBottomBarProps {
dropdownRef: React.RefObject<HTMLDivElement>; dropdownRef: React.RefObject<HTMLDivElement>;
setView: (view: View) => void; setView: (view: View) => void;
alerts: Alert[]; alerts: Alert[];
sessionModel?: string | null;
sessionProvider?: string | null;
onModelChanged: (override: { model: string; provider: string }) => void;
sessionLoaded?: boolean;
} }
export default function ModelsBottomBar({ export default function ModelsBottomBar({
@@ -31,8 +35,17 @@ export default function ModelsBottomBar({
dropdownRef, dropdownRef,
setView, setView,
alerts, alerts,
sessionModel,
sessionProvider,
onModelChanged,
sessionLoaded,
}: ModelsBottomBarProps) { }: 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 currentModelInfo = useCurrentModelInfo();
const { read, getProviders } = useConfig(); const { read, getProviders } = useConfig();
const [displayProvider, setDisplayProvider] = useState<string | null>(null); const [displayProvider, setDisplayProvider] = useState<string | null>(null);
@@ -101,6 +114,9 @@ export default function ModelsBottomBar({
: undefined; : undefined;
// Determine which model to display - activeModel takes priority when lead/worker is active // 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 = const displayModel =
isLeadWorkerActive && currentModelInfo?.model isLeadWorkerActive && currentModelInfo?.model
? currentModelInfo.model ? currentModelInfo.model
@@ -139,6 +155,10 @@ export default function ModelsBottomBar({
setDisplayModelName(getModelDisplayName(currentModel)); setDisplayModelName(getModelDisplayName(currentModel));
}, [currentModel]); }, [currentModel]);
const handleModelSelected = (model: string, provider: string) => {
onModelChanged({ model, provider });
};
return ( return (
<div className="relative flex items-center" ref={dropdownRef}> <div className="relative flex items-center" ref={dropdownRef}>
<BottomMenuAlertPopover alerts={alerts} /> <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"> <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"> <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" /> <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} {displayModel}
{isLeadWorkerActive && modelMode && ( {isLeadWorkerActive && modelMode && (
<span className="ml-1 text-[10px] opacity-60">({modelMode})</span> <span className="ml-1 text-[10px] opacity-60">({modelMode})</span>
@@ -182,6 +202,9 @@ export default function ModelsBottomBar({
sessionId={sessionId} sessionId={sessionId}
setView={setView} setView={setView}
onClose={() => setIsAddModelModalOpen(false)} onClose={() => setIsAddModelModalOpen(false)}
sessionModel={currentModel}
sessionProvider={currentProvider}
onModelSelected={(model, provider) => handleModelSelected(model, provider)}
/> />
) : null} ) : null}
@@ -23,13 +23,6 @@ vi.mock('../../../ConfigContext', () => ({
}), }),
})); }));
// Minimal mock for useModelAndProvider
vi.mock('../../../ModelAndProviderContext', () => ({
useModelAndProvider: () => ({
currentModel: null,
}),
}));
describe('LeadWorkerSettings', () => { describe('LeadWorkerSettings', () => {
beforeEach(() => { beforeEach(() => {
vi.clearAllMocks(); vi.clearAllMocks();
@@ -1,6 +1,5 @@
import { useState, useEffect } from 'react'; import { useState, useEffect } from 'react';
import { useConfig } from '../../../ConfigContext'; import { useConfig } from '../../../ConfigContext';
import { useModelAndProvider } from '../../../ModelAndProviderContext';
import { Button } from '../../../ui/button'; import { Button } from '../../../ui/button';
import { Select } from '../../../ui/Select'; import { Select } from '../../../ui/Select';
import { Input } from '../../../ui/input'; import { Input } from '../../../ui/input';
@@ -15,7 +14,6 @@ interface LeadWorkerSettingsProps {
export function LeadWorkerSettings({ isOpen, onClose }: LeadWorkerSettingsProps) { export function LeadWorkerSettings({ isOpen, onClose }: LeadWorkerSettingsProps) {
const { read, upsert, getProviders, remove } = useConfig(); const { read, upsert, getProviders, remove } = useConfig();
const { currentModel } = useModelAndProvider();
const [leadModel, setLeadModel] = useState<string>(''); const [leadModel, setLeadModel] = useState<string>('');
const [workerModel, setWorkerModel] = useState<string>(''); const [workerModel, setWorkerModel] = useState<string>('');
const [leadProvider, setLeadProvider] = useState<string>(''); const [leadProvider, setLeadProvider] = useState<string>('');
@@ -69,12 +67,10 @@ export function LeadWorkerSettings({ isOpen, onClose }: LeadWorkerSettingsProps)
if (fallbackTurnsConfig) setFallbackTurns(Number(fallbackTurnsConfig)); if (fallbackTurnsConfig) setFallbackTurns(Number(fallbackTurnsConfig));
else setFallbackTurns(2); else setFallbackTurns(2);
// Set worker model to current model or from config // Set worker model from config
const workerModelConfig = await read('GOOSE_MODEL', false); const workerModelConfig = await read('GOOSE_MODEL', false);
if (workerModelConfig) { if (workerModelConfig) {
setWorkerModel(workerModelConfig as string); setWorkerModel(workerModelConfig as string);
} else if (currentModel) {
setWorkerModel(currentModel as string);
} else { } else {
setWorkerModel(''); setWorkerModel('');
} }
@@ -140,7 +136,7 @@ export function LeadWorkerSettings({ isOpen, onClose }: LeadWorkerSettingsProps)
}; };
loadConfig(); 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 // If current models are not in the list (e.g., previously set to custom), switch to custom mode
useEffect(() => { 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 { Bot, ExternalLink } from 'lucide-react';
import { import {
@@ -84,9 +84,11 @@ type SwitchModelModalProps = {
sessionId: string | null; sessionId: string | null;
onClose: () => void; onClose: () => void;
setView: (view: View) => void; setView: (view: View) => void;
onModelSelected?: (model: string) => void; onModelSelected?: (model: string, provider: string) => void;
initialProvider?: string | null; initialProvider?: string | null;
titleOverride?: string; titleOverride?: string;
sessionModel?: string | null;
sessionProvider?: string | null;
}; };
export const SwitchModelModal = ({ export const SwitchModelModal = ({
sessionId, sessionId,
@@ -95,9 +97,14 @@ export const SwitchModelModal = ({
onModelSelected, onModelSelected,
initialProvider, initialProvider,
titleOverride, titleOverride,
sessionModel,
sessionProvider,
}: SwitchModelModalProps) => { }: SwitchModelModalProps) => {
const { getProviders, read, upsert } = useConfig(); 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 }[]>([]); const [providerOptions, setProviderOptions] = useState<{ value: string; label: string }[]>([]);
type ModelOption = { value: string; label: string; provider: string; isDisabled?: boolean }; type ModelOption = { value: string; label: string; provider: string; isDisabled?: boolean };
const [modelOptions, setModelOptions] = useState<{ options: ModelOption[] }[]>([]); const [modelOptions, setModelOptions] = useState<{ options: ModelOption[] }[]>([]);
@@ -242,10 +249,11 @@ export const SwitchModelModal = ({
} }
} }
await changeModel(sessionId, modelObj); const success = await changeModel(sessionId, modelObj);
onModelSelected?.(modelObj.name); if (success) {
onModelSelected?.(modelObj.name, modelObj.provider || '');
trackModelChanged(modelObj.provider || '', modelObj.name); trackModelChanged(modelObj.provider || '', modelObj.name);
}
onClose(); onClose();
} }
@@ -258,24 +266,37 @@ export const SwitchModelModal = ({
} }
}, [attemptedSubmit, validateForm]); }, [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(() => { useEffect(() => {
// Load predefined models if enabled // Load predefined models if enabled
if (usePredefinedModels) { if (usePredefinedModels) {
const models = getPredefinedModelsFromEnv(); const models = getPredefinedModelsFromEnv();
setPredefinedModels(models); 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 // Load providers for manual model selection
+4 -2
View File
@@ -1,5 +1,4 @@
import { useEffect, useRef, useState } from 'react'; import { useEffect, useRef, useState } from 'react';
import { useModelAndProvider } from '../components/ModelAndProviderContext';
import { fetchCanonicalModelInfo } from '../utils/canonical'; import { fetchCanonicalModelInfo } from '../utils/canonical';
import { Session } from '../api'; 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 prevModelRef = useRef<string | undefined>(undefined);
const prevProviderRef = useRef<string | undefined>(undefined); const prevProviderRef = useRef<string | undefined>(undefined);
// Handle model changes and accumulate costs // Handle model changes and accumulate costs
useEffect(() => { useEffect(() => {
if (!currentModel || !currentProvider) return;
const handleModelChange = async () => { const handleModelChange = async () => {
if ( if (
prevModelRef.current !== undefined && prevModelRef.current !== undefined &&