Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 11 additions & 0 deletions open-sse/config/providerModels.js
Original file line number Diff line number Diff line change
Expand Up @@ -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" },
Expand Down
31 changes: 25 additions & 6 deletions open-sse/handlers/imageGenerationCore.js
Original file line number Diff line number Diff line change
Expand Up @@ -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`.
Expand Down Expand Up @@ -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)}..."`);

Expand All @@ -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);
Expand Down Expand Up @@ -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`);
Expand All @@ -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) {
Expand Down
178 changes: 178 additions & 0 deletions open-sse/handlers/imageProviders/cloudflareAi.js
Original file line number Diff line number Diff line change
@@ -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,
};
2 changes: 2 additions & 0 deletions open-sse/handlers/imageProviders/index.js
Original file line number Diff line number Diff line change
Expand Up @@ -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"),
Expand All @@ -26,6 +27,7 @@ const ADAPTERS = {
"stability-ai": stabilityAi,
"black-forest-labs": blackForestLabs,
runwayml,
"cloudflare-ai": cloudflareAi,
};

export function getImageAdapter(provider) {
Expand Down
Loading