diff --git a/LAWS/CHAT.md b/LAWS/CHAT.md index 917e91747..6687c4345 100644 --- a/LAWS/CHAT.md +++ b/LAWS/CHAT.md @@ -37,3 +37,9 @@ - A session's subagent activity MUST appear in the chat transcript with the subagent identity when known. - A session's subagent activity MUST appear in the chat transcript with the delegated task when known. + +## Session configuration + +- A session’s provider MUST support its model, and its harness MUST support that provider. +- A session MUST have exactly one effective configuration. +- Berd MUST show the configuration that the session uses. diff --git a/src/app/AppShell.berdctl.test.tsx b/src/app/AppShell.berdctl.test.tsx index f85f08521..8cc92f52b 100644 --- a/src/app/AppShell.berdctl.test.tsx +++ b/src/app/AppShell.berdctl.test.tsx @@ -80,6 +80,7 @@ vi.mock("@/app/views/NavigationPanesView", () => ({ })); vi.mock("@/shared/api/acp", () => ({ + reserveAcpSessionConfiguration: () => ({ sequence: 0, clear: () => {} }), acpCreateSession: (...args: unknown[]) => mockAcpCreateSession(...args), acpListSessionsPage: (...args: unknown[]) => mockAcpListSessionsPage(...args), acpLoadSession: (...args: unknown[]) => mockAcpLoadSession(...args), diff --git a/src/app/AppShell.navigation.test.tsx b/src/app/AppShell.navigation.test.tsx index 095c4cb48..abb34d596 100644 --- a/src/app/AppShell.navigation.test.tsx +++ b/src/app/AppShell.navigation.test.tsx @@ -317,6 +317,19 @@ function seedProviderModels( ], ]), ); + // Simulate a successful live inventory response: seeding a provider's + // display candidates alone is advisory and never establishes proof. + useProviderModelCacheStore.setState((state) => { + const providers = new Map(state.providers); + const existing = providers.get(providerId); + if (existing) { + providers.set(providerId, { + ...existing, + provenModelIds: models.map((model) => model.id), + }); + } + return { providers }; + }); } vi.mock("@/shared/profile/buildProfile", () => ({ @@ -440,6 +453,7 @@ vi.mock("@/shared/api/acp", () => ({ acpListSessionsPage: (...args: unknown[]) => mockAcpListSessionsPage(...args), acpLoadSession: (...args: unknown[]) => mockAcpLoadSession(...args), discoverAcpProviders: vi.fn().mockResolvedValue([]), + reserveAcpSessionConfiguration: () => ({ sequence: 0, clear: () => {} }), })); vi.mock("@/shared/api/acpApi", () => ({ @@ -1412,6 +1426,7 @@ describe("AppShell global navigation", () => { "openai", "~/goose artifacts", expect.any(Object), + expect.objectContaining({ clear: expect.any(Function) }), ); }); expect( @@ -3435,6 +3450,7 @@ describe("AppShell global navigation", () => { "codex-acp", "~/goose artifacts", expect.objectContaining({ modelId: "gpt-5.4-mini" }), + expect.objectContaining({ clear: expect.any(Function) }), ); }); act(() => pendingPrepare.resolve({})); @@ -3497,6 +3513,7 @@ describe("AppShell global navigation", () => { "databricks_v2", "~/goose artifacts", expect.objectContaining({ modelId: "goose-gpt-5-5" }), + expect.objectContaining({ clear: expect.any(Function) }), ); }); await waitFor(() => { diff --git a/src/app/AppShell.tsx b/src/app/AppShell.tsx index 4606476c0..662acc4c7 100644 --- a/src/app/AppShell.tsx +++ b/src/app/AppShell.tsx @@ -2700,10 +2700,16 @@ export function AppShell({ const persona = agentState.personas.find( (candidate) => candidate.id === agentId, ); - const cachedModels = [ - ...useProviderModelCacheStore.getState().providers, - ].flatMap(([providerId, entry]) => - entry.models.map((model) => ({ + const modelCache = useProviderModelCacheStore.getState(); + const cachedModels = [...modelCache.providers].flatMap( + ([providerId, entry]) => + entry.models.map((model) => ({ + ...model, + providerId: model.providerId ?? providerId, + })), + ); + const provenModels = [...modelCache.providers].flatMap(([providerId]) => + modelCache.getProvenModelsForProvider(providerId).map((model) => ({ ...model, providerId: model.providerId ?? providerId, })), @@ -2711,6 +2717,11 @@ export function AppShell({ const executionTarget = personaExecutionTarget(persona, { providers: agentState.providers, models: cachedModels, + getProvenModelsForHarness: () => provenModels, + isModelInventoryAuthoritative: (providerId) => + useProviderModelCacheStore + .getState() + .isModelInventoryAuthoritative(providerId), catalogEntries: getProviderCatalog(), }); diff --git a/src/app/lib/chatRuntimeStartup.test.ts b/src/app/lib/chatRuntimeStartup.test.ts index 3bf339aff..e5b8dc1e2 100644 --- a/src/app/lib/chatRuntimeStartup.test.ts +++ b/src/app/lib/chatRuntimeStartup.test.ts @@ -5,6 +5,18 @@ const mockLoadPersistedMessageQueues = vi.hoisted(() => ); const mockGetClient = vi.hoisted(() => vi.fn<() => Promise>()); const mockRefreshAllModelProviders = vi.hoisted(() => vi.fn()); +const mockMigratePersonaTargetIfUnchanged = vi.hoisted(() => vi.fn()); +const mockAgentState = vi.hoisted(() => ({ + personas: [] as Array>, + providers: [] as Array<{ id: string; label?: string }>, +})); +const mockModelCacheState = vi.hoisted(() => ({ + providers: new Map< + string, + { models: Array>; provenModelIds?: string[] } + >(), + runtimeManagedProviderIds: new Set(), +})); // The latch under test wraps startChatRuntime, whose body touches most of the // startup module graph. Everything it reaches is stubbed inert (resolved, @@ -18,6 +30,7 @@ vi.mock("@/features/agents/stores/agentStore", () => ({ setProviders: () => {}, setPersonas: () => {}, setPersonasLoading: () => {}, + ...mockAgentState, }), }, })); @@ -48,12 +61,15 @@ vi.mock("@/features/providers/runtimeProviderConstraints", () => ({ })); vi.mock("@/features/providers/modelCacheRefresh", () => ({ - getModelCacheRefreshProviderIds: () => [], + getModelCacheRefreshProviderIds: () => ["claude-acp"], })); vi.mock("@/features/providers/providerCatalog", () => ({ + canonicalProviderCatalogIdFromEntries: (_entries: unknown, id: string) => id, getModelProviders: () => [], getProviderCatalog: () => [], + resolveAgentProviderCatalogIdStrictFromEntries: () => null, + resolveModelProviderCatalogIdStrictFromEntries: () => null, })); vi.mock("@/features/providers/runtimeProviderConfig", () => ({ @@ -110,8 +126,7 @@ vi.mock("@/features/providers/stores/defaultProviderReadinessStore", () => ({ vi.mock("@/features/providers/stores/providerModelCacheStore", () => ({ useProviderModelCacheStore: { getState: () => ({ - providers: new Map(), - runtimeManagedProviderIds: new Set(), + ...mockModelCacheState, loadPersisted: () => {}, refreshAllModelProviders: (...args: unknown[]) => mockRefreshAllModelProviders(...args), @@ -166,8 +181,9 @@ vi.mock("@/shared/api/distro", () => ({ })); vi.mock("@/shared/api/agents", () => ({ - listPersonas: async () => [], - migratePersonaTargetIfUnchanged: async () => null, + listPersonas: async () => mockAgentState.personas, + migratePersonaTargetIfUnchanged: (...args: unknown[]) => + mockMigratePersonaTargetIfUnchanged(...args), })); function deferred() { @@ -190,6 +206,12 @@ describe("runChatRuntimeStartup", () => { mockGetClient.mockResolvedValue({}); mockRefreshAllModelProviders.mockReset(); mockRefreshAllModelProviders.mockResolvedValue(undefined); + mockMigratePersonaTargetIfUnchanged.mockReset(); + mockMigratePersonaTargetIfUnchanged.mockResolvedValue(null); + mockAgentState.personas = []; + mockAgentState.providers = []; + mockModelCacheState.providers = new Map(); + mockModelCacheState.runtimeManagedProviderIds = new Set(); }); it("collapses concurrent callers onto one startup run", async () => { @@ -221,6 +243,64 @@ describe("runChatRuntimeStartup", () => { inventoryRefresh.resolve(); }); + it("does not migrate a runtime-managed configuration seed before live discovery", async () => { + mockAgentState.providers = [{ id: "claude-acp", label: "Claude Code" }]; + mockAgentState.personas = [ + { + id: "persona-1", + displayName: "Configured Claude", + systemPrompt: "Help.", + provider: "claude-acp", + modelProviderId: "claude-acp", + model: "configured-model", + isBuiltin: false, + writable: true, + }, + ]; + mockModelCacheState.runtimeManagedProviderIds = new Set(["claude-acp"]); + mockModelCacheState.providers = new Map([ + [ + "claude-acp", + { + models: [{ id: "configured-model", providerId: "claude-acp" }], + }, + ], + ]); + + const { runChatRuntimeStartup } = await import("./chatRuntimeStartup"); + await runChatRuntimeStartup(); + + expect(mockMigratePersonaTargetIfUnchanged).not.toHaveBeenCalled(); + }); + + it("leaves an authoritative unsupported persona persisted for explicit repair", async () => { + mockAgentState.providers = [{ id: "claude-acp", label: "Claude Code" }]; + mockAgentState.personas = [ + { + id: "persona-1", + displayName: "Legacy Claude", + systemPrompt: "Help.", + provider: "claude-acp", + modelProviderId: "openai", + model: "gpt-5", + isBuiltin: false, + writable: true, + }, + ]; + mockModelCacheState.providers = new Map([ + ["claude-acp", { models: [], provenModelIds: [] }], + ]); + mockRefreshAllModelProviders.mockImplementation(async () => { + mockModelCacheState.providers = new Map([ + ["claude-acp", { models: [], provenModelIds: [] }], + ]); + }); + + const { runChatRuntimeStartup } = await import("./chatRuntimeStartup"); + await runChatRuntimeStartup(); + expect(mockMigratePersonaTargetIfUnchanged).not.toHaveBeenCalled(); + }); + it("stays latched after a successful run", async () => { const { runChatRuntimeStartup } = await import("./chatRuntimeStartup"); const first = runChatRuntimeStartup(); diff --git a/src/app/lib/chatRuntimeStartup.ts b/src/app/lib/chatRuntimeStartup.ts index 72c74573b..f42469b18 100644 --- a/src/app/lib/chatRuntimeStartup.ts +++ b/src/app/lib/chatRuntimeStartup.ts @@ -259,15 +259,19 @@ async function startChatRuntime( const cachedModels = [...modelState.providers].flatMap( ([providerId, entry]) => authoritativeProviderIds.has(providerId) - ? entry.models.map((model) => ({ - ...model, - providerId: model.providerId ?? providerId, - })) + ? entry.models + .filter((model) => entry.provenModelIds?.includes(model.id)) + .map((model) => ({ + ...model, + providerId: model.providerId ?? providerId, + })) : [], ); const targetContext = { providers: useAgentStore.getState().providers, models: cachedModels, + isModelInventoryAuthoritative: (providerId: string) => + authoritativeProviderIds.has(providerId), catalogEntries: getProviderCatalog(), }; const personas = useAgentStore.getState().personas; @@ -311,13 +315,14 @@ async function startChatRuntime( ); await modelCacheStore.refreshAllModelProviders(refreshProviderIds); const modelState = useProviderModelCacheStore.getState(); - return new Set([ - ...modelState.runtimeManagedProviderIds, - ...refreshProviderIds.filter((providerId) => { + return new Set( + refreshProviderIds.filter((providerId) => { const entry = modelState.providers.get(providerId); - return entry != null && !entry.error; + return ( + entry != null && !entry.error && entry.provenModelIds !== undefined + ); }), - ]); + ); }; const loadSessionState = async () => { diff --git a/src/app/views/__tests__/NavigationPanesView.test.tsx b/src/app/views/__tests__/NavigationPanesView.test.tsx index d40ffed73..36a84cb16 100644 --- a/src/app/views/__tests__/NavigationPanesView.test.tsx +++ b/src/app/views/__tests__/NavigationPanesView.test.tsx @@ -315,6 +315,7 @@ vi.mock("@/features/chat/stores/chatStore", () => ({ })); vi.mock("@/shared/api/acp", () => ({ + reserveAcpSessionConfiguration: () => ({ sequence: 0, clear: () => {} }), acpSearchSessions: (...args: unknown[]) => mockAcpSearchSessions(...args), })); diff --git a/src/features/agents/lib/__tests__/personaExecutionTarget.test.ts b/src/features/agents/lib/__tests__/personaExecutionTarget.test.ts index d5e05acc9..fa3db52ba 100644 --- a/src/features/agents/lib/__tests__/personaExecutionTarget.test.ts +++ b/src/features/agents/lib/__tests__/personaExecutionTarget.test.ts @@ -1,5 +1,6 @@ import { describe, expect, it } from "vitest"; import type { ProviderCatalogEntry } from "@/shared/types/providers"; +import { gooseServeSelectionFromExecutionTarget } from "@/features/chat/lib/gooseServeExecutionTarget"; import { personaExecutionTarget, personaTargetMigration, @@ -18,12 +19,17 @@ const catalog = (id: string, category: "agent" | "model", aliases?: string[]) => const context = ( models: Array<{ id: string; providerId?: string; displayName?: string }> = [], + authoritativeProviderIds: readonly string[] = [], + provenModels = models, ) => ({ providers: [ { id: "goose", label: "Goose" }, { id: "claude-acp", label: "Claude Code" }, ], models, + getProvenModelsForHarness: () => provenModels, + isModelInventoryAuthoritative: (providerId: string) => + authoritativeProviderIds.includes(providerId), catalogEntries: [ catalog("goose", "agent"), catalog("claude-acp", "agent", ["claude"]), @@ -52,6 +58,200 @@ describe("personaExecutionTarget", () => { }); }); + it("does not treat an advisory-only display candidate as live proof", () => { + const persona = { + provider: "goose", + modelProviderId: "openai", + model: "advisory-only", + }; + const targetContext = context( + [{ id: "advisory-only", providerId: "openai" }], + ["openai"], + [], + ); + + expect(personaExecutionTarget(persona, targetContext)).toBeUndefined(); + expect(personaTargetMigration(persona, targetContext)).toBeNull(); + }); + + it("rejects an agent harness persisted as a Goose model provider", () => { + const persona = { + provider: "goose", + modelProviderId: "claude-acp", + model: "sonnet", + }; + + expect(personaExecutionTarget(persona, context())).toBeUndefined(); + expect(personaTargetMigration(persona, context())).toEqual({ + provider: null, + modelProviderId: null, + model: null, + }); + }); + + it("never materializes an agent provider as a Goose target without a model", () => { + const persona = { provider: "goose", modelProviderId: "claude-acp" }; + + expect(personaExecutionTarget(persona, context())).toEqual({ + harnessId: "goose", + }); + expect(personaTargetMigration(persona, context())).toEqual({ + provider: "goose", + modelProviderId: null, + model: null, + }); + }); + + it.each([ + { + name: "Goose canonical provider with a supported model", + persona: { provider: "goose", modelProviderId: "openai", model: "gpt-5" }, + models: [{ id: "gpt-5", providerId: "openai" }], + authoritativeProviderIds: ["openai"], + target: { + harnessId: "goose", + modelProviderId: "openai", + modelId: "gpt-5", + modelName: "gpt-5", + }, + migration: null, + }, + { + name: "Goose canonical provider with an unsupported model", + persona: { provider: "goose", modelProviderId: "openai", model: "gpt-5" }, + models: [], + authoritativeProviderIds: ["openai"], + target: undefined, + migration: null, + }, + { + name: "external harness with foreign provider and supported model", + persona: { + provider: "claude-acp", + modelProviderId: "openai", + model: "sonnet", + }, + models: [{ id: "sonnet", displayName: "Sonnet" }], + authoritativeProviderIds: ["claude-acp"], + target: { + harnessId: "claude-acp", + modelProviderId: "claude-acp", + modelId: "sonnet", + modelName: "Sonnet", + }, + migration: { + provider: "claude-acp", + modelProviderId: "claude-acp", + model: "sonnet", + }, + }, + { + name: "external harness with foreign provider and unsupported model", + persona: { + provider: "claude-acp", + modelProviderId: "openai", + model: "gpt-5", + }, + models: [], + authoritativeProviderIds: ["claude-acp"], + target: undefined, + migration: null, + }, + { + name: "external harness with unavailable inventory", + persona: { + provider: "claude-acp", + modelProviderId: "openai", + model: "gpt-5", + }, + models: [], + authoritativeProviderIds: [], + target: { + harnessId: "claude-acp", + modelProviderId: "claude-acp", + modelId: "gpt-5", + modelName: "gpt-5", + }, + migration: { + provider: "claude-acp", + modelProviderId: "claude-acp", + model: "gpt-5", + }, + }, + ])("canonicalizes $name across persisted target, migration, and wire selection", ({ + persona, + models, + authoritativeProviderIds, + target, + migration, + }) => { + const targetContext = context(models, authoritativeProviderIds); + + expect(personaExecutionTarget(persona, targetContext)).toEqual(target); + expect(personaTargetMigration(persona, targetContext)).toEqual(migration); + const actualTarget = personaExecutionTarget(persona, targetContext); + if (!target) { + expect(gooseServeSelectionFromExecutionTarget(actualTarget)).toEqual({}); + return; + } + const wireProviderId = + target.harnessId === "goose" ? target.modelProviderId : target.harnessId; + expect(gooseServeSelectionFromExecutionTarget(actualTarget)).toEqual({ + providerId: wireProviderId, + modelId: target.modelId, + modelName: target.modelName, + }); + }); + + it("owns an external harness model provider and repairs legacy display metadata", () => { + const persona = { + provider: "claude-acp", + modelProviderId: "openai", + model: "sonnet", + }; + const target = personaExecutionTarget( + persona, + context([{ id: "sonnet", displayName: "Sonnet" }]), + ); + + expect(target).toEqual({ + harnessId: "claude-acp", + modelProviderId: "claude-acp", + modelId: "sonnet", + modelName: "Sonnet", + }); + expect(gooseServeSelectionFromExecutionTarget(target)).toEqual({ + providerId: "claude-acp", + modelId: "sonnet", + modelName: "Sonnet", + }); + expect( + personaTargetMigration( + persona, + context([{ id: "sonnet", displayName: "Sonnet" }]), + ), + ).toEqual({ + provider: "claude-acp", + modelProviderId: "claude-acp", + model: "sonnet", + }); + }); + + it("returns no target for an unknown persisted harness", () => { + const persona = { + provider: "deleted-harness", + modelProviderId: "openai", + model: "gpt-5", + }; + + expect(personaExecutionTarget(persona, context())).toBeUndefined(); + expect(personaTargetMigration(persona, context())).toEqual({ + provider: null, + modelProviderId: null, + model: null, + }); + }); + it("uses an external harness as the runtime provider boundary", () => { expect( personaExecutionTarget( diff --git a/src/features/agents/lib/personaExecutionTarget.ts b/src/features/agents/lib/personaExecutionTarget.ts index 44c0fa251..73b93c64b 100644 --- a/src/features/agents/lib/personaExecutionTarget.ts +++ b/src/features/agents/lib/personaExecutionTarget.ts @@ -28,6 +28,10 @@ export interface PersonaTargetContext { providers: readonly AvailableHarness[]; models: readonly AvailableModel[]; getModelsForHarness?: (harnessId: string) => readonly AvailableModel[]; + /** Live inventory models, separate from display/advisory candidates. */ + getProvenModelsForHarness?: (harnessId: string) => readonly AvailableModel[]; + /** Whether the model inventory for a provider/harness is authoritative. */ + isModelInventoryAuthoritative?: (providerId: string) => boolean; catalogEntries: ProviderCatalogEntry[]; } @@ -75,17 +79,37 @@ function harnessIdForPersona( ); } +function isAgentProviderId( + providerId: string, + catalogEntries: ProviderCatalogEntry[], +): boolean { + return ( + resolveAgentProviderCatalogIdStrictFromEntries( + catalogEntries, + providerId, + ) !== null + ); +} + function persistedModelProviderId( persona: Pick, harnessId: string, catalogEntries: ProviderCatalogEntry[], ): string | undefined { - if (persona.modelProviderId?.trim()) { - return canonicalModelProviderId(persona.modelProviderId, catalogEntries); + // A non-Goose harness is its own provider boundary. Its persisted model + // provider is display metadata from an older representation, never an + // independent provider that may be sent to Goose. + if (harnessId !== "goose") return harnessId; + + const persistedProviderId = persona.modelProviderId?.trim(); + if ( + persistedProviderId && + !isAgentProviderId(persistedProviderId, catalogEntries) + ) { + return canonicalModelProviderId(persistedProviderId, catalogEntries); } if ( persona.provider?.trim() && - harnessId === "goose" && (INTERNAL_DATABRICKS_KEYS.has(normalizeProviderKey(persona.provider)) || resolveModelProviderCatalogIdStrictFromEntries( catalogEntries, @@ -94,7 +118,7 @@ function persistedModelProviderId( ) { return canonicalModelProviderId(persona.provider, catalogEntries); } - return harnessId === "goose" ? undefined : harnessId; + return undefined; } /** @@ -110,6 +134,8 @@ export function personaExecutionTarget( providers, models, getModelsForHarness, + getProvenModelsForHarness, + isModelInventoryAuthoritative, catalogEntries, }: PersonaTargetContext, ): SessionExecutionTarget | undefined { @@ -121,6 +147,8 @@ export function personaExecutionTarget( if (!harnessId) return undefined; const availableModels = getModelsForHarness?.(harnessId) ?? models; + const provenModels = + getProvenModelsForHarness?.(harnessId) ?? availableModels; const modelId = normalizeConcreteModelId(persona?.model); let modelProviderId = persistedModelProviderId( persona ?? {}, @@ -131,7 +159,7 @@ export function personaExecutionTarget( // Compatibility read until the migration write completes. if (modelId && !modelProviderId && harnessId === "goose") { const matches = new Set( - availableModels.flatMap((model) => + provenModels.flatMap((model) => model.id === modelId && model.providerId ? [canonicalModelProviderId(model.providerId, catalogEntries)] : [], @@ -151,6 +179,21 @@ export function personaExecutionTarget( canonicalModelProviderId(model.providerId, catalogEntries) === modelProviderId), ); + const provenModel = provenModels.find( + (model) => + model.id === modelId && + (!model.providerId || + canonicalModelProviderId(model.providerId, catalogEntries) === + modelProviderId), + ); + const inventoryIsAuthoritative = + isModelInventoryAuthoritative?.(modelProviderId ?? harnessId) ?? false; + + if (modelId && !provenModel && inventoryIsAuthoritative) { + // Preserve the saved configuration for repair, but do not expose an + // incomplete execution target: agents require both provider and model. + return undefined; + } return normalizeSessionExecutionTarget({ harnessId, @@ -190,10 +233,22 @@ export function personaTargetMigration( context.providers, context.catalogEntries, ); + const persistedAgentProviderForGoose = + harnessIdForPersona( + persona.provider, + context.providers, + context.catalogEntries, + ) === "goose" && + Boolean( + persona.modelProviderId && + isAgentProviderId(persona.modelProviderId, context.catalogEntries), + ); // Clear only when the saved data itself proves it cannot form one target. // No inventory match may be a transient availability problem, so preserve // that legacy metadata until a later authoritative refresh can repair it. - return unknownHarness || matchingProviderIds.size > 1 + return unknownHarness || + persistedAgentProviderForGoose || + matchingProviderIds.size > 1 ? { provider: null, modelProviderId: null, model: null } : null; } diff --git a/src/features/berdctl/__tests__/commands/commands.test.ts b/src/features/berdctl/__tests__/commands/commands.test.ts index 0ae4fa33d..b99b14dd8 100644 --- a/src/features/berdctl/__tests__/commands/commands.test.ts +++ b/src/features/berdctl/__tests__/commands/commands.test.ts @@ -74,6 +74,7 @@ const mocks = vi.hoisted(() => ({ })); vi.mock("@/shared/api/acp", () => ({ + reserveAcpSessionConfiguration: () => ({ sequence: 0, clear: () => {} }), acpCreateSession: (...args: unknown[]) => mocks.acpCreateSession(...args), acpDuplicateSession: (...args: unknown[]) => mocks.acpDuplicateSession(...args), @@ -289,6 +290,9 @@ function seedModelCache(cacheKey: string, modelIds: string[]): void { providerId: cacheKey, models: modelIds.map((id) => ({ id, name: id })), fetchedAt: Date.now(), + // Simulate a successful live inventory response: proof is what keeps + // this cache entry from being treated as stale and re-fetched. + provenModelIds: modelIds, }); return { providers }; }); @@ -840,7 +844,9 @@ describe("sessions.create", () => { expect(mocks.acpCreateSession).toHaveBeenCalledWith( "codex-acp", "/resolved/cwd", - expect.objectContaining({ modelId: "gpt-6" }), + expect.objectContaining({ + modelId: "gpt-6", + }), ); }); @@ -1005,7 +1011,8 @@ describe("sessions.send", () => { "session-1", "codex-acp", "/resolved/cwd", - { modelId: "gpt-6" }, + { modelId: "gpt-6", selectionAlreadyResolved: true }, + expect.objectContaining({ clear: expect.any(Function) }), ); expect(controller.openSession).not.toHaveBeenCalled(); @@ -1082,7 +1089,8 @@ describe("sessions.send", () => { "session-1", "codex-acp", "/resolved/cwd", - {}, + { selectionAlreadyResolved: true }, + expect.objectContaining({ clear: expect.any(Function) }), ); expect(getPendingSessionWorkspaceActivation("session-1")).toBeNull(); }); @@ -1131,7 +1139,8 @@ describe("sessions.send", () => { "session-1", "old-provider", "/resolved/cwd", - { modelId: "old-model" }, + { modelId: "old-model", selectionAlreadyResolved: true }, + expect.objectContaining({ clear: expect.any(Function) }), ); await vi.waitFor(() => { expect(mocks.acpSendMessage).toHaveBeenCalledWith( diff --git a/src/features/berdctl/commands/runtime/sessionSend.test.ts b/src/features/berdctl/commands/runtime/sessionSend.test.ts index 864ca273e..99dcc5cc0 100644 --- a/src/features/berdctl/commands/runtime/sessionSend.test.ts +++ b/src/features/berdctl/commands/runtime/sessionSend.test.ts @@ -37,6 +37,7 @@ const mocks = vi.hoisted(() => ({ })); vi.mock("@/shared/api/acp", () => ({ + reserveAcpSessionConfiguration: () => ({ sequence: 0, clear: () => {} }), acpGetSessionInfo: (...args: unknown[]) => mocks.acpGetSessionInfo(...args), acpLoadSession: (...args: unknown[]) => mocks.acpLoadSession(...args), acpPrepareSession: (...args: unknown[]) => { @@ -580,7 +581,8 @@ describe("sendPromptToExistingSessionInBackground", () => { SESSION_ID, "claude-acp", expect.any(String), - { modelId: "claude-sonnet-4" }, + { modelId: "claude-sonnet-4", selectionAlreadyResolved: true }, + expect.objectContaining({ clear: expect.any(Function) }), ); expect(mocks.acpSendMessage).toHaveBeenCalledWith( SESSION_ID, @@ -655,7 +657,8 @@ describe("sendPromptToExistingSessionInBackground", () => { SESSION_ID, UPDATED_TARGET.modelProviderId, "/tmp/project", - { modelId: UPDATED_TARGET.modelId }, + { modelId: UPDATED_TARGET.modelId, selectionAlreadyResolved: true }, + expect.objectContaining({ clear: expect.any(Function) }), ]); expect(mocks.acpSendMessage).toHaveBeenCalledTimes(1); expect( @@ -720,7 +723,8 @@ describe("sendPromptToExistingSessionInBackground", () => { SESSION_ID, INITIAL_TARGET.modelProviderId, "/tmp/project", - { modelId: INITIAL_TARGET.modelId }, + { modelId: INITIAL_TARGET.modelId, selectionAlreadyResolved: true }, + expect.objectContaining({ clear: expect.any(Function) }), ]); expect(mocks.transportProviders).toEqual([INITIAL_TARGET.modelProviderId]); expect( @@ -865,7 +869,8 @@ describe("sendPromptToExistingSessionInBackground", () => { SESSION_ID, INITIAL_TARGET.modelProviderId, "/tmp/project", - { modelId: INITIAL_TARGET.modelId }, + { modelId: INITIAL_TARGET.modelId, selectionAlreadyResolved: true }, + expect.objectContaining({ clear: expect.any(Function) }), ); expect(mocks.acpPrepareSession.mock.calls.at(-1)).toEqual([ SESSION_ID, @@ -873,8 +878,10 @@ describe("sendPromptToExistingSessionInBackground", () => { "/tmp/project", expect.objectContaining({ modelId: UPDATED_TARGET.modelId, + selectionAlreadyResolved: true, requestId: "select-updated-during-prepare", }), + expect.objectContaining({ clear: expect.any(Function) }), ]); expect(mocks.acpSendMessage).toHaveBeenCalledTimes(1); expect(mocks.transportProviders).toEqual([INITIAL_TARGET.modelProviderId]); @@ -915,7 +922,8 @@ describe("sendPromptToExistingSessionInBackground", () => { SESSION_ID, INITIAL_TARGET.modelProviderId, "/tmp/project", - { modelId: INITIAL_TARGET.modelId }, + { modelId: INITIAL_TARGET.modelId, selectionAlreadyResolved: true }, + expect.objectContaining({ clear: expect.any(Function) }), ]); expect(mocks.acpPrepareSession.mock.calls.at(-1)).toEqual([ SESSION_ID, @@ -923,8 +931,10 @@ describe("sendPromptToExistingSessionInBackground", () => { "/tmp/project", expect.objectContaining({ modelId: UPDATED_TARGET.modelId, + selectionAlreadyResolved: true, requestId: "select-updated-during-cwd", }), + expect.objectContaining({ clear: expect.any(Function) }), ]); expect(mocks.acpSendMessage).toHaveBeenCalledTimes(1); expect(mocks.transportProviders).toEqual([INITIAL_TARGET.modelProviderId]); diff --git a/src/features/chat/hooks/__tests__/useChat.attachments.test.ts b/src/features/chat/hooks/__tests__/useChat.attachments.test.ts index 6dde38e17..6019c0904 100644 --- a/src/features/chat/hooks/__tests__/useChat.attachments.test.ts +++ b/src/features/chat/hooks/__tests__/useChat.attachments.test.ts @@ -9,6 +9,7 @@ const mockAcpCancelSession = vi.fn(); const mockAcpPrepareSession = vi.fn(); vi.mock("@/shared/api/acp", () => ({ + reserveAcpSessionConfiguration: () => ({ sequence: 0, clear: () => {} }), acpSendMessage: (...args: unknown[]) => { const result = mockAcpSendMessage(...args); const options = args[2] as diff --git a/src/features/chat/hooks/__tests__/useChat.compaction.test.ts b/src/features/chat/hooks/__tests__/useChat.compaction.test.ts index 03f1d53fc..451125154 100644 --- a/src/features/chat/hooks/__tests__/useChat.compaction.test.ts +++ b/src/features/chat/hooks/__tests__/useChat.compaction.test.ts @@ -12,6 +12,7 @@ const mockAcpSendMessage = vi.fn(); const mockAcpLoadSession = vi.fn(); vi.mock("@/shared/api/acp", () => ({ + reserveAcpSessionConfiguration: () => ({ sequence: 0, clear: () => {} }), acpSendMessage: (...args: unknown[]) => mockAcpSendMessage(...args), acpCancelSession: vi.fn(), acpLoadSession: (...args: unknown[]) => mockAcpLoadSession(...args), diff --git a/src/features/chat/hooks/__tests__/useChat.personaPreparation.test.ts b/src/features/chat/hooks/__tests__/useChat.personaPreparation.test.ts index 37e22ab1f..0c64c6f20 100644 --- a/src/features/chat/hooks/__tests__/useChat.personaPreparation.test.ts +++ b/src/features/chat/hooks/__tests__/useChat.personaPreparation.test.ts @@ -10,6 +10,7 @@ const mockAcpCancelSession = vi.fn(); const mockAcpLoadSession = vi.fn(); vi.mock("@/shared/api/acp", () => ({ + reserveAcpSessionConfiguration: () => ({ sequence: 0, clear: () => {} }), acpSendMessage: (...args: unknown[]) => { const result = mockAcpSendMessage(...args); const options = args[2] as diff --git a/src/features/chat/hooks/__tests__/useChat.skillChips.test.ts b/src/features/chat/hooks/__tests__/useChat.skillChips.test.ts index 6d506431f..3c8762c9a 100644 --- a/src/features/chat/hooks/__tests__/useChat.skillChips.test.ts +++ b/src/features/chat/hooks/__tests__/useChat.skillChips.test.ts @@ -10,6 +10,7 @@ const mockAcpLoadSession = vi.fn(); const mockAcpPrepareSession = vi.fn(); vi.mock("@/shared/api/acp", () => ({ + reserveAcpSessionConfiguration: () => ({ sequence: 0, clear: () => {} }), acpSendMessage: (...args: unknown[]) => { const result = mockAcpSendMessage(...args); const options = args[2] as diff --git a/src/features/chat/hooks/__tests__/useChat.test.ts b/src/features/chat/hooks/__tests__/useChat.test.ts index 3a7ebb280..d70bd0591 100644 --- a/src/features/chat/hooks/__tests__/useChat.test.ts +++ b/src/features/chat/hooks/__tests__/useChat.test.ts @@ -24,6 +24,7 @@ const mockAcpPrepareSession = vi.fn(); let mockAcpDispatches = true; vi.mock("@/shared/api/acp", () => ({ + reserveAcpSessionConfiguration: () => ({ sequence: 0, clear: () => {} }), acpSendMessage: (...args: unknown[]) => { const result = mockAcpSendMessage(...args); const options = args[2] as diff --git a/src/features/chat/hooks/__tests__/useChatSessionController.targetLeaseCompaction.test.ts b/src/features/chat/hooks/__tests__/useChatSessionController.targetLeaseCompaction.test.ts index 71afa9c07..e2984c356 100644 --- a/src/features/chat/hooks/__tests__/useChatSessionController.targetLeaseCompaction.test.ts +++ b/src/features/chat/hooks/__tests__/useChatSessionController.targetLeaseCompaction.test.ts @@ -31,6 +31,7 @@ function deferred() { } vi.mock("@/shared/api/acp", () => ({ + reserveAcpSessionConfiguration: () => ({ sequence: 0, clear: () => {} }), acpPrepareSession: async (...args: unknown[]) => { const result = await mockAcpPrepareSession(...args); preparedProviderBySession.set(args[0] as string, args[1] as string); diff --git a/src/features/chat/hooks/__tests__/useChatSessionController.test.ts b/src/features/chat/hooks/__tests__/useChatSessionController.test.ts index d12c33d75..08bc549c8 100644 --- a/src/features/chat/hooks/__tests__/useChatSessionController.test.ts +++ b/src/features/chat/hooks/__tests__/useChatSessionController.test.ts @@ -115,6 +115,7 @@ function deferred() { } vi.mock("@/shared/api/acp", () => ({ + reserveAcpSessionConfiguration: () => ({ sequence: 0, clear: () => {} }), acpPrepareSession: (...args: unknown[]) => mockAcpPrepareSession(...args), acpSetSessionConfigOption: (...args: unknown[]) => mockAcpSetSessionConfigOption(...args), @@ -255,6 +256,9 @@ vi.mock("../useAgentModelPickerState", () => ({ getModelsForAgent: (agentId: string) => mockPickerState.modelsByAgent.get(agentId) ?? mockPickerState.availableModels, + getProvenModelsForAgent: (agentId: string) => + mockPickerState.modelsByAgent.get(agentId) ?? + mockPickerState.availableModels, isModelInventoryAuthoritative: () => false, modelsLoading: mockPickerState.modelsLoading, modelStatusMessage: mockPickerState.modelStatusMessage, @@ -323,6 +327,7 @@ function expectSessionPreparation({ ...(modelId ? { modelId } : {}), ...(forceConfigRefresh ? { forceConfigRefresh: true } : {}), }), + expect.objectContaining({ clear: expect.any(Function) }), ); } @@ -4808,6 +4813,7 @@ describe("useChatSessionController", () => { "openai", "/tmp/project", expect.objectContaining({ requestId: expect.any(String) }), + expect.objectContaining({ clear: expect.any(Function) }), ); }); diff --git a/src/features/chat/hooks/__tests__/useResolvedAgentModelPicker.test.ts b/src/features/chat/hooks/__tests__/useResolvedAgentModelPicker.test.ts index dd04d0bbc..35321ad43 100644 --- a/src/features/chat/hooks/__tests__/useResolvedAgentModelPicker.test.ts +++ b/src/features/chat/hooks/__tests__/useResolvedAgentModelPicker.test.ts @@ -58,6 +58,7 @@ function renderModelPicker(overrides: Partial = {}) { vi.mock("../useAgentModelPickerState", () => ({ useAgentModelPickerState: (args: unknown) => ({ getModelsForAgent: () => [], + getProvenModelsForAgent: () => [], isModelInventoryAuthoritative: () => false, ...mockUseAgentModelPickerState(args), }), @@ -68,6 +69,7 @@ vi.mock("@/shared/api/acpConnection", () => ({ })); vi.mock("@/shared/api/acp", () => ({ + reserveAcpSessionConfiguration: () => ({ sequence: 0, clear: () => {} }), acpPrepareSession: (...args: unknown[]) => mockPrepareSession(...args), })); @@ -211,6 +213,7 @@ describe("useResolvedAgentModelPicker", () => { "openai", "/w", expect.objectContaining({ modelId: "next" }), + expect.objectContaining({ clear: expect.any(Function) }), ); }); @@ -469,6 +472,14 @@ describe("useResolvedAgentModelPicker", () => { providerId: "anthropic", }, ], + getProvenModelsForAgent: () => [ + { id: "gpt-5.4", name: "GPT-5.4", providerId: "openai" }, + { + id: "claude-sonnet-4", + name: "Claude Sonnet 4", + providerId: "anthropic", + }, + ], modelsLoading: false, modelStatusMessage: null, handleProviderChange: (providerId: string) => @@ -556,6 +567,19 @@ describe("useResolvedAgentModelPicker", () => { recommended: true, }, ], + getProvenModelsForAgent: () => [ + { + id: "gpt-5.4-mini", + name: "GPT Mini 5.4", + providerId: "codex-acp", + }, + { + id: "gpt-5.5", + name: "GPT 5.5", + providerId: "codex-acp", + recommended: true, + }, + ], modelsLoading: false, modelStatusMessage: null, handleProviderChange: vi.fn(), @@ -1149,6 +1173,7 @@ describe("useResolvedAgentModelPicker", () => { recommended: true, }, ], + getProvenModelsForAgent: () => [], isModelInventoryAuthoritative: () => false, modelsLoading: true, modelStatusMessage: null, @@ -1201,6 +1226,14 @@ describe("useResolvedAgentModelPicker", () => { recommended: true, }, ], + getProvenModelsForAgent: () => [ + { + id: "gpt-5.6", + name: "GPT-5.6", + providerId: "openai", + recommended: true, + }, + ], isModelInventoryAuthoritative: (providerId: string) => providerId === "openai", modelsLoading: false, @@ -1220,6 +1253,52 @@ describe("useResolvedAgentModelPicker", () => { }); }); + it("does not auto-select an advisory model while inventory proof is unavailable", () => { + mockUseAgentModelPickerState.mockImplementation(() => ({ + pickerAgents: [{ id: "goose", label: "Goose" }], + availableModels: [ + { id: "advisory", name: "Advisory", recommended: true }, + ], + getProvenModelsForAgent: () => [], + isModelInventoryAuthoritative: () => false, + modelsLoading: true, + modelStatusMessage: null, + handleProviderChange: vi.fn(), + handleModelChange: vi.fn(), + })); + + const { result } = renderModelPicker({ + selectedProvider: "openai", + sessionId: null, + session: undefined, + }); + + expect(result.current.effectiveModelSelection).toBeNull(); + }); + + it("does not auto-select an unqualified advisory model from authoritative empty inventory", () => { + mockUseAgentModelPickerState.mockImplementation(() => ({ + pickerAgents: [{ id: "goose", label: "Goose" }], + availableModels: [ + { id: "advisory", name: "Advisory", recommended: true }, + ], + getProvenModelsForAgent: () => [], + isModelInventoryAuthoritative: () => true, + modelsLoading: false, + modelStatusMessage: null, + handleProviderChange: vi.fn(), + handleModelChange: vi.fn(), + })); + + const { result } = renderModelPicker({ + selectedProvider: "openai", + sessionId: null, + session: undefined, + }); + + expect(result.current.effectiveModelSelection).toBeNull(); + }); + it("ignores a stored model missing from an authoritative populated inventory", () => { window.localStorage.setItem( "goose:preferredModelsByAgent", @@ -1241,6 +1320,9 @@ describe("useResolvedAgentModelPicker", () => { recommended: true, }, ], + getProvenModelsForAgent: () => [ + { id: "gpt-5.6", name: "GPT-5.6", providerId: "openai" }, + ], isModelInventoryAuthoritative: () => true, modelsLoading: true, modelStatusMessage: null, diff --git a/src/features/chat/hooks/useAgentModelPickerState.ts b/src/features/chat/hooks/useAgentModelPickerState.ts index 6293edcf1..f42c15e67 100644 --- a/src/features/chat/hooks/useAgentModelPickerState.ts +++ b/src/features/chat/hooks/useAgentModelPickerState.ts @@ -29,6 +29,7 @@ export function useAgentModelPickerState({ configuredModelProviderIds, modelCacheRefreshProviderIds, getModelsForAgent, + getProvenModelsForAgent, isModelInventoryAuthoritative: isProviderModelInventoryAuthoritative, refreshAllModelProviders, isRefreshingProvider, @@ -212,6 +213,7 @@ export function useAgentModelPickerState({ pickerAgents, availableModels, getModelsForAgent, + getProvenModelsForAgent, isModelInventoryAuthoritative, modelsLoading, modelStatusMessage, diff --git a/src/features/chat/hooks/useChatSessionController.ts b/src/features/chat/hooks/useChatSessionController.ts index 3cbe97a42..85d54e357 100644 --- a/src/features/chat/hooks/useChatSessionController.ts +++ b/src/features/chat/hooks/useChatSessionController.ts @@ -1093,6 +1093,8 @@ export function useChatSessionController({ pickerAgents, availableModels, getModelsForAgent, + getProvenModelsForAgent, + isModelInventoryAuthoritative, modelsLoading, modelStatusMessage, handleProviderChange, @@ -1208,9 +1210,17 @@ export function useChatSessionController({ providers, models: getModelsForAgent("goose"), getModelsForHarness: getModelsForAgent, + getProvenModelsForHarness: getProvenModelsForAgent, + isModelInventoryAuthoritative, catalogEntries, }), - [catalogEntries, getModelsForAgent, providers], + [ + catalogEntries, + getModelsForAgent, + getProvenModelsForAgent, + isModelInventoryAuthoritative, + providers, + ], ); const prepareSessionForCurrentSelection = useCallback( async ( diff --git a/src/features/chat/hooks/useResolvedAgentModelPicker.ts b/src/features/chat/hooks/useResolvedAgentModelPicker.ts index d21d8a4b9..85a4bb808 100644 --- a/src/features/chat/hooks/useResolvedAgentModelPicker.ts +++ b/src/features/chat/hooks/useResolvedAgentModelPicker.ts @@ -133,6 +133,7 @@ function getPreferredSelectionForAgent( function resolveAvailableSelection( selection: PreferredModelSelection, models: readonly ModelOption[], + provenModels: readonly ModelOption[], selectedModelProviderId: string | null, isInventoryAuthoritative: (providerId: string) => boolean, ): PreferredModelSelection | null { @@ -143,7 +144,10 @@ function resolveAvailableSelection( return null; } - const matchingModel = models.find( + const candidates = isInventoryAuthoritative(selection.modelProviderId) + ? provenModels + : models; + const matchingModel = candidates.find( (model) => model.id === selection.id && (!model.providerId || model.providerId === selection.modelProviderId), @@ -337,6 +341,7 @@ export function useResolvedAgentModelPicker({ pickerAgents, availableModels, getModelsForAgent, + getProvenModelsForAgent, isModelInventoryAuthoritative, modelsLoading, modelStatusMessage, @@ -695,6 +700,7 @@ export function useResolvedAgentModelPicker({ const availableStoredSelection = resolveAvailableSelection( storedSelection, availableModels, + getProvenModelsForAgent(selectedAgentId), concreteSelectedProviderId, isModelInventoryAuthoritative, ); @@ -723,6 +729,7 @@ export function useResolvedAgentModelPicker({ modelProviderId: defaultModelProviderId, }, availableModels, + getProvenModelsForAgent(selectedAgentId), concreteSelectedProviderId, isModelInventoryAuthoritative, ); @@ -730,6 +737,7 @@ export function useResolvedAgentModelPicker({ availableModels, catalogEntries, concreteSelectedProviderId, + getProvenModelsForAgent, gooseDefaultSelection, isModelInventoryAuthoritative, selectedAgentId, @@ -767,7 +775,11 @@ export function useResolvedAgentModelPicker({ }; } - if (isModelAlias(executionTarget.modelId)) { + if ( + isModelAlias(executionTarget.modelId) || + (executionTarget.modelProviderId && + isModelInventoryAuthoritative(executionTarget.modelProviderId)) + ) { return null; } @@ -777,17 +789,28 @@ export function useResolvedAgentModelPicker({ modelProviderId: executionTarget.modelProviderId, source: "explicit", }; - }, [availableModels, session]); + }, [availableModels, isModelInventoryAuthoritative, session]); const availableDefaultModelSelection = useMemo(() => { + const provenModels = getProvenModelsForAgent(selectedAgentId); + const selectableModels = availableModels.filter((model) => { + const providerId = model.providerId ?? concreteSelectedProviderId; + return provenModels.some( + (proven) => + proven.id === model.id && + (!providerId || + !proven.providerId || + proven.providerId === providerId), + ); + }); const compatibleModels = concreteSelectedProviderId - ? availableModels.filter( + ? selectableModels.filter( (model) => !model.providerId || model.providerId === concreteSelectedProviderId, ) - : availableModels; + : selectableModels; const defaultModel = compatibleModels.find((model) => model.recommended) ?? compatibleModels[0]; @@ -802,7 +825,13 @@ export function useResolvedAgentModelPicker({ modelProviderId: defaultModel.providerId ?? selectedProvider, source: defaultModel.recommended ? "default" : "explicit", }; - }, [availableModels, concreteSelectedProviderId, selectedProvider]); + }, [ + availableModels, + concreteSelectedProviderId, + getProvenModelsForAgent, + selectedAgentId, + selectedProvider, + ]); const fallbackModelSelection = session ? null @@ -817,6 +846,8 @@ export function useResolvedAgentModelPicker({ pickerAgents, availableModels, getModelsForAgent, + getProvenModelsForAgent, + isModelInventoryAuthoritative, modelsLoading, modelStatusMessage, handleProviderChange, diff --git a/src/features/chat/lib/__tests__/sessionActivation.test.ts b/src/features/chat/lib/__tests__/sessionActivation.test.ts index 8e8979f58..3e0b24729 100644 --- a/src/features/chat/lib/__tests__/sessionActivation.test.ts +++ b/src/features/chat/lib/__tests__/sessionActivation.test.ts @@ -37,6 +37,7 @@ const resolvePath = vi.hoisted(() => vi.fn()); const checkDirectoriesExist = vi.hoisted(() => vi.fn()); vi.mock("@/shared/api/acp", () => ({ + reserveAcpSessionConfiguration: () => ({ sequence: 0, clear: () => {} }), acpGetSessionInfo: (...args: unknown[]) => acpGetSessionInfo(...args), acpLoadSession: (...args: unknown[]) => acpLoadSession(...args), acpPrepareSession: (...args: unknown[]) => acpPrepareSession(...args), @@ -571,7 +572,8 @@ describe("loadSessionMessages", () => { "s-selection-race", "databricks_v2", "/resolved/existing/session", - { modelId: "goose-gpt-5-6-sol" }, + { modelId: "goose-gpt-5-6-sol", selectionAlreadyResolved: true }, + expect.objectContaining({ clear: expect.any(Function) }), ); expect(acpLoadSession.mock.invocationCallOrder[0]).toBeLessThan( acpPrepareSession.mock.invocationCallOrder[0], diff --git a/src/features/chat/lib/__tests__/sessionWorkspaceCleanup.test.ts b/src/features/chat/lib/__tests__/sessionWorkspaceCleanup.test.ts index 8216387f1..0c0df1745 100644 --- a/src/features/chat/lib/__tests__/sessionWorkspaceCleanup.test.ts +++ b/src/features/chat/lib/__tests__/sessionWorkspaceCleanup.test.ts @@ -23,6 +23,7 @@ const mocks = vi.hoisted(() => ({ })); vi.mock("@/shared/api/acp", () => ({ + reserveAcpSessionConfiguration: () => ({ sequence: 0, clear: () => {} }), acpListSessionsPage: mocks.acpListSessionsPage, })); diff --git a/src/features/chat/lib/__tests__/steerCore.test.ts b/src/features/chat/lib/__tests__/steerCore.test.ts index 8f5602679..204458370 100644 --- a/src/features/chat/lib/__tests__/steerCore.test.ts +++ b/src/features/chat/lib/__tests__/steerCore.test.ts @@ -5,6 +5,7 @@ import { MAX_PROMPT_ATTACHMENT_BYTES } from "../attachmentPayloadBudget"; const mockAcpSteerMessage = vi.fn(); vi.mock("@/shared/api/acp", () => ({ + reserveAcpSessionConfiguration: () => ({ sequence: 0, clear: () => {} }), acpSteerMessage: (...args: unknown[]) => mockAcpSteerMessage(...args), })); diff --git a/src/features/chat/lib/queuedSessionSend.test.ts b/src/features/chat/lib/queuedSessionSend.test.ts index bfdb2777b..f37d12394 100644 --- a/src/features/chat/lib/queuedSessionSend.test.ts +++ b/src/features/chat/lib/queuedSessionSend.test.ts @@ -31,6 +31,7 @@ const mocks = vi.hoisted(() => ({ })); vi.mock("@/shared/api/acp", () => ({ + reserveAcpSessionConfiguration: () => ({ sequence: 0, clear: () => {} }), acpGetSessionInfo: (...args: unknown[]) => mocks.acpGetSessionInfo(...args), acpLoadSession: (...args: unknown[]) => mocks.acpLoadSession(...args), acpPrepareSession: (...args: unknown[]) => mocks.acpPrepareSession(...args), diff --git a/src/features/chat/lib/sendCore.test.ts b/src/features/chat/lib/sendCore.test.ts index 167bf5476..2f7b47e57 100644 --- a/src/features/chat/lib/sendCore.test.ts +++ b/src/features/chat/lib/sendCore.test.ts @@ -9,6 +9,7 @@ const mocks = vi.hoisted(() => ({ })); vi.mock("@/shared/api/acp", () => ({ + reserveAcpSessionConfiguration: () => ({ sequence: 0, clear: () => {} }), acpSendMessage: (...args: unknown[]) => mocks.acpSendMessage(...args), })); diff --git a/src/features/chat/lib/sessionModelPreference.test.ts b/src/features/chat/lib/sessionModelPreference.test.ts index db4b3d4e6..d561dc3e0 100644 --- a/src/features/chat/lib/sessionModelPreference.test.ts +++ b/src/features/chat/lib/sessionModelPreference.test.ts @@ -95,6 +95,19 @@ describe("resolveSessionModelPreference", () => { }); }); + it("drops a stored model when an authoritative inventory is empty", () => { + expect( + sanitizeSessionModelPreference( + { + providerId: "openai", + modelId: "gpt-5.4", + modelName: "GPT-5.4", + }, + { models: [] }, + ), + ).toEqual({ providerId: "openai" }); + }); + it("drops a stored model when the provider model list no longer contains it", () => { expect( sanitizeSessionModelPreference( diff --git a/src/features/chat/lib/sessionModelPreference.ts b/src/features/chat/lib/sessionModelPreference.ts index 80ab31deb..d0c65a980 100644 --- a/src/features/chat/lib/sessionModelPreference.ts +++ b/src/features/chat/lib/sessionModelPreference.ts @@ -67,10 +67,6 @@ export function sanitizeSessionModelPreference( return preference; } - if (providerModels.models.length === 0) { - return preference; - } - if (providerModels.models.some((model) => model.id === preference.modelId)) { return preference; } diff --git a/src/features/chat/lib/sessionTargetCoordinator.test.ts b/src/features/chat/lib/sessionTargetCoordinator.test.ts index 93854799b..829ddd480 100644 --- a/src/features/chat/lib/sessionTargetCoordinator.test.ts +++ b/src/features/chat/lib/sessionTargetCoordinator.test.ts @@ -17,6 +17,7 @@ import { const mockPrepare = vi.fn(); vi.mock("@/shared/api/acp", () => ({ + reserveAcpSessionConfiguration: () => ({ sequence: 0, clear: () => {} }), acpPrepareSession: (...args: unknown[]) => mockPrepare(...args), })); @@ -126,9 +127,13 @@ describe("session target coordinator", () => { target: target("c"), }); expect(mockPrepare).toHaveBeenCalledTimes(1); - expect(mockPrepare).toHaveBeenCalledWith("s", "openai", "/w", { - modelId: "c", - }); + expect(mockPrepare).toHaveBeenCalledWith( + "s", + "openai", + "/w", + { modelId: "c", selectionAlreadyResolved: true }, + expect.objectContaining({ clear: expect.any(Function) }), + ); }); it("prevents an on-wire stale operation from committing over the winner", async () => { @@ -405,10 +410,17 @@ describe("session target coordinator", () => { requireReasoningEffort: true, }); - expect(mockPrepare).toHaveBeenCalledWith("s", "openai", "/w", { - modelId: "a", - forceConfigRefresh: true, - }); + expect(mockPrepare).toHaveBeenCalledWith( + "s", + "openai", + "/w", + { + modelId: "a", + forceConfigRefresh: true, + selectionAlreadyResolved: true, + }, + expect.objectContaining({ clear: expect.any(Function) }), + ); }); it("defers external hydration until the dispatch lease releases", () => { diff --git a/src/features/chat/lib/sessionTargetCoordinator.ts b/src/features/chat/lib/sessionTargetCoordinator.ts index 4720276b2..7a4b8f891 100644 --- a/src/features/chat/lib/sessionTargetCoordinator.ts +++ b/src/features/chat/lib/sessionTargetCoordinator.ts @@ -1,4 +1,7 @@ -import { acpPrepareSession } from "@/shared/api/acp"; +import { + acpPrepareSession, + reserveAcpSessionConfiguration, +} from "@/shared/api/acp"; import type { AcpModelConfigSnapshot, AcpReasoningEffortConfigSnapshot, @@ -269,6 +272,7 @@ async function execute( operation: PendingOperation, ): Promise { const { request, operationId } = operation; + const intent = reserveAcpSessionConfiguration(request.sessionId); try { const effective = await resolveEffectiveTarget(request.target); const liveTarget = useChatSessionStore @@ -325,9 +329,10 @@ async function execute( throw new Error("Session execution target requires a provider boundary."); } const forceConfigRefresh = - request.requireReasoningEffort && - !useChatSessionStore.getState().getSession(request.sessionId) - ?.reasoningEffort; + (request.requireReasoningEffort && + !useChatSessionStore.getState().getSession(request.sessionId) + ?.reasoningEffort) || + (request.target.modelId !== undefined && effective.modelId === undefined); const snapshot = await acpPrepareSession( request.sessionId, selection.providerId, @@ -335,10 +340,12 @@ async function execute( { ...(selection.modelId ? { modelId: selection.modelId } : {}), ...(forceConfigRefresh ? { forceConfigRefresh: true } : {}), + selectionAlreadyResolved: true, ...(request.operationId || request.requestId ? { requestId: operationId } : {}), }, + intent, ); if (!currentOperation(actor, operation)) { resolveSuperseded(actor, operation); @@ -356,11 +363,14 @@ async function execute( settleOperation(operation, { status: "session-missing", applied: false }); return; } - const acknowledged = - !effective.modelId && snapshot?.model - ? (materializeSessionExecutionModel(effective, snapshot.model) ?? - effective) - : effective; + // The ACP response is the acknowledgement of the configuration actually + // applied. Always reconcile from it: inventory can change between + // preflight resolution and preparation, and a provider reset can select a + // replacement model. + const acknowledged = snapshot?.model + ? (materializeSessionExecutionModel(effective, snapshot.model) ?? + effective) + : effective; const legacyIntent = actor.selection ? { requestId: actor.selection.operationId, @@ -441,6 +451,8 @@ async function execute( error, fallback, }); + } finally { + intent.clear(); } } diff --git a/src/features/chat/lib/sessionTargetTransition.integration.test.ts b/src/features/chat/lib/sessionTargetTransition.integration.test.ts index 43b94ebdc..3518d0933 100644 --- a/src/features/chat/lib/sessionTargetTransition.integration.test.ts +++ b/src/features/chat/lib/sessionTargetTransition.integration.test.ts @@ -6,6 +6,32 @@ const mockSetProvider = vi.fn(); const mockSetModel = vi.fn(); const mockGetClient = vi.fn(); +function deferred() { + let resolve!: (value: T) => void; + let reject!: (error: unknown) => void; + const promise = new Promise((resolvePromise, rejectPromise) => { + resolve = resolvePromise; + reject = rejectPromise; + }); + return { promise, resolve, reject }; +} + +function executionConfigResponse(providerId: string, modelId: string) { + return { + configOptions: [ + { + id: "provider", + kind: { type: "select", currentValue: providerId, options: [] }, + }, + { + id: "model", + category: "model", + kind: { type: "select", currentValue: modelId, options: [] }, + }, + ], + }; +} + vi.mock("@/shared/api/acpApi", () => ({ loadSession: (...args: unknown[]) => mockLoadSession(...args), setProvider: (...args: unknown[]) => mockSetProvider(...args), @@ -245,6 +271,11 @@ describe("transitionSessionTarget with managed Goose models", () => { "./sessionTargetCoordinator" ); + mockSetProvider.mockResolvedValueOnce({ + model: { modelId: "backend-fallback", modelName: "Backend fallback" }, + reasoningEffort: null, + }); + await expect( transitionSessionTarget({ sessionId: "no-default-session", @@ -261,18 +292,149 @@ describe("transitionSessionTarget with managed Goose models", () => { resolvedTarget: { harnessId: "goose", modelProviderId: "databricks_v2", + modelId: "backend-fallback", }, }); + expect(mockSetProvider).toHaveBeenCalledWith( + "no-default-session", + "databricks_v2", + { requestId: undefined }, + ); expect(mockSetModel).not.toHaveBeenCalled(); + expect(mockGetClient).toHaveBeenCalledTimes(1); expect( useChatSessionStore.getState().getSession("no-default-session"), ).toMatchObject({ executionTarget: { harnessId: "goose", modelProviderId: "databricks_v2", + modelId: "backend-fallback", + }, + }); + const { requireSessionInvocationSelection } = await import( + "@/shared/api/acpSessionRegistry" + ); + expect(requireSessionInvocationSelection("no-default-session")).toEqual({ + providerId: "databricks_v2", + modelId: "backend-fallback", + }); + }); + + it("suppresses a concurrent load while migration proof is pending", async () => { + const supportedModels = deferred<{ models: string[] }>(); + const supportedModelsList = vi + .fn() + .mockReturnValue(supportedModels.promise); + mockGetClient.mockResolvedValue({ + goose: { GooseUnstableProvidersSupportedModelsList: supportedModelsList }, + }); + mockLoadSession.mockResolvedValueOnce( + executionConfigResponse("legacy-provider", "legacy-model"), + ); + const applyModelConfigSnapshot = vi.fn(); + const { setSessionConfigSnapshotHandlers } = await import( + "@/shared/api/acpSessionConfigSnapshots" + ); + setSessionConfigSnapshotHandlers({ applyModelConfigSnapshot }); + const { resetManagedModelSelectionRepairCacheForTests } = await import( + "@/features/providers/lib/managedModelSelectionRepair" + ); + resetManagedModelSelectionRepairCacheForTests(); + const { acpLoadSession } = await import("@/shared/api/acp"); + const { transitionSessionTarget } = await import( + "./sessionTargetCoordinator" + ); + + const transition = transitionSessionTarget({ + sessionId: "migration-proof-session", + target: { + harnessId: "goose", + modelProviderId: "legacy-provider", + modelId: "legacy-model", + modelName: "Legacy model", + }, + workingDir: "/tmp/project", + }); + await vi.waitFor(() => + expect(supportedModelsList).toHaveBeenCalledTimes(1), + ); + + await acpLoadSession("migration-proof-session", "/tmp/project"); + expect(applyModelConfigSnapshot).not.toHaveBeenCalled(); + + supportedModels.resolve({ models: ["goose-gpt-5-5"] }); + await expect(transition).resolves.toMatchObject({ + applied: true, + resolvedTarget: { + harnessId: "goose", + modelProviderId: "databricks_v2", + modelId: "goose-gpt-5-5", }, }); + expect(applyModelConfigSnapshot).not.toHaveBeenCalled(); + }); + + it("materializes the prepared model for a matching provider-only transition", async () => { + const sessionId = "provider-only-prepared-session"; + const { useChatSessionStore } = await import( + "@/features/chat/stores/chatSessionStore" + ); + useChatSessionStore.setState({ + sessions: [ + { + id: sessionId, + title: "Prepared session", + executionTarget: { + harnessId: "goose", + modelProviderId: "databricks_v2", + }, + createdAt: "2026-08-01T00:00:00.000Z", + updatedAt: "2026-08-01T00:00:00.000Z", + messageCount: 1, + }, + ], + }); + const registry = await import("@/shared/api/acpSessionRegistry"); + registry.registerPreparedSession( + sessionId, + "databricks_v2", + "/tmp/project", + "goose-gpt-5-5", + ); + const { transitionSessionTarget } = await import( + "./sessionTargetCoordinator" + ); + + await expect( + transitionSessionTarget({ + sessionId, + target: { + harnessId: "goose", + modelProviderId: "databricks_v2", + }, + workingDir: "/tmp/project", + }), + ).resolves.toMatchObject({ + status: "committed", + target: { + modelProviderId: "databricks_v2", + modelId: "goose-gpt-5-5", + }, + }); + + expect(mockSetProvider).not.toHaveBeenCalled(); + expect(mockSetModel).not.toHaveBeenCalled(); + expect( + useChatSessionStore.getState().getSession(sessionId)?.executionTarget, + ).toMatchObject({ + modelProviderId: "databricks_v2", + modelId: "goose-gpt-5-5", + }); + expect(registry.requireSessionInvocationSelection(sessionId)).toEqual({ + providerId: "databricks_v2", + modelId: "goose-gpt-5-5", + }); }); it("finishes on the explicitly selected model instead of the managed default", async () => { diff --git a/src/features/chat/lib/sessionTargetTransition.test.ts b/src/features/chat/lib/sessionTargetTransition.test.ts index 40c6c7d81..958c4c4ef 100644 --- a/src/features/chat/lib/sessionTargetTransition.test.ts +++ b/src/features/chat/lib/sessionTargetTransition.test.ts @@ -5,6 +5,7 @@ import { resetSessionTargetCoordinatorsForTests } from "./sessionTargetCoordinat const mockAcpPrepareSession = vi.fn(); vi.mock("@/shared/api/acp", () => ({ + reserveAcpSessionConfiguration: () => ({ sequence: 0, clear: () => {} }), acpPrepareSession: (...args: unknown[]) => mockAcpPrepareSession(...args), })); @@ -38,7 +39,8 @@ describe("transitionSessionTarget", () => { "session-latest", "new-provider", "/new", - { modelId: "new-model" }, + { modelId: "new-model", selectionAlreadyResolved: true }, + expect.objectContaining({ clear: expect.any(Function) }), ); }); @@ -63,8 +65,10 @@ describe("transitionSessionTarget", () => { { modelId: "goose-gpt-5-6-sol", forceConfigRefresh: true, + selectionAlreadyResolved: true, requestId: "request-5-6", }, + expect.objectContaining({ clear: expect.any(Function) }), ); }); }); diff --git a/src/features/chat/stores/__tests__/chatSessionStore.test.ts b/src/features/chat/stores/__tests__/chatSessionStore.test.ts index 61c3842fb..6193eccaa 100644 --- a/src/features/chat/stores/__tests__/chatSessionStore.test.ts +++ b/src/features/chat/stores/__tests__/chatSessionStore.test.ts @@ -29,6 +29,7 @@ const mocks = vi.hoisted(() => ({ })); vi.mock("@/shared/api/acp", () => ({ + reserveAcpSessionConfiguration: () => ({ sequence: 0, clear: () => {} }), acpCreateSession: (...args: unknown[]) => mocks.acpCreateSession(...args), acpListSessionsPage: (...args: unknown[]) => mocks.acpListSessionsPage(...args), diff --git a/src/features/chat/ui/AgentModelPicker.tsx b/src/features/chat/ui/AgentModelPicker.tsx index f88687bd2..8ce856b01 100644 --- a/src/features/chat/ui/AgentModelPicker.tsx +++ b/src/features/chat/ui/AgentModelPicker.tsx @@ -55,6 +55,8 @@ interface AgentModelPickerProps { loading?: boolean; isCompact?: boolean; showSelectedModelInTrigger?: boolean; + /** A provider-only target must not synthesize a default model for display. */ + showDefaultModelInTrigger?: boolean; triggerTabIndex?: number; triggerIconOnly?: boolean; open?: boolean; @@ -175,6 +177,7 @@ export function AgentModelPicker({ loading = false, isCompact = false, showSelectedModelInTrigger = true, + showDefaultModelInTrigger = true, triggerTabIndex, triggerIconOnly = false, open: controlledOpen, @@ -275,15 +278,16 @@ export function AgentModelPicker({ displayModelLabel, selectedAgentId, ]); - const triggerLabel = showSelectedModelInTrigger - ? resolvePickerTriggerLabel({ - currentModelId, - currentModelName, - currentModelProviderId, - availableModels: displayedModels, - selectedAgentLabel, - }) - : selectedAgentLabel; + const triggerLabel = + showSelectedModelInTrigger && (currentModelId || showDefaultModelInTrigger) + ? resolvePickerTriggerLabel({ + currentModelId, + currentModelName, + currentModelProviderId, + availableModels: displayedModels, + selectedAgentLabel, + }) + : selectedAgentLabel; const triggerTitle = triggerLabel ?? (loading ? t("toolbar.loading") : undefined); const triggerButtonSize = triggerIconOnly ? "icon-pill-sm" : "sm"; diff --git a/src/features/home/ui/HomeScreen.test.tsx b/src/features/home/ui/HomeScreen.test.tsx index bd56861a7..64ad3b1b5 100644 --- a/src/features/home/ui/HomeScreen.test.tsx +++ b/src/features/home/ui/HomeScreen.test.tsx @@ -127,6 +127,7 @@ vi.mock("@/features/chat/hooks/useMentionHandlers", () => ({ })); vi.mock("@/shared/api/acp", () => ({ + reserveAcpSessionConfiguration: () => ({ sequence: 0, clear: () => {} }), discoverAcpProviders: vi.fn().mockResolvedValue([ { id: "goose", label: "Goose" }, { id: "claude-acp", label: "Claude Code" }, diff --git a/src/features/providers/defaultProviderConfig.test.ts b/src/features/providers/defaultProviderConfig.test.ts index ab8e429cd..e5369adf9 100644 --- a/src/features/providers/defaultProviderConfig.test.ts +++ b/src/features/providers/defaultProviderConfig.test.ts @@ -63,6 +63,36 @@ describe("reconcileManagedDefaultProviderSelection", () => { mockGetStoredModelPreference.mockReturnValue(null); }); + it("does not wait for model inventory for provider-only defaults", async () => { + useRuntimeConfigStore.setState({ + loaded: true, + config: managedRuntimeConfig, + result: { + status: "ready", + source: "bundledFile", + config: managedRuntimeConfig, + }, + }); + const supportedModelsList = vi.fn().mockReturnValue(new Promise(() => {})); + mockGetClient.mockResolvedValue({ + goose: { + GooseUnstableDefaultsRead: vi.fn().mockResolvedValue({ + providerId: "databricks_v2", + modelId: undefined, + }), + GooseUnstableDefaultsSave: defaultsSave, + GooseUnstableProvidersSupportedModelsList: supportedModelsList, + }, + } as never); + + await expect(reconcileManagedDefaultProviderSelection()).resolves.toEqual({ + providerId: "databricks_v2", + modelId: undefined, + }); + expect(supportedModelsList).not.toHaveBeenCalled(); + expect(defaultsSave).not.toHaveBeenCalled(); + }); + it("repairs a persisted Goose harness sentinel to the managed default", async () => { useRuntimeConfigStore.setState({ loaded: true, @@ -80,6 +110,9 @@ describe("reconcileManagedDefaultProviderSelection", () => { modelId: "goose", }), GooseUnstableDefaultsSave: defaultsSave, + GooseUnstableProvidersSupportedModelsList: vi.fn().mockResolvedValue({ + models: ["goose-gpt-5-5"], + }), }, } as never); @@ -119,6 +152,30 @@ describe("saveDefaultProviderSelection", () => { }); }); + it("does not persist an advisory recommendation excluded from live proof", async () => { + const refreshProviderModels = vi.fn().mockImplementation((providerId) => { + useProviderModelCacheStore.setState({ + providers: new Map([ + [ + providerId, + { + providerId, + fetchedAt: Date.now(), + provenModelIds: [], + models: [{ id: "advisory", name: "Advisory", recommended: true }], + }, + ], + ]), + }); + }); + useProviderModelCacheStore.setState({ refreshProviderModels }); + + await expect(saveDefaultProviderSelection("openai")).rejects.toThrow( + "Could not load models for provider", + ); + expect(defaultsSave).not.toHaveBeenCalled(); + }); + it("saves backend defaults, local goose preference, and readiness", async () => { const refreshProviderModels = vi.fn().mockImplementation((providerId) => { useProviderModelCacheStore.setState({ @@ -128,6 +185,7 @@ describe("saveDefaultProviderSelection", () => { { providerId, fetchedAt: Date.now(), + provenModelIds: ["gpt-4o"], models: [{ id: "gpt-4o", name: "gpt-4o", recommended: true }], }, ], @@ -271,6 +329,7 @@ describe("saveDefaultProviderSelectionFromConfiguredProvider", () => { { providerId, fetchedAt: Date.now(), + provenModelIds: models.map((model) => model.id), models, }, ], diff --git a/src/features/providers/defaultProviderConfig.ts b/src/features/providers/defaultProviderConfig.ts index 58ad083af..753c05671 100644 --- a/src/features/providers/defaultProviderConfig.ts +++ b/src/features/providers/defaultProviderConfig.ts @@ -137,9 +137,13 @@ export async function saveDefaultProviderSelection( const modelCacheStore = useProviderModelCacheStore.getState(); await modelCacheStore.refreshProviderModels(providerId, { force: true }); - const models = useProviderModelCacheStore - .getState() - .getModelsForProvider(providerId); + const cache = useProviderModelCacheStore.getState(); + if (!cache.isModelInventoryAuthoritative(providerId)) { + throw new Error( + "Could not prove models for provider. Check provider setup and try again.", + ); + } + const models = cache.getProvenModelsForProvider(providerId); const runtimeDefaultModelId = providerId === getDefaultGooseModelProviderId() ? getDefaultGooseModelId() diff --git a/src/features/providers/hooks/useNewSessionTarget.test.tsx b/src/features/providers/hooks/useNewSessionTarget.test.tsx index 72db5f68a..eaa149365 100644 --- a/src/features/providers/hooks/useNewSessionTarget.test.tsx +++ b/src/features/providers/hooks/useNewSessionTarget.test.tsx @@ -60,6 +60,7 @@ describe("useNewSessionTarget", () => { providerId: "anthropic", }, ], + provenModelIds: ["claude-sonnet-4"], fetchedAt: Date.now(), }, ], @@ -95,6 +96,70 @@ describe("useNewSessionTarget", () => { }); }); + it.each([ + { + name: "proof is absent", + provenModelIds: undefined, + expectedModelId: "removed-model", + }, + { + name: "live proof is empty", + provenModelIds: [], + expectedModelId: "goose-gpt-5-5", + }, + { + name: "live proof omits the stored model", + provenModelIds: ["claude-sonnet-4"], + expectedModelId: "goose-gpt-5-5", + }, + ])("uses $name when resolving a stored model for a new chat", async ({ + provenModelIds, + expectedModelId, + }) => { + useProviderModelCacheStore.setState({ + providers: new Map([ + [ + "anthropic", + { + providerId: "anthropic", + models: [ + { + id: "claude-sonnet-4", + name: "Claude Sonnet 4", + providerId: "anthropic", + }, + ], + ...(provenModelIds !== undefined ? { provenModelIds } : {}), + fetchedAt: Date.now(), + }, + ], + ]), + refreshingProviderIds: new Set(), + }); + window.localStorage.setItem( + "goose:preferredModelsByAgent", + JSON.stringify({ + goose: { + modelId: "removed-model", + modelName: "Removed model", + providerId: "anthropic", + }, + }), + ); + + const { result } = renderHook(() => useNewSessionTarget()); + let target: Awaited> | undefined; + await act(async () => { + target = await result.current(); + }); + + expect(target).toMatchObject({ + status: "ready", + providerId: "goose", + modelId: expectedModelId, + }); + }); + it("drops an unavailable stored model before resolving the new chat", async () => { window.localStorage.setItem( "goose:preferredModelsByAgent", diff --git a/src/features/providers/hooks/useProviderModels.test.tsx b/src/features/providers/hooks/useProviderModels.test.tsx index f0a1da316..498a94cd1 100644 --- a/src/features/providers/hooks/useProviderModels.test.tsx +++ b/src/features/providers/hooks/useProviderModels.test.tsx @@ -104,6 +104,7 @@ describe("useProviderModels", () => { { providerId: "databricks_v2", models, + provenModelIds: models.map((model) => model.id), fetchedAt: Date.now(), }, ], diff --git a/src/features/providers/hooks/useProviderModels.ts b/src/features/providers/hooks/useProviderModels.ts index 869ccfe07..2de572b4f 100644 --- a/src/features/providers/hooks/useProviderModels.ts +++ b/src/features/providers/hooks/useProviderModels.ts @@ -82,6 +82,16 @@ export function useProviderModels() { [providers], ); + const getProvenModelsForProvider = useCallback( + (providerId: string) => { + const entry = providers.get(providerId); + if (!entry?.provenModelIds) return EMPTY_MODELS; + const provenIds = new Set(entry.provenModelIds); + return entry.models.filter((model) => provenIds.has(model.id)); + }, + [providers], + ); + const isModelInventoryAuthoritative = useCallback( (providerId: string) => isCachedModelInventoryAuthoritative(providers.get(providerId)), @@ -123,6 +133,14 @@ export function useProviderModels() { ], ); + const getProvenModelsForAgent = useCallback( + (agentId: string) => + agentId === "goose" + ? configuredModelProviderIds.flatMap(getProvenModelsForProvider) + : getProvenModelsForProvider(agentId), + [configuredModelProviderIds, getProvenModelsForProvider], + ); + const isRefreshingProvider = useCallback( (providerId: string) => refreshingProviderIds.has(providerId), [refreshingProviderIds], @@ -141,6 +159,7 @@ export function useProviderModels() { modelCacheRefreshProviderIds, getModelsForAgent, getModelsForProvider, + getProvenModelsForAgent, isModelInventoryAuthoritative, refreshProviderModels, refreshAllModelProviders, diff --git a/src/features/providers/lib/managedModelSelectionRepair.test.ts b/src/features/providers/lib/managedModelSelectionRepair.test.ts index ecc4a6fdc..f2fd17937 100644 --- a/src/features/providers/lib/managedModelSelectionRepair.test.ts +++ b/src/features/providers/lib/managedModelSelectionRepair.test.ts @@ -11,6 +11,7 @@ import { notifyProviderModelInventoryInvalidated } from "./providerModelInventor vi.mock("@/shared/api/acpConnection", () => ({ getClient: vi.fn(), + invalidateClientConnection: vi.fn().mockResolvedValue(undefined), })); const managedConfig: RuntimeConfig = { @@ -62,6 +63,34 @@ describe("repairManagedGooseModelSelection", () => { }); }); + it("keeps an authoritative-empty same-provider target provider-only", async () => { + vi.mocked(getClient).mockResolvedValue({ + goose: { + GooseUnstableProvidersSupportedModelsList: vi.fn().mockResolvedValue({ + models: [], + }), + }, + } as never); + + await expect( + repairManagedGooseModelSelection( + { providerId: "databricks_v2" }, + "session", + ), + ).resolves.toEqual({ providerId: "databricks_v2", modelId: undefined }); + }); + + it("preserves same-provider model-free intent when live proof cannot be read", async () => { + vi.mocked(getClient).mockRejectedValue(new Error("offline")); + + await expect( + repairManagedGooseModelSelection( + { providerId: "databricks_v2" }, + "session", + ), + ).resolves.toEqual({ providerId: "databricks_v2", modelId: undefined }); + }); + it("repairs any model absent from the live target-provider inventory", async () => { vi.mocked(getClient).mockResolvedValue({ goose: { @@ -172,6 +201,69 @@ describe("repairManagedGooseModelSelection", () => { expect(supportedModelsList).toHaveBeenCalledTimes(2); }); + it("releases same-provider proof after stalled client acquisition", async () => { + vi.useFakeTimers(); + vi.mocked(getClient).mockReturnValue(new Promise(() => {})); + + const repair = repairManagedGooseModelSelection( + { providerId: "databricks_v2", modelId: "future-model" }, + "session", + ); + await vi.advanceTimersByTimeAsync(60_000); + await expect(repair).resolves.toEqual({ + providerId: "databricks_v2", + modelId: "future-model", + }); + + vi.mocked(getClient).mockResolvedValue({ + goose: { + GooseUnstableProvidersSupportedModelsList: vi.fn().mockResolvedValue({ + models: ["future-model"], + }), + }, + } as never); + await expect( + repairManagedGooseModelSelection( + { providerId: "databricks_v2", modelId: "future-model" }, + "session", + ), + ).resolves.toEqual({ + providerId: "databricks_v2", + modelId: "future-model", + }); + }); + + it("releases same-provider proof after stalled inventory RPC", async () => { + vi.useFakeTimers(); + const stalledInventory = new Promise<{ models: string[] }>(() => {}); + const supportedModelsList = vi + .fn() + .mockReturnValueOnce(stalledInventory) + .mockResolvedValueOnce({ models: ["future-model"] }); + vi.mocked(getClient).mockResolvedValue({ + goose: { GooseUnstableProvidersSupportedModelsList: supportedModelsList }, + } as never); + + const repair = repairManagedGooseModelSelection( + { providerId: "databricks_v2", modelId: "future-model" }, + "session", + ); + await vi.advanceTimersByTimeAsync(60_000); + await expect(repair).resolves.toEqual({ + providerId: "databricks_v2", + modelId: "future-model", + }); + await expect( + repairManagedGooseModelSelection( + { providerId: "databricks_v2", modelId: "future-model" }, + "session", + ), + ).resolves.toEqual({ + providerId: "databricks_v2", + modelId: "future-model", + }); + expect(supportedModelsList).toHaveBeenCalledTimes(2); + }); it("preserves the selected model when live inventory cannot be read", async () => { vi.mocked(getClient).mockRejectedValue(new Error("offline")); diff --git a/src/features/providers/lib/managedModelSelectionRepair.ts b/src/features/providers/lib/managedModelSelectionRepair.ts index d4a9b752c..ea5fb1388 100644 --- a/src/features/providers/lib/managedModelSelectionRepair.ts +++ b/src/features/providers/lib/managedModelSelectionRepair.ts @@ -1,15 +1,18 @@ import packageJson from "../../../../package.json"; -import { getClient } from "@/shared/api/acpConnection"; import { useRuntimeConfigStore } from "@/shared/runtime-config/runtimeConfigStore"; import { resolveAgentProviderCatalogIdStrict } from "@/features/providers/providerCatalog"; -import { subscribeToProviderModelInventoryInvalidation } from "./providerModelInventoryEvents"; import { + providerModelInventoryGeneration, + subscribeToProviderModelInventoryInvalidation, +} from "@/shared/runtime-config/providerModelInventoryInvalidation"; +import { + readBoundedProvenModelInventory, resolveManagedGooseProviderSelection, + resolveValidatedManagedGooseProviderSelection, type GooseProviderSelection, type ManagedGooseProviderSelection, } from "@/shared/runtime-config/modelProviderPolicy"; -const DATABRICKS_V2_PROVIDER_ID = "databricks_v2"; const VALIDATED_INVENTORY_TTL_MS = 5 * 60 * 1000; const validatedInventories = new Map< string, @@ -19,16 +22,9 @@ const inventoryRequests = new Map< string, Promise | null> >(); -const inventoryGenerations = new Map(); - -function inventoryGeneration(providerId: string): number { - return inventoryGenerations.get(providerId) ?? 0; -} - subscribeToProviderModelInventoryInvalidation((providerId) => { validatedInventories.delete(providerId); inventoryRequests.delete(providerId); - inventoryGenerations.set(providerId, inventoryGeneration(providerId) + 1); }); export type ManagedModelRepairSource = @@ -50,17 +46,12 @@ async function validatedModelIds( const existing = inventoryRequests.get(providerId); if (existing) return existing; - const generationAtStart = inventoryGeneration(providerId); + const generationAtStart = providerModelInventoryGeneration(providerId); let request!: Promise | null>; request = (async () => { try { - const client = await getClient(); - const response = - await client.goose.GooseUnstableProvidersSupportedModelsList({ - providerId, - }); - const modelIds = new Set(response.models as string[]); - if (generationAtStart !== inventoryGeneration(providerId)) { + const modelIds = await readBoundedProvenModelInventory(providerId); + if (generationAtStart !== providerModelInventoryGeneration(providerId)) { return validatedModelIds(providerId); } validatedInventories.set(providerId, { @@ -69,6 +60,9 @@ async function validatedModelIds( }); return modelIds; } catch (error) { + if (generationAtStart !== providerModelInventoryGeneration(providerId)) { + return validatedModelIds(providerId); + } console.warn("Could not validate managed provider model inventory", { providerId, error: error instanceof Error ? error.message : String(error), @@ -102,11 +96,14 @@ export async function repairManagedGooseModelSelection( const config = useRuntimeConfigStore.getState().config; const initial = resolveManagedGooseProviderSelection(config, selection); if (!initial) return null; + if (initial.providerId !== selection.providerId) { + return resolveValidatedManagedGooseProviderSelection(config, selection); + } + if (!selection.modelId) { + return initial; + } - const targetModelIds = - initial.providerId === DATABRICKS_V2_PROVIDER_ID && selection.modelId - ? await validatedModelIds(initial.providerId) - : null; + const targetModelIds = await validatedModelIds(initial.providerId); const repaired = resolveManagedGooseProviderSelection(config, selection, { ...(targetModelIds ? { targetModelIds } : {}), targetInventoryValidated: targetModelIds !== null, @@ -132,5 +129,4 @@ export async function repairManagedGooseModelSelection( export function resetManagedModelSelectionRepairCacheForTests(): void { validatedInventories.clear(); inventoryRequests.clear(); - inventoryGenerations.clear(); } diff --git a/src/features/providers/lib/providerModelInventoryEvents.ts b/src/features/providers/lib/providerModelInventoryEvents.ts index 3c17f91e8..244fb9aff 100644 --- a/src/features/providers/lib/providerModelInventoryEvents.ts +++ b/src/features/providers/lib/providerModelInventoryEvents.ts @@ -1,19 +1,5 @@ -type ProviderModelInventoryInvalidationListener = (providerId: string) => void; - -const invalidationListeners = - new Set(); - -export function notifyProviderModelInventoryInvalidated( - providerId: string, -): void { - for (const listener of invalidationListeners) { - listener(providerId); - } -} - -export function subscribeToProviderModelInventoryInvalidation( - listener: ProviderModelInventoryInvalidationListener, -): () => void { - invalidationListeners.add(listener); - return () => invalidationListeners.delete(listener); -} +export { + notifyProviderModelInventoryInvalidated, + providerModelInventoryGeneration, + subscribeToProviderModelInventoryInvalidation, +} from "@/shared/runtime-config/providerModelInventoryInvalidation"; diff --git a/src/features/providers/lib/resolveSessionModelPreference.test.ts b/src/features/providers/lib/resolveSessionModelPreference.test.ts index cff241f2c..5c6d3f660 100644 --- a/src/features/providers/lib/resolveSessionModelPreference.test.ts +++ b/src/features/providers/lib/resolveSessionModelPreference.test.ts @@ -13,8 +13,9 @@ const mockCheckAllProviderStatus = vi.mocked(checkAllProviderStatus); function setCachedModels( providerId: string, models: string[], - fetchedAt = Date.now(), + options: { fetchedAt?: number; proven?: boolean } = {}, ) { + const { fetchedAt = Date.now(), proven = true } = options; useProviderModelCacheStore.setState({ providers: new Map([ [ @@ -22,6 +23,7 @@ function setCachedModels( { providerId, models: models.map((id) => ({ id, name: id, providerId })), + ...(proven ? { provenModelIds: models } : {}), fetchedAt, }, ], @@ -63,6 +65,7 @@ describe("resolveSupportedSessionModelPreference", () => { modelId: "gpt-5.4", }, }); + setCachedModels("openai", ["gpt-5.4"]); await expect( resolveSupportedSessionModelPreference("goose"), @@ -204,50 +207,72 @@ describe("resolveSupportedSessionModelPreference", () => { }); }); - it("preserves the selected model when the model cache has no model list", async () => { - setCachedModels("openai", []); - - await expect( - resolveSupportedSessionModelPreference("openai", "gpt-5.4"), - ).resolves.toEqual({ - providerId: "openai", - modelId: "gpt-5.4", - modelName: "gpt-5.4", - }); - }); - - it("drops an unsupported model when populated model cache is available", async () => { - setCachedModels("openai", ["gpt-5.3"]); + it.each([ + { + name: "proof is absent", + models: ["gpt-5.3"], + options: { proven: false }, + expected: { + providerId: "openai", + modelId: "gpt-5.4", + modelName: "gpt-5.4", + }, + }, + { + name: "live proof is empty", + models: [], + options: {}, + expected: { providerId: "openai" }, + }, + { + name: "live proof contains the preferred model", + models: ["gpt-5.4"], + options: {}, + expected: { + providerId: "openai", + modelId: "gpt-5.4", + modelName: "gpt-5.4", + }, + }, + { + name: "live proof omits the preferred model", + models: ["gpt-5.3"], + options: {}, + expected: { providerId: "openai" }, + }, + ])("uses only live proof when $name", async ({ + models, + options, + expected, + }) => { + setCachedModels("openai", models, options); await expect( resolveSupportedSessionModelPreference("openai", "gpt-5.4"), - ).resolves.toEqual({ - providerId: "openai", - }); + ).resolves.toEqual(expected); }); - it("preserves a selected model while a populated cache is provisional", async () => { - setCachedModels("openai", ["gpt-5.3"], 0); - - await expect( - resolveSupportedSessionModelPreference("openai", "gpt-5.4"), - ).resolves.toEqual({ - providerId: "openai", - modelId: "gpt-5.4", - modelName: "gpt-5.4", + it("rejects an advisory display candidate absent from live proof", async () => { + useProviderModelCacheStore.setState({ + providers: new Map([ + [ + "openai", + { + providerId: "openai", + models: [ + { id: "gpt-5.3", name: "gpt-5.3", providerId: "openai" }, + { id: "gpt-5.4", name: "gpt-5.4", providerId: "openai" }, + ], + provenModelIds: ["gpt-5.3"], + fetchedAt: Date.now(), + }, + ], + ]), }); - }); - - it("keeps a supported model when populated model cache is available", async () => { - setCachedModels("openai", ["gpt-5.4"]); await expect( resolveSupportedSessionModelPreference("openai", "gpt-5.4"), - ).resolves.toEqual({ - providerId: "openai", - modelId: "gpt-5.4", - modelName: "gpt-5.4", - }); + ).resolves.toEqual({ providerId: "openai" }); }); it("drops a stored model whose provider is disconnected", async () => { @@ -272,6 +297,7 @@ describe("resolveSupportedSessionModelPreference", () => { modelId: "goose-gpt-5-5", }, }); + setCachedModels("databricks_v2", ["goose-gpt-5-5"]); mockCheckAllProviderStatus.mockResolvedValue([ { providerId: "openai", isConfigured: false }, ]); diff --git a/src/features/providers/lib/resolveSessionModelPreference.ts b/src/features/providers/lib/resolveSessionModelPreference.ts index cd6193511..dd4a5c2af 100644 --- a/src/features/providers/lib/resolveSessionModelPreference.ts +++ b/src/features/providers/lib/resolveSessionModelPreference.ts @@ -112,23 +112,31 @@ export async function resolveSupportedSessionModelPreference( }; } + const modelCache = useProviderModelCacheStore.getState(); + + // A configured default is synthesized intent. Unlike an explicit preference, + // it cannot survive without a successful inventory proof. if (providerId === "goose" && !sessionModelPreference.modelId) { - sessionModelPreference = gooseDefaultPreference() ?? sessionModelPreference; + const fallback = gooseDefaultPreference(); + if (!fallback) return sessionModelPreference; + if (!modelCache.isModelInventoryAuthoritative(fallback.providerId)) { + return { providerId }; + } + return sanitizeSessionModelPreference(fallback, { + models: modelCache.getProvenModelsForProvider(fallback.providerId), + }); } if (!sessionModelPreference.modelId) { return sessionModelPreference; } - const modelCache = useProviderModelCacheStore.getState(); - const models = modelCache.getModelsForProvider( - sessionModelPreference.providerId, - ); + const modelProviderId = sessionModelPreference.providerId; - if ( - modelCache.isModelInventoryAuthoritative(sessionModelPreference.providerId) - ) { - return sanitizeSessionModelPreference(sessionModelPreference, { models }); + if (modelCache.isModelInventoryAuthoritative(modelProviderId)) { + return sanitizeSessionModelPreference(sessionModelPreference, { + models: modelCache.getProvenModelsForProvider(modelProviderId), + }); } if (!(await isProviderDisconnected(sessionModelPreference.providerId))) { @@ -137,12 +145,13 @@ export async function resolveSupportedSessionModelPreference( if (providerId === "goose") { const fallback = gooseDefaultPreference(); - if (fallback && fallback.providerId !== sessionModelPreference.providerId) { - const fallbackModels = useProviderModelCacheStore - .getState() - .getModelsForProvider(fallback.providerId); + if ( + fallback && + fallback.providerId !== sessionModelPreference.providerId && + modelCache.isModelInventoryAuthoritative(fallback.providerId) + ) { return sanitizeSessionModelPreference(fallback, { - models: fallbackModels, + models: modelCache.getProvenModelsForProvider(fallback.providerId), }); } } diff --git a/src/features/providers/modelCacheRefresh.ts b/src/features/providers/modelCacheRefresh.ts index fa99a093c..8777240fc 100644 --- a/src/features/providers/modelCacheRefresh.ts +++ b/src/features/providers/modelCacheRefresh.ts @@ -11,7 +11,10 @@ import { getModelProviders, getModelProvidersFromEntries, } from "./providerCatalog"; -import { runtimeRefreshableModelProviderIds } from "./runtimeProviderConfig"; +import { + runtimeManagedModelProviderIds, + runtimeRefreshableModelProviderIds, +} from "./runtimeProviderConfig"; export function getModelCacheRefreshProviderIds( runtimeConfig: RuntimeConfig | null | undefined, @@ -33,10 +36,13 @@ export function getModelCacheRefreshProviderIds( ? new Set(configuredProviderIds) : null; - for (const providerId of runtimeRefreshableModelProviderIds( - runtimeConfig, - defaultModelInventoryMode, - )) { + for (const providerId of [ + ...runtimeManagedModelProviderIds(runtimeConfig, defaultModelInventoryMode), + ...runtimeRefreshableModelProviderIds( + runtimeConfig, + defaultModelInventoryMode, + ), + ]) { ids.add(providerId); } diff --git a/src/features/providers/runtimeProviderConfig.test.ts b/src/features/providers/runtimeProviderConfig.test.ts index 1a8e6dce9..c6f2b4bc7 100644 --- a/src/features/providers/runtimeProviderConfig.test.ts +++ b/src/features/providers/runtimeProviderConfig.test.ts @@ -436,10 +436,21 @@ describe("getModelCacheRefreshProviderIds", () => { ]); }); - it("excludes runtime-managed model providers from startup refresh", () => { - expect(getModelCacheRefreshProviderIds(DEFAULT_RUNTIME_CONFIG)).toEqual([ - "codex-acp", - ]); + it("includes runtime-managed model providers so live discovery can establish proof", () => { + expect( + getModelCacheRefreshProviderIds({ + ...MANAGED_RUNTIME_CONFIG, + goose: { + ...MANAGED_RUNTIME_CONFIG.goose, + modelProviders: [ + { + ...MANAGED_RUNTIME_CONFIG.goose.modelProviders[0], + modelInventoryMode: "authoritative", + }, + ], + }, + }), + ).toEqual(["databricks_v2", "codex-acp"]); }); it("includes model providers for bundled appDefault refresh", () => { diff --git a/src/features/providers/runtimeProviderConfig.ts b/src/features/providers/runtimeProviderConfig.ts index 55f2beced..06adb5d4c 100644 --- a/src/features/providers/runtimeProviderConfig.ts +++ b/src/features/providers/runtimeProviderConfig.ts @@ -164,16 +164,16 @@ function modelInventoryMode( } export function runtimeManagedModelProviderIds( - runtimeConfig: RuntimeConfig, + runtimeConfig: RuntimeConfig | null | undefined, defaultMode: RuntimeModelInventoryMode = DEFAULT_MODEL_INVENTORY_MODE, ): Set { return new Set( - runtimeConfig.goose.modelProviders + runtimeConfig?.goose.modelProviders .filter( (provider) => modelInventoryMode(provider, defaultMode) === "authoritative", ) - .map((provider) => provider.id), + .map((provider) => provider.id) ?? [], ); } diff --git a/src/features/providers/stores/providerModelCacheStore.test.ts b/src/features/providers/stores/providerModelCacheStore.test.ts index f61cf3848..a07199c46 100644 --- a/src/features/providers/stores/providerModelCacheStore.test.ts +++ b/src/features/providers/stores/providerModelCacheStore.test.ts @@ -41,49 +41,73 @@ describe("providerModelCacheStore", () => { }); }); - it("seeds runtime models as authoritative runtime-managed entries", async () => { + it("keeps runtime-managed configuration seeds advisory until live discovery succeeds", async () => { const model = seededModel({ contextLimit: 128000, recommended: true, featured: true, sortOrder: 0, }); - useProviderModelCacheStore .getState() .seedRuntimeModels(new Map([["databricks_v2", [model]]])); + + expect( + useProviderModelCacheStore + .getState() + .isModelInventoryAuthoritative("databricks_v2"), + ).toBe(false); + expect( + useProviderModelCacheStore + .getState() + .getProvenModelsForProvider("databricks_v2"), + ).toEqual([]); + + mocks.supportedModelsList.mockResolvedValueOnce({ models: [] }); await useProviderModelCacheStore .getState() - .refreshAllModelProviders(["databricks_v2"]); - await useProviderModelCacheStore - .getState() - .refreshProviderModels("databricks_v2", { force: true }); + .refreshProviderModels("databricks_v2"); const entry = useProviderModelCacheStore .getState() .providers.get("databricks_v2"); expect(entry?.runtimeManaged).toBe(true); + expect(entry?.provenModelIds).toEqual([]); expect( useProviderModelCacheStore .getState() .getModelsForProvider("databricks_v2"), ).toEqual([model]); - expect(mocks.supportedModelsList).not.toHaveBeenCalled(); - }); - - it("preserves runtime-managed models after invalidation and forced refresh", async () => { - const model = seededModel({ - contextLimit: 128000, - recommended: true, - featured: true, - sortOrder: 0, + expect(mocks.supportedModelsList).toHaveBeenCalledWith({ + providerId: "databricks_v2", }); + }); + it("invalidates runtime-managed proof without discarding its display seed", async () => { + const model = seededModel({ recommended: true, featured: true }); useProviderModelCacheStore .getState() .seedRuntimeModels(new Map([["databricks_v2", [model]]])); + mocks.supportedModelsList + .mockResolvedValueOnce({ models: ["goose-gpt-5-5"] }) + .mockResolvedValueOnce({ models: ["goose-gpt-5-5"] }); + + await useProviderModelCacheStore + .getState() + .refreshProviderModels("databricks_v2"); useProviderModelCacheStore.getState().invalidateProvider("databricks_v2"); + const invalidatedEntry = useProviderModelCacheStore + .getState() + .providers.get("databricks_v2"); + expect(invalidatedEntry?.provenModelIds).toBeUndefined(); + expect(invalidatedEntry?.models).toEqual([model]); + expect( + useProviderModelCacheStore + .getState() + .isModelInventoryAuthoritative("databricks_v2"), + ).toBe(false); + await useProviderModelCacheStore .getState() .refreshProviderModels("databricks_v2", { force: true }); @@ -92,12 +116,16 @@ describe("providerModelCacheStore", () => { .getState() .providers.get("databricks_v2"); expect(entry?.runtimeManaged).toBe(true); + expect(entry?.provenModelIds).toEqual(["goose-gpt-5-5"]); expect( useProviderModelCacheStore .getState() - .getModelsForProvider("databricks_v2"), - ).toEqual([model]); - expect(mocks.supportedModelsList).not.toHaveBeenCalled(); + .getModelsForProvider("databricks_v2") + .map((candidate) => candidate.id), + ).toEqual(["goose-gpt-5-5", "seeded-model"]); + expect(mocks.supportedModelsList).toHaveBeenLastCalledWith({ + providerId: "databricks_v2", + }); }); it("keeps refreshable runtime models provisional until discovery succeeds", async () => { @@ -206,6 +234,12 @@ describe("providerModelCacheStore", () => { expect(models.find((model) => model.id === "goose-gpt-5-6-sol")).toEqual( expect.objectContaining(configuredModel), ); + expect( + useProviderModelCacheStore + .getState() + .getProvenModelsForProvider("databricks_v2") + .map((model) => model.id), + ).toEqual(["goose-gpt-5-5"]); }); it("keeps configured models after a failed refresh and retry", async () => { @@ -239,6 +273,50 @@ describe("providerModelCacheStore", () => { ).toEqual(["goose-gpt-5-5", "goose-gpt-5-6-sol"]); }); + it.each([ + { provenModelIds: ["supported-model"], expected: ["supported-model"] }, + { provenModelIds: [], expected: [] }, + ])("preserves authoritative proof and runtime policy after refresh failure: $provenModelIds", async ({ + provenModelIds, + expected, + }) => { + const model = seededModel({ id: "configured-model" }); + useProviderModelCacheStore.setState({ + providers: new Map([ + [ + "databricks_v2", + { + providerId: "databricks_v2", + models: [model], + configuredModels: [model], + provenModelIds, + fetchedAt: 123, + runtimeManaged: true, + }, + ], + ]), + runtimeManagedProviderIds: new Set(["databricks_v2"]), + }); + mocks.supportedModelsList.mockRejectedValueOnce(new Error("offline")); + + await useProviderModelCacheStore + .getState() + .refreshProviderModels("databricks_v2", { force: true }); + + const entry = useProviderModelCacheStore + .getState() + .providers.get("databricks_v2"); + expect(entry?.provenModelIds).toEqual(expected); + expect(entry?.runtimeManaged).toBe(true); + expect(entry?.configuredModels).toEqual([model]); + expect(entry?.fetchedAt).toBe(123); + expect( + useProviderModelCacheStore + .getState() + .isModelInventoryAuthoritative("databricks_v2"), + ).toBe(true); + }); + it("removes stale runtime-managed providers when runtime config changes", () => { const model = seededModel(); @@ -271,6 +349,56 @@ describe("providerModelCacheStore", () => { ).toBe(true); }); + it("keeps configuration-only runtime seeds provisional across restart", () => { + const model = seededModel({ recommended: true, featured: true }); + useProviderModelCacheStore + .getState() + .seedRuntimeModels(new Map([["databricks_v2", [model]]])); + + useProviderModelCacheStore.setState({ + providers: new Map(), + runtimeManagedProviderIds: new Set(), + }); + useProviderModelCacheStore.getState().loadPersisted(); + + const entry = useProviderModelCacheStore + .getState() + .providers.get("databricks_v2"); + expect(entry?.provenModelIds).toBeUndefined(); + expect( + useProviderModelCacheStore + .getState() + .isModelInventoryAuthoritative("databricks_v2"), + ).toBe(false); + }); + + it("persists authority only after a successful live response", async () => { + const model = seededModel({ recommended: true, featured: true }); + useProviderModelCacheStore + .getState() + .seedRuntimeModels(new Map([["databricks_v2", [model]]])); + mocks.supportedModelsList.mockResolvedValueOnce({ models: [] }); + + await useProviderModelCacheStore + .getState() + .refreshProviderModels("databricks_v2"); + useProviderModelCacheStore.setState({ + providers: new Map(), + runtimeManagedProviderIds: new Set(), + }); + useProviderModelCacheStore.getState().loadPersisted(); + + const entry = useProviderModelCacheStore + .getState() + .providers.get("databricks_v2"); + expect(entry?.provenModelIds).toEqual([]); + expect( + useProviderModelCacheStore + .getState() + .isModelInventoryAuthoritative("databricks_v2"), + ).toBe(true); + }); + it("runs a forced refresh after an in-flight refresh finishes", async () => { let rejectInitialRefresh!: (error: Error) => void; const initialRefresh = new Promise<{ models: string[] }>( diff --git a/src/features/providers/stores/providerModelCacheStore.ts b/src/features/providers/stores/providerModelCacheStore.ts index b69ab3361..5590b3b45 100644 --- a/src/features/providers/stores/providerModelCacheStore.ts +++ b/src/features/providers/stores/providerModelCacheStore.ts @@ -13,8 +13,12 @@ const providerRefreshVersions = new Map(); export interface CachedProviderModels { providerId: string; + /** Display candidates: live models plus configured recommendations. */ models: ModelOption[]; + /** IDs returned by a successful live inventory response; the only proof. */ + provenModelIds?: string[]; fetchedAt: number; + /** Runtime configuration policy; it does not establish model proof. */ runtimeManaged?: boolean; configuredModels?: ModelOption[]; error?: string; @@ -33,6 +37,7 @@ interface ProviderModelCacheActions { options?: { fresh?: boolean; runtimeManagedProviderIds?: Set }, ) => void; getModelsForProvider: (providerId: string) => ModelOption[]; + getProvenModelsForProvider: (providerId: string) => ModelOption[]; isModelInventoryAuthoritative: (providerId: string) => boolean; getError: (providerId: string) => string | null; refreshProviderModels: ( @@ -121,19 +126,62 @@ async function fetchProviderSupportedModels( return response.models; } +function mergeDisplayModels( + discoveredModels: ModelOption[], + configuredModels: ModelOption[], +): ModelOption[] { + const configuredModelsById = new Map( + configuredModels.map((model) => [model.id, model]), + ); + const hasConfiguredFeaturedModel = configuredModels.some( + (model) => model.featured, + ); + const discoveredModelIds = new Set(discoveredModels.map((model) => model.id)); + return [ + ...discoveredModels.map((model) => ({ + ...model, + ...(hasConfiguredFeaturedModel ? { featured: false } : {}), + ...configuredModelsById.get(model.id), + })), + ...configuredModels.filter((model) => !discoveredModelIds.has(model.id)), + ]; +} + +function getProvenModels( + entry: CachedProviderModels | undefined, +): ModelOption[] { + if (!entry?.provenModelIds) return []; + const provenIds = new Set(entry.provenModelIds); + return entry.models.filter((model) => provenIds.has(model.id)); +} + export function isCachedModelInventoryAuthoritative( entry: CachedProviderModels | undefined, ): boolean { - return entry != null && (entry.runtimeManaged || entry.fetchedAt > 0); + return entry != null && Array.isArray(entry.provenModelIds); +} + +/** + * Return whether cached inventory still permits a concrete provider/model pair. + * Missing or provisional inventory cannot disprove a prepared selection; a + * successful authoritative response can. + */ +export function isModelSelectionAllowedByCachedInventory( + providerId: string, + modelId: string, +): boolean { + const entry = useProviderModelCacheStore.getState().providers.get(providerId); + if (!isCachedModelInventoryAuthoritative(entry)) { + return true; + } + return entry?.provenModelIds?.includes(modelId) === true; } function isStale(entry: CachedProviderModels | undefined): boolean { if (!entry || !isCachedModelInventoryAuthoritative(entry)) { return true; } - return ( - !entry.runtimeManaged && Date.now() - entry.fetchedAt > MODEL_CACHE_TTL_MS - ); + return Date.now() - entry.fetchedAt > MODEL_CACHE_TTL_MS; } function refreshVersion(providerId: string): number { @@ -166,15 +214,22 @@ export const useProviderModelCacheStore = create( for (const providerId of runtimeProviderIds) { bumpRefreshVersion(providerId); - const models = modelsByProviderId.get(providerId) ?? []; + const configuredModels = modelsByProviderId.get(providerId) ?? []; const runtimeManaged = runtimeManagedProviderIds.has(providerId); + const existing = providers.get(providerId); + const provenModels = getProvenModels(existing); + const provenModelIds = existing?.provenModelIds; + const hasLiveProof = Array.isArray(provenModelIds); providers.set(providerId, { providerId, - models, - fetchedAt: runtimeManaged || options.fresh ? Date.now() : 0, - ...(runtimeManaged - ? { runtimeManaged } - : { configuredModels: models }), + // Every runtime seed is advisory, including providers whose + // connection policy is runtime-managed. A prior successful live + // response stays proof, but the seed can neither create nor renew it. + models: mergeDisplayModels(provenModels, configuredModels), + fetchedAt: existing?.fetchedAt ?? 0, + configuredModels, + ...(hasLiveProof ? { provenModelIds } : {}), + ...(runtimeManaged ? { runtimeManaged } : {}), }); if (runtimeManaged) { nextRuntimeManagedProviderIds.add(providerId); @@ -202,6 +257,13 @@ export const useProviderModelCacheStore = create( getModelsForProvider: (providerId) => get().providers.get(providerId)?.models ?? [], + getProvenModelsForProvider: (providerId) => { + const entry = get().providers.get(providerId); + if (!entry?.provenModelIds) return []; + const provenIds = new Set(entry.provenModelIds); + return entry.models.filter((model) => provenIds.has(model.id)); + }, + isModelInventoryAuthoritative: (providerId) => isCachedModelInventoryAuthoritative(get().providers.get(providerId)), @@ -210,12 +272,6 @@ export const useProviderModelCacheStore = create( refreshProviderModels: async (providerId, options = {}) => { const current = get(); const existing = current.providers.get(providerId); - if ( - existing?.runtimeManaged || - current.runtimeManagedProviderIds.has(providerId) - ) { - return; - } if (!options.force && !isStale(existing)) { return; } @@ -260,29 +316,13 @@ export const useProviderModelCacheStore = create( const ids = await fetchProviderSupportedModels(providerId); const discoveredModels = providerModelOptionsFromIds(providerId, ids); const configuredModels = existing?.configuredModels ?? []; - const configuredModelsById = new Map( - configuredModels.map((model) => [model.id, model]), - ); - const hasConfiguredFeaturedModel = configuredModels.some( - (model) => model.featured, - ); - const discoveredModelIds = new Set( - discoveredModels.map((model) => model.id), - ); - const models = [ - ...discoveredModels.map((model) => ({ - ...model, - ...(hasConfiguredFeaturedModel ? { featured: false } : {}), - ...configuredModelsById.get(model.id), - })), - ...configuredModels.filter( - (model) => !discoveredModelIds.has(model.id), - ), - ]; + const models = mergeDisplayModels(discoveredModels, configuredModels); const entry: CachedProviderModels = { providerId, models, fetchedAt: Date.now(), + provenModelIds: ids, + ...(existing?.runtimeManaged ? { runtimeManaged: true } : {}), ...(configuredModels.length > 0 ? { configuredModels } : {}), }; if (versionAtStart !== refreshVersion(providerId)) { @@ -305,6 +345,12 @@ export const useProviderModelCacheStore = create( providerId, models: existing?.models ?? [], fetchedAt: existing?.fetchedAt ?? 0, + ...(existing?.provenModelIds + ? { provenModelIds: existing.provenModelIds } + : isCachedModelInventoryAuthoritative(existing) + ? { provenModelIds: [] } + : {}), + ...(existing?.runtimeManaged ? { runtimeManaged: true } : {}), ...(existing?.configuredModels ? { configuredModels: existing.configuredModels } : {}), @@ -341,13 +387,17 @@ export const useProviderModelCacheStore = create( invalidateProvider: (providerId) => { bumpRefreshVersion(providerId); set((state) => { - if (state.runtimeManagedProviderIds.has(providerId)) { - const existing = state.providers.get(providerId); - if (!existing || existing.runtimeManaged) { - return {}; - } + const existing = state.providers.get(providerId); + if (state.runtimeManagedProviderIds.has(providerId) && existing) { const providers = new Map(state.providers); - providers.set(providerId, { ...existing, runtimeManaged: true }); + // Keep the configured display seed but remove proof: an invalidation + // means the old live response can no longer justify compatibility. + providers.set(providerId, { + ...existing, + models: existing.configuredModels ?? existing.models, + fetchedAt: 0, + provenModelIds: undefined, + }); persistModels(providers); return { providers }; } diff --git a/src/features/sessions/hooks/__tests__/useAutoArchiveSessions.test.ts b/src/features/sessions/hooks/__tests__/useAutoArchiveSessions.test.ts index 543c0cd3e..930349a22 100644 --- a/src/features/sessions/hooks/__tests__/useAutoArchiveSessions.test.ts +++ b/src/features/sessions/hooks/__tests__/useAutoArchiveSessions.test.ts @@ -14,6 +14,7 @@ const mocks = vi.hoisted(() => ({ })); vi.mock("@/shared/api/acp", () => ({ + reserveAcpSessionConfiguration: () => ({ sequence: 0, clear: () => {} }), acpGetSessionInfo: (...args: unknown[]) => mocks.getSessionInfo(...args), })); diff --git a/src/features/sessions/hooks/__tests__/useSessionSearch.test.ts b/src/features/sessions/hooks/__tests__/useSessionSearch.test.ts index 10d574d72..06d1943fb 100644 --- a/src/features/sessions/hooks/__tests__/useSessionSearch.test.ts +++ b/src/features/sessions/hooks/__tests__/useSessionSearch.test.ts @@ -47,6 +47,7 @@ function sweep( } vi.mock("@/shared/api/acp", () => ({ + reserveAcpSessionConfiguration: () => ({ sequence: 0, clear: () => {} }), acpSearchSessions: (...args: unknown[]) => mockAcpSearchSessions(...args), })); diff --git a/src/features/sessions/ui/__tests__/SessionHistoryView.test.tsx b/src/features/sessions/ui/__tests__/SessionHistoryView.test.tsx index b07a6cd5d..be8035709 100644 --- a/src/features/sessions/ui/__tests__/SessionHistoryView.test.tsx +++ b/src/features/sessions/ui/__tests__/SessionHistoryView.test.tsx @@ -39,6 +39,7 @@ const mocks = vi.hoisted(() => ({ })); vi.mock("@/shared/api/acp", () => ({ + reserveAcpSessionConfiguration: () => ({ sequence: 0, clear: () => {} }), acpExportSession: (...args: unknown[]) => mocks.acpExportSession(...args), acpImportSession: (...args: unknown[]) => mocks.acpImportSession(...args), acpSearchSessions: (...args: unknown[]) => mocks.acpSearchSessions(...args), diff --git a/src/shared/api/__tests__/acp.test.ts b/src/shared/api/__tests__/acp.test.ts index 31b8d224e..45a9f62c6 100644 --- a/src/shared/api/__tests__/acp.test.ts +++ b/src/shared/api/__tests__/acp.test.ts @@ -140,6 +140,16 @@ vi.mock("@/features/berdctl/appPreamble", () => ({ getBerdctlPreamble: () => mockGetBerdctlPreamble(), })); +const mockSupportedModelsList = vi.hoisted(() => vi.fn()); +vi.mock("../acpConnection", () => ({ + getClient: () => + Promise.resolve({ + goose: { + GooseUnstableProvidersSupportedModelsList: mockSupportedModelsList, + }, + }), +})); + vi.mock("../acpActiveMessageTracking", () => ({ setActiveMessageId: vi.fn(), clearActiveMessageId: vi.fn(), @@ -199,6 +209,134 @@ describe("acpSendMessage", () => { expect(mockPrompt).not.toHaveBeenCalled(); }); + it("blocks a prepared model disproved by cached authoritative inventory without network I/O", async () => { + await setRuntimeConfig(managedRuntimeConfig); + const { useProviderModelCacheStore } = await import( + "@/features/providers/stores/providerModelCacheStore" + ); + const sessionRegistry = await import("../acpSessionRegistry"); + const { acpSendMessage } = await import("../acp"); + sessionRegistry.registerPreparedSession( + "acp-session-invalidated-model", + "databricks_v2", + "/tmp/project", + "removed-model", + ); + useProviderModelCacheStore.setState({ + providers: new Map([ + [ + "databricks_v2", + { + providerId: "databricks_v2", + models: [ + { + id: "supported-model", + name: "Supported model", + providerId: "databricks_v2", + }, + ], + provenModelIds: ["supported-model"], + fetchedAt: Date.now(), + }, + ], + ]), + }); + + await expect( + acpSendMessage("acp-session-invalidated-model", "hello"), + ).rejects.toThrow("removed-model is no longer supported"); + + expect(mockSupportedModelsList).not.toHaveBeenCalled(); + expect(mockPrompt).not.toHaveBeenCalled(); + }); + + it("keeps a disproved prepared model blocked after a failed forced refresh", async () => { + await setRuntimeConfig(managedRuntimeConfig); + const { useProviderModelCacheStore } = await import( + "@/features/providers/stores/providerModelCacheStore" + ); + const sessionRegistry = await import("../acpSessionRegistry"); + const { acpSendMessage } = await import("../acp"); + sessionRegistry.registerPreparedSession( + "acp-session-model-disproved-before-refresh-failure", + "databricks_v2", + "/tmp/project", + "removed-model", + ); + mockSupportedModelsList.mockResolvedValueOnce({ + models: ["supported-model"], + }); + await useProviderModelCacheStore + .getState() + .refreshProviderModels("databricks_v2", { force: true }); + + await expect( + acpSendMessage( + "acp-session-model-disproved-before-refresh-failure", + "before failure", + ), + ).rejects.toThrow("removed-model is no longer supported"); + + mockSupportedModelsList.mockRejectedValueOnce(new Error("offline")); + await useProviderModelCacheStore + .getState() + .refreshProviderModels("databricks_v2", { force: true }); + const inventoryCallsBeforeSend = mockSupportedModelsList.mock.calls.length; + + await expect( + acpSendMessage( + "acp-session-model-disproved-before-refresh-failure", + "after failure", + ), + ).rejects.toThrow("removed-model is no longer supported"); + + expect(mockSupportedModelsList).toHaveBeenCalledTimes( + inventoryCallsBeforeSend, + ); + expect(mockPrompt).not.toHaveBeenCalled(); + expect( + useProviderModelCacheStore.getState().providers.get("databricks_v2") + ?.provenModelIds, + ).toEqual(["supported-model"]); + useProviderModelCacheStore.getState().invalidateProvider("databricks_v2"); + }); + + it("admits a managed-provider prompt without reading live model inventory", async () => { + await setRuntimeConfig(managedRuntimeConfig); + const sessionRegistry = await import("../acpSessionRegistry"); + const { acpSendMessage } = await import("../acp"); + sessionRegistry.registerPreparedSession( + "acp-session-managed-send", + "databricks_v2", + "/tmp/project", + "goose-gpt-5-5", + ); + + await acpSendMessage("acp-session-managed-send", "hello"); + + expect(mockSupportedModelsList).not.toHaveBeenCalled(); + expect(mockPrompt).toHaveBeenCalledOnce(); + }); + + it("rejects an out-of-policy provider without reading live model inventory", async () => { + await setRuntimeConfig(managedRuntimeConfig); + const sessionRegistry = await import("../acpSessionRegistry"); + const { acpSendMessage } = await import("../acp"); + sessionRegistry.registerPreparedSession( + "acp-session-outside-policy", + "outside-policy", + "/tmp/project", + "outside-model", + ); + + await expect( + acpSendMessage("acp-session-outside-policy", "hello"), + ).rejects.toThrow("outside the managed Goose provider policy"); + + expect(mockSupportedModelsList).not.toHaveBeenCalled(); + expect(mockPrompt).not.toHaveBeenCalled(); + }); + it("reports dispatch only after ACP setup reaches the transport boundary", async () => { const sessionRegistry = await import("../acpSessionRegistry"); const { acpSendMessage } = await import("../acp"); @@ -752,6 +890,302 @@ describe("acpLoadSession", () => { ); }); + it("does not dispatch a load snapshot when provider-changing prepare is awaiting inventory proof", async () => { + await setRuntimeConfig(managedRuntimeConfig); + const supportedModels = deferred<{ models: string[] }>(); + const loadResponse = deferred>(); + mockSupportedModelsList.mockReturnValueOnce(supportedModels.promise); + mockLoadSession.mockReturnValueOnce(loadResponse.promise); + mockSetProvider.mockResolvedValueOnce({ + model: null, + reasoningEffort: null, + }); + mockSetModel.mockResolvedValueOnce({ model: null, reasoningEffort: null }); + const applyModelConfigSnapshot = vi.fn(); + const { setSessionConfigSnapshotHandlers } = await import( + "../acpSessionConfigSnapshots" + ); + setSessionConfigSnapshotHandlers({ applyModelConfigSnapshot }); + const { acpLoadSession, acpPrepareSession } = await import("../acp"); + + const configure = acpPrepareSession( + "acp-session-preflight-race", + "goose", + "/tmp/replay", + { modelId: "other-model" }, + ); + await vi.waitFor(() => expect(mockSupportedModelsList).toHaveBeenCalled()); + + const load = acpLoadSession("acp-session-preflight-race", "/tmp/replay"); + await vi.waitFor(() => expect(mockLoadSession).toHaveBeenCalledTimes(1)); + loadResponse.resolve( + executionConfigResponse("other-managed", "other-model"), + ); + await load; + + expect(applyModelConfigSnapshot).not.toHaveBeenCalled(); + expect(mockSetProvider).not.toHaveBeenCalled(); + expect(mockSetModel).not.toHaveBeenCalled(); + + supportedModels.resolve({ models: ["goose-gpt-5-5"] }); + await configure; + + expect(mockSetProvider).toHaveBeenCalledWith( + "acp-session-preflight-race", + "databricks_v2", + noRequestProviderContext, + ); + expect(mockSetModel).toHaveBeenCalledWith( + "acp-session-preflight-race", + "goose-gpt-5-5", + noRequestModelContext("databricks_v2"), + ); + }); + + it("publishes a deferred authoritative load when preflight rejects before mutation", async () => { + await setRuntimeConfig(managedRuntimeConfig); + const loadResponse = deferred>(); + const supportedModels = deferred<{ models: string[] }>(); + mockLoadSession.mockReturnValueOnce(loadResponse.promise); + mockSupportedModelsList.mockReturnValueOnce(supportedModels.promise); + const applyModelConfigSnapshot = vi.fn(); + const { setSessionConfigSnapshotHandlers } = await import( + "../acpSessionConfigSnapshots" + ); + setSessionConfigSnapshotHandlers({ applyModelConfigSnapshot }); + const { acpLoadSession, acpPrepareSession } = await import("../acp"); + + const load = acpLoadSession( + "acp-session-load-before-rejected-preflight", + "/tmp/replay", + ); + await vi.waitFor(() => expect(mockLoadSession).toHaveBeenCalledTimes(1)); + const configure = acpPrepareSession( + "acp-session-load-before-rejected-preflight", + "goose", + "/tmp/replay", + { modelId: "other-model" }, + ); + await vi.waitFor(() => expect(mockSupportedModelsList).toHaveBeenCalled()); + + loadResponse.resolve( + executionConfigResponse("other-managed", "other-model"), + ); + await load; + supportedModels.reject(new Error("offline")); + await expect(configure).rejects.toThrow( + "Cannot verify models for migrated provider", + ); + await Promise.resolve(); + + expect(applyModelConfigSnapshot).toHaveBeenCalledWith( + "acp-session-load-before-rejected-preflight", + { modelId: "other-model", modelName: "other-model" }, + { origin: "response" }, + ); + expect(mockSetProvider).not.toHaveBeenCalled(); + expect(mockSetModel).not.toHaveBeenCalled(); + const { requireSessionInvocationSelection } = await import( + "../acpSessionRegistry" + ); + expect( + requireSessionInvocationSelection( + "acp-session-load-before-rejected-preflight", + ), + ).toEqual({ providerId: "other-managed", modelId: "other-model" }); + }); + + it.each([ + { + name: "setProvider", + load: executionConfigResponse("other-managed", "other-model"), + modelId: "other-model", + reject: () => mockSetProvider.mockRejectedValueOnce(new Error("offline")), + expected: () => expect(mockSetProvider).toHaveBeenCalledTimes(1), + }, + { + name: "setModel", + load: executionConfigResponse("databricks_v2", "old-model"), + modelId: "goose-gpt-5-5", + reject: () => mockSetModel.mockRejectedValueOnce(new Error("offline")), + expected: () => expect(mockSetModel).toHaveBeenCalledTimes(1), + }, + ])("does not publish a deferred load after attempted $name fails", async ({ + name, + load, + modelId, + reject, + expected, + }) => { + await setRuntimeConfig(managedRuntimeConfig); + const supportedModels = deferred<{ models: string[] }>(); + mockSupportedModelsList.mockReturnValueOnce(supportedModels.promise); + mockLoadSession.mockResolvedValueOnce(load); + reject(); + const applyModelConfigSnapshot = vi.fn(); + const { setSessionConfigSnapshotHandlers } = await import( + "../acpSessionConfigSnapshots" + ); + setSessionConfigSnapshotHandlers({ applyModelConfigSnapshot }); + const sessionRegistry = await import("../acpSessionRegistry"); + const { acpLoadSession, acpPrepareSession } = await import("../acp"); + const sessionId = `acp-session-failed-${modelId}`; + + const configure = acpPrepareSession(sessionId, "goose", "/tmp/replay", { + modelId, + }); + await vi.waitFor(() => expect(mockSupportedModelsList).toHaveBeenCalled()); + await acpLoadSession(sessionId, "/tmp/replay"); + expect(applyModelConfigSnapshot).not.toHaveBeenCalled(); + + supportedModels.resolve({ models: ["goose-gpt-5-5"] }); + await expect(configure).rejects.toThrow("offline"); + expected(); + if (name === "setProvider") { + expect(sessionRegistry.isSessionPrepared(sessionId)).toBe(false); + } + await Promise.resolve(); + expect(applyModelConfigSnapshot).not.toHaveBeenCalled(); + }); + + it("does not let a stale preflight consume a newer preflight intent or publish a load", async () => { + await setRuntimeConfig(managedRuntimeConfig); + const firstInventory = deferred<{ models: string[] }>(); + const secondInventory = deferred<{ models: string[] }>(); + mockSupportedModelsList + .mockReturnValueOnce(firstInventory.promise) + .mockReturnValueOnce(secondInventory.promise); + mockSetProvider.mockResolvedValue({ model: null, reasoningEffort: null }); + mockSetModel.mockResolvedValue({ model: null, reasoningEffort: null }); + const applyModelConfigSnapshot = vi.fn(); + const { setSessionConfigSnapshotHandlers } = await import( + "../acpSessionConfigSnapshots" + ); + setSessionConfigSnapshotHandlers({ applyModelConfigSnapshot }); + const { acpLoadSession, acpPrepareSession } = await import("../acp"); + + const first = acpPrepareSession( + "acp-session-two-preflights", + "goose", + "/tmp/replay", + { + modelId: "other-model", + }, + ); + await vi.waitFor(() => + expect(mockSupportedModelsList).toHaveBeenCalledTimes(1), + ); + const second = acpPrepareSession( + "acp-session-two-preflights", + "goose", + "/tmp/replay", + { + modelId: "other-model", + }, + ); + await vi.waitFor(() => + expect(mockSupportedModelsList).toHaveBeenCalledTimes(2), + ); + + firstInventory.resolve({ models: ["goose-gpt-5-5"] }); + await first; + expect(mockSetProvider).not.toHaveBeenCalled(); + expect(mockSetModel).not.toHaveBeenCalled(); + + await acpLoadSession("acp-session-two-preflights", "/tmp/replay"); + expect(applyModelConfigSnapshot).not.toHaveBeenCalled(); + + secondInventory.resolve({ models: ["goose-gpt-5-5"] }); + await second; + expect(mockSetProvider).toHaveBeenCalledTimes(1); + expect(mockSetModel).toHaveBeenCalledTimes(1); + }); + + it("releases only a rejected latest preflight intent", async () => { + await setRuntimeConfig(managedRuntimeConfig); + const firstInventory = deferred<{ models: string[] }>(); + mockSupportedModelsList + .mockReturnValueOnce(firstInventory.promise) + .mockRejectedValueOnce(new Error("offline")); + mockLoadSession.mockResolvedValueOnce( + executionConfigResponse("other-managed", "other-model"), + ); + const applyModelConfigSnapshot = vi.fn(); + const { setSessionConfigSnapshotHandlers } = await import( + "../acpSessionConfigSnapshots" + ); + setSessionConfigSnapshotHandlers({ applyModelConfigSnapshot }); + const { acpLoadSession, acpPrepareSession } = await import("../acp"); + + const first = acpPrepareSession( + "acp-session-rejected-latest", + "goose", + "/tmp/replay", + { + modelId: "other-model", + }, + ); + await vi.waitFor(() => + expect(mockSupportedModelsList).toHaveBeenCalledTimes(1), + ); + await expect( + acpPrepareSession("acp-session-rejected-latest", "goose", "/tmp/replay", { + modelId: "other-model", + }), + ).rejects.toThrow("Cannot verify models for migrated provider"); + + firstInventory.resolve({ models: ["goose-gpt-5-5"] }); + await first; + await acpLoadSession("acp-session-rejected-latest", "/tmp/replay"); + + expect(applyModelConfigSnapshot).toHaveBeenCalledTimes(1); + expect(mockSetProvider).not.toHaveBeenCalled(); + expect(mockSetModel).not.toHaveBeenCalled(); + }); + + it("keeps the latest preflight when it resolves before an older one", async () => { + await setRuntimeConfig(managedRuntimeConfig); + const firstInventory = deferred<{ models: string[] }>(); + const secondInventory = deferred<{ models: string[] }>(); + mockSupportedModelsList + .mockReturnValueOnce(firstInventory.promise) + .mockReturnValueOnce(secondInventory.promise); + mockSetProvider.mockResolvedValue({ model: null, reasoningEffort: null }); + mockSetModel.mockResolvedValue({ model: null, reasoningEffort: null }); + const { acpPrepareSession } = await import("../acp"); + + const first = acpPrepareSession( + "acp-session-reverse-preflights", + "goose", + "/tmp/replay", + { + modelId: "other-model", + }, + ); + await vi.waitFor(() => + expect(mockSupportedModelsList).toHaveBeenCalledTimes(1), + ); + const second = acpPrepareSession( + "acp-session-reverse-preflights", + "goose", + "/tmp/replay", + { + modelId: "other-model", + }, + ); + await vi.waitFor(() => + expect(mockSupportedModelsList).toHaveBeenCalledTimes(2), + ); + + secondInventory.resolve({ models: ["goose-gpt-5-5"] }); + await second; + firstInventory.resolve({ models: ["goose-gpt-5-5"] }); + await first; + + expect(mockSetProvider).toHaveBeenCalledTimes(1); + expect(mockSetModel).toHaveBeenCalledTimes(1); + }); + it("does not dispatch a load snapshot superseded by a UI configuration", async () => { const loadResponse = deferred>(); mockLoadSession.mockReturnValueOnce(loadResponse.promise); @@ -799,6 +1233,8 @@ describe("acpCreateSession", () => { mockSetProvider.mockResolvedValue({ model: null, reasoningEffort: null }); mockSetModel.mockReset(); mockSetModel.mockResolvedValue({ model: null, reasoningEffort: null }); + mockSupportedModelsList.mockReset(); + mockSupportedModelsList.mockResolvedValue({ models: ["goose-gpt-5-5"] }); await setRuntimeConfig(DEFAULT_RUNTIME_CONFIG); }); @@ -1032,6 +1468,64 @@ describe("acpCreateSession", () => { }); }); + it("applies the complete resolved migration pair for provider-only input", async () => { + await setRuntimeConfig(managedRuntimeConfig); + mockSupportedModelsList.mockResolvedValueOnce({ + models: ["goose-gpt-5-5"], + }); + mockNewSession.mockResolvedValue({ sessionId: "migrated-session" }); + const { acpCreateSession } = await import("../acp"); + + await acpCreateSession("goose", "/tmp/project"); + + expect(mockSetProvider).toHaveBeenCalledWith( + "migrated-session", + "databricks_v2", + ); + expect(mockSetModel).toHaveBeenCalledWith( + "migrated-session", + "goose-gpt-5-5", + noRequestModelContext("databricks_v2"), + ); + }); + + it("does not mutate ACP when managed provider migration cannot prove support", async () => { + await setRuntimeConfig(managedRuntimeConfig); + mockSupportedModelsList.mockRejectedValueOnce(new Error("offline")); + const { acpCreateSession } = await import("../acp"); + + await expect( + acpCreateSession("goose", "/tmp/project", { modelId: "other-model" }), + ).rejects.toThrow( + "Cannot verify models for migrated provider databricks_v2", + ); + + expect(mockNewSession).not.toHaveBeenCalled(); + expect(mockSetProvider).not.toHaveBeenCalled(); + expect(mockSetModel).not.toHaveBeenCalled(); + }); + + it("uses a proven default instead of an unsupported migrated model", async () => { + await setRuntimeConfig(managedRuntimeConfig); + mockSupportedModelsList.mockResolvedValueOnce({ + models: ["goose-gpt-5-5"], + }); + mockNewSession.mockResolvedValue({ sessionId: "migrated-session" }); + const { acpCreateSession } = await import("../acp"); + + await acpCreateSession("goose", "/tmp/project", { modelId: "other-model" }); + + expect(mockSetProvider).toHaveBeenCalledWith( + "migrated-session", + "databricks_v2", + ); + expect(mockSetModel).toHaveBeenCalledWith( + "migrated-session", + "goose-gpt-5-5", + noRequestModelContext("databricks_v2"), + ); + }); + it("rejects an explicit provider outside managed policy before creating", async () => { await setRuntimeConfig(managedRuntimeConfig); const { acpCreateSession } = await import("../acp"); @@ -1059,6 +1553,34 @@ describe("acpCreateSession", () => { expect(mockSetModel).not.toHaveBeenCalled(); }); + it("does not send a same-provider model disproved by live inventory", async () => { + await setRuntimeConfig(managedRuntimeConfig); + mockSupportedModelsList.mockResolvedValueOnce({ + models: ["goose-gpt-5-5"], + }); + mockNewSession.mockResolvedValue({ sessionId: "managed-session" }); + const { acpCreateSession } = await import("../acp"); + + await acpCreateSession("databricks_v2", "/tmp/project", { + modelId: "retired-model", + }); + + expect(mockSetProvider).toHaveBeenCalledWith( + "managed-session", + "databricks_v2", + ); + expect(mockSetModel).toHaveBeenCalledWith( + "managed-session", + "goose-gpt-5-5", + noRequestModelContext("databricks_v2"), + ); + expect(mockSetModel).not.toHaveBeenCalledWith( + "managed-session", + "retired-model", + expect.anything(), + ); + }); + it.each([ "claude-acp", "codex-acp", @@ -1177,6 +1699,8 @@ describe("acpPrepareSession", () => { beforeEach(async () => { vi.clearAllMocks(); vi.resetModules(); + mockSupportedModelsList.mockReset(); + mockSupportedModelsList.mockResolvedValue({ models: ["goose-gpt-5-5"] }); await setRuntimeConfig(DEFAULT_RUNTIME_CONFIG); }); @@ -1203,6 +1727,162 @@ describe("acpPrepareSession", () => { expect(sessionRegistry.isSessionPrepared("acp-session-1")).toBe(true); }); + it("keeps a caller-owned configuration intent after prepare rejects", async () => { + await setRuntimeConfig(managedRuntimeConfig); + mockSupportedModelsList.mockRejectedValueOnce(new Error("offline")); + const applyModelConfigSnapshot = vi.fn(); + const { setSessionConfigSnapshotHandlers } = await import( + "../acpSessionConfigSnapshots" + ); + setSessionConfigSnapshotHandlers({ applyModelConfigSnapshot }); + const { + acpLoadSession, + acpPrepareSession, + reserveAcpSessionConfiguration, + } = await import("../acp"); + const sessionId = "acp-session-caller-owned-intent"; + const intent = reserveAcpSessionConfiguration(sessionId); + + await expect( + acpPrepareSession( + sessionId, + "goose", + "/tmp/project", + { + modelId: "other-model", + }, + intent, + ), + ).rejects.toThrow("Cannot verify models for migrated provider"); + mockLoadSession.mockResolvedValueOnce( + executionConfigResponse("openai", "gpt-5.5"), + ); + await acpLoadSession(sessionId, "/tmp/project"); + expect(applyModelConfigSnapshot).not.toHaveBeenCalled(); + + intent.clear(); + mockLoadSession.mockResolvedValueOnce( + executionConfigResponse("openai", "gpt-5.5"), + ); + await acpLoadSession(sessionId, "/tmp/project"); + expect(applyModelConfigSnapshot).toHaveBeenCalledTimes(1); + }); + + it("releases timed-out inventory intent so fresh load and retry can reconcile", async () => { + await setRuntimeConfig(managedRuntimeConfig); + vi.useFakeTimers(); + const timedOutInventory = deferred<{ models: string[] }>(); + mockSupportedModelsList + .mockReturnValueOnce(timedOutInventory.promise) + .mockResolvedValueOnce({ models: ["goose-gpt-5-5"] }); + mockLoadSession.mockResolvedValue( + executionConfigResponse("openai", "gpt-5.5"), + ); + mockSetProvider.mockResolvedValue({ model: null, reasoningEffort: null }); + mockSetModel.mockResolvedValue({ model: null, reasoningEffort: null }); + const applyModelConfigSnapshot = vi.fn(); + const { setSessionConfigSnapshotHandlers } = await import( + "../acpSessionConfigSnapshots" + ); + setSessionConfigSnapshotHandlers({ applyModelConfigSnapshot }); + const { acpLoadSession, acpPrepareSession } = await import("../acp"); + const sessionId = "acp-session-timeout-retry"; + + const timedOutPrepare = acpPrepareSession( + sessionId, + "goose", + "/tmp/project", + { + modelId: "other-model", + }, + ); + const rejectedPrepare = expect(timedOutPrepare).rejects.toThrow( + "Cannot verify models for migrated provider", + ); + await vi.advanceTimersByTimeAsync(60_000); + await rejectedPrepare; + + expect(mockLoadSession).not.toHaveBeenCalled(); + expect(mockSetProvider).not.toHaveBeenCalled(); + expect(mockSetModel).not.toHaveBeenCalled(); + timedOutInventory.resolve({ models: ["goose-gpt-5-5"] }); + await Promise.resolve(); + expect(mockSetProvider).not.toHaveBeenCalled(); + expect(mockSetModel).not.toHaveBeenCalled(); + + await acpLoadSession(sessionId, "/tmp/project"); + expect(applyModelConfigSnapshot).toHaveBeenCalledTimes(1); + await acpPrepareSession(sessionId, "goose", "/tmp/project", { + modelId: "other-model", + }); + expect(mockSetProvider).toHaveBeenCalledWith( + sessionId, + "databricks_v2", + noRequestProviderContext, + ); + expect(mockSetModel).toHaveBeenCalledWith( + sessionId, + "goose-gpt-5-5", + noRequestModelContext("databricks_v2"), + ); + vi.useRealTimers(); + }); + + it("rejects invalidated inventory proof and requires a fresh proof before mutation", async () => { + await setRuntimeConfig(managedRuntimeConfig); + const staleInventory = deferred<{ models: string[] }>(); + mockSupportedModelsList + .mockReturnValueOnce(staleInventory.promise) + .mockResolvedValueOnce({ models: ["goose-gpt-5-5"] }); + mockSetProvider.mockResolvedValue({ model: null, reasoningEffort: null }); + mockSetModel.mockResolvedValue({ model: null, reasoningEffort: null }); + const { acpPrepareSession } = await import("../acp"); + const { notifyProviderModelInventoryInvalidated } = await import( + "@/shared/runtime-config/providerModelInventoryInvalidation" + ); + const sessionId = "acp-session-invalidated-proof"; + + const stalePrepare = acpPrepareSession(sessionId, "goose", "/tmp/project", { + modelId: "other-model", + }); + await vi.waitFor(() => + expect(mockSupportedModelsList).toHaveBeenCalledTimes(1), + ); + notifyProviderModelInventoryInvalidated("databricks_v2"); + staleInventory.resolve({ models: ["goose-gpt-5-5"] }); + await expect(stalePrepare).rejects.toThrow( + "Cannot verify models for migrated provider", + ); + + expect(mockLoadSession).not.toHaveBeenCalled(); + expect(mockSetProvider).not.toHaveBeenCalled(); + expect(mockSetModel).not.toHaveBeenCalled(); + await acpPrepareSession(sessionId, "goose", "/tmp/project", { + modelId: "other-model", + }); + expect(mockSupportedModelsList).toHaveBeenCalledTimes(2); + expect(mockSetProvider).toHaveBeenCalledTimes(1); + expect(mockSetModel).toHaveBeenCalledTimes(1); + }); + + it("does not load or mutate a session when managed migration cannot prove support", async () => { + await setRuntimeConfig(managedRuntimeConfig); + mockSupportedModelsList.mockRejectedValueOnce(new Error("offline")); + const { acpPrepareSession } = await import("../acp"); + + await expect( + acpPrepareSession("legacy-session", "goose", "/tmp/project", { + modelId: "other-model", + }), + ).rejects.toThrow( + "Cannot verify models for migrated provider databricks_v2", + ); + + expect(mockLoadSession).not.toHaveBeenCalled(); + expect(mockSetProvider).not.toHaveBeenCalled(); + expect(mockSetModel).not.toHaveBeenCalled(); + }); + it("rejects a provider outside managed policy before loading the session", async () => { await setRuntimeConfig(managedRuntimeConfig); const { acpPrepareSession } = await import("../acp"); @@ -1231,6 +1911,9 @@ describe("acpPrepareSession", () => { it("allows upstream models omitted from recommendation metadata", async () => { await setRuntimeConfig(managedRuntimeConfig); + mockSupportedModelsList.mockResolvedValueOnce({ + models: ["new-upstream-model"], + }); const { acpPrepareSession } = await import("../acp"); await acpPrepareSession("other-session", "other-managed", "/tmp/project", { diff --git a/src/shared/api/__tests__/acpSessionRegistry.test.ts b/src/shared/api/__tests__/acpSessionRegistry.test.ts index 38c34bb20..46d170d12 100644 --- a/src/shared/api/__tests__/acpSessionRegistry.test.ts +++ b/src/shared/api/__tests__/acpSessionRegistry.test.ts @@ -305,6 +305,42 @@ describe("applySessionModel", () => { expect(prompt).toHaveBeenCalledWith("anthropic"); }); + it("holds prompt transport behind pending configuration intent until it clears", async () => { + const registry = await importPreparedRegistry("openai", "gpt-4.1"); + const supersession = registry.supersedeSessionMutation("session-1"); + const prompt = vi.fn().mockResolvedValue("complete"); + + const result = registry.runPreparedSessionPrompt("session-1", prompt); + await Promise.resolve(); + expect(prompt).not.toHaveBeenCalled(); + + supersession.clear(); + await expect(result).resolves.toBe("complete"); + expect(prompt).toHaveBeenCalledWith("openai"); + }); + + it("holds prompt transport until pending configuration is consumed", async () => { + const registry = await importPreparedRegistry("openai", "gpt-4.1"); + const supersession = registry.supersedeSessionMutation("session-1"); + const prompt = vi.fn().mockResolvedValue("complete"); + + const result = registry.runPreparedSessionPrompt("session-1", prompt); + const configure = registry.configureSession( + "session-1", + "anthropic", + "/project", + "claude-fable", + {}, + supersession, + ); + await expect(configure).resolves.toEqual({ + model: { modelId: "claude-fable", modelName: "claude-fable" }, + reasoningEffort: null, + }); + await expect(result).resolves.toBe("complete"); + expect(prompt).toHaveBeenCalledWith("anthropic"); + }); + it("does not time out a long-running prompt or admit config work mid-turn", async () => { vi.useFakeTimers(); try { @@ -382,6 +418,43 @@ describe("applySessionModel", () => { ); }); + it("keeps a mutation behind an owned preflight configuration after its old tail settles", async () => { + const registry = await importPreparedRegistry("openai", "gpt-4.1"); + const firstProviderResponse = deferred(); + mockSetProvider.mockReturnValueOnce(firstProviderResponse.promise); + + const supersession = registry.supersedeSessionMutation("session-1"); + const first = registry.prepareSession( + "session-1", + "anthropic", + "/project", + {}, + supersession, + ); + + await vi.waitFor(() => expect(mockSetProvider).toHaveBeenCalledTimes(1)); + + // The owned intent was consumed before the real mutation was appended. A + // cleanup registered against the formerly resolved tail must not reclaim + // the queue while this provider call is still in flight. + const second = registry.prepareSession("session-1", "gemini", "/project"); + await Promise.resolve(); + expect(mockSetProvider).toHaveBeenCalledTimes(1); + + firstProviderResponse.resolve( + modelConfigResponse("claude-default", "Claude Default"), + ); + await first; + await second; + + expect(mockSetProvider).toHaveBeenCalledTimes(2); + expect(mockSetProvider).toHaveBeenLastCalledWith( + "session-1", + "gemini", + noRequestProviderContext, + ); + }); + it("does not run a load between one provider and model configuration", async () => { const registry = await importPreparedRegistry("openai", "gpt-4.1"); const providerResponse = deferred(); @@ -420,6 +493,26 @@ describe("applySessionModel", () => { expect(mockSetModel).toHaveBeenCalledTimes(1); }); + it("returns the acknowledged requested model when setModel omits a snapshot", async () => { + const registry = await importPreparedRegistry("openai", "gpt-4.1"); + mockSetProvider.mockResolvedValueOnce( + modelConfigResponse("claude-sonnet", "Claude Sonnet"), + ); + mockSetModel.mockResolvedValueOnce(undefined); + + await expect( + registry.configureSession( + "session-1", + "anthropic", + "/project", + "claude-fable", + ), + ).resolves.toEqual({ + model: { modelId: "claude-fable", modelName: "claude-fable" }, + reasoningEffort: null, + }); + }); + it("returns the final model snapshot without provider-default fields", async () => { const registry = await importPreparedRegistry("openai", "gpt-5.5"); mockSetProvider.mockResolvedValueOnce( diff --git a/src/shared/api/acp.ts b/src/shared/api/acp.ts index 4817b03fb..259ba8335 100644 --- a/src/shared/api/acp.ts +++ b/src/shared/api/acp.ts @@ -28,7 +28,10 @@ import { type PersonaHandoffClaim, } from "./acpPersonaHandoff"; import { useRuntimeConfigStore } from "@/shared/runtime-config/runtimeConfigStore"; -import { resolveManagedGooseProviderSelection } from "@/shared/runtime-config/modelProviderPolicy"; +import { + resolveManagedGooseProviderSelection, + resolveValidatedManagedGooseProviderSelection, +} from "@/shared/runtime-config/modelProviderPolicy"; import { getStyleGuidelinesPrompt } from "@/shared/preferences/styleGuidelinesPreference"; import { getBerdctlPreamble } from "@/features/berdctl/appPreamble"; import { INTERACTION_NORMS_PREAMBLE } from "@/shared/api/interactionNorms"; @@ -78,6 +81,21 @@ export interface AcpSessionConfigApplyOptions { modelId?: string | null; /** UI selection intent that owns any response snapshots. */ requestId?: string; + /** The coordinator already resolved this complete selection from inventory. */ + selectionAlreadyResolved?: boolean; +} + +export type AcpSessionConfigurationIntent = + sessionRegistry.SessionMutationSupersession; + +/** + * Reserve a session's configuration ordering before asynchronous resolution. + * The owner must pass it to acpPrepareSession and clear it when finished. + */ +export function reserveAcpSessionConfiguration( + sessionId: string, +): AcpSessionConfigurationIntent { + return sessionRegistry.supersedeSessionMutation(sessionId); } export interface AcpCreateSessionResult { @@ -169,11 +187,25 @@ async function acpSendMessageNow( const sid = sessionId.slice(0, 8); const tStart = performance.now(); - const resolvedProvider = resolveGooseSessionSelection(providerId).providerId; - if (resolvedProvider !== providerId) { - throw new Error( - `Session provider ${providerId} is outside the managed Goose provider policy. Re-prepare the session before prompting.`, - ); + if ( + providerId === "goose" || + CURATED_PROVIDER_CATALOG_BY_ID.get(providerId)?.category !== "agent" + ) { + const runtimeConfigState = useRuntimeConfigStore.getState(); + if (runtimeConfigState.result.status === "unavailable") { + throw new Error( + `Goose provider policy is unavailable: ${runtimeConfigState.result.message}`, + ); + } + const resolvedProvider = resolveManagedGooseProviderSelection( + runtimeConfigState.config, + { providerId }, + )?.providerId; + if (resolvedProvider && resolvedProvider !== providerId) { + throw new Error( + `Session provider ${providerId} is outside the managed Goose provider policy. Re-prepare the session before prompting.`, + ); + } } // Goose owns prompt assembly and accepts a real system prompt via its ACP @@ -324,10 +356,10 @@ export async function acpSteerMessage( ); } -function resolveGooseSessionSelection( +async function resolveGooseSessionSelection( providerId: string, modelId?: string | null, -): { providerId: string; modelId?: string } { +): Promise<{ providerId: string; modelId?: string }> { if (modelId === "goose") { throw new Error(`Invalid model id: ${modelId}`); } @@ -356,7 +388,7 @@ function resolveGooseSessionSelection( providerId, ...(concreteModelId ? { modelId: concreteModelId } : {}), }; - const managedSelection = resolveManagedGooseProviderSelection( + const managedSelection = await resolveValidatedManagedGooseProviderSelection( runtimeConfigState.config, requestedSelection, ); @@ -368,9 +400,36 @@ function resolveGooseSessionSelection( ); } - // A concrete provider is renderer-owned. Policy may validate it, but must - // not replace its provider or inject a different provider's default model. - return requestedSelection; + // A concrete provider stays renderer-owned; validation may remove or replace + // only its model with a proven result. + return managedSelection; +} + +/** Apply a caller-resolved session selection without performing another inventory read. */ +async function applyResolvedSessionSelection( + sessionId: string, + selection: { providerId: string; modelId?: string }, + workingDir: string, + options: AcpSessionConfigApplyOptions, + supersession: AcpSessionConfigurationIntent, +): Promise { + const applyResolvedModel = Boolean(selection.modelId); + return applyResolvedModel && selection.modelId + ? sessionRegistry.configureSession( + sessionId, + selection.providerId, + workingDir, + selection.modelId, + options, + supersession, + ) + : sessionRegistry.prepareSession( + sessionId, + selection.providerId, + workingDir, + options, + supersession, + ); } /** Prepare or warm an ACP session ahead of the first prompt. */ @@ -379,34 +438,37 @@ export async function acpPrepareSession( providerId: string, workingDir: string, options: AcpSessionConfigApplyOptions = {}, + intent?: AcpSessionConfigurationIntent, ): Promise { const sid = sessionId.slice(0, 8); const t0 = performance.now(); perfLog( `[perf:prepare] ${sid} acpPrepareSession start (provider=${providerId})`, ); - const selection = resolveGooseSessionSelection(providerId, options.modelId); - const applyResolvedModel = - Boolean(options.modelId) || selection.providerId !== providerId; - const snapshots = - applyResolvedModel && selection.modelId - ? await sessionRegistry.configureSession( - sessionId, - selection.providerId, - workingDir, - selection.modelId, - options, - ) - : await sessionRegistry.prepareSession( - sessionId, - selection.providerId, - workingDir, - options, - ); - perfLog( - `[perf:prepare] ${sid} acpPrepareSession done in ${(performance.now() - t0).toFixed(1)}ms`, - ); - return snapshots; + const supersession = intent ?? reserveAcpSessionConfiguration(sessionId); + const ownsSupersession = intent === undefined; + try { + const resolvedModelId = normalizeConcreteModelId(options.modelId); + const selection = options.selectionAlreadyResolved + ? { + providerId, + ...(resolvedModelId ? { modelId: resolvedModelId } : {}), + } + : await resolveGooseSessionSelection(providerId, options.modelId); + const snapshots = await applyResolvedSessionSelection( + sessionId, + selection, + workingDir, + options, + supersession, + ); + perfLog( + `[perf:prepare] ${sid} acpPrepareSession done in ${(performance.now() - t0).toFixed(1)}ms`, + ); + return snapshots; + } finally { + if (ownsSupersession) supersession.clear(); + } } export async function acpCreateSession( @@ -414,7 +476,10 @@ export async function acpCreateSession( workingDir: string, options: AcpCreateSessionOptions = {}, ): Promise { - const selection = resolveGooseSessionSelection(providerId, options.modelId); + const selection = await resolveGooseSessionSelection( + providerId, + options.modelId, + ); providerId = selection.providerId; options = { ...options, modelId: selection.modelId }; // Only the "goose" sentinel should rely on backend defaults. Concrete @@ -563,29 +628,37 @@ export async function acpLoadSession( sessionId: shortLogId(sessionId), }); perfLog(`[perf:load] ${sid} acpLoadSession → client.loadSession`); - const { response, isCurrent, executionSelection } = + const { response, isCurrent, deferredCurrent, executionSelection } = await sessionRegistry.loadSession(sessionId, effectiveWorkingDir); + const publish = () => { + const snapshots = readSessionConfigOptionsSnapshots(response); + logReasoningEffortInfo("acpLoadSession response", { + sessionId: shortLogId(sessionId), + hasReasoningEffortSnapshot: Boolean(snapshots.reasoningEffort), + ...reasoningEffortConfigLogFields( + "reasoningEffort", + snapshots.reasoningEffort, + ), + }); + applySessionConfigOptionsSnapshot(sessionId, response, { + origin: "response", + }); + perfLog( + `[perf:load] ${sid} client.loadSession resolved in ${(performance.now() - t0).toFixed(1)}ms`, + ); + }; if (!isCurrent) { + if (deferredCurrent) { + void deferredCurrent.then((becameCurrent) => { + if (becameCurrent) publish(); + }); + } perfLog( - `[perf:load] ${sid} dropped superseded load snapshot in ${(performance.now() - t0).toFixed(1)}ms`, + `[perf:load] ${sid} deferred or dropped superseded load snapshot in ${(performance.now() - t0).toFixed(1)}ms`, ); return undefined; } - const snapshots = readSessionConfigOptionsSnapshots(response); - logReasoningEffortInfo("acpLoadSession response", { - sessionId: shortLogId(sessionId), - hasReasoningEffortSnapshot: Boolean(snapshots.reasoningEffort), - ...reasoningEffortConfigLogFields( - "reasoningEffort", - snapshots.reasoningEffort, - ), - }); - applySessionConfigOptionsSnapshot(sessionId, response, { - origin: "response", - }); - perfLog( - `[perf:load] ${sid} client.loadSession resolved in ${(performance.now() - t0).toFixed(1)}ms`, - ); + publish(); return executionSelection; } diff --git a/src/shared/api/acpConnection.test.ts b/src/shared/api/acpConnection.test.ts new file mode 100644 index 000000000..631dd2974 --- /dev/null +++ b/src/shared/api/acpConnection.test.ts @@ -0,0 +1,125 @@ +import { beforeEach, describe, expect, it, vi } from "vitest"; + +const mocks = vi.hoisted(() => { + const urlRequests: Array<{ resolve: (url: string) => void }> = []; + const initializations: Array<{ + resolve: () => void; + reject: (error: Error) => void; + }> = []; + const streams: Array<{ writable: { abort: ReturnType } }> = []; + class MockGooseClient { + closed = new Promise(() => {}); + async initialize(): Promise { + await new Promise((resolve, reject) => + initializations.push({ resolve, reject }), + ); + } + } + return { urlRequests, initializations, streams, MockGooseClient }; +}); +vi.mock("@tauri-apps/api/core", () => ({ + invoke: vi.fn( + () => new Promise((resolve) => mocks.urlRequests.push({ resolve })), + ), +})); +vi.mock("@aaif/goose-sdk", () => ({ + DEFAULT_GOOSE_MCP_HOST_CAPABILITIES: {}, + GooseClient: mocks.MockGooseClient, +})); +vi.mock("./createWebSocketStream", () => ({ + createWebSocketStream: (url: string) => { + const stream = { + url, + writable: { abort: vi.fn().mockResolvedValue(undefined) }, + }; + mocks.streams.push(stream); + return stream; + }, +})); +describe("ACP connection lifecycle", () => { + beforeEach(() => { + mocks.urlRequests.length = 0; + mocks.initializations.length = 0; + mocks.streams.length = 0; + vi.resetModules(); + }); + it("rejects a retired URL lookup before opening its transport", async () => { + const connection = await import("./acpConnection"); + const stale = connection.getClient(); + await vi.waitFor(() => expect(mocks.urlRequests).toHaveLength(1)); + + await connection.invalidateClientConnection(); + const retry = connection.getClient(); + await vi.waitFor(() => expect(mocks.urlRequests).toHaveLength(2)); + mocks.urlRequests[0]?.resolve("ws://stale"); + await expect(stale).rejects.toThrow("initialization was superseded"); + expect(mocks.streams).toHaveLength(0); + + mocks.urlRequests[1]?.resolve("ws://current"); + await vi.waitFor(() => expect(mocks.initializations).toHaveLength(1)); + mocks.initializations[0]?.resolve(); + await expect(retry).resolves.toBeTruthy(); + expect(mocks.streams).toHaveLength(1); + expect(mocks.streams[0]).toMatchObject({ url: "ws://current" }); + }); + + it("retires and aborts an initializing transport before retry", async () => { + const connection = await import("./acpConnection"); + const stale = connection.getClient(); + await vi.waitFor(() => expect(mocks.urlRequests).toHaveLength(1)); + mocks.urlRequests[0]?.resolve("ws://goose"); + await vi.waitFor(() => expect(mocks.initializations).toHaveLength(1)); + await connection.invalidateClientConnection(); + expect(mocks.streams[0]?.writable.abort).toHaveBeenCalledOnce(); + const retry = connection.getClient(); + await vi.waitFor(() => expect(mocks.urlRequests).toHaveLength(2)); + mocks.urlRequests[1]?.resolve("ws://goose-retry"); + await vi.waitFor(() => expect(mocks.initializations).toHaveLength(2)); + mocks.initializations[1]?.resolve(); + const client = await retry; + mocks.initializations[0]?.resolve(); + await expect(stale).rejects.toThrow( + "ACP connection initialization was superseded", + ); + expect(connection.getClientSync()).toBe(client); + expect(mocks.streams[1]?.writable.abort).not.toHaveBeenCalled(); + }); + + it("rejects every waiter for a retired initialization", async () => { + const connection = await import("./acpConnection"); + const first = connection.getClient(); + const second = connection.getClient(); + await vi.waitFor(() => expect(mocks.urlRequests).toHaveLength(1)); + mocks.urlRequests[0]?.resolve("ws://goose"); + await vi.waitFor(() => expect(mocks.initializations).toHaveLength(1)); + + await connection.invalidateClientConnection(); + mocks.initializations[0]?.resolve(); + + await expect(first).rejects.toThrow("initialization was superseded"); + await expect(second).rejects.toThrow("initialization was superseded"); + expect(mocks.streams[0]?.writable.abort).toHaveBeenCalledOnce(); + }); + + it("aborts a transport when initialization rejects and preserves the error", async () => { + const connection = await import("./acpConnection"); + const failed = connection.getClient(); + await vi.waitFor(() => expect(mocks.urlRequests).toHaveLength(1)); + mocks.urlRequests[0]?.resolve("ws://goose"); + await vi.waitFor(() => expect(mocks.initializations).toHaveLength(1)); + + mocks.initializations[0]?.reject(new Error("handshake failed")); + + await expect(failed).rejects.toThrow("handshake failed"); + expect(mocks.streams[0]?.writable.abort).toHaveBeenCalledOnce(); + expect(connection.getClientSync()).toBeNull(); + + const retry = connection.getClient(); + await vi.waitFor(() => expect(mocks.urlRequests).toHaveLength(2)); + mocks.urlRequests[1]?.resolve("ws://goose-retry"); + await vi.waitFor(() => expect(mocks.initializations).toHaveLength(2)); + mocks.initializations[1]?.resolve(); + await expect(retry).resolves.toBeTruthy(); + expect(mocks.streams[1]?.writable.abort).not.toHaveBeenCalled(); + }); +}); diff --git a/src/shared/api/acpConnection.ts b/src/shared/api/acpConnection.ts index 0f69884f5..b37071809 100644 --- a/src/shared/api/acpConnection.ts +++ b/src/shared/api/acpConnection.ts @@ -61,6 +61,21 @@ export function setPermissionHandler(handler: PermissionRequestHandler): void { let clientPromise: Promise | null = null; let resolvedClient: GooseClient | null = null; let activeStream: ReturnType | null = null; +let connectionGeneration = 0; + +interface ConnectionAttempt { + generation: number; + stream: ReturnType | null; + streamAborted: boolean; +} + +let currentAttempt: ConnectionAttempt | null = null; + +async function abortAttemptStream(attempt: ConnectionAttempt): Promise { + if (!attempt.stream || attempt.streamAborted) return; + attempt.streamAborted = true; + await attempt.stream.writable.abort(); +} function createClientCallbacks(): () => Client { return () => ({ @@ -95,14 +110,16 @@ function createClientCallbacks(): () => Client { function monitorConnection( client: GooseClient, stream: ReturnType, + attempt: ConnectionAttempt, ): void { const clearCurrentConnection = () => { - if (activeStream !== stream) { + if (currentAttempt !== attempt || activeStream !== stream) { return; } resolvedClient = null; clientPromise = null; activeStream = null; + currentAttempt = null; }; client.closed .then(() => { @@ -125,16 +142,28 @@ function monitorConnection( * safer than allowing later mutations to race work still running remotely. */ export async function invalidateClientConnection(): Promise { - const stream = activeStream; + connectionGeneration += 1; + const attempt = currentAttempt; + currentAttempt = null; + const stream = attempt?.stream ?? activeStream; activeStream = null; resolvedClient = null; clientPromise = null; - if (stream) { + if (attempt) { + await abortAttemptStream(attempt); + } else if (stream) { await stream.writable.abort(); } } -async function initializeConnection(): Promise { +interface InitializedConnection { + client: GooseClient; + stream: ReturnType; +} + +async function initializeConnection( + attempt: ConnectionAttempt, +): Promise { // Dev-only: inject a real failure into startup so the WARP probe runs // for real against kgoose. `VITE_DEV_STARTUP_ERROR=warp just dev` lets // us experience the diagnostic UI with whatever real WARP state the @@ -164,10 +193,18 @@ async function initializeConnection(): Promise { perfLog( `[perf:conn] get_goose_serve_url in ${(performance.now() - tStart).toFixed(1)}ms`, ); + if ( + currentAttempt !== attempt || + attempt.generation !== connectionGeneration + ) { + throw new Error( + "ACP connection initialization was superseded; retry the operation.", + ); + } const tStream = performance.now(); const stream = createWebSocketStream(wsUrl); - activeStream = stream; + attempt.stream = stream; const client = new GooseClient(createClientCallbacks(), stream); perfLog( @@ -194,31 +231,54 @@ async function initializeConnection(): Promise { `[perf:conn] client.initialize in ${(performance.now() - tInit).toFixed(1)}ms (total ${(performance.now() - tStart).toFixed(1)}ms)`, ); - monitorConnection(client, stream); - - return client; + return { client, stream }; } export async function getClient(): Promise { - if (resolvedClient) { - return resolvedClient; - } - + if (resolvedClient) return resolvedClient; if (!clientPromise) { perfLog("[perf:conn] getClient() → initializing new ACP connection"); - clientPromise = initializeConnection() - .then((client) => { - resolvedClient = client; - return client; + const attempt: ConnectionAttempt = { + generation: connectionGeneration, + stream: null, + streamAborted: false, + }; + currentAttempt = attempt; + const initialization = initializeConnection(attempt) + .then(async ({ client, stream }) => { + if ( + currentAttempt === attempt && + attempt.generation === connectionGeneration + ) { + activeStream = stream; + resolvedClient = client; + monitorConnection(client, stream, attempt); + return client; + } + await abortAttemptStream(attempt); + throw new Error( + "ACP connection initialization was superseded; retry the operation.", + ); }) - .catch((error) => { - clientPromise = null; + .catch(async (error) => { + // initializeConnection may fail after opening the transport. Retire it + // before dropping the attempt while preserving the original failure. + await abortAttemptStream(attempt).catch((abortError) => { + console.warn( + "[acp] Failed to abort rejected connection attempt.", + abortError, + ); + }); + if (currentAttempt === attempt) { + currentAttempt = null; + clientPromise = null; + } throw error; }); + clientPromise = initialization; } else { perfLog("[perf:conn] getClient() awaiting in-flight initializeConnection"); } - return clientPromise; } diff --git a/src/shared/api/acpSessionRegistry.ts b/src/shared/api/acpSessionRegistry.ts index 83c6d0c33..b9b2db819 100644 --- a/src/shared/api/acpSessionRegistry.ts +++ b/src/shared/api/acpSessionRegistry.ts @@ -11,6 +11,7 @@ import { shortLogId, } from "@/shared/lib/reasoningEffortDiagnostics"; import { normalizeConcreteModelId } from "@/shared/lib/modelIdentity"; +import { isModelSelectionAllowedByCachedInventory } from "@/features/providers/stores/providerModelCacheStore"; export interface AcpSessionExecutionSelection { providerId: string; @@ -30,13 +31,44 @@ interface SessionConfigMutationOptions { const SESSION_MUTATION_TIMEOUT_MS = 60_000; +interface SessionMutationQueue { + latestSequence: number; + tail: Promise; + /** Configuration intent awaiting async preflight before it can enqueue. */ + pendingSupersession?: { + sequence: number; + previousSequence: number; + settled: Promise; + resolve: () => void; + }; +} + +/** Opaque ownership of a configuration intent that is awaiting preflight. */ +export interface SessionMutationSupersession { + readonly sequence: number; + clear(): void; +} + const prepared = new Map(); -const mutationQueues = new Map< - string, - { latestSequence: number; tail: Promise } ->(); +const mutationQueues = new Map(); let nextMutationSequence = 1; +function scheduleQueueCleanup( + sessionId: string, + queue: SessionMutationQueue, +): void { + const tail = queue.tail; + void tail.then(() => { + if ( + mutationQueues.get(sessionId) === queue && + queue.tail === tail && + queue.pendingSupersession === undefined + ) { + mutationQueues.delete(sessionId); + } + }); +} + function clonePreparedSession( entry: PreparedSession | undefined, ): PreparedSession | undefined { @@ -101,7 +133,11 @@ async function runBoundedSessionMutation( function serializeSessionMutation( sessionId: string, - mutation: (isLatest: () => boolean) => Promise, + mutation: ( + isLatest: () => boolean, + sequence: number, + queue: SessionMutationQueue, + ) => Promise, bounded = true, ): Promise { let queue = mutationQueues.get(sessionId); @@ -112,7 +148,14 @@ function serializeSessionMutation( const sequence = nextMutationSequence++; queue.latestSequence = sequence; - const execute = () => mutation(() => queue?.latestSequence === sequence); + const execute = () => + mutation( + () => + queue?.pendingSupersession === undefined && + queue?.latestSequence === sequence, + sequence, + queue, + ); const result = queue.tail.then(() => bounded ? runBoundedSessionMutation(sessionId, execute()) : execute(), ); @@ -121,23 +164,80 @@ function serializeSessionMutation( () => undefined, ); queue.tail = tail; - void tail.then(() => { - if (mutationQueues.get(sessionId)?.tail === tail) { - mutationQueues.delete(sessionId); - } - }); + scheduleQueueCleanup(sessionId, queue); return result; } +function consumeSessionSupersession( + sessionId: string, + supersession: SessionMutationSupersession | undefined, +): boolean { + if (!supersession) return true; + const queue = mutationQueues.get(sessionId); + if (queue?.pendingSupersession?.sequence !== supersession.sequence) { + return false; + } + const pending = queue.pendingSupersession; + queue.pendingSupersession = undefined; + pending.resolve(); + scheduleQueueCleanup(sessionId, queue); + return true; +} + +export function supersedeSessionMutation( + sessionId: string, +): SessionMutationSupersession { + let queue = mutationQueues.get(sessionId); + if (!queue) { + queue = { latestSequence: 0, tail: Promise.resolve() }; + mutationQueues.set(sessionId, queue); + } + // Retain preflight intent before its authoritative I/O completes so a load + // cannot publish a snapshot that predates the requested configuration. + const sequence = nextMutationSequence++; + const previousSequence = + queue.pendingSupersession?.previousSequence ?? queue.latestSequence; + queue.pendingSupersession?.resolve(); + let resolveSettled!: () => void; + const settled = new Promise((resolve) => { + resolveSettled = resolve; + }); + queue.latestSequence = sequence; + queue.pendingSupersession = { + sequence, + previousSequence, + settled, + resolve: resolveSettled, + }; + scheduleQueueCleanup(sessionId, queue); + + return { + sequence, + clear() { + if (queue?.pendingSupersession?.sequence !== sequence) return; + const pending = queue.pendingSupersession; + queue.pendingSupersession = undefined; + if (queue.latestSequence === sequence) { + queue.latestSequence = pending.previousSequence; + } + pending.resolve(); + scheduleQueueCleanup(sessionId, queue); + }, + }; +} + export async function prepareSession( sessionId: string, providerId: string, workingDir: string, options: SessionConfigMutationOptions = {}, + supersession?: SessionMutationSupersession, ): Promise { - return serializeSessionMutation(sessionId, () => + if (!consumeSessionSupersession(sessionId, supersession)) return; + const snapshots = await serializeSessionMutation(sessionId, () => prepareSessionNow(sessionId, providerId, workingDir, options), ); + return snapshots; } async function prepareSessionNow( @@ -192,6 +292,15 @@ async function prepareSessionNow( perfLog( `[perf:prepare] ${sid} reuse existing session (updates=${changed}) in ${(performance.now() - tReuse).toFixed(1)}ms`, ); + if (!snapshots && existing.executionSelection?.modelId) { + return { + model: { + modelId: existing.executionSelection.modelId, + modelName: existing.executionSelection.modelId, + }, + reasoningEffort: null, + }; + } return snapshots; } @@ -314,12 +423,14 @@ export async function configureSession( workingDir: string, modelId?: string, options: SessionConfigMutationOptions = {}, + supersession?: SessionMutationSupersession, ): Promise { const concreteModelId = normalizeConcreteModelId(modelId); if (modelId && !concreteModelId) { throw new Error(`Invalid model id: ${modelId}`); } - return serializeSessionMutation(sessionId, async () => { + if (!consumeSessionSupersession(sessionId, supersession)) return; + const snapshots = await serializeSessionMutation(sessionId, async () => { let snapshots = await prepareSessionNow( sessionId, providerId, @@ -327,12 +438,19 @@ export async function configureSession( concreteModelId ? {} : options, ); if (concreteModelId) { - snapshots = - (await applySessionModelNow(sessionId, concreteModelId, options)) ?? - snapshots; + const modelSnapshots = await applySessionModelNow( + sessionId, + concreteModelId, + options, + ); + snapshots = modelSnapshots ?? { + model: { modelId: concreteModelId, modelName: concreteModelId }, + reasoningEffort: null, + }; } return snapshots; }); + return snapshots; } export function applySessionConfigOption( @@ -365,14 +483,29 @@ export function requireSessionInvocationSelection( "Session requires a configured provider and model before prompting. Re-prepare the session after completing provider setup.", ); } + if ( + !isModelSelectionAllowedByCachedInventory( + selection.providerId, + selection.modelId, + ) + ) { + throw new Error( + `Session model ${selection.modelId} is no longer supported by provider ${selection.providerId}. Re-prepare the session before prompting.`, + ); + } return { ...selection, modelId: selection.modelId }; } /** Run prompt setup and transport without allowing session config to interleave. */ -export function runPreparedSessionPrompt( +export async function runPreparedSessionPrompt( sessionId: string, prompt: (providerId: string) => Promise, ): Promise { + let pending = mutationQueues.get(sessionId)?.pendingSupersession; + while (pending) { + await pending.settled; + pending = mutationQueues.get(sessionId)?.pendingSupersession; + } return serializeSessionMutation( sessionId, () => prompt(requireSessionInvocationSelection(sessionId).providerId), @@ -386,13 +519,26 @@ export async function loadSession( ): Promise<{ response: Awaited>; isCurrent: boolean; + deferredCurrent?: Promise; executionSelection?: AcpSessionExecutionSelection; }> { return serializeSessionMutation( sessionId, - async (isLatest) => { + async (isLatest, _sequence, queue) => { const response = await acpApi.loadSession(sessionId, workingDir); + const pendingAtResponse = queue.pendingSupersession; const isCurrentResult = isLatest(); + const deferredCurrent = pendingAtResponse + ? (async () => { + let pending: SessionMutationQueue["pendingSupersession"] = + pendingAtResponse; + while (pending) { + await pending.settled; + pending = queue.pendingSupersession; + } + return isLatest(); + })() + : undefined; const executionSnapshot = readSessionExecutionConfigSnapshot(response); prepared.set(sessionId, { workingDir, @@ -401,6 +547,7 @@ export async function loadSession( return { response, isCurrent: isCurrentResult, + ...(deferredCurrent ? { deferredCurrent } : {}), executionSelection: executionSnapshot ?? undefined, }; }, diff --git a/src/shared/runtime-config/modelProviderPolicy.test.ts b/src/shared/runtime-config/modelProviderPolicy.test.ts index 4195073ab..6e937e21c 100644 --- a/src/shared/runtime-config/modelProviderPolicy.test.ts +++ b/src/shared/runtime-config/modelProviderPolicy.test.ts @@ -1,10 +1,28 @@ -import { describe, expect, it } from "vitest"; +import { beforeEach, describe, expect, it, vi } from "vitest"; import type { RuntimeConfig } from "./schema"; +import { notifyProviderModelInventoryInvalidated } from "./providerModelInventoryInvalidation"; import { managedGooseSelectionChanged, resolveManagedGooseProviderSelection, + resolveValidatedManagedGooseProviderSelection, } from "./modelProviderPolicy"; +const mockGetClient = vi.hoisted(() => vi.fn()); +const mockInvalidateClientConnection = vi.hoisted(() => vi.fn()); +const mockSupportedModelsList = vi.hoisted(() => vi.fn()); +vi.mock("@/shared/api/acpConnection", () => ({ + getClient: () => mockGetClient(), + invalidateClientConnection: () => mockInvalidateClientConnection(), +})); + +function deferred() { + let resolve!: (value: T) => void; + const promise = new Promise((resolvePromise) => { + resolve = resolvePromise; + }); + return { promise, resolve }; +} + const managedConfig: RuntimeConfig = { schemaVersion: 1, goose: { @@ -29,6 +47,18 @@ const managedConfig: RuntimeConfig = { }; describe("resolveManagedGooseProviderSelection", () => { + beforeEach(() => { + vi.useRealTimers(); + mockGetClient.mockReset(); + mockInvalidateClientConnection.mockReset(); + mockInvalidateClientConnection.mockResolvedValue(undefined); + mockGetClient.mockResolvedValue({ + goose: { + GooseUnstableProvidersSupportedModelsList: mockSupportedModelsList, + }, + }); + mockSupportedModelsList.mockReset(); + }); it("returns unrestricted for an empty provider list", () => { expect( resolveManagedGooseProviderSelection( @@ -58,7 +88,7 @@ describe("resolveManagedGooseProviderSelection", () => { ).toEqual({ providerId: "databricks_v2", modelId: "shared-model" }); }); - it("repairs the legacy Goose model sentinel without live inventory", () => { + it("keeps a legacy model sentinel model-free without live proof", () => { expect( resolveManagedGooseProviderSelection(managedConfig, { providerId: "databricks", @@ -66,7 +96,7 @@ describe("resolveManagedGooseProviderSelection", () => { }), ).toEqual({ providerId: "databricks_v2", - modelId: "goose-gpt-5-5", + modelId: undefined, }); }); @@ -151,14 +181,183 @@ describe("resolveManagedGooseProviderSelection", () => { }); }); - it("uses the configured default only when no model is selected", () => { + it("keeps a migrated provider model-free without live proof", () => { expect( resolveManagedGooseProviderSelection(managedConfig, { providerId: "databricks", }), ).toEqual({ + providerId: "databricks_v2", + modelId: undefined, + }); + }); + + it("validates a model retained across a provider migration", async () => { + mockSupportedModelsList.mockResolvedValue({ models: ["shared-model"] }); + + await expect( + resolveValidatedManagedGooseProviderSelection(managedConfig, { + providerId: "disallowed", + modelId: "shared-model", + }), + ).resolves.toEqual({ + providerId: "databricks_v2", + modelId: "shared-model", + }); + }); + + it("replaces an unsupported migrated model only with a proven default", async () => { + mockSupportedModelsList.mockResolvedValue({ models: ["goose-gpt-5-5"] }); + + await expect( + resolveValidatedManagedGooseProviderSelection(managedConfig, { + providerId: "disallowed", + modelId: "other-model", + }), + ).resolves.toEqual({ providerId: "databricks_v2", modelId: "goose-gpt-5-5", }); }); + + it("uses a deterministic proven inventory model when a migration has no default model", async () => { + mockSupportedModelsList.mockResolvedValue({ + models: ["z-model", "a-model"], + }); + const configWithoutDefault: RuntimeConfig = { + ...managedConfig, + goose: { ...managedConfig.goose, defaultModelId: undefined }, + }; + + await expect( + resolveValidatedManagedGooseProviderSelection(configWithoutDefault, { + providerId: "disallowed", + }), + ).resolves.toEqual({ + providerId: "databricks_v2", + modelId: "a-model", + }); + }); + + it("rejects a migration with no selected default when target inventory is empty", async () => { + mockSupportedModelsList.mockResolvedValue({ models: [] }); + const configWithoutDefault: RuntimeConfig = { + ...managedConfig, + goose: { ...managedConfig.goose, defaultModelId: undefined }, + }; + + await expect( + resolveValidatedManagedGooseProviderSelection(configWithoutDefault, { + providerId: "disallowed", + }), + ).rejects.toThrow( + "No supported model is available for migrated provider databricks_v2", + ); + }); + + it("times out a never-settling inventory proof without accepting its late result", async () => { + vi.useFakeTimers(); + let resolveInventory!: (value: { models: string[] }) => void; + mockSupportedModelsList.mockReturnValue( + new Promise<{ models: string[] }>((resolve) => { + resolveInventory = resolve; + }), + ); + + const migration = resolveValidatedManagedGooseProviderSelection( + managedConfig, + { + providerId: "disallowed", + }, + ); + const rejectedMigration = expect(migration).rejects.toThrow( + "Cannot verify models for migrated provider databricks_v2", + ); + await vi.advanceTimersByTimeAsync(60_000); + await rejectedMigration; + expect(mockInvalidateClientConnection).not.toHaveBeenCalled(); + resolveInventory({ models: ["goose-gpt-5-5"] }); + await Promise.resolve(); + }); + + it("does not abort a concurrent prompt when inventory proof times out", async () => { + vi.useFakeTimers(); + const prompt = deferred(); + mockSupportedModelsList.mockReturnValue(new Promise(() => {})); + mockGetClient.mockResolvedValue({ + goose: { + GooseUnstableProvidersSupportedModelsList: mockSupportedModelsList, + prompt: () => prompt.promise, + }, + }); + + const client = await mockGetClient(); + const activePrompt = client.goose.prompt(); + const migration = resolveValidatedManagedGooseProviderSelection( + managedConfig, + { providerId: "disallowed" }, + ); + const rejectedMigration = expect(migration).rejects.toThrow( + "Cannot verify models for migrated provider databricks_v2", + ); + await vi.advanceTimersByTimeAsync(60_000); + await rejectedMigration; + + expect(mockInvalidateClientConnection).not.toHaveBeenCalled(); + prompt.resolve("complete"); + await expect(activePrompt).resolves.toBe("complete"); + }); + + it("times out stalled ACP client acquisition", async () => { + vi.useFakeTimers(); + mockGetClient.mockReturnValue(new Promise(() => {})); + + const migration = resolveValidatedManagedGooseProviderSelection( + managedConfig, + { providerId: "disallowed" }, + ); + const rejectedMigration = expect(migration).rejects.toThrow( + "Cannot verify models for migrated provider databricks_v2", + ); + await vi.advanceTimersByTimeAsync(60_000); + await rejectedMigration; + expect(mockInvalidateClientConnection).not.toHaveBeenCalled(); + expect(mockSupportedModelsList).not.toHaveBeenCalled(); + }); + + it("rejects an inventory proof invalidated while it is in flight", async () => { + let resolveInventory!: (value: { models: string[] }) => void; + mockSupportedModelsList.mockReturnValue( + new Promise<{ models: string[] }>((resolve) => { + resolveInventory = resolve; + }), + ); + + const migration = resolveValidatedManagedGooseProviderSelection( + managedConfig, + { + providerId: "disallowed", + }, + ); + await vi.waitFor(() => expect(mockSupportedModelsList).toHaveBeenCalled()); + notifyProviderModelInventoryInvalidated("databricks_v2"); + resolveInventory({ models: ["goose-gpt-5-5"] }); + + await expect(migration).rejects.toThrow( + "Cannot verify models for migrated provider databricks_v2", + ); + }); + + it("rejects a provider migration when support cannot be proved", async () => { + mockSupportedModelsList.mockRejectedValue(new Error("offline")); + + await expect( + resolveValidatedManagedGooseProviderSelection(managedConfig, { + providerId: "disallowed", + modelId: "other-model", + }), + ).rejects.toThrow( + "Cannot verify models for migrated provider databricks_v2", + ); + }); }); diff --git a/src/shared/runtime-config/modelProviderPolicy.ts b/src/shared/runtime-config/modelProviderPolicy.ts index 59e16dc35..c29ccb5e0 100644 --- a/src/shared/runtime-config/modelProviderPolicy.ts +++ b/src/shared/runtime-config/modelProviderPolicy.ts @@ -1,5 +1,7 @@ import type { RuntimeConfig, RuntimeGooseConfig } from "./schema"; import { normalizeConcreteModelId } from "@/shared/lib/modelIdentity"; +import { getClient } from "@/shared/api/acpConnection"; +import { providerModelInventoryGeneration } from "./providerModelInventoryInvalidation"; export interface GooseProviderSelection { providerId?: string | null; @@ -18,7 +20,7 @@ export interface ManagedGooseProviderResolutionContext { targetInventoryValidated?: boolean; } -const DATABRICKS_V2_PROVIDER_ID = "databricks_v2"; +const INVENTORY_PROOF_TIMEOUT_MS = 60_000; /** * Runtime model providers define provider policy and curated model metadata. @@ -48,8 +50,8 @@ function defaultManagedProviderId(goose: RuntimeGooseConfig): string { * - `null` means policy is unrestricted; the caller must preserve its values. * - Allowed providers and all of their upstream-discovered models stay selected. * - Disallowed/missing providers move to the runtime default provider. - * - Existing model selections survive provider migration. A missing model uses - * the configured default, whose inventory entry is recommendation metadata. + * - Existing concrete model selections survive while proof is unavailable. + * - A configured default is synthesized only from successful inventory proof. */ export function resolveManagedGooseProviderSelection( config: Pick, @@ -65,20 +67,102 @@ export function resolveManagedGooseProviderSelection( (provider) => provider.id === selection.providerId, )?.id; const providerId = configuredProviderId ?? defaultManagedProviderId(goose); - let modelId = - normalizeConcreteModelId(selection.modelId) ?? - normalizeConcreteModelId(goose.defaultModelId); + const providerWasMigrated = configuredProviderId === undefined; + const selectedModelId = normalizeConcreteModelId(selection.modelId); + const defaultModelId = normalizeConcreteModelId(goose.defaultModelId); + if (context.targetInventoryValidated === true) { + const provenModelIds = context.targetModelIds ?? new Set(); + const needsModelRepair = providerWasMigrated || selection.modelId != null; + const modelId = + (selectedModelId && provenModelIds.has(selectedModelId) + ? selectedModelId + : undefined) ?? + // A default is a synthesized fallback. Do not turn same-provider, + // provider-only intent into a concrete selection merely because live + // inventory happened to be available. It may repair an existing concrete + // selection (including a legacy sentinel) or a provider migration. + (needsModelRepair && defaultModelId && provenModelIds.has(defaultModelId) + ? defaultModelId + : undefined); + return { providerId, modelId }; + } + + return { providerId, modelId: selectedModelId }; +} - if ( - providerId === DATABRICKS_V2_PROVIDER_ID && - modelId && - context.targetInventoryValidated === true && - !context.targetModelIds?.has(modelId) - ) { - modelId = goose.defaultModelId; +/** + * Read a provider's live model inventory as authoritative evidence for managed + * configuration decisions. Both ACP acquisition and the inventory RPC share + * one deadline. A timed-out proof is abandoned without invalidating the shared + * ACP transport, so unrelated active prompts remain intact. Results from an + * invalidated inventory generation are never accepted. + */ +export async function readBoundedProvenModelInventory( + providerId: string, +): Promise> { + const generationAtStart = providerModelInventoryGeneration(providerId); + let timeoutId: ReturnType | undefined; + try { + const response = await Promise.race([ + getClient().then((client) => + client.goose.GooseUnstableProvidersSupportedModelsList({ providerId }), + ), + new Promise((_, reject) => { + timeoutId = setTimeout(() => { + reject( + new Error(`Timed out proving models for provider ${providerId}.`), + ); + }, INVENTORY_PROOF_TIMEOUT_MS); + }), + ]); + if (generationAtStart !== providerModelInventoryGeneration(providerId)) { + throw new Error( + `Model inventory changed while proving provider ${providerId}.`, + ); + } + return new Set(response.models as string[]); + } finally { + if (timeoutId !== undefined) clearTimeout(timeoutId); } +} + +export async function resolveValidatedManagedGooseProviderSelection( + config: Pick, + selection: GooseProviderSelection, +): Promise { + const resolved = resolveManagedGooseProviderSelection(config, selection); + if (!resolved) return null; - return { providerId, modelId: modelId ?? undefined }; + let supportedModelIds: ReadonlySet; + try { + supportedModelIds = await readBoundedProvenModelInventory( + resolved.providerId, + ); + } catch (error) { + if (resolved.providerId === selection.providerId) return resolved; + throw new Error( + `Cannot verify models for migrated provider ${resolved.providerId}; provider selection was not changed.`, + { cause: error }, + ); + } + + const proven = resolveManagedGooseProviderSelection(config, selection, { + targetModelIds: supportedModelIds, + targetInventoryValidated: true, + }); + if (resolved.providerId === selection.providerId) return proven; + if (proven?.modelId) return proven; + + const provenInventoryFallback = [...supportedModelIds].sort()[0]; + if (provenInventoryFallback) { + return { + providerId: resolved.providerId, + modelId: provenInventoryFallback, + }; + } + throw new Error( + `No supported model is available for migrated provider ${resolved.providerId}; provider selection was not changed.`, + ); } export function managedGooseSelectionChanged( diff --git a/src/shared/runtime-config/providerModelInventoryInvalidation.ts b/src/shared/runtime-config/providerModelInventoryInvalidation.ts new file mode 100644 index 000000000..0ced94916 --- /dev/null +++ b/src/shared/runtime-config/providerModelInventoryInvalidation.ts @@ -0,0 +1,28 @@ +type ProviderModelInventoryInvalidationListener = (providerId: string) => void; + +const invalidationListeners = + new Set(); +const inventoryGenerations = new Map(); + +export function providerModelInventoryGeneration(providerId: string): number { + return inventoryGenerations.get(providerId) ?? 0; +} + +export function notifyProviderModelInventoryInvalidated( + providerId: string, +): void { + inventoryGenerations.set( + providerId, + providerModelInventoryGeneration(providerId) + 1, + ); + for (const listener of invalidationListeners) { + listener(providerId); + } +} + +export function subscribeToProviderModelInventoryInvalidation( + listener: ProviderModelInventoryInvalidationListener, +): () => void { + invalidationListeners.add(listener); + return () => invalidationListeners.delete(listener); +} diff --git a/src/shared/ui/GlobalComposerPill.test.tsx b/src/shared/ui/GlobalComposerPill.test.tsx index 02e8d841b..7fe647e12 100644 --- a/src/shared/ui/GlobalComposerPill.test.tsx +++ b/src/shared/ui/GlobalComposerPill.test.tsx @@ -24,6 +24,7 @@ const mockNormalizeImageBase64 = vi.fn(); const mockSearchFilesForMentions = vi.fn(); const mockResizeImage = vi.fn(); const mockGetModelsForAgent = vi.fn(); +const mockGetProvenModelsForAgent = vi.fn(); const pendingMentionLoad = new Promise(() => {}); const mockRefreshAllModelProviders = vi.fn(); const mockRefreshAgentProviderStatus = vi.fn(); @@ -89,6 +90,8 @@ vi.mock("@/features/providers/hooks/useProviderModels", () => ({ configuredModelProviderIds: ["openai", "anthropic"], modelCacheRefreshProviderIds: ["openai", "anthropic"], getModelsForAgent: (agentId: string) => mockGetModelsForAgent(agentId), + getProvenModelsForAgent: (agentId: string) => + mockGetProvenModelsForAgent(agentId), isModelInventoryAuthoritative: () => mockProviderModelsState.inventoryAuthoritative, refreshAllModelProviders: (...args: unknown[]) => @@ -255,6 +258,10 @@ describe("GlobalComposerPill", () => { vi.mocked(listSkills).mockImplementation(() => pendingMentionLoad); mockGetModelsForAgent.mockReset(); mockGetModelsForAgent.mockReturnValue([]); + mockGetProvenModelsForAgent.mockReset(); + mockGetProvenModelsForAgent.mockImplementation((agentId: string) => + mockGetModelsForAgent(agentId), + ); mockRefreshAllModelProviders.mockReset(); mockRefreshAllModelProviders.mockResolvedValue(undefined); mockRefreshAgentProviderStatus.mockReset(); @@ -776,7 +783,7 @@ describe("GlobalComposerPill", () => { }); }); - it("keeps the Composer target when a persona has no plausible target", async () => { + it("blocks a persona that has saved model metadata but no plausible target", async () => { const user = userEvent.setup(); useAgentStore.setState({ personas: [ @@ -795,11 +802,11 @@ describe("GlobalComposerPill", () => { }); await user.type(screen.getByRole("textbox"), "Hello"); - await user.click(screen.getByRole("button", { name: /send message/i })); - expectSent(onSend, "Hello", { - personaId: "persona-1", - }); + expect( + screen.getByRole("button", { name: /send message/i }), + ).toBeDisabled(); + expect(onSend).not.toHaveBeenCalled(); }); it("applies a legacy persona target when inventory arrives", async () => { @@ -880,6 +887,119 @@ describe("GlobalComposerPill", () => { }); }); + await user.type(screen.getByRole("textbox"), "Hello"); + + // The authoritative inventory disproves this persona's saved model, so it + // has no runnable target until the user explicitly selects a supported one. + expect( + screen.getByRole("button", { name: /send message/i }), + ).toBeDisabled(); + expect(onSend).not.toHaveBeenCalled(); + }); + + it("fails closed when authoritative inventory invalidates the active session model", async () => { + const user = userEvent.setup(); + const activeTarget = { + harnessId: "goose", + modelProviderId: "databricks_v2", + modelId: "model-a", + modelName: "Model A", + }; + mockGetModelsForAgent.mockReturnValue([ + { + id: "model-a", + name: "Model A", + providerId: "databricks_v2", + }, + ]); + mockGetProvenModelsForAgent.mockReturnValue([ + { + id: "model-a", + name: "Model A", + providerId: "databricks_v2", + }, + ]); + const onSend = vi.fn(); + const { rerender } = render( + , + ); + + await user.type(screen.getByRole("textbox"), "Hello"); + expect(screen.getByRole("button", { name: /send message/i })).toBeEnabled(); + expect(screen.getByText("Model A")).toBeInTheDocument(); + + // A successful refresh is now authoritative and excludes model A. + mockGetModelsForAgent.mockReturnValue([ + { + id: "model-b", + name: "Model B", + providerId: "databricks_v2", + }, + ]); + mockGetProvenModelsForAgent.mockReturnValue([ + { + id: "model-b", + name: "Model B", + providerId: "databricks_v2", + }, + ]); + rerender( + , + ); + + expect(screen.queryByText("Model A")).not.toBeInTheDocument(); + const send = screen.getByRole("button", { name: /send message/i }); + expect(send).toBeDisabled(); + await user.click(send); + expect(onSend).not.toHaveBeenCalled(); + }); + + it("does not synthesize an advisory model while inventory proof is unavailable", async () => { + const user = userEvent.setup(); + mockProviderModelsState.inventoryAuthoritative = false; + useAgentStore.setState({ selectedProvider: "databricks_v2" }); + mockGetModelsForAgent.mockReturnValue([ + { id: "advisory", name: "Advisory", recommended: true }, + ]); + mockGetProvenModelsForAgent.mockReturnValue([]); + const onSend = renderGlobalComposer(vi.fn(), { + currentExecutionTarget: { + harnessId: "goose", + modelProviderId: "databricks_v2", + }, + }); + + await user.type(screen.getByRole("textbox"), "Hello"); + await user.click(screen.getByRole("button", { name: /send message/i })); + + expectSent(onSend, "Hello", { + executionTarget: { + harnessId: "goose", + modelProviderId: "databricks_v2", + }, + }); + }); + + it("does not synthesize an unqualified advisory model from authoritative empty inventory", async () => { + const user = userEvent.setup(); + useAgentStore.setState({ selectedProvider: "databricks_v2" }); + mockGetModelsForAgent.mockReturnValue([ + { id: "advisory", name: "Advisory", recommended: true }, + ]); + mockGetProvenModelsForAgent.mockReturnValue([]); + const onSend = renderGlobalComposer(vi.fn(), { + currentExecutionTarget: { + harnessId: "goose", + modelProviderId: "databricks_v2", + }, + }); + await user.type(screen.getByRole("textbox"), "Hello"); await user.click(screen.getByRole("button", { name: /send message/i })); @@ -887,10 +1007,7 @@ describe("GlobalComposerPill", () => { executionTarget: { harnessId: "goose", modelProviderId: "databricks_v2", - modelId: "goose-claude-opus-4-8", - modelName: "goose-claude-opus-4-8", }, - personaId: "persona-1", }); }); @@ -1156,17 +1273,15 @@ describe("GlobalComposerPill", () => { expect(screen.getByText("UX Critic")).toBeInTheDocument(); await user.type(screen.getByRole("textbox"), "Hello"); - await user.click(screen.getByRole("button", { name: /send message/i })); - expectSent(onSend, "Hello", { - executionTarget: { - harnessId: "goose", - modelProviderId: "databricks_v2", - modelId: "goose-default", - modelName: "goose-default", - }, - personaId: "persona-2", - }); + // `goose-default` is absent from the authoritative mock inventory, so the + // newly selected persona cannot dispatch through an implicit fallback. + expect(screen.queryByText("GPT-5.5")).not.toBeInTheDocument(); + expect(screen.getByText("Goose")).toBeInTheDocument(); + expect( + screen.getByRole("button", { name: /send message/i }), + ).toBeDisabled(); + expect(onSend).not.toHaveBeenCalled(); }); it("keeps the selected provider/model after clearing the suggested persona", async () => { @@ -1194,13 +1309,10 @@ describe("GlobalComposerPill", () => { await user.type(screen.getByRole("textbox"), "Hello"); await user.click(screen.getByRole("button", { name: /send message/i })); + // The unsupported persona target never becomes a live selection, so + // clearing it retains the composer fallback without stale metadata. expectSent(onSend, "Hello", { - executionTarget: { - harnessId: "claude-acp", - modelProviderId: "claude-acp", - modelId: "claude-sonnet-4", - modelName: "claude-sonnet-4", - }, + executionTarget: { harnessId: "goose" }, }); }); @@ -1577,6 +1689,13 @@ describe("GlobalComposerPill", () => { it("expands with the controlled Home model", async () => { const user = userEvent.setup(); const onExpand = vi.fn().mockResolvedValue(true); + mockGetModelsForAgent.mockReturnValue([ + { + id: "goose-claude-fable", + name: "Claude Fable", + providerId: "anthropic", + }, + ]); renderGlobalComposer(vi.fn(), { onExpand, currentExecutionTarget: { @@ -1705,6 +1824,17 @@ describe("GlobalComposerPill", () => { it("uses a controlled external harness when the global provider differs", async () => { const user = userEvent.setup(); + mockGetModelsForAgent.mockImplementation((agentId: string) => + agentId === "claude-acp" + ? [ + { + id: "claude-opus-4-1", + name: "Claude Opus 4.1", + providerId: "claude-acp", + }, + ] + : [], + ); const onSend = renderGlobalComposer(vi.fn(), { currentExecutionTarget: { harnessId: "claude-acp", diff --git a/src/shared/ui/GlobalComposerPill.tsx b/src/shared/ui/GlobalComposerPill.tsx index 86ae76ed0..dae7f1351 100644 --- a/src/shared/ui/GlobalComposerPill.tsx +++ b/src/shared/ui/GlobalComposerPill.tsx @@ -454,6 +454,7 @@ export function GlobalComposerPill({ pickerAgents, availableModels, getModelsForAgent, + getProvenModelsForAgent, isModelInventoryAuthoritative, modelsLoading, modelStatusMessage, @@ -503,9 +504,18 @@ export function GlobalComposerPill({ providers, models: getModelsForAgent("goose"), getModelsForHarness: getModelsForAgent, + getProvenModelsForHarness: getProvenModelsForAgent, + isModelInventoryAuthoritative, catalogEntries, }), - [catalogEntries, getModelsForAgent, providers, selectedPersona], + [ + catalogEntries, + getModelsForAgent, + getProvenModelsForAgent, + isModelInventoryAuthoritative, + providers, + selectedPersona, + ], ); useEffect(() => { @@ -570,10 +580,21 @@ export function GlobalComposerPill({ ? selectedProviderForPicker : null; const defaultModelSelection = useMemo(() => { + const provenModels = getProvenModelsForAgent(selectedAgentId); + const selectableModels = availableModels.filter((model) => { + const providerId = model.providerId ?? concreteSelectedProviderId; + return provenModels.some( + (proven) => + proven.id === model.id && + (!providerId || + !proven.providerId || + proven.providerId === providerId), + ); + }); const storedPreference = getStoredModelPreference(selectedAgentId); if (storedPreference) { const matchingModel = findMatchingModel( - availableModels, + selectableModels, storedPreference.modelId, storedPreference.providerId, ); @@ -609,7 +630,7 @@ export function GlobalComposerPill({ gooseDefaultSelection.modelProviderId === concreteSelectedProviderId) ) { const matchingDefault = findMatchingModel( - availableModels, + selectableModels, gooseDefaultSelection.modelId, gooseDefaultSelection.modelProviderId, ); @@ -619,25 +640,21 @@ export function GlobalComposerPill({ selectedProviderForPicker, ); } - if ( - !isModelInventoryAuthoritative(gooseDefaultSelection.modelProviderId) - ) { - return gooseDefaultSelection; - } } const compatibleModels = concreteSelectedProviderId - ? availableModels.filter( + ? selectableModels.filter( (model) => !model.providerId || model.providerId === concreteSelectedProviderId, ) - : availableModels; + : selectableModels; return getPreferredModel(compatibleModels, selectedProviderForPicker); }, [ availableModels, concreteSelectedProviderId, + getProvenModelsForAgent, gooseDefaultSelection, isModelInventoryAuthoritative, selectedAgentId, @@ -651,20 +668,51 @@ export function GlobalComposerPill({ if (!currentExecutionTarget?.modelId) { return null; } + const modelProviderId = currentExecutionTarget.modelProviderId; + if ( + modelProviderId && + isModelInventoryAuthoritative(modelProviderId) && + !getProvenModelsForAgent(currentExecutionTarget.harnessId).some( + (model) => + model.id === currentExecutionTarget.modelId && + (!model.providerId || model.providerId === modelProviderId), + ) + ) { + return null; + } return { modelProviderId: currentExecutionTarget.modelProviderId, modelId: currentExecutionTarget.modelId, modelName: currentExecutionTarget.modelName, }; - }, [currentExecutionTarget]); + }, [ + currentExecutionTarget, + getProvenModelsForAgent, + isModelInventoryAuthoritative, + ]); const hasLocalExecutionOverride = providerOverride !== null || modelOverride !== null; + const personaSelectionOverridden = + personaOverrideUserOverrideForRef.current === selectedPersonaId; + // A persona target is the configuration sent to the runtime. Materialize the + // picker from that exact target: a provider-only target deliberately has no + // model selection and must not borrow a default model for display. + const personaModelSelection = + !personaSelectionOverridden && personaTarget?.modelId + ? { + modelProviderId: personaTarget.modelProviderId, + modelId: personaTarget.modelId, + modelName: personaTarget.modelName, + } + : null; const effectiveModelSelection = - modelOverride ?? - (!hasLocalExecutionOverride && currentExecutionTarget !== undefined - ? controlledModelSelection - : defaultModelSelection); + !personaSelectionOverridden && personaTarget + ? personaModelSelection + : (modelOverride ?? + (!hasLocalExecutionOverride && currentExecutionTarget !== undefined + ? controlledModelSelection + : defaultModelSelection)); const localExecutionTarget = useMemo( () => hasLocalExecutionOverride || currentExecutionTarget === undefined @@ -682,12 +730,19 @@ export function GlobalComposerPill({ selectedProviderForPicker, ], ); - const personaSelectionOverridden = - personaOverrideUserOverrideForRef.current === selectedPersonaId; + const personaHasSavedExecutionTarget = Boolean( + selectedPersona?.provider || + selectedPersona?.modelProviderId || + selectedPersona?.model, + ); + const controlledTargetInvalidated = + currentExecutionTarget?.modelId != null && controlledModelSelection == null; const effectiveExecutionTarget = - !personaSelectionOverridden && personaTarget + !personaSelectionOverridden && personaHasSavedExecutionTarget ? personaTarget - : (localExecutionTarget ?? currentExecutionTarget ?? undefined); + : (localExecutionTarget ?? + (controlledTargetInvalidated ? undefined : currentExecutionTarget) ?? + undefined); const canSend = hasSendableContent && Boolean(effectiveExecutionTarget) && @@ -1505,6 +1560,9 @@ export function GlobalComposerPill({ availableModels={availableModels} modelsLoading={modelsLoading} modelStatusMessage={modelStatusMessage} + showDefaultModelInTrigger={ + personaSelectionOverridden || !personaTarget + } onModelChange={handleModelChange} onOpen={handlePickerOpen} onOpenChange={setModelPickerOpen}