diff --git a/api/app/clients/AnthropicClient.js b/api/app/clients/AnthropicClient.js index aa5f913367..2373a321f5 100644 --- a/api/app/clients/AnthropicClient.js +++ b/api/app/clients/AnthropicClient.js @@ -1,4 +1,5 @@ const Anthropic = require('@anthropic-ai/sdk'); +const { HttpsProxyAgent } = require('https-proxy-agent'); const { encoding_for_model: encodingForModel, get_encoding: getEncoding } = require('tiktoken'); const { getResponseSender, @@ -123,9 +124,14 @@ class AnthropicClient extends BaseClient { getClient() { /** @type {Anthropic.default.RequestOptions} */ const options = { + fetch: this.fetch, apiKey: this.apiKey, }; + if (this.options.proxy) { + options.httpAgent = new HttpsProxyAgent(this.options.proxy); + } + if (this.options.reverseProxyUrl) { options.baseURL = this.options.reverseProxyUrl; } diff --git a/api/app/clients/BaseClient.js b/api/app/clients/BaseClient.js index 39215d917b..f2b3e7c59b 100644 --- a/api/app/clients/BaseClient.js +++ b/api/app/clients/BaseClient.js @@ -1,4 +1,5 @@ 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'); @@ -17,6 +18,7 @@ class BaseClient { month: 'long', day: 'numeric', }); + this.fetch = this.fetch.bind(this); } setOptions() { @@ -54,6 +56,22 @@ class BaseClient { }); } + /** + * Makes an HTTP request and logs the process. + * + * @param {RequestInfo} url - The URL to make the request to. Can be a string or a Request object. + * @param {RequestInit} [init] - Optional init options for the request. + * @returns {Promise} - A promise that resolves to the response of the fetch request. + */ + async fetch(_url, init) { + let url = _url; + if (this.options.directEndpoint) { + url = this.options.reverseProxyUrl; + } + logger.debug(`Making request to ${url}`); + return await fetch(url, init); + } + getBuildMessagesOptions() { throw new Error('Subclasses must implement getBuildMessagesOptions'); } diff --git a/api/app/clients/OpenAIClient.js b/api/app/clients/OpenAIClient.js index 1635e600ca..ced2387bd5 100644 --- a/api/app/clients/OpenAIClient.js +++ b/api/app/clients/OpenAIClient.js @@ -589,7 +589,7 @@ class OpenAIClient extends BaseClient { let streamResult = null; this.modelOptions.user = this.user; const invalidBaseUrl = this.completionsUrl && extractBaseURL(this.completionsUrl) === null; - const useOldMethod = !!(invalidBaseUrl || !this.isChatCompletion || typeof Bun !== 'undefined'); + const useOldMethod = !!(invalidBaseUrl || !this.isChatCompletion); if (typeof opts.onProgress === 'function' && useOldMethod) { const completionResult = await this.getCompletion( payload, @@ -829,7 +829,7 @@ class OpenAIClient extends BaseClient { const instructionsPayload = [ { - role: 'system', + role: this.options.titleMessageRole ?? 'system', content: `Please generate ${titleInstruction} ${convo} @@ -1134,6 +1134,7 @@ ${convo} let chatCompletion; /** @type {OpenAI} */ const openai = new OpenAI({ + fetch: this.fetch, apiKey: this.apiKey, ...opts, }); diff --git a/api/app/clients/PluginsClient.js b/api/app/clients/PluginsClient.js index da1ee1eb46..e321cb351e 100644 --- a/api/app/clients/PluginsClient.js +++ b/api/app/clients/PluginsClient.js @@ -268,7 +268,7 @@ class PluginsClient extends OpenAIClient { if (opts.progressCallback) { opts.onProgress = opts.progressCallback.call(null, { ...(opts.progressOptions ?? {}), - parentMessageId: opts.progressOptions?.parentMessageId ?? userMessage.messageId, + parentMessageId: userMessage.messageId, messageId: responseMessageId, }); } diff --git a/api/models/Message.js b/api/models/Message.js index 0df0a7c2da..6435fa6fe8 100644 --- a/api/models/Message.js +++ b/api/models/Message.js @@ -129,6 +129,14 @@ module.exports = { throw new Error('Failed to save message.'); } }, + async updateMessageText({ messageId, text }) { + try { + await Message.updateOne({ messageId }, { text }); + } catch (err) { + logger.error('Error updating message text:', err); + throw new Error('Failed to update message text.'); + } + }, async updateMessage(message) { try { const { messageId, ...update } = message; diff --git a/api/server/controllers/assistants/chatV2.js b/api/server/controllers/assistants/chatV2.js index 664037d762..3b73d1520f 100644 --- a/api/server/controllers/assistants/chatV2.js +++ b/api/server/controllers/assistants/chatV2.js @@ -496,7 +496,7 @@ const chatV2 = async (req, res) => { handlers, thread_id, attachedFileIds, - parentMessageId, + parentMessageId: userMessageId, responseMessage: openai.responseMessage, // streamOptions: { diff --git a/api/server/middleware/abortRun.js b/api/server/middleware/abortRun.js index 6522d6746d..512554aec9 100644 --- a/api/server/middleware/abortRun.js +++ b/api/server/middleware/abortRun.js @@ -1,6 +1,7 @@ const { CacheKeys, RunStatus, isUUID } = require('librechat-data-provider'); const { initializeClient } = require('~/server/services/Endpoints/assistants'); const { checkMessageGaps, recordUsage } = require('~/server/services/Threads'); +const { deleteMessages } = require('~/models/Message'); const { getConvo } = require('~/models/Conversation'); const getLogStores = require('~/cache/getLogStores'); const { sendMessage } = require('~/server/utils'); @@ -66,13 +67,19 @@ async function abortRun(req, res) { logger.error('[abortRun] Error fetching or processing run', error); } + /* TODO: a reconciling strategy between the existing intermediate message would be more optimal than deleting it */ + await deleteMessages({ + user: req.user.id, + unfinished: true, + conversationId, + }); runMessages = await checkMessageGaps({ openai, + run_id, endpoint, thread_id, - run_id, - latestMessageId, conversationId, + latestMessageId, }); const finalEvent = { diff --git a/api/server/routes/ask/gptPlugins.js b/api/server/routes/ask/gptPlugins.js index 2fab6188af..2acf4ce592 100644 --- a/api/server/routes/ask/gptPlugins.js +++ b/api/server/routes/ask/gptPlugins.js @@ -106,7 +106,11 @@ router.post( const pluginMap = new Map(); const onAgentAction = async (action, runId) => { pluginMap.set(runId, action.tool); - sendIntermediateMessage(res, { plugins }); + sendIntermediateMessage(res, { + plugins, + parentMessageId: userMessage.messageId, + messageId: responseMessageId, + }); }; const onToolStart = async (tool, input, runId, parentRunId) => { @@ -124,7 +128,11 @@ router.post( } const extraTokens = ':::plugin:::\n'; plugins.push(latestPlugin); - sendIntermediateMessage(res, { plugins }, extraTokens); + sendIntermediateMessage( + res, + { plugins, parentMessageId: userMessage.messageId, messageId: responseMessageId }, + extraTokens, + ); }; const onToolEnd = async (output, runId) => { @@ -142,7 +150,11 @@ router.post( const onChainEnd = () => { saveMessage({ ...userMessage, user }); - sendIntermediateMessage(res, { plugins }); + sendIntermediateMessage(res, { + plugins, + parentMessageId: userMessage.messageId, + messageId: responseMessageId, + }); }; const getAbortData = () => ({ diff --git a/api/server/routes/edit/gptPlugins.js b/api/server/routes/edit/gptPlugins.js index f1b0cba248..cf71d487dc 100644 --- a/api/server/routes/edit/gptPlugins.js +++ b/api/server/routes/edit/gptPlugins.js @@ -110,7 +110,11 @@ router.post( if (!start) { saveMessage({ ...userMessage, user }); } - sendIntermediateMessage(res, { plugin }); + sendIntermediateMessage(res, { + plugin, + parentMessageId: userMessage.messageId, + messageId: responseMessageId, + }); // logger.debug('PLUGIN ACTION', formattedAction); }; @@ -119,7 +123,11 @@ router.post( plugin.outputs = steps && steps[0].action ? formatSteps(steps) : 'An error occurred.'; plugin.loading = false; saveMessage({ ...userMessage, user }); - sendIntermediateMessage(res, { plugin }); + sendIntermediateMessage(res, { + plugin, + parentMessageId: userMessage.messageId, + messageId: responseMessageId, + }); // logger.debug('CHAIN END', plugin.outputs); }; diff --git a/api/server/routes/files/index.js b/api/server/routes/files/index.js index e74b167d45..1490b4ec90 100644 --- a/api/server/routes/files/index.js +++ b/api/server/routes/files/index.js @@ -14,6 +14,10 @@ const initialize = async () => { router.use(checkBan); router.use(uaParser); + /* Important: stt/tts routes must be added before the upload limiters */ + router.use('/stt', stt); + router.use('/tts', tts); + const upload = await createMulterInstance(); const { fileUploadIpLimiter, fileUploadUserLimiter } = createFileLimiters(); router.post('*', fileUploadIpLimiter, fileUploadUserLimiter); diff --git a/api/server/services/Endpoints/custom/initializeClient.js b/api/server/services/Endpoints/custom/initializeClient.js index a5da778282..9fb6bfd1af 100644 --- a/api/server/services/Endpoints/custom/initializeClient.js +++ b/api/server/services/Endpoints/custom/initializeClient.js @@ -112,6 +112,8 @@ const initializeClient = async ({ req, res, endpointOption }) => { modelDisplayLabel: endpointConfig.modelDisplayLabel, titleMethod: endpointConfig.titleMethod ?? 'completion', contextStrategy: endpointConfig.summarize ? 'summarize' : null, + directEndpoint: endpointConfig.directEndpoint, + titleMessageRole: endpointConfig.titleMessageRole, endpointTokenConfig, }; diff --git a/api/server/services/Runs/StreamRunManager.js b/api/server/services/Runs/StreamRunManager.js index bcae609c7b..bea8042aef 100644 --- a/api/server/services/Runs/StreamRunManager.js +++ b/api/server/services/Runs/StreamRunManager.js @@ -9,9 +9,9 @@ const { } = 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 { processMessages } = require('~/server/services/Threads'); -const { saveMessage } = require('~/models'); const { logger } = require('~/config'); /** @@ -68,6 +68,8 @@ class StreamRunManager { this.attachedFileIds = fields.attachedFileIds; /** @type {undefined | Promise} */ this.visionPromise = fields.visionPromise; + /** @type {boolean} */ + this.savedInitialMessage = false; /** * @type {Object. Promise>} @@ -129,6 +131,33 @@ class StreamRunManager { sendMessage(this.res, contentData); } + /* <------------------ Misc. Helpers ------------------> */ + /** Returns the latest intermediate text + * @returns {string} + */ + getText() { + return this.intermediateText; + } + + /** Saves the initial intermediate message + * @returns {Promise} + */ + async saveInitialMessage() { + return saveMessage({ + conversationId: this.finalMessage.conversationId, + messageId: this.finalMessage.messageId, + parentMessageId: this.parentMessageId, + model: this.req.body.assistant_id, + endpoint: this.req.body.endpoint, + isCreatedByUser: false, + user: this.req.user.id, + text: this.getText(), + sender: 'Assistant', + unfinished: true, + error: false, + }); + } + /* <------------------ Main Event Handlers ------------------> */ /** @@ -530,23 +559,20 @@ class StreamRunManager { const stepKey = message_creation.message_id; const index = this.getStepIndex(stepKey); this.orderedRunSteps.set(index, message_creation); - const getText = () => this.intermediateText; + // Create the Factory Function to stream the message const { onProgress: progressCallback } = createOnProgress({ onProgress: throttle( () => { - const text = getText(); - saveMessage({ - messageId: this.finalMessage.messageId, - conversationId: this.finalMessage.conversationId, - parentMessageId: this.parentMessageId, - model: this.req.body.model, - user: this.req.user.id, - sender: 'Assistant', - unfinished: true, - error: false, - text, - }); + if (!this.savedInitialMessage) { + this.saveInitialMessage(); + this.savedInitialMessage = true; + } else { + updateMessageText({ + messageId: this.finalMessage.messageId, + text: this.getText(), + }); + } }, 2000, { trailing: false }, diff --git a/bun.lockb b/bun.lockb index 3f37e7a782..1351a98068 100755 Binary files a/bun.lockb and b/bun.lockb differ diff --git a/client/src/hooks/Audio/usePauseGlobalAudio.ts b/client/src/hooks/Audio/usePauseGlobalAudio.ts index a36f66c89f..06fc7b730b 100644 --- a/client/src/hooks/Audio/usePauseGlobalAudio.ts +++ b/client/src/hooks/Audio/usePauseGlobalAudio.ts @@ -6,9 +6,9 @@ import store from '~/store'; function usePauseGlobalAudio(index = 0) { /* Global Audio Variables */ const setAudioRunId = useSetRecoilState(store.audioRunFamily(index)); + const setGlobalIsPlaying = useSetRecoilState(store.globalAudioPlayingFamily(index)); const setIsGlobalAudioFetching = useSetRecoilState(store.globalAudioFetchingFamily(index)); const [globalAudioURL, setGlobalAudioURL] = useRecoilState(store.globalAudioURLFamily(index)); - const setGlobalIsPlaying = useSetRecoilState(store.globalAudioPlayingFamily(index)); const pauseGlobalAudio = useCallback(() => { if (globalAudioURL) { diff --git a/client/src/hooks/SSE/useSSE.ts b/client/src/hooks/SSE/useSSE.ts index 50f88c891a..4b821ba403 100644 --- a/client/src/hooks/SSE/useSSE.ts +++ b/client/src/hooks/SSE/useSSE.ts @@ -282,6 +282,12 @@ export default function useSSE(submission: TSubmission | null, index = 0) { setShowStopButton(false); setCompleted((prev) => new Set(prev.add(submission?.initialResponse?.messageId))); + const currentMessages = getMessages(); + // Early return if messages are empty; i.e., the user navigated away + if (!currentMessages?.length) { + return setIsSubmitting(false); + } + // update the messages; if assistants endpoint, client doesn't receive responseMessage if (runMessages) { setMessages([...runMessages]); @@ -323,7 +329,15 @@ export default function useSSE(submission: TSubmission | null, index = 0) { setIsSubmitting(false); }, - [genTitle, queryClient, setMessages, setConversation, setIsSubmitting, setShowStopButton], + [ + genTitle, + queryClient, + getMessages, + setMessages, + setConversation, + setIsSubmitting, + setShowStopButton, + ], ); const errorHandler = useCallback( diff --git a/packages/data-provider/package.json b/packages/data-provider/package.json index 8f638a7585..d04a592a8c 100644 --- a/packages/data-provider/package.json +++ b/packages/data-provider/package.json @@ -1,6 +1,6 @@ { "name": "librechat-data-provider", - "version": "0.6.5", + "version": "0.6.6", "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 bbcd073807..c842326c5a 100644 --- a/packages/data-provider/src/config.ts +++ b/packages/data-provider/src/config.ts @@ -198,6 +198,8 @@ export const endpointSchema = z.object({ 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; @@ -747,7 +749,7 @@ export enum Constants { /** Key for the app's version. */ VERSION = 'v0.7.2', /** Key for the Custom Config's version (librechat.yaml). */ - CONFIG_VERSION = '1.1.2', + CONFIG_VERSION = '1.1.3', /** Standard value for the first message's `parentMessageId` value, to indicate no parent exists. */ NO_PARENT = '00000000-0000-0000-0000-000000000000', /** Fixed, encoded domain length for Azure OpenAI Assistants Function name parsing. */