diff --git a/api/app/clients/AnthropicClient.js b/api/app/clients/AnthropicClient.js index 6e9d4accc2..8dc0e40d56 100644 --- a/api/app/clients/AnthropicClient.js +++ b/api/app/clients/AnthropicClient.js @@ -1,6 +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 { Constants, EModelEndpoint, @@ -19,6 +18,7 @@ const { } = require('./prompts'); const { getModelMaxTokens, getModelMaxOutputTokens, matchModelName } = require('~/utils'); const { spendTokens, spendStructuredTokens } = require('~/models/spendTokens'); +const Tokenizer = require('~/server/services/Tokenizer'); const { sleep } = require('~/server/utils'); const BaseClient = require('./BaseClient'); const { logger } = require('~/config'); @@ -26,8 +26,6 @@ const { logger } = require('~/config'); const HUMAN_PROMPT = '\n\nHuman:'; const AI_PROMPT = '\n\nAssistant:'; -const tokenizersCache = {}; - /** Helper function to introduce a delay before retrying */ function delayBeforeRetry(attempts, baseDelay = 1000) { return new Promise((resolve) => setTimeout(resolve, baseDelay * attempts)); @@ -149,7 +147,6 @@ class AnthropicClient extends BaseClient { this.startToken = '||>'; this.endToken = ''; - this.gptEncoder = this.constructor.getTokenizer('cl100k_base'); return this; } @@ -849,22 +846,18 @@ class AnthropicClient extends BaseClient { logger.debug('AnthropicClient doesn\'t use getBuildMessagesOptions'); } - static getTokenizer(encoding, isModelName = false, extendSpecialTokens = {}) { - if (tokenizersCache[encoding]) { - return tokenizersCache[encoding]; - } - let tokenizer; - if (isModelName) { - tokenizer = encodingForModel(encoding, extendSpecialTokens); - } else { - tokenizer = getEncoding(encoding, extendSpecialTokens); - } - tokenizersCache[encoding] = tokenizer; - return tokenizer; + getEncoding() { + return 'cl100k_base'; } + /** + * Returns the token count of a given text. It also checks and resets the tokenizers if necessary. + * @param {string} text - The text to get the token count for. + * @returns {number} The token count of the given text. + */ getTokenCount(text) { - return this.gptEncoder.encode(text, 'all').length; + const encoding = this.getEncoding(); + return Tokenizer.getTokenCount(text, encoding); } /** diff --git a/api/app/clients/GoogleClient.js b/api/app/clients/GoogleClient.js index f6fca55112..7e34a65c02 100644 --- a/api/app/clients/GoogleClient.js +++ b/api/app/clients/GoogleClient.js @@ -6,7 +6,6 @@ const { ChatGoogleVertexAI } = require('@langchain/google-vertexai'); const { ChatGoogleGenerativeAI } = require('@langchain/google-genai'); const { GoogleGenerativeAI: GenAI } = require('@google/generative-ai'); const { AIMessage, HumanMessage, SystemMessage } = require('@langchain/core/messages'); -const { encoding_for_model: encodingForModel, get_encoding: getEncoding } = require('tiktoken'); const { validateVisionModel, getResponseSender, @@ -17,6 +16,7 @@ const { AuthKeys, } = require('librechat-data-provider'); const { encodeAndFormat } = require('~/server/services/Files/images'); +const Tokenizer = require('~/server/services/Tokenizer'); const { getModelMaxTokens } = require('~/utils'); const { sleep } = require('~/server/utils'); const { logger } = require('~/config'); @@ -31,7 +31,6 @@ const BaseClient = require('./BaseClient'); const loc = process.env.GOOGLE_LOC || 'us-central1'; const publisher = 'google'; const endpointPrefix = `${loc}-aiplatform.googleapis.com`; -const tokenizersCache = {}; const settings = endpointSettings[EModelEndpoint.google]; const EXCLUDED_GENAI_MODELS = /gemini-(?:1\.0|1-0|pro)/; @@ -177,25 +176,15 @@ class GoogleClient extends BaseClient { // without tripping the stop sequences, so I'm using "||>" instead. this.startToken = '||>'; this.endToken = ''; - this.gptEncoder = this.constructor.getTokenizer('cl100k_base'); } else if (isTextModel) { this.startToken = '||>'; this.endToken = ''; - this.gptEncoder = this.constructor.getTokenizer('text-davinci-003', true, { - '<|im_start|>': 100264, - '<|im_end|>': 100265, - }); } else { // Previously I was trying to use "<|endoftext|>" but there seems to be some bug with OpenAI's token counting // system that causes only the first "<|endoftext|>" to be counted as 1 token, and the rest are not treated // as a single token. So we're using this instead. this.startToken = '||>'; this.endToken = ''; - try { - this.gptEncoder = this.constructor.getTokenizer(this.modelOptions.model, true); - } catch { - this.gptEncoder = this.constructor.getTokenizer('text-davinci-003', true); - } } if (!this.modelOptions.stop) { @@ -926,23 +915,18 @@ class GoogleClient extends BaseClient { ]; } - /* TO-DO: Handle tokens with Google tokenization NOTE: these are required */ - static getTokenizer(encoding, isModelName = false, extendSpecialTokens = {}) { - if (tokenizersCache[encoding]) { - return tokenizersCache[encoding]; - } - let tokenizer; - if (isModelName) { - tokenizer = encodingForModel(encoding, extendSpecialTokens); - } else { - tokenizer = getEncoding(encoding, extendSpecialTokens); - } - tokenizersCache[encoding] = tokenizer; - return tokenizer; + getEncoding() { + return 'cl100k_base'; } + /** + * Returns the token count of a given text. It also checks and resets the tokenizers if necessary. + * @param {string} text - The text to get the token count for. + * @returns {number} The token count of the given text. + */ getTokenCount(text) { - return this.gptEncoder.encode(text, 'all').length; + const encoding = this.getEncoding(); + return Tokenizer.getTokenCount(text, encoding); } } diff --git a/api/app/clients/OpenAIClient.js b/api/app/clients/OpenAIClient.js index f81a4bc92c..15fd20aefe 100644 --- a/api/app/clients/OpenAIClient.js +++ b/api/app/clients/OpenAIClient.js @@ -13,7 +13,6 @@ const { validateVisionModel, mapModelToAzureConfig, } = require('librechat-data-provider'); -const { encoding_for_model: encodingForModel, get_encoding: getEncoding } = require('tiktoken'); const { extractBaseURL, constructAzureURL, @@ -29,6 +28,7 @@ const { createContextHandlers, } = require('./prompts'); const { encodeAndFormat } = require('~/server/services/Files/images/encode'); +const Tokenizer = require('~/server/services/Tokenizer'); const { spendTokens } = require('~/models/spendTokens'); const { isEnabled, sleep } = require('~/server/utils'); const { handleOpenAIErrors } = require('./tools/util'); @@ -40,11 +40,6 @@ const { tokenSplit } = require('./document'); const BaseClient = require('./BaseClient'); const { logger } = require('~/config'); -// Cache to store Tiktoken instances -const tokenizersCache = {}; -// Counter for keeping track of the number of tokenizer calls -let tokenizerCallsCount = 0; - class OpenAIClient extends BaseClient { constructor(apiKey, options = {}) { super(apiKey, options); @@ -307,75 +302,8 @@ class OpenAIClient extends BaseClient { } } - // Selects an appropriate tokenizer based on the current configuration of the client instance. - // It takes into account factors such as whether it's a chat completion, an unofficial chat GPT model, etc. - selectTokenizer() { - let tokenizer; - this.encoding = 'text-davinci-003'; - if (this.isChatCompletion) { - this.encoding = this.modelOptions.model.includes('gpt-4o') ? 'o200k_base' : 'cl100k_base'; - tokenizer = this.constructor.getTokenizer(this.encoding); - } else if (this.isUnofficialChatGptModel) { - const extendSpecialTokens = { - '<|im_start|>': 100264, - '<|im_end|>': 100265, - }; - tokenizer = this.constructor.getTokenizer(this.encoding, true, extendSpecialTokens); - } else { - try { - const { model } = this.modelOptions; - this.encoding = model.includes('instruct') ? 'text-davinci-003' : model; - tokenizer = this.constructor.getTokenizer(this.encoding, true); - } catch { - tokenizer = this.constructor.getTokenizer('text-davinci-003', true); - } - } - - return tokenizer; - } - - // Retrieves a tokenizer either from the cache or creates a new one if one doesn't exist in the cache. - // If a tokenizer is being created, it's also added to the cache. - static getTokenizer(encoding, isModelName = false, extendSpecialTokens = {}) { - let tokenizer; - if (tokenizersCache[encoding]) { - tokenizer = tokenizersCache[encoding]; - } else { - if (isModelName) { - tokenizer = encodingForModel(encoding, extendSpecialTokens); - } else { - tokenizer = getEncoding(encoding, extendSpecialTokens); - } - tokenizersCache[encoding] = tokenizer; - } - return tokenizer; - } - - // Frees all encoders in the cache and resets the count. - static freeAndResetAllEncoders() { - try { - Object.keys(tokenizersCache).forEach((key) => { - if (tokenizersCache[key]) { - tokenizersCache[key].free(); - delete tokenizersCache[key]; - } - }); - // Reset count - tokenizerCallsCount = 1; - } catch (error) { - logger.error('[OpenAIClient] Free and reset encoders error', error); - } - } - - // Checks if the cache of tokenizers has reached a certain size. If it has, it frees and resets all tokenizers. - resetTokenizersIfNecessary() { - if (tokenizerCallsCount >= 25) { - if (this.options.debug) { - logger.debug('[OpenAIClient] freeAndResetAllEncoders: reached 25 encodings, resetting...'); - } - this.constructor.freeAndResetAllEncoders(); - } - tokenizerCallsCount++; + getEncoding() { + return this.model?.includes('gpt-4o') ? 'o200k_base' : 'cl100k_base'; } /** @@ -384,15 +312,8 @@ class OpenAIClient extends BaseClient { * @returns {number} The token count of the given text. */ getTokenCount(text) { - this.resetTokenizersIfNecessary(); - try { - const tokenizer = this.selectTokenizer(); - return tokenizer.encode(text, 'all').length; - } catch (error) { - this.constructor.freeAndResetAllEncoders(); - const tokenizer = this.selectTokenizer(); - return tokenizer.encode(text, 'all').length; - } + const encoding = this.getEncoding(); + return Tokenizer.getTokenCount(text, encoding); } /** diff --git a/api/app/clients/specs/OpenAIClient.test.js b/api/app/clients/specs/OpenAIClient.test.js index ab70d4c907..2aaec518eb 100644 --- a/api/app/clients/specs/OpenAIClient.test.js +++ b/api/app/clients/specs/OpenAIClient.test.js @@ -1,5 +1,7 @@ +jest.mock('~/cache/getLogStores'); require('dotenv').config(); const OpenAI = require('openai'); +const getLogStores = require('~/cache/getLogStores'); const { fetchEventSource } = require('@waylaidwanderer/fetch-event-source'); const { genAzureChatCompletion } = require('~/utils/azureUtils'); const OpenAIClient = require('../OpenAIClient'); @@ -134,7 +136,13 @@ OpenAI.mockImplementation(() => ({ })); describe('OpenAIClient', () => { - let client, client2; + const mockSet = jest.fn(); + const mockCache = { set: mockSet }; + + beforeEach(() => { + getLogStores.mockReturnValue(mockCache); + }); + let client; const model = 'gpt-4'; const parentMessageId = '1'; const messages = [ @@ -176,7 +184,6 @@ describe('OpenAIClient', () => { beforeEach(() => { const options = { ...defaultOptions }; client = new OpenAIClient('test-api-key', options); - client2 = new OpenAIClient('test-api-key', options); client.summarizeMessages = jest.fn().mockResolvedValue({ role: 'assistant', content: 'Refined answer', @@ -185,7 +192,6 @@ describe('OpenAIClient', () => { client.buildPrompt = jest .fn() .mockResolvedValue({ prompt: messages.map((m) => m.text).join('\n') }); - client.constructor.freeAndResetAllEncoders(); client.getMessages = jest.fn().mockResolvedValue([]); }); @@ -335,77 +341,11 @@ describe('OpenAIClient', () => { }); }); - describe('selectTokenizer', () => { - it('should get the correct tokenizer based on the instance state', () => { - const tokenizer = client.selectTokenizer(); - expect(tokenizer).toBeDefined(); - }); - }); - - describe('freeAllTokenizers', () => { - it('should free all tokenizers', () => { - // Create a tokenizer - const tokenizer = client.selectTokenizer(); - - // Mock 'free' method on the tokenizer - tokenizer.free = jest.fn(); - - client.constructor.freeAndResetAllEncoders(); - - // Check if 'free' method has been called on the tokenizer - expect(tokenizer.free).toHaveBeenCalled(); - }); - }); - describe('getTokenCount', () => { it('should return the correct token count', () => { const count = client.getTokenCount('Hello, world!'); expect(count).toBeGreaterThan(0); }); - - it('should reset the encoder and count when count reaches 25', () => { - const freeAndResetEncoderSpy = jest.spyOn(client.constructor, 'freeAndResetAllEncoders'); - - // Call getTokenCount 25 times - for (let i = 0; i < 25; i++) { - client.getTokenCount('test text'); - } - - expect(freeAndResetEncoderSpy).toHaveBeenCalled(); - }); - - it('should not reset the encoder and count when count is less than 25', () => { - const freeAndResetEncoderSpy = jest.spyOn(client.constructor, 'freeAndResetAllEncoders'); - freeAndResetEncoderSpy.mockClear(); - - // Call getTokenCount 24 times - for (let i = 0; i < 24; i++) { - client.getTokenCount('test text'); - } - - expect(freeAndResetEncoderSpy).not.toHaveBeenCalled(); - }); - - it('should handle errors and reset the encoder', () => { - const freeAndResetEncoderSpy = jest.spyOn(client.constructor, 'freeAndResetAllEncoders'); - - // Mock encode function to throw an error - client.selectTokenizer().encode = jest.fn().mockImplementation(() => { - throw new Error('Test error'); - }); - - client.getTokenCount('test text'); - - expect(freeAndResetEncoderSpy).toHaveBeenCalled(); - }); - - it('should not throw null pointer error when freeing the same encoder twice', () => { - client.constructor.freeAndResetAllEncoders(); - client2.constructor.freeAndResetAllEncoders(); - - const count = client2.getTokenCount('test text'); - expect(count).toBeGreaterThan(0); - }); }); describe('getSaveOptions', () => { @@ -548,7 +488,6 @@ describe('OpenAIClient', () => { testCases.forEach((testCase) => { it(`should return ${testCase.expected} tokens for model ${testCase.model}`, () => { client.modelOptions.model = testCase.model; - client.selectTokenizer(); // 3 tokens for assistant label let totalTokens = 3; for (let message of example_messages) { @@ -582,7 +521,6 @@ describe('OpenAIClient', () => { it(`should return ${expectedTokens} tokens for model ${visionModel} (Vision Request)`, () => { client.modelOptions.model = visionModel; - client.selectTokenizer(); // 3 tokens for assistant label let totalTokens = 3; for (let message of vision_request) { diff --git a/api/cache/getLogStores.js b/api/cache/getLogStores.js index eec530ad18..b7ff50150e 100644 --- a/api/cache/getLogStores.js +++ b/api/cache/getLogStores.js @@ -5,7 +5,7 @@ const { math, isEnabled } = require('~/server/utils'); const keyvRedis = require('./keyvRedis'); const keyvMongo = require('./keyvMongo'); -const { BAN_DURATION, USE_REDIS, DEBUG_MEMORY_CACHE } = process.env ?? {}; +const { BAN_DURATION, USE_REDIS, DEBUG_MEMORY_CACHE, CI } = process.env ?? {}; const duration = math(BAN_DURATION, 7200000); const isRedisEnabled = isEnabled(USE_REDIS); @@ -95,10 +95,8 @@ const namespaces = { * @returns {Keyv[]} */ function getTTLStores() { - return Object.values(namespaces).filter((store) => - store instanceof Keyv && - typeof store.opts?.ttl === 'number' && - store.opts.ttl > 0, + return Object.values(namespaces).filter( + (store) => store instanceof Keyv && typeof store.opts?.ttl === 'number' && store.opts.ttl > 0, ); } @@ -125,23 +123,28 @@ async function clearExpiredFromCache(cache) { for (const key of keys) { try { const raw = cache.opts.store.get(key); - if (!raw) {continue;} + if (!raw) { + continue; + } const data = cache.opts.deserialize(raw); // Check if the entry is older than TTL if (data?.expires && data.expires <= expiryTime) { const deleted = await cache.opts.store.delete(key); if (!deleted) { - debugMemoryCache && console.warn(`[Cache] Error deleting entry: ${key} from ${cache.opts.namespace}`); + debugMemoryCache && + console.warn(`[Cache] Error deleting entry: ${key} from ${cache.opts.namespace}`); continue; } cleared++; } } catch (error) { - debugMemoryCache && console.log(`[Cache] Error processing entry from ${cache.opts.namespace}:`, error); + debugMemoryCache && + console.log(`[Cache] Error processing entry from ${cache.opts.namespace}:`, error); const deleted = await cache.opts.store.delete(key); if (!deleted) { - debugMemoryCache && console.warn(`[Cache] Error deleting entry: ${key} from ${cache.opts.namespace}`); + debugMemoryCache && + console.warn(`[Cache] Error deleting entry: ${key} from ${cache.opts.namespace}`); continue; } cleared++; @@ -149,7 +152,10 @@ async function clearExpiredFromCache(cache) { } if (cleared > 0) { - debugMemoryCache && console.log(`[Cache] Cleared ${cleared} entries older than ${ttl}ms from ${cache.opts.namespace}`); + debugMemoryCache && + console.log( + `[Cache] Cleared ${cleared} entries older than ${ttl}ms from ${cache.opts.namespace}`, + ); } } @@ -157,7 +163,7 @@ const auditCache = () => { const ttlStores = getTTLStores(); console.log('[Cache] Starting audit'); - ttlStores.forEach(store => { + ttlStores.forEach((store) => { if (!store?.opts?.store?.entries) { return; } @@ -166,21 +172,20 @@ const auditCache = () => { count: store.opts.store.size, ttl: store.opts.ttl, keys: Array.from(store.opts.store.keys()), - entriesWithTimestamps: Array.from(store.opts.store.entries()) - .map(([key, value]) => ({ - key, - value, - })), + entriesWithTimestamps: Array.from(store.opts.store.entries()).map(([key, value]) => ({ + key, + value, + })), }); }); }; /** - * Clears expired entries from all TTL-enabled stores - */ + * Clears expired entries from all TTL-enabled stores + */ async function clearAllExpiredFromCache() { const ttlStores = getTTLStores(); - await Promise.all(ttlStores.map(store => clearExpiredFromCache(store))); + await Promise.all(ttlStores.map((store) => clearExpiredFromCache(store))); // Force garbage collection if available (Node.js with --expose-gc flag) if (global.gc) { @@ -188,7 +193,7 @@ async function clearAllExpiredFromCache() { } } -if (!isRedisEnabled) { +if (!isRedisEnabled && !isEnabled(CI)) { /** @type {Set} */ const cleanupIntervals = new Set(); @@ -221,7 +226,7 @@ if (!isRedisEnabled) { const dispose = () => { debugMemoryCache && console.log('[Cache] Cleaning up and shutting down...'); - cleanupIntervals.forEach(interval => clearInterval(interval)); + cleanupIntervals.forEach((interval) => clearInterval(interval)); cleanupIntervals.clear(); // One final cleanup before exit diff --git a/api/server/routes/__tests__/config.spec.js b/api/server/routes/__tests__/config.spec.js index a19919b31c..13af53f299 100644 --- a/api/server/routes/__tests__/config.spec.js +++ b/api/server/routes/__tests__/config.spec.js @@ -1,3 +1,4 @@ +jest.mock('~/cache/getLogStores'); const request = require('supertest'); const express = require('express'); const routes = require('../'); diff --git a/api/server/services/Endpoints/gptPlugins/initialize.spec.js b/api/server/services/Endpoints/gptPlugins/initialize.spec.js index 54dfffc795..02199c9397 100644 --- a/api/server/services/Endpoints/gptPlugins/initialize.spec.js +++ b/api/server/services/Endpoints/gptPlugins/initialize.spec.js @@ -1,4 +1,5 @@ // gptPlugins/initializeClient.spec.js +jest.mock('~/cache/getLogStores'); const { EModelEndpoint, ErrorTypes, validateAzureGroups } = require('librechat-data-provider'); const { getUserKey, getUserKeyValues } = require('~/server/services/UserService'); const initializeClient = require('./initialize'); diff --git a/api/server/services/Endpoints/openAI/initialize.spec.js b/api/server/services/Endpoints/openAI/initialize.spec.js index b1a702e995..16563e4b26 100644 --- a/api/server/services/Endpoints/openAI/initialize.spec.js +++ b/api/server/services/Endpoints/openAI/initialize.spec.js @@ -1,3 +1,4 @@ +jest.mock('~/cache/getLogStores'); const { EModelEndpoint, ErrorTypes, validateAzureGroups } = require('librechat-data-provider'); const { getUserKey, getUserKeyValues } = require('~/server/services/UserService'); const initializeClient = require('./initialize'); diff --git a/api/server/services/Tokenizer.js b/api/server/services/Tokenizer.js index b88d5f8856..ac620defa8 100644 --- a/api/server/services/Tokenizer.js +++ b/api/server/services/Tokenizer.js @@ -59,6 +59,6 @@ class Tokenizer { } } -const tokenizerService = new Tokenizer(); +const TokenizerSingleton = new Tokenizer(); -module.exports = tokenizerService; +module.exports = TokenizerSingleton; diff --git a/api/server/services/Tokenizer.spec.js b/api/server/services/Tokenizer.spec.js new file mode 100644 index 0000000000..2f93489dcc --- /dev/null +++ b/api/server/services/Tokenizer.spec.js @@ -0,0 +1,136 @@ +/** + * @file Tokenizer.spec.cjs + * + * Tests the real TokenizerSingleton (no mocking of `tiktoken`). + * Make sure to install `tiktoken` and have it configured properly. + */ + +const Tokenizer = require('./Tokenizer'); // <-- Adjust path to your singleton file +const { logger } = require('~/config'); + +describe('Tokenizer', () => { + it('should be a singleton (same instance)', () => { + const AnotherTokenizer = require('./Tokenizer'); // same path + expect(Tokenizer).toBe(AnotherTokenizer); + }); + + describe('getTokenizer', () => { + it('should create an encoder for an explicit model name (e.g., "gpt-4")', () => { + // The real `encoding_for_model` will be called internally + // as soon as we pass isModelName = true. + const tokenizer = Tokenizer.getTokenizer('gpt-4', true); + + // Basic sanity checks + expect(tokenizer).toBeDefined(); + // You can optionally check certain properties from `tiktoken` if they exist + // e.g., expect(typeof tokenizer.encode).toBe('function'); + }); + + it('should create an encoder for a known encoding (e.g., "cl100k_base")', () => { + // The real `get_encoding` will be called internally + // as soon as we pass isModelName = false. + const tokenizer = Tokenizer.getTokenizer('cl100k_base', false); + + expect(tokenizer).toBeDefined(); + // e.g., expect(typeof tokenizer.encode).toBe('function'); + }); + + it('should return cached tokenizer if previously fetched', () => { + const tokenizer1 = Tokenizer.getTokenizer('cl100k_base', false); + const tokenizer2 = Tokenizer.getTokenizer('cl100k_base', false); + // Should be the exact same instance from the cache + expect(tokenizer1).toBe(tokenizer2); + }); + }); + + describe('freeAndResetAllEncoders', () => { + beforeEach(() => { + jest.clearAllMocks(); + }); + + it('should free all encoders and reset tokenizerCallsCount to 1', () => { + // By creating two different encodings, we populate the cache + Tokenizer.getTokenizer('cl100k_base', false); + Tokenizer.getTokenizer('r50k_base', false); + + // Now free them + Tokenizer.freeAndResetAllEncoders(); + + // The internal cache is cleared + expect(Tokenizer.tokenizersCache['cl100k_base']).toBeUndefined(); + expect(Tokenizer.tokenizersCache['r50k_base']).toBeUndefined(); + + // tokenizerCallsCount is reset to 1 + expect(Tokenizer.tokenizerCallsCount).toBe(1); + }); + + it('should catch and log errors if freeing fails', () => { + // Mock logger.error before the test + const mockLoggerError = jest.spyOn(logger, 'error'); + + // Set up a problematic tokenizer in the cache + Tokenizer.tokenizersCache['cl100k_base'] = { + free() { + throw new Error('Intentional free error'); + }, + }; + + // Should not throw uncaught errors + Tokenizer.freeAndResetAllEncoders(); + + // Verify logger.error was called with correct arguments + expect(mockLoggerError).toHaveBeenCalledWith( + '[Tokenizer] Free and reset encoders error', + expect.any(Error), + ); + + // Clean up + mockLoggerError.mockRestore(); + Tokenizer.tokenizersCache = {}; + }); + }); + + describe('getTokenCount', () => { + beforeEach(() => { + jest.clearAllMocks(); + Tokenizer.freeAndResetAllEncoders(); + }); + + it('should return the number of tokens in the given text', () => { + const text = 'Hello, world!'; + const count = Tokenizer.getTokenCount(text, 'cl100k_base'); + expect(count).toBeGreaterThan(0); + }); + + it('should reset encoders if an error is thrown', () => { + // We can simulate an error by temporarily overriding the selected tokenizer’s `encode` method. + const tokenizer = Tokenizer.getTokenizer('cl100k_base', false); + const originalEncode = tokenizer.encode; + tokenizer.encode = () => { + throw new Error('Forced error'); + }; + + // Despite the forced error, the code should catch and reset, then re-encode + const count = Tokenizer.getTokenCount('Hello again', 'cl100k_base'); + expect(count).toBeGreaterThan(0); + + // Restore the original encode + tokenizer.encode = originalEncode; + }); + + it('should reset tokenizers after 25 calls', () => { + // Spy on freeAndResetAllEncoders + const resetSpy = jest.spyOn(Tokenizer, 'freeAndResetAllEncoders'); + + // Make 24 calls; should NOT reset yet + for (let i = 0; i < 24; i++) { + Tokenizer.getTokenCount('test text', 'cl100k_base'); + } + expect(resetSpy).not.toHaveBeenCalled(); + + // 25th call triggers the reset + Tokenizer.getTokenCount('the 25th call!', 'cl100k_base'); + expect(resetSpy).toHaveBeenCalledTimes(1); + }); + }); +}); diff --git a/api/test/jestSetup.js b/api/test/jestSetup.js index cc6f61177c..f84b90743a 100644 --- a/api/test/jestSetup.js +++ b/api/test/jestSetup.js @@ -5,3 +5,4 @@ process.env.MONGO_URI = 'mongodb://127.0.0.1:27017/dummy-uri'; process.env.BAN_VIOLATIONS = 'true'; process.env.BAN_DURATION = '7200000'; process.env.BAN_INTERVAL = '20'; +process.env.CI = 'true';