391 lines
16 KiB
TypeScript
391 lines
16 KiB
TypeScript
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 });
|
||
}
|