Add Cloudflare Workers AI image generation (#973)

This commit is contained in:
Aleksei
2026-05-09 05:53:39 +03:00
committed by GitHub
parent dd15d162fc
commit 787d248030
7 changed files with 433 additions and 18 deletions

View File

@@ -393,6 +393,17 @@ export const PROVIDER_MODELS = {
{ id: "@cf/zai-org/glm-4.7-flash", name: "GLM 4.7 Flash" },
{ id: "@cf/qwen/qwq-32b", name: "QwQ 32B" },
{ id: "@cf/qwen/qwen2.5-coder-32b-instruct", name: "Qwen 2.5 Coder 32B Instruct" },
{ id: "@cf/black-forest-labs/flux-2-klein-9b", name: "FLUX.2 Klein 9B", type: "image", params: ["size"] },
{ id: "@cf/black-forest-labs/flux-2-klein-4b", name: "FLUX.2 Klein 4B", type: "image", params: ["size"] },
{ id: "@cf/black-forest-labs/flux-2-dev", name: "FLUX.2 Dev", type: "image", params: ["size"] },
{ id: "@cf/leonardo/lucid-origin", name: "Lucid Origin", type: "image", params: ["size"] },
{ id: "@cf/leonardo/phoenix-1.0", name: "Phoenix 1.0", type: "image", params: ["size"] },
{ id: "@cf/black-forest-labs/flux-1-schnell", name: "FLUX.1 Schnell", type: "image", params: ["size"] },
{ id: "@cf/bytedance/stable-diffusion-xl-lightning", name: "SDXL Lightning", type: "image", params: ["size"] },
{ id: "@cf/lykon/dreamshaper-8-lcm", name: "DreamShaper 8 LCM", type: "image", params: ["size"] },
{ id: "@cf/runwayml/stable-diffusion-v1-5-img2img", name: "Stable Diffusion v1.5 Img2Img", type: "image", params: ["size"], capabilities: ["edit"] },
{ id: "@cf/runwayml/stable-diffusion-v1-5-inpainting", name: "Stable Diffusion v1.5 Inpainting", type: "image", params: ["size"], capabilities: ["edit", "mask"] },
{ id: "@cf/stabilityai/stable-diffusion-xl-base-1.0", name: "SDXL Base 1.0", type: "image", params: ["size"] },
],
byteplus: [
{ id: "seed-2-0-pro-260328", name: "Seed 2.0 Pro" },

View File

@@ -5,6 +5,12 @@ import { getExecutor } from "../executors/index.js";
import { getImageAdapter } from "./imageProviders/index.js";
import { urlToBase64 } from "./imageProviders/_base.js";
function serializeRequestBody(requestBody) {
if (typeof FormData !== "undefined" && requestBody instanceof FormData) return requestBody;
if (typeof requestBody === "string") return requestBody;
return JSON.stringify(requestBody);
}
/**
* Core image generation handler — orchestrator only.
* Provider-specific URL/headers/body/parse/normalize live in `./imageProviders/{id}.js`.
@@ -44,9 +50,17 @@ export async function handleImageGenerationCore({
);
}
const url = adapter.buildUrl(model, credentials);
const headers = adapter.buildHeaders(credentials);
const requestBody = adapter.buildBody(model, body);
let url;
let headers;
let requestBody;
try {
url = adapter.buildUrl(model, credentials);
requestBody = await adapter.buildBody(model, body);
headers = adapter.buildHeaders(credentials, requestBody, model, body);
} catch (error) {
return createErrorResult(HTTP_STATUS.BAD_REQUEST, error.message || `Invalid ${provider} image request`);
}
log?.debug?.("IMAGE", `${provider.toUpperCase()} | ${model} | prompt="${body.prompt.slice(0, 50)}..."`);
@@ -55,7 +69,7 @@ export async function handleImageGenerationCore({
providerResponse = await fetch(url, {
method: "POST",
headers,
body: JSON.stringify(requestBody),
body: serializeRequestBody(requestBody),
});
} catch (error) {
const errMsg = formatProviderError(error, provider, model, HTTP_STATUS.BAD_GATEWAY);
@@ -83,12 +97,13 @@ export async function handleImageGenerationCore({
if (onCredentialsRefreshed) await onCredentialsRefreshed(newCredentials);
try {
const retryHeaders = adapter.buildHeaders(credentials);
const retryBody = await adapter.buildBody(model, body);
const retryHeaders = adapter.buildHeaders(credentials, retryBody, model, body);
const retryUrl = adapter.buildUrl(model, credentials);
providerResponse = await fetch(retryUrl, {
method: "POST",
headers: retryHeaders,
body: JSON.stringify(requestBody),
body: serializeRequestBody(retryBody),
});
} catch {
log?.warn?.("TOKEN", `${provider.toUpperCase()} | retry after refresh failed`);
@@ -114,6 +129,10 @@ export async function handleImageGenerationCore({
log,
streamToClient,
onRequestSuccess,
url,
requestBody,
model,
body,
});
// Codex streaming case: returns an SSE Response directly
if (parsed?.sseResponse) {

View File

@@ -0,0 +1,178 @@
import { nowSec, urlToBase64 } from "./_base.js";
const BASE_URL = "https://api.cloudflare.com/client/v4/accounts";
const MULTIPART_MODELS = new Set([
"@cf/black-forest-labs/flux-2-dev",
"@cf/black-forest-labs/flux-2-klein-4b",
"@cf/black-forest-labs/flux-2-klein-9b",
]);
const OPTIONAL_FIELDS = [
"negative_prompt",
"guidance",
"seed",
"num_steps",
"steps",
"strength",
];
function sizeToDimensions(size) {
const match = /^(\d+)x(\d+)$/.exec(String(size || ""));
if (!match) return {};
return {
width: Number(match[1]),
height: Number(match[2]),
};
}
function getDimensions(body) {
return {
...sizeToDimensions(body.size),
...(Number.isFinite(Number(body.width)) ? { width: Number(body.width) } : {}),
...(Number.isFinite(Number(body.height)) ? { height: Number(body.height) } : {}),
};
}
async function resolveImageInput(value) {
if (Array.isArray(value)) {
return { bytes: value, b64: Buffer.from(value).toString("base64") };
}
if (typeof value !== "string") return null;
const trimmed = value.trim();
if (!trimmed) return null;
if (/^https?:\/\//i.test(trimmed)) {
const b64 = await urlToBase64(trimmed);
return { bytes: base64ToBytes(b64), b64 };
}
const match = /^data:image\/[^;]+;base64,(.+)$/i.exec(trimmed);
const b64 = match ? match[1] : trimmed;
return { bytes: base64ToBytes(b64), b64 };
}
function base64ToBytes(value) {
try {
return Array.from(Buffer.from(value, "base64"));
} catch {
return value;
}
}
function addOptionalFields(target, body, append) {
for (const key of OPTIONAL_FIELDS) {
const value = body[key];
if (value === undefined || value === null || value === "") continue;
append(target, key, value);
}
}
async function buildJsonBody(body) {
const req = { prompt: body.prompt, ...getDimensions(body) };
addOptionalFields(req, body, (target, key, value) => {
target[key] = value;
});
const imageData = await resolveImageInput(body.image);
if (imageData) {
req.image_b64 = imageData.b64;
req.image = imageData.bytes;
}
const maskData = await resolveImageInput(body.mask_image || body.maskImage || body.mask);
if (maskData) {
req.mask_b64 = maskData.b64;
req.mask = maskData.bytes;
req.mask_image = maskData.bytes;
}
return req;
}
function buildMultipartBody(body) {
const form = new FormData();
form.append("prompt", body.prompt);
const dimensions = getDimensions(body);
for (const [key, value] of Object.entries(dimensions)) {
form.append(key, String(value));
}
addOptionalFields(form, body, (target, key, value) => {
target.append(key, String(value));
});
return form;
}
function imageItemFromString(value) {
if (typeof value !== "string" || !value) return null;
if (/^data:image\/[^;]+;base64,/i.test(value)) {
return { b64_json: value.replace(/^data:image\/[^;]+;base64,/i, "") };
}
if (/^https?:\/\//i.test(value)) return { url: value };
return { b64_json: value };
}
function normalizeCloudflareResponse(responseBody) {
if (responseBody?.created && Array.isArray(responseBody?.data)) return responseBody;
const result = responseBody?.result ?? responseBody;
const queuedResponse = Array.isArray(result?.responses)
? result.responses.find((item) => item?.success !== false)?.result
: null;
if (queuedResponse) return normalizeCloudflareResponse(queuedResponse);
const image =
(typeof result === "string" ? result : null) ||
result?.image ||
result?.data?.[0]?.b64_json ||
result?.data?.[0]?.url;
const item = imageItemFromString(image);
return {
created: nowSec(),
data: item ? [item] : [],
};
}
export default {
buildUrl: (model, creds) => {
const accountId = creds?.providerSpecificData?.accountId;
if (!accountId) throw new Error("cloudflare-ai requires accountId in providerSpecificData");
return `${BASE_URL}/${accountId}/ai/run/${model}`;
},
buildHeaders: (creds, requestBody) => {
const headers = {};
const isMultipart = typeof FormData !== "undefined" && requestBody instanceof FormData;
if (!isMultipart) {
headers["Content-Type"] = "application/json";
}
const key = creds?.apiKey || creds?.accessToken;
if (key) headers.Authorization = `Bearer ${key}`;
return headers;
},
buildBody: async (model, body) => (
MULTIPART_MODELS.has(model)
? buildMultipartBody(body)
: await buildJsonBody(body)
),
async parseResponse(response) {
const contentType = (response.headers.get("Content-Type") || "").toLowerCase();
if (contentType.startsWith("image/")) {
const buf = await response.arrayBuffer();
return {
created: nowSec(),
data: [{ b64_json: Buffer.from(buf).toString("base64") }],
};
}
const json = await response.json();
return normalizeCloudflareResponse(json);
},
normalize: normalizeCloudflareResponse,
};

View File

@@ -10,6 +10,7 @@ import falAi from "./falAi.js";
import stabilityAi from "./stabilityAi.js";
import blackForestLabs from "./blackForestLabs.js";
import runwayml from "./runwayml.js";
import cloudflareAi from "./cloudflareAi.js";
const ADAPTERS = {
openai: createOpenAIAdapter("openai"),
@@ -26,6 +27,7 @@ const ADAPTERS = {
"stability-ai": stabilityAi,
"black-forest-labs": blackForestLabs,
runwayml,
"cloudflare-ai": cloudflareAi,
};
export function getImageAdapter(provider) {