Show modal selector after configuring a provider (#6005)

This commit is contained in:
Zane
2025-12-08 13:27:53 -08:00
committed by GitHub
parent 13d1f3077e
commit 92e75ac460
5 changed files with 201 additions and 108 deletions
+33 -27
View File
@@ -1,4 +1,4 @@
import { useEffect, useState } from 'react'; import { useEffect, useState, useMemo } from 'react';
import { useNavigate } from 'react-router-dom'; import { useNavigate } from 'react-router-dom';
import { useConfig } from './ConfigContext'; import { useConfig } from './ConfigContext';
import { SetupModal } from './SetupModal'; import { SetupModal } from './SetupModal';
@@ -8,6 +8,8 @@ import WelcomeGooseLogo from './WelcomeGooseLogo';
import { toastService } from '../toasts'; import { toastService } from '../toasts';
import { OllamaSetup } from './OllamaSetup'; import { OllamaSetup } from './OllamaSetup';
import ApiKeyTester from './ApiKeyTester'; import ApiKeyTester from './ApiKeyTester';
import { SwitchModelModal } from './settings/models/subcomponents/SwitchModelModal';
import { createNavigationHandler } from '../utils/navigationUtils';
import { Goose, OpenRouter, Tetrate } from './icons'; import { Goose, OpenRouter, Tetrate } from './icons';
@@ -24,6 +26,10 @@ export default function ProviderGuard({ didSelectProvider, children }: ProviderG
const [showFirstTimeSetup, setShowFirstTimeSetup] = useState(false); const [showFirstTimeSetup, setShowFirstTimeSetup] = useState(false);
const [showOllamaSetup, setShowOllamaSetup] = useState(false); const [showOllamaSetup, setShowOllamaSetup] = useState(false);
const [userInActiveSetup, setUserInActiveSetup] = useState(false); const [userInActiveSetup, setUserInActiveSetup] = useState(false);
const [showSwitchModelModal, setShowSwitchModelModal] = useState(false);
const [switchModelProvider, setSwitchModelProvider] = useState<string | null>(null);
const setView = useMemo(() => createNavigationHandler(navigate), [navigate]);
const [openRouterSetupState, setOpenRouterSetupState] = useState<{ const [openRouterSetupState, setOpenRouterSetupState] = useState<{
show: boolean; show: boolean;
@@ -45,18 +51,8 @@ export default function ProviderGuard({ didSelectProvider, children }: ProviderG
try { try {
const result = await startTetrateSetup(); const result = await startTetrateSetup();
if (result.success) { if (result.success) {
setTetrateSetupState({ setSwitchModelProvider('tetrate');
show: true, setShowSwitchModelModal(true);
title: 'Setup Complete!',
message: result.message,
showRetry: false,
autoClose: 3000,
});
setTimeout(() => {
setShowFirstTimeSetup(false);
setHasProvider(true);
navigate('/', { replace: true });
}, 3000);
} else { } else {
setTetrateSetupState({ setTetrateSetupState({
show: true, show: true,
@@ -76,34 +72,33 @@ export default function ProviderGuard({ didSelectProvider, children }: ProviderG
} }
}; };
const handleApiKeySuccess = async (provider: string, model: string, apiKey: string) => { const handleApiKeySuccess = async (provider: string, _model: string, apiKey: string) => {
const keyName = `${provider.toUpperCase()}_API_KEY`; const keyName = `${provider.toUpperCase()}_API_KEY`;
await upsert(keyName, apiKey, true); await upsert(keyName, apiKey, true);
await upsert('GOOSE_PROVIDER', provider, false); await upsert('GOOSE_PROVIDER', provider, false);
await upsert('GOOSE_MODEL', model, false);
setSwitchModelProvider(provider);
setShowSwitchModelModal(true);
};
const handleModelSelected = () => {
setShowSwitchModelModal(false);
setUserInActiveSetup(false); setUserInActiveSetup(false);
setShowFirstTimeSetup(false); setShowFirstTimeSetup(false);
setHasProvider(true); setHasProvider(true);
navigate('/', { replace: true }); navigate('/', { replace: true });
}; };
const handleSwitchModelClose = () => {
setShowSwitchModelModal(false);
};
const handleOpenRouterSetup = async () => { const handleOpenRouterSetup = async () => {
try { try {
const result = await startOpenRouterSetup(); const result = await startOpenRouterSetup();
if (result.success) { if (result.success) {
setOpenRouterSetupState({ setSwitchModelProvider('openrouter');
show: true, setShowSwitchModelModal(true);
title: 'Setup Complete!',
message: result.message,
showRetry: false,
autoClose: 3000,
});
setTimeout(() => {
setShowFirstTimeSetup(false);
setHasProvider(true);
navigate('/', { replace: true });
}, 3000);
} else { } else {
setOpenRouterSetupState({ setOpenRouterSetupState({
show: true, show: true,
@@ -337,6 +332,17 @@ export default function ProviderGuard({ didSelectProvider, children }: ProviderG
autoClose={tetrateSetupState.autoClose} autoClose={tetrateSetupState.autoClose}
/> />
)} )}
{showSwitchModelModal && (
<SwitchModelModal
sessionId={null}
onClose={handleSwitchModelClose}
setView={setView}
onModelSelected={handleModelSelected}
initialProvider={switchModelProvider}
titleOverride="Choose Model"
/>
)}
</div> </div>
); );
} }
@@ -1,5 +1,5 @@
import { useEffect, useState, useCallback } from 'react'; import { useEffect, useState, useCallback } from 'react';
import { ArrowLeftRight, ExternalLink } from 'lucide-react'; import { Bot, ExternalLink } from 'lucide-react';
import { import {
Dialog, Dialog,
@@ -20,18 +20,66 @@ import Model, { getProviderMetadata, fetchModelsForProviders } from '../modelInt
import { getPredefinedModelsFromEnv, shouldShowPredefinedModels } from '../predefinedModelsUtils'; import { getPredefinedModelsFromEnv, shouldShowPredefinedModels } from '../predefinedModelsUtils';
import { ProviderType } from '../../../../api'; import { ProviderType } from '../../../../api';
const PREFERRED_MODEL_PATTERNS = [
/claude-sonnet-4/i,
/claude-4/i,
/gpt-4o(?!-mini)/i,
/claude-3-5-sonnet/i,
/claude-3\.5-sonnet/i,
/gpt-4-turbo/i,
/gpt-4(?!-|o)/i,
/claude-3-opus/i,
/claude-3-sonnet/i,
/gemini-pro/i,
/llama-3/i,
/gpt-4o-mini/i,
/claude-3-haiku/i,
/gemini/i,
];
function findPreferredModel(
models: { value: string; label: string; provider: string }[]
): string | null {
if (models.length === 0) return null;
const validModels = models.filter(
(m) => m.value !== 'custom' && m.value !== '__loading__' && !m.value.startsWith('__')
);
if (validModels.length === 0) return null;
for (const pattern of PREFERRED_MODEL_PATTERNS) {
const match = validModels.find((m) => pattern.test(m.value));
if (match) {
return match.value;
}
}
return validModels[0].value;
}
type SwitchModelModalProps = { type SwitchModelModalProps = {
sessionId: string | null; sessionId: string | null;
onClose: () => void; onClose: () => void;
setView: (view: View) => void; setView: (view: View) => void;
onModelSelected?: () => void;
initialProvider?: string | null;
titleOverride?: string;
}; };
export const SwitchModelModal = ({ sessionId, onClose, setView }: SwitchModelModalProps) => { export const SwitchModelModal = ({
sessionId,
onClose,
setView,
onModelSelected,
initialProvider,
titleOverride,
}: SwitchModelModalProps) => {
const { getProviders, getProviderModels, read } = useConfig(); const { getProviders, getProviderModels, read } = useConfig();
const { changeModel } = useModelAndProvider(); const { changeModel } = useModelAndProvider();
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[] }[]>([]);
const [provider, setProvider] = useState<string | null>(null); const [provider, setProvider] = useState<string | null>(initialProvider || null);
const [model, setModel] = useState<string>(''); const [model, setModel] = useState<string>('');
const [isCustomModel, setIsCustomModel] = useState(false); const [isCustomModel, setIsCustomModel] = useState(false);
const [validationErrors, setValidationErrors] = useState({ const [validationErrors, setValidationErrors] = useState({
@@ -95,6 +143,9 @@ export const SwitchModelModal = ({ sessionId, onClose, setView }: SwitchModelMod
} }
await changeModel(sessionId, modelObj); await changeModel(sessionId, modelObj);
if (onModelSelected) {
onModelSelected();
}
onClose(); onClose();
} }
}; };
@@ -209,11 +260,25 @@ export const SwitchModelModal = ({ sessionId, onClose, setView }: SwitchModelMod
})(); })();
}, [getProviders, getProviderModels, usePredefinedModels, read]); }, [getProviders, getProviderModels, usePredefinedModels, read]);
// Filter model options based on selected provider
const filteredModelOptions = provider const filteredModelOptions = provider
? modelOptions.filter((group) => group.options[0]?.provider === provider) ? modelOptions.filter((group) => group.options[0]?.provider === provider)
: []; : [];
useEffect(() => {
if (!provider || loadingModels || model || isCustomModel) return;
const providerModels = modelOptions
.filter((group) => group.options[0]?.provider === provider)
.flatMap((group) => group.options);
if (providerModels.length > 0) {
const preferredModel = findPreferredModel(providerModels);
if (preferredModel) {
setModel(preferredModel);
}
}
}, [provider, modelOptions, loadingModels, model, isCustomModel]);
// Handle model selection change // Handle model selection change
const handleModelChange = (newValue: unknown) => { const handleModelChange = (newValue: unknown) => {
const selectedOption = newValue as { value: string; label: string; provider: string } | null; const selectedOption = newValue as { value: string; label: string; provider: string } | null;
@@ -277,30 +342,16 @@ export const SwitchModelModal = ({ sessionId, onClose, setView }: SwitchModelMod
<DialogContent className="sm:max-w-[500px]"> <DialogContent className="sm:max-w-[500px]">
<DialogHeader> <DialogHeader>
<DialogTitle className="flex items-center gap-2"> <DialogTitle className="flex items-center gap-2">
<ArrowLeftRight size={24} className="text-textStandard" /> <Bot size={24} className="text-textStandard" />
Switch models {titleOverride || 'Switch models'}
</DialogTitle> </DialogTitle>
<DialogDescription> <DialogDescription>
Configure your AI model providers by adding their API keys. Your keys are stored Select a provider and model to use for your conversations.
securely and encrypted locally.
</DialogDescription> </DialogDescription>
</DialogHeader> </DialogHeader>
<div className="flex flex-col gap-4 py-4"> <div className="flex flex-col gap-4 py-4">
<div>
<a
href={QUICKSTART_GUIDE_URL}
target="_blank"
rel="noopener noreferrer"
className="flex items-center text-textStandard font-medium text-sm"
>
<ExternalLink size={16} className="mr-1" />
View quick start guide
</a>
</div>
{usePredefinedModels ? ( {usePredefinedModels ? (
/* Predefined Models Section */
<div className="w-full flex flex-col gap-4"> <div className="w-full flex flex-col gap-4">
<div className="flex justify-between items-center"> <div className="flex justify-between items-center">
<label className="text-sm font-medium text-textStandard">Choose a model:</label> <label className="text-sm font-medium text-textStandard">Choose a model:</label>
@@ -448,13 +499,24 @@ export const SwitchModelModal = ({ sessionId, onClose, setView }: SwitchModelMod
)} )}
</div> </div>
<DialogFooter className="pt-2"> <DialogFooter className="pt-4 flex-col sm:flex-row gap-3">
<Button variant="outline" onClick={handleClose} type="button"> <a
Cancel href={QUICKSTART_GUIDE_URL}
</Button> target="_blank"
<Button onClick={handleSubmit} disabled={!isValid}> rel="noopener noreferrer"
Select model className="inline-flex items-center text-text-muted hover:text-textStandard text-sm mr-auto"
</Button> >
<ExternalLink size={14} className="mr-1" />
Quick start guide
</a>
<div className="flex gap-2">
<Button variant="outline" onClick={handleClose} type="button">
Cancel
</Button>
<Button onClick={handleSubmit} disabled={!isValid}>
Select model
</Button>
</div>
</DialogFooter> </DialogFooter>
</DialogContent> </DialogContent>
</Dialog> </Dialog>
@@ -10,6 +10,8 @@ import {
import { Plus } from 'lucide-react'; import { Plus } from 'lucide-react';
import { Dialog, DialogContent, DialogHeader, DialogTitle } from '../../ui/dialog'; import { Dialog, DialogContent, DialogHeader, DialogTitle } from '../../ui/dialog';
import CustomProviderForm from './modal/subcomponents/forms/CustomProviderForm'; import CustomProviderForm from './modal/subcomponents/forms/CustomProviderForm';
import { SwitchModelModal } from '../models/subcomponents/SwitchModelModal';
import type { View } from '../../../utils/navigationUtils';
const GridLayout = memo(function GridLayout({ children }: { children: React.ReactNode }) { const GridLayout = memo(function GridLayout({ children }: { children: React.ReactNode }) {
return ( return (
@@ -50,21 +52,30 @@ function ProviderCards({
providers, providers,
isOnboarding, isOnboarding,
refreshProviders, refreshProviders,
onProviderLaunch, setView,
onModelSelected,
}: { }: {
providers: ProviderDetails[]; providers: ProviderDetails[];
isOnboarding: boolean; isOnboarding: boolean;
refreshProviders?: () => void; refreshProviders?: () => void;
onProviderLaunch: (provider: ProviderDetails) => void; setView?: (view: View) => void;
onModelSelected?: () => void;
}) { }) {
const [configuringProvider, setConfiguringProvider] = useState<ProviderDetails | null>(null); const [configuringProvider, setConfiguringProvider] = useState<ProviderDetails | null>(null);
const [showCustomProviderModal, setShowCustomProviderModal] = useState(false); const [showCustomProviderModal, setShowCustomProviderModal] = useState(false);
const [showSwitchModelModal, setShowSwitchModelModal] = useState(false);
const [switchModelProvider, setSwitchModelProvider] = useState<string | null>(null);
const [editingProvider, setEditingProvider] = useState<{ const [editingProvider, setEditingProvider] = useState<{
id: string; id: string;
config: DeclarativeProviderConfig; config: DeclarativeProviderConfig;
isEditable: boolean; isEditable: boolean;
} | null>(null); } | null>(null);
const handleProviderLaunchWithModelSelection = useCallback((provider: ProviderDetails) => {
setSwitchModelProvider(provider.name);
setShowSwitchModelModal(true);
}, []);
const openModal = useCallback( const openModal = useCallback(
(provider: ProviderDetails) => setConfiguringProvider(provider), (provider: ProviderDetails) => setConfiguringProvider(provider),
[] []
@@ -101,11 +112,14 @@ function ProviderCards({
body: data, body: data,
throwOnError: true, throwOnError: true,
}); });
const providerId = editingProvider.id;
setShowCustomProviderModal(false); setShowCustomProviderModal(false);
setEditingProvider(null); setEditingProvider(null);
if (refreshProviders) { if (refreshProviders) {
refreshProviders(); refreshProviders();
} }
setSwitchModelProvider(providerId);
setShowSwitchModelModal(true);
}, },
[editingProvider, refreshProviders] [editingProvider, refreshProviders]
); );
@@ -122,6 +136,32 @@ function ProviderCards({
} }
}, [refreshProviders]); }, [refreshProviders]);
const onProviderConfigured = useCallback(
(provider: ProviderDetails) => {
setConfiguringProvider(null);
if (refreshProviders) {
refreshProviders();
}
setSwitchModelProvider(provider.name);
setShowSwitchModelModal(true);
},
[refreshProviders]
);
const onCloseSwitchModelModal = useCallback(() => {
setShowSwitchModelModal(false);
}, []);
const handleSetView = useCallback(
(view: View) => {
setShowSwitchModelModal(false);
if (setView) {
setView(view);
}
},
[setView]
);
const handleCreateCustomProvider = useCallback( const handleCreateCustomProvider = useCallback(
async (data: UpdateCustomProviderRequest) => { async (data: UpdateCustomProviderRequest) => {
const { createCustomProvider } = await import('../../../api'); const { createCustomProvider } = await import('../../../api');
@@ -130,6 +170,7 @@ function ProviderCards({
if (refreshProviders) { if (refreshProviders) {
refreshProviders(); refreshProviders();
} }
setShowSwitchModelModal(true);
}, },
[refreshProviders] [refreshProviders]
); );
@@ -144,7 +185,7 @@ function ProviderCards({
key={provider.name} key={provider.name}
provider={provider} provider={provider}
onConfigure={() => configureProviderViaModal(provider)} onConfigure={() => configureProviderViaModal(provider)}
onLaunch={() => onProviderLaunch(provider)} onLaunch={() => handleProviderLaunchWithModelSelection(provider)}
isOnboarding={isOnboarding} isOnboarding={isOnboarding}
/> />
)); ));
@@ -154,7 +195,7 @@ function ProviderCards({
); );
return cards; return cards;
}, [providers, isOnboarding, configureProviderViaModal, onProviderLaunch]); }, [providers, isOnboarding, configureProviderViaModal, handleProviderLaunchWithModelSelection]);
const initialData = editingProvider && { const initialData = editingProvider && {
engine: editingProvider.config.engine.toLowerCase() + '_compatible', engine: editingProvider.config.engine.toLowerCase() + '_compatible',
@@ -187,6 +228,17 @@ function ProviderCards({
<ProviderConfigurationModal <ProviderConfigurationModal
provider={configuringProvider} provider={configuringProvider}
onClose={onCloseProviderConfig} onClose={onCloseProviderConfig}
onConfigured={onProviderConfigured}
/>
)}
{showSwitchModelModal && (
<SwitchModelModal
sessionId={null}
onClose={onCloseSwitchModelModal}
setView={handleSetView}
onModelSelected={onModelSelected}
initialProvider={switchModelProvider}
titleOverride="Choose Model"
/> />
)} )}
</> </>
@@ -197,12 +249,14 @@ export default function ProviderGrid({
providers, providers,
isOnboarding, isOnboarding,
refreshProviders, refreshProviders,
onProviderLaunch, setView,
onModelSelected,
}: { }: {
providers: ProviderDetails[]; providers: ProviderDetails[];
isOnboarding: boolean; isOnboarding: boolean;
refreshProviders?: () => void; refreshProviders?: () => void;
onProviderLaunch?: (provider: ProviderDetails) => void; setView?: (view: View) => void;
onModelSelected?: () => void;
}) { }) {
return ( return (
<GridLayout> <GridLayout>
@@ -210,7 +264,8 @@ export default function ProviderGrid({
providers={providers} providers={providers}
isOnboarding={isOnboarding} isOnboarding={isOnboarding}
refreshProviders={refreshProviders} refreshProviders={refreshProviders}
onProviderLaunch={onProviderLaunch || (() => {})} setView={setView}
onModelSelected={onModelSelected}
/> />
</GridLayout> </GridLayout>
); );
@@ -1,10 +1,11 @@
import { useEffect, useState, useCallback, useRef } from 'react'; import { useEffect, useState, useCallback, useRef, useMemo } from 'react';
import { useNavigate } from 'react-router-dom';
import { ScrollArea } from '../../ui/scroll-area'; import { ScrollArea } from '../../ui/scroll-area';
import BackButton from '../../ui/BackButton'; import BackButton from '../../ui/BackButton';
import ProviderGrid from './ProviderGrid'; import ProviderGrid from './ProviderGrid';
import { useConfig } from '../../ConfigContext'; import { useConfig } from '../../ConfigContext';
import { ProviderDetails, setConfigProvider } from '../../../api'; import { ProviderDetails } from '../../../api';
import { toastService } from '../../../toasts'; import { createNavigationHandler } from '../../../utils/navigationUtils';
interface ProviderSettingsProps { interface ProviderSettingsProps {
onClose: () => void; onClose: () => void;
@@ -18,10 +19,13 @@ export default function ProviderSettings({
onProviderLaunched, onProviderLaunched,
}: ProviderSettingsProps) { }: ProviderSettingsProps) {
const { getProviders } = useConfig(); const { getProviders } = useConfig();
const navigate = useNavigate();
const [loading, setLoading] = useState(true); const [loading, setLoading] = useState(true);
const [providers, setProviders] = useState<ProviderDetails[]>([]); const [providers, setProviders] = useState<ProviderDetails[]>([]);
const initialLoadDone = useRef(false); const initialLoadDone = useRef(false);
const setView = useMemo(() => createNavigationHandler(navigate), [navigate]);
// Create a function to load providers that can be called multiple times // Create a function to load providers that can be called multiple times
const loadProviders = useCallback(async () => { const loadProviders = useCallback(async () => {
setLoading(true); setLoading(true);
@@ -54,47 +58,6 @@ export default function ProviderSettings({
} }
}, [getProviders]); }, [getProviders]);
// Handler for when a provider is launched if this component is used as part of onboarding page
const handleProviderLaunch = useCallback(
async (provider: ProviderDetails) => {
const provider_name = provider.name;
const model = provider.metadata.default_model;
try {
await setConfigProvider({
body: {
provider: provider_name,
model,
},
throwOnError: true,
});
toastService.configure({ silent: false });
toastService.success({
title: 'Success!',
msg: `Started goose with ${model} by ${provider.metadata.display_name}. You can change the model via the dropdown.`,
});
if (onProviderLaunched) {
onProviderLaunched();
} else {
onClose();
}
} catch (error) {
console.error(`Failed to initialize with provider ${provider_name}:`, error);
// Show error toast
toastService.configure({ silent: false });
toastService.error({
title: 'Initialization Failed',
msg: `Failed to initialize with ${provider.metadata.display_name}: ${error instanceof Error ? error.message : String(error)}`,
traceback: error instanceof Error ? error.stack || '' : '',
});
}
},
[onClose, onProviderLaunched]
);
return ( return (
<div className="h-screen w-full flex flex-col bg-background-default text-text-default"> <div className="h-screen w-full flex flex-col bg-background-default text-text-default">
<ScrollArea className="flex-1 w-full"> <ScrollArea className="flex-1 w-full">
@@ -127,8 +90,9 @@ export default function ProviderSettings({
<ProviderGrid <ProviderGrid
providers={providers} providers={providers}
isOnboarding={isOnboarding} isOnboarding={isOnboarding}
onProviderLaunch={handleProviderLaunch}
refreshProviders={refreshProviders} refreshProviders={refreshProviders}
setView={setView}
onModelSelected={onProviderLaunched}
/> />
)} )}
</div> </div>
@@ -23,11 +23,13 @@ import { Button } from '../../../../components/ui/button';
interface ProviderConfigurationModalProps { interface ProviderConfigurationModalProps {
provider: ProviderDetails; provider: ProviderDetails;
onClose: () => void; onClose: () => void;
onConfigured?: (provider: ProviderDetails) => void;
} }
export default function ProviderConfigurationModal({ export default function ProviderConfigurationModal({
provider, provider,
onClose, onClose,
onConfigured,
}: ProviderConfigurationModalProps) { }: ProviderConfigurationModalProps) {
const [validationErrors, setValidationErrors] = useState<Record<string, string>>({}); const [validationErrors, setValidationErrors] = useState<Record<string, string>>({});
const { upsert, remove } = useConfig(); const { upsert, remove } = useConfig();
@@ -83,7 +85,11 @@ export default function ProviderConfigurationModal({
try { try {
await providerConfigSubmitHandler(upsert, provider, toSubmit); await providerConfigSubmitHandler(upsert, provider, toSubmit);
onClose(); if (onConfigured) {
onConfigured(provider);
} else {
onClose();
}
} catch (error) { } catch (error) {
setError(`${error}`); setError(`${error}`);
} }