Add Cloudflare Workers AI image generation (#973)
This commit is contained in:
@@ -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" },
|
||||
|
||||
@@ -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) {
|
||||
|
||||
178
open-sse/handlers/imageProviders/cloudflareAi.js
Normal file
178
open-sse/handlers/imageProviders/cloudflareAi.js
Normal 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,
|
||||
};
|
||||
@@ -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) {
|
||||
|
||||
Reference in New Issue
Block a user