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
206 changes: 190 additions & 16 deletions packages/desktop/src/common/chat/imageGenCore.ts
Original file line number Diff line number Diff line change
Expand Up @@ -15,11 +15,15 @@
import { jsonrepair } from 'jsonrepair';
import type OpenAI from 'openai';
import { ClientFactory, type RotatingClient } from '@/common/api/ClientFactory';
import type { OpenAIRotatingClient } from '@/common/api/OpenAIRotatingClient';
import type { OpenAIChatCompletionParams } from '@/common/api/OpenAI2GeminiConverter';
import type { TProviderWithModel } from '@/common/config/storage';
import type { UnifiedChatCompletionResponse } from '@/common/api/RotatingApiClient';
import { IMAGE_EXTENSIONS, MIME_TYPE_MAP, MIME_TO_EXT_MAP, DEFAULT_IMAGE_EXTENSION } from '@/common/config/constants';
import { resolveImageGenerationApiMode } from '@/common/utils/imageModelAllowlist';

const API_TIMEOUT_MS = 120000; // 2 minutes for image generation API calls
const API_TIMEOUT_MS = 180000; // High-quality image renders can take longer than ordinary chat calls
const MAX_IMAGE_DOWNLOAD_BYTES = 64 * 1024 * 1024;

type ImageExtension = (typeof IMAGE_EXTENSIONS)[number];

Expand Down Expand Up @@ -115,14 +119,59 @@
return DEFAULT_IMAGE_EXTENSION;
}

