diff --git a/proto/cline/models.proto b/proto/cline/models.proto index 44c0ccdae78..52b78363492 100644 --- a/proto/cline/models.proto +++ b/proto/cline/models.proto @@ -359,6 +359,7 @@ message ModelsApiConfiguration { optional OpenRouterModelInfo plan_mode_vercel_ai_gateway_model_info = 130; optional string plan_mode_oca_model_id = 131; optional OcaModelInfo plan_mode_oca_model_info = 132; + repeated string plan_mode_oca_vector_ids = 133; // Act mode configurations @@ -395,4 +396,5 @@ message ModelsApiConfiguration { optional OpenRouterModelInfo act_mode_vercel_ai_gateway_model_info = 230; optional string act_mode_oca_model_id = 231; optional OcaModelInfo act_mode_oca_model_info = 232; + repeated string act_mode_oca_vector_ids = 233; } diff --git a/proto/cline/vectors.proto b/proto/cline/vectors.proto new file mode 100644 index 00000000000..f17b3d376e9 --- /dev/null +++ b/proto/cline/vectors.proto @@ -0,0 +1,23 @@ +syntax = "proto3"; + +package cline; +import "cline/common.proto"; +option java_package = "bot.cline.proto"; +option java_multiple_files = true; + +// Service for vector-related operations +service VectorsService { + // Fetches available vectors from OCA + rpc refreshOcaVectors(StringRequest) returns (VectorStores); +} + +message VectorStoreInfo { + string id = 1; + string name = 2; + string description = 3; +} + +message VectorStores { + map vectors = 1; + optional string error = 2; +} diff --git a/src/core/api/index.ts b/src/core/api/index.ts index a02508e554f..b7b494e57b9 100644 --- a/src/core/api/index.ts +++ b/src/core/api/index.ts @@ -386,6 +386,7 @@ function createHandlerForProvider( ? options.planModeOcaModelInfo?.supportsPromptCache : options.actModeOcaModelInfo?.supportsPromptCache, taskId: options.ulid, + vectorIds: mode === "plan" ? options.planModeOcaVectorIds : options.actModeOcaVectorIds, }) default: return new AnthropicHandler({ diff --git a/src/core/api/providers/oca.ts b/src/core/api/providers/oca.ts index bbc36c0933d..53037613c92 100644 --- a/src/core/api/providers/oca.ts +++ b/src/core/api/providers/oca.ts @@ -18,6 +18,7 @@ export interface OcaHandlerOptions extends CommonApiHandlerOptions { thinkingBudgetTokens?: number ocaUsePromptCache?: boolean taskId?: string + vectorIds?: string[] } export class OcaHandler implements ApiHandler { @@ -180,7 +181,15 @@ export class OcaHandler implements ApiHandler { return message }) - const stream = await client.chat.completions.create({ + const tools: OpenAI.Chat.Completions.ChatCompletionTool[] = [] + if (this.getVectorStores().length > 0) { + tools.push({ + type: "file_search", + vector_store_ids: this.getVectorStores(), + } as any) + } + + const requestObject: OpenAI.Chat.Completions.ChatCompletionCreateParamsStreaming = { model: this.options.ocaModelId || liteLlmDefaultModelId, messages: [enhancedSystemMessage, ...enhancedMessages], temperature, @@ -192,7 +201,12 @@ export class OcaHandler implements ApiHandler { ...(this.options.taskId && { litellm_session_id: `cline-${this.options.taskId}`, }), // Add session ID for LiteLLM tracking - }) + tools, + } + + console.log("Input to OCA chat completions: ", requestObject) + + const stream = await client.chat.completions.create(requestObject) const inputCost = (await this.calculateCost(1e6, 0)) || 0 const outputCost = (await this.calculateCost(0, 1e6)) || 0 @@ -258,4 +272,12 @@ export class OcaHandler implements ApiHandler { info: this.options.ocaModelInfo || liteLlmModelInfoSaneDefaults, } } + + getVectorStores() { + if (this.options.vectorIds) { + return this.options.vectorIds + } else { + return [] + } + } } diff --git a/src/core/controller/models/refreshOcaModels.ts b/src/core/controller/models/refreshOcaModels.ts index d07587d2320..26fc9d3889c 100644 --- a/src/core/controller/models/refreshOcaModels.ts +++ b/src/core/controller/models/refreshOcaModels.ts @@ -39,6 +39,7 @@ export async function refreshOcaModels(controller: Controller, request: StringRe Logger.log(`Making refresh oca model request with customer opc-request-id: ${headers["opc-request-id"]}`) const response = await axios.get(modelsUrl, { headers, ...getAxiosSettings() }) if (response.data?.data) { + console.log("Model response: ", response.data) if (response.data.data.length === 0) { HostProvider.window.showMessage({ type: ShowMessageType.ERROR, diff --git a/src/core/controller/vectors/refreshOcaVectors.ts b/src/core/controller/vectors/refreshOcaVectors.ts new file mode 100644 index 00000000000..b412e6864cf --- /dev/null +++ b/src/core/controller/vectors/refreshOcaVectors.ts @@ -0,0 +1,103 @@ +import { StringRequest } from "@shared/proto/cline/common" +import { VectorStoreInfo, VectorStores } from "@shared/proto/cline/vectors" +import axios from "axios" +import { HostProvider } from "@/hosts/host-provider" +import { OcaAuthService } from "@/services/auth/oca/OcaAuthService" +import { DEFAULT_OCA_BASE_URL } from "@/services/auth/oca/utils/constants" +import { createOcaHeaders, getAxiosSettings } from "@/services/auth/oca/utils/utils" +import { Logger } from "@/services/logging/Logger" +import { ShowMessageType } from "@/shared/proto/index.host" +import { Controller } from ".." + +/** + * Refreshes the Oca models and returns the updated model list + * @param controller The controller instance + * @param request Empty request object + * @returns Response containing the Oca models + */ +export async function refreshOcaVectors(controller: Controller, request: StringRequest): Promise { + const vectors: Record = {} + const ocaAccessToken = await OcaAuthService.getInstance().getAuthToken() + if (!ocaAccessToken) { + HostProvider.window.showMessage({ + type: ShowMessageType.ERROR, + message: "Not authenticated with OCA. Please sign in first.", + }) + return VectorStores.create({ error: "Not authenticated with OCA" }) + } + const baseUrl = request.value || DEFAULT_OCA_BASE_URL + const vectorsUrl = `${baseUrl}/vector_store/list` + const headers = await createOcaHeaders(ocaAccessToken!, "vectors-refresh") + try { + Logger.log(`Making refresh oca vector request with customer opc-request-id: ${headers["opc-request-id"]}`) + const response = await axios.get(vectorsUrl, { headers, ...getAxiosSettings() }) + if (response.data && response.data.data) { + const vectorIds: string[] = [] + for (const vectorStore of response.data.data) { + const vectorStoreId = vectorStore.vector_store_id + if (typeof vectorStoreId !== "string" || !vectorStoreId) { + continue + } + vectors[vectorStoreId] = VectorStoreInfo.create({ + id: vectorStoreId, + name: vectorStore.vector_store_name, + description: vectorStore.vector_store_description, + }) + vectorIds.push(vectorStoreId) + } + console.log("Oca vectors fetched", vectors) + + // Fetch current config + const apiConfiguration = controller.stateManager.getApiConfiguration() + const updatedConfig = { ...apiConfiguration } + + // Which mode(s) to update? + const planActSeparateModelsSetting = controller.stateManager.getGlobalSettingsKey("planActSeparateModelsSetting") + const currentMode = controller.stateManager.getGlobalSettingsKey("mode") + const planModeSelectedVectorId: string[] = apiConfiguration?.planModeOcaVectorIds + ? apiConfiguration?.planModeOcaVectorIds.filter( + (vectorId) => vectorIds.filter((secondVectorId) => vectorId === secondVectorId).length >= 1, + ) + : [] + const actModeSelectedVectorId: string[] = apiConfiguration?.actModeOcaVectorIds + ? apiConfiguration?.actModeOcaVectorIds.filter( + (vectorId) => vectorIds.filter((secondVectorId) => vectorId === secondVectorId).length >= 1, + ) + : [] + + // Save new model selection(s) to configuration object, per plan/act mode setting + if (planActSeparateModelsSetting) { + if (currentMode === "plan") { + updatedConfig.planModeOcaVectorIds = planModeSelectedVectorId + } else { + updatedConfig.actModeOcaVectorIds = actModeSelectedVectorId + } + } else { + updatedConfig.planModeOcaVectorIds = planModeSelectedVectorId + updatedConfig.actModeOcaVectorIds = actModeSelectedVectorId + } + + controller.stateManager.setApiConfiguration(updatedConfig) + + HostProvider.window.showMessage({ + type: ShowMessageType.INFORMATION, + message: `Refreshed Oca knowledge bases from ${baseUrl}`, + }) + } else { + console.error("Invalid response from oca API") + HostProvider.window.showMessage({ + type: ShowMessageType.INFORMATION, + message: `Failed to fetch Oca vectors. Please check your configuration from ${baseUrl}`, + }) + } + } catch (error) { + console.error("Error fetching oca vectors:", error) + const errorMsg = error.message || "Error refreshing Oca knowledge bases" + HostProvider.window.showMessage({ + type: ShowMessageType.ERROR, + message: errorMsg, + }) + return VectorStores.create({ error: errorMsg }) + } + return VectorStores.create({ vectors }) +} diff --git a/src/core/storage/StateManager.ts b/src/core/storage/StateManager.ts index 87b8d095b35..d84114239a0 100644 --- a/src/core/storage/StateManager.ts +++ b/src/core/storage/StateManager.ts @@ -446,6 +446,7 @@ export class StateManager { planModeVercelAiGatewayModelInfo, planModeOcaModelId, planModeOcaModelInfo, + planModeOcaVectorIds, // Act mode configurations actModeApiProvider, actModeApiModelId, @@ -480,6 +481,7 @@ export class StateManager { actModeVercelAiGatewayModelInfo, actModeOcaModelId, actModeOcaModelInfo, + actModeOcaVectorIds, } = apiConfiguration // Batch update global state keys @@ -518,6 +520,7 @@ export class StateManager { planModeVercelAiGatewayModelInfo, planModeOcaModelId, planModeOcaModelInfo, + planModeOcaVectorIds, // Act mode configuration updates actModeApiProvider, @@ -553,6 +556,7 @@ export class StateManager { actModeVercelAiGatewayModelInfo, actModeOcaModelId, actModeOcaModelInfo, + actModeOcaVectorIds, // Global state updates awsRegion, @@ -993,6 +997,7 @@ export class StateManager { this.globalStateCache["planModeVercelAiGatewayModelInfo"], planModeOcaModelId: this.globalStateCache["planModeOcaModelId"], planModeOcaModelInfo: this.globalStateCache["planModeOcaModelInfo"], + planModeOcaVectorIds: this.globalStateCache["planModeOcaVectorIds"], // Act mode configurations actModeApiProvider: this.taskStateCache["actModeApiProvider"] || this.globalStateCache["actModeApiProvider"], @@ -1055,6 +1060,7 @@ export class StateManager { this.globalStateCache["actModeVercelAiGatewayModelInfo"], actModeOcaModelId: this.globalStateCache["actModeOcaModelId"], actModeOcaModelInfo: this.globalStateCache["actModeOcaModelInfo"], + actModeOcaVectorIds: this.globalStateCache["actModeOcaVectorIds"], } } } diff --git a/src/core/storage/state-keys.ts b/src/core/storage/state-keys.ts index 992100302eb..19b6750319f 100644 --- a/src/core/storage/state-keys.ts +++ b/src/core/storage/state-keys.ts @@ -135,6 +135,7 @@ export interface Settings { planModeHuaweiCloudMaasModelInfo: ModelInfo | undefined planModeOcaModelId: string | undefined planModeOcaModelInfo: OcaModelInfo | undefined + planModeOcaVectorIds: string[] // Act mode configurations actModeApiProvider: ApiProvider actModeApiModelId: string | undefined @@ -171,6 +172,7 @@ export interface Settings { actModeVercelAiGatewayModelInfo: ModelInfo | undefined actModeOcaModelId: string | undefined actModeOcaModelInfo: OcaModelInfo | undefined + actModeOcaVectorIds: string[] } export interface Secrets { diff --git a/src/core/storage/utils/state-helpers.ts b/src/core/storage/utils/state-helpers.ts index e5d5f228eb0..f13dfb71189 100644 --- a/src/core/storage/utils/state-helpers.ts +++ b/src/core/storage/utils/state-helpers.ts @@ -314,6 +314,7 @@ export async function readGlobalStateFromDisk(context: ExtensionContext): Promis >("planModeVercelAiGatewayModelInfo") const planModeOcaModelId = context.globalState.get("planModeOcaModelId") as string | undefined const planModeOcaModelInfo = context.globalState.get("planModeOcaModelInfo") as OcaModelInfo | undefined + const planModeOcaVectorIds = context.globalState.get("planModeOcaVectorIds") as string[] | [] // Act mode configurations const actModeApiProvider = context.globalState.get("actModeApiProvider") const actModeApiModelId = context.globalState.get("actModeApiModelId") @@ -380,6 +381,7 @@ export async function readGlobalStateFromDisk(context: ExtensionContext): Promis >("actModeVercelAiGatewayModelInfo") const actModeOcaModelId = context.globalState.get("actModeOcaModelId") as string | undefined const actModeOcaModelInfo = context.globalState.get("actModeOcaModelInfo") as OcaModelInfo | undefined + const actModeOcaVectorIds = context.globalState.get("actModeOcaVectorIds") as string[] | [] const sapAiCoreUseOrchestrationMode = context.globalState.get("sapAiCoreUseOrchestrationMode") @@ -492,6 +494,7 @@ export async function readGlobalStateFromDisk(context: ExtensionContext): Promis planModeVercelAiGatewayModelInfo, planModeOcaModelId, planModeOcaModelInfo, + planModeOcaVectorIds, // Act mode configurations actModeApiProvider: actModeApiProvider || apiProvider, actModeApiModelId, @@ -526,6 +529,7 @@ export async function readGlobalStateFromDisk(context: ExtensionContext): Promis actModeVercelAiGatewayModelInfo, actModeOcaModelId, actModeOcaModelInfo, + actModeOcaVectorIds, // Other global fields focusChainSettings: focusChainSettings || DEFAULT_FOCUS_CHAIN_SETTINGS, diff --git a/src/services/auth/oca/providers/OcaAuthProvider.ts b/src/services/auth/oca/providers/OcaAuthProvider.ts index 7e2087b5d8a..f81901e9127 100644 --- a/src/services/auth/oca/providers/OcaAuthProvider.ts +++ b/src/services/auth/oca/providers/OcaAuthProvider.ts @@ -104,10 +104,12 @@ export class OcaAuthProvider { headers: { "Content-Type": "application/x-www-form-urlencoded" }, ...getAxiosSettings(), }) + console.log("Successful response: ", tokenResponse) const accessToken = tokenResponse.data.access_token const userInfo: OcaUserInfo = await this.getUserAccountInfo(accessToken) return { user: userInfo, apiKey: accessToken } } catch (err: unknown) { + console.log(err) const isAxios = (axios as any)?.isAxiosError?.(err) const status = isAxios ? (err as any).response?.status : undefined const data: any = isAxios ? (err as any).response?.data : undefined @@ -161,6 +163,7 @@ export class OcaAuthProvider { OcaAuthProvider.pkceStateMap.delete(state) const discovery = await axios.get(`${idcs_url}/.well-known/openid-configuration`, { ...getAxiosSettings() }) const tokenEndpoint = discovery.data.token_endpoint + console.log("Sign In Token Endpoint: ", tokenEndpoint) const params: any = { grant_type: "authorization_code", code, @@ -183,7 +186,7 @@ export class OcaAuthProvider { throw new Error("No ID token received from OCA") } - // Step 2: Get access_token (this is what you'll use for APIs) + //Step 2: Get access_token (this is what you'll use for APIs) const accessToken = tokenResponse.data.access_token const refreshToken = tokenResponse.data.refresh_token if (refreshToken) { diff --git a/src/shared/api.ts b/src/shared/api.ts index cdaeced85e0..348f33eab05 100644 --- a/src/shared/api.ts +++ b/src/shared/api.ts @@ -153,6 +153,7 @@ export interface ApiHandlerOptions { planModeVercelAiGatewayModelInfo?: ModelInfo planModeOcaModelId?: string planModeOcaModelInfo?: OcaModelInfo + planModeOcaVectorIds?: string[] // Act mode configurations // Act mode configurations @@ -188,6 +189,7 @@ export interface ApiHandlerOptions { actModeVercelAiGatewayModelInfo?: ModelInfo actModeOcaModelId?: string actModeOcaModelInfo?: OcaModelInfo + actModeOcaVectorIds?: string[] } export type ApiConfiguration = ApiHandlerOptions & diff --git a/src/shared/proto-conversions/models/api-configuration-conversion.ts b/src/shared/proto-conversions/models/api-configuration-conversion.ts index 92ea00032cd..e29c71f5acb 100644 --- a/src/shared/proto-conversions/models/api-configuration-conversion.ts +++ b/src/shared/proto-conversions/models/api-configuration-conversion.ts @@ -134,6 +134,13 @@ function convertProtoOcaModelInfoToOcaModelInfo(info: ProtoOcaModelInfo | undefi } } +function convertOcaVectorIdsToProtoOcaVectorIds(vectorIds: string[] | undefined): string[] { + if (!vectorIds) { + return [] + } + return vectorIds +} + // Convert application LiteLLMModelInfo to proto LiteLLMModelInfo function convertLiteLLMModelInfoToProto(info: AppLiteLLMModelInfo | undefined): LiteLLMModelInfo | undefined { if (!info) { @@ -503,6 +510,7 @@ export function convertApiConfigurationToProto(config: ApiConfiguration): ProtoA planModeVercelAiGatewayModelInfo: convertModelInfoToProtoOpenRouter(config.planModeVercelAiGatewayModelInfo), planModeOcaModelId: config.planModeOcaModelId, planModeOcaModelInfo: convertOcaModelInfoToProtoOcaModelInfo(config.planModeOcaModelInfo), + planModeOcaVectorIds: convertOcaVectorIdsToProtoOcaVectorIds(config.planModeOcaVectorIds), // Act mode configurations actModeApiProvider: config.actModeApiProvider ? convertApiProviderToProto(config.actModeApiProvider) : undefined, @@ -538,6 +546,7 @@ export function convertApiConfigurationToProto(config: ApiConfiguration): ProtoA actModeVercelAiGatewayModelInfo: convertModelInfoToProtoOpenRouter(config.actModeVercelAiGatewayModelInfo), actModeOcaModelId: config.actModeOcaModelId, actModeOcaModelInfo: convertOcaModelInfoToProtoOcaModelInfo(config.actModeOcaModelInfo), + actModeOcaVectorIds: convertOcaVectorIdsToProtoOcaVectorIds(config.actModeOcaVectorIds), } } @@ -655,6 +664,7 @@ export function convertProtoToApiConfiguration(protoConfig: ProtoApiConfiguratio planModeVercelAiGatewayModelInfo: convertProtoToModelInfo(protoConfig.planModeVercelAiGatewayModelInfo), planModeOcaModelId: protoConfig.planModeOcaModelId, planModeOcaModelInfo: convertProtoOcaModelInfoToOcaModelInfo(protoConfig.planModeOcaModelInfo), + planModeOcaVectorIds: protoConfig.planModeOcaVectorIds, // Act mode configurations actModeApiProvider: @@ -691,5 +701,6 @@ export function convertProtoToApiConfiguration(protoConfig: ProtoApiConfiguratio actModeVercelAiGatewayModelInfo: convertProtoToModelInfo(protoConfig.actModeVercelAiGatewayModelInfo), actModeOcaModelId: protoConfig.actModeOcaModelId, actModeOcaModelInfo: convertProtoOcaModelInfoToOcaModelInfo(protoConfig.actModeOcaModelInfo), + actModeOcaVectorIds: protoConfig.actModeOcaVectorIds, } } diff --git a/webview-ui/src/components/settings/providers/OcaProvider.tsx b/webview-ui/src/components/settings/providers/OcaProvider.tsx index ebcfa701b89..b8c83a49d65 100644 --- a/webview-ui/src/components/settings/providers/OcaProvider.tsx +++ b/webview-ui/src/components/settings/providers/OcaProvider.tsx @@ -1,11 +1,11 @@ import type { OcaModelInfo } from "@shared/api" -import type { OcaAuthState, OcaUserInfo } from "@shared/proto/index.cline" +import type { OcaAuthState, OcaUserInfo, VectorStoreInfo } from "@shared/proto/index.cline" import { EmptyRequest, StringRequest } from "@shared/proto/index.cline" import { Mode } from "@shared/storage/types" import { VSCodeButton, VSCodeLink, VSCodeProgressRing } from "@vscode/webview-ui-toolkit/react" import React, { useCallback, useEffect, useRef, useState } from "react" import { useExtensionState } from "@/context/ExtensionStateContext" -import { ModelsServiceClient, OcaAccountServiceClient } from "@/services/grpc-client" +import { ModelsServiceClient, OcaAccountServiceClient, VectorsServiceClient } from "@/services/grpc-client" import { VSC_BUTTON_BACKGROUND, VSC_BUTTON_FOREGROUND, @@ -16,6 +16,7 @@ import { import { BaseUrlField } from "../common/BaseUrlField" import { useApiConfigurationHandlers } from "../utils/useApiConfigurationHandlers" import OcaModelPicker from "./OcaModelPicker" +import OcaVectorPicker from "./OcaVectorPicker" /** * Props for the OcaProvider component @@ -112,7 +113,7 @@ function useOcaAuth() { * - Debounces base URL changes to avoid unnecessary calls. * - Guards against race conditions with a requestId and unmount checks. */ -function useOcaModels({ +function useOcaModelsAndKbs({ isAuthenticated, baseUrl, login, @@ -122,77 +123,135 @@ function useOcaModels({ login: () => Promise }) { const [models, setModels] = useState>({}) - const [loading, setLoading] = useState(false) - const [hasError, setHasError] = useState(false) - const [lastRefreshedAt, setLastRefreshedAt] = useState(null) - - const reqIdRef = useRef(0) - const unmountedRef = useRef(false) - const debounceTimerRef = useRef(null) - - const doRefresh = useCallback(async (url: string) => { - const myReqId = ++reqIdRef.current - setLoading(true) - setHasError(false) + const [kbs, setKbs] = useState>({}) + const [modelsLoading, setModelsLoading] = useState(false) + const [modelsHasError, setModelsHasError] = useState(false) + const [modelsLastRefreshedAt, setModelsLastRefreshedAt] = useState(null) + const [kbsLoading, setKbsLoading] = useState(false) + const [kbsHasError, setKbsHasError] = useState(false) + const [kbsLastRefreshedAt, setKbsLastRefreshedAt] = useState(null) + + const modelReqIdRef = useRef(0) + const modelUnmountedRef = useRef(false) + const modelDebounceTimerRef = useRef(null) + + const kbReqIdRef = useRef(0) + const kbUnmountedRef = useRef(false) + const kbDebounceTimerRef = useRef(null) + + const doRefreshModels = useCallback(async (url: string) => { + const myReqId = ++modelReqIdRef.current + setModelsLoading(true) + setModelsHasError(false) try { const resp = await ModelsServiceClient.refreshOcaModels(StringRequest.create({ value: url || "" })) + console.log(resp) // Only apply if still latest and still mounted - if (!unmountedRef.current && myReqId === reqIdRef.current) { + if (!modelUnmountedRef.current && myReqId === modelReqIdRef.current) { if (resp.error) { - setHasError(true) + setModelsHasError(true) } else { setModels(resp.models || {}) - setHasError(false) - setLastRefreshedAt(Date.now()) + setModelsHasError(false) + setModelsLastRefreshedAt(Date.now()) } } } catch (err) { - if (!unmountedRef.current && myReqId === reqIdRef.current) { + if (!modelUnmountedRef.current && myReqId === modelReqIdRef.current) { console.error("Failed to refresh Oca models:", err) - setHasError(true) + setModelsHasError(true) + } + } finally { + if (!modelUnmountedRef.current && myReqId === modelReqIdRef.current) { + setModelsLoading(false) + } + } + }, []) + + const doRefreshKbs = useCallback(async (url: string) => { + const myReqId = ++kbReqIdRef.current + setKbsLoading(true) + setKbsHasError(false) + try { + const resp = await VectorsServiceClient.refreshOcaVectors(StringRequest.create({ value: url || "" })) + // Only apply if still latest and still mounted + if (!kbUnmountedRef.current && myReqId === kbReqIdRef.current) { + if (resp.error) { + setKbsHasError(true) + } else { + setKbs(resp.vectors || {}) + setKbsHasError(false) + setKbsLastRefreshedAt(Date.now()) + } + } + } catch (err) { + if (!kbUnmountedRef.current && myReqId === kbReqIdRef.current) { + console.error("Failed to refresh Oca knowledge bases:", err) + setKbsHasError(true) } } finally { - if (!unmountedRef.current && myReqId === reqIdRef.current) { - setLoading(false) + if (!kbUnmountedRef.current && myReqId === kbReqIdRef.current) { + setKbsLoading(false) } } }, []) // Debounce changes to baseUrl or auth useEffect(() => { - unmountedRef.current = false - if (debounceTimerRef.current) { - window.clearTimeout(debounceTimerRef.current) - debounceTimerRef.current = null + modelUnmountedRef.current = false + kbUnmountedRef.current = false + if (modelDebounceTimerRef.current) { + window.clearTimeout(modelDebounceTimerRef.current) + modelDebounceTimerRef.current = null + } + if (kbDebounceTimerRef.current) { + window.clearTimeout(kbDebounceTimerRef.current) + kbDebounceTimerRef.current = null } if (!isAuthenticated) { // Clear models if logged out; prevent stale data setModels({}) - setLoading(false) - setHasError(false) + setModelsLoading(false) + setModelsHasError(false) + + setKbs({}) + setKbsLoading(false) + setKbsHasError(false) return } - debounceTimerRef.current = window.setTimeout(() => { - void doRefresh(baseUrl || "") + modelDebounceTimerRef.current = window.setTimeout(() => { + void doRefreshModels(baseUrl || "") + }, 250) + + kbDebounceTimerRef.current = window.setTimeout(() => { + void doRefreshKbs(baseUrl || "") }, 250) return () => { - unmountedRef.current = true - if (debounceTimerRef.current) { - window.clearTimeout(debounceTimerRef.current) - debounceTimerRef.current = null + modelUnmountedRef.current = true + if (modelDebounceTimerRef.current) { + window.clearTimeout(modelDebounceTimerRef.current) + modelDebounceTimerRef.current = null } // bump reqId so any in-flight result is ignored - reqIdRef.current++ + modelReqIdRef.current++ + + kbUnmountedRef.current = true + if (kbDebounceTimerRef.current) { + window.clearTimeout(kbDebounceTimerRef.current) + kbDebounceTimerRef.current = null + } + // bump reqId so any in-flight result is ignored + kbReqIdRef.current++ } - }, [isAuthenticated, baseUrl, doRefresh]) + }, [isAuthenticated, baseUrl, doRefreshModels, doRefreshKbs]) // User-initiated refresh with auto login + single retry on failure const refreshModels = useCallback(async () => { - setLoading(true) - setHasError(false) + setModelsLoading(true) + setModelsHasError(false) async function tryRefresh(retry = false): Promise { try { @@ -201,26 +260,67 @@ function useOcaModels({ throw new Error(resp.error) } setModels(resp.models || {}) - setHasError(false) - setLastRefreshedAt(Date.now()) + setModelsHasError(false) + setModelsLastRefreshedAt(Date.now()) return true } catch (_err) { if (!retry) { await login() // prompt login return tryRefresh(true) // retry once } else { - setHasError(true) + setModelsHasError(true) } return false } finally { - setLoading(false) + setModelsLoading(false) } } await tryRefresh() }, [baseUrl, login]) - return { models, loading, hasError, refreshModels, lastRefreshedAt } + const refreshKbs = useCallback(async () => { + setKbsLoading(true) + setKbsHasError(false) + + async function tryRefresh(retry = false): Promise { + try { + const resp = await VectorsServiceClient.refreshOcaVectors(StringRequest.create({ value: baseUrl || "" })) + if (resp.error) { + throw new Error(resp.error) + } + setKbs(resp.vectors || {}) + setKbsHasError(false) + setKbsLastRefreshedAt(Date.now()) + return true + } catch (_err) { + if (!retry) { + await login() // prompt login + return tryRefresh(true) // retry once + } else { + setKbsHasError(true) + } + return false + } finally { + setKbsLoading(false) + } + } + + await tryRefresh() + }, [baseUrl, login]) + + return { + models, + modelsLoading, + modelsHasError, + modelsLastRefreshedAt, + kbs, + kbsLoading, + kbsHasError, + kbsLastRefreshedAt, + refreshModels, + refreshKbs, + } } /** @@ -238,19 +338,28 @@ export const OcaProvider = ({ isPopup, currentMode }: OcaProviderProps) => { const { models: ocaModels, refreshModels, - hasError: ocaHasError, - loading: ocaLoading, - lastRefreshedAt, - } = useOcaModels({ + modelsHasError: ocaModelsHasError, + modelsLoading: ocaModelsLoading, + modelsLastRefreshedAt, + kbs: ocaKbs, + refreshKbs, + kbsHasError: ocaKbsHasError, + kbsLoading: ocaKbsLoading, + kbsLastRefreshedAt, + } = useOcaModelsAndKbs({ isAuthenticated, baseUrl: ocaBaseUrl, login, }) - const handleRefresh = useCallback(async () => { + const handleModelRefresh = useCallback(async () => { await refreshModels() }, [refreshModels]) + const handleKbRefresh = useCallback(async () => { + await refreshKbs() + }, [refreshModels]) + // On first subscription result: if user exists, refresh models once. const didInitialAuthCheckRef = useRef(false) useEffect(() => { @@ -260,9 +369,10 @@ export const OcaProvider = ({ isPopup, currentMode }: OcaProviderProps) => { didInitialAuthCheckRef.current = true if (isAuthenticated) { void refreshModels() + void refreshKbs() } // If user empty, do nothing (no auto login, no refresh) - }, [ready, isAuthenticated, refreshModels]) + }, [ready, isAuthenticated, refreshModels, refreshKbs]) return (
@@ -345,20 +455,50 @@ export const OcaProvider = ({ isPopup, currentMode }: OcaProviderProps) => { apiConfiguration={apiConfiguration} currentMode={currentMode} isPopup={isPopup} - lastRefreshedAt={lastRefreshedAt} - loading={ocaLoading} + lastRefreshedAt={modelsLastRefreshedAt} + loading={ocaModelsLoading} ocaModels={ocaModels} - onRefresh={handleRefresh} + onRefresh={handleModelRefresh} /> - {isAuthenticated && ocaHasError && ( + {isAuthenticated && ocaModelsHasError && (
Failed to refresh models. Check your session or network.
- + + Retry + + { + await login() + }}> + Sign in again + +
+
+ )} + + + + {isAuthenticated && ocaKbsHasError && ( +
+
Failed to refresh knowledge bases. Check your session or network.
+
+ Retry + onRefresh: () => void | Promise + loading?: boolean + lastRefreshedAt?: number | null +} + +const OcaVectorPicker: React.FC = ({ + apiConfiguration, + currentMode, + ocaKbs, + onRefresh, + loading, + lastRefreshedAt, +}: OcaVectorPickerProps) => { + const { handleModeFieldChange } = useApiConfigurationHandlers() + + const handleKbChange = async (kbs: string[]) => { + await handleModeFieldChange({ plan: "planModeOcaVectorIds", act: "actModeOcaVectorIds" }, kbs, currentMode) + } + + const handleRefreshToken = async () => { + await onRefresh?.() + } + + const { selectedVectorIds } = useMemo(() => { + return normalizeApiConfiguration(apiConfiguration, currentMode) + }, [apiConfiguration, currentMode]) + + const kbIds = useMemo(() => { + return Object.keys(ocaKbs || []).sort((a, b) => ocaKbs[a].name.localeCompare(ocaKbs[b].name)) + }, [ocaKbs]) + + const lastRefreshedText = useMemo(() => { + return typeof lastRefreshedAt === "number" ? new Date(lastRefreshedAt).toLocaleTimeString() : null + }, [lastRefreshedAt]) + + const toggleOption = async (option: { id: string; name: string }) => { + const prevVectorIds = selectedVectorIds || [] + const newVectorIds = prevVectorIds.includes(option.id) + ? prevVectorIds.filter((o) => o !== option.id) + : [...prevVectorIds, option.id] + await handleKbChange(newVectorIds) + } + + return ( +
+ + +
+ { + return { + id: kbId, + name: ocaKbs[kbId].name, + } + })} + selectedIds={selectedVectorIds || []} + toggleOption={toggleOption} + /> + + {loading ? "Refreshing…" : "Refresh"} + +
+ {lastRefreshedText ? ( +
+ Last refreshed at {lastRefreshedText} +
+ ) : null} +
+ ) +} + +export default OcaVectorPicker + +interface MultiSelectDropdownProps { + className?: string + id?: string + options: { + id: string + name: string + }[] + selectedIds: string[] + toggleOption: (option: { id: string; name: string }) => Promise +} + +const MultiSelectDropdown: React.FC = ({ options, selectedIds, toggleOption, className, id }) => { + const [open, setOpen] = useState(false) + const wrapperRef = useRef(null) + + useEffect(() => { + function handleClickOutside(event: MouseEvent) { + if (wrapperRef.current && !wrapperRef.current.contains(event.target as Node)) { + setOpen(false) + } + } + document.addEventListener("mousedown", handleClickOutside) + return () => document.removeEventListener("mousedown", handleClickOutside) + }, []) + + const selectedLabel = + selectedIds.length === 0 || options.length === 0 + ? "Select options..." + : selectedIds.map((selectedId) => options.filter((option) => selectedId === option.id)[0].name).join(", ") + + return ( +
+ {/* VSCode style dropdown button */} +
options.length > 0 && setOpen((o) => !o)} + onKeyDown={(e) => { + if (e.key === "Escape") { + setOpen(false) + } + if ((e.key === " " || e.key === "Enter") && options.length > 0) { + setOpen((o) => !o) + } + }} + style={{ + display: "flex", + alignItems: "center", + position: "absolute", + top: 0, + left: 0, + bottom: 0, + right: 0, + minHeight: "100%", + border: "1px solid var(--vscode-dropdown-border, #3c3c3c)", + borderRadius: "calc(var(--corner-radius-round) * 1px)", + background: "var(--vscode-dropdown-background, #1e1e1e)", + color: "var(--vscode-dropdown-foreground, #cccccc)", + minWidth: 0, + padding: "2px 6px 2px 8px", + cursor: options.length == 0 ? "not-allowed" : "pointer", + fontSize: 12, + fontFamily: "var(--vscode-font-family, inherit)", + boxSizing: "border-box", + outline: open ? "2px solid var(--vscode-focusBorder, #0078d4)" : "none", + margin: 0, + opacity: options.length == 0 ? 0.6 : 1, + }} + tabIndex={0}> + + {selectedLabel} + + + + +
+ + {open && ( +
+ {options.map((option) => { + const checked = selectedIds.includes(option.id) + return ( +
{ + e.stopPropagation() + await toggleOption(option) + }} + onKeyDown={async (e) => { + if (e.key === " " || e.key === "Enter") { + await toggleOption(option) + } + }} + role="option" + style={{ + display: "flex", + alignItems: "center", + padding: "4px 8px", + margin: 0, + borderRadius: 3, + cursor: "pointer", + background: checked ? "var(--vscode-list-activeSelectionBackground, #094771)" : "transparent", + color: checked + ? "var(--vscode-list-activeSelectionForeground, #fff)" + : "var(--vscode-dropdown-foreground, #cccccc)", + }} + tabIndex={0}> + + + {option.name} + +
+ ) + })} +
+ )} +
+ ) +} diff --git a/webview-ui/src/components/settings/utils/providerUtils.ts b/webview-ui/src/components/settings/utils/providerUtils.ts index 6f919694251..f6e89dbf7a0 100644 --- a/webview-ui/src/components/settings/utils/providerUtils.ts +++ b/webview-ui/src/components/settings/utils/providerUtils.ts @@ -72,6 +72,7 @@ export interface NormalizedApiConfig { selectedProvider: ApiProvider selectedModelId: string selectedModelInfo: ModelInfo + selectedVectorIds?: string[] } /** @@ -352,10 +353,14 @@ export function normalizeApiConfiguration( const ocaModelId = currentMode === "plan" ? apiConfiguration?.planModeOcaModelId : apiConfiguration?.actModeOcaModelId const ocaModelInfo = currentMode === "plan" ? apiConfiguration?.planModeOcaModelInfo : apiConfiguration?.actModeOcaModelInfo + const ocaVectorIds = + currentMode === "plan" ? apiConfiguration?.planModeOcaVectorIds : apiConfiguration?.actModeOcaVectorIds + return { selectedProvider: provider, selectedModelId: ocaModelId || "", selectedModelInfo: ocaModelInfo || liteLlmModelInfoSaneDefaults, + selectedVectorIds: ocaVectorIds || [], } default: return getProviderData(anthropicModels, anthropicDefaultModelId)