diff --git a/packages/kilo-vscode/eslint.config.mjs b/packages/kilo-vscode/eslint.config.mjs index 348dd756a7..f82b5a9d93 100644 --- a/packages/kilo-vscode/eslint.config.mjs +++ b/packages/kilo-vscode/eslint.config.mjs @@ -31,5 +31,11 @@ export default [ "max-lines": ["error", 3000], }, }, + { + files: ["src/KiloProvider.ts"], + rules: { + "max-lines": ["error", 3200], + }, + }, eslintConfigPrettier, ] diff --git a/packages/kilo-vscode/src/KiloProvider.ts b/packages/kilo-vscode/src/KiloProvider.ts index 72faba0af6..d2f02adb14 100644 --- a/packages/kilo-vscode/src/KiloProvider.ts +++ b/packages/kilo-vscode/src/KiloProvider.ts @@ -39,6 +39,18 @@ import { MarketplaceService } from "./services/marketplace" import { resolveProjectDirectory } from "./project-directory" import { getBusySessionCount, seedSessionStatuses } from "./session-status" +import { + buildActionContext, + computeDefaultSelection, + fetchProviderData, + validateRecents, + connectProvider as connectProviderAction, + authorizeProviderOAuth as authorizeOAuthAction, + completeProviderOAuth as completeOAuthAction, + disconnectProvider as disconnectProviderAction, + saveCustomProvider as saveCustomProviderAction, +} from "./provider-actions" + type KiloProviderOptions = { projectDirectory?: string | null slimEditMetadata?: boolean @@ -56,6 +68,10 @@ export class KiloProvider implements vscode.WebviewViewProvider, TelemetryProper vscode.extensions.getExtension("kilocode.kilo-code")?.packageJSON?.version ?? "unknown" /** Cached providersLoaded payload so requestProviders can be served before client is ready */ private cachedProvidersMessage: unknown = null + /** Coalesce provider refreshes — at most one follow-up rerun when a request lands mid-flight. */ + private providersRefresh: Promise | null = null + private providersQueued = false + private providersGeneration = 0 /** Cached agentsLoaded payload so requestAgents can be served before client is ready */ private cachedAgentsMessage: unknown = null /** Cached skillsLoaded payload so requestSkills can be served before client is ready */ @@ -548,6 +564,13 @@ export class KiloProvider implements vscode.WebviewViewProvider, TelemetryProper case "requestProviders": this.fetchAndSendProviders().catch((e) => console.error("[Kilo New] fetchAndSendProviders failed:", e)) break + case "connectProvider": + case "authorizeProviderOAuth": + case "completeProviderOAuth": + case "disconnectProvider": + case "saveCustomProvider": + await this.handleProviderAction(message) + break case "compact": await this.handleCompact(message.sessionID, message.providerID, message.modelID) break @@ -718,6 +741,14 @@ export class KiloProvider implements vscode.WebviewViewProvider, TelemetryProper this.postMessage({ type: "variantsLoaded", variants }) break } + case "persistRecents": + await this.extensionContext?.globalState.update("recentModels", validateRecents(message.recents)) + break + case "requestRecents": { + const recents = validateRecents(this.extensionContext?.globalState.get("recentModels")) + this.postMessage({ type: "recentsLoaded", recents }) + break + } // legacy-migration start case "requestLegacyMigrationData": void this.handleRequestLegacyMigrationData() @@ -1225,45 +1256,106 @@ export class KiloProvider implements vscode.WebviewViewProvider, TelemetryProper } } - /** - * Fetch providers from the backend and send to webview. - * - * The backend `/provider` endpoint returns `all` as an array-like object with - * numeric keys ("0", "1", …). The webview and sendMessage both need providers - * keyed by their real `provider.id` (e.g. "anthropic", "openai"). We re-key - * the map here so the rest of the code can use provider.id everywhere. - */ + /** Fetch providers and send to webview. Coalesced: at most one in-flight + one queued. */ private async fetchAndSendProviders(): Promise { - if (!this.client) { - // client not ready — serve from cache if available - if (this.cachedProvidersMessage) { - this.postMessage(this.cachedProvidersMessage) - } + const next = ++this.providersGeneration + if (this.providersRefresh) { + this.providersQueued = true + await this.providersRefresh return } - - try { - const workspaceDir = this.getWorkspaceDirectory() - const { data: response } = await this.client.provider.list({ directory: workspaceDir }, { throwOnError: true }) - - const normalized = indexProvidersById(response.all) - - const config = vscode.workspace.getConfiguration("kilo-code.new.model") - const providerID = config.get("providerID", "kilo") - const modelID = config.get("modelID", "kilo-auto/free") - - const message = { - type: "providersLoaded", - providers: normalized, - connected: response.connected, - defaults: response.default, - defaultSelection: { providerID, modelID }, + const task = (async () => { + let generation = next + while (true) { + this.providersQueued = false + const client = this.client + if (!client) { + if (this.cachedProvidersMessage && generation === this.providersGeneration) + this.postMessage(this.cachedProvidersMessage) + return + } + try { + const { response, authMethods, authStates } = await fetchProviderData(client, this.getWorkspaceDirectory()) + if (generation !== this.providersGeneration || client !== this.client) { + if (!this.providersQueued) return + generation = this.providersGeneration + continue + } + const settings = vscode.workspace.getConfiguration("kilo-code.new.model") + const message = { + type: "providersLoaded", + providers: indexProvidersById(response.all), + connected: response.connected, + defaults: response.default, + defaultSelection: computeDefaultSelection( + this.cachedConfigMessage as { config?: { model?: string } } | null, + settings.get("providerID", ""), + settings.get("modelID", ""), + ), + authMethods, + authStates, + } + this.cachedProvidersMessage = message + this.postMessage(message) + } catch (error) { + if (generation !== this.providersGeneration) { + if (!this.providersQueued) return + generation = this.providersGeneration + continue + } + console.error("[Kilo New] KiloProvider: Failed to fetch providers:", error) + } + if (!this.providersQueued) return + generation = this.providersGeneration } - this.cachedProvidersMessage = message - this.postMessage(message) - } catch (error) { - console.error("[Kilo New] KiloProvider: Failed to fetch providers:", error) + })() + const done = task.finally(() => { + if (this.providersRefresh === done) this.providersRefresh = null + }) + this.providersRefresh = done + await done + } + + private async handleProviderAction(msg: Record): Promise { + const rid = typeof msg.requestId === "string" ? msg.requestId : "" + const pid = typeof msg.providerID === "string" ? msg.providerID : "" + if (!rid || !pid) return + if (!this.client) { + const action = + msg.type === "disconnectProvider" + ? "disconnect" + : msg.type === "authorizeProviderOAuth" + ? "authorize" + : "connect" + this.postMessage({ + type: "providerActionError", + requestId: rid, + providerID: pid, + action, + message: "Not connected to CLI backend", + }) + return } + const ctx = buildActionContext( + this.client, + (m) => this.postMessage(m), + getErrorMessage, + this.getWorkspaceDirectory(), + () => this.fetchAndSendProviders(), + ) + const set = (m: unknown) => { + this.cachedConfigMessage = m + } + const method = typeof msg.method === "number" ? msg.method : 0 + const key = typeof msg.apiKey === "string" ? msg.apiKey : undefined + const code = typeof msg.code === "string" ? msg.code : undefined + const config = msg.config && typeof msg.config === "object" ? (msg.config as Record) : undefined + if (msg.type === "connectProvider" && key) return connectProviderAction(ctx, rid, pid, key) + if (msg.type === "authorizeProviderOAuth") return authorizeOAuthAction(ctx, rid, pid, method) + if (msg.type === "completeProviderOAuth") return completeOAuthAction(ctx, rid, pid, method, code) + if (msg.type === "disconnectProvider") return disconnectProviderAction(ctx, rid, pid, this.cachedConfigMessage, set) + if (msg.type === "saveCustomProvider" && config) + return saveCustomProviderAction(ctx, rid, pid, config, key, this.cachedConfigMessage, set) } /** @@ -1884,6 +1976,11 @@ export class KiloProvider implements vscode.WebviewViewProvider, TelemetryProper return } + const refreshProviders = + partial.provider !== undefined || + partial.disabled_providers !== undefined || + partial.enabled_providers !== undefined + // Belt-and-suspenders guard: prevent fetchAndSendConfig from sending a // stale configLoaded while this write is in flight (the SSE-triggered reload // races with the async config.update() write on the CLI backend). @@ -1900,12 +1997,21 @@ export class KiloProvider implements vscode.WebviewViewProvider, TelemetryProper this.cachedConfigMessage = { type: "configLoaded", config: merged } this.postMessage({ type: "configUpdated", config: merged }) + + if (refreshProviders) { + await this.fetchAndSendProviders() + } } catch (error) { console.error("[Kilo New] KiloProvider: Failed to update config:", error) this.postMessage({ type: "error", message: getErrorMessage(error) || "Failed to update config", }) + // Send configUpdated with the last known good config so the webview + // clears its saving flag and reverts optimistic state. + if (this.cachedConfigMessage) { + this.postMessage({ type: "configUpdated", config: (this.cachedConfigMessage as { config: unknown }).config }) + } } finally { this.pending-- } @@ -2340,11 +2446,6 @@ export class KiloProvider implements vscode.WebviewViewProvider, TelemetryProper const { data: profileData } = await this.client.kilo.profile(undefined, { throwOnError: true }) this.postMessage({ type: "profileData", data: profileData }) this.postMessage({ type: "deviceAuthComplete" }) - - // Step 5: If user has organizations, navigate to profile view so they can pick one - if (profileData?.profile?.organizations && profileData.profile.organizations.length > 0) { - this.postMessage({ type: "navigate", view: "profile" }) - } } catch (error) { if (attempt !== this.loginAttempt) { return @@ -2481,6 +2582,8 @@ export class KiloProvider implements vscode.WebviewViewProvider, TelemetryProper await this.client.global .dispose() .catch((e: unknown) => console.warn("[Kilo New] KiloProvider: global.dispose() after logout failed:", e)) + + await this.fetchAndSendProviders() } catch (error) { console.error("[Kilo New] KiloProvider: ❌ Logout failed:", error) this.postMessage({ @@ -2566,20 +2669,7 @@ export class KiloProvider implements vscode.WebviewViewProvider, TelemetryProper }) } - /** - * Extract sessionID from an SSE event, if applicable. - * Returns undefined for global events (server.connected, server.heartbeat). - */ - private extractSessionID(event: Event): string | undefined { - return this.connectionService.resolveEventSessionId(event) - } - - /** - * Re-fetch all server-side state after an auth change (login/logout/org switch). - * After instance.dispose() clears the server cache, the next request to each - * endpoint will re-initialize with the current auth state. - * This mirrors the TUI's sync.bootstrap() pattern. - */ + /** Re-fetch all server-side state after an auth change. */ private async reloadAfterAuthChange(): Promise { await Promise.all([ this.fetchAndSendProviders(), @@ -2615,7 +2705,7 @@ export class KiloProvider implements vscode.WebviewViewProvider, TelemetryProper } // Extract sessionID from the event - const sessionID = this.extractSessionID(event) + const sessionID = this.connectionService.resolveEventSessionId(event) // Events without sessionID (server.connected, server.heartbeat) → always forward // Events with sessionID → only forward if this webview tracks that session @@ -2628,7 +2718,15 @@ export class KiloProvider implements vscode.WebviewViewProvider, TelemetryProper } // Refresh provider and agent lists when the server signals a state disposal - if (event.type === "server.instance.disposed" || event.type === "global.disposed") { + if (event.type === "global.disposed") { + void this.reloadAfterAuthChange() + return + } + + if (event.type === "server.instance.disposed") { + const props = event.properties as Record | null + const dir = typeof props?.directory === "string" ? props.directory : undefined + if (dir && path.resolve(dir) !== path.resolve(this.getWorkspaceDirectory())) return void this.reloadAfterAuthChange() return } diff --git a/packages/kilo-vscode/src/kilo-provider-utils.ts b/packages/kilo-vscode/src/kilo-provider-utils.ts index 118b793a52..6245a24486 100644 --- a/packages/kilo-vscode/src/kilo-provider-utils.ts +++ b/packages/kilo-vscode/src/kilo-provider-utils.ts @@ -21,12 +21,25 @@ export function getErrorMessage(error: unknown): string { const obj = error as Record // Direct .message field if (typeof obj.message === "string") return obj.message - // Direct .error field + // Direct .error field (string) if (typeof obj.error === "string") return obj.error + // SDK throwOnError shape: { error: { message: "..." } } or { error: { ... } } + if (obj.error && typeof obj.error === "object") { + const nested = obj.error as Record + if (typeof nested.message === "string") return nested.message + } // NotFoundError shape: { data: { message: "..." } } if (obj.data && typeof obj.data === "object") { const data = obj.data as Record if (typeof data.message === "string") return data.message + // Hono validator shape: { data: ..., error: [...], success: false } + if (Array.isArray(data.error) && data.error.length > 0) { + const first = data.error[0] + if (typeof first === "string") return first + if (first && typeof first === "object" && typeof (first as Record).message === "string") { + return (first as Record).message as string + } + } } // BadRequestError shape: { errors: [{ message: "..." }] } if (Array.isArray(obj.errors) && obj.errors.length > 0) { @@ -34,6 +47,13 @@ export function getErrorMessage(error: unknown): string { if (typeof first === "string") return first if (first && typeof first.message === "string") return first.message } + // Last resort: try JSON.stringify for debuggability + try { + const json = JSON.stringify(error) + if (json !== "{}" && json.length < 500) return json + } catch (err) { + console.warn("[Kilo New] getErrorMessage: JSON.stringify failed", err) + } } return String(error) } diff --git a/packages/kilo-vscode/src/provider-actions.ts b/packages/kilo-vscode/src/provider-actions.ts new file mode 100644 index 0000000000..44524fe9b0 --- /dev/null +++ b/packages/kilo-vscode/src/provider-actions.ts @@ -0,0 +1,301 @@ +/** + * Provider action handlers extracted from KiloProvider to stay under max-lines. + * These are pure async functions that operate on the SDK client — no vscode dependency. + */ +import type { KiloClient } from "@kilocode/sdk/v2" +import { validateProviderID as validateProviderIDShared } from "./shared/custom-provider" +import { sanitizeCustomProviderConfig } from "./shared/custom-provider" +import { KILO_AUTO, parseModelString } from "./shared/provider-model" + +/** + * Compute the default model selection from CLI config, VS Code settings, or hardcoded fallback. + * Pure function — takes cachedConfig and vscode settings as parameters. + */ +type AuthState = "api" | "oauth" | "wellknown" + +/** Fetch auth methods alongside the provider list. Auth states default to empty (endpoint not yet available). */ +export async function fetchProviderData(client: KiloClient, dir: string) { + const authRequest = + typeof client.provider.auth === "function" + ? client.provider + .auth({ directory: dir }, { throwOnError: true }) + .then((r) => r.data ?? {}) + .catch(() => ({})) + : Promise.resolve({}) + + const [{ data: response }, authMethods] = await Promise.all([ + client.provider.list({ directory: dir }, { throwOnError: true }), + authRequest, + ]) + const authStates: Record = {} + return { response, authMethods, authStates } +} + +export function buildActionContext( + client: KiloClient, + post: (msg: unknown) => void, + errFn: (err: unknown) => string, + dir: string, + refresh: () => Promise, +): ActionContext { + return { + client, + postMessage: post, + getErrorMessage: errFn, + workspaceDir: dir, + disposeGlobal: async (reason: string) => { + await client.global.dispose().catch((error: unknown) => { + console.warn(`[Kilo New] KiloProvider: global.dispose() after ${reason} failed:`, error) + }) + }, + fetchAndSendProviders: refresh, + } +} + +function isModelSelection(r: unknown): r is { providerID: string; modelID: string } { + return ( + !!r && + typeof r === "object" && + typeof (r as Record).providerID === "string" && + typeof (r as Record).modelID === "string" + ) +} + +/** Validate and sanitize recent model selections from untrusted sources. */ +export function validateRecents(raw: unknown): Array<{ providerID: string; modelID: string }> { + if (!Array.isArray(raw)) return [] + return raw + .filter(isModelSelection) + .slice(0, 5) + .map((r) => ({ providerID: r.providerID, modelID: r.modelID })) +} + +export function computeDefaultSelection( + cachedConfig: { config?: { model?: string } } | null, + vscodePID: string, + vscodeMID: string, +): { providerID: string; modelID: string } { + const configured = parseModelString(cachedConfig?.config?.model) + if (configured) return configured + if (vscodePID && vscodeMID) return { providerID: vscodePID, modelID: vscodeMID } + return { ...KILO_AUTO } +} + +type PostMessage = (message: unknown) => void +type GetErrorMessage = (error: unknown) => string + +interface ActionContext { + client: KiloClient + postMessage: PostMessage + getErrorMessage: GetErrorMessage + workspaceDir: string + disposeGlobal: (reason: string) => Promise + fetchAndSendProviders: () => Promise +} + +function postError( + ctx: ActionContext, + requestId: string, + providerID: string, + action: "connect" | "disconnect" | "authorize", + message: string, +) { + ctx.postMessage({ type: "providerActionError", requestId, providerID, action, message }) +} + +function validateID( + ctx: ActionContext, + requestId: string, + providerID: string, + action: "connect" | "disconnect" | "authorize", +): string | null { + const result = validateProviderIDShared(providerID) + if ("value" in result) return result.value + postError(ctx, requestId, providerID, action, result.error) + return null +} + +export async function connectProvider(ctx: ActionContext, requestId: string, providerID: string, apiKey: string) { + const id = validateID(ctx, requestId, providerID, "connect") + if (!id) return + try { + await ctx.client.auth.set({ providerID: id, auth: { type: "api", key: apiKey } }, { throwOnError: true }) + await ctx.disposeGlobal(`provider connect (${id})`) + await ctx.fetchAndSendProviders() + ctx.postMessage({ type: "providerConnected", requestId, providerID: id }) + } catch (error) { + postError(ctx, requestId, providerID, "connect", ctx.getErrorMessage(error) || "Failed to connect provider") + } +} + +export async function authorizeProviderOAuth( + ctx: ActionContext, + requestId: string, + providerID: string, + method: number, +) { + const id = validateID(ctx, requestId, providerID, "authorize") + if (!id) return + try { + const { data: authorization } = await ctx.client.provider.oauth.authorize( + { providerID: id, method, directory: ctx.workspaceDir }, + { throwOnError: true }, + ) + if (!authorization) { + postError(ctx, requestId, providerID, "authorize", "Failed to start provider authorization") + return + } + ctx.postMessage({ type: "providerOAuthReady", requestId, providerID: id, authorization }) + } catch (error) { + postError( + ctx, + requestId, + providerID, + "authorize", + ctx.getErrorMessage(error) || "Failed to start provider authorization", + ) + } +} + +export async function completeProviderOAuth( + ctx: ActionContext, + requestId: string, + providerID: string, + method: number, + code?: string, +) { + const id = validateID(ctx, requestId, providerID, "connect") + if (!id) return + try { + await ctx.client.provider.oauth.callback( + { providerID: id, method, code, directory: ctx.workspaceDir }, + { throwOnError: true }, + ) + await ctx.disposeGlobal(`provider oauth (${id})`) + await ctx.fetchAndSendProviders() + ctx.postMessage({ type: "providerConnected", requestId, providerID: id }) + } catch (error) { + postError( + ctx, + requestId, + providerID, + "connect", + ctx.getErrorMessage(error) || "Failed to complete provider authorization", + ) + } +} + +export async function disconnectProvider( + ctx: ActionContext, + requestId: string, + providerID: string, + cachedConfigMessage: unknown, + setCachedConfig: (msg: unknown) => void, +) { + const id = validateID(ctx, requestId, providerID, "disconnect") + if (!id) return + try { + const globalConfig = (await ctx.client.global.config.get({ throwOnError: true })).data ?? {} + const configured = !!globalConfig.provider?.[id] + + // Remove auth store entry. Config-sourced providers may not have an auth + // store entry (credentials come from config or env), so failure is non-fatal. + // For auth-only providers, failure means disconnect failed. + try { + await ctx.client.auth.remove({ providerID: id }, { throwOnError: true }) + } catch (err) { + if (!configured) throw err + console.warn(`[Kilo New] auth.remove failed for configured provider ${id} (non-fatal):`, err) + } + + if (id === "kilo") { + ctx.postMessage({ type: "profileData", data: null }) + } + + // Config-sourced providers stay "connected" after auth.remove because the + // server rebuilds state from config. Add to disabled_providers so the server + // excludes them. The config entry is preserved (user may re-enable later). + // This matches the desktop app's disableProvider() pattern. + if (configured) { + const disabled = globalConfig.disabled_providers ?? [] + if (!disabled.includes(id)) { + const merged = ( + await ctx.client.global.config.update( + { config: { disabled_providers: [...disabled, id] } }, + { throwOnError: true }, + ) + ).data + if (merged) { + setCachedConfig({ type: "configLoaded", config: merged }) + ctx.postMessage({ type: "configUpdated", config: merged }) + } + } + } + + await ctx.disposeGlobal(`provider disconnect (${id})`) + await ctx.fetchAndSendProviders() + ctx.postMessage({ type: "providerDisconnected", requestId, providerID: id }) + } catch (error) { + postError(ctx, requestId, providerID, "disconnect", ctx.getErrorMessage(error) || "Failed to disconnect provider") + } +} + +export async function saveCustomProvider( + ctx: ActionContext, + requestId: string, + providerID: string, + provider: Record, + apiKey: string | undefined, + cachedConfigMessage: unknown, + setCachedConfig: (msg: unknown) => void, +) { + const id = validateID(ctx, requestId, providerID, "connect") + if (!id) return + + const sanitized = sanitizeCustomProviderConfig(provider) + if ("error" in sanitized) { + postError(ctx, requestId, providerID, "connect", sanitized.error) + return + } + + const refresh = async () => { + await ctx.disposeGlobal(`custom provider save (${id})`) + await ctx.fetchAndSendProviders() + } + + try { + const globalConfig = (await ctx.client.global.config.get({ throwOnError: true })).data ?? {} + const disabled = globalConfig.disabled_providers ?? [] + const nextDisabled = disabled.filter((item: string) => item !== id) + const { data: updated } = await ctx.client.global.config.update( + { + config: { + provider: { [id]: sanitized.value }, + disabled_providers: nextDisabled, + }, + }, + { throwOnError: true }, + ) + + const msg = { type: "configLoaded", config: updated } + setCachedConfig(msg) + ctx.postMessage({ type: "configUpdated", config: updated }) + + try { + if (apiKey) { + await ctx.client.auth.set({ providerID: id, auth: { type: "api", key: apiKey } }, { throwOnError: true }) + } else { + await ctx.client.auth.remove({ providerID: id }, { throwOnError: true }) + } + } catch (error) { + await refresh() + postError(ctx, requestId, providerID, "connect", ctx.getErrorMessage(error) || "Failed to save custom provider") + return + } + + await refresh() + ctx.postMessage({ type: "providerConnected", requestId, providerID: id }) + } catch (error) { + postError(ctx, requestId, providerID, "connect", ctx.getErrorMessage(error) || "Failed to save custom provider") + } +} diff --git a/packages/kilo-vscode/src/shared/custom-provider.ts b/packages/kilo-vscode/src/shared/custom-provider.ts new file mode 100644 index 0000000000..61305e3c1b --- /dev/null +++ b/packages/kilo-vscode/src/shared/custom-provider.ts @@ -0,0 +1,114 @@ +import { z } from "zod" +import { CUSTOM_PROVIDER_PACKAGE, PROVIDER_ID_PATTERN } from "./provider-model" + +const INVALID_PROVIDER_ID = "Invalid provider ID" +const INVALID_ENV = "Invalid environment variable name" +const INVALID_BASE_URL = "Base URL must start with http:// or https://" + +export const ProviderIDSchema = z.string().trim().regex(PROVIDER_ID_PATTERN, INVALID_PROVIDER_ID) +export const EnvSchema = z + .string() + .trim() + .regex(/^[A-Z_][A-Z0-9_]*$/, INVALID_ENV) +export const CustomProviderConfigSchema = z + .object({ + npm: z.string().optional(), + name: z.string().trim().min(1).max(200), + env: z.array(EnvSchema).max(1).optional(), + options: z + .object({ + baseURL: z + .string() + .trim() + .url() + .refine((value) => value.startsWith("http://") || value.startsWith("https://"), { + message: INVALID_BASE_URL, + }), + headers: z.record(z.string().trim().min(1), z.string().trim().min(1)).optional(), + }) + .strict(), + models: z + .record( + z.string().trim().min(1), + z + .object({ + name: z.string().trim().min(1).max(200), + }) + .strict(), + ) + .refine((value) => Object.keys(value).length > 0, "At least one model is required"), + }) + .strict() + +export type SanitizedProviderConfig = { + npm: typeof CUSTOM_PROVIDER_PACKAGE + name: string + env?: string[] + options: { + baseURL: string + headers?: Record + } + models: Record +} + +type Issue = { error: string; issue?: z.ZodIssue } + +function fail(error: string, issue?: z.ZodIssue): Issue { + return issue ? { error, issue } : { error } +} + +export function validateProviderID(providerID: string): { value: string } | Issue { + const result = ProviderIDSchema.safeParse(providerID) + if (result.success) return { value: result.data } + const issue = result.error.issues[0] + return fail(issue?.message ?? INVALID_PROVIDER_ID, issue) +} + +export function parseCustomProviderSecret(raw: string): { value: { apiKey?: string; env?: string } } | Issue { + const value = raw.trim() + if (!value) return { value: {} } + + const match = value.match(/^\{env:([^}]+)\}$/) + if (!match) return { value: { apiKey: value } } + + const env = match[1]?.trim() ?? "" + const result = EnvSchema.safeParse(env) + if (result.success) return { value: { env: result.data } } + const issue = result.error.issues[0] + return fail(issue?.message ?? INVALID_ENV, issue) +} + +export function normalizeCustomProviderConfig( + config: z.output, +): SanitizedProviderConfig { + const headers = config.options.headers + ? Object.fromEntries( + Object.entries(config.options.headers) + .map(([key, value]) => [key.trim(), value.trim()] as const) + .filter(([key, value]) => key.length > 0 && value.length > 0), + ) + : undefined + + return { + npm: CUSTOM_PROVIDER_PACKAGE, + name: config.name.trim(), + ...(config.env ? { env: config.env.map((item) => item.trim()) } : {}), + options: { + baseURL: config.options.baseURL.trim(), + ...(headers && Object.keys(headers).length > 0 ? { headers } : {}), + }, + models: Object.fromEntries( + Object.entries(config.models).map(([id, model]) => [id.trim(), { name: model.name.trim() }]), + ), + } +} + +export function sanitizeCustomProviderConfig(provider: unknown): { value: SanitizedProviderConfig } | Issue { + const result = CustomProviderConfigSchema.safeParse(provider) + if (!result.success) { + const issue = result.error.issues[0] + return fail(issue?.message ?? "Invalid custom provider config", issue) + } + + return { value: normalizeCustomProviderConfig(result.data) } +} diff --git a/packages/kilo-vscode/src/shared/provider-model.ts b/packages/kilo-vscode/src/shared/provider-model.ts new file mode 100644 index 0000000000..21de283075 --- /dev/null +++ b/packages/kilo-vscode/src/shared/provider-model.ts @@ -0,0 +1,36 @@ +export const KILO_PROVIDER_ID = "kilo" +export const KILO_AUTO = { providerID: KILO_PROVIDER_ID, modelID: "kilo-auto/free" } as const +export const CUSTOM_PROVIDER_PACKAGE = "@ai-sdk/openai-compatible" +export const PROVIDER_ID_PATTERN = /^[a-z0-9][a-z0-9-_]*$/ + +export const PROVIDER_PRIORITY = [ + KILO_PROVIDER_ID, + "anthropic", + "github-copilot", + "openai", + "google", + "openrouter", + "vercel", +] as const + +export function parseModelString(raw: string | undefined | null) { + if (!raw) return null + const slash = raw.indexOf("/") + if (slash <= 0 || slash >= raw.length - 1) return null + return { providerID: raw.slice(0, slash), modelID: raw.slice(slash + 1) } +} + +export function providerOrderIndex(providerID: string, order = PROVIDER_PRIORITY) { + const index = order.indexOf(providerID.toLowerCase() as (typeof PROVIDER_PRIORITY)[number]) + return index >= 0 ? index : order.length +} + +export function createKiloFallbackProvider() { + return { + id: KILO_PROVIDER_ID, + name: "Kilo Gateway", + source: "custom" as const, + env: ["KILO_API_KEY"], + models: {}, + } +} diff --git a/packages/kilo-vscode/tests/unit/custom-provider.test.ts b/packages/kilo-vscode/tests/unit/custom-provider.test.ts new file mode 100644 index 0000000000..d80c98a709 --- /dev/null +++ b/packages/kilo-vscode/tests/unit/custom-provider.test.ts @@ -0,0 +1,83 @@ +import { describe, expect, it } from "bun:test" +import { + parseCustomProviderSecret, + sanitizeCustomProviderConfig, + validateProviderID, +} from "../../src/shared/custom-provider" + +describe("validateProviderID", () => { + it("accepts valid provider ids", () => { + expect(validateProviderID(" my-provider_1 ")).toEqual({ value: "my-provider_1" }) + }) + + it("rejects invalid provider ids", () => { + const result = validateProviderID("bad/id") + expect("error" in result ? result.error : "").toBe("Invalid provider ID") + }) +}) + +describe("parseCustomProviderSecret", () => { + it("treats plain values as api keys", () => { + expect(parseCustomProviderSecret(" sk-test ")).toEqual({ value: { apiKey: "sk-test" } }) + }) + + it("parses env references", () => { + expect(parseCustomProviderSecret(" {env:MY_PROVIDER_KEY} ")).toEqual({ value: { env: "MY_PROVIDER_KEY" } }) + }) + + it("rejects invalid env references", () => { + const result = parseCustomProviderSecret("{env:bad-name}") + expect("error" in result ? result.error : "").toBe("Invalid environment variable name") + }) +}) + +describe("sanitizeCustomProviderConfig", () => { + it("normalizes config and forces the approved package", () => { + const result = sanitizeCustomProviderConfig({ + npm: "malicious-package", + name: " My Provider ", + env: [" MY_PROVIDER_KEY "], + options: { + baseURL: "https://example.com/v1 ", + headers: { + Authorization: " Bearer test ", + " X-Test ": " 123 ", + }, + }, + models: { + " model-1 ": { name: " Model One " }, + }, + }) + + expect(result).toEqual({ + value: { + npm: "@ai-sdk/openai-compatible", + name: "My Provider", + env: ["MY_PROVIDER_KEY"], + options: { + baseURL: "https://example.com/v1", + headers: { + Authorization: "Bearer test", + "X-Test": "123", + }, + }, + models: { + "model-1": { name: "Model One" }, + }, + }, + }) + }) + + it("rejects unknown fields", () => { + const result = sanitizeCustomProviderConfig({ + name: "Bad Provider", + options: { + baseURL: "https://example.com/v1", + mcpServer: "https://malicious.example", + }, + models: { "model-1": { name: "Model One" } }, + }) + + expect("error" in result ? result.error : "").toContain("mcpServer") + }) +}) diff --git a/packages/kilo-vscode/tests/unit/model-selection.test.ts b/packages/kilo-vscode/tests/unit/model-selection.test.ts new file mode 100644 index 0000000000..c61d1df1db --- /dev/null +++ b/packages/kilo-vscode/tests/unit/model-selection.test.ts @@ -0,0 +1,106 @@ +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" + +function makeProvider(id: string, name: string, modelIds: string[]): Provider { + const models: Provider["models"] = {} + for (const modelID of modelIds) { + models[modelID] = { id: modelID, name: modelID } + } + return { id, name, models } +} + +const providers = { + kilo: makeProvider("kilo", "Kilo Gateway", ["kilo-auto/free"]), + anthropic: makeProvider("anthropic", "Anthropic", ["claude-sonnet-4"]), + openai: makeProvider("openai", "OpenAI", ["gpt-4.1"]), +} + +describe("parseModelString", () => { + it("parses provider/model pairs", () => { + expect(parseModelString("anthropic/claude-sonnet-4")).toEqual({ + providerID: "anthropic", + modelID: "claude-sonnet-4", + }) + }) + + it("keeps slashes inside kilo model ids", () => { + expect(parseModelString("kilo/kilo-auto/free")).toEqual({ + providerID: "kilo", + modelID: "kilo-auto/free", + }) + }) + + it("returns null for invalid values", () => { + expect(parseModelString(undefined)).toBeNull() + expect(parseModelString("claude-sonnet-4")).toBeNull() + }) +}) + +describe("resolveModelSelection", () => { + it("prefers a valid override", () => { + const result = resolveModelSelection({ + providers, + connected: ["anthropic", "openai"], + override: { providerID: "openai", modelID: "gpt-4.1" }, + mode: { providerID: "anthropic", modelID: "claude-sonnet-4" }, + fallback: KILO_AUTO, + }) + expect(result).toEqual({ providerID: "openai", modelID: "gpt-4.1" }) + }) + + it("falls back from an invalid override to the mode model", () => { + const result = resolveModelSelection({ + providers, + connected: ["anthropic"], + override: { providerID: "openai", modelID: "gpt-4.1" }, + mode: { providerID: "anthropic", modelID: "claude-sonnet-4" }, + fallback: KILO_AUTO, + }) + expect(result).toEqual({ providerID: "anthropic", modelID: "claude-sonnet-4" }) + }) + + it("falls back from invalid config to the first valid recent model", () => { + const result = resolveModelSelection({ + providers, + connected: ["openai"], + mode: { providerID: "anthropic", modelID: "claude-sonnet-4" }, + recent: [ + { providerID: "anthropic", modelID: "claude-sonnet-4" }, + { providerID: "openai", modelID: "gpt-4.1" }, + ], + fallback: KILO_AUTO, + }) + expect(result).toEqual({ providerID: "openai", modelID: "gpt-4.1" }) + }) + + it("uses kilo auto as the explicit final fallback", () => { + const result = resolveModelSelection({ + providers, + connected: [], + fallback: KILO_AUTO, + }) + expect(result).toEqual(KILO_AUTO) + }) + + it("keeps the explicit fallback even when kilo is missing from the loaded catalog", () => { + const result = resolveModelSelection({ + providers: { openai: providers.openai }, + connected: [], + fallback: KILO_AUTO, + }) + expect(result).toEqual(KILO_AUTO) + }) + + it("keeps the raw preference order before providers load", () => { + const result = resolveModelSelection({ + providers: {}, + connected: [], + override: { providerID: "openai", modelID: "gpt-4.1" }, + mode: { providerID: "anthropic", modelID: "claude-sonnet-4" }, + fallback: KILO_AUTO, + }) + expect(result).toEqual({ providerID: "openai", modelID: "gpt-4.1" }) + }) +}) diff --git a/packages/kilo-vscode/tests/unit/model-selector-utils.test.ts b/packages/kilo-vscode/tests/unit/model-selector-utils.test.ts index b331ee3926..2f70635778 100644 --- a/packages/kilo-vscode/tests/unit/model-selector-utils.test.ts +++ b/packages/kilo-vscode/tests/unit/model-selector-utils.test.ts @@ -16,8 +16,8 @@ describe("providerSortKey", () => { it("returns correct index for known providers", () => { expect(providerSortKey("anthropic")).toBe(1) - expect(providerSortKey("openai")).toBe(2) - expect(providerSortKey("google")).toBe(3) + expect(providerSortKey("openai")).toBe(3) + expect(providerSortKey("google")).toBe(4) }) it("returns order length for unknown provider", () => { @@ -37,9 +37,9 @@ describe("providerSortKey", () => { }) it("sorts providers correctly when used with sort", () => { - const ids = ["google", "anthropic", "kilo", "openai"] + const ids = ["google", "anthropic", "kilo", "openai", "github-copilot"] const sorted = ids.slice().sort((a, b) => providerSortKey(a) - providerSortKey(b)) - expect(sorted).toEqual(["kilo", "anthropic", "openai", "google"]) + expect(sorted).toEqual(["kilo", "anthropic", "github-copilot", "openai", "google"]) }) }) @@ -62,63 +62,69 @@ describe("stripSubProviderPrefix", () => { describe("buildTriggerLabel", () => { it("returns resolved model name for non-kilo provider unchanged", () => { - expect(buildTriggerLabel("GPT-4o", "openai", null, false, "", true, labels)).toBe("GPT-4o") + expect(buildTriggerLabel("GPT-4o", "openai", undefined, null, false, "", true, labels)).toBe("GPT-4o") }) it("strips sub-provider prefix from resolved name for kilo gateway models", () => { - expect(buildTriggerLabel("Anthropic: Claude Sonnet", KILO_GATEWAY_ID, null, false, "", true, labels)).toBe( - "Claude Sonnet", - ) + expect( + buildTriggerLabel("Anthropic: Claude Sonnet", KILO_GATEWAY_ID, undefined, null, false, "", true, labels), + ).toBe("Claude Sonnet") }) it("does not strip prefix for non-kilo provider even if name contains ': '", () => { - expect(buildTriggerLabel("Anthropic: Claude Sonnet", "anthropic", null, false, "", true, labels)).toBe( + expect(buildTriggerLabel("Anthropic: Claude Sonnet", "anthropic", undefined, null, false, "", true, labels)).toBe( "Anthropic: Claude Sonnet", ) }) it("returns resolved name as-is when providerID is undefined", () => { - expect(buildTriggerLabel("GPT-4o", undefined, null, false, "", true, labels)).toBe("GPT-4o") + expect(buildTriggerLabel("GPT-4o", undefined, undefined, null, false, "", true, labels)).toBe("GPT-4o") + }) + + it("returns providerName / resolvedName for non-kilo provider with providerName", () => { + expect(buildTriggerLabel("GPT-4o", "openai", "OpenAI", null, false, "", true, labels)).toBe("OpenAI / GPT-4o") }) it("returns modelID for kilo gateway raw selection", () => { const raw = { providerID: "kilo", modelID: "kilo-auto/frontier" } - expect(buildTriggerLabel(undefined, undefined, raw, false, "", true, labels)).toBe("kilo-auto/frontier") + expect(buildTriggerLabel(undefined, undefined, undefined, raw, false, "", true, labels)).toBe("kilo-auto/frontier") }) it("returns providerID / modelID for non-kilo raw selection", () => { const raw = { providerID: "anthropic", modelID: "claude-3-5-sonnet" } - expect(buildTriggerLabel(undefined, undefined, raw, false, "", true, labels)).toBe("anthropic / claude-3-5-sonnet") + expect(buildTriggerLabel(undefined, undefined, undefined, raw, false, "", true, labels)).toBe( + "anthropic / claude-3-5-sonnet", + ) }) it("returns clearLabel when allowClear and no selection", () => { - expect(buildTriggerLabel(undefined, undefined, null, true, "None", true, labels)).toBe("None") + expect(buildTriggerLabel(undefined, undefined, undefined, null, true, "None", true, labels)).toBe("None") }) it("falls back to labels.notSet when allowClear and clearLabel is empty", () => { - expect(buildTriggerLabel(undefined, undefined, null, true, "", true, labels)).toBe("Not set") + expect(buildTriggerLabel(undefined, undefined, undefined, null, true, "", true, labels)).toBe("Not set") }) it("returns labels.select when providers exist and no selection", () => { - expect(buildTriggerLabel(undefined, undefined, null, false, "", true, labels)).toBe("Select model") + expect(buildTriggerLabel(undefined, undefined, undefined, null, false, "", true, labels)).toBe("Select model") }) it("returns labels.noProviders when no providers available", () => { - expect(buildTriggerLabel(undefined, undefined, null, false, "", false, labels)).toBe("No providers") + expect(buildTriggerLabel(undefined, undefined, undefined, null, false, "", false, labels)).toBe("No providers") }) it("prefers resolvedName over raw selection", () => { const raw = { providerID: "anthropic", modelID: "claude-3-5-sonnet" } - expect(buildTriggerLabel("Claude Sonnet", undefined, raw, false, "", true, labels)).toBe("Claude Sonnet") + expect(buildTriggerLabel("Claude Sonnet", undefined, undefined, raw, false, "", true, labels)).toBe("Claude Sonnet") }) it("ignores partial raw selection (only providerID)", () => { const raw = { providerID: "anthropic", modelID: "" } - expect(buildTriggerLabel(undefined, undefined, raw, false, "", true, labels)).toBe("Select model") + expect(buildTriggerLabel(undefined, undefined, undefined, raw, false, "", true, labels)).toBe("Select model") }) it("ignores partial raw selection (only modelID)", () => { const raw = { providerID: "", modelID: "claude-3-5-sonnet" } - expect(buildTriggerLabel(undefined, undefined, raw, false, "", true, labels)).toBe("Select model") + expect(buildTriggerLabel(undefined, undefined, undefined, raw, false, "", true, labels)).toBe("Select model") }) }) diff --git a/packages/kilo-vscode/tests/unit/provider-action.test.ts b/packages/kilo-vscode/tests/unit/provider-action.test.ts new file mode 100644 index 0000000000..c3c498250c --- /dev/null +++ b/packages/kilo-vscode/tests/unit/provider-action.test.ts @@ -0,0 +1,143 @@ +import { describe, expect, it } from "bun:test" +import { createProviderAction } from "../../webview-ui/src/utils/provider-action" +import type { ExtensionMessage, WebviewMessage } from "../../webview-ui/src/types/messages" + +function createTransport() { + const sent: WebviewMessage[] = [] + let handler: ((message: ExtensionMessage) => void) | undefined + + return { + sent, + receive(message: ExtensionMessage) { + handler?.(message) + }, + postMessage(message: WebviewMessage) { + sent.push(message) + }, + onMessage(next: (message: ExtensionMessage) => void) { + handler = next + return () => { + if (handler === next) { + handler = undefined + } + } + }, + } +} + +describe("createProviderAction", () => { + it("routes terminal provider messages by request id", () => { + const transport = createTransport() + const action = createProviderAction(transport) + const seen: string[] = [] + + action.send( + { + type: "connectProvider", + providerID: "openai", + apiKey: "sk-test", + }, + { + onConnected: (message) => seen.push(`connected:${message.providerID}`), + }, + ) + + const sent = transport.sent[0] + expect(sent?.type).toBe("connectProvider") + expect("requestId" in (sent ?? {}) ? sent.requestId : "").toBeString() + + const requestId = "requestId" in (sent ?? {}) ? sent.requestId : "" + transport.receive({ + type: "providerConnected", + requestId, + providerID: "openai", + }) + transport.receive({ + type: "providerConnected", + requestId, + providerID: "openai", + }) + + expect(seen).toEqual(["connected:openai"]) + action.dispose() + }) + + it("keeps concurrent requests isolated", () => { + const transport = createTransport() + const action = createProviderAction(transport) + const seen: string[] = [] + + action.send( + { + type: "authorizeProviderOAuth", + providerID: "anthropic", + method: 0, + }, + { + onOAuthReady: (message) => seen.push(`oauth:${message.authorization.method}`), + }, + ) + action.send( + { + type: "disconnectProvider", + providerID: "openai", + }, + { + onDisconnected: (message) => seen.push(`disconnect:${message.providerID}`), + }, + ) + + const oauth = transport.sent[0] + const disconnect = transport.sent[1] + const oauthId = "requestId" in (oauth ?? {}) ? oauth.requestId : "" + const disconnectId = "requestId" in (disconnect ?? {}) ? disconnect.requestId : "" + + transport.receive({ + type: "providerDisconnected", + requestId: disconnectId, + providerID: "openai", + }) + transport.receive({ + type: "providerOAuthReady", + requestId: oauthId, + providerID: "anthropic", + authorization: { url: "https://example.com", method: "code", instructions: "Code: 1234" }, + }) + + expect(seen).toEqual(["disconnect:openai", "oauth:code"]) + action.dispose() + }) + + it("can drop stale requests", () => { + const transport = createTransport() + const action = createProviderAction(transport) + const seen: string[] = [] + + const requestId = action.send( + { + type: "saveCustomProvider", + providerID: "myprovider", + config: { + name: "My Provider", + options: { baseURL: "https://example.com/v1" }, + models: { "model-1": { name: "Model One" } }, + }, + }, + { + onError: (message) => seen.push(message.message), + }, + ) + + action.clear(requestId) + transport.receive({ + type: "providerActionError", + requestId, + providerID: "myprovider", + action: "connect", + message: "boom", + }) + + expect(seen).toEqual([]) + action.dispose() + }) +}) diff --git a/packages/kilo-vscode/tests/unit/provider-utils.test.ts b/packages/kilo-vscode/tests/unit/provider-utils.test.ts index 315a589f9a..c8eab7e588 100644 --- a/packages/kilo-vscode/tests/unit/provider-utils.test.ts +++ b/packages/kilo-vscode/tests/unit/provider-utils.test.ts @@ -1,5 +1,5 @@ import { describe, it, expect } from "bun:test" -import { flattenModels, findModel } from "../../webview-ui/src/context/provider-utils" +import { flattenModels, findModel, isModelValid } from "../../webview-ui/src/context/provider-utils" import type { Provider } from "../../webview-ui/src/types/messages" function makeProvider(id: string, name: string, modelIds: string[]): Provider { @@ -78,3 +78,26 @@ describe("findModel", () => { expect(findModel([], { providerID: "openai", modelID: "gpt-4" })).toBeUndefined() }) }) + +describe("isModelValid", () => { + const providers = { + kilo: makeProvider("kilo", "Kilo Gateway", ["kilo-auto/free"]), + openai: makeProvider("openai", "OpenAI", ["gpt-4o"]), + } + + it("accepts a connected provider model", () => { + expect(isModelValid(providers, ["openai"], { providerID: "openai", modelID: "gpt-4o" })).toBe(true) + }) + + it("rejects a disconnected non-kilo provider", () => { + expect(isModelValid(providers, [], { providerID: "openai", modelID: "gpt-4o" })).toBe(false) + }) + + it("accepts kilo models when present in the catalog", () => { + expect(isModelValid(providers, [], { providerID: "kilo", modelID: "kilo-auto/free" })).toBe(true) + }) + + it("rejects unknown models", () => { + expect(isModelValid(providers, ["openai"], { providerID: "openai", modelID: "missing" })).toBe(false) + }) +}) diff --git a/packages/kilo-vscode/tests/unit/provider-visibility.test.ts b/packages/kilo-vscode/tests/unit/provider-visibility.test.ts new file mode 100644 index 0000000000..5b71ad7ef9 --- /dev/null +++ b/packages/kilo-vscode/tests/unit/provider-visibility.test.ts @@ -0,0 +1,23 @@ +import { describe, expect, it } from "bun:test" + +import { visibleConnectedIds } from "../../webview-ui/src/components/settings/provider-visibility" + +describe("visibleConnectedIds", () => { + it("hides Kilo from the connected list when auth is missing", () => { + const ids = visibleConnectedIds(["kilo", "openrouter"], { openrouter: "api" }) + + expect(ids).toEqual(["openrouter"]) + }) + + it("keeps Kilo in the connected list when auth exists", () => { + const ids = visibleConnectedIds(["kilo", "openrouter"], { kilo: "oauth", openrouter: "api" }) + + expect(ids).toEqual(["kilo", "openrouter"]) + }) + + it("leaves non-Kilo providers untouched", () => { + const ids = visibleConnectedIds(["anthropic"], {}) + + expect(ids).toEqual(["anthropic"]) + }) +}) diff --git a/packages/kilo-vscode/tests/visual-regression.spec.ts-snapshots/settings/providers-configure-chromium-linux.png b/packages/kilo-vscode/tests/visual-regression.spec.ts-snapshots/settings/providers-configure-chromium-linux.png index 58ecfff2a4..d4443c43e7 100644 --- a/packages/kilo-vscode/tests/visual-regression.spec.ts-snapshots/settings/providers-configure-chromium-linux.png +++ b/packages/kilo-vscode/tests/visual-regression.spec.ts-snapshots/settings/providers-configure-chromium-linux.png @@ -1,3 +1,3 @@ version https://git-lfs.github.com/spec/v1 -oid sha256:b7fe365d1a0ffdd713d3146c66112fad2207c2030132ab117543a6a430197aca -size 36247 +oid sha256:19d68a55c4a80d83ba70770744764ab33b6a5291ff02a34106a12d0935203ffa +size 20185 diff --git a/packages/kilo-vscode/tests/visual-regression.spec.ts-snapshots/settings/settings-panel-chromium-linux.png b/packages/kilo-vscode/tests/visual-regression.spec.ts-snapshots/settings/settings-panel-chromium-linux.png index 0412526b18..ed6030cc17 100644 --- a/packages/kilo-vscode/tests/visual-regression.spec.ts-snapshots/settings/settings-panel-chromium-linux.png +++ b/packages/kilo-vscode/tests/visual-regression.spec.ts-snapshots/settings/settings-panel-chromium-linux.png @@ -1,3 +1,3 @@ version https://git-lfs.github.com/spec/v1 -oid sha256:b83a16d340c4af266d4dd0134a840c7c07bed50837a58460d3cebab61ca2bd06 -size 45545 +oid sha256:3ad2458e2fe4e735c09fe89292c238f4b3113d2fa7dab5353889f0fc885400a7 +size 30038 diff --git a/packages/kilo-vscode/webview-ui/src/components/profile/DeviceAuthCard.tsx b/packages/kilo-vscode/webview-ui/src/components/profile/DeviceAuthCard.tsx index 4befc8fc06..aebcdec4f5 100644 --- a/packages/kilo-vscode/webview-ui/src/components/profile/DeviceAuthCard.tsx +++ b/packages/kilo-vscode/webview-ui/src/components/profile/DeviceAuthCard.tsx @@ -3,6 +3,8 @@ import { Button } from "@kilocode/kilo-ui/button" import { Card } from "@kilocode/kilo-ui/card" import { Spinner } from "@kilocode/kilo-ui/spinner" import { showToast } from "@kilocode/kilo-ui/toast" +import { useDialog } from "@kilocode/kilo-ui/context/dialog" +import { Dialog } from "@kilocode/kilo-ui/dialog" import { useVSCode } from "../../context/vscode" import { useLanguage } from "../../context/language" import { generateQRCode } from "../../utils/qrcode" @@ -18,15 +20,41 @@ interface DeviceAuthCardProps { onRetry: () => void } +const ERROR_LIMIT = 180 + const formatTime = (seconds: number): string => { const m = Math.floor(seconds / 60) const s = seconds % 60 return `${m}:${s.toString().padStart(2, "0")}` } +function compactError(error: string | undefined, fallback: string) { + if (!error) return fallback + + const text = error.replace(/\s+/g, " ").trim() + const html = text.search(/= 0 + ? text + .slice(0, html) + .trim() + .replace(/[\s:,-]+$/, "") + : text + + if (head.length > 0) { + if (head.length <= ERROR_LIMIT) return head + return `${head.slice(0, ERROR_LIMIT).trimEnd()}...` + } + + const status = text.match(/\b([45]\d{2})\b/)?.[1] + if (status) return `${fallback} (${status})` + return fallback +} + const DeviceAuthCard: Component = (props) => { const vscode = useVSCode() const language = useLanguage() + const dialog = useDialog() const [timeRemaining, setTimeRemaining] = createSignal(props.expiresIn ?? 900) const [qrDataUrl, setQrDataUrl] = createSignal("") @@ -70,6 +98,56 @@ const DeviceAuthCard: Component = (props) => { } } + const errorSummary = () => compactError(props.error, language.t("deviceAuth.status.failed")) + + const hasErrorDetails = () => { + if (!props.error) return false + return errorSummary() !== props.error.replace(/\s+/g, " ").trim() + } + + const handleCopyError = () => { + if (!props.error) return + navigator.clipboard.writeText(props.error) + showToast({ variant: "success", title: language.t("deviceAuth.toast.errorCopied") }) + } + + const handleShowError = () => { + if (!props.error) return + + dialog.show(() => ( + +
+
+            {props.error}
+          
+
+ + +
+
+
+ )) + } + return ( {/* Initiating state */} @@ -281,11 +359,23 @@ const DeviceAuthCard: Component = (props) => { margin: "8px 0 12px 0", }} > - {props.error || language.t("deviceAuth.status.failed")} + {errorSummary()}

