From 8c077fbee7fa253c53201dbbf2161b1423cce724 Mon Sep 17 00:00:00 2001 From: webreflection Date: Wed, 2 Sep 2026 12:43:45 +0200 Subject: [PATCH] fix(vscode): implement org-level default model selection --- .changeset/org-default-model-selection.md | 7 + packages/kilo-gateway/src/api/models.ts | 3 +- packages/kilo-gateway/src/api/profile.ts | 14 +- packages/kilo-gateway/test/api/models.test.ts | 29 +- .../kilo-gateway/test/api/profile.test.ts | 54 ++- packages/kilo-vscode/src/KiloProvider.ts | 13 +- .../src/kilo-provider/handlers/auth.ts | 4 + packages/kilo-vscode/src/provider-actions.ts | 42 ++- .../fixtures/session-provider-activity.tsx | 318 +++++++++++++++++- .../tests/unit/kilo-provider-catalog.test.ts | 116 +++++++ .../tests/unit/model-selection.test.ts | 112 +++++- .../unit/new-worktree-dialog-sandbox.test.ts | 231 +++++++++++++ .../tests/unit/provider-actions-save.test.ts | 113 +++++++ .../tests/unit/session-model-store.test.ts | 94 ++++++ .../unit/session-provider-activity.test.ts | 2 +- .../agent-manager/NewWorktreeDialog.tsx | 36 +- .../agent-manager/new-worktree-models.ts | 33 ++ .../src/components/chat/PromptInput.tsx | 6 +- .../webview-ui/src/context/model-selection.ts | 55 ++- .../webview-ui/src/context/provider.tsx | 42 ++- .../src/context/session-model-store.ts | 37 +- .../webview-ui/src/context/session-types.ts | 4 +- .../webview-ui/src/context/session.tsx | 189 ++++------- .../webview-ui/src/stories/StoryProviders.tsx | 6 +- .../src/stories/history.stories.tsx | 2 +- .../src/types/messages/extension-messages.ts | 3 + .../opencode/src/kilocode/provider/catalog.ts | 31 ++ .../server/httpapi/groups/kilo-gateway.ts | 1 + .../server/httpapi/handlers/kilo-gateway.ts | 9 +- packages/opencode/src/provider/model-cache.ts | 9 +- packages/opencode/src/provider/models.ts | 8 +- .../instance/httpapi/handlers/config.ts | 8 +- .../instance/httpapi/handlers/provider.ts | 5 + .../test/kilocode/kilo-loader-auth.test.ts | 133 ++++++-- .../test/kilocode/model-cache-org.test.ts | 75 ++++- .../kilocode/server/httpapi-public.test.ts | 1 + .../server/kilo-gateway-statuses.test.ts | 55 ++- .../prompt-training-model-filter.test.ts | 145 +++++++- packages/sdk/js/src/v2/gen/types.gen.ts | 1 + packages/sdk/openapi.json | 3 + 40 files changed, 1769 insertions(+), 280 deletions(-) create mode 100644 .changeset/org-default-model-selection.md create mode 100644 packages/kilo-vscode/tests/unit/kilo-provider-catalog.test.ts create mode 100644 packages/kilo-vscode/webview-ui/agent-manager/new-worktree-models.ts create mode 100644 packages/opencode/src/kilocode/provider/catalog.ts diff --git a/.changeset/org-default-model-selection.md b/.changeset/org-default-model-selection.md new file mode 100644 index 00000000000..6a7d63dab9f --- /dev/null +++ b/.changeset/org-default-model-selection.md @@ -0,0 +1,7 @@ +--- +"kilo-code": patch +"@kilocode/cli": patch +"@kilocode/kilo-gateway": patch +--- + +Use organization model defaults after login and team switching while preserving valid model preferences. Keep unavailable organization catalogs from falling back to public models. diff --git a/packages/kilo-gateway/src/api/models.ts b/packages/kilo-gateway/src/api/models.ts index 3ccb4ff2534..7508d24a429 100644 --- a/packages/kilo-gateway/src/api/models.ts +++ b/packages/kilo-gateway/src/api/models.ts @@ -240,8 +240,7 @@ async function fetchRawKiloModels(options?: { } if (!response.ok) { - // 401 with auth credentials: fall back to unauthenticated public endpoint - if (response.status === 401 && (token || organizationId)) { + if (response.status === 401 && token && !organizationId && !baseURL.includes("/api/organizations/")) { return fetchRawKiloModels({}) } const kind = response.status === 401 || response.status === 403 ? "unauthorized" : "http" diff --git a/packages/kilo-gateway/src/api/profile.ts b/packages/kilo-gateway/src/api/profile.ts index 0540af4eeff..87caa463598 100644 --- a/packages/kilo-gateway/src/api/profile.ts +++ b/packages/kilo-gateway/src/api/profile.ts @@ -99,7 +99,11 @@ export const getKiloBalance = fetchBalance * When token is provided, returns the authenticated user's default model * When no token is provided, returns the default free model for anonymous usage */ -export async function fetchDefaultModel(token?: string, organizationId?: string): Promise { +export async function fetchDefaultModel( + token?: string, + organizationId?: string, + fallback = token ? DEFAULT_MODEL : DEFAULT_FREE_MODEL, +): Promise { const path = organizationId ? `/api/organizations/${organizationId}/defaults` : `/api/defaults` const url = `${KILO_API_BASE}${path}` @@ -114,16 +118,16 @@ export async function fetchDefaultModel(token?: string, organizationId?: string) const response = await fetch(url, { headers }) if (!response.ok) { - return token ? DEFAULT_MODEL : DEFAULT_FREE_MODEL + return fallback } const data = (await response.json()) as { defaultModel?: string; defaultFreeModel?: string } if (token) { - return data.defaultModel || DEFAULT_MODEL + return data.defaultModel || fallback } - return data.defaultFreeModel || DEFAULT_FREE_MODEL + return data.defaultFreeModel || fallback } catch { - return token ? DEFAULT_MODEL : DEFAULT_FREE_MODEL + return fallback } } diff --git a/packages/kilo-gateway/test/api/models.test.ts b/packages/kilo-gateway/test/api/models.test.ts index 7d73b3d1b5e..20fab7516c0 100644 --- a/packages/kilo-gateway/test/api/models.test.ts +++ b/packages/kilo-gateway/test/api/models.test.ts @@ -1,6 +1,6 @@ // Verifies fetchKiloModels typed result and 401 fallback behaviour. -import { test, expect } from "bun:test" +import { test, expect, spyOn } from "bun:test" import { fetchKiloModels, fetchKiloTranscriptionModels } from "../../src/api/models.js" const VALID_RESPONSE = JSON.stringify({ @@ -113,7 +113,6 @@ test("falls back to public endpoint on 401 and returns models", async () => { const result = await fetchKiloModels({ kilocodeToken: "expired-token", - kilocodeOrganizationId: "org-123", }) ;(globalThis as any).fetch = orig @@ -123,6 +122,31 @@ test("falls back to public endpoint on 401 and returns models", async () => { expect(Object.keys(result.models).length).toBeGreaterThan(0) }) +test.each([ + { kilocodeToken: "expired-token", kilocodeOrganizationId: "org-123" }, + { kilocodeOrganizationId: "org-123" }, + { kilocodeToken: "expired-token", baseURL: "https://api.kilo.ai/api/organizations/org-123" }, + { kilocodeToken: "expired-token", baseURL: "https://gateway.test/api/organizations/org-123" }, +])("never retries an organization-scoped 401 against the public catalog: %j", async (options) => { + const fetch = spyOn(globalThis, "fetch").mockResolvedValue(new Response(null, { status: 401 })) + try { + expect(await fetchKiloModels(options)).toEqual({ models: {}, error: { kind: "unauthorized", status: 401 } }) + expect(fetch).toHaveBeenCalledTimes(1) + } finally { + fetch.mockRestore() + } +}) + +test("preserves a successful empty organization catalog", async () => { + const fetch = spyOn(globalThis, "fetch").mockResolvedValue(Response.json({ data: [] })) + try { + expect(await fetchKiloModels({ kilocodeToken: "token", kilocodeOrganizationId: "org-123" })).toEqual({ models: {} }) + expect(fetch).toHaveBeenCalledTimes(1) + } finally { + fetch.mockRestore() + } +}) + test("returns error with kind=network on fetch exception", async () => { const orig = globalThis.fetch stubFetch(async () => { @@ -445,4 +469,3 @@ test("omits cost when pricing contains negative values (dynamic/auto-routed pric cache_read: 0.3, }) }) - diff --git a/packages/kilo-gateway/test/api/profile.test.ts b/packages/kilo-gateway/test/api/profile.test.ts index 5ae6b2aaf5c..f4d4ee919aa 100644 --- a/packages/kilo-gateway/test/api/profile.test.ts +++ b/packages/kilo-gateway/test/api/profile.test.ts @@ -1,7 +1,57 @@ -import { describe, expect, test } from "bun:test" -import { defaultOrganizationId } from "../../src/api/profile.js" +import { describe, expect, spyOn, test } from "bun:test" +import { defaultOrganizationId, fetchDefaultModel } from "../../src/api/profile.js" +import { DEFAULT_FREE_MODEL, DEFAULT_MODEL, KILO_API_BASE } from "../../src/api/constants.js" import type { KilocodeProfile } from "../../src/types.js" +describe("fetchDefaultModel", () => { + const failures: [string, () => Promise][] = [ + ["missing", async () => Response.json({})], + ["empty", async () => Response.json({ defaultModel: "", defaultFreeModel: "" })], + ["unauthorized", async () => new Response(null, { status: 401 })], + ["server error", async () => new Response(null, { status: 500 })], + ["invalid JSON", async () => new Response("invalid")], + ["network error", async () => Promise.reject(new Error("offline"))], + ] + + test.each(failures)("preserves old defaults and accepts an Org fallback: %s", async (_, response) => { + const fetch = spyOn(globalThis, "fetch").mockImplementation( + Object.assign(response, { preconnect: globalThis.fetch.preconnect }), + ) + try { + expect(await fetchDefaultModel()).toBe(DEFAULT_FREE_MODEL) + expect(await fetchDefaultModel("token")).toBe(DEFAULT_MODEL) + expect(await fetchDefaultModel("token", "org")).toBe(DEFAULT_MODEL) + expect(await fetchDefaultModel("token", "org", "allowed/first")).toBe("allowed/first") + expect(await fetchDefaultModel("token", "org", "")).toBe("") + } finally { + fetch.mockRestore() + } + }) + + test("uses the API default ahead of the supplied fallback", async () => { + const fetch = spyOn(globalThis, "fetch").mockResolvedValue( + Response.json({ defaultModel: "allowed/default", defaultFreeModel: "public/free" }), + ) + try { + expect(await fetchDefaultModel("token", "org", "allowed/first")).toBe("allowed/default") + expect(fetch.mock.calls.at(0)?.at(0)).toBe(`${KILO_API_BASE}/api/organizations/org/defaults`) + } finally { + fetch.mockRestore() + } + }) + + test("keeps anonymous API defaults", async () => { + const fetch = spyOn(globalThis, "fetch").mockResolvedValue( + Response.json({ defaultModel: "paid", defaultFreeModel: "public/free" }), + ) + try { + expect(await fetchDefaultModel()).toBe("public/free") + } finally { + fetch.mockRestore() + } + }) +}) + const profile = (input: Partial = {}): KilocodeProfile => ({ email: "user@example.com", organizations: [{ id: "org_1", name: "Acme", role: "MEMBER" }], diff --git a/packages/kilo-vscode/src/KiloProvider.ts b/packages/kilo-vscode/src/KiloProvider.ts index 8ea71df3aea..041f1158836 100644 --- a/packages/kilo-vscode/src/KiloProvider.ts +++ b/packages/kilo-vscode/src/KiloProvider.ts @@ -2529,6 +2529,13 @@ export class KiloProvider implements vscode.WebviewViewProvider, TelemetryProper this.postMessage(message) } + private invalidateProviders(): void { + this.providersGeneration++ + this.providersQueued = false + this.cachedProvidersMessage = null + this.postMessage({ type: "providersLoading" }) + } + /** Fetch providers and send to webview. Coalesced: at most one in-flight + one queued. */ private async fetchAndSendProviders(): Promise { const next = ++this.providersGeneration @@ -2548,7 +2555,7 @@ export class KiloProvider implements vscode.WebviewViewProvider, TelemetryProper return } try { - const { response, authMethods, authStates, storedKeys } = await fetchProviderData( + const { response, authMethods, authStates, storedKeys, organizationId, ready } = await fetchProviderData( client, this.getWorkspaceDirectory(), ) @@ -2564,6 +2571,8 @@ export class KiloProvider implements vscode.WebviewViewProvider, TelemetryProper providers: indexProvidersById(response.all), connected: response.connected, defaults: response.default, + organizationId, + ready, defaultSelection: computeDefaultSelection( this.cachedConfigMessage as { config?: { model?: string } } | null, settings.get("providerID", ""), @@ -4383,6 +4392,7 @@ export class KiloProvider implements vscode.WebviewViewProvider, TelemetryProper getWorkspaceDirectory: () => this.getWorkspaceDirectory(), disposeGlobal: () => this.disposeGlobal(), invalidateProviderUsage: () => this.invalidateProviderUsage(), + invalidateProviders: () => this.invalidateProviders(), fetchAndSendProviders: () => this.fetchAndSendProviders(), fetchAndSendAgents: () => this.fetchAndSendAgents(), fetchAndSendSpeechToTextModels: () => this.fetchAndSendSpeechToTextModels(), @@ -4536,6 +4546,7 @@ export class KiloProvider implements vscode.WebviewViewProvider, TelemetryProper /** Re-fetch all server-side state after an auth change. */ private async reloadAfterAuthChange(): Promise { this.invalidateProviderUsage() + this.invalidateProviders() await this.fetchAndSendConfig() await Promise.all([ this.fetchAndSendProviders(), diff --git a/packages/kilo-vscode/src/kilo-provider/handlers/auth.ts b/packages/kilo-vscode/src/kilo-provider/handlers/auth.ts index 573a1053eca..3520e177c99 100644 --- a/packages/kilo-vscode/src/kilo-provider/handlers/auth.ts +++ b/packages/kilo-vscode/src/kilo-provider/handlers/auth.ts @@ -14,6 +14,7 @@ export interface AuthContext { getWorkspaceDirectory(): string disposeGlobal(): Promise invalidateProviderUsage(): void + invalidateProviders(): void fetchAndSendProviders(): Promise fetchAndSendAgents(): Promise fetchAndSendSpeechToTextModels(): Promise @@ -62,6 +63,7 @@ export async function handleLogin(ctx: AuthContext, attempt: number, getAttempt: console.log("[Kilo New] KiloProvider: 🔐 Login successful") ctx.invalidateProviderUsage() + ctx.invalidateProviders() await ctx.disposeGlobal() // Step 3: Fetch profile and push to webview @@ -88,6 +90,7 @@ export async function handleLogout(ctx: AuthContext): Promise { ctx.postMessage({ type: "profileData", data: null }) ctx.invalidateProviderUsage() + ctx.invalidateProviders() await ctx.disposeGlobal() await ctx.fetchAndSendProviders() @@ -123,6 +126,7 @@ export async function handleSetOrganization(ctx: AuthContext, organizationId: st } ctx.invalidateProviderUsage() + ctx.invalidateProviders() await ctx.disposeGlobal() // Org switch succeeded — refresh profile and providers independently (best-effort) diff --git a/packages/kilo-vscode/src/provider-actions.ts b/packages/kilo-vscode/src/provider-actions.ts index 9d67b2797f9..86f92a8759d 100644 --- a/packages/kilo-vscode/src/provider-actions.ts +++ b/packages/kilo-vscode/src/provider-actions.ts @@ -61,13 +61,24 @@ export async function fetchProviderData(client: KiloClient, dir: string) { : Promise.resolve({}) const kiloRequest = client.kilo .authStatus({ directory: dir }, { throwOnError: true }) - .then((r) => (r.data?.authenticated ? (r.data.type ?? null) : null)) - .catch(() => null) + .then((r) => r.data) + .catch(() => undefined) + const recommendation = kiloRequest.then(async (auth) => { + if (!auth?.organizationId) return undefined + return client.config + .providers({ directory: dir }, { throwOnError: true }) + .then((r) => r.data?.default.kilo) + .catch((error: unknown) => { + console.warn("[Kilo New] Failed to fetch organization model default:", error) + return undefined + }) + }) - const [{ data: response }, authMethods, kiloAuth] = await Promise.all([ + const [{ data: response }, authMethods, kiloAuth, recommended] = await Promise.all([ client.provider.list({ directory: dir }, { throwOnError: true }), authRequest, kiloRequest, + recommendation, ]) const authStates: Record = {} const storedKeys: Record = {} @@ -89,8 +100,29 @@ export async function fetchProviderData(client: KiloClient, dir: string) { return next as (typeof response.all)[number] }) delete authStates[KILO_PROVIDER_ID] - if (kiloAuth) authStates[KILO_PROVIDER_ID] = kiloAuth - return { response: { ...response, all }, authMethods, authStates, storedKeys } + if (kiloAuth?.authenticated && kiloAuth.type) authStates[KILO_PROVIDER_ID] = kiloAuth.type + const organizationId = kiloAuth ? (kiloAuth.organizationId ?? null) : undefined + const defaults = { ...response.default } + if (organizationId) { + const models = all.find((item) => item.id === KILO_PROVIDER_ID)?.models ?? {} + const model = recommended && models[recommended] ? recommended : Object.keys(models).at(0) + if (model) defaults[KILO_PROVIDER_ID] = model + if (!model) delete defaults[KILO_PROVIDER_ID] + } + if (!kiloAuth) delete defaults[KILO_PROVIDER_ID] + return { + response: { + ...response, + all: kiloAuth ? all : all.filter((item) => item.id !== KILO_PROVIDER_ID), + connected: kiloAuth ? response.connected : response.connected.filter((id) => id !== KILO_PROVIDER_ID), + default: defaults, + }, + authMethods, + authStates, + storedKeys, + organizationId, + ready: !!kiloAuth, + } } /** diff --git a/packages/kilo-vscode/tests/fixtures/session-provider-activity.tsx b/packages/kilo-vscode/tests/fixtures/session-provider-activity.tsx index 39656ebdc6f..94f45c40ee9 100644 --- a/packages/kilo-vscode/tests/fixtures/session-provider-activity.tsx +++ b/packages/kilo-vscode/tests/fixtures/session-provider-activity.tsx @@ -1,11 +1,12 @@ import assert from "node:assert/strict" import { Window } from "happy-dom" +import type { ModelSelection, WebviewMessage } from "../../webview-ui/src/types/messages" const window = new Window({ url: "http://localhost" }) Object.defineProperty(window, "origin", { value: window.location.origin }) -const sent: unknown[] = [] +const sent: WebviewMessage[] = [] const api = { - postMessage: (message: unknown) => sent.push(message), + postMessage: (message: WebviewMessage) => sent.push(message), getState: () => undefined, setState: () => {}, } @@ -37,7 +38,7 @@ Object.assign(globalThis, { }) const { render } = await import("solid-js/web") -const { For, Show, createSignal } = await import("solid-js") +const { For, Show, createEffect, createSignal } = await import("solid-js") const { WorktreeItem } = await import("../../webview-ui/agent-manager/WorktreeItem") const { SubagentPanel } = await import("../../webview-ui/agent-manager/SubagentPanel") const { DragDropProvider, SortableProvider } = await import("@thisbeyond/solid-dnd") @@ -47,24 +48,19 @@ const { ServerProvider } = await import("../../webview-ui/src/context/server") const { ConfigContext } = await import("../../webview-ui/src/context/config") const { LanguageContext } = await import("../../webview-ui/src/context/language") const { NotificationsProvider } = await import("../../webview-ui/src/context/notifications") -const { ProviderContext } = await import("../../webview-ui/src/context/provider") +const { ProviderProvider } = await import("../../webview-ui/src/context/provider") const { SessionProvider, useSession } = await import("../../webview-ui/src/context/session") const { post } = await import("../../webview-ui/src/utils/webview-message") const { terminal } = await import("../../webview-ui/src/context/session-outcome") +const { PromptInput } = await import("../../webview-ui/src/components/chat/PromptInput") +const { IndexingProvider } = await import("../../webview-ui/src/context/indexing") +const { MemoryProvider } = await import("../../webview-ui/src/context/memory") +const { SpeechToTextModelsProvider } = await import("../../webview-ui/src/context/speech-to-text-models") +const { drafts, imageDrafts, savePromptDraft } = await import("../../webview-ui/src/utils/draft-store") -const provider = { - providers: () => ({}), - connected: () => [], - defaults: () => ({}), - defaultSelection: () => ({ providerID: "kilocode", modelID: "auto" }), - models: () => [], - findModel: () => undefined, - authMethods: () => ({}), - authStates: () => ({}), - isModelValid: () => true, -} +const [settings, setSettings] = createSignal<{ model?: string; agent?: Record }>({}) const config = { - config: () => ({}), + config: settings, globalConfig: () => ({}), globalDraft: () => ({}), projectConfig: () => ({}), @@ -91,13 +87,16 @@ const language = { } const ref = { value: undefined as ReturnType | undefined } +const observed: (ModelSelection | null)[] = [] const [operation, setOperation] = createSignal(false) const [run, setRun] = createSignal(false) const [inspector, setInspector] = createSignal(false) +const [composer, setComposer] = createSignal(false) const [active, setActive] = createSignal("task-child") const Probe = () => { const session = useSession() ref.value = session + createEffect(() => observed.push(session.selected())) const ids = ["root", "background"] const deps = { terms: { activeId: () => undefined }, @@ -167,6 +166,15 @@ const Probe = () => { onClosePanel={() => setInspector(false)} /> + + + + + + + + + ) } @@ -179,7 +187,7 @@ const dispose = render( () => ( - + @@ -189,7 +197,7 @@ const dispose = render( - + ), @@ -278,6 +286,280 @@ try { const value = ref.value assert(value) + const auto = { providerID: "kilo", modelID: "kilo-auto/free" } + const personal = { providerID: "kilo", modelID: "personal" } + const first = { providerID: "kilo", modelID: "z-first" } + const recommended = { providerID: "kilo", modelID: "a-recommended" } + const external = { providerID: "openai", modelID: "external" } + const choice = (actual: ModelSelection | null, expected: ModelSelection) => { + assert.equal(actual?.providerID, expected.providerID) + assert.equal(actual?.modelID, expected.modelID) + } + const writes = () => sent.filter((item) => item.type === "persistModelSelection" || item.type === "persistRecents") + const requests = () => + sent.filter((item) => ["sendMessage", "sendCommand", "importAndSend", "compact"].includes(item.type)) + const catalog = async (organizationId: string | null, ids: string[], model?: string, ready = true) => { + await emit({ + type: "providersLoaded", + organizationId, + ready, + providers: { + kilo: { id: "kilo", name: "Kilo", models: Object.fromEntries(ids.map((id) => [id, { id, name: id }])) }, + openai: { id: "openai", name: "OpenAI", models: { external: { id: "external", name: "External" } } }, + }, + connected: ["kilo", "openai"], + defaults: model ? { kilo: model } : {}, + defaultSelection: auto, + authMethods: {}, + authStates: {}, + }) + } + + assert.equal(value.selected(), null) + value.sendMessage("initial pending") + assert.equal(requests().length, 0) + await emit({ type: "agentsLoaded", agents: [{ name: "code" }, { name: "ask" }], defaultAgent: "code" }) + await emit({ type: "recentsLoaded", recents: [auto, first, external] }) + await catalog("org-a", [first.modelID, recommended.modelID, auto.modelID], recommended.modelID) + choice(value.selected(), recommended) + choice(value.selected("selection"), recommended) + choice(value.modelForAgent("ask"), recommended) + assert.deepEqual(writes(), []) + + observed.length = 0 + await catalog("org-b", [first.modelID, recommended.modelID, auto.modelID], first.modelID) + choice(value.selected(), first) + assert(observed.length > 0) + assert(observed.every((selection) => selection?.modelID === first.modelID)) + for (const model of [undefined, "disallowed"]) { + await catalog("org-a", [first.modelID, recommended.modelID], model) + choice(value.selected(), first) + } + assert.deepEqual(writes(), []) + + await catalog(null, [auto.modelID, personal.modelID]) + choice(value.selected(), auto) + value.selectModel(personal.providerID, personal.modelID) + await settle() + assert.equal(writes().length, 2) + choice(value.selected(), personal) + value.setSessionModel("selection", personal.providerID, personal.modelID) + value.setCurrentSessionID("selection") + const remembered = writes().slice() + await catalog("org-a", [first.modelID, recommended.modelID], recommended.modelID) + choice(value.selected(), recommended) + choice(value.selected("selection"), recommended) + choice(value.modelForAgent("code"), recommended) + await catalog(null, [auto.modelID, personal.modelID]) + choice(value.selected(), personal) + choice(value.modelForAgent("code"), personal) + assert.deepEqual(writes(), remembered) + + await emit({ type: "modelSelectionsLoaded", selections: {} }) + await emit({ type: "recentsLoaded", recents: [auto] }) + choice(value.selected(), personal) + setSettings({ model: "kilo/personal" }) + await settle() + setSettings({ model: "kilo/a-recommended" }) + await catalog("org-a", [first.modelID, recommended.modelID], recommended.modelID) + choice(value.selected(), recommended) + setSettings({}) + await catalog(null, [auto.modelID, personal.modelID]) + choice(value.selected(), personal) + assert.deepEqual(writes(), remembered) + + await catalog("org-a", [first.modelID, recommended.modelID], recommended.modelID) + await emit({ + type: "messagesLoaded", + sessionID: "history", + messages: [ + { + id: "history-message", + sessionID: "history", + role: "user", + model: personal, + createdAt: info("history").createdAt, + }, + ], + }) + choice(value.selected("history"), recommended) + await catalog(null, [auto.modelID, personal.modelID]) + choice(value.selected("history"), personal) + assert.deepEqual(writes(), remembered) + + for (const pending of ["retained", "loading", "empty"]) { + if (pending === "retained") + await catalog("org-a", [personal.modelID, recommended.modelID], recommended.modelID, false) + if (pending === "loading") await emit({ type: "providersLoading" }) + if (pending === "empty") await catalog("org-a", [], recommended.modelID) + assert.equal(value.selected(), null) + assert.equal(value.selected("selection"), null) + assert.equal(value.modelForAgent("code"), null) + const before = requests().length + assert.equal(value.sendMessage("blocked"), false) + assert.equal(value.sendMessage("blocked explicit", personal.providerID, personal.modelID), false) + assert.equal(value.sendCommand("blocked", ""), false) + value.compact() + assert.equal(requests().length, before) + assert.deepEqual(writes(), remembered) + } + + value.setSessionModel("external", external.providerID, external.modelID) + await emit({ type: "providersLoading" }) + choice(value.selected("external"), external) + await catalog("org-a", [first.modelID, recommended.modelID], recommended.modelID) + const before = requests().length + value.sendMessage("invalid explicit", personal.providerID, personal.modelID) + value.sendCommand("invalid", "", undefined, undefined, undefined, undefined, undefined, undefined, { + model: "kilo/personal", + }) + assert.equal(requests().length, before) + assert.deepEqual(writes(), remembered) + assert.equal(value.sendMessage("effective model"), true) + const message = requests().at(-1) + assert(message?.type === "sendMessage") + assert.equal(message.providerID, recommended.providerID) + assert.equal(message.modelID, recommended.modelID) + await emit({ + type: "messageCreated", + message: { + id: message.messageID, + sessionID: "selection", + role: "user", + model: recommended, + createdAt: info("selection").createdAt, + }, + }) + await catalog(null, [auto.modelID, personal.modelID]) + choice(value.selected(), personal) + await catalog("org-a", [first.modelID, recommended.modelID], recommended.modelID) + await emit({ type: "sessionStatus", sessionID: "selection", status: "idle" }) + assert.equal(value.sendCommand("effective", ""), true) + const command = requests().at(-1) + assert(command?.type === "sendCommand") + assert.equal(command.providerID, recommended.providerID) + assert.equal(command.modelID, recommended.modelID) + value.setCurrentSessionID("cloud:preview") + assert.equal(value.sendMessage("cloud effective model"), true) + const cloud = requests().at(-1) + assert(cloud?.type === "importAndSend") + assert.equal(cloud.providerID, recommended.providerID) + assert.equal(cloud.modelID, recommended.modelID) + assert.equal(value.sendCommand("cloud", ""), true) + const imported = requests().at(-1) + assert(imported?.type === "importAndSend") + assert.equal(imported.providerID, recommended.providerID) + assert.equal(imported.modelID, recommended.modelID) + await catalog("org-a", []) + const blocked = requests().length + assert.equal(value.sendMessage("cloud unavailable"), false) + assert.equal(value.sendCommand("cloud", "unavailable"), false) + assert.equal(requests().length, blocked) + await catalog("org-a", [first.modelID, recommended.modelID], recommended.modelID) + assert.deepEqual(writes(), remembered) + + value.setCurrentSessionID(undefined) + await emit({ type: "modelSelectionsLoaded", selections: { code: personal } }) + choice(value.selected(), recommended) + await catalog(null, [auto.modelID, personal.modelID]) + choice(value.selected(), personal) + await catalog("org-a", [auto.modelID, first.modelID, recommended.modelID], recommended.modelID) + await emit({ type: "modelSelectionsLoaded", selections: { code: auto } }) + choice(value.selected(), auto) + setSettings({ agent: { code: { model: "kilo/z-first" } } }) + await settle() + choice(value.modelForAgent("code"), first) + choice(value.selected(), auto) + setSettings({}) + await emit({ type: "modelSelectionsLoaded", selections: {} }) + choice(value.selected(), recommended) + assert.deepEqual(writes(), remembered) + + const key = "acceptance:session:composer" + const image = { id: "image", filename: "image.png", mime: "image/png", dataUrl: "data:image/png;base64,cGl4ZWw=" } + const input = () => { + const element = host.querySelector("textarea.prompt-input") + assert(element) + return element + } + const seed = async (text: string) => { + setComposer(false) + await settle() + value.setCurrentSessionID("composer") + await emit({ type: "sessionStatus", sessionID: "composer", status: "idle" }) + savePromptDraft(key, text, [], [image]) + setComposer(true) + await settle() + await emit({ + type: "commandsLoaded", + commands: [ + { name: "review-test", description: "Test command", hints: [] }, + { name: "unavailable-test", description: "Unavailable command", hints: [], model: "kilo/unavailable" }, + ], + }) + assert.equal(input().value, text) + } + const submit = (enter: boolean) => { + if (enter) { + input().dispatchEvent(new window.KeyboardEvent("keydown", { key: "Enter", bubbles: true })) + return + } + const button = host.querySelector('[aria-label="prompt.action.send"]') + assert(button) + button.click() + } + const retained = (text: string, count: number) => { + assert.equal(requests().length, count) + assert.equal(input().value, text) + assert.equal(drafts.get(key), text) + assert.deepEqual(imageDrafts.get(key), [image]) + assert(host.querySelector('img[src="data:image/png;base64,cGl4ZWw="]')) + } + for (const text of ["preserve this draft", "/review-test preserve this draft"]) { + for (const empty of [false, true]) { + await seed(text) + if (empty) await catalog("org-a", []) + if (!empty) await emit({ type: "providersLoading" }) + const count = requests().length + submit(empty) + await settle() + retained(text, count) + } + await catalog("org-a", [recommended.modelID], recommended.modelID) + const count = requests().length + submit(false) + await settle() + assert.equal(requests().length, count + 1) + const request = requests().at(-1) + assert(request?.type === (text.startsWith("/") ? "sendCommand" : "sendMessage")) + assert.deepEqual(request.files, [{ mime: image.mime, url: image.dataUrl, filename: image.filename }]) + assert.equal(input().value, "") + assert.equal(drafts.has(key), false) + assert.equal(imageDrafts.has(key), false) + assert.equal(host.querySelector('img[src="data:image/png;base64,cGl4ZWw="]'), null) + } + await seed("/unavailable-test preserve command") + const rejected = requests().length + submit(false) + await settle() + retained("/unavailable-test preserve command", rejected) + + for (const text of ["prepare @terminal", "/review-test prepare @terminal"]) { + await catalog("org-a", [recommended.modelID], recommended.modelID) + await seed(text) + const count = requests().length + const start = sent.length + submit(true) + const request = sent.slice(start).find((message) => message.type === "requestTerminalContext") + assert(request?.type === "requestTerminalContext") + await emit({ type: "providersLoading" }) + await emit({ type: "terminalContextResult", requestId: request.requestId, content: "terminal output" }) + retained(text, count) + } + setComposer(false) + await settle() + await catalog("org-a", [recommended.modelID], recommended.modelID) + value.setCurrentSessionID("root") await check("root", "idle") await check("background", "idle") diff --git a/packages/kilo-vscode/tests/unit/kilo-provider-catalog.test.ts b/packages/kilo-vscode/tests/unit/kilo-provider-catalog.test.ts new file mode 100644 index 00000000000..04619a51720 --- /dev/null +++ b/packages/kilo-vscode/tests/unit/kilo-provider-catalog.test.ts @@ -0,0 +1,116 @@ +import { describe, expect, it } from "bun:test" + +const { KiloProvider } = await import("../../src/KiloProvider") + +const catalog = (org: string) => ({ + data: { + all: [{ id: "kilo", name: "Kilo Gateway", models: { [`${org}/model`]: { id: `${org}/model` } } }], + connected: ["kilo"], + default: { kilo: "kilo-auto/free" }, + }, +}) + +type Internals = { + connectionState: string + cachedProvidersMessage: unknown + fetchAndSendProviders(): Promise + invalidateProviders(): void +} + +function setup(list: () => Promise>, org: () => string) { + const client = { + provider: { list, auth: async () => ({ data: {} }) }, + kilo: { authStatus: async () => ({ data: { authenticated: true, type: "oauth", organizationId: org() } }) }, + config: { providers: async () => ({ data: { default: { kilo: `${org()}/model` } } }) }, + } + const provider = new KiloProvider({} as never, { getClient: () => client } as never) + const internal = provider as unknown as Internals + internal.connectionState = "connected" + const messages: Array> = [] + provider.postMessage = (message) => void messages.push(message as Record) + return { internal, messages } +} + +describe("KiloProvider catalog refresh", () => { + it("invalidates cached Kilo data before another account refresh", async () => { + const { internal, messages } = setup( + async () => catalog("org"), + () => "org", + ) + await internal.fetchAndSendProviders() + expect(internal.cachedProvidersMessage).toMatchObject({ organizationId: "org", ready: true }) + + internal.invalidateProviders() + + expect(internal.cachedProvidersMessage).toBeNull() + expect(messages.at(-1)).toEqual({ type: "providersLoading" }) + }) + + it("publishes only the newest catalog and recommendation after a queued switch", async () => { + const first = Promise.withResolvers>() + const started = Promise.withResolvers() + let org = "a" + let calls = 0 + const { internal, messages } = setup( + async () => { + calls++ + if (calls !== 1) return catalog(org) + started.resolve() + return first.promise + }, + () => org, + ) + + const before = internal.fetchAndSendProviders() + await started.promise + org = "b" + const after = internal.fetchAndSendProviders() + first.resolve(catalog("a")) + await Promise.all([before, after]) + + expect(calls).toBe(2) + expect(messages).toHaveLength(1) + expect(messages.at(0)).toMatchObject({ + type: "providersLoaded", + organizationId: "b", + ready: true, + defaults: { kilo: "b/model" }, + providers: { kilo: { models: { "b/model": { id: "b/model" } } } }, + }) + }) + + it("cannot republish an in-flight old catalog after invalidation", async () => { + const first = Promise.withResolvers>() + const { internal, messages } = setup( + () => first.promise, + () => "old", + ) + const pending = internal.fetchAndSendProviders() + + internal.invalidateProviders() + first.resolve(catalog("old")) + await pending + + expect(messages).toEqual([{ type: "providersLoading" }]) + expect(internal.cachedProvidersMessage).toBeNull() + }) + + it("does not restore an old catalog when the new account cannot load", async () => { + let fail = false + const { internal, messages } = setup( + async () => { + if (fail) throw new Error("Catalog unavailable") + return catalog("old") + }, + () => "old", + ) + await internal.fetchAndSendProviders() + internal.invalidateProviders() + fail = true + await internal.fetchAndSendProviders() + + expect(messages.at(-1)).toEqual({ type: "providersLoading" }) + expect(messages.filter((message) => message.type === "providersLoaded")).toHaveLength(1) + expect(internal.cachedProvidersMessage).toBeNull() + }) +}) diff --git a/packages/kilo-vscode/tests/unit/model-selection.test.ts b/packages/kilo-vscode/tests/unit/model-selection.test.ts index c61d1df1db6..3fd9fe8d6e7 100644 --- a/packages/kilo-vscode/tests/unit/model-selection.test.ts +++ b/packages/kilo-vscode/tests/unit/model-selection.test.ts @@ -1,7 +1,7 @@ import { describe, expect, it } from "bun:test" import { resolveModelSelection } from "../../webview-ui/src/context/model-selection" import { KILO_AUTO, parseModelString } from "../../src/shared/provider-model" -import type { Provider } from "../../webview-ui/src/types/messages" +import type { ModelSelection, Provider } from "../../webview-ui/src/types/messages" function makeProvider(id: string, name: string, modelIds: string[]): Provider { const models: Provider["models"] = {} @@ -84,16 +84,16 @@ describe("resolveModelSelection", () => { expect(result).toEqual(KILO_AUTO) }) - it("keeps the explicit fallback even when kilo is missing from the loaded catalog", () => { + it("rejects a fallback missing from the loaded catalog", () => { const result = resolveModelSelection({ providers: { openai: providers.openai }, connected: [], fallback: KILO_AUTO, }) - expect(result).toEqual(KILO_AUTO) + expect(result).toBeNull() }) - it("keeps the raw preference order before providers load", () => { + it("does not treat an empty catalog as unvalidated preferences", () => { const result = resolveModelSelection({ providers: {}, connected: [], @@ -101,6 +101,108 @@ describe("resolveModelSelection", () => { mode: { providerID: "anthropic", modelID: "claude-sonnet-4" }, fallback: KILO_AUTO, }) - expect(result).toEqual({ providerID: "openai", modelID: "gpt-4.1" }) + expect(result).toBeNull() + }) +}) + +describe("organization model selection", () => { + const first = { providerID: "kilo", modelID: "z-first" } + const recommendation = { providerID: "kilo", modelID: "a-default" } + const recent = { providerID: "kilo", modelID: "older-recent" } + const external = { providerID: "openai", modelID: "gpt-4.1" } + const input = { + providers: { + ...providers, + kilo: makeProvider("kilo", "Kilo Gateway", [ + first.modelID, + recommendation.modelID, + recent.modelID, + "kilo-auto/free", + ]), + }, + connected: ["openai"], + ready: true, + organizationId: "org-a", + defaults: { kilo: recommendation.modelID }, + recent: [{ providerID: "kilo", modelID: "missing-recent" }, recent, external], + fallback: KILO_AUTO, + } + + it("uses the recommendation for fresh Org login instead of recents or the generic fallback", () => { + expect(resolveModelSelection(input)).toEqual(recommendation) + }) + + it.each([undefined, "", "unavailable"])("uses catalog order for an absent or invalid default %s", (model) => { + expect(resolveModelSelection({ ...input, defaults: model === undefined ? {} : { kilo: model } })).toEqual(first) + }) + + it.each(["session", "override", "mode", "global"] as const)( + "preserves a valid %s before the recommendation", + (key) => { + expect(resolveModelSelection({ ...input, [key]: KILO_AUTO })).toEqual(KILO_AUTO) + }, + ) + + it("validates session, manual, mode, and global preferences in order", () => { + const missing = { providerID: "kilo", modelID: "missing" } + const choices = { session: KILO_AUTO, override: recent, mode: first, global: external } + expect(resolveModelSelection({ ...input, ...choices })).toEqual(KILO_AUTO) + expect(resolveModelSelection({ ...input, ...choices, session: missing })).toEqual(recent) + expect(resolveModelSelection({ ...input, ...choices, session: missing, override: missing })).toEqual(first) + expect(resolveModelSelection({ ...input, ...choices, session: missing, override: missing, mode: missing })).toEqual( + external, + ) + expect( + resolveModelSelection({ ...input, session: missing, override: missing, mode: missing, global: missing }), + ).toEqual(recommendation) + }) + + it("preserves explicitly configured external providers only while connected", () => { + expect(resolveModelSelection({ ...input, override: external })).toEqual(external) + expect(resolveModelSelection({ ...input, connected: [], override: external })).toEqual(recommendation) + }) + + it.each([{}, { kilo: makeProvider("kilo", "Kilo Gateway", []) }, { openai: providers.openai }])( + "does not fall back to free models or external recents for an empty Org catalog", + (catalog) => { + expect(resolveModelSelection({ ...input, providers: catalog, override: KILO_AUTO })).toBeNull() + }, + ) + + it("keeps explicit external models available with an empty Org catalog", () => { + expect(resolveModelSelection({ ...input, providers: { openai: providers.openai }, override: external })).toEqual( + external, + ) + }) + + it("does not trust a retained Kilo catalog while refresh or auth context is pending", () => { + for (const pending of [{ ready: false }, { organizationId: undefined }]) { + expect(resolveModelSelection({ ...input, ...pending, override: KILO_AUTO })).toBeNull() + expect(resolveModelSelection({ ...input, ...pending, override: external })).toEqual(external) + } + }) + + it("keeps Personal recents ahead of defaults and validates its final fallback", () => { + expect(resolveModelSelection({ ...input, organizationId: null })).toEqual(recent) + expect(resolveModelSelection({ ...input, organizationId: null, recent: [] })).toEqual(KILO_AUTO) + expect( + resolveModelSelection({ ...input, organizationId: null, recent: [], fallback: external, connected: [] }), + ).toBeNull() + }) + + it("restores the same explicit choice through Personal, Org A, Org B, and Personal", () => { + const override: ModelSelection = { providerID: "kilo", modelID: "personal" } + const personal = { + ...input, + organizationId: null, + providers: { kilo: makeProvider("kilo", "Kilo", [override.modelID]) }, + } + expect(resolveModelSelection({ ...personal, override })).toEqual(override) + expect(resolveModelSelection({ ...input, override })).toEqual(recommendation) + expect( + resolveModelSelection({ ...input, organizationId: "org-b", defaults: { kilo: first.modelID }, override }), + ).toEqual(first) + expect(resolveModelSelection({ ...personal, override })).toEqual(override) + expect(override).toEqual({ providerID: "kilo", modelID: "personal" }) }) }) diff --git a/packages/kilo-vscode/tests/unit/new-worktree-dialog-sandbox.test.ts b/packages/kilo-vscode/tests/unit/new-worktree-dialog-sandbox.test.ts index 1bb96d3f751..1e7233186ea 100644 --- a/packages/kilo-vscode/tests/unit/new-worktree-dialog-sandbox.test.ts +++ b/packages/kilo-vscode/tests/unit/new-worktree-dialog-sandbox.test.ts @@ -35,3 +35,234 @@ describe("NewWorktreeDialog base branch", () => { expect(src).not.toContain("baseBranch: advanced ? (baseBranch() ?? undefined) : undefined") }) }) + +function check(code: string) { + const cwd = join(__dirname, "..", "..", "webview-ui") + const script = ` + import assert from "node:assert/strict" + import { dirname, join } from "node:path" + import { plugin } from "bun" + import { isModelValid } from "./src/context/provider-utils.ts" + import { toggleModel, setAllocationVariant } from "./agent-manager/multi-model-utils.ts" + + const solid = join(dirname(require.resolve("solid-js")), "solid.js") + plugin({ + name: "solid-browser", + setup(build) { + build.onResolve({ filter: /^solid-js$/ }, () => ({ path: solid })) + }, + }) + const { batch, createComputed, createRoot, createSignal } = await import("solid-js") + const { createDialogModels } = await import("./agent-manager/new-worktree-models.ts") + + const x = { providerID: "kilo", modelID: "x" } + const y = { providerID: "kilo", modelID: "y" } + const z = { providerID: "kilo", modelID: "z" } + const free = { providerID: "kilo", modelID: "kilo-auto/free" } + const external = { providerID: "external", modelID: "custom" } + const catalog = (...models) => Object.fromEntries( + [...new Set(models.map((model) => model.providerID))].map((id) => [id, { + id, + name: id, + models: Object.fromEntries(models.filter((model) => model.providerID === id).map((model) => [ + model.modelID, + { id: model.modelID, name: model.modelID, variants: { high: {} } }, + ])), + }]), + ) + function scene(saved, initial = { providers: catalog(x, y), fallback: y, ready: true, connected: [] }) { + const [snapshot, refresh] = createSignal(initial) + const [agent, switchAgent] = createSignal("code") + const state = createDialogModels({ + saved, + ready: () => snapshot().ready, + valid: (value) => isModelValid(snapshot().providers, snapshot().connected, value), + variants: (value) => Object.keys(snapshot().providers[value.providerID]?.models[value.modelID]?.variants ?? {}), + fallback: () => agent() === "code" ? snapshot().fallback : snapshot().alternate ?? null, + }) + const seen = [] + createComputed(() => seen.push(state.model())) + return { state, snapshot, refresh: (update) => refresh((current) => ({ ...current, ...update })), switchAgent, seen } + } + createRoot((dispose) => { + try { + ${code} + } finally { + dispose() + } + }) + ` + const child = Bun.spawnSync([process.execPath, "--conditions=browser", "-e", script], { + cwd, + stdout: "pipe", + stderr: "pipe", + }) + expect(child.exitCode, child.stdout.toString() + child.stderr.toString()).toBe(0) +} + +describe("NewWorktreeDialog models", () => { + it("persists only the saved choice and wires the effective model to display, variants, and guarded submission", () => { + expect(src).toContain("saved: saved.model,") + expect(src).toContain("fallback: () => session.modelForAgent(agent()),") + expect(src).toContain("ready: provider.ready,") + expect(src).toContain("const model = selection.model") + expect(src).toContain("model: selection.choice(),") + expect(src).not.toContain("model: model(),") + expect(src).toContain("selection.select(undefined)") + expect(src).not.toContain("setModel(") + expect(src).toContain("selection.select(next)") + expect(src).toContain("value={model()}") + expect(src).toContain("const sel = model()") + expect(src).toContain("session.variantForAgent(agent(), model())") + expect(src).toContain("const sel = isCompare ? null : model()") + expect(src).toContain("return selection.canSubmit(compareMode() ? modelAllocations() : undefined)") + expect(src).toContain("if (!canSubmit()) return") + expect(src).toContain("disabled={!canSubmit()}") + }) + + it("keeps saved X through reactive X to Y to X catalog changes", () => { + check(` + const { state, refresh, seen } = scene(x) + assert.deepEqual(state.model(), x) + refresh({ providers: catalog(y) }) + assert.deepEqual(state.model(), y) + assert.deepEqual(state.choice(), x) + assert.equal(state.canSubmit(), true) + refresh({ providers: catalog(x, y) }) + assert.deepEqual(state.choice(), x) + assert.deepEqual(seen, [x, y, x]) + `) + }) + + it("restores an initially unavailable cached X without replacing it with Y", () => { + check(` + const { state, refresh } = scene(x, { providers: catalog(y), fallback: y, ready: true, connected: [] }) + assert.deepEqual(state.model(), y) + assert.deepEqual(state.choice(), x) + const reopened = scene(state.choice(), { providers: catalog(y), fallback: y, ready: true, connected: [] }) + assert.deepEqual(reopened.state.model(), y) + reopened.refresh({ providers: catalog(x, y) }) + assert.deepEqual(reopened.state.model(), x) + state.select(y) + refresh({ providers: catalog(x, y) }) + assert.deepEqual(state.model(), y) + assert.deepEqual(state.choice(), y) + `) + }) + + it("never saves automatic initial, agent, or refreshed organization defaults", () => { + check(` + const { state, refresh, switchAgent, seen } = scene(undefined) + assert.deepEqual(state.model(), y) + assert.equal(state.choice(), undefined) + state.select(y) + assert.deepEqual(state.choice(), y) + refresh({ providers: catalog(y, z), alternate: z }) + batch(() => { + switchAgent("plan") + state.select(undefined) + }) + assert.deepEqual(state.model(), z) + assert.equal(state.choice(), undefined) + refresh({ providers: catalog(x), alternate: x }) + assert.deepEqual(seen, [y, z, x]) + assert.equal(state.choice(), undefined) + `) + }) + + it("retains explicit legacy free and connected external models", () => { + check(` + const initial = { providers: catalog(free, external, y), fallback: y, ready: true, connected: ["external"] } + assert.deepEqual(scene(free, initial).state.model(), free) + const { state, refresh } = scene(external, initial) + assert.deepEqual(state.model(), external) + refresh({ ready: false, providers: catalog(external) }) + assert.deepEqual(state.model(), external) + assert.equal(state.canSubmit(), true) + refresh({ ready: true, providers: catalog(external, y), connected: [] }) + assert.deepEqual(state.model(), y) + assert.deepEqual(state.choice(), external) + refresh({ connected: ["external"] }) + assert.deepEqual(state.model(), external) + `) + }) + + it("keeps external-only comparisons usable while a Kilo catalog refresh blocks mixed comparisons", () => { + check(` + const { state, refresh } = scene(x, { + providers: catalog(x, external, y), fallback: y, ready: true, connected: ["external"], + }) + const solo = toggleModel(new Map(), "external", "custom", "Custom") + const mixed = toggleModel(solo, "kilo", "x", "X") + const original = [...mixed.values()].map((entry) => ({ ...entry })) + assert.equal(state.canSubmit(solo), true) + assert.equal(state.canSubmit(mixed), true) + refresh({ ready: false, providers: catalog(external), fallback: null }) + assert.equal(state.model(), null) + assert.equal(state.canSubmit(), false) + assert.equal(state.canSubmit(solo), true) + assert.equal(state.canSubmit(mixed), false) + assert.deepEqual(state.choice(), x) + assert.deepEqual([...mixed.values()], original) + refresh({ ready: true, providers: catalog(x, external, y), fallback: y }) + assert.deepEqual(state.model(), x) + assert.equal(state.canSubmit(mixed), true) + `) + }) + + it("blocks pending, empty, and invalid fallback catalogs without clearing a saved choice", () => { + check(` + const { state, refresh, seen } = scene(x) + refresh({ ready: false }) + assert.equal(state.model(), null) + assert.equal(state.canSubmit(), false) + assert.deepEqual(state.choice(), x) + refresh({ ready: true, providers: {} }) + assert.equal(state.model(), null) + assert.equal(state.canSubmit(), false) + refresh({ providers: catalog(y), fallback: x }) + assert.equal(state.canSubmit(), false) + refresh({ fallback: null }) + assert.equal(state.canSubmit(), false) + refresh({ providers: catalog(x) }) + assert.deepEqual(seen, [x, null, x]) + assert.deepEqual(state.choice(), x) + assert.equal(state.canSubmit(), true) + `) + }) + + it("blocks invalid comparison models and variants without rewriting explicit allocations", () => { + check(` + const { state, refresh } = scene(x) + const first = toggleModel(new Map(), "kilo", "x", "X") + const allocations = toggleModel(first, "kilo", "y", "Y") + const original = [...allocations.values()].map((entry) => ({ ...entry })) + const [allowed, setAllowed] = createSignal(false) + createComputed(() => setAllowed(state.canSubmit(allocations))) + assert.equal(allowed(), true) + refresh({ providers: catalog(y) }) + assert.deepEqual(state.model(), y) + assert.equal(allowed(), false) + assert.deepEqual([...allocations.values()], original) + refresh({ providers: catalog(x, y) }) + assert.equal(allowed(), true) + refresh({ ready: false }) + assert.equal(allowed(), false) + refresh({ ready: true }) + const variants = setAllocationVariant(allocations, "kilo", "x", "high") + assert.equal(state.canSubmit(variants), true) + refresh({ providers: { kilo: { id: "kilo", name: "kilo", models: { + x: { id: "x", name: "X", variants: { low: {} } }, + y: { id: "y", name: "Y" }, + } } } }) + assert.equal(state.canSubmit(variants), false) + assert.equal(variants.get("kilo/x").variant, "high") + assert.equal(state.canSubmit(new Map()), false) + const disconnected = toggleModel(new Map(), "external", "custom", "Custom") + refresh({ providers: catalog(external), connected: [] }) + assert.equal(state.canSubmit(disconnected), false) + refresh({ connected: ["external"] }) + assert.equal(state.canSubmit(disconnected), true) + `) + }) +}) diff --git a/packages/kilo-vscode/tests/unit/provider-actions-save.test.ts b/packages/kilo-vscode/tests/unit/provider-actions-save.test.ts index 8f0e10b60a8..9ae72b6b65c 100644 --- a/packages/kilo-vscode/tests/unit/provider-actions-save.test.ts +++ b/packages/kilo-vscode/tests/unit/provider-actions-save.test.ts @@ -438,6 +438,119 @@ describe("disconnectProvider", () => { }) describe("fetchProviderData", () => { + for (const item of [ + { name: "uses the allowed organization default", recommended: "org/default", expected: "org/default" }, + { name: "uses the first allowed model when no default exists", recommended: undefined, expected: "org/first" }, + { + name: "ignores a default outside the organization catalog", + recommended: "kilo-auto/free", + expected: "org/first", + }, + { name: "uses the first allowed model when defaults cannot load", error: true, expected: "org/first" }, + { + name: "does not invent a default for an empty catalog", + empty: true, + recommended: "org/default", + expected: undefined, + }, + ]) { + it(item.name, async () => { + const directories: string[] = [] + const client = { + provider: { + list: async () => ({ + data: { + all: [ + { + id: "kilo", + name: "Kilo Gateway", + models: item.empty ? {} : { "org/first": { id: "org/first" }, "org/default": { id: "org/default" } }, + }, + { id: "anthropic", name: "Anthropic", models: { claude: { id: "claude" } }, metadata: { priority: 1 } }, + ], + connected: ["kilo", "anthropic"], + default: { kilo: "kilo-auto/free", anthropic: "claude" }, + }, + }), + auth: async () => ({ data: {} }), + }, + kilo: { + authStatus: async () => ({ data: { authenticated: true, type: "oauth", organizationId: "org" } }), + }, + config: { + providers: async (input: { directory: string }) => { + directories.push(input.directory) + if (item.error) throw new Error("Defaults unavailable") + return { data: { default: { kilo: item.recommended, anthropic: "unrelated" } } } + }, + }, + } as unknown as Parameters[0] + + const result = await fetchProviderData(client, "/workspace") + expect(result.response.default.kilo).toBe(item.expected) + expect(result.response.default.anthropic).toBe("claude") + expect(result.response.all.find((provider) => provider.id === "anthropic")).toMatchObject({ + metadata: { priority: 1 }, + }) + expect(result.organizationId).toBe("org") + expect(result.ready).toBe(true) + expect(directories).toEqual(["/workspace"]) + }) + } + + it("distinguishes failed auth context from Personal and removes unverified Kilo models", async () => { + const client = { + provider: { + list: async () => ({ + data: { + all: [ + { id: "kilo", models: { "kilo-auto/free": {} } }, + { id: "external", models: { model: {} } }, + ], + connected: ["kilo", "external"], + default: { kilo: "kilo-auto/free", external: "model" }, + }, + }), + auth: async () => ({ data: {} }), + }, + kilo: { + authStatus: async () => { + throw new Error("Context unavailable") + }, + }, + } as unknown as Parameters[0] + + const result = await fetchProviderData(client, "/workspace") + expect(result.ready).toBe(false) + expect(result.organizationId).toBeUndefined() + expect(result.response.all.map((provider) => provider.id)).toEqual(["external"]) + expect(result.response.connected).toEqual(["external"]) + expect(result.response.default).toEqual({ external: "model" }) + }) + + it("retains Personal defaults without fetching organization recommendations", async () => { + let calls = 0 + const client = { + provider: { + list: async () => ({ data: { all: [], connected: [], default: { kilo: "kilo-auto/free" } } }), + auth: async () => ({ data: {} }), + }, + kilo: { authStatus: async () => ({ data: { authenticated: true, type: "oauth" } }) }, + config: { + providers: async () => { + calls++ + return { data: { default: { kilo: "unexpected" } } } + }, + }, + } as unknown as Parameters[0] + + const result = await fetchProviderData(client, "/workspace") + expect(result.ready).toBe(true) + expect(result.organizationId).toBeNull() + expect(calls).toBe(0) + expect(result.response.default).toEqual({ kilo: "kilo-auto/free" }) + }) + it("derives api auth state and strips keys from provider payloads", async () => { const client = { provider: { diff --git a/packages/kilo-vscode/tests/unit/session-model-store.test.ts b/packages/kilo-vscode/tests/unit/session-model-store.test.ts index c83f7616582..e866efc19ca 100644 --- a/packages/kilo-vscode/tests/unit/session-model-store.test.ts +++ b/packages/kilo-vscode/tests/unit/session-model-store.test.ts @@ -259,3 +259,97 @@ describe("per-mode model memory", () => { expect(getSelected(switched, configured, "session-a", "code")).toEqual(gpt) }) }) + +describe("organization model store", () => { + const first = { providerID: "kilo", modelID: "first" } + const recommendation = { providerID: "kilo", modelID: "org-default" } + const organization: ResolveEnv = { + ...env(), + ready: true, + organizationId: "org-a", + providers: { ...providers, kilo: makeProvider("kilo", [first.modelID, recommendation.modelID, KILO_AUTO.modelID]) }, + defaults: { kilo: recommendation.modelID }, + } + + it("ignores implicit mode memory and generic recents across every accessor", () => { + const store: ModelStore = { + ...emptyStore(), + modelSelections: { code: KILO_AUTO, ask: first }, + recentModels: [KILO_AUTO, gpt], + } + const before = structuredClone(store) + expect(getSelected(store, organization, undefined, "code")).toEqual(recommendation) + expect(getSelected(store, organization, "session-a", "code")).toEqual(recommendation) + expect(getSessionModel(store, organization, "session-a", "code")).toEqual(recommendation) + expect(getAgentModel(store, organization, "ask")).toEqual(recommendation) + expect(store).toEqual(before) + }) + + it.each([undefined, "session-a"])("preserves explicit free selections in scope %s", (scope) => { + const store = emptyStore() + const updated = { ...store, ...applyModel(store, "code", KILO_AUTO, scope) } + expect(getSelected(updated, organization, scope, "code")).toEqual(KILO_AUTO) + expect(getSessionModel(updated, organization, "session-a", "code")).toEqual(KILO_AUTO) + if (!scope) expect(getAgentModel(updated, organization, "code")).toEqual(KILO_AUTO) + }) + + it.each([undefined, "session-a"])("restores explicit X through X to Y to X in scope %s without writes", (scope) => { + const store = emptyStore() + const updated = { ...store, ...applyModel(store, "code", KILO_AUTO, scope) } + const before = structuredClone(updated) + const restricted = { ...organization, providers: { kilo: makeProvider("kilo", [recommendation.modelID]) } } + expect(getSelected(updated, env(), scope, "code")).toEqual(KILO_AUTO) + expect(getSelected(updated, restricted, scope, "code")).toEqual(recommendation) + expect(getSessionModel(updated, restricted, "session-a", "code")).toEqual(recommendation) + expect(getSelected(updated, env(), scope, "code")).toEqual(KILO_AUTO) + expect(getSessionModel(updated, env(), "session-a", "code")).toEqual(KILO_AUTO) + if (!scope) { + expect(getAgentModel(updated, restricted, "code")).toEqual(recommendation) + expect(getAgentModel(updated, env(), "code")).toEqual(KILO_AUTO) + } + expect(updated).toEqual(before) + }) + + it("falls through an unavailable session override to a valid explicit manual choice", () => { + const store = { + ...emptyStore(), + modelSelections: { code: gpt }, + userSetAgents: { code: true }, + sessionOverrides: { "session-a": { providerID: "kilo", modelID: "missing" } }, + } + expect(getSelected(store, organization, "session-a", "code")).toEqual(gpt) + expect(getSessionModel(store, organization, "session-a", "code")).toEqual(gpt) + }) + + it("validates history overrides without deleting them when the catalog is empty or pending", () => { + const store = { ...emptyStore(), sessionOverrides: { "session-a": KILO_AUTO } } + const before = structuredClone(store) + for (const pending of [{ ready: false }, { providers: {} }, { organizationId: undefined }]) { + expect(getSelected(store, { ...organization, ...pending }, "session-a", "code")).toBeNull() + expect(getSessionModel(store, { ...organization, ...pending }, "session-a", "code")).toBeNull() + } + expect(getSessionModel(store, organization, "session-a", "code")).toEqual(KILO_AUTO) + expect(store).toEqual(before) + }) + + it("preserves connected external session choices while Kilo refreshes", () => { + const store = { ...emptyStore(), sessionOverrides: { "session-a": gpt } } + expect(getSessionModel(store, { ...organization, ready: false }, "session-a", "code")).toEqual(gpt) + expect(getSessionModel(store, { ...organization, connected: [] }, "session-a", "code")).toEqual(recommendation) + }) + + it("keeps Agent Manager mode configuration precedence without destroying the manual choice", () => { + const store = { ...emptyStore(), modelSelections: { code: KILO_AUTO }, userSetAgents: { code: true } } + const configured = { ...organization, getModeModel: () => first, getGlobalModel: () => gpt } + expect(getAgentModel(store, configured, "code")).toEqual(first) + expect(getSelected(store, configured, undefined, "code")).toEqual(KILO_AUTO) + expect(store.modelSelections.code).toEqual(KILO_AUTO) + expect(getAgentModel(store, organization, "code")).toEqual(KILO_AUTO) + }) + + it("uses valid mode and global config before the recommendation when implicit memory is stale", () => { + const store = { ...emptyStore(), modelSelections: { code: KILO_AUTO } } + expect(getAgentModel(store, { ...organization, getModeModel: () => first }, "code")).toEqual(first) + expect(getSelected(store, { ...organization, getGlobalModel: () => gpt }, undefined, "code")).toEqual(gpt) + }) +}) diff --git a/packages/kilo-vscode/tests/unit/session-provider-activity.test.ts b/packages/kilo-vscode/tests/unit/session-provider-activity.test.ts index a3ad97fc106..4074699634a 100644 --- a/packages/kilo-vscode/tests/unit/session-provider-activity.test.ts +++ b/packages/kilo-vscode/tests/unit/session-provider-activity.test.ts @@ -9,7 +9,7 @@ const webview = path.join(root, "webview-ui") const fixture = path.join(root, "tests/fixtures/session-provider-activity.tsx") describe("SessionProvider activity", () => { - it("covers real session activity lifecycle messages", async () => { + it("covers real session activity and composer send acceptance", async () => { const solid = path.dirname(Bun.resolveSync("solid-js/package.json", webview)) const aliases: Record = { "solid-js": path.join(solid, "dist/solid.js"), diff --git a/packages/kilo-vscode/webview-ui/agent-manager/NewWorktreeDialog.tsx b/packages/kilo-vscode/webview-ui/agent-manager/NewWorktreeDialog.tsx index 6ddf415391c..2ba6f0650f5 100644 --- a/packages/kilo-vscode/webview-ui/agent-manager/NewWorktreeDialog.tsx +++ b/packages/kilo-vscode/webview-ui/agent-manager/NewWorktreeDialog.tsx @@ -51,6 +51,7 @@ import { tracker } from "./telemetry" import { cycleAgent } from "../src/context/session-agent" import type { ModeRouter } from "./mode-router" import { ProjectSelect } from "./ProjectSelect" +import { createDialogModels } from "./new-worktree-models" type VersionCount = 1 | 2 | 3 | 4 const VERSION_OPTIONS: VersionCount[] = [1, 2, 3, 4] @@ -90,16 +91,6 @@ function restoreAgent(value: string | undefined, list: Array<{ name: string }>, return list.some((item) => item.name === value) ? value : base } -function restoreModel(value: Model | undefined, providers: Record, valid: (value: Model) => boolean) { - if (!value) return undefined - if (Object.keys(providers).length === 0) return value - return valid(value) ? value : undefined -} - -function fallback(value: T | undefined, get: () => T): T { - return value === undefined ? get() : value -} - const isMac = typeof navigator !== "undefined" && /Mac|iPhone|iPad/.test(navigator.userAgent) function sanitizeSegment(text: string, maxLength = 50): string { @@ -168,14 +159,17 @@ export const NewWorktreeDialog: Component<{ const saved = readDialogSelections(cached?.advancedDialogSelections) const [versions, setVersions] = createSignal(1) const initialAgent = restoreAgent(saved.agent, session.agents(), session.selectedAgent()) - const initialModel = fallback( - restoreModel(saved.model, provider.providers(), (value) => provider.isModelValid(value)), - () => session.modelForAgent(initialAgent), - ) - const [model, setModel] = createSignal(initialModel) + const [agent, setAgent] = createSignal(initialAgent) + const selection = createDialogModels({ + saved: saved.model, + fallback: () => session.modelForAgent(agent()), + ready: provider.ready, + valid: provider.isModelValid, + variants: (value) => Object.keys(provider.findModel(value)?.variants ?? {}), + }) + const model = selection.model const [compareMode, setCompareMode] = createSignal(false) const [modelAllocations, setModelAllocations] = createSignal(new Map()) - const [agent, setAgent] = createSignal(initialAgent) const [starting, setStarting] = createSignal(false) const [enhancing, setEnhancing] = createSignal(false) const [showAdvanced, setShowAdvanced] = createSignal(false) @@ -207,8 +201,7 @@ export const NewWorktreeDialog: Component<{ const selectAgent = (name: string) => { setAgent(name) - const sel = session.modelForAgent(name) - setModel(sel) + selection.select(undefined) setVariant(undefined) } @@ -329,7 +322,7 @@ export const NewWorktreeDialog: Component<{ ...state, advancedDialogSelections: { agent: agent(), - model: model(), + model: selection.choice(), variant: variant(), sandbox: sandbox(), }, @@ -417,8 +410,7 @@ export const NewWorktreeDialog: Component<{ const canSubmit = () => { if (starting()) return false if (speech.active()) return false - if (compareMode() && totalAllocations(modelAllocations()) === 0) return false - return true + return selection.canSubmit(compareMode() ? modelAllocations() : undefined) } const total = () => (compareMode() ? totalAllocations(modelAllocations()) : versions()) const mode = () => (compareMode() ? "compare_models" : versions() > 1 ? "multiple_versions" : "single") @@ -850,7 +842,7 @@ export const NewWorktreeDialog: Component<{ const current = effectiveVariant() const next = { providerID: pid, modelID: mid } const list = Object.keys(provider.findModel(next)?.variants ?? {}) - setModel(next) + selection.select(next) setVariant(preserveVariant(current, list) ?? DEFAULT_VARIANT) }} onPick={restorePrompt} diff --git a/packages/kilo-vscode/webview-ui/agent-manager/new-worktree-models.ts b/packages/kilo-vscode/webview-ui/agent-manager/new-worktree-models.ts new file mode 100644 index 00000000000..7fafbdd6111 --- /dev/null +++ b/packages/kilo-vscode/webview-ui/agent-manager/new-worktree-models.ts @@ -0,0 +1,33 @@ +import { createMemo, createSignal } from "solid-js" +import type { ModelSelection } from "../src/types/messages" +import { type ModelAllocations, MAX_MULTI_VERSIONS, totalAllocations } from "./multi-model-utils" + +export function createDialogModels(opts: { + saved?: ModelSelection + fallback: () => ModelSelection | null + ready: () => boolean + valid: (model: ModelSelection) => boolean + variants: (model: ModelSelection) => string[] +}) { + const [choice, select] = createSignal(opts.saved) + const valid = (value: ModelSelection) => (value.providerID !== "kilo" || opts.ready()) && opts.valid(value) + const model = createMemo(() => { + const saved = choice() + if (saved && valid(saved)) return saved + const fallback = opts.fallback() + return fallback && valid(fallback) ? fallback : null + }) + const canSubmit = (allocations?: ModelAllocations) => { + if (!allocations) return model() !== null + const total = totalAllocations(allocations) + if (total < 1 || total > MAX_MULTI_VERSIONS) return false + return [...allocations.values()].every( + (entry) => + Number.isInteger(entry.count) && + entry.count > 0 && + valid(entry) && + (entry.variant === undefined || opts.variants(entry).includes(entry.variant)), + ) + } + return { choice, select, model, canSubmit } +} diff --git a/packages/kilo-vscode/webview-ui/src/components/chat/PromptInput.tsx b/packages/kilo-vscode/webview-ui/src/components/chat/PromptInput.tsx index d0ec11875fc..f748e5ceae1 100644 --- a/packages/kilo-vscode/webview-ui/src/components/chat/PromptInput.tsx +++ b/packages/kilo-vscode/webview-ui/src/components/chat/PromptInput.tsx @@ -1341,7 +1341,7 @@ export const PromptInput: Component = (props) => { // Server-side slash command (cmdMatch/matched already computed above) if (matched && !data && !browserData) { const args = draft.slice(cmdMatch![0].length).trim() - session.sendCommand( + const accepted = session.sendCommand( matched.name, args, sel?.providerID, @@ -1356,8 +1356,9 @@ export const PromptInput: Component = (props) => { variant: matched.variant, }, ) + if (!accepted) return } else { - session.sendMessage( + const accepted = session.sendMessage( message, sel?.providerID, sel?.modelID, @@ -1368,6 +1369,7 @@ export const PromptInput: Component = (props) => { origin ?? null, browserData, ) + if (!accepted) return } drafts.delete(key) diff --git a/packages/kilo-vscode/webview-ui/src/context/model-selection.ts b/packages/kilo-vscode/webview-ui/src/context/model-selection.ts index 86d9a255b98..efea49c065d 100644 --- a/packages/kilo-vscode/webview-ui/src/context/model-selection.ts +++ b/packages/kilo-vscode/webview-ui/src/context/model-selection.ts @@ -1,43 +1,38 @@ import type { ModelSelection, Provider } from "../types/messages" import { isModelValid } from "./provider-utils" -function validate( - providers: Record, - connected: string[], - selection: ModelSelection | null | undefined, -): ModelSelection | null { - if (!selection) return null - if (Object.keys(providers).length === 0) return selection - return isModelValid(providers, connected, selection) ? selection : null -} - -function recent( - providers: Record, - connected: string[], - selections: ModelSelection[] | undefined, -): ModelSelection | null { - for (const item of selections ?? []) { - const selection = validate(providers, connected, item) - if (selection) return selection - } - return null -} - export function resolveModelSelection(input: { providers: Record connected: string[] + ready?: boolean + organizationId?: string | null + defaults?: Record + session?: ModelSelection | null override?: ModelSelection | null mode?: ModelSelection | null global?: ModelSelection | null recent?: ModelSelection[] fallback?: ModelSelection | null }): ModelSelection | null { - return ( - validate(input.providers, input.connected, input.override) ?? - validate(input.providers, input.connected, input.mode) ?? - validate(input.providers, input.connected, input.global) ?? - recent(input.providers, input.connected, input.recent) ?? - input.fallback ?? - null - ) + const pending = input.ready === false || (input.ready !== undefined && input.organizationId === undefined) + const validate = (selection: ModelSelection | null | undefined) => { + if (!selection || (pending && selection.providerID === "kilo")) return null + return isModelValid(input.providers, input.connected, selection) ? selection : null + } + const preference = + validate(input.session) ?? validate(input.override) ?? validate(input.mode) ?? validate(input.global) + if (preference) return preference + if (pending) return null + if (input.organizationId) { + const recommendation = input.defaults?.kilo + const selection = recommendation ? validate({ providerID: "kilo", modelID: recommendation }) : null + if (selection) return selection + const first = Object.keys(input.providers.kilo?.models ?? {}).at(0) + return first ? validate({ providerID: "kilo", modelID: first }) : null + } + for (const selection of input.recent ?? []) { + const model = validate(selection) + if (model) return model + } + return validate(input.fallback) } diff --git a/packages/kilo-vscode/webview-ui/src/context/provider.tsx b/packages/kilo-vscode/webview-ui/src/context/provider.tsx index 4b173b8baa0..fc8860aa1e3 100644 --- a/packages/kilo-vscode/webview-ui/src/context/provider.tsx +++ b/packages/kilo-vscode/webview-ui/src/context/provider.tsx @@ -4,7 +4,7 @@ * Selection is now per-session — see session.tsx. */ -import { createContext, useContext, createSignal, createMemo, onCleanup } from "solid-js" +import { batch, createContext, useContext, createSignal, createMemo, onCleanup } from "solid-js" import type { ParentComponent, Accessor } from "solid-js" import { useVSCode } from "./vscode" import type { Provider, ProviderModel, ModelSelection, ExtensionMessage, ProviderAuthState } from "../types/messages" @@ -18,6 +18,8 @@ interface ProviderContextValue { providers: Accessor> connected: Accessor defaults: Accessor> + organizationId: Accessor + ready: Accessor defaultSelection: Accessor models: Accessor findModel: (selection: ModelSelection | null) => EnrichedModel | undefined @@ -34,6 +36,8 @@ export const ProviderProvider: ParentComponent = (props) => { const [providers, setProviders] = createSignal>({}) const [connected, setConnected] = createSignal([]) const [defaults, setDefaults] = createSignal>({}) + const [organizationId, setOrganizationId] = createSignal() + const [ready, setReady] = createSignal(false) const [defaultSelection, setDefaultSelection] = createSignal(KILO_AUTO) const [authMethods, setAuthMethods] = createSignal>({}) const [authStates, setAuthStates] = createSignal>({}) @@ -51,16 +55,36 @@ export const ProviderProvider: ParentComponent = (props) => { // Register handler immediately (not in onMount) so we never miss // a providersLoaded message that arrives before the DOM mount. const unsubscribe = vscode.onMessage((message: ExtensionMessage) => { - if (message.type !== "providersLoaded") { + if (message.type === "providersLoading") { + batch(() => { + setReady(false) + setOrganizationId(undefined) + setProviders((prev) => { + const next = { ...prev } + delete next.kilo + return next + }) + setDefaults((prev) => { + const next = { ...prev } + delete next.kilo + return next + }) + setConnected((prev) => prev.filter((id) => id !== "kilo")) + }) return } + if (message.type !== "providersLoaded") return - setProviders(message.providers) - setConnected(message.connected) - setDefaults(message.defaults) - setDefaultSelection(message.defaultSelection) - setAuthMethods(message.authMethods) - setAuthStates(message.authStates) + batch(() => { + setProviders(message.providers) + setConnected(message.connected) + setDefaults(message.defaults) + setOrganizationId(message.ready === false ? undefined : (message.organizationId ?? null)) + setReady(message.ready ?? true) + setDefaultSelection(message.defaultSelection) + setAuthMethods(message.authMethods) + setAuthStates(message.authStates) + }) }) onCleanup(unsubscribe) @@ -93,6 +117,8 @@ export const ProviderProvider: ParentComponent = (props) => { providers, connected, defaults, + organizationId, + ready, defaultSelection, models, findModel, diff --git a/packages/kilo-vscode/webview-ui/src/context/session-model-store.ts b/packages/kilo-vscode/webview-ui/src/context/session-model-store.ts index 704e03fc648..105d60fa373 100644 --- a/packages/kilo-vscode/webview-ui/src/context/session-model-store.ts +++ b/packages/kilo-vscode/webview-ui/src/context/session-model-store.ts @@ -16,11 +16,15 @@ export interface ModelStore { /** sessionID -> agent name */ agentSelections: Record recentModels: ModelSelection[] + userSetAgents?: Record } export interface ResolveEnv { providers: Record connected: string[] + ready?: boolean + organizationId?: string | null + defaults?: Record fallback: ModelSelection | null getModeModel: (agentName: string) => ModelSelection | null getGlobalModel: () => ModelSelection | null @@ -31,10 +35,15 @@ function resolveModel( agentName: string, override?: ModelSelection | null, recents?: ModelSelection[], + session?: ModelSelection, ): ModelSelection | null { return resolveModelSelection({ providers: env.providers, connected: env.connected, + ready: env.ready, + organizationId: env.organizationId, + defaults: env.defaults, + session, override, mode: env.getModeModel(agentName), global: env.getGlobalModel(), @@ -54,10 +63,8 @@ export function getSessionModel( sessionID: string, defaultAgent: string, ): ModelSelection | null { - const override = store.sessionOverrides[sessionID] - if (override) return override const agentName = store.agentSelections[sessionID] ?? defaultAgent - return resolveModel(env, agentName, store.modelSelections[agentName], store.recentModels) + return getSelected(store, env, sessionID, agentName) } /** @@ -71,11 +78,14 @@ export function getSelected( sessionID: string | undefined, agentName: string, ): ModelSelection | null { - if (sessionID) { - const session = store.sessionOverrides[sessionID] - if (session) return session - } - return resolveModel(env, agentName, store.modelSelections[agentName], store.recentModels) + const override = env.organizationId && !store.userSetAgents?.[agentName] ? null : store.modelSelections[agentName] + return resolveModel( + env, + agentName, + override, + store.recentModels, + sessionID ? store.sessionOverrides[sessionID] : undefined, + ) } /** Returns the effective model for a mode outside a session scope. */ @@ -83,15 +93,19 @@ export function getAgentModel( store: ModelStore, env: ResolveEnv, agentName: string, - userSet = false, + userSet = store.userSetAgents?.[agentName] === true, ): ModelSelection | null { - const override = env.getModeModel(agentName) && userSet ? null : store.modelSelections[agentName] + const override = + (env.getModeModel(agentName) && userSet) || (env.organizationId && !userSet) + ? null + : store.modelSelections[agentName] return resolveModel(env, agentName, override, store.recentModels) } export interface ApplyResult { modelSelections: Record sessionOverrides: Record + userSetAgents: Record } /** @@ -116,5 +130,6 @@ export function applyModel( sessionOverrides[sessionID] = selection } - return { modelSelections, sessionOverrides } + const userSetAgents = sessionID ? { ...store.userSetAgents } : { ...store.userSetAgents, [agentName]: true } + return { modelSelections, sessionOverrides, userSetAgents } } diff --git a/packages/kilo-vscode/webview-ui/src/context/session-types.ts b/packages/kilo-vscode/webview-ui/src/context/session-types.ts index 3cbd21d7c39..4e7c3acb10b 100644 --- a/packages/kilo-vscode/webview-ui/src/context/session-types.ts +++ b/packages/kilo-vscode/webview-ui/src/context/session-types.ts @@ -173,7 +173,7 @@ export interface SessionContextValue { review?: ReviewMessageData, origin?: string | null, browserFeedback?: BrowserFeedbackData, - ) => void + ) => boolean sendCommand: ( command: string, args: string, @@ -184,7 +184,7 @@ export interface SessionContextValue { context?: string, origin?: string | null, overrides?: { agent?: string; model?: string; variant?: string }, - ) => void + ) => boolean abort: () => void compact: () => void respondToPermission: ( diff --git a/packages/kilo-vscode/webview-ui/src/context/session.tsx b/packages/kilo-vscode/webview-ui/src/context/session.tsx index 59f775af01c..f12df7749fb 100644 --- a/packages/kilo-vscode/webview-ui/src/context/session.tsx +++ b/packages/kilo-vscode/webview-ui/src/context/session.tsx @@ -78,7 +78,7 @@ import { } from "./session-utils" import { Identifier } from "../utils/id" import { resolveModelSelection } from "./model-selection" -import { getAgentModel } from "./session-model-store" +import { getAgentModel, getSelected, getSessionModel } from "./session-model-store" import { resolveMessagePrefs } from "./session-preferences" import { errorIDs, preserveSessionErrors, withoutResolvedSessionErrors } from "./session-errors" import { PartStash } from "./part-stash" @@ -420,15 +420,35 @@ export const SessionProvider: ParentComponent = (props) => { return parseModelString(config().model) } - function resolveModel(agentName: string, override?: ModelSelection | null): ModelSelection | null { - return resolveModelSelection({ + function environment() { + return { providers: provider.providers(), connected: provider.connected(), - override, + ready: provider.ready(), + organizationId: provider.organizationId(), + defaults: provider.defaults(), + getModeModel, + getGlobalModel, + fallback: KILO_AUTO, + } + } + + function preferences() { + return { + modelSelections: store.modelSelections, + sessionOverrides: store.sessionOverrides, + agentSelections: store.agentSelections, + recentModels: store.recentModels, + userSetAgents: userSetAgents(), + } + } + + function resolveModel(agentName: string): ModelSelection | null { + return resolveModelSelection({ + ...environment(), mode: getModeModel(agentName), global: getGlobalModel(), recent: store.recentModels, - fallback: KILO_AUTO, }) } @@ -441,23 +461,14 @@ export const SessionProvider: ParentComponent = (props) => { setStore("modelSelections", agentName, sel) }) - const currentSelected = createMemo(() => { - const sid = currentSessionID() - if (sid) { - const session = store.sessionOverrides[sid] - if (session) return session - } - const agentName = selectedAgentName() - return resolveModel(agentName, store.modelSelections[agentName]) - }) + const currentSelected = createMemo(() => + getSelected(preferences(), environment(), currentSessionID(), selectedAgentName()), + ) // Precedence: scoped override > per-agent global/default > config/default. function selected(sessionID?: string): ModelSelection | null { if (!sessionID) return currentSelected() - const session = store.sessionOverrides[sessionID] - if (session) return session - const agentName = agentForScope(sessionID) - return resolveModel(agentName, store.modelSelections[agentName]) + return getSessionModel(preferences(), environment(), sessionID, defaultAgent()) } function pushRecent(selection: ModelSelection) { @@ -598,23 +609,7 @@ export const SessionProvider: ParentComponent = (props) => { } function modelForAgent(agentName: string): ModelSelection | null { - return getAgentModel( - { - modelSelections: store.modelSelections, - sessionOverrides: store.sessionOverrides, - agentSelections: store.agentSelections, - recentModels: store.recentModels, - }, - { - providers: provider.providers(), - connected: provider.connected(), - getModeModel, - getGlobalModel, - fallback: KILO_AUTO, - }, - agentName, - userSetAgents()[agentName] === true, - ) + return getAgentModel(preferences(), environment(), agentName) } // Handle agentsLoaded immediately (not in onMount) so we never miss @@ -732,12 +727,14 @@ export const SessionProvider: ParentComponent = (props) => { // Uses replace semantics so an empty payload clears old entries. const unsubSelections = vscode.onMessage((message: ExtensionMessage) => { if (message.type !== "modelSelectionsLoaded") return - setStore("modelSelections", reconcile(message.selections)) - const flags: Record = {} - for (const name of Object.keys(message.selections)) { - flags[name] = true - } - setUserSetAgents(flags) + batch(() => { + setStore("modelSelections", reconcile(message.selections)) + const flags: Record = {} + for (const name of Object.keys(message.selections)) { + flags[name] = true + } + setUserSetAgents(flags) + }) }) vscode.postMessage({ type: "requestModelSelections" }) onCleanup(unsubSelections) @@ -764,64 +761,6 @@ export const SessionProvider: ParentComponent = (props) => { vscode.postMessage({ type: "requestFavorites" }) onCleanup(unsubFavorites) - // Clear model overrides that match the previous config model (not intentional user overrides). - // When config.model changes, old overrides that were just default values should be cleared - // so sessions fall through to resolveModel() and pick up the new config model. - const [lastConfigModel, setLastConfigModel] = createSignal(getGlobalModel()) - createEffect(() => { - const newConfigModel = getGlobalModel() - // Use untrack to read previous value without making this effect re-trigger on its own updates - const oldConfigModel = untrack(() => lastConfigModel()) - if (oldConfigModel) { - // Also clear when newConfigModel is null (user removed model from config) - if (newConfigModel) { - const modelChanged = - oldConfigModel.providerID !== newConfigModel.providerID || oldConfigModel.modelID !== newConfigModel.modelID - if (modelChanged) { - // Clear overrides that match the OLD config model - these were likely defaults, - // not intentional user overrides. Overrides that differ from both old and new - // config are preserved (intentional user selections). - setStore( - "sessionOverrides", - produce((overrides) => { - for (const sid of Object.keys(overrides)) { - const override = overrides[sid] - if ( - override && - override.providerID === oldConfigModel.providerID && - override.modelID === oldConfigModel.modelID - ) { - delete overrides[sid] - } - } - }), - ) - } - } else { - // newConfigModel is null - clear all overrides that matched the old config model - // since the config no longer specifies a model. This ensures sessions fall through - // to provider defaults rather than using a stale removed model. - setStore( - "sessionOverrides", - produce((overrides) => { - for (const sid of Object.keys(overrides)) { - const override = overrides[sid] - if ( - override && - override.providerID === oldConfigModel.providerID && - override.modelID === oldConfigModel.modelID - ) { - delete overrides[sid] - } - } - }), - ) - } - } - // Update the tracked config model - setLastConfigModel(newConfigModel) - }) - function handleError(message: Extract) { if (!message.sessionID || message.sessionID === currentSessionID()) setLoading(false) if (message.sessionID) patchPage(message.sessionID, { loadingInitial: false, loadingOlder: false }) @@ -2121,6 +2060,14 @@ export const SessionProvider: ParentComponent = (props) => { queueMicrotask(() => window.dispatchEvent(new CustomEvent("resumeAutoScroll"))) } + function available(selection: ModelSelection | null): selection is ModelSelection { + const resolved = resolveModelSelection({ ...environment(), override: selection }) + if (selection && resolved?.providerID === selection.providerID && resolved.modelID === selection.modelID) + return true + showToast({ variant: "error", title: language.t("dialog.model.select.title") }) + return false + } + function sendMessage( text: string, providerID?: string, @@ -2131,17 +2078,18 @@ export const SessionProvider: ParentComponent = (props) => { review?: ReviewMessageData, origin?: string | null, browserFeedback?: BrowserFeedbackData, - ) { + ): boolean { if (!server.isConnected()) { console.warn("[Kilo New] Cannot send message: not connected") - return + return false } const messageID = Identifier.ascending("message") const sid = origin === undefined ? currentSessionID() : (origin ?? undefined) - const selection = providerID && modelID ? { providerID, modelID } : selected(sid) - recordModelUsage(selection?.providerID, selection?.modelID) + const selection = providerID && modelID ? { providerID, modelID } : selected(draftID ?? sid) + if (!available(selection)) return false + recordModelUsage(selection.providerID, selection.modelID) const preview = sid?.startsWith("cloud:") ? sid.slice("cloud:".length) : origin === undefined @@ -2155,15 +2103,15 @@ export const SessionProvider: ParentComponent = (props) => { cloudSessionId: preview, text, messageID, - providerID, - modelID, + providerID: selection.providerID, + modelID: selection.modelID, agent, variant: variants.request(scope), files, review, browserFeedback, }) - return + return true } const suggestion = scopedSuggestions(sid)[0] @@ -2192,8 +2140,8 @@ export const SessionProvider: ParentComponent = (props) => { messageID, sessionID: sid, draftID: effectiveDraftID, - providerID, - modelID, + providerID: selection.providerID, + modelID: selection.modelID, agent, variant: variants.request(scope), files, @@ -2201,6 +2149,7 @@ export const SessionProvider: ParentComponent = (props) => { browserFeedback, agentManagerContext: context, }) + return true } function sendCommand( @@ -2213,10 +2162,10 @@ export const SessionProvider: ParentComponent = (props) => { context?: string, origin?: string | null, overrides?: { agent?: string; model?: string; variant?: string }, - ) { + ): boolean { if (!server.isConnected()) { console.warn("[Kilo New] Cannot send command: not connected") - return + return false } const sid = origin === undefined ? currentSessionID() : (origin ?? undefined) @@ -2229,17 +2178,17 @@ export const SessionProvider: ParentComponent = (props) => { } if (overrides?.model) { const parsed = parseModelString(overrides.model) - if (parsed) { - selectModel(parsed.providerID, parsed.modelID, scope) - } + if (!available(parsed)) return false + selectModel(parsed.providerID, parsed.modelID, scope) } if (overrides?.variant) { selectVariant(overrides.variant, scope) } - const effectiveSelection = selected(scope) - const effectiveProvider = effectiveSelection?.providerID ?? providerID - const effectiveModel = effectiveSelection?.modelID ?? modelID + const effectiveSelection = selected(scope) ?? (providerID && modelID ? { providerID, modelID } : null) + if (!available(effectiveSelection)) return false + const effectiveProvider = effectiveSelection.providerID + const effectiveModel = effectiveSelection.modelID recordModelUsage(effectiveProvider, effectiveModel) // Cloud previews need import-then-command; post importAndSend with command metadata @@ -2263,7 +2212,7 @@ export const SessionProvider: ParentComponent = (props) => { command, commandArgs: args, }) - return + return true } const messageID = Identifier.ascending("message") @@ -2298,6 +2247,7 @@ export const SessionProvider: ParentComponent = (props) => { files, agentManagerContext: context, }) + return true } const resumable = () => @@ -2356,11 +2306,12 @@ export const SessionProvider: ParentComponent = (props) => { } const sel = selected() + if (!available(sel)) return vscode.postMessage({ type: "compact", sessionID, - providerID: sel?.providerID, - modelID: sel?.modelID, + providerID: sel.providerID, + modelID: sel.modelID, }) } diff --git a/packages/kilo-vscode/webview-ui/src/stories/StoryProviders.tsx b/packages/kilo-vscode/webview-ui/src/stories/StoryProviders.tsx index aec116a739c..ab8f430c96c 100644 --- a/packages/kilo-vscode/webview-ui/src/stories/StoryProviders.tsx +++ b/packages/kilo-vscode/webview-ui/src/stories/StoryProviders.tsx @@ -110,6 +110,8 @@ const MockProviderProvider: ParentComponent<{ kiloAuth?: boolean; training?: boo providers: () => MOCK_PROVIDERS as any, connected: () => ["kilo"], defaults: () => ({}), + organizationId: () => null, + ready: () => true, defaultSelection: () => ({ providerID: "kilo", modelID: "anthropic/claude-sonnet-4-6" }), models, findModel: (sel: any) => _findModel(models(), sel), @@ -264,8 +266,8 @@ export function mockSessionValue(overrides?: { currentVariant: () => undefined, variantForAgent: () => undefined, selectVariant: noop, - sendMessage: noop, - sendCommand: noop, + sendMessage: () => true, + sendCommand: () => true, abort: noop, compact: noop, respondToPermission: noop, diff --git a/packages/kilo-vscode/webview-ui/src/stories/history.stories.tsx b/packages/kilo-vscode/webview-ui/src/stories/history.stories.tsx index 3f54314e6b2..1243cc64b8a 100644 --- a/packages/kilo-vscode/webview-ui/src/stories/history.stories.tsx +++ b/packages/kilo-vscode/webview-ui/src/stories/history.stories.tsx @@ -87,7 +87,7 @@ const WithSessions: ParentComponent<{ sessions?: typeof mockSessions }> = (props variantList: () => [], currentVariant: () => undefined, selectVariant: noop, - sendMessage: noop, + sendMessage: () => true, abort: noop, compact: noop, respondToPermission: noop, diff --git a/packages/kilo-vscode/webview-ui/src/types/messages/extension-messages.ts b/packages/kilo-vscode/webview-ui/src/types/messages/extension-messages.ts index cc28f510509..24da4fc5ecb 100644 --- a/packages/kilo-vscode/webview-ui/src/types/messages/extension-messages.ts +++ b/packages/kilo-vscode/webview-ui/src/types/messages/extension-messages.ts @@ -510,6 +510,8 @@ export interface ProvidersLoadedMessage { providers: Record connected: string[] defaults: Record + organizationId?: string | null + ready?: boolean defaultSelection: ModelSelection authMethods: Record authStates: Record @@ -1544,6 +1546,7 @@ export type ExtensionMessage = | ImageModelsLoadedMessage | SpeechToTextModelsLoadedMessage | ProvidersLoadedMessage + | { type: "providersLoading" } | AgentsLoadedMessage | SkillsLoadedMessage | CommandsLoadedMessage diff --git a/packages/opencode/src/kilocode/provider/catalog.ts b/packages/opencode/src/kilocode/provider/catalog.ts new file mode 100644 index 00000000000..2f67f25276c --- /dev/null +++ b/packages/opencode/src/kilocode/provider/catalog.ts @@ -0,0 +1,31 @@ +import type { Auth } from "@/auth" +import type { Provider } from "@/provider/provider" +import { fetchDefaultModel } from "@kilocode/kilo-gateway" + +export function organization( + options: { kilocodeOrganizationId?: string; baseURL?: string } | undefined, + info: Auth.Info | undefined, +) { + return ( + options?.kilocodeOrganizationId ?? + URL.parse(options?.baseURL ?? "") + ?.pathname.match(/\/api\/organizations\/([^/]+)/) + ?.at(1) ?? + (info?.type === "oauth" ? info.accountId : undefined) ?? + process.env.KILO_ORG_ID + ) +} + +export async function recommend( + models: Provider.Info["models"], + options: { kilocodeOrganizationId?: string; baseURL?: string; apiKey?: string } | undefined, + info: Auth.Info | undefined, +) { + const org = organization(options, info) + const stored = info?.type === "oauth" ? info.access : info?.key + const token = org ? process.env.KILO_API_KEY || stored || options?.apiKey : stored + const fallback = org ? Object.keys(models).at(0) : undefined + if (org && !fallback) return undefined + const model = await fetchDefaultModel(token, org, fallback) + return Object.hasOwn(models, model) ? model : fallback +} diff --git a/packages/opencode/src/kilocode/server/httpapi/groups/kilo-gateway.ts b/packages/opencode/src/kilocode/server/httpapi/groups/kilo-gateway.ts index f77ad9c077a..35322bfbe8a 100644 --- a/packages/opencode/src/kilocode/server/httpapi/groups/kilo-gateway.ts +++ b/packages/opencode/src/kilocode/server/httpapi/groups/kilo-gateway.ts @@ -46,6 +46,7 @@ export const ProfileWithBalance = Schema.Struct({ export const AuthStatus = Schema.Struct({ authenticated: Schema.Boolean, type: Schema.optional(Schema.Literals(["api", "oauth"])), + organizationId: Schema.optional(Schema.String), }) export const NotificationAction = Schema.Struct({ diff --git a/packages/opencode/src/kilocode/server/httpapi/handlers/kilo-gateway.ts b/packages/opencode/src/kilocode/server/httpapi/handlers/kilo-gateway.ts index 9ea67eb455e..2042e1fb5f6 100644 --- a/packages/opencode/src/kilocode/server/httpapi/handlers/kilo-gateway.ts +++ b/packages/opencode/src/kilocode/server/httpapi/handlers/kilo-gateway.ts @@ -34,6 +34,8 @@ import { Flag } from "@opencode-ai/core/flag/flag" import { Database } from "@opencode-ai/core/database/database" import { KilocodeConfig } from "@/kilocode/config/config" import { Auth } from "@/auth" +import { Config } from "@/config/config" +import { organization as catalogOrganization } from "@/kilocode/provider/catalog" import { EventV2Bridge } from "@/event-v2-bridge" import { Storage } from "@/storage/storage" import { Instance } from "@/kilocode/instance" @@ -56,6 +58,7 @@ function logError(route: string, err: unknown) { export const kiloGatewayHandlers = HttpApiBuilder.group(InstanceHttpApi, "kilo", (handlers) => Effect.gen(function* () { const auth = yield* Auth.Service + const config = yield* Config.Service const store = yield* InstanceStore.Service const cache = yield* ModelCache.Service const events = yield* EventV2Bridge.Service @@ -81,9 +84,11 @@ export const kiloGatewayHandlers = HttpApiBuilder.group(InstanceHttpApi, "kilo", const authStatus = Effect.fn("KiloGatewayHttpApi.authStatus")(function* () { const info = yield* auth.get("kilo").pipe(Effect.mapError(() => new HttpApiError.BadRequest({}))) + const cfg = yield* config.get() + const organizationId = catalogOrganization(cfg.provider?.kilo?.options, info) const type = getToken(info) && (info?.type === "api" || info?.type === "oauth") ? info.type : undefined - if (!type) return { authenticated: false } - return { authenticated: true, type } + if (!type) return { authenticated: false, organizationId } + return { authenticated: true, type, organizationId } }) const proxyAuth = Effect.fn("KiloGatewayHttpApi.proxyAuth")(function* () { diff --git a/packages/opencode/src/provider/model-cache.ts b/packages/opencode/src/provider/model-cache.ts index 40b266edc99..b5b755e9c6b 100644 --- a/packages/opencode/src/provider/model-cache.ts +++ b/packages/opencode/src/provider/model-cache.ts @@ -4,6 +4,7 @@ import { Context, Deferred, Duration, Effect, Exit, Layer, Schema, Scope } from import { FetchHttpClient, HttpClient, HttpClientRequest, HttpClientResponse } from "effect/unstable/http" import { Config } from "../config/config" import { Auth } from "../auth" +import { organization } from "@/kilocode/provider/catalog" import type { Provider } from "@opencode-ai/core/models-dev" import * as Log from "@opencode-ai/core/util/log" import { LayerNode } from "@opencode-ai/core/effect/layer-node" @@ -126,17 +127,13 @@ export const layer: Layer.Layer< if (providerID === "kilo") { const item = config.provider?.[providerID] if (item?.options?.apiKey) options.kilocodeToken = item.options.apiKey - if (item?.options?.kilocodeOrganizationId) options.kilocodeOrganizationId = item.options.kilocodeOrganizationId const info = yield* auth.get(providerID) + options.kilocodeOrganizationId = organization(item?.options, info) if (info?.type === "api") options.kilocodeToken = info.key - if (info?.type === "oauth") { - options.kilocodeToken = info.access - if (info.accountId) options.kilocodeOrganizationId = info.accountId - } + if (info?.type === "oauth") options.kilocodeToken = info.access if (process.env.KILO_API_KEY) options.kilocodeToken = process.env.KILO_API_KEY - if (process.env.KILO_ORG_ID) options.kilocodeOrganizationId = process.env.KILO_ORG_ID log.debug("auth options resolved", { providerID, hasToken: !!options.kilocodeToken, diff --git a/packages/opencode/src/provider/models.ts b/packages/opencode/src/provider/models.ts index faa8b65f2c1..61b2de420c0 100644 --- a/packages/opencode/src/provider/models.ts +++ b/packages/opencode/src/provider/models.ts @@ -6,6 +6,7 @@ import * as Core from "@opencode-ai/core/models-dev" import { Context, Effect, Layer } from "effect" import { AI_SDK_PROVIDERS, KILO_OPENROUTER_BASE, PROMPTS } from "@kilocode/kilo-gateway" import { overlay } from "@/kilocode/anaconda-desktop/provider" +import { organization } from "@/kilocode/provider/catalog" import { LayerNode } from "@opencode-ai/core/effect/layer-node" import { AppNodeBuilder } from "@opencode-ai/core/effect/app-node-builder" // kilocode_change @@ -77,14 +78,14 @@ export const layer: Layer.Layer Effect.succeed(undefined))) - const org = opts?.kilocodeOrganizationId ?? (info?.type === "oauth" ? info.accountId : undefined) + const org = organization(opts, info) const url = baseURL(opts?.baseURL, org) const fetch = { ...(url ? { baseURL: url } : {}), ...(org ? { kilocodeOrganizationId: org } : {}), } const fetched = yield* cache.fetch("kilo", fetch).pipe(Effect.catch(() => Effect.succeed({}))) - const models = Object.keys(fetched).length > 0 ? fetched : (fallback?.models ?? {}) + const models = org || Object.keys(fetched).length > 0 ? fetched : (fallback?.models ?? {}) providers.kilo = { id: "kilo", name: "Kilo Gateway", @@ -93,7 +94,8 @@ export const layer: Layer.Layer new HttpApiError.Unauthorized({}))) // kilocode_change - const token = info?.type === "oauth" ? info.access : info?.key - const organizationId = info?.type === "oauth" ? info.accountId : undefined - const model = yield* Effect.promise(() => fetchDefaultModel(token, organizationId)) + const model = yield* Effect.promise(() => + recommend(providers[ProviderV2.ID.kilo].models, config.provider?.kilo?.options, info), + ) if (model && providers[ProviderV2.ID.kilo]?.models[model]) defaults[ProviderV2.ID.kilo] = ModelV2.ID.make(model) } // kilocode_change end diff --git a/packages/opencode/src/server/routes/instance/httpapi/handlers/provider.ts b/packages/opencode/src/server/routes/instance/httpapi/handlers/provider.ts index 28b639624ed..ef79747b21c 100644 --- a/packages/opencode/src/server/routes/instance/httpapi/handlers/provider.ts +++ b/packages/opencode/src/server/routes/instance/httpapi/handlers/provider.ts @@ -5,6 +5,8 @@ import { Provider } from "@/provider/provider" import { mapValues, pickBy } from "remeda" // kilocode_change import { ModelCache } from "@/provider/model-cache" // kilocode_change +import { Auth } from "@/auth" // kilocode_change +import { organization } from "@/kilocode/provider/catalog" // kilocode_change import { disposeAllInstancesAfterProviderAuthCallback, invalidatePresence, @@ -45,6 +47,7 @@ export const providerHandlers = HttpApiBuilder.group(InstanceHttpApi, "provider" const provider = yield* Provider.Service const svc = yield* ProviderAuth.Service const cache = yield* ModelCache.Service // kilocode_change + const access = yield* Auth.Service // kilocode_change const list = Effect.fn("ProviderHttpApi.list")(function* () { const config = yield* cfg.get() @@ -57,6 +60,8 @@ export const providerHandlers = HttpApiBuilder.group(InstanceHttpApi, "provider" } const connected = yield* provider.list() // kilocode_change start + const info = yield* access.get("kilo").pipe(Effect.orDie) + if (organization(config.provider?.kilo?.options, info)) delete filtered.kilo const providers = filterPromptTrainingModels( Object.assign( mapValues(filtered, (item) => Provider.fromModelsDevProvider(item)), diff --git a/packages/opencode/test/kilocode/kilo-loader-auth.test.ts b/packages/opencode/test/kilocode/kilo-loader-auth.test.ts index 42edf9fdcad..2991f67be05 100644 --- a/packages/opencode/test/kilocode/kilo-loader-auth.test.ts +++ b/packages/opencode/test/kilocode/kilo-loader-auth.test.ts @@ -9,6 +9,7 @@ import { Effect, Layer } from "effect" import { FetchHttpClient } from "effect/unstable/http" import { kiloCustomLoaders, patchKiloProviderPrivacy } from "../../src/kilocode/provider/provider" import { Auth } from "../../src/auth" +import type { Config } from "../../src/config/config" import { ModelCache } from "../../src/provider/model-cache" import { Provider } from "../../src/provider/provider" import { TestConfig } from "../fixture/config" @@ -17,24 +18,36 @@ import { provideInstance, testInstanceStoreLayer } from "../fixture/fixture" const input = { id: "kilo", + name: "Kilo Gateway", env: ["KILO_API_KEY"], models: { "free-model": { id: "free-model", name: "Free Model", + release_date: "", + attachment: false, + reasoning: false, + temperature: true, + tool_call: true, cost: { input: 0, output: 0 }, limit: { context: 128000, output: 4096 }, }, "paid-model": { id: "paid-model", name: "Paid Model", + release_date: "", + attachment: false, + reasoning: false, + temperature: true, + tool_call: true, cost: { input: 1, output: 2 }, limit: { context: 128000, output: 4096 }, }, }, -} +} satisfies ModelsDev.Provider const seed: Record = { + kilo: input, apertis: { id: "apertis", name: "Apertis", @@ -68,36 +81,39 @@ function load(data?: { auth?: object; config?: object; env?: Record Effect.succeed(options?.config ?? {}) }) + const access = options?.info ? Layer.mock(Auth.Service)({ get: () => Effect.succeed(options.info) }) : auth const models = Layer.succeed( ModelCache.KiloModelsService, ModelCache.KiloModelsService.of({ - fetch: () => - Effect.succeed({ - models: { - "free-model": { - id: "free-model", - name: "Free Model", - cost: { input: 0, output: 0 }, - limit: { context: 128000, output: 4096 }, + fetch: + options?.fetch ?? + (() => + Effect.succeed({ + models: { + "free-model": { + id: "free-model", + name: "Free Model", + cost: { input: 0, output: 0 }, + limit: { context: 128000, output: 4096 }, + }, + "paid-model": { + id: "paid-model", + name: "Paid Model", + cost: { input: 1, output: 2 }, + isFree: false, + mayTrainOnYourPrompts: true, + limit: { context: 128000, output: 4096 }, + }, }, - "paid-model": { - id: "paid-model", - name: "Paid Model", - cost: { input: 1, output: 2 }, - isFree: false, - mayTrainOnYourPrompts: true, - limit: { context: 128000, output: 4096 }, - }, - }, - }), + })), }), ) const cache = Layer.fresh(ModelCache.layer).pipe( Layer.provide(FetchHttpClient.layer), Layer.provide(cfg), - Layer.provide(auth), + Layer.provide(access), Layer.provide(models), ) const core = Layer.succeed( @@ -112,7 +128,7 @@ function layer() { Layer.provide(FetchHttpClient.layer), Layer.provide(files), Layer.provide(cfg), - Layer.provide(auth), + Layer.provide(access), Layer.provide(cache), ) } @@ -149,6 +165,77 @@ it.live("does not infer free status from zero catalog prices", () => }), ) +for (const context of ["config", "oauth", "env", "url"] as const) { + for (const outcome of ["empty", "unauthorized", "network", "throw"] as const) { + it.live(`keeps ${context} Org ${outcome} catalogs unavailable without public fallback or detached refresh`, () => + Effect.gen(function* () { + const env = process.env.KILO_ORG_ID + yield* Effect.acquireRelease( + Effect.sync(() => { + process.env.KILO_ORG_ID = "org-env" + }), + () => + Effect.sync(() => { + if (env === undefined) delete process.env.KILO_ORG_ID + else process.env.KILO_ORG_ID = env + }), + ) + const calls: Parameters[0][] = [] + const config: Config.Info = + context === "config" + ? { provider: { kilo: { options: { kilocodeOrganizationId: "org-config" } } } } + : context === "url" + ? { provider: { kilo: { options: { baseURL: "https://gateway.test/api/organizations/org-url" } } } } + : {} + const info = + context === "env" + ? undefined + : new Auth.Oauth({ + type: "oauth", + access: "test-token", + refresh: "test-refresh", + expires: 0, + accountId: "org-oauth", + }) + const fetch: ModelCache.KiloModels["fetch"] = (options) => + Effect.gen(function* () { + calls.push(options) + if (outcome === "throw") return yield* Effect.fail(new Error("offline")) + return { models: {}, ...(outcome === "empty" ? {} : { error: { kind: outcome } }) } + }) + yield* ModelsDev.Service.use((models) => + Effect.gen(function* () { + expect((yield* models.get()).kilo.models).toEqual({}) + expect((yield* models.get()).kilo.models).toEqual({}) + expect(calls).toHaveLength(outcome === "throw" ? 2 : 1) + expect(calls.at(0)?.kilocodeOrganizationId).toBe(`org-${context}`) + }), + ).pipe(Effect.provide(layer({ config, info, fetch })), provideInstance(process.cwd())) + }), + ) + } +} + +it.live("preserves Personal public snapshot fallback", () => + Effect.gen(function* () { + const env = process.env.KILO_ORG_ID + yield* Effect.acquireRelease( + Effect.sync(() => { + delete process.env.KILO_ORG_ID + }), + () => + Effect.sync(() => { + if (env !== undefined) process.env.KILO_ORG_ID = env + }), + ) + const providers = yield* ModelsDev.Service.use((models) => models.get()).pipe( + Effect.provide(layer({ fetch: () => Effect.succeed({ models: {} }) })), + provideInstance(process.cwd()), + ) + expect(providers.kilo.models).toEqual(input.models) + }), +) + it.effect("enables a paid catalog anonymously without auth", () => Effect.gen(function* () { const result = yield* load() diff --git a/packages/opencode/test/kilocode/model-cache-org.test.ts b/packages/opencode/test/kilocode/model-cache-org.test.ts index cc27b6c9fe0..14ba2158376 100644 --- a/packages/opencode/test/kilocode/model-cache-org.test.ts +++ b/packages/opencode/test/kilocode/model-cache-org.test.ts @@ -3,7 +3,7 @@ // should use the organization-specific endpoint, not the personal endpoint. import { expect } from "bun:test" -import { Effect, Layer, Ref } from "effect" +import { Deferred, Effect, Fiber, Layer, Ref } from "effect" import { FetchHttpClient } from "effect/unstable/http" import * as Log from "@opencode-ai/core/util/log" @@ -48,6 +48,79 @@ function layer(info: Auth.Info | undefined, captured: Ref.Ref + Effect.gen(function* () { + const account = yield* Ref.make(undefined) + const started = yield* Deferred.make() + const wait = yield* Deferred.make() + const calls: Options[] = [] + const auth = Layer.mock(Auth.Service)({ + get: () => + Ref.get(account).pipe( + Effect.map( + (accountId) => + new Auth.Oauth({ + type: "oauth", + access: "test-token", + refresh: "test-refresh", + expires: 0, + accountId, + }), + ), + ), + }) + const models = Layer.succeed( + ModelCache.KiloModelsService, + ModelCache.KiloModelsService.of({ + fetch: (options) => + Effect.gen(function* () { + calls.push(options) + if (calls.length === 2) { + yield* Deferred.succeed(started, undefined) + yield* Deferred.await(wait) + } + const id = options.kilocodeOrganizationId ?? "personal" + return { models: { [id]: { id, name: id, limit: { context: 128000, output: 4096 } } } } + }), + }), + ) + const cache = Layer.fresh(ModelCache.layer).pipe( + Layer.provide(FetchHttpClient.layer), + Layer.provide(TestConfig.layer()), + Layer.provide(auth), + Layer.provide(models), + ) + yield* ModelCache.Service.use((cache) => + Effect.gen(function* () { + expect(Object.keys(yield* cache.fetch("kilo"))).toEqual(["personal"]) + const pending = yield* cache.refresh("kilo").pipe(Effect.forkChild) + yield* Deferred.await(started) + yield* Ref.set(account, "org-a") + yield* cache.clear("kilo") + expect(yield* cache.get("kilo")).toBeUndefined() + expect(Object.keys(yield* cache.fetch("kilo"))).toEqual(["org-a"]) + yield* Deferred.succeed(wait, undefined) + yield* Fiber.join(pending) + expect(Object.keys((yield* cache.get("kilo")) ?? {})).toEqual(["org-a"]) + expect(yield* cache.getFailure("kilo")).toBeUndefined() + yield* Ref.set(account, "org-b") + yield* cache.clear("kilo") + expect(Object.keys(yield* cache.fetch("kilo"))).toEqual(["org-b"]) + yield* Ref.set(account, undefined) + yield* cache.clear("kilo") + expect(Object.keys(yield* cache.fetch("kilo"))).toEqual(["personal"]) + expect(calls.map((options) => options.kilocodeOrganizationId)).toEqual([ + undefined, + undefined, + "org-a", + "org-b", + undefined, + ]) + }), + ).pipe(Effect.provide(cache)) + }), +) + it.live("model fetch uses accountId from OAuth auth as kilocodeOrganizationId", () => Effect.gen(function* () { const captured = yield* Ref.make(undefined) diff --git a/packages/opencode/test/kilocode/server/httpapi-public.test.ts b/packages/opencode/test/kilocode/server/httpapi-public.test.ts index c5ac7218696..3c2c6d4af33 100644 --- a/packages/opencode/test/kilocode/server/httpapi-public.test.ts +++ b/packages/opencode/test/kilocode/server/httpapi-public.test.ts @@ -214,6 +214,7 @@ describe("Kilo PublicApi OpenAPI contract", () => { expect(auth).toEqual({ authenticated: { type: "boolean" }, type: { type: "string", enum: ["api", "oauth"] }, + organizationId: { type: "string" }, }) const sessions = response(KiloGatewayPaths.cloudSessions)?.properties diff --git a/packages/opencode/test/kilocode/server/kilo-gateway-statuses.test.ts b/packages/opencode/test/kilocode/server/kilo-gateway-statuses.test.ts index 60d9f5eff6d..62847277b6d 100644 --- a/packages/opencode/test/kilocode/server/kilo-gateway-statuses.test.ts +++ b/packages/opencode/test/kilocode/server/kilo-gateway-statuses.test.ts @@ -6,6 +6,8 @@ import { Effect, Layer } from "effect" import { HttpClient, HttpClientRequest, HttpRouter } from "effect/unstable/http" import { HttpApi, HttpApiBuilder } from "effect/unstable/httpapi" import { Auth } from "../../../src/auth" +import type { Config } from "../../../src/config/config" +import { TestConfig } from "../../fixture/config" import { KiloGatewayApi, KiloGatewayPaths } from "../../../src/kilocode/server/httpapi/groups/kilo-gateway" import { kiloGatewayHandlers } from "../../../src/kilocode/server/httpapi/handlers/kilo-gateway" import { InstanceStore } from "../../../src/project/instance-store" @@ -23,9 +25,14 @@ import { import { testEffect } from "../../lib/effect" const TestHttpApi = HttpApi.make("opencode-instance").addHttpApi(KiloGatewayApi) +const state: { info: Auth.Info | undefined; config: Config.Info } = { + info: new Auth.Api({ type: "api", key: "test-token" }), + config: {}, +} const auth = Layer.mock(Auth.Service)({ - get: () => Effect.succeed(new Auth.Api({ type: "api", key: "test-token" })), + get: () => Effect.sync(() => state.info), }) +const config = TestConfig.layer({ get: () => Effect.sync(() => state.config) }) const store = Layer.mock(InstanceStore.Service)({}) const cache = Layer.mock(ModelCache.Service)({}) const session = Layer.mock(Session.Service)({}) @@ -53,6 +60,7 @@ const layer = HttpRouter.serve( passthroughInstanceContext, testWorkspaceRouting, auth, + config, store, cache, session, @@ -106,6 +114,51 @@ describe("Kilo gateway HttpApi statuses", () => { }), ) + for (const context of ["config", "oauth", "env", "url", "personal", "anonymous"] as const) { + it.live(`reports ${context} organization context locally without secrets`, () => + Effect.gen(function* () { + const previous = { ...state } + const env = process.env.KILO_ORG_ID + yield* Effect.acquireRelease( + Effect.sync(() => { + state.config = + context === "config" + ? { provider: { kilo: { options: { kilocodeOrganizationId: "org-config" } } } } + : context === "url" + ? { provider: { kilo: { options: { baseURL: "https://gateway.test/api/organizations/org-url" } } } } + : {} + state.info = + context === "anonymous" + ? undefined + : new Auth.Oauth({ + type: "oauth", + access: "test-token", + refresh: "private-refresh", + expires: Date.now() + 3600000, + ...(["config", "oauth", "url"].includes(context) ? { accountId: "org-oauth" } : {}), + }) + if (context === "personal") delete process.env.KILO_ORG_ID + else process.env.KILO_ORG_ID = "org-env" + }), + () => + Effect.sync(() => { + Object.assign(state, previous) + if (env === undefined) delete process.env.KILO_ORG_ID + else process.env.KILO_ORG_ID = env + }), + ) + yield* stub(() => Promise.reject(new Error("unexpected Gateway request"))) + const response = yield* HttpClient.get(KiloGatewayPaths.authStatus) + expect(response.status).toBe(200) + expect(yield* response.json).toEqual({ + authenticated: context !== "anonymous", + ...(context !== "anonymous" ? { type: "oauth" } : {}), + ...(context !== "personal" ? { organizationId: `org-${context === "anonymous" ? "env" : context}` } : {}), + }) + }), + ) + } + it.live("preserves cloud session list rate limits", () => Effect.gen(function* () { yield* stub(() => new Response("rate limited", { status: 429 })) diff --git a/packages/opencode/test/kilocode/server/prompt-training-model-filter.test.ts b/packages/opencode/test/kilocode/server/prompt-training-model-filter.test.ts index f0728844710..f31992c1aa9 100644 --- a/packages/opencode/test/kilocode/server/prompt-training-model-filter.test.ts +++ b/packages/opencode/test/kilocode/server/prompt-training-model-filter.test.ts @@ -1,10 +1,14 @@ import { afterEach, expect } from "bun:test" import { Effect } from "effect" +import { AppNodeBuilder } from "@opencode-ai/core/effect/app-node-builder" +import { ModelCache } from "../../../src/provider/model-cache" import { Server } from "../../../src/server/server" import * as Log from "@opencode-ai/core/util/log" import { disposeAllInstances, tmpdir } from "../../fixture/fixture" import { resetDatabase } from "../../fixture/db" -import { it } from "../../lib/effect" +import { testEffectShared } from "../../lib/effect" + +const it = testEffectShared(AppNodeBuilder.build(ModelCache.node)) void Log.init({ print: false }) @@ -31,6 +35,143 @@ const response = { ], } +for (const scenario of [ + "valid", + "missing", + "empty-default", + "disallowed", + "default-error", + "empty", + "error", + "unauthorized", + "filtered", +] as const) { + it.live(`keeps Org catalogs and recommendations safe: ${scenario}`, () => + Effect.gen(function* () { + const cache = yield* ModelCache.Service + yield* cache.clear("kilo") + const env = { + KILO_AUTH_CONTENT: process.env.KILO_AUTH_CONTENT, + KILO_API_KEY: process.env.KILO_API_KEY, + KILO_ORG_ID: process.env.KILO_ORG_ID, + } + yield* Effect.acquireRelease( + Effect.sync(() => { + process.env.KILO_AUTH_CONTENT = JSON.stringify({ + kilo: { + type: "oauth", + access: "test-token", + refresh: "test-refresh", + expires: 0, + accountId: "org-oauth", + }, + }) + delete process.env.KILO_API_KEY + process.env.KILO_ORG_ID = "org-env" + }), + () => + Effect.sync(() => { + for (const [key, value] of Object.entries(env)) { + if (value === undefined) delete process.env[key] + else process.env[key] = value + } + }), + ) + const paths: string[] = [] + const original = globalThis.fetch + let active = true + yield* Effect.acquireRelease( + Effect.sync(() => { + globalThis.fetch = Object.assign( + async (input: RequestInfo | URL, init?: RequestInit) => { + if (!active) return original(input, init) + const url = new URL(typeof input === "string" ? input : input instanceof URL ? input.href : input.url) + if (url.pathname.endsWith("/modes")) return new Response(null, { status: 404 }) + if (!url.pathname.endsWith("/models") && !url.pathname.endsWith("/defaults")) return original(input, init) + paths.push(url.pathname) + if (url.pathname.endsWith("/defaults")) { + if (scenario === "default-error") return new Response(null, { status: 500 }) + return Response.json({ + defaultModel: + scenario === "valid" + ? "test/z-last" + : scenario === "disallowed" + ? "test/training" + : scenario === "empty-default" + ? "" + : undefined, + }) + } + if (url.pathname === "/api/organizations/org-config/models") { + if (scenario === "unauthorized") return new Response(null, { status: 401 }) + if (scenario === "error") return new Response(null, { status: 500 }) + if (scenario === "empty") return Response.json({ data: [] }) + return Response.json({ + data: [ + ...response.data, + { ...response.data.at(1), id: "test/z-last", name: "Last", preferredIndex: 0 }, + ], + }) + } + return Response.json({ data: [{ ...response.data.at(1), id: "public/leak" }] }) + }, + { preconnect: original.preconnect }, + ) + }), + () => + Effect.sync(() => { + active = false + globalThis.fetch = original + }), + ) + const tmp = yield* Effect.acquireRelease( + Effect.promise(() => + tmpdir({ + config: { + formatter: false, + lsp: false, + enabled_providers: ["kilo", "external"], + hide_prompt_training_models: true, + provider: { + kilo: { + options: { kilocodeOrganizationId: "org-config" }, + ...(scenario === "filtered" ? { whitelist: ["test/training"] } : {}), + }, + external: { + npm: "@ai-sdk/openai-compatible", + options: { apiKey: "external-test-key" }, + models: { independent: { name: "Independent", limit: { context: 128000, output: 4096 } } }, + }, + }, + }, + }), + ), + (tmp) => Effect.promise(() => tmp[Symbol.asyncDispose]()), + ) + const all = yield* request("/provider", tmp.path) + const connected = yield* request("/config/providers", tmp.path) + expect(yield* request("/kilo/auth-status", tmp.path)).toEqual({ + authenticated: true, + type: "oauth", + organizationId: "org-config", + }) + const unavailable = ["empty", "error", "unauthorized", "filtered"].includes(scenario) + expect(models(all, "all")).toEqual(unavailable ? [] : ["test/private", "test/z-last"]) + expect(models(connected, "providers")).toEqual(unavailable ? [] : ["test/private", "test/z-last"]) + expect(connected.default.kilo).toBe( + unavailable ? undefined : scenario === "valid" ? "test/z-last" : "test/private", + ) + expect(connected.default.external).toBe("independent") + expect(all.default.external).toBe("independent") + expect(all.connected).toContain("external") + expect(paths.filter((path) => path.endsWith("/models"))).toEqual(["/api/organizations/org-config/models"]) + expect(paths.filter((path) => path.endsWith("/defaults"))).toEqual( + unavailable ? [] : ["/api/organizations/org-config/defaults"], + ) + }), + ) +} + function record(input: unknown): input is Record { return typeof input === "object" && input !== null && !Array.isArray(input) } @@ -58,6 +199,8 @@ afterEach(async () => { it.live( "filters prompt-training models from both provider catalogs", Effect.gen(function* () { + const cache = yield* ModelCache.Service + yield* cache.clear("kilo") const server = yield* Effect.acquireRelease( Effect.sync(() => Bun.serve({ diff --git a/packages/sdk/js/src/v2/gen/types.gen.ts b/packages/sdk/js/src/v2/gen/types.gen.ts index 9a93fccecbf..790ba099931 100644 --- a/packages/sdk/js/src/v2/gen/types.gen.ts +++ b/packages/sdk/js/src/v2/gen/types.gen.ts @@ -16219,6 +16219,7 @@ export type KiloAuthStatusResponses = { 200: { authenticated: boolean type?: "api" | "oauth" + organizationId?: string } } diff --git a/packages/sdk/openapi.json b/packages/sdk/openapi.json index 78bd13cd5d9..803d9fe6e03 100644 --- a/packages/sdk/openapi.json +++ b/packages/sdk/openapi.json @@ -13536,6 +13536,9 @@ "type": { "type": "string", "enum": ["api", "oauth"] + }, + "organizationId": { + "type": "string" } }, "required": ["authenticated"],