From 45b4f5bffd7a1d3e204dda7755be297f75fc48e1 Mon Sep 17 00:00:00 2001 From: wangguanghao Date: Mon, 7 Sep 2026 15:30:39 +0800 Subject: [PATCH 1/3] feat(models): support provider-owned image models --- .../src/bun/remote/remote-runtime-client.ts | 6 +- .../src/components/settings/models-page.tsx | 53 +++-- apps/desktop/src/shared/rpc.ts | 4 +- .../core/src/types/models/image-generation.ts | 64 ++++-- .../core/src/types/models/provider-group.ts | 6 +- .../src/models/ark-image-generation.ts | 186 +++++++++++------- packages/runtime/src/models/index.ts | 7 +- packages/runtime/src/models/model-manager.ts | 156 ++++++++------- packages/runtime/src/models/types.ts | 6 +- packages/runtime/src/runtime/model-groups.ts | 5 +- packages/runtime/src/runtime/types.ts | 12 +- packages/runtime/src/tools/built-in/media.ts | 3 +- packages/runtime/src/tools/tool-registry.ts | 3 +- .../tests/models/ark-image-generation.test.ts | 117 +++++++++++ .../tests/models/model-manager.test.ts | 117 +++++++++++ .../built-in/built-in-tools-module.test.ts | 2 +- packages/ui/src/components/model-provider.tsx | 8 +- .../tool/built-in-tool-import-dialog.tsx | 125 +++++++++--- packages/ui/src/host/types.ts | 4 +- .../tool/tool-executor.test.ts | 30 +++ 20 files changed, 692 insertions(+), 222 deletions(-) diff --git a/apps/desktop/src/bun/remote/remote-runtime-client.ts b/apps/desktop/src/bun/remote/remote-runtime-client.ts index 71b5c9e1..2d195786 100644 --- a/apps/desktop/src/bun/remote/remote-runtime-client.ts +++ b/apps/desktop/src/bun/remote/remote-runtime-client.ts @@ -1,5 +1,5 @@ import type { - ArkImageGenerationConfig, + ImageGenerationConfig, AgentEvent, BuiltinTool, CustomModel, @@ -135,7 +135,7 @@ export class RemoteRuntimeClient implements RuntimeClient { } createSubagentThread( - input: Parameters[0], + input: Parameters[0] ) { return this._rpc< Awaited> @@ -315,7 +315,7 @@ export class RemoteRuntimeClient implements RuntimeClient { api?: "anthropic-messages" | "openai-completions" | "openai-responses" | null; icon?: string | null; - imageGeneration?: ArkImageGenerationConfig; + imageGeneration?: ImageGenerationConfig; }) { return this._rpc("models.updateProvider", input); } diff --git a/apps/desktop/src/components/settings/models-page.tsx b/apps/desktop/src/components/settings/models-page.tsx index 0d7d6add..7feba426 100644 --- a/apps/desktop/src/components/settings/models-page.tsx +++ b/apps/desktop/src/components/settings/models-page.tsx @@ -3,8 +3,10 @@ import { formatProviderProfileLabel, getArkImageModelDefinitions, - type ArkImageGenerationConfig, + getImageModelDefinitions, type CustomModel, + type ImageGenerationApi, + type ImageGenerationConfig, type ModelProviderGroup, type ProviderProfile, type SeedreamImageModelDefinition, @@ -499,9 +501,7 @@ function ProviderListItem({ role="button" tabIndex={0} aria-label={`Select ${provider.name} provider`} - aria-expanded={ - provider.profiles.length > 1 ? expanded : undefined - } + aria-expanded={provider.profiles.length > 1 ? expanded : undefined} onClick={handleGroupClick} onKeyDown={(e) => { if (e.key === "Enter" || e.key === " ") { @@ -613,7 +613,7 @@ function ProviderListItem({ @@ -1073,8 +1073,8 @@ function ProviderEditor({ ) : null} - {provider.id === "ark" && canManageModels ? ( - <_ArkImageGenerationEditor provider={provider} /> + {(provider.id === "ark" || !isBuiltin) && canManageModels ? ( + <_ImageGenerationEditor provider={provider} /> ) : null} @@ -1173,8 +1173,8 @@ function _ProviderProfileEditor({ {isOfficial ? (
- Leave it blank to use the official {provider.name}{" "} - environment variable + Leave it blank to use the official {provider.name} environment + variable
) : null} @@ -1214,15 +1214,18 @@ function _ProviderProfileEditor({ ); } -/** Chat-model-parity inventory management for Ark image models. */ -function _ArkImageGenerationEditor({ +/** Chat-model-parity inventory management for provider-owned image models. */ +function _ImageGenerationEditor({ provider, }: { provider: ModelProviderGroup; }) { const updateProvider = useUpdateProvider(); const config = provider.imageGeneration ?? {}; - const models = getArkImageModelDefinitions(config); + const models = + provider.id === "ark" + ? getArkImageModelDefinitions(config) + : getImageModelDefinitions(config); const disabledModels = new Set(config.disabledModels ?? []); const enabledModels = models.filter((model) => !disabledModels.has(model.id)); const customModels = new Set((config.models ?? []).map((model) => model.id)); @@ -1240,7 +1243,7 @@ function _ArkImageGenerationEditor({ return true; }); - const update = (imageGeneration: ArkImageGenerationConfig) => { + const update = (imageGeneration: ImageGenerationConfig) => { void updateProvider(provider.id, { imageGeneration }).catch((error) => { toast.error("Failed to update image generation", { description: @@ -1307,6 +1310,30 @@ function _ArkImageGenerationEditor({ return ( <> + {provider.id !== "ark" ? ( +
+ Image API type + +
+ ) : null}
Image models diff --git a/apps/desktop/src/shared/rpc.ts b/apps/desktop/src/shared/rpc.ts index 15387bf9..ad734950 100644 --- a/apps/desktop/src/shared/rpc.ts +++ b/apps/desktop/src/shared/rpc.ts @@ -1,5 +1,5 @@ import type { - ArkImageGenerationConfig, + ImageGenerationConfig, AgentEvent, AgentStreamRequest, BuiltinTool, @@ -208,7 +208,7 @@ export interface DesktopRPCType { | "openai-responses" | null; icon?: string | null; - imageGeneration?: ArkImageGenerationConfig; + imageGeneration?: ImageGenerationConfig; }; response: ModelProviderGroup[]; }; diff --git a/packages/core/src/types/models/image-generation.ts b/packages/core/src/types/models/image-generation.ts index bdd58cf1..f33c8637 100644 --- a/packages/core/src/types/models/image-generation.ts +++ b/packages/core/src/types/models/image-generation.ts @@ -2,7 +2,7 @@ export const SEEDREAM_IMAGE_SIZES = ["1K", "2K", "3K", "4K"] as const; export type SeedreamImageSize = (typeof SEEDREAM_IMAGE_SIZES)[number]; -export interface SeedreamImageModelDefinition { +export interface ImageModelDefinition { id: string; name: string; supportedSizes: readonly SeedreamImageSize[]; @@ -11,6 +11,9 @@ export interface SeedreamImageModelDefinition { icon?: string; } +/** @deprecated Use ImageModelDefinition. */ +export type SeedreamImageModelDefinition = ImageModelDefinition; + /** Curated Ark Seedream catalog used by settings and runtime validation. */ export const SEEDREAM_IMAGE_MODELS = [ { @@ -37,17 +40,25 @@ export const SEEDREAM_IMAGE_MODELS = [ supportedSizes: ["1K", "2K", "4K"], defaultSize: "2K", }, -] as const satisfies readonly SeedreamImageModelDefinition[]; +] as const satisfies readonly ImageModelDefinition[]; export type SeedreamImageModelId = (typeof SEEDREAM_IMAGE_MODELS)[number]["id"]; -export interface ArkImageGenerationConfig { - /** User-added Ark image models layered on top of the curated catalog. */ - models?: SeedreamImageModelDefinition[]; +export type ImageGenerationApi = + "ark-images" | "openai-images" | "openai-images-extra-body"; + +export interface ImageGenerationConfig { + /** Request protocol. Ark defaults to its native API; custom providers use OpenAI Images. */ + api?: ImageGenerationApi; + /** User-added image models layered on top of an optional provider catalog. */ + models?: ImageModelDefinition[]; /** Image-model ids disabled in Settings. Absent means every model is enabled. */ disabledModels?: string[]; } +/** @deprecated Use ImageGenerationConfig. */ +export type ArkImageGenerationConfig = ImageGenerationConfig; + /** Per-Thread configuration owned by one `generate_image` tool instance. */ export interface GenerateImageToolConfig { model: string; @@ -55,7 +66,7 @@ export interface GenerateImageToolConfig { watermark: boolean; } -export const DEFAULT_ARK_IMAGE_GENERATION_CONFIG: ArkImageGenerationConfig = {}; +export const DEFAULT_ARK_IMAGE_GENERATION_CONFIG: ImageGenerationConfig = {}; /** Find one curated Seedream model definition by its stable Ark model id. */ export function getSeedreamImageModelDefinition( @@ -66,14 +77,22 @@ export function getSeedreamImageModelDefinition( /** Merge the curated Seedream catalog with user-added Ark image models. */ export function getArkImageModelDefinitions( - config: ArkImageGenerationConfig -): readonly SeedreamImageModelDefinition[] { - return [...SEEDREAM_IMAGE_MODELS, ...(config.models ?? [])]; + config: ImageGenerationConfig +): readonly ImageModelDefinition[] { + return getImageModelDefinitions(config, SEEDREAM_IMAGE_MODELS); +} + +/** Merge a provider's optional built-in catalog with its user-owned models. */ +export function getImageModelDefinitions( + config: ImageGenerationConfig, + catalog: readonly ImageModelDefinition[] = [] +): readonly ImageModelDefinition[] { + return [...catalog, ...(config.models ?? [])]; } /** Resolve a curated or user-added Ark image model by id. */ export function getArkImageModelDefinition( - config: ArkImageGenerationConfig, + config: ImageGenerationConfig, modelId: string ): SeedreamImageModelDefinition | undefined { return getArkImageModelDefinitions(config).find( @@ -81,6 +100,17 @@ export function getArkImageModelDefinition( ); } +/** Resolve one configured image model using the owning provider's catalog. */ +export function getImageModelDefinition( + config: ImageGenerationConfig, + modelId: string, + catalog: readonly ImageModelDefinition[] = [] +): ImageModelDefinition | undefined { + return getImageModelDefinitions(config, catalog).find( + (model) => model.id === modelId + ); +} + /** Whether a size preset is supported by the selected Seedream model. */ export function isSeedreamImageSizeSupported( modelId: string, @@ -95,12 +125,22 @@ export function isSeedreamImageSizeSupported( /** Whether a curated or user-added Ark model supports a size preset. */ export function isArkImageSizeSupported( - config: ArkImageGenerationConfig, + config: ImageGenerationConfig, modelId: string, size: string +): boolean { + return isImageSizeSupported(config, modelId, size, SEEDREAM_IMAGE_MODELS); +} + +/** Whether a configured provider image model supports a size preset. */ +export function isImageSizeSupported( + config: ImageGenerationConfig, + modelId: string, + size: string, + catalog: readonly ImageModelDefinition[] = [] ): boolean { return Boolean( - getArkImageModelDefinition(config, modelId)?.supportedSizes.some( + getImageModelDefinition(config, modelId, catalog)?.supportedSizes.some( (supported) => supported === size ) ); diff --git a/packages/core/src/types/models/provider-group.ts b/packages/core/src/types/models/provider-group.ts index af326213..a67c8927 100644 --- a/packages/core/src/types/models/provider-group.ts +++ b/packages/core/src/types/models/provider-group.ts @@ -1,6 +1,6 @@ import * as pi from "@earendil-works/pi-ai"; -import type { ArkImageGenerationConfig } from "./image-generation"; +import type { ImageGenerationConfig } from "./image-generation"; import type { ProviderProfile } from "./provider-profile"; /** @@ -25,8 +25,8 @@ export interface ModelProviderGroup { disabledModels?: string[]; /** Ids of the user-added custom models within this provider. */ customModels?: string[]; - /** Native Ark image-model inventory. Present only on the Ark provider. */ - imageGeneration?: ArkImageGenerationConfig; + /** Provider-owned image-model inventory and request protocol. */ + imageGeneration?: ImageGenerationConfig; websiteLink?: string; websiteURL?: string; apiKeyURL?: string; diff --git a/packages/runtime/src/models/ark-image-generation.ts b/packages/runtime/src/models/ark-image-generation.ts index 6c6b393a..53db7a69 100644 --- a/packages/runtime/src/models/ark-image-generation.ts +++ b/packages/runtime/src/models/ark-image-generation.ts @@ -7,10 +7,13 @@ import { type ImagesOptions, } from "@earendil-works/pi-ai"; import { - getArkImageModelDefinition, - getArkImageModelDefinitions, - isArkImageSizeSupported, - type ArkImageGenerationConfig, + getImageModelDefinition, + getImageModelDefinitions, + isImageSizeSupported, + SEEDREAM_IMAGE_MODELS, + type ImageGenerationApi, + type ImageGenerationConfig, + type ImageModelDefinition, type ProviderConnectionRef, type SeedreamImageSize, } from "@llm-space/core"; @@ -26,7 +29,7 @@ type FetchLike = ( ) => Promise; export interface ArkImageGenerationDependencies { - getConfig(): ArkImageGenerationConfig | undefined; + getConfig(providerId: string): ImageGenerationConfig | undefined; resolveConnection( connection: ProviderConnectionRef ): Promise; @@ -87,53 +90,59 @@ export function createArkImageGenerator( if (!prompt) { throw new Error("prompt must be a non-empty string."); } - const config = dependencies.getConfig(); + const connectionRef = input.connection ?? { providerId: "ark" }; + const providerId = connectionRef.providerId; + const config = dependencies.getConfig(providerId); if (!config) { throw new Error( - "Configure Image generation in Settings → Models → VolcEngine Ark before calling generate_image." + providerId === "ark" + ? "Configure Image generation in Settings → Models → VolcEngine Ark before calling generate_image." + : `Configure image generation for provider "${providerId}" before calling generate_image.` ); } - const modelDefinition = getArkImageModelDefinition(config, input.model); + const catalog = providerId === "ark" ? SEEDREAM_IMAGE_MODELS : []; + const modelDefinition = getImageModelDefinition( + config, + input.model, + catalog + ); if (!modelDefinition) { throw new Error( - `The configured Ark image model "${input.model}" is no longer available. Choose an enabled model for generate_image.` + `The configured image model "${input.model}" is no longer available on provider "${providerId}". Choose an enabled model for generate_image.` ); } if (config.disabledModels?.includes(input.model)) { throw new Error( - `The configured Ark image model "${modelDefinition.name}" is disabled. Choose an enabled model for generate_image.` + `The configured ${providerId === "ark" ? "Ark " : ""}image model "${modelDefinition.name}" is disabled. Choose an enabled model for generate_image.` ); } - if (!isArkImageSizeSupported(config, input.model, input.size)) { + if (!isImageSizeSupported(config, input.model, input.size, catalog)) { throw new Error( `${modelDefinition.name} does not support the ${input.size} size preset.` ); } - const connectionRef = input.connection ?? { providerId: "ark" }; - if (connectionRef.providerId !== "ark") { - throw new Error( - `Ark image generation cannot use provider: ${connectionRef.providerId}` - ); - } const connection = await dependencies.resolveConnection(connectionRef); const apiKey = connection.apiKey; if (!apiKey) { throw new Error( - "Configure an Ark API key in Settings → Models → VolcEngine Ark before calling generate_image." + `Configure an API key for provider "${providerId}" before calling generate_image.` ); } const imagesModels = createImagesModels(); imagesModels.setProvider( - _createArkImagesProvider({ - baseUrl: connection.baseUrl ?? ARK_BASE_URL, + _createImagesProvider({ + providerId, + baseUrl: + connection.baseUrl ?? (providerId === "ark" ? ARK_BASE_URL : ""), config, + catalog, fetch: fetchImpl, }) ); - const model = imagesModels.getModel("ark", input.model); + const model = imagesModels.getModel(providerId, input.model); if (!model) { - throw new Error(`Unsupported Seedream model: ${input.model}`); + throw new Error(`Unsupported image model: ${input.model}`); } const generated = (await imagesModels.generateImages( model, @@ -161,7 +170,7 @@ export function createArkImageGenerator( }; } -/** Bind Ark generation to the shared model connection resolver. */ +/** Bind provider-owned image generation to the shared model connection resolver. */ export function createConfiguredArkImageGenerator({ modelManager, env, @@ -170,59 +179,75 @@ export function createConfiguredArkImageGenerator({ env: Record; }) { return createArkImageGenerator({ - getConfig: () => modelManager.getArkImageGenerationConfig(), + getConfig: (providerId) => + modelManager.getImageGenerationConfig(providerId), resolveConnection: (connection) => modelManager.resolveConnection(connection, { - fallbackApiKey: env.ARK_API_KEY, + fallbackApiKey: + connection.providerId === "ark" ? env.ARK_API_KEY : undefined, }), }); } -/** Build the pi-ai image provider around Ark's native generation endpoint. */ -function _createArkImagesProvider({ +/** Build a pi-ai image provider around the configured request protocol. */ +function _createImagesProvider({ + providerId, baseUrl, config, + catalog, fetch, }: { + providerId: string; baseUrl: string; - config: ArkImageGenerationConfig; + config: ImageGenerationConfig; + catalog: readonly ImageModelDefinition[]; fetch: FetchLike; }) { - const models: ImagesModel[] = - getArkImageModelDefinitions(config).map((definition) => ({ - id: definition.id, - name: definition.name, - api: ARK_IMAGES_API, - provider: "ark", - baseUrl, - input: ["text"], - output: ["image"], - cost: { - input: 0, - output: 0, - cacheRead: 0, - cacheWrite: 0, - }, - })); + if (!baseUrl) { + throw new Error(`Configure a Base URL for provider "${providerId}".`); + } + const api = + config.api ?? (providerId === "ark" ? "ark-images" : "openai-images"); + const models: ImagesModel[] = getImageModelDefinitions( + config, + catalog + ).map((definition) => ({ + id: definition.id, + name: definition.name, + api: ARK_IMAGES_API, + provider: providerId, + baseUrl, + input: ["text"], + output: ["image"], + cost: { + input: 0, + output: 0, + cacheRead: 0, + cacheWrite: 0, + }, + })); return createImagesProvider({ - id: "ark", - name: "VolcEngine Ark", + id: providerId, + name: providerId, auth: { apiKey: envApiKeyAuth("ARK_API_KEY", ["ARK_API_KEY"]) }, models, api: { generateImages: (model, context, options) => - _generateArkImages(model, context.input, options, fetch), + _generateImages(model, context.input, options, fetch, api), }, }); } -/** Execute one synchronous Ark request and normalize it to pi image content. */ -async function _generateArkImages( +/** Execute one synchronous image request and normalize it to pi image content. */ +async function _generateImages( model: ImagesModel, input: { type: string; text?: string }[], options: ImagesOptions | undefined, - fetch: FetchLike + fetch: FetchLike, + api: ImageGenerationApi ): Promise { + const label = + api === "ark-images" ? "Ark image generation" : "Image generation"; const output: ArkAssistantImages = { api: model.api, provider: model.provider, @@ -233,7 +258,7 @@ async function _generateArkImages( }; try { if (!options?.apiKey) { - throw new Error("No API key for provider: ark"); + throw new Error(`No API key for provider: ${model.provider}`); } const prompt = input .filter((item) => item.type === "text" && typeof item.text === "string") @@ -244,9 +269,18 @@ async function _generateArkImages( model: model.id, prompt, size: metadata.size, - watermark: metadata.watermark, - response_format: "b64_json", - stream: false, + ...(api === "ark-images" + ? { + response_format: "b64_json", + watermark: metadata.watermark, + stream: false, + } + : api === "openai-images-extra-body" + ? { + return_base64: true, + extra_body: { response_format: "b64_json" }, + } + : { response_format: "b64_json" }), }; const transformed = await options.onPayload?.(payload, model); if (transformed !== undefined) { @@ -269,12 +303,16 @@ async function _generateArkImages( }, model ); - const body = await _readArkResponse(response); + const body = await _readImageResponse(response, label); if (!response.ok) { - throw _arkProviderError(body.error ?? body, `HTTP ${response.status}`); + throw _providerError( + body.error ?? body, + `HTTP ${response.status}`, + label + ); } if (body.error) { - throw _arkProviderError(body.error, "Provider error"); + throw _providerError(body.error, "Provider error", label); } const items = Array.isArray(body.data) ? (body.data as ArkImageResponseItem[]) @@ -285,11 +323,11 @@ async function _generateArkImages( if (!succeeded) { const failed = items.find((item) => item.error)?.error; if (failed) { - throw _arkProviderError(failed, "Image generation failed"); + throw _providerError(failed, "Image generation failed", label); } - throw new Error("Ark image generation returned no image data."); + throw new Error(`${label} returned no image data.`); } - const image = _normalizeBase64Image(succeeded.b64_json as string); + const image = _normalizeBase64Image(succeeded.b64_json as string, label); output.output.push({ type: "image", ...image }); output.generatedModel = typeof body.model === "string" ? body.model : model.id; @@ -299,16 +337,19 @@ async function _generateArkImages( } catch (error) { output.stopReason = options?.signal?.aborted ? "aborted" : "error"; output.errorMessage = options?.signal?.aborted - ? "Ark image generation was aborted." + ? `${label} was aborted.` : error instanceof Error ? error.message - : "Ark image generation failed."; + : `${label} failed.`; return output; } } /** Read JSON without exposing a provider's raw body in malformed-response errors. */ -async function _readArkResponse(response: Response): Promise { +async function _readImageResponse( + response: Response, + label: string +): Promise { const text = await response.text(); try { const value = JSON.parse(text) as unknown; @@ -318,13 +359,17 @@ async function _readArkResponse(response: Response): Promise { return value; } catch { throw new Error( - `Ark image generation returned invalid JSON (HTTP ${response.status}).` + `${label} returned invalid JSON (HTTP ${response.status}).` ); } } /** Keep Ark's machine-readable code while avoiding raw response serialization. */ -function _arkProviderError(value: unknown, fallback: string): Error { +function _providerError( + value: unknown, + fallback: string, + label: string +): Error { const candidate = value && typeof value === "object" ? (value as { code?: unknown; message?: unknown }) @@ -337,7 +382,7 @@ function _arkProviderError(value: unknown, fallback: string): Error { typeof candidate.message === "string" && candidate.message.trim() ? candidate.message.trim() : "Request failed."; - return new Error(`Ark image generation failed (${code}): ${message}`); + return new Error(`${label} failed (${code}): ${message}`); } /** Resolve and validate the provider-specific options carried in pi metadata. */ @@ -371,7 +416,10 @@ function _requestHeaders( } /** Normalize either a raw base64 value or a data URL and infer its MIME type. */ -function _normalizeBase64Image(value: string): { +function _normalizeBase64Image( + value: string, + label: string +): { data: string; mimeType: string; } { @@ -379,11 +427,11 @@ function _normalizeBase64Image(value: string): { const mimeType = dataUrl?.[1]; const data = (dataUrl?.[2] ?? value).replace(/\s/g, ""); if (!data || !/^[A-Za-z0-9+/]*={0,2}$/.test(data)) { - throw new Error("Ark image generation returned malformed base64 data."); + throw new Error(`${label} returned malformed base64 data.`); } const bytes = Buffer.from(data, "base64"); if (bytes.length === 0) { - throw new Error("Ark image generation returned empty image data."); + throw new Error(`${label} returned empty image data.`); } return { data, diff --git a/packages/runtime/src/models/index.ts b/packages/runtime/src/models/index.ts index 00d06a02..f0aae037 100644 --- a/packages/runtime/src/models/index.ts +++ b/packages/runtime/src/models/index.ts @@ -1,12 +1,11 @@ export { + createArkImageGenerator as createImageGenerator, createArkImageGenerator, + createConfiguredArkImageGenerator as createConfiguredImageGenerator, createConfiguredArkImageGenerator, type ArkImageGenerationDependencies, type ArkImageGenerationInput, type ArkImageGenerationResult, } from "./ark-image-generation"; -export { - ModelManager, - type ResolvedProviderConnection, -} from "./model-manager"; +export { ModelManager, type ResolvedProviderConnection } from "./model-manager"; export type { ModelsConfig, ProviderConfig } from "./types"; diff --git a/packages/runtime/src/models/model-manager.ts b/packages/runtime/src/models/model-manager.ts index e472e886..bf4398b9 100644 --- a/packages/runtime/src/models/model-manager.ts +++ b/packages/runtime/src/models/model-manager.ts @@ -11,12 +11,12 @@ import { } from "@earendil-works/pi-ai"; import { DEFAULT_ARK_IMAGE_GENERATION_CONFIG, - getArkImageModelDefinitions, + getImageModelDefinitions, ModelConfig, SEEDREAM_IMAGE_MODELS, SEEDREAM_IMAGE_SIZES, - type ArkImageGenerationConfig, type CustomModel, + type ImageGenerationConfig, type ModelProviderGroup, type ProviderConnectionRef, type ProviderProfile, @@ -87,17 +87,6 @@ const PROVIDER_PROFILE_FILE_SCHEMA = z.object({ baseUrl: z.string().optional(), headers: z.record(z.string(), z.string()).optional(), }); -const ARK_IMAGE_MODEL_FILE_SCHEMA = z.object({ - id: z.string(), - name: z.string(), - supportedSizes: z.array(z.enum(SEEDREAM_IMAGE_SIZES)), - defaultSize: z.enum(SEEDREAM_IMAGE_SIZES), - icon: z.string().optional(), -}); -const ARK_IMAGE_GENERATION_FILE_SCHEMA = z.object({ - models: z.array(ARK_IMAGE_MODEL_FILE_SCHEMA).optional(), - disabledModels: z.array(z.string()).optional(), -}); const ProviderConfigFileSchema = z.object({ id: z.string(), name: z.string().optional(), @@ -113,7 +102,9 @@ const ProviderConfigFileSchema = z.object({ disabledModels: z.array(z.string()).optional(), models: z.array(CustomModelFileSchema).optional(), customModels: z.array(z.string()).optional(), - imageGeneration: ARK_IMAGE_GENERATION_FILE_SCHEMA.optional(), + // Parse this field leniently so one damaged image inventory cannot discard + // every otherwise valid provider before field-level normalization runs. + imageGeneration: z.unknown().optional(), }); const ModelsConfigFileSchema = z.object({ providers: z.array(ProviderConfigFileSchema), @@ -154,7 +145,7 @@ export class ModelManager { // renderer always sees which models are user-added, then persist any change. const providersChanged = this._normalizeCustomProviders(); const modelsChanged = this._normalizeCustomModels(); - const imageGenerationChanged = this._normalizeArkImageGeneration(); + const imageGenerationChanged = this._normalizeImageGeneration(); const profilesChanged = this._normalizeProviderProfiles(); if ( providersChanged || @@ -304,6 +295,7 @@ export class ModelManager { id, name, api, + imageGeneration: { api: "openai-images" }, profiles: [ { id: uuid(), @@ -334,7 +326,7 @@ export class ModelManager { name?: string | null; api?: CustomProviderApi | null; icon?: string | null; - imageGeneration?: ArkImageGenerationConfig; + imageGeneration?: ImageGenerationConfig; } ): void { const entry = this._config.providers.find( @@ -356,12 +348,17 @@ export class ModelManager { else entry.icon = icon; } if (imageGeneration !== undefined) { - if (providerId !== "ark" || entry.builtin !== true) { + if (entry.builtin === true && entry.id !== "ark") { throw new Error( - "Image generation can only be configured on the builtin Ark provider." + "Image generation can only be configured on Ark or a custom provider." ); } - _assertArkImageGenerationConfig(imageGeneration); + _assertImageGenerationConfig( + imageGeneration, + entry.id === "ark" && entry.builtin === true + ? SEEDREAM_IMAGE_MODELS + : [] + ); entry.imageGeneration = { ...imageGeneration }; } // Rebuild the registry so a cleared baseUrl restores the model's default @@ -531,10 +528,15 @@ export class ModelManager { } /** The saved Ark image-model inventory, or undefined without Ark. */ - getArkImageGenerationConfig(): ArkImageGenerationConfig | undefined { - const config = this._config.providers.find( - (entry) => entry.id === "ark" && entry.builtin === true - )?.imageGeneration; + getArkImageGenerationConfig(): ImageGenerationConfig | undefined { + return this.getImageGenerationConfig("ark"); + } + + /** The provider-owned image inventory, if image generation is configured. */ + getImageGenerationConfig( + providerId: string + ): ImageGenerationConfig | undefined { + const config = this._findProvider(providerId)?.imageGeneration; return config ? { ...config } : undefined; } @@ -883,26 +885,32 @@ export class ModelManager { return changed; } - /** - * Normalize Ark's image-model inventory and remove legacy provider defaults. - * This keeps upgrades readable without treating image models as chat models - * or rejecting the whole settings file. - */ - private _normalizeArkImageGeneration(): boolean { - const entry = this._config.providers.find( - (provider) => provider.id === "ark" && provider.builtin === true - ); - if (!entry) { - return false; - } - const normalized = _normalizeArkImageGenerationConfig( - entry.imageGeneration - ); - if (JSON.stringify(entry.imageGeneration) === JSON.stringify(normalized)) { - return false; + /** Normalize every provider's image inventory while preserving Ark legacy data. */ + private _normalizeImageGeneration(): boolean { + let changed = false; + for (const entry of this._config.providers) { + if (entry.imageGeneration === undefined) { + continue; + } + const isArk = entry.id === "ark" && entry.builtin === true; + if (entry.builtin === true && !isArk) { + delete entry.imageGeneration; + changed = true; + continue; + } + const normalized = _normalizeImageGenerationConfig( + entry.imageGeneration, + isArk ? SEEDREAM_IMAGE_MODELS : [], + isArk ? "ark-images" : "openai-images" + ); + if ( + JSON.stringify(entry.imageGeneration) !== JSON.stringify(normalized) + ) { + entry.imageGeneration = normalized; + changed = true; + } } - entry.imageGeneration = normalized; - return true; + return changed; } /** @@ -1059,12 +1067,14 @@ export class ModelManager { } } -/** Validate a renderer-supplied Ark image configuration before persisting it. */ -function _assertArkImageGenerationConfig( - config: ArkImageGenerationConfig +/** Validate a renderer-supplied provider image configuration before persisting it. */ +function _assertImageGenerationConfig( + config: ImageGenerationConfig, + catalog: readonly SeedreamImageModelDefinition[] ): void { - _assertCustomArkImageModels(config.models); - const models = getArkImageModelDefinitions(config); + const scope = catalog.length > 0 ? "Ark image" : "image"; + _assertCustomImageModels(config.models, catalog); + const models = getImageModelDefinitions(config, catalog); const modelIds = new Set(models.map((model) => model.id)); const disabledModels = config.disabledModels ?? []; if ( @@ -1072,25 +1082,33 @@ function _assertArkImageGenerationConfig( new Set(disabledModels).size !== disabledModels.length ) { throw new Error( - "Disabled Ark image models must reference unique model ids." + `Disabled ${scope} models must reference unique model ids.` ); } } /** Normalize untrusted JSON from older or manually edited settings files. */ -function _normalizeArkImageGenerationConfig( - value: unknown -): ArkImageGenerationConfig { +function _normalizeImageGenerationConfig( + value: unknown, + catalog: readonly SeedreamImageModelDefinition[], + defaultApi: "ark-images" | "openai-images" +): ImageGenerationConfig { const candidate = value && typeof value === "object" - ? (value as Partial) + ? (value as Partial) : {}; - const models = _normalizeCustomArkImageModels(candidate.models); - const withModels: ArkImageGenerationConfig = { + const models = _normalizeCustomImageModels(candidate.models, catalog); + const api = + candidate.api === "ark-images" || + candidate.api === "openai-images" || + candidate.api === "openai-images-extra-body" + ? candidate.api + : defaultApi; + const withModels: ImageGenerationConfig = { ...(models.length > 0 ? { models } : {}), }; const modelIds = new Set( - getArkImageModelDefinitions(withModels).map((model) => model.id) + getImageModelDefinitions(withModels, catalog).map((model) => model.id) ); const disabledModels = Array.isArray(candidate.disabledModels) ? [ @@ -1103,22 +1121,25 @@ function _normalizeArkImageGenerationConfig( ] : []; return { + ...(api !== "ark-images" ? { api } : {}), ...(models.length > 0 ? { models } : {}), ...(disabledModels.length > 0 ? { disabledModels } : {}), }; } /** Reject invalid custom model definitions supplied through renderer RPC. */ -function _assertCustomArkImageModels( - models: ArkImageGenerationConfig["models"] +function _assertCustomImageModels( + models: ImageGenerationConfig["models"], + catalog: readonly SeedreamImageModelDefinition[] ): void { + const scope = catalog.length > 0 ? "Ark image" : "image"; if (models === undefined) { return; } if (!Array.isArray(models)) { - throw new Error("Custom Ark image models must be an array."); + throw new Error(`Custom ${scope} models must be an array.`); } - const seen = new Set(SEEDREAM_IMAGE_MODELS.map((model) => model.id)); + const seen = new Set(catalog.map((model) => model.id)); for (const model of models) { if ( !model || @@ -1129,10 +1150,10 @@ function _assertCustomArkImageModels( model.name.trim() !== model.name || model.name === "" ) { - throw new Error("Custom Ark image models require a valid id and name."); + throw new Error(`Custom ${scope} models require a valid id and name.`); } if (seen.has(model.id)) { - throw new Error(`Duplicate Ark image model id: ${model.id}`); + throw new Error(`Duplicate ${scope} model id: ${model.id}`); } seen.add(model.id); const sizes = model.supportedSizes; @@ -1144,25 +1165,24 @@ function _assertCustomArkImageModels( !sizes.includes(model.defaultSize) ) { throw new Error( - `Custom Ark image model ${model.id} has invalid size presets.` + `Custom ${scope} model ${model.id} has invalid size presets.` ); } if (model.icon !== undefined && typeof model.icon !== "string") { - throw new Error( - `Custom Ark image model ${model.id} has an invalid icon.` - ); + throw new Error(`Custom ${scope} model ${model.id} has an invalid icon.`); } } } /** Repair user-edited JSON by keeping only complete, unique custom models. */ -function _normalizeCustomArkImageModels( - value: unknown +function _normalizeCustomImageModels( + value: unknown, + catalog: readonly SeedreamImageModelDefinition[] ): SeedreamImageModelDefinition[] { if (!Array.isArray(value)) { return []; } - const seen = new Set(SEEDREAM_IMAGE_MODELS.map((model) => model.id)); + const seen = new Set(catalog.map((model) => model.id)); const models: SeedreamImageModelDefinition[] = []; for (const raw of value) { if (!raw || typeof raw !== "object") { diff --git a/packages/runtime/src/models/types.ts b/packages/runtime/src/models/types.ts index 1f9ccc80..e398850f 100644 --- a/packages/runtime/src/models/types.ts +++ b/packages/runtime/src/models/types.ts @@ -1,6 +1,6 @@ import type { - ArkImageGenerationConfig, CustomModel, + ImageGenerationConfig, ModelConfig, ProviderProfile, } from "@llm-space/core"; @@ -58,8 +58,8 @@ export interface ProviderConfig { * so these models can later be singled out for deletion. */ customModels?: string[]; - /** Native Ark image-model inventory; only valid on the builtin Ark provider. */ - imageGeneration?: ArkImageGenerationConfig; + /** Provider-owned image-model inventory and request protocol. */ + imageGeneration?: ImageGenerationConfig; } /** Shape of `settings/models.json`. */ diff --git a/packages/runtime/src/runtime/model-groups.ts b/packages/runtime/src/runtime/model-groups.ts index 5563c27b..8b080494 100644 --- a/packages/runtime/src/runtime/model-groups.ts +++ b/packages/runtime/src/runtime/model-groups.ts @@ -16,10 +16,7 @@ export async function getModelProviderGroups( api: modelManager.getApi(provider.id), disabledModels: modelManager.getDisabledModels(provider.id), customModels: modelManager.getCustomModels(provider.id), - imageGeneration: - provider.id === "ark" - ? modelManager.getArkImageGenerationConfig() - : undefined, + imageGeneration: modelManager.getImageGenerationConfig(provider.id), websiteLink: modelManager.getWebsiteLink(provider.id), icon: modelManager.getProviderIcon(provider.id), })); diff --git a/packages/runtime/src/runtime/types.ts b/packages/runtime/src/runtime/types.ts index b5f7cfd3..7e5770d7 100644 --- a/packages/runtime/src/runtime/types.ts +++ b/packages/runtime/src/runtime/types.ts @@ -1,7 +1,7 @@ import type { AgentEvent, AgentStreamRequest, - ArkImageGenerationConfig, + ImageGenerationConfig, BuiltinTool, BuiltinToolCallResponse, CustomModel, @@ -30,7 +30,6 @@ import type { CreateSubagentThreadResult, } from "@llm-space/core/thread"; - import type { TraceConnectedProjectInput, TraceImportFile, @@ -119,7 +118,7 @@ export interface RuntimeClient { api?: "anthropic-messages" | "openai-completions" | "openai-responses" | null; icon?: string | null; - imageGeneration?: ArkImageGenerationConfig; + imageGeneration?: ImageGenerationConfig; }): Promise; setModelEnabled(input: { providerId: string; @@ -154,7 +153,7 @@ export interface RuntimeClient { }): Promise; createSubagentThread( - input: CreateSubagentThreadInput, + input: CreateSubagentThreadInput ): Promise; fsLs(path: string): Promise; @@ -168,10 +167,7 @@ export interface RuntimeClient { path: string, run: ThreadRunSnapshot & { id: string } ): Promise; - fsReadRunSnapshot( - path: string, - snapshotRef: string - ): Promise; + fsReadRunSnapshot(path: string, snapshotRef: string): Promise; fsRealpath(path: string): Promise; /** Read arbitrary prompt text (`~` expands on this runtime). */ readTextFile(path: string): Promise; diff --git a/packages/runtime/src/tools/built-in/media.ts b/packages/runtime/src/tools/built-in/media.ts index 3f8b1d9f..2ef29cf1 100644 --- a/packages/runtime/src/tools/built-in/media.ts +++ b/packages/runtime/src/tools/built-in/media.ts @@ -33,9 +33,8 @@ export const generateImageTool: BuiltinTool = { type: "builtin", name: "generate_image", icon: "image", - connection: { providerId: "ark" }, description: - "Generate one image with this tool's selected Ark image model. Use the configured default size unless the user requests a supported 1K, 2K, 3K, or 4K preset.", + "Generate one image with this tool's selected provider and image model. Use the configured default size unless the user requests a supported 1K, 2K, 3K, or 4K preset.", strict: true, parameters: { type: "object", diff --git a/packages/runtime/src/tools/tool-registry.ts b/packages/runtime/src/tools/tool-registry.ts index c9356d36..2902ab3d 100644 --- a/packages/runtime/src/tools/tool-registry.ts +++ b/packages/runtime/src/tools/tool-registry.ts @@ -124,7 +124,8 @@ export class ToolRegistry { } if ( connection && - entry.tool.connection?.providerId !== connection.providerId + entry.tool.connection && + entry.tool.connection.providerId !== connection.providerId ) { throw new Error( `Built-in tool ${name} does not use provider: ${connection.providerId}` diff --git a/packages/runtime/tests/models/ark-image-generation.test.ts b/packages/runtime/tests/models/ark-image-generation.test.ts index 7945b690..22341ade 100644 --- a/packages/runtime/tests/models/ark-image-generation.test.ts +++ b/packages/runtime/tests/models/ark-image-generation.test.ts @@ -4,8 +4,10 @@ import { type ArkImageGenerationConfig } from "@llm-space/core"; import { createArkImageGenerator, + createConfiguredArkImageGenerator, type ArkImageGenerationDependencies, } from "../../src/models/ark-image-generation"; +import type { ModelManager } from "../../src/models/model-manager"; const PNG_BASE64 = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNk+M/wHwAF/gL+X2NDWQAAAABJRU5ErkJggg=="; @@ -265,4 +267,119 @@ describe("Ark image generation", () => { ).rejects.toThrow("Ark image generation was aborted"); expect(receivedSignal).toBe(controller.signal); }); + + test("uses a custom provider connection and Agnes-compatible payload", async () => { + let requestUrl = ""; + let requestHeaders: HeadersInit | undefined; + let requestBody: unknown; + const generate = createArkImageGenerator( + _dependencies({ + getConfig: (providerId) => + providerId === "agnes" + ? { + api: "openai-images-extra-body", + models: [ + { + id: "agnes-image-2.5-flash", + name: "Agnes Image 2.5 Flash", + supportedSizes: ["1K", "2K", "3K", "4K"], + defaultSize: "2K", + }, + ], + } + : undefined, + resolveConnection: (connection) => { + expect(connection).toEqual({ + providerId: "agnes", + profileId: "agnes-default", + }); + return Promise.resolve({ + apiKey: "agnes-key", + baseUrl: "https://api.agnes-ai.cn/v1/", + headers: { "X-Tenant": "test" }, + }); + }, + fetch: (input, init) => { + requestUrl = + typeof input === "string" + ? input + : input instanceof URL + ? input.href + : input.url; + requestHeaders = init?.headers; + requestBody = + typeof init?.body === "string" ? JSON.parse(init.body) : undefined; + return Promise.resolve( + Response.json({ + model: "agnes-image-2.5-flash", + data: [{ b64_json: PNG_BASE64, size: "2048x2048" }], + }) + ); + }, + }) + ); + + await generate({ + prompt: "A red circle", + model: "agnes-image-2.5-flash", + size: "2K", + watermark: false, + connection: { + providerId: "agnes", + profileId: "agnes-default", + }, + }); + + expect(requestUrl).toBe("https://api.agnes-ai.cn/v1/images/generations"); + expect(requestHeaders).toMatchObject({ + Authorization: "Bearer agnes-key", + "X-Tenant": "test", + }); + expect(requestBody).toEqual({ + model: "agnes-image-2.5-flash", + prompt: "A red circle", + size: "2K", + return_base64: true, + extra_body: { response_format: "b64_json" }, + }); + }); + + test("never falls back to the Ark key for a custom provider", async () => { + let fallbackApiKey: string | undefined; + const modelManager = { + getImageGenerationConfig: () => ({ + api: "openai-images", + models: [ + { + id: "fixture", + name: "Fixture", + supportedSizes: ["1K"], + defaultSize: "1K", + }, + ], + }), + resolveConnection: ( + _connection: unknown, + options: { fallbackApiKey?: string } + ) => { + fallbackApiKey = options.fallbackApiKey; + return Promise.resolve({}); + }, + } as unknown as ModelManager; + const generate = createConfiguredArkImageGenerator({ + modelManager, + env: { ARK_API_KEY: "ark-secret-canary" }, + }); + + expect( + generate({ + prompt: "fixture", + model: "fixture", + size: "1K", + watermark: false, + connection: { providerId: "other-images" }, + }) + ).rejects.toThrow('Configure an API key for provider "other-images"'); + expect(fallbackApiKey).toBeUndefined(); + }); }); diff --git a/packages/runtime/tests/models/model-manager.test.ts b/packages/runtime/tests/models/model-manager.test.ts index ab92fe96..2e6bd69a 100644 --- a/packages/runtime/tests/models/model-manager.test.ts +++ b/packages/runtime/tests/models/model-manager.test.ts @@ -367,3 +367,120 @@ describe("ModelManager provider profiles", () => { ).toBeUndefined(); }); }); + +describe("ModelManager provider-owned image generation", () => { + test("repairs a damaged image inventory without discarding providers", async () => { + const settingsDir = await _settingsDir({ + providers: [ + { + id: "fixture-images", + name: "Fixture Images", + baseUrl: "https://images.example/v1", + imageGeneration: { + api: "openai-images", + models: "damaged", + disabledModels: ["missing"], + }, + }, + { + id: "openai", + builtin: true, + imageGeneration: { + api: "openai-images", + models: [], + }, + }, + ], + }); + + const manager = new ModelManager({ settingsDir }); + + expect(manager.getImageGenerationConfig("fixture-images")).toEqual({ + api: "openai-images", + }); + expect(manager.getImageGenerationConfig("openai")).toBeUndefined(); + expect(manager.getProfiles("fixture-images")[0]?.baseUrl).toBe( + "https://images.example/v1" + ); + const persisted = JSON.parse( + await readFile(path.join(settingsDir, "models.json"), "utf8") + ) as { + providers: { + id: string; + imageGeneration?: unknown; + }[]; + }; + expect(persisted.providers.map((provider) => provider.id)).toEqual([ + "fixture-images", + "openai", + ]); + expect( + persisted.providers.find((provider) => provider.id === "openai") + ?.imageGeneration + ).toBeUndefined(); + }); + + test("rejects image configuration on non-Ark built-in providers", async () => { + const settingsDir = await _settingsDir({ providers: [] }); + const manager = new ModelManager({ settingsDir }); + manager.addBuiltInProvider({ id: "openai" }); + + expect(() => + manager.updateProvider("openai", { + imageGeneration: { + api: "openai-images", + models: [ + { + id: "unsupported-builtin-image", + name: "Unsupported Builtin Image", + supportedSizes: ["1K"], + defaultSize: "1K", + }, + ], + }, + }) + ).toThrow( + "Image generation can only be configured on Ark or a custom provider" + ); + }); + + test("persists an image-only custom provider without registering chat models", async () => { + const settingsDir = await _settingsDir({ providers: [] }); + const manager = new ModelManager({ settingsDir }); + manager.addCustomProvider({ + id: "agnes", + name: "Agnes", + baseUrl: "https://api.agnes-ai.cn/v1", + }); + manager.updateProvider("agnes", { + imageGeneration: { + api: "openai-images", + models: [ + { + id: "agnes-image-2.5-flash", + name: "Agnes Image 2.5 Flash", + supportedSizes: ["1K", "2K", "3K", "4K"], + defaultSize: "2K", + }, + ], + }, + }); + + const reloaded = new ModelManager({ settingsDir }); + expect(reloaded.getImageGenerationConfig("agnes")).toEqual({ + api: "openai-images", + models: [ + { + id: "agnes-image-2.5-flash", + name: "Agnes Image 2.5 Flash", + supportedSizes: ["1K", "2K", "3K", "4K"], + defaultSize: "2K", + }, + ], + }); + const provider = (await reloaded.getAvailableModels()) + .getProviders() + .find((candidate) => candidate.id === "agnes"); + expect(provider?.getModels()).toEqual([]); + }); +}); diff --git a/packages/runtime/tests/tools/built-in/built-in-tools-module.test.ts b/packages/runtime/tests/tools/built-in/built-in-tools-module.test.ts index a1ffe2a7..4c121750 100644 --- a/packages/runtime/tests/tools/built-in/built-in-tools-module.test.ts +++ b/packages/runtime/tests/tools/built-in/built-in-tools-module.test.ts @@ -107,7 +107,7 @@ describe("built-in tools module", () => { expect( tools.listTools().find((tool) => tool.name === "generate_image") ?.connection - ).toEqual({ providerId: "ark" }); + ).toBeUndefined(); expect( await tools.call({ name: "skill", diff --git a/packages/ui/src/components/model-provider.tsx b/packages/ui/src/components/model-provider.tsx index 8065d8cf..c79100c5 100644 --- a/packages/ui/src/components/model-provider.tsx +++ b/packages/ui/src/components/model-provider.tsx @@ -2,7 +2,7 @@ import type * as pi from "@earendil-works/pi-ai"; import type { - ArkImageGenerationConfig, + ImageGenerationConfig, CustomModel, ModelConfig, ModelProviderGroup, @@ -46,7 +46,7 @@ interface ModelContextValue { api?: "anthropic-messages" | "openai-completions" | "openai-responses" | null; icon?: string | null; - imageGeneration?: ArkImageGenerationConfig; + imageGeneration?: ImageGenerationConfig; } ) => Promise; setModelEnabled: ( @@ -323,7 +323,7 @@ export function ModelProvider({ | "openai-responses" | null; icon?: string | null; - imageGeneration?: ArkImageGenerationConfig; + imageGeneration?: ImageGenerationConfig; } ) => { const result = await enqueueMutation(client, () => @@ -608,7 +608,7 @@ export function useUpdateProvider(): ( api?: "anthropic-messages" | "openai-completions" | "openai-responses" | null; icon?: string | null; - imageGeneration?: ArkImageGenerationConfig; + imageGeneration?: ImageGenerationConfig; } ) => Promise { return useModelProvider().updateProvider; diff --git a/packages/ui/src/components/thread-playground/tool/built-in-tool-import-dialog.tsx b/packages/ui/src/components/thread-playground/tool/built-in-tool-import-dialog.tsx index d38f78a3..be10fe40 100644 --- a/packages/ui/src/components/thread-playground/tool/built-in-tool-import-dialog.tsx +++ b/packages/ui/src/components/thread-playground/tool/built-in-tool-import-dialog.tsx @@ -2,9 +2,11 @@ import { getArkImageModelDefinitions, + getImageModelDefinitions, SEEDREAM_IMAGE_SIZES, type BuiltinTool, type GenerateImageToolConfig, + type ModelProviderGroup, type SeedreamImageModelDefinition, type SeedreamImageSize, } from "@llm-space/core"; @@ -87,6 +89,22 @@ const MEDIA_TOOL_NAMES = new Set([ "stop_speaking", ]); +/** Resolve enabled image models without mixing them into chat inventory. */ +function _enabledImageModels( + provider: ModelProviderGroup +): readonly SeedreamImageModelDefinition[] { + const config = provider.imageGeneration; + if (!config) { + return []; + } + const disabled = new Set(config.disabledModels ?? []); + const models = + provider.id === "ark" + ? getArkImageModelDefinitions(config) + : getImageModelDefinitions(config); + return models.filter((model) => !disabled.has(model.id)); +} + function _BuiltInToolImportDialog({ existingToolNames, existingTools, @@ -117,28 +135,20 @@ function _BuiltInToolImportDialog({ ); const [generateImageConfig, setGenerateImageConfig] = useState(null); + const [imageProviderId, setImageProviderId] = useState(); const toolRowRefs = useRef(new Map()); const { builtinTools } = useHostServices(); const providers = useModels(); - const generateImageTool = tools.find( - (tool) => tool.name === "generate_image" + const imageProviders = useMemo( + () => providers.filter((provider) => provider.imageGeneration), + [providers] ); - const imageProviderId = generateImageTool - ? getToolConnectionProviderId(generateImageTool) - : undefined; const imageProvider = useMemo( () => providers.find((provider) => provider.id === imageProviderId), [imageProviderId, providers] ); const enabledImageModels = useMemo(() => { - const config = imageProvider?.imageGeneration; - if (!config) { - return []; - } - const disabled = new Set(config.disabledModels ?? []); - return getArkImageModelDefinitions(config).filter( - (model) => !disabled.has(model.id) - ); + return imageProvider ? _enabledImageModels(imageProvider) : []; }, [imageProvider]); const loadTools = useCallback(async () => { @@ -189,11 +199,20 @@ function _BuiltInToolImportDialog({ return; } const existing = existingTools.get("generate_image"); + const providerId = existing + ? getToolConnectionProviderId(existing) + : imageProviders.find( + (provider) => _enabledImageModels(provider).length > 0 + )?.id; + setImageProviderId(providerId); if (existing) { setGenerateImageConfig(_readGenerateImageConfig(existing.config)); return; } - const first = enabledImageModels[0]; + const provider = imageProviders.find( + (candidate) => candidate.id === providerId + ); + const first = provider ? _enabledImageModels(provider)[0] : undefined; setGenerateImageConfig( first ? { @@ -203,14 +222,41 @@ function _BuiltInToolImportDialog({ } : null ); - }, [enabledImageModels, existingTools, open]); + }, [existingTools, imageProviders, open]); /** Persist config immediately for an existing tool or keep it as an add draft. */ - const handleGenerateImageConfigChange = (config: GenerateImageToolConfig) => { + const handleGenerateImageConfigChange = ( + config: GenerateImageToolConfig, + providerId = imageProviderId + ) => { setGenerateImageConfig(config); const existing = existingTools.get("generate_image"); if (existing) { - onUpdate(existing.name, { ...existing, config: { ...config } }); + onUpdate(existing.name, { + ...existing, + config: { ...config }, + ...(providerId ? { connection: { providerId } } : {}), + }); + } + }; + + const handleImageProviderChange = (providerId: string) => { + setImageProviderId(providerId); + const provider = imageProviders.find( + (candidate) => candidate.id === providerId + ); + const first = provider ? _enabledImageModels(provider)[0] : undefined; + if (first) { + handleGenerateImageConfigChange( + { + model: first.id, + size: first.defaultSize, + watermark: true, + }, + providerId + ); + } else { + setGenerateImageConfig(null); } }; @@ -228,12 +274,18 @@ function _BuiltInToolImportDialog({ ); if (!model || !generateImageConfig) { toast.error("Choose an enabled image model", { - description: - "Enable an Ark image model in Settings, then select it here.", + description: "Enable an image model in Settings, then select it here.", }); return; } - onAdd({ ...tool, config: { ...generateImageConfig } }); + if (!imageProviderId) { + return; + } + onAdd({ + ...tool, + connection: { providerId: imageProviderId }, + config: { ...generateImageConfig }, + }); }; const filteredTools = useMemo(() => { const q = query.trim().toLowerCase(); @@ -403,11 +455,13 @@ function _BuiltInToolImportDialog({ <_GenerateImageConfigFields config={generateImageConfig} enabledModels={enabledImageModels} + providers={imageProviders} showProfileSelector={ (imageProvider?.profiles.length ?? 0) > 1 } providerId={imageProviderId} selectedModel={configuredImageModel} + onProviderChange={handleImageProviderChange} onChange={handleGenerateImageConfigChange} /> )} @@ -439,16 +493,20 @@ export const BuiltInToolImportDialog = memo(_BuiltInToolImportDialog); function _GenerateImageConfigFields({ config, enabledModels, + providers, showProfileSelector, providerId, selectedModel, + onProviderChange, onChange, }: { config: GenerateImageToolConfig | null; enabledModels: readonly SeedreamImageModelDefinition[]; + providers: readonly ModelProviderGroup[]; showProfileSelector: boolean; providerId?: string; selectedModel?: SeedreamImageModelDefinition; + onProviderChange: (providerId: string) => void; onChange: (config: GenerateImageToolConfig) => void; }) { const handleModelChange = (modelId: string) => { @@ -471,10 +529,30 @@ function _GenerateImageConfigFields({ className={cn( "mt-3 grid gap-3", showProfileSelector - ? "grid-cols-[minmax(0,1fr)_8rem_7rem_auto]" - : "grid-cols-[minmax(0,1fr)_7rem_auto]" + ? "grid-cols-[9rem_minmax(0,1fr)_8rem_7rem_auto]" + : "grid-cols-[9rem_minmax(0,1fr)_7rem_auto]" )} > +
+ Provider + +
+
Model update({ ...config, api: api as ImageGenerationApi }) } @@ -1326,6 +1329,7 @@ function _ImageGenerationEditor({ + Ark Images OpenAI Images OpenAI Images with extra_body @@ -1430,6 +1434,7 @@ function _ImageGenerationEditor({ model.id)} onSave={handleSaveCustomModel} @@ -1449,7 +1454,7 @@ function _ImageModelListItem({ onDelete, }: { providerName: string; - model: SeedreamImageModelDefinition; + model: ImageModelDefinition; enabled: boolean; isCustom: boolean; onToggle: (enabled: boolean) => void; diff --git a/packages/core/src/types/models/image-generation.ts b/packages/core/src/types/models/image-generation.ts index f33c8637..c2395f58 100644 --- a/packages/core/src/types/models/image-generation.ts +++ b/packages/core/src/types/models/image-generation.ts @@ -2,11 +2,31 @@ export const SEEDREAM_IMAGE_SIZES = ["1K", "2K", "3K", "4K"] as const; export type SeedreamImageSize = (typeof SEEDREAM_IMAGE_SIZES)[number]; +export const OPENAI_IMAGE_SIZES = [ + "auto", + "256x256", + "512x512", + "1024x1024", + "1536x1024", + "1024x1536", + "1792x1024", + "1024x1792", +] as const; + +export type OpenAIImageSize = (typeof OPENAI_IMAGE_SIZES)[number]; + +export const IMAGE_SIZES = [ + ...SEEDREAM_IMAGE_SIZES, + ...OPENAI_IMAGE_SIZES, +] as const; + +export type ImageSize = (typeof IMAGE_SIZES)[number]; + export interface ImageModelDefinition { id: string; name: string; - supportedSizes: readonly SeedreamImageSize[]; - defaultSize: SeedreamImageSize; + supportedSizes: readonly ImageSize[]; + defaultSize: ImageSize; /** Optional `@lobehub/icons` keyword for a user-added image model. */ icon?: string; } @@ -62,7 +82,7 @@ export type ArkImageGenerationConfig = ImageGenerationConfig; /** Per-Thread configuration owned by one `generate_image` tool instance. */ export interface GenerateImageToolConfig { model: string; - size: SeedreamImageSize; + size: ImageSize; watermark: boolean; } @@ -145,3 +165,11 @@ export function isImageSizeSupported( ) ); } + +/** Narrow untrusted values to a supported provider image-size option. */ +export function isImageSize(value: unknown): value is ImageSize { + return ( + typeof value === "string" && + (IMAGE_SIZES as readonly string[]).includes(value) + ); +} diff --git a/packages/runtime/src/models/ark-image-generation.ts b/packages/runtime/src/models/ark-image-generation.ts index 53db7a69..ebf24cb6 100644 --- a/packages/runtime/src/models/ark-image-generation.ts +++ b/packages/runtime/src/models/ark-image-generation.ts @@ -9,13 +9,16 @@ import { import { getImageModelDefinition, getImageModelDefinitions, + isImageSize, isImageSizeSupported, + OPENAI_IMAGE_SIZES, SEEDREAM_IMAGE_MODELS, type ImageGenerationApi, type ImageGenerationConfig, type ImageModelDefinition, + type ImageSize, + type OpenAIImageSize, type ProviderConnectionRef, - type SeedreamImageSize, } from "@llm-space/core"; import type { ModelManager, ResolvedProviderConnection } from "./model-manager"; @@ -39,7 +42,7 @@ export interface ArkImageGenerationDependencies { export interface ArkImageGenerationInput { prompt: string; model: string; - size: SeedreamImageSize; + size: ImageSize; watermark: boolean; connection?: ProviderConnectionRef; signal?: AbortSignal; @@ -57,8 +60,8 @@ interface ArkAssistantImages extends AssistantImages { generatedSize?: string; } -interface ArkImagesMetadata { - size: SeedreamImageSize; +interface ImageGenerationMetadata { + size: ImageSize; watermark: boolean; } @@ -264,23 +267,24 @@ async function _generateImages( .filter((item) => item.type === "text" && typeof item.text === "string") .map((item) => item.text) .join("\n"); - const metadata = _arkMetadata(options.metadata); + const metadata = _imageMetadata(options.metadata); let payload: unknown = { model: model.id, prompt, - size: metadata.size, ...(api === "ark-images" ? { + size: metadata.size, response_format: "b64_json", watermark: metadata.watermark, stream: false, } : api === "openai-images-extra-body" ? { + size: metadata.size, return_base64: true, extra_body: { response_format: "b64_json" }, } - : { response_format: "b64_json" }), + : _openAIImagesPayload(model.id, metadata.size)), }; const transformed = await options.onPayload?.(payload, model); if (transformed !== undefined) { @@ -345,6 +349,52 @@ async function _generateImages( } } +/** Build only parameters accepted by the selected standard OpenAI image model. */ +function _openAIImagesPayload( + modelId: string, + configuredSize: ImageSize +): { size: OpenAIImageSize; response_format?: "b64_json" } { + const size = _openAIImageSize(modelId, configuredSize); + return { + size, + ...(_isDallEModel(modelId) ? { response_format: "b64_json" as const } : {}), + }; +} + +/** Validate standard protocol sizes, with a narrow migration for legacy 1K configs. */ +function _openAIImageSize( + modelId: string, + configuredSize: ImageSize +): OpenAIImageSize { + const size = configuredSize === "1K" ? "1024x1024" : configuredSize; + if (!(OPENAI_IMAGE_SIZES as readonly string[]).includes(size)) { + throw new Error( + `OpenAI Images requires an explicit pixel size; configure ${modelId} with an OpenAI-compatible size instead of ${configuredSize}.` + ); + } + if (modelId === "dall-e-2") { + const supported = ["256x256", "512x512", "1024x1024"]; + if (!supported.includes(size)) { + throw new Error(`dall-e-2 does not support image size ${size}.`); + } + } else if (modelId === "dall-e-3") { + const supported = ["1024x1024", "1792x1024", "1024x1792"]; + if (!supported.includes(size)) { + throw new Error(`dall-e-3 does not support image size ${size}.`); + } + } else if (modelId.startsWith("gpt-image-")) { + const supported = ["auto", "1024x1024", "1536x1024", "1024x1536"]; + if (!supported.includes(size)) { + throw new Error(`${modelId} does not support image size ${size}.`); + } + } + return size as OpenAIImageSize; +} + +function _isDallEModel(modelId: string): boolean { + return modelId === "dall-e-2" || modelId === "dall-e-3"; +} + /** Read JSON without exposing a provider's raw body in malformed-response errors. */ async function _readImageResponse( response: Response, @@ -386,15 +436,15 @@ function _providerError( } /** Resolve and validate the provider-specific options carried in pi metadata. */ -function _arkMetadata( +function _imageMetadata( metadata: Record | undefined -): ArkImagesMetadata { +): ImageGenerationMetadata { const size = metadata?.size; - if (size !== "1K" && size !== "2K" && size !== "3K" && size !== "4K") { - throw new Error("Ark image generation size metadata is invalid."); + if (!isImageSize(size)) { + throw new Error("Image generation size metadata is invalid."); } if (typeof metadata?.watermark !== "boolean") { - throw new Error("Ark image generation watermark metadata is invalid."); + throw new Error("Image generation watermark metadata is invalid."); } return { size, watermark: metadata.watermark }; } diff --git a/packages/runtime/src/models/model-manager.ts b/packages/runtime/src/models/model-manager.ts index bf4398b9..4ecc5f5a 100644 --- a/packages/runtime/src/models/model-manager.ts +++ b/packages/runtime/src/models/model-manager.ts @@ -12,9 +12,9 @@ import { import { DEFAULT_ARK_IMAGE_GENERATION_CONFIG, getImageModelDefinitions, + isImageSize, ModelConfig, SEEDREAM_IMAGE_MODELS, - SEEDREAM_IMAGE_SIZES, type CustomModel, type ImageGenerationConfig, type ModelProviderGroup, @@ -22,7 +22,6 @@ import { type ProviderProfile, type ProviderProfilePatch, type SeedreamImageModelDefinition, - type SeedreamImageSize, uuid, } from "@llm-space/core"; import { @@ -1121,7 +1120,7 @@ function _normalizeImageGenerationConfig( ] : []; return { - ...(api !== "ark-images" ? { api } : {}), + ...(api !== defaultApi ? { api } : {}), ...(models.length > 0 ? { models } : {}), ...(disabledModels.length > 0 ? { disabledModels } : {}), }; @@ -1160,7 +1159,7 @@ function _assertCustomImageModels( if ( !Array.isArray(sizes) || sizes.length === 0 || - sizes.some((size) => !_isSeedreamImageSize(size)) || + sizes.some((size) => !isImageSize(size)) || new Set(sizes).size !== sizes.length || !sizes.includes(model.defaultSize) ) { @@ -1196,13 +1195,13 @@ function _normalizeCustomImageModels( continue; } const supportedSizes = Array.isArray(candidate.supportedSizes) - ? [...new Set(candidate.supportedSizes.filter(_isSeedreamImageSize))] + ? [...new Set(candidate.supportedSizes.filter(isImageSize))] : []; if (supportedSizes.length === 0) { continue; } const defaultSize = - _isSeedreamImageSize(candidate.defaultSize) && + isImageSize(candidate.defaultSize) && supportedSizes.includes(candidate.defaultSize) ? candidate.defaultSize : supportedSizes[0]; @@ -1221,11 +1220,3 @@ function _normalizeCustomImageModels( } return models; } - -/** Narrow unknown persisted values to Ark's supported size presets. */ -function _isSeedreamImageSize(value: unknown): value is SeedreamImageSize { - return ( - typeof value === "string" && - (SEEDREAM_IMAGE_SIZES as readonly string[]).includes(value) - ); -} diff --git a/packages/runtime/src/tools/built-in/media.ts b/packages/runtime/src/tools/built-in/media.ts index 2ef29cf1..4d69b0b7 100644 --- a/packages/runtime/src/tools/built-in/media.ts +++ b/packages/runtime/src/tools/built-in/media.ts @@ -4,11 +4,11 @@ import os from "node:os"; import path from "node:path"; import { - SEEDREAM_IMAGE_SIZES, + IMAGE_SIZES, type BuiltinTool, type GenerateImageToolConfig, + type ImageSize, type ProviderConnectionRef, - type SeedreamImageSize, } from "@llm-space/core"; import { expandHomePath } from "@llm-space/core/server"; @@ -18,7 +18,7 @@ export interface MediaBuiltInToolsDependencies { generateImage(input: { prompt: string; model: string; - size: SeedreamImageSize; + size: ImageSize; watermark: boolean; connection?: ProviderConnectionRef; }): Promise<{ @@ -34,7 +34,7 @@ export const generateImageTool: BuiltinTool = { name: "generate_image", icon: "image", description: - "Generate one image with this tool's selected provider and image model. Use the configured default size unless the user requests a supported 1K, 2K, 3K, or 4K preset.", + "Generate one image with this tool's selected provider and image model. Use the configured default size unless the user requests another size supported by that model.", strict: true, parameters: { type: "object", @@ -47,7 +47,7 @@ export const generateImageTool: BuiltinTool = { }, size: { type: "string", - enum: [...SEEDREAM_IMAGE_SIZES], + enum: [...IMAGE_SIZES], description: "Optional resolution preset. Omit it to use the configured default; unsupported presets for the configured model return an error.", }, @@ -80,16 +80,16 @@ export function createMediaBuiltInTools( const size = args.size; if ( size !== undefined && - !SEEDREAM_IMAGE_SIZES.some((candidate) => candidate === size) + !IMAGE_SIZES.some((candidate) => candidate === size) ) { - throw new Error("size must be one of 1K, 2K, 3K, or 4K."); + throw new Error("size must be a supported image size."); } const outputDirectory = args.output_directory; const config = _generateImageConfig(configValue); const result = await dependencies.generateImage({ prompt, model: config.model, - size: (size as SeedreamImageSize | undefined) ?? config.size, + size: (size as ImageSize | undefined) ?? config.size, watermark: config.watermark, connection: context.connection, }); @@ -224,7 +224,7 @@ function _generateImageConfig( "Choose an enabled image model for generate_image in Add built-in tools." ); } - if (!SEEDREAM_IMAGE_SIZES.some((candidate) => candidate === size)) { + if (!IMAGE_SIZES.some((candidate) => candidate === size)) { throw new Error( "Choose a valid default size for generate_image in Add built-in tools." ); @@ -234,5 +234,5 @@ function _generateImageConfig( "Choose a watermark policy for generate_image in Add built-in tools." ); } - return { model, size: size as SeedreamImageSize, watermark }; + return { model, size: size as ImageSize, watermark }; } diff --git a/packages/runtime/tests/models/ark-image-generation.test.ts b/packages/runtime/tests/models/ark-image-generation.test.ts index 22341ade..ff5f0d32 100644 --- a/packages/runtime/tests/models/ark-image-generation.test.ts +++ b/packages/runtime/tests/models/ark-image-generation.test.ts @@ -344,6 +344,135 @@ describe("Ark image generation", () => { }); }); + test("uses standard OpenAI sizes and omits response_format for GPT Image", async () => { + let requestBody: Record | undefined; + const generate = createArkImageGenerator( + _dependencies({ + getConfig: () => ({ + api: "openai-images", + models: [ + { + id: "gpt-image-1", + name: "GPT Image 1", + supportedSizes: ["1024x1024"], + defaultSize: "1024x1024", + }, + ], + }), + fetch: (_input, init) => { + const parsed: unknown = + typeof init?.body === "string" ? JSON.parse(init.body) : undefined; + requestBody = + parsed && typeof parsed === "object" + ? (parsed as Record) + : undefined; + return Promise.resolve( + Response.json({ + model: "gpt-image-1", + data: [{ b64_json: PNG_BASE64, size: "1024x1024" }], + }) + ); + }, + }) + ); + + await generate({ + prompt: "A red circle", + model: "gpt-image-1", + size: "1024x1024", + watermark: false, + connection: { providerId: "openai-images" }, + }); + + expect(requestBody).toEqual({ + model: "gpt-image-1", + prompt: "A red circle", + size: "1024x1024", + }); + }); + + test("requests base64 explicitly for DALL-E and migrates legacy 1K", async () => { + let requestBody: Record | undefined; + const generate = createArkImageGenerator( + _dependencies({ + getConfig: () => ({ + api: "openai-images", + models: [ + { + id: "dall-e-3", + name: "DALL-E 3", + supportedSizes: ["1K"], + defaultSize: "1K", + }, + ], + }), + fetch: (_input, init) => { + const parsed: unknown = + typeof init?.body === "string" ? JSON.parse(init.body) : undefined; + requestBody = + parsed && typeof parsed === "object" + ? (parsed as Record) + : undefined; + return Promise.resolve( + Response.json({ + model: "dall-e-3", + data: [{ b64_json: PNG_BASE64, size: "1024x1024" }], + }) + ); + }, + }) + ); + + await generate({ + prompt: "A red circle", + model: "dall-e-3", + size: "1K", + watermark: false, + connection: { providerId: "openai-images" }, + }); + + expect(requestBody).toEqual({ + model: "dall-e-3", + prompt: "A red circle", + response_format: "b64_json", + size: "1024x1024", + }); + }); + + test("rejects ambiguous non-1K presets in standard OpenAI mode", async () => { + let calls = 0; + const generate = createArkImageGenerator( + _dependencies({ + getConfig: () => ({ + api: "openai-images", + models: [ + { + id: "gpt-image-1", + name: "GPT Image 1", + supportedSizes: ["2K"], + defaultSize: "2K", + }, + ], + }), + fetch: () => { + calls += 1; + return Promise.resolve(Response.json({})); + }, + }) + ); + + expect( + generate({ + prompt: "A red circle", + model: "gpt-image-1", + size: "2K", + watermark: false, + connection: { providerId: "openai-images" }, + }) + ).rejects.toThrow("requires an explicit pixel size"); + expect(calls).toBe(0); + }); + test("never falls back to the Ark key for a custom provider", async () => { let fallbackApiKey: string | undefined; const modelManager = { diff --git a/packages/runtime/tests/models/model-manager.test.ts b/packages/runtime/tests/models/model-manager.test.ts index 2e6bd69a..6ad5f218 100644 --- a/packages/runtime/tests/models/model-manager.test.ts +++ b/packages/runtime/tests/models/model-manager.test.ts @@ -395,9 +395,7 @@ describe("ModelManager provider-owned image generation", () => { const manager = new ModelManager({ settingsDir }); - expect(manager.getImageGenerationConfig("fixture-images")).toEqual({ - api: "openai-images", - }); + expect(manager.getImageGenerationConfig("fixture-images")).toEqual({}); expect(manager.getImageGenerationConfig("openai")).toBeUndefined(); expect(manager.getProfiles("fixture-images")[0]?.baseUrl).toBe( "https://images.example/v1" @@ -468,7 +466,6 @@ describe("ModelManager provider-owned image generation", () => { const reloaded = new ModelManager({ settingsDir }); expect(reloaded.getImageGenerationConfig("agnes")).toEqual({ - api: "openai-images", models: [ { id: "agnes-image-2.5-flash", @@ -483,4 +480,42 @@ describe("ModelManager provider-owned image generation", () => { .find((candidate) => candidate.id === "agnes"); expect(provider?.getModels()).toEqual([]); }); + + test("preserves an explicit Ark image protocol on a custom provider", async () => { + const settingsDir = await _settingsDir({ providers: [] }); + const manager = new ModelManager({ settingsDir }); + manager.addCustomProvider({ + id: "ark-gateway", + name: "Ark Gateway", + baseUrl: "https://gateway.example/v3", + }); + manager.updateProvider("ark-gateway", { + imageGeneration: { + api: "ark-images", + models: [ + { + id: "custom-seedream", + name: "Custom Seedream", + supportedSizes: ["1K", "2K"], + defaultSize: "1K", + }, + ], + }, + }); + + const reloaded = new ModelManager({ settingsDir }); + + expect(reloaded.getImageGenerationConfig("ark-gateway")).toMatchObject({ + api: "ark-images", + }); + const persisted = JSON.parse( + await readFile(path.join(settingsDir, "models.json"), "utf8") + ) as { + providers: { id: string; imageGeneration?: { api?: string } }[]; + }; + expect( + persisted.providers.find((provider) => provider.id === "ark-gateway") + ?.imageGeneration?.api + ).toBe("ark-images"); + }); }); diff --git a/packages/ui/src/components/thread-playground/tool/built-in-tool-import-dialog.tsx b/packages/ui/src/components/thread-playground/tool/built-in-tool-import-dialog.tsx index be10fe40..d7c5a668 100644 --- a/packages/ui/src/components/thread-playground/tool/built-in-tool-import-dialog.tsx +++ b/packages/ui/src/components/thread-playground/tool/built-in-tool-import-dialog.tsx @@ -3,12 +3,12 @@ import { getArkImageModelDefinitions, getImageModelDefinitions, - SEEDREAM_IMAGE_SIZES, + IMAGE_SIZES, type BuiltinTool, type GenerateImageToolConfig, + type ImageSize, type ModelProviderGroup, type SeedreamImageModelDefinition, - type SeedreamImageSize, } from "@llm-space/core"; import { CloudSunIcon, @@ -595,7 +595,7 @@ function _GenerateImageConfigFields({ disabled={!selectedModel} onValueChange={(size) => { if (config) { - onChange({ ...config, size: size as SeedreamImageSize }); + onChange({ ...config, size: size as ImageSize }); } }} > @@ -652,12 +652,12 @@ function _readGenerateImageConfig( const watermark = value?.watermark; if ( typeof model !== "string" || - !SEEDREAM_IMAGE_SIZES.some((candidate) => candidate === size) || + !IMAGE_SIZES.some((candidate) => candidate === size) || typeof watermark !== "boolean" ) { return null; } - return { model, size: size as SeedreamImageSize, watermark }; + return { model, size: size as ImageSize, watermark }; } function _categoryForTool(toolName: string): BuiltInToolCategoryId { From 814298dc35647fe4685acbdbbfbd8e4b259046e2 Mon Sep 17 00:00:00 2001 From: wangguanghao Date: Tue, 8 Sep 2026 23:14:04 +0800 Subject: [PATCH 3/3] fix(models): preserve image aliases and legacy size compatibility --- .../settings/image-model-editor-dialog.tsx | 87 +++- .../provider-image-settings.fixture.tsx | 444 ++++++++++++++++++ .../settings/provider-image-settings.test.ts | 20 + apps/desktop/tsconfig.json | 2 +- .../core/src/types/models/image-generation.ts | 33 +- .../src/models/ark-image-generation.ts | 58 ++- packages/runtime/src/models/model-manager.ts | 11 + .../tests/models/ark-image-generation.test.ts | 225 +++++++++ .../tests/models/model-manager.test.ts | 71 +++ .../tool/built-in-tool-import-dialog.tsx | 44 +- 10 files changed, 939 insertions(+), 56 deletions(-) create mode 100644 apps/desktop/tests/components/settings/provider-image-settings.fixture.tsx create mode 100644 apps/desktop/tests/components/settings/provider-image-settings.test.ts diff --git a/apps/desktop/src/components/settings/image-model-editor-dialog.tsx b/apps/desktop/src/components/settings/image-model-editor-dialog.tsx index ed4aed4c..ad4ada84 100644 --- a/apps/desktop/src/components/settings/image-model-editor-dialog.tsx +++ b/apps/desktop/src/components/settings/image-model-editor-dialog.tsx @@ -1,7 +1,8 @@ "use client"; import { - OPENAI_IMAGE_SIZES, + getOpenAIImageSizes, + normalizeImageSize, SEEDREAM_IMAGE_SIZES, type ImageGenerationApi, type ImageModelDefinition, @@ -34,35 +35,36 @@ interface ImageModelFormState { icon: string; supportedSizes: ImageSize[]; defaultSize: ImageSize; + responseFormat: "auto" | "b64_json"; } /** Create the editable form state for a new or existing image model. */ function _initialState( model: ImageModelDefinition | null | undefined, - sizeOptions: readonly ImageSize[] + api: ImageGenerationApi ): ImageModelFormState { if (model) { - const supportedSizes = model.supportedSizes.filter((size) => - sizeOptions.includes(size) - ); - const normalizedSizes = - supportedSizes.length > 0 ? supportedSizes : [...sizeOptions]; return { id: model.id, name: model.name, icon: model.icon ?? "", - supportedSizes: normalizedSizes, - defaultSize: normalizedSizes.includes(model.defaultSize) - ? model.defaultSize - : normalizedSizes[0], + supportedSizes: [ + ...new Set( + model.supportedSizes.map((size) => normalizeImageSize(size, api)) + ), + ], + defaultSize: normalizeImageSize(model.defaultSize, api), + responseFormat: model.responseFormat ?? "auto", }; } return { id: "", name: "", icon: "", - supportedSizes: [...sizeOptions], - defaultSize: sizeOptions.includes("1024x1024") ? "1024x1024" : "2K", + supportedSizes: + api === "openai-images" ? ["1024x1024"] : [...SEEDREAM_IMAGE_SIZES], + defaultSize: api === "openai-images" ? "1024x1024" : "2K", + responseFormat: "auto", }; } @@ -82,30 +84,38 @@ export function ImageModelEditorDialog({ existingIds: readonly string[]; onSave: (model: ImageModelDefinition, originalId?: string) => void; }) { - const sizeOptions = - api === "openai-images" ? OPENAI_IMAGE_SIZES : SEEDREAM_IMAGE_SIZES; const [form, setForm] = useState(() => - _initialState(model, sizeOptions) + _initialState(model, api) ); useEffect(() => { if (open) { - setForm(_initialState(model, sizeOptions)); + setForm(_initialState(model, api)); } - }, [model, open, sizeOptions]); + }, [model, open, api]); const id = form.id.trim(); + const sizeOptions: readonly ImageSize[] = + api === "openai-images" ? getOpenAIImageSizes(id) : SEEDREAM_IMAGE_SIZES; + const unsupportedSizes = form.supportedSizes.filter( + (size) => !sizeOptions.includes(size) + ); + const displayedSizes = [...sizeOptions, ...unsupportedSizes]; const duplicateId = existingIds.some( (candidate) => candidate === id && candidate !== model?.id ); const canSave = - id.length > 0 && form.supportedSizes.length > 0 && !duplicateId; + id.length > 0 && + form.supportedSizes.length > 0 && + form.supportedSizes.includes(form.defaultSize) && + unsupportedSizes.length === 0 && + !duplicateId; /** Keep the default size valid while the supported-size set changes. */ const handleSizeToggle = (size: ImageSize, enabled: boolean) => { setForm((current) => { const supportedSizes = enabled - ? sizeOptions.filter( + ? displayedSizes.filter( (candidate) => current.supportedSizes.includes(candidate) || candidate === size ) @@ -132,6 +142,9 @@ export function ImageModelEditorDialog({ name: form.name.trim() || id, supportedSizes: form.supportedSizes, defaultSize: form.defaultSize, + ...(form.responseFormat === "b64_json" + ? { responseFormat: "b64_json" as const } + : {}), ...(icon ? { icon } : {}), }, model?.id @@ -210,7 +223,7 @@ export function ImageModelEditorDialog({ <_Field label="Supported sizes">
- {sizeOptions.map((size) => ( + {displayedSizes.map((size) => (
))}
+ {unsupportedSizes.length > 0 && ( +

+ Unsupported sizes for this model and API:{" "} + {unsupportedSizes.join(", ")}. +

+ )} <_Field label="Default size"> @@ -252,6 +271,32 @@ export function ImageModelEditorDialog({ + + {api === "openai-images" && ( + <_Field label="Response format"> + + + )}
diff --git a/apps/desktop/tests/components/settings/provider-image-settings.fixture.tsx b/apps/desktop/tests/components/settings/provider-image-settings.fixture.tsx new file mode 100644 index 00000000..2ae60cf8 --- /dev/null +++ b/apps/desktop/tests/components/settings/provider-image-settings.fixture.tsx @@ -0,0 +1,444 @@ +import { afterAll, afterEach, expect, mock, test } from "bun:test"; +import { mkdtemp, rm } from "node:fs/promises"; +import os from "node:os"; +import path from "node:path"; + +import type { + BuiltinTool, + ImageGenerationApi, + ImageModelDefinition, + ModelProviderGroup, +} from "@llm-space/core"; +import { + act, + createContext, + useContext, + type ComponentProps, + type ReactElement, + type ReactNode, +} from "react"; +import { createRoot, type Root } from "react-dom/client"; + +import { createArkImageGenerator } from "../../../../../packages/runtime/src/models/ark-image-generation"; +import { ModelManager } from "../../../../../packages/runtime/src/models/model-manager"; +import { + installReactTestDom, + TestEvent, + type TestElement, +} from "../../../src/test/react-test-dom"; + +const DOM = installReactTestDom(); +const MOUNTS: { root: Root; container: TestElement }[] = []; +const TEMP_DIRS: string[] = []; +const PNG_BASE64 = + "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNk+M/wHwAF/gL+X2NDWQAAAABJRU5ErkJggg=="; +const LEGACY_MODEL: ImageModelDefinition = { + id: "dall-e-3", + name: "Original", + supportedSizes: ["1K"], + defaultSize: "1K", +}; +const TOOL: BuiltinTool = { + type: "builtin", + name: "generate_image", + description: "Generate an image.", + parameters: { type: "object", properties: {} }, +}; +let providers: ModelProviderGroup[] = []; +const HOST = { builtinTools: { list: () => Promise.resolve([TOOL]) } }; + +function _Container({ children }: { children?: ReactNode }) { + return
{children}
; +} + +const SELECT_CONTEXT = createContext<{ + value?: string; + disabled?: boolean; + onValueChange?: (value: string) => void; +}>({}); + +await mock.module("@llm-space/ui/ui/dialog", () => ({ + Dialog: ({ open, children }: { open: boolean; children?: ReactNode }) => + open ?
{children}
: null, + DialogContent: _Container, + DialogHeader: _Container, + DialogFooter: _Container, + DialogTitle: _Container, + DialogDescription: _Container, +})); +await mock.module("@llm-space/ui/ui/select", () => ({ + Select: ({ + children, + ...state + }: { + children?: ReactNode; + value?: string; + disabled?: boolean; + onValueChange?: (value: string) => void; + }) => ( + {children} + ), + SelectTrigger: function TestSelectTrigger({ + "aria-label": label, + }: { + "aria-label"?: string; + }) { + const { value, disabled } = useContext(SELECT_CONTEXT); + return + ); + }, +})); +await mock.module("@llm-space/ui/ui/switch", () => ({ + Switch: ({ + checked, + disabled, + onCheckedChange, + "aria-label": label, + }: { + checked?: boolean; + disabled?: boolean; + onCheckedChange?: (checked: boolean) => void; + "aria-label"?: string; + }) => ( +