From d6a28ec3acc334acd7818423cb8769370de205f2 Mon Sep 17 00:00:00 2001 From: Logan Yang Date: Fri, 26 Jan 2024 14:26:04 -0800 Subject: [PATCH] [v2.4.14] Make openai key not required for other chat and embedding models (#261) --- manifest.json | 2 +- package-lock.json | 4 +- package.json | 2 +- src/LLMProviders/chainManager.ts | 29 +++++++--- src/LLMProviders/embeddingManager.ts | 69 ++++++++++++++---------- src/settings/SettingsPage.tsx | 1 + src/settings/components/ApiSettings.tsx | 12 +++++ src/settings/components/QASettings.tsx | 1 + src/settings/components/SettingsMain.tsx | 4 ++ versions.json | 3 +- 10 files changed, 87 insertions(+), 40 deletions(-) diff --git a/manifest.json b/manifest.json index e195d67d..2af86281 100644 --- a/manifest.json +++ b/manifest.json @@ -1,7 +1,7 @@ { "id": "copilot", "name": "Copilot", - "version": "2.4.13", + "version": "2.4.14", "minAppVersion": "0.15.0", "description": "A ChatGPT Copilot in Obsidian.", "author": "Logan Yang", diff --git a/package-lock.json b/package-lock.json index 61c07d00..aa5642d9 100644 --- a/package-lock.json +++ b/package-lock.json @@ -1,12 +1,12 @@ { "name": "obsidian-copilot", - "version": "2.4.13", + "version": "2.4.14", "lockfileVersion": 3, "requires": true, "packages": { "": { "name": "obsidian-copilot", - "version": "2.4.13", + "version": "2.4.14", "license": "AGPL-3.0", "dependencies": { "@huggingface/inference": "^2.6.4", diff --git a/package.json b/package.json index b2dc2227..0921511c 100644 --- a/package.json +++ b/package.json @@ -1,6 +1,6 @@ { "name": "obsidian-copilot", - "version": "2.4.13", + "version": "2.4.14", "description": "ChatGPT integration for Obsidian", "main": "main.js", "scripts": { diff --git a/src/LLMProviders/chainManager.ts b/src/LLMProviders/chainManager.ts index ce1b5e4a..6b56797a 100644 --- a/src/LLMProviders/chainManager.ts +++ b/src/LLMProviders/chainManager.ts @@ -55,7 +55,6 @@ export default class ChainManager { this.memoryManager = MemoryManager.getInstance(this.langChainParams); this.chatModelManager = ChatModelManager.getInstance(this.langChainParams); this.promptManager = PromptManager.getInstance(this.langChainParams); - this.embeddingsManager = EmbeddingsManager.getInstance(this.langChainParams); this.createChainWithNewModel(this.langChainParams.modelDisplayName); } @@ -134,12 +133,15 @@ export default class ChainManager { return; } this.validateChainType(chainType); + // MUST set embeddingsManager when switching to QA mode + if (chainType === ChainType.RETRIEVAL_QA_CHAIN) { + this.embeddingsManager = EmbeddingsManager.getInstance(this.langChainParams); + } // Get chatModel, memory, prompt, and embeddingAPI from respective managers const chatModel = this.chatModelManager.getChatModel(); const memory = this.memoryManager.getMemory(); const chatPrompt = this.promptManager.getChatPrompt(); - const embeddingsAPI = this.embeddingsManager.getEmbeddingsAPI(); switch (chainType) { case ChainType.LLM_CHAIN: { @@ -182,8 +184,14 @@ export default class ChainManager { const parsedMemoryVectors: MemoryVector[] | undefined = await VectorDBManager.getMemoryVectors(docHash); if (parsedMemoryVectors) { // Index already exists + const embeddingsAPI = this.embeddingsManager.getEmbeddingsAPI(); + if (!embeddingsAPI) { + console.error('Error getting embeddings API. Please check your settings.'); + return; + } const vectorStore = await VectorDBManager.rebuildMemoryVectorStore( - parsedMemoryVectors, embeddingsAPI + parsedMemoryVectors, + embeddingsAPI, ); // Create new conversational retrieval chain @@ -391,14 +399,19 @@ export default class ChainManager { } async buildIndex(noteContent: string, docHash: string): Promise { - const textSplitter = new RecursiveCharacterTextSplitter({ chunkSize: 1000 }); - - const docs = await textSplitter.createDocuments([noteContent]); - const embeddingsAPI = this.embeddingsManager.getEmbeddingsAPI(); - // Note: HF can give 503 errors frequently (it's free) console.log('Creating vector store...'); try { + const textSplitter = new RecursiveCharacterTextSplitter({ chunkSize: 1000 }); + + const docs = await textSplitter.createDocuments([noteContent]); + const embeddingsAPI = this.embeddingsManager.getEmbeddingsAPI(); + if (!embeddingsAPI) { + const errorMsg = 'Failed to create vector store, embedding API is not set correctly, please check your settings.'; + new Notice(errorMsg); + console.error(errorMsg); + return; + } this.vectorStore = await MemoryVectorStore.fromDocuments( docs, embeddingsAPI, ); diff --git a/src/LLMProviders/embeddingManager.ts b/src/LLMProviders/embeddingManager.ts index e2d23324..e9a4561a 100644 --- a/src/LLMProviders/embeddingManager.ts +++ b/src/LLMProviders/embeddingManager.ts @@ -21,7 +21,7 @@ export default class EmbeddingManager { return EmbeddingManager.instance; } - getEmbeddingsAPI(): Embeddings { + getEmbeddingsAPI(): Embeddings | undefined { const { openAIApiKey, azureOpenAIApiKey, @@ -33,26 +33,32 @@ export default class EmbeddingManager { // Note that openAIProxyBaseUrl has the highest priority. // If openAIProxyBaseUrl is set, it overrides both chat and embedding models. - const OpenAIEmbeddingsAPI = openAIProxyBaseUrl ? - new ProxyOpenAIEmbeddings({ - modelName: this.langChainParams.embeddingModel, - openAIApiKey, - maxRetries: 3, - maxConcurrency: 3, - timeout: 10000, - openAIProxyBaseUrl, - }): - new OpenAIEmbeddings({ - modelName: this.langChainParams.embeddingModel, - openAIApiKey, - maxRetries: 3, - maxConcurrency: 3, - timeout: 10000, - }); + const OpenAIEmbeddingsAPI = openAIApiKey ? ( + openAIProxyBaseUrl ? + new ProxyOpenAIEmbeddings({ + modelName: this.langChainParams.embeddingModel, + openAIApiKey, + maxRetries: 3, + maxConcurrency: 3, + timeout: 10000, + openAIProxyBaseUrl, + }) : + new OpenAIEmbeddings({ + modelName: this.langChainParams.embeddingModel, + openAIApiKey, + maxRetries: 3, + maxConcurrency: 3, + timeout: 10000, + }) + ) : null; switch(this.langChainParams.embeddingProvider) { case ModelProviders.OPENAI: - return OpenAIEmbeddingsAPI + if (OpenAIEmbeddingsAPI) { + return OpenAIEmbeddingsAPI; + } + console.error('OpenAI API key is not provided for the embedding model.'); + break; case ModelProviders.HUGGINGFACE: return new HuggingFaceInferenceEmbeddings({ apiKey: this.langChainParams.huggingfaceApiKey, @@ -66,18 +72,27 @@ export default class EmbeddingManager { maxConcurrency: 3, }); case ModelProviders.AZURE_OPENAI: - return new OpenAIEmbeddings({ - azureOpenAIApiKey, - azureOpenAIApiInstanceName, - azureOpenAIApiDeploymentName: azureOpenAIApiEmbeddingDeploymentName, - azureOpenAIApiVersion, + if (azureOpenAIApiKey) { + return new OpenAIEmbeddings({ + azureOpenAIApiKey, + azureOpenAIApiInstanceName, + azureOpenAIApiDeploymentName: azureOpenAIApiEmbeddingDeploymentName, + azureOpenAIApiVersion, + maxRetries: 3, + maxConcurrency: 3, + }); + } + console.error('Azure OpenAI API key is not provided for the embedding model.'); + break; + default: + console.error('No embedding provider set or no valid API key provided. Defaulting to OpenAI.'); + return OpenAIEmbeddingsAPI || new OpenAIEmbeddings({ + modelName: this.langChainParams.embeddingModel, + openAIApiKey: 'default-key', maxRetries: 3, maxConcurrency: 3, + timeout: 10000, }); - default: - console.error('No embedding provider set. Using OpenAI.'); - return OpenAIEmbeddingsAPI; } } - } \ No newline at end of file diff --git a/src/settings/SettingsPage.tsx b/src/settings/SettingsPage.tsx index f8bcaebd..9fbfe1e2 100644 --- a/src/settings/SettingsPage.tsx +++ b/src/settings/SettingsPage.tsx @@ -49,6 +49,7 @@ export class CopilotSettingTab extends PluginSettingTab { await this.plugin.saveSettings(); // Reload the plugin + // eslint-disable-next-line @typescript-eslint/no-explicit-any const app = (this.plugin.app as any); await app.plugins.disablePlugin("copilot"); await app.plugins.enablePlugin("copilot"); diff --git a/src/settings/components/ApiSettings.tsx b/src/settings/components/ApiSettings.tsx index 55d33a3f..50b9c427 100644 --- a/src/settings/components/ApiSettings.tsx +++ b/src/settings/components/ApiSettings.tsx @@ -20,6 +20,8 @@ interface ApiSettingsProps { setAzureOpenAIApiDeploymentName: (value: string) => void; azureOpenAIApiVersion: string; setAzureOpenAIApiVersion: (value: string) => void; + azureOpenAIApiEmbeddingDeploymentName: string; + setAzureOpenAIApiEmbeddingDeploymentName: (value: string) => void; } const ApiSettings: React.FC = ({ @@ -39,6 +41,8 @@ const ApiSettings: React.FC = ({ setAzureOpenAIApiDeploymentName, azureOpenAIApiVersion, setAzureOpenAIApiVersion, + azureOpenAIApiEmbeddingDeploymentName, + setAzureOpenAIApiEmbeddingDeploymentName, }) => { return (
@@ -154,6 +158,14 @@ const ApiSettings: React.FC = ({ placeholder="Enter Azure OpenAI API Version" type="text" /> +
diff --git a/src/settings/components/QASettings.tsx b/src/settings/components/QASettings.tsx index e7b14778..1a10da70 100644 --- a/src/settings/components/QASettings.tsx +++ b/src/settings/components/QASettings.tsx @@ -50,6 +50,7 @@ const QASettings: React.FC = ({ />