diff --git a/api/app/clients/AnthropicClient.js b/api/app/clients/AnthropicClient.js index 2373a321f5..16b21ea2e3 100644 --- a/api/app/clients/AnthropicClient.js +++ b/api/app/clients/AnthropicClient.js @@ -2,8 +2,9 @@ const Anthropic = require('@anthropic-ai/sdk'); const { HttpsProxyAgent } = require('https-proxy-agent'); const { encoding_for_model: encodingForModel, get_encoding: getEncoding } = require('tiktoken'); const { - getResponseSender, + Constants, EModelEndpoint, + getResponseSender, validateVisionModel, } = require('librechat-data-provider'); const { encodeAndFormat } = require('~/server/services/Files/images/encode'); @@ -16,6 +17,7 @@ const { } = require('./prompts'); const spendTokens = require('~/models/spendTokens'); const { getModelMaxTokens } = require('~/utils'); +const { sleep } = require('~/server/utils'); const BaseClient = require('./BaseClient'); const { logger } = require('~/config'); @@ -605,6 +607,7 @@ class AnthropicClient extends BaseClient { }; const maxRetries = 3; + const streamRate = this.options.streamRate ?? Constants.DEFAULT_STREAM_RATE; async function processResponse() { let attempts = 0; @@ -627,6 +630,8 @@ class AnthropicClient extends BaseClient { } else if (completion.completion) { handleChunk(completion.completion); } + + await sleep(streamRate); } // Successful processing, exit loop diff --git a/api/app/clients/BaseClient.js b/api/app/clients/BaseClient.js index 31932a1887..b09a6a5d95 100644 --- a/api/app/clients/BaseClient.js +++ b/api/app/clients/BaseClient.js @@ -1,10 +1,11 @@ const crypto = require('crypto'); const fetch = require('node-fetch'); -const { supportsBalanceCheck, Constants } = require('librechat-data-provider'); +const { supportsBalanceCheck, Constants, CacheKeys, Time } = require('librechat-data-provider'); const { getMessages, saveMessage, updateMessage, saveConvo } = require('~/models'); const { addSpaceIfNeeded, isEnabled } = require('~/server/utils'); const checkBalance = require('~/models/checkBalance'); const { getFiles } = require('~/models/File'); +const { getLogStores } = require('~/cache'); const TextStream = require('./TextStream'); const { logger } = require('~/config'); @@ -540,6 +541,15 @@ class BaseClient { await this.recordTokenUsage({ promptTokens, completionTokens }); } this.responsePromise = this.saveMessageToDatabase(responseMessage, saveOptions, user); + const messageCache = getLogStores(CacheKeys.MESSAGES); + messageCache.set( + responseMessageId, + { + text: responseMessage.text, + complete: true, + }, + Time.FIVE_MINUTES, + ); delete responseMessage.tokenCount; return responseMessage; } diff --git a/api/app/clients/GoogleClient.js b/api/app/clients/GoogleClient.js index a01df71841..e115ab1db8 100644 --- a/api/app/clients/GoogleClient.js +++ b/api/app/clients/GoogleClient.js @@ -13,10 +13,12 @@ const { endpointSettings, EModelEndpoint, VisionModes, + Constants, AuthKeys, } = require('librechat-data-provider'); const { encodeAndFormat } = require('~/server/services/Files/images'); const { getModelMaxTokens } = require('~/utils'); +const { sleep } = require('~/server/utils'); const { logger } = require('~/config'); const { formatMessage, @@ -620,8 +622,9 @@ class GoogleClient extends BaseClient { } async getCompletion(_payload, options = {}) { - const { onProgress, abortController } = options; const { parameters, instances } = _payload; + const { onProgress, abortController } = options; + const streamRate = this.options.streamRate ?? Constants.DEFAULT_STREAM_RATE; const { messages: _messages, context, examples: _examples } = instances?.[0] ?? {}; let examples; @@ -701,6 +704,7 @@ class GoogleClient extends BaseClient { delay, }); reply += chunkText; + await sleep(streamRate); } return reply; } @@ -712,10 +716,17 @@ class GoogleClient extends BaseClient { safetySettings: safetySettings, }); - let delay = this.isGenerativeModel ? 12 : 8; - if (modelName.includes('flash')) { - delay = 5; + let delay = this.options.streamRate || 8; + + if (!this.options.streamRate) { + if (this.isGenerativeModel) { + delay = 12; + } + if (modelName.includes('flash')) { + delay = 5; + } } + for await (const chunk of stream) { const chunkText = chunk?.content ?? chunk; await this.generateTextStream(chunkText, onProgress, { diff --git a/api/app/clients/OllamaClient.js b/api/app/clients/OllamaClient.js index 57bc8754fb..c88ef72d58 100644 --- a/api/app/clients/OllamaClient.js +++ b/api/app/clients/OllamaClient.js @@ -1,7 +1,9 @@ const { z } = require('zod'); const axios = require('axios'); const { Ollama } = require('ollama'); +const { Constants } = require('librechat-data-provider'); const { deriveBaseURL } = require('~/utils'); +const { sleep } = require('~/server/utils'); const { logger } = require('~/config'); const ollamaPayloadSchema = z.object({ @@ -40,6 +42,7 @@ const getValidBase64 = (imageUrl) => { class OllamaClient { constructor(options = {}) { const host = deriveBaseURL(options.baseURL ?? 'http://localhost:11434'); + this.streamRate = options.streamRate ?? Constants.DEFAULT_STREAM_RATE; /** @type {Ollama} */ this.client = new Ollama({ host }); } @@ -136,6 +139,8 @@ class OllamaClient { stream.controller.abort(); break; } + + await sleep(this.streamRate); } } // TODO: regular completion diff --git a/api/app/clients/OpenAIClient.js b/api/app/clients/OpenAIClient.js index 7520cbb897..ccc5165fc7 100644 --- a/api/app/clients/OpenAIClient.js +++ b/api/app/clients/OpenAIClient.js @@ -1182,8 +1182,10 @@ ${convo} }); } + const streamRate = this.options.streamRate ?? Constants.DEFAULT_STREAM_RATE; + if (this.message_file_map && this.isOllama) { - const ollamaClient = new OllamaClient({ baseURL }); + const ollamaClient = new OllamaClient({ baseURL, streamRate }); return await ollamaClient.chatCompletion({ payload: modelOptions, onProgress, @@ -1221,8 +1223,6 @@ ${convo} } }); - const azureDelay = this.modelOptions.model?.includes('gpt-4') ? 30 : 17; - for await (const chunk of stream) { const token = chunk.choices[0]?.delta?.content || ''; intermediateReply += token; @@ -1232,9 +1232,7 @@ ${convo} break; } - if (this.azure) { - await sleep(azureDelay); - } + await sleep(streamRate); } if (!UnexpectedRoleError) { diff --git a/api/app/clients/PluginsClient.js b/api/app/clients/PluginsClient.js index 2ce0ece4e7..a23fb019ba 100644 --- a/api/app/clients/PluginsClient.js +++ b/api/app/clients/PluginsClient.js @@ -1,5 +1,6 @@ const OpenAIClient = require('./OpenAIClient'); const { CallbackManager } = require('langchain/callbacks'); +const { CacheKeys, Time } = require('librechat-data-provider'); const { BufferMemory, ChatMessageHistory } = require('langchain/memory'); const { initializeCustomAgent, initializeFunctionsAgent } = require('./agents'); const { addImages, buildErrorInput, buildPromptPrefix } = require('./output_parsers'); @@ -11,6 +12,7 @@ const { SelfReflectionTool } = require('./tools'); const { isEnabled } = require('~/server/utils'); const { extractBaseURL } = require('~/utils'); const { loadTools } = require('./tools/util'); +const { getLogStores } = require('~/cache'); const { logger } = require('~/config'); class PluginsClient extends OpenAIClient { @@ -220,6 +222,13 @@ class PluginsClient extends OpenAIClient { } } + /** + * + * @param {TMessage} responseMessage + * @param {Partial} saveOptions + * @param {string} user + * @returns + */ async handleResponseMessage(responseMessage, saveOptions, user) { const { output, errorMessage, ...result } = this.result; logger.debug('[PluginsClient][handleResponseMessage] Output:', { @@ -239,6 +248,15 @@ class PluginsClient extends OpenAIClient { } this.responsePromise = this.saveMessageToDatabase(responseMessage, saveOptions, user); + const messageCache = getLogStores(CacheKeys.MESSAGES); + messageCache.set( + responseMessage.messageId, + { + text: responseMessage.text, + complete: true, + }, + Time.FIVE_MINUTES, + ); delete responseMessage.tokenCount; return { ...responseMessage, ...result }; } diff --git a/api/cache/getLogStores.js b/api/cache/getLogStores.js index 9a7282e25a..2b33751a04 100644 --- a/api/cache/getLogStores.js +++ b/api/cache/getLogStores.js @@ -1,13 +1,11 @@ const Keyv = require('keyv'); -const { CacheKeys, ViolationTypes } = require('librechat-data-provider'); +const { CacheKeys, ViolationTypes, Time } = require('librechat-data-provider'); const { logFile, violationFile } = require('./keyvFiles'); const { math, isEnabled } = require('~/server/utils'); const keyvRedis = require('./keyvRedis'); const keyvMongo = require('./keyvMongo'); const { BAN_DURATION, USE_REDIS } = process.env ?? {}; -const THIRTY_MINUTES = 1800000; -const TEN_MINUTES = 600000; const duration = math(BAN_DURATION, 7200000); @@ -29,17 +27,21 @@ const roles = isEnabled(USE_REDIS) ? new Keyv({ store: keyvRedis }) : new Keyv({ namespace: CacheKeys.ROLES }); -const audioRuns = isEnabled(USE_REDIS) // ttl: 30 minutes - ? new Keyv({ store: keyvRedis, ttl: TEN_MINUTES }) - : new Keyv({ namespace: CacheKeys.AUDIO_RUNS, ttl: TEN_MINUTES }); +const audioRuns = isEnabled(USE_REDIS) + ? new Keyv({ store: keyvRedis, ttl: Time.TEN_MINUTES }) + : new Keyv({ namespace: CacheKeys.AUDIO_RUNS, ttl: Time.TEN_MINUTES }); + +const messages = isEnabled(USE_REDIS) + ? new Keyv({ store: keyvRedis, ttl: Time.FIVE_MINUTES }) + : new Keyv({ namespace: CacheKeys.MESSAGES, ttl: Time.FIVE_MINUTES }); const tokenConfig = isEnabled(USE_REDIS) // ttl: 30 minutes - ? new Keyv({ store: keyvRedis, ttl: THIRTY_MINUTES }) - : new Keyv({ namespace: CacheKeys.TOKEN_CONFIG, ttl: THIRTY_MINUTES }); + ? new Keyv({ store: keyvRedis, ttl: Time.THIRTY_MINUTES }) + : new Keyv({ namespace: CacheKeys.TOKEN_CONFIG, ttl: Time.THIRTY_MINUTES }); const genTitle = isEnabled(USE_REDIS) // ttl: 2 minutes - ? new Keyv({ store: keyvRedis, ttl: 120000 }) - : new Keyv({ namespace: CacheKeys.GEN_TITLE, ttl: 120000 }); + ? new Keyv({ store: keyvRedis, ttl: Time.TWO_MINUTES }) + : new Keyv({ namespace: CacheKeys.GEN_TITLE, ttl: Time.TWO_MINUTES }); const modelQueries = isEnabled(process.env.USE_REDIS) ? new Keyv({ store: keyvRedis }) @@ -47,7 +49,7 @@ const modelQueries = isEnabled(process.env.USE_REDIS) const abortKeys = isEnabled(USE_REDIS) ? new Keyv({ store: keyvRedis }) - : new Keyv({ namespace: CacheKeys.ABORT_KEYS, ttl: 600000 }); + : new Keyv({ namespace: CacheKeys.ABORT_KEYS, ttl: Time.TEN_MINUTES }); const namespaces = { [CacheKeys.ROLES]: roles, @@ -81,6 +83,7 @@ const namespaces = { [CacheKeys.GEN_TITLE]: genTitle, [CacheKeys.MODEL_QUERIES]: modelQueries, [CacheKeys.AUDIO_RUNS]: audioRuns, + [CacheKeys.MESSAGES]: messages, }; /** diff --git a/api/server/controllers/AskController.js b/api/server/controllers/AskController.js index 5b211d66b9..674c22a834 100644 --- a/api/server/controllers/AskController.js +++ b/api/server/controllers/AskController.js @@ -1,7 +1,8 @@ const throttle = require('lodash/throttle'); -const { getResponseSender, Constants, EModelEndpoint } = require('librechat-data-provider'); +const { getResponseSender, Constants, CacheKeys, Time } = require('librechat-data-provider'); const { createAbortController, handleAbortError } = require('~/server/middleware'); const { sendMessage, createOnProgress } = require('~/server/utils'); +const { getLogStores } = require('~/cache'); const { saveMessage } = require('~/models'); const { logger } = require('~/config'); @@ -51,11 +52,13 @@ const AskController = async (req, res, next, initializeClient, addTitle) => { try { const { client } = await initializeClient({ req, res, endpointOption }); - const unfinished = endpointOption.endpoint === EModelEndpoint.google ? false : true; + const messageCache = getLogStores(CacheKeys.MESSAGES); const { onProgress: progressCallback, getPartialText } = createOnProgress({ onProgress: throttle( ({ text: partialText }) => { - saveMessage(req, { + /* + const unfinished = endpointOption.endpoint === EModelEndpoint.google ? false : true; + messageCache.set(responseMessageId, { messageId: responseMessageId, sender, conversationId, @@ -65,7 +68,10 @@ const AskController = async (req, res, next, initializeClient, addTitle) => { unfinished, error: false, user, - }); + }, Time.FIVE_MINUTES); + */ + + messageCache.set(responseMessageId, partialText, Time.FIVE_MINUTES); }, 3000, { trailing: false }, diff --git a/api/server/controllers/EditController.js b/api/server/controllers/EditController.js index 66114b0f6c..e8be7f3e7a 100644 --- a/api/server/controllers/EditController.js +++ b/api/server/controllers/EditController.js @@ -1,7 +1,8 @@ const throttle = require('lodash/throttle'); -const { getResponseSender, EModelEndpoint } = require('librechat-data-provider'); +const { getResponseSender, CacheKeys, Time } = require('librechat-data-provider'); const { createAbortController, handleAbortError } = require('~/server/middleware'); const { sendMessage, createOnProgress } = require('~/server/utils'); +const { getLogStores } = require('~/cache'); const { saveMessage } = require('~/models'); const { logger } = require('~/config'); @@ -51,12 +52,14 @@ const EditController = async (req, res, next, initializeClient) => { } }; - const unfinished = endpointOption.endpoint === EModelEndpoint.google ? false : true; + const messageCache = getLogStores(CacheKeys.MESSAGES); const { onProgress: progressCallback, getPartialText } = createOnProgress({ generation, onProgress: throttle( ({ text: partialText }) => { - saveMessage(req, { + /* + const unfinished = endpointOption.endpoint === EModelEndpoint.google ? false : true; + { messageId: responseMessageId, sender, conversationId, @@ -67,7 +70,8 @@ const EditController = async (req, res, next, initializeClient) => { isEdited: true, error: false, user, - }); + } */ + messageCache.set(responseMessageId, partialText, Time.FIVE_MINUTES); }, 3000, { trailing: false }, diff --git a/api/server/controllers/assistants/chatV2.js b/api/server/controllers/assistants/chatV2.js index 4e10201f73..67e106ca0d 100644 --- a/api/server/controllers/assistants/chatV2.js +++ b/api/server/controllers/assistants/chatV2.js @@ -1,12 +1,12 @@ const { v4 } = require('uuid'); const { + Time, Constants, RunStatus, CacheKeys, ContentTypes, ToolCallTypes, EModelEndpoint, - ViolationTypes, retrievalMimeTypes, AssistantStreamEvents, } = require('librechat-data-provider'); @@ -14,12 +14,12 @@ const { initThread, recordUsage, saveUserMessage, - checkMessageGaps, addThreadMetadata, saveAssistantMessage, } = require('~/server/services/Threads'); -const { sendResponse, sendMessage, sleep, isEnabled, countTokens } = require('~/server/utils'); const { runAssistant, createOnTextProgress } = require('~/server/services/AssistantService'); +const { sendMessage, sleep, isEnabled, countTokens } = require('~/server/utils'); +const { createErrorHandler } = require('~/server/controllers/assistants/errors'); const validateAuthor = require('~/server/middleware/assistants/validateAuthor'); const { createRun, StreamRunManager } = require('~/server/services/Runs'); const { addTitle } = require('~/server/services/Endpoints/assistants'); @@ -44,7 +44,7 @@ const ten_minutes = 1000 * 60 * 10; const chatV2 = async (req, res) => { logger.debug('[/assistants/chat/] req.body', req.body); - /** @type {{ files: MongoFile[]}} */ + /** @type {{files: MongoFile[]}} */ const { text, model, @@ -90,140 +90,20 @@ const chatV2 = async (req, res) => { /** @type {Run | undefined} - The completed run, undefined if incomplete */ let completedRun; - const handleError = async (error) => { - const defaultErrorMessage = - 'The Assistant run failed to initialize. Try sending a message in a new conversation.'; - const messageData = { - thread_id, - assistant_id, - conversationId, - parentMessageId, - sender: 'System', - user: req.user.id, - shouldSaveMessage: false, - messageId: responseMessageId, - endpoint, - }; + const getContext = () => ({ + openai, + run_id, + endpoint, + cacheKey, + thread_id, + completedRun, + assistant_id, + conversationId, + parentMessageId, + responseMessageId, + }); - if (error.message === 'Run cancelled') { - return res.end(); - } else if (error.message === 'Request closed' && completedRun) { - return; - } else if (error.message === 'Request closed') { - logger.debug('[/assistants/chat/] Request aborted on close'); - } else if (/Files.*are invalid/.test(error.message)) { - const errorMessage = `Files are invalid, or may not have uploaded yet.${ - endpoint === EModelEndpoint.azureAssistants - ? ' If using Azure OpenAI, files are only available in the region of the assistant\'s model at the time of upload.' - : '' - }`; - return sendResponse(req, res, messageData, errorMessage); - } else if (error?.message?.includes('string too long')) { - return sendResponse( - req, - res, - messageData, - 'Message too long. The Assistants API has a limit of 32,768 characters per message. Please shorten it and try again.', - ); - } else if (error?.message?.includes(ViolationTypes.TOKEN_BALANCE)) { - return sendResponse(req, res, messageData, error.message); - } else { - logger.error('[/assistants/chat/]', error); - } - - if (!openai || !thread_id || !run_id) { - return sendResponse(req, res, messageData, defaultErrorMessage); - } - - await sleep(2000); - - try { - const status = await cache.get(cacheKey); - if (status === 'cancelled') { - logger.debug('[/assistants/chat/] Run already cancelled'); - return res.end(); - } - await cache.delete(cacheKey); - const cancelledRun = await openai.beta.threads.runs.cancel(thread_id, run_id); - logger.debug('[/assistants/chat/] Cancelled run:', cancelledRun); - } catch (error) { - logger.error('[/assistants/chat/] Error cancelling run', error); - } - - await sleep(2000); - - let run; - try { - run = await openai.beta.threads.runs.retrieve(thread_id, run_id); - await recordUsage({ - ...run.usage, - model: run.model, - user: req.user.id, - conversationId, - }); - } catch (error) { - logger.error('[/assistants/chat/] Error fetching or processing run', error); - } - - let finalEvent; - try { - const runMessages = await checkMessageGaps({ - openai, - run_id, - endpoint, - thread_id, - conversationId, - latestMessageId: responseMessageId, - }); - - const errorContentPart = { - text: { - value: - error?.message ?? 'There was an error processing your request. Please try again later.', - }, - type: ContentTypes.ERROR, - }; - - if (!Array.isArray(runMessages[runMessages.length - 1]?.content)) { - runMessages[runMessages.length - 1].content = [errorContentPart]; - } else { - const contentParts = runMessages[runMessages.length - 1].content; - for (let i = 0; i < contentParts.length; i++) { - const currentPart = contentParts[i]; - /** @type {CodeToolCall | RetrievalToolCall | FunctionToolCall | undefined} */ - const toolCall = currentPart?.[ContentTypes.TOOL_CALL]; - if ( - toolCall && - toolCall?.function && - !(toolCall?.function?.output || toolCall?.function?.output?.length) - ) { - contentParts[i] = { - ...currentPart, - [ContentTypes.TOOL_CALL]: { - ...toolCall, - function: { - ...toolCall.function, - output: 'error processing tool', - }, - }, - }; - } - } - runMessages[runMessages.length - 1].content.push(errorContentPart); - } - - finalEvent = { - final: true, - conversation: await getConvo(req.user.id, conversationId), - runMessages, - }; - } catch (error) { - logger.error('[/assistants/chat/] Error finalizing error process', error); - return sendResponse(req, res, messageData, 'The Assistant run failed'); - } - - return sendResponse(req, res, finalEvent); - }; + const handleError = createErrorHandler({ req, res, getContext }); try { res.on('close', async () => { @@ -490,6 +370,11 @@ const chatV2 = async (req, res) => { }, }; + /** @type {undefined | TAssistantEndpoint} */ + const config = req.app.locals[endpoint] ?? {}; + /** @type {undefined | TBaseEndpoint} */ + const allConfig = req.app.locals.all; + const streamRunManager = new StreamRunManager({ req, res, @@ -499,6 +384,7 @@ const chatV2 = async (req, res) => { attachedFileIds, parentMessageId: userMessageId, responseMessage: openai.responseMessage, + streamRate: allConfig?.streamRate ?? config.streamRate, // streamOptions: { // }, @@ -511,6 +397,16 @@ const chatV2 = async (req, res) => { response = streamRunManager; response.text = streamRunManager.intermediateText; + + const messageCache = getLogStores(CacheKeys.MESSAGES); + messageCache.set( + responseMessageId, + { + complete: true, + text: response.text, + }, + Time.FIVE_MINUTES, + ); }; await processRun(); diff --git a/api/server/controllers/assistants/errors.js b/api/server/controllers/assistants/errors.js new file mode 100644 index 0000000000..a4b880bf04 --- /dev/null +++ b/api/server/controllers/assistants/errors.js @@ -0,0 +1,193 @@ +// errorHandler.js +const { sendResponse } = require('~/server/utils'); +const { logger } = require('~/config'); +const getLogStores = require('~/cache/getLogStores'); +const { CacheKeys, ViolationTypes, ContentTypes } = require('librechat-data-provider'); +const { getConvo } = require('~/models/Conversation'); +const { recordUsage, checkMessageGaps } = require('~/server/services/Threads'); + +/** + * @typedef {Object} ErrorHandlerContext + * @property {OpenAIClient} openai - The OpenAI client + * @property {string} thread_id - The thread ID + * @property {string} run_id - The run ID + * @property {boolean} completedRun - Whether the run has completed + * @property {string} assistant_id - The assistant ID + * @property {string} conversationId - The conversation ID + * @property {string} parentMessageId - The parent message ID + * @property {string} responseMessageId - The response message ID + * @property {string} endpoint - The endpoint being used + * @property {string} cacheKey - The cache key for the current request + */ + +/** + * @typedef {Object} ErrorHandlerDependencies + * @property {Express.Request} req - The Express request object + * @property {Express.Response} res - The Express response object + * @property {() => ErrorHandlerContext} getContext - Function to get the current context + * @property {string} [originPath] - The origin path for the error handler + */ + +/** + * Creates an error handler function with the given dependencies + * @param {ErrorHandlerDependencies} dependencies - The dependencies for the error handler + * @returns {(error: Error) => Promise} The error handler function + */ +const createErrorHandler = ({ req, res, getContext, originPath = '/assistants/chat/' }) => { + const cache = getLogStores(CacheKeys.ABORT_KEYS); + + /** + * Handles errors that occur during the chat process + * @param {Error} error - The error that occurred + * @returns {Promise} + */ + return async (error) => { + const { + openai, + run_id, + endpoint, + cacheKey, + thread_id, + completedRun, + assistant_id, + conversationId, + parentMessageId, + responseMessageId, + } = getContext(); + + const defaultErrorMessage = + 'The Assistant run failed to initialize. Try sending a message in a new conversation.'; + const messageData = { + thread_id, + assistant_id, + conversationId, + parentMessageId, + sender: 'System', + user: req.user.id, + shouldSaveMessage: false, + messageId: responseMessageId, + endpoint, + }; + + if (error.message === 'Run cancelled') { + return res.end(); + } else if (error.message === 'Request closed' && completedRun) { + return; + } else if (error.message === 'Request closed') { + logger.debug(`[${originPath}] Request aborted on close`); + } else if (/Files.*are invalid/.test(error.message)) { + const errorMessage = `Files are invalid, or may not have uploaded yet.${ + endpoint === 'azureAssistants' + ? ' If using Azure OpenAI, files are only available in the region of the assistant\'s model at the time of upload.' + : '' + }`; + return sendResponse(req, res, messageData, errorMessage); + } else if (error?.message?.includes('string too long')) { + return sendResponse( + req, + res, + messageData, + 'Message too long. The Assistants API has a limit of 32,768 characters per message. Please shorten it and try again.', + ); + } else if (error?.message?.includes(ViolationTypes.TOKEN_BALANCE)) { + return sendResponse(req, res, messageData, error.message); + } else { + logger.error(`[${originPath}]`, error); + } + + if (!openai || !thread_id || !run_id) { + return sendResponse(req, res, messageData, defaultErrorMessage); + } + + await new Promise((resolve) => setTimeout(resolve, 2000)); + + try { + const status = await cache.get(cacheKey); + if (status === 'cancelled') { + logger.debug(`[${originPath}] Run already cancelled`); + return res.end(); + } + await cache.delete(cacheKey); + const cancelledRun = await openai.beta.threads.runs.cancel(thread_id, run_id); + logger.debug(`[${originPath}] Cancelled run:`, cancelledRun); + } catch (error) { + logger.error(`[${originPath}] Error cancelling run`, error); + } + + await new Promise((resolve) => setTimeout(resolve, 2000)); + + let run; + try { + run = await openai.beta.threads.runs.retrieve(thread_id, run_id); + await recordUsage({ + ...run.usage, + model: run.model, + user: req.user.id, + conversationId, + }); + } catch (error) { + logger.error(`[${originPath}] Error fetching or processing run`, error); + } + + let finalEvent; + try { + const runMessages = await checkMessageGaps({ + openai, + run_id, + endpoint, + thread_id, + conversationId, + latestMessageId: responseMessageId, + }); + + const errorContentPart = { + text: { + value: + error?.message ?? 'There was an error processing your request. Please try again later.', + }, + type: ContentTypes.ERROR, + }; + + if (!Array.isArray(runMessages[runMessages.length - 1]?.content)) { + runMessages[runMessages.length - 1].content = [errorContentPart]; + } else { + const contentParts = runMessages[runMessages.length - 1].content; + for (let i = 0; i < contentParts.length; i++) { + const currentPart = contentParts[i]; + /** @type {CodeToolCall | RetrievalToolCall | FunctionToolCall | undefined} */ + const toolCall = currentPart?.[ContentTypes.TOOL_CALL]; + if ( + toolCall && + toolCall?.function && + !(toolCall?.function?.output || toolCall?.function?.output?.length) + ) { + contentParts[i] = { + ...currentPart, + [ContentTypes.TOOL_CALL]: { + ...toolCall, + function: { + ...toolCall.function, + output: 'error processing tool', + }, + }, + }; + } + } + runMessages[runMessages.length - 1].content.push(errorContentPart); + } + + finalEvent = { + final: true, + conversation: await getConvo(req.user.id, conversationId), + runMessages, + }; + } catch (error) { + logger.error(`[${originPath}] Error finalizing error process`, error); + return sendResponse(req, res, messageData, 'The Assistant run failed'); + } + + return sendResponse(req, res, finalEvent); + }; +}; + +module.exports = { createErrorHandler }; diff --git a/api/server/middleware/abortMiddleware.js b/api/server/middleware/abortMiddleware.js index 4ee5684a2f..a8ef269c9f 100644 --- a/api/server/middleware/abortMiddleware.js +++ b/api/server/middleware/abortMiddleware.js @@ -30,7 +30,10 @@ async function abortMessage(req, res) { return res.status(204).send({ message: 'Request not found' }); } const finalEvent = await abortController.abortCompletion(); - logger.info('[abortMessage] Aborted request', { abortKey }); + logger.debug( + `[abortMessage] ID: ${req.user.id} | ${req.user.email} | Aborted request: ` + + JSON.stringify({ abortKey }), + ); abortControllers.delete(abortKey); if (res.headersSent && finalEvent) { diff --git a/api/server/routes/ask/gptPlugins.js b/api/server/routes/ask/gptPlugins.js index 299bb199f2..602ff25086 100644 --- a/api/server/routes/ask/gptPlugins.js +++ b/api/server/routes/ask/gptPlugins.js @@ -1,10 +1,11 @@ const express = require('express'); const throttle = require('lodash/throttle'); -const { getResponseSender, Constants } = require('librechat-data-provider'); +const { getResponseSender, Constants, CacheKeys, Time } = require('librechat-data-provider'); const { initializeClient } = require('~/server/services/Endpoints/gptPlugins'); const { sendMessage, createOnProgress } = require('~/server/utils'); const { addTitle } = require('~/server/services/Endpoints/openAI'); const { saveMessage } = require('~/models'); +const { getLogStores } = require('~/cache'); const { handleAbort, createAbortController, @@ -71,7 +72,8 @@ router.post( } }; - const throttledSaveMessage = throttle(saveMessage, 3000, { trailing: false }); + const messageCache = getLogStores(CacheKeys.MESSAGES); + const throttledSetMessage = throttle(messageCache.set, 3000, { trailing: false }); let streaming = null; let timer = null; @@ -85,7 +87,8 @@ router.post( clearTimeout(timer); } - throttledSaveMessage(req, { + /* + { messageId: responseMessageId, sender, conversationId, @@ -96,7 +99,9 @@ router.post( error: false, plugins, user, - }); + } + */ + throttledSetMessage(responseMessageId, partialText, Time.FIVE_MINUTES); streaming = new Promise((resolve) => { timer = setTimeout(() => { diff --git a/api/server/routes/edit/gptPlugins.js b/api/server/routes/edit/gptPlugins.js index 0e4a77567b..926c8e4f5f 100644 --- a/api/server/routes/edit/gptPlugins.js +++ b/api/server/routes/edit/gptPlugins.js @@ -1,19 +1,20 @@ const express = require('express'); const throttle = require('lodash/throttle'); -const { getResponseSender } = require('librechat-data-provider'); +const { getResponseSender, CacheKeys, Time } = require('librechat-data-provider'); const { - handleAbort, - createAbortController, - handleAbortError, setHeaders, + handleAbort, + moderateText, validateModel, + handleAbortError, validateEndpoint, buildEndpointOption, - moderateText, + createAbortController, } = require('~/server/middleware'); const { sendMessage, createOnProgress, formatSteps, formatAction } = require('~/server/utils'); const { initializeClient } = require('~/server/services/Endpoints/gptPlugins'); const { saveMessage } = require('~/models'); +const { getLogStores } = require('~/cache'); const { validateTools } = require('~/app'); const { logger } = require('~/config'); @@ -79,7 +80,8 @@ router.post( } }; - const throttledSaveMessage = throttle(saveMessage, 3000, { trailing: false }); + const messageCache = getLogStores(CacheKeys.MESSAGES); + const throttledSetMessage = throttle(messageCache.set, 3000, { trailing: false }); const { onProgress: progressCallback, sendIntermediateMessage, @@ -91,7 +93,8 @@ router.post( plugin.loading = false; } - throttledSaveMessage(req, { + /* + { messageId: responseMessageId, sender, conversationId, @@ -102,7 +105,9 @@ router.post( isEdited: true, error: false, user, - }); + } + */ + throttledSetMessage(responseMessageId, partialText, Time.FIVE_MINUTES); }, }); diff --git a/api/server/services/AppService.js b/api/server/services/AppService.js index e416d5f6e7..d776aa63b7 100644 --- a/api/server/services/AppService.js +++ b/api/server/services/AppService.js @@ -67,17 +67,18 @@ const AppService = async (app) => { handleRateLimits(config?.rateLimits); const endpointLocals = {}; + const endpoints = config?.endpoints; - if (config?.endpoints?.[EModelEndpoint.azureOpenAI]) { + if (endpoints?.[EModelEndpoint.azureOpenAI]) { endpointLocals[EModelEndpoint.azureOpenAI] = azureConfigSetup(config); checkAzureVariables(); } - if (config?.endpoints?.[EModelEndpoint.azureOpenAI]?.assistants) { + if (endpoints?.[EModelEndpoint.azureOpenAI]?.assistants) { endpointLocals[EModelEndpoint.azureAssistants] = azureAssistantsDefaults(); } - if (config?.endpoints?.[EModelEndpoint.azureAssistants]) { + if (endpoints?.[EModelEndpoint.azureAssistants]) { endpointLocals[EModelEndpoint.azureAssistants] = assistantsConfigSetup( config, EModelEndpoint.azureAssistants, @@ -85,7 +86,7 @@ const AppService = async (app) => { ); } - if (config?.endpoints?.[EModelEndpoint.assistants]) { + if (endpoints?.[EModelEndpoint.assistants]) { endpointLocals[EModelEndpoint.assistants] = assistantsConfigSetup( config, EModelEndpoint.assistants, @@ -93,6 +94,19 @@ const AppService = async (app) => { ); } + if (endpoints?.[EModelEndpoint.openAI]) { + endpointLocals[EModelEndpoint.openAI] = endpoints[EModelEndpoint.openAI]; + } + if (endpoints?.[EModelEndpoint.google]) { + endpointLocals[EModelEndpoint.google] = endpoints[EModelEndpoint.google]; + } + if (endpoints?.[EModelEndpoint.anthropic]) { + endpointLocals[EModelEndpoint.anthropic] = endpoints[EModelEndpoint.anthropic]; + } + if (endpoints?.[EModelEndpoint.gptPlugins]) { + endpointLocals[EModelEndpoint.gptPlugins] = endpoints[EModelEndpoint.gptPlugins]; + } + app.locals = { ...defaultLocals, modelSpecs: config.modelSpecs, diff --git a/api/server/services/Endpoints/anthropic/initializeClient.js b/api/server/services/Endpoints/anthropic/initializeClient.js index c5d6696b3e..42b902b1fc 100644 --- a/api/server/services/Endpoints/anthropic/initializeClient.js +++ b/api/server/services/Endpoints/anthropic/initializeClient.js @@ -19,11 +19,27 @@ const initializeClient = async ({ req, res, endpointOption }) => { checkUserKeyExpiry(expiresAt, EModelEndpoint.anthropic); } + const clientOptions = {}; + + /** @type {undefined | TBaseEndpoint} */ + const anthropicConfig = req.app.locals[EModelEndpoint.anthropic]; + + if (anthropicConfig) { + clientOptions.streamRate = anthropicConfig.streamRate; + } + + /** @type {undefined | TBaseEndpoint} */ + const allConfig = req.app.locals.all; + if (allConfig) { + clientOptions.streamRate = allConfig.streamRate; + } + const client = new AnthropicClient(anthropicApiKey, { req, res, reverseProxyUrl: ANTHROPIC_REVERSE_PROXY ?? null, proxy: PROXY ?? null, + ...clientOptions, ...endpointOption, }); diff --git a/api/server/services/Endpoints/custom/initializeClient.js b/api/server/services/Endpoints/custom/initializeClient.js index 9fb6bfd1af..dbc7a769fb 100644 --- a/api/server/services/Endpoints/custom/initializeClient.js +++ b/api/server/services/Endpoints/custom/initializeClient.js @@ -114,9 +114,16 @@ const initializeClient = async ({ req, res, endpointOption }) => { contextStrategy: endpointConfig.summarize ? 'summarize' : null, directEndpoint: endpointConfig.directEndpoint, titleMessageRole: endpointConfig.titleMessageRole, + streamRate: endpointConfig.streamRate, endpointTokenConfig, }; + /** @type {undefined | TBaseEndpoint} */ + const allConfig = req.app.locals.all; + if (allConfig) { + customOptions.streamRate = allConfig.streamRate; + } + const clientOptions = { reverseProxyUrl: baseURL ?? null, proxy: PROXY ?? null, diff --git a/api/server/services/Endpoints/google/initializeClient.js b/api/server/services/Endpoints/google/initializeClient.js index d2099edcf5..788375e1e7 100644 --- a/api/server/services/Endpoints/google/initializeClient.js +++ b/api/server/services/Endpoints/google/initializeClient.js @@ -27,11 +27,27 @@ const initializeClient = async ({ req, res, endpointOption }) => { [AuthKeys.GOOGLE_API_KEY]: GOOGLE_KEY, }; + const clientOptions = {}; + + /** @type {undefined | TBaseEndpoint} */ + const allConfig = req.app.locals.all; + /** @type {undefined | TBaseEndpoint} */ + const googleConfig = req.app.locals[EModelEndpoint.google]; + + if (googleConfig) { + clientOptions.streamRate = googleConfig.streamRate; + } + + if (allConfig) { + clientOptions.streamRate = allConfig.streamRate; + } + const client = new GoogleClient(credentials, { req, res, reverseProxyUrl: GOOGLE_REVERSE_PROXY ?? null, proxy: PROXY ?? null, + ...clientOptions, ...endpointOption, }); diff --git a/api/server/services/Endpoints/google/initializeClient.spec.js b/api/server/services/Endpoints/google/initializeClient.spec.js index b46a535618..657dcbcaa8 100644 --- a/api/server/services/Endpoints/google/initializeClient.spec.js +++ b/api/server/services/Endpoints/google/initializeClient.spec.js @@ -8,6 +8,8 @@ jest.mock('~/server/services/UserService', () => ({ getUserKey: jest.fn().mockImplementation(() => ({})), })); +const app = { locals: {} }; + describe('google/initializeClient', () => { afterEach(() => { jest.clearAllMocks(); @@ -23,6 +25,7 @@ describe('google/initializeClient', () => { const req = { body: { key: expiresAt }, user: { id: '123' }, + app, }; const res = {}; const endpointOption = { modelOptions: { model: 'default-model' } }; @@ -44,6 +47,7 @@ describe('google/initializeClient', () => { const req = { body: { key: null }, user: { id: '123' }, + app, }; const res = {}; const endpointOption = { modelOptions: { model: 'default-model' } }; @@ -66,6 +70,7 @@ describe('google/initializeClient', () => { const req = { body: { key: expiresAt }, user: { id: '123' }, + app, }; const res = {}; const endpointOption = { modelOptions: { model: 'default-model' } }; diff --git a/api/server/services/Endpoints/gptPlugins/initializeClient.js b/api/server/services/Endpoints/gptPlugins/initializeClient.js index 312b23eb67..7e79d42564 100644 --- a/api/server/services/Endpoints/gptPlugins/initializeClient.js +++ b/api/server/services/Endpoints/gptPlugins/initializeClient.js @@ -86,6 +86,9 @@ const initializeClient = async ({ req, res, endpointOption }) => { clientOptions.titleModel = azureConfig.titleModel; clientOptions.titleMethod = azureConfig.titleMethod ?? 'completion'; + const azureRate = modelName.includes('gpt-4') ? 30 : 17; + clientOptions.streamRate = azureConfig.streamRate ?? azureRate; + const groupName = modelGroupMap[modelName].group; clientOptions.addParams = azureConfig.groupMap[groupName].addParams; clientOptions.dropParams = azureConfig.groupMap[groupName].dropParams; @@ -98,6 +101,19 @@ const initializeClient = async ({ req, res, endpointOption }) => { apiKey = clientOptions.azure.azureOpenAIApiKey; } + /** @type {undefined | TBaseEndpoint} */ + const pluginsConfig = req.app.locals[EModelEndpoint.gptPlugins]; + + if (!useAzure && pluginsConfig) { + clientOptions.streamRate = pluginsConfig.streamRate; + } + + /** @type {undefined | TBaseEndpoint} */ + const allConfig = req.app.locals.all; + if (allConfig) { + clientOptions.streamRate = allConfig.streamRate; + } + if (!apiKey) { throw new Error(`${endpoint} API key not provided. Please provide it again.`); } diff --git a/api/server/services/Endpoints/openAI/initializeClient.js b/api/server/services/Endpoints/openAI/initializeClient.js index 9a3a5c4189..1518cba028 100644 --- a/api/server/services/Endpoints/openAI/initializeClient.js +++ b/api/server/services/Endpoints/openAI/initializeClient.js @@ -76,6 +76,10 @@ const initializeClient = async ({ req, res, endpointOption }) => { clientOptions.titleConvo = azureConfig.titleConvo; clientOptions.titleModel = azureConfig.titleModel; + + const azureRate = modelName.includes('gpt-4') ? 30 : 17; + clientOptions.streamRate = azureConfig.streamRate ?? azureRate; + clientOptions.titleMethod = azureConfig.titleMethod ?? 'completion'; const groupName = modelGroupMap[modelName].group; @@ -90,6 +94,19 @@ const initializeClient = async ({ req, res, endpointOption }) => { apiKey = clientOptions.azure.azureOpenAIApiKey; } + /** @type {undefined | TBaseEndpoint} */ + const openAIConfig = req.app.locals[EModelEndpoint.openAI]; + + if (!isAzureOpenAI && openAIConfig) { + clientOptions.streamRate = openAIConfig.streamRate; + } + + /** @type {undefined | TBaseEndpoint} */ + const allConfig = req.app.locals.all; + if (allConfig) { + clientOptions.streamRate = allConfig.streamRate; + } + if (userProvidesKey & !apiKey) { throw new Error( JSON.stringify({ diff --git a/api/server/services/Files/Audio/streamAudio.js b/api/server/services/Files/Audio/streamAudio.js index 9f301e710b..eb8134e958 100644 --- a/api/server/services/Files/Audio/streamAudio.js +++ b/api/server/services/Files/Audio/streamAudio.js @@ -1,5 +1,6 @@ const WebSocket = require('ws'); -const { Message } = require('~/models/Message'); +const { CacheKeys } = require('librechat-data-provider'); +const { getLogStores } = require('~/cache'); /** * @param {string[]} voiceIds - Array of voice IDs @@ -104,6 +105,8 @@ function createChunkProcessor(messageId) { throw new Error('Message ID is required'); } + const messageCache = getLogStores(CacheKeys.MESSAGES); + /** * @returns {Promise<{ text: string, isFinished: boolean }[] | string>} */ @@ -116,14 +119,17 @@ function createChunkProcessor(messageId) { return `No change in message after ${MAX_NO_CHANGE_COUNT} attempts`; } - const message = await Message.findOne({ messageId }, 'text unfinished').lean(); + /** @type { string | { text: string; complete: boolean } } */ + const message = await messageCache.get(messageId); - if (!message || !message.text) { + if (!message) { notFoundCount++; return []; } - const { text, unfinished } = message; + const text = typeof message === 'string' ? message : message.text; + const complete = typeof message === 'string' ? false : message.complete; + if (text === processedText) { noChangeCount++; } @@ -131,7 +137,7 @@ function createChunkProcessor(messageId) { const remainingText = text.slice(processedText.length); const chunks = []; - if (unfinished && remainingText.length >= 20) { + if (!complete && remainingText.length >= 20) { const separatorIndex = findLastSeparatorIndex(remainingText); if (separatorIndex !== -1) { const chunkText = remainingText.slice(0, separatorIndex + 1); @@ -141,7 +147,7 @@ function createChunkProcessor(messageId) { chunks.push({ text: remainingText, isFinished: false }); processedText = text; } - } else if (!unfinished && remainingText.trim().length > 0) { + } else if (complete && remainingText.trim().length > 0) { chunks.push({ text: remainingText.trim(), isFinished: true }); processedText = text; } diff --git a/api/server/services/Files/Audio/streamAudio.spec.js b/api/server/services/Files/Audio/streamAudio.spec.js index 7aff8dbfa7..501e252c14 100644 --- a/api/server/services/Files/Audio/streamAudio.spec.js +++ b/api/server/services/Files/Audio/streamAudio.spec.js @@ -1,89 +1,145 @@ const { createChunkProcessor, splitTextIntoChunks } = require('./streamAudio'); -const { Message } = require('~/models/Message'); -jest.mock('~/models/Message', () => ({ - Message: { - findOne: jest.fn().mockReturnValue({ - lean: jest.fn(), - }), - }, -})); +jest.mock('keyv'); + +const globalCache = {}; +jest.mock('~/cache/getLogStores', () => { + return jest.fn().mockImplementation(() => { + const EventEmitter = require('events'); + const { CacheKeys } = require('librechat-data-provider'); + + class KeyvMongo extends EventEmitter { + constructor(url = 'mongodb://127.0.0.1:27017', options) { + super(); + this.ttlSupport = false; + url = url ?? {}; + if (typeof url === 'string') { + url = { url }; + } + if (url.uri) { + url = { url: url.uri, ...url }; + } + this.opts = { + url, + collection: 'keyv', + ...url, + ...options, + }; + } + + get = async (key) => { + return new Promise((resolve) => { + resolve(globalCache[key] || null); + }); + }; + + set = async (key, value) => { + return new Promise((resolve) => { + globalCache[key] = value; + resolve(true); + }); + }; + } + + return new KeyvMongo('', { + namespace: CacheKeys.MESSAGES, + ttl: 0, + }); + }); +}); describe('processChunks', () => { let processChunks; + let mockMessageCache; beforeEach(() => { + jest.resetAllMocks(); + mockMessageCache = { + get: jest.fn(), + }; + require('~/cache/getLogStores').mockReturnValue(mockMessageCache); processChunks = createChunkProcessor('message-id'); - Message.findOne.mockClear(); - Message.findOne().lean.mockClear(); }); it('should return an empty array when the message is not found', async () => { - Message.findOne().lean.mockResolvedValueOnce(null); + mockMessageCache.get.mockResolvedValueOnce(null); const result = await processChunks(); expect(result).toEqual([]); - expect(Message.findOne).toHaveBeenCalledWith({ messageId: 'message-id' }, 'text unfinished'); - expect(Message.findOne().lean).toHaveBeenCalled(); + expect(mockMessageCache.get).toHaveBeenCalledWith('message-id'); }); - it('should return an empty array when the message does not have a text property', async () => { - Message.findOne().lean.mockResolvedValueOnce({ unfinished: true }); + it('should return an error message after MAX_NOT_FOUND_COUNT attempts', async () => { + mockMessageCache.get.mockResolvedValue(null); + for (let i = 0; i < 6; i++) { + await processChunks(); + } const result = await processChunks(); - expect(result).toEqual([]); - expect(Message.findOne).toHaveBeenCalledWith({ messageId: 'message-id' }, 'text unfinished'); - expect(Message.findOne().lean).toHaveBeenCalled(); + expect(result).toBe('Message not found after 6 attempts'); }); - it('should return chunks for an unfinished message with separators', async () => { + it('should return chunks for an incomplete message with separators', async () => { const messageText = 'This is a long message. It should be split into chunks. Lol hi mom'; - Message.findOne().lean.mockResolvedValueOnce({ text: messageText, unfinished: true }); + mockMessageCache.get.mockResolvedValueOnce({ text: messageText, complete: false }); const result = await processChunks(); expect(result).toEqual([ { text: 'This is a long message. It should be split into chunks.', isFinished: false }, ]); - expect(Message.findOne).toHaveBeenCalledWith({ messageId: 'message-id' }, 'text unfinished'); - expect(Message.findOne().lean).toHaveBeenCalled(); }); - it('should return chunks for an unfinished message without separators', async () => { + it('should return chunks for an incomplete message without separators', async () => { const messageText = 'This is a long message without separators hello there my friend'; - Message.findOne().lean.mockResolvedValueOnce({ text: messageText, unfinished: true }); + mockMessageCache.get.mockResolvedValueOnce({ text: messageText, complete: false }); const result = await processChunks(); expect(result).toEqual([{ text: messageText, isFinished: false }]); - expect(Message.findOne).toHaveBeenCalledWith({ messageId: 'message-id' }, 'text unfinished'); - expect(Message.findOne().lean).toHaveBeenCalled(); }); - it('should return the remaining text as a chunk for a finished message', async () => { + it('should return the remaining text as a chunk for a complete message', async () => { const messageText = 'This is a finished message.'; - Message.findOne().lean.mockResolvedValueOnce({ text: messageText, unfinished: false }); + mockMessageCache.get.mockResolvedValueOnce({ text: messageText, complete: true }); const result = await processChunks(); expect(result).toEqual([{ text: messageText, isFinished: true }]); - expect(Message.findOne).toHaveBeenCalledWith({ messageId: 'message-id' }, 'text unfinished'); - expect(Message.findOne().lean).toHaveBeenCalled(); }); - it('should return an empty array for a finished message with no remaining text', async () => { + it('should return an empty array for a complete message with no remaining text', async () => { const messageText = 'This is a finished message.'; - Message.findOne().lean.mockResolvedValueOnce({ text: messageText, unfinished: false }); + mockMessageCache.get.mockResolvedValueOnce({ text: messageText, complete: true }); await processChunks(); - Message.findOne().lean.mockResolvedValueOnce({ text: messageText, unfinished: false }); + mockMessageCache.get.mockResolvedValueOnce({ text: messageText, complete: true }); const result = await processChunks(); expect(result).toEqual([]); - expect(Message.findOne).toHaveBeenCalledWith({ messageId: 'message-id' }, 'text unfinished'); - expect(Message.findOne().lean).toHaveBeenCalledTimes(2); + }); + + it('should return an error message after MAX_NO_CHANGE_COUNT attempts with no change', async () => { + const messageText = 'This is a message that does not change.'; + mockMessageCache.get.mockResolvedValue({ text: messageText, complete: false }); + + for (let i = 0; i < 11; i++) { + await processChunks(); + } + const result = await processChunks(); + + expect(result).toBe('No change in message after 10 attempts'); + }); + + it('should handle string messages as incomplete', async () => { + const messageText = 'This is a message as a string.'; + mockMessageCache.get.mockResolvedValueOnce(messageText); + + const result = await processChunks(); + + expect(result).toEqual([{ text: messageText, isFinished: false }]); }); }); diff --git a/api/server/services/Runs/StreamRunManager.js b/api/server/services/Runs/StreamRunManager.js index 71eb0b0100..951818bb6f 100644 --- a/api/server/services/Runs/StreamRunManager.js +++ b/api/server/services/Runs/StreamRunManager.js @@ -1,17 +1,19 @@ const throttle = require('lodash/throttle'); const { + Time, + CacheKeys, StepTypes, ContentTypes, ToolCallTypes, - // StepStatus, MessageContentTypes, AssistantStreamEvents, + Constants, } = require('librechat-data-provider'); const { retrieveAndProcessFile } = require('~/server/services/Files/process'); const { processRequiredActions } = require('~/server/services/ToolService'); -const { saveMessage, updateMessageText } = require('~/models/Message'); -const { createOnProgress, sendMessage } = require('~/server/utils'); +const { createOnProgress, sendMessage, sleep } = require('~/server/utils'); const { processMessages } = require('~/server/services/Threads'); +const { getLogStores } = require('~/cache'); const { logger } = require('~/config'); /** @@ -68,8 +70,8 @@ class StreamRunManager { this.attachedFileIds = fields.attachedFileIds; /** @type {undefined | Promise} */ this.visionPromise = fields.visionPromise; - /** @type {boolean} */ - this.savedInitialMessage = false; + /** @type {number} */ + this.streamRate = fields.streamRate ?? Constants.DEFAULT_STREAM_RATE; /** * @type {Object. Promise>} @@ -139,11 +141,11 @@ class StreamRunManager { return this.intermediateText; } - /** Saves the initial intermediate message - * @returns {Promise} + /** Returns the current, intermediate message + * @returns {TMessage} */ - async saveInitialMessage() { - return saveMessage(this.req, { + getIntermediateMessage() { + return { conversationId: this.finalMessage.conversationId, messageId: this.finalMessage.messageId, parentMessageId: this.parentMessageId, @@ -155,7 +157,7 @@ class StreamRunManager { sender: 'Assistant', unfinished: true, error: false, - }); + }; } /* <------------------ Main Event Handlers ------------------> */ @@ -347,6 +349,8 @@ class StreamRunManager { type: ContentTypes.TOOL_CALL, index, }); + + await sleep(this.streamRate); } }; @@ -444,6 +448,7 @@ class StreamRunManager { if (content && content.type === MessageContentTypes.TEXT) { this.intermediateText += content.text.value; onProgress(content.text.value); + await sleep(this.streamRate); } } @@ -589,21 +594,14 @@ class StreamRunManager { const index = this.getStepIndex(stepKey); this.orderedRunSteps.set(index, message_creation); + const messageCache = getLogStores(CacheKeys.MESSAGES); // Create the Factory Function to stream the message const { onProgress: progressCallback } = createOnProgress({ onProgress: throttle( () => { - if (!this.savedInitialMessage) { - this.saveInitialMessage(); - this.savedInitialMessage = true; - } else { - updateMessageText({ - messageId: this.finalMessage.messageId, - text: this.getText(), - }); - } + messageCache.set(this.finalMessage.messageId, this.getText(), Time.FIVE_MINUTES); }, - 2000, + 3000, { trailing: false }, ), }); diff --git a/api/server/services/start/assistants.js b/api/server/services/start/assistants.js index ab96db8701..b46edc676b 100644 --- a/api/server/services/start/assistants.js +++ b/api/server/services/start/assistants.js @@ -51,6 +51,7 @@ function assistantsConfigSetup(config, assistantsEndpoint, prevConfig = {}) { excludedIds: parsedConfig.excludedIds, privateAssistants: parsedConfig.privateAssistants, timeoutMs: parsedConfig.timeoutMs, + streamRate: parsedConfig.streamRate, }; } diff --git a/api/typedefs.js b/api/typedefs.js index ecf78c1374..c8f46c6d9b 100644 --- a/api/typedefs.js +++ b/api/typedefs.js @@ -465,6 +465,12 @@ * @memberof typedefs */ +/** + * @exports TBaseEndpoint + * @typedef {import('librechat-data-provider').TBaseEndpoint} TBaseEndpoint + * @memberof typedefs + */ + /** * @exports TEndpoint * @typedef {import('librechat-data-provider').TEndpoint} TEndpoint diff --git a/package-lock.json b/package-lock.json index 5074826fdb..9f2fae80f6 100644 --- a/package-lock.json +++ b/package-lock.json @@ -29437,7 +29437,7 @@ }, "packages/data-provider": { "name": "librechat-data-provider", - "version": "0.7.1", + "version": "0.7.2", "license": "ISC", "dependencies": { "@types/js-yaml": "^4.0.9", diff --git a/packages/data-provider/package.json b/packages/data-provider/package.json index c3622a3c32..393dacd051 100644 --- a/packages/data-provider/package.json +++ b/packages/data-provider/package.json @@ -1,6 +1,6 @@ { "name": "librechat-data-provider", - "version": "0.7.1", + "version": "0.7.2", "description": "data services for librechat apps", "main": "dist/index.js", "module": "dist/index.es.js", diff --git a/packages/data-provider/src/config.ts b/packages/data-provider/src/config.ts index 555d8af400..1f149aa638 100644 --- a/packages/data-provider/src/config.ts +++ b/packages/data-provider/src/config.ts @@ -136,71 +136,81 @@ export const defaultAssistantsVersion = { [EModelEndpoint.azureAssistants]: 1, }; -export const assistantEndpointSchema = z.object({ - /* assistants specific */ - disableBuilder: z.boolean().optional(), - pollIntervalMs: z.number().optional(), - timeoutMs: z.number().optional(), - version: z.union([z.string(), z.number()]).default(2), - supportedIds: z.array(z.string()).min(1).optional(), - excludedIds: z.array(z.string()).min(1).optional(), - privateAssistants: z.boolean().optional(), - retrievalModels: z.array(z.string()).min(1).optional().default(defaultRetrievalModels), - capabilities: z - .array(z.nativeEnum(Capabilities)) - .optional() - .default([ - Capabilities.code_interpreter, - Capabilities.image_vision, - Capabilities.retrieval, - Capabilities.actions, - Capabilities.tools, - ]), - /* general */ - apiKey: z.string().optional(), - baseURL: z.string().optional(), - models: z - .object({ - default: z.array(z.string()).min(1), - fetch: z.boolean().optional(), - userIdQuery: z.boolean().optional(), - }) - .optional(), - titleConvo: z.boolean().optional(), - titleMethod: z.union([z.literal('completion'), z.literal('functions')]).optional(), - titleModel: z.string().optional(), - headers: z.record(z.any()).optional(), +export const baseEndpointSchema = z.object({ + streamRate: z.number().optional(), }); +export type TBaseEndpoint = z.infer; + +export const assistantEndpointSchema = baseEndpointSchema.merge( + z.object({ + /* assistants specific */ + disableBuilder: z.boolean().optional(), + pollIntervalMs: z.number().optional(), + timeoutMs: z.number().optional(), + version: z.union([z.string(), z.number()]).default(2), + supportedIds: z.array(z.string()).min(1).optional(), + excludedIds: z.array(z.string()).min(1).optional(), + privateAssistants: z.boolean().optional(), + retrievalModels: z.array(z.string()).min(1).optional().default(defaultRetrievalModels), + capabilities: z + .array(z.nativeEnum(Capabilities)) + .optional() + .default([ + Capabilities.code_interpreter, + Capabilities.image_vision, + Capabilities.retrieval, + Capabilities.actions, + Capabilities.tools, + ]), + /* general */ + apiKey: z.string().optional(), + baseURL: z.string().optional(), + models: z + .object({ + default: z.array(z.string()).min(1), + fetch: z.boolean().optional(), + userIdQuery: z.boolean().optional(), + }) + .optional(), + titleConvo: z.boolean().optional(), + titleMethod: z.union([z.literal('completion'), z.literal('functions')]).optional(), + titleModel: z.string().optional(), + headers: z.record(z.any()).optional(), + }), +); + export type TAssistantEndpoint = z.infer; -export const endpointSchema = z.object({ - name: z.string().refine((value) => !eModelEndpointSchema.safeParse(value).success, { - message: `Value cannot be one of the default endpoint (EModelEndpoint) values: ${Object.values( - EModelEndpoint, - ).join(', ')}`, +export const endpointSchema = baseEndpointSchema.merge( + z.object({ + name: z.string().refine((value) => !eModelEndpointSchema.safeParse(value).success, { + message: `Value cannot be one of the default endpoint (EModelEndpoint) values: ${Object.values( + EModelEndpoint, + ).join(', ')}`, + }), + apiKey: z.string(), + baseURL: z.string(), + models: z.object({ + default: z.array(z.string()).min(1), + fetch: z.boolean().optional(), + userIdQuery: z.boolean().optional(), + }), + titleConvo: z.boolean().optional(), + titleMethod: z.union([z.literal('completion'), z.literal('functions')]).optional(), + titleModel: z.string().optional(), + summarize: z.boolean().optional(), + summaryModel: z.string().optional(), + forcePrompt: z.boolean().optional(), + modelDisplayLabel: z.string().optional(), + headers: z.record(z.any()).optional(), + addParams: z.record(z.any()).optional(), + dropParams: z.array(z.string()).optional(), + customOrder: z.number().optional(), + directEndpoint: z.boolean().optional(), + titleMessageRole: z.string().optional(), }), - apiKey: z.string(), - baseURL: z.string(), - models: z.object({ - default: z.array(z.string()).min(1), - fetch: z.boolean().optional(), - userIdQuery: z.boolean().optional(), - }), - titleConvo: z.boolean().optional(), - titleMethod: z.union([z.literal('completion'), z.literal('functions')]).optional(), - titleModel: z.string().optional(), - summarize: z.boolean().optional(), - summaryModel: z.string().optional(), - forcePrompt: z.boolean().optional(), - modelDisplayLabel: z.string().optional(), - headers: z.record(z.any()).optional(), - addParams: z.record(z.any()).optional(), - dropParams: z.array(z.string()).optional(), - customOrder: z.number().optional(), - directEndpoint: z.boolean().optional(), - titleMessageRole: z.string().optional(), -}); +); export type TEndpoint = z.infer; @@ -213,6 +223,7 @@ export const azureEndpointSchema = z .and( endpointSchema .pick({ + streamRate: true, titleConvo: true, titleMethod: true, titleModel: true, @@ -426,10 +437,15 @@ export const configSchema = z.object({ modelSpecs: specsConfigSchema.optional(), endpoints: z .object({ + all: baseEndpointSchema.optional(), + [EModelEndpoint.openAI]: baseEndpointSchema.optional(), + [EModelEndpoint.google]: baseEndpointSchema.optional(), + [EModelEndpoint.anthropic]: baseEndpointSchema.optional(), + [EModelEndpoint.gptPlugins]: baseEndpointSchema.optional(), [EModelEndpoint.azureOpenAI]: azureEndpointSchema.optional(), [EModelEndpoint.azureAssistants]: assistantEndpointSchema.optional(), [EModelEndpoint.assistants]: assistantEndpointSchema.optional(), - custom: z.array(endpointSchema.partial()).optional(), + [EModelEndpoint.custom]: z.array(endpointSchema.partial()).optional(), }) .strict() .refine((data) => Object.keys(data).length > 0, { @@ -657,6 +673,16 @@ export enum InfiniteCollections { SHARED_LINKS = 'sharedLinks', } +/** + * Enum for time intervals + */ +export enum Time { + THIRTY_MINUTES = 1800000, + TEN_MINUTES = 600000, + FIVE_MINUTES = 300000, + TWO_MINUTES = 120000, +} + /** * Enum for cache keys. */ @@ -727,6 +753,10 @@ export enum CacheKeys { * Key for the cached audio run Ids. */ AUDIO_RUNS = 'audioRuns', + /** + * Key for in-progress messages. + */ + MESSAGES = 'messages', } /** @@ -911,6 +941,8 @@ export enum Constants { COMMON_DIVIDER = '__', /** Max length for commands */ COMMANDS_MAX_LENGTH = 56, + /** Default Stream Rate (ms) */ + DEFAULT_STREAM_RATE = 1, } export enum LocalStorageKeys {