mirror of
https://github.com/logancyang/obsidian-copilot.git
synced 2026-07-22 07:50:24 +00:00
Support more chain types, implement vector search powered by huggingface inference api (#34)
* Implement ChainFactory for chain singletons * Add in-memory vector search powered by huggingface inference api * Add todo items for unlimited context search
This commit is contained in:
parent
24defc6504
commit
ab2b5e5387
10 changed files with 242 additions and 36 deletions
|
|
@ -1,5 +1,5 @@
|
|||
# 🔍 Copilot for Obsidian
|
||||
 
|
||||
 
|
||||
|
||||
|
||||
Copilot for Obsidian is a ChatGPT interface right inside Obsidian. It has a minimalistic design and is straightforward to use.
|
||||
|
|
|
|||
9
package-lock.json
generated
9
package-lock.json
generated
|
|
@ -9,6 +9,7 @@
|
|||
"version": "2.0.0",
|
||||
"license": "AGPL-3.0",
|
||||
"dependencies": {
|
||||
"@huggingface/inference": "^1.8.0",
|
||||
"@tabler/icons-react": "^2.14.0",
|
||||
"axios": "^1.3.4",
|
||||
"esbuild-plugin-svg": "^0.1.0",
|
||||
|
|
@ -1107,6 +1108,14 @@
|
|||
"node": ">=16.15"
|
||||
}
|
||||
},
|
||||
"node_modules/@huggingface/inference": {
|
||||
"version": "1.8.0",
|
||||
"resolved": "https://registry.npmjs.org/@huggingface/inference/-/inference-1.8.0.tgz",
|
||||
"integrity": "sha512-Dkh7PiyMf6TINRocQsdceiR5LcqJiUHgWjaBMRpCUOCbs+GZA122VH9q+wodoSptj6rIQf7wIwtDsof+/gd0WA==",
|
||||
"engines": {
|
||||
"node": ">=18"
|
||||
}
|
||||
},
|
||||
"node_modules/@humanwhocodes/config-array": {
|
||||
"version": "0.11.8",
|
||||
"resolved": "https://registry.npmjs.org/@humanwhocodes/config-array/-/config-array-0.11.8.tgz",
|
||||
|
|
|
|||
|
|
@ -33,6 +33,7 @@
|
|||
"typescript": "4.7.4"
|
||||
},
|
||||
"dependencies": {
|
||||
"@huggingface/inference": "^1.8.0",
|
||||
"@tabler/icons-react": "^2.14.0",
|
||||
"axios": "^1.3.4",
|
||||
"esbuild-plugin-svg": "^0.1.0",
|
||||
|
|
|
|||
143
src/aiState.ts
143
src/aiState.ts
|
|
@ -1,7 +1,20 @@
|
|||
import { AI_SENDER, DEFAULT_SYSTEM_PROMPT, USER_SENDER } from '@/constants';
|
||||
import ChainFactory, {
|
||||
CONVERSATIONAL_RETRIEVAL_QA_CHAIN,
|
||||
LLM_CHAIN,
|
||||
} from '@/chainFactory';
|
||||
import {
|
||||
AI_SENDER,
|
||||
DEFAULT_SYSTEM_PROMPT,
|
||||
USER_SENDER
|
||||
} from '@/constants';
|
||||
import { ChatMessage } from '@/sharedState';
|
||||
import { ConversationChain } from "langchain/chains";
|
||||
import {
|
||||
BaseChain,
|
||||
ConversationChain,
|
||||
ConversationalRetrievalQAChain
|
||||
} from "langchain/chains";
|
||||
import { ChatOpenAI } from 'langchain/chat_models/openai';
|
||||
import { HuggingFaceInferenceEmbeddings } from "langchain/embeddings/hf";
|
||||
import { BufferWindowMemory } from "langchain/memory";
|
||||
import {
|
||||
ChatPromptTemplate,
|
||||
|
|
@ -10,10 +23,13 @@ import {
|
|||
SystemMessagePromptTemplate,
|
||||
} from "langchain/prompts";
|
||||
import { AIChatMessage, HumanChatMessage, SystemChatMessage } from 'langchain/schema';
|
||||
import { RecursiveCharacterTextSplitter } from "langchain/text_splitter";
|
||||
import { MemoryVectorStore } from "langchain/vectorstores/memory";
|
||||
import { useState } from 'react';
|
||||
|
||||
export interface LangChainParams {
|
||||
key: string,
|
||||
huggingfaceApiKey: string,
|
||||
model: string,
|
||||
temperature: number,
|
||||
maxTokens: number,
|
||||
|
|
@ -21,9 +37,16 @@ export interface LangChainParams {
|
|||
chatContextTurns: number,
|
||||
}
|
||||
|
||||
interface SetChainOptions {
|
||||
prompt?: ChatPromptTemplate;
|
||||
noteContent?: string;
|
||||
}
|
||||
|
||||
class AIState {
|
||||
static chatOpenAI: ChatOpenAI;
|
||||
static chain: ConversationChain;
|
||||
static chain: BaseChain;
|
||||
static retrievalChain: ConversationalRetrievalQAChain;
|
||||
static useChain: string;
|
||||
memory: BufferWindowMemory;
|
||||
langChainParams: LangChainParams;
|
||||
|
||||
|
|
@ -35,22 +58,23 @@ class AIState {
|
|||
returnMessages: true,
|
||||
});
|
||||
|
||||
this.createNewChain();
|
||||
this.createNewChain(LLM_CHAIN);
|
||||
}
|
||||
|
||||
clearChatMemory(): void {
|
||||
console.log('clearing chat memory');
|
||||
this.memory.clear();
|
||||
this.createNewChain();
|
||||
this.createNewChain(LLM_CHAIN);
|
||||
AIState.useChain = LLM_CHAIN;
|
||||
}
|
||||
|
||||
setModel(newModel: string): void {
|
||||
console.log('setting model to', newModel);
|
||||
this.langChainParams.model = newModel;
|
||||
this.createNewChain();
|
||||
this.createNewChain(LLM_CHAIN);
|
||||
}
|
||||
|
||||
createNewChain(): void {
|
||||
createNewChain(chainType: string): void {
|
||||
const {
|
||||
key, model, temperature, maxTokens, systemMessage,
|
||||
} = this.langChainParams;
|
||||
|
|
@ -69,17 +93,47 @@ class AIState {
|
|||
streaming: true,
|
||||
});
|
||||
|
||||
this.setChain(chainType, {prompt: chatPrompt});
|
||||
}
|
||||
|
||||
async setChain(
|
||||
chainType: string,
|
||||
options: SetChainOptions = {},
|
||||
): Promise<void> {
|
||||
// TODO: Use this once https://github.com/hwchase17/langchainjs/issues/1327 is resolved
|
||||
AIState.chain = new ConversationChain({
|
||||
llm: AIState.chatOpenAI,
|
||||
memory: this.memory,
|
||||
prompt: chatPrompt,
|
||||
});
|
||||
if (chainType === LLM_CHAIN && options.prompt) {
|
||||
AIState.chain = ChainFactory.getLLMChain({
|
||||
llm: AIState.chatOpenAI,
|
||||
memory: this.memory,
|
||||
prompt: options.prompt,
|
||||
}) as ConversationChain;
|
||||
AIState.useChain = LLM_CHAIN;
|
||||
console.log('Set chain:', LLM_CHAIN);
|
||||
} else if (chainType === CONVERSATIONAL_RETRIEVAL_QA_CHAIN && options.noteContent) {
|
||||
const textSplitter = new RecursiveCharacterTextSplitter({ chunkSize: 1000 });
|
||||
const docs = await textSplitter.createDocuments([options.noteContent]);
|
||||
console.log('docs:', docs);
|
||||
const vectorStore = await MemoryVectorStore.fromDocuments(
|
||||
docs,
|
||||
new HuggingFaceInferenceEmbeddings({
|
||||
apiKey: this.langChainParams.huggingfaceApiKey,
|
||||
}),
|
||||
);
|
||||
/* Create or retrieve the chain */
|
||||
AIState.retrievalChain = ChainFactory.getRetrievalChain({
|
||||
llm: AIState.chatOpenAI,
|
||||
retriever: vectorStore.asRetriever(),
|
||||
});
|
||||
// Issue where conversational retrieval chain gives rephrased question
|
||||
// when streaming: https://github.com/hwchase17/langchainjs/issues/754#issuecomment-1540257078
|
||||
// Temp workaround triggers CORS issue 'refused to set header user-agent'
|
||||
// Wait for official fix.
|
||||
AIState.useChain = CONVERSATIONAL_RETRIEVAL_QA_CHAIN;
|
||||
console.log('Set chain:', CONVERSATIONAL_RETRIEVAL_QA_CHAIN);
|
||||
}
|
||||
}
|
||||
|
||||
async countTokens(inputStr: string): Promise<number> {
|
||||
// TODO: This is currently falling back to an approximation. Follow up with LangchainJS:
|
||||
// https://github.com/hwchase17/langchainjs/issues/985
|
||||
return AIState.chatOpenAI.getNumTokens(inputStr);
|
||||
}
|
||||
|
||||
|
|
@ -135,32 +189,55 @@ class AIState {
|
|||
|
||||
async runChain(
|
||||
userMessage: string,
|
||||
chatContext: ChatMessage[],
|
||||
abortController: AbortController,
|
||||
updateCurrentAiMessage: (message: string) => void,
|
||||
addMessage: (message: ChatMessage) => void,
|
||||
debug = false,
|
||||
) {
|
||||
if (debug) {
|
||||
console.log('Chat memory:', this.memory);
|
||||
}
|
||||
let fullAIResponse = '';
|
||||
// TODO: chain.call stop signal gives error:
|
||||
// "input values have 2 keys, you must specify an input key or pass only 1 key as input".
|
||||
// Follow up with LangchainJS: https://github.com/hwchase17/langchainjs/issues/1327
|
||||
await AIState.chain.call(
|
||||
{
|
||||
input: userMessage,
|
||||
// signal: abortController.signal,
|
||||
},
|
||||
[
|
||||
{
|
||||
handleLLMNewToken: (token) => {
|
||||
fullAIResponse += token;
|
||||
updateCurrentAiMessage(fullAIResponse);
|
||||
}
|
||||
switch(AIState.useChain) {
|
||||
case LLM_CHAIN:
|
||||
if (debug) {
|
||||
console.log('Chat memory:', this.memory);
|
||||
}
|
||||
]
|
||||
);
|
||||
// TODO: chain.call stop signal gives error:
|
||||
// "input values have 2 keys, you must specify an input key or pass only 1 key as input".
|
||||
// Follow up with LangchainJS: https://github.com/hwchase17/langchainjs/issues/1327
|
||||
await AIState.chain.call(
|
||||
{
|
||||
input: userMessage,
|
||||
// signal: abortController.signal,
|
||||
},
|
||||
[
|
||||
{
|
||||
handleLLMNewToken: (token) => {
|
||||
fullAIResponse += token;
|
||||
updateCurrentAiMessage(fullAIResponse);
|
||||
}
|
||||
}
|
||||
]
|
||||
);
|
||||
break;
|
||||
case CONVERSATIONAL_RETRIEVAL_QA_CHAIN:
|
||||
await AIState.retrievalChain.call(
|
||||
{
|
||||
question: userMessage,
|
||||
chat_history: chatContext,
|
||||
},
|
||||
[
|
||||
{
|
||||
handleLLMNewToken: (token) => {
|
||||
fullAIResponse += token;
|
||||
updateCurrentAiMessage(fullAIResponse);
|
||||
}
|
||||
}
|
||||
]
|
||||
);
|
||||
break;
|
||||
default:
|
||||
console.error('Chain type not supported:', AIState.useChain);
|
||||
}
|
||||
|
||||
addMessage({
|
||||
message: fullAIResponse,
|
||||
|
|
|
|||
58
src/chainFactory.ts
Normal file
58
src/chainFactory.ts
Normal file
|
|
@ -0,0 +1,58 @@
|
|||
import { BaseLanguageModel } from "langchain/base_language";
|
||||
import {
|
||||
BaseChain,
|
||||
ConversationChain,
|
||||
ConversationalRetrievalQAChain,
|
||||
LLMChainInput,
|
||||
} from "langchain/chains";
|
||||
import { BaseRetriever } from "langchain/schema";
|
||||
|
||||
|
||||
export interface ConversationalRetrievalChainParams {
|
||||
llm: BaseLanguageModel;
|
||||
retriever: BaseRetriever;
|
||||
options?: {
|
||||
questionGeneratorTemplate?: string;
|
||||
qaTemplate?: string;
|
||||
returnSourceDocuments?: boolean;
|
||||
}
|
||||
}
|
||||
|
||||
// Add new chain types here
|
||||
export const LLM_CHAIN = 'llm_chain';
|
||||
export const CONVERSATIONAL_RETRIEVAL_QA_CHAIN = 'conversational_retrieval_chain';
|
||||
export const SUPPORTED_CHAIN_TYPES = new Set([
|
||||
LLM_CHAIN,
|
||||
CONVERSATIONAL_RETRIEVAL_QA_CHAIN,
|
||||
]);
|
||||
|
||||
class ChainFactory {
|
||||
private static instances: Map<string, BaseChain> = new Map();
|
||||
|
||||
public static getLLMChain(args: LLMChainInput): BaseChain {
|
||||
let instance = ChainFactory.instances.get(LLM_CHAIN);
|
||||
if (!instance) {
|
||||
instance = new ConversationChain(args as LLMChainInput);
|
||||
console.log('New chain created: ', instance._chainType());
|
||||
ChainFactory.instances.set(LLM_CHAIN, instance);
|
||||
}
|
||||
return instance;
|
||||
}
|
||||
|
||||
public static getRetrievalChain(
|
||||
args: ConversationalRetrievalChainParams
|
||||
): ConversationalRetrievalQAChain {
|
||||
let instance = ChainFactory.instances.get(CONVERSATIONAL_RETRIEVAL_QA_CHAIN);
|
||||
if (!instance) {
|
||||
const argsRetrieval = args as ConversationalRetrievalChainParams;
|
||||
instance = ConversationalRetrievalQAChain.fromLLM(
|
||||
argsRetrieval.llm, argsRetrieval.retriever, argsRetrieval.options
|
||||
);
|
||||
console.log('New chain created: ', instance._chainType());
|
||||
ChainFactory.instances.set(CONVERSATIONAL_RETRIEVAL_QA_CHAIN, instance);
|
||||
}
|
||||
return instance as ConversationalRetrievalQAChain;
|
||||
}
|
||||
}
|
||||
|
||||
export default ChainFactory;
|
||||
|
|
@ -28,10 +28,10 @@ import {
|
|||
simplifyPrompt,
|
||||
summarizePrompt,
|
||||
tocPrompt,
|
||||
useNoteAsContextPrompt
|
||||
useNoteAsContextPrompt,
|
||||
} from '@/utils';
|
||||
import { EventEmitter } from 'events';
|
||||
import { TFile } from 'obsidian';
|
||||
import { Notice, TFile } from 'obsidian';
|
||||
import React, {
|
||||
useContext,
|
||||
useEffect,
|
||||
|
|
@ -124,12 +124,31 @@ const Chat: React.FC<ChatProps> = ({
|
|||
|
||||
const file = app.workspace.getActiveFile();
|
||||
if (!file) {
|
||||
new Notice('No active note found.');
|
||||
console.error('No active note found.');
|
||||
return;
|
||||
}
|
||||
const noteContent = await getFileContent(file);
|
||||
const noteName = getFileName(file);
|
||||
|
||||
/* TODO: Make a switch for unlimited context search, on and off. When turned on, this
|
||||
message is shown in both notice and console: Unlimited Context Enabled!
|
||||
*/
|
||||
// const activeNoteOnMessage: ChatMessage = {
|
||||
// sender: AI_SENDER,
|
||||
// message: `OK please ask me questions about [[${noteName}]]`,
|
||||
// isVisible: true,
|
||||
// };
|
||||
// addMessage(activeNoteOnMessage);
|
||||
|
||||
// if (noteContent) {
|
||||
// aiState.setChain(CONVERSATIONAL_RETRIEVAL_QA_CHAIN, { noteContent });
|
||||
// } else {
|
||||
// new Notice('No note content found.');
|
||||
// console.error('No note content found.');
|
||||
// return;
|
||||
// }
|
||||
|
||||
// Set the context based on the noteContent
|
||||
const prompt = useNoteAsContextPrompt(noteName, noteContent);
|
||||
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ export const AI_SENDER = 'ai';
|
|||
export const DEFAULT_SYSTEM_PROMPT = 'You are Obsidian Copilot, a helpful assistant that integrates AI to Obsidian note-taking.';
|
||||
export const DEFAULT_SETTINGS: CopilotSettings = {
|
||||
openAiApiKey: '',
|
||||
huggingfaceApiKey: '',
|
||||
defaultModel: 'gpt-3.5-turbo',
|
||||
temperature: '0.7',
|
||||
maxTokens: '1000',
|
||||
|
|
|
|||
|
|
@ -37,7 +37,17 @@ export const getAIResponse = async (
|
|||
abortController,
|
||||
updateCurrentAiMessage,
|
||||
addMessage,
|
||||
debug,
|
||||
);
|
||||
|
||||
// await aiState.runChain(
|
||||
// userMessage.message,
|
||||
// chatContext,
|
||||
// abortController,
|
||||
// updateCurrentAiMessage,
|
||||
// addMessage,
|
||||
// debug,
|
||||
// );
|
||||
} catch (error) {
|
||||
const errorData = error?.response?.data?.error || error;
|
||||
const errorCode = errorData?.code || error;
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ import { Editor, Notice, Plugin, WorkspaceLeaf } from 'obsidian';
|
|||
|
||||
export interface CopilotSettings {
|
||||
openAiApiKey: string;
|
||||
huggingfaceApiKey: string;
|
||||
defaultModel: string;
|
||||
temperature: string;
|
||||
maxTokens: string;
|
||||
|
|
@ -257,12 +258,14 @@ export default class CopilotPlugin extends Plugin {
|
|||
getAIStateParams(): LangChainParams {
|
||||
const {
|
||||
openAiApiKey,
|
||||
huggingfaceApiKey,
|
||||
temperature,
|
||||
maxTokens,
|
||||
contextTurns,
|
||||
} = sanitizeSettings(this.settings);
|
||||
return {
|
||||
key: openAiApiKey,
|
||||
huggingfaceApiKey: huggingfaceApiKey,
|
||||
model: this.settings.defaultModel,
|
||||
temperature: Number(temperature),
|
||||
maxTokens: Number(maxTokens),
|
||||
|
|
|
|||
|
|
@ -173,6 +173,34 @@ export class CopilotSettingTab extends PluginSettingTab {
|
|||
// })
|
||||
// );
|
||||
|
||||
// containerEl.createEl('h4', {text: 'Other API Settings'});
|
||||
|
||||
// new Setting(containerEl)
|
||||
// .setName("Your Huggingface Inference API key")
|
||||
// .setDesc(
|
||||
// createFragment((frag) => {
|
||||
// frag.appendText("You can find your API key at ");
|
||||
// frag.createEl('a', {
|
||||
// text: "https://hf.co/settings/tokens",
|
||||
// href: "https://hf.co/settings/tokens"
|
||||
// });
|
||||
// frag.createEl('br');
|
||||
// frag.appendText("It is used to make requests to Huggingface Inference API for vector search (BETA).");
|
||||
// })
|
||||
// )
|
||||
// .addText((text) =>{
|
||||
// text.inputEl.type = "password";
|
||||
// text.inputEl.style.width = "80%";
|
||||
// text
|
||||
// .setPlaceholder("Huggingface Inference API key")
|
||||
// .setValue(this.plugin.settings.huggingfaceApiKey)
|
||||
// .onChange(async (value) => {
|
||||
// this.plugin.settings.huggingfaceApiKey = value;
|
||||
// await this.plugin.saveSettings();
|
||||
// })
|
||||
// }
|
||||
// );
|
||||
|
||||
containerEl.createEl('h4', {text: 'Advanced Settings'});
|
||||
|
||||
new Setting(containerEl)
|
||||
|
|
|
|||
Loading…
Reference in a new issue