From 61014aec52a6b9aa6d55ea6def1147bab197e0bd Mon Sep 17 00:00:00 2001 From: rinseaid Date: Mon, 10 Aug 2026 02:54:13 -0400 Subject: [PATCH] fix(image): support Fal reference-image edits (#9933) Co-authored-by: rinseaid Co-authored-by: rinseaid --- open-sse/handlers/imageGeneration.ts | 6 +- .../handlers/imageGeneration/providers/fal.ts | 115 ++++++++++++++++++ src/app/api/v1/images/edits/route.ts | 69 +++++++++-- tests/unit/fal-image-edit.test.ts | 68 +++++++++++ 4 files changed, 248 insertions(+), 10 deletions(-) create mode 100644 open-sse/handlers/imageGeneration/providers/fal.ts create mode 100644 tests/unit/fal-image-edit.test.ts diff --git a/open-sse/handlers/imageGeneration.ts b/open-sse/handlers/imageGeneration.ts index 9f77478d8c..76770bf679 100644 --- a/open-sse/handlers/imageGeneration.ts +++ b/open-sse/handlers/imageGeneration.ts @@ -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) => { diff --git a/open-sse/handlers/imageGeneration/providers/fal.ts b/open-sse/handlers/imageGeneration/providers/fal.ts new file mode 100644 index 0000000000..2a72d5c77a --- /dev/null +++ b/open-sse/handlers/imageGeneration/providers/fal.ts @@ -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; + 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 = { + 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", + }); + } +} diff --git a/src/app/api/v1/images/edits/route.ts b/src/app/api/v1/images/edits/route.ts index 33b9497cd7..9c8871d6bf 100644 --- a/src/app/api/v1/images/edits/route.ts +++ b/src/app/api/v1/images/edits/route.ts @@ -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({ diff --git a/tests/unit/fal-image-edit.test.ts b/tests/unit/fal-image-edit.test.ts new file mode 100644 index 0000000000..ab0e540f1a --- /dev/null +++ b/tests/unit/fal-image-edit.test.ts @@ -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; + } +});