diff --git a/packages/app/src/components/agent-status-bar.tsx b/packages/app/src/components/agent-status-bar.tsx index a531b0af5..0f40f0046 100644 --- a/packages/app/src/components/agent-status-bar.tsx +++ b/packages/app/src/components/agent-status-bar.tsx @@ -880,15 +880,15 @@ export const AgentStatusBar = memo(function AgentStatusBar({ refetchIfStale: refetchSnapshotIfStale, } = useProvidersSnapshot(serverId, agent?.cwd); - const snapshotModels = useMemo(() => { + const snapshotSelectedEntry = useMemo(() => { if (!snapshotEntries || !agent?.provider) { return null; } - const entry = snapshotEntries.find((e) => e.provider === agent.provider); - return entry?.models ?? null; + return snapshotEntries.find((e) => e.provider === agent.provider) ?? null; }, [snapshotEntries, agent?.provider]); - const models = snapshotModels; + const models = snapshotSelectedEntry?.models ?? null; + const selectedProviderIsLoading = snapshotSelectedEntry?.status === "loading"; const agentProviderDefinitions = useMemo(() => { const definition = agent?.provider @@ -899,11 +899,11 @@ export const AgentStatusBar = memo(function AgentStatusBar({ const agentProviderModels = useMemo(() => { const map = new Map(); - if (agent?.provider && snapshotModels) { - map.set(agent.provider, snapshotModels); + if (agent?.provider && models) { + map.set(agent.provider, models); } return map; - }, [agent?.provider, snapshotModels]); + }, [agent?.provider, models]); const displayMode = availableModes.find((mode) => mode.id === agent?.currentModeId)?.label || @@ -1047,8 +1047,8 @@ export const AgentStatusBar = memo(function AgentStatusBar({ toast.error(toErrorMessage(error)); }); }} - isModelLoading={snapshotIsLoading} - onModelSelectorOpen={refetchSnapshotIfStale} + isModelLoading={snapshotIsLoading || selectedProviderIsLoading} + onModelSelectorOpen={() => refetchSnapshotIfStale(agent?.provider)} onDropdownClose={onDropdownClose} disabled={!client} /> diff --git a/packages/app/src/hooks/use-agent-form-state.ts b/packages/app/src/hooks/use-agent-form-state.ts index 6e6c2276b..d2bbbeb2b 100644 --- a/packages/app/src/hooks/use-agent-form-state.ts +++ b/packages/app/src/hooks/use-agent-form-state.ts @@ -446,6 +446,7 @@ export function useAgentFormState(options: UseAgentFormStateOptions = {}): UseAg [formState.provider, snapshotEntries], ); const snapshotSelectedProviderModels = snapshotSelectedEntry?.models ?? null; + const selectedProviderIsLoading = snapshotSelectedEntry?.status === "loading"; const snapshotSelectedProviderModes = snapshotSelectedEntry?.modes ?? (formState.provider ? snapshotProviderDefinitionMap.get(formState.provider)?.modes : []) ?? @@ -456,7 +457,7 @@ export function useAgentFormState(options: UseAgentFormStateOptions = {}): UseAg const allProviderModels = snapshotAllProviderModels; const availableModels = snapshotSelectedProviderModels; const modeOptions = snapshotSelectedProviderModes; - const isAllModelsLoading = snapshotIsLoading; + const isAllModelsLoading = snapshotIsLoading || selectedProviderIsLoading; // Combine initialValues with initialServerId for resolution const combinedInitialValues = useMemo((): FormInitialValues | undefined => { @@ -669,7 +670,7 @@ export function useAgentFormState(options: UseAgentFormStateOptions = {}): UseAg }, [refreshSnapshot]); const refetchProviderModelsIfStale = useCallback(() => { - refetchSnapshotIfStale(); + refetchSnapshotIfStale(formStateRef.current.provider); }, [refetchSnapshotIfStale]); const persistFormPreferences = useCallback(async () => { @@ -712,7 +713,7 @@ export function useAgentFormState(options: UseAgentFormStateOptions = {}): UseAg const effectiveModel = resolveEffectiveModel(availableModels, formState.model); const resolvedModelId = effectiveModel?.id ?? formState.model; const availableThinkingOptions = effectiveModel?.thinkingOptions ?? []; - const isModelLoading = snapshotIsLoading; + const isModelLoading = snapshotIsLoading || selectedProviderIsLoading; const modelError = snapshotError; const workingDirIsEmpty = !formState.workingDir.trim(); diff --git a/packages/app/src/hooks/use-providers-snapshot.test.ts b/packages/app/src/hooks/use-providers-snapshot.test.ts index c40684925..eb5183795 100644 --- a/packages/app/src/hooks/use-providers-snapshot.test.ts +++ b/packages/app/src/hooks/use-providers-snapshot.test.ts @@ -7,6 +7,7 @@ import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; import { afterEach, describe, expect, it, vi } from "vitest"; import type { DaemonClient } from "@server/client/daemon-client"; +import type { ProviderSnapshotEntry } from "@server/server/agent/agent-sdk-types"; import { useSessionStore } from "@/stores/session-store"; import { providersSnapshotQueryKey, @@ -14,11 +15,31 @@ import { useProvidersSnapshot, } from "./use-providers-snapshot"; -const { mockClient, mockRuntime } = vi.hoisted(() => { +type ProviderSnapshotUpdateMessage = { + type: "providers_snapshot_update"; + payload: { + cwd: string; + entries: ProviderSnapshotEntry[]; + generatedAt: string; + }; +}; +type ProviderSnapshotUpdateListener = (message: ProviderSnapshotUpdateMessage) => void; +type ProvidersSnapshot = { + entries: ProviderSnapshotEntry[]; + generatedAt: string; + requestId: string; +}; +type HookResult = ReturnType["result"]; + +const { mockClient, mockRuntime, snapshotUpdateListeners } = vi.hoisted(() => { + const snapshotUpdateListeners: ProviderSnapshotUpdateListener[] = []; const mockClient = { getProvidersSnapshot: vi.fn(), refreshProvidersSnapshot: vi.fn(), - on: vi.fn(() => () => {}), + on: vi.fn((_event: string, listener: ProviderSnapshotUpdateListener) => { + snapshotUpdateListeners.push(listener); + return () => {}; + }), }; return { mockClient, @@ -26,6 +47,7 @@ const { mockClient, mockRuntime } = vi.hoisted(() => { client: mockClient, isConnected: true, }, + snapshotUpdateListeners, }; }); @@ -66,11 +88,70 @@ function renderProvidersSnapshotHook(cwd?: string | null) { return renderHook(() => useProvidersSnapshot(serverId, cwd), { wrapper }); } +const readyCodexModel = { provider: "codex", id: "gpt-5.4", label: "GPT-5.4" } as const; + +function providersSnapshot(entries: ProviderSnapshotEntry[]): ProvidersSnapshot { + return { + entries, + generatedAt: "2026-01-01T00:00:00.000Z", + requestId: "snapshot", + }; +} + +function codexEntry( + status: ProviderSnapshotEntry["status"], + models?: ProviderSnapshotEntry["models"], +): ProviderSnapshotEntry { + return { + provider: "codex", + status, + ...(models ? { models } : {}), + }; +} + +async function waitForSnapshotReads(count: number): Promise { + await waitFor(() => { + expect(mockClient.getProvidersSnapshot).toHaveBeenCalledTimes(count); + }); +} + +async function waitForSnapshotEntries( + result: HookResult, + entries: ProviderSnapshotEntry[], +): Promise { + await waitFor(() => { + expect(result.current.entries).toEqual(entries); + }); +} + +async function emitProvidersSnapshotUpdate(entries: ProviderSnapshotEntry[]): Promise { + const listener = snapshotUpdateListeners.at(-1); + expect(listener).toBeDefined(); + + await act(async () => { + listener?.({ + type: "providers_snapshot_update", + payload: { + cwd: "/repo", + entries, + generatedAt: "2026-01-01T00:00:01.000Z", + }, + }); + }); +} + +async function openSelectorForSelectedProvider(result: HookResult): Promise { + await act(async () => { + result.current.refetchIfStale("codex"); + }); +} + afterEach(() => { act(() => { useSessionStore.getState().clearSession(serverId); }); vi.clearAllMocks(); + snapshotUpdateListeners.length = 0; }); describe("providers snapshot hook cache scope", () => { @@ -95,11 +176,7 @@ describe("providers snapshot hook cache scope", () => { it("sends no cwd for settings snapshot loads and refreshes", async () => { enableProvidersSnapshot(); - mockClient.getProvidersSnapshot.mockResolvedValue({ - entries: [], - generatedAt: "2026-01-01T00:00:00.000Z", - requestId: "settings", - }); + mockClient.getProvidersSnapshot.mockResolvedValue(providersSnapshot([])); mockClient.refreshProvidersSnapshot.mockResolvedValue({ acknowledged: true, requestId: "settings-refresh", @@ -121,11 +198,7 @@ describe("providers snapshot hook cache scope", () => { it("sends cwd for workspace snapshot loads and refreshes", async () => { enableProvidersSnapshot(); - mockClient.getProvidersSnapshot.mockResolvedValue({ - entries: [], - generatedAt: "2026-01-01T00:00:00.000Z", - requestId: "workspace", - }); + mockClient.getProvidersSnapshot.mockResolvedValue(providersSnapshot([])); mockClient.refreshProvidersSnapshot.mockResolvedValue({ acknowledged: true, requestId: "workspace-refresh", @@ -147,4 +220,58 @@ describe("providers snapshot hook cache scope", () => { }); expect(mockClient.getProvidersSnapshot).toHaveBeenLastCalledWith({ cwd: "/repo" }); }); + + it("refetches loading snapshot updates through the read path but ignores empty updates", async () => { + enableProvidersSnapshot(); + mockClient.getProvidersSnapshot + .mockResolvedValueOnce(providersSnapshot([codexEntry("ready", [])])) + .mockResolvedValueOnce(providersSnapshot([codexEntry("ready", [readyCodexModel])])); + + renderProvidersSnapshotHook("/repo"); + + await waitForSnapshotReads(1); + await emitProvidersSnapshotUpdate([]); + + expect(mockClient.getProvidersSnapshot).toHaveBeenCalledTimes(1); + + await emitProvidersSnapshotUpdate([codexEntry("loading")]); + await waitForSnapshotReads(2); + + expect(mockClient.getProvidersSnapshot).toHaveBeenLastCalledWith({ cwd: "/repo" }); + expect(mockClient.refreshProvidersSnapshot).not.toHaveBeenCalled(); + }); + + it.each([ + { name: "missing", entries: [] }, + { name: "loading", entries: [codexEntry("loading")] }, + ])("ensures a selected provider snapshot on selector open when it is $name", async ({ + entries, + }) => { + enableProvidersSnapshot(); + mockClient.getProvidersSnapshot + .mockResolvedValueOnce(providersSnapshot(entries)) + .mockResolvedValueOnce(providersSnapshot([codexEntry("ready", [readyCodexModel])])); + + const { result } = renderProvidersSnapshotHook("/repo"); + + await waitForSnapshotEntries(result, entries); + await openSelectorForSelectedProvider(result); + await waitForSnapshotReads(2); + + expect(mockClient.getProvidersSnapshot).toHaveBeenLastCalledWith({ cwd: "/repo" }); + expect(mockClient.refreshProvidersSnapshot).not.toHaveBeenCalled(); + }); + + it("does not ensure a selected provider snapshot on selector open when the provider is ready with no models", async () => { + enableProvidersSnapshot(); + mockClient.getProvidersSnapshot.mockResolvedValue(providersSnapshot([codexEntry("ready", [])])); + + const { result } = renderProvidersSnapshotHook("/repo"); + + await waitForSnapshotEntries(result, [codexEntry("ready", [])]); + await openSelectorForSelectedProvider(result); + + expect(mockClient.getProvidersSnapshot).toHaveBeenCalledTimes(1); + expect(mockClient.refreshProvidersSnapshot).not.toHaveBeenCalled(); + }); }); diff --git a/packages/app/src/hooks/use-providers-snapshot.ts b/packages/app/src/hooks/use-providers-snapshot.ts index bf2b62710..a1d34a244 100644 --- a/packages/app/src/hooks/use-providers-snapshot.ts +++ b/packages/app/src/hooks/use-providers-snapshot.ts @@ -50,7 +50,7 @@ interface UseProvidersSnapshotResult { error: string | null; supportsSnapshot: boolean; refresh: (providers?: AgentProvider[]) => Promise; - refetchIfStale: () => void; + refetchIfStale: (selectedProvider?: AgentProvider | null) => void; } interface UseProvidersSnapshotOptions { @@ -118,6 +118,14 @@ export function useProvidersSnapshot( generatedAt: message.payload.generatedAt, requestId: "providers_snapshot_update", }); + const shouldRefetch = message.payload.entries.some((entry) => entry.status === "loading"); + if (shouldRefetch) { + void queryClient.invalidateQueries({ + queryKey, + exact: true, + refetchType: "active", + }); + } }); }, [ client, @@ -143,9 +151,26 @@ export function useProvidersSnapshot( [client, normalizedCwd, queryClient, queryKey, refreshSnapshot], ); - const refetchIfStale = useCallback(() => { - void queryClient.refetchQueries({ queryKey, type: "active", stale: true }); - }, [queryClient, queryKey]); + const refetchIfStale = useCallback( + (selectedProvider?: AgentProvider | null) => { + if (!selectedProvider) { + void queryClient.refetchQueries({ queryKey, type: "active", stale: true }); + return; + } + + const selectedEntry = snapshotQuery.data?.entries.find( + (entry) => entry.provider === selectedProvider, + ); + + if (!selectedEntry || selectedEntry.status === "loading") { + void queryClient.refetchQueries({ queryKey, type: "active" }); + return; + } + + void queryClient.refetchQueries({ queryKey, type: "active", stale: true }); + }, + [queryClient, queryKey, snapshotQuery.data?.entries], + ); return { entries: snapshotQuery.data?.entries ?? undefined, diff --git a/packages/server/src/server/agent/provider-snapshot-manager.test.ts b/packages/server/src/server/agent/provider-snapshot-manager.test.ts index 8dd36584b..bc7b247bf 100644 --- a/packages/server/src/server/agent/provider-snapshot-manager.test.ts +++ b/packages/server/src/server/agent/provider-snapshot-manager.test.ts @@ -588,6 +588,50 @@ describe("ProviderSnapshotManager", () => { manager.destroy(); }); + test("settings refresh invalidation self-heals workspace snapshots through the next read without force", async () => { + const fetchModels = vi + .fn<(cwd: string, force: boolean) => Promise>() + .mockImplementation(async (cwd) => [createModel("codex", cwd)]); + const { registry } = createRegistry([ + createMockProvider({ + provider: "codex", + fetchModels: async (cwd, force) => fetchModels(cwd, force), + fetchModes: async () => [createMode("auto")], + }), + ]); + const manager = new ProviderSnapshotManager(registry, createTestLogger()); + + manager.getSnapshot(projectCwd); + + await vi.waitFor(() => { + expect(getProviderEntry(manager.getSnapshot(projectCwd), "codex")?.status).toBe("ready"); + }); + + await manager.refreshSettingsSnapshot({ providers: ["codex"] }); + + expect(fetchModels.mock.calls.map(([cwd]) => cwd)).toEqual([projectCwd, homedir()]); + expect(fetchModels.mock.calls.map(([, force]) => force)).toEqual([false, true]); + + const invalidatedSnapshot = manager.getSnapshot(projectCwd); + + expect(getProviderEntry(invalidatedSnapshot, "codex")).toMatchObject({ + provider: "codex", + status: "loading", + }); + + await vi.waitFor(() => { + expect(fetchModels).toHaveBeenCalledTimes(3); + }); + + expect(fetchModels.mock.calls[2]).toEqual([projectCwd, false]); + + await vi.waitFor(() => { + expect(getProviderEntry(manager.getSnapshot(projectCwd), "codex")?.status).toBe("ready"); + }); + + manager.destroy(); + }); + test("refresh marks a slow provider as error after the timeout", async () => { const fetchModels = deferred(); const { registry } = createRegistry([ diff --git a/packages/server/src/server/agent/provider-snapshot-manager.ts b/packages/server/src/server/agent/provider-snapshot-manager.ts index f3ad03aaf..0d71c2c51 100644 --- a/packages/server/src/server/agent/provider-snapshot-manager.ts +++ b/packages/server/src/server/agent/provider-snapshot-manager.ts @@ -63,6 +63,13 @@ export class ProviderSnapshotManager { this.resetSnapshotToLoading(resolvedCwd, missingProviders); void this.warmUp(resolvedCwd, missingProviders); } + const providerLoads = this.providerLoads.get(resolvedCwd); + const loadingProviders = Array.from(entries.values()) + .filter((entry) => entry.status === "loading" && !providerLoads?.has(entry.provider)) + .map((entry) => entry.provider); + if (loadingProviders.length > 0) { + void this.warmUp(resolvedCwd, loadingProviders); + } if (this.shouldRevalidate(resolvedCwd)) { void this.warmUp(resolvedCwd); } @@ -392,10 +399,10 @@ export class ProviderSnapshotManager { } if (!options.providers) { - this.snapshots.delete(cwd); + this.resetSnapshotToLoading(cwd); this.lastCheckedAts.delete(cwd); this.providerLoads.delete(cwd); - this.events.emit("change", [], cwd); + this.emitChange(cwd); continue; } @@ -405,9 +412,10 @@ export class ProviderSnapshotManager { } let changed = false; for (const provider of options.providers) { - changed = snapshot.delete(provider) || changed; + changed = snapshot.has(provider) || changed; this.providerLoads.get(cwd)?.delete(provider); } + this.resetSnapshotToLoading(cwd, options.providers); if (this.providerLoads.get(cwd)?.size === 0) { this.providerLoads.delete(cwd); }