Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions .env.example
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
11 changes: 2 additions & 9 deletions api/app/clients/tools/structured/OpenAIImageTools.js
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@ const {
extractBaseURL,
getProxyDispatcher,
applyAxiosProxyConfig,
getImageGenClientOptions,
} = require('@librechat/api');
const { getStrategyFunctions } = require('~/server/services/Files/strategies');
const { getFiles } = require('~/models');
Expand Down Expand Up @@ -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' &&
Expand Down
58 changes: 58 additions & 0 deletions api/test/app/clients/tools/structured/OpenAIImageTools.test.js
Original file line number Diff line number Diff line change
Expand Up @@ -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', () => ({
Expand Down Expand Up @@ -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');
});
});
2 changes: 2 additions & 0 deletions librechat.example.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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}'
Expand Down
48 changes: 47 additions & 1 deletion packages/api/src/tools/toolkits/oai.spec.ts
Original file line number Diff line number Diff line change
@@ -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],
Expand Down Expand Up @@ -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({});
});
});
39 changes: 39 additions & 0 deletions packages/api/src/tools/toolkits/oai.ts
Original file line number Diff line number Diff line change
@@ -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 =
Expand Down Expand Up @@ -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;
}
29 changes: 29 additions & 0 deletions packages/api/src/web/web.spec.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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'],
Expand Down Expand Up @@ -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<string, string> = {};
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', () => {
Expand Down
1 change: 1 addition & 0 deletions packages/api/src/web/web.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down
7 changes: 7 additions & 0 deletions packages/data-provider/src/config.spec.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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: {
Expand Down
1 change: 1 addition & 0 deletions packages/data-provider/src/config.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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({
Expand Down