From cfa2778b4d07a9f529f41943192d1cb824d4c036 Mon Sep 17 00:00:00 2001 From: Abhijay Jain Date: Thu, 12 Feb 2026 02:56:00 +0530 Subject: [PATCH] feat: load provider/model specified inside the recipe config (#6884) Signed-off-by: Abhijay007 Co-authored-by: Zane Staggs --- crates/goose-server/src/routes/agent.rs | 31 +++++++++------ crates/goose/src/agents/agent.rs | 3 +- ui/desktop/src/components/BaseChat.tsx | 8 ++++ .../components/ModelAndProviderContext.tsx | 8 ++++ .../models/bottom_bar/ModelsBottomBar.tsx | 39 +++++++------------ .../models/subcomponents/SwitchModelModal.tsx | 19 +++------ 6 files changed, 56 insertions(+), 52 deletions(-) diff --git a/crates/goose-server/src/routes/agent.rs b/crates/goose-server/src/routes/agent.rs index 6ef557d4..c8ceb98d 100644 --- a/crates/goose-server/src/routes/agent.rs +++ b/crates/goose-server/src/routes/agent.rs @@ -254,18 +254,27 @@ async fn start_agent( } if let Some(recipe) = original_recipe { - manager - .update(&session.id) - .recipe(Some(recipe)) - .apply() - .await - .map_err(|err| { - error!("Failed to update session with recipe: {}", err); - ErrorResponse { - message: format!("Failed to update session with recipe: {}", err), - status: StatusCode::INTERNAL_SERVER_ERROR, + let mut update = manager.update(&session.id).recipe(Some(recipe.clone())); + + if let Some(ref settings) = recipe.settings { + if let Some(ref provider) = settings.goose_provider { + update = update.provider_name(provider); + + if let Some(ref model) = settings.goose_model { + if let Ok(model_config) = ModelConfig::new(model) { + update = update.model_config(model_config); + } } - })?; + } + } + + update.apply().await.map_err(|err| { + error!("Failed to update session with recipe: {}", err); + ErrorResponse { + message: format!("Failed to update session with recipe: {}", err), + status: StatusCode::INTERNAL_SERVER_ERROR, + } + })?; } // Refetch session to get all updates diff --git a/crates/goose/src/agents/agent.rs b/crates/goose/src/agents/agent.rs index 309378e1..3d476ead 100644 --- a/crates/goose/src/agents/agent.rs +++ b/crates/goose/src/agents/agent.rs @@ -1577,7 +1577,8 @@ impl Agent { None => { let model_name = config .get_goose_model() - .map_err(|_| anyhow!("Could not configure agent: missing model"))?; + .ok() + .ok_or_else(|| anyhow!("Could not configure agent: missing model"))?; crate::model::ModelConfig::new(&model_name) .map_err(|e| anyhow!("Could not configure agent: invalid model {}", e))? } diff --git a/ui/desktop/src/components/BaseChat.tsx b/ui/desktop/src/components/BaseChat.tsx index 7b335ff1..2278d262 100644 --- a/ui/desktop/src/components/BaseChat.tsx +++ b/ui/desktop/src/components/BaseChat.tsx @@ -35,6 +35,7 @@ import { useToolCount } from './alerts/useToolCount'; import { getThinkingMessage, getTextAndImageContent } from '../types/message'; import ParameterInputModal from './ParameterInputModal'; import { substituteParameters } from '../utils/providerUtils'; +import { useModelAndProvider } from './ModelAndProviderContext'; import CreateRecipeFromSessionModal from './recipes/CreateRecipeFromSessionModal'; import { toastSuccess } from '../toasts'; import { Recipe } from '../recipe'; @@ -176,6 +177,13 @@ export default function BaseChat({ }); const recipe = session?.recipe; + 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]); useEffect(() => { if (!recipe) return; diff --git a/ui/desktop/src/components/ModelAndProviderContext.tsx b/ui/desktop/src/components/ModelAndProviderContext.tsx index 04e279e9..9b5e2fab 100644 --- a/ui/desktop/src/components/ModelAndProviderContext.tsx +++ b/ui/desktop/src/components/ModelAndProviderContext.tsx @@ -26,6 +26,7 @@ interface ModelAndProviderContextType { getCurrentModelDisplayName: () => Promise; getCurrentProviderDisplayName: () => Promise; // Gets provider display name from subtext refreshCurrentModelAndProvider: () => Promise; + setProviderAndModel: (provider: string, model: string) => void; } interface ModelAndProviderProviderProps { @@ -173,6 +174,11 @@ 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(); @@ -189,6 +195,7 @@ export const ModelAndProviderProvider: React.FC = getCurrentModelDisplayName, getCurrentProviderDisplayName, refreshCurrentModelAndProvider, + setProviderAndModel, }), [ currentModel, @@ -200,6 +207,7 @@ export const ModelAndProviderProvider: React.FC = getCurrentModelDisplayName, getCurrentProviderDisplayName, refreshCurrentModelAndProvider, + setProviderAndModel, ] ); 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 190dd799..f7c97c2d 100644 --- a/ui/desktop/src/components/settings/models/bottom_bar/ModelsBottomBar.tsx +++ b/ui/desktop/src/components/settings/models/bottom_bar/ModelsBottomBar.tsx @@ -13,6 +13,7 @@ import { import { useCurrentModelInfo } from '../../../BaseChat'; import { useConfig } from '../../../ConfigContext'; import { getProviderMetadata } from '../modelInterface'; +import { getModelDisplayName } from '../predefinedModelsUtils'; import { Alert } from '../../../alerts'; import BottomMenuAlertPopover from '../../../bottom_menu/BottomMenuAlertPopover'; @@ -32,9 +33,6 @@ export default function ModelsBottomBar({ const { currentModel, currentProvider, - getCurrentModelAndProviderForDisplay, - getCurrentModelDisplayName, - getCurrentProviderDisplayName, } = useModelAndProvider(); const currentModelInfo = useCurrentModelInfo(); const { read, getProviders } = useConfig(); @@ -62,7 +60,6 @@ export default function ModelsBottomBar({ // Refresh lead/worker status when modal closes const handleLeadWorkerModalClose = () => { setIsLeadWorkerModalOpen(false); - // Refresh the lead/worker status after modal closes const checkLeadWorker = async () => { try { const leadModel = await read('GOOSE_LEAD_MODEL', false); @@ -78,8 +75,6 @@ export default function ModelsBottomBar({ checkLeadWorker(); }; - // Since currentModelInfo.mode is not working, let's determine mode differently - // We'll need to get the lead model and compare it with the current model const [leadModelName, setLeadModelName] = useState(''); const [currentActiveModel, setCurrentActiveModel] = useState(''); @@ -111,20 +106,16 @@ export default function ModelsBottomBar({ ? currentModelInfo.model : currentModel || providerDefaultModel || displayModelName; - // Update display provider when current provider changes useEffect(() => { - if (currentProvider) { - (async () => { - const providerDisplayName = await getCurrentProviderDisplayName(); - if (providerDisplayName) { - setDisplayProvider(providerDisplayName); - } else { - const modelProvider = await getCurrentModelAndProviderForDisplay(); - setDisplayProvider(modelProvider.provider); - } - })(); - } - }, [currentProvider, getCurrentProviderDisplayName, getCurrentModelAndProviderForDisplay]); + if (!currentProvider) return; + getProviderMetadata(currentProvider, getProviders) + .then((metadata) => { + setDisplayProvider(metadata.display_name || currentProvider); + }) + .catch(() => { + setDisplayProvider(currentProvider); + }); + }, [currentProvider, currentModel, getProviders]); // Fetch provider default model when provider changes and no current model useEffect(() => { @@ -139,18 +130,14 @@ export default function ModelsBottomBar({ } })(); } else if (currentModel) { - // Clear provider default when we have a current model setProviderDefaultModel(null); } }, [currentProvider, currentModel, getProviders]); - // Update display model name when current model changes useEffect(() => { - (async () => { - const displayName = await getCurrentModelDisplayName(); - setDisplayModelName(displayName); - })(); - }, [currentModel, getCurrentModelDisplayName]); + if (!currentModel) return; + setDisplayModelName(getModelDisplayName(currentModel)); + }, [currentModel]); return (
diff --git a/ui/desktop/src/components/settings/models/subcomponents/SwitchModelModal.tsx b/ui/desktop/src/components/settings/models/subcomponents/SwitchModelModal.tsx index 477167f6..79af485e 100644 --- a/ui/desktop/src/components/settings/models/subcomponents/SwitchModelModal.tsx +++ b/ui/desktop/src/components/settings/models/subcomponents/SwitchModelModal.tsx @@ -85,14 +85,8 @@ export const SwitchModelModal = ({ const [providerOptions, setProviderOptions] = useState<{ value: string; label: string }[]>([]); type ModelOption = { value: string; label: string; provider: string; isDisabled?: boolean }; const [modelOptions, setModelOptions] = useState<{ options: ModelOption[] }[]>([]); - const [provider, setProvider] = useState( - initialProvider || currentProvider || null - ); - // Only use currentModel if we're not switching to a different provider - // Otherwise, let the auto-select logic pick an appropriate model for the new provider - const [model, setModel] = useState( - initialProvider && initialProvider !== currentProvider ? '' : currentModel || '' - ); + const [provider, setProvider] = useState(initialProvider || currentProvider || null); + const [model, setModel] = useState(currentModel || ''); const [isCustomModel, setIsCustomModel] = useState(false); const [validationErrors, setValidationErrors] = useState({ provider: '', @@ -172,12 +166,10 @@ export const SwitchModelModal = ({ } await changeModel(sessionId, modelObj); + onModelSelected?.(modelObj.name); trackModelChanged(modelObj.provider || '', modelObj.name); - if (onModelSelected) { - onModelSelected(modelObj.name); - } onClose(); } }; @@ -212,8 +204,7 @@ export const SwitchModelModal = ({ // Load providers for manual model selection (async () => { try { - // Force refresh if initialProvider is set (OAuth flow needs fresh data) - const providersResponse = await getProviders(!!initialProvider); + const providersResponse = await getProviders(false); const activeProviders = providersResponse.filter((provider) => provider.is_configured); // Create provider options and add "Use other provider" option setProviderOptions([ @@ -282,7 +273,7 @@ export const SwitchModelModal = ({ setLoadingModels(false); } })(); - }, [getProviders, usePredefinedModels, read, initialProvider]); + }, [getProviders, usePredefinedModels, read]); const filteredModelOptions = provider ? modelOptions.filter((group) => group.options[0]?.provider === provider)