export async function saveGeneratedImage(base64Data: string, workspaceDir: string): Promise<string> {
async function downloadGeneratedImage(url: string, workspaceDir: string, signal?: AbortSignal): Promise<string> {
const response = await fetch(url, { signal });
if (!response.ok) {
throw new Error(`Failed to download generated image: HTTP ${response.status} ${response.statusText}`);
}

const declaredLength = Number(response.headers.get('content-length'));
if (Number.isFinite(declaredLength) && declaredLength > MAX_IMAGE_DOWNLOAD_BYTES) {
throw new Error(`Generated image exceeds the ${MAX_IMAGE_DOWNLOAD_BYTES / 1024 / 1024} MB download limit.`);
}

const contentType = (response.headers.get('content-type') || '').split(';')[0].trim().toLowerCase();
if (contentType && !contentType.startsWith('image/') && contentType !== 'application/octet-stream') {
throw new Error(`Generated image URL returned unsupported content type: ${contentType}`);
}

const imageBuffer = Buffer.from(await response.arrayBuffer());
if (imageBuffer.byteLength > MAX_IMAGE_DOWNLOAD_BYTES) {
throw new Error(`Generated image exceeds the ${MAX_IMAGE_DOWNLOAD_BYTES / 1024 / 1024} MB download limit.`);
}

let fileExtension = contentType.startsWith('image/')
? MIME_TO_EXT_MAP[contentType.slice('image/'.length)]
: undefined;
if (!fileExtension) {
const urlExtension = path.extname(new URL(url).pathname).toLowerCase();
if (IMAGE_EXTENSIONS.includes(urlExtension as ImageExtension)) {
fileExtension = urlExtension;
}
}

const fileName = `img-${Date.now()}${fileExtension || DEFAULT_IMAGE_EXTENSION}`;
const filePath = path.join(path.resolve(workspaceDir), fileName);
await fs.promises.writeFile(filePath, imageBuffer);
return filePath;
}

export async function saveGeneratedImage(
imageSource: string,
workspaceDir: string,
signal?: AbortSignal
): Promise<string> {
if (isHttpUrl(imageSource)) {
return await downloadGeneratedImage(imageSource, workspaceDir, signal);
}

const timestamp = Date.now();
const fileExtension = getFileExtensionFromDataUrl(base64Data);
const fileExtension = getFileExtensionFromDataUrl(imageSource);
const file_name = `img-${timestamp}${fileExtension}`;
const resolvedDir = path.resolve(workspaceDir);
const file_path = path.join(resolvedDir, file_name);

const base64WithoutPrefix = base64Data.replace(/^data:image\/[^;]+;base64,/, '');
const base64WithoutPrefix = imageSource.replace(/^data:image\/[^;]+;base64,/, '');
const imageBuffer = Buffer.from(base64WithoutPrefix, 'base64');

try {
Expand Down Expand Up @@ -208,6 +257,97 @@
error?: string;
}

type OpenAIImageClient = RotatingClient & Pick<OpenAIRotatingClient, 'createImage'>;
type ImageResultItem = { b64_json?: string; url?: string; revised_prompt?: string };

type ImagesApiResponse = OpenAI.Images.ImagesResponse & {
images?: ImageResultItem[];
output_format?: string;
};

function supportsOpenAIImagesApi(client: RotatingClient): client is OpenAIImageClient {
return 'createImage' in client && typeof client.createImage === 'function';
}

function extractImageResultItems(response: ImagesApiResponse): ImageResultItem[] {
if (Array.isArray(response.data) && response.data.length > 0) {
return response.data;
}
return Array.isArray(response.images) ? response.images : [];
}

function isMicrosoftMaiProvider(provider: TProviderWithModel): boolean {
const baseUrl = provider.base_url.toLowerCase();
return (
baseUrl.includes('services.ai.azure.com/mai/v1') ||
(baseUrl.includes('services.ai.azure.com') && /^mai[-_/ ]?image/i.test(provider.use_model))
);
}

function ensureVersionedImagesBaseUrl(provider: TProviderWithModel): string {
const trimmed = provider.base_url.replace(/\/+$/, '');
if (isMicrosoftMaiProvider(provider) && !/\/mai\/v1$/i.test(trimmed)) {
return `${trimmed}/mai/v1`;
}
if (!trimmed || /\/v\d+(?:beta)?$/i.test(trimmed) || /\/openai\/deployments\//i.test(trimmed)) {
return trimmed;
}
return `${trimmed}/v1`;
}

async function executeOpenAIImagesGeneration(
prompt: string,
provider: TProviderWithModel,
rotatingClient: RotatingClient,
workspaceDir: string,
signal?: AbortSignal
): Promise<ImageGenResult> {
if (!supportsOpenAIImagesApi(rotatingClient)) {
return {
success: false,
text: `Model ${provider.use_model} requires an OpenAI-compatible Images API provider.`,
error: 'OpenAI Images API is not available for the selected provider.',
};
}

const generationParams: Record<string, unknown> = { model: provider.use_model, prompt };
if (isMicrosoftMaiProvider(provider)) {
generationParams.width = 1024;
generationParams.height = 1024;
}

const response = (await rotatingClient.createImage(generationParams as unknown as OpenAI.Images.ImageGenerateParams, {
signal,
timeout: API_TIMEOUT_MS,
})) as ImagesApiResponse;
const image = extractImageResultItems(response)[0];
if (!image?.b64_json && !image?.url) {
return {
success: false,
text: 'Image generation API did not return image data or an image URL.',
error: 'No image data returned.',
};
}

let imagePath: string;
if (image.b64_json) {
const outputFormat = String(response.output_format || 'png');
const mimeSubtype = outputFormat === 'jpg' ? 'jpeg' : outputFormat;
imagePath = await saveGeneratedImage(`data:image/${mimeSubtype};base64,${image.b64_json}`, workspaceDir, signal);
} else {
imagePath = await saveGeneratedImage(image.url!, workspaceDir, signal);
}
const relativeImagePath = path.relative(workspaceDir, imagePath);
const revisedPrompt = image.revised_prompt ? `\n\nRevised prompt: ${image.revised_prompt}` : '';

return {
success: true,
text: `Image generated successfully.${revisedPrompt}\n\nGenerated image saved to: ${imagePath}`,
imagePath,
relativeImagePath,
};
}

/**
* Core image generation function shared between MCP server and Gemini tool.
*/
Expand Down Expand Up @@ -257,14 +397,41 @@
}

const hasImages = imageUris.length > 0;
const apiMode = resolveImageGenerationApiMode(provider, provider.use_model) || 'chat-completions';
if (apiMode === 'openai-images' && hasImages) {
return {
success: false,
text: `Image editing is not yet supported for ${provider.use_model}. Generate a new image without image_uris instead.`,
error: 'OpenAI Images API editing is not implemented.',
};
}

const clientOptions = {
proxy,
rotatingOptions: { maxRetries: 3, retryDelay: 1000 },
...(apiMode === 'openai-images'
? {
baseConfig: {
baseURL: ensureVersionedImagesBaseUrl(provider),
...(isMicrosoftMaiProvider(provider) ? { defaultHeaders: { 'api-key': provider.api_key } } : {}),
},
}
: {}),
};
const rotatingClient: RotatingClient = await ClientFactory.createRotatingClient(provider, clientOptions);

if (apiMode === 'openai-images') {
return await executeOpenAIImagesGeneration(params.prompt, provider, rotatingClient, resolvedWorkspaceDir, signal);
}

let enhancedPrompt: string;
if (hasImages) {
enhancedPrompt = `Analyze/Edit image: ${params.prompt}`;
} else {
enhancedPrompt = `Generate image: ${params.prompt}`;
}

const contentParts: OpenAI.Chat.Completions.ChatCompletionContentPart[] = [{ type: 'text', text: enhancedPrompt }];
const contentParts: Array<{ type: 'text'; text: string } | ImageContent> = [{ type: 'text', text: enhancedPrompt }];

// Process image URIs
if (hasImages) {
Expand Down Expand Up @@ -294,17 +461,24 @@
}
}

