mirror of
https://github.com/diegosouzapw/OmniRoute.git
synced 2026-08-11 17:52:31 +03:00
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:
@@ -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) => {
|
||||
|
||||
115
open-sse/handlers/imageGeneration/providers/fal.ts
Normal file
115
open-sse/handlers/imageGeneration/providers/fal.ts
Normal 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",
|
||||
});
|
||||
}
|
||||
}
|
||||
@@ -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({
|
||||
|
||||
68
tests/unit/fal-image-edit.test.ts
Normal file
68
tests/unit/fal-image-edit.test.ts
Normal 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;
|
||||
}
|
||||
});
|
||||
Reference in New Issue
Block a user