feat(model): enhance model configuration and script handling for audio and video capabilities

This commit is contained in:
HouYunFei
2026-07-15 10:13:32 +08:00
parent a4dcc679c9
commit c57f7d61a7
13 changed files with 690 additions and 266 deletions
+66 -70
View File
@@ -4,6 +4,13 @@ import { persist } from "zustand/middleware";
import { nanoid } from "nanoid";
export type ApiCallFormat = "openai" | "gemini";
export type ModelCapability = "image" | "video" | "text" | "audio";
export type ChannelModel = {
name: string;
capability: ModelCapability;
script?: string;
};
export type ModelChannel = {
id: string;
@@ -11,7 +18,7 @@ export type ModelChannel = {
baseUrl: string;
apiKey: string;
apiFormat: ApiCallFormat;
models: string[];
models: ChannelModel[];
};
export type AiConfig = {
@@ -35,10 +42,6 @@ export type AiConfig = {
videoWatermark: string;
systemPrompt: string;
models: string[];
imageModels: string[];
videoModels: string[];
textModels: string[];
audioModels: string[];
quality: string;
size: string;
count: string;
@@ -52,10 +55,9 @@ export type WebdavSyncConfig = {
directory: string;
lastSyncedAt: string;
};
export type ConfigTabKey = "channels" | "models" | "preferences" | "webdav";
export type ConfigTabKey = "channels" | "preferences" | "webdav";
export const CONFIG_STORE_KEY = "infinite-canvas:ai_config_store";
export type ModelCapability = "image" | "video" | "text" | "audio";
const CHANNEL_MODEL_SEPARATOR = "::";
const OPENAI_BASE_URL = "https://api.openai.com";
const GEMINI_BASE_URL = "https://generativelanguage.googleapis.com";
@@ -72,7 +74,12 @@ export const defaultConfig: AiConfig = {
baseUrl: OPENAI_BASE_URL,
apiKey: "",
apiFormat: "openai",
models: ["gpt-image-2", "grok-imagine-video", "gpt-5.5", "gpt-4o-mini-tts"],
models: [
{ name: "gpt-image-2", capability: "image" },
{ name: "grok-imagine-video", capability: "video" },
{ name: "gpt-5.5", capability: "text" },
{ name: "gpt-4o-mini-tts", capability: "audio" },
],
},
],
model: "default::gpt-image-2",
@@ -90,10 +97,6 @@ export const defaultConfig: AiConfig = {
videoWatermark: "false",
systemPrompt: "",
models: ["default::gpt-image-2", "default::grok-imagine-video", "default::gpt-5.5", "default::gpt-4o-mini-tts"],
imageModels: ["default::gpt-image-2"],
videoModels: ["default::grok-imagine-video"],
textModels: ["default::gpt-5.5"],
audioModels: ["default::gpt-4o-mini-tts"],
quality: "auto",
size: "1:1",
count: "1",
@@ -122,44 +125,44 @@ type ConfigStore = {
clearPromptContinue: () => void;
};
function isVideoModelName(model: string) {
const value = modelOptionName(model).toLowerCase();
return value.includes("seedance") || value.includes("video") || value.includes("sora") || value.includes("veo") || value.includes("kling") || value.includes("wan") || value.includes("hailuo");
const VIDEO_KEYWORDS = ["seedance", "video", "sora", "veo", "kling", "wan", "hailuo"];
const AUDIO_KEYWORDS = ["audio", "tts", "speech", "voice", "music", "sound"];
const IMAGE_KEYWORDS = ["seedream", "gpt-image", "image", "dall-e", "dalle", "imagen", "flux", "sdxl", "stable-diffusion", "midjourney"];
/** Best-effort default capability for a freshly fetched model name; user can override in the channel editor. */
export function guessCapability(name: string): ModelCapability {
const value = name.toLowerCase();
if (VIDEO_KEYWORDS.some((keyword) => value.includes(keyword))) return "video";
if (AUDIO_KEYWORDS.some((keyword) => value.includes(keyword))) return "audio";
if (IMAGE_KEYWORDS.some((keyword) => value.includes(keyword))) return "image";
return "text";
}
function isImageModelName(model: string) {
const value = modelOptionName(model).toLowerCase();
return !isVideoModelName(model) && !isAudioModelName(model) && (value.includes("seedream") || value.includes("gpt-image") || value.includes("image") || value.includes("dall-e") || value.includes("dalle") || value.includes("imagen") || value.includes("flux") || value.includes("sdxl") || value.includes("stable-diffusion") || value.includes("midjourney"));
function findChannelModel(config: AiConfig, value: string): { channel: ModelChannel; model: ChannelModel } | null {
const decoded = decodeChannelModel(value);
const name = decoded?.model || value;
const channel = decoded ? config.channels.find((item) => item.id === decoded.channelId) : config.channels.find((item) => item.models.some((model) => model.name === name));
const model = channel?.models.find((item) => item.name === name);
return channel && model ? { channel, model } : null;
}
function isAudioModelName(model: string) {
const value = modelOptionName(model).toLowerCase();
return value.includes("audio") || value.includes("tts") || value.includes("speech") || value.includes("voice") || value.includes("music") || value.includes("sound");
export function modelCapabilityOf(config: AiConfig, value: string): ModelCapability | undefined {
return findChannelModel(config, value)?.model.capability;
}
function isTextModelName(model: string) {
return !isImageModelName(model) && !isVideoModelName(model) && !isAudioModelName(model);
}
export function modelMatchesCapability(model: string, capability?: ModelCapability) {
export function modelMatchesCapability(config: AiConfig, value: string, capability?: ModelCapability) {
if (!capability) return true;
if (capability === "image") return isImageModelName(model);
if (capability === "video") return isVideoModelName(model);
if (capability === "audio") return isAudioModelName(model);
return isTextModelName(model);
}
export function filterModelsByCapability(models: string[], capability?: ModelCapability) {
return capability ? models.filter((model) => modelMatchesCapability(model, capability)) : models;
return modelCapabilityOf(config, value) === capability;
}
export function selectableModelsByCapability(config: AiConfig, capability?: ModelCapability) {
if (!capability) return config.models;
return config[modelListKey(capability)];
return config.channels.flatMap((channel) => channel.models.filter((model) => model.capability === capability).map((model) => encodeChannelModel(channel.id, model.name)));
}
function modelListKey(capability: ModelCapability) {
return `${capability}Models` as "imageModels" | "videoModels" | "textModels" | "audioModels";
/** The user script (if any) attached to a model; empty string means use the system default call. */
export function resolveModelScript(config: AiConfig, value: string) {
return findChannelModel(config, value)?.model.script?.trim() || "";
}
function isAiConfigReady(config: AiConfig, model: string) {
@@ -215,7 +218,7 @@ export const useConfigStore = create<ConfigStore>()(
channels,
models,
imageModel: normalizeModelOptionValue(config.imageModel || config.model, channels),
videoModel: normalizeModelOptionValue(config.videoModel || "grok-imagine-video", channels),
videoModel: normalizeModelOptionValue(config.videoModel, channels),
textModel: normalizeModelOptionValue(config.textModel || config.model, channels),
audioModel: normalizeModelOptionValue(config.audioModel || defaultConfig.audioModel, channels),
audioVoice: config.audioVoice || defaultConfig.audioVoice,
@@ -227,10 +230,6 @@ export const useConfigStore = create<ConfigStore>()(
videoGenerateAudio: config.videoGenerateAudio || "true",
videoWatermark: config.videoWatermark || "false",
canvasImageCount: config.canvasImageCount || "3",
imageModels: Array.isArray(persistedConfig.imageModels) ? normalizeModelList(config.imageModels, channels) : filterModelsByCapability(models, "image"),
videoModels: Array.isArray(persistedConfig.videoModels) ? normalizeModelList(config.videoModels, channels) : filterModelsByCapability(models, "video"),
textModels: Array.isArray(persistedConfig.textModels) ? normalizeModelList(config.textModels, channels) : filterModelsByCapability(models, "text"),
audioModels: Array.isArray(persistedConfig.audioModels) ? normalizeModelList(config.audioModels, channels) : filterModelsByCapability(models, "audio"),
},
};
},
@@ -238,18 +237,26 @@ export const useConfigStore = create<ConfigStore>()(
),
);
function normalizeModelList(models: string[], channels: ModelChannel[]) {
const allModelOptions = channels.flatMap((channel) => channel.models.map((model) => encodeChannelModel(channel.id, model)));
return Array.from(new Set((models || []).map((model) => model.trim()).filter(Boolean)))
.map((model) => normalizeModelOptionValue(model, channels))
.filter((model) => !allModelOptions.length || allModelOptions.includes(model) || !isChannelModelValue(model));
}
export function useEffectiveConfig() {
const config = useConfigStore((state) => state.config);
return useMemo(() => ({ ...config, channelMode: "local" as const }), [config]);
}
/** Normalize a mixed list of raw model names or model objects into deduped ChannelModel entries. */
export function normalizeChannelModels(models: Array<string | ChannelModel> | undefined): ChannelModel[] {
const seen = new Set<string>();
const result: ChannelModel[] = [];
for (const item of models || []) {
const name = (typeof item === "string" ? item : item?.name || "").trim();
if (!name || seen.has(name)) continue;
seen.add(name);
const capability = typeof item === "string" ? guessCapability(name) : item.capability || guessCapability(name);
const script = typeof item === "string" ? undefined : item.script?.trim() || undefined;
result.push({ name, capability, script });
}
return result;
}
export function createModelChannel(channel?: Partial<ModelChannel>): ModelChannel {
const apiFormat = normalizeApiFormat(channel?.apiFormat);
return {
@@ -258,7 +265,7 @@ export function createModelChannel(channel?: Partial<ModelChannel>): ModelChanne
baseUrl: channel?.baseUrl?.trim() || defaultBaseUrlForApiFormat(apiFormat),
apiKey: channel?.apiKey || "",
apiFormat,
models: uniqueRawModels(channel?.models || []),
models: normalizeChannelModels(channel?.models),
};
}
@@ -288,7 +295,7 @@ export function modelOptionLabel(config: AiConfig, value: string) {
}
export function modelOptionsFromChannels(channels: ModelChannel[]) {
return uniqueModelOptions(channels.flatMap((channel) => channel.models.map((model) => encodeChannelModel(channel.id, model))));
return uniqueModelOptions(channels.flatMap((channel) => channel.models.map((model) => encodeChannelModel(channel.id, model.name))));
}
export function normalizeModelOptionValue(value: string | undefined, channels: ModelChannel[]) {
@@ -297,17 +304,17 @@ export function normalizeModelOptionValue(value: string | undefined, channels: M
const decoded = decodeChannelModel(model);
if (decoded) {
const channel = channels.find((item) => item.id === decoded.channelId);
return channel && channel.models.includes(decoded.model) ? model : "";
return channel && channel.models.some((item) => item.name === decoded.model) ? model : "";
}
const channel = channels.find((item) => item.models.includes(model)) || channels[0];
return channel && channel.models.includes(model) ? encodeChannelModel(channel.id, model) : model;
const channel = channels.find((item) => item.models.some((entry) => entry.name === model)) || channels[0];
return channel && channel.models.some((item) => item.name === model) ? encodeChannelModel(channel.id, model) : model;
}
export function resolveModelChannel(config: AiConfig, value: string) {
const decoded = decodeChannelModel(value);
const model = decoded?.model || value;
const matched = decoded ? config.channels.find((channel) => channel.id === decoded.channelId) : config.channels.find((channel) => channel.models.includes(model));
return matched || config.channels[0] || createModelChannel({ id: "default", name: "默认渠道", baseUrl: config.baseUrl, apiKey: config.apiKey, apiFormat: config.apiFormat, models: config.models.map(modelOptionName) });
const matched = decoded ? config.channels.find((channel) => channel.id === decoded.channelId) : config.channels.find((channel) => channel.models.some((item) => item.name === model));
return matched || config.channels[0] || createModelChannel({ id: "default", name: "默认渠道", baseUrl: config.baseUrl, apiKey: config.apiKey, apiFormat: config.apiFormat, models: config.models.map(modelOptionName).map((name) => ({ name, capability: guessCapability(name) })) });
}
export function resolveModelRequestConfig(config: AiConfig, value: string) {
@@ -328,7 +335,7 @@ function normalizeChannels(config: AiConfig) {
...channel,
id: channel.id || (index === 0 ? "default" : `channel-${index + 1}`),
name: channel.name || (index === 0 ? "默认渠道" : `渠道 ${index + 1}`),
models: uniqueRawModels(channel.models || []),
models: normalizeChannelModels(channel.models),
}),
);
if (!channels.length) {
@@ -339,18 +346,11 @@ function normalizeChannels(config: AiConfig) {
baseUrl: config.baseUrl || defaultConfig.baseUrl,
apiKey: config.apiKey || "",
apiFormat: config.apiFormat || defaultConfig.apiFormat,
models: uniqueRawModels([
...(config.models || []),
config.model,
config.imageModel,
config.videoModel,
config.textModel,
config.audioModel,
]),
models: normalizeChannelModels([config.model, config.imageModel, config.videoModel, config.textModel, config.audioModel].map(modelOptionName)),
}),
);
}
return channels.map((channel) => ({ ...channel, models: uniqueRawModels(channel.models) }));
return channels;
}
export function defaultBaseUrlForApiFormat(apiFormat: ApiCallFormat) {
@@ -361,10 +361,6 @@ function normalizeApiFormat(apiFormat: unknown): ApiCallFormat {
return apiFormat === "gemini" ? "gemini" : "openai";
}
function uniqueRawModels(models: string[]) {
return Array.from(new Set((models || []).map((model) => modelOptionName(model).trim()).filter(Boolean)));
}
function uniqueModelOptions(models: string[]) {
return Array.from(new Set((models || []).map((model) => model.trim()).filter(Boolean)));
}