const messages: OpenAI.Chat.Completions.ChatCompletionMessageParam[] = [{ role: 'user', content: contentParts }];

// Create client and call API
const rotatingClient: RotatingClient = await ClientFactory.createRotatingClient(provider, {
proxy,
rotatingOptions: { maxRetries: 3, retryDelay: 1000 },
});

const messages = [{ role: 'user' as const, content: contentParts }];
const completionParams: OpenAIChatCompletionParams &
Omit<OpenAI.Chat.Completions.ChatCompletionCreateParamsNonStreaming, 'modalities'> & {
modalities?: Array<'image' | 'text'>;
} = {
model: provider.use_model,
messages,
...(provider.base_url.toLowerCase().includes('openrouter.ai') ? { modalities: ['image', 'text'] } : {}),
};
const completion: UnifiedChatCompletionResponse = await rotatingClient.createChatCompletion(
{ model: provider.use_model, messages: messages as any },
{ signal, timeout: API_TIMEOUT_MS }
// OpenRouter extends OpenAI's modalities union with `image`; the OpenAI
// SDK type only declares `text | audio`, although it forwards this field.
completionParams as unknown as OpenAIChatCompletionParams &
OpenAI.Chat.Completions.ChatCompletionCreateParamsNonStreaming,
{
signal,
timeout: API_TIMEOUT_MS,
}
);

const choice = completion.choices[0];
Expand Down Expand Up @@ -358,7 +532,7 @@

const firstImage = images[0];
if (firstImage.type === 'image_url' && firstImage.image_url?.url) {
const imagePath = await saveGeneratedImage(firstImage.image_url.url, resolvedWorkspaceDir);
const imagePath = await saveGeneratedImage(firstImage.image_url.url, resolvedWorkspaceDir, signal);
const relativeImagePath = path.relative(resolvedWorkspaceDir, imagePath);

// Strip any inline base64 data URLs from the human-readable text before
Expand Down
109 changes: 93 additions & 16 deletions packages/desktop/src/common/utils/imageModelAllowlist.ts
Original file line number Diff line number Diff line change
Expand Up @@ -5,29 +5,69 @@
*/

/**
* Allowlist for built-in image generation tool.
* Capability routing for the built-in image generation tool.
*
* The tool currently only supports "form B" — OpenAI chat completions multimodal
* output (model returns images via `message.images` or markdown). It does NOT
* support "form A" (`/v1/images/generations` endpoint) or async/polling APIs.
*
* Model selection therefore must be a platform+model allowlist of providers
* known to work, rather than a coarse name-substring match. Otherwise users
* see options like `gpt-image-1` / `dall-e-3` / `sd-3.5` in the dropdown that
* are guaranteed to fail at runtime.
*
* Rules below mirror `useConfigModelListWithImage.ts` — the same providers we
* auto-supplement with default image models. When #6 lands a form-A adapter,
* extend this list accordingly.
* Chat-style providers return images from a chat completion (`message.images`
* or markdown). Dedicated image models use the OpenAI-compatible Images API.
* Keep this routing model-specific: a provider can expose both ordinary chat
* models and image models, and provider names alone are not sufficient.
*/

import { AuthType, type AuthType as ProviderAuthType } from '@/common/types/provider/authType';
import { getProviderAuthType } from '@/common/utils/platformAuthType';

type ProviderShape = {
platform?: string;
base_url?: string;
name?: string;
auth_type?: ProviderAuthType;
model_protocols?: Record<string, string>;
};

const IMAGE_NAME_PATTERN = /(image|banana|imagine)/i;
const OPENAI_IMAGES_MODEL_PATTERNS = [
/^gpt-image/i,
/^dall-e-[23]/i,
/^grok[-_/ ]?imagine[-_/ ]?image/i,
/seedream/i,
/flux/i,
/(stable[-_/ ]?(?:diffusion|image)|(?:^|[-_/])sd(?:xl|[-_.]?3(?:[-_.]?5)?)(?:$|[-_/]))/i,
/recraft/i,
/qwen[-_/ ]?image/i,
/(?:^|[-_/])z[-_/ ]?image/i,
/ideogram/i,
/(?:^|[-_/])mai[-_/ ]?image/i,
/hidream/i,
/cogview/i,
/hunyuan[-_/ ]?image/i,
/kolors/i,
/^krea(?:[-_/ ]|$)/i,
/^reve(?:[-_/ ]|$)/i,
/firefly[-_/ ]?image/i,
/midjourney/i,
];

/**
* Official endpoints on these hosts use a vendor-specific payload, auth flow,
* or submit-and-poll job protocol. The same model families can still be used
* when they are exposed by an OpenAI-compatible gateway.
*/
const NATIVE_IMAGE_API_HOST_MARKERS = [
'api.stability.ai',
'api.replicate.com',
'fal.run',
'api.bfl.ai',
'api.bfl.ml',
'api.ideogram.ai',
'dashscope.aliyuncs.com',
'dashscope-intl.aliyuncs.com',
'api.dev.runwayml.com',
'firefly-api.adobe.io',
'cloud.leonardo.ai',
'api.midjourney.com',
];

type ImageGenerationApiMode = 'chat-completions' | 'openai-images';

const RULES: Array<{
id: string;
Expand All @@ -47,7 +87,44 @@ const RULES: Array<{
},
];

export const isImageGenSupported = (provider: ProviderShape, modelName: string): boolean => {
if (!IMAGE_NAME_PATTERN.test(modelName)) return false;
return RULES.some((rule) => rule.match(provider));
const includesIgnoreCase = (value: string | undefined, markers: string[]): boolean => {
const lowerValue = value?.toLowerCase();
return !!lowerValue && markers.some((marker) => lowerValue.includes(marker.toLowerCase()));
};

const isOpenAIImagesModel = (modelName: string): boolean =>
OPENAI_IMAGES_MODEL_PATTERNS.some((pattern) => pattern.test(modelName));

const isMicrosoftMaiImagesEndpoint = (baseUrl: string | undefined): boolean =>
includesIgnoreCase(baseUrl, ['services.ai.azure.com/mai/v1']);

export const resolveImageGenerationApiMode = (
provider: ProviderShape,
modelName: string
): ImageGenerationApiMode | undefined => {
if (
(IMAGE_NAME_PATTERN.test(modelName) || isOpenAIImagesModel(modelName)) &&
RULES.some((rule) => rule.match(provider))
) {
return 'chat-completions';
}

const authType = getProviderAuthType({
platform: provider.platform || 'custom',
auth_type: provider.auth_type,
model_protocols: provider.model_protocols,
use_model: modelName,
});
if (
(isOpenAIImagesModel(modelName) || isMicrosoftMaiImagesEndpoint(provider.base_url)) &&
authType === AuthType.USE_OPENAI &&
!includesIgnoreCase(provider.base_url, NATIVE_IMAGE_API_HOST_MARKERS)
) {
return 'openai-images';
}

return undefined;
};

export const isImageGenSupported = (provider: ProviderShape, modelName: string): boolean =>
resolveImageGenerationApiMode(provider, modelName) !== undefined;
Loading
Loading