Add baoyu-skills package
This commit is contained in:
@@ -0,0 +1,213 @@
|
||||
import assert from "node:assert/strict";
|
||||
import fs from "node:fs/promises";
|
||||
import os from "node:os";
|
||||
import path from "node:path";
|
||||
import test, { type TestContext } from "node:test";
|
||||
|
||||
import type { CliArgs } from "../types.ts";
|
||||
import {
|
||||
buildRequestBody,
|
||||
extractImageFromResponse,
|
||||
parseAspectRatio,
|
||||
resolveReferenceImages,
|
||||
resolveSize,
|
||||
snapDim,
|
||||
validateArgs,
|
||||
} from "./agnes.ts";
|
||||
|
||||
function useEnv(
|
||||
t: TestContext,
|
||||
values: Record<string, string | null>,
|
||||
): void {
|
||||
const previous = new Map<string, string | undefined>();
|
||||
for (const [key, value] of Object.entries(values)) {
|
||||
previous.set(key, process.env[key]);
|
||||
if (value == null) {
|
||||
delete process.env[key];
|
||||
} else {
|
||||
process.env[key] = value;
|
||||
}
|
||||
}
|
||||
|
||||
t.after(() => {
|
||||
for (const [key, value] of previous.entries()) {
|
||||
if (value == null) {
|
||||
delete process.env[key];
|
||||
} else {
|
||||
process.env[key] = value;
|
||||
}
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
function makeArgs(overrides: Partial<CliArgs> = {}): CliArgs {
|
||||
return {
|
||||
prompt: null,
|
||||
promptFiles: [],
|
||||
imagePath: null,
|
||||
provider: null,
|
||||
model: null,
|
||||
aspectRatio: null,
|
||||
size: null,
|
||||
quality: null,
|
||||
imageSize: null,
|
||||
imageApiDialect: null,
|
||||
referenceImages: [],
|
||||
n: 1,
|
||||
batchFile: null,
|
||||
jobs: null,
|
||||
json: false,
|
||||
help: false,
|
||||
responseFormat: null,
|
||||
...overrides,
|
||||
};
|
||||
}
|
||||
|
||||
test("snapDim rounds to the nearest multiple of 32", () => {
|
||||
assert.equal(snapDim(767), 768);
|
||||
assert.equal(snapDim(1023), 1024);
|
||||
assert.equal(snapDim(1024), 1024);
|
||||
assert.equal(snapDim(32), 32);
|
||||
assert.equal(snapDim(0), 32);
|
||||
assert.equal(snapDim(16), 32);
|
||||
assert.equal(snapDim(48), 64);
|
||||
});
|
||||
|
||||
test("parseAspectRatio parses valid ratios and rejects invalid inputs", () => {
|
||||
assert.deepEqual(parseAspectRatio("3:4"), { width: 3, height: 4 });
|
||||
assert.deepEqual(parseAspectRatio("16:9"), { width: 16, height: 9 });
|
||||
assert.deepEqual(parseAspectRatio("1:1"), { width: 1, height: 1 });
|
||||
assert.deepEqual(parseAspectRatio("1.5:1"), { width: 1.5, height: 1 });
|
||||
|
||||
assert.equal(parseAspectRatio(""), null);
|
||||
assert.equal(parseAspectRatio("invalid"), null);
|
||||
assert.equal(parseAspectRatio("3x4"), null);
|
||||
assert.equal(parseAspectRatio("0:1"), null);
|
||||
assert.equal(parseAspectRatio("1:0"), null);
|
||||
});
|
||||
|
||||
test("resolveSize returns explicit --size directly", () => {
|
||||
assert.equal(resolveSize({ size: "1024x1024" }), "1024x1024");
|
||||
assert.equal(resolveSize({ size: "768x1024", aspectRatio: "16:9" }), "768x1024");
|
||||
});
|
||||
|
||||
test("resolveSize returns default 1024x1024 when no size or ratio given", () => {
|
||||
assert.equal(resolveSize({}), "1024x1024");
|
||||
assert.equal(resolveSize({ size: null, aspectRatio: null }), "1024x1024");
|
||||
});
|
||||
|
||||
test("resolveSize computes 32-aligned size within 2048 max edge", () => {
|
||||
assert.equal(resolveSize({ aspectRatio: "1:1" }), "1024x1024");
|
||||
assert.equal(resolveSize({ aspectRatio: "16:9" }), "2048x1152");
|
||||
assert.equal(resolveSize({ aspectRatio: "4:3" }), "2048x1536");
|
||||
assert.equal(resolveSize({ aspectRatio: "3:4" }), "1536x2048");
|
||||
assert.equal(resolveSize({ aspectRatio: "9:16" }), "1152x2048");
|
||||
});
|
||||
|
||||
test("resolveSize aligns to 32 and respects max edge", () => {
|
||||
assert.equal(resolveSize({ aspectRatio: "3:1" }), "2048x672");
|
||||
assert.equal(resolveSize({ aspectRatio: "1:3" }), "672x2048");
|
||||
});
|
||||
|
||||
test("validateArgs rejects --n > 1", () => {
|
||||
assert.throws(
|
||||
() => validateArgs("agnes-image-2.5-flash", makeArgs({ n: 2 })),
|
||||
/returns a single image per request/,
|
||||
);
|
||||
assert.doesNotThrow(() =>
|
||||
validateArgs("agnes-image-2.5-flash", makeArgs({ n: 1 })),
|
||||
);
|
||||
});
|
||||
|
||||
test("buildRequestBody maps prompt, model, size, and reference images", () => {
|
||||
const body = buildRequestBody("a cat", "agnes-image-2.5-flash", {
|
||||
size: "1024x1024",
|
||||
aspectRatio: null,
|
||||
referenceImages: [],
|
||||
});
|
||||
assert.equal(body.model, "agnes-image-2.5-flash");
|
||||
assert.equal(body.prompt, "a cat");
|
||||
assert.equal(body.size, "1024x1024");
|
||||
assert.deepEqual(body.extra_body, { response_format: "url" });
|
||||
|
||||
const bodyWithRef = buildRequestBody("a cat", "agnes-image-2.5-flash", {
|
||||
size: null,
|
||||
aspectRatio: "3:4",
|
||||
referenceImages: ["https://example.com/ref.jpg"],
|
||||
});
|
||||
assert.equal(bodyWithRef.size, "1536x2048");
|
||||
assert.deepEqual(bodyWithRef.image, ["https://example.com/ref.jpg"]);
|
||||
});
|
||||
|
||||
test("extractImageFromResponse decodes b64_json payloads", async () => {
|
||||
const fromBase64 = await extractImageFromResponse({
|
||||
data: [{ b64_json: Buffer.from("hello").toString("base64") }],
|
||||
});
|
||||
assert.equal(Buffer.from(fromBase64).toString("utf8"), "hello");
|
||||
});
|
||||
|
||||
test("extractImageFromResponse downloads URL payloads", async (t) => {
|
||||
const originalFetch = globalThis.fetch;
|
||||
t.after(() => {
|
||||
globalThis.fetch = originalFetch;
|
||||
});
|
||||
|
||||
globalThis.fetch = async () =>
|
||||
new Response(Uint8Array.from([1, 2, 3]), {
|
||||
status: 200,
|
||||
headers: { "Content-Type": "image/png" },
|
||||
});
|
||||
|
||||
const fromUrl = await extractImageFromResponse({
|
||||
data: [{ url: "https://example.com/output.png" }],
|
||||
});
|
||||
assert.deepEqual([...fromUrl], [1, 2, 3]);
|
||||
});
|
||||
|
||||
test("extractImageFromResponse throws on empty data", async () => {
|
||||
await assert.rejects(
|
||||
() => extractImageFromResponse({ data: [] }),
|
||||
/No image/,
|
||||
);
|
||||
await assert.rejects(
|
||||
() => extractImageFromResponse({ data: [{}] }),
|
||||
/No image/,
|
||||
);
|
||||
});
|
||||
|
||||
test("resolveReferenceImages converts local files to data URIs and passes URLs through", async (t) => {
|
||||
const dir = await fs.mkdtemp(path.join(os.tmpdir(), "agnes-ref-"));
|
||||
t.after(() => fs.rm(dir, { recursive: true, force: true }));
|
||||
|
||||
const localPath = path.join(dir, "ref.png");
|
||||
const localBytes = Buffer.from([0x89, 0x50, 0x4e, 0x47]);
|
||||
await fs.writeFile(localPath, localBytes);
|
||||
|
||||
const jpegPath = path.join(dir, "photo.jpeg");
|
||||
await fs.writeFile(jpegPath, Buffer.from([0xff, 0xd8]));
|
||||
|
||||
const results = await resolveReferenceImages([
|
||||
localPath,
|
||||
"https://example.com/remote.jpg",
|
||||
jpegPath,
|
||||
]);
|
||||
|
||||
assert.equal(results.length, 3);
|
||||
assert.match(results[0]!, /^data:image\/png;base64,/);
|
||||
assert.match(results[1]!, /^https:\/\/example.com\/remote.jpg$/);
|
||||
assert.match(results[2]!, /^data:image\/jpeg;base64,/);
|
||||
});
|
||||
|
||||
test("resolveReferenceImages detects gif and webp mime types", async (t) => {
|
||||
const dir = await fs.mkdtemp(path.join(os.tmpdir(), "agnes-mime-"));
|
||||
t.after(() => fs.rm(dir, { recursive: true, force: true }));
|
||||
|
||||
const webpPath = path.join(dir, "ref.webp");
|
||||
const gifPath = path.join(dir, "ref.gif");
|
||||
await fs.writeFile(webpPath, Buffer.from([0x00]));
|
||||
await fs.writeFile(gifPath, Buffer.from([0x00]));
|
||||
|
||||
const results = await resolveReferenceImages([webpPath, gifPath]);
|
||||
assert.match(results[0]!, /^data:image\/webp;base64,/);
|
||||
assert.match(results[1]!, /^data:image\/gif;base64,/);
|
||||
});
|
||||
175
baoyu-skills/skills/baoyu-image-gen/scripts/providers/agnes.ts
Normal file
175
baoyu-skills/skills/baoyu-image-gen/scripts/providers/agnes.ts
Normal file
@@ -0,0 +1,175 @@
|
||||
import { readFile } from "node:fs/promises";
|
||||
import path from "node:path";
|
||||
import type { CliArgs } from "../types";
|
||||
|
||||
const DEFAULT_MODEL = "agnes-image-2.5-flash";
|
||||
const DEFAULT_BASE_URL = "https://apihub.agnes-ai.com/v1";
|
||||
const DEFAULT_SIZE = "1024x1024";
|
||||
|
||||
type AgnesResponse = {
|
||||
created?: number;
|
||||
data: Array<{ url?: string; b64_json?: string }>;
|
||||
};
|
||||
|
||||
export function getDefaultModel(): string {
|
||||
return process.env.AGNES_IMAGE_MODEL || DEFAULT_MODEL;
|
||||
}
|
||||
|
||||
function getApiKey(): string {
|
||||
const key = process.env.AGNES_API_KEY;
|
||||
if (!key) {
|
||||
throw new Error("AGNES_API_KEY is required. Get one from https://apihub.agnes-ai.com.");
|
||||
}
|
||||
return key;
|
||||
}
|
||||
|
||||
function getBaseUrl(): string {
|
||||
return (process.env.AGNES_BASE_URL || DEFAULT_BASE_URL).replace(/\/+$/, "");
|
||||
}
|
||||
|
||||
export function parseAspectRatio(ar: string): { width: number; height: number } | null {
|
||||
const match = ar.match(/^(\d+(?:\.\d+)?):(\d+(?:\.\d+)?)$/);
|
||||
if (!match) return null;
|
||||
const w = parseFloat(match[1]!);
|
||||
const h = parseFloat(match[2]!);
|
||||
if (w <= 0 || h <= 0) return null;
|
||||
return { width: w, height: h };
|
||||
}
|
||||
|
||||
export function snapDim(n: number): number {
|
||||
return Math.max(32, Math.round(n / 32) * 32);
|
||||
}
|
||||
|
||||
export function resolveSize(args: Pick<CliArgs, "size" | "aspectRatio">): string {
|
||||
if (args.size) return args.size;
|
||||
|
||||
if (args.aspectRatio) {
|
||||
const parsed = parseAspectRatio(args.aspectRatio);
|
||||
if (parsed) {
|
||||
if (parsed.width === 1 && parsed.height === 1) return "1024x1024";
|
||||
const maxEdge = 2048;
|
||||
const scale = Math.max(1, Math.floor(maxEdge / Math.max(parsed.width, parsed.height)));
|
||||
const width = parsed.width * scale;
|
||||
const height = parsed.height * scale;
|
||||
return `${snapDim(width)}x${snapDim(height)}`;
|
||||
}
|
||||
}
|
||||
|
||||
return DEFAULT_SIZE;
|
||||
}
|
||||
|
||||
function isRemoteUrl(refPath: string): boolean {
|
||||
return /^https?:\/\//i.test(refPath);
|
||||
}
|
||||
|
||||
export async function resolveReferenceImages(
|
||||
referenceImages: string[]
|
||||
): Promise<string[]> {
|
||||
const result: string[] = [];
|
||||
for (const refPath of referenceImages) {
|
||||
if (isRemoteUrl(refPath)) {
|
||||
result.push(refPath);
|
||||
continue;
|
||||
}
|
||||
const bytes = await readFile(refPath);
|
||||
const ext = path.extname(refPath).toLowerCase();
|
||||
let mime = "image/png";
|
||||
if (ext === ".jpg" || ext === ".jpeg") mime = "image/jpeg";
|
||||
else if (ext === ".webp") mime = "image/webp";
|
||||
else if (ext === ".gif") mime = "image/gif";
|
||||
const b64 = Buffer.from(bytes).toString("base64");
|
||||
result.push(`data:${mime};base64,${b64}`);
|
||||
}
|
||||
return result;
|
||||
}
|
||||
|
||||
export function validateArgs(_model: string, args: CliArgs): void {
|
||||
if (args.n > 1) {
|
||||
throw new Error("Agnes image generation currently returns a single image per request. Set --n 1 or omit --n.");
|
||||
}
|
||||
}
|
||||
|
||||
export function getDefaultOutputExtension(_model: string, args: CliArgs): string {
|
||||
return args.responseFormat === "url" ? ".txt" : ".png";
|
||||
}
|
||||
|
||||
export function buildRequestBody(
|
||||
prompt: string,
|
||||
model: string,
|
||||
args: Pick<CliArgs, "size" | "aspectRatio" | "referenceImages">
|
||||
): Record<string, unknown> {
|
||||
const body: Record<string, unknown> = {
|
||||
model,
|
||||
prompt,
|
||||
size: resolveSize(args),
|
||||
};
|
||||
|
||||
if (args.referenceImages.length > 0) {
|
||||
body.image = args.referenceImages;
|
||||
}
|
||||
|
||||
body.extra_body = { response_format: "url" };
|
||||
|
||||
return body;
|
||||
}
|
||||
|
||||
export async function extractImageFromResponse(result: AgnesResponse): Promise<Uint8Array> {
|
||||
const img = result.data[0];
|
||||
|
||||
if (img?.b64_json) {
|
||||
return Uint8Array.from(Buffer.from(img.b64_json, "base64"));
|
||||
}
|
||||
|
||||
if (img?.url) {
|
||||
const imgRes = await fetch(img.url);
|
||||
if (!imgRes.ok) throw new Error(`Failed to download image from Agnes: ${imgRes.status}`);
|
||||
return new Uint8Array(await imgRes.arrayBuffer());
|
||||
}
|
||||
|
||||
throw new Error("No image in Agnes response");
|
||||
}
|
||||
|
||||
export async function generateImage(
|
||||
prompt: string,
|
||||
model: string,
|
||||
args: CliArgs
|
||||
): Promise<Uint8Array> {
|
||||
const apiKey = getApiKey();
|
||||
const baseUrl = getBaseUrl();
|
||||
|
||||
const referenceImages = await resolveReferenceImages(args.referenceImages);
|
||||
|
||||
const body = buildRequestBody(prompt, model, { ...args, referenceImages });
|
||||
|
||||
const controller = new AbortController();
|
||||
const timeout = setTimeout(() => controller.abort(), 120_000);
|
||||
|
||||
try {
|
||||
const res = await fetch(`${baseUrl}/images/generations`, {
|
||||
method: "POST",
|
||||
headers: {
|
||||
"Content-Type": "application/json",
|
||||
Authorization: `Bearer ${apiKey}`,
|
||||
},
|
||||
body: JSON.stringify(body),
|
||||
signal: controller.signal,
|
||||
});
|
||||
|
||||
if (!res.ok) {
|
||||
const err = await res.text();
|
||||
throw new Error(`Agnes API error (${res.status}): ${err}`);
|
||||
}
|
||||
|
||||
const result = (await res.json()) as AgnesResponse;
|
||||
|
||||
if (args.responseFormat === "url") {
|
||||
const url = result.data[0]?.url;
|
||||
if (!url) throw new Error("No URL in Agnes response");
|
||||
return new Uint8Array(Buffer.from(url, "utf-8"));
|
||||
}
|
||||
|
||||
return extractImageFromResponse(result);
|
||||
} finally {
|
||||
clearTimeout(timeout);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,189 @@
|
||||
import assert from "node:assert/strict";
|
||||
import fs from "node:fs/promises";
|
||||
import os from "node:os";
|
||||
import path from "node:path";
|
||||
import test, { type TestContext } from "node:test";
|
||||
|
||||
import type { CliArgs } from "../types.ts";
|
||||
import {
|
||||
generateImage,
|
||||
getDefaultModel,
|
||||
parseAzureBaseURL,
|
||||
validateArgs,
|
||||
} from "./azure.ts";
|
||||
|
||||
function useEnv(
|
||||
t: TestContext,
|
||||
values: Record<string, string | null>,
|
||||
): void {
|
||||
const previous = new Map<string, string | undefined>();
|
||||
for (const [key, value] of Object.entries(values)) {
|
||||
previous.set(key, process.env[key]);
|
||||
if (value == null) {
|
||||
delete process.env[key];
|
||||
} else {
|
||||
process.env[key] = value;
|
||||
}
|
||||
}
|
||||
|
||||
t.after(() => {
|
||||
for (const [key, value] of previous.entries()) {
|
||||
if (value == null) {
|
||||
delete process.env[key];
|
||||
} else {
|
||||
process.env[key] = value;
|
||||
}
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
function makeArgs(overrides: Partial<CliArgs> = {}): CliArgs {
|
||||
return {
|
||||
prompt: null,
|
||||
promptFiles: [],
|
||||
imagePath: null,
|
||||
provider: null,
|
||||
model: null,
|
||||
aspectRatio: null,
|
||||
size: null,
|
||||
quality: null,
|
||||
imageSize: null,
|
||||
imageApiDialect: null,
|
||||
referenceImages: [],
|
||||
n: 1,
|
||||
batchFile: null,
|
||||
jobs: null,
|
||||
json: false,
|
||||
help: false,
|
||||
...overrides,
|
||||
};
|
||||
}
|
||||
|
||||
async function makeTempDir(prefix: string): Promise<string> {
|
||||
return fs.mkdtemp(path.join(os.tmpdir(), prefix));
|
||||
}
|
||||
|
||||
test("Azure endpoint parsing and default deployment selection follow env precedence", (t) => {
|
||||
assert.deepEqual(parseAzureBaseURL("https://example.openai.azure.com"), {
|
||||
resourceBaseURL: "https://example.openai.azure.com/openai",
|
||||
deployment: null,
|
||||
});
|
||||
assert.deepEqual(
|
||||
parseAzureBaseURL("https://example.openai.azure.com/openai/deployments/from-url"),
|
||||
{
|
||||
resourceBaseURL: "https://example.openai.azure.com/openai",
|
||||
deployment: "from-url",
|
||||
},
|
||||
);
|
||||
|
||||
useEnv(t, {
|
||||
AZURE_OPENAI_BASE_URL: "https://example.openai.azure.com/openai/deployments/from-url",
|
||||
AZURE_OPENAI_DEPLOYMENT: "explicit-deploy",
|
||||
AZURE_OPENAI_IMAGE_MODEL: "env-fallback",
|
||||
});
|
||||
assert.equal(getDefaultModel(), "explicit-deploy");
|
||||
});
|
||||
|
||||
test("Azure validateArgs rejects unsupported edit input formats before the API call", () => {
|
||||
assert.doesNotThrow(() =>
|
||||
validateArgs("demo-deployment", makeArgs({ referenceImages: ["hero.png", "photo.jpeg"] })),
|
||||
);
|
||||
assert.throws(
|
||||
() => validateArgs("demo-deployment", makeArgs({ referenceImages: ["hero.webp"] })),
|
||||
/PNG or JPG\/JPEG/,
|
||||
);
|
||||
});
|
||||
|
||||
test("Azure image generation routes model to deployment and sends mapped quality", async (t) => {
|
||||
useEnv(t, {
|
||||
AZURE_OPENAI_API_KEY: "azure-key",
|
||||
AZURE_OPENAI_BASE_URL: "https://example.openai.azure.com/openai/deployments/default-deploy",
|
||||
AZURE_API_VERSION: null,
|
||||
AZURE_OPENAI_DEPLOYMENT: null,
|
||||
AZURE_OPENAI_IMAGE_MODEL: null,
|
||||
});
|
||||
|
||||
const originalFetch = globalThis.fetch;
|
||||
t.after(() => {
|
||||
globalThis.fetch = originalFetch;
|
||||
});
|
||||
|
||||
const calls: Array<{ url: string; body: string }> = [];
|
||||
globalThis.fetch = async (input, init) => {
|
||||
calls.push({
|
||||
url: String(input),
|
||||
body: String(init?.body ?? ""),
|
||||
});
|
||||
return Response.json({
|
||||
data: [{ b64_json: Buffer.from("azure-image").toString("base64") }],
|
||||
});
|
||||
};
|
||||
|
||||
const bytes = await generateImage(
|
||||
"A calm lake at sunset",
|
||||
"custom-deploy",
|
||||
makeArgs({ quality: "normal" }),
|
||||
);
|
||||
|
||||
assert.equal(Buffer.from(bytes).toString("utf8"), "azure-image");
|
||||
assert.equal(
|
||||
calls[0]?.url,
|
||||
"https://example.openai.azure.com/openai/deployments/custom-deploy/images/generations?api-version=2025-04-01-preview",
|
||||
);
|
||||
|
||||
const body = JSON.parse(calls[0]!.body) as Record<string, string>;
|
||||
assert.equal(body.quality, "medium");
|
||||
assert.equal(body.size, "1024x1024");
|
||||
});
|
||||
|
||||
test("Azure image edits include quality in multipart requests", async (t) => {
|
||||
const root = await makeTempDir("baoyu-image-gen-azure-");
|
||||
t.after(() => fs.rm(root, { recursive: true, force: true }));
|
||||
|
||||
const pngPath = path.join(root, "ref.png");
|
||||
const jpgPath = path.join(root, "ref.jpg");
|
||||
await fs.writeFile(pngPath, "png-bytes");
|
||||
await fs.writeFile(jpgPath, "jpg-bytes");
|
||||
|
||||
useEnv(t, {
|
||||
AZURE_OPENAI_API_KEY: "azure-key",
|
||||
AZURE_OPENAI_BASE_URL: "https://example.openai.azure.com",
|
||||
AZURE_API_VERSION: "2025-04-01-preview",
|
||||
AZURE_OPENAI_DEPLOYMENT: null,
|
||||
AZURE_OPENAI_IMAGE_MODEL: null,
|
||||
});
|
||||
|
||||
const originalFetch = globalThis.fetch;
|
||||
t.after(() => {
|
||||
globalThis.fetch = originalFetch;
|
||||
});
|
||||
|
||||
const calls: Array<{ url: string; form: FormData }> = [];
|
||||
globalThis.fetch = async (input, init) => {
|
||||
calls.push({
|
||||
url: String(input),
|
||||
form: init?.body as FormData,
|
||||
});
|
||||
return Response.json({
|
||||
data: [{ b64_json: Buffer.from("edited-image").toString("base64") }],
|
||||
});
|
||||
};
|
||||
|
||||
const bytes = await generateImage(
|
||||
"Add warm lighting",
|
||||
"edit-deploy",
|
||||
makeArgs({
|
||||
quality: "2k",
|
||||
referenceImages: [pngPath, jpgPath],
|
||||
}),
|
||||
);
|
||||
|
||||
assert.equal(Buffer.from(bytes).toString("utf8"), "edited-image");
|
||||
assert.equal(
|
||||
calls[0]?.url,
|
||||
"https://example.openai.azure.com/openai/deployments/edit-deploy/images/edits?api-version=2025-04-01-preview",
|
||||
);
|
||||
assert.equal(calls[0]?.form.get("quality"), "high");
|
||||
assert.equal(calls[0]?.form.get("size"), "1024x1024");
|
||||
assert.equal(calls[0]?.form.getAll("image[]").length, 2);
|
||||
});
|
||||
192
baoyu-skills/skills/baoyu-image-gen/scripts/providers/azure.ts
Normal file
192
baoyu-skills/skills/baoyu-image-gen/scripts/providers/azure.ts
Normal file
@@ -0,0 +1,192 @@
|
||||
import path from "node:path";
|
||||
import { readFile } from "node:fs/promises";
|
||||
import type { CliArgs } from "../types";
|
||||
import { getOpenAISize, extractImageFromResponse } from "./openai.ts";
|
||||
|
||||
type OpenAIImageResponse = { data: Array<{ url?: string; b64_json?: string }> };
|
||||
type AzureEndpoint = {
|
||||
resourceBaseURL: string;
|
||||
deployment: string | null;
|
||||
};
|
||||
|
||||
const DEFAULT_AZURE_API_VERSION = "2025-04-01-preview";
|
||||
const AZURE_EDIT_IMAGE_EXTENSIONS = new Set([".png", ".jpg", ".jpeg"]);
|
||||
|
||||
export function parseAzureBaseURL(url: string): AzureEndpoint {
|
||||
const parsed = new URL(url);
|
||||
const trimmedPath = parsed.pathname.replace(/\/+$/, "");
|
||||
const deploymentMatch = trimmedPath.match(/^(.*?)(?:\/openai)?\/deployments\/([^/]+)$/);
|
||||
|
||||
if (deploymentMatch) {
|
||||
parsed.pathname = `${deploymentMatch[1] || ""}/openai`;
|
||||
return {
|
||||
resourceBaseURL: parsed.toString().replace(/\/+$/, ""),
|
||||
deployment: decodeURIComponent(deploymentMatch[2]!),
|
||||
};
|
||||
}
|
||||
|
||||
parsed.pathname = trimmedPath.endsWith("/openai") ? trimmedPath : `${trimmedPath}/openai`;
|
||||
return {
|
||||
resourceBaseURL: parsed.toString().replace(/\/+$/, ""),
|
||||
deployment: null,
|
||||
};
|
||||
}
|
||||
|
||||
export function getDefaultModel(): string {
|
||||
const explicitDeployment = process.env.AZURE_OPENAI_DEPLOYMENT?.trim();
|
||||
if (explicitDeployment) return explicitDeployment;
|
||||
|
||||
const baseURL = process.env.AZURE_OPENAI_BASE_URL;
|
||||
if (baseURL) {
|
||||
try {
|
||||
const { deployment } = parseAzureBaseURL(baseURL);
|
||||
if (deployment) return deployment;
|
||||
} catch {
|
||||
// Ignore invalid URLs here so the required-env check can raise the user-facing error later.
|
||||
}
|
||||
}
|
||||
|
||||
return process.env.AZURE_OPENAI_IMAGE_MODEL || "gpt-image-2.5-flare";
|
||||
}
|
||||
|
||||
function getEndpoint(): AzureEndpoint {
|
||||
const url = process.env.AZURE_OPENAI_BASE_URL;
|
||||
if (!url) {
|
||||
throw new Error(
|
||||
"AZURE_OPENAI_BASE_URL is required. Set it to your Azure resource or deployment endpoint, e.g.: https://your-resource.openai.azure.com or https://your-resource.openai.azure.com/openai/deployments/your-deployment"
|
||||
);
|
||||
}
|
||||
return parseAzureBaseURL(url);
|
||||
}
|
||||
|
||||
function getApiKey(): string {
|
||||
const key = process.env.AZURE_OPENAI_API_KEY;
|
||||
if (!key) {
|
||||
throw new Error(
|
||||
"AZURE_OPENAI_API_KEY is required. Get it from Azure Portal → your OpenAI resource → Keys and Endpoint."
|
||||
);
|
||||
}
|
||||
return key;
|
||||
}
|
||||
|
||||
function getApiVersion(): string {
|
||||
return process.env.AZURE_API_VERSION || DEFAULT_AZURE_API_VERSION;
|
||||
}
|
||||
|
||||
function getDeployment(model: string): string {
|
||||
const deployment = model.trim();
|
||||
if (!deployment) {
|
||||
throw new Error(
|
||||
"Azure deployment name is required. Use --model <deployment>, AZURE_OPENAI_DEPLOYMENT, AZURE_OPENAI_IMAGE_MODEL, or embed the deployment in AZURE_OPENAI_BASE_URL."
|
||||
);
|
||||
}
|
||||
return deployment;
|
||||
}
|
||||
|
||||
function buildURL(deployment: string, pathSuffix: string): string {
|
||||
const { resourceBaseURL } = getEndpoint();
|
||||
return `${resourceBaseURL}/deployments/${encodeURIComponent(deployment)}${pathSuffix}?api-version=${getApiVersion()}`;
|
||||
}
|
||||
|
||||
function authHeaders(): Record<string, string> {
|
||||
return { "api-key": getApiKey() };
|
||||
}
|
||||
|
||||
function getAzureQuality(quality: CliArgs["quality"]): "medium" | "high" {
|
||||
return quality === "2k" ? "high" : "medium";
|
||||
}
|
||||
|
||||
export function validateArgs(_model: string, args: CliArgs): void {
|
||||
for (const refPath of args.referenceImages) {
|
||||
const ext = path.extname(refPath).toLowerCase();
|
||||
if (!AZURE_EDIT_IMAGE_EXTENSIONS.has(ext)) {
|
||||
throw new Error(
|
||||
`Azure OpenAI reference images must be PNG or JPG/JPEG. Unsupported file: ${refPath}`
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
export async function generateImage(
|
||||
prompt: string,
|
||||
model: string,
|
||||
args: CliArgs
|
||||
): Promise<Uint8Array> {
|
||||
const deployment = getDeployment(model);
|
||||
const size = args.size || getOpenAISize(model, args.aspectRatio, args.quality);
|
||||
|
||||
if (args.referenceImages.length > 0) {
|
||||
return generateWithAzureEdits(prompt, deployment, size, args.referenceImages, args.quality);
|
||||
}
|
||||
|
||||
return generateWithAzureGenerations(prompt, deployment, size, args.quality);
|
||||
}
|
||||
|
||||
async function generateWithAzureGenerations(
|
||||
prompt: string,
|
||||
deployment: string,
|
||||
size: string,
|
||||
quality: CliArgs["quality"]
|
||||
): Promise<Uint8Array> {
|
||||
const body: Record<string, any> = {
|
||||
prompt,
|
||||
size,
|
||||
n: 1,
|
||||
quality: getAzureQuality(quality),
|
||||
};
|
||||
|
||||
const res = await fetch(buildURL(deployment, "/images/generations"), {
|
||||
method: "POST",
|
||||
headers: {
|
||||
"Content-Type": "application/json",
|
||||
...authHeaders(),
|
||||
},
|
||||
body: JSON.stringify(body),
|
||||
});
|
||||
|
||||
if (!res.ok) {
|
||||
const err = await res.text();
|
||||
throw new Error(`Azure OpenAI API error: ${err}`);
|
||||
}
|
||||
|
||||
const result = (await res.json()) as OpenAIImageResponse;
|
||||
return extractImageFromResponse(result);
|
||||
}
|
||||
|
||||
async function generateWithAzureEdits(
|
||||
prompt: string,
|
||||
deployment: string,
|
||||
size: string,
|
||||
referenceImages: string[],
|
||||
quality: CliArgs["quality"]
|
||||
): Promise<Uint8Array> {
|
||||
const form = new FormData();
|
||||
form.append("prompt", prompt);
|
||||
form.append("size", size);
|
||||
form.append("n", "1");
|
||||
form.append("quality", getAzureQuality(quality));
|
||||
|
||||
for (const refPath of referenceImages) {
|
||||
const bytes = await readFile(refPath);
|
||||
const filename = path.basename(refPath);
|
||||
const mimeType = path.extname(filename).toLowerCase() === ".png" ? "image/png" : "image/jpeg";
|
||||
const blob = new Blob([bytes], { type: mimeType });
|
||||
form.append("image[]", blob, filename);
|
||||
}
|
||||
|
||||
const res = await fetch(buildURL(deployment, "/images/edits"), {
|
||||
method: "POST",
|
||||
headers: {
|
||||
...authHeaders(),
|
||||
},
|
||||
body: form,
|
||||
});
|
||||
|
||||
if (!res.ok) {
|
||||
const err = await res.text();
|
||||
throw new Error(`Azure OpenAI edits API error: ${err}`);
|
||||
}
|
||||
|
||||
const result = (await res.json()) as OpenAIImageResponse;
|
||||
return extractImageFromResponse(result);
|
||||
}
|
||||
@@ -0,0 +1,62 @@
|
||||
import assert from "node:assert/strict";
|
||||
import test from "node:test";
|
||||
|
||||
import type { CliArgs } from "../types.ts";
|
||||
import {
|
||||
getDefaultModel,
|
||||
getDefaultOutputExtension,
|
||||
validateArgs,
|
||||
} from "./codex-cli.ts";
|
||||
|
||||
function makeArgs(overrides: Partial<CliArgs> = {}): CliArgs {
|
||||
return {
|
||||
prompt: null,
|
||||
promptFiles: [],
|
||||
imagePath: null,
|
||||
provider: "codex-cli",
|
||||
model: null,
|
||||
aspectRatio: null,
|
||||
aspectRatioSource: null,
|
||||
size: null,
|
||||
quality: "2k",
|
||||
imageSize: null,
|
||||
imageSizeSource: null,
|
||||
imageApiDialect: null,
|
||||
referenceImages: [],
|
||||
n: 1,
|
||||
batchFile: null,
|
||||
jobs: null,
|
||||
json: false,
|
||||
help: false,
|
||||
...overrides,
|
||||
};
|
||||
}
|
||||
|
||||
test("codex-cli defaults to codex-image-gen model and PNG output", () => {
|
||||
assert.equal(getDefaultModel(), "codex-image-gen");
|
||||
assert.equal(getDefaultOutputExtension(), ".png");
|
||||
});
|
||||
|
||||
test("codex-cli validateArgs rejects n>1 with a non-retryable message", () => {
|
||||
assert.throws(
|
||||
() => validateArgs("codex-image-gen", makeArgs({ n: 2 })),
|
||||
/supports only n=1/,
|
||||
);
|
||||
});
|
||||
|
||||
test("codex-cli validateArgs rejects ratio-metadata dialect", () => {
|
||||
assert.throws(
|
||||
() => validateArgs("codex-image-gen", makeArgs({ imageApiDialect: "ratio-metadata" })),
|
||||
/Invalid imageApiDialect/,
|
||||
);
|
||||
});
|
||||
|
||||
test("codex-cli validateArgs accepts default n=1 with no dialect", () => {
|
||||
assert.doesNotThrow(() => validateArgs("codex-image-gen", makeArgs()));
|
||||
});
|
||||
|
||||
test("codex-cli validateArgs accepts reference images (Codex image_gen supports refs)", () => {
|
||||
assert.doesNotThrow(() =>
|
||||
validateArgs("codex-image-gen", makeArgs({ referenceImages: ["/tmp/a.png", "/tmp/b.png"] })),
|
||||
);
|
||||
});
|
||||
@@ -0,0 +1,197 @@
|
||||
import path from "node:path";
|
||||
import { spawn } from "node:child_process";
|
||||
import { fileURLToPath } from "node:url";
|
||||
import { tmpdir } from "node:os";
|
||||
import { mkdir, readFile, rm, writeFile, access } from "node:fs/promises";
|
||||
import { randomBytes } from "node:crypto";
|
||||
import type { CliArgs } from "../types";
|
||||
|
||||
const PROVIDER_FILE = fileURLToPath(import.meta.url);
|
||||
const SCRIPTS_DIR = path.resolve(path.dirname(PROVIDER_FILE), "..");
|
||||
const BUNDLED_WRAPPER = path.join(SCRIPTS_DIR, "codex-imagegen", "main.ts");
|
||||
|
||||
type WrapperOkResult = {
|
||||
status: "ok";
|
||||
path: string;
|
||||
bytes: number;
|
||||
elapsed_seconds: number;
|
||||
thread_id: string | null;
|
||||
attempts: number;
|
||||
cached: boolean;
|
||||
};
|
||||
|
||||
type WrapperErrorResult = {
|
||||
status: "error";
|
||||
path: string;
|
||||
bytes: number;
|
||||
error: string;
|
||||
error_kind: string;
|
||||
};
|
||||
|
||||
type WrapperResult = WrapperOkResult | WrapperErrorResult;
|
||||
|
||||
export function getDefaultModel(): string {
|
||||
return "codex-image-gen";
|
||||
}
|
||||
|
||||
export function getDefaultOutputExtension(): string {
|
||||
return ".png";
|
||||
}
|
||||
|
||||
export function validateArgs(_model: string, args: CliArgs): void {
|
||||
if (args.n > 1) {
|
||||
throw new Error(
|
||||
"codex-cli provider supports only n=1 (Codex image_gen returns a single image per call).",
|
||||
);
|
||||
}
|
||||
if (args.imageApiDialect && args.imageApiDialect !== "openai-native") {
|
||||
throw new Error(
|
||||
`Invalid imageApiDialect for codex-cli: ${args.imageApiDialect}. codex-cli does not use OpenAI Images API dialects.`,
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
async function exists(filePath: string): Promise<boolean> {
|
||||
try {
|
||||
await access(filePath);
|
||||
return true;
|
||||
} catch {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
async function resolveWrapperPath(): Promise<string> {
|
||||
const override = process.env.BAOYU_CODEX_IMAGEGEN_BIN;
|
||||
if (override) {
|
||||
if (!(await exists(override))) {
|
||||
throw new Error(
|
||||
`Invalid BAOYU_CODEX_IMAGEGEN_BIN: ${override} does not exist.`,
|
||||
);
|
||||
}
|
||||
return override;
|
||||
}
|
||||
if (await exists(BUNDLED_WRAPPER)) return BUNDLED_WRAPPER;
|
||||
throw new Error(
|
||||
`codex-cli wrapper not found at ${BUNDLED_WRAPPER}. ` +
|
||||
`Reinstall baoyu-image-gen, or set BAOYU_CODEX_IMAGEGEN_BIN to a codex-imagegen main.ts (or .sh) path.`,
|
||||
);
|
||||
}
|
||||
|
||||
type SpawnResult = {
|
||||
stdout: string;
|
||||
stderr: string;
|
||||
code: number;
|
||||
};
|
||||
|
||||
async function spawnWrapper(wrapperPath: string, cliArgs: string[]): Promise<SpawnResult> {
|
||||
return new Promise((resolve, reject) => {
|
||||
const isTs = wrapperPath.endsWith(".ts");
|
||||
const command = isTs ? "bun" : wrapperPath;
|
||||
const args = isTs ? [wrapperPath, ...cliArgs] : cliArgs;
|
||||
const child = spawn(command, args, { stdio: ["ignore", "pipe", "pipe"] });
|
||||
let stdout = "";
|
||||
let stderr = "";
|
||||
child.stdout.on("data", (chunk: Buffer) => {
|
||||
stdout += chunk.toString("utf8");
|
||||
});
|
||||
child.stderr.on("data", (chunk: Buffer) => {
|
||||
const text = chunk.toString("utf8");
|
||||
stderr += text;
|
||||
process.stderr.write(text);
|
||||
});
|
||||
child.on("error", (err) => reject(err));
|
||||
child.on("close", (code) => resolve({ stdout, stderr, code: code ?? 1 }));
|
||||
});
|
||||
}
|
||||
|
||||
function parseWrapperJson(stdout: string): WrapperResult {
|
||||
const trimmed = stdout.trim();
|
||||
if (!trimmed) {
|
||||
throw new Error("Invalid codex-cli response: empty stdout from wrapper.");
|
||||
}
|
||||
const lastLine = trimmed.split(/\r?\n/).pop() ?? trimmed;
|
||||
try {
|
||||
return JSON.parse(lastLine) as WrapperResult;
|
||||
} catch (parseErr) {
|
||||
throw new Error(
|
||||
`Invalid codex-cli response: could not parse JSON from wrapper stdout (${(parseErr as Error).message}).`,
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
function parsePositiveInt(value: string | undefined): number | null {
|
||||
if (!value) return null;
|
||||
const parsed = parseInt(value, 10);
|
||||
return Number.isFinite(parsed) && parsed > 0 ? parsed : null;
|
||||
}
|
||||
|
||||
function getEnvOverride(name: string): string | null {
|
||||
const value = process.env[name];
|
||||
return value && value.length > 0 ? value : null;
|
||||
}
|
||||
|
||||
export async function generateImage(
|
||||
prompt: string,
|
||||
_model: string,
|
||||
args: CliArgs,
|
||||
): Promise<Uint8Array> {
|
||||
const wrapperPath = await resolveWrapperPath();
|
||||
|
||||
const sessionDir = path.join(tmpdir(), "baoyu-image-gen-codex-cli");
|
||||
await mkdir(sessionDir, { recursive: true });
|
||||
const token = randomBytes(8).toString("hex");
|
||||
const tmpOutput = path.join(sessionDir, `out-${token}.png`);
|
||||
const tmpPrompt = path.join(sessionDir, `prompt-${token}.md`);
|
||||
await writeFile(tmpPrompt, prompt, "utf8");
|
||||
|
||||
const aspect = args.aspectRatio ?? "1:1";
|
||||
const cliArgs: string[] = [
|
||||
"--image",
|
||||
tmpOutput,
|
||||
"--prompt-file",
|
||||
tmpPrompt,
|
||||
"--aspect",
|
||||
aspect,
|
||||
];
|
||||
|
||||
for (const ref of args.referenceImages) {
|
||||
cliArgs.push("--ref", path.resolve(ref));
|
||||
}
|
||||
|
||||
const cacheDir = getEnvOverride("BAOYU_CODEX_IMAGEGEN_CACHE_DIR");
|
||||
if (cacheDir) cliArgs.push("--cache-dir", cacheDir);
|
||||
|
||||
const timeoutMs = parsePositiveInt(process.env.BAOYU_CODEX_IMAGEGEN_TIMEOUT_MS);
|
||||
if (timeoutMs) cliArgs.push("--timeout", String(timeoutMs));
|
||||
|
||||
const retries = parsePositiveInt(process.env.BAOYU_CODEX_IMAGEGEN_RETRIES);
|
||||
if (retries !== null) cliArgs.push("--retries", String(retries));
|
||||
|
||||
const logFile = getEnvOverride("BAOYU_CODEX_IMAGEGEN_LOG_FILE");
|
||||
if (logFile) cliArgs.push("--log-file", logFile);
|
||||
|
||||
try {
|
||||
const spawnResult = await spawnWrapper(wrapperPath, cliArgs);
|
||||
const parsed = parseWrapperJson(spawnResult.stdout);
|
||||
|
||||
if (parsed.status === "error") {
|
||||
throw new Error(
|
||||
`Invalid codex-cli result (${parsed.error_kind}): ${parsed.error}`,
|
||||
);
|
||||
}
|
||||
|
||||
if (spawnResult.code !== 0) {
|
||||
throw new Error(
|
||||
`Invalid codex-cli result: wrapper exited with code ${spawnResult.code} despite reporting status=ok.`,
|
||||
);
|
||||
}
|
||||
|
||||
const bytes = await readFile(parsed.path ?? tmpOutput);
|
||||
return new Uint8Array(bytes);
|
||||
} finally {
|
||||
await Promise.allSettled([
|
||||
rm(tmpOutput, { force: true }),
|
||||
rm(tmpPrompt, { force: true }),
|
||||
]);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,394 @@
|
||||
import assert from "node:assert/strict";
|
||||
import test, { type TestContext } from "node:test";
|
||||
|
||||
import {
|
||||
generateImage,
|
||||
getDefaultModel,
|
||||
getModelFamily,
|
||||
getQwen2SizeFromAspectRatio,
|
||||
getSizeFromAspectRatio,
|
||||
getWan27SizeFromAspectRatio,
|
||||
normalizeSize,
|
||||
parseAspectRatio,
|
||||
parseSize,
|
||||
resolveSizeForModel,
|
||||
} from "./dashscope.ts";
|
||||
import type { CliArgs } from "../types.ts";
|
||||
|
||||
function makeCliArgs(overrides: Partial<CliArgs> = {}): CliArgs {
|
||||
return {
|
||||
prompt: null,
|
||||
promptFiles: [],
|
||||
imagePath: null,
|
||||
provider: "dashscope",
|
||||
model: null,
|
||||
aspectRatio: null,
|
||||
aspectRatioSource: null,
|
||||
size: null,
|
||||
quality: "2k",
|
||||
imageSize: null,
|
||||
imageSizeSource: null,
|
||||
imageApiDialect: null,
|
||||
referenceImages: [],
|
||||
n: 1,
|
||||
batchFile: null,
|
||||
jobs: null,
|
||||
json: false,
|
||||
help: false,
|
||||
...overrides,
|
||||
};
|
||||
}
|
||||
|
||||
function useEnv(
|
||||
t: TestContext,
|
||||
values: Record<string, string | null>,
|
||||
): void {
|
||||
const previous = new Map<string, string | undefined>();
|
||||
for (const [key, value] of Object.entries(values)) {
|
||||
previous.set(key, process.env[key]);
|
||||
if (value == null) {
|
||||
delete process.env[key];
|
||||
} else {
|
||||
process.env[key] = value;
|
||||
}
|
||||
}
|
||||
|
||||
t.after(() => {
|
||||
for (const [key, value] of previous.entries()) {
|
||||
if (value == null) {
|
||||
delete process.env[key];
|
||||
} else {
|
||||
process.env[key] = value;
|
||||
}
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
test("DashScope default model prefers env override and otherwise uses qwen-image-2.0-pro", (t) => {
|
||||
useEnv(t, { DASHSCOPE_IMAGE_MODEL: null });
|
||||
assert.equal(getDefaultModel(), "qwen-image-2.0-pro");
|
||||
|
||||
process.env.DASHSCOPE_IMAGE_MODEL = "qwen-image-max";
|
||||
assert.equal(getDefaultModel(), "qwen-image-max");
|
||||
});
|
||||
|
||||
test("DashScope aspect-ratio parsing accepts numeric ratios only", () => {
|
||||
assert.deepEqual(parseAspectRatio("3:2"), { width: 3, height: 2 });
|
||||
assert.equal(parseAspectRatio("square"), null);
|
||||
assert.equal(parseAspectRatio("-1:2"), null);
|
||||
});
|
||||
|
||||
test("DashScope model family routing distinguishes qwen-2.0, fixed-size qwen, wan2.7, and legacy models", () => {
|
||||
assert.equal(getModelFamily("qwen-image-3.0-pro"), "qwen2");
|
||||
assert.equal(getModelFamily("qwen-image-2.0-pro"), "qwen2");
|
||||
assert.equal(getModelFamily("qwen-image-2.0-pro-2026-04-22"), "qwen2");
|
||||
assert.equal(getModelFamily("qwen-image"), "qwenFixed");
|
||||
assert.equal(getModelFamily("wan2.7-image"), "wan27");
|
||||
assert.equal(getModelFamily("wan2.7-image-pro"), "wan27");
|
||||
assert.equal(getModelFamily("z-image-turbo"), "legacy");
|
||||
assert.equal(getModelFamily("wanx-v1"), "legacy");
|
||||
});
|
||||
|
||||
test("Legacy DashScope size selection keeps the previous quality-based heuristic", () => {
|
||||
assert.equal(getSizeFromAspectRatio(null, "normal"), "1024*1024");
|
||||
assert.equal(getSizeFromAspectRatio("16:9", "normal"), "1280*720");
|
||||
assert.equal(getSizeFromAspectRatio("16:9", "2k"), "2048*1152");
|
||||
assert.equal(getSizeFromAspectRatio("invalid", "2k"), "1536*1536");
|
||||
});
|
||||
|
||||
test("Qwen 2.0 recommended sizes follow the official common-ratio table", () => {
|
||||
assert.equal(getQwen2SizeFromAspectRatio(null, "normal"), "1024*1024");
|
||||
assert.equal(getQwen2SizeFromAspectRatio(null, "2k"), "1536*1536");
|
||||
assert.equal(getQwen2SizeFromAspectRatio("16:9", "normal"), "1280*720");
|
||||
assert.equal(getQwen2SizeFromAspectRatio("21:9", "2k"), "2048*872");
|
||||
});
|
||||
|
||||
test("Qwen 2.0 derives free-form sizes within pixel budget for uncommon ratios", () => {
|
||||
const size = getQwen2SizeFromAspectRatio("5:2", "normal");
|
||||
const parsed = parseSize(size);
|
||||
assert.ok(parsed);
|
||||
assert.ok(parsed.width * parsed.height >= 512 * 512);
|
||||
assert.ok(parsed.width * parsed.height <= 2048 * 2048);
|
||||
assert.ok(Math.abs(parsed.width / parsed.height - 2.5) < 0.08);
|
||||
});
|
||||
|
||||
test("resolveSizeForModel validates explicit qwen-image-2.0 sizes by total pixels", () => {
|
||||
assert.equal(
|
||||
resolveSizeForModel("qwen-image-2.0-pro", {
|
||||
size: "2048x872",
|
||||
aspectRatio: null,
|
||||
quality: "2k",
|
||||
}),
|
||||
"2048*872",
|
||||
);
|
||||
|
||||
assert.throws(
|
||||
() =>
|
||||
resolveSizeForModel("qwen-image-2.0-pro", {
|
||||
size: "4096x4096",
|
||||
aspectRatio: null,
|
||||
quality: "2k",
|
||||
}),
|
||||
/total pixels between/,
|
||||
);
|
||||
});
|
||||
|
||||
test("resolveSizeForModel enforces fixed sizes for qwen-image-max/plus/image", () => {
|
||||
assert.equal(
|
||||
resolveSizeForModel("qwen-image-max", {
|
||||
size: null,
|
||||
aspectRatio: "1:1",
|
||||
quality: "2k",
|
||||
}),
|
||||
"1328*1328",
|
||||
);
|
||||
|
||||
assert.equal(
|
||||
resolveSizeForModel("qwen-image", {
|
||||
size: "1664x928",
|
||||
aspectRatio: "9:16",
|
||||
quality: "normal",
|
||||
}),
|
||||
"1664*928",
|
||||
);
|
||||
|
||||
assert.throws(
|
||||
() =>
|
||||
resolveSizeForModel("qwen-image-max", {
|
||||
size: null,
|
||||
aspectRatio: "21:9",
|
||||
quality: "2k",
|
||||
}),
|
||||
/supports only fixed ratios/,
|
||||
);
|
||||
|
||||
assert.throws(
|
||||
() =>
|
||||
resolveSizeForModel("qwen-image-plus", {
|
||||
size: "1024x1024",
|
||||
aspectRatio: null,
|
||||
quality: "2k",
|
||||
}),
|
||||
/support only these sizes/,
|
||||
);
|
||||
});
|
||||
|
||||
test("DashScope size normalization converts WxH into provider format", () => {
|
||||
assert.equal(normalizeSize("1024x1024"), "1024*1024");
|
||||
assert.equal(normalizeSize("2048*1152"), "2048*1152");
|
||||
});
|
||||
|
||||
test("Wan 2.7 derives sizes that match the requested ratio at the chosen pixel budget", () => {
|
||||
const square2k = getWan27SizeFromAspectRatio(null, "2k", 2048 * 2048);
|
||||
const parsedSquare = parseSize(square2k);
|
||||
assert.ok(parsedSquare);
|
||||
assert.equal(parsedSquare.width, parsedSquare.height);
|
||||
assert.ok(parsedSquare.width * parsedSquare.height <= 2048 * 2048);
|
||||
|
||||
const widescreen = getWan27SizeFromAspectRatio("16:9", "2k", 2048 * 2048);
|
||||
const parsedWide = parseSize(widescreen);
|
||||
assert.ok(parsedWide);
|
||||
assert.ok(Math.abs(parsedWide.width / parsedWide.height - 16 / 9) < 0.05);
|
||||
assert.ok(parsedWide.width * parsedWide.height <= 2048 * 2048);
|
||||
|
||||
const pro4k = getWan27SizeFromAspectRatio("16:9", "2k", 4096 * 4096);
|
||||
const parsed4k = parseSize(pro4k);
|
||||
assert.ok(parsed4k);
|
||||
assert.ok(parsed4k.width * parsed4k.height > 2048 * 2048);
|
||||
assert.ok(parsed4k.width * parsed4k.height <= 4096 * 4096);
|
||||
});
|
||||
|
||||
test("Wan 2.7 rejects aspect ratios outside the [1:8, 8:1] range", () => {
|
||||
assert.throws(
|
||||
() => getWan27SizeFromAspectRatio("9:1", "2k", 2048 * 2048),
|
||||
/1:8, 8:1/,
|
||||
);
|
||||
assert.throws(
|
||||
() => getWan27SizeFromAspectRatio("1:9", "normal", 2048 * 2048),
|
||||
/1:8, 8:1/,
|
||||
);
|
||||
});
|
||||
|
||||
test("Wan 2.7 derived sizes stay inside the boundary ratio limits after rounding", () => {
|
||||
for (const ar of ["8:1", "1:8"]) {
|
||||
const size = getWan27SizeFromAspectRatio(ar, "2k", 2048 * 2048);
|
||||
const parsed = parseSize(size);
|
||||
assert.ok(parsed);
|
||||
const ratio = parsed.width / parsed.height;
|
||||
assert.ok(ratio >= 1 / 8);
|
||||
assert.ok(ratio <= 8);
|
||||
assert.ok(parsed.width * parsed.height <= 2048 * 2048);
|
||||
}
|
||||
});
|
||||
|
||||
test("resolveSizeForModel routes wan2.7-image to the 2K-capped derivation", () => {
|
||||
const size = resolveSizeForModel("wan2.7-image", {
|
||||
size: null,
|
||||
aspectRatio: "16:9",
|
||||
quality: "2k",
|
||||
});
|
||||
const parsed = parseSize(size);
|
||||
assert.ok(parsed);
|
||||
assert.ok(parsed.width * parsed.height <= 2048 * 2048);
|
||||
assert.ok(Math.abs(parsed.width / parsed.height - 16 / 9) < 0.05);
|
||||
});
|
||||
|
||||
test("resolveSizeForModel allows wan2.7-image-pro 4K only when there are no reference images", () => {
|
||||
assert.equal(
|
||||
resolveSizeForModel("wan2.7-image-pro", {
|
||||
size: "4096*4096",
|
||||
aspectRatio: null,
|
||||
quality: "2k",
|
||||
}),
|
||||
"4096*4096",
|
||||
);
|
||||
|
||||
assert.throws(
|
||||
() =>
|
||||
resolveSizeForModel("wan2.7-image-pro", {
|
||||
size: "4096*4096",
|
||||
aspectRatio: null,
|
||||
quality: "2k",
|
||||
referenceImages: ["a.png"],
|
||||
}),
|
||||
/total pixels between 768\*768 and 2048\*2048/,
|
||||
);
|
||||
|
||||
const proWithRef = resolveSizeForModel("wan2.7-image-pro", {
|
||||
size: null,
|
||||
aspectRatio: "1:1",
|
||||
quality: "2k",
|
||||
referenceImages: ["a.png"],
|
||||
});
|
||||
const parsedRef = parseSize(proWithRef);
|
||||
assert.ok(parsedRef);
|
||||
assert.ok(parsedRef.width * parsedRef.height <= 2048 * 2048);
|
||||
});
|
||||
|
||||
test("Wan 2.7 request body forces n=1 and omits prompt_extend / negative_prompt", async (t) => {
|
||||
useEnv(t, { DASHSCOPE_API_KEY: "fake-key" });
|
||||
|
||||
const originalFetch = globalThis.fetch;
|
||||
let capturedBody: any = null;
|
||||
globalThis.fetch = (async (_url: string, init?: RequestInit) => {
|
||||
capturedBody = JSON.parse(String(init?.body));
|
||||
return new Response(
|
||||
JSON.stringify({
|
||||
output: {
|
||||
choices: [
|
||||
{
|
||||
message: {
|
||||
content: [{ image: "data:image/png;base64,iVBORw0KGgo=" }],
|
||||
},
|
||||
},
|
||||
],
|
||||
},
|
||||
}),
|
||||
{ status: 200, headers: { "content-type": "application/json" } },
|
||||
);
|
||||
}) as typeof fetch;
|
||||
t.after(() => {
|
||||
globalThis.fetch = originalFetch;
|
||||
});
|
||||
|
||||
await generateImage("hello", "wan2.7-image-pro", makeCliArgs({ aspectRatio: "1:1" }));
|
||||
|
||||
assert.equal(capturedBody.model, "wan2.7-image-pro");
|
||||
assert.deepEqual(Object.keys(capturedBody.parameters).sort(), ["n", "size", "watermark"]);
|
||||
assert.equal(capturedBody.parameters.n, 1);
|
||||
assert.equal(capturedBody.parameters.watermark, false);
|
||||
assert.equal(typeof capturedBody.parameters.size, "string");
|
||||
assert.ok(!("prompt_extend" in capturedBody.parameters));
|
||||
assert.ok(!("negative_prompt" in capturedBody.parameters));
|
||||
|
||||
assert.deepEqual(capturedBody.input.messages[0].content, [{ text: "hello" }]);
|
||||
});
|
||||
|
||||
test("Wan 2.7 request body forwards remote reference image URLs", async (t) => {
|
||||
useEnv(t, { DASHSCOPE_API_KEY: "fake-key" });
|
||||
|
||||
const originalFetch = globalThis.fetch;
|
||||
let capturedBody: any = null;
|
||||
globalThis.fetch = (async (_url: string, init?: RequestInit) => {
|
||||
capturedBody = JSON.parse(String(init?.body));
|
||||
return new Response(
|
||||
JSON.stringify({
|
||||
output: {
|
||||
choices: [
|
||||
{
|
||||
message: {
|
||||
content: [{ image: "data:image/png;base64,iVBORw0KGgo=" }],
|
||||
},
|
||||
},
|
||||
],
|
||||
},
|
||||
}),
|
||||
{ status: 200, headers: { "content-type": "application/json" } },
|
||||
);
|
||||
}) as typeof fetch;
|
||||
t.after(() => {
|
||||
globalThis.fetch = originalFetch;
|
||||
});
|
||||
|
||||
await generateImage(
|
||||
"combine these",
|
||||
"wan2.7-image-pro",
|
||||
makeCliArgs({ referenceImages: ["https://example.com/ref.png"] }),
|
||||
);
|
||||
|
||||
assert.deepEqual(capturedBody.input.messages[0].content, [
|
||||
{ image: "https://example.com/ref.png" },
|
||||
{ text: "combine these" },
|
||||
]);
|
||||
});
|
||||
|
||||
test("Wan 2.7 rejects --n > 1 to prevent silent multi-image billing", async (t) => {
|
||||
useEnv(t, { DASHSCOPE_API_KEY: "fake-key" });
|
||||
|
||||
await assert.rejects(
|
||||
() => generateImage("hi", "wan2.7-image-pro", makeCliArgs({ n: 2 })),
|
||||
/support exactly one output image/,
|
||||
);
|
||||
});
|
||||
|
||||
test("resolveSizeForModel validates explicit wan2.7 sizes by pixel budget and ratio", () => {
|
||||
assert.equal(
|
||||
resolveSizeForModel("wan2.7-image-pro", {
|
||||
size: "3840x2160",
|
||||
aspectRatio: null,
|
||||
quality: "2k",
|
||||
}),
|
||||
"3840*2160",
|
||||
);
|
||||
|
||||
assert.throws(
|
||||
() =>
|
||||
resolveSizeForModel("wan2.7-image-pro", {
|
||||
size: "3840x2160",
|
||||
aspectRatio: null,
|
||||
quality: "2k",
|
||||
referenceImages: ["a.png"],
|
||||
}),
|
||||
/total pixels between 768\*768 and 2048\*2048/,
|
||||
);
|
||||
|
||||
assert.throws(
|
||||
() =>
|
||||
resolveSizeForModel("wan2.7-image", {
|
||||
size: "4096x4096",
|
||||
aspectRatio: null,
|
||||
quality: "2k",
|
||||
}),
|
||||
/total pixels between 768\*768 and 2048\*2048/,
|
||||
);
|
||||
|
||||
assert.throws(
|
||||
() =>
|
||||
resolveSizeForModel("wan2.7-image-pro", {
|
||||
size: "3072*256",
|
||||
aspectRatio: null,
|
||||
quality: "2k",
|
||||
}),
|
||||
/1:8, 8:1/,
|
||||
);
|
||||
});
|
||||
@@ -0,0 +1,627 @@
|
||||
import path from "node:path";
|
||||
import { readFile } from "node:fs/promises";
|
||||
import type { CliArgs, Quality } from "../types";
|
||||
|
||||
type DashScopeModelFamily = "qwen2" | "qwenFixed" | "wan27" | "legacy";
|
||||
|
||||
type DashScopeModelSpec = {
|
||||
family: DashScopeModelFamily;
|
||||
defaultSize: string;
|
||||
};
|
||||
|
||||
const DEFAULT_MODEL = "qwen-image-2.0-pro";
|
||||
const MIN_QWEN_2_TOTAL_PIXELS = 512 * 512;
|
||||
const MAX_QWEN_2_TOTAL_PIXELS = 2048 * 2048;
|
||||
const SIZE_STEP = 16;
|
||||
const QWEN_NEGATIVE_PROMPT =
|
||||
"低分辨率,低画质,肢体畸形,手指畸形,画面过饱和,蜡像感,人脸无细节,过度光滑,画面具有AI感,构图混乱,文字模糊,扭曲";
|
||||
|
||||
const QWEN_2_TARGET_PIXELS: Record<Quality, number> = {
|
||||
normal: 1024 * 1024,
|
||||
"2k": 1536 * 1536,
|
||||
};
|
||||
|
||||
const MIN_WAN27_TOTAL_PIXELS = 768 * 768;
|
||||
const MAX_WAN27_PRO_T2I_PIXELS = 4096 * 4096;
|
||||
const MAX_WAN27_GENERAL_PIXELS = 2048 * 2048;
|
||||
const WAN27_MAX_REFERENCE_IMAGES = 9;
|
||||
|
||||
const WAN27_TARGET_PIXELS: Record<Quality, number> = {
|
||||
normal: 1024 * 1024,
|
||||
"2k": 2048 * 2048,
|
||||
};
|
||||
|
||||
const QWEN_2_RECOMMENDED: Record<string, Record<Quality, string>> = {
|
||||
"1:1": { normal: "1024*1024", "2k": "1536*1536" },
|
||||
"2:3": { normal: "768*1152", "2k": "1024*1536" },
|
||||
"3:2": { normal: "1152*768", "2k": "1536*1024" },
|
||||
"3:4": { normal: "960*1280", "2k": "1080*1440" },
|
||||
"4:3": { normal: "1280*960", "2k": "1440*1080" },
|
||||
"9:16": { normal: "720*1280", "2k": "1080*1920" },
|
||||
"16:9": { normal: "1280*720", "2k": "1920*1080" },
|
||||
"21:9": { normal: "1344*576", "2k": "2048*872" },
|
||||
};
|
||||
|
||||
const QWEN_FIXED_SIZES_BY_RATIO: Record<string, string> = {
|
||||
"16:9": "1664*928",
|
||||
"4:3": "1472*1104",
|
||||
"1:1": "1328*1328",
|
||||
"3:4": "1104*1472",
|
||||
"9:16": "928*1664",
|
||||
};
|
||||
|
||||
const QWEN_FIXED_SIZES = Object.values(QWEN_FIXED_SIZES_BY_RATIO);
|
||||
|
||||
const LEGACY_STANDARD_SIZES: [number, number][] = [
|
||||
[1024, 1024],
|
||||
[1280, 720],
|
||||
[720, 1280],
|
||||
[1024, 768],
|
||||
[768, 1024],
|
||||
[1536, 1024],
|
||||
[1024, 1536],
|
||||
[1536, 864],
|
||||
[864, 1536],
|
||||
];
|
||||
|
||||
const LEGACY_STANDARD_SIZES_2K: [number, number][] = [
|
||||
[1536, 1536],
|
||||
[2048, 1152],
|
||||
[1152, 2048],
|
||||
[1536, 1024],
|
||||
[1024, 1536],
|
||||
[1536, 864],
|
||||
[864, 1536],
|
||||
[2048, 2048],
|
||||
];
|
||||
|
||||
const QWEN_2_SPEC: DashScopeModelSpec = {
|
||||
family: "qwen2",
|
||||
defaultSize: "1024*1024",
|
||||
};
|
||||
|
||||
const QWEN_FIXED_SPEC: DashScopeModelSpec = {
|
||||
family: "qwenFixed",
|
||||
defaultSize: QWEN_FIXED_SIZES_BY_RATIO["16:9"],
|
||||
};
|
||||
|
||||
const WAN27_SPEC: DashScopeModelSpec = {
|
||||
family: "wan27",
|
||||
defaultSize: "2048*2048",
|
||||
};
|
||||
|
||||
const LEGACY_SPEC: DashScopeModelSpec = {
|
||||
family: "legacy",
|
||||
defaultSize: "1536*1536",
|
||||
};
|
||||
|
||||
const MODEL_SPEC_ALIASES: Record<string, DashScopeModelSpec> = {
|
||||
"qwen-image-3.0-pro": QWEN_2_SPEC,
|
||||
"qwen-image-2.0-pro": QWEN_2_SPEC,
|
||||
"qwen-image-2.0-pro-2026-04-22": QWEN_2_SPEC,
|
||||
"qwen-image-2.0-pro-2026-03-03": QWEN_2_SPEC,
|
||||
"qwen-image-2.0": QWEN_2_SPEC,
|
||||
"qwen-image-2.0-2026-03-03": QWEN_2_SPEC,
|
||||
"qwen-image-max": QWEN_FIXED_SPEC,
|
||||
"qwen-image-max-2025-12-30": QWEN_FIXED_SPEC,
|
||||
"qwen-image-plus": QWEN_FIXED_SPEC,
|
||||
"qwen-image-plus-2026-01-09": QWEN_FIXED_SPEC,
|
||||
"qwen-image": QWEN_FIXED_SPEC,
|
||||
"wan2.7-image-pro": WAN27_SPEC,
|
||||
"wan2.7-image": WAN27_SPEC,
|
||||
};
|
||||
|
||||
export function getDefaultModel(): string {
|
||||
return process.env.DASHSCOPE_IMAGE_MODEL || DEFAULT_MODEL;
|
||||
}
|
||||
|
||||
function getReferenceImageMime(filePath: string): string {
|
||||
const ext = path.extname(filePath).toLowerCase();
|
||||
if (ext === ".jpg" || ext === ".jpeg") return "image/jpeg";
|
||||
if (ext === ".webp") return "image/webp";
|
||||
if (ext === ".bmp") return "image/bmp";
|
||||
return "image/png";
|
||||
}
|
||||
|
||||
async function loadReferenceImage(refPath: string): Promise<string> {
|
||||
if (/^https?:\/\//i.test(refPath)) {
|
||||
return refPath;
|
||||
}
|
||||
const fullPath = path.resolve(refPath);
|
||||
const bytes = await readFile(fullPath);
|
||||
return `data:${getReferenceImageMime(fullPath)};base64,${bytes.toString("base64")}`;
|
||||
}
|
||||
|
||||
function getApiKey(): string | null {
|
||||
return process.env.DASHSCOPE_API_KEY || null;
|
||||
}
|
||||
|
||||
function getBaseUrl(): string {
|
||||
const base = process.env.DASHSCOPE_BASE_URL || "https://dashscope.aliyuncs.com";
|
||||
return base.replace(/\/+$/g, "");
|
||||
}
|
||||
|
||||
function getModelSpec(model: string): DashScopeModelSpec {
|
||||
return MODEL_SPEC_ALIASES[model.trim().toLowerCase()] || LEGACY_SPEC;
|
||||
}
|
||||
|
||||
export function getModelFamily(model: string): DashScopeModelFamily {
|
||||
return getModelSpec(model).family;
|
||||
}
|
||||
|
||||
function normalizeQuality(quality: CliArgs["quality"]): Quality {
|
||||
return quality === "normal" ? "normal" : "2k";
|
||||
}
|
||||
|
||||
export function parseAspectRatio(ar: string): { width: number; height: number } | null {
|
||||
const match = ar.match(/^(\d+(?:\.\d+)?):(\d+(?:\.\d+)?)$/);
|
||||
if (!match) return null;
|
||||
const w = parseFloat(match[1]!);
|
||||
const h = parseFloat(match[2]!);
|
||||
if (w <= 0 || h <= 0) return null;
|
||||
return { width: w, height: h };
|
||||
}
|
||||
|
||||
export function normalizeSize(size: string): string {
|
||||
return size.replace("x", "*");
|
||||
}
|
||||
|
||||
export function parseSize(size: string): { width: number; height: number } | null {
|
||||
const match = normalizeSize(size).match(/^(\d+)\*(\d+)$/);
|
||||
if (!match) return null;
|
||||
const width = Number(match[1]);
|
||||
const height = Number(match[2]);
|
||||
if (!Number.isFinite(width) || !Number.isFinite(height) || width <= 0 || height <= 0) {
|
||||
return null;
|
||||
}
|
||||
return { width, height };
|
||||
}
|
||||
|
||||
function formatSize(width: number, height: number): string {
|
||||
return `${width}*${height}`;
|
||||
}
|
||||
|
||||
function getRatioValue(ar: string): number | null {
|
||||
const parsed = parseAspectRatio(ar);
|
||||
if (!parsed) return null;
|
||||
return parsed.width / parsed.height;
|
||||
}
|
||||
|
||||
function findKnownRatioKey(ar: string, candidates: string[], tolerance = 0.02): string | null {
|
||||
const targetRatio = getRatioValue(ar);
|
||||
if (targetRatio == null) return null;
|
||||
|
||||
let bestKey: string | null = null;
|
||||
let bestDiff = Infinity;
|
||||
|
||||
for (const candidate of candidates) {
|
||||
const candidateRatio = getRatioValue(candidate);
|
||||
if (candidateRatio == null) continue;
|
||||
const diff = Math.abs(candidateRatio - targetRatio);
|
||||
if (diff < bestDiff) {
|
||||
bestDiff = diff;
|
||||
bestKey = candidate;
|
||||
}
|
||||
}
|
||||
|
||||
return bestDiff <= tolerance ? bestKey : null;
|
||||
}
|
||||
|
||||
function roundToStep(value: number): number {
|
||||
return Math.max(SIZE_STEP, Math.round(value / SIZE_STEP) * SIZE_STEP);
|
||||
}
|
||||
|
||||
function floorToStep(value: number): number {
|
||||
return Math.max(SIZE_STEP, Math.floor(value / SIZE_STEP) * SIZE_STEP);
|
||||
}
|
||||
|
||||
function fitToPixelBudget(
|
||||
width: number,
|
||||
height: number,
|
||||
minPixels: number,
|
||||
maxPixels: number,
|
||||
): { width: number; height: number } {
|
||||
let nextWidth = width;
|
||||
let nextHeight = height;
|
||||
let pixels = nextWidth * nextHeight;
|
||||
|
||||
if (pixels > maxPixels) {
|
||||
const scale = Math.sqrt(maxPixels / pixels);
|
||||
nextWidth *= scale;
|
||||
nextHeight *= scale;
|
||||
} else if (pixels < minPixels) {
|
||||
const scale = Math.sqrt(minPixels / pixels);
|
||||
nextWidth *= scale;
|
||||
nextHeight *= scale;
|
||||
}
|
||||
|
||||
let roundedWidth = roundToStep(nextWidth);
|
||||
let roundedHeight = roundToStep(nextHeight);
|
||||
pixels = roundedWidth * roundedHeight;
|
||||
|
||||
while (pixels > maxPixels && (roundedWidth > SIZE_STEP || roundedHeight > SIZE_STEP)) {
|
||||
if (roundedWidth >= roundedHeight && roundedWidth > SIZE_STEP) {
|
||||
roundedWidth -= SIZE_STEP;
|
||||
} else if (roundedHeight > SIZE_STEP) {
|
||||
roundedHeight -= SIZE_STEP;
|
||||
} else {
|
||||
break;
|
||||
}
|
||||
pixels = roundedWidth * roundedHeight;
|
||||
}
|
||||
|
||||
while (pixels < minPixels) {
|
||||
if (roundedWidth <= roundedHeight) {
|
||||
roundedWidth += SIZE_STEP;
|
||||
} else {
|
||||
roundedHeight += SIZE_STEP;
|
||||
}
|
||||
pixels = roundedWidth * roundedHeight;
|
||||
}
|
||||
|
||||
return { width: roundedWidth, height: roundedHeight };
|
||||
}
|
||||
|
||||
function clampWan27DerivedSizeToRatioBounds(
|
||||
size: { width: number; height: number },
|
||||
): { width: number; height: number } {
|
||||
let { width, height } = size;
|
||||
const ratio = width / height;
|
||||
|
||||
if (ratio > 8) {
|
||||
width = floorToStep(height * 8);
|
||||
} else if (ratio < 1 / 8) {
|
||||
height = floorToStep(width * 8);
|
||||
}
|
||||
|
||||
return { width, height };
|
||||
}
|
||||
|
||||
export function getSizeFromAspectRatio(ar: string | null, quality: CliArgs["quality"]): string {
|
||||
const normalizedQuality = normalizeQuality(quality);
|
||||
const sizes = normalizedQuality === "2k" ? LEGACY_STANDARD_SIZES_2K : LEGACY_STANDARD_SIZES;
|
||||
const defaultSize = normalizedQuality === "2k" ? "1536*1536" : "1024*1024";
|
||||
|
||||
if (!ar) return defaultSize;
|
||||
|
||||
const parsed = parseAspectRatio(ar);
|
||||
if (!parsed) return defaultSize;
|
||||
|
||||
const targetRatio = parsed.width / parsed.height;
|
||||
let best = defaultSize;
|
||||
let bestDiff = Infinity;
|
||||
|
||||
for (const [width, height] of sizes) {
|
||||
const diff = Math.abs(width / height - targetRatio);
|
||||
if (diff < bestDiff) {
|
||||
bestDiff = diff;
|
||||
best = formatSize(width, height);
|
||||
}
|
||||
}
|
||||
|
||||
return best;
|
||||
}
|
||||
|
||||
export function getQwen2SizeFromAspectRatio(ar: string | null, quality: CliArgs["quality"]): string {
|
||||
const normalizedQuality = normalizeQuality(quality);
|
||||
|
||||
if (!ar) {
|
||||
return QWEN_2_RECOMMENDED["1:1"][normalizedQuality];
|
||||
}
|
||||
|
||||
const recommendedRatio = findKnownRatioKey(ar, Object.keys(QWEN_2_RECOMMENDED));
|
||||
if (recommendedRatio) {
|
||||
return QWEN_2_RECOMMENDED[recommendedRatio][normalizedQuality];
|
||||
}
|
||||
|
||||
const parsed = parseAspectRatio(ar);
|
||||
if (!parsed) {
|
||||
return QWEN_2_RECOMMENDED["1:1"][normalizedQuality];
|
||||
}
|
||||
|
||||
const targetRatio = parsed.width / parsed.height;
|
||||
const targetPixels = QWEN_2_TARGET_PIXELS[normalizedQuality];
|
||||
const rawWidth = Math.sqrt(targetPixels * targetRatio);
|
||||
const rawHeight = Math.sqrt(targetPixels / targetRatio);
|
||||
const fitted = fitToPixelBudget(
|
||||
rawWidth,
|
||||
rawHeight,
|
||||
MIN_QWEN_2_TOTAL_PIXELS,
|
||||
MAX_QWEN_2_TOTAL_PIXELS,
|
||||
);
|
||||
|
||||
return formatSize(fitted.width, fitted.height);
|
||||
}
|
||||
|
||||
function isWan27ProModel(model: string): boolean {
|
||||
return model.trim().toLowerCase() === "wan2.7-image-pro";
|
||||
}
|
||||
|
||||
function getWan27MaxPixels(model: string, hasReferenceImages: boolean): number {
|
||||
if (isWan27ProModel(model) && !hasReferenceImages) {
|
||||
return MAX_WAN27_PRO_T2I_PIXELS;
|
||||
}
|
||||
return MAX_WAN27_GENERAL_PIXELS;
|
||||
}
|
||||
|
||||
export function getWan27SizeFromAspectRatio(
|
||||
ar: string | null,
|
||||
quality: CliArgs["quality"],
|
||||
maxPixels: number,
|
||||
): string {
|
||||
const normalizedQuality = normalizeQuality(quality);
|
||||
const targetPixels = Math.min(WAN27_TARGET_PIXELS[normalizedQuality], maxPixels);
|
||||
|
||||
if (!ar) {
|
||||
const side = roundToStep(Math.sqrt(targetPixels));
|
||||
return formatSize(side, side);
|
||||
}
|
||||
|
||||
const parsed = parseAspectRatio(ar);
|
||||
if (!parsed) {
|
||||
const side = roundToStep(Math.sqrt(targetPixels));
|
||||
return formatSize(side, side);
|
||||
}
|
||||
|
||||
const ratio = parsed.width / parsed.height;
|
||||
if (ratio < 1 / 8 || ratio > 8) {
|
||||
throw new Error(
|
||||
`DashScope wan2.7 image models support aspect ratios in [1:8, 8:1]. Received "${ar}".`
|
||||
);
|
||||
}
|
||||
|
||||
const rawWidth = Math.sqrt(targetPixels * ratio);
|
||||
const rawHeight = Math.sqrt(targetPixels / ratio);
|
||||
const fitted = fitToPixelBudget(
|
||||
rawWidth,
|
||||
rawHeight,
|
||||
MIN_WAN27_TOTAL_PIXELS,
|
||||
maxPixels,
|
||||
);
|
||||
const bounded = clampWan27DerivedSizeToRatioBounds(fitted);
|
||||
|
||||
return formatSize(bounded.width, bounded.height);
|
||||
}
|
||||
|
||||
function validateWan27Size(size: string, maxPixels: number, model: string): string {
|
||||
const normalized = normalizeSize(size);
|
||||
const parsed = validateSizeFormat(normalized);
|
||||
const totalPixels = parsed.width * parsed.height;
|
||||
if (totalPixels < MIN_WAN27_TOTAL_PIXELS || totalPixels > maxPixels) {
|
||||
const limit = maxPixels === MAX_WAN27_PRO_T2I_PIXELS ? "4096*4096" : "2048*2048";
|
||||
throw new Error(
|
||||
`DashScope ${model} requires total pixels between 768*768 and ${limit} ` +
|
||||
`for the current request. Received ${normalized} (${totalPixels} pixels).`
|
||||
);
|
||||
}
|
||||
const ratio = parsed.width / parsed.height;
|
||||
if (ratio < 1 / 8 || ratio > 8) {
|
||||
throw new Error(
|
||||
`DashScope wan2.7 image models support aspect ratios in [1:8, 8:1]. ` +
|
||||
`Received ${normalized} (ratio ${ratio.toFixed(3)}).`
|
||||
);
|
||||
}
|
||||
return normalized;
|
||||
}
|
||||
|
||||
function getQwenFixedSizeFromAspectRatio(ar: string | null, quality: CliArgs["quality"]): string {
|
||||
if (quality === "normal") {
|
||||
console.warn(
|
||||
"DashScope qwen-image-max/plus/image models use fixed output sizes; --quality normal does not change the generated resolution."
|
||||
);
|
||||
}
|
||||
|
||||
if (!ar) return QWEN_FIXED_SPEC.defaultSize;
|
||||
|
||||
const ratioKey = findKnownRatioKey(ar, Object.keys(QWEN_FIXED_SIZES_BY_RATIO));
|
||||
if (!ratioKey) {
|
||||
throw new Error(
|
||||
`DashScope model supports only fixed ratios ${Object.keys(QWEN_FIXED_SIZES_BY_RATIO).join(", ")}. ` +
|
||||
`For custom ratios like "${ar}", use --model qwen-image-2.0-pro.`
|
||||
);
|
||||
}
|
||||
|
||||
return QWEN_FIXED_SIZES_BY_RATIO[ratioKey]!;
|
||||
}
|
||||
|
||||
function validateSizeFormat(size: string): { width: number; height: number } {
|
||||
const parsed = parseSize(size);
|
||||
if (!parsed) {
|
||||
throw new Error(`Invalid DashScope size "${size}". Expected <width>x<height> or <width>*<height>.`);
|
||||
}
|
||||
return parsed;
|
||||
}
|
||||
|
||||
function validateQwen2Size(size: string): string {
|
||||
const normalized = normalizeSize(size);
|
||||
const parsed = validateSizeFormat(normalized);
|
||||
const totalPixels = parsed.width * parsed.height;
|
||||
if (totalPixels < MIN_QWEN_2_TOTAL_PIXELS || totalPixels > MAX_QWEN_2_TOTAL_PIXELS) {
|
||||
throw new Error(
|
||||
`DashScope qwen-image-2.0* models require total pixels between ${MIN_QWEN_2_TOTAL_PIXELS} ` +
|
||||
`and ${MAX_QWEN_2_TOTAL_PIXELS}. Received ${normalized} (${totalPixels} pixels).`
|
||||
);
|
||||
}
|
||||
return normalized;
|
||||
}
|
||||
|
||||
function validateQwenFixedSize(size: string): string {
|
||||
const normalized = normalizeSize(size);
|
||||
validateSizeFormat(normalized);
|
||||
if (!QWEN_FIXED_SIZES.includes(normalized)) {
|
||||
throw new Error(
|
||||
`DashScope qwen-image-max/plus/image models support only these sizes: ${QWEN_FIXED_SIZES.join(", ")}. ` +
|
||||
`Received ${normalized}.`
|
||||
);
|
||||
}
|
||||
return normalized;
|
||||
}
|
||||
|
||||
export function resolveSizeForModel(
|
||||
model: string,
|
||||
args: Pick<CliArgs, "size" | "aspectRatio" | "quality"> & { referenceImages?: string[] },
|
||||
): string {
|
||||
const spec = getModelSpec(model);
|
||||
const referenceCount = args.referenceImages?.length ?? 0;
|
||||
|
||||
if (spec.family === "wan27") {
|
||||
const maxPixels = getWan27MaxPixels(model, referenceCount > 0);
|
||||
if (args.size) return validateWan27Size(args.size, maxPixels, model);
|
||||
return getWan27SizeFromAspectRatio(args.aspectRatio, args.quality, maxPixels);
|
||||
}
|
||||
|
||||
if (args.size) {
|
||||
if (spec.family === "qwen2") return validateQwen2Size(args.size);
|
||||
if (spec.family === "qwenFixed") return validateQwenFixedSize(args.size);
|
||||
validateSizeFormat(args.size);
|
||||
return normalizeSize(args.size);
|
||||
}
|
||||
|
||||
if (spec.family === "qwen2") {
|
||||
return getQwen2SizeFromAspectRatio(args.aspectRatio, args.quality);
|
||||
}
|
||||
|
||||
if (spec.family === "qwenFixed") {
|
||||
return getQwenFixedSizeFromAspectRatio(args.aspectRatio, args.quality);
|
||||
}
|
||||
|
||||
return getSizeFromAspectRatio(args.aspectRatio, args.quality);
|
||||
}
|
||||
|
||||
function buildParameters(
|
||||
family: DashScopeModelFamily,
|
||||
size: string,
|
||||
): Record<string, unknown> {
|
||||
if (family === "wan27") {
|
||||
return {
|
||||
size,
|
||||
n: 1,
|
||||
watermark: false,
|
||||
};
|
||||
}
|
||||
|
||||
const parameters: Record<string, unknown> = {
|
||||
prompt_extend: false,
|
||||
size,
|
||||
};
|
||||
|
||||
if (family === "qwen2" || family === "qwenFixed") {
|
||||
parameters.watermark = false;
|
||||
parameters.negative_prompt = QWEN_NEGATIVE_PROMPT;
|
||||
}
|
||||
|
||||
return parameters;
|
||||
}
|
||||
|
||||
type DashScopeResponse = {
|
||||
output?: {
|
||||
result_image?: string;
|
||||
choices?: Array<{
|
||||
message?: {
|
||||
content?: Array<{ image?: string }>;
|
||||
};
|
||||
}>;
|
||||
};
|
||||
};
|
||||
|
||||
async function extractImageFromResponse(result: DashScopeResponse): Promise<Uint8Array> {
|
||||
let imageData: string | null = null;
|
||||
|
||||
if (result.output?.result_image) {
|
||||
imageData = result.output.result_image;
|
||||
} else if (result.output?.choices?.[0]?.message?.content) {
|
||||
const content = result.output.choices[0].message.content;
|
||||
for (const item of content) {
|
||||
if (item.image) {
|
||||
imageData = item.image;
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (!imageData) {
|
||||
console.error("Response:", JSON.stringify(result, null, 2));
|
||||
throw new Error("No image in response");
|
||||
}
|
||||
|
||||
if (imageData.startsWith("http://") || imageData.startsWith("https://")) {
|
||||
const imgRes = await fetch(imageData);
|
||||
if (!imgRes.ok) throw new Error("Failed to download image");
|
||||
const buf = await imgRes.arrayBuffer();
|
||||
return new Uint8Array(buf);
|
||||
}
|
||||
|
||||
return Uint8Array.from(Buffer.from(imageData, "base64"));
|
||||
}
|
||||
|
||||
export async function generateImage(
|
||||
prompt: string,
|
||||
model: string,
|
||||
args: CliArgs
|
||||
): Promise<Uint8Array> {
|
||||
const apiKey = getApiKey();
|
||||
if (!apiKey) throw new Error("DASHSCOPE_API_KEY is required");
|
||||
|
||||
const spec = getModelSpec(model);
|
||||
|
||||
if (args.referenceImages.length > 0 && spec.family !== "wan27") {
|
||||
throw new Error(
|
||||
"Reference images are not supported with this DashScope model. Use a wan2.7 image model (--model wan2.7-image-pro or wan2.7-image), or switch to --provider google with a Gemini multimodal model."
|
||||
);
|
||||
}
|
||||
|
||||
if (args.referenceImages.length > WAN27_MAX_REFERENCE_IMAGES) {
|
||||
throw new Error(
|
||||
`DashScope wan2.7 image models accept at most ${WAN27_MAX_REFERENCE_IMAGES} reference images. Received ${args.referenceImages.length}.`
|
||||
);
|
||||
}
|
||||
|
||||
if (spec.family === "wan27" && args.n !== 1) {
|
||||
throw new Error(
|
||||
"DashScope wan2.7 image models in baoyu-image-gen support exactly one output image per request (extra images would be billed but discarded). Remove --n or use --n 1."
|
||||
);
|
||||
}
|
||||
|
||||
const size = resolveSizeForModel(model, args);
|
||||
const url = `${getBaseUrl()}/api/v1/services/aigc/multimodal-generation/generation`;
|
||||
|
||||
const content: Array<Record<string, unknown>> = [];
|
||||
if (spec.family === "wan27" && args.referenceImages.length > 0) {
|
||||
for (const refPath of args.referenceImages) {
|
||||
content.push({ image: await loadReferenceImage(refPath) });
|
||||
}
|
||||
}
|
||||
content.push({ text: prompt });
|
||||
|
||||
const body = {
|
||||
model,
|
||||
input: {
|
||||
messages: [
|
||||
{
|
||||
role: "user",
|
||||
content,
|
||||
},
|
||||
],
|
||||
},
|
||||
parameters: buildParameters(spec.family, size),
|
||||
};
|
||||
|
||||
console.log(`Generating image with DashScope (${model})...`, { family: spec.family, size });
|
||||
|
||||
const res = await fetch(url, {
|
||||
method: "POST",
|
||||
headers: {
|
||||
"Content-Type": "application/json",
|
||||
Authorization: `Bearer ${apiKey}`,
|
||||
},
|
||||
body: JSON.stringify(body),
|
||||
});
|
||||
|
||||
if (!res.ok) {
|
||||
const err = await res.text();
|
||||
throw new Error(`DashScope API error (${res.status}): ${err}`);
|
||||
}
|
||||
|
||||
const result = await res.json() as DashScopeResponse;
|
||||
return extractImageFromResponse(result);
|
||||
}
|
||||
@@ -0,0 +1,146 @@
|
||||
import assert from "node:assert/strict";
|
||||
import test, { type TestContext } from "node:test";
|
||||
|
||||
import type { CliArgs } from "../types.ts";
|
||||
import {
|
||||
addAspectRatioToPrompt,
|
||||
buildGoogleUrl,
|
||||
buildPromptWithAspect,
|
||||
extractInlineImageData,
|
||||
extractPredictedImageData,
|
||||
getGoogleImageSize,
|
||||
isGoogleImagen,
|
||||
is1KOnlyGoogleModel,
|
||||
isGoogleMultimodal,
|
||||
normalizeGoogleModelId,
|
||||
resolveGeminiImageSize,
|
||||
} from "./google.ts";
|
||||
|
||||
function useEnv(
|
||||
t: TestContext,
|
||||
values: Record<string, string | null>,
|
||||
): void {
|
||||
const previous = new Map<string, string | undefined>();
|
||||
for (const [key, value] of Object.entries(values)) {
|
||||
previous.set(key, process.env[key]);
|
||||
if (value == null) {
|
||||
delete process.env[key];
|
||||
} else {
|
||||
process.env[key] = value;
|
||||
}
|
||||
}
|
||||
|
||||
t.after(() => {
|
||||
for (const [key, value] of previous.entries()) {
|
||||
if (value == null) {
|
||||
delete process.env[key];
|
||||
} else {
|
||||
process.env[key] = value;
|
||||
}
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
function makeArgs(overrides: Partial<CliArgs> = {}): CliArgs {
|
||||
return {
|
||||
prompt: null,
|
||||
promptFiles: [],
|
||||
imagePath: null,
|
||||
provider: null,
|
||||
model: null,
|
||||
aspectRatio: null,
|
||||
size: null,
|
||||
quality: null,
|
||||
imageSize: null,
|
||||
imageApiDialect: null,
|
||||
referenceImages: [],
|
||||
n: 1,
|
||||
batchFile: null,
|
||||
jobs: null,
|
||||
json: false,
|
||||
help: false,
|
||||
...overrides,
|
||||
};
|
||||
}
|
||||
|
||||
test("Google provider helpers normalize model IDs and select image size defaults", () => {
|
||||
assert.equal(
|
||||
normalizeGoogleModelId("models/gemini-3.1-flash-image-preview"),
|
||||
"gemini-3.1-flash-image-preview",
|
||||
);
|
||||
assert.equal(isGoogleMultimodal("models/gemini-3-pro-image-preview"), true);
|
||||
assert.equal(isGoogleMultimodal("gemini-3-pro-image"), true);
|
||||
assert.equal(isGoogleMultimodal("gemini-3.1-flash-image"), true);
|
||||
assert.equal(isGoogleMultimodal("models/gemini-3-pro-image"), true);
|
||||
assert.equal(isGoogleImagen("imagen-3.0-generate-002"), true);
|
||||
assert.equal(getGoogleImageSize(makeArgs({ imageSize: null, quality: "2k" })), "2K");
|
||||
assert.equal(getGoogleImageSize(makeArgs({ imageSize: "4K", quality: "normal" })), "4K");
|
||||
});
|
||||
|
||||
test("Google clamps 1K-only models to 1K output", () => {
|
||||
assert.equal(isGoogleMultimodal("gemini-3.1-flash-lite-image"), true);
|
||||
assert.equal(is1KOnlyGoogleModel("gemini-3.1-flash-lite-image"), true);
|
||||
assert.equal(is1KOnlyGoogleModel("gemini-3.1-flash-image"), false);
|
||||
assert.equal(
|
||||
resolveGeminiImageSize("gemini-3.1-flash-lite-image", makeArgs({ imageSize: "4K", quality: "2k" })),
|
||||
"1K",
|
||||
);
|
||||
assert.equal(
|
||||
resolveGeminiImageSize("gemini-3-pro-image", makeArgs({ imageSize: null, quality: "2k" })),
|
||||
"2K",
|
||||
);
|
||||
});
|
||||
|
||||
test("Google URL builder appends v1beta when the base URL does not already include it", (t) => {
|
||||
useEnv(t, { GOOGLE_BASE_URL: "https://generativelanguage.googleapis.com" });
|
||||
assert.equal(
|
||||
buildGoogleUrl("models/demo:generateContent"),
|
||||
"https://generativelanguage.googleapis.com/v1beta/models/demo:generateContent",
|
||||
);
|
||||
});
|
||||
|
||||
test("Google URL and prompt helpers preserve existing v1beta paths and aspect hints", (t) => {
|
||||
useEnv(t, { GOOGLE_BASE_URL: "https://example.com/custom/v1beta/" });
|
||||
assert.equal(
|
||||
buildGoogleUrl("/models/demo:predict"),
|
||||
"https://example.com/custom/v1beta/models/demo:predict",
|
||||
);
|
||||
|
||||
assert.equal(
|
||||
addAspectRatioToPrompt("A city skyline", "16:9"),
|
||||
"A city skyline Aspect ratio: 16:9.",
|
||||
);
|
||||
assert.equal(
|
||||
buildPromptWithAspect("A city skyline", "16:9", "2k"),
|
||||
"A city skyline Aspect ratio: 16:9. High resolution 2048px.",
|
||||
);
|
||||
});
|
||||
|
||||
test("Google response extractors find inline and predicted image payloads", () => {
|
||||
assert.equal(
|
||||
extractInlineImageData({
|
||||
candidates: [
|
||||
{
|
||||
content: {
|
||||
parts: [{ inlineData: { data: "inline-base64" } }],
|
||||
},
|
||||
},
|
||||
],
|
||||
}),
|
||||
"inline-base64",
|
||||
);
|
||||
|
||||
assert.equal(
|
||||
extractPredictedImageData({
|
||||
predictions: [{ image: { imageBytes: "predicted-base64" } }],
|
||||
}),
|
||||
"predicted-base64",
|
||||
);
|
||||
|
||||
assert.equal(
|
||||
extractPredictedImageData({
|
||||
generatedImages: [{ bytesBase64Encoded: "generated-base64" }],
|
||||
}),
|
||||
"generated-base64",
|
||||
);
|
||||
});
|
||||
372
baoyu-skills/skills/baoyu-image-gen/scripts/providers/google.ts
Normal file
372
baoyu-skills/skills/baoyu-image-gen/scripts/providers/google.ts
Normal file
@@ -0,0 +1,372 @@
|
||||
import path from "node:path";
|
||||
import { readFile } from "node:fs/promises";
|
||||
import { execFileSync } from "node:child_process";
|
||||
import type { CliArgs } from "../types";
|
||||
|
||||
const GOOGLE_MULTIMODAL_MODELS = [
|
||||
"gemini-3-pro-image",
|
||||
"gemini-3.1-flash-image",
|
||||
"gemini-3.1-flash-lite-image",
|
||||
"gemini-3-pro-image-preview",
|
||||
"gemini-3-flash-preview",
|
||||
"gemini-3.1-flash-image-preview",
|
||||
];
|
||||
const GOOGLE_1K_ONLY_MODELS = ["gemini-3.1-flash-lite-image"];
|
||||
const GOOGLE_IMAGEN_MODELS = [
|
||||
"imagen-3.0-generate-002",
|
||||
"imagen-3.0-generate-001",
|
||||
];
|
||||
|
||||
export function getDefaultModel(): string {
|
||||
return process.env.GOOGLE_IMAGE_MODEL || "gemini-3-pro-image";
|
||||
}
|
||||
|
||||
export function normalizeGoogleModelId(model: string): string {
|
||||
return model.startsWith("models/") ? model.slice("models/".length) : model;
|
||||
}
|
||||
|
||||
export function isGoogleMultimodal(model: string): boolean {
|
||||
const normalized = normalizeGoogleModelId(model);
|
||||
return GOOGLE_MULTIMODAL_MODELS.some((m) => normalized.includes(m));
|
||||
}
|
||||
|
||||
export function isGoogleImagen(model: string): boolean {
|
||||
const normalized = normalizeGoogleModelId(model);
|
||||
return GOOGLE_IMAGEN_MODELS.some((m) => normalized.includes(m));
|
||||
}
|
||||
|
||||
function getGoogleApiKey(): string | null {
|
||||
return process.env.GOOGLE_API_KEY || process.env.GEMINI_API_KEY || null;
|
||||
}
|
||||
|
||||
export function getGoogleImageSize(args: CliArgs): "1K" | "2K" | "4K" {
|
||||
if (args.imageSize) return args.imageSize as "1K" | "2K" | "4K";
|
||||
return args.quality === "2k" ? "2K" : "1K";
|
||||
}
|
||||
|
||||
export function is1KOnlyGoogleModel(model: string): boolean {
|
||||
const normalized = normalizeGoogleModelId(model);
|
||||
return GOOGLE_1K_ONLY_MODELS.some((m) => normalized.includes(m));
|
||||
}
|
||||
|
||||
export function resolveGeminiImageSize(
|
||||
model: string,
|
||||
args: CliArgs,
|
||||
): "1K" | "2K" | "4K" {
|
||||
const size = getGoogleImageSize(args);
|
||||
if (size !== "1K" && is1KOnlyGoogleModel(model)) {
|
||||
console.error(
|
||||
`Warning: ${normalizeGoogleModelId(model)} only supports 1K output, using 1K instead of ${size}.`,
|
||||
);
|
||||
return "1K";
|
||||
}
|
||||
return size;
|
||||
}
|
||||
|
||||
function getGoogleBaseUrl(): string {
|
||||
const base =
|
||||
process.env.GOOGLE_BASE_URL || "https://generativelanguage.googleapis.com";
|
||||
return base.replace(/\/+$/g, "");
|
||||
}
|
||||
|
||||
export function buildGoogleUrl(pathname: string): string {
|
||||
const base = getGoogleBaseUrl();
|
||||
const cleanedPath = pathname.replace(/^\/+/g, "");
|
||||
if (base.endsWith("/v1beta")) return `${base}/${cleanedPath}`;
|
||||
return `${base}/v1beta/${cleanedPath}`;
|
||||
}
|
||||
|
||||
function toModelPath(model: string): string {
|
||||
const modelId = normalizeGoogleModelId(model);
|
||||
return `models/${modelId}`;
|
||||
}
|
||||
|
||||
function getHttpProxy(): string | null {
|
||||
return (
|
||||
process.env.https_proxy ||
|
||||
process.env.HTTPS_PROXY ||
|
||||
process.env.http_proxy ||
|
||||
process.env.HTTP_PROXY ||
|
||||
process.env.ALL_PROXY ||
|
||||
null
|
||||
);
|
||||
}
|
||||
|
||||
async function postGoogleJsonViaCurl<T>(
|
||||
url: string,
|
||||
apiKey: string,
|
||||
body: unknown,
|
||||
): Promise<T> {
|
||||
const proxy = getHttpProxy();
|
||||
const bodyStr = JSON.stringify(body);
|
||||
const args = [
|
||||
"-s",
|
||||
"--connect-timeout",
|
||||
"30",
|
||||
"--max-time",
|
||||
"300",
|
||||
...(proxy ? ["-x", proxy] : []),
|
||||
url,
|
||||
"-H",
|
||||
"Content-Type: application/json",
|
||||
"-H",
|
||||
`x-goog-api-key: ${apiKey}`,
|
||||
"-d",
|
||||
"@-",
|
||||
];
|
||||
|
||||
let result = "";
|
||||
try {
|
||||
result = execFileSync("curl", args, {
|
||||
input: bodyStr,
|
||||
encoding: "utf8",
|
||||
maxBuffer: 100 * 1024 * 1024,
|
||||
timeout: 310000,
|
||||
});
|
||||
} catch (error) {
|
||||
const e = error as { message?: string; stderr?: string | Buffer };
|
||||
const stderrText =
|
||||
typeof e.stderr === "string"
|
||||
? e.stderr
|
||||
: e.stderr
|
||||
? e.stderr.toString("utf8")
|
||||
: "";
|
||||
const details = stderrText.trim() || e.message || "curl request failed";
|
||||
throw new Error(`Google API request failed via curl: ${details}`);
|
||||
}
|
||||
|
||||
const parsed = JSON.parse(result) as any;
|
||||
if (parsed.error) {
|
||||
throw new Error(
|
||||
`Google API error (${parsed.error.code}): ${parsed.error.message}`,
|
||||
);
|
||||
}
|
||||
return parsed as T;
|
||||
}
|
||||
|
||||
async function postGoogleJsonViaFetch<T>(
|
||||
url: string,
|
||||
apiKey: string,
|
||||
body: unknown,
|
||||
): Promise<T> {
|
||||
const res = await fetch(url, {
|
||||
method: "POST",
|
||||
headers: {
|
||||
"Content-Type": "application/json",
|
||||
"x-goog-api-key": apiKey,
|
||||
},
|
||||
body: JSON.stringify(body),
|
||||
});
|
||||
|
||||
if (!res.ok) {
|
||||
const err = await res.text();
|
||||
throw new Error(`Google API error (${res.status}): ${err}`);
|
||||
}
|
||||
|
||||
return (await res.json()) as T;
|
||||
}
|
||||
|
||||
async function postGoogleJson<T>(pathname: string, body: unknown): Promise<T> {
|
||||
const apiKey = getGoogleApiKey();
|
||||
if (!apiKey) throw new Error("GOOGLE_API_KEY or GEMINI_API_KEY is required");
|
||||
|
||||
const url = buildGoogleUrl(pathname);
|
||||
const proxy = getHttpProxy();
|
||||
|
||||
// When an HTTP proxy is detected, use curl instead of fetch.
|
||||
// Bun's fetch has a known issue where long-lived connections through
|
||||
// HTTP proxies get their sockets closed unexpectedly, causing image
|
||||
// generation requests to fail with "socket connection was closed
|
||||
// unexpectedly". Using curl as the HTTP client works around this.
|
||||
if (proxy) {
|
||||
return postGoogleJsonViaCurl<T>(url, apiKey, body);
|
||||
}
|
||||
|
||||
return postGoogleJsonViaFetch<T>(url, apiKey, body);
|
||||
}
|
||||
|
||||
export function buildPromptWithAspect(
|
||||
prompt: string,
|
||||
ar: string | null,
|
||||
quality: CliArgs["quality"],
|
||||
): string {
|
||||
let result = prompt;
|
||||
if (ar) {
|
||||
result += ` Aspect ratio: ${ar}.`;
|
||||
}
|
||||
if (quality === "2k") {
|
||||
result += " High resolution 2048px.";
|
||||
}
|
||||
return result;
|
||||
}
|
||||
|
||||
export function addAspectRatioToPrompt(prompt: string, ar: string | null): string {
|
||||
if (!ar) return prompt;
|
||||
return `${prompt} Aspect ratio: ${ar}.`;
|
||||
}
|
||||
|
||||
async function readImageAsBase64(
|
||||
p: string,
|
||||
): Promise<{ data: string; mimeType: string }> {
|
||||
const buf = await readFile(p);
|
||||
const ext = path.extname(p).toLowerCase();
|
||||
let mimeType = "image/png";
|
||||
if (ext === ".jpg" || ext === ".jpeg") mimeType = "image/jpeg";
|
||||
else if (ext === ".gif") mimeType = "image/gif";
|
||||
else if (ext === ".webp") mimeType = "image/webp";
|
||||
return { data: buf.toString("base64"), mimeType };
|
||||
}
|
||||
|
||||
export function extractInlineImageData(response: {
|
||||
candidates?: Array<{
|
||||
content?: { parts?: Array<{ inlineData?: { data?: string } }> };
|
||||
}>;
|
||||
}): string | null {
|
||||
for (const candidate of response.candidates || []) {
|
||||
for (const part of candidate.content?.parts || []) {
|
||||
const data = part.inlineData?.data;
|
||||
if (typeof data === "string" && data.length > 0) return data;
|
||||
}
|
||||
}
|
||||
return null;
|
||||
}
|
||||
|
||||
export function extractPredictedImageData(response: {
|
||||
predictions?: Array<any>;
|
||||
generatedImages?: Array<any>;
|
||||
}): string | null {
|
||||
const candidates = [
|
||||
...(response.predictions || []),
|
||||
...(response.generatedImages || []),
|
||||
];
|
||||
for (const candidate of candidates) {
|
||||
if (!candidate || typeof candidate !== "object") continue;
|
||||
if (typeof candidate.imageBytes === "string") return candidate.imageBytes;
|
||||
if (typeof candidate.bytesBase64Encoded === "string")
|
||||
return candidate.bytesBase64Encoded;
|
||||
if (typeof candidate.data === "string") return candidate.data;
|
||||
const image = candidate.image;
|
||||
if (image && typeof image === "object") {
|
||||
if (typeof image.imageBytes === "string") return image.imageBytes;
|
||||
if (typeof image.bytesBase64Encoded === "string")
|
||||
return image.bytesBase64Encoded;
|
||||
if (typeof image.data === "string") return image.data;
|
||||
}
|
||||
}
|
||||
return null;
|
||||
}
|
||||
|
||||
async function generateWithGemini(
|
||||
prompt: string,
|
||||
model: string,
|
||||
args: CliArgs,
|
||||
): Promise<Uint8Array> {
|
||||
const promptWithAspect = addAspectRatioToPrompt(prompt, args.aspectRatio);
|
||||
const parts: Array<{
|
||||
text?: string;
|
||||
inlineData?: { data: string; mimeType: string };
|
||||
}> = [];
|
||||
for (const refPath of args.referenceImages) {
|
||||
const { data, mimeType } = await readImageAsBase64(refPath);
|
||||
parts.push({ inlineData: { data, mimeType } });
|
||||
}
|
||||
parts.push({ text: promptWithAspect });
|
||||
|
||||
const imageConfig: { imageSize: "1K" | "2K" | "4K" } = {
|
||||
imageSize: resolveGeminiImageSize(model, args),
|
||||
};
|
||||
|
||||
console.log("Generating image with Gemini...", imageConfig);
|
||||
const response = await postGoogleJson<{
|
||||
candidates?: Array<{
|
||||
content?: { parts?: Array<{ inlineData?: { data?: string } }> };
|
||||
}>;
|
||||
}>(`${toModelPath(model)}:generateContent`, {
|
||||
contents: [
|
||||
{
|
||||
role: "user",
|
||||
parts,
|
||||
},
|
||||
],
|
||||
generationConfig: {
|
||||
responseModalities: ["IMAGE"],
|
||||
imageConfig,
|
||||
},
|
||||
});
|
||||
console.log("Generation completed.");
|
||||
|
||||
const imageData = extractInlineImageData(response);
|
||||
if (imageData) return Uint8Array.from(Buffer.from(imageData, "base64"));
|
||||
|
||||
throw new Error("No image in response");
|
||||
}
|
||||
|
||||
async function generateWithImagen(
|
||||
prompt: string,
|
||||
model: string,
|
||||
args: CliArgs,
|
||||
): Promise<Uint8Array> {
|
||||
const fullPrompt = buildPromptWithAspect(
|
||||
prompt,
|
||||
args.aspectRatio,
|
||||
args.quality,
|
||||
);
|
||||
const imageSize = getGoogleImageSize(args);
|
||||
if (imageSize === "4K") {
|
||||
console.error(
|
||||
"Warning: Imagen models do not support 4K imageSize, using 2K instead.",
|
||||
);
|
||||
}
|
||||
|
||||
const parameters: Record<string, unknown> = {
|
||||
sampleCount: args.n,
|
||||
};
|
||||
if (args.aspectRatio) {
|
||||
parameters.aspectRatio = args.aspectRatio;
|
||||
}
|
||||
if (imageSize === "1K" || imageSize === "2K") {
|
||||
parameters.imageSize = imageSize;
|
||||
} else {
|
||||
parameters.imageSize = "2K";
|
||||
}
|
||||
|
||||
const response = await postGoogleJson<{
|
||||
predictions?: Array<any>;
|
||||
generatedImages?: Array<any>;
|
||||
}>(`${toModelPath(model)}:predict`, {
|
||||
instances: [
|
||||
{
|
||||
prompt: fullPrompt,
|
||||
},
|
||||
],
|
||||
parameters,
|
||||
});
|
||||
|
||||
const imageData = extractPredictedImageData(response);
|
||||
if (imageData) return Uint8Array.from(Buffer.from(imageData, "base64"));
|
||||
|
||||
throw new Error("No image in response");
|
||||
}
|
||||
|
||||
export async function generateImage(
|
||||
prompt: string,
|
||||
model: string,
|
||||
args: CliArgs,
|
||||
): Promise<Uint8Array> {
|
||||
if (isGoogleImagen(model)) {
|
||||
if (args.referenceImages.length > 0) {
|
||||
throw new Error(
|
||||
"Reference images are not supported with Imagen models. Use a Gemini multimodal model such as gemini-3-pro-image, gemini-3.1-flash-image, gemini-3.1-flash-lite-image, gemini-3-pro-image-preview, gemini-3-flash-preview, or gemini-3.1-flash-image-preview.",
|
||||
);
|
||||
}
|
||||
return generateWithImagen(prompt, model, args);
|
||||
}
|
||||
|
||||
if (!isGoogleMultimodal(model) && args.referenceImages.length > 0) {
|
||||
throw new Error(
|
||||
"Reference images are only supported with Gemini multimodal models such as gemini-3-pro-image, gemini-3.1-flash-image, gemini-3.1-flash-lite-image, gemini-3-pro-image-preview, gemini-3-flash-preview, or gemini-3.1-flash-image-preview.",
|
||||
);
|
||||
}
|
||||
|
||||
return generateWithGemini(prompt, model, args);
|
||||
}
|
||||
@@ -0,0 +1,115 @@
|
||||
import assert from "node:assert/strict";
|
||||
import test, { type TestContext } from "node:test";
|
||||
|
||||
import type { CliArgs } from "../types.ts";
|
||||
import { generateImage } from "./jimeng.ts";
|
||||
|
||||
function makeArgs(overrides: Partial<CliArgs> = {}): CliArgs {
|
||||
return {
|
||||
prompt: null,
|
||||
promptFiles: [],
|
||||
imagePath: null,
|
||||
provider: null,
|
||||
model: null,
|
||||
aspectRatio: null,
|
||||
size: null,
|
||||
quality: null,
|
||||
imageSize: null,
|
||||
imageApiDialect: null,
|
||||
referenceImages: [],
|
||||
n: 1,
|
||||
batchFile: null,
|
||||
jobs: null,
|
||||
json: false,
|
||||
help: false,
|
||||
...overrides,
|
||||
};
|
||||
}
|
||||
|
||||
function useEnv(
|
||||
t: TestContext,
|
||||
values: Record<string, string | null>,
|
||||
): void {
|
||||
const previous = new Map<string, string | undefined>();
|
||||
for (const [key, value] of Object.entries(values)) {
|
||||
previous.set(key, process.env[key]);
|
||||
if (value == null) {
|
||||
delete process.env[key];
|
||||
} else {
|
||||
process.env[key] = value;
|
||||
}
|
||||
}
|
||||
|
||||
t.after(() => {
|
||||
for (const [key, value] of previous.entries()) {
|
||||
if (value == null) {
|
||||
delete process.env[key];
|
||||
} else {
|
||||
process.env[key] = value;
|
||||
}
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
test("Jimeng submit request uses prompt field expected by current API", async (t) => {
|
||||
useEnv(t, {
|
||||
JIMENG_ACCESS_KEY_ID: "test-access-key",
|
||||
JIMENG_SECRET_ACCESS_KEY: "test-secret-key",
|
||||
JIMENG_BASE_URL: null,
|
||||
JIMENG_REGION: null,
|
||||
});
|
||||
|
||||
const originalFetch = globalThis.fetch;
|
||||
t.after(() => {
|
||||
globalThis.fetch = originalFetch;
|
||||
});
|
||||
|
||||
const calls: Array<{
|
||||
input: string;
|
||||
init?: RequestInit;
|
||||
}> = [];
|
||||
|
||||
globalThis.fetch = async (input, init) => {
|
||||
calls.push({
|
||||
input: String(input),
|
||||
init,
|
||||
});
|
||||
|
||||
if (calls.length === 1) {
|
||||
return Response.json({
|
||||
code: 10000,
|
||||
data: {
|
||||
task_id: "task-123",
|
||||
},
|
||||
});
|
||||
}
|
||||
|
||||
return Response.json({
|
||||
code: 10000,
|
||||
data: {
|
||||
status: "done",
|
||||
binary_data_base64: [Buffer.from("jimeng-image").toString("base64")],
|
||||
},
|
||||
});
|
||||
};
|
||||
|
||||
const image = await generateImage(
|
||||
"A quiet bamboo forest",
|
||||
"jimeng_t2i_v40",
|
||||
makeArgs({ quality: "normal" }),
|
||||
);
|
||||
|
||||
assert.equal(Buffer.from(image).toString("utf8"), "jimeng-image");
|
||||
assert.equal(calls.length, 2);
|
||||
assert.equal(
|
||||
calls[0]?.input,
|
||||
"https://visual.volcengineapi.com/?Action=CVSync2AsyncSubmitTask&Version=2022-08-31",
|
||||
);
|
||||
|
||||
const submitBody = JSON.parse(String(calls[0]?.init?.body)) as Record<string, unknown>;
|
||||
assert.equal(submitBody.req_key, "jimeng_t2i_v40");
|
||||
assert.equal(submitBody.prompt, "A quiet bamboo forest");
|
||||
assert.ok(!("prompt_text" in submitBody));
|
||||
assert.equal(submitBody.width, 1024);
|
||||
assert.equal(submitBody.height, 1024);
|
||||
});
|
||||
467
baoyu-skills/skills/baoyu-image-gen/scripts/providers/jimeng.ts
Normal file
467
baoyu-skills/skills/baoyu-image-gen/scripts/providers/jimeng.ts
Normal file
@@ -0,0 +1,467 @@
|
||||
import type { CliArgs } from "../types";
|
||||
import * as crypto from "node:crypto";
|
||||
|
||||
type JimengSizePreset = "normal" | "2k" | "4k";
|
||||
|
||||
export function getDefaultModel(): string {
|
||||
return process.env.JIMENG_IMAGE_MODEL || "jimeng_t2i_v40";
|
||||
}
|
||||
|
||||
function getAccessKey(): string | null {
|
||||
return process.env.JIMENG_ACCESS_KEY_ID || null;
|
||||
}
|
||||
|
||||
function getSecretKey(): string | null {
|
||||
return process.env.JIMENG_SECRET_ACCESS_KEY || null;
|
||||
}
|
||||
|
||||
function getRegion(): string {
|
||||
return process.env.JIMENG_REGION || "cn-north-1";
|
||||
}
|
||||
|
||||
function getBaseUrl(): string {
|
||||
return process.env.JIMENG_BASE_URL || "https://visual.volcengineapi.com";
|
||||
}
|
||||
|
||||
function resolveEndpoint(query: Record<string, string>): {
|
||||
url: string;
|
||||
host: string;
|
||||
canonicalUri: string;
|
||||
} {
|
||||
let baseUrl: URL;
|
||||
try {
|
||||
baseUrl = new URL(getBaseUrl());
|
||||
} catch {
|
||||
throw new Error(`Invalid JIMENG_BASE_URL: ${getBaseUrl()}`);
|
||||
}
|
||||
|
||||
baseUrl.search = "";
|
||||
for (const [key, value] of Object.entries(query).sort(([a], [b]) => a.localeCompare(b))) {
|
||||
baseUrl.searchParams.set(key, value);
|
||||
}
|
||||
|
||||
return {
|
||||
url: baseUrl.toString(),
|
||||
host: baseUrl.host,
|
||||
canonicalUri: baseUrl.pathname || "/",
|
||||
};
|
||||
}
|
||||
|
||||
/**
|
||||
* Volcengine HMAC-SHA256 signature generation
|
||||
* Following the official documentation at:
|
||||
* https://www.volcengine.com/docs/85621/1817045
|
||||
*/
|
||||
function generateSignature(
|
||||
method: string,
|
||||
query: Record<string, string>,
|
||||
headers: Record<string, string>,
|
||||
body: string,
|
||||
accessKey: string,
|
||||
secretKey: string,
|
||||
region: string,
|
||||
service: string,
|
||||
canonicalUri: string
|
||||
): string {
|
||||
// 1. Create canonical request
|
||||
// Sort query parameters alphabetically
|
||||
const sortedQuery = Object.entries(query)
|
||||
.sort(([a], [b]) => a.localeCompare(b))
|
||||
.map(([k, v]) => `${encodeURIComponent(k)}=${encodeURIComponent(v)}`)
|
||||
.join("&");
|
||||
|
||||
// Sort headers alphabetically and create canonical headers
|
||||
const sortedHeaders = Object.entries(headers)
|
||||
.sort(([a], [b]) => a.localeCompare(b))
|
||||
.map(([k, v]) => `${k.toLowerCase()}:${v.trim()}\n`)
|
||||
.join("");
|
||||
|
||||
const signedHeaders = Object.keys(headers)
|
||||
.sort()
|
||||
.map(k => k.toLowerCase())
|
||||
.join(";");
|
||||
|
||||
const hashedPayload = crypto.createHash("sha256").update(body, "utf8").digest("hex");
|
||||
|
||||
const canonicalRequest = [
|
||||
method,
|
||||
canonicalUri,
|
||||
sortedQuery,
|
||||
sortedHeaders,
|
||||
signedHeaders,
|
||||
hashedPayload,
|
||||
].join("\n");
|
||||
|
||||
const hashedCanonicalRequest = crypto
|
||||
.createHash("sha256")
|
||||
.update(canonicalRequest, "utf8")
|
||||
.digest("hex");
|
||||
|
||||
// 2. Create string to sign
|
||||
const algorithm = "HMAC-SHA256";
|
||||
const timestamp = headers["X-Date"] || headers["x-date"];
|
||||
if (!timestamp) {
|
||||
throw new Error("Jimeng signature generation requires an X-Date header.");
|
||||
}
|
||||
const dateStamp = timestamp.slice(0, 8);
|
||||
|
||||
const credentialScope = `${dateStamp}/${region}/${service}/request`;
|
||||
|
||||
const stringToSign = [
|
||||
algorithm,
|
||||
timestamp,
|
||||
credentialScope,
|
||||
hashedCanonicalRequest,
|
||||
].join("\n");
|
||||
|
||||
// 3. Calculate signature
|
||||
const kDate = crypto
|
||||
.createHmac("sha256", secretKey)
|
||||
.update(dateStamp)
|
||||
.digest();
|
||||
|
||||
const kRegion = crypto.createHmac("sha256", kDate).update(region).digest();
|
||||
const kService = crypto.createHmac("sha256", kRegion).update(service).digest();
|
||||
const kSigning = crypto.createHmac("sha256", kService).update("request").digest();
|
||||
|
||||
const signature = crypto
|
||||
.createHmac("sha256", kSigning)
|
||||
.update(stringToSign)
|
||||
.digest("hex");
|
||||
|
||||
// 4. Create authorization header
|
||||
return `${algorithm} Credential=${accessKey}/${credentialScope}, SignedHeaders=${signedHeaders}, Signature=${signature}`;
|
||||
}
|
||||
|
||||
/**
|
||||
* Parse aspect ratio string like "16:9", "1:1", "4:3" into width and height
|
||||
*/
|
||||
function parseAspectRatio(ar: string): { width: number; height: number } | null {
|
||||
const match = ar.match(/^(\d+(?:\.\d+)?):(\d+(?:\.\d+)?)$/);
|
||||
if (!match) return null;
|
||||
const w = parseFloat(match[1]!);
|
||||
const h = parseFloat(match[2]!);
|
||||
if (w <= 0 || h <= 0) return null;
|
||||
return { width: w, height: h };
|
||||
}
|
||||
|
||||
/**
|
||||
* Supported size presets for different quality levels
|
||||
* Based on Volcengine Jimeng documentation
|
||||
*/
|
||||
const SIZE_PRESETS: Record<string, Record<string, string>> = {
|
||||
normal: {
|
||||
"1:1": "1024x1024",
|
||||
"4:3": "1360x1020",
|
||||
"16:9": "1536x864",
|
||||
"3:2": "1440x960",
|
||||
"21:9": "1920x824",
|
||||
},
|
||||
"2k": {
|
||||
"1:1": "2048x2048",
|
||||
"4:3": "2304x1728",
|
||||
"16:9": "2560x1440",
|
||||
"3:2": "2496x1664",
|
||||
"21:9": "3024x1296",
|
||||
},
|
||||
"4k": {
|
||||
"1:1": "4096x4096",
|
||||
"4:3": "4694x3520",
|
||||
"16:9": "5404x3040",
|
||||
"3:2": "4992x3328",
|
||||
"21:9": "6198x2656",
|
||||
},
|
||||
};
|
||||
|
||||
function normalizeDimensions(value: string): string | null {
|
||||
const match = value.trim().match(/^(\d+)\s*[xX*]\s*(\d+)$/);
|
||||
if (!match) return null;
|
||||
return `${match[1]}x${match[2]}`;
|
||||
}
|
||||
|
||||
function getClosestPresetSize(ar: string | null, qualityLevel: JimengSizePreset): string {
|
||||
const presets = SIZE_PRESETS[qualityLevel];
|
||||
const defaultSize = presets["1:1"]!;
|
||||
|
||||
if (!ar) return defaultSize;
|
||||
|
||||
const parsed = parseAspectRatio(ar);
|
||||
if (!parsed) return defaultSize;
|
||||
|
||||
const targetRatio = parsed.width / parsed.height;
|
||||
let bestMatch = defaultSize;
|
||||
let bestDiff = Infinity;
|
||||
|
||||
for (const [ratio, size] of Object.entries(presets)) {
|
||||
const [w, h] = ratio.split(":").map(Number);
|
||||
const presetRatio = w / h;
|
||||
const diff = Math.abs(presetRatio - targetRatio);
|
||||
if (diff < bestDiff) {
|
||||
bestDiff = diff;
|
||||
bestMatch = size;
|
||||
}
|
||||
}
|
||||
|
||||
return bestMatch;
|
||||
}
|
||||
|
||||
function normalizeImageSizePreset(imageSize: string, ar: string | null): string | null {
|
||||
const preset = imageSize.trim().toUpperCase();
|
||||
if (preset === "1K") return getClosestPresetSize(ar, "normal");
|
||||
if (preset === "2K") return getClosestPresetSize(ar, "2k");
|
||||
if (preset === "4K") return getClosestPresetSize(ar, "4k");
|
||||
return normalizeDimensions(imageSize);
|
||||
}
|
||||
|
||||
function getImageSize(ar: string | null, quality: CliArgs["quality"], imageSize?: string | null): string {
|
||||
if (imageSize) {
|
||||
const normalizedSize = normalizeImageSizePreset(imageSize, ar);
|
||||
if (normalizedSize) return normalizedSize;
|
||||
}
|
||||
|
||||
// Default to 2K quality if not specified
|
||||
const qualityLevel: JimengSizePreset = quality === "normal" ? "normal" : "2k";
|
||||
return getClosestPresetSize(ar, qualityLevel);
|
||||
}
|
||||
|
||||
/**
|
||||
* Step 1: Submit async task to Volcengine Jimeng API
|
||||
*/
|
||||
async function submitTask(
|
||||
prompt: string,
|
||||
model: string,
|
||||
size: string,
|
||||
accessKey: string,
|
||||
secretKey: string,
|
||||
region: string
|
||||
): Promise<string> {
|
||||
// Query parameters for submit endpoint
|
||||
const query = {
|
||||
Action: "CVSync2AsyncSubmitTask",
|
||||
Version: "2022-08-31",
|
||||
};
|
||||
const endpoint = resolveEndpoint(query);
|
||||
|
||||
// Request body - Jimeng API expects width/height as separate integers
|
||||
const [width, height] = size.split("x").map(Number);
|
||||
const bodyObj = {
|
||||
req_key: model,
|
||||
prompt,
|
||||
// Use separate width and height parameters instead of size string
|
||||
width: width,
|
||||
height: height,
|
||||
// Optional: seed for reproducibility
|
||||
// seed: Math.floor(Math.random() * 999999),
|
||||
};
|
||||
|
||||
const body = JSON.stringify(bodyObj);
|
||||
|
||||
// Headers
|
||||
const timestampHeader = new Date().toISOString().replace(/[:\-]|\.\d{3}/g, "");
|
||||
const headers = {
|
||||
"Content-Type": "application/json",
|
||||
"X-Date": timestampHeader,
|
||||
"Host": endpoint.host,
|
||||
};
|
||||
|
||||
// Generate signature
|
||||
const authorization = generateSignature(
|
||||
"POST",
|
||||
query,
|
||||
headers,
|
||||
body,
|
||||
accessKey,
|
||||
secretKey,
|
||||
region,
|
||||
"cv",
|
||||
endpoint.canonicalUri
|
||||
);
|
||||
|
||||
console.error(`Submitting task to Jimeng (${model})...`, { width, height });
|
||||
|
||||
const res = await fetch(endpoint.url, {
|
||||
method: "POST",
|
||||
headers: {
|
||||
...headers,
|
||||
"Authorization": authorization,
|
||||
},
|
||||
body,
|
||||
});
|
||||
|
||||
if (!res.ok) {
|
||||
const err = await res.text();
|
||||
throw new Error(`Jimeng API submit error (${res.status}): ${err}`);
|
||||
}
|
||||
|
||||
const result = (await res.json()) as {
|
||||
code?: number;
|
||||
message?: string;
|
||||
data?: {
|
||||
task_id?: string;
|
||||
};
|
||||
};
|
||||
|
||||
// Volcengine API returns code 10000 for success
|
||||
if (result.code !== 10000 || !result.data?.task_id) {
|
||||
console.error("Submit response:", JSON.stringify(result, null, 2));
|
||||
throw new Error(`Failed to submit task: ${result.message || "Unknown error"}`);
|
||||
}
|
||||
|
||||
return result.data.task_id;
|
||||
}
|
||||
|
||||
/**
|
||||
* Step 2: Poll for task result
|
||||
* Returns image data directly as Uint8Array
|
||||
*/
|
||||
async function pollForResult(
|
||||
taskId: string,
|
||||
model: string,
|
||||
accessKey: string,
|
||||
secretKey: string,
|
||||
region: string
|
||||
): Promise<Uint8Array> {
|
||||
const maxAttempts = 60;
|
||||
const pollIntervalMs = 2000;
|
||||
|
||||
for (let attempt = 0; attempt < maxAttempts; attempt++) {
|
||||
// Query parameters for result endpoint
|
||||
const query = {
|
||||
Action: "CVSync2AsyncGetResult",
|
||||
Version: "2022-08-31",
|
||||
};
|
||||
const endpoint = resolveEndpoint(query);
|
||||
|
||||
// Request body - include req_key and task_id
|
||||
const bodyObj = {
|
||||
req_key: model,
|
||||
task_id: taskId,
|
||||
};
|
||||
|
||||
const body = JSON.stringify(bodyObj);
|
||||
|
||||
// Headers
|
||||
const timestampHeader = new Date().toISOString().replace(/[:\-]|\.\d{3}/g, "");
|
||||
const headers = {
|
||||
"Content-Type": "application/json",
|
||||
"X-Date": timestampHeader,
|
||||
"Host": endpoint.host,
|
||||
};
|
||||
|
||||
// Generate signature
|
||||
const authorization = generateSignature(
|
||||
"POST",
|
||||
query,
|
||||
headers,
|
||||
body,
|
||||
accessKey,
|
||||
secretKey,
|
||||
region,
|
||||
"cv",
|
||||
endpoint.canonicalUri
|
||||
);
|
||||
|
||||
const res = await fetch(endpoint.url, {
|
||||
method: "POST",
|
||||
headers: {
|
||||
...headers,
|
||||
"Authorization": authorization,
|
||||
},
|
||||
body,
|
||||
});
|
||||
|
||||
if (!res.ok) {
|
||||
const err = await res.text();
|
||||
throw new Error(`Jimeng API poll error (${res.status}): ${err}`);
|
||||
}
|
||||
|
||||
const result = (await res.json()) as {
|
||||
code?: number;
|
||||
message?: string;
|
||||
data?: {
|
||||
status?: string;
|
||||
image_urls?: string[];
|
||||
binary_data_base64?: string[];
|
||||
};
|
||||
};
|
||||
|
||||
// Volcengine API returns code 10000 for success
|
||||
if (result.code === 10000 && result.data) {
|
||||
const { status, image_urls, binary_data_base64 } = result.data;
|
||||
|
||||
// Check for base64 image data (preferred by Jimeng)
|
||||
if (binary_data_base64 && binary_data_base64.length > 0) {
|
||||
console.error("Image received as base64 data");
|
||||
const base64Data = binary_data_base64[0]!;
|
||||
// Convert base64 to Uint8Array
|
||||
const binaryString = Buffer.from(base64Data, "base64").toString("binary");
|
||||
const bytes = new Uint8Array(binaryString.length);
|
||||
for (let i = 0; i < binaryString.length; i++) {
|
||||
bytes[i] = binaryString.charCodeAt(i);
|
||||
}
|
||||
return bytes;
|
||||
}
|
||||
|
||||
// Fallback to URL format
|
||||
if (status === "done" && image_urls && image_urls.length > 0) {
|
||||
// Download from URL
|
||||
console.error(`Downloading image from ${image_urls[0]}...`);
|
||||
const imgRes = await fetch(image_urls[0]!);
|
||||
if (!imgRes.ok) {
|
||||
throw new Error(`Failed to download image from ${image_urls[0]}`);
|
||||
}
|
||||
const buffer = await imgRes.arrayBuffer();
|
||||
return new Uint8Array(buffer);
|
||||
}
|
||||
|
||||
if (status === "in_queue" || status === "generating") {
|
||||
console.error(`Task status: ${status} (${attempt + 1}/${maxAttempts})`);
|
||||
await new Promise(resolve => setTimeout(resolve, pollIntervalMs));
|
||||
continue;
|
||||
}
|
||||
|
||||
if (status === "fail") {
|
||||
throw new Error(`Jimeng task failed: ${result.message || "Generation failed"}`);
|
||||
}
|
||||
}
|
||||
|
||||
console.error("Poll response:", JSON.stringify(result, null, 2));
|
||||
throw new Error(`Unexpected response during polling: ${result.message || "Unknown error"}`);
|
||||
}
|
||||
|
||||
throw new Error("Task timeout: image generation took too long");
|
||||
}
|
||||
|
||||
export async function generateImage(
|
||||
prompt: string,
|
||||
model: string,
|
||||
args: CliArgs
|
||||
): Promise<Uint8Array> {
|
||||
if (args.referenceImages.length > 0) {
|
||||
throw new Error(
|
||||
"Jimeng does not support reference images. Use --provider google, openai, openrouter, or replicate."
|
||||
);
|
||||
}
|
||||
|
||||
const accessKey = getAccessKey();
|
||||
const secretKey = getSecretKey();
|
||||
const region = getRegion();
|
||||
|
||||
if (!accessKey || !secretKey) {
|
||||
throw new Error(
|
||||
"JIMENG_ACCESS_KEY_ID and JIMENG_SECRET_ACCESS_KEY are required. " +
|
||||
"Get your credentials from https://console.volcengine.com/iam/keymanage"
|
||||
);
|
||||
}
|
||||
|
||||
const size = getImageSize(args.aspectRatio, args.quality, args.imageSize);
|
||||
|
||||
// Step 1: Submit task
|
||||
const taskId = await submitTask(prompt, model, size, accessKey, secretKey, region);
|
||||
|
||||
// Step 2: Poll for result (returns image data directly)
|
||||
const imageData = await pollForResult(taskId, model, accessKey, secretKey, region);
|
||||
|
||||
console.error("Image generation complete!");
|
||||
return imageData;
|
||||
}
|
||||
@@ -0,0 +1,175 @@
|
||||
import assert from "node:assert/strict";
|
||||
import fs from "node:fs/promises";
|
||||
import os from "node:os";
|
||||
import path from "node:path";
|
||||
import test, { type TestContext } from "node:test";
|
||||
|
||||
import type { CliArgs } from "../types.ts";
|
||||
import {
|
||||
buildMinimaxUrl,
|
||||
buildRequestBody,
|
||||
buildSubjectReference,
|
||||
extractImageFromResponse,
|
||||
parsePixelSize,
|
||||
validateArgs,
|
||||
} from "./minimax.ts";
|
||||
|
||||
function useEnv(
|
||||
t: TestContext,
|
||||
values: Record<string, string | null>,
|
||||
): void {
|
||||
const previous = new Map<string, string | undefined>();
|
||||
for (const [key, value] of Object.entries(values)) {
|
||||
previous.set(key, process.env[key]);
|
||||
if (value == null) {
|
||||
delete process.env[key];
|
||||
} else {
|
||||
process.env[key] = value;
|
||||
}
|
||||
}
|
||||
|
||||
t.after(() => {
|
||||
for (const [key, value] of previous.entries()) {
|
||||
if (value == null) {
|
||||
delete process.env[key];
|
||||
} else {
|
||||
process.env[key] = value;
|
||||
}
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
function makeArgs(overrides: Partial<CliArgs> = {}): CliArgs {
|
||||
return {
|
||||
prompt: null,
|
||||
promptFiles: [],
|
||||
imagePath: null,
|
||||
provider: null,
|
||||
model: null,
|
||||
aspectRatio: null,
|
||||
size: null,
|
||||
quality: null,
|
||||
imageSize: null,
|
||||
imageApiDialect: null,
|
||||
referenceImages: [],
|
||||
n: 1,
|
||||
batchFile: null,
|
||||
jobs: null,
|
||||
json: false,
|
||||
help: false,
|
||||
...overrides,
|
||||
};
|
||||
}
|
||||
|
||||
test("MiniMax URL builder uses documented default and normalizes /v1 suffixes", (t) => {
|
||||
useEnv(t, { MINIMAX_BASE_URL: null });
|
||||
assert.equal(buildMinimaxUrl(), "https://api.minimaxi.com/v1/image_generation");
|
||||
|
||||
process.env.MINIMAX_BASE_URL = "https://api.minimax.io";
|
||||
assert.equal(buildMinimaxUrl(), "https://api.minimax.io/v1/image_generation");
|
||||
|
||||
process.env.MINIMAX_BASE_URL = "https://proxy.example.com/custom/v1/";
|
||||
assert.equal(buildMinimaxUrl(), "https://proxy.example.com/custom/v1/image_generation");
|
||||
});
|
||||
|
||||
test("MiniMax size parsing and validation follow documented constraints", () => {
|
||||
assert.deepEqual(parsePixelSize("1536x1024"), { width: 1536, height: 1024 });
|
||||
assert.deepEqual(parsePixelSize("1536*1024"), { width: 1536, height: 1024 });
|
||||
assert.equal(parsePixelSize("wide"), null);
|
||||
|
||||
validateArgs("image-01", makeArgs({ size: "1536x1024", n: 9 }));
|
||||
|
||||
assert.throws(
|
||||
() => validateArgs("image-01-live", makeArgs({ size: "1536x1024" })),
|
||||
/only supported with model image-01/,
|
||||
);
|
||||
assert.throws(
|
||||
() => validateArgs("image-01", makeArgs({ size: "1537x1024" })),
|
||||
/divisible by 8/,
|
||||
);
|
||||
assert.throws(
|
||||
() => validateArgs("image-01", makeArgs({ aspectRatio: "2.35:1" })),
|
||||
/aspect_ratio must be one of/,
|
||||
);
|
||||
assert.throws(
|
||||
() => validateArgs("image-01", makeArgs({ n: 10 })),
|
||||
/at most 9 images/,
|
||||
);
|
||||
});
|
||||
|
||||
test("MiniMax request body maps aspect ratio, size, n, and subject references", async (t) => {
|
||||
const dir = await fs.mkdtemp(path.join(os.tmpdir(), "minimax-test-"));
|
||||
t.after(() => fs.rm(dir, { recursive: true, force: true }));
|
||||
|
||||
const refPath = path.join(dir, "portrait.png");
|
||||
await fs.writeFile(refPath, Buffer.from("portrait"));
|
||||
|
||||
const ratioBody = await buildRequestBody(
|
||||
"A portrait by the window",
|
||||
"image-01",
|
||||
makeArgs({ aspectRatio: "16:9", n: 2, referenceImages: [refPath] }),
|
||||
);
|
||||
assert.equal(ratioBody.aspect_ratio, "16:9");
|
||||
assert.equal(ratioBody.n, 2);
|
||||
assert.equal(ratioBody.response_format, "base64");
|
||||
assert.match(ratioBody.subject_reference?.[0]?.image_file || "", /^data:image\/png;base64,/);
|
||||
|
||||
const sizeBody = await buildRequestBody(
|
||||
"A portrait by the window",
|
||||
"image-01",
|
||||
makeArgs({ size: "1536x1024" }),
|
||||
);
|
||||
assert.equal(sizeBody.width, 1536);
|
||||
assert.equal(sizeBody.height, 1024);
|
||||
assert.equal(sizeBody.aspect_ratio, undefined);
|
||||
});
|
||||
|
||||
test("MiniMax subject references require supported file types", async (t) => {
|
||||
const dir = await fs.mkdtemp(path.join(os.tmpdir(), "minimax-ref-"));
|
||||
t.after(() => fs.rm(dir, { recursive: true, force: true }));
|
||||
|
||||
const good = path.join(dir, "portrait.jpg");
|
||||
const bad = path.join(dir, "portrait.webp");
|
||||
await fs.writeFile(good, Buffer.from("portrait"));
|
||||
await fs.writeFile(bad, Buffer.from("portrait"));
|
||||
|
||||
const subjectReference = await buildSubjectReference([good]);
|
||||
assert.equal(subjectReference?.[0]?.type, "character");
|
||||
|
||||
await assert.rejects(
|
||||
() => buildSubjectReference([bad]),
|
||||
/only supports JPG, JPEG, or PNG/,
|
||||
);
|
||||
});
|
||||
|
||||
test("MiniMax response extraction supports base64 and URL payloads", async (t) => {
|
||||
const originalFetch = globalThis.fetch;
|
||||
t.after(() => {
|
||||
globalThis.fetch = originalFetch;
|
||||
});
|
||||
|
||||
const fromBase64 = await extractImageFromResponse({
|
||||
data: {
|
||||
image_base64: [Buffer.from("hello").toString("base64")],
|
||||
},
|
||||
});
|
||||
assert.equal(Buffer.from(fromBase64).toString("utf8"), "hello");
|
||||
|
||||
globalThis.fetch = async () =>
|
||||
new Response(Uint8Array.from([1, 2, 3]), {
|
||||
status: 200,
|
||||
headers: { "Content-Type": "image/jpeg" },
|
||||
});
|
||||
|
||||
const fromUrl = await extractImageFromResponse({
|
||||
data: {
|
||||
image_urls: ["https://example.com/output.jpg"],
|
||||
},
|
||||
});
|
||||
assert.deepEqual([...fromUrl], [1, 2, 3]);
|
||||
|
||||
await assert.rejects(
|
||||
() => extractImageFromResponse({ base_resp: { status_code: 1001, status_msg: "blocked" } }),
|
||||
/blocked/,
|
||||
);
|
||||
});
|
||||
220
baoyu-skills/skills/baoyu-image-gen/scripts/providers/minimax.ts
Normal file
220
baoyu-skills/skills/baoyu-image-gen/scripts/providers/minimax.ts
Normal file
@@ -0,0 +1,220 @@
|
||||
import path from "node:path";
|
||||
import { readFile } from "node:fs/promises";
|
||||
|
||||
import type { CliArgs } from "../types";
|
||||
|
||||
const DEFAULT_MODEL = "image-01";
|
||||
const MAX_REFERENCE_IMAGE_BYTES = 10 * 1024 * 1024;
|
||||
const SUPPORTED_ASPECT_RATIOS = new Set(["1:1", "16:9", "4:3", "3:2", "2:3", "3:4", "9:16", "21:9"]);
|
||||
|
||||
type MinimaxSubjectReference = {
|
||||
type: "character";
|
||||
image_file: string;
|
||||
};
|
||||
|
||||
type MinimaxRequestBody = {
|
||||
model: string;
|
||||
prompt: string;
|
||||
response_format: "base64";
|
||||
aspect_ratio?: string;
|
||||
width?: number;
|
||||
height?: number;
|
||||
n?: number;
|
||||
subject_reference?: MinimaxSubjectReference[];
|
||||
};
|
||||
|
||||
type MinimaxResponse = {
|
||||
id?: string;
|
||||
data?: {
|
||||
image_urls?: string[];
|
||||
image_base64?: string[];
|
||||
};
|
||||
base_resp?: {
|
||||
status_code?: number;
|
||||
status_msg?: string;
|
||||
};
|
||||
};
|
||||
|
||||
export function getDefaultModel(): string {
|
||||
return process.env.MINIMAX_IMAGE_MODEL || DEFAULT_MODEL;
|
||||
}
|
||||
|
||||
function getApiKey(): string | null {
|
||||
return process.env.MINIMAX_API_KEY || null;
|
||||
}
|
||||
|
||||
export function buildMinimaxUrl(): string {
|
||||
const base = (process.env.MINIMAX_BASE_URL || "https://api.minimaxi.com").replace(/\/+$/g, "");
|
||||
return base.endsWith("/v1") ? `${base}/image_generation` : `${base}/v1/image_generation`;
|
||||
}
|
||||
|
||||
function getMimeType(filename: string): "image/jpeg" | "image/png" {
|
||||
const ext = path.extname(filename).toLowerCase();
|
||||
if (ext === ".jpg" || ext === ".jpeg") return "image/jpeg";
|
||||
if (ext === ".png") return "image/png";
|
||||
throw new Error(
|
||||
`MiniMax subject_reference only supports JPG, JPEG, or PNG files: ${filename}`
|
||||
);
|
||||
}
|
||||
|
||||
export function parsePixelSize(size: string): { width: number; height: number } | null {
|
||||
const match = size.trim().match(/^(\d+)\s*[xX*]\s*(\d+)$/);
|
||||
if (!match) return null;
|
||||
|
||||
const width = parseInt(match[1]!, 10);
|
||||
const height = parseInt(match[2]!, 10);
|
||||
if (!Number.isFinite(width) || !Number.isFinite(height) || width <= 0 || height <= 0) {
|
||||
return null;
|
||||
}
|
||||
|
||||
return { width, height };
|
||||
}
|
||||
|
||||
function validatePixelSize(width: number, height: number): void {
|
||||
if (width < 512 || width > 2048 || height < 512 || height > 2048) {
|
||||
throw new Error("MiniMax custom size must keep width and height between 512 and 2048.");
|
||||
}
|
||||
if (width % 8 !== 0 || height % 8 !== 0) {
|
||||
throw new Error("MiniMax custom size requires width and height divisible by 8.");
|
||||
}
|
||||
}
|
||||
|
||||
export function validateArgs(model: string, args: CliArgs): void {
|
||||
if (args.n > 9) {
|
||||
throw new Error("MiniMax supports at most 9 images per request.");
|
||||
}
|
||||
|
||||
if (args.aspectRatio && !SUPPORTED_ASPECT_RATIOS.has(args.aspectRatio)) {
|
||||
throw new Error(
|
||||
`MiniMax aspect_ratio must be one of: ${Array.from(SUPPORTED_ASPECT_RATIOS).join(", ")}.`
|
||||
);
|
||||
}
|
||||
|
||||
if (args.size && !args.aspectRatio) {
|
||||
if (model !== "image-01") {
|
||||
throw new Error("MiniMax custom --size is only supported with model image-01. Use --model image-01 or pass --ar instead.");
|
||||
}
|
||||
const parsed = parsePixelSize(args.size);
|
||||
if (!parsed) {
|
||||
throw new Error("MiniMax --size must be in WxH format, for example 1536x1024.");
|
||||
}
|
||||
validatePixelSize(parsed.width, parsed.height);
|
||||
}
|
||||
}
|
||||
|
||||
export async function buildSubjectReference(
|
||||
referenceImages: string[],
|
||||
): Promise<MinimaxSubjectReference[] | undefined> {
|
||||
if (referenceImages.length === 0) return undefined;
|
||||
|
||||
const subjectReference: MinimaxSubjectReference[] = [];
|
||||
for (const refPath of referenceImages) {
|
||||
const bytes = await readFile(refPath);
|
||||
if (bytes.length > MAX_REFERENCE_IMAGE_BYTES) {
|
||||
throw new Error(`MiniMax subject_reference images must be smaller than 10MB: ${refPath}`);
|
||||
}
|
||||
|
||||
subjectReference.push({
|
||||
type: "character",
|
||||
image_file: `data:${getMimeType(refPath)};base64,${bytes.toString("base64")}`,
|
||||
});
|
||||
}
|
||||
|
||||
return subjectReference;
|
||||
}
|
||||
|
||||
export async function buildRequestBody(
|
||||
prompt: string,
|
||||
model: string,
|
||||
args: CliArgs,
|
||||
): Promise<MinimaxRequestBody> {
|
||||
validateArgs(model, args);
|
||||
|
||||
const body: MinimaxRequestBody = {
|
||||
model,
|
||||
prompt,
|
||||
response_format: "base64",
|
||||
};
|
||||
|
||||
if (args.aspectRatio) {
|
||||
body.aspect_ratio = args.aspectRatio;
|
||||
} else if (args.size) {
|
||||
const parsed = parsePixelSize(args.size);
|
||||
if (!parsed) {
|
||||
throw new Error("MiniMax --size must be in WxH format, for example 1536x1024.");
|
||||
}
|
||||
body.width = parsed.width;
|
||||
body.height = parsed.height;
|
||||
}
|
||||
|
||||
if (args.n > 1) {
|
||||
body.n = args.n;
|
||||
}
|
||||
|
||||
const subjectReference = await buildSubjectReference(args.referenceImages);
|
||||
if (subjectReference) {
|
||||
body.subject_reference = subjectReference;
|
||||
}
|
||||
|
||||
return body;
|
||||
}
|
||||
|
||||
async function downloadImage(url: string): Promise<Uint8Array> {
|
||||
const response = await fetch(url);
|
||||
if (!response.ok) {
|
||||
throw new Error(`Failed to download image from MiniMax: ${response.status}`);
|
||||
}
|
||||
return new Uint8Array(await response.arrayBuffer());
|
||||
}
|
||||
|
||||
export async function extractImageFromResponse(result: MinimaxResponse): Promise<Uint8Array> {
|
||||
const baseResp = result.base_resp;
|
||||
if (baseResp && baseResp.status_code !== undefined && baseResp.status_code !== 0) {
|
||||
throw new Error(baseResp.status_msg || `MiniMax API returned status_code=${baseResp.status_code}`);
|
||||
}
|
||||
|
||||
const base64Image = result.data?.image_base64?.[0];
|
||||
if (base64Image) {
|
||||
return Uint8Array.from(Buffer.from(base64Image, "base64"));
|
||||
}
|
||||
|
||||
const url = result.data?.image_urls?.[0];
|
||||
if (url) {
|
||||
return downloadImage(url);
|
||||
}
|
||||
|
||||
throw new Error("No image data in MiniMax response");
|
||||
}
|
||||
|
||||
export function getDefaultOutputExtension(): ".jpg" {
|
||||
return ".jpg";
|
||||
}
|
||||
|
||||
export async function generateImage(
|
||||
prompt: string,
|
||||
model: string,
|
||||
args: CliArgs
|
||||
): Promise<Uint8Array> {
|
||||
const apiKey = getApiKey();
|
||||
if (!apiKey) {
|
||||
throw new Error("MINIMAX_API_KEY is required. Get one from https://platform.minimaxi.com/");
|
||||
}
|
||||
|
||||
const body = await buildRequestBody(prompt, model, args);
|
||||
const response = await fetch(buildMinimaxUrl(), {
|
||||
method: "POST",
|
||||
headers: {
|
||||
"Content-Type": "application/json",
|
||||
Authorization: `Bearer ${apiKey}`,
|
||||
},
|
||||
body: JSON.stringify(body),
|
||||
});
|
||||
|
||||
if (!response.ok) {
|
||||
const err = await response.text();
|
||||
throw new Error(`MiniMax API error (${response.status}): ${err}`);
|
||||
}
|
||||
|
||||
const result = (await response.json()) as MinimaxResponse;
|
||||
return extractImageFromResponse(result);
|
||||
}
|
||||
@@ -0,0 +1,192 @@
|
||||
import assert from "node:assert/strict";
|
||||
import test from "node:test";
|
||||
|
||||
import type { CliArgs } from "../types.ts";
|
||||
import {
|
||||
buildOpenAIGenerationsBody,
|
||||
extractImageFromResponse,
|
||||
getDefaultModel,
|
||||
getOpenAIAspectRatio,
|
||||
getOpenAIImageApiDialect,
|
||||
getOpenAIResolution,
|
||||
getMimeType,
|
||||
getOpenAISize,
|
||||
getOrientationFromAspectRatio,
|
||||
inferAspectRatioFromSize,
|
||||
inferResolutionFromSize,
|
||||
parseAspectRatio,
|
||||
validateArgs,
|
||||
} from "./openai.ts";
|
||||
|
||||
function makeArgs(overrides: Partial<CliArgs> = {}): CliArgs {
|
||||
return {
|
||||
prompt: null,
|
||||
promptFiles: [],
|
||||
imagePath: null,
|
||||
provider: null,
|
||||
model: null,
|
||||
aspectRatio: null,
|
||||
size: null,
|
||||
quality: "2k",
|
||||
imageSize: null,
|
||||
imageApiDialect: null,
|
||||
referenceImages: [],
|
||||
n: 1,
|
||||
batchFile: null,
|
||||
jobs: null,
|
||||
json: false,
|
||||
help: false,
|
||||
...overrides,
|
||||
};
|
||||
}
|
||||
|
||||
test("OpenAI aspect-ratio parsing and size selection match model families", () => {
|
||||
assert.equal(getDefaultModel(), "gpt-image-2.5-flare");
|
||||
assert.deepEqual(parseAspectRatio("16:9"), { width: 16, height: 9 });
|
||||
assert.equal(parseAspectRatio("wide"), null);
|
||||
assert.equal(parseAspectRatio("0:1"), null);
|
||||
|
||||
assert.equal(getOpenAISize("dall-e-3", "16:9", "2k"), "1792x1024");
|
||||
assert.equal(getOpenAISize("dall-e-3", "9:16", "normal"), "1024x1792");
|
||||
assert.equal(getOpenAISize("dall-e-2", "16:9", "2k"), "1024x1024");
|
||||
assert.equal(getOpenAISize("gpt-image-1.5", "16:9", "2k"), "1536x1024");
|
||||
assert.equal(getOpenAISize("gpt-image-1.5", "4:3", "2k"), "1024x1024");
|
||||
assert.equal(getOpenAISize("gpt-image-2", "16:9", "2k"), "2048x1152");
|
||||
assert.equal(getOpenAISize("gpt-image-2", "9:16", "2k"), "1152x2048");
|
||||
assert.equal(getOpenAISize("gpt-image-2", "4:3", "2k"), "2048x1536");
|
||||
assert.equal(getOpenAISize("gpt-image-2", "2.35:1", "normal"), "1248x528");
|
||||
assert.equal(getOpenAISize("gpt-image-2.5-flare", "16:9", "2k"), "2048x1152");
|
||||
assert.equal(getOpenAISize("gpt-image-2.5-sunburst", "9:16", "normal"), "608x1088");
|
||||
assert.equal(getOpenAISize("gpt-image-2.5-flare-2026-09-08", "1:1", "2k"), "2048x2048");
|
||||
assert.equal(inferAspectRatioFromSize("1536x1024"), "3:2");
|
||||
assert.equal(inferResolutionFromSize("1536x1024"), "2K");
|
||||
assert.equal(getOpenAIAspectRatio({ aspectRatio: null, size: "2048x1152" }), "16:9");
|
||||
assert.equal(getOpenAIResolution({ imageSize: null, size: "2048x1152", quality: "normal" }), "2K");
|
||||
assert.equal(getOrientationFromAspectRatio("16:9"), "landscape");
|
||||
assert.equal(getOrientationFromAspectRatio("9:16"), "portrait");
|
||||
assert.equal(getOrientationFromAspectRatio("1:1"), null);
|
||||
assert.equal(getOpenAIImageApiDialect({ imageApiDialect: null }), "openai-native");
|
||||
});
|
||||
|
||||
test("OpenAI generations body switches between native and ratio-metadata dialects", () => {
|
||||
assert.deepEqual(
|
||||
buildOpenAIGenerationsBody("Draw a skyline", "gpt-image-2", {
|
||||
aspectRatio: "16:9",
|
||||
size: null,
|
||||
quality: "2k",
|
||||
imageSize: null,
|
||||
imageApiDialect: null,
|
||||
}),
|
||||
{
|
||||
model: "gpt-image-2",
|
||||
prompt: "Draw a skyline",
|
||||
size: "2048x1152",
|
||||
quality: "high",
|
||||
},
|
||||
);
|
||||
|
||||
assert.deepEqual(
|
||||
buildOpenAIGenerationsBody("Draw a skyline", "gemini-3-pro-image-preview", {
|
||||
aspectRatio: "16:9",
|
||||
size: null,
|
||||
quality: "2k",
|
||||
imageSize: null,
|
||||
imageApiDialect: "ratio-metadata",
|
||||
}),
|
||||
{
|
||||
model: "gemini-3-pro-image-preview",
|
||||
prompt: "Draw a skyline",
|
||||
size: "16:9",
|
||||
metadata: {
|
||||
resolution: "2K",
|
||||
orientation: "landscape",
|
||||
},
|
||||
},
|
||||
);
|
||||
|
||||
assert.deepEqual(
|
||||
buildOpenAIGenerationsBody("Draw a portrait", "gemini-3-pro-image-preview", {
|
||||
aspectRatio: null,
|
||||
size: "1152x2048",
|
||||
quality: "normal",
|
||||
imageSize: null,
|
||||
imageApiDialect: "ratio-metadata",
|
||||
}),
|
||||
{
|
||||
model: "gemini-3-pro-image-preview",
|
||||
prompt: "Draw a portrait",
|
||||
size: "9:16",
|
||||
metadata: {
|
||||
resolution: "2K",
|
||||
orientation: "portrait",
|
||||
},
|
||||
},
|
||||
);
|
||||
});
|
||||
|
||||
test("OpenAI validates GPT Image custom size constraints", () => {
|
||||
assert.doesNotThrow(() =>
|
||||
validateArgs("gpt-image-2", makeArgs({ size: "3840x2160" })),
|
||||
);
|
||||
assert.doesNotThrow(() =>
|
||||
validateArgs("gpt-image-2-2026-04-21", makeArgs({ aspectRatio: "2.35:1" })),
|
||||
);
|
||||
|
||||
assert.throws(
|
||||
() => validateArgs("gpt-image-2", makeArgs({ size: "1024x576" })),
|
||||
/total pixels/,
|
||||
);
|
||||
assert.throws(
|
||||
() => validateArgs("gpt-image-2", makeArgs({ size: "1025x1024" })),
|
||||
/multiples of 16px/,
|
||||
);
|
||||
assert.throws(
|
||||
() => validateArgs("gpt-image-2", makeArgs({ aspectRatio: "4:1" })),
|
||||
/must not exceed 3:1/,
|
||||
);
|
||||
assert.doesNotThrow(() =>
|
||||
validateArgs("gpt-image-2.5-flare", makeArgs({ size: "3840x2160" })),
|
||||
);
|
||||
assert.throws(
|
||||
() => validateArgs("gpt-image-2.5-sunburst", makeArgs({ aspectRatio: "4:1" })),
|
||||
/gpt-image-2\.5-sunburst aspect ratio must not exceed 3:1/,
|
||||
);
|
||||
assert.doesNotThrow(() =>
|
||||
validateArgs("gpt-image-1.5", makeArgs({ size: "1536x1024" })),
|
||||
);
|
||||
});
|
||||
|
||||
test("OpenAI mime-type detection covers supported reference image extensions", () => {
|
||||
assert.equal(getMimeType("frame.png"), "image/png");
|
||||
assert.equal(getMimeType("frame.jpg"), "image/jpeg");
|
||||
assert.equal(getMimeType("frame.webp"), "image/webp");
|
||||
assert.equal(getMimeType("frame.gif"), "image/gif");
|
||||
});
|
||||
|
||||
test("OpenAI response extraction supports base64 and URL download flows", async (t) => {
|
||||
const originalFetch = globalThis.fetch;
|
||||
t.after(() => {
|
||||
globalThis.fetch = originalFetch;
|
||||
});
|
||||
|
||||
const fromBase64 = await extractImageFromResponse({
|
||||
data: [{ b64_json: Buffer.from("hello").toString("base64") }],
|
||||
});
|
||||
assert.equal(Buffer.from(fromBase64).toString("utf8"), "hello");
|
||||
|
||||
globalThis.fetch = async () =>
|
||||
new Response(Uint8Array.from([1, 2, 3]), {
|
||||
status: 200,
|
||||
headers: { "Content-Type": "application/octet-stream" },
|
||||
});
|
||||
|
||||
const fromUrl = await extractImageFromResponse({
|
||||
data: [{ url: "https://example.com/image.png" }],
|
||||
});
|
||||
assert.deepEqual([...fromUrl], [1, 2, 3]);
|
||||
|
||||
await assert.rejects(
|
||||
() => extractImageFromResponse({ data: [{}] }),
|
||||
/No image in response/,
|
||||
);
|
||||
});
|
||||
444
baoyu-skills/skills/baoyu-image-gen/scripts/providers/openai.ts
Normal file
444
baoyu-skills/skills/baoyu-image-gen/scripts/providers/openai.ts
Normal file
@@ -0,0 +1,444 @@
|
||||
import path from "node:path";
|
||||
import { readFile } from "node:fs/promises";
|
||||
import type { CliArgs, OpenAIImageApiDialect } from "../types";
|
||||
|
||||
export function getDefaultModel(): string {
|
||||
return process.env.OPENAI_IMAGE_MODEL || "gpt-image-2.5-flare";
|
||||
}
|
||||
|
||||
type OpenAIImageResponse = { data: Array<{ url?: string; b64_json?: string }> };
|
||||
|
||||
export function parseAspectRatio(ar: string): { width: number; height: number } | null {
|
||||
const match = ar.match(/^(\d+(?:\.\d+)?):(\d+(?:\.\d+)?)$/);
|
||||
if (!match) return null;
|
||||
const w = parseFloat(match[1]!);
|
||||
const h = parseFloat(match[2]!);
|
||||
if (w <= 0 || h <= 0) return null;
|
||||
return { width: w, height: h };
|
||||
}
|
||||
|
||||
type SizeMapping = {
|
||||
square: string;
|
||||
landscape: string;
|
||||
portrait: string;
|
||||
};
|
||||
|
||||
type OpenAIGenerationsBody = Record<string, unknown>;
|
||||
|
||||
function isGptImageModel(model: string): boolean {
|
||||
return model.includes("gpt-image");
|
||||
}
|
||||
|
||||
function isGptImageCustomSizeModel(model: string): boolean {
|
||||
return model.includes("gpt-image-2");
|
||||
}
|
||||
|
||||
function roundToMultiple(value: number, multiple: number): number {
|
||||
return Math.max(multiple, Math.round(value / multiple) * multiple);
|
||||
}
|
||||
|
||||
function buildCustomSizeFromAspectRatio(
|
||||
ar: string | null,
|
||||
quality: CliArgs["quality"],
|
||||
): string {
|
||||
const parsed = ar ? parseAspectRatio(ar) : null;
|
||||
const ratio = parsed ? parsed.width / parsed.height : 1;
|
||||
|
||||
if (!parsed || Math.abs(ratio - 1) < 0.1) {
|
||||
const edge = quality === "2k" ? 2048 : 1024;
|
||||
return `${edge}x${edge}`;
|
||||
}
|
||||
|
||||
const targetLongEdge = quality === "2k" ? 2048 : 1024;
|
||||
let width: number;
|
||||
let height: number;
|
||||
|
||||
if (ratio > 1) {
|
||||
width = targetLongEdge;
|
||||
height = roundToMultiple(width / ratio, 16);
|
||||
} else {
|
||||
height = targetLongEdge;
|
||||
width = roundToMultiple(height * ratio, 16);
|
||||
}
|
||||
|
||||
while (width * height < 655_360) {
|
||||
if (ratio > 1) {
|
||||
width += 16;
|
||||
height = roundToMultiple(width / ratio, 16);
|
||||
} else {
|
||||
height += 16;
|
||||
width = roundToMultiple(height * ratio, 16);
|
||||
}
|
||||
}
|
||||
|
||||
return `${width}x${height}`;
|
||||
}
|
||||
|
||||
export function getOpenAISize(
|
||||
model: string,
|
||||
ar: string | null,
|
||||
quality: CliArgs["quality"]
|
||||
): string {
|
||||
const isDalle3 = model.includes("dall-e-3");
|
||||
const isDalle2 = model.includes("dall-e-2");
|
||||
|
||||
if (isDalle2) {
|
||||
return "1024x1024";
|
||||
}
|
||||
|
||||
if (isGptImageCustomSizeModel(model)) {
|
||||
return buildCustomSizeFromAspectRatio(ar, quality);
|
||||
}
|
||||
|
||||
const sizes: SizeMapping = isDalle3
|
||||
? {
|
||||
square: "1024x1024",
|
||||
landscape: "1792x1024",
|
||||
portrait: "1024x1792",
|
||||
}
|
||||
: {
|
||||
square: "1024x1024",
|
||||
landscape: "1536x1024",
|
||||
portrait: "1024x1536",
|
||||
};
|
||||
|
||||
if (!ar) return sizes.square;
|
||||
|
||||
const parsed = parseAspectRatio(ar);
|
||||
if (!parsed) return sizes.square;
|
||||
|
||||
const ratio = parsed.width / parsed.height;
|
||||
|
||||
if (Math.abs(ratio - 1) < 0.1) return sizes.square;
|
||||
if (ratio > 1.5) return sizes.landscape;
|
||||
if (ratio < 0.67) return sizes.portrait;
|
||||
return sizes.square;
|
||||
}
|
||||
|
||||
function parsePixelSize(value: string): { width: number; height: number } | null {
|
||||
const match = value.match(/^(\d+)\s*[xX]\s*(\d+)$/);
|
||||
if (!match) return null;
|
||||
|
||||
const width = parseInt(match[1]!, 10);
|
||||
const height = parseInt(match[2]!, 10);
|
||||
if (!Number.isFinite(width) || !Number.isFinite(height) || width <= 0 || height <= 0) {
|
||||
return null;
|
||||
}
|
||||
|
||||
return { width, height };
|
||||
}
|
||||
|
||||
function gcd(a: number, b: number): number {
|
||||
let x = Math.abs(a);
|
||||
let y = Math.abs(b);
|
||||
while (y !== 0) {
|
||||
const next = x % y;
|
||||
x = y;
|
||||
y = next;
|
||||
}
|
||||
return x || 1;
|
||||
}
|
||||
|
||||
export function getOpenAIImageApiDialect(args: Pick<CliArgs, "imageApiDialect">): OpenAIImageApiDialect {
|
||||
return args.imageApiDialect ?? "openai-native";
|
||||
}
|
||||
|
||||
export function inferAspectRatioFromSize(size: string | null): string | null {
|
||||
if (!size) return null;
|
||||
const parsed = parsePixelSize(size);
|
||||
if (!parsed) return null;
|
||||
|
||||
const divisor = gcd(parsed.width, parsed.height);
|
||||
return `${parsed.width / divisor}:${parsed.height / divisor}`;
|
||||
}
|
||||
|
||||
export function inferResolutionFromSize(size: string | null): "1K" | "2K" | "4K" | null {
|
||||
if (!size) return null;
|
||||
const parsed = parsePixelSize(size);
|
||||
if (!parsed) return null;
|
||||
|
||||
const longestEdge = Math.max(parsed.width, parsed.height);
|
||||
if (longestEdge <= 1024) return "1K";
|
||||
if (longestEdge <= 2048) return "2K";
|
||||
return "4K";
|
||||
}
|
||||
|
||||
export function getOpenAIAspectRatio(args: Pick<CliArgs, "aspectRatio" | "size">): string {
|
||||
return args.aspectRatio ?? inferAspectRatioFromSize(args.size) ?? "1:1";
|
||||
}
|
||||
|
||||
export function getOpenAIResolution(
|
||||
args: Pick<CliArgs, "imageSize" | "size" | "quality">
|
||||
): "1K" | "2K" | "4K" {
|
||||
if (args.imageSize === "1K" || args.imageSize === "2K" || args.imageSize === "4K") {
|
||||
return args.imageSize;
|
||||
}
|
||||
|
||||
const inferred = inferResolutionFromSize(args.size);
|
||||
if (inferred) return inferred;
|
||||
|
||||
return args.quality === "normal" ? "1K" : "2K";
|
||||
}
|
||||
|
||||
function getOpenAIQuality(model: string, quality: CliArgs["quality"]): "standard" | "hd" | "medium" | "high" | null {
|
||||
if (model.includes("dall-e-3")) {
|
||||
return quality === "2k" ? "hd" : "standard";
|
||||
}
|
||||
|
||||
if (isGptImageModel(model)) {
|
||||
return quality === "2k" ? "high" : "medium";
|
||||
}
|
||||
|
||||
return null;
|
||||
}
|
||||
|
||||
export function getOrientationFromAspectRatio(ar: string): "landscape" | "portrait" | null {
|
||||
const parsed = parseAspectRatio(ar);
|
||||
if (!parsed) return null;
|
||||
|
||||
const ratio = parsed.width / parsed.height;
|
||||
if (Math.abs(ratio - 1) < 0.1) return null;
|
||||
return ratio > 1 ? "landscape" : "portrait";
|
||||
}
|
||||
|
||||
export function buildOpenAIGenerationsBody(
|
||||
prompt: string,
|
||||
model: string,
|
||||
args: Pick<CliArgs, "aspectRatio" | "size" | "quality" | "imageSize" | "imageApiDialect">
|
||||
): OpenAIGenerationsBody {
|
||||
if (getOpenAIImageApiDialect(args) === "ratio-metadata") {
|
||||
const aspectRatio = getOpenAIAspectRatio(args);
|
||||
const metadata: Record<string, string> = {
|
||||
resolution: getOpenAIResolution(args),
|
||||
};
|
||||
const orientation = getOrientationFromAspectRatio(aspectRatio);
|
||||
if (orientation) metadata.orientation = orientation;
|
||||
|
||||
return {
|
||||
model,
|
||||
prompt,
|
||||
size: aspectRatio,
|
||||
metadata,
|
||||
};
|
||||
}
|
||||
|
||||
const body: OpenAIGenerationsBody = {
|
||||
model,
|
||||
prompt,
|
||||
size: args.size || getOpenAISize(model, args.aspectRatio, args.quality),
|
||||
};
|
||||
|
||||
const quality = getOpenAIQuality(model, args.quality);
|
||||
if (quality) {
|
||||
body.quality = quality;
|
||||
}
|
||||
|
||||
return body;
|
||||
}
|
||||
|
||||
export function validateArgs(model: string, args: CliArgs): void {
|
||||
if (!isGptImageCustomSizeModel(model)) return;
|
||||
|
||||
if (args.aspectRatio && !args.size) {
|
||||
const parsed = parseAspectRatio(args.aspectRatio);
|
||||
if (!parsed) {
|
||||
throw new Error(`Invalid ${model} aspect ratio: ${args.aspectRatio}`);
|
||||
}
|
||||
const ratio = parsed.width / parsed.height;
|
||||
if (Math.max(ratio, 1 / ratio) > 3) {
|
||||
throw new Error(`${model} aspect ratio must not exceed 3:1.`);
|
||||
}
|
||||
}
|
||||
|
||||
if (!args.size) return;
|
||||
|
||||
const parsedSize = parsePixelSize(args.size);
|
||||
if (!parsedSize) {
|
||||
throw new Error(`Invalid ${model} --size: ${args.size}. Expected <width>x<height>.`);
|
||||
}
|
||||
|
||||
const { width, height } = parsedSize;
|
||||
const totalPixels = width * height;
|
||||
const ratio = Math.max(width, height) / Math.min(width, height);
|
||||
|
||||
if (Math.max(width, height) > 3840) {
|
||||
throw new Error(`${model} --size maximum edge length must be 3840px or less.`);
|
||||
}
|
||||
if (width % 16 !== 0 || height % 16 !== 0) {
|
||||
throw new Error(`${model} --size width and height must both be multiples of 16px.`);
|
||||
}
|
||||
if (ratio > 3) {
|
||||
throw new Error(`${model} --size long edge to short edge ratio must not exceed 3:1.`);
|
||||
}
|
||||
if (totalPixels < 655_360 || totalPixels > 8_294_400) {
|
||||
throw new Error(`${model} --size total pixels must be between 655,360 and 8,294,400.`);
|
||||
}
|
||||
}
|
||||
|
||||
export async function generateImage(
|
||||
prompt: string,
|
||||
model: string,
|
||||
args: CliArgs
|
||||
): Promise<Uint8Array> {
|
||||
const baseURL = process.env.OPENAI_BASE_URL || "https://api.openai.com/v1";
|
||||
const apiKey = process.env.OPENAI_API_KEY;
|
||||
|
||||
if (!apiKey) {
|
||||
throw new Error(
|
||||
"OPENAI_API_KEY is required. Codex/ChatGPT desktop login does not automatically grant OpenAI Images API access to this script."
|
||||
);
|
||||
}
|
||||
|
||||
if (process.env.OPENAI_IMAGE_USE_CHAT === "true") {
|
||||
return generateWithChatCompletions(baseURL, apiKey, prompt, model);
|
||||
}
|
||||
|
||||
const imageApiDialect = getOpenAIImageApiDialect(args);
|
||||
|
||||
if (args.referenceImages.length > 0) {
|
||||
if (imageApiDialect !== "openai-native") {
|
||||
throw new Error(
|
||||
"Reference images are not supported with the ratio-metadata OpenAI dialect yet. Use openai-native, Google, Azure, OpenRouter, MiniMax, Seedream, or Replicate for image-edit workflows."
|
||||
);
|
||||
}
|
||||
if (model.includes("dall-e-2") || model.includes("dall-e-3")) {
|
||||
throw new Error(
|
||||
"Reference images with OpenAI in this skill require GPT Image models. Use --model gpt-image-2.5-flare (or another gpt-image model)."
|
||||
);
|
||||
}
|
||||
const size = args.size || getOpenAISize(model, args.aspectRatio, args.quality);
|
||||
return generateWithOpenAIEdits(baseURL, apiKey, prompt, model, size, args.referenceImages, args.quality);
|
||||
}
|
||||
|
||||
return generateWithOpenAIGenerations(
|
||||
baseURL,
|
||||
apiKey,
|
||||
buildOpenAIGenerationsBody(prompt, model, args)
|
||||
);
|
||||
}
|
||||
|
||||
async function generateWithChatCompletions(
|
||||
baseURL: string,
|
||||
apiKey: string,
|
||||
prompt: string,
|
||||
model: string
|
||||
): Promise<Uint8Array> {
|
||||
const res = await fetch(`${baseURL}/chat/completions`, {
|
||||
method: "POST",
|
||||
headers: {
|
||||
"Content-Type": "application/json",
|
||||
Authorization: `Bearer ${apiKey}`,
|
||||
},
|
||||
body: JSON.stringify({
|
||||
model,
|
||||
messages: [{ role: "user", content: prompt }],
|
||||
}),
|
||||
});
|
||||
|
||||
if (!res.ok) {
|
||||
const err = await res.text();
|
||||
throw new Error(`OpenAI API error: ${err}`);
|
||||
}
|
||||
|
||||
const result = (await res.json()) as { choices: Array<{ message: { content: string } }> };
|
||||
const content = result.choices[0]?.message?.content ?? "";
|
||||
|
||||
const match = content.match(/data:image\/[^;]+;base64,([A-Za-z0-9+/=]+)/);
|
||||
if (match) {
|
||||
return Uint8Array.from(Buffer.from(match[1]!, "base64"));
|
||||
}
|
||||
|
||||
throw new Error("No image found in chat completions response");
|
||||
}
|
||||
|
||||
async function generateWithOpenAIGenerations(
|
||||
baseURL: string,
|
||||
apiKey: string,
|
||||
body: OpenAIGenerationsBody
|
||||
): Promise<Uint8Array> {
|
||||
const res = await fetch(`${baseURL}/images/generations`, {
|
||||
method: "POST",
|
||||
headers: {
|
||||
"Content-Type": "application/json",
|
||||
Authorization: `Bearer ${apiKey}`,
|
||||
},
|
||||
body: JSON.stringify(body),
|
||||
});
|
||||
|
||||
if (!res.ok) {
|
||||
const err = await res.text();
|
||||
throw new Error(`OpenAI API error: ${err}`);
|
||||
}
|
||||
|
||||
const result = (await res.json()) as OpenAIImageResponse;
|
||||
return extractImageFromResponse(result);
|
||||
}
|
||||
|
||||
async function generateWithOpenAIEdits(
|
||||
baseURL: string,
|
||||
apiKey: string,
|
||||
prompt: string,
|
||||
model: string,
|
||||
size: string,
|
||||
referenceImages: string[],
|
||||
quality: CliArgs["quality"]
|
||||
): Promise<Uint8Array> {
|
||||
const form = new FormData();
|
||||
form.append("model", model);
|
||||
form.append("prompt", prompt);
|
||||
form.append("size", size);
|
||||
|
||||
const openAIQuality = getOpenAIQuality(model, quality);
|
||||
if (openAIQuality && openAIQuality !== "standard" && openAIQuality !== "hd") {
|
||||
form.append("quality", openAIQuality);
|
||||
}
|
||||
|
||||
for (const refPath of referenceImages) {
|
||||
const bytes = await readFile(refPath);
|
||||
const filename = path.basename(refPath);
|
||||
const mimeType = getMimeType(filename);
|
||||
const blob = new Blob([bytes], { type: mimeType });
|
||||
form.append("image[]", blob, filename);
|
||||
}
|
||||
|
||||
const res = await fetch(`${baseURL}/images/edits`, {
|
||||
method: "POST",
|
||||
headers: {
|
||||
Authorization: `Bearer ${apiKey}`,
|
||||
},
|
||||
body: form,
|
||||
});
|
||||
|
||||
if (!res.ok) {
|
||||
const err = await res.text();
|
||||
throw new Error(`OpenAI edits API error: ${err}`);
|
||||
}
|
||||
|
||||
const result = (await res.json()) as OpenAIImageResponse;
|
||||
return extractImageFromResponse(result);
|
||||
}
|
||||
|
||||
export function getMimeType(filename: string): string {
|
||||
const ext = path.extname(filename).toLowerCase();
|
||||
if (ext === ".jpg" || ext === ".jpeg") return "image/jpeg";
|
||||
if (ext === ".webp") return "image/webp";
|
||||
if (ext === ".gif") return "image/gif";
|
||||
return "image/png";
|
||||
}
|
||||
|
||||
export async function extractImageFromResponse(result: OpenAIImageResponse): Promise<Uint8Array> {
|
||||
const img = result.data[0];
|
||||
|
||||
if (img?.b64_json) {
|
||||
return Uint8Array.from(Buffer.from(img.b64_json, "base64"));
|
||||
}
|
||||
|
||||
if (img?.url) {
|
||||
const imgRes = await fetch(img.url);
|
||||
if (!imgRes.ok) throw new Error("Failed to download image");
|
||||
const buf = await imgRes.arrayBuffer();
|
||||
return new Uint8Array(buf);
|
||||
}
|
||||
|
||||
throw new Error("No image in response");
|
||||
}
|
||||
@@ -0,0 +1,169 @@
|
||||
import assert from "node:assert/strict";
|
||||
import test from "node:test";
|
||||
|
||||
import type { CliArgs } from "../types.ts";
|
||||
import {
|
||||
buildContent,
|
||||
buildRequestBody,
|
||||
extractImageFromResponse,
|
||||
getAspectRatio,
|
||||
getImageSize,
|
||||
validateArgs,
|
||||
} from "./openrouter.ts";
|
||||
|
||||
const GEMINI_MODEL = "google/gemini-3.1-flash-image-preview";
|
||||
const GEMINI_25_MODEL = "google/gemini-2.5-flash-image";
|
||||
const GPT_5_IMAGE_MODEL = "openai/gpt-5-image";
|
||||
const OPENROUTER_AUTO_MODEL = "openrouter/auto";
|
||||
const FLUX_MODEL = "black-forest-labs/flux.2-pro";
|
||||
|
||||
function makeArgs(overrides: Partial<CliArgs> = {}): CliArgs {
|
||||
return {
|
||||
prompt: null,
|
||||
promptFiles: [],
|
||||
imagePath: null,
|
||||
provider: null,
|
||||
model: null,
|
||||
aspectRatio: null,
|
||||
size: null,
|
||||
quality: null,
|
||||
imageSize: null,
|
||||
imageApiDialect: null,
|
||||
referenceImages: [],
|
||||
n: 1,
|
||||
batchFile: null,
|
||||
jobs: null,
|
||||
json: false,
|
||||
help: false,
|
||||
...overrides,
|
||||
};
|
||||
}
|
||||
|
||||
test("OpenRouter request body uses image_config and string content for text-only prompts", () => {
|
||||
const args = makeArgs({ aspectRatio: "16:9", quality: "2k" });
|
||||
const body = buildRequestBody("hello", GEMINI_MODEL, args, []);
|
||||
|
||||
assert.deepEqual(body.image_config, {
|
||||
image_size: "2K",
|
||||
aspect_ratio: "16:9",
|
||||
});
|
||||
assert.deepEqual(body.provider, {
|
||||
require_parameters: true,
|
||||
});
|
||||
assert.deepEqual(body.modalities, ["image", "text"]);
|
||||
assert.equal(body.stream, false);
|
||||
assert.equal(body.messages[0].content, "hello");
|
||||
});
|
||||
|
||||
test("OpenRouter request body keeps text+image modalities for current text+image models", () => {
|
||||
for (const model of [GEMINI_MODEL, GEMINI_25_MODEL, GPT_5_IMAGE_MODEL, OPENROUTER_AUTO_MODEL]) {
|
||||
const body = buildRequestBody("hello", model, makeArgs({ quality: "2k" }), []);
|
||||
|
||||
assert.deepEqual(body.image_config, {
|
||||
image_size: "2K",
|
||||
});
|
||||
assert.deepEqual(body.provider, {
|
||||
require_parameters: true,
|
||||
});
|
||||
assert.deepEqual(body.modalities, ["image", "text"]);
|
||||
assert.equal(body.messages[0].content, "hello");
|
||||
}
|
||||
});
|
||||
|
||||
test("OpenRouter request body uses image-only modalities for image-only models under CLI defaults", () => {
|
||||
const body = buildRequestBody("hello", FLUX_MODEL, makeArgs({ quality: "2k" }), []);
|
||||
|
||||
assert.deepEqual(body.image_config, {
|
||||
image_size: "2K",
|
||||
});
|
||||
assert.deepEqual(body.provider, {
|
||||
require_parameters: true,
|
||||
});
|
||||
assert.deepEqual(body.modalities, ["image"]);
|
||||
assert.equal(body.stream, false);
|
||||
assert.equal(body.messages[0].content, "hello");
|
||||
});
|
||||
|
||||
test("OpenRouter helper omits image_config when no size or quality is passed", () => {
|
||||
const body = buildRequestBody("hello", FLUX_MODEL, makeArgs(), []);
|
||||
|
||||
assert.equal(body.image_config, undefined);
|
||||
assert.equal(body.provider, undefined);
|
||||
assert.deepEqual(body.modalities, ["image"]);
|
||||
assert.equal(body.stream, false);
|
||||
assert.equal(body.messages[0].content, "hello");
|
||||
});
|
||||
|
||||
test("OpenRouter request body keeps multimodal array content when references are provided", () => {
|
||||
const content = buildContent("hello", ["data:image/png;base64,abc"]);
|
||||
assert.ok(Array.isArray(content));
|
||||
assert.deepEqual(content[0], { type: "text", text: "hello" });
|
||||
assert.deepEqual(content[1], {
|
||||
type: "image_url",
|
||||
image_url: { url: "data:image/png;base64,abc" },
|
||||
});
|
||||
});
|
||||
|
||||
test("OpenRouter size and aspect helpers infer supported values", () => {
|
||||
assert.equal(getImageSize(makeArgs()), null);
|
||||
assert.equal(getImageSize(makeArgs({ quality: "normal" })), "1K");
|
||||
assert.equal(getImageSize(makeArgs({ size: "2048x1024" })), "2K");
|
||||
assert.equal(getAspectRatio(GEMINI_MODEL, makeArgs({ size: "1600x900" })), "16:9");
|
||||
assert.equal(getAspectRatio(GEMINI_MODEL, makeArgs({ size: "1024x4096" })), "1:4");
|
||||
assert.equal(getAspectRatio(GEMINI_25_MODEL, makeArgs({ size: "1600x900" })), "16:9");
|
||||
assert.equal(getAspectRatio(FLUX_MODEL, makeArgs({ size: "1024x4096" })), null);
|
||||
});
|
||||
|
||||
test("OpenRouter validates explicit aspect ratios and inferred size ratios against model support", () => {
|
||||
assert.doesNotThrow(() =>
|
||||
validateArgs(GEMINI_MODEL, makeArgs({ aspectRatio: "1:4" })),
|
||||
);
|
||||
assert.doesNotThrow(() =>
|
||||
validateArgs(GEMINI_MODEL, makeArgs({ size: "1024x4096" })),
|
||||
);
|
||||
assert.throws(
|
||||
() => validateArgs(GEMINI_25_MODEL, makeArgs({ aspectRatio: "1:4" })),
|
||||
/does not support aspect ratio 1:4/,
|
||||
);
|
||||
assert.throws(
|
||||
() => validateArgs(FLUX_MODEL, makeArgs({ aspectRatio: "1:4" })),
|
||||
/does not support aspect ratio 1:4/,
|
||||
);
|
||||
assert.throws(
|
||||
() => validateArgs(GEMINI_MODEL, makeArgs({ size: "2048x1024" })),
|
||||
/does not support size 2048x1024 \(aspect ratio 2:1\)/,
|
||||
);
|
||||
});
|
||||
|
||||
test("OpenRouter response extraction supports inline image data and finish_reason errors", async () => {
|
||||
const bytes = await extractImageFromResponse({
|
||||
choices: [
|
||||
{
|
||||
message: {
|
||||
images: [
|
||||
{
|
||||
image_url: {
|
||||
url: `data:image/png;base64,${Buffer.from("hello").toString("base64")}`,
|
||||
},
|
||||
},
|
||||
],
|
||||
},
|
||||
},
|
||||
],
|
||||
});
|
||||
assert.equal(Buffer.from(bytes).toString("utf8"), "hello");
|
||||
|
||||
await assert.rejects(
|
||||
() =>
|
||||
extractImageFromResponse({
|
||||
choices: [
|
||||
{
|
||||
finish_reason: "error",
|
||||
native_finish_reason: "MALFORMED_FUNCTION_CALL",
|
||||
message: { content: null },
|
||||
},
|
||||
],
|
||||
}),
|
||||
/finish_reason=MALFORMED_FUNCTION_CALL/,
|
||||
);
|
||||
});
|
||||
@@ -0,0 +1,372 @@
|
||||
import path from "node:path";
|
||||
import { readFile } from "node:fs/promises";
|
||||
import type { CliArgs } from "../types";
|
||||
|
||||
const DEFAULT_MODEL = "google/gemini-3.1-flash-image";
|
||||
const COMMON_ASPECT_RATIOS = [
|
||||
"1:1",
|
||||
"2:3",
|
||||
"3:2",
|
||||
"3:4",
|
||||
"4:3",
|
||||
"4:5",
|
||||
"5:4",
|
||||
"9:16",
|
||||
"16:9",
|
||||
"21:9",
|
||||
];
|
||||
const GEMINI_EXTENDED_ASPECT_RATIOS = ["1:4", "4:1", "1:8", "8:1"];
|
||||
|
||||
type OpenRouterImageEntry = {
|
||||
image_url?: string | { url?: string | null } | null;
|
||||
imageUrl?: string | { url?: string | null } | null;
|
||||
};
|
||||
|
||||
type OpenRouterMessagePart = {
|
||||
type?: string;
|
||||
text?: string;
|
||||
image_url?: string | { url?: string | null } | null;
|
||||
imageUrl?: string | { url?: string | null } | null;
|
||||
};
|
||||
|
||||
type OpenRouterResponse = {
|
||||
choices?: Array<{
|
||||
finish_reason?: string | null;
|
||||
native_finish_reason?: string | null;
|
||||
message?: {
|
||||
images?: OpenRouterImageEntry[];
|
||||
content?: string | OpenRouterMessagePart[] | null;
|
||||
};
|
||||
}>;
|
||||
};
|
||||
|
||||
export function getDefaultModel(): string {
|
||||
return process.env.OPENROUTER_IMAGE_MODEL || DEFAULT_MODEL;
|
||||
}
|
||||
|
||||
function normalizeModelId(model: string): string {
|
||||
return model.trim().toLowerCase().split(":")[0]!;
|
||||
}
|
||||
|
||||
function isTextAndImageModel(model: string): boolean {
|
||||
const normalized = normalizeModelId(model);
|
||||
if (normalized === "openrouter/auto") {
|
||||
return true;
|
||||
}
|
||||
|
||||
if (normalized.startsWith("google/gemini-") && normalized.includes("image")) {
|
||||
return true;
|
||||
}
|
||||
|
||||
if (normalized.startsWith("openai/gpt-") && normalized.includes("image")) {
|
||||
return true;
|
||||
}
|
||||
|
||||
return false;
|
||||
}
|
||||
|
||||
function getSupportedAspectRatios(model: string): Set<string> {
|
||||
const normalized = normalizeModelId(model);
|
||||
if (
|
||||
normalized !== "google/gemini-3.1-flash-image" &&
|
||||
normalized !== "google/gemini-3.1-flash-image-preview"
|
||||
) {
|
||||
return new Set(COMMON_ASPECT_RATIOS);
|
||||
}
|
||||
|
||||
return new Set([...COMMON_ASPECT_RATIOS, ...GEMINI_EXTENDED_ASPECT_RATIOS]);
|
||||
}
|
||||
|
||||
function getApiKey(): string | null {
|
||||
return process.env.OPENROUTER_API_KEY || null;
|
||||
}
|
||||
|
||||
function getBaseUrl(): string {
|
||||
const base = process.env.OPENROUTER_BASE_URL || "https://openrouter.ai/api/v1";
|
||||
return base.replace(/\/+$/g, "");
|
||||
}
|
||||
|
||||
function getHeaders(apiKey: string): Record<string, string> {
|
||||
const headers: Record<string, string> = {
|
||||
"Content-Type": "application/json",
|
||||
Authorization: `Bearer ${apiKey}`,
|
||||
};
|
||||
|
||||
const referer = process.env.OPENROUTER_HTTP_REFERER?.trim();
|
||||
if (referer) {
|
||||
headers["HTTP-Referer"] = referer;
|
||||
}
|
||||
|
||||
const title = process.env.OPENROUTER_TITLE?.trim();
|
||||
if (title) {
|
||||
headers["X-OpenRouter-Title"] = title;
|
||||
headers["X-Title"] = title;
|
||||
}
|
||||
|
||||
return headers;
|
||||
}
|
||||
|
||||
function parsePixelSize(value: string): { width: number; height: number } | null {
|
||||
const match = value.match(/^(\d+)\s*[xX]\s*(\d+)$/);
|
||||
if (!match) return null;
|
||||
|
||||
const width = parseInt(match[1]!, 10);
|
||||
const height = parseInt(match[2]!, 10);
|
||||
|
||||
if (!Number.isFinite(width) || !Number.isFinite(height) || width <= 0 || height <= 0) {
|
||||
return null;
|
||||
}
|
||||
|
||||
return { width, height };
|
||||
}
|
||||
|
||||
function gcd(a: number, b: number): number {
|
||||
let x = Math.abs(a);
|
||||
let y = Math.abs(b);
|
||||
while (y !== 0) {
|
||||
const next = x % y;
|
||||
x = y;
|
||||
y = next;
|
||||
}
|
||||
return x || 1;
|
||||
}
|
||||
|
||||
function inferAspectRatio(size: string | null): string | null {
|
||||
if (!size) return null;
|
||||
const parsed = parsePixelSize(size);
|
||||
if (!parsed) return null;
|
||||
|
||||
const divisor = gcd(parsed.width, parsed.height);
|
||||
return `${parsed.width / divisor}:${parsed.height / divisor}`;
|
||||
}
|
||||
|
||||
function inferImageSize(size: string | null): "1K" | "2K" | "4K" | null {
|
||||
if (!size) return null;
|
||||
const parsed = parsePixelSize(size);
|
||||
if (!parsed) return null;
|
||||
|
||||
const longestEdge = Math.max(parsed.width, parsed.height);
|
||||
if (longestEdge <= 1024) return "1K";
|
||||
if (longestEdge <= 2048) return "2K";
|
||||
return "4K";
|
||||
}
|
||||
|
||||
export function getImageSize(args: CliArgs): "1K" | "2K" | "4K" | null {
|
||||
if (args.imageSize) return args.imageSize as "1K" | "2K" | "4K";
|
||||
|
||||
const inferredFromSize = inferImageSize(args.size);
|
||||
if (inferredFromSize) return inferredFromSize;
|
||||
|
||||
if (args.quality === "normal") return "1K";
|
||||
if (args.quality === "2k") return "2K";
|
||||
return null;
|
||||
}
|
||||
|
||||
export function getAspectRatio(model: string, args: CliArgs): string | null {
|
||||
if (args.aspectRatio) return args.aspectRatio;
|
||||
|
||||
const inferred = inferAspectRatio(args.size);
|
||||
if (!inferred || !getSupportedAspectRatios(model).has(inferred)) {
|
||||
return null;
|
||||
}
|
||||
|
||||
return inferred;
|
||||
}
|
||||
|
||||
function getModalities(model: string): string[] {
|
||||
return isTextAndImageModel(model) ? ["image", "text"] : ["image"];
|
||||
}
|
||||
|
||||
export function validateArgs(model: string, args: CliArgs): void {
|
||||
const requestedAspectRatio = args.aspectRatio || inferAspectRatio(args.size);
|
||||
if (!requestedAspectRatio) {
|
||||
return;
|
||||
}
|
||||
|
||||
const supported = getSupportedAspectRatios(model);
|
||||
if (supported.has(requestedAspectRatio)) {
|
||||
return;
|
||||
}
|
||||
|
||||
const requestedValue = args.aspectRatio
|
||||
? `aspect ratio ${requestedAspectRatio}`
|
||||
: `size ${args.size} (aspect ratio ${requestedAspectRatio})`;
|
||||
|
||||
throw new Error(
|
||||
`OpenRouter model ${model} does not support ${requestedValue}. Supported values: ${Array.from(supported).join(", ")}`
|
||||
);
|
||||
}
|
||||
|
||||
function getMimeType(filename: string): string {
|
||||
const ext = path.extname(filename).toLowerCase();
|
||||
if (ext === ".jpg" || ext === ".jpeg") return "image/jpeg";
|
||||
if (ext === ".webp") return "image/webp";
|
||||
if (ext === ".gif") return "image/gif";
|
||||
return "image/png";
|
||||
}
|
||||
|
||||
async function readImageAsDataUrl(filePath: string): Promise<string> {
|
||||
const bytes = await readFile(filePath);
|
||||
return `data:${getMimeType(filePath)};base64,${bytes.toString("base64")}`;
|
||||
}
|
||||
|
||||
export function buildContent(
|
||||
prompt: string,
|
||||
referenceImages: string[],
|
||||
): string | Array<Record<string, unknown>> {
|
||||
if (referenceImages.length === 0) {
|
||||
return prompt;
|
||||
}
|
||||
|
||||
const content: Array<Record<string, unknown>> = [{ type: "text", text: prompt }];
|
||||
|
||||
for (const imageUrl of referenceImages) {
|
||||
content.push({
|
||||
type: "image_url",
|
||||
image_url: { url: imageUrl },
|
||||
});
|
||||
}
|
||||
|
||||
return content;
|
||||
}
|
||||
|
||||
function extractImageUrl(entry: OpenRouterImageEntry | OpenRouterMessagePart): string | null {
|
||||
const value = entry.image_url ?? entry.imageUrl;
|
||||
if (!value) return null;
|
||||
if (typeof value === "string") return value;
|
||||
return value.url ?? null;
|
||||
}
|
||||
|
||||
function decodeDataUrl(value: string): Uint8Array | null {
|
||||
const match = value.match(/^data:image\/[^;]+;base64,([A-Za-z0-9+/=]+)$/);
|
||||
if (!match) return null;
|
||||
return Uint8Array.from(Buffer.from(match[1]!, "base64"));
|
||||
}
|
||||
|
||||
async function downloadImage(value: string): Promise<Uint8Array> {
|
||||
const inline = decodeDataUrl(value);
|
||||
if (inline) return inline;
|
||||
|
||||
if (value.startsWith("http://") || value.startsWith("https://")) {
|
||||
const response = await fetch(value);
|
||||
if (!response.ok) {
|
||||
throw new Error(`Failed to download OpenRouter image: ${response.status}`);
|
||||
}
|
||||
const buffer = await response.arrayBuffer();
|
||||
return new Uint8Array(buffer);
|
||||
}
|
||||
|
||||
return Uint8Array.from(Buffer.from(value, "base64"));
|
||||
}
|
||||
|
||||
export async function extractImageFromResponse(result: OpenRouterResponse): Promise<Uint8Array> {
|
||||
const choice = result.choices?.[0];
|
||||
const message = choice?.message;
|
||||
|
||||
for (const image of message?.images ?? []) {
|
||||
const imageUrl = extractImageUrl(image);
|
||||
if (imageUrl) return downloadImage(imageUrl);
|
||||
}
|
||||
|
||||
if (Array.isArray(message?.content)) {
|
||||
for (const item of message.content) {
|
||||
const imageUrl = extractImageUrl(item);
|
||||
if (imageUrl) return downloadImage(imageUrl);
|
||||
|
||||
if (item.type === "text" && item.text) {
|
||||
const inline = decodeDataUrl(item.text);
|
||||
if (inline) return inline;
|
||||
}
|
||||
}
|
||||
} else if (typeof message?.content === "string") {
|
||||
const inline = decodeDataUrl(message.content);
|
||||
if (inline) return inline;
|
||||
}
|
||||
|
||||
const finishReason =
|
||||
choice?.native_finish_reason || choice?.finish_reason || "unknown";
|
||||
throw new Error(
|
||||
`No image in OpenRouter response (finish_reason=${finishReason})`,
|
||||
);
|
||||
}
|
||||
|
||||
export function buildRequestBody(
|
||||
prompt: string,
|
||||
model: string,
|
||||
args: CliArgs,
|
||||
referenceImages: string[],
|
||||
): Record<string, unknown> {
|
||||
validateArgs(model, args);
|
||||
|
||||
const imageConfig: Record<string, string> = {};
|
||||
|
||||
const imageSize = getImageSize(args);
|
||||
if (imageSize) {
|
||||
imageConfig.image_size = imageSize;
|
||||
}
|
||||
|
||||
const aspectRatio = getAspectRatio(model, args);
|
||||
if (aspectRatio) {
|
||||
imageConfig.aspect_ratio = aspectRatio;
|
||||
}
|
||||
|
||||
const body: Record<string, unknown> = {
|
||||
messages: [
|
||||
{
|
||||
role: "user",
|
||||
content: buildContent(prompt, referenceImages),
|
||||
},
|
||||
],
|
||||
modalities: getModalities(model),
|
||||
stream: false,
|
||||
};
|
||||
|
||||
if (Object.keys(imageConfig).length > 0) {
|
||||
body.image_config = imageConfig;
|
||||
body.provider = {
|
||||
require_parameters: true,
|
||||
};
|
||||
}
|
||||
|
||||
return body;
|
||||
}
|
||||
|
||||
export async function generateImage(
|
||||
prompt: string,
|
||||
model: string,
|
||||
args: CliArgs
|
||||
): Promise<Uint8Array> {
|
||||
const apiKey = getApiKey();
|
||||
if (!apiKey) {
|
||||
throw new Error("OPENROUTER_API_KEY is required. Get one at https://openrouter.ai/settings/keys");
|
||||
}
|
||||
|
||||
const referenceImages: string[] = [];
|
||||
for (const refPath of args.referenceImages) {
|
||||
referenceImages.push(await readImageAsDataUrl(refPath));
|
||||
}
|
||||
|
||||
const body = {
|
||||
model,
|
||||
...buildRequestBody(prompt, model, args, referenceImages),
|
||||
};
|
||||
|
||||
console.log(
|
||||
`Generating image with OpenRouter (${model})...`,
|
||||
(body.image_config as Record<string, string>),
|
||||
);
|
||||
|
||||
const response = await fetch(`${getBaseUrl()}/chat/completions`, {
|
||||
method: "POST",
|
||||
headers: getHeaders(apiKey),
|
||||
body: JSON.stringify(body),
|
||||
});
|
||||
|
||||
if (!response.ok) {
|
||||
const errorText = await response.text();
|
||||
throw new Error(`OpenRouter API error (${response.status}): ${errorText}`);
|
||||
}
|
||||
|
||||
const result = (await response.json()) as OpenRouterResponse;
|
||||
return extractImageFromResponse(result);
|
||||
}
|
||||
@@ -0,0 +1,310 @@
|
||||
import assert from "node:assert/strict";
|
||||
import test from "node:test";
|
||||
|
||||
import type { CliArgs } from "../types.ts";
|
||||
import {
|
||||
buildInput,
|
||||
extractOutputUrl,
|
||||
getDefaultModel,
|
||||
getModelFamily,
|
||||
parseModelId,
|
||||
validateArgs,
|
||||
} from "./replicate.ts";
|
||||
|
||||
function makeArgs(overrides: Partial<CliArgs> = {}): CliArgs {
|
||||
return {
|
||||
prompt: null,
|
||||
promptFiles: [],
|
||||
imagePath: null,
|
||||
provider: null,
|
||||
model: null,
|
||||
aspectRatio: null,
|
||||
aspectRatioSource: null,
|
||||
size: null,
|
||||
quality: null,
|
||||
imageSize: null,
|
||||
imageSizeSource: null,
|
||||
imageApiDialect: null,
|
||||
referenceImages: [],
|
||||
n: 1,
|
||||
batchFile: null,
|
||||
jobs: null,
|
||||
json: false,
|
||||
help: false,
|
||||
...overrides,
|
||||
};
|
||||
}
|
||||
|
||||
test("Replicate default model now points at nano-banana-2", () => {
|
||||
const previous = process.env.REPLICATE_IMAGE_MODEL;
|
||||
delete process.env.REPLICATE_IMAGE_MODEL;
|
||||
try {
|
||||
assert.equal(getDefaultModel(), "google/nano-banana-2");
|
||||
} finally {
|
||||
if (previous == null) {
|
||||
delete process.env.REPLICATE_IMAGE_MODEL;
|
||||
} else {
|
||||
process.env.REPLICATE_IMAGE_MODEL = previous;
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
test("Replicate model parsing and family detection accept supported official ids", () => {
|
||||
assert.deepEqual(parseModelId("google/nano-banana-2"), {
|
||||
owner: "google",
|
||||
name: "nano-banana-2",
|
||||
version: null,
|
||||
});
|
||||
assert.deepEqual(parseModelId("owner/model:abc123"), {
|
||||
owner: "owner",
|
||||
name: "model",
|
||||
version: "abc123",
|
||||
});
|
||||
|
||||
assert.equal(getModelFamily("google/nano-banana-pro"), "nano-banana");
|
||||
assert.equal(getModelFamily("bytedance/seedream-4.5"), "seedream45");
|
||||
assert.equal(getModelFamily("bytedance/seedream-5-lite"), "seedream5lite");
|
||||
assert.equal(getModelFamily("wan-video/wan-2.7-image"), "wan27image");
|
||||
assert.equal(getModelFamily("wan-video/wan-2.7-image-pro"), "wan27imagepro");
|
||||
assert.equal(getModelFamily("stability-ai/sdxl"), "unknown");
|
||||
|
||||
assert.throws(
|
||||
() => parseModelId("just-a-model-name"),
|
||||
/Invalid Replicate model format/,
|
||||
);
|
||||
});
|
||||
|
||||
test("Replicate nano-banana input builder maps refs, aspect ratio, and quality presets", () => {
|
||||
assert.deepEqual(
|
||||
buildInput(
|
||||
"google/nano-banana-2",
|
||||
"A robot painter",
|
||||
makeArgs({
|
||||
aspectRatio: "16:9",
|
||||
quality: "2k",
|
||||
}),
|
||||
["data:image/png;base64,AAAA"],
|
||||
),
|
||||
{
|
||||
prompt: "A robot painter",
|
||||
resolution: "2K",
|
||||
output_format: "png",
|
||||
aspect_ratio: "16:9",
|
||||
image_input: ["data:image/png;base64,AAAA"],
|
||||
},
|
||||
);
|
||||
|
||||
assert.deepEqual(
|
||||
buildInput(
|
||||
"google/nano-banana-2",
|
||||
"A robot painter",
|
||||
makeArgs({ size: "1024x1024", quality: "normal" }),
|
||||
[],
|
||||
),
|
||||
{
|
||||
prompt: "A robot painter",
|
||||
resolution: "1K",
|
||||
output_format: "png",
|
||||
aspect_ratio: "1:1",
|
||||
},
|
||||
);
|
||||
});
|
||||
|
||||
test("Replicate Seedream and Wan inputs use family-specific request fields", () => {
|
||||
assert.deepEqual(
|
||||
buildInput(
|
||||
"bytedance/seedream-4.5",
|
||||
"A cinematic portrait",
|
||||
makeArgs({ quality: "2k", referenceImages: ["local.png"] }),
|
||||
["data:image/png;base64,AAAA"],
|
||||
),
|
||||
{
|
||||
prompt: "A cinematic portrait",
|
||||
size: "4K",
|
||||
image_input: ["data:image/png;base64,AAAA"],
|
||||
aspect_ratio: "match_input_image",
|
||||
},
|
||||
);
|
||||
|
||||
assert.deepEqual(
|
||||
buildInput(
|
||||
"bytedance/seedream-4.5",
|
||||
"A cinematic portrait",
|
||||
makeArgs({ size: "1536x1024" }),
|
||||
[],
|
||||
),
|
||||
{
|
||||
prompt: "A cinematic portrait",
|
||||
size: "custom",
|
||||
width: 1536,
|
||||
height: 1024,
|
||||
},
|
||||
);
|
||||
|
||||
assert.deepEqual(
|
||||
buildInput(
|
||||
"bytedance/seedream-5-lite",
|
||||
"A poster",
|
||||
makeArgs({ aspectRatio: "21:9", quality: "2k" }),
|
||||
[],
|
||||
),
|
||||
{
|
||||
prompt: "A poster",
|
||||
size: "3K",
|
||||
aspect_ratio: "21:9",
|
||||
},
|
||||
);
|
||||
|
||||
assert.deepEqual(
|
||||
buildInput(
|
||||
"wan-video/wan-2.7-image",
|
||||
"A storyboard frame",
|
||||
makeArgs({ aspectRatio: "16:9", quality: "2k" }),
|
||||
[],
|
||||
),
|
||||
{
|
||||
prompt: "A storyboard frame",
|
||||
size: "2048*1152",
|
||||
},
|
||||
);
|
||||
|
||||
assert.deepEqual(
|
||||
buildInput(
|
||||
"wan-video/wan-2.7-image-pro",
|
||||
"Blend these references",
|
||||
makeArgs({ size: "2K", referenceImages: ["a.png", "b.png"] }),
|
||||
["ref-a", "ref-b"],
|
||||
),
|
||||
{
|
||||
prompt: "Blend these references",
|
||||
size: "2K",
|
||||
images: ["ref-a", "ref-b"],
|
||||
},
|
||||
);
|
||||
});
|
||||
|
||||
test("Replicate validateArgs blocks misleading multi-output and unsupported family options locally", () => {
|
||||
assert.throws(
|
||||
() =>
|
||||
validateArgs(
|
||||
"google/nano-banana-2",
|
||||
makeArgs({ n: 2 }),
|
||||
),
|
||||
/exactly one output image/,
|
||||
);
|
||||
|
||||
assert.throws(
|
||||
() =>
|
||||
validateArgs(
|
||||
"bytedance/seedream-4.5",
|
||||
makeArgs({ size: "1K" }),
|
||||
),
|
||||
/2K, 4K, or an explicit WxH size/,
|
||||
);
|
||||
|
||||
assert.throws(
|
||||
() =>
|
||||
validateArgs(
|
||||
"bytedance/seedream-5-lite",
|
||||
makeArgs({ size: "4K" }),
|
||||
),
|
||||
/supports 2K or 3K output/,
|
||||
);
|
||||
|
||||
assert.throws(
|
||||
() =>
|
||||
validateArgs(
|
||||
"wan-video/wan-2.7-image",
|
||||
makeArgs({ referenceImages: new Array(10).fill("ref.png") }),
|
||||
),
|
||||
/at most 9 reference images/,
|
||||
);
|
||||
|
||||
assert.throws(
|
||||
() =>
|
||||
validateArgs(
|
||||
"wan-video/wan-2.7-image-pro",
|
||||
makeArgs({ referenceImages: ["ref.png"], size: "4K" }),
|
||||
),
|
||||
/only supports 4K text-to-image/,
|
||||
);
|
||||
|
||||
assert.throws(
|
||||
() =>
|
||||
validateArgs(
|
||||
"stability-ai/sdxl",
|
||||
makeArgs({ aspectRatio: "16:9" }),
|
||||
),
|
||||
/compatibility list/,
|
||||
);
|
||||
|
||||
assert.doesNotThrow(() =>
|
||||
validateArgs(
|
||||
"google/nano-banana-2",
|
||||
makeArgs({ imageSize: "2K", imageSizeSource: "config" }),
|
||||
),
|
||||
);
|
||||
|
||||
assert.throws(
|
||||
() =>
|
||||
validateArgs(
|
||||
"google/nano-banana-2",
|
||||
makeArgs({ imageSize: "2K", imageSizeSource: "cli" }),
|
||||
),
|
||||
/do not use --imageSize/,
|
||||
);
|
||||
|
||||
assert.doesNotThrow(() =>
|
||||
validateArgs(
|
||||
"stability-ai/sdxl",
|
||||
makeArgs({ aspectRatio: "16:9", aspectRatioSource: "config" }),
|
||||
),
|
||||
);
|
||||
|
||||
assert.throws(
|
||||
() =>
|
||||
validateArgs(
|
||||
"stability-ai/sdxl",
|
||||
makeArgs({ aspectRatio: "16:9", aspectRatioSource: "cli" }),
|
||||
),
|
||||
/compatibility list/,
|
||||
);
|
||||
|
||||
assert.doesNotThrow(() =>
|
||||
validateArgs(
|
||||
"stability-ai/sdxl",
|
||||
makeArgs(),
|
||||
),
|
||||
);
|
||||
});
|
||||
|
||||
test("Replicate output extraction supports single outputs and rejects silent multi-image drops", () => {
|
||||
assert.equal(
|
||||
extractOutputUrl({ output: "https://example.com/a.png" } as never),
|
||||
"https://example.com/a.png",
|
||||
);
|
||||
assert.equal(
|
||||
extractOutputUrl({ output: ["https://example.com/b.png"] } as never),
|
||||
"https://example.com/b.png",
|
||||
);
|
||||
assert.equal(
|
||||
extractOutputUrl({ output: { url: "https://example.com/c.png" } } as never),
|
||||
"https://example.com/c.png",
|
||||
);
|
||||
|
||||
assert.throws(
|
||||
() =>
|
||||
extractOutputUrl({
|
||||
output: [
|
||||
"https://example.com/one.png",
|
||||
"https://example.com/two.png",
|
||||
],
|
||||
} as never),
|
||||
/supports saving exactly one image/,
|
||||
);
|
||||
|
||||
assert.throws(
|
||||
() => extractOutputUrl({ output: { invalid: true } } as never),
|
||||
/Unexpected Replicate output format/,
|
||||
);
|
||||
});
|
||||
@@ -0,0 +1,616 @@
|
||||
import path from "node:path";
|
||||
import { readFile } from "node:fs/promises";
|
||||
import type { CliArgs } from "../types";
|
||||
|
||||
const DEFAULT_MODEL = "google/nano-banana-2";
|
||||
const SYNC_WAIT_SECONDS = 60;
|
||||
const POLL_INTERVAL_MS = 2000;
|
||||
const MAX_POLL_MS = 300_000;
|
||||
const DOCUMENTED_REPLICATE_ASPECT_RATIOS = new Set([
|
||||
"1:1",
|
||||
"2:3",
|
||||
"3:2",
|
||||
"3:4",
|
||||
"4:3",
|
||||
"5:4",
|
||||
"4:5",
|
||||
"9:16",
|
||||
"16:9",
|
||||
"21:9",
|
||||
]);
|
||||
|
||||
export type ReplicateModelFamily =
|
||||
| "nano-banana"
|
||||
| "seedream45"
|
||||
| "seedream5lite"
|
||||
| "wan27image"
|
||||
| "wan27imagepro"
|
||||
| "unknown";
|
||||
|
||||
type PixelSize = {
|
||||
width: number;
|
||||
height: number;
|
||||
};
|
||||
|
||||
type Seedream45Size = "2K" | "4K" | { width: number; height: number };
|
||||
|
||||
export function getDefaultModel(): string {
|
||||
return process.env.REPLICATE_IMAGE_MODEL || DEFAULT_MODEL;
|
||||
}
|
||||
|
||||
function getApiToken(): string | null {
|
||||
return process.env.REPLICATE_API_TOKEN || null;
|
||||
}
|
||||
|
||||
function getBaseUrl(): string {
|
||||
const base = process.env.REPLICATE_BASE_URL || "https://api.replicate.com";
|
||||
return base.replace(/\/+$/g, "");
|
||||
}
|
||||
|
||||
function normalizeModelId(model: string): string {
|
||||
return model.trim().toLowerCase().split(":")[0]!;
|
||||
}
|
||||
|
||||
export function getModelFamily(model: string): ReplicateModelFamily {
|
||||
const normalized = normalizeModelId(model);
|
||||
|
||||
if (
|
||||
normalized === "google/nano-banana" ||
|
||||
normalized === "google/nano-banana-pro" ||
|
||||
normalized === "google/nano-banana-2"
|
||||
) {
|
||||
return "nano-banana";
|
||||
}
|
||||
|
||||
if (normalized === "bytedance/seedream-4.5") {
|
||||
return "seedream45";
|
||||
}
|
||||
|
||||
if (normalized === "bytedance/seedream-5-lite") {
|
||||
return "seedream5lite";
|
||||
}
|
||||
|
||||
if (normalized === "wan-video/wan-2.7-image") {
|
||||
return "wan27image";
|
||||
}
|
||||
|
||||
if (normalized === "wan-video/wan-2.7-image-pro") {
|
||||
return "wan27imagepro";
|
||||
}
|
||||
|
||||
return "unknown";
|
||||
}
|
||||
|
||||
export function parseModelId(model: string): { owner: string; name: string; version: string | null } {
|
||||
const [ownerName, version] = model.split(":");
|
||||
const parts = ownerName!.split("/");
|
||||
if (parts.length !== 2 || !parts[0] || !parts[1]) {
|
||||
throw new Error(
|
||||
`Invalid Replicate model format: "${model}". Expected "owner/name" or "owner/name:version".`
|
||||
);
|
||||
}
|
||||
return { owner: parts[0], name: parts[1], version: version || null };
|
||||
}
|
||||
|
||||
function parsePixelSize(value: string): PixelSize | null {
|
||||
const match = value.trim().match(/^(\d+)\s*[xX*]\s*(\d+)$/);
|
||||
if (!match) return null;
|
||||
|
||||
const width = parseInt(match[1]!, 10);
|
||||
const height = parseInt(match[2]!, 10);
|
||||
if (!Number.isFinite(width) || !Number.isFinite(height) || width <= 0 || height <= 0) {
|
||||
return null;
|
||||
}
|
||||
|
||||
return { width, height };
|
||||
}
|
||||
|
||||
function parseAspectRatio(value: string): PixelSize | null {
|
||||
const match = value.trim().match(/^(\d+)\s*:\s*(\d+)$/);
|
||||
if (!match) return null;
|
||||
|
||||
const width = parseInt(match[1]!, 10);
|
||||
const height = parseInt(match[2]!, 10);
|
||||
if (!Number.isFinite(width) || !Number.isFinite(height) || width <= 0 || height <= 0) {
|
||||
return null;
|
||||
}
|
||||
|
||||
return { width, height };
|
||||
}
|
||||
|
||||
function gcd(a: number, b: number): number {
|
||||
let x = Math.abs(a);
|
||||
let y = Math.abs(b);
|
||||
|
||||
while (y !== 0) {
|
||||
const next = x % y;
|
||||
x = y;
|
||||
y = next;
|
||||
}
|
||||
|
||||
return x || 1;
|
||||
}
|
||||
|
||||
function inferAspectRatioFromSize(size: string): string | null {
|
||||
const parsed = parsePixelSize(size);
|
||||
if (!parsed) return null;
|
||||
|
||||
const divisor = gcd(parsed.width, parsed.height);
|
||||
const normalized = `${parsed.width / divisor}:${parsed.height / divisor}`;
|
||||
if (!DOCUMENTED_REPLICATE_ASPECT_RATIOS.has(normalized)) {
|
||||
return null;
|
||||
}
|
||||
|
||||
return normalized;
|
||||
}
|
||||
|
||||
function getQualityPreset(args: CliArgs): "normal" | "2k" {
|
||||
return args.quality === "normal" ? "normal" : "2k";
|
||||
}
|
||||
|
||||
function validateDocumentedAspectRatio(model: string, aspectRatio: string): void {
|
||||
if (aspectRatio === "match_input_image") {
|
||||
return;
|
||||
}
|
||||
|
||||
if (DOCUMENTED_REPLICATE_ASPECT_RATIOS.has(aspectRatio)) {
|
||||
return;
|
||||
}
|
||||
|
||||
throw new Error(
|
||||
`Replicate model ${model} does not support aspect ratio ${aspectRatio}. Supported values: ${Array.from(DOCUMENTED_REPLICATE_ASPECT_RATIOS).join(", ")}`
|
||||
);
|
||||
}
|
||||
|
||||
function getRequestedAspectRatio(model: string, args: CliArgs): string | null {
|
||||
if (args.aspectRatio) {
|
||||
validateDocumentedAspectRatio(model, args.aspectRatio);
|
||||
return args.aspectRatio;
|
||||
}
|
||||
|
||||
if (!args.size) return null;
|
||||
|
||||
const inferred = inferAspectRatioFromSize(args.size);
|
||||
if (!inferred) {
|
||||
throw new Error(
|
||||
`Replicate model ${model} cannot derive a supported aspect ratio from --size ${args.size}. Use one of: ${Array.from(DOCUMENTED_REPLICATE_ASPECT_RATIOS).join(", ")}`
|
||||
);
|
||||
}
|
||||
|
||||
return inferred;
|
||||
}
|
||||
|
||||
function getNanoBananaResolution(args: CliArgs): "1K" | "2K" {
|
||||
if (args.size) {
|
||||
const parsed = parsePixelSize(args.size);
|
||||
if (!parsed) {
|
||||
throw new Error("Replicate nano-banana --size must be in WxH format, for example 1536x1024.");
|
||||
}
|
||||
|
||||
const longestEdge = Math.max(parsed.width, parsed.height);
|
||||
if (longestEdge <= 1024) return "1K";
|
||||
if (longestEdge <= 2048) return "2K";
|
||||
throw new Error("Replicate nano-banana only supports sizes that map to 1K or 2K output.");
|
||||
}
|
||||
|
||||
return getQualityPreset(args) === "normal" ? "1K" : "2K";
|
||||
}
|
||||
|
||||
function resolveSeedream45Size(args: CliArgs): Seedream45Size {
|
||||
if (args.size) {
|
||||
const upper = args.size.trim().toUpperCase();
|
||||
if (upper === "2K" || upper === "4K") {
|
||||
return upper;
|
||||
}
|
||||
|
||||
const parsed = parsePixelSize(args.size);
|
||||
if (!parsed) {
|
||||
throw new Error("Replicate Seedream 4.5 --size must be 2K, 4K, or an explicit WxH size.");
|
||||
}
|
||||
if (parsed.width < 1024 || parsed.width > 4096 || parsed.height < 1024 || parsed.height > 4096) {
|
||||
throw new Error("Replicate Seedream 4.5 custom --size must keep width and height between 1024 and 4096.");
|
||||
}
|
||||
return parsed;
|
||||
}
|
||||
|
||||
return getQualityPreset(args) === "normal" ? "2K" : "4K";
|
||||
}
|
||||
|
||||
function resolveSeedream5LiteSize(args: CliArgs): "2K" | "3K" {
|
||||
if (args.size) {
|
||||
const upper = args.size.trim().toUpperCase();
|
||||
if (upper === "2K" || upper === "3K") {
|
||||
return upper;
|
||||
}
|
||||
|
||||
throw new Error("Replicate Seedream 5 Lite currently supports 2K or 3K output in this tool.");
|
||||
}
|
||||
|
||||
return getQualityPreset(args) === "normal" ? "2K" : "3K";
|
||||
}
|
||||
|
||||
function formatCustomWanSize(size: PixelSize): string {
|
||||
return `${size.width}*${size.height}`;
|
||||
}
|
||||
|
||||
function resolveWanSizeFromAspectRatio(
|
||||
aspectRatio: string,
|
||||
maxDimension: number,
|
||||
): string {
|
||||
const parsedRatio = parseAspectRatio(aspectRatio);
|
||||
if (!parsedRatio) {
|
||||
throw new Error(`Replicate Wan aspect ratio must be in W:H format, got ${aspectRatio}.`);
|
||||
}
|
||||
|
||||
const scale = Math.min(maxDimension / parsedRatio.width, maxDimension / parsedRatio.height);
|
||||
const width = Math.max(1, Math.floor(parsedRatio.width * scale));
|
||||
const height = Math.max(1, Math.floor(parsedRatio.height * scale));
|
||||
return formatCustomWanSize({ width, height });
|
||||
}
|
||||
|
||||
function resolveWanSize(family: "wan27image" | "wan27imagepro", args: CliArgs): "1K" | "2K" | "4K" | string {
|
||||
const referenceMode = args.referenceImages.length > 0;
|
||||
const maxDimension = family === "wan27imagepro" && !referenceMode ? 4096 : 2048;
|
||||
|
||||
if (args.size) {
|
||||
const upper = args.size.trim().toUpperCase();
|
||||
if (upper === "1K" || upper === "2K" || upper === "4K") {
|
||||
if (upper === "4K" && family !== "wan27imagepro") {
|
||||
throw new Error("Replicate Wan 2.7 Image only supports 1K, 2K, or custom sizes up to 2048px.");
|
||||
}
|
||||
if (upper === "4K" && referenceMode) {
|
||||
throw new Error("Replicate Wan 2.7 Image Pro only supports 4K text-to-image. Remove --ref or lower the size.");
|
||||
}
|
||||
return upper;
|
||||
}
|
||||
|
||||
const parsed = parsePixelSize(args.size);
|
||||
if (!parsed) {
|
||||
throw new Error("Replicate Wan --size must be 1K, 2K, 4K, or an explicit WxH size.");
|
||||
}
|
||||
if (parsed.width > maxDimension || parsed.height > maxDimension) {
|
||||
throw new Error(
|
||||
`Replicate ${family === "wan27imagepro" ? "Wan 2.7 Image Pro" : "Wan 2.7 Image"} custom --size must keep width and height at or below ${maxDimension}px in the current mode.`
|
||||
);
|
||||
}
|
||||
return formatCustomWanSize(parsed);
|
||||
}
|
||||
|
||||
if (args.aspectRatio) {
|
||||
return resolveWanSizeFromAspectRatio(
|
||||
args.aspectRatio,
|
||||
getQualityPreset(args) === "normal" ? 1024 : 2048,
|
||||
);
|
||||
}
|
||||
|
||||
return getQualityPreset(args) === "normal" ? "1K" : "2K";
|
||||
}
|
||||
|
||||
function buildNanoBananaInput(
|
||||
prompt: string,
|
||||
model: string,
|
||||
args: CliArgs,
|
||||
referenceImages: string[],
|
||||
): Record<string, unknown> {
|
||||
const input: Record<string, unknown> = {
|
||||
prompt,
|
||||
resolution: getNanoBananaResolution(args),
|
||||
output_format: "png",
|
||||
};
|
||||
|
||||
const aspectRatio = getRequestedAspectRatio(model, args);
|
||||
if (aspectRatio) {
|
||||
input.aspect_ratio = aspectRatio;
|
||||
} else if (referenceImages.length > 0) {
|
||||
input.aspect_ratio = "match_input_image";
|
||||
}
|
||||
|
||||
if (referenceImages.length > 0) {
|
||||
input.image_input = referenceImages;
|
||||
}
|
||||
|
||||
return input;
|
||||
}
|
||||
|
||||
function buildSeedreamInput(
|
||||
family: "seedream45" | "seedream5lite",
|
||||
prompt: string,
|
||||
model: string,
|
||||
args: CliArgs,
|
||||
referenceImages: string[],
|
||||
): Record<string, unknown> {
|
||||
const size = family === "seedream45" ? resolveSeedream45Size(args) : resolveSeedream5LiteSize(args);
|
||||
const input: Record<string, unknown> = {
|
||||
prompt,
|
||||
};
|
||||
|
||||
if (family === "seedream45" && typeof size === "object") {
|
||||
input.size = "custom";
|
||||
input.width = size.width;
|
||||
input.height = size.height;
|
||||
} else {
|
||||
input.size = size;
|
||||
}
|
||||
|
||||
if (referenceImages.length > 0) {
|
||||
input.image_input = referenceImages;
|
||||
}
|
||||
|
||||
if (args.aspectRatio) {
|
||||
validateDocumentedAspectRatio(model, args.aspectRatio);
|
||||
input.aspect_ratio = args.aspectRatio;
|
||||
} else if (referenceImages.length > 0 && family === "seedream45") {
|
||||
input.aspect_ratio = "match_input_image";
|
||||
}
|
||||
|
||||
return input;
|
||||
}
|
||||
|
||||
function buildWanInput(
|
||||
family: "wan27image" | "wan27imagepro",
|
||||
prompt: string,
|
||||
args: CliArgs,
|
||||
referenceImages: string[],
|
||||
): Record<string, unknown> {
|
||||
const input: Record<string, unknown> = {
|
||||
prompt,
|
||||
size: resolveWanSize(family, args),
|
||||
};
|
||||
|
||||
if (referenceImages.length > 0) {
|
||||
input.images = referenceImages;
|
||||
}
|
||||
|
||||
return input;
|
||||
}
|
||||
|
||||
export function validateArgs(model: string, args: CliArgs): void {
|
||||
parseModelId(model);
|
||||
|
||||
if (args.n !== 1) {
|
||||
throw new Error("Replicate integration currently supports exactly one output image per request. Remove --n or use --n 1.");
|
||||
}
|
||||
|
||||
if (args.imageSize && args.imageSizeSource !== "config") {
|
||||
throw new Error("Replicate models in baoyu-image-gen do not use --imageSize. Use --quality, --ar, or --size instead.");
|
||||
}
|
||||
|
||||
const family = getModelFamily(model);
|
||||
|
||||
if (family === "nano-banana") {
|
||||
if (args.referenceImages.length > 14) {
|
||||
throw new Error("Replicate nano-banana supports at most 14 reference images.");
|
||||
}
|
||||
if (args.aspectRatio) {
|
||||
validateDocumentedAspectRatio(model, args.aspectRatio);
|
||||
}
|
||||
if (args.size) {
|
||||
getRequestedAspectRatio(model, args);
|
||||
getNanoBananaResolution(args);
|
||||
}
|
||||
return;
|
||||
}
|
||||
|
||||
if (family === "seedream45") {
|
||||
if (args.referenceImages.length > 14) {
|
||||
throw new Error("Replicate Seedream 4.5 supports at most 14 reference images.");
|
||||
}
|
||||
if (args.aspectRatio) {
|
||||
validateDocumentedAspectRatio(model, args.aspectRatio);
|
||||
}
|
||||
resolveSeedream45Size(args);
|
||||
return;
|
||||
}
|
||||
|
||||
if (family === "seedream5lite") {
|
||||
if (args.referenceImages.length > 14) {
|
||||
throw new Error("Replicate Seedream 5 Lite supports at most 14 reference images.");
|
||||
}
|
||||
if (args.aspectRatio) {
|
||||
validateDocumentedAspectRatio(model, args.aspectRatio);
|
||||
}
|
||||
resolveSeedream5LiteSize(args);
|
||||
return;
|
||||
}
|
||||
|
||||
if (family === "wan27image" || family === "wan27imagepro") {
|
||||
if (args.referenceImages.length > 9) {
|
||||
throw new Error("Replicate Wan 2.7 image models support at most 9 reference images.");
|
||||
}
|
||||
if (args.aspectRatio) {
|
||||
const parsed = parseAspectRatio(args.aspectRatio);
|
||||
if (!parsed) {
|
||||
throw new Error(`Replicate Wan aspect ratio must be in W:H format, got ${args.aspectRatio}.`);
|
||||
}
|
||||
}
|
||||
resolveWanSize(family, args);
|
||||
return;
|
||||
}
|
||||
|
||||
const hasExplicitAspectRatio = !!args.aspectRatio && args.aspectRatioSource !== "config";
|
||||
|
||||
if (args.referenceImages.length > 0 || hasExplicitAspectRatio || args.size) {
|
||||
throw new Error(
|
||||
`Replicate model ${model} is not in the baoyu-image-gen compatibility list. Supported families: google/nano-banana*, bytedance/seedream-4.5, bytedance/seedream-5-lite, wan-video/wan-2.7-image, wan-video/wan-2.7-image-pro.`
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
export function getDefaultOutputExtension(model: string): ".png" {
|
||||
const _family = getModelFamily(model);
|
||||
return ".png";
|
||||
}
|
||||
|
||||
export function buildInput(
|
||||
model: string,
|
||||
prompt: string,
|
||||
args: CliArgs,
|
||||
referenceImages: string[],
|
||||
): Record<string, unknown> {
|
||||
const family = getModelFamily(model);
|
||||
|
||||
if (family === "nano-banana") {
|
||||
return buildNanoBananaInput(prompt, model, args, referenceImages);
|
||||
}
|
||||
|
||||
if (family === "seedream45" || family === "seedream5lite") {
|
||||
return buildSeedreamInput(family, prompt, model, args, referenceImages);
|
||||
}
|
||||
|
||||
if (family === "wan27image" || family === "wan27imagepro") {
|
||||
return buildWanInput(family, prompt, args, referenceImages);
|
||||
}
|
||||
|
||||
return { prompt };
|
||||
}
|
||||
|
||||
async function readImageAsDataUrl(p: string): Promise<string> {
|
||||
const buf = await readFile(p);
|
||||
const ext = path.extname(p).toLowerCase();
|
||||
let mimeType = "image/png";
|
||||
if (ext === ".jpg" || ext === ".jpeg") mimeType = "image/jpeg";
|
||||
else if (ext === ".gif") mimeType = "image/gif";
|
||||
else if (ext === ".webp") mimeType = "image/webp";
|
||||
return `data:${mimeType};base64,${buf.toString("base64")}`;
|
||||
}
|
||||
|
||||
type PredictionResponse = {
|
||||
id: string;
|
||||
status: string;
|
||||
output: unknown;
|
||||
error: string | null;
|
||||
urls?: { get?: string };
|
||||
};
|
||||
|
||||
async function createPrediction(
|
||||
apiToken: string,
|
||||
model: { owner: string; name: string; version: string | null },
|
||||
input: Record<string, unknown>,
|
||||
sync: boolean
|
||||
): Promise<PredictionResponse> {
|
||||
const baseUrl = getBaseUrl();
|
||||
|
||||
let url: string;
|
||||
const body: Record<string, unknown> = { input };
|
||||
|
||||
if (model.version) {
|
||||
url = `${baseUrl}/v1/predictions`;
|
||||
body.version = model.version;
|
||||
} else {
|
||||
url = `${baseUrl}/v1/models/${model.owner}/${model.name}/predictions`;
|
||||
}
|
||||
|
||||
const headers: Record<string, string> = {
|
||||
Authorization: `Bearer ${apiToken}`,
|
||||
"Content-Type": "application/json",
|
||||
};
|
||||
|
||||
if (sync) {
|
||||
headers["Prefer"] = `wait=${SYNC_WAIT_SECONDS}`;
|
||||
}
|
||||
|
||||
const res = await fetch(url, {
|
||||
method: "POST",
|
||||
headers,
|
||||
body: JSON.stringify(body),
|
||||
});
|
||||
|
||||
if (!res.ok) {
|
||||
const err = await res.text();
|
||||
throw new Error(`Replicate API error (${res.status}): ${err}`);
|
||||
}
|
||||
|
||||
return (await res.json()) as PredictionResponse;
|
||||
}
|
||||
|
||||
async function pollPrediction(apiToken: string, getUrl: string): Promise<PredictionResponse> {
|
||||
const start = Date.now();
|
||||
|
||||
while (Date.now() - start < MAX_POLL_MS) {
|
||||
const res = await fetch(getUrl, {
|
||||
headers: { Authorization: `Bearer ${apiToken}` },
|
||||
});
|
||||
|
||||
if (!res.ok) {
|
||||
const err = await res.text();
|
||||
throw new Error(`Replicate poll error (${res.status}): ${err}`);
|
||||
}
|
||||
|
||||
const prediction = (await res.json()) as PredictionResponse;
|
||||
|
||||
if (prediction.status === "succeeded") return prediction;
|
||||
if (prediction.status === "failed" || prediction.status === "canceled") {
|
||||
throw new Error(`Replicate prediction ${prediction.status}: ${prediction.error || "unknown error"}`);
|
||||
}
|
||||
|
||||
await new Promise((r) => setTimeout(r, POLL_INTERVAL_MS));
|
||||
}
|
||||
|
||||
throw new Error(`Replicate prediction timed out after ${MAX_POLL_MS / 1000}s`);
|
||||
}
|
||||
|
||||
export function extractOutputUrl(prediction: PredictionResponse): string {
|
||||
const output = prediction.output;
|
||||
|
||||
if (typeof output === "string") return output;
|
||||
|
||||
if (Array.isArray(output)) {
|
||||
if (output.length !== 1) {
|
||||
throw new Error(
|
||||
`Replicate returned ${output.length} outputs, but baoyu-image-gen currently supports saving exactly one image per request.`
|
||||
);
|
||||
}
|
||||
const first = output[0];
|
||||
if (typeof first === "string") return first;
|
||||
}
|
||||
|
||||
if (output && typeof output === "object" && "url" in output) {
|
||||
const url = (output as Record<string, unknown>).url;
|
||||
if (typeof url === "string") return url;
|
||||
}
|
||||
|
||||
throw new Error(`Unexpected Replicate output format: ${JSON.stringify(output)}`);
|
||||
}
|
||||
|
||||
async function downloadImage(url: string): Promise<Uint8Array> {
|
||||
const res = await fetch(url);
|
||||
if (!res.ok) throw new Error(`Failed to download image from Replicate: ${res.status}`);
|
||||
const buf = await res.arrayBuffer();
|
||||
return new Uint8Array(buf);
|
||||
}
|
||||
|
||||
export async function generateImage(
|
||||
prompt: string,
|
||||
model: string,
|
||||
args: CliArgs
|
||||
): Promise<Uint8Array> {
|
||||
const apiToken = getApiToken();
|
||||
if (!apiToken) throw new Error("REPLICATE_API_TOKEN is required. Get one at https://replicate.com/account/api-tokens");
|
||||
|
||||
const parsedModel = parseModelId(model);
|
||||
validateArgs(model, args);
|
||||
|
||||
const refDataUrls: string[] = [];
|
||||
for (const refPath of args.referenceImages) {
|
||||
refDataUrls.push(await readImageAsDataUrl(refPath));
|
||||
}
|
||||
|
||||
const input = buildInput(model, prompt, args, refDataUrls);
|
||||
|
||||
console.log(`Generating image with Replicate (${model})...`);
|
||||
|
||||
let prediction = await createPrediction(apiToken, parsedModel, input, true);
|
||||
|
||||
if (prediction.status !== "succeeded") {
|
||||
if (!prediction.urls?.get) {
|
||||
throw new Error("Replicate prediction did not return a poll URL");
|
||||
}
|
||||
console.log("Waiting for prediction to complete...");
|
||||
prediction = await pollPrediction(apiToken, prediction.urls.get);
|
||||
}
|
||||
|
||||
console.log("Generation completed.");
|
||||
|
||||
const outputUrl = extractOutputUrl(prediction);
|
||||
return downloadImage(outputUrl);
|
||||
}
|
||||
@@ -0,0 +1,245 @@
|
||||
import assert from "node:assert/strict";
|
||||
import fs from "node:fs/promises";
|
||||
import os from "node:os";
|
||||
import path from "node:path";
|
||||
import test, { type TestContext } from "node:test";
|
||||
|
||||
import type { CliArgs } from "../types.ts";
|
||||
import {
|
||||
buildImageInput,
|
||||
buildRequestBody,
|
||||
generateImage,
|
||||
getDefaultOutputExtension,
|
||||
resolveSeedreamSize,
|
||||
validateArgs,
|
||||
} from "./seedream.ts";
|
||||
|
||||
function makeArgs(overrides: Partial<CliArgs> = {}): CliArgs {
|
||||
return {
|
||||
prompt: null,
|
||||
promptFiles: [],
|
||||
imagePath: null,
|
||||
provider: null,
|
||||
model: null,
|
||||
aspectRatio: null,
|
||||
size: null,
|
||||
quality: null,
|
||||
imageSize: null,
|
||||
imageApiDialect: null,
|
||||
referenceImages: [],
|
||||
n: 1,
|
||||
batchFile: null,
|
||||
jobs: null,
|
||||
json: false,
|
||||
help: false,
|
||||
...overrides,
|
||||
};
|
||||
}
|
||||
|
||||
function useEnv(
|
||||
t: TestContext,
|
||||
values: Record<string, string | null>,
|
||||
): void {
|
||||
const previous = new Map<string, string | undefined>();
|
||||
for (const [key, value] of Object.entries(values)) {
|
||||
previous.set(key, process.env[key]);
|
||||
if (value == null) {
|
||||
delete process.env[key];
|
||||
} else {
|
||||
process.env[key] = value;
|
||||
}
|
||||
}
|
||||
|
||||
t.after(() => {
|
||||
for (const [key, value] of previous.entries()) {
|
||||
if (value == null) {
|
||||
delete process.env[key];
|
||||
} else {
|
||||
process.env[key] = value;
|
||||
}
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
async function makeTempPng(t: TestContext, name: string): Promise<string> {
|
||||
const dir = await fs.mkdtemp(path.join(os.tmpdir(), "seedream-test-"));
|
||||
t.after(() => fs.rm(dir, { recursive: true, force: true }));
|
||||
|
||||
const filePath = path.join(dir, name);
|
||||
const png1x1 =
|
||||
"iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR42mP8/x8AAwMCAO+a7m0AAAAASUVORK5CYII=";
|
||||
await fs.writeFile(filePath, Buffer.from(png1x1, "base64"));
|
||||
return filePath;
|
||||
}
|
||||
|
||||
test("Seedream request body and default extensions follow official model capabilities", () => {
|
||||
const five = buildRequestBody(
|
||||
"A robot illustrator",
|
||||
"doubao-seedream-5-0-260128",
|
||||
makeArgs(),
|
||||
);
|
||||
assert.equal(five.size, "2K");
|
||||
assert.equal(five.response_format, "url");
|
||||
assert.equal(five.output_format, "png");
|
||||
assert.equal(getDefaultOutputExtension("doubao-seedream-5-0-260128"), ".png");
|
||||
|
||||
const fourFive = buildRequestBody(
|
||||
"A robot illustrator",
|
||||
"doubao-seedream-4-5-251128",
|
||||
makeArgs(),
|
||||
);
|
||||
assert.equal(fourFive.size, "2K");
|
||||
assert.equal(fourFive.response_format, "url");
|
||||
assert.ok(!("output_format" in fourFive));
|
||||
assert.equal(getDefaultOutputExtension("doubao-seedream-4-5-251128"), ".jpg");
|
||||
|
||||
assert.throws(
|
||||
() =>
|
||||
buildRequestBody(
|
||||
"Change the bubbles into hearts",
|
||||
"doubao-seededit-3-0-i2i-250628",
|
||||
makeArgs({ referenceImages: ["ref.png"] }),
|
||||
"data:image/png;base64,AAAA",
|
||||
),
|
||||
/no longer supported/,
|
||||
);
|
||||
});
|
||||
|
||||
test("Seedream size selection validates model-specific presets", () => {
|
||||
assert.equal(
|
||||
resolveSeedreamSize("doubao-seedream-4-0-250828", makeArgs({ quality: "normal" })),
|
||||
"1K",
|
||||
);
|
||||
assert.equal(
|
||||
resolveSeedreamSize("doubao-seedream-3-0-t2i-250415", makeArgs({ quality: "2k" })),
|
||||
"2048x2048",
|
||||
);
|
||||
|
||||
assert.throws(
|
||||
() =>
|
||||
resolveSeedreamSize("doubao-seedream-5-0-260128", makeArgs({ size: "4K" })),
|
||||
/only supports 2K, 3K/,
|
||||
);
|
||||
assert.throws(
|
||||
() =>
|
||||
resolveSeedreamSize("doubao-seedream-3-0-t2i-250415", makeArgs({ imageSize: "2K" })),
|
||||
/only supports explicit WxH sizes/,
|
||||
);
|
||||
assert.throws(
|
||||
() =>
|
||||
resolveSeedreamSize("doubao-seededit-3-0-i2i-250628", makeArgs({ size: "1024x1024" })),
|
||||
/no longer supported/,
|
||||
);
|
||||
});
|
||||
|
||||
test("Seedream reference-image support is model-specific", () => {
|
||||
assert.doesNotThrow(() =>
|
||||
validateArgs(
|
||||
"doubao-seedream-5-0-260128",
|
||||
makeArgs({ referenceImages: ["a.png", "b.png"] }),
|
||||
),
|
||||
);
|
||||
|
||||
assert.throws(
|
||||
() =>
|
||||
validateArgs(
|
||||
"doubao-seedream-3-0-t2i-250415",
|
||||
makeArgs({ referenceImages: ["a.png"] }),
|
||||
),
|
||||
/does not support reference images/,
|
||||
);
|
||||
|
||||
assert.throws(
|
||||
() =>
|
||||
validateArgs(
|
||||
"doubao-seededit-3-0-i2i-250628",
|
||||
makeArgs(),
|
||||
),
|
||||
/no longer supported/,
|
||||
);
|
||||
|
||||
assert.throws(
|
||||
() =>
|
||||
validateArgs(
|
||||
"ep-20260315171508-t8br2",
|
||||
makeArgs({ referenceImages: ["a.png"] }),
|
||||
),
|
||||
/require a known model ID/,
|
||||
);
|
||||
});
|
||||
|
||||
test("Seedream image input encodes local references as data URLs", async (t) => {
|
||||
const refOne = await makeTempPng(t, "one.png");
|
||||
const refTwo = await makeTempPng(t, "two.png");
|
||||
|
||||
const single = await buildImageInput("doubao-seedream-4-5-251128", [refOne]);
|
||||
assert.match(String(single), /^data:image\/png;base64,/);
|
||||
|
||||
const multiple = await buildImageInput("doubao-seedream-5-0-260128", [refOne, refTwo]);
|
||||
assert.ok(Array.isArray(multiple));
|
||||
assert.equal(multiple.length, 2);
|
||||
});
|
||||
|
||||
test("Seedream generateImage posts the documented response_format and downloads the returned URL", async (t) => {
|
||||
useEnv(t, { ARK_API_KEY: "test-key", SEEDREAM_BASE_URL: null });
|
||||
|
||||
const originalFetch = globalThis.fetch;
|
||||
t.after(() => {
|
||||
globalThis.fetch = originalFetch;
|
||||
});
|
||||
|
||||
const calls: Array<{
|
||||
input: string;
|
||||
init?: RequestInit;
|
||||
}> = [];
|
||||
|
||||
globalThis.fetch = async (input, init) => {
|
||||
calls.push({
|
||||
input: String(input),
|
||||
init,
|
||||
});
|
||||
|
||||
if (calls.length === 1) {
|
||||
return Response.json({
|
||||
model: "doubao-seedream-4-5-251128",
|
||||
created: 1740000000,
|
||||
data: [
|
||||
{
|
||||
url: "https://example.com/generated-image",
|
||||
size: "2048x2048",
|
||||
},
|
||||
],
|
||||
usage: {
|
||||
generated_images: 1,
|
||||
output_tokens: 1,
|
||||
total_tokens: 1,
|
||||
},
|
||||
});
|
||||
}
|
||||
|
||||
return new Response(Uint8Array.from([7, 8, 9]), {
|
||||
status: 200,
|
||||
headers: { "Content-Type": "image/jpeg" },
|
||||
});
|
||||
};
|
||||
|
||||
const image = await generateImage(
|
||||
"A robot illustrator",
|
||||
"doubao-seedream-4-5-251128",
|
||||
makeArgs(),
|
||||
);
|
||||
|
||||
assert.deepEqual([...image], [7, 8, 9]);
|
||||
assert.equal(calls.length, 2);
|
||||
assert.equal(
|
||||
calls[0]?.input,
|
||||
"https://ark.cn-beijing.volces.com/api/v3/images/generations",
|
||||
);
|
||||
|
||||
const requestBody = JSON.parse(String(calls[0]?.init?.body)) as Record<string, unknown>;
|
||||
assert.equal(requestBody.model, "doubao-seedream-4-5-251128");
|
||||
assert.equal(requestBody.size, "2K");
|
||||
assert.equal(requestBody.response_format, "url");
|
||||
assert.ok(!("output_format" in requestBody));
|
||||
assert.equal(calls[1]?.input, "https://example.com/generated-image");
|
||||
});
|
||||
@@ -0,0 +1,341 @@
|
||||
import path from "node:path";
|
||||
import { readFile } from "node:fs/promises";
|
||||
|
||||
import type { CliArgs } from "../types";
|
||||
|
||||
export type SeedreamModelFamily =
|
||||
| "seedream5"
|
||||
| "seedream45"
|
||||
| "seedream40"
|
||||
| "seedream30"
|
||||
| "unknown";
|
||||
|
||||
type SeedreamRequestImage = string | string[];
|
||||
|
||||
type SeedreamRequestBody = {
|
||||
model: string;
|
||||
prompt: string;
|
||||
size: string;
|
||||
response_format: "url";
|
||||
watermark: boolean;
|
||||
image?: SeedreamRequestImage;
|
||||
output_format?: "png";
|
||||
};
|
||||
|
||||
type SeedreamImageResponse = {
|
||||
model?: string;
|
||||
created?: number;
|
||||
data?: Array<{
|
||||
url?: string;
|
||||
b64_json?: string;
|
||||
size?: string;
|
||||
error?: {
|
||||
code?: string;
|
||||
message?: string;
|
||||
};
|
||||
}>;
|
||||
usage?: {
|
||||
generated_images: number;
|
||||
output_tokens: number;
|
||||
total_tokens: number;
|
||||
};
|
||||
error?: {
|
||||
code?: string;
|
||||
message?: string;
|
||||
};
|
||||
};
|
||||
|
||||
export function getDefaultModel(): string {
|
||||
return process.env.SEEDREAM_IMAGE_MODEL || "doubao-seedream-5-0-260128";
|
||||
}
|
||||
|
||||
function getApiKey(): string | null {
|
||||
return process.env.ARK_API_KEY || null;
|
||||
}
|
||||
|
||||
function getBaseUrl(): string {
|
||||
return process.env.SEEDREAM_BASE_URL || "https://ark.cn-beijing.volces.com/api/v3";
|
||||
}
|
||||
|
||||
function parsePixelSize(value: string): { width: number; height: number } | null {
|
||||
const match = value.trim().match(/^(\d+)\s*[xX]\s*(\d+)$/);
|
||||
if (!match) return null;
|
||||
|
||||
const width = parseInt(match[1]!, 10);
|
||||
const height = parseInt(match[2]!, 10);
|
||||
if (!Number.isFinite(width) || !Number.isFinite(height) || width <= 0 || height <= 0) {
|
||||
return null;
|
||||
}
|
||||
|
||||
return { width, height };
|
||||
}
|
||||
|
||||
function normalizePixelSize(value: string): string | null {
|
||||
const parsed = parsePixelSize(value);
|
||||
if (!parsed) return null;
|
||||
return `${parsed.width}x${parsed.height}`;
|
||||
}
|
||||
|
||||
function normalizeSizePreset(value: string): string | null {
|
||||
const upper = value.trim().toUpperCase();
|
||||
if (upper === "ADAPTIVE") return "adaptive";
|
||||
if (upper === "1K" || upper === "2K" || upper === "3K" || upper === "4K") return upper;
|
||||
return null;
|
||||
}
|
||||
|
||||
function normalizeSizeValue(value: string): string | null {
|
||||
return normalizeSizePreset(value) ?? normalizePixelSize(value);
|
||||
}
|
||||
|
||||
function getMimeType(filename: string): string {
|
||||
const ext = path.extname(filename).toLowerCase();
|
||||
if (ext === ".jpg" || ext === ".jpeg") return "image/jpeg";
|
||||
if (ext === ".webp") return "image/webp";
|
||||
if (ext === ".gif") return "image/gif";
|
||||
if (ext === ".bmp") return "image/bmp";
|
||||
if (ext === ".tiff" || ext === ".tif") return "image/tiff";
|
||||
return "image/png";
|
||||
}
|
||||
|
||||
async function readImageAsDataUrl(filePath: string): Promise<string> {
|
||||
const bytes = await readFile(filePath);
|
||||
return `data:${getMimeType(filePath)};base64,${bytes.toString("base64")}`;
|
||||
}
|
||||
|
||||
export function getModelFamily(model: string): SeedreamModelFamily {
|
||||
const normalized = model.trim();
|
||||
if (/^doubao-seedream-5-0(?:-lite)?-\d+$/.test(normalized)) return "seedream5";
|
||||
if (/^doubao-seedream-4-5-\d+$/.test(normalized)) return "seedream45";
|
||||
if (/^doubao-seedream-4-0-\d+$/.test(normalized)) return "seedream40";
|
||||
if (/^doubao-seedream-3-0-t2i-\d+$/.test(normalized)) return "seedream30";
|
||||
return "unknown";
|
||||
}
|
||||
|
||||
function isRemovedSeededitModel(model: string): boolean {
|
||||
return /^doubao-seededit-3-0-i2i-\d+$/.test(model.trim());
|
||||
}
|
||||
|
||||
function assertSupportedModel(model: string): void {
|
||||
if (isRemovedSeededitModel(model)) {
|
||||
throw new Error(
|
||||
`${model} is no longer supported. SeedEdit 3.0 support has been removed from this tool; use Seedream 5.0/4.5/4.0/3.0 instead.`
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
export function supportsReferenceImages(model: string): boolean {
|
||||
const family = getModelFamily(model);
|
||||
return family === "seedream5" || family === "seedream45" || family === "seedream40";
|
||||
}
|
||||
|
||||
function supportsOutputFormat(model: string): boolean {
|
||||
return getModelFamily(model) === "seedream5";
|
||||
}
|
||||
|
||||
export function getDefaultOutputExtension(model: string): ".png" | ".jpg" {
|
||||
assertSupportedModel(model);
|
||||
return supportsOutputFormat(model) ? ".png" : ".jpg";
|
||||
}
|
||||
|
||||
export function getDefaultSeedreamSize(model: string, args: CliArgs): string {
|
||||
assertSupportedModel(model);
|
||||
const family = getModelFamily(model);
|
||||
|
||||
if (family === "seedream5") return "2K";
|
||||
if (family === "seedream45") return "2K";
|
||||
if (family === "seedream40") return args.quality === "normal" ? "1K" : "2K";
|
||||
if (family === "seedream30") return args.quality === "2k" ? "2048x2048" : "1024x1024";
|
||||
return "2K";
|
||||
}
|
||||
|
||||
export function resolveSeedreamSize(model: string, args: CliArgs): string {
|
||||
assertSupportedModel(model);
|
||||
const family = getModelFamily(model);
|
||||
const requested = args.size || args.imageSize || null;
|
||||
const normalized = requested ? normalizeSizeValue(requested) : null;
|
||||
|
||||
if (!normalized) {
|
||||
return getDefaultSeedreamSize(model, args);
|
||||
}
|
||||
|
||||
if (family === "seedream30") {
|
||||
const pixelSize = normalizePixelSize(normalized);
|
||||
if (!pixelSize) {
|
||||
throw new Error("Seedream 3.0 only supports explicit WxH sizes such as 1024x1024.");
|
||||
}
|
||||
return pixelSize;
|
||||
}
|
||||
|
||||
if (family === "seedream5") {
|
||||
if (normalized === "4K" || normalized === "1K" || normalized === "adaptive") {
|
||||
throw new Error("Seedream 5.0 only supports 2K, 3K, or explicit WxH sizes.");
|
||||
}
|
||||
return normalized;
|
||||
}
|
||||
|
||||
if (family === "seedream45") {
|
||||
if (normalized === "1K" || normalized === "3K" || normalized === "adaptive") {
|
||||
throw new Error("Seedream 4.5 only supports 2K, 4K, or explicit WxH sizes.");
|
||||
}
|
||||
return normalized;
|
||||
}
|
||||
|
||||
if (family === "seedream40") {
|
||||
if (normalized === "3K" || normalized === "adaptive") {
|
||||
throw new Error("Seedream 4.0 only supports 1K, 2K, 4K, or explicit WxH sizes.");
|
||||
}
|
||||
return normalized;
|
||||
}
|
||||
|
||||
if (normalized === "adaptive") {
|
||||
throw new Error("Adaptive size is not supported by Seedream image generation.");
|
||||
}
|
||||
|
||||
if (normalized === "1K" || normalized === "3K" || normalized === "4K") {
|
||||
throw new Error(
|
||||
"Unknown Seedream model ID. Use a documented model ID or pass an explicit WxH size instead of preset imageSize."
|
||||
);
|
||||
}
|
||||
|
||||
return normalized;
|
||||
}
|
||||
|
||||
export function validateArgs(model: string, args: CliArgs): void {
|
||||
assertSupportedModel(model);
|
||||
const family = getModelFamily(model);
|
||||
const refCount = args.referenceImages.length;
|
||||
|
||||
if (refCount === 0) {
|
||||
resolveSeedreamSize(model, args);
|
||||
return;
|
||||
}
|
||||
|
||||
if (family === "unknown") {
|
||||
throw new Error(
|
||||
"Reference images with Seedream require a known model ID. Use Seedream 5.0/4.5/4.0 model IDs instead of an endpoint ID."
|
||||
);
|
||||
}
|
||||
|
||||
if (!supportsReferenceImages(model)) {
|
||||
throw new Error(`${model} does not support reference images.`);
|
||||
}
|
||||
|
||||
if ((family === "seedream5" || family === "seedream45" || family === "seedream40") && refCount > 14) {
|
||||
throw new Error(`${model} supports at most 14 reference images.`);
|
||||
}
|
||||
|
||||
resolveSeedreamSize(model, args);
|
||||
}
|
||||
|
||||
export async function buildImageInput(
|
||||
model: string,
|
||||
referenceImages: string[],
|
||||
): Promise<SeedreamRequestImage | undefined> {
|
||||
if (referenceImages.length === 0) return undefined;
|
||||
assertSupportedModel(model);
|
||||
|
||||
const encoded = await Promise.all(referenceImages.map((refPath) => readImageAsDataUrl(refPath)));
|
||||
|
||||
return encoded.length === 1 ? encoded[0]! : encoded;
|
||||
}
|
||||
|
||||
export function buildRequestBody(
|
||||
prompt: string,
|
||||
model: string,
|
||||
args: CliArgs,
|
||||
imageInput?: SeedreamRequestImage,
|
||||
): SeedreamRequestBody {
|
||||
validateArgs(model, args);
|
||||
|
||||
const requestBody: SeedreamRequestBody = {
|
||||
model,
|
||||
prompt,
|
||||
size: resolveSeedreamSize(model, args),
|
||||
response_format: "url",
|
||||
watermark: false,
|
||||
};
|
||||
|
||||
if (imageInput) {
|
||||
requestBody.image = imageInput;
|
||||
}
|
||||
|
||||
if (supportsOutputFormat(model)) {
|
||||
requestBody.output_format = "png";
|
||||
}
|
||||
|
||||
return requestBody;
|
||||
}
|
||||
|
||||
async function downloadImage(url: string): Promise<Uint8Array> {
|
||||
const imgResponse = await fetch(url);
|
||||
if (!imgResponse.ok) {
|
||||
throw new Error(`Failed to download image from ${url}`);
|
||||
}
|
||||
|
||||
const buffer = await imgResponse.arrayBuffer();
|
||||
return new Uint8Array(buffer);
|
||||
}
|
||||
|
||||
export async function extractImageFromResponse(result: SeedreamImageResponse): Promise<Uint8Array> {
|
||||
const first = result.data?.find((item) => item.url || item.b64_json || item.error);
|
||||
|
||||
if (!first) {
|
||||
throw new Error("No image data in Seedream response");
|
||||
}
|
||||
|
||||
if (first.error) {
|
||||
throw new Error(first.error.message || "Seedream returned an image generation error");
|
||||
}
|
||||
|
||||
if (first.b64_json) {
|
||||
return Uint8Array.from(Buffer.from(first.b64_json, "base64"));
|
||||
}
|
||||
|
||||
if (first.url) {
|
||||
console.error(`Downloading image from ${first.url}...`);
|
||||
return downloadImage(first.url);
|
||||
}
|
||||
|
||||
throw new Error("No image URL or base64 data in Seedream response");
|
||||
}
|
||||
|
||||
export async function generateImage(
|
||||
prompt: string,
|
||||
model: string,
|
||||
args: CliArgs,
|
||||
): Promise<Uint8Array> {
|
||||
const apiKey = getApiKey();
|
||||
if (!apiKey) {
|
||||
throw new Error(
|
||||
"ARK_API_KEY is required. " +
|
||||
"Get your API key from https://console.volcengine.com/ark"
|
||||
);
|
||||
}
|
||||
|
||||
validateArgs(model, args);
|
||||
const imageInput = await buildImageInput(model, args.referenceImages);
|
||||
const requestBody = buildRequestBody(prompt, model, args, imageInput);
|
||||
|
||||
console.error(`Calling Seedream API (${model}) with size: ${requestBody.size}`);
|
||||
|
||||
const response = await fetch(`${getBaseUrl()}/images/generations`, {
|
||||
method: "POST",
|
||||
headers: {
|
||||
"Content-Type": "application/json",
|
||||
Authorization: `Bearer ${apiKey}`,
|
||||
},
|
||||
body: JSON.stringify(requestBody),
|
||||
});
|
||||
|
||||
if (!response.ok) {
|
||||
const err = await response.text();
|
||||
throw new Error(`Seedream API error (${response.status}): ${err}`);
|
||||
}
|
||||
|
||||
const result = (await response.json()) as SeedreamImageResponse;
|
||||
if (result.error) {
|
||||
throw new Error(result.error.message || "Seedream API returned an error");
|
||||
}
|
||||
|
||||
return extractImageFromResponse(result);
|
||||
}
|
||||
@@ -0,0 +1,181 @@
|
||||
import assert from "node:assert/strict";
|
||||
import test, { type TestContext } from "node:test";
|
||||
|
||||
import type { CliArgs } from "../types.ts";
|
||||
import {
|
||||
buildRequestBody,
|
||||
buildZaiUrl,
|
||||
extractImageFromResponse,
|
||||
getDefaultModel,
|
||||
getModelFamily,
|
||||
parseAspectRatio,
|
||||
parseSize,
|
||||
resolveSizeForModel,
|
||||
validateArgs,
|
||||
} from "./zai.ts";
|
||||
|
||||
function makeArgs(overrides: Partial<CliArgs> = {}): CliArgs {
|
||||
return {
|
||||
prompt: null,
|
||||
promptFiles: [],
|
||||
imagePath: null,
|
||||
provider: null,
|
||||
model: null,
|
||||
aspectRatio: null,
|
||||
size: null,
|
||||
quality: null,
|
||||
imageSize: null,
|
||||
imageApiDialect: null,
|
||||
referenceImages: [],
|
||||
n: 1,
|
||||
batchFile: null,
|
||||
jobs: null,
|
||||
json: false,
|
||||
help: false,
|
||||
...overrides,
|
||||
};
|
||||
}
|
||||
|
||||
function useEnv(
|
||||
t: TestContext,
|
||||
values: Record<string, string | null>,
|
||||
): void {
|
||||
const previous = new Map<string, string | undefined>();
|
||||
for (const [key, value] of Object.entries(values)) {
|
||||
previous.set(key, process.env[key]);
|
||||
if (value == null) {
|
||||
delete process.env[key];
|
||||
} else {
|
||||
process.env[key] = value;
|
||||
}
|
||||
}
|
||||
|
||||
t.after(() => {
|
||||
for (const [key, value] of previous.entries()) {
|
||||
if (value == null) {
|
||||
delete process.env[key];
|
||||
} else {
|
||||
process.env[key] = value;
|
||||
}
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
test("Z.AI default model prefers env override and otherwise uses glm-image", (t) => {
|
||||
useEnv(t, {
|
||||
ZAI_IMAGE_MODEL: null,
|
||||
BIGMODEL_IMAGE_MODEL: null,
|
||||
});
|
||||
assert.equal(getDefaultModel(), "glm-image");
|
||||
|
||||
process.env.BIGMODEL_IMAGE_MODEL = "cogview-4-250304";
|
||||
assert.equal(getDefaultModel(), "cogview-4-250304");
|
||||
});
|
||||
|
||||
test("Z.AI URL builder normalizes host, v4 base, and full endpoint inputs", (t) => {
|
||||
useEnv(t, { ZAI_BASE_URL: "https://api.z.ai" });
|
||||
assert.equal(buildZaiUrl(), "https://api.z.ai/api/paas/v4/images/generations");
|
||||
|
||||
process.env.ZAI_BASE_URL = "https://proxy.example.com/api/paas/v4/";
|
||||
assert.equal(buildZaiUrl(), "https://proxy.example.com/api/paas/v4/images/generations");
|
||||
|
||||
process.env.ZAI_BASE_URL = "https://proxy.example.com/custom/images/generations";
|
||||
assert.equal(buildZaiUrl(), "https://proxy.example.com/custom/images/generations");
|
||||
});
|
||||
|
||||
test("Z.AI model family and parsing helpers recognize documented formats", () => {
|
||||
assert.equal(getModelFamily("glm-image"), "glm");
|
||||
assert.equal(getModelFamily("cogview-4-250304"), "legacy");
|
||||
assert.deepEqual(parseAspectRatio("16:9"), { width: 16, height: 9 });
|
||||
assert.equal(parseAspectRatio("wide"), null);
|
||||
assert.deepEqual(parseSize("1280x1280"), { width: 1280, height: 1280 });
|
||||
assert.deepEqual(parseSize("1472*1088"), { width: 1472, height: 1088 });
|
||||
assert.equal(parseSize("big"), null);
|
||||
});
|
||||
|
||||
test("Z.AI size resolution follows documented recommended ratios and validates custom sizes", () => {
|
||||
assert.equal(
|
||||
resolveSizeForModel("glm-image", makeArgs({ aspectRatio: "16:9", quality: "2k" })),
|
||||
"1728x960",
|
||||
);
|
||||
assert.equal(
|
||||
resolveSizeForModel("cogview-4-250304", makeArgs({ aspectRatio: "4:3", quality: "normal" })),
|
||||
"1152x864",
|
||||
);
|
||||
assert.equal(
|
||||
resolveSizeForModel("glm-image", makeArgs({ size: "1568x1056", quality: "2k" })),
|
||||
"1568x1056",
|
||||
);
|
||||
|
||||
const uncommon = resolveSizeForModel(
|
||||
"glm-image",
|
||||
makeArgs({ aspectRatio: "5:2", quality: "normal" }),
|
||||
);
|
||||
const parsed = parseSize(uncommon);
|
||||
assert.ok(parsed);
|
||||
assert.ok(parsed.width % 32 === 0);
|
||||
assert.ok(parsed.height % 32 === 0);
|
||||
assert.ok(parsed.width * parsed.height <= 2 ** 22);
|
||||
|
||||
assert.throws(
|
||||
() => resolveSizeForModel("glm-image", makeArgs({ size: "1000x1000", quality: "2k" })),
|
||||
/between 1024 and 2048/,
|
||||
);
|
||||
assert.throws(
|
||||
() => resolveSizeForModel("glm-image", makeArgs({ size: "1280x1260", quality: "2k" })),
|
||||
/divisible by 32/,
|
||||
);
|
||||
assert.throws(
|
||||
() => resolveSizeForModel("cogview-4-250304", makeArgs({ size: "2048x2048", quality: "2k" })),
|
||||
/must not exceed 2\^21 total pixels/,
|
||||
);
|
||||
});
|
||||
|
||||
test("Z.AI validation rejects unsupported refs and multi-image requests", () => {
|
||||
assert.throws(
|
||||
() => validateArgs("glm-image", makeArgs({ referenceImages: ["ref.png"] })),
|
||||
/text-to-image only/,
|
||||
);
|
||||
assert.throws(
|
||||
() => validateArgs("glm-image", makeArgs({ n: 2 })),
|
||||
/single image per request/,
|
||||
);
|
||||
});
|
||||
|
||||
test("Z.AI request body maps skill quality and resolved size into provider fields", () => {
|
||||
const body = buildRequestBody(
|
||||
"A cinematic science poster",
|
||||
"glm-image",
|
||||
makeArgs({ aspectRatio: "4:3", quality: "normal" }),
|
||||
);
|
||||
|
||||
assert.deepEqual(body, {
|
||||
model: "glm-image",
|
||||
prompt: "A cinematic science poster",
|
||||
quality: "standard",
|
||||
size: "1472x1088",
|
||||
});
|
||||
});
|
||||
|
||||
test("Z.AI response extraction downloads the returned image URL", async (t) => {
|
||||
const originalFetch = globalThis.fetch;
|
||||
t.after(() => {
|
||||
globalThis.fetch = originalFetch;
|
||||
});
|
||||
|
||||
globalThis.fetch = async () =>
|
||||
new Response(Uint8Array.from([1, 2, 3]), {
|
||||
status: 200,
|
||||
headers: { "Content-Type": "image/png" },
|
||||
});
|
||||
|
||||
const image = await extractImageFromResponse({
|
||||
data: [{ url: "https://cdn.example.com/glm-image.png" }],
|
||||
});
|
||||
assert.deepEqual([...image], [1, 2, 3]);
|
||||
|
||||
await assert.rejects(
|
||||
() => extractImageFromResponse({ data: [{}] }),
|
||||
/No image URL/,
|
||||
);
|
||||
});
|
||||
306
baoyu-skills/skills/baoyu-image-gen/scripts/providers/zai.ts
Normal file
306
baoyu-skills/skills/baoyu-image-gen/scripts/providers/zai.ts
Normal file
@@ -0,0 +1,306 @@
|
||||
import type { CliArgs, Quality } from "../types";
|
||||
|
||||
type ZaiModelFamily = "glm" | "legacy";
|
||||
|
||||
type ZaiRequestBody = {
|
||||
model: string;
|
||||
prompt: string;
|
||||
quality: "hd" | "standard";
|
||||
size: string;
|
||||
};
|
||||
|
||||
type ZaiResponse = {
|
||||
data?: Array<{ url?: string }>;
|
||||
};
|
||||
|
||||
const DEFAULT_MODEL = "glm-image";
|
||||
const GLM_MAX_PIXELS = 2 ** 22;
|
||||
const LEGACY_MAX_PIXELS = 2 ** 21;
|
||||
const GLM_SIZE_STEP = 32;
|
||||
const LEGACY_SIZE_STEP = 16;
|
||||
|
||||
const GLM_RECOMMENDED_SIZES: Record<string, string> = {
|
||||
"1:1": "1280x1280",
|
||||
"3:2": "1568x1056",
|
||||
"2:3": "1056x1568",
|
||||
"4:3": "1472x1088",
|
||||
"3:4": "1088x1472",
|
||||
"16:9": "1728x960",
|
||||
"9:16": "960x1728",
|
||||
};
|
||||
|
||||
const LEGACY_RECOMMENDED_SIZES: Record<string, string> = {
|
||||
"1:1": "1024x1024",
|
||||
"9:16": "768x1344",
|
||||
"3:4": "864x1152",
|
||||
"16:9": "1344x768",
|
||||
"4:3": "1152x864",
|
||||
"2:1": "1440x720",
|
||||
"1:2": "720x1440",
|
||||
};
|
||||
|
||||
export function getDefaultModel(): string {
|
||||
return process.env.ZAI_IMAGE_MODEL || process.env.BIGMODEL_IMAGE_MODEL || DEFAULT_MODEL;
|
||||
}
|
||||
|
||||
function getApiKey(): string | null {
|
||||
return process.env.ZAI_API_KEY || process.env.BIGMODEL_API_KEY || null;
|
||||
}
|
||||
|
||||
export function buildZaiUrl(): string {
|
||||
const base = (process.env.ZAI_BASE_URL || process.env.BIGMODEL_BASE_URL || "https://api.z.ai/api/paas/v4")
|
||||
.replace(/\/+$/g, "");
|
||||
if (base.endsWith("/images/generations")) return base;
|
||||
if (base.endsWith("/api/paas/v4")) return `${base}/images/generations`;
|
||||
if (base.endsWith("/v4")) return `${base}/images/generations`;
|
||||
return `${base}/api/paas/v4/images/generations`;
|
||||
}
|
||||
|
||||
export function getModelFamily(model: string): ZaiModelFamily {
|
||||
return model.trim().toLowerCase() === "glm-image" ? "glm" : "legacy";
|
||||
}
|
||||
|
||||
export function parseAspectRatio(ar: string): { width: number; height: number } | null {
|
||||
const match = ar.match(/^(\d+(?:\.\d+)?):(\d+(?:\.\d+)?)$/);
|
||||
if (!match) return null;
|
||||
const width = Number(match[1]);
|
||||
const height = Number(match[2]);
|
||||
if (!Number.isFinite(width) || !Number.isFinite(height) || width <= 0 || height <= 0) {
|
||||
return null;
|
||||
}
|
||||
return { width, height };
|
||||
}
|
||||
|
||||
export function parseSize(size: string): { width: number; height: number } | null {
|
||||
const match = size.trim().match(/^(\d+)\s*[xX*]\s*(\d+)$/);
|
||||
if (!match) return null;
|
||||
const width = parseInt(match[1]!, 10);
|
||||
const height = parseInt(match[2]!, 10);
|
||||
if (!Number.isFinite(width) || !Number.isFinite(height) || width <= 0 || height <= 0) {
|
||||
return null;
|
||||
}
|
||||
return { width, height };
|
||||
}
|
||||
|
||||
function formatSize(width: number, height: number): string {
|
||||
return `${width}x${height}`;
|
||||
}
|
||||
|
||||
function roundToStep(value: number, step: number): number {
|
||||
return Math.max(step, Math.round(value / step) * step);
|
||||
}
|
||||
|
||||
function getRatioValue(ar: string): number | null {
|
||||
const parsed = parseAspectRatio(ar);
|
||||
if (!parsed) return null;
|
||||
return parsed.width / parsed.height;
|
||||
}
|
||||
|
||||
function findClosestRatioKey(ar: string, candidates: string[]): string | null {
|
||||
const targetRatio = getRatioValue(ar);
|
||||
if (targetRatio == null) return null;
|
||||
|
||||
let bestKey: string | null = null;
|
||||
let bestDiff = Infinity;
|
||||
for (const candidate of candidates) {
|
||||
const candidateRatio = getRatioValue(candidate);
|
||||
if (candidateRatio == null) continue;
|
||||
const diff = Math.abs(candidateRatio - targetRatio);
|
||||
if (diff < bestDiff) {
|
||||
bestDiff = diff;
|
||||
bestKey = candidate;
|
||||
}
|
||||
}
|
||||
|
||||
return bestDiff <= 0.05 ? bestKey : null;
|
||||
}
|
||||
|
||||
function getTargetPixels(quality: Quality): number {
|
||||
return quality === "normal" ? 1024 * 1024 : 1536 * 1536;
|
||||
}
|
||||
|
||||
function fitToPixelBudget(
|
||||
width: number,
|
||||
height: number,
|
||||
targetPixels: number,
|
||||
maxPixels: number,
|
||||
step: number,
|
||||
): { width: number; height: number } {
|
||||
let nextWidth = width;
|
||||
let nextHeight = height;
|
||||
const pixels = nextWidth * nextHeight;
|
||||
|
||||
if (pixels > maxPixels) {
|
||||
const scale = Math.sqrt(maxPixels / pixels);
|
||||
nextWidth *= scale;
|
||||
nextHeight *= scale;
|
||||
} else {
|
||||
const scale = Math.sqrt(targetPixels / pixels);
|
||||
nextWidth *= scale;
|
||||
nextHeight *= scale;
|
||||
}
|
||||
|
||||
let roundedWidth = roundToStep(nextWidth, step);
|
||||
let roundedHeight = roundToStep(nextHeight, step);
|
||||
let roundedPixels = roundedWidth * roundedHeight;
|
||||
|
||||
while (roundedPixels > maxPixels && (roundedWidth > step || roundedHeight > step)) {
|
||||
if (roundedWidth >= roundedHeight && roundedWidth > step) {
|
||||
roundedWidth -= step;
|
||||
} else if (roundedHeight > step) {
|
||||
roundedHeight -= step;
|
||||
} else {
|
||||
break;
|
||||
}
|
||||
roundedPixels = roundedWidth * roundedHeight;
|
||||
}
|
||||
|
||||
return { width: roundedWidth, height: roundedHeight };
|
||||
}
|
||||
|
||||
function validateCustomSize(
|
||||
size: string,
|
||||
family: ZaiModelFamily,
|
||||
): string {
|
||||
const parsed = parseSize(size);
|
||||
if (!parsed) {
|
||||
throw new Error("Z.AI --size must be in WxH format, for example 1280x1280.");
|
||||
}
|
||||
|
||||
const widthStep = family === "glm" ? GLM_SIZE_STEP : LEGACY_SIZE_STEP;
|
||||
const minEdge = family === "glm" ? 1024 : 512;
|
||||
const maxPixels = family === "glm" ? GLM_MAX_PIXELS : LEGACY_MAX_PIXELS;
|
||||
|
||||
if (parsed.width < minEdge || parsed.width > 2048 || parsed.height < minEdge || parsed.height > 2048) {
|
||||
throw new Error(
|
||||
family === "glm"
|
||||
? "GLM-image custom size requires width and height between 1024 and 2048."
|
||||
: "Z.AI legacy image models require width and height between 512 and 2048."
|
||||
);
|
||||
}
|
||||
|
||||
if (parsed.width % widthStep !== 0 || parsed.height % widthStep !== 0) {
|
||||
throw new Error(
|
||||
family === "glm"
|
||||
? "GLM-image custom size requires width and height divisible by 32."
|
||||
: "Z.AI legacy image models require width and height divisible by 16."
|
||||
);
|
||||
}
|
||||
|
||||
if (parsed.width * parsed.height > maxPixels) {
|
||||
throw new Error(
|
||||
family === "glm"
|
||||
? "GLM-image custom size must not exceed 2^22 total pixels."
|
||||
: "Z.AI legacy image size must not exceed 2^21 total pixels."
|
||||
);
|
||||
}
|
||||
|
||||
return formatSize(parsed.width, parsed.height);
|
||||
}
|
||||
|
||||
export function resolveSizeForModel(
|
||||
model: string,
|
||||
args: Pick<CliArgs, "size" | "aspectRatio" | "quality">,
|
||||
): string {
|
||||
const family = getModelFamily(model);
|
||||
const quality = args.quality === "normal" ? "normal" : "2k";
|
||||
|
||||
if (args.size) {
|
||||
return validateCustomSize(args.size, family);
|
||||
}
|
||||
|
||||
const recommended = family === "glm" ? GLM_RECOMMENDED_SIZES : LEGACY_RECOMMENDED_SIZES;
|
||||
const defaultSize = family === "glm" ? "1280x1280" : "1024x1024";
|
||||
|
||||
if (!args.aspectRatio) return defaultSize;
|
||||
|
||||
const recommendedRatio = findClosestRatioKey(args.aspectRatio, Object.keys(recommended));
|
||||
if (recommendedRatio) {
|
||||
return recommended[recommendedRatio]!;
|
||||
}
|
||||
|
||||
const parsedRatio = parseAspectRatio(args.aspectRatio);
|
||||
if (!parsedRatio) return defaultSize;
|
||||
|
||||
const targetPixels = getTargetPixels(quality);
|
||||
const maxPixels = family === "glm" ? GLM_MAX_PIXELS : LEGACY_MAX_PIXELS;
|
||||
const step = family === "glm" ? GLM_SIZE_STEP : LEGACY_SIZE_STEP;
|
||||
const fit = fitToPixelBudget(
|
||||
parsedRatio.width,
|
||||
parsedRatio.height,
|
||||
targetPixels,
|
||||
maxPixels,
|
||||
step,
|
||||
);
|
||||
return formatSize(fit.width, fit.height);
|
||||
}
|
||||
|
||||
function getZaiQuality(quality: CliArgs["quality"]): "hd" | "standard" {
|
||||
return quality === "normal" ? "standard" : "hd";
|
||||
}
|
||||
|
||||
export function validateArgs(_model: string, args: CliArgs): void {
|
||||
if (args.referenceImages.length > 0) {
|
||||
throw new Error("Z.AI GLM-image currently supports text-to-image only in baoyu-image-gen. Remove --ref or choose another provider.");
|
||||
}
|
||||
|
||||
if (args.n > 1) {
|
||||
throw new Error("Z.AI image generation currently returns a single image per request in baoyu-image-gen.");
|
||||
}
|
||||
}
|
||||
|
||||
export function buildRequestBody(
|
||||
prompt: string,
|
||||
model: string,
|
||||
args: CliArgs,
|
||||
): ZaiRequestBody {
|
||||
validateArgs(model, args);
|
||||
return {
|
||||
model,
|
||||
prompt,
|
||||
quality: getZaiQuality(args.quality),
|
||||
size: resolveSizeForModel(model, args),
|
||||
};
|
||||
}
|
||||
|
||||
export async function extractImageFromResponse(result: ZaiResponse): Promise<Uint8Array> {
|
||||
const url = result.data?.[0]?.url;
|
||||
if (!url) {
|
||||
throw new Error("No image URL in Z.AI response");
|
||||
}
|
||||
|
||||
const imageResponse = await fetch(url);
|
||||
if (!imageResponse.ok) {
|
||||
throw new Error(`Failed to download image from Z.AI: ${imageResponse.status}`);
|
||||
}
|
||||
|
||||
return new Uint8Array(await imageResponse.arrayBuffer());
|
||||
}
|
||||
|
||||
export async function generateImage(
|
||||
prompt: string,
|
||||
model: string,
|
||||
args: CliArgs,
|
||||
): Promise<Uint8Array> {
|
||||
const apiKey = getApiKey();
|
||||
if (!apiKey) {
|
||||
throw new Error("ZAI_API_KEY is required. Get one from https://docs.z.ai/.");
|
||||
}
|
||||
|
||||
const response = await fetch(buildZaiUrl(), {
|
||||
method: "POST",
|
||||
headers: {
|
||||
"Content-Type": "application/json",
|
||||
Authorization: `Bearer ${apiKey}`,
|
||||
},
|
||||
body: JSON.stringify(buildRequestBody(prompt, model, args)),
|
||||
});
|
||||
|
||||
if (!response.ok) {
|
||||
const err = await response.text();
|
||||
throw new Error(`Z.AI API error (${response.status}): ${err}`);
|
||||
}
|
||||
|
||||
const result = (await response.json()) as ZaiResponse;
|
||||
return extractImageFromResponse(result);
|
||||
}
|
||||
Reference in New Issue
Block a user