fix: V2 settings carry extensions over during model change (#1944)

This commit is contained in:
Alex Hancock
2025-03-31 15:05:04 -04:00
committed by GitHub
parent cd0e65177b
commit 41f3478893
6 changed files with 80 additions and 88 deletions
@@ -21,6 +21,8 @@ import type {
} from '../api/types.gen';
import { removeShims } from './settings_v2/extensions/utils';
export type { ExtensionConfig } from '../api/types.gen';
// Define a local version that matches the structure of the imported one
export type FixedExtensionEntry = ExtensionConfig & {
enabled: boolean;
@@ -1,8 +1,9 @@
import { initializeAgent } from '../../../agent/index';
import { initializeSystem } from '../../../utils/providerUtils';
import { toastError, toastSuccess } from '../../../toasts';
import { ProviderDetails } from '@/src/api';
import { getProviderMetadata } from './modelInterface';
import { ProviderMetadata } from '../../../api';
import type { ExtensionConfig, FixedExtensionEntry } from '../../ConfigContext';
// titles
const CHANGE_MODEL_TOAST_TITLE = 'Model selected';
@@ -23,12 +24,23 @@ interface changeModelProps {
model: string;
provider: string;
writeToConfig: (key: string, value: unknown, is_secret: boolean) => Promise<void>;
getExtensions?: (b: boolean) => Promise<FixedExtensionEntry[]>;
addExtension?: (name: string, config: ExtensionConfig, enabled: boolean) => Promise<void>;
}
// TODO: error handling
export async function changeModel({ model, provider, writeToConfig }: changeModelProps) {
export async function changeModel({
model,
provider,
writeToConfig,
getExtensions,
addExtension,
}: changeModelProps) {
try {
await initializeAgent({ model: model, provider: provider });
await initializeSystem(provider, model, {
getExtensions,
addExtension,
});
} catch (error) {
console.error(`Failed to change model at agent step -- ${model} ${provider}`);
toastError({
@@ -61,49 +73,6 @@ export async function changeModel({ model, provider, writeToConfig }: changeMode
}
}
interface startAgentFromConfigProps {
readFromConfig: (key: string, is_secret: boolean) => Promise<unknown>;
}
// starts agent with the values for GOOSE_PROVIDER and GOOSE_MODEL that are in the config
export async function startAgentFromConfig({ readFromConfig }: startAgentFromConfigProps) {
let modelProvider: { model: string; provider: string };
// read from config
try {
modelProvider = await getCurrentModelAndProvider({ readFromConfig: readFromConfig });
} catch (error) {
toastError({
title: START_AGENT_TITLE,
msg: CONFIG_READ_MODEL_ERROR_MSG,
traceback: error,
});
return;
}
const model = modelProvider.model;
const provider = modelProvider.provider;
console.log(`Starting agent with GOOSE_MODEL=${model} and GOOSE_PROVIDER=${provider}`);
try {
await initializeAgent({ model: model, provider: provider });
} catch (error) {
console.error(`Failed to change model at agent step -- ${model} ${provider}`);
toastError({
title: CHANGE_MODEL_TOAST_TITLE,
msg: SWITCH_MODEL_AGENT_ERROR_MSG,
traceback: error,
});
return;
} finally {
toastSuccess({
title: CHANGE_MODEL_TOAST_TITLE,
msg: `${INITIALIZE_SYSTEM_WITH_MODEL_SUCCESS_MSG} with ${model} from ${provider}`,
});
}
}
interface getCurrentModelAndProviderProps {
readFromConfig: (key: string, is_secret: boolean) => Promise<unknown>;
}
@@ -7,7 +7,7 @@ import { QUICKSTART_GUIDE_URL } from '../../providers/modal/constants';
import { Input } from '../../../ui/input';
import { Select } from '../../../ui/Select';
import { useConfig } from '../../../ConfigContext';
import { changeModel as switchModel } from '../index';
import { changeModel } from '../index';
import type { View } from '../../../../App';
const ModalButtons = ({ onSubmit, onCancel, isValid, validationErrors }) => (
@@ -36,7 +36,7 @@ type AddModelModalProps = {
setView: (view: View) => void;
};
export const AddModelModal = ({ onClose, setView }: AddModelModalProps) => {
const { getProviders, upsert } = useConfig();
const { getProviders, upsert, getExtensions, addExtension } = useConfig();
const [providerOptions, setProviderOptions] = useState([]);
const [modelOptions, setModelOptions] = useState([]);
const [provider, setProvider] = useState<string | null>(null);
@@ -72,12 +72,18 @@ export const AddModelModal = ({ onClose, setView }: AddModelModalProps) => {
return formIsValid;
};
const changeModel = async () => {
const onSubmit = async () => {
setAttemptedSubmit(true);
const isFormValid = validateForm();
if (isFormValid) {
await switchModel({ model: model, provider: provider, writeToConfig: upsert });
await changeModel({
model: model,
provider: provider,
writeToConfig: upsert,
getExtensions,
addExtension,
});
onClose();
}
};
@@ -159,7 +165,7 @@ export const AddModelModal = ({ onClose, setView }: AddModelModalProps) => {
onClose={onClose}
footer={
<ModalButtons
onSubmit={changeModel}
onSubmit={onSubmit}
onCancel={onClose}
isValid={isValid}
validationErrors={validationErrors}
@@ -4,7 +4,7 @@ import BackButton from '../../ui/BackButton';
import ProviderGrid from './ProviderGrid';
import { useConfig } from '../../ConfigContext';
import { ProviderDetails } from '../../../api/types.gen';
import { initializeAgent } from '../../../agent/';
import { initializeSystem } from '../../../utils/providerUtils';
import WelcomeGooseLogo from '../../WelcomeGooseLogo';
interface ProviderSettingsProps {
@@ -13,7 +13,7 @@ interface ProviderSettingsProps {
}
export default function ProviderSettings({ onClose, isOnboarding }: ProviderSettingsProps) {
const { getProviders, upsert } = useConfig();
const { getProviders, upsert, getExtensions, addExtension } = useConfig();
const [loading, setLoading] = useState(true);
const [providers, setProviders] = useState<ProviderDetails[]>([]);
const initialLoadDone = useRef(false);
@@ -69,13 +69,16 @@ export default function ProviderSettings({ onClose, isOnboarding }: ProviderSett
);
// initialize agent
await initializeAgent({ provider: provider.name, model });
await initializeSystem(provider.name, model, {
getExtensions,
addExtension,
});
} catch (error) {
console.error(`Failed to initialize with provider ${provider_name}:`, error);
}
onClose();
},
[initializeAgent, onClose, upsert]
[onClose, upsert]
);
return (