fix(image): support Fal reference-image edits (#9933)

Co-authored-by: rinseaid <rinseaid@rinseaid.net>
Co-authored-by: rinseaid <rinseaid@users.noreply.github.com>
This commit is contained in:
rinseaid
2026-08-10 02:54:13 -04:00
committed by GitHub
parent 0cb7410ca6
commit 61014aec52
4 changed files with 248 additions and 10 deletions

View File

@@ -2149,7 +2149,7 @@ function parseSizeToDimensions(size, fallback = 1024) {
};
}
function normalizeRequestedImageFormat(
export function normalizeRequestedImageFormat(
body,
fallback = "png",
allowedFormats = ["jpeg", "png", "webp"]
@@ -2169,7 +2169,7 @@ function normalizeRequestedImageFormat(
return fallback;
}
function mapFalImageSize(size, fallback = "square_hd") {
export function mapFalImageSize(size, fallback = "square_hd") {
if (typeof size !== "string") return fallback;
if (FAL_PRESET_SIZES[size]) return FAL_PRESET_SIZES[size];
if (size.includes("x")) {
@@ -2200,7 +2200,7 @@ function shouldIncludeStabilityMask(model) {
]).has(model);
}
async function normalizeProviderImagePayload(payload, body, log, defaultFormat) {
export async function normalizeProviderImagePayload(payload, body, log, defaultFormat) {
const candidates = [];
const pushCandidate = (value) => {

View File

@@ -0,0 +1,115 @@
import type { ExecutorLog, ProviderCredentials } from "../../../executors/base.ts";
import {
mapFalImageSize,
normalizeProviderImagePayload,
normalizeRequestedImageFormat,
saveImageErrorResult,
saveImageSuccessResult,
} from "../../imageGeneration.ts";
import { sanitizeErrorMessage } from "../../../utils/error.ts";
export const FAL_IMAGE_EDIT_MODELS = new Set([
"fal-ai/flux-2-flex",
"fal-ai/flux-2-pro",
"fal-ai/flux-2-max",
]);
export const FAL_IMAGE_EDIT_MAX_REFERENCES = 10;
export function isFalImageEditModel(model: string | null): boolean {
return typeof model === "string" && FAL_IMAGE_EDIT_MODELS.has(model);
}
type FalAIImageEditOptions = {
model: string;
provider: string;
providerConfig: { baseUrl: string };
body: Record<string, unknown>;
images: Array<{ bytes: Buffer; mime: string }>;
credentials: ProviderCredentials;
log: ExecutorLog | null | undefined;
};
export async function handleFalAIImageEdit({
model,
provider,
providerConfig,
body,
images,
credentials,
log,
}: FalAIImageEditOptions) {
const startTime = Date.now();
const editModel = `${model}/edit`;
const outputFormat = normalizeRequestedImageFormat(body, "png", ["jpeg", "png"]);
const upstreamBody: Record<string, unknown> = {
prompt: body.prompt,
image_urls: images.map(
({ bytes, mime }) => `data:${mime || "image/png"};base64,${bytes.toString("base64")}`
),
image_size: mapFalImageSize(body.size, "auto"),
output_format: outputFormat,
sync_mode: body.sync_mode ?? true,
};
if (body.n !== undefined) upstreamBody.num_images = Number(body.n) || 1;
if (body.seed !== undefined) upstreamBody.seed = body.seed;
if (log) {
const promptPreview = String(body.prompt ?? "").slice(0, 60);
log.info("IMAGE", `${provider}/${editModel} (fal-ai edit) | prompt: "${promptPreview}..."`);
}
try {
const token = credentials.apiKey || credentials.accessToken;
const response = await fetch(`${providerConfig.baseUrl.replace(/\/$/, "")}/${editModel}`, {
method: "POST",
headers: {
"Content-Type": "application/json",
Authorization: `Key ${token}`,
},
body: JSON.stringify(upstreamBody),
});
if (!response.ok) {
const errorText = await response.text();
if (log)
log.error("IMAGE", `${provider} error ${response.status}: ${errorText.slice(0, 200)}`);
return saveImageErrorResult({
provider,
model: editModel,
status: response.status,
startTime,
error: errorText,
requestBody: upstreamBody,
path: "/v1/images/edits",
});
}
const payload = await response.json();
const normalizedBody =
body.response_format === undefined ? { ...body, response_format: "b64_json" } : body;
const imagesOut = await normalizeProviderImagePayload(payload, normalizedBody, log);
return saveImageSuccessResult({
provider,
model: editModel,
startTime,
requestBody: upstreamBody,
responseBody: { images_count: imagesOut.length },
created: payload.created,
images: imagesOut,
path: "/v1/images/edits",
});
} catch (err) {
const message = err instanceof Error ? err.message : String(err);
if (log) log.error("IMAGE", `${provider} fetch error: ${message}`);
return saveImageErrorResult({
provider,
model: editModel,
status: 502,
startTime,
error: `Image provider error: ${sanitizeErrorMessage(message || err)}`,
path: "/v1/images/edits",
});
}
}

View File

@@ -4,6 +4,11 @@ import {
handleImageEdit,
handleOpenAIImageEdit,
} from "@omniroute/open-sse/handlers/imageGeneration.ts";
import {
handleFalAIImageEdit,
FAL_IMAGE_EDIT_MAX_REFERENCES,
isFalImageEditModel,
} from "@omniroute/open-sse/handlers/imageGeneration/providers/fal.ts";
import { createInjectionGuard } from "@/middleware/promptInjectionGuard";
import {
getProviderCredentialsWithQuotaPreflight,
@@ -207,7 +212,8 @@ function buildAdobeFireflyEditDataUrls(
}
}
if (dataUrls.length === 0 && imageBytes && imageBytes.length > 0) {
const mime = typeof imageMime === "string" && imageMime.startsWith("image/") ? imageMime : "image/png";
const mime =
typeof imageMime === "string" && imageMime.startsWith("image/") ? imageMime : "image/png";
dataUrls.push(`data:${mime};base64,${imageBytes.toString("base64")}`);
}
return dataUrls;
@@ -250,7 +256,10 @@ async function handleAdobeFireflyEditRequest(params: {
resolvedModel
);
if (!credentials) {
return errorResponse(HTTP_STATUS.UNAUTHORIZED, `No credentials for provider: ${parsed.provider}`);
return errorResponse(
HTTP_STATUS.UNAUTHORIZED,
`No credentials for provider: ${parsed.provider}`
);
}
if (credentials.allRateLimited) {
return unavailableResponse(
@@ -362,11 +371,10 @@ async function postHandler(request: Request, _context?: unknown) {
? 4
: providerConfig?.format === "codex-responses"
? Number.POSITIVE_INFINITY
: MAX_NON_CODEX_IMAGE_EDIT_REFERENCES;
if (
providerConfig?.format !== "codex-responses" &&
imageInputCount > maxRefsForProvider
) {
: providerConfig?.format === "fal-ai" && isFalImageEditModel(parsed.model)
? FAL_IMAGE_EDIT_MAX_REFERENCES
: MAX_NON_CODEX_IMAGE_EDIT_REFERENCES;
if (providerConfig?.format !== "codex-responses" && imageInputCount > maxRefsForProvider) {
return errorResponse(
HTTP_STATUS.BAD_REQUEST,
providerConfig?.format === "adobe-firefly-image"
@@ -514,6 +522,53 @@ async function postHandler(request: Request, _context?: unknown) {
);
}
if (providerConfig?.format === "fal-ai" && isFalImageEditModel(parsed.model)) {
const credentials = await getProviderCredentialsWithQuotaPreflight(
parsed.provider,
null,
allowedConnections,
resolvedModel
);
if (!credentials) {
return errorResponse(
HTTP_STATUS.UNAUTHORIZED,
`No credentials for provider: ${parsed.provider}`
);
}
if (credentials.allRateLimited) {
return unavailableResponse(
HTTP_STATUS.RATE_LIMITED,
`[${parsed.provider}] All accounts rate limited`,
credentials.retryAfter,
credentials.retryAfterHuman
);
}
const result = await handleFalAIImageEdit({
provider: parsed.provider,
model: parsed.model,
providerConfig,
body: {
prompt,
size: size ?? undefined,
response_format: responseFormat ?? undefined,
n: 1,
},
images,
credentials,
log,
});
if (result.success) {
await clearRecoveredProviderState(credentials);
return jsonResponse(result.data);
}
return jsonResponse(
toJsonErrorPayload(result.error, "Image edit provider error"),
result.status
);
}
// Adobe Firefly: edit = storage upload + generate-async referenceBlobs (same as i2i generate).
if (providerConfig?.format === "adobe-firefly-image") {
return handleAdobeFireflyEditRequest({

View File

@@ -0,0 +1,68 @@
import test from "node:test";
import assert from "node:assert/strict";
import dns from "node:dns";
import { mkdtempSync } from "node:fs";
import { tmpdir } from "node:os";
import { join } from "node:path";
process.env.DATA_DIR = mkdtempSync(join(tmpdir(), "omniroute-fal-images-"));
const originalDnsLookup = dns.promises.lookup;
(dns.promises as { lookup: unknown }).lookup = (async (
_hostname: string,
options?: { all?: boolean }
) => {
const record = { address: "203.0.113.1", family: 4 };
return options?.all ? [record] : record;
}) as typeof dns.promises.lookup;
process.on("exit", () => {
(dns.promises as { lookup: unknown }).lookup = originalDnsLookup;
});
const { handleFalAIImageEdit } =
await import("../../open-sse/handlers/imageGeneration/providers/fal.ts");
test("handleFalAIImageEdit forwards multiple references to the Fal edit endpoint", async () => {
const originalFetch = globalThis.fetch;
let captured;
globalThis.fetch = async (url, options = {}) => {
const stringUrl = String(url);
if (stringUrl === "https://fal.run/fal-ai/flux-2-flex/edit") {
captured = {
headers: options.headers,
body: JSON.parse(String(options.body || "{}")),
};
return new Response(JSON.stringify({ images: [{ url: "data:image/png;base64,CAkK" }] }), {
status: 200,
headers: { "content-type": "application/json" },
});
}
throw new Error(`Unexpected URL: ${stringUrl}`);
};
try {
const result = await handleFalAIImageEdit({
model: "fal-ai/flux-2-flex",
provider: "fal-ai",
providerConfig: { baseUrl: "https://fal.run" },
body: { prompt: "make the dog match the reference" },
images: [
{ bytes: Buffer.from([1, 2, 3]), mime: "image/png" },
{ bytes: Buffer.from([4, 5, 6]), mime: "image/jpeg" },
],
credentials: { apiKey: "fal-key" },
log: null,
});
assert.equal(result.success, true);
assert.equal(captured.headers.Authorization, "Key fal-key");
assert.deepEqual(captured.body.image_urls, [
"data:image/png;base64,AQID",
"data:image/jpeg;base64,BAUG",
]);
assert.equal(captured.body.prompt, "make the dog match the reference");
assert.equal(result.data.data[0].b64_json, "CAkK");
} finally {
globalThis.fetch = originalFetch;
}
});