diff --git a/api/app/clients/OpenAIClient.js b/api/app/clients/OpenAIClient.js index ce39311f3c..f0dbc366bb 100644 --- a/api/app/clients/OpenAIClient.js +++ b/api/app/clients/OpenAIClient.js @@ -847,7 +847,7 @@ ${convo} err?.message?.includes('abort') || (err instanceof OpenAI.APIError && err?.message?.includes('abort')) ) { - return ''; + return intermediateReply; } if ( err?.message?.includes( diff --git a/api/models/Message.js b/api/models/Message.js index 7cb9bdc377..7accf9285a 100644 --- a/api/models/Message.js +++ b/api/models/Message.js @@ -15,16 +15,16 @@ module.exports = { parentMessageId, sender, text, - isCreatedByUser = false, + isCreatedByUser, error, unfinished, files, - isEdited = false, - finish_reason = null, - tokenCount = null, - plugin = null, - plugins = null, - model = null, + isEdited, + finish_reason, + tokenCount, + plugin, + plugins, + model, }) { try { const validConvoId = idSchema.safeParse(conversationId); diff --git a/api/models/schema/messageSchema.js b/api/models/schema/messageSchema.js index bcc28cd23b..33d799544b 100644 --- a/api/models/schema/messageSchema.js +++ b/api/models/schema/messageSchema.js @@ -21,6 +21,7 @@ const messageSchema = mongoose.Schema( }, model: { type: String, + default: null, }, conversationSignature: { type: String, diff --git a/api/server/controllers/AskController.js b/api/server/controllers/AskController.js index ffaa10938c..78933feebc 100644 --- a/api/server/controllers/AskController.js +++ b/api/server/controllers/AskController.js @@ -118,16 +118,19 @@ const AskController = async (req, res, next, initializeClient, addTitle) => { delete userMessage.image_urls; } - sendMessage(res, { - title: await getConvoTitle(user, conversationId), - final: true, - conversation: await getConvo(user, conversationId), - requestMessage: userMessage, - responseMessage: response, - }); - res.end(); + if (!abortController.signal.aborted) { + sendMessage(res, { + title: await getConvoTitle(user, conversationId), + final: true, + conversation: await getConvo(user, conversationId), + requestMessage: userMessage, + responseMessage: response, + }); + res.end(); + + await saveMessage({ ...response, user }); + } - await saveMessage({ ...response, user }); await saveMessage(userMessage); if (addTitle && parentMessageId === '00000000-0000-0000-0000-000000000000' && newConvo) { diff --git a/api/server/controllers/EditController.js b/api/server/controllers/EditController.js index ecc1461260..72ee58026a 100644 --- a/api/server/controllers/EditController.js +++ b/api/server/controllers/EditController.js @@ -112,16 +112,18 @@ const EditController = async (req, res, next, initializeClient) => { response = { ...response, ...metadata }; } - await saveMessage({ ...response, user }); + if (!abortController.signal.aborted) { + sendMessage(res, { + title: await getConvoTitle(user, conversationId), + final: true, + conversation: await getConvo(user, conversationId), + requestMessage: userMessage, + responseMessage: response, + }); + res.end(); - sendMessage(res, { - title: await getConvoTitle(user, conversationId), - final: true, - conversation: await getConvo(user, conversationId), - requestMessage: userMessage, - responseMessage: response, - }); - res.end(); + await saveMessage({ ...response, user }); + } } catch (error) { const partialText = getPartialText(); handleAbortError(res, req, error, { diff --git a/api/server/utils/countTokens.js b/api/server/utils/countTokens.js index 9c8c98e76a..34c070aa8c 100644 --- a/api/server/utils/countTokens.js +++ b/api/server/utils/countTokens.js @@ -1,13 +1,12 @@ -const { load } = require('tiktoken/load'); const { Tiktoken } = require('tiktoken/lite'); -const registry = require('tiktoken/registry.json'); -const models = require('tiktoken/model_to_encoding.json'); +const p50k_base = require('tiktoken/encoders/p50k_base.json'); +const cl100k_base = require('tiktoken/encoders/cl100k_base.json'); const logger = require('~/config/winston'); const countTokens = async (text = '', modelName = 'gpt-3.5-turbo') => { let encoder = null; try { - const model = await load(registry[models[modelName]]); + const model = modelName.includes('text-davinci-003') ? p50k_base : cl100k_base; encoder = new Tiktoken(model.bpe_ranks, model.special_tokens, model.pat_str); const tokens = encoder.encode(text); encoder.free(); diff --git a/config/helpers.js b/config/helpers.js index a86d562eb3..2b634612d4 100644 --- a/config/helpers.js +++ b/config/helpers.js @@ -6,7 +6,8 @@ const fs = require('fs'); const path = require('path'); const readline = require('readline'); const { execSync } = require('child_process'); -const { connectDb } = require('@librechat/backend/lib/db'); +require('module-alias')({ base: path.resolve(__dirname, '..', 'api') }); +const connectDb = require('~/lib/db/connectDb'); const askQuestion = (query) => { const rl = readline.createInterface({