diff --git a/.env.example b/.env.example index 4146dc3e5a7..3b6ca1b4566 100644 --- a/.env.example +++ b/.env.example @@ -671,6 +671,8 @@ AZURE_AI_SEARCH_SEARCH_OPTION_SELECT= # IMAGE_GEN_OAI_BASEURL= # Custom OpenAI base URL for image generation tool # IMAGE_GEN_OAI_AZURE_API_VERSION= # Custom Azure OpenAI deployments # IMAGE_GEN_OAI_MODEL=gpt-image-1 # OpenAI image model (e.g., gpt-image-1, gpt-image-1.5) +# IMAGE_GEN_OAI_TIMEOUT_MS=600000 # Image generation request timeout in milliseconds +# IMAGE_GEN_OAI_MAX_RETRIES=2 # Image generation request retries # IMAGE_GEN_OAI_DESCRIPTION= # IMAGE_GEN_OAI_DESCRIPTION_WITH_FILES=Custom description for image generation tool when files are present # IMAGE_GEN_OAI_DESCRIPTION_NO_FILES=Custom description for image generation tool when no files are present diff --git a/api/app/clients/tools/structured/OpenAIImageTools.js b/api/app/clients/tools/structured/OpenAIImageTools.js index d92d17b77e6..e13dab4a43f 100644 --- a/api/app/clients/tools/structured/OpenAIImageTools.js +++ b/api/app/clients/tools/structured/OpenAIImageTools.js @@ -11,6 +11,7 @@ const { extractBaseURL, getProxyDispatcher, applyAxiosProxyConfig, + getImageGenClientOptions, } = require('@librechat/api'); const { getStrategyFunctions } = require('~/server/services/Files/strategies'); const { getFiles } = require('~/models'); @@ -126,16 +127,8 @@ function createOpenAIImageTools(fields = {}) { if (!prompt) { throw new Error('Missing required field: prompt'); } - const clientConfig = { ...closureConfig }; - const proxyDispatcher = getProxyDispatcher(); - if (proxyDispatcher) { - clientConfig.fetchOptions = { - dispatcher: proxyDispatcher, - }; - } - /** @type {OpenAI} */ - const openai = new OpenAI(clientConfig); + const openai = new OpenAI({ ...closureConfig, ...getImageGenClientOptions() }); let output_format = imageOutputType; if ( background === 'transparent' && diff --git a/api/test/app/clients/tools/structured/OpenAIImageTools.test.js b/api/test/app/clients/tools/structured/OpenAIImageTools.test.js index b83ed5335c6..7d11729c89d 100644 --- a/api/test/app/clients/tools/structured/OpenAIImageTools.test.js +++ b/api/test/app/clients/tools/structured/OpenAIImageTools.test.js @@ -27,6 +27,7 @@ jest.mock('@librechat/api', () => ({ extractBaseURL: jest.fn((url) => url), getProxyDispatcher: jest.fn(() => undefined), applyAxiosProxyConfig: jest.fn(), + getImageGenClientOptions: jest.requireActual('@librechat/api').getImageGenClientOptions, })); jest.mock('~/server/services/Files/strategies', () => ({ @@ -162,3 +163,60 @@ describe('OpenAIImageTools - IMAGE_GEN_OAI_MODEL environment variable', () => { ); }); }); + +describe('OpenAIImageTools - client timeout and retries', () => { + let originalEnv; + + beforeEach(() => { + jest.clearAllMocks(); + originalEnv = { ...process.env }; + + process.env.IMAGE_GEN_OAI_API_KEY = 'test-api-key'; + delete process.env.IMAGE_GEN_OAI_TIMEOUT_MS; + delete process.env.IMAGE_GEN_OAI_MAX_RETRIES; + + OpenAI.mockImplementation(() => ({ + images: { + generate: jest.fn().mockResolvedValue({ + data: [{ b64_json: 'base64-encoded-image-data' }], + }), + }, + })); + }); + + afterEach(() => { + process.env = originalEnv; + }); + + const createTools = () => + createOpenAIImageTools({ + isAgent: true, + override: false, + req: { user: { id: 'test-user' } }, + }); + + it('should pass IMAGE_GEN_OAI_TIMEOUT_MS and IMAGE_GEN_OAI_MAX_RETRIES to the OpenAI client', async () => { + process.env.IMAGE_GEN_OAI_TIMEOUT_MS = '1800000'; + process.env.IMAGE_GEN_OAI_MAX_RETRIES = '0'; + + const [imageGenTool] = createTools(); + await imageGenTool.func({ prompt: 'test prompt' }); + + expect(OpenAI).toHaveBeenCalledWith( + expect.objectContaining({ + timeout: 1800000, + maxRetries: 0, + fetchOptions: { dispatcher: expect.anything() }, + }), + ); + }); + + it('should keep the OpenAI client defaults when the env vars are unset', async () => { + const [imageGenTool] = createTools(); + await imageGenTool.func({ prompt: 'test prompt' }); + + const [clientConfig] = OpenAI.mock.calls[0]; + expect(clientConfig).not.toHaveProperty('timeout'); + expect(clientConfig).not.toHaveProperty('maxRetries'); + }); +}); diff --git a/librechat.example.yaml b/librechat.example.yaml index 982f7fd9d1c..bdbb0249a87 100644 --- a/librechat.example.yaml +++ b/librechat.example.yaml @@ -1415,6 +1415,8 @@ endpoints: # jinaApiUrl: '${JINA_API_URL}' # Custom Jina API URL (optional, defaults to https://api.jina.ai/v1/rerank) # # Other rerankers # cohereApiKey: '${COHERE_API_KEY}' +# # Rerank request timeout in milliseconds, applied per page (default: 10000) +# rerankerTimeout: 10000 # # Search providers # serperApiKey: '${SERPER_API_KEY}' # searxngInstanceUrl: '${SEARXNG_INSTANCE_URL}' diff --git a/packages/api/src/tools/toolkits/oai.spec.ts b/packages/api/src/tools/toolkits/oai.spec.ts index e82975d4226..fca00a1319e 100644 --- a/packages/api/src/tools/toolkits/oai.spec.ts +++ b/packages/api/src/tools/toolkits/oai.spec.ts @@ -1,4 +1,5 @@ -import { oaiToolkit, IMAGE_SIZE_PATTERN } from './oai'; +import { oaiToolkit, IMAGE_SIZE_PATTERN, getImageGenClientOptions } from './oai'; +import { getDirectDispatcher, getProxyDispatcher } from '~/utils/proxy'; const sizeSchemas = [ ['image_gen_oai', oaiToolkit.image_gen_oai.schema.properties?.size], @@ -34,3 +35,48 @@ describe('OpenAI image toolkit size schema', () => { }, ); }); + +describe('getImageGenClientOptions', () => { + const originalEnv = process.env; + const transport = { headersTimeout: 1800000, bodyTimeout: 1800000 }; + + afterEach(() => { + process.env = originalEnv; + }); + + it('returns no options when nothing is configured', () => { + process.env = {}; + expect(getImageGenClientOptions()).toEqual({}); + }); + + it('sets the timeout and retries, with matching undici timeouts on the dispatcher', () => { + process.env = { IMAGE_GEN_OAI_TIMEOUT_MS: '1800000', IMAGE_GEN_OAI_MAX_RETRIES: '0' }; + const options = getImageGenClientOptions(); + expect(options).toMatchObject({ timeout: 1800000, maxRetries: 0 }); + expect(options.fetchOptions?.dispatcher).toBe(getDirectDispatcher(transport)); + }); + + it('routes through the proxy dispatcher when a proxy is configured', () => { + process.env = { PROXY: 'http://proxy.test:8080' }; + expect(getImageGenClientOptions().fetchOptions?.dispatcher).toBe(getProxyDispatcher()); + + process.env.IMAGE_GEN_OAI_TIMEOUT_MS = '1800000'; + expect(getImageGenClientOptions().fetchOptions?.dispatcher).toBe( + getProxyDispatcher(undefined, transport), + ); + }); + + it.each([ + ['IMAGE_GEN_OAI_TIMEOUT_MS', ''], + ['IMAGE_GEN_OAI_TIMEOUT_MS', '0'], + ['IMAGE_GEN_OAI_TIMEOUT_MS', '-1'], + ['IMAGE_GEN_OAI_TIMEOUT_MS', '1.5'], + ['IMAGE_GEN_OAI_TIMEOUT_MS', '10s'], + ['IMAGE_GEN_OAI_MAX_RETRIES', ''], + ['IMAGE_GEN_OAI_MAX_RETRIES', '-1'], + ['IMAGE_GEN_OAI_MAX_RETRIES', 'two'], + ])('ignores %s=%j', (key, value) => { + process.env = { [key]: value }; + expect(getImageGenClientOptions()).toEqual({}); + }); +}); diff --git a/packages/api/src/tools/toolkits/oai.ts b/packages/api/src/tools/toolkits/oai.ts index 19fd2eec547..3d8451502bd 100644 --- a/packages/api/src/tools/toolkits/oai.ts +++ b/packages/api/src/tools/toolkits/oai.ts @@ -1,4 +1,6 @@ +import type { Dispatcher } from 'undici'; import type { ExtendedJsonSchema } from '../registry/schema'; +import { getDirectDispatcher, getProxyDispatcher } from '~/utils/proxy'; /** Default descriptions for image generation tool */ const DEFAULT_IMAGE_GEN_DESCRIPTION = @@ -161,3 +163,40 @@ export const oaiToolkit: { responseFormat: 'content_and_artifact' as const, }, } as const; + +type ImageGenClientOptions = { + timeout?: number; + maxRetries?: number; + fetchOptions?: { dispatcher: Dispatcher }; +}; + +const parseInteger = (value?: string): number | undefined => { + if (value == null || value.trim() === '') { + return undefined; + } + const parsed = Number(value); + return Number.isSafeInteger(parsed) ? parsed : undefined; +}; + +/** OpenAI client options for `image_gen_oai`; unset or invalid values keep the client defaults. */ +export function getImageGenClientOptions(): ImageGenClientOptions { + const options: ImageGenClientOptions = {}; + const maxRetries = parseInteger(process.env.IMAGE_GEN_OAI_MAX_RETRIES); + if (maxRetries != null && maxRetries >= 0) { + options.maxRetries = maxRetries; + } + + const timeout = parseInteger(process.env.IMAGE_GEN_OAI_TIMEOUT_MS); + if (timeout == null || timeout <= 0) { + const dispatcher = getProxyDispatcher(); + return dispatcher ? { ...options, fetchOptions: { dispatcher } } : options; + } + + /** undici's default 300 s headers/body timeouts would otherwise end a longer request first */ + const transport = { headersTimeout: timeout, bodyTimeout: timeout }; + options.timeout = timeout; + options.fetchOptions = { + dispatcher: getProxyDispatcher(undefined, transport) ?? getDirectDispatcher(transport), + }; + return options; +} diff --git a/packages/api/src/web/web.spec.ts b/packages/api/src/web/web.spec.ts index 1a8ae5ceda6..ebabab1bccd 100644 --- a/packages/api/src/web/web.spec.ts +++ b/packages/api/src/web/web.spec.ts @@ -2512,6 +2512,7 @@ describe('web.ts', () => { expect(result.authenticated).toBe(true); expect(result.authResult.scraperTimeout).toBe(7500); // Should use default timeout + expect(result.authResult.rerankerTimeout).toBeUndefined(); expect(result.authResult.firecrawlOptions).toEqual({ includeTags: ['p'], formats: ['markdown'], @@ -2734,6 +2735,34 @@ describe('web.ts', () => { expect(result.authResult.scraperTimeout).toBe(7500); // Should use default timeout expect(result.authResult.firecrawlOptions).toBeUndefined(); // Should be undefined }); + + it('should pass rerankerTimeout through to the auth result', async () => { + const webSearchConfig = { + serperApiKey: '${SERPER_API_KEY}', + firecrawlApiKey: '${FIRECRAWL_API_KEY}', + jinaApiKey: '${JINA_API_KEY}', + safeSearch: SafeSearchTypes.MODERATE, + rerankerTimeout: 30000, + } as TCustomConfig['webSearch']; + + mockLoadAuthValues.mockImplementation(({ authFields }) => { + const result: Record = {}; + authFields.forEach((field: string) => { + result[field] = 'test-api-key'; + }); + return Promise.resolve(result); + }); + + const result = await loadWebSearchAuth({ + userId, + webSearchConfig, + loadAuthValues: mockLoadAuthValues, + }); + + expect(result.authenticated).toBe(true); + expect(result.authResult.rerankerType).toBe(RerankerTypes.JINA); + expect(result.authResult.rerankerTimeout).toBe(30000); + }); }); describe('SSRF protection for user-provided URLs', () => { diff --git a/packages/api/src/web/web.ts b/packages/api/src/web/web.ts index ca3372a1c59..100d80b706b 100644 --- a/packages/api/src/web/web.ts +++ b/packages/api/src/web/web.ts @@ -656,6 +656,7 @@ export async function loadWebSearchAuth({ authResult.safeSearch = webSearchConfig?.safeSearch ?? SafeSearchTypes.MODERATE; } authResult.scraperTimeout = webSearchConfig?.scraperTimeout ?? scraperOptionsTimeout ?? 7500; + authResult.rerankerTimeout = webSearchConfig?.rerankerTimeout; authResult.firecrawlOptions = webSearchConfig?.firecrawlOptions; authResult.searxngSearchOptions = webSearchConfig?.searxngSearchOptions; authResult.tavilySearchOptions = webSearchConfig?.tavilySearchOptions; diff --git a/packages/data-provider/src/config.spec.ts b/packages/data-provider/src/config.spec.ts index 99fb5de8b05..0b2d7e0cb33 100644 --- a/packages/data-provider/src/config.spec.ts +++ b/packages/data-provider/src/config.spec.ts @@ -2052,6 +2052,13 @@ describe('allowedAddressesSchema', () => { }); describe('webSearchSchema', () => { + it('accepts a reranker timeout and rejects invalid values', () => { + expect(webSearchSchema.parse({ rerankerTimeout: 30000 }).rerankerTimeout).toBe(30000); + expect(webSearchSchema.parse({}).rerankerTimeout).toBeUndefined(); + expect(() => webSearchSchema.parse({ rerankerTimeout: -1 })).toThrow(); + expect(() => webSearchSchema.parse({ rerankerTimeout: 1.5 })).toThrow(); + }); + it('accepts Tavily string modes for answer and raw content options', () => { const result = webSearchSchema.parse({ tavilySearchOptions: { diff --git a/packages/data-provider/src/config.ts b/packages/data-provider/src/config.ts index 915bfbe7751..ac2bd99b254 100644 --- a/packages/data-provider/src/config.ts +++ b/packages/data-provider/src/config.ts @@ -3359,6 +3359,7 @@ export const webSearchSchema = z.object({ scraperProvider: z.nativeEnum(ScraperProviders).optional(), rerankerType: z.nativeEnum(RerankerTypes).optional(), scraperTimeout: z.number().int().nonnegative().optional(), + rerankerTimeout: z.number().int().nonnegative().optional(), safeSearch: z.nativeEnum(SafeSearchTypes).default(SafeSearchTypes.MODERATE), firecrawlOptions: z .object({