diff --git a/ui/desktop/src/App.test.tsx b/ui/desktop/src/App.test.tsx index d594441eda..f45e44588c 100644 --- a/ui/desktop/src/App.test.tsx +++ b/ui/desktop/src/App.test.tsx @@ -87,6 +87,8 @@ vi.mock('./components/ModelAndProviderContext', () => ({ provider: null, model: null, getCurrentModelAndProvider: vi.fn(), + getFallbackModelAndProvider: vi.fn().mockResolvedValue({ provider: '', model: '' }), + refreshCurrentModelAndProvider: vi.fn().mockResolvedValue(undefined), setCurrentModelAndProvider: vi.fn(), }), })); @@ -205,8 +207,6 @@ describe('App Component - Brand New State', () => { window.location.hash = ''; window.location.search = ''; window.location.pathname = '/'; - window.sessionStorage?.clear?.(); - window.localStorage?.clear?.(); }); afterEach(() => { diff --git a/ui/desktop/src/components/onboarding/OnboardingGuard.tsx b/ui/desktop/src/components/onboarding/OnboardingGuard.tsx index cb7098180a..59e543679a 100644 --- a/ui/desktop/src/components/onboarding/OnboardingGuard.tsx +++ b/ui/desktop/src/components/onboarding/OnboardingGuard.tsx @@ -48,7 +48,7 @@ export default function OnboardingGuard({ children }: OnboardingGuardProps) { const intl = useIntl(); const navigate = useNavigate(); const { read, upsert, getProviders } = useConfig(); - const { refreshCurrentModelAndProvider } = useModelAndProvider(); + const { getFallbackModelAndProvider, refreshCurrentModelAndProvider } = useModelAndProvider(); const [isCheckingProvider, setIsCheckingProvider] = useState(true); const [hasProvider, setHasProvider] = useState(false); @@ -67,7 +67,25 @@ export default function OnboardingGuard({ children }: OnboardingGuardProps) { for (let attempt = 0; attempt <= retries; attempt++) { try { const provider = (await read('GOOSE_PROVIDER', false, { throwOnError: true })) as string | null; - setHasProvider(!!provider?.trim()); + if (provider?.trim()) { + setHasProvider(true); + setIsCheckingProvider(false); + return; + } + + const fallback = await getFallbackModelAndProvider(); + if (fallback.provider?.trim() && fallback.model?.trim()) { + const configuredProvider = (await read('GOOSE_PROVIDER', false)) as string | null; + const configuredModel = (await read('GOOSE_MODEL', false)) as string | null; + if (configuredProvider?.trim() && configuredModel?.trim()) { + await refreshCurrentModelAndProvider(); + setHasProvider(true); + setIsCheckingProvider(false); + return; + } + } + + setHasProvider(false); setIsCheckingProvider(false); return; } catch (error) {