- +
+ + + + + + + +
diff --git a/packages/kilo-vscode/webview-ui/src/components/settings/CustomProviderDialog.tsx b/packages/kilo-vscode/webview-ui/src/components/settings/CustomProviderDialog.tsx new file mode 100644 index 0000000000..a1debefb01 --- /dev/null +++ b/packages/kilo-vscode/webview-ui/src/components/settings/CustomProviderDialog.tsx @@ -0,0 +1,458 @@ +import { Button } from "@kilocode/kilo-ui/button" +import { useDialog } from "@kilocode/kilo-ui/context/dialog" +import { Dialog } from "@kilocode/kilo-ui/dialog" +import { IconButton } from "@kilocode/kilo-ui/icon-button" +import { ProviderIcon } from "@kilocode/kilo-ui/provider-icon" +import { TextField } from "@kilocode/kilo-ui/text-field" +import { showToast } from "@kilocode/kilo-ui/toast" +import { For, onCleanup } from "solid-js" +import { createStore } from "solid-js/store" +import { useConfig } from "../../context/config" +import { useLanguage } from "../../context/language" +import { useProvider } from "../../context/provider" +import { useVSCode } from "../../context/vscode" +import { createProviderAction } from "../../utils/provider-action" + +const PROVIDER_ID = /^[a-z0-9][a-z0-9-_]*$/ +const OPENAI_COMPATIBLE = "@ai-sdk/openai-compatible" + +type Translator = ReturnType["t"] + +type ModelRow = { + id: string + name: string +} + +type HeaderRow = { + key: string + value: string +} + +type FormState = { + providerID: string + name: string + baseURL: string + apiKey: string + models: ModelRow[] + headers: HeaderRow[] + saving: boolean +} + +type FormErrors = { + providerID: string | undefined + name: string | undefined + baseURL: string | undefined + models: Array<{ id?: string; name?: string }> + headers: Array<{ key?: string; value?: string }> +} + +type ValidateArgs = { + form: FormState + t: Translator + disabledProviders: string[] + existingProviderIDs: Set +} + +function validateCustomProvider(input: ValidateArgs) { + const providerID = input.form.providerID.trim() + const name = input.form.name.trim() + const baseURL = input.form.baseURL.trim() + const apiKey = input.form.apiKey.trim() + + const env = apiKey.match(/^\{env:([^}]+)\}$/)?.[1]?.trim() + const key = apiKey && !env ? apiKey : undefined + + const idError = !providerID + ? input.t("provider.custom.error.providerID.required") + : !PROVIDER_ID.test(providerID) + ? input.t("provider.custom.error.providerID.format") + : undefined + + const nameError = !name ? input.t("provider.custom.error.name.required") : undefined + const urlError = !baseURL + ? input.t("provider.custom.error.baseURL.required") + : !/^https?:\/\//.test(baseURL) + ? input.t("provider.custom.error.baseURL.format") + : undefined + + const disabled = input.disabledProviders.includes(providerID) + const existsError = idError + ? undefined + : input.existingProviderIDs.has(providerID) && !disabled + ? input.t("provider.custom.error.providerID.exists") + : undefined + + const seenModels = new Set() + const modelErrors = input.form.models.map((m) => { + const id = m.id.trim() + const modelIdError = !id + ? input.t("provider.custom.error.required") + : seenModels.has(id) + ? input.t("provider.custom.error.duplicate") + : (() => { + seenModels.add(id) + return undefined + })() + const modelNameError = !m.name.trim() ? input.t("provider.custom.error.required") : undefined + return { id: modelIdError, name: modelNameError } + }) + const modelsValid = modelErrors.every((m) => !m.id && !m.name) + const models = Object.fromEntries(input.form.models.map((m) => [m.id.trim(), { name: m.name.trim() }])) + + const seenHeaders = new Set() + const headerErrors = input.form.headers.map((h) => { + const key = h.key.trim() + const value = h.value.trim() + + if (!key && !value) return {} + const keyError = !key + ? input.t("provider.custom.error.required") + : seenHeaders.has(key.toLowerCase()) + ? input.t("provider.custom.error.duplicate") + : (() => { + seenHeaders.add(key.toLowerCase()) + return undefined + })() + const valueError = !value ? input.t("provider.custom.error.required") : undefined + return { key: keyError, value: valueError } + }) + const headersValid = headerErrors.every((h) => !h.key && !h.value) + const headers = Object.fromEntries( + input.form.headers + .map((h) => ({ key: h.key.trim(), value: h.value.trim() })) + .filter((h) => !!h.key && !!h.value) + .map((h) => [h.key, h.value]), + ) + + const errors: FormErrors = { + providerID: idError ?? existsError, + name: nameError, + baseURL: urlError, + models: modelErrors, + headers: headerErrors, + } + + const ok = !idError && !existsError && !nameError && !urlError && modelsValid && headersValid + if (!ok) return { errors } + + const options = { + baseURL, + ...(Object.keys(headers).length ? { headers } : {}), + } + + return { + errors, + result: { + providerID, + name, + key, + config: { + npm: OPENAI_COMPATIBLE, + name, + ...(env ? { env: [env] } : {}), + options, + models, + }, + }, + } +} + +interface CustomProviderDialogProps { + onBack?: () => void +} + +const CustomProviderDialog = (props: CustomProviderDialogProps) => { + const dialog = useDialog() + const { config } = useConfig() + const provider = useProvider() + const language = useLanguage() + const vscode = useVSCode() + const action = createProviderAction(vscode) + onCleanup(action.dispose) + + const [form, setForm] = createStore({ + providerID: "", + name: "", + baseURL: "", + apiKey: "", + models: [{ id: "", name: "" }], + headers: [{ key: "", value: "" }], + saving: false, + }) + + const [errors, setErrors] = createStore({ + providerID: undefined, + name: undefined, + baseURL: undefined, + models: [{}], + headers: [{}], + }) + + function goBack() { + if (props.onBack) { + props.onBack() + return + } + dialog.close() + } + + function addModel() { + setForm("models", (v) => [...v, { id: "", name: "" }]) + setErrors("models", (v) => [...v, {}]) + } + + function removeModel(index: number) { + if (form.models.length <= 1) return + setForm("models", (v) => v.filter((_, i) => i !== index)) + setErrors("models", (v) => v.filter((_, i) => i !== index)) + } + + function addHeader() { + setForm("headers", (v) => [...v, { key: "", value: "" }]) + setErrors("headers", (v) => [...v, {}]) + } + + function removeHeader(index: number) { + if (form.headers.length <= 1) return + setForm("headers", (v) => v.filter((_, i) => i !== index)) + setErrors("headers", (v) => v.filter((_, i) => i !== index)) + } + + function validate() { + const output = validateCustomProvider({ + form, + t: language.t, + disabledProviders: config().disabled_providers ?? [], + existingProviderIDs: new Set(Object.keys(provider.providers())), + }) + setErrors(output.errors) + return output.result + } + + function save(e: SubmitEvent) { + e.preventDefault() + if (form.saving) return + + const result = validate() + if (!result) return + + setForm("saving", true) + + action.send( + { + type: "saveCustomProvider", + providerID: result.providerID, + config: result.config, + apiKey: result.key, + }, + { + onConnected: () => { + setForm("saving", false) + dialog.close() + showToast({ + variant: "success", + icon: "circle-check", + title: language.t("provider.connect.toast.connected.title", { provider: result.name }), + description: language.t("provider.connect.toast.connected.description", { provider: result.name }), + }) + }, + onError: (message) => { + setForm("saving", false) + showToast({ title: language.t("common.requestFailed"), description: message.message }) + }, + }, + ) + } + + return ( + + } + transition + > +
+
+ +
+ {language.t("provider.custom.title")} +
+
+ +
+ + +
+ setForm("providerID", v)} + validationState={errors.providerID ? "invalid" : undefined} + error={errors.providerID} + /> + setForm("name", v)} + validationState={errors.name ? "invalid" : undefined} + error={errors.name} + /> + setForm("baseURL", v)} + validationState={errors.baseURL ? "invalid" : undefined} + error={errors.baseURL} + /> + setForm("apiKey", v)} + /> +
+ + {/* Models */} +
+ + + {(m, i) => ( +
+
+ setForm("models", i(), "id", v)} + validationState={errors.models[i()]?.id ? "invalid" : undefined} + error={errors.models[i()]?.id} + /> +
+
+ setForm("models", i(), "name", v)} + validationState={errors.models[i()]?.name ? "invalid" : undefined} + error={errors.models[i()]?.name} + /> +
+ removeModel(i())} + disabled={form.models.length <= 1} + aria-label={language.t("provider.custom.models.remove")} + style={{ "margin-top": "6px" }} + /> +
+ )} +
+ +
+ + {/* Headers */} +
+ + + {(h, i) => ( +
+
+ setForm("headers", i(), "key", v)} + validationState={errors.headers[i()]?.key ? "invalid" : undefined} + error={errors.headers[i()]?.key} + /> +
+
+ setForm("headers", i(), "value", v)} + validationState={errors.headers[i()]?.value ? "invalid" : undefined} + error={errors.headers[i()]?.value} + /> +
+ removeHeader(i())} + disabled={form.headers.length <= 1} + aria-label={language.t("provider.custom.headers.remove")} + style={{ "margin-top": "6px" }} + /> +
+ )} +
+ +
+ + +
+
+
+ ) +} + +export default CustomProviderDialog diff --git a/packages/kilo-vscode/webview-ui/src/components/settings/ModelsTab.tsx b/packages/kilo-vscode/webview-ui/src/components/settings/ModelsTab.tsx new file mode 100644 index 0000000000..8e0afb12fd --- /dev/null +++ b/packages/kilo-vscode/webview-ui/src/components/settings/ModelsTab.tsx @@ -0,0 +1,90 @@ +import { Component, For, createMemo } from "solid-js" +import { Card } from "@kilocode/kilo-ui/card" +import { useConfig } from "../../context/config" +import { useLanguage } from "../../context/language" +import { useSession } from "../../context/session" +import { parseModelString } from "../../../../src/shared/provider-model" +import { ModelSelectorBase } from "../shared/ModelSelector" +import SettingsRow from "./SettingsRow" + +const ModelsTab: Component = () => { + const { config, updateConfig } = useConfig() + const language = useLanguage() + const session = useSession() + + function handleModelSelect(configKey: "model" | "small_model") { + return (providerID: string, modelID: string) => { + if (!providerID || !modelID) { + updateConfig({ [configKey]: null }) + return + } + updateConfig({ [configKey]: `${providerID}/${modelID}` }) + } + } + + const allAgents = createMemo(() => session.agents()) + + function handleModeModelSelect(agentName: string) { + return (providerID: string, modelID: string) => { + if (!providerID || !modelID) { + updateConfig({ agent: { [agentName]: { model: null } } }) + return + } + updateConfig({ agent: { [agentName]: { model: `${providerID}/${modelID}` } } }) + } + } + + return ( +
+ + + + + + + + + +

{language.t("settings.providers.modeModels")}

+ + + {(agent, index) => ( + + + + )} + + +
+ ) +} + +export default ModelsTab diff --git a/packages/kilo-vscode/webview-ui/src/components/settings/ProviderConnectDialog.tsx b/packages/kilo-vscode/webview-ui/src/components/settings/ProviderConnectDialog.tsx new file mode 100644 index 0000000000..145580f7ab --- /dev/null +++ b/packages/kilo-vscode/webview-ui/src/components/settings/ProviderConnectDialog.tsx @@ -0,0 +1,406 @@ +import { Button } from "@kilocode/kilo-ui/button" +import { useDialog } from "@kilocode/kilo-ui/context/dialog" +import { Dialog } from "@kilocode/kilo-ui/dialog" +import { Spinner } from "@kilocode/kilo-ui/spinner" +import { TextField } from "@kilocode/kilo-ui/text-field" +import { showToast } from "@kilocode/kilo-ui/toast" +import type { ProviderAuthAuthorization, ProviderAuthMethod } from "@kilocode/sdk/v2/client" +import { Component, For, Match, Show, Switch, createMemo, createSignal, onCleanup, onMount } from "solid-js" +import { createStore } from "solid-js/store" +import { useLanguage } from "../../context/language" +import { useProvider } from "../../context/provider" +import { useVSCode } from "../../context/vscode" +import { createProviderAction } from "../../utils/provider-action" + +interface ProviderConnectDialogProps { + providerID: string +} + +interface ViewState { + methodIndex?: number + authorization?: ProviderAuthAuthorization + phase?: "authorizing" | "connecting" + error?: string + failed?: string +} + +function fallbackMethods(label: string): ProviderAuthMethod[] { + return [{ type: "api", label }] +} + +function formatError(value: unknown, fallback: string): string { + if (value && typeof value === "object" && "message" in value) { + const message = (value as { message?: unknown }).message + if (typeof message === "string" && message) return message + } + if (typeof value === "string" && value) return value + return fallback +} + +const ProviderConnectDialog: Component = (props) => { + const dialog = useDialog() + const language = useLanguage() + const provider = useProvider() + const vscode = useVSCode() + const action = createProviderAction(vscode) + + const [state, setState] = createStore({}) + + const item = createMemo(() => provider.providers()[props.providerID]) + const name = () => item()?.name ?? props.providerID + const methods = createMemo(() => { + return provider.authMethods()[props.providerID] ?? fallbackMethods(language.t("provider.connect.method.apiKey")) + }) + const method = createMemo(() => { + const index = state.methodIndex + return index === undefined ? undefined : methods()[index] + }) + + onCleanup(action.dispose) + + onMount(() => { + if (methods().length !== 1) return + selectMethod(0) + }) + + function openExternal(url: string) { + vscode.postMessage({ type: "openExternal", url }) + } + + function reset() { + action.clear() + setState({ + methodIndex: undefined, + authorization: undefined, + phase: undefined, + error: undefined, + failed: undefined, + }) + } + + function fail(message: string) { + const failed = state.authorization?.method === "auto" || state.phase === "authorizing" + setState({ + ...state, + phase: undefined, + error: failed ? undefined : message, + failed: failed ? message : undefined, + }) + } + + function succeed() { + showToast({ + variant: "success", + icon: "circle-check", + title: language.t("provider.connect.toast.connected.title", { provider: name() }), + description: language.t("provider.connect.toast.connected.description", { provider: name() }), + }) + dialog.close() + } + + function selectMethod(index: number) { + const current = methods()[index] + action.clear() + setState({ + methodIndex: index, + authorization: undefined, + phase: current?.type === "oauth" ? "authorizing" : undefined, + error: undefined, + failed: undefined, + }) + if (current?.type !== "oauth") return + + action.send( + { + type: "authorizeProviderOAuth", + providerID: props.providerID, + method: index, + }, + { + onOAuthReady: (message) => { + setState({ + ...state, + authorization: message.authorization, + phase: undefined, + error: undefined, + failed: undefined, + }) + }, + onError: (message) => fail(message.message), + }, + ) + } + + function connect(apiKey: string) { + setState({ + ...state, + phase: "connecting", + error: undefined, + failed: undefined, + }) + action.send( + { + type: "connectProvider", + providerID: props.providerID, + apiKey, + }, + { + onConnected: succeed, + onError: (message) => fail(message.message), + }, + ) + } + + function complete(code?: string) { + const index = state.methodIndex + if (index === undefined) return + + setState({ + ...state, + phase: "connecting", + error: undefined, + failed: undefined, + }) + action.send( + { + type: "completeProviderOAuth", + providerID: props.providerID, + method: index, + code, + }, + { + onConnected: succeed, + onError: (message) => fail(message.message), + }, + ) + } + + const title = () => language.t("provider.connect.title", { provider: name() }) + + const MethodSelection: Component = () => { + return ( +
+
{language.t("provider.connect.selectMethod", { provider: name() })}
+
+ + {(item, index) => ( + + )} + +
+
+ +
+
+ ) + } + + const ApiView: Component = () => { + const [value, setValue] = createSignal("") + + function submit(e: SubmitEvent) { + e.preventDefault() + const apiKey = value().trim() + if (!apiKey) { + setState({ ...state, error: language.t("provider.connect.apiKey.required") }) + return + } + connect(apiKey) + } + + return ( +
+
+ {language.t("provider.connect.apiKey.description", { provider: name() })} +
+ +
+ + +
+ + ) + } + + const OAuthCodeView: Component = () => { + const [value, setValue] = createSignal("") + + onMount(() => { + if (!state.authorization?.url) return + openExternal(state.authorization.url) + }) + + function submit(e: SubmitEvent) { + e.preventDefault() + const code = value().trim() + if (!code) { + setState({ ...state, error: language.t("provider.connect.oauth.code.required") }) + return + } + complete(code) + } + + return ( +
+
+ {language.t("provider.connect.oauth.code.visit.prefix")} + { + e.preventDefault() + if (!state.authorization?.url) return + openExternal(state.authorization.url) + }} + > + {language.t("provider.connect.oauth.code.visit.link")} + + {language.t("provider.connect.oauth.code.visit.suffix", { provider: name() })} +
+ +
+ + +
+ + ) + } + + const OAuthAutoView: Component = () => { + const code = createMemo(() => { + const instructions = state.authorization?.instructions + if (!instructions) return "" + if (!instructions.includes(":")) return instructions + return instructions.split(":")[1]?.trim() ?? instructions + }) + + onMount(() => { + if (state.authorization?.url) openExternal(state.authorization.url) + complete() + }) + + return ( +
+
+ {language.t("provider.connect.oauth.auto.visit.prefix")} + { + e.preventDefault() + if (!state.authorization?.url) return + openExternal(state.authorization.url) + }} + > + {language.t("provider.connect.oauth.auto.visit.link")} + + {language.t("provider.connect.oauth.auto.visit.suffix", { provider: name() })} +
+ +
+
{language.t("provider.connect.oauth.auto.confirmationCode")}
+
{code()}
+
+
+
+ + + {state.error + ? language.t("provider.connect.status.failed", { error: state.error }) + : language.t("provider.connect.status.waiting")} + +
+
+ +
+
+ ) + } + + return ( + + + + + + +
+
+ + {language.t("provider.connect.status.inProgress")} +
+
+
+ +
+
{formatError(state.failed, language.t("common.requestFailed"))}
+
+ +
+
+
+ + + + + + + + + + +
+
{formatError(state.error ?? state.failed, language.t("common.requestFailed"))}
+
+ +
+
+
+
+
+ ) +} + +export default ProviderConnectDialog diff --git a/packages/kilo-vscode/webview-ui/src/components/settings/ProviderSelectDialog.tsx b/packages/kilo-vscode/webview-ui/src/components/settings/ProviderSelectDialog.tsx new file mode 100644 index 0000000000..e9085add8d --- /dev/null +++ b/packages/kilo-vscode/webview-ui/src/components/settings/ProviderSelectDialog.tsx @@ -0,0 +1,134 @@ +import { useDialog } from "@kilocode/kilo-ui/context/dialog" +import { Dialog } from "@kilocode/kilo-ui/dialog" +import { List } from "@kilocode/kilo-ui/list" +import { ProviderIcon } from "@kilocode/kilo-ui/provider-icon" +import { Tag } from "@kilocode/kilo-ui/tag" +import { Show, createMemo } from "solid-js" +import { useConfig } from "../../context/config" +import { useLanguage } from "../../context/language" +import { useProvider } from "../../context/provider" +import { useServer } from "../../context/server" +import type { Provider } from "../../types/messages" +import ProviderConnectDialog from "./ProviderConnectDialog" +import { + CUSTOM_PROVIDER_ID, + isPopularProvider, + kiloFallbackProvider, + popularProviderIndex, + providerIcon, +} from "./provider-catalog" +import CustomProviderDialog from "./CustomProviderDialog" +import { KILO_PROVIDER_ID } from "../../../../src/shared/provider-model" + +type ProviderItem = { + id: string + name: string +} + +const ProviderSelectDialog = () => { + const dialog = useDialog() + const { config } = useConfig() + const provider = useProvider() + const server = useServer() + const language = useLanguage() + + const items = createMemo(() => { + language.locale() + + const disabled = new Set(config().disabled_providers ?? []) + const connected = new Set(provider.connected()) + const all = Object.values(provider.providers()) + const withKilo = all.some((item) => item.id === KILO_PROVIDER_ID) ? all : [kiloFallbackProvider(), ...all] + const available = withKilo.filter((item) => !disabled.has(item.id) && !connected.has(item.id)) + + return [ + { + id: CUSTOM_PROVIDER_ID, + name: language.t("settings.providers.tag.customProvider"), + }, + ...available.map((item) => ({ + id: item.id, + name: item.name, + })), + ] + }) + + function open(item: ProviderItem) { + if (item.id === CUSTOM_PROVIDER_ID) { + dialog.show(() => dialog.show(() => )} />) + return + } + + if (item.id === KILO_PROVIDER_ID) { + dialog.close() + server.startLogin() + return + } + + dialog.show(() => ) + } + + return ( + + + search={{ placeholder: language.t("dialog.provider.search.placeholder"), autofocus: true }} + emptyMessage={language.t("dialog.provider.empty")} + activeIcon="plus-small" + key={(item) => item.id} + items={items()} + filterKeys={["id", "name"]} + groupBy={(item) => + item.id !== CUSTOM_PROVIDER_ID && isPopularProvider(item.id) + ? language.t("dialog.provider.group.recommended") + : language.t("dialog.provider.group.other") + } + sortBy={(a, b) => { + if (a.id === CUSTOM_PROVIDER_ID) return -1 + if (b.id === CUSTOM_PROVIDER_ID) return 1 + + const rank = popularProviderIndex(a.id) - popularProviderIndex(b.id) + if (rank !== 0) return rank + return a.name.localeCompare(b.name) + }} + sortGroupsBy={(a, b) => { + const recommended = language.t("dialog.provider.group.recommended") + if (a.category === recommended && b.category !== recommended) return -1 + if (b.category === recommended && a.category !== recommended) return 1 + return 0 + }} + onSelect={(item) => { + if (!item) return + open(item) + }} + > + {(item) => ( +
+ +
+ + {item.name} + + + {language.t("dialog.provider.tag.recommended")} + + + {language.t("settings.providers.tag.custom")} + +
+
+ )} + +
+ ) +} + +export default ProviderSelectDialog diff --git a/packages/kilo-vscode/webview-ui/src/components/settings/ProvidersTab.tsx b/packages/kilo-vscode/webview-ui/src/components/settings/ProvidersTab.tsx index cf7f642782..4f4408b773 100644 --- a/packages/kilo-vscode/webview-ui/src/components/settings/ProvidersTab.tsx +++ b/packages/kilo-vscode/webview-ui/src/components/settings/ProvidersTab.tsx @@ -1,218 +1,315 @@ -import { Component, For, createSignal, createMemo } from "solid-js" -import { Select } from "@kilocode/kilo-ui/select" -import { Card } from "@kilocode/kilo-ui/card" import { Button } from "@kilocode/kilo-ui/button" -import { IconButton } from "@kilocode/kilo-ui/icon-button" +import { Card } from "@kilocode/kilo-ui/card" +import { useDialog } from "@kilocode/kilo-ui/context/dialog" import { Icon } from "@kilocode/kilo-ui/icon" +import { ProviderIcon } from "@kilocode/kilo-ui/provider-icon" +import { Tag } from "@kilocode/kilo-ui/tag" +import { showToast } from "@kilocode/kilo-ui/toast" +import { Component, For, Show, createMemo, onCleanup } from "solid-js" import { useConfig } from "../../context/config" -import { useProvider } from "../../context/provider" import { useLanguage } from "../../context/language" -import { useSession } from "../../context/session" -import { ModelSelectorBase } from "../shared/ModelSelector" -import type { ModelSelection } from "../../types/messages" -import SettingsRow from "./SettingsRow" +import { useProvider } from "../../context/provider" +import { useServer } from "../../context/server" +import { useVSCode } from "../../context/vscode" +import type { Provider } from "../../types/messages" +import CustomProviderDialog from "./CustomProviderDialog" +import ProviderConnectDialog from "./ProviderConnectDialog" +import ProviderSelectDialog from "./ProviderSelectDialog" +import { CUSTOM_PROVIDER_ID, isPopularProvider, providerIcon, providerNoteKey, sortProviders } from "./provider-catalog" +import { visibleConnectedIds } from "./provider-visibility" +import { KILO_PROVIDER_ID } from "../../../../src/shared/provider-model" +import { createProviderAction } from "../../utils/provider-action" -interface ProviderOption { - value: string - label: string -} - -/** Parse a "provider/model" config string into a ModelSelection (or null). */ -function parseModelConfig(raw: string | undefined): ModelSelection | null { - if (!raw) { - return null - } - const slash = raw.indexOf("/") - if (slash <= 0) { - return null - } - return { providerID: raw.slice(0, slash), modelID: raw.slice(slash + 1) } -} +type ProviderSource = "env" | "api" | "config" | "custom" const ProvidersTab: Component = () => { - const { config, updateConfig } = useConfig() + const dialog = useDialog() + const { config } = useConfig() const provider = useProvider() const language = useLanguage() - const session = useSession() + const server = useServer() + const vscode = useVSCode() + const action = createProviderAction(vscode) - const providerOptions = createMemo(() => - Object.keys(provider.providers()) - .sort() - .map((id) => ({ value: id, label: id })), - ) + onCleanup(action.dispose) - const [newDisabled, setNewDisabled] = createSignal() + const kiloLoggedIn = createMemo(() => !!server.profileData()) - const disabledProviders = () => config().disabled_providers ?? [] + const connectedProviders = createMemo(() => { + const ids = visibleConnectedIds(provider.connected(), provider.authStates()) + const all = provider.providers() + return ids + .filter((id) => id !== KILO_PROVIDER_ID) + .map((id) => all[id]) + .filter((item): item is Provider => !!item) + }) - const addDisabled = (value: string) => { - const current = [...disabledProviders()] - if (value && !current.includes(value)) { - current.push(value) - updateConfig({ disabled_providers: current }) - } + const popularProviders = createMemo(() => { + const connected = new Set(provider.connected()) + const disabled = new Set(config().disabled_providers ?? []) + const all = Object.values(provider.providers()) + return sortProviders( + all.filter( + (item) => + item.id !== KILO_PROVIDER_ID && + isPopularProvider(item.id) && + !connected.has(item.id) && + !disabled.has(item.id), + ), + ) + }) + + function source(item: Provider): ProviderSource | undefined { + if (!("source" in item)) return + const value = (item as Provider & { source?: string }).source + if (value === "env" || value === "api" || value === "config" || value === "custom") return value + return } - const removeDisabled = (index: number) => { - const current = [...disabledProviders()] - current.splice(index, 1) - updateConfig({ disabled_providers: current }) + function sourceTag(item: Provider) { + const current = source(item) + if (current === "env") return language.t("settings.providers.tag.environment") + if (current === "api") return language.t("provider.connect.method.apiKey") + if (current === "config") { + const cfg = config().provider?.[item.id] + if (cfg?.npm === "@ai-sdk/openai-compatible") return language.t("settings.providers.tag.custom") + return language.t("settings.providers.tag.config") + } + if (current === "custom") return language.t("settings.providers.tag.custom") + return language.t("settings.providers.tag.other") } - function handleModelSelect(configKey: "model" | "small_model") { - return (providerID: string, modelID: string) => { - if (!providerID || !modelID) { - updateConfig({ [configKey]: null }) - } else { - updateConfig({ [configKey]: `${providerID}/${modelID}` }) - } - } + function canDisconnect(item: Provider) { + return source(item) !== "env" } - const allAgents = createMemo(() => session.agents()) + function disconnect(providerID: string, name: string) { + action.send( + { type: "disconnectProvider", providerID }, + { + onDisconnected: () => { + showToast({ + variant: "success", + icon: "circle-check", + title: language.t("provider.disconnect.toast.disconnected.title", { provider: name }), + description: language.t("provider.disconnect.toast.disconnected.description", { provider: name }), + }) + }, + onError: (message) => { + showToast({ title: language.t("common.requestFailed"), description: message.message }) + }, + }, + ) + } - function handleModeModelSelect(agentName: string) { - return (providerID: string, modelID: string) => { - if (!providerID || !modelID) { - updateConfig({ agent: { [agentName]: { model: null } } }) - } else { - updateConfig({ agent: { [agentName]: { model: `${providerID}/${modelID}` } } }) - } + function connectProvider(item: Provider) { + if (item.id === KILO_PROVIDER_ID) { + server.startLogin() + return } + dialog.show(() => ) } return (
- {/* Model selection */} + {/* Kilo Gateway — always at the top, not editable */} - - - - - - - - - {/* Model per Mode */} -

{language.t("settings.providers.modeModels")}

- - - {(agent, index) => ( - - - - )} - - - - {/* Beta notice */} - - -

{language.t("settings.providers.betaNotice")}

-
- - {/* Disabled providers */} -

{language.t("settings.providers.disabled")}

- -
- {language.t("settings.providers.disabled.description")} -
0 ? "1px solid var(--border-weak-base)" : "none", + gap: "12px", + "min-height": "56px", + padding: "12px 0", }} > -
-