diff --git a/.github/workflows/backend-review.yml b/.github/workflows/backend-review.yml index 7c79cadce8..e5c86caa4c 100644 --- a/.github/workflows/backend-review.yml +++ b/.github/workflows/backend-review.yml @@ -1,10 +1,5 @@ name: Backend Unit Tests on: - # push: - # branches: - # - main - # - dev - # - release/* pull_request: branches: - main @@ -23,6 +18,7 @@ jobs: JWT_SECRET: ${{ secrets.JWT_SECRET }} CREDS_KEY: ${{ secrets.CREDS_KEY }} CREDS_IV: ${{ secrets.CREDS_IV }} + NODE_ENV: ci steps: - uses: actions/checkout@v2 - name: Use Node.js 20.x @@ -34,8 +30,8 @@ jobs: - name: Install dependencies run: npm ci - # - name: Install Linux X64 Sharp - # run: npm install --platform=linux --arch=x64 --verbose sharp + - name: Install Data Provider + run: npm run build:data-provider - name: Run unit tests run: cd api && npm run test:ci diff --git a/api/app/clients/AnthropicClient.js b/api/app/clients/AnthropicClient.js index cf9571c69b..ebec514040 100644 --- a/api/app/clients/AnthropicClient.js +++ b/api/app/clients/AnthropicClient.js @@ -1,4 +1,3 @@ -const Keyv = require('keyv'); // const { Agent, ProxyAgent } = require('undici'); const BaseClient = require('./BaseClient'); const { @@ -15,8 +14,6 @@ const tokenizersCache = {}; class AnthropicClient extends BaseClient { constructor(apiKey, options = {}, cacheOptions = {}) { super(apiKey, options, cacheOptions); - cacheOptions.namespace = cacheOptions.namespace || 'anthropic'; - this.conversationsCache = new Keyv(cacheOptions); this.apiKey = apiKey || process.env.ANTHROPIC_API_KEY; this.sender = 'Anthropic'; this.userLabel = HUMAN_PROMPT; @@ -107,6 +104,23 @@ class AnthropicClient extends BaseClient { content: message?.content ?? message.text, })); + let lastAuthor = ''; + let groupedMessages = []; + + for (let message of formattedMessages) { + // If last author is not same as current author, add to new group + if (lastAuthor !== message.author) { + groupedMessages.push({ + author: message.author, + content: [message.content], + }); + lastAuthor = message.author; + // If same author, append content to the last group + } else { + groupedMessages[groupedMessages.length - 1].content.push(message.content); + } + } + let identityPrefix = ''; if (this.options.userLabel) { identityPrefix = `\nHuman's name: ${this.options.userLabel}`; @@ -129,8 +143,12 @@ class AnthropicClient extends BaseClient { promptPrefix = `${identityPrefix}${promptPrefix}`; } - const promptSuffix = `${promptPrefix}${this.assistantLabel}\n`; // Prompt AI to respond. - let currentTokenCount = this.getTokenCount(promptSuffix); + // Prompt AI to respond, empty if last message was from AI + let isEdited = lastAuthor === this.assistantLabel; + const promptSuffix = isEdited ? '' : `${promptPrefix}${this.assistantLabel}\n`; + let currentTokenCount = isEdited + ? this.getTokenCount(promptPrefix) + : this.getTokenCount(promptSuffix); let promptBody = ''; const maxTokenCount = this.maxPromptTokens; @@ -148,10 +166,13 @@ class AnthropicClient extends BaseClient { }; const buildPromptBody = async () => { - if (currentTokenCount < maxTokenCount && formattedMessages.length > 0) { - const message = formattedMessages.pop(); + if (currentTokenCount < maxTokenCount && groupedMessages.length > 0) { + const message = groupedMessages.pop(); const isCreatedByUser = message.author === this.userLabel; - const messageString = `${message.author}\n${message.content}${this.endToken}\n`; + // Use promptPrefix if message is edited assistant' + const messagePrefix = + isCreatedByUser || !isEdited ? message.author : `${promptPrefix}${message.author}`; + const messageString = `${messagePrefix}\n${message.content}${this.endToken}\n`; let newPromptBody = `${messageString}${promptBody}`; context.unshift(message); @@ -182,6 +203,12 @@ class AnthropicClient extends BaseClient { } promptBody = newPromptBody; currentTokenCount = newTokenCount; + + // Switch off isEdited after using it for the first time + if (isEdited) { + isEdited = false; + } + // wait for next tick to avoid blocking the event loop await new Promise((resolve) => setImmediate(resolve)); return buildPromptBody(); @@ -197,7 +224,8 @@ class AnthropicClient extends BaseClient { context.shift(); } - const prompt = `${promptBody}${promptSuffix}`; + let prompt = `${promptBody}${promptSuffix}`; + // Add 2 tokens for metadata after all messages have been counted. currentTokenCount += 2; diff --git a/api/app/clients/BaseClient.js b/api/app/clients/BaseClient.js index baaa0990d3..5d55d33fdc 100644 --- a/api/app/clients/BaseClient.js +++ b/api/app/clients/BaseClient.js @@ -5,11 +5,12 @@ const { ChatOpenAI } = require('langchain/chat_models/openai'); const { loadSummarizationChain } = require('langchain/chains'); const { refinePrompt } = require('./prompts/refinePrompt'); const { getConvo, getMessages, saveMessage, updateMessage, saveConvo } = require('../../models'); +const { addSpaceIfNeeded } = require('../../server/utils'); class BaseClient { constructor(apiKey, options = {}) { this.apiKey = apiKey; - this.sender = options.sender || 'AI'; + this.sender = options.sender ?? 'AI'; this.contextStrategy = null; this.currentDateString = new Date().toLocaleDateString('en-us', { year: 'numeric', @@ -51,18 +52,20 @@ class BaseClient { if (opts && typeof opts === 'object') { this.setOptions(opts); } - const user = opts.user || null; - const conversationId = opts.conversationId || crypto.randomUUID(); - const parentMessageId = opts.parentMessageId || '00000000-0000-0000-0000-000000000000'; - const userMessageId = opts.overrideParentMessageId || crypto.randomUUID(); - const responseMessageId = crypto.randomUUID(); + const user = opts.user ?? null; + const conversationId = opts.conversationId ?? crypto.randomUUID(); + const parentMessageId = opts.parentMessageId ?? '00000000-0000-0000-0000-000000000000'; + const userMessageId = opts.overrideParentMessageId ?? crypto.randomUUID(); + const responseMessageId = opts.responseMessageId ?? crypto.randomUUID(); const saveOptions = this.getSaveOptions(); - this.abortController = opts.abortController || new AbortController(); - this.currentMessages = (await this.loadHistory(conversationId, parentMessageId)) ?? []; + const head = opts.isEdited ? responseMessageId : parentMessageId; + this.currentMessages = (await this.loadHistory(conversationId, head)) ?? []; + this.abortController = opts.abortController ?? new AbortController(); return { ...opts, user, + head, conversationId, parentMessageId, userMessageId, @@ -72,7 +75,7 @@ class BaseClient { } createUserMessage({ messageId, parentMessageId, conversationId, text }) { - const userMessage = { + return { messageId, parentMessageId, conversationId, @@ -80,19 +83,27 @@ class BaseClient { text, isCreatedByUser: true, }; - return userMessage; } async handleStartMethods(message, opts) { - const { user, conversationId, parentMessageId, userMessageId, responseMessageId, saveOptions } = - await this.setMessageOptions(opts); - - const userMessage = this.createUserMessage({ - messageId: userMessageId, - parentMessageId, + const { + user, + head, conversationId, - text: message, - }); + parentMessageId, + userMessageId, + responseMessageId, + saveOptions, + } = await this.setMessageOptions(opts); + + const userMessage = opts.isEdited + ? this.currentMessages[this.currentMessages.length - 2] + : this.createUserMessage({ + messageId: userMessageId, + parentMessageId, + conversationId, + text: message, + }); if (typeof opts?.getIds === 'function') { opts.getIds({ @@ -109,6 +120,7 @@ class BaseClient { return { ...opts, user, + head, conversationId, responseMessageId, saveOptions, @@ -373,7 +385,7 @@ class BaseClient { if (this.options.debug) { console.debug('<-------------------------PAYLOAD/TOKEN COUNT MAP------------------------->'); - console.debug('Payload:', payload); + // console.debug('Payload:', payload); console.debug('Token Count Map:', tokenCountMap); console.debug('Prompt Tokens', promptTokens, remainingContextTokens, this.maxContextTokens); } @@ -382,13 +394,16 @@ class BaseClient { } async sendMessage(message, opts = {}) { - const { user, conversationId, responseMessageId, saveOptions, userMessage } = + const { user, head, isEdited, conversationId, responseMessageId, saveOptions, userMessage } = await this.handleStartMethods(message, opts); this.user = user; // It's not necessary to push to currentMessages // depending on subclass implementation of handling messages - this.currentMessages.push(userMessage); + // When this is an edit, all messages are already in currentMessages, both user and response + if (!isEdited) { + this.currentMessages.push(userMessage); + } let { prompt: payload, @@ -398,13 +413,13 @@ class BaseClient { this.currentMessages, // When the userMessage is pushed to currentMessages, the parentMessage is the userMessageId. // this only matters when buildMessages is utilizing the parentMessageId, and may vary on implementation - userMessage.messageId, + isEdited ? head : userMessage.messageId, this.getBuildMessagesOptions(opts), ); if (this.options.debug) { console.debug('payload'); - console.debug(payload); + // console.debug(payload); } if (tokenCountMap) { @@ -423,7 +438,11 @@ class BaseClient { this.handleTokenCountMap(tokenCountMap); } - await this.saveMessageToDatabase(userMessage, saveOptions, user); + if (!isEdited) { + await this.saveMessageToDatabase(userMessage, saveOptions, user); + } + + const generation = isEdited ? this.currentMessages[this.currentMessages.length - 1].text : ''; const responseMessage = { messageId: responseMessageId, conversationId, @@ -431,7 +450,7 @@ class BaseClient { isCreatedByUser: false, model: this.modelOptions.model, sender: this.sender, - text: await this.sendCompletion(payload, opts), + text: addSpaceIfNeeded(generation) + (await this.sendCompletion(payload, opts)), promptTokens, }; @@ -453,7 +472,7 @@ class BaseClient { console.debug('Loading history for conversation', conversationId, parentMessageId); } - const messages = (await getMessages({ conversationId })) || []; + const messages = (await getMessages({ conversationId })) ?? []; if (messages.length === 0) { return []; diff --git a/api/app/clients/OpenAIClient.js b/api/app/clients/OpenAIClient.js index 53f4815d74..87b25b1226 100644 --- a/api/app/clients/OpenAIClient.js +++ b/api/app/clients/OpenAIClient.js @@ -314,6 +314,7 @@ class OpenAIClient extends BaseClient { async sendCompletion(payload, opts = {}) { let reply = ''; let result = null; + let streamResult = null; if (typeof opts.onProgress === 'function') { await this.getCompletion( payload, @@ -321,6 +322,10 @@ class OpenAIClient extends BaseClient { if (progressMessage === '[DONE]') { return; } + + if (progressMessage.choices) { + streamResult = progressMessage; + } const token = this.isChatCompletion ? progressMessage.choices?.[0]?.delta?.content : progressMessage.choices?.[0]?.text; @@ -355,6 +360,10 @@ class OpenAIClient extends BaseClient { } } + if (streamResult && typeof opts.addMetadata === 'function') { + const { finish_reason } = streamResult.choices[0]; + opts.addMetadata({ finish_reason }); + } return reply.trim(); } diff --git a/api/app/clients/PluginsClient.js b/api/app/clients/PluginsClient.js index d29f03517a..dc427b4ae4 100644 --- a/api/app/clients/PluginsClient.js +++ b/api/app/clients/PluginsClient.js @@ -345,7 +345,8 @@ Only respond with your conversational reply to the following User Message: } async sendMessage(message, opts = {}) { - const completionMode = this.options.tools.length === 0; + // If a message is edited, no tools can be used. + const completionMode = this.options.tools.length === 0 || opts.isEdited; if (completionMode) { this.setOptions(opts); return super.sendMessage(message, opts); diff --git a/api/app/clients/specs/AnthropicClient.test.js b/api/app/clients/specs/AnthropicClient.test.js new file mode 100644 index 0000000000..52324914b9 --- /dev/null +++ b/api/app/clients/specs/AnthropicClient.test.js @@ -0,0 +1,139 @@ +const AnthropicClient = require('../AnthropicClient'); +const HUMAN_PROMPT = '\n\nHuman:'; +const AI_PROMPT = '\n\nAssistant:'; + +describe('AnthropicClient', () => { + let client; + const model = 'claude-2'; + const parentMessageId = '1'; + const messages = [ + { role: 'user', isCreatedByUser: true, text: 'Hello', messageId: parentMessageId }, + { role: 'assistant', isCreatedByUser: false, text: 'Hi', messageId: '2', parentMessageId }, + { + role: 'user', + isCreatedByUser: true, + text: 'What\'s up', + messageId: '3', + parentMessageId: '2', + }, + ]; + + beforeEach(() => { + const options = { + modelOptions: { + model, + temperature: 0.7, + }, + }; + client = new AnthropicClient('test-api-key'); + client.setOptions(options); + }); + + describe('setOptions', () => { + it('should set the options correctly', () => { + expect(client.apiKey).toBe('test-api-key'); + expect(client.modelOptions.model).toBe(model); + expect(client.modelOptions.temperature).toBe(0.7); + }); + }); + + describe('getSaveOptions', () => { + it('should return the correct save options', () => { + const options = client.getSaveOptions(); + expect(options).toHaveProperty('modelLabel'); + expect(options).toHaveProperty('promptPrefix'); + }); + }); + + describe('buildMessages', () => { + it('should handle promptPrefix from options when promptPrefix argument is not provided', async () => { + client.options.promptPrefix = 'Test Prefix from options'; + const result = await client.buildMessages(messages, parentMessageId); + const { prompt } = result; + expect(prompt).toContain('Test Prefix from options'); + }); + + it('should build messages correctly for chat completion', async () => { + const result = await client.buildMessages(messages, '2'); + expect(result).toHaveProperty('prompt'); + expect(result.prompt).toContain(HUMAN_PROMPT); + expect(result.prompt).toContain('Hello'); + expect(result.prompt).toContain(AI_PROMPT); + expect(result.prompt).toContain('Hi'); + }); + + it('should group messages by the same author', async () => { + const groupedMessages = messages.map((m) => ({ ...m, isCreatedByUser: true, role: 'user' })); + const result = await client.buildMessages(groupedMessages, '3'); + expect(result.context).toHaveLength(1); + + // Check that HUMAN_PROMPT appears only once in the prompt + const matches = result.prompt.match(new RegExp(HUMAN_PROMPT, 'g')); + expect(matches).toHaveLength(1); + + groupedMessages.push({ + role: 'assistant', + isCreatedByUser: false, + text: 'I heard you the first time', + messageId: '4', + parentMessageId: '3', + }); + + const result2 = await client.buildMessages(groupedMessages, '4'); + expect(result2.context).toHaveLength(2); + + // Check that HUMAN_PROMPT appears only once in the prompt + const human_matches = result2.prompt.match(new RegExp(HUMAN_PROMPT, 'g')); + const ai_matches = result2.prompt.match(new RegExp(AI_PROMPT, 'g')); + expect(human_matches).toHaveLength(1); + expect(ai_matches).toHaveLength(1); + }); + + it('should handle isEdited condition', async () => { + const editedMessages = [ + { role: 'user', isCreatedByUser: true, text: 'Hello', messageId: '1' }, + { role: 'assistant', isCreatedByUser: false, text: 'Hi', messageId: '2', parentMessageId }, + ]; + + const trimmedLabel = AI_PROMPT.trim(); + const result = await client.buildMessages(editedMessages, '2'); + expect(result.prompt.trim().endsWith(trimmedLabel)).toBeFalsy(); + + // Add a human message at the end to test the opposite + editedMessages.push({ + role: 'user', + isCreatedByUser: true, + text: 'Hi again', + messageId: '3', + parentMessageId: '2', + }); + const result2 = await client.buildMessages(editedMessages, '3'); + expect(result2.prompt.trim().endsWith(trimmedLabel)).toBeTruthy(); + }); + + it('should build messages correctly with a promptPrefix', async () => { + const promptPrefix = 'Test Prefix'; + client.options.promptPrefix = promptPrefix; + const result = await client.buildMessages(messages, parentMessageId); + const { prompt } = result; + expect(prompt).toBeDefined(); + expect(prompt).toContain(promptPrefix); + const textAfterPrefix = prompt.split(promptPrefix)[1]; + expect(textAfterPrefix).toContain(AI_PROMPT); + + const editedMessages = messages.slice(0, -1); + const result2 = await client.buildMessages(editedMessages, parentMessageId); + const textAfterPrefix2 = result2.prompt.split(promptPrefix)[1]; + expect(textAfterPrefix2).toContain(AI_PROMPT); + }); + + it('should handle identityPrefix from options', async () => { + client.options.userLabel = 'John'; + client.options.modelLabel = 'Claude-2'; + const result = await client.buildMessages(messages, parentMessageId); + const { prompt } = result; + expect(prompt).toContain('Human\'s name: John'); + expect(prompt).toContain('You are Claude-2'); + }); + }); +}); diff --git a/api/app/clients/specs/BaseClient.test.js b/api/app/clients/specs/BaseClient.test.js index d81bfe6274..10d5868cb0 100644 --- a/api/app/clients/specs/BaseClient.test.js +++ b/api/app/clients/specs/BaseClient.test.js @@ -45,6 +45,18 @@ const fakeMessages = []; const userMessage = 'Hello, ChatGPT!'; const apiKey = 'fake-api-key'; +const messageHistory = [ + { role: 'user', isCreatedByUser: true, text: 'Hello', messageId: '1' }, + { role: 'assistant', isCreatedByUser: false, text: 'Hi', messageId: '2', parentMessageId: '1' }, + { + role: 'user', + isCreatedByUser: true, + text: 'What\'s up', + messageId: '3', + parentMessageId: '2', + }, +]; + describe('BaseClient', () => { let TestClient; const options = { @@ -277,9 +289,54 @@ describe('BaseClient', () => { }); test('should return chat history', async () => { - const chatMessages = await TestClient.loadHistory(conversationId, parentMessageId); - expect(TestClient.currentMessages).toHaveLength(4); - expect(chatMessages[0].text).toEqual(userMessage); + TestClient = initializeFakeClient(apiKey, options, messageHistory); + const chatMessages = await TestClient.loadHistory(conversationId, '2'); + expect(TestClient.currentMessages).toHaveLength(2); + expect(chatMessages[0].text).toEqual('Hello'); + + const chatMessages2 = await TestClient.loadHistory(conversationId, '3'); + expect(TestClient.currentMessages).toHaveLength(3); + expect(chatMessages2[chatMessages2.length - 1].text).toEqual('What\'s up'); + }); + + /* Most of the new sendMessage logic revolving around edited/continued AI messages + * can be summarized by the following test. The condition will load the entire history up to + * the message that is being edited, which will trigger the AI API to 'continue' the response. + * The 'userMessage' is only passed by convention and is not necessary for the generation. + */ + it('should not push userMessage to currentMessages when isEdited is true and vice versa', async () => { + const overrideParentMessageId = 'user-message-id'; + const responseMessageId = 'response-message-id'; + const newHistory = messageHistory.slice(); + newHistory.push({ + role: 'assistant', + isCreatedByUser: false, + text: 'test message', + messageId: responseMessageId, + parentMessageId: '3', + }); + + TestClient = initializeFakeClient(apiKey, options, newHistory); + const sendMessageOptions = { + isEdited: true, + overrideParentMessageId, + parentMessageId: '3', + responseMessageId, + }; + + await TestClient.sendMessage('test message', sendMessageOptions); + const currentMessages = TestClient.currentMessages; + expect(currentMessages[currentMessages.length - 1].messageId).not.toEqual( + overrideParentMessageId, + ); + + // Test the opposite case + sendMessageOptions.isEdited = false; + await TestClient.sendMessage('test message', sendMessageOptions); + const currentMessages2 = TestClient.currentMessages; + expect(currentMessages2[currentMessages2.length - 1].messageId).toEqual( + overrideParentMessageId, + ); }); test('setOptions is called with the correct arguments', async () => { diff --git a/api/app/clients/specs/FakeClient.js b/api/app/clients/specs/FakeClient.js index 5cd7556bcf..e46cee687e 100644 --- a/api/app/clients/specs/FakeClient.js +++ b/api/app/clients/specs/FakeClient.js @@ -1,4 +1,3 @@ -const crypto = require('crypto'); const BaseClient = require('../BaseClient'); const { maxTokensMap } = require('../../../utils'); @@ -87,86 +86,6 @@ const initializeFakeClient = (apiKey, options, fakeMessages) => { return 'Mock response text'; }); - TestClient.sendMessage = jest.fn().mockImplementation(async (message, opts = {}) => { - if (opts && typeof opts === 'object') { - TestClient.setOptions(opts); - } - - const user = opts.user || null; - const conversationId = opts.conversationId || crypto.randomUUID(); - const parentMessageId = opts.parentMessageId || '00000000-0000-0000-0000-000000000000'; - const userMessageId = opts.overrideParentMessageId || crypto.randomUUID(); - const saveOptions = TestClient.getSaveOptions(); - - this.pastMessages = await TestClient.loadHistory( - conversationId, - TestClient.options?.parentMessageId, - ); - - const userMessage = { - text: message, - sender: TestClient.sender, - isCreatedByUser: true, - messageId: userMessageId, - parentMessageId, - conversationId, - }; - - const response = { - sender: TestClient.sender, - text: 'Hello, User!', - isCreatedByUser: false, - messageId: crypto.randomUUID(), - parentMessageId: userMessage.messageId, - conversationId, - }; - - fakeMessages.push(userMessage); - fakeMessages.push(response); - - if (typeof opts.getIds === 'function') { - opts.getIds({ - userMessage, - conversationId, - responseMessageId: response.messageId, - }); - } - - if (typeof opts.onStart === 'function') { - opts.onStart(userMessage); - } - - let { prompt: payload, tokenCountMap } = await TestClient.buildMessages( - this.currentMessages, - userMessage.messageId, - TestClient.getBuildMessagesOptions(opts), - ); - - if (tokenCountMap) { - payload = payload.map((message, i) => { - const { tokenCount, ...messageWithoutTokenCount } = message; - // userMessage is always the last one in the payload - if (i === payload.length - 1) { - userMessage.tokenCount = message.tokenCount; - console.debug( - `Token count for user message: ${tokenCount}`, - `Instruction Tokens: ${tokenCountMap.instructions || 'N/A'}`, - ); - } - return messageWithoutTokenCount; - }); - TestClient.handleTokenCountMap(tokenCountMap); - } - - await TestClient.saveMessageToDatabase(userMessage, saveOptions, user); - response.text = await TestClient.sendCompletion(payload, opts); - if (tokenCountMap && TestClient.getTokenCountForResponse) { - response.tokenCount = TestClient.getTokenCountForResponse(response); - } - await TestClient.saveMessageToDatabase(response, saveOptions, user); - return response; - }); - TestClient.buildMessages = jest.fn(async (messages, parentMessageId) => { const orderedMessages = TestClient.constructor.getMessagesForConversation( messages, diff --git a/api/app/clients/specs/OpenAIClient.test.js b/api/app/clients/specs/OpenAIClient.test.js index 41aeb4f3b4..dd4de5cc7c 100644 --- a/api/app/clients/specs/OpenAIClient.test.js +++ b/api/app/clients/specs/OpenAIClient.test.js @@ -1,5 +1,7 @@ const OpenAIClient = require('../OpenAIClient'); +jest.mock('meilisearch'); + describe('OpenAIClient', () => { let client, client2; const model = 'gpt-4'; @@ -25,6 +27,9 @@ describe('OpenAIClient', () => { content: 'Refined answer', tokenCount: 30, }); + client.buildPrompt = jest + .fn() + .mockResolvedValue({ prompt: messages.map((m) => m.text).join('\n') }); client.constructor.freeAndResetAllEncoders(); }); diff --git a/api/app/clients/specs/PluginsClient.test.js b/api/app/clients/specs/PluginsClient.test.js index 59218c6206..a11d490b22 100644 --- a/api/app/clients/specs/PluginsClient.test.js +++ b/api/app/clients/specs/PluginsClient.test.js @@ -111,7 +111,6 @@ describe('PluginsClient', () => { }); const response = await TestAgent.sendMessage(userMessage); - console.log(response); parentMessageId = response.messageId; conversationId = response.conversationId; expect(response).toEqual(expectedResult); diff --git a/api/app/clients/tools/util/addOpenAPISpecs.js b/api/app/clients/tools/util/addOpenAPISpecs.js index 2d5756f194..8b87be9941 100644 --- a/api/app/clients/tools/util/addOpenAPISpecs.js +++ b/api/app/clients/tools/util/addOpenAPISpecs.js @@ -20,7 +20,6 @@ async function addOpenAPISpecs(availableTools) { } return availableTools; } catch (error) { - console.log('addOpenAPISpecs error', error); return availableTools; } } diff --git a/api/app/clients/tools/util/handleTools.test.js b/api/app/clients/tools/util/handleTools.test.js index 674543ba29..f0c6ee6601 100644 --- a/api/app/clients/tools/util/handleTools.test.js +++ b/api/app/clients/tools/util/handleTools.test.js @@ -83,7 +83,6 @@ describe('Tool Handlers', () => { it('returns valid tools given input tools and user authentication', async () => { const validTools = await validateTools(fakeUser._id, initialTools); expect(validTools).toBeDefined(); - console.log('validateTools: validTools', validTools); expect(validTools.some((tool) => tool === pluginKey)).toBeTruthy(); expect(validTools.length).toBeGreaterThan(0); }); diff --git a/api/app/index.js b/api/app/index.js index 95624829a9..fe1462331e 100644 --- a/api/app/index.js +++ b/api/app/index.js @@ -3,15 +3,11 @@ const { askBing } = require('./bingai'); const clients = require('./clients'); const titleConvo = require('./titleConvo'); const titleConvoBing = require('./titleConvoBing'); -const getCitations = require('../lib/parse/getCitations'); -const citeText = require('../lib/parse/citeText'); module.exports = { browserClient, askBing, titleConvo, titleConvoBing, - getCitations, - citeText, ...clients, }; diff --git a/api/lib/parse/getCitations.js b/api/lib/parse/getCitations.js deleted file mode 100644 index f99363d145..0000000000 --- a/api/lib/parse/getCitations.js +++ /dev/null @@ -1,18 +0,0 @@ -// const regex = / \[\d+\..*?\]\(.*?\)/g; -const regex = / \[.*?]\(.*?\)/g; - -const getCitations = (res) => { - const adaptiveCards = res.details.adaptiveCards; - const textBlocks = adaptiveCards && adaptiveCards[0].body; - if (!textBlocks) { - return ''; - } - let links = textBlocks[textBlocks.length - 1]?.text.match(regex); - if (links?.length === 0 || !links) { - return ''; - } - links = links.map((link) => link.trim()); - return links.join('\n - '); -}; - -module.exports = getCitations; diff --git a/api/models/Message.js b/api/models/Message.js index 20d06486b8..0fa2265e88 100644 --- a/api/models/Message.js +++ b/api/models/Message.js @@ -14,6 +14,7 @@ module.exports = { error, unfinished, cancelled, + finish_reason = null, tokenCount = null, plugin = null, model = null, @@ -29,6 +30,7 @@ module.exports = { sender, text, isCreatedByUser, + finish_reason, error, unfinished, cancelled, diff --git a/api/models/schema/messageSchema.js b/api/models/schema/messageSchema.js index 792b2d545a..98d6cef239 100644 --- a/api/models/schema/messageSchema.js +++ b/api/models/schema/messageSchema.js @@ -67,6 +67,9 @@ const messageSchema = mongoose.Schema( type: Boolean, default: false, }, + finish_reason: { + type: String, + }, _meiliIndex: { type: Boolean, required: false, diff --git a/api/server/index.js b/api/server/index.js index 2480dc25f5..6ff02656a3 100644 --- a/api/server/index.js +++ b/api/server/index.js @@ -84,6 +84,7 @@ config.validate(); // Validate the config app.use('/api/user', routes.user); app.use('/api/search', routes.search); app.use('/api/ask', routes.ask); + app.use('/api/edit', routes.edit); app.use('/api/messages', routes.messages); app.use('/api/convos', routes.convos); app.use('/api/presets', routes.presets); diff --git a/api/server/middleware/abortControllers.js b/api/server/middleware/abortControllers.js new file mode 100644 index 0000000000..31acbfe389 --- /dev/null +++ b/api/server/middleware/abortControllers.js @@ -0,0 +1,2 @@ +// abortControllers.js +module.exports = new Map(); diff --git a/api/server/middleware/abortMiddleware.js b/api/server/middleware/abortMiddleware.js new file mode 100644 index 0000000000..f678e44868 --- /dev/null +++ b/api/server/middleware/abortMiddleware.js @@ -0,0 +1,106 @@ +const { saveMessage, getConvo, getConvoTitle } = require('../../models'); +const { sendMessage, handleError } = require('../utils'); +const abortControllers = require('./abortControllers'); + +async function abortMessage(req, res) { + const { abortKey } = req.body; + + if (!abortControllers.has(abortKey) && !res.headersSent) { + return res.status(404).send('Request not found'); + } + + const { abortController } = abortControllers.get(abortKey); + const ret = await abortController.abortCompletion(); + console.log('Aborted request', abortKey); + abortControllers.delete(abortKey); + res.send(JSON.stringify(ret)); +} + +const handleAbort = () => { + return async (req, res) => { + try { + return await abortMessage(req, res); + } catch (err) { + console.error(err); + } + }; +}; + +const createAbortController = (res, req, endpointOption, getAbortData) => { + const abortController = new AbortController(); + const onStart = (userMessage) => { + sendMessage(res, { message: userMessage, created: true }); + abortControllers.set(userMessage.conversationId, { abortController, ...endpointOption }); + + res.on('finish', function () { + abortControllers.delete(userMessage.conversationId); + }); + }; + + abortController.abortCompletion = async function () { + abortController.abort(); + const { conversationId, userMessage, ...responseData } = getAbortData(); + + const responseMessage = { + ...responseData, + finish_reason: 'incomplete', + model: endpointOption.modelOptions.model, + unfinished: false, + cancelled: true, + error: false, + }; + + saveMessage(responseMessage); + + return { + title: await getConvoTitle(req.user.id, conversationId), + final: true, + conversation: await getConvo(req.user.id, conversationId), + requestMessage: userMessage, + responseMessage: responseMessage, + }; + }; + + return { abortController, onStart }; +}; + +const handleAbortError = async (res, req, error, data) => { + console.error(error); + const { sender, conversationId, messageId, parentMessageId, partialText } = data; + + const respondWithError = async () => { + const errorMessage = { + sender, + messageId, + conversationId, + parentMessageId, + unfinished: false, + cancelled: false, + error: true, + text: error.message, + }; + if (abortControllers.has(conversationId)) { + const { abortController } = abortControllers.get(conversationId); + abortController.abort(); + abortControllers.delete(conversationId); + } + await saveMessage(errorMessage); + handleError(res, errorMessage); + }; + + if (partialText?.length > 2) { + try { + return await abortMessage(req, res); + } catch (err) { + return respondWithError(); + } + } else { + return respondWithError(); + } +}; + +module.exports = { + handleAbort, + createAbortController, + handleAbortError, +}; diff --git a/api/server/middleware/buildEndpointOption.js b/api/server/middleware/buildEndpointOption.js new file mode 100644 index 0000000000..ea6ad637e8 --- /dev/null +++ b/api/server/middleware/buildEndpointOption.js @@ -0,0 +1,20 @@ +const openAI = require('../routes/endpoints/openAI'); +const gptPlugins = require('../routes/endpoints/gptPlugins'); +const anthropic = require('../routes/endpoints/anthropic'); +const { parseConvo } = require('../routes/endpoints/schemas'); + +const buildFunction = { + openAI: openAI.buildOptions, + azureOpenAI: openAI.buildOptions, + gptPlugins: gptPlugins.buildOptions, + anthropic: anthropic.buildOptions, +}; + +function buildEndpointOption(req, res, next) { + const { endpoint } = req.body; + const parsedBody = parseConvo(endpoint, req.body); + req.body.endpointOption = buildFunction[endpoint](endpoint, parsedBody); + next(); +} + +module.exports = buildEndpointOption; diff --git a/api/server/middleware/index.js b/api/server/middleware/index.js new file mode 100644 index 0000000000..1426260086 --- /dev/null +++ b/api/server/middleware/index.js @@ -0,0 +1,17 @@ +const abortMiddleware = require('./abortMiddleware'); +const setHeaders = require('./setHeaders'); +const requireJwtAuth = require('./requireJwtAuth'); +const requireLocalAuth = require('./requireLocalAuth'); +const validateEndpoint = require('./validateEndpoint'); +const buildEndpointOption = require('./buildEndpointOption'); +const validateRegistration = require('./validateRegistration'); + +module.exports = { + ...abortMiddleware, + setHeaders, + requireJwtAuth, + requireLocalAuth, + validateEndpoint, + buildEndpointOption, + validateRegistration, +}; diff --git a/api/middleware/requireJwtAuth.js b/api/server/middleware/requireJwtAuth.js similarity index 100% rename from api/middleware/requireJwtAuth.js rename to api/server/middleware/requireJwtAuth.js diff --git a/api/middleware/requireLocalAuth.js b/api/server/middleware/requireLocalAuth.js similarity index 92% rename from api/middleware/requireLocalAuth.js rename to api/server/middleware/requireLocalAuth.js index b8700412bd..107d370e85 100644 --- a/api/middleware/requireLocalAuth.js +++ b/api/server/middleware/requireLocalAuth.js @@ -1,5 +1,5 @@ const passport = require('passport'); -const DebugControl = require('../utils/debug.js'); +const DebugControl = require('../../utils/debug.js'); function log({ title, parameters }) { DebugControl.log.functionName(title); diff --git a/api/server/middleware/setHeaders.js b/api/server/middleware/setHeaders.js new file mode 100644 index 0000000000..c1b58e2a5a --- /dev/null +++ b/api/server/middleware/setHeaders.js @@ -0,0 +1,12 @@ +function setHeaders(req, res, next) { + res.writeHead(200, { + Connection: 'keep-alive', + 'Content-Type': 'text/event-stream', + 'Cache-Control': 'no-cache, no-transform', + 'Access-Control-Allow-Origin': '*', + 'X-Accel-Buffering': 'no', + }); + next(); +} + +module.exports = setHeaders; diff --git a/api/server/middleware/validateEndpoint.js b/api/server/middleware/validateEndpoint.js new file mode 100644 index 0000000000..6e9c914c8e --- /dev/null +++ b/api/server/middleware/validateEndpoint.js @@ -0,0 +1,19 @@ +const { handleError } = require('../utils'); + +function validateEndpoint(req, res, next) { + const { endpoint } = req.body; + + if (!req.body.text || req.body.text.length === 0) { + return handleError(res, { text: 'Prompt empty or too short' }); + } + + const pathEndpoint = req.baseUrl.split('/')[3]; + + if (endpoint !== pathEndpoint) { + return handleError(res, { text: 'Illegal request: Endpoint mismatch' }); + } + + next(); +} + +module.exports = validateEndpoint; diff --git a/api/middleware/validateRegistration.js b/api/server/middleware/validateRegistration.js similarity index 100% rename from api/middleware/validateRegistration.js rename to api/server/middleware/validateRegistration.js diff --git a/api/server/routes/ask/anthropic.js b/api/server/routes/ask/anthropic.js index 6ede82c4d2..beca07b85a 100644 --- a/api/server/routes/ask/anthropic.js +++ b/api/server/routes/ask/anthropic.js @@ -1,72 +1,43 @@ const express = require('express'); const router = express.Router(); -const crypto = require('crypto'); -const { titleConvo, AnthropicClient } = require('../../../app'); -const requireJwtAuth = require('../../../middleware/requireJwtAuth'); -const { abortMessage } = require('../../../utils'); +const { getResponseSender } = require('../endpoints/schemas'); +const { initializeClient } = require('../endpoints/anthropic'); +const { + handleAbort, + createAbortController, + handleAbortError, + setHeaders, + requireJwtAuth, + validateEndpoint, + buildEndpointOption, +} = require('../../middleware'); const { saveMessage, getConvoTitle, saveConvo, getConvo } = require('../../../models'); -const { handleError, sendMessage, createOnProgress } = require('./handlers'); +const { sendMessage, createOnProgress } = require('../../utils'); -const abortControllers = new Map(); +router.post('/abort', requireJwtAuth, handleAbort()); -router.post('/abort', requireJwtAuth, async (req, res) => { - try { - return await abortMessage(req, res, abortControllers); - } catch (err) { - console.error(err); - } -}); +router.post( + '/', + requireJwtAuth, + validateEndpoint, + buildEndpointOption, + setHeaders, + async (req, res) => { + let { + text, + endpointOption, + conversationId, + parentMessageId = null, + overrideParentMessageId = null, + } = req.body; + console.log('ask log'); + console.dir({ text, conversationId, endpointOption }, { depth: null }); + let userMessage; + let userMessageId; + let responseMessageId; + let lastSavedTimestamp = 0; + let saveDelay = 100; -router.post('/', requireJwtAuth, async (req, res) => { - const { endpoint, text, parentMessageId, conversationId: oldConversationId } = req.body; - if (text.length === 0) { - return handleError(res, { text: 'Prompt empty or too short' }); - } - if (endpoint !== 'anthropic') { - return handleError(res, { text: 'Illegal request' }); - } - - const endpointOption = { - promptPrefix: req.body?.promptPrefix ?? null, - modelLabel: req.body?.modelLabel ?? null, - token: req.body?.token ?? null, - modelOptions: { - model: req.body?.model ?? 'claude-1', - temperature: req.body?.temperature ?? 1, - maxOutputTokens: req.body?.maxOutputTokens ?? 1024, - topP: req.body?.topP ?? 0.7, - topK: req.body?.topK ?? 5, - }, - }; - - const conversationId = oldConversationId || crypto.randomUUID(); - - return await ask({ - text, - endpointOption, - conversationId, - parentMessageId, - req, - res, - }); -}); - -const ask = async ({ text, endpointOption, parentMessageId = null, conversationId, req, res }) => { - res.writeHead(200, { - Connection: 'keep-alive', - 'Content-Type': 'text/event-stream', - 'Cache-Control': 'no-cache, no-transform', - 'Access-Control-Allow-Origin': '*', - 'X-Accel-Buffering': 'no', - }); - - let userMessage; - let userMessageId; - let responseMessageId; - let lastSavedTimestamp = 0; - const { overrideParentMessageId = null } = req.body; - - try { const getIds = (data) => { userMessage = data.userMessage; userMessageId = data.userMessage.messageId; @@ -79,116 +50,95 @@ const ask = async ({ text, endpointOption, parentMessageId = null, conversationI const { onProgress: progressCallback, getPartialText } = createOnProgress({ onProgress: ({ text: partialText }) => { const currentTimestamp = Date.now(); - if (currentTimestamp - lastSavedTimestamp > 500) { + + if (currentTimestamp - lastSavedTimestamp > saveDelay) { lastSavedTimestamp = currentTimestamp; saveMessage({ messageId: responseMessageId, - sender: 'Anthropic', + sender: getResponseSender(endpointOption), conversationId, - parentMessageId: overrideParentMessageId || userMessageId, + parentMessageId: overrideParentMessageId ?? userMessageId, text: partialText, unfinished: true, cancelled: false, error: false, }); } + + if (saveDelay < 500) { + saveDelay = 500; + } }, }); - - const abortController = new AbortController(); - abortController.abortAsk = async function () { - this.abort(); - - const responseMessage = { - messageId: responseMessageId, - sender: 'Anthropic', + try { + const getAbortData = () => ({ conversationId, - parentMessageId: overrideParentMessageId || userMessageId, + messageId: responseMessageId, + sender: getResponseSender(endpointOption), + parentMessageId: overrideParentMessageId ?? userMessageId, text: getPartialText(), - model: endpointOption.modelOptions.model, - unfinished: false, - cancelled: true, - error: false, - }; + userMessage, + }); - saveMessage(responseMessage); + const { abortController, onStart } = createAbortController( + res, + req, + endpointOption, + getAbortData, + ); - return { + const { client } = initializeClient(req, endpointOption); + + let response = await client.sendMessage(text, { + getIds, + debug: false, + user: req.user.id, + conversationId, + parentMessageId, + overrideParentMessageId, + ...endpointOption, + onProgress: progressCallback.call(null, { + res, + text, + parentMessageId: overrideParentMessageId ?? userMessageId, + }), + onStart, + abortController, + }); + + if (overrideParentMessageId) { + response.parentMessageId = overrideParentMessageId; + } + + await saveConvo(req.user.id, { + ...endpointOption, + ...endpointOption.modelOptions, + conversationId, + endpoint: 'anthropic', + }); + + await saveMessage(response); + sendMessage(res, { title: await getConvoTitle(req.user.id, conversationId), final: true, conversation: await getConvo(req.user.id, conversationId), requestMessage: userMessage, - responseMessage: responseMessage, - }; - }; + responseMessage: response, + }); + res.end(); - const onStart = (userMessage) => { - sendMessage(res, { message: userMessage, created: true }); - abortControllers.set(userMessage.conversationId, { abortController, ...endpointOption }); - }; - - const client = new AnthropicClient(endpointOption.token); - - let response = await client.sendMessage(text, { - getIds, - debug: false, - user: req.user.id, - conversationId, - parentMessageId, - overrideParentMessageId, - ...endpointOption, - onProgress: progressCallback.call(null, { - res, - text, - parentMessageId: overrideParentMessageId || userMessageId, - }), - onStart, - abortController, - }); - - if (overrideParentMessageId) { - response.parentMessageId = overrideParentMessageId; - } - - await saveConvo(req.user.id, { - ...endpointOption, - ...endpointOption.modelOptions, - conversationId, - endpoint: 'anthropic', - }); - - await saveMessage(response); - sendMessage(res, { - title: await getConvoTitle(req.user.id, conversationId), - final: true, - conversation: await getConvo(req.user.id, conversationId), - requestMessage: userMessage, - responseMessage: response, - }); - res.end(); - - if (parentMessageId == '00000000-0000-0000-0000-000000000000') { - const title = await titleConvo({ text, response }); - await saveConvo(req.user.id, { + // TODO: add anthropic titling + } catch (error) { + const partialText = getPartialText(); + handleAbortError(res, req, error, { + partialText, conversationId, - title, + sender: getResponseSender(endpointOption), + messageId: responseMessageId, + parentMessageId: userMessageId, }); } - } catch (error) { - console.error(error); - const errorMessage = { - messageId: responseMessageId, - sender: 'Anthropic', - conversationId, - parentMessageId, - unfinished: false, - cancelled: false, - error: true, - text: error.message, - }; - await saveMessage(errorMessage); - handleError(res, errorMessage); - } -}; + }, +); module.exports = router; diff --git a/api/server/routes/ask/askChatGPTBrowser.js b/api/server/routes/ask/askChatGPTBrowser.js index 576f581081..2b2472b89e 100644 --- a/api/server/routes/ask/askChatGPTBrowser.js +++ b/api/server/routes/ask/askChatGPTBrowser.js @@ -1,13 +1,12 @@ const express = require('express'); const crypto = require('crypto'); const router = express.Router(); -// const { getChatGPTBrowserModels } = require('../endpoints'); const { browserClient } = require('../../../app/'); const { saveMessage, getConvoTitle, saveConvo, getConvo } = require('../../../models'); -const { handleError, sendMessage, createOnProgress, handleText } = require('./handlers'); -const requireJwtAuth = require('../../../middleware/requireJwtAuth'); +const { handleError, sendMessage, createOnProgress, handleText } = require('../../utils'); +const { requireJwtAuth, setHeaders } = require('../../middleware'); -router.post('/', requireJwtAuth, async (req, res) => { +router.post('/', requireJwtAuth, setHeaders, async (req, res) => { const { endpoint, text, @@ -86,15 +85,6 @@ const ask = async ({ }) => { let { text, parentMessageId: userParentMessageId, messageId: userMessageId } = userMessage; const userId = req.user.id; - - res.writeHead(200, { - Connection: 'keep-alive', - 'Content-Type': 'text/event-stream', - 'Cache-Control': 'no-cache, no-transform', - 'Access-Control-Allow-Origin': '*', - 'X-Accel-Buffering': 'no', - }); - let responseMessageId = crypto.randomUUID(); let getPartialMessage = null; try { diff --git a/api/server/routes/ask/bingAI.js b/api/server/routes/ask/bingAI.js index ced293105a..f8e834fb97 100644 --- a/api/server/routes/ask/bingAI.js +++ b/api/server/routes/ask/bingAI.js @@ -3,10 +3,10 @@ const crypto = require('crypto'); const router = express.Router(); const { titleConvoBing, askBing } = require('../../../app'); const { saveMessage, getConvoTitle, saveConvo, getConvo } = require('../../../models'); -const { handleError, sendMessage, createOnProgress, handleText } = require('./handlers'); -const requireJwtAuth = require('../../../middleware/requireJwtAuth'); +const { handleError, sendMessage, createOnProgress, handleText } = require('../../utils'); +const { requireJwtAuth, setHeaders } = require('../../middleware'); -router.post('/', requireJwtAuth, async (req, res) => { +router.post('/', requireJwtAuth, setHeaders, async (req, res) => { const { endpoint, text, @@ -103,14 +103,6 @@ const ask = async ({ let responseMessageId = crypto.randomUUID(); - res.writeHead(200, { - Connection: 'keep-alive', - 'Content-Type': 'text/event-stream', - 'Cache-Control': 'no-cache, no-transform', - 'Access-Control-Allow-Origin': '*', - 'X-Accel-Buffering': 'no', - }); - if (preSendRequest) { sendMessage(res, { message: userMessage, created: true }); } diff --git a/api/server/routes/ask/google.js b/api/server/routes/ask/google.js index f3d25cbcd4..17775e2f38 100644 --- a/api/server/routes/ask/google.js +++ b/api/server/routes/ask/google.js @@ -2,12 +2,11 @@ const express = require('express'); const router = express.Router(); const crypto = require('crypto'); const { titleConvo, GoogleClient } = require('../../../app'); -// const GoogleClient = require('../../../app/google/GoogleClient'); const { saveMessage, getConvoTitle, saveConvo, getConvo } = require('../../../models'); -const { handleError, sendMessage, createOnProgress } = require('./handlers'); -const requireJwtAuth = require('../../../middleware/requireJwtAuth'); +const { handleError, sendMessage, createOnProgress } = require('../../utils'); +const { requireJwtAuth, setHeaders } = require('../../middleware'); -router.post('/', requireJwtAuth, async (req, res) => { +router.post('/', requireJwtAuth, setHeaders, async (req, res) => { const { endpoint, text, parentMessageId, conversationId: oldConversationId } = req.body; if (text.length === 0) { return handleError(res, { text: 'Prompt empty or too short' }); @@ -50,13 +49,6 @@ router.post('/', requireJwtAuth, async (req, res) => { }); const ask = async ({ text, endpointOption, parentMessageId = null, conversationId, req, res }) => { - res.writeHead(200, { - Connection: 'keep-alive', - 'Content-Type': 'text/event-stream', - 'Cache-Control': 'no-cache, no-transform', - 'Access-Control-Allow-Origin': '*', - 'X-Accel-Buffering': 'no', - }); let userMessage; let userMessageId; let responseMessageId; diff --git a/api/server/routes/ask/gptPlugins.js b/api/server/routes/ask/gptPlugins.js index 7a336fe97d..2ce3d21d8e 100644 --- a/api/server/routes/ask/gptPlugins.js +++ b/api/server/routes/ask/gptPlugins.js @@ -1,112 +1,56 @@ const express = require('express'); const router = express.Router(); -const { titleConvo, validateTools, PluginsClient } = require('../../../app'); -const { abortMessage, getAzureCredentials } = require('../../../utils'); -const { saveMessage, getConvoTitle, saveConvo, getConvo } = require('../../../models'); +const { getResponseSender } = require('../endpoints/schemas'); +const { validateTools } = require('../../../app'); +const { addTitle } = require('../endpoints/openAI'); +const { initializeClient } = require('../endpoints/gptPlugins'); +const { saveMessage, getConvoTitle, getConvo } = require('../../../models'); +const { sendMessage, createOnProgress, formatSteps, formatAction } = require('../../utils'); const { - handleError, - sendMessage, - createOnProgress, - formatSteps, - formatAction, -} = require('./handlers'); -const requireJwtAuth = require('../../../middleware/requireJwtAuth'); + handleAbort, + createAbortController, + handleAbortError, + setHeaders, + requireJwtAuth, + validateEndpoint, + buildEndpointOption, +} = require('../../middleware'); -const abortControllers = new Map(); +router.post('/abort', requireJwtAuth, handleAbort()); -router.post('/abort', requireJwtAuth, async (req, res) => { - try { - return await abortMessage(req, res, abortControllers); - } catch (err) { - console.error(err); - } -}); +router.post( + '/', + requireJwtAuth, + validateEndpoint, + buildEndpointOption, + setHeaders, + async (req, res) => { + let { + text, + endpointOption, + conversationId, + parentMessageId = null, + overrideParentMessageId = null, + } = req.body; + console.log('ask log'); + console.dir({ text, conversationId, endpointOption }, { depth: null }); + let metadata; + let userMessage; + let userMessageId; + let responseMessageId; + let lastSavedTimestamp = 0; + let saveDelay = 100; + const newConvo = !conversationId; + const user = req.user.id; -router.post('/', requireJwtAuth, async (req, res) => { - const { endpoint, text, parentMessageId, conversationId } = req.body; - if (text.length === 0) { - return handleError(res, { text: 'Prompt empty or too short' }); - } - if (endpoint !== 'gptPlugins') { - return handleError(res, { text: 'Illegal request' }); - } + const plugin = { + loading: true, + inputs: [], + latest: null, + outputs: null, + }; - const agentOptions = req.body?.agentOptions ?? { - agent: 'functions', - skipCompletion: true, - model: 'gpt-3.5-turbo', - temperature: 0, - // top_p: 1, - // presence_penalty: 0, - // frequency_penalty: 0 - }; - - const tools = req.body?.tools.map((tool) => tool.pluginKey) ?? []; - // build endpoint option - const endpointOption = { - chatGptLabel: tools.length === 0 ? req.body?.chatGptLabel ?? null : null, - promptPrefix: tools.length === 0 ? req.body?.promptPrefix ?? null : null, - tools, - modelOptions: { - model: req.body?.model ?? 'gpt-4', - temperature: req.body?.temperature ?? 0, - top_p: req.body?.top_p ?? 1, - presence_penalty: req.body?.presence_penalty ?? 0, - frequency_penalty: req.body?.frequency_penalty ?? 0, - }, - agentOptions: { - ...agentOptions, - // agent: 'functions' - }, - }; - - console.log('ask log'); - console.dir({ text, conversationId, endpointOption }, { depth: null }); - - // eslint-disable-next-line no-use-before-define - return await ask({ - text, - endpoint, - endpointOption, - conversationId, - parentMessageId, - req, - res, - }); -}); - -const ask = async ({ - text, - endpoint, - endpointOption, - parentMessageId = null, - conversationId, - req, - res, -}) => { - res.writeHead(200, { - Connection: 'keep-alive', - 'Content-Type': 'text/event-stream', - 'Cache-Control': 'no-cache, no-transform', - 'Access-Control-Allow-Origin': '*', - 'X-Accel-Buffering': 'no', - }); - let userMessage; - let userMessageId; - let responseMessageId; - let lastSavedTimestamp = 0; - const newConvo = !conversationId; - const { overrideParentMessageId = null } = req.body; - const user = req.user.id; - - const plugin = { - loading: true, - inputs: [], - latest: null, - outputs: null, - }; - - try { + const addMetadata = (data) => (metadata = data); const getIds = (data) => { userMessage = data.userMessage; userMessageId = userMessage.messageId; @@ -128,11 +72,11 @@ const ask = async ({ plugin.loading = false; } - if (currentTimestamp - lastSavedTimestamp > 500) { + if (currentTimestamp - lastSavedTimestamp > saveDelay) { lastSavedTimestamp = currentTimestamp; saveMessage({ messageId: responseMessageId, - sender: 'ChatGPT', + sender: getResponseSender(endpointOption), conversationId, parentMessageId: overrideParentMessageId || userMessageId, text: partialText, @@ -142,63 +86,13 @@ const ask = async ({ error: false, }); } + + if (saveDelay < 500) { + saveDelay = 500; + } }, }); - const abortController = new AbortController(); - abortController.abortAsk = async function () { - this.abort(); - - const responseMessage = { - messageId: responseMessageId, - sender: endpointOption?.chatGptLabel || 'ChatGPT', - conversationId, - parentMessageId: overrideParentMessageId || userMessageId, - text: getPartialText(), - plugin: { ...plugin, loading: false }, - model: endpointOption.modelOptions.model, - unfinished: false, - cancelled: true, - error: false, - }; - - saveMessage(responseMessage); - - return { - title: await getConvoTitle(req.user.id, conversationId), - final: true, - conversation: await getConvo(req.user.id, conversationId), - requestMessage: userMessage, - responseMessage: responseMessage, - }; - }; - - const onStart = (userMessage) => { - sendMessage(res, { message: userMessage, created: true }); - abortControllers.set(userMessage.conversationId, { abortController, ...endpointOption }); - }; - - endpointOption.tools = await validateTools(user, endpointOption.tools); - const clientOptions = { - debug: true, - endpoint, - reverseProxyUrl: process.env.OPENAI_REVERSE_PROXY || null, - proxy: process.env.PROXY || null, - ...endpointOption, - }; - - let openAIApiKey = req.body?.token ?? process.env.OPENAI_API_KEY; - if (process.env.PLUGINS_USE_AZURE) { - clientOptions.azure = getAzureCredentials(); - openAIApiKey = clientOptions.azure.azureOpenAIApiKey; - } - - if (openAIApiKey && openAIApiKey.includes('azure') && !clientOptions.azure) { - clientOptions.azure = JSON.parse(req.body?.token) ?? getAzureCredentials(); - openAIApiKey = clientOptions.azure.azureOpenAIApiKey; - } - const chatAgent = new PluginsClient(openAIApiKey, clientOptions); - const onAgentAction = (action, start = false) => { const formattedAction = formatAction(action); plugin.inputs.push(formattedAction); @@ -219,70 +113,86 @@ const ask = async ({ // console.log('CHAIN END', plugin.outputs); }; - let response = await chatAgent.sendMessage(text, { - getIds, - user, - parentMessageId, + const getAbortData = () => ({ + sender: getResponseSender(endpointOption), conversationId, - overrideParentMessageId, - onAgentAction, - onChainEnd, - onStart, - ...endpointOption, - onProgress: progressCallback.call(null, { - res, - text, - plugin, - parentMessageId: overrideParentMessageId || userMessageId, - }), - abortController, + messageId: responseMessageId, + parentMessageId: overrideParentMessageId ?? userMessageId, + text: getPartialText(), + plugin: { ...plugin, loading: false }, + userMessage, }); + const { abortController, onStart } = createAbortController( + res, + req, + endpointOption, + getAbortData, + ); - if (overrideParentMessageId) { - response.parentMessageId = overrideParentMessageId; - } + try { + endpointOption.tools = await validateTools(user, endpointOption.tools); + const { client, azure, openAIApiKey } = initializeClient(req, endpointOption); - console.log('CLIENT RESPONSE'); - console.dir(response, { depth: null }); - response.plugin = { ...plugin, loading: false }; - await saveMessage(response); + let response = await client.sendMessage(text, { + user, + conversationId, + parentMessageId, + overrideParentMessageId, + getIds, + onAgentAction, + onChainEnd, + onStart, + addMetadata, + ...endpointOption, + onProgress: progressCallback.call(null, { + res, + text, + plugin, + parentMessageId: overrideParentMessageId || userMessageId, + }), + abortController, + }); - sendMessage(res, { - title: await getConvoTitle(req.user.id, conversationId), - final: true, - conversation: await getConvo(req.user.id, conversationId), - requestMessage: userMessage, - responseMessage: response, - }); - res.end(); + if (overrideParentMessageId) { + response.parentMessageId = overrideParentMessageId; + } - if (parentMessageId == '00000000-0000-0000-0000-000000000000' && newConvo) { - const title = await titleConvo({ + if (metadata) { + response = { ...response, ...metadata }; + } + + console.log('CLIENT RESPONSE'); + console.dir(response, { depth: null }); + response.plugin = { ...plugin, loading: false }; + await saveMessage(response); + + sendMessage(res, { + title: await getConvoTitle(req.user.id, conversationId), + final: true, + conversation: await getConvo(req.user.id, conversationId), + requestMessage: userMessage, + responseMessage: response, + }); + res.end(); + addTitle(req, { text, + newConvo, response, openAIApiKey, - azure: !!clientOptions.azure, + parentMessageId, + azure: !!azure, }); - await saveConvo(req.user.id, { - conversationId: conversationId, - title, + } catch (error) { + const partialText = getPartialText(); + handleAbortError(res, req, error, { + partialText, + conversationId, + sender: getResponseSender(endpointOption), + messageId: responseMessageId, + parentMessageId: userMessageId, }); } - } catch (error) { - console.error(error); - const errorMessage = { - messageId: responseMessageId, - sender: 'ChatGPT', - conversationId, - parentMessageId: userMessageId, - unfinished: false, - cancelled: false, - error: true, - text: error.message, - }; - await saveMessage(errorMessage); - handleError(res, errorMessage); - } -}; + }, +); module.exports = router; diff --git a/api/server/routes/ask/index.js b/api/server/routes/ask/index.js index d088d97b17..77da50f68a 100644 --- a/api/server/routes/ask/index.js +++ b/api/server/routes/ask/index.js @@ -1,7 +1,5 @@ const express = require('express'); const router = express.Router(); -// const askAzureOpenAI = require('./askAzureOpenAI';) -// const askOpenAI = require('./askOpenAI'); const openAI = require('./openAI'); const google = require('./google'); const bingAI = require('./bingAI'); @@ -9,7 +7,6 @@ const gptPlugins = require('./gptPlugins'); const askChatGPTBrowser = require('./askChatGPTBrowser'); const anthropic = require('./anthropic'); -// router.use('/azureOpenAI', askAzureOpenAI); router.use(['/azureOpenAI', '/openAI'], openAI); router.use('/google', google); router.use('/bingAI', bingAI); diff --git a/api/server/routes/ask/openAI.js b/api/server/routes/ask/openAI.js index db8e3a3cbd..be236956af 100644 --- a/api/server/routes/ask/openAI.js +++ b/api/server/routes/ask/openAI.js @@ -1,231 +1,160 @@ const express = require('express'); const router = express.Router(); -const { titleConvo, OpenAIClient } = require('../../../app'); -const { getAzureCredentials, abortMessage } = require('../../../utils'); -const { saveMessage, getConvoTitle, saveConvo, getConvo } = require('../../../models'); -const { handleError, sendMessage, createOnProgress } = require('./handlers'); -const requireJwtAuth = require('../../../middleware/requireJwtAuth'); +const { getResponseSender } = require('../endpoints/schemas'); +const { sendMessage, createOnProgress } = require('../../utils'); +const { addTitle, initializeClient } = require('../endpoints/openAI'); +const { saveMessage, getConvoTitle, getConvo } = require('../../../models'); +const { + handleAbort, + createAbortController, + handleAbortError, + setHeaders, + requireJwtAuth, + validateEndpoint, + buildEndpointOption, +} = require('../../middleware'); -const abortControllers = new Map(); +router.post('/abort', requireJwtAuth, handleAbort()); -router.post('/abort', requireJwtAuth, async (req, res) => { - try { - return await abortMessage(req, res, abortControllers); - } catch (err) { - console.error(err); - } -}); +router.post( + '/', + requireJwtAuth, + validateEndpoint, + buildEndpointOption, + setHeaders, + async (req, res) => { + let { + text, + endpointOption, + conversationId, + parentMessageId = null, + overrideParentMessageId = null, + } = req.body; + console.log('ask log'); + console.dir({ text, conversationId, endpointOption }, { depth: null }); + let metadata; + let userMessage; + let userMessageId; + let responseMessageId; + let lastSavedTimestamp = 0; + let saveDelay = 100; + const newConvo = !conversationId; + const user = req.user.id; -router.post('/', requireJwtAuth, async (req, res) => { - const { endpoint, text, parentMessageId, conversationId } = req.body; - if (text.length === 0) { - return handleError(res, { text: 'Prompt empty or too short' }); - } - const isOpenAI = endpoint === 'openAI' || endpoint === 'azureOpenAI'; - if (!isOpenAI) { - return handleError(res, { text: 'Illegal request' }); - } + const addMetadata = (data) => (metadata = data); - // build endpoint option - const endpointOption = { - chatGptLabel: req.body?.chatGptLabel ?? null, - promptPrefix: req.body?.promptPrefix ?? null, - modelOptions: { - model: req.body?.model ?? 'gpt-3.5-turbo', - temperature: req.body?.temperature ?? 1, - top_p: req.body?.top_p ?? 1, - presence_penalty: req.body?.presence_penalty ?? 0, - frequency_penalty: req.body?.frequency_penalty ?? 0, - }, - }; - - console.log('ask log'); - console.dir({ text, conversationId, endpointOption }, { depth: null }); - - // eslint-disable-next-line no-use-before-define - return await ask({ - text, - endpointOption, - conversationId, - parentMessageId, - endpoint, - req, - res, - }); -}); - -const ask = async ({ - text, - endpointOption, - parentMessageId = null, - endpoint, - conversationId, - req, - res, -}) => { - res.writeHead(200, { - Connection: 'keep-alive', - 'Content-Type': 'text/event-stream', - 'Cache-Control': 'no-cache, no-transform', - 'Access-Control-Allow-Origin': '*', - 'X-Accel-Buffering': 'no', - }); - let userMessage; - let userMessageId; - let responseMessageId; - let lastSavedTimestamp = 0; - const newConvo = !conversationId; - const { overrideParentMessageId = null } = req.body; - const user = req.user.id; - - const getIds = (data) => { - userMessage = data.userMessage; - userMessageId = userMessage.messageId; - responseMessageId = data.responseMessageId; - if (!conversationId) { - conversationId = data.conversationId; - } - }; - - const { onProgress: progressCallback, getPartialText } = createOnProgress({ - onProgress: ({ text: partialText }) => { - const currentTimestamp = Date.now(); - - if (currentTimestamp - lastSavedTimestamp > 500) { - lastSavedTimestamp = currentTimestamp; - saveMessage({ - messageId: responseMessageId, - sender: 'ChatGPT', - conversationId, - parentMessageId: overrideParentMessageId || userMessageId, - text: partialText, - model: endpointOption.modelOptions.model, - unfinished: true, - cancelled: false, - error: false, - }); + const getIds = (data) => { + userMessage = data.userMessage; + userMessageId = userMessage.messageId; + responseMessageId = data.responseMessageId; + if (!conversationId) { + conversationId = data.conversationId; } - }, - }); + }; - const abortController = new AbortController(); - abortController.abortAsk = async function () { - this.abort(); + const { onProgress: progressCallback, getPartialText } = createOnProgress({ + onProgress: ({ text: partialText }) => { + const currentTimestamp = Date.now(); - const responseMessage = { + if (currentTimestamp - lastSavedTimestamp > saveDelay) { + lastSavedTimestamp = currentTimestamp; + saveMessage({ + messageId: responseMessageId, + sender: getResponseSender(endpointOption), + conversationId, + parentMessageId: overrideParentMessageId ?? userMessageId, + text: partialText, + model: endpointOption.modelOptions.model, + unfinished: true, + cancelled: false, + error: false, + }); + } + + if (saveDelay < 500) { + saveDelay = 500; + } + }, + }); + + const getAbortData = () => ({ + sender: getResponseSender(endpointOption), + conversationId, messageId: responseMessageId, - sender: endpointOption?.chatGptLabel || 'ChatGPT', - conversationId, - parentMessageId: overrideParentMessageId || userMessageId, + parentMessageId: overrideParentMessageId ?? userMessageId, text: getPartialText(), - model: endpointOption.modelOptions.model, - unfinished: false, - cancelled: true, - error: false, - }; - - saveMessage(responseMessage); - - return { - title: await getConvoTitle(req.user.id, conversationId), - final: true, - conversation: await getConvo(req.user.id, conversationId), - requestMessage: userMessage, - responseMessage: responseMessage, - }; - }; - - const onStart = (userMessage) => { - sendMessage(res, { message: userMessage, created: true }); - abortControllers.set(userMessage.conversationId, { abortController, ...endpointOption }); - }; - - try { - const clientOptions = { - // debug: true, - // contextStrategy: 'refine', - reverseProxyUrl: process.env.OPENAI_REVERSE_PROXY || null, - proxy: process.env.PROXY || null, - endpoint, - ...endpointOption, - }; - - let openAIApiKey = req.body?.token ?? process.env.OPENAI_API_KEY; - - if (process.env.AZURE_API_KEY && endpoint === 'azureOpenAI') { - clientOptions.azure = JSON.parse(req.body?.token) ?? getAzureCredentials(); - openAIApiKey = clientOptions.azure.azureOpenAIApiKey; - } - - const client = new OpenAIClient(openAIApiKey, clientOptions); - - let response = await client.sendMessage(text, { - user, - parentMessageId, - conversationId, - overrideParentMessageId, - getIds, - onStart, - onProgress: progressCallback.call(null, { - res, - text, - parentMessageId: overrideParentMessageId || userMessageId, - }), - abortController, + userMessage, }); - if (overrideParentMessageId) { - response.parentMessageId = overrideParentMessageId; - } - - console.log( - 'promptTokens, completionTokens:', - response.promptTokens, - response.completionTokens, + const { abortController, onStart } = createAbortController( + res, + req, + endpointOption, + getAbortData, ); - await saveMessage(response); - sendMessage(res, { - title: await getConvoTitle(req.user.id, conversationId), - final: true, - conversation: await getConvo(req.user.id, conversationId), - requestMessage: userMessage, - responseMessage: response, - }); - res.end(); + try { + const { client, openAIApiKey } = initializeClient(req, endpointOption); - if (parentMessageId == '00000000-0000-0000-0000-000000000000' && newConvo) { - const title = await titleConvo({ + let response = await client.sendMessage(text, { + user, + parentMessageId, + conversationId, + overrideParentMessageId, + getIds, + onStart, + addMetadata, + abortController, + onProgress: progressCallback.call(null, { + res, + text, + parentMessageId: overrideParentMessageId || userMessageId, + }), + }); + + if (overrideParentMessageId) { + response.parentMessageId = overrideParentMessageId; + } + + if (metadata) { + response = { ...response, ...metadata }; + } + + console.log( + 'promptTokens, completionTokens:', + response.promptTokens, + response.completionTokens, + ); + await saveMessage(response); + + sendMessage(res, { + title: await getConvoTitle(req.user.id, conversationId), + final: true, + conversation: await getConvo(req.user.id, conversationId), + requestMessage: userMessage, + responseMessage: response, + }); + res.end(); + + addTitle(req, { text, + newConvo, response, openAIApiKey, - azure: endpoint === 'azureOpenAI', + parentMessageId, + azure: endpointOption.endpoint === 'azureOpenAI', }); - await saveConvo(req.user.id, { + } catch (error) { + const partialText = getPartialText(); + handleAbortError(res, req, error, { + partialText, conversationId, - title, - }); - } - } catch (error) { - console.error(error); - const partialText = getPartialText(); - if (partialText?.length > 2) { - return await abortMessage(req, res, abortControllers); - } else { - const errorMessage = { + sender: getResponseSender(endpointOption), messageId: responseMessageId, - sender: 'ChatGPT', - conversationId, parentMessageId: userMessageId, - unfinished: false, - cancelled: false, - error: true, - text: error.message, - }; - await saveMessage(errorMessage); - handleError(res, errorMessage); + }); } - } -}; + }, +); module.exports = router; diff --git a/api/server/routes/auth.js b/api/server/routes/auth.js index 1f0c660c67..ff2e7cab0e 100644 --- a/api/server/routes/auth.js +++ b/api/server/routes/auth.js @@ -7,9 +7,7 @@ const { } = require('../controllers/AuthController'); const { loginController } = require('../controllers/auth/LoginController'); const { logoutController } = require('../controllers/auth/LogoutController'); -const requireJwtAuth = require('../../middleware/requireJwtAuth'); -const requireLocalAuth = require('../../middleware/requireLocalAuth'); -const validateRegistration = require('../../middleware/validateRegistration'); +const { requireJwtAuth, requireLocalAuth, validateRegistration } = require('../middleware'); const router = express.Router(); diff --git a/api/server/routes/convos.js b/api/server/routes/convos.js index 28e29bea70..66b3ffc0ac 100644 --- a/api/server/routes/convos.js +++ b/api/server/routes/convos.js @@ -2,7 +2,7 @@ const express = require('express'); const router = express.Router(); const { getConvo, saveConvo } = require('../../models'); const { getConvosByPage, deleteConvos } = require('../../models/Conversation'); -const requireJwtAuth = require('../../middleware/requireJwtAuth'); +const requireJwtAuth = require('../middleware/requireJwtAuth'); router.get('/', requireJwtAuth, async (req, res) => { const pageNumber = req.query.pageNumber || 1; diff --git a/api/server/routes/edit/anthropic.js b/api/server/routes/edit/anthropic.js new file mode 100644 index 0000000000..7e141cbdfa --- /dev/null +++ b/api/server/routes/edit/anthropic.js @@ -0,0 +1,139 @@ +const express = require('express'); +const router = express.Router(); +const { getResponseSender } = require('../endpoints/schemas'); +const { initializeClient } = require('../endpoints/anthropic'); +const { + handleAbort, + createAbortController, + handleAbortError, + setHeaders, + requireJwtAuth, + validateEndpoint, + buildEndpointOption, +} = require('../../middleware'); +const { saveMessage, getConvoTitle, getConvo } = require('../../../models'); +const { sendMessage, createOnProgress } = require('../../utils'); + +router.post('/abort', requireJwtAuth, handleAbort()); + +router.post( + '/', + requireJwtAuth, + validateEndpoint, + buildEndpointOption, + setHeaders, + async (req, res) => { + let { + text, + generation, + endpointOption, + conversationId, + responseMessageId, + parentMessageId = null, + overrideParentMessageId = null, + } = req.body; + console.log('edit log'); + console.dir({ text, conversationId, endpointOption }, { depth: null }); + let metadata; + let userMessage; + let lastSavedTimestamp = 0; + let saveDelay = 100; + const userMessageId = parentMessageId; + + const addMetadata = (data) => (metadata = data); + const getIds = (data) => (userMessage = data.userMessage); + + const { onProgress: progressCallback, getPartialText } = createOnProgress({ + generation, + onProgress: ({ text: partialText }) => { + const currentTimestamp = Date.now(); + if (currentTimestamp - lastSavedTimestamp > saveDelay) { + lastSavedTimestamp = currentTimestamp; + saveMessage({ + messageId: responseMessageId, + sender: getResponseSender(endpointOption), + conversationId, + parentMessageId: overrideParentMessageId ?? userMessageId, + text: partialText, + unfinished: true, + cancelled: false, + error: false, + }); + } + + if (saveDelay < 500) { + saveDelay = 500; + } + }, + }); + try { + const getAbortData = () => ({ + conversationId, + messageId: responseMessageId, + sender: getResponseSender(endpointOption), + parentMessageId: overrideParentMessageId ?? userMessageId, + text: getPartialText(), + userMessage, + }); + + const { abortController, onStart } = createAbortController( + res, + req, + endpointOption, + getAbortData, + ); + + const { client } = initializeClient(req, endpointOption); + + let response = await client.sendMessage(text, { + user: req.user.id, + isEdited: true, + conversationId, + parentMessageId, + responseMessageId, + overrideParentMessageId, + ...endpointOption, + onProgress: progressCallback.call(null, { + res, + text, + parentMessageId: overrideParentMessageId ?? userMessageId, + }), + getIds, + onStart, + addMetadata, + abortController, + }); + + if (metadata) { + response = { ...response, ...metadata }; + } + + if (overrideParentMessageId) { + response.parentMessageId = overrideParentMessageId; + } + + await saveMessage(response); + sendMessage(res, { + title: await getConvoTitle(req.user.id, conversationId), + final: true, + conversation: await getConvo(req.user.id, conversationId), + requestMessage: userMessage, + responseMessage: response, + }); + res.end(); + + // TODO: add anthropic titling + } catch (error) { + const partialText = getPartialText(); + handleAbortError(res, req, error, { + partialText, + conversationId, + sender: getResponseSender(endpointOption), + messageId: responseMessageId, + parentMessageId: userMessageId, + }); + } + }, +); + +module.exports = router; diff --git a/api/server/routes/edit/gptPlugins.js b/api/server/routes/edit/gptPlugins.js new file mode 100644 index 0000000000..a0c81d46ca --- /dev/null +++ b/api/server/routes/edit/gptPlugins.js @@ -0,0 +1,185 @@ +const express = require('express'); +const router = express.Router(); +const { getResponseSender } = require('../endpoints/schemas'); +const { validateTools } = require('../../../app'); +const { initializeClient } = require('../endpoints/gptPlugins'); +const { saveMessage, getConvoTitle, getConvo } = require('../../../models'); +const { sendMessage, createOnProgress, formatSteps, formatAction } = require('../../utils'); +const { + handleAbort, + createAbortController, + handleAbortError, + setHeaders, + requireJwtAuth, + validateEndpoint, + buildEndpointOption, +} = require('../../middleware'); + +router.post('/abort', requireJwtAuth, handleAbort()); + +router.post( + '/', + requireJwtAuth, + validateEndpoint, + buildEndpointOption, + setHeaders, + async (req, res) => { + let { + text, + generation, + endpointOption, + conversationId, + responseMessageId, + parentMessageId = null, + overrideParentMessageId = null, + } = req.body; + console.log('edit log'); + console.dir({ text, conversationId, endpointOption }, { depth: null }); + let metadata; + let userMessage; + let lastSavedTimestamp = 0; + let saveDelay = 100; + const userMessageId = parentMessageId; + const user = req.user.id; + + const plugin = { + loading: true, + inputs: [], + latest: null, + outputs: null, + }; + + const addMetadata = (data) => (metadata = data); + const getIds = (data) => (userMessage = data.userMessage); + + const { + onProgress: progressCallback, + sendIntermediateMessage, + getPartialText, + } = createOnProgress({ + generation, + onProgress: ({ text: partialText }) => { + const currentTimestamp = Date.now(); + + if (plugin.loading === true) { + plugin.loading = false; + } + + if (currentTimestamp - lastSavedTimestamp > saveDelay) { + lastSavedTimestamp = currentTimestamp; + saveMessage({ + messageId: responseMessageId, + sender: getResponseSender(endpointOption), + conversationId, + parentMessageId: overrideParentMessageId || userMessageId, + text: partialText, + model: endpointOption.modelOptions.model, + unfinished: true, + cancelled: false, + error: false, + }); + } + + if (saveDelay < 500) { + saveDelay = 500; + } + }, + }); + + const onAgentAction = (action, start = false) => { + const formattedAction = formatAction(action); + plugin.inputs.push(formattedAction); + plugin.latest = formattedAction.plugin; + if (!start) { + saveMessage(userMessage); + } + sendIntermediateMessage(res, { plugin }); + // console.log('PLUGIN ACTION', formattedAction); + }; + + const onChainEnd = (data) => { + let { intermediateSteps: steps } = data; + plugin.outputs = steps && steps[0].action ? formatSteps(steps) : 'An error occurred.'; + plugin.loading = false; + saveMessage(userMessage); + sendIntermediateMessage(res, { plugin }); + // console.log('CHAIN END', plugin.outputs); + }; + + const getAbortData = () => ({ + sender: getResponseSender(endpointOption), + conversationId, + messageId: responseMessageId, + parentMessageId: overrideParentMessageId ?? userMessageId, + text: getPartialText(), + plugin: { ...plugin, loading: false }, + userMessage, + }); + const { abortController, onStart } = createAbortController( + res, + req, + endpointOption, + getAbortData, + ); + + try { + endpointOption.tools = await validateTools(user, endpointOption.tools); + const { client } = initializeClient(req, endpointOption); + + let response = await client.sendMessage(text, { + user, + isEdited: true, + conversationId, + parentMessageId, + responseMessageId, + overrideParentMessageId, + getIds, + onAgentAction, + onChainEnd, + onStart, + addMetadata, + ...endpointOption, + onProgress: progressCallback.call(null, { + res, + text, + plugin, + parentMessageId: overrideParentMessageId || userMessageId, + }), + abortController, + }); + + if (overrideParentMessageId) { + response.parentMessageId = overrideParentMessageId; + } + + if (metadata) { + response = { ...response, ...metadata }; + } + + console.log('CLIENT RESPONSE'); + console.dir(response, { depth: null }); + response.plugin = { ...plugin, loading: false }; + await saveMessage(response); + + sendMessage(res, { + title: await getConvoTitle(req.user.id, conversationId), + final: true, + conversation: await getConvo(req.user.id, conversationId), + requestMessage: userMessage, + responseMessage: response, + }); + res.end(); + } catch (error) { + const partialText = getPartialText(); + handleAbortError(res, req, error, { + partialText, + conversationId, + sender: getResponseSender(endpointOption), + messageId: responseMessageId, + parentMessageId: userMessageId, + }); + } + }, +); + +module.exports = router; diff --git a/api/server/routes/edit/index.js b/api/server/routes/edit/index.js new file mode 100644 index 0000000000..7eda18b8a1 --- /dev/null +++ b/api/server/routes/edit/index.js @@ -0,0 +1,13 @@ +const express = require('express'); +const router = express.Router(); +const openAI = require('./openAI'); +const gptPlugins = require('./gptPlugins'); +const anthropic = require('./anthropic'); +// const google = require('./google'); + +router.use(['/azureOpenAI', '/openAI'], openAI); +router.use('/gptPlugins', gptPlugins); +router.use('/anthropic', anthropic); +// router.use('/google', google); + +module.exports = router; diff --git a/api/server/routes/edit/openAI.js b/api/server/routes/edit/openAI.js new file mode 100644 index 0000000000..1ad0bcbebf --- /dev/null +++ b/api/server/routes/edit/openAI.js @@ -0,0 +1,141 @@ +const express = require('express'); +const router = express.Router(); +const { getResponseSender } = require('../endpoints/schemas'); +const { initializeClient } = require('../endpoints/openAI'); +const { saveMessage, getConvoTitle, getConvo } = require('../../../models'); +const { sendMessage, createOnProgress } = require('../../utils'); +const { + handleAbort, + createAbortController, + handleAbortError, + setHeaders, + requireJwtAuth, + validateEndpoint, + buildEndpointOption, +} = require('../../middleware'); + +router.post('/abort', requireJwtAuth, handleAbort()); + +router.post( + '/', + requireJwtAuth, + validateEndpoint, + buildEndpointOption, + setHeaders, + async (req, res) => { + let { + text, + generation, + endpointOption, + conversationId, + responseMessageId, + parentMessageId = null, + overrideParentMessageId = null, + } = req.body; + console.log('edit log'); + console.dir({ text, conversationId, endpointOption }, { depth: null }); + let metadata; + let userMessage; + let lastSavedTimestamp = 0; + let saveDelay = 100; + const userMessageId = parentMessageId; + + const addMetadata = (data) => (metadata = data); + const getIds = (data) => (userMessage = data.userMessage); + + const { onProgress: progressCallback, getPartialText } = createOnProgress({ + generation, + onProgress: ({ text: partialText }) => { + const currentTimestamp = Date.now(); + + if (currentTimestamp - lastSavedTimestamp > saveDelay) { + lastSavedTimestamp = currentTimestamp; + saveMessage({ + messageId: responseMessageId, + sender: getResponseSender(endpointOption), + conversationId, + parentMessageId: overrideParentMessageId || userMessageId, + text: partialText, + model: endpointOption.modelOptions.model, + unfinished: true, + cancelled: false, + error: false, + }); + } + + if (saveDelay < 500) { + saveDelay = 500; + } + }, + }); + + const getAbortData = () => ({ + sender: getResponseSender(endpointOption), + conversationId, + messageId: responseMessageId, + parentMessageId: overrideParentMessageId ?? userMessageId, + text: getPartialText(), + userMessage, + }); + + const { abortController, onStart } = createAbortController( + res, + req, + endpointOption, + getAbortData, + ); + + try { + const { client } = initializeClient(req, endpointOption); + + let response = await client.sendMessage(text, { + user: req.user.id, + isEdited: true, + conversationId, + parentMessageId, + responseMessageId, + overrideParentMessageId, + getIds, + onStart, + addMetadata, + abortController, + onProgress: progressCallback.call(null, { + res, + text, + parentMessageId: overrideParentMessageId || userMessageId, + }), + }); + + if (metadata) { + response = { ...response, ...metadata }; + } + + console.log( + 'promptTokens, completionTokens:', + response.promptTokens, + response.completionTokens, + ); + await saveMessage(response); + + sendMessage(res, { + title: await getConvoTitle(req.user.id, conversationId), + final: true, + conversation: await getConvo(req.user.id, conversationId), + requestMessage: userMessage, + responseMessage: response, + }); + res.end(); + } catch (error) { + const partialText = getPartialText(); + handleAbortError(res, req, error, { + partialText, + conversationId, + sender: getResponseSender(endpointOption), + messageId: responseMessageId, + parentMessageId: userMessageId, + }); + } + }, +); + +module.exports = router; diff --git a/api/server/routes/endpoints.js b/api/server/routes/endpoints.js index b6016244c8..bee04998e1 100644 --- a/api/server/routes/endpoints.js +++ b/api/server/routes/endpoints.js @@ -3,37 +3,45 @@ const express = require('express'); const router = express.Router(); const { availableTools } = require('../../app/clients/tools'); const { addOpenAPISpecs } = require('../../app/clients/tools/util/addOpenAPISpecs'); +// const { getAzureCredentials, genAzureChatCompletion } = require('../../utils/'); const openAIApiKey = process.env.OPENAI_API_KEY; const azureOpenAIApiKey = process.env.AZURE_API_KEY; +const useAzurePlugins = !!process.env.PLUGINS_USE_AZURE; const userProvidedOpenAI = openAIApiKey ? openAIApiKey === 'user_provided' : azureOpenAIApiKey === 'user_provided'; const fetchOpenAIModels = async (opts = { azure: false, plugins: false }, _models = []) => { let models = _models.slice() ?? []; + let apiKey = openAIApiKey; + let basePath = 'https://api.openai.com/v1'; if (opts.azure) { - /* TODO: Add Azure models from api/models */ return models; + // const azure = getAzureCredentials(); + // basePath = (genAzureChatCompletion(azure)) + // .split('/deployments')[0] + // .concat(`/models?api-version=${azure.azureOpenAIApiVersion}`); + // apiKey = azureOpenAIApiKey; } - let basePath = 'https://api.openai.com/v1/'; const reverseProxyUrl = process.env.OPENAI_REVERSE_PROXY; if (reverseProxyUrl) { basePath = reverseProxyUrl.match(/.*v1/)[0]; } - if (basePath.includes('v1')) { + if (basePath.includes('v1') || opts.azure) { try { - const res = await axios.get(`${basePath}/models`, { + const res = await axios.get(`${basePath}${opts.azure ? '' : '/models'}`, { headers: { - Authorization: `Bearer ${openAIApiKey}`, + Authorization: `Bearer ${apiKey}`, }, }); models = res.data.data.map((item) => item.id); + // console.log(`Fetched ${models.length} models from ${opts.azure ? 'Azure ' : ''}OpenAI API`); } catch (err) { - console.error(err); + console.log(`Failed to fetch models from ${opts.azure ? 'Azure ' : ''}OpenAI API`); } } @@ -149,7 +157,7 @@ router.get('/', async function (req, res) { const gptPlugins = openAIApiKey || azureOpenAIApiKey ? { - availableModels: await getOpenAIModels({ plugins: true }), + availableModels: await getOpenAIModels({ azure: useAzurePlugins, plugins: true }), plugins, availableAgents: ['classic', 'functions'], userProvide: userProvidedOpenAI, diff --git a/api/server/routes/endpoints/anthropic/buildOptions.js b/api/server/routes/endpoints/anthropic/buildOptions.js new file mode 100644 index 0000000000..2b0143d2b0 --- /dev/null +++ b/api/server/routes/endpoints/anthropic/buildOptions.js @@ -0,0 +1,15 @@ +const buildOptions = (endpoint, parsedBody) => { + const { modelLabel, promptPrefix, ...rest } = parsedBody; + const endpointOption = { + endpoint, + modelLabel, + promptPrefix, + modelOptions: { + ...rest, + }, + }; + + return endpointOption; +}; + +module.exports = buildOptions; diff --git a/api/server/routes/endpoints/anthropic/index.js b/api/server/routes/endpoints/anthropic/index.js new file mode 100644 index 0000000000..84e4bd5973 --- /dev/null +++ b/api/server/routes/endpoints/anthropic/index.js @@ -0,0 +1,8 @@ +const buildOptions = require('./buildOptions'); +const initializeClient = require('./initializeClient'); + +module.exports = { + // addTitle, // todo + buildOptions, + initializeClient, +}; diff --git a/api/server/routes/endpoints/anthropic/initializeClient.js b/api/server/routes/endpoints/anthropic/initializeClient.js new file mode 100644 index 0000000000..eab0487f95 --- /dev/null +++ b/api/server/routes/endpoints/anthropic/initializeClient.js @@ -0,0 +1,12 @@ +const { AnthropicClient } = require('../../../../app'); + +const initializeClient = (req) => { + let anthropicApiKey = req.body?.token ?? process.env.ANTHROPIC_API_KEY; + const client = new AnthropicClient(anthropicApiKey); + return { + client, + anthropicApiKey, + }; +}; + +module.exports = initializeClient; diff --git a/api/server/routes/endpoints/gptPlugins/buildOptions.js b/api/server/routes/endpoints/gptPlugins/buildOptions.js new file mode 100644 index 0000000000..ebf4116ec3 --- /dev/null +++ b/api/server/routes/endpoints/gptPlugins/buildOptions.js @@ -0,0 +1,31 @@ +const buildOptions = (endpoint, parsedBody) => { + const { + chatGptLabel, + promptPrefix, + agentOptions, + tools, + model, + temperature, + top_p, + presence_penalty, + frequency_penalty, + } = parsedBody; + const endpointOption = { + endpoint, + tools: tools.map((tool) => tool.pluginKey) ?? [], + chatGptLabel, + promptPrefix, + agentOptions, + modelOptions: { + model, + temperature, + top_p, + presence_penalty, + frequency_penalty, + }, + }; + + return endpointOption; +}; + +module.exports = buildOptions; diff --git a/api/server/routes/endpoints/gptPlugins/index.js b/api/server/routes/endpoints/gptPlugins/index.js new file mode 100644 index 0000000000..3994468306 --- /dev/null +++ b/api/server/routes/endpoints/gptPlugins/index.js @@ -0,0 +1,7 @@ +const buildOptions = require('./buildOptions'); +const initializeClient = require('./initializeClient'); + +module.exports = { + buildOptions, + initializeClient, +}; diff --git a/api/server/routes/endpoints/gptPlugins/initializeClient.js b/api/server/routes/endpoints/gptPlugins/initializeClient.js new file mode 100644 index 0000000000..6035a6f5a4 --- /dev/null +++ b/api/server/routes/endpoints/gptPlugins/initializeClient.js @@ -0,0 +1,30 @@ +const { PluginsClient } = require('../../../../app'); +const { getAzureCredentials } = require('../../../../utils'); + +const initializeClient = (req, endpointOption) => { + const clientOptions = { + debug: true, + reverseProxyUrl: process.env.OPENAI_REVERSE_PROXY || null, + proxy: process.env.PROXY || null, + ...endpointOption, + }; + + let openAIApiKey = req.body?.token ?? process.env.OPENAI_API_KEY; + if (process.env.PLUGINS_USE_AZURE) { + clientOptions.azure = getAzureCredentials(); + openAIApiKey = clientOptions.azure.azureOpenAIApiKey; + } + + if (openAIApiKey && openAIApiKey.includes('azure') && !clientOptions.azure) { + clientOptions.azure = JSON.parse(req.body?.token) ?? getAzureCredentials(); + openAIApiKey = clientOptions.azure.azureOpenAIApiKey; + } + const client = new PluginsClient(openAIApiKey, clientOptions); + return { + client, + azure: clientOptions.azure, + openAIApiKey, + }; +}; + +module.exports = initializeClient; diff --git a/api/server/routes/endpoints/openAI/addTitle.js b/api/server/routes/endpoints/openAI/addTitle.js new file mode 100644 index 0000000000..c3f6c2ad2a --- /dev/null +++ b/api/server/routes/endpoints/openAI/addTitle.js @@ -0,0 +1,22 @@ +const { titleConvo } = require('../../../../app'); +const { saveConvo } = require('../../../../models'); + +const addTitle = async ( + req, + { text, azure, response, newConvo, parentMessageId, openAIApiKey }, +) => { + if (parentMessageId == '00000000-0000-0000-0000-000000000000' && newConvo) { + const title = await titleConvo({ + text, + azure, + response, + openAIApiKey, + }); + await saveConvo(req.user.id, { + conversationId: response.conversationId, + title, + }); + } +}; + +module.exports = addTitle; diff --git a/api/server/routes/endpoints/openAI/buildOptions.js b/api/server/routes/endpoints/openAI/buildOptions.js new file mode 100644 index 0000000000..a1ad232bb7 --- /dev/null +++ b/api/server/routes/endpoints/openAI/buildOptions.js @@ -0,0 +1,15 @@ +const buildOptions = (endpoint, parsedBody) => { + const { chatGptLabel, promptPrefix, ...rest } = parsedBody; + const endpointOption = { + endpoint, + chatGptLabel, + promptPrefix, + modelOptions: { + ...rest, + }, + }; + + return endpointOption; +}; + +module.exports = buildOptions; diff --git a/api/server/routes/endpoints/openAI/index.js b/api/server/routes/endpoints/openAI/index.js new file mode 100644 index 0000000000..772b1efb11 --- /dev/null +++ b/api/server/routes/endpoints/openAI/index.js @@ -0,0 +1,9 @@ +const addTitle = require('./addTitle'); +const buildOptions = require('./buildOptions'); +const initializeClient = require('./initializeClient'); + +module.exports = { + addTitle, + buildOptions, + initializeClient, +}; diff --git a/api/server/routes/endpoints/openAI/initializeClient.js b/api/server/routes/endpoints/openAI/initializeClient.js new file mode 100644 index 0000000000..4e4910c5b0 --- /dev/null +++ b/api/server/routes/endpoints/openAI/initializeClient.js @@ -0,0 +1,27 @@ +const { OpenAIClient } = require('../../../../app'); +const { getAzureCredentials } = require('../../../../utils'); + +const initializeClient = (req, endpointOption) => { + const clientOptions = { + // debug: true, + // contextStrategy: 'refine', + reverseProxyUrl: process.env.OPENAI_REVERSE_PROXY || null, + proxy: process.env.PROXY || null, + ...endpointOption, + }; + + let openAIApiKey = req.body?.token ?? process.env.OPENAI_API_KEY; + + if (process.env.AZURE_API_KEY && endpointOption.endpoint === 'azureOpenAI') { + clientOptions.azure = JSON.parse(req.body?.token) ?? getAzureCredentials(); + openAIApiKey = clientOptions.azure.azureOpenAIApiKey; + } + + const client = new OpenAIClient(openAIApiKey, clientOptions); + return { + client, + openAIApiKey, + }; +}; + +module.exports = initializeClient; diff --git a/api/server/routes/endpoints/schemas.js b/api/server/routes/endpoints/schemas.js new file mode 100644 index 0000000000..e7e9f30e8d --- /dev/null +++ b/api/server/routes/endpoints/schemas.js @@ -0,0 +1,369 @@ +const { z } = require('zod'); + +const EModelEndpoint = { + azureOpenAI: 'azureOpenAI', + openAI: 'openAI', + bingAI: 'bingAI', + chatGPTBrowser: 'chatGPTBrowser', + google: 'google', + gptPlugins: 'gptPlugins', + anthropic: 'anthropic', +}; + +const eModelEndpointSchema = z.nativeEnum(EModelEndpoint); + +/* +const tMessageSchema = z.object({ + messageId: z.string(), + clientId: z.string().nullable().optional(), + conversationId: z.string().nullable(), + parentMessageId: z.string().nullable(), + sender: z.string(), + text: z.string(), + isCreatedByUser: z.boolean(), + error: z.boolean(), + createdAt: z + .string() + .optional() + .default(() => new Date().toISOString()), + updatedAt: z + .string() + .optional() + .default(() => new Date().toISOString()), + current: z.boolean().optional(), + unfinished: z.boolean().optional(), + submitting: z.boolean().optional(), + searchResult: z.boolean().optional(), + finish_reason: z.string().optional(), +}); + +const tPresetSchema = tConversationSchema + .omit({ + conversationId: true, + createdAt: true, + updatedAt: true, + title: true, + }) + .merge( + z.object({ + conversationId: z.string().optional(), + presetId: z.string().nullable().optional(), + title: z.string().nullable().optional(), + }), + ); +*/ + +const tPluginAuthConfigSchema = z.object({ + authField: z.string(), + label: z.string(), + description: z.string(), +}); + +const tPluginSchema = z.object({ + name: z.string(), + pluginKey: z.string(), + description: z.string(), + icon: z.string(), + authConfig: z.array(tPluginAuthConfigSchema), + authenticated: z.boolean().optional(), + isButton: z.boolean().optional(), +}); + +const tExampleSchema = z.object({ + input: z.object({ + content: z.string(), + }), + output: z.object({ + content: z.string(), + }), +}); + +const tAgentOptionsSchema = z.object({ + agent: z.string(), + skipCompletion: z.boolean(), + model: z.string(), + temperature: z.number(), +}); + +const tConversationSchema = z.object({ + conversationId: z.string().nullable(), + title: z.string(), + user: z.string().optional(), + endpoint: eModelEndpointSchema.nullable(), + suggestions: z.array(z.string()).optional(), + messages: z.array(z.string()).optional(), + tools: z.array(tPluginSchema).optional(), + createdAt: z.string(), + updatedAt: z.string(), + systemMessage: z.string().nullable().optional(), + modelLabel: z.string().nullable().optional(), + examples: z.array(tExampleSchema).optional(), + chatGptLabel: z.string().nullable().optional(), + userLabel: z.string().optional(), + model: z.string().nullable().optional(), + promptPrefix: z.string().nullable().optional(), + temperature: z.number().optional(), + topP: z.number().optional(), + topK: z.number().optional(), + context: z.string().nullable().optional(), + top_p: z.number().optional(), + frequency_penalty: z.number().optional(), + presence_penalty: z.number().optional(), + jailbreak: z.boolean().optional(), + jailbreakConversationId: z.string().nullable().optional(), + conversationSignature: z.string().nullable().optional(), + parentMessageId: z.string().optional(), + clientId: z.string().nullable().optional(), + invocationId: z.number().nullable().optional(), + toneStyle: z.string().nullable().optional(), + maxOutputTokens: z.number().optional(), + agentOptions: tAgentOptionsSchema.nullable().optional(), +}); + +const openAISchema = tConversationSchema + .pick({ + model: true, + chatGptLabel: true, + promptPrefix: true, + temperature: true, + top_p: true, + presence_penalty: true, + frequency_penalty: true, + }) + .transform((obj) => ({ + ...obj, + model: obj.model ?? 'gpt-3.5-turbo', + chatGptLabel: obj.chatGptLabel ?? null, + promptPrefix: obj.promptPrefix ?? null, + temperature: obj.temperature ?? 1, + top_p: obj.top_p ?? 1, + presence_penalty: obj.presence_penalty ?? 0, + frequency_penalty: obj.frequency_penalty ?? 0, + })) + .catch(() => ({ + model: 'gpt-3.5-turbo', + chatGptLabel: null, + promptPrefix: null, + temperature: 1, + top_p: 1, + presence_penalty: 0, + frequency_penalty: 0, + })); + +const googleSchema = tConversationSchema + .pick({ + model: true, + modelLabel: true, + promptPrefix: true, + examples: true, + temperature: true, + maxOutputTokens: true, + topP: true, + topK: true, + }) + .transform((obj) => ({ + ...obj, + model: obj.model ?? 'chat-bison', + modelLabel: obj.modelLabel ?? null, + promptPrefix: obj.promptPrefix ?? null, + temperature: obj.temperature ?? 0.2, + maxOutputTokens: obj.maxOutputTokens ?? 1024, + topP: obj.topP ?? 0.95, + topK: obj.topK ?? 40, + })) + .catch(() => ({ + model: 'chat-bison', + modelLabel: null, + promptPrefix: null, + temperature: 0.2, + maxOutputTokens: 1024, + topP: 0.95, + topK: 40, + })); + +const bingAISchema = tConversationSchema + .pick({ + jailbreak: true, + systemMessage: true, + context: true, + toneStyle: true, + jailbreakConversationId: true, + conversationSignature: true, + clientId: true, + invocationId: true, + }) + .transform((obj) => ({ + ...obj, + model: '', + jailbreak: obj.jailbreak ?? false, + systemMessage: obj.systemMessage ?? null, + context: obj.context ?? null, + toneStyle: obj.toneStyle ?? 'creative', + jailbreakConversationId: obj.jailbreakConversationId ?? null, + conversationSignature: obj.conversationSignature ?? null, + clientId: obj.clientId ?? null, + invocationId: obj.invocationId ?? 1, + })) + .catch(() => ({ + model: '', + jailbreak: false, + systemMessage: null, + context: null, + toneStyle: 'creative', + jailbreakConversationId: null, + conversationSignature: null, + clientId: null, + invocationId: 1, + })); + +const anthropicSchema = tConversationSchema + .pick({ + model: true, + modelLabel: true, + promptPrefix: true, + temperature: true, + maxOutputTokens: true, + topP: true, + topK: true, + }) + .transform((obj) => ({ + ...obj, + model: obj.model ?? 'claude-1', + modelLabel: obj.modelLabel ?? null, + promptPrefix: obj.promptPrefix ?? null, + temperature: obj.temperature ?? 1, + maxOutputTokens: obj.maxOutputTokens ?? 1024, + topP: obj.topP ?? 0.7, + topK: obj.topK ?? 5, + })) + .catch(() => ({ + model: 'claude-1', + modelLabel: null, + promptPrefix: null, + temperature: 1, + maxOutputTokens: 1024, + topP: 0.7, + topK: 5, + })); + +const chatGPTBrowserSchema = tConversationSchema + .pick({ + model: true, + }) + .transform((obj) => ({ + ...obj, + model: obj.model ?? 'text-davinci-002-render-sha', + })) + .catch(() => ({ + model: 'text-davinci-002-render-sha', + })); + +const gptPluginsSchema = tConversationSchema + .pick({ + model: true, + chatGptLabel: true, + promptPrefix: true, + temperature: true, + top_p: true, + presence_penalty: true, + frequency_penalty: true, + tools: true, + agentOptions: true, + }) + .transform((obj) => ({ + ...obj, + model: obj.model ?? 'gpt-3.5-turbo', + chatGptLabel: obj.chatGptLabel ?? null, + promptPrefix: obj.promptPrefix ?? null, + temperature: obj.temperature ?? 0.8, + top_p: obj.top_p ?? 1, + presence_penalty: obj.presence_penalty ?? 0, + frequency_penalty: obj.frequency_penalty ?? 0, + tools: obj.tools ?? [], + agentOptions: obj.agentOptions ?? { + agent: 'functions', + skipCompletion: true, + model: 'gpt-3.5-turbo', + temperature: 0, + }, + })) + .catch(() => ({ + model: 'gpt-3.5-turbo', + chatGptLabel: null, + promptPrefix: null, + temperature: 0.8, + top_p: 1, + presence_penalty: 0, + frequency_penalty: 0, + tools: [], + agentOptions: { + agent: 'functions', + skipCompletion: true, + model: 'gpt-3.5-turbo', + temperature: 0, + }, + })); + +const endpointSchemas = { + openAI: openAISchema, + azureOpenAI: openAISchema, + google: googleSchema, + bingAI: bingAISchema, + anthropic: anthropicSchema, + chatGPTBrowser: chatGPTBrowserSchema, + gptPlugins: gptPluginsSchema, +}; + +function getFirstDefinedValue(possibleValues) { + let returnValue; + for (const value of possibleValues) { + if (value) { + returnValue = value; + break; + } + } + return returnValue; +} + +const parseConvo = (endpoint, conversation, possibleValues) => { + const schema = endpointSchemas[endpoint]; + + if (!schema) { + throw new Error(`Unknown endpoint: ${endpoint}`); + } + + const convo = schema.parse(conversation); + + if (possibleValues && convo) { + convo.model = getFirstDefinedValue(possibleValues.model) ?? convo.model; + } + + return convo; +}; + +const getResponseSender = (endpointOption) => { + const { endpoint, chatGptLabel, modelLabel, jailbreak } = endpointOption; + + if (['openAI', 'azureOpenAI', 'gptPlugins', 'chatGPTBrowser'].includes(endpoint)) { + return chatGptLabel ?? 'ChatGPT'; + } + + if (endpoint === 'bingAI') { + return jailbreak ? 'Sydney' : 'BingAI'; + } + + if (endpoint === 'anthropic') { + return modelLabel ?? 'Anthropic'; + } + + if (endpoint === 'google') { + return modelLabel ?? 'PaLM2'; + } + + return ''; +}; + +module.exports = { + parseConvo, + getResponseSender, +}; diff --git a/api/server/routes/index.js b/api/server/routes/index.js index 18d2a44fc4..c78af3fcff 100644 --- a/api/server/routes/index.js +++ b/api/server/routes/index.js @@ -1,4 +1,5 @@ const ask = require('./ask'); +const edit = require('./edit'); const messages = require('./messages'); const convos = require('./convos'); const presets = require('./presets'); @@ -15,6 +16,7 @@ const config = require('./config'); module.exports = { search, ask, + edit, messages, convos, presets, diff --git a/api/server/routes/messages.js b/api/server/routes/messages.js index a13b4272bc..0530ebc263 100644 --- a/api/server/routes/messages.js +++ b/api/server/routes/messages.js @@ -1,7 +1,7 @@ const express = require('express'); const router = express.Router(); const { getMessages } = require('../../models/Message'); -const requireJwtAuth = require('../../middleware/requireJwtAuth'); +const requireJwtAuth = require('../middleware/requireJwtAuth'); router.get('/:conversationId', requireJwtAuth, async (req, res) => { const { conversationId } = req.params; diff --git a/api/server/routes/plugins.js b/api/server/routes/plugins.js index cb93163242..4a7715a618 100644 --- a/api/server/routes/plugins.js +++ b/api/server/routes/plugins.js @@ -1,6 +1,6 @@ const express = require('express'); const { getAvailablePluginsController } = require('../controllers/PluginController'); -const requireJwtAuth = require('../../middleware/requireJwtAuth'); +const requireJwtAuth = require('../middleware/requireJwtAuth'); const router = express.Router(); diff --git a/api/server/routes/presets.js b/api/server/routes/presets.js index 8a08f0b509..127a8e5b6b 100644 --- a/api/server/routes/presets.js +++ b/api/server/routes/presets.js @@ -2,7 +2,7 @@ const express = require('express'); const router = express.Router(); const { getPresets, savePreset, deletePresets } = require('../../models'); const crypto = require('crypto'); -const requireJwtAuth = require('../../middleware/requireJwtAuth'); +const requireJwtAuth = require('../middleware/requireJwtAuth'); router.get('/', requireJwtAuth, async (req, res) => { const presets = (await getPresets(req.user.id)).map((preset) => { diff --git a/api/server/routes/search.js b/api/server/routes/search.js index 9f495bb386..955e58da97 100644 --- a/api/server/routes/search.js +++ b/api/server/routes/search.js @@ -5,7 +5,7 @@ const { Message } = require('../../models/Message'); const { Conversation, getConvosQueried } = require('../../models/Conversation'); const { reduceHits } = require('../../lib/utils/reduceHits'); const { cleanUpPrimaryKeyValue } = require('../../lib/utils/misc'); -const requireJwtAuth = require('../../middleware/requireJwtAuth'); +const requireJwtAuth = require('../middleware/requireJwtAuth'); const cache = new Map(); diff --git a/api/server/routes/tokenizer.js b/api/server/routes/tokenizer.js index 995263b0d0..9fd79d0acf 100644 --- a/api/server/routes/tokenizer.js +++ b/api/server/routes/tokenizer.js @@ -4,7 +4,7 @@ const { Tiktoken } = require('@dqbd/tiktoken/lite'); const { load } = require('@dqbd/tiktoken/load'); const registry = require('@dqbd/tiktoken/registry.json'); const models = require('@dqbd/tiktoken/model_to_encoding.json'); -const requireJwtAuth = require('../../middleware/requireJwtAuth'); +const requireJwtAuth = require('../middleware/requireJwtAuth'); router.post('/', requireJwtAuth, async (req, res) => { try { diff --git a/api/server/routes/user.js b/api/server/routes/user.js index 293ce4cf63..b90e3d965b 100644 --- a/api/server/routes/user.js +++ b/api/server/routes/user.js @@ -1,5 +1,5 @@ const express = require('express'); -const requireJwtAuth = require('../../middleware/requireJwtAuth'); +const requireJwtAuth = require('../middleware/requireJwtAuth'); const { getUserController, updateUserPluginsController } = require('../controllers/UserController'); const router = express.Router(); diff --git a/api/server/services/AuthService.js b/api/server/services/AuthService.js index 309a15a2ba..52380bc8a1 100644 --- a/api/server/services/AuthService.js +++ b/api/server/services/AuthService.js @@ -1,10 +1,10 @@ -const User = require('../../models/User'); -const Token = require('../../models/schema/tokenSchema'); const crypto = require('crypto'); const bcrypt = require('bcryptjs'); +const User = require('../../models/User'); +const Token = require('../../models/schema/tokenSchema'); const { registerSchema } = require('../../strategies/validators'); -const { sendEmail } = require('../../utils'); const config = require('../../../config/loader'); +const { sendEmail } = require('../utils'); const domains = config.domains; /** diff --git a/api/server/services/PluginService.js b/api/server/services/PluginService.js index 970f16f6d9..b70005dffa 100644 --- a/api/server/services/PluginService.js +++ b/api/server/services/PluginService.js @@ -1,5 +1,5 @@ const PluginAuth = require('../../models/schema/pluginAuthSchema'); -const { encrypt, decrypt } = require('../../utils/'); +const { encrypt, decrypt } = require('../utils/'); const getUserPluginAuthValue = async (user, authField) => { try { diff --git a/api/lib/parse/citeText.js b/api/server/utils/citations.js similarity index 67% rename from api/lib/parse/citeText.js rename to api/server/utils/citations.js index 8fc1cea8b4..33136c18b8 100644 --- a/api/lib/parse/citeText.js +++ b/api/server/utils/citations.js @@ -1,4 +1,19 @@ const citationRegex = /\[\^\d+?\^\]/g; +const regex = / \[.*?]\(.*?\)/g; + +const getCitations = (res) => { + const adaptiveCards = res.details.adaptiveCards; + const textBlocks = adaptiveCards && adaptiveCards[0].body; + if (!textBlocks) { + return ''; + } + let links = textBlocks[textBlocks.length - 1]?.text.match(regex); + if (links?.length === 0 || !links) { + return ''; + } + links = links.map((link) => link.trim()); + return links.join('\n - '); +}; const citeText = (res, noLinks = false) => { let result = res.text || res; @@ -32,4 +47,4 @@ const citeText = (res, noLinks = false) => { return result; }; -module.exports = citeText; +module.exports = { getCitations, citeText }; diff --git a/api/utils/crypto.js b/api/server/utils/crypto.js similarity index 100% rename from api/utils/crypto.js rename to api/server/utils/crypto.js diff --git a/api/utils/emails/passwordReset.handlebars b/api/server/utils/emails/passwordReset.handlebars similarity index 100% rename from api/utils/emails/passwordReset.handlebars rename to api/server/utils/emails/passwordReset.handlebars diff --git a/api/utils/emails/requestPasswordReset.handlebars b/api/server/utils/emails/requestPasswordReset.handlebars similarity index 100% rename from api/utils/emails/requestPasswordReset.handlebars rename to api/server/utils/emails/requestPasswordReset.handlebars diff --git a/api/server/routes/ask/handlers.js b/api/server/utils/handleText.js similarity index 93% rename from api/server/routes/ask/handlers.js rename to api/server/utils/handleText.js index d917c65ca4..b5efa08d87 100644 --- a/api/server/routes/ask/handlers.js +++ b/api/server/utils/handleText.js @@ -1,8 +1,10 @@ const _ = require('lodash'); const citationRegex = /\[\^\d+?\^]/g; -const { getCitations, citeText } = require('../../../app'); +const { getCitations, citeText } = require('./citations'); const cursor = ''; +const addSpaceIfNeeded = (text) => (text.length > 0 && !text.endsWith(' ') ? text + ' ' : text); + const handleError = (res, message) => { res.write(`event: error\ndata: ${JSON.stringify(message)}\n\n`); res.end(); @@ -15,12 +17,12 @@ const sendMessage = (res, message, event = 'message') => { res.write(`event: ${event}\ndata: ${JSON.stringify(message)}\n\n`); }; -const createOnProgress = ({ onProgress: _onProgress }) => { +const createOnProgress = ({ generation = '', onProgress: _onProgress }) => { let i = 0; let code = ''; - let tokens = ''; let precode = ''; let codeBlock = false; + let tokens = addSpaceIfNeeded(generation); const progressCallback = async (partial, { res, text, plugin, bing = false, ...rest }) => { let chunk = partial === text ? '' : partial; @@ -155,4 +157,5 @@ module.exports = { handleText, formatSteps, formatAction, + addSpaceIfNeeded, }; diff --git a/api/server/utils/index.js b/api/server/utils/index.js new file mode 100644 index 0000000000..e76d5b4365 --- /dev/null +++ b/api/server/utils/index.js @@ -0,0 +1,11 @@ +const cryptoUtils = require('./crypto'); +const handleText = require('./handleText'); +const citations = require('./citations'); +const sendEmail = require('./sendEmail'); + +module.exports = { + ...cryptoUtils, + ...handleText, + ...citations, + sendEmail, +}; diff --git a/api/utils/sendEmail.js b/api/server/utils/sendEmail.js similarity index 100% rename from api/utils/sendEmail.js rename to api/server/utils/sendEmail.js diff --git a/api/utils/abortMessage.js b/api/utils/abortMessage.js deleted file mode 100644 index fea33eb4c7..0000000000 --- a/api/utils/abortMessage.js +++ /dev/null @@ -1,18 +0,0 @@ -async function abortMessage(req, res, abortControllers) { - const { abortKey } = req.body; - console.log('req.body', req.body); - if (!abortControllers.has(abortKey)) { - return res.status(404).send('Request not found'); - } - - const { abortController } = abortControllers.get(abortKey); - - abortControllers.delete(abortKey); - const ret = await abortController.abortAsk(); - console.log('Aborted request', abortKey); - console.log('Aborted message:', ret); - - res.send(JSON.stringify(ret)); -} - -module.exports = abortMessage; diff --git a/api/utils/index.js b/api/utils/index.js index 0a4dd75bf5..7de983c0f0 100644 --- a/api/utils/index.js +++ b/api/utils/index.js @@ -1,16 +1,10 @@ const azureUtils = require('./azureUtils'); -const cryptoUtils = require('./crypto'); const { tiktokenModels, maxTokensMap } = require('./tokens'); -const sendEmail = require('./sendEmail'); -const abortMessage = require('./abortMessage'); const findMessageContent = require('./findMessageContent'); module.exports = { - ...cryptoUtils, ...azureUtils, maxTokensMap, tiktokenModels, - sendEmail, - abortMessage, findMessageContent, }; diff --git a/client/src/common/types.ts b/client/src/common/types.ts index dfcf5d5446..9507782358 100644 --- a/client/src/common/types.ts +++ b/client/src/common/types.ts @@ -48,3 +48,17 @@ export type TSetOptionsPayload = { checkPluginSelection: (value: string) => boolean; setTools: (newValue: string) => void; }; + +export type TPresetItemProps = { + preset: TPreset; + value: TPreset; + onSelect: (preset: TPreset) => void; + onChangePreset: (preset: TPreset) => void; + onDeletePreset: (preset: TPreset) => void; +}; + +export type TOnClick = (e: React.MouseEvent) => void; + +export type TGenButtonProps = { + onClick: TOnClick; +}; diff --git a/client/src/components/Input/EndpointMenu/EndpointMenu.jsx b/client/src/components/Input/EndpointMenu/EndpointMenu.jsx index d7bbf439e9..cfc2c25bc4 100644 --- a/client/src/components/Input/EndpointMenu/EndpointMenu.jsx +++ b/client/src/components/Input/EndpointMenu/EndpointMenu.jsx @@ -103,10 +103,18 @@ export default function NewConversationMenu() { }; // set the current model + const isModular = modularEndpoints.has(endpoint); const onSelectPreset = (newPreset) => { setMenuOpen(false); + if (!newPreset) { + return; + } - if (modularEndpoints.has(endpoint) && modularEndpoints.has(newPreset?.endpoint)) { + if ( + isModular && + modularEndpoints.has(newPreset?.endpoint) && + endpoint === newPreset?.endpoint + ) { const currentConvo = getDefaultConversation({ conversation, endpointsConfig, @@ -118,10 +126,6 @@ export default function NewConversationMenu() { return; } - if (!newPreset) { - return; - } - newConversation({}, newPreset); }; diff --git a/client/src/components/Input/EndpointMenu/PresetItem.jsx b/client/src/components/Input/EndpointMenu/PresetItem.tsx similarity index 87% rename from client/src/components/Input/EndpointMenu/PresetItem.jsx rename to client/src/components/Input/EndpointMenu/PresetItem.tsx index ca5d6b86a7..47ac0fb9f9 100644 --- a/client/src/components/Input/EndpointMenu/PresetItem.jsx +++ b/client/src/components/Input/EndpointMenu/PresetItem.tsx @@ -1,7 +1,14 @@ +import type { TPresetItemProps } from '~/common'; +import type { TPreset } from 'librechat-data-provider'; import { DropdownMenuRadioItem, EditIcon, TrashIcon } from '~/components'; import { getIcon } from '~/components/Endpoints'; -export default function PresetItem({ preset = {}, value, onChangePreset, onDeletePreset }) { +export default function PresetItem({ + preset = {} as TPreset, + value, + onChangePreset, + onDeletePreset, +}: TPresetItemProps) { const { endpoint } = preset; const icon = getIcon({ @@ -14,9 +21,9 @@ export default function PresetItem({ preset = {}, value, onChangePreset, onDelet const getPresetTitle = () => { let _title = `${endpoint}`; + const { chatGptLabel, modelLabel, model, jailbreak, toneStyle } = preset; if (endpoint === 'azureOpenAI' || endpoint === 'openAI') { - const { chatGptLabel, model } = preset; if (model) { _title += `: ${model}`; } @@ -24,7 +31,6 @@ export default function PresetItem({ preset = {}, value, onChangePreset, onDelet _title += ` as ${chatGptLabel}`; } } else if (endpoint === 'google') { - const { modelLabel, model } = preset; if (model) { _title += `: ${model}`; } @@ -32,7 +38,6 @@ export default function PresetItem({ preset = {}, value, onChangePreset, onDelet _title += ` as ${modelLabel}`; } } else if (endpoint === 'bingAI') { - const { jailbreak, toneStyle } = preset; if (toneStyle) { _title += `: ${toneStyle}`; } @@ -40,12 +45,10 @@ export default function PresetItem({ preset = {}, value, onChangePreset, onDelet _title += ' as Sydney'; } } else if (endpoint === 'chatGPTBrowser') { - const { model } = preset; if (model) { _title += `: ${model}`; } } else if (endpoint === 'gptPlugins') { - const { model } = preset; if (model) { _title += `: ${model}`; } @@ -60,6 +63,7 @@ export default function PresetItem({ preset = {}, value, onChangePreset, onDelet // regular model return ( diff --git a/client/src/components/Input/EndpointMenu/PresetItems.jsx b/client/src/components/Input/EndpointMenu/PresetItems.tsx similarity index 82% rename from client/src/components/Input/EndpointMenu/PresetItems.jsx rename to client/src/components/Input/EndpointMenu/PresetItems.tsx index a7048bc685..5e6e47b509 100644 --- a/client/src/components/Input/EndpointMenu/PresetItems.jsx +++ b/client/src/components/Input/EndpointMenu/PresetItems.tsx @@ -1,10 +1,11 @@ import React from 'react'; import PresetItem from './PresetItem'; +import type { TPreset } from 'librechat-data-provider'; export default function PresetItems({ presets, onSelect, onChangePreset, onDeletePreset }) { return ( <> - {presets.map((preset) => ( + {presets.map((preset: TPreset) => ( ) => void; + className?: string; +}) { + return ( + + ); +} diff --git a/client/src/components/Input/Generations/Continue.tsx b/client/src/components/Input/Generations/Continue.tsx new file mode 100644 index 0000000000..7d472b30f6 --- /dev/null +++ b/client/src/components/Input/Generations/Continue.tsx @@ -0,0 +1,12 @@ +import type { TGenButtonProps } from '~/common'; +import { ContinueIcon } from '~/components/svg'; +import Button from './Button'; + +export default function Continue({ onClick }: TGenButtonProps) { + return ( + + ); +} diff --git a/client/src/components/Input/Generations/GenerationButtons.tsx b/client/src/components/Input/Generations/GenerationButtons.tsx new file mode 100644 index 0000000000..7c51f59e09 --- /dev/null +++ b/client/src/components/Input/Generations/GenerationButtons.tsx @@ -0,0 +1,61 @@ +import type { TMessage } from 'librechat-data-provider'; +import { useMessageHandler, useMediaQuery, useGenerations } from '~/hooks'; +import { cn } from '~/utils'; +import Regenerate from './Regenerate'; +import Continue from './Continue'; +import Stop from './Stop'; + +type GenerationButtonsProps = { + endpoint: string; + showPopover: boolean; + opacityClass: string; +}; + +export default function GenerationButtons({ + endpoint, + showPopover, + opacityClass, +}: GenerationButtonsProps) { + const { + messages, + isSubmitting, + latestMessage, + handleContinue, + handleRegenerate, + handleStopGenerating, + } = useMessageHandler(); + const isSmallScreen = useMediaQuery('(max-width: 768px)'); + const { continueSupported, regenerateEnabled } = useGenerations({ + endpoint, + message: latestMessage as TMessage, + isSubmitting, + }); + + if (isSmallScreen) { + return null; + } + + let button: React.ReactNode = null; + + if (isSubmitting) { + button = ; + } else if (continueSupported) { + button = ; + } else if (messages && messages.length > 0 && regenerateEnabled) { + button = ; + } + + return ( +
+
+
+
+ {button} +
+
+
+ ); +} diff --git a/client/src/components/Input/Generations/Regenerate.tsx b/client/src/components/Input/Generations/Regenerate.tsx new file mode 100644 index 0000000000..8187ca3260 --- /dev/null +++ b/client/src/components/Input/Generations/Regenerate.tsx @@ -0,0 +1,12 @@ +import type { TGenButtonProps } from '~/common'; +import { RegenerateIcon } from '~/components/svg'; +import Button from './Button'; + +export default function Regenerate({ onClick }: TGenButtonProps) { + return ( + + ); +} diff --git a/client/src/components/Input/Generations/Stop.tsx b/client/src/components/Input/Generations/Stop.tsx new file mode 100644 index 0000000000..41579d6c5f --- /dev/null +++ b/client/src/components/Input/Generations/Stop.tsx @@ -0,0 +1,12 @@ +import type { TGenButtonProps } from '~/common'; +import { StopGeneratingIcon } from '~/components/svg'; +import Button from './Button'; + +export default function Stop({ onClick }: TGenButtonProps) { + return ( + + ); +} diff --git a/client/src/components/Input/Generations/index.ts b/client/src/components/Input/Generations/index.ts new file mode 100644 index 0000000000..bbf5aeb41b --- /dev/null +++ b/client/src/components/Input/Generations/index.ts @@ -0,0 +1 @@ +export { default as GenerationButtons } from './GenerationButtons'; diff --git a/client/src/components/Input/OptionsBar.tsx b/client/src/components/Input/OptionsBar.tsx index 12dfc63360..988271c46b 100644 --- a/client/src/components/Input/OptionsBar.tsx +++ b/client/src/components/Input/OptionsBar.tsx @@ -12,7 +12,7 @@ import { Button } from '~/components/ui'; import { cn, cardStyle } from '~/utils/'; import { useSetOptions } from '~/hooks'; import { ModelSelect } from './ModelSelect'; -import GenerationButtons from './GenerationButtons'; +import { GenerationButtons } from './Generations'; import store from '~/store'; export default function OptionsBar() { @@ -76,7 +76,11 @@ export default function OptionsBar() { : () => setShowPopover((prev) => !prev); return (
- +
- {children} - {text} - {/* -Regenerate response */} - - ); -} diff --git a/client/src/components/Input/TextChat.jsx b/client/src/components/Input/TextChat.jsx index e35caca9cb..ec2a4f2265 100644 --- a/client/src/components/Input/TextChat.jsx +++ b/client/src/components/Input/TextChat.jsx @@ -10,22 +10,18 @@ import { cn } from '~/utils'; import store from '~/store'; export default function TextChat({ isSearchView = false }) { - const inputRef = useRef(null); - const isComposing = useRef(false); - + const { ask, isSubmitting, handleStopGenerating, latestMessage, endpointsConfig } = + useMessageHandler(); + const conversation = useRecoilValue(store.conversation); + const setShowBingToneSetting = useSetRecoilState(store.showBingToneSetting); const [text, setText] = useRecoilState(store.text); const { theme } = useContext(ThemeContext); - const conversation = useRecoilValue(store.conversation); - const latestMessage = useRecoilValue(store.latestMessage); - - const endpointsConfig = useRecoilValue(store.endpointsConfig); - const isSubmitting = useRecoilValue(store.isSubmitting); - const setShowBingToneSetting = useSetRecoilState(store.showBingToneSetting); + const isComposing = useRef(false); + const inputRef = useRef(null); // TODO: do we need this? const disabled = false; - const { ask, stopGenerating } = useMessageHandler(); const isNotAppendable = latestMessage?.unfinished & !isSubmitting || latestMessage?.error; const { conversationId, jailbreak } = conversation || {}; @@ -60,11 +56,6 @@ export default function TextChat({ isSearchView = false }) { setText(''); }; - const handleStopGenerating = (e) => { - e.preventDefault(); - stopGenerating(); - }; - const handleKeyDown = (e) => { if (e.key === 'Enter' && isSubmitting) { return; diff --git a/client/src/components/Messages/HoverButtons.jsx b/client/src/components/Messages/HoverButtons.jsx deleted file mode 100644 index 4d481c7a3e..0000000000 --- a/client/src/components/Messages/HoverButtons.jsx +++ /dev/null @@ -1,87 +0,0 @@ -import React from 'react'; -import { cn } from '~/utils/'; -import Clipboard from '../svg/Clipboard'; -import CheckMark from '../svg/CheckMark'; -import EditIcon from '../svg/EditIcon'; -import RegenerateIcon from '../svg/RegenerateIcon'; - -export default function HoverButtons({ - isEditting, - enterEdit, - copyToClipboard, - conversation, - isSubmitting, - message, - regenerate, -}) { - const { endpoint } = conversation; - const [isCopied, setIsCopied] = React.useState(false); - - const branchingSupported = - // azureOpenAI, openAI, chatGPTBrowser support branching, so edit enabled // 5/21/23: Bing is allowing editing and Message regenerating - !![ - 'azureOpenAI', - 'openAI', - 'chatGPTBrowser', - 'google', - 'bingAI', - 'gptPlugins', - 'anthropic', - ].find((e) => e === endpoint); - // Sydney in bingAI supports branching, so edit enabled - - const editEnabled = - !message?.error && - message?.isCreatedByUser && - !message?.searchResult && - !isEditting && - branchingSupported; - - // for now, once branching is supported, regerate will be enabled - let regenerateEnabled = - // !message?.error && - !message?.isCreatedByUser && - !message?.searchResult && - !isEditting && - !isSubmitting && - branchingSupported; - - return ( -
- {editEnabled ? ( - - ) : null} - {regenerateEnabled ? ( - - ) : null} - - -
- ); -} diff --git a/client/src/components/Messages/HoverButtons.tsx b/client/src/components/Messages/HoverButtons.tsx new file mode 100644 index 0000000000..5eec313ddb --- /dev/null +++ b/client/src/components/Messages/HoverButtons.tsx @@ -0,0 +1,86 @@ +import { useState } from 'react'; +import type { TConversation, TMessage } from 'librechat-data-provider'; +import { Clipboard, CheckMark, EditIcon, RegenerateIcon, ContinueIcon } from '~/components/svg'; +import { useGenerations } from '~/hooks'; +import { cn } from '~/utils'; + +type THoverButtons = { + isEditing: boolean; + enterEdit: () => void; + copyToClipboard: (setIsCopied: (isCopied: boolean) => void) => void; + conversation: TConversation; + isSubmitting: boolean; + message: TMessage; + regenerate: () => void; + handleContinue: (e: React.MouseEvent) => void; +}; + +export default function HoverButtons({ + isEditing, + enterEdit, + copyToClipboard, + conversation, + isSubmitting, + message, + regenerate, + handleContinue, +}: THoverButtons) { + const { endpoint } = conversation; + const [isCopied, setIsCopied] = useState(false); + const { editEnabled, regenerateEnabled, continueSupported } = useGenerations({ + isEditing, + isSubmitting, + message, + endpoint: endpoint ?? '', + }); + + return ( +
+ + + {regenerateEnabled ? ( + + ) : null} + {continueSupported ? ( + + ) : null} +
+ ); +} diff --git a/client/src/components/Messages/Message.jsx b/client/src/components/Messages/Message.jsx index 87ce0c4034..a27e6bbc43 100644 --- a/client/src/components/Messages/Message.jsx +++ b/client/src/components/Messages/Message.jsx @@ -1,6 +1,6 @@ /* eslint-disable react-hooks/exhaustive-deps */ import { useState, useEffect, useRef } from 'react'; -import { useRecoilValue, useSetRecoilState } from 'recoil'; +import { useSetRecoilState } from 'recoil'; import copy from 'copy-to-clipboard'; import Plugin from './Plugin'; import SubRow from './Content/SubRow'; @@ -25,13 +25,12 @@ export default function Message({ setSiblingIdx, }) { const { text, searchResult, isCreatedByUser, error, submitting, unfinished } = message; - const isSubmitting = useRecoilValue(store.isSubmitting); const setLatestMessage = useSetRecoilState(store.latestMessage); const [abortScroll, setAbort] = useState(false); const textEditor = useRef(null); const last = !message?.children?.length; const edit = message.messageId == currentEditId; - const { ask, regenerate } = useMessageHandler(); + const { isSubmitting, ask, regenerate, handleContinue } = useMessageHandler(); const { switchToConversation } = store.useConversation(); const blinker = submitting && isSubmitting; const getConversationQuery = useGetConversationByIdQuery(message.conversationId, { @@ -223,12 +222,13 @@ export default function Message({ )}
enterEdit()} regenerate={() => regenerateMessage()} + handleContinue={handleContinue} copyToClipboard={copyToClipboard} /> diff --git a/client/src/components/Messages/ScrollToBottom.jsx b/client/src/components/Messages/ScrollToBottom.tsx similarity index 79% rename from client/src/components/Messages/ScrollToBottom.jsx rename to client/src/components/Messages/ScrollToBottom.tsx index 281f0d420f..3c76555ff2 100644 --- a/client/src/components/Messages/ScrollToBottom.jsx +++ b/client/src/components/Messages/ScrollToBottom.tsx @@ -1,10 +1,12 @@ -import React from 'react'; +type Props = { + scrollHandler: React.MouseEventHandler; +}; -export default function ScrollToBottom({ scrollHandler }) { +export default function ScrollToBottom({ scrollHandler }: Props) { return (