fix: provider-scoped snapshot refresh with self-healing invalidation

This commit is contained in:
Mohamed Boudra
2026-04-21 13:14:59 +07:00
parent b01e274cce
commit 3832053736
6 changed files with 236 additions and 31 deletions

View File

@@ -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<string, AgentModelDefinition[]>();
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}
/>

View File

@@ -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();

View File

@@ -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<typeof renderProvidersSnapshotHook>["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<void> {
await waitFor(() => {
expect(mockClient.getProvidersSnapshot).toHaveBeenCalledTimes(count);
});
}
async function waitForSnapshotEntries(
result: HookResult,
entries: ProviderSnapshotEntry[],
): Promise<void> {
await waitFor(() => {
expect(result.current.entries).toEqual(entries);
});
}
async function emitProvidersSnapshotUpdate(entries: ProviderSnapshotEntry[]): Promise<void> {
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<void> {
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();
});
});

View File

@@ -50,7 +50,7 @@ interface UseProvidersSnapshotResult {
error: string | null;
supportsSnapshot: boolean;
refresh: (providers?: AgentProvider[]) => Promise<void>;
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,

View File

@@ -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<AgentModelDefinition[]>>()
.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<AgentModelDefinition[]>();
const { registry } = createRegistry([

View File

@@ -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);
}