| const crypto = require('crypto'); |
| const fetch = require('node-fetch'); |
| const { supportsBalanceCheck, Constants } = require('librechat-data-provider'); |
| const { getConvo, getMessages, saveMessage, updateMessage, saveConvo } = require('~/models'); |
| const { addSpaceIfNeeded, isEnabled } = require('~/server/utils'); |
| const checkBalance = require('~/models/checkBalance'); |
| const { getFiles } = require('~/models/File'); |
| const TextStream = require('./TextStream'); |
| const { logger } = require('~/config'); |
|
|
| class BaseClient { |
| constructor(apiKey, options = {}) { |
| this.apiKey = apiKey; |
| this.sender = options.sender ?? 'AI'; |
| this.contextStrategy = null; |
| this.currentDateString = new Date().toLocaleDateString('en-us', { |
| year: 'numeric', |
| month: 'long', |
| day: 'numeric', |
| }); |
| this.fetch = this.fetch.bind(this); |
| |
| this.skipSaveConvo = false; |
| |
| this.skipSaveUserMessage = false; |
| } |
|
|
| setOptions() { |
| throw new Error('Method \'setOptions\' must be implemented.'); |
| } |
|
|
| async getCompletion() { |
| throw new Error('Method \'getCompletion\' must be implemented.'); |
| } |
|
|
| async sendCompletion() { |
| throw new Error('Method \'sendCompletion\' must be implemented.'); |
| } |
|
|
| getSaveOptions() { |
| throw new Error('Subclasses must implement getSaveOptions'); |
| } |
|
|
| async buildMessages() { |
| throw new Error('Subclasses must implement buildMessages'); |
| } |
|
|
| async summarizeMessages() { |
| throw new Error('Subclasses attempted to call summarizeMessages without implementing it'); |
| } |
|
|
| async getTokenCountForResponse(response) { |
| logger.debug('`[BaseClient] recordTokenUsage` not implemented.', response); |
| } |
|
|
| async recordTokenUsage({ promptTokens, completionTokens }) { |
| logger.debug('`[BaseClient] recordTokenUsage` not implemented.', { |
| promptTokens, |
| completionTokens, |
| }); |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| async fetch(_url, init) { |
| let url = _url; |
| if (this.options.directEndpoint) { |
| url = this.options.reverseProxyUrl; |
| } |
| logger.debug(`Making request to ${url}`); |
| if (typeof Bun !== 'undefined') { |
| return await fetch(url, init); |
| } |
| return await fetch(url, init); |
| } |
|
|
| getBuildMessagesOptions() { |
| throw new Error('Subclasses must implement getBuildMessagesOptions'); |
| } |
|
|
| async generateTextStream(text, onProgress, options = {}) { |
| const stream = new TextStream(text, options); |
| await stream.processTextStream(onProgress); |
| } |
|
|
| |
| |
| |
| processOverideIds() { |
| |
| let { overrideConvoId, overrideUserMessageId } = this.options?.req?.body ?? {}; |
| if (overrideConvoId) { |
| const [conversationId, index] = overrideConvoId.split(Constants.COMMON_DIVIDER); |
| overrideConvoId = conversationId; |
| if (index !== '0') { |
| this.skipSaveConvo = true; |
| } |
| } |
| if (overrideUserMessageId) { |
| const [userMessageId, index] = overrideUserMessageId.split(Constants.COMMON_DIVIDER); |
| overrideUserMessageId = userMessageId; |
| if (index !== '0') { |
| this.skipSaveUserMessage = true; |
| } |
| } |
|
|
| return [overrideConvoId, overrideUserMessageId]; |
| } |
|
|
| async setMessageOptions(opts = {}) { |
| if (opts && opts.replaceOptions) { |
| this.setOptions(opts); |
| } |
|
|
| const [overrideConvoId, overrideUserMessageId] = this.processOverideIds(); |
| const { isEdited, isContinued } = opts; |
| const user = opts.user ?? null; |
| this.user = user; |
| const saveOptions = this.getSaveOptions(); |
| this.abortController = opts.abortController ?? new AbortController(); |
| const conversationId = overrideConvoId ?? opts.conversationId ?? crypto.randomUUID(); |
| const parentMessageId = opts.parentMessageId ?? Constants.NO_PARENT; |
| const userMessageId = |
| overrideUserMessageId ?? opts.overrideParentMessageId ?? crypto.randomUUID(); |
| let responseMessageId = opts.responseMessageId ?? crypto.randomUUID(); |
| let head = isEdited ? responseMessageId : parentMessageId; |
| this.currentMessages = (await this.loadHistory(conversationId, head)) ?? []; |
| this.conversationId = conversationId; |
|
|
| if (isEdited && !isContinued) { |
| responseMessageId = crypto.randomUUID(); |
| head = responseMessageId; |
| this.currentMessages[this.currentMessages.length - 1].messageId = head; |
| } |
|
|
| return { |
| ...opts, |
| user, |
| head, |
| conversationId, |
| parentMessageId, |
| userMessageId, |
| responseMessageId, |
| saveOptions, |
| }; |
| } |
|
|
| createUserMessage({ messageId, parentMessageId, conversationId, text }) { |
| return { |
| messageId, |
| parentMessageId, |
| conversationId, |
| sender: 'User', |
| text, |
| isCreatedByUser: true, |
| }; |
| } |
|
|
| async handleStartMethods(message, opts) { |
| const { |
| user, |
| head, |
| conversationId, |
| 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?.getReqData === 'function') { |
| opts.getReqData({ |
| userMessage, |
| conversationId, |
| responseMessageId, |
| }); |
| } |
|
|
| if (typeof opts?.onStart === 'function') { |
| opts.onStart(userMessage, responseMessageId); |
| } |
|
|
| return { |
| ...opts, |
| user, |
| head, |
| conversationId, |
| responseMessageId, |
| saveOptions, |
| userMessage, |
| }; |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| addInstructions(messages, instructions) { |
| const payload = []; |
| if (!instructions || Object.keys(instructions).length === 0) { |
| return messages; |
| } |
| if (messages.length > 1) { |
| payload.push(...messages.slice(0, -1)); |
| } |
|
|
| payload.push(instructions); |
|
|
| if (messages.length > 0) { |
| payload.push(messages[messages.length - 1]); |
| } |
|
|
| return payload; |
| } |
|
|
| async handleTokenCountMap(tokenCountMap) { |
| if (this.currentMessages.length === 0) { |
| return; |
| } |
|
|
| for (let i = 0; i < this.currentMessages.length; i++) { |
| |
| if (i === this.currentMessages.length - 1) { |
| break; |
| } |
|
|
| const message = this.currentMessages[i]; |
| const { messageId } = message; |
| const update = {}; |
|
|
| if (messageId === tokenCountMap.summaryMessage?.messageId) { |
| logger.debug(`[BaseClient] Adding summary props to ${messageId}.`); |
|
|
| update.summary = tokenCountMap.summaryMessage.content; |
| update.summaryTokenCount = tokenCountMap.summaryMessage.tokenCount; |
| } |
|
|
| if (message.tokenCount && !update.summaryTokenCount) { |
| logger.debug(`[BaseClient] Skipping ${messageId}: already had a token count.`); |
| continue; |
| } |
|
|
| const tokenCount = tokenCountMap[messageId]; |
| if (tokenCount) { |
| message.tokenCount = tokenCount; |
| update.tokenCount = tokenCount; |
| await this.updateMessageInDatabase({ messageId, ...update }); |
| } |
| } |
| } |
|
|
| concatenateMessages(messages) { |
| return messages.reduce((acc, message) => { |
| const nameOrRole = message.name ?? message.role; |
| return acc + `${nameOrRole}:\n${message.content}\n\n`; |
| }, ''); |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| async getMessagesWithinTokenLimit(_messages, maxContextTokens) { |
| |
| |
| let currentTokenCount = 3; |
| let summaryIndex = -1; |
| let remainingContextTokens = maxContextTokens ?? this.maxContextTokens; |
| const messages = [..._messages]; |
|
|
| const context = []; |
| if (currentTokenCount < remainingContextTokens) { |
| while (messages.length > 0 && currentTokenCount < remainingContextTokens) { |
| const poppedMessage = messages.pop(); |
| const { tokenCount } = poppedMessage; |
|
|
| if (poppedMessage && currentTokenCount + tokenCount <= remainingContextTokens) { |
| context.push(poppedMessage); |
| currentTokenCount += tokenCount; |
| } else { |
| messages.push(poppedMessage); |
| break; |
| } |
| } |
| } |
|
|
| const prunedMemory = messages; |
| summaryIndex = prunedMemory.length - 1; |
| remainingContextTokens -= currentTokenCount; |
|
|
| return { |
| context: context.reverse(), |
| remainingContextTokens, |
| messagesToRefine: prunedMemory, |
| summaryIndex, |
| }; |
| } |
|
|
| async handleContextStrategy({ instructions, orderedMessages, formattedMessages }) { |
| let _instructions; |
| let tokenCount; |
|
|
| if (instructions) { |
| ({ tokenCount, ..._instructions } = instructions); |
| } |
| _instructions && logger.debug('[BaseClient] instructions tokenCount: ' + tokenCount); |
| let payload = this.addInstructions(formattedMessages, _instructions); |
| let orderedWithInstructions = this.addInstructions(orderedMessages, instructions); |
|
|
| let { context, remainingContextTokens, messagesToRefine, summaryIndex } = |
| await this.getMessagesWithinTokenLimit(orderedWithInstructions); |
|
|
| logger.debug('[BaseClient] Context Count (1/2)', { |
| remainingContextTokens, |
| maxContextTokens: this.maxContextTokens, |
| }); |
|
|
| let summaryMessage; |
| let summaryTokenCount; |
| let { shouldSummarize } = this; |
|
|
| |
| const { length } = payload; |
| const diff = length - context.length; |
| const firstMessage = orderedWithInstructions[0]; |
| const usePrevSummary = |
| shouldSummarize && |
| diff === 1 && |
| firstMessage?.summary && |
| this.previous_summary.messageId === firstMessage.messageId; |
|
|
| if (diff > 0) { |
| payload = payload.slice(diff); |
| logger.debug( |
| `[BaseClient] Difference between original payload (${length}) and context (${context.length}): ${diff}`, |
| ); |
| } |
|
|
| const latestMessage = orderedWithInstructions[orderedWithInstructions.length - 1]; |
| if (payload.length === 0 && !shouldSummarize && latestMessage) { |
| throw new Error( |
| `Prompt token count of ${latestMessage.tokenCount} exceeds max token count of ${this.maxContextTokens}.`, |
| ); |
| } |
|
|
| if (usePrevSummary) { |
| summaryMessage = { role: 'system', content: firstMessage.summary }; |
| summaryTokenCount = firstMessage.summaryTokenCount; |
| payload.unshift(summaryMessage); |
| remainingContextTokens -= summaryTokenCount; |
| } else if (shouldSummarize && messagesToRefine.length > 0) { |
| ({ summaryMessage, summaryTokenCount } = await this.summarizeMessages({ |
| messagesToRefine, |
| remainingContextTokens, |
| })); |
| summaryMessage && payload.unshift(summaryMessage); |
| remainingContextTokens -= summaryTokenCount; |
| } |
|
|
| |
| shouldSummarize = summaryMessage && shouldSummarize; |
|
|
| logger.debug('[BaseClient] Context Count (2/2)', { |
| remainingContextTokens, |
| maxContextTokens: this.maxContextTokens, |
| }); |
|
|
| let tokenCountMap = orderedWithInstructions.reduce((map, message, index) => { |
| const { messageId } = message; |
| if (!messageId) { |
| return map; |
| } |
|
|
| if (shouldSummarize && index === summaryIndex && !usePrevSummary) { |
| map.summaryMessage = { ...summaryMessage, messageId, tokenCount: summaryTokenCount }; |
| } |
|
|
| map[messageId] = orderedWithInstructions[index].tokenCount; |
| return map; |
| }, {}); |
|
|
| const promptTokens = this.maxContextTokens - remainingContextTokens; |
|
|
| logger.debug('[BaseClient] tokenCountMap:', tokenCountMap); |
| logger.debug('[BaseClient]', { |
| promptTokens, |
| remainingContextTokens, |
| payloadSize: payload.length, |
| maxContextTokens: this.maxContextTokens, |
| }); |
|
|
| return { payload, tokenCountMap, promptTokens, messages: orderedWithInstructions }; |
| } |
|
|
| async sendMessage(message, opts = {}) { |
| const { user, head, isEdited, conversationId, responseMessageId, saveOptions, userMessage } = |
| await this.handleStartMethods(message, opts); |
|
|
| if (opts.progressCallback) { |
| opts.onProgress = opts.progressCallback.call(null, { |
| ...(opts.progressOptions ?? {}), |
| parentMessageId: userMessage.messageId, |
| messageId: responseMessageId, |
| }); |
| } |
|
|
| const { generation = '' } = opts; |
|
|
| |
| |
| |
| if (isEdited) { |
| let latestMessage = this.currentMessages[this.currentMessages.length - 1]; |
| if (!latestMessage) { |
| latestMessage = { |
| messageId: responseMessageId, |
| conversationId, |
| parentMessageId: userMessage.messageId, |
| isCreatedByUser: false, |
| model: this.modelOptions.model, |
| sender: this.sender, |
| text: generation, |
| }; |
| this.currentMessages.push(userMessage, latestMessage); |
| } else { |
| latestMessage.text = generation; |
| } |
| } else { |
| this.currentMessages.push(userMessage); |
| } |
|
|
| let { |
| prompt: payload, |
| tokenCountMap, |
| promptTokens, |
| } = await this.buildMessages( |
| this.currentMessages, |
| |
| |
| isEdited ? head : userMessage.messageId, |
| this.getBuildMessagesOptions(opts), |
| opts, |
| ); |
|
|
| if (tokenCountMap) { |
| logger.debug('[BaseClient] tokenCountMap', tokenCountMap); |
| if (tokenCountMap[userMessage.messageId]) { |
| userMessage.tokenCount = tokenCountMap[userMessage.messageId]; |
| logger.debug('[BaseClient] userMessage', userMessage); |
| } |
|
|
| this.handleTokenCountMap(tokenCountMap); |
| } |
|
|
| if (!isEdited && !this.skipSaveUserMessage) { |
| await this.saveMessageToDatabase(userMessage, saveOptions, user); |
| } |
|
|
| if ( |
| isEnabled(process.env.CHECK_BALANCE) && |
| supportsBalanceCheck[this.options.endpointType ?? this.options.endpoint] |
| ) { |
| await checkBalance({ |
| req: this.options.req, |
| res: this.options.res, |
| txData: { |
| user: this.user, |
| tokenType: 'prompt', |
| amount: promptTokens, |
| model: this.modelOptions.model, |
| endpoint: this.options.endpoint, |
| endpointTokenConfig: this.options.endpointTokenConfig, |
| }, |
| }); |
| } |
|
|
| const completion = await this.sendCompletion(payload, opts); |
| this.abortController.requestCompleted = true; |
|
|
| const responseMessage = { |
| messageId: responseMessageId, |
| conversationId, |
| parentMessageId: userMessage.messageId, |
| isCreatedByUser: false, |
| isEdited, |
| model: this.modelOptions.model, |
| sender: this.sender, |
| text: addSpaceIfNeeded(generation) + completion, |
| promptTokens, |
| iconURL: this.options.iconURL, |
| endpoint: this.options.endpoint, |
| ...(this.metadata ?? {}), |
| }; |
|
|
| if ( |
| tokenCountMap && |
| this.recordTokenUsage && |
| this.getTokenCountForResponse && |
| this.getTokenCount |
| ) { |
| responseMessage.tokenCount = this.getTokenCountForResponse(responseMessage); |
| const completionTokens = this.getTokenCount(completion); |
| await this.recordTokenUsage({ promptTokens, completionTokens }); |
| } |
| await this.saveMessageToDatabase(responseMessage, saveOptions, user); |
| delete responseMessage.tokenCount; |
| return responseMessage; |
| } |
|
|
| async getConversation(conversationId, user = null) { |
| return await getConvo(user, conversationId); |
| } |
|
|
| async loadHistory(conversationId, parentMessageId = null) { |
| logger.debug('[BaseClient] Loading history:', { conversationId, parentMessageId }); |
|
|
| const messages = (await getMessages({ conversationId })) ?? []; |
|
|
| if (messages.length === 0) { |
| return []; |
| } |
|
|
| let mapMethod = null; |
| if (this.getMessageMapMethod) { |
| mapMethod = this.getMessageMapMethod(); |
| } |
|
|
| let _messages = this.constructor.getMessagesForConversation({ |
| messages, |
| parentMessageId, |
| mapMethod, |
| }); |
|
|
| _messages = await this.addPreviousAttachments(_messages); |
|
|
| if (!this.shouldSummarize) { |
| return _messages; |
| } |
|
|
| |
| for (let i = _messages.length - 1; i >= 0; i--) { |
| if (_messages[i]?.summary) { |
| this.previous_summary = _messages[i]; |
| break; |
| } |
| } |
|
|
| if (this.previous_summary) { |
| const { messageId, summary, tokenCount, summaryTokenCount } = this.previous_summary; |
| logger.debug('[BaseClient] Previous summary:', { |
| messageId, |
| summary, |
| tokenCount, |
| summaryTokenCount, |
| }); |
| } |
|
|
| return _messages; |
| } |
|
|
| |
| |
| |
| |
| |
| |
| async saveMessageToDatabase(message, endpointOptions, user = null) { |
| await saveMessage({ |
| ...message, |
| endpoint: this.options.endpoint, |
| unfinished: false, |
| user, |
| }); |
|
|
| if (this.skipSaveConvo) { |
| return; |
| } |
| await saveConvo(user, { |
| conversationId: message.conversationId, |
| endpoint: this.options.endpoint, |
| endpointType: this.options.endpointType, |
| ...endpointOptions, |
| }); |
| } |
|
|
| async updateMessageInDatabase(message) { |
| await updateMessage(message); |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| static getMessagesForConversation({ |
| messages, |
| parentMessageId, |
| mapMethod = null, |
| summary = false, |
| }) { |
| if (!messages || messages.length === 0) { |
| return []; |
| } |
|
|
| const orderedMessages = []; |
| let currentMessageId = parentMessageId; |
| const visitedMessageIds = new Set(); |
|
|
| while (currentMessageId) { |
| if (visitedMessageIds.has(currentMessageId)) { |
| break; |
| } |
| const message = messages.find((msg) => { |
| const messageId = msg.messageId ?? msg.id; |
| return messageId === currentMessageId; |
| }); |
|
|
| visitedMessageIds.add(currentMessageId); |
|
|
| if (!message) { |
| break; |
| } |
|
|
| if (summary && message.summary) { |
| message.role = 'system'; |
| message.text = message.summary; |
| } |
|
|
| if (summary && message.summaryTokenCount) { |
| message.tokenCount = message.summaryTokenCount; |
| } |
|
|
| orderedMessages.push(message); |
|
|
| if (summary && message.summary) { |
| break; |
| } |
|
|
| currentMessageId = |
| message.parentMessageId === Constants.NO_PARENT ? null : message.parentMessageId; |
| } |
|
|
| orderedMessages.reverse(); |
|
|
| if (mapMethod) { |
| return orderedMessages.map(mapMethod); |
| } |
|
|
| return orderedMessages; |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| getTokenCountForMessage(message) { |
| |
| let tokensPerMessage = 3; |
| let tokensPerName = 1; |
|
|
| if (this.modelOptions.model === 'gpt-3.5-turbo-0301') { |
| tokensPerMessage = 4; |
| tokensPerName = -1; |
| } |
|
|
| const processValue = (value) => { |
| if (Array.isArray(value)) { |
| for (let item of value) { |
| if (!item || !item.type || item.type === 'image_url') { |
| continue; |
| } |
|
|
| const nestedValue = item[item.type]; |
|
|
| if (!nestedValue) { |
| continue; |
| } |
|
|
| processValue(nestedValue); |
| } |
| } else { |
| numTokens += this.getTokenCount(value); |
| } |
| }; |
|
|
| let numTokens = tokensPerMessage; |
| for (let [key, value] of Object.entries(message)) { |
| processValue(value); |
|
|
| if (key === 'name') { |
| numTokens += tokensPerName; |
| } |
| } |
| return numTokens; |
| } |
|
|
| async sendPayload(payload, opts = {}) { |
| if (opts && typeof opts === 'object') { |
| this.setOptions(opts); |
| } |
|
|
| return await this.sendCompletion(payload, opts); |
| } |
|
|
| |
| |
| |
| |
| |
| async addPreviousAttachments(_messages) { |
| if (!this.options.resendFiles) { |
| return _messages; |
| } |
|
|
| |
| |
| |
| |
| const processMessage = async (message) => { |
| if (!this.message_file_map) { |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |