Files
MiragenFlow/web/src/services/api/image.ts
T

391 lines
16 KiB
TypeScript
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
import axios from "axios";
import { buildApiUrl, resolveModelRequestConfig, type AiConfig, type ModelChannel } from "@/stores/use-config-store";
import { nanoid } from "nanoid";
import { dataUrlToFile } from "@/lib/image-utils";
import { buildImageReferencePromptText } from "@/lib/image-reference-prompt";
import { imageToDataUrl } from "@/services/image-storage";
import type { ReferenceImage } from "@/types/image";
export type AiTextMessage = {
role: "system" | "user" | "assistant";
content: string | Array<{ type: "text"; text: string } | { type: "image_url"; image_url: { url: string } }>;
};
export type ResponseToolCall = {
id: string;
type: "function";
function: { name: string; arguments: string };
};
export type ResponseInputMessage =
| AiTextMessage
| { type: "function_call"; call_id: string; name: string; arguments: string }
| { role: "tool"; tool_call_id: string; content: string };
export type ResponseFunctionTool = {
type: "function";
function: {
name: string;
description?: string;
parameters: Record<string, unknown>;
strict?: boolean;
};
};
export type ToolResponseResult = {
content: string;
toolCalls: ResponseToolCall[];
};
type ToolChoice = "auto" | "required" | { type: "function"; name: string };
type ResponseMessageContent = AiTextMessage["content"] | string;
type ResponseInputContent = { type: "input_text"; text: string } | { type: "input_image"; image_url: string };
type ResponseInputItem =
| { role: "system" | "user" | "assistant"; content: string | ResponseInputContent[] }
| { type: "function_call"; call_id: string; name: string; arguments: string }
| { type: "function_call_output"; call_id: string; output: string };
type ResponseApiToolDefinition = {
type: "function";
name: string;
description?: string;
parameters: Record<string, unknown>;
strict?: boolean;
};
type ResponseApiOutputItem =
| { type?: "message"; content?: Array<{ type?: string; text?: string }> }
| { type?: "function_call"; id?: string; call_id?: string; name?: string; arguments?: string };
type ResponseApiPayload = {
id?: string;
output?: ResponseApiOutputItem[];
output_text?: string;
error?: { message?: string };
code?: number;
msg?: string;
};
type ImageApiResponse = {
data?: Array<Record<string, unknown>>;
error?: { message?: string };
code?: number;
msg?: string;
};
const QUALITY_BASE: Record<string, number> = {
low: 1024,
medium: 2048,
high: 2880,
standard: 1024,
hd: 2048,
};
const QUALITY_ALIASES: Record<string, string> = {
"1k": "low",
"2k": "medium",
"4k": "high",
};
const DEFAULT_IMAGE_SHORT_SIDE = 1024;
const IMAGE_SIZE_STEP = 16;
const IMAGE_MIN_PIXELS = 655360;
const IMAGE_MAX_PIXELS = 8294400;
const IMAGE_MAX_EDGE = 3840;
const IMAGE_MAX_RATIO = 3;
const IMAGE_OUTPUT_FORMAT = "png";
function normalizeQuality(quality: string) {
const value = quality.trim().toLowerCase();
const normalized = QUALITY_ALIASES[value] || value;
return QUALITY_BASE[normalized] ? normalized : undefined;
}
/** Map "quality + ratio" to an explicit pixel dimension like "3840x2160". */
function resolveSize(quality: string | undefined, ratio: string): string {
const parsedRatio = parseImageRatio(ratio);
const basePixels = quality ? QUALITY_BASE[quality] : undefined;
const isLandscape = parsedRatio.width >= parsedRatio.height;
const longRatio = isLandscape ? parsedRatio.width / parsedRatio.height : parsedRatio.height / parsedRatio.width;
let longSide: number;
let shortSide: number;
if (basePixels) {
const targetPixels = basePixels * basePixels;
const longSideRaw = Math.sqrt(targetPixels * longRatio);
longSide = Math.floor(longSideRaw / IMAGE_SIZE_STEP) * IMAGE_SIZE_STEP;
shortSide = Math.round(longSide / longRatio / IMAGE_SIZE_STEP) * IMAGE_SIZE_STEP;
} else {
shortSide = DEFAULT_IMAGE_SHORT_SIDE;
longSide = Math.round((shortSide * longRatio) / IMAGE_SIZE_STEP) * IMAGE_SIZE_STEP;
}
const width = isLandscape ? longSide : shortSide;
const height = isLandscape ? shortSide : longSide;
validateImageSize(width, height);
return `${width}x${height}`;
}
function parseImageRatio(value: string) {
const parts = value.split(":");
if (parts.length !== 2) throw new Error("图像尺寸格式不支持,请使用 auto、9:16 或 1024x1024");
const w = Number(parts[0]);
const h = Number(parts[1]);
if (!Number.isFinite(w) || !Number.isFinite(h) || w <= 0 || h <= 0) throw new Error("图像比例必须是正数,例如 9:16");
if (Math.max(w, h) / Math.min(w, h) > IMAGE_MAX_RATIO) throw new Error("图像宽高比不能超过 3:1,请调整尺寸");
return { width: w, height: h };
}
function parseImageDimensions(value: string) {
const match = value.match(/^(\d+)x(\d+)$/i);
if (!match) return null;
return { width: Number(match[1]), height: Number(match[2]) };
}
function validateImageSize(width: number, height: number) {
if (!Number.isInteger(width) || !Number.isInteger(height) || width <= 0 || height <= 0) throw new Error("图像尺寸必须是正整数,例如 1024x1024");
if (width % IMAGE_SIZE_STEP !== 0 || height % IMAGE_SIZE_STEP !== 0) throw new Error("图像尺寸的宽高必须是 16 的倍数,请调整尺寸");
if (Math.max(width, height) > IMAGE_MAX_EDGE) throw new Error("图像尺寸最长边不能超过 3840px,请调整尺寸");
if (Math.max(width, height) / Math.min(width, height) > IMAGE_MAX_RATIO) throw new Error("图像宽高比不能超过 3:1,请调整尺寸");
const pixels = width * height;
if (pixels < IMAGE_MIN_PIXELS || pixels > IMAGE_MAX_PIXELS) throw new Error("图像总像素需在 655360 到 8294400 之间,请调整尺寸");
}
function resolveRequestSize(quality: string | undefined, size: string) {
const value = size.trim();
if (!value || value.toLowerCase() === "auto") return undefined;
const dimensions = parseImageDimensions(value);
if (dimensions) {
validateImageSize(dimensions.width, dimensions.height);
return `${dimensions.width}x${dimensions.height}`;
}
if (value.includes(":")) return resolveSize(quality, value);
throw new Error("图像尺寸格式不支持,请使用 auto、9:16 或 1024x1024");
}
function resolveImageDataUrl(item: Record<string, unknown>) {
if (typeof item.b64_json === "string" && item.b64_json) {
return `data:image/png;base64,${item.b64_json}`;
}
if (typeof item.url === "string" && item.url) {
return item.url;
}
return null;
}
function parseImagePayload(payload: ImageApiResponse) {
if (typeof payload.code === "number" && payload.code !== 0) {
throw new Error(payload.msg || "请求失败");
}
const images =
payload.data
?.map(resolveImageDataUrl)
.filter((value): value is string => Boolean(value))
.map((dataUrl) => ({ id: nanoid(), dataUrl })) || [];
if (images.length === 0) {
throw new Error("接口没有返回图片");
}
return images;
}
function readAxiosError(error: unknown, fallback: string) {
if (axios.isAxiosError<{ error?: { message?: string }; msg?: string; code?: number }>(error)) {
const responseData = error.response?.data;
return responseData?.msg || responseData?.error?.message || readStatusError(error.response?.status, fallback);
}
return error instanceof Error ? error.message : fallback;
}
function readStatusError(status: number | undefined, fallback: string) {
if (status === 401 || status === 403) return "鉴权失败,请检查 API Key、套餐权限或模型权限";
if (status === 429) return "请求被限流或额度不足,请稍后重试";
return status ? `${fallback}:${status}` : fallback;
}
function withSystemPrompt(config: AiConfig, prompt: string) {
const systemPrompt = config.systemPrompt.trim();
return systemPrompt ? `${systemPrompt}\n\n${prompt}` : prompt;
}
function aiApiUrl(config: AiConfig, path: string) {
return buildApiUrl(config.baseUrl, path);
}
function aiHeaders(config: AiConfig, contentType?: string) {
return {
Authorization: `Bearer ${config.apiKey}`,
...(contentType ? { "Content-Type": contentType } : {}),
};
}
function withSystemMessage<T extends ResponseInputMessage>(config: AiConfig, messages: T[]): ResponseInputMessage[] {
const systemPrompt = config.systemPrompt.trim();
return systemPrompt ? [{ role: "system" as const, content: systemPrompt }, ...messages] : messages;
}
function toResponseInput(messages: ResponseInputMessage[]): ResponseInputItem[] {
return messages.flatMap((message): ResponseInputItem[] => {
if ("type" in message) return [message];
if (message.role === "tool") return [{ type: "function_call_output", call_id: message.tool_call_id, output: message.content }];
return [{ role: message.role, content: toResponseContent(message.content || "") }];
});
}
function toResponseContent(content: ResponseMessageContent): string | ResponseInputContent[] {
if (!Array.isArray(content)) return String(content || "");
return content.map((item) => (item.type === "text" ? { type: "input_text" as const, text: item.text } : { type: "input_image" as const, image_url: item.image_url.url }));
}
function toResponseTool(tool: ResponseFunctionTool): ResponseApiToolDefinition {
return {
type: "function",
name: tool.function.name,
description: tool.function.description,
parameters: tool.function.parameters,
strict: tool.function.strict,
};
}
function parseToolResponse(payload: ResponseApiPayload): ToolResponseResult {
const output = payload.output || [];
const content =
payload.output_text ||
output
.flatMap((item) => (item.type === "message" ? item.content || [] : []))
.map((item) => item.text || "")
.join("");
const toolCalls = output
.filter((item): item is Extract<ResponseApiOutputItem, { type?: "function_call" }> => item.type === "function_call")
.map((item) => ({
id: item.call_id || item.id || "",
type: "function" as const,
function: { name: item.name || "", arguments: item.arguments || "{}" },
}))
.filter((item) => item.id && item.function.name);
return { content, toolCalls };
}
export async function requestGeneration(config: AiConfig, prompt: string) {
const requestConfig = resolveModelRequestConfig(config, config.model || config.imageModel);
const n = Math.max(1, Math.min(15, Math.floor(Math.abs(Number(config.count)) || 1)));
const quality = normalizeQuality(config.quality);
const requestSize = resolveRequestSize(quality, config.size);
try {
const response = await axios.post<ImageApiResponse>(
aiApiUrl(requestConfig, "/images/generations"),
{
model: requestConfig.model,
prompt: withSystemPrompt(requestConfig, prompt),
n,
...(quality ? { quality } : {}),
...(requestSize ? { size: requestSize } : {}),
response_format: "b64_json",
output_format: IMAGE_OUTPUT_FORMAT,
},
{
headers: aiHeaders(requestConfig, "application/json"),
},
);
const images = parseImagePayload(response.data);
return images;
} catch (error) {
throw new Error(readAxiosError(error, "请求失败"));
}
}
export async function requestEdit(config: AiConfig, prompt: string, references: ReferenceImage[], mask?: ReferenceImage) {
const requestConfig = resolveModelRequestConfig(config, config.model || config.imageModel);
const n = Math.max(1, Math.min(15, Math.floor(Math.abs(Number(config.count)) || 1)));
const quality = normalizeQuality(config.quality);
const requestSize = resolveRequestSize(quality, config.size);
const requestPrompt = buildImageReferencePromptText(prompt, references);
const formData = new FormData();
formData.set("model", requestConfig.model);
formData.set("prompt", withSystemPrompt(requestConfig, requestPrompt));
formData.set("n", String(n));
formData.set("response_format", "b64_json");
formData.set("output_format", IMAGE_OUTPUT_FORMAT);
if (quality) {
formData.set("quality", quality);
}
if (requestSize) {
formData.set("size", requestSize);
}
const files = await Promise.all(references.map(async (image) => dataUrlToFile({ ...image, dataUrl: await imageToDataUrl(image) })));
files.forEach((file) => formData.append("image", file));
if (mask) formData.set("mask", dataUrlToFile(mask));
try {
const response = await axios.post<ImageApiResponse>(aiApiUrl(requestConfig, "/images/edits"), formData, { headers: aiHeaders(requestConfig) });
const images = parseImagePayload(response.data);
return images;
} catch (error) {
throw new Error(readAxiosError(error, "请求失败"));
}
}
export async function requestImageQuestion(config: AiConfig, messages: AiTextMessage[], onDelta: (text: string) => void) {
const requestConfig = resolveModelRequestConfig(config, config.model || config.textModel);
try {
const response = await axios.post<ResponseApiPayload>(
aiApiUrl(requestConfig, "/responses"),
{
model: requestConfig.model,
input: toResponseInput(withSystemMessage(requestConfig, messages)),
},
{
headers: aiHeaders(requestConfig, "application/json"),
},
);
if (typeof response.data.code === "number" && response.data.code !== 0) throw new Error(response.data.msg || "请求失败");
if (response.data.error?.message) throw new Error(response.data.error.message);
const answer = parseToolResponse(response.data).content || "没有返回内容";
onDelta(answer);
return answer;
} catch (error) {
throw new Error(readAxiosError(error, "请求失败"));
}
}
export async function requestToolResponse(config: AiConfig, messages: ResponseInputMessage[], tools: ResponseFunctionTool[], toolChoice: ToolChoice = "auto"): Promise<ToolResponseResult> {
const requestConfig = resolveModelRequestConfig(config, config.model || config.textModel);
try {
const response = await axios.post<ResponseApiPayload>(
aiApiUrl(requestConfig, "/responses"),
{
model: requestConfig.model,
input: toResponseInput(withSystemMessage(requestConfig, messages)),
tools: tools.map(toResponseTool),
tool_choice: toolChoice,
parallel_tool_calls: false,
},
{
headers: aiHeaders(requestConfig, "application/json"),
},
);
if (typeof response.data.code === "number" && response.data.code !== 0) throw new Error(response.data.msg || "请求失败");
if (response.data.error?.message) throw new Error(response.data.error.message);
return parseToolResponse(response.data);
} catch (error) {
throw new Error(readAxiosError(error, "请求失败"));
}
}
export async function fetchImageModels(config: Pick<AiConfig, "baseUrl" | "apiKey">) {
try {
const response = await axios.get<{ data?: Array<{ id?: string }>; error?: { message?: string } }>(buildApiUrl(config.baseUrl, "/models"), {
headers: {
Authorization: `Bearer ${config.apiKey}`,
},
});
return (response.data.data || [])
.map((model) => model.id)
.filter((id): id is string => Boolean(id))
.sort((a, b) => a.localeCompare(b));
} catch (error) {
throw new Error(readAxiosError(error, "读取模型失败"));
}
}
export async function fetchChannelModels(channel: ModelChannel) {
return fetchImageModels({ baseUrl: channel.baseUrl, apiKey: channel.apiKey });
}