diff --git a/src/LLMProviders/chainManager.ts b/src/LLMProviders/chainManager.ts index 8540feb9..69e98f1f 100644 --- a/src/LLMProviders/chainManager.ts +++ b/src/LLMProviders/chainManager.ts @@ -208,6 +208,7 @@ export default class ChainManager { retriever: vectorStore.asRetriever(undefined, (doc) => { return doc.metadata.path === options.noteFile?.path; }), + systemMessage: this.langChainParams.systemMessage, }, ChainManager.storeRetrieverDocuments.bind(ChainManager), options.debug @@ -238,6 +239,7 @@ export default class ChainManager { { llm: chatModel, retriever: retriever, + systemMessage: this.langChainParams.systemMessage, }, ChainManager.storeRetrieverDocuments.bind(ChainManager), options.debug @@ -281,6 +283,7 @@ export default class ChainManager { { llm: chatModel, retriever: retriever, + systemMessage: this.langChainParams.systemMessage, }, ChainManager.storeRetrieverDocuments.bind(ChainManager), options.debug @@ -477,15 +480,17 @@ export default class ChainManager { // expand and reveal the chunk if (this.langChainParams.chainType === ChainType.VAULT_QA_CHAIN) { const docTitles = extractUniqueTitlesFromDocs(ChainManager.retrievedDocuments); - const markdownLinks = docTitles - .map( - (title) => - `- [${title}](obsidian://open?vault=${encodeURIComponent(this.app.vault.getName())}&file=${encodeURIComponent( - title - )})` - ) - .join("\n"); - fullAIResponse += "\n\n#### Sources:\n" + markdownLinks; + if (docTitles.length > 0) { + const markdownLinks = docTitles + .map( + (title) => + `- [${title}](obsidian://open?vault=${encodeURIComponent(this.app.vault.getName())}&file=${encodeURIComponent( + title + )})` + ) + .join("\n"); + fullAIResponse += "\n\n#### Sources:\n" + markdownLinks; + } } return fullAIResponse; diff --git a/src/LLMProviders/embeddingManager.ts b/src/LLMProviders/embeddingManager.ts index 9ad46f49..e926f8ff 100644 --- a/src/LLMProviders/embeddingManager.ts +++ b/src/LLMProviders/embeddingManager.ts @@ -2,12 +2,14 @@ import { CustomModel, LangChainParams } from "@/aiParams"; import { EmbeddingModelProviders } from "@/constants"; import EncryptionService from "@/encryptionService"; +import { CustomError } from "@/error"; import { ProxyOpenAIEmbeddings } from "@/langchainWrappers"; import { CohereEmbeddings } from "@langchain/cohere"; import { Embeddings } from "@langchain/core/embeddings"; import { GoogleGenerativeAIEmbeddings } from "@langchain/google-genai"; import { OllamaEmbeddings } from "@langchain/ollama"; import { OpenAIEmbeddings } from "@langchain/openai"; + export default class EmbeddingManager { private encryptionService: EncryptionService; private activeEmbeddingModels: CustomModel[]; @@ -122,14 +124,14 @@ export default class EmbeddingManager { const { embeddingModelKey } = this.getLangChainParams(); if (!EmbeddingManager.modelMap.hasOwnProperty(embeddingModelKey)) { - console.error(`No embedding model found for: ${embeddingModelKey}`); - return; + throw new CustomError(`No embedding model found for: ${embeddingModelKey}`); } const selectedModel = EmbeddingManager.modelMap[embeddingModelKey]; if (!selectedModel.hasApiKey) { - console.error(`API key is not provided for the embedding model: ${embeddingModelKey}`); - return; + throw new CustomError( + `API key is not provided for the embedding model: ${embeddingModelKey}` + ); } const customModel = this.getCustomModel(embeddingModelKey); @@ -139,7 +141,9 @@ export default class EmbeddingManager { EmbeddingManager.embeddingModel = new selectedModel.EmbeddingConstructor(config); return EmbeddingManager.embeddingModel; } catch (error) { - console.error(`Error creating embedding model: ${embeddingModelKey}`, error); + throw new CustomError( + `Error creating embedding model: ${embeddingModelKey}. ${error.message}` + ); } } diff --git a/src/chainFactory.ts b/src/chainFactory.ts index 43e83082..d76b310f 100644 --- a/src/chainFactory.ts +++ b/src/chainFactory.ts @@ -1,9 +1,9 @@ +import { BaseLanguageModel } from "@langchain/core/language_models/base"; import { StringOutputParser } from "@langchain/core/output_parsers"; +import { ChatPromptTemplate, PromptTemplate } from "@langchain/core/prompts"; import { BaseRetriever } from "@langchain/core/retrievers"; import { RunnablePassthrough, RunnableSequence } from "@langchain/core/runnables"; -import { BaseLanguageModel } from "@langchain/core/language_models/base"; import { BaseChatMemory } from "langchain/memory"; -import { ChatPromptTemplate, PromptTemplate } from "@langchain/core/prompts"; import { formatDocumentsAsString } from "langchain/util/document"; export interface LLMChainInput { @@ -24,6 +24,7 @@ export interface RetrievalChainParams { export interface ConversationalRetrievalChainParams { llm: BaseLanguageModel; retriever: BaseRetriever; + systemMessage: string; options?: { returnSourceDocuments?: boolean; questionGeneratorTemplate?: string; @@ -128,13 +129,13 @@ class ChainFactory { onDocumentsRetrieved: (documents: Document[]) => void, debug?: boolean ): RunnableSequence { - const { llm, retriever } = args; + const { llm, retriever, systemMessage } = args; // NOTE: This is a tricky part of the Conversational RAG. Weaker models may fail this instruction // and lose the follow up question altogether. const condenseQuestionTemplate = `Given the following conversation and a follow up question, summarize the conversation as context and keep the follow up question unchanged, in its original language. - If the follow up question is unrelated to its preceding messages, return the question directly. + If the follow up question is unrelated to its preceding messages, return this follow up question directly. If it is related, then combine the summary and the follow up question to construct a standalone question. Make sure to keep any [[]] wrapped note titles in the question unchanged. @@ -144,11 +145,13 @@ class ChainFactory { Standalone question:`; const CONDENSE_QUESTION_PROMPT = PromptTemplate.fromTemplate(condenseQuestionTemplate); - const answerTemplate = `Answer the question with as detailed as possible based only on the following context: - {context} + const answerTemplate = `{system_message} - Question: {question} - `; +Answer the question with as detailed as possible based only on the following context: +{context} + +Question: {question} +`; const ANSWER_PROMPT = PromptTemplate.fromTemplate(answerTemplate); const formatChatHistory = (chatHistory: [string, string][]) => { @@ -173,6 +176,10 @@ class ChainFactory { CONDENSE_QUESTION_PROMPT, llm, new StringOutputParser(), + (output) => { + if (debug) console.log("Standalone Question: ", output); + return output; + }, ]); const formatDocumentsAsStringAndStore = async (documents: Document[]) => { @@ -185,6 +192,7 @@ class ChainFactory { { context: retriever.pipe(formatDocumentsAsStringAndStore), question: new RunnablePassthrough(), + system_message: () => systemMessage, }, ANSWER_PROMPT, llm, diff --git a/src/components/Chat.tsx b/src/components/Chat.tsx index 8a4616ce..977a2fe1 100644 --- a/src/components/Chat.tsx +++ b/src/components/Chat.tsx @@ -373,8 +373,13 @@ ${chatContent}`; return; } - await plugin.indexVaultToVectorStore(); - new Notice("Vault index refreshed."); + try { + await plugin.indexVaultToVectorStore(); + new Notice("Vault index refreshed."); + } catch (error) { + console.error("Error refreshing vault index:", error); + new Notice("Failed to refresh vault index. Check console for details."); + } }; const clearCurrentAiMessage = () => { @@ -681,7 +686,6 @@ ${chatContent}`; clearMessages(); clearChatMemory(); clearCurrentAiMessage(); - console.log(Date.now()); }} onSaveAsNote={() => handleSaveAsNote(true)} onSendActiveNoteToPrompt={handleSendActiveNoteToPrompt} diff --git a/src/components/ChatComponents/ChatIcons.tsx b/src/components/ChatComponents/ChatIcons.tsx index db107c4e..4dd0d757 100644 --- a/src/components/ChatComponents/ChatIcons.tsx +++ b/src/components/ChatComponents/ChatIcons.tsx @@ -1,5 +1,6 @@ import { CustomModel, SetChainOptions } from "@/aiParams"; import { AI_SENDER, VAULT_VECTOR_STORE_STRATEGY } from "@/constants"; +import { CustomError } from "@/error"; import { CopilotSettings } from "@/settings/SettingsPage"; import { ChatMessage } from "@/sharedState"; import { formatDateTime, getFileContent, getFileName } from "@/utils"; @@ -112,7 +113,19 @@ const ChatIcons: React.FC = ({ addMessage(activeNoteOnMessage); } - setCurrentChain(selectedChain, { debug }); + try { + await setCurrentChain(selectedChain, { debug }); + } catch (error) { + if (error instanceof CustomError) { + console.error("Error setting QA chain:", error.msg); + new Notice(`Error: ${error.msg}. Please check your embedding model settings.`); + } else { + console.error("Unexpected error setting QA chain:", error); + new Notice( + "An unexpected error occurred while setting up the QA chain. Please check the console for details." + ); + } + } }; handleChainSelection(); diff --git a/src/main.ts b/src/main.ts index 7ac075ab..72cc29e9 100644 --- a/src/main.ts +++ b/src/main.ts @@ -460,110 +460,125 @@ export default class CopilotPlugin extends Plugin { } async indexVaultToVectorStore(overwrite?: boolean): Promise { - const embeddingInstance = this.embeddingsManager.getEmbeddingsAPI(); - if (!embeddingInstance) { - throw new Error("Embedding instance not found."); - } - - // Check if embedding model has changed - const prevEmbeddingModel = await VectorDBManager.checkEmbeddingModel(this.dbVectorStores); - // TODO: Remove this when Ollama model is dynamically set - const currEmbeddingModel = EmbeddingsManager.getModelName(embeddingInstance); - - if (this.settings.debug) { - console.log( - `\nVault QA exclusion paths: ${this.settings.qaExclusionPaths ? this.settings.qaExclusionPaths : "None"}` - ); - console.log("Prev vs Current embedding models:", prevEmbeddingModel, currEmbeddingModel); - } - - if (!areEmbeddingModelsSame(prevEmbeddingModel, currEmbeddingModel)) { - // Model has changed, clear DB and reindex from scratch - overwrite = true; - // Clear the current vector store with mixed embeddings - try { - // Clear the vectorstore db - await this.dbVectorStores.destroy(); - // Reinitialize the database - this.dbVectorStores = new PouchDB( - `copilot_vector_stores_${this.getVaultIdentifier()}` - ); - new Notice("Detected change in embedding model. Rebuild vector store from scratch."); - console.log("Detected change in embedding model. Rebuild vector store from scratch."); - } catch (err) { - console.error("Error clearing vector store for reindexing:", err); - new Notice("Error clearing vector store for reindexing."); + try { + const embeddingInstance = this.embeddingsManager.getEmbeddingsAPI(); + if (!embeddingInstance) { + throw new CustomError("Embedding instance not found."); } - } - const latestMtime = await VectorDBManager.getLatestFileMtime(this.dbVectorStores); + // Check if embedding model has changed + const prevEmbeddingModel = await VectorDBManager.checkEmbeddingModel(this.dbVectorStores); + // TODO: Remove this when Ollama model is dynamically set + const currEmbeddingModel = EmbeddingsManager.getModelName(embeddingInstance); - const files = this.app.vault - .getMarkdownFiles() - .filter((file) => { - if (!latestMtime || overwrite) return true; - return file.stat.mtime > latestMtime; - }) - // file not in qaExclusionPaths - .filter((file) => { - if (!this.settings.qaExclusionPaths) return true; - return !isPathInList(file.path, this.settings.qaExclusionPaths); + if (this.settings.debug) { + console.log( + `\nVault QA exclusion paths: ${this.settings.qaExclusionPaths ? this.settings.qaExclusionPaths : "None"}` + ); + console.log("Prev vs Current embedding models:", prevEmbeddingModel, currEmbeddingModel); + } + + if (!areEmbeddingModelsSame(prevEmbeddingModel, currEmbeddingModel)) { + // Model has changed, clear DB and reindex from scratch + overwrite = true; + // Clear the current vector store with mixed embeddings + try { + // Clear the vectorstore db + await this.dbVectorStores.destroy(); + // Reinitialize the database + this.dbVectorStores = new PouchDB( + `copilot_vector_stores_${this.getVaultIdentifier()}` + ); + new Notice("Detected change in embedding model. Rebuild vector store from scratch."); + console.log("Detected change in embedding model. Rebuild vector store from scratch."); + } catch (err) { + console.error("Error clearing vector store for reindexing:", err); + new Notice("Error clearing vector store for reindexing."); + } + } + + const latestMtime = await VectorDBManager.getLatestFileMtime(this.dbVectorStores); + + const files = this.app.vault + .getMarkdownFiles() + .filter((file) => { + if (!latestMtime || overwrite) return true; + return file.stat.mtime > latestMtime; + }) + // file not in qaExclusionPaths + .filter((file) => { + if (!this.settings.qaExclusionPaths) return true; + return !isPathInList(file.path, this.settings.qaExclusionPaths); + }); + + const fileContents: string[] = await Promise.all( + files.map((file) => this.app.vault.cachedRead(file)) + ); + const fileMetadatas = files.map((file) => this.app.metadataCache.getFileCache(file)); + + const totalFiles = files.length; + if (totalFiles === 0) { + new Notice("Copilot vault index is up-to-date."); + return 0; + } + + let indexedCount = 0; + const indexNotice = new Notice( + `Copilot is indexing your vault...\n0/${totalFiles} files processed.`, + 0 + ); + + const errors: string[] = []; + const loadPromises = files.map(async (file, index) => { + try { + const noteFile = { + basename: file.basename, + path: file.path, + mtime: file.stat.mtime, + content: fileContents[index], + metadata: fileMetadatas[index]?.frontmatter ?? {}, + }; + const result = await VectorDBManager.indexFile( + this.dbVectorStores, + embeddingInstance, + noteFile + ); + + indexedCount++; + indexNotice.setMessage( + `Copilot is indexing your vault...\n${indexedCount}/${totalFiles} files processed.` + ); + return result; + } catch (err) { + console.error("Error indexing file:", err); + errors.push(`Error indexing file: ${file.basename}`); + } }); - const fileContents: string[] = await Promise.all( - files.map((file) => this.app.vault.cachedRead(file)) - ); - const fileMetadatas = files.map((file) => this.app.metadataCache.getFileCache(file)); + await Promise.all(loadPromises); + setTimeout(() => { + indexNotice.hide(); + }, 2000); - const totalFiles = files.length; - if (totalFiles === 0) { - new Notice("Copilot vault index is up-to-date."); + if (errors.length > 0) { + new Notice(`Indexing completed with errors. Check the console for details.`); + console.log("Indexing Errors:", errors.join("\n")); + } + return files.length; + } catch (error) { + if (error instanceof CustomError) { + console.error("Error indexing vault to vector store:", error.msg); + new Notice( + `Error indexing vault: ${error.msg}. Please check your embedding model settings.` + ); + } else { + console.error("Unexpected error indexing vault to vector store:", error); + new Notice( + "An unexpected error occurred while indexing the vault. Please check the console for details." + ); + } return 0; } - - let indexedCount = 0; - const indexNotice = new Notice( - `Copilot is indexing your vault...\n0/${totalFiles} files processed.`, - 0 - ); - - const errors: string[] = []; - const loadPromises = files.map(async (file, index) => { - try { - const noteFile = { - basename: file.basename, - path: file.path, - mtime: file.stat.mtime, - content: fileContents[index], - metadata: fileMetadatas[index]?.frontmatter ?? {}, - }; - const result = await VectorDBManager.indexFile( - this.dbVectorStores, - embeddingInstance, - noteFile - ); - - indexedCount++; - indexNotice.setMessage( - `Copilot is indexing your vault...\n${indexedCount}/${totalFiles} files processed.` - ); - return result; - } catch (err) { - console.error("Error indexing file:", err); - errors.push(`Error indexing file: ${file.basename}`); - } - }); - - await Promise.all(loadPromises); - setTimeout(() => { - indexNotice.hide(); - }, 2000); - - if (errors.length > 0) { - new Notice(`Indexing completed with errors. Check the console for details.`); - console.log("Indexing Errors:", errors.join("\n")); - } - return files.length; } async processText(