Compare commits

..
66 Commits
Author SHA1 Message Date
Jodh Singh ef7c121415 Merge pull request #86 from jetjodh/jetjodh/seedance25-video-to-video
Add curated Seedance 2.5 video-to-video node
2026-09-23 12:54:46 -07:00
Jodh Singh 28a045b72c Add curated Seedance 2.5 video editing node 2026-09-23 12:51:51 -07:00
Jodh Singh 4a6c7b934e Merge pull request #87 from jetjodh/jetjodh/fix-h3-control-ci
Fix H3 shared-control CI coverage
2026-09-23 12:51:18 -07:00
Jodh Singh be596c6831 Fix H3 shared control test selection 2026-09-23 11:38:56 -07:00
github-actions[bot] 991ceedfc6 chore: refresh fal model registry 2026-09-22 09:23:16 +00:00
github-actions[bot] 2cff2ef7dd chore: refresh fal model registry 2026-09-21 10:05:33 +00:00
github-actions[bot] 7414b7c2bd chore: refresh fal model registry 2026-09-20 09:20:30 +00:00
github-actions[bot] d71da89c09 chore: refresh fal model registry 2026-09-19 08:51:33 +00:00
github-actions[bot] 432dfd19a0 chore: refresh fal model registry 2026-09-18 09:03:45 +00:00
Gökay Aydoğan 647a38f265 Fix schema controls and prevent registry regressions across all models (#85)
* fix: preserve and validate model controls across the fal catalog

* fix: preserve enum constraints through nested normalization
2026-09-17 19:11:50 +03:00
github-actions[bot] 82934d7764 chore: refresh fal model registry 2026-09-17 09:31:41 +00:00
github-actions[bot] 7f8af1d013 chore: refresh fal model registry 2026-09-16 09:24:17 +00:00
github-actions[bot] 68eba7e44c chore: refresh fal model registry 2026-09-15 09:29:35 +00:00
github-actions[bot] bb00c3cd94 chore: refresh fal model registry 2026-09-14 09:57:33 +00:00
github-actions[bot] e81ac14244 chore: refresh fal model registry 2026-09-13 09:37:10 +00:00
github-actions[bot] d481082a1a chore: refresh fal model registry 2026-09-12 08:41:25 +00:00
github-actions[bot] 6747ca2e7a chore: refresh fal model registry 2026-09-11 08:58:31 +00:00
github-actions[bot] 356f1159ac chore: refresh fal model registry 2026-09-10 09:00:21 +00:00
github-actions[bot] 381c0012e4 chore: refresh fal model registry 2026-09-09 09:03:53 +00:00
github-actions[bot] 20bc1e301a chore: refresh fal model registry 2026-09-08 08:58:53 +00:00
github-actions[bot] 6870e00ef7 chore: refresh fal model registry 2026-09-07 09:26:30 +00:00
github-actions[bot] 39fc5ae8df chore: refresh fal model registry 2026-09-06 08:43:20 +00:00
github-actions[bot] e85a103192 chore: refresh fal model registry 2026-09-05 08:24:33 +00:00
github-actions[bot] c7872ac509 chore: refresh fal model registry 2026-09-04 08:54:34 +00:00
github-actions[bot] 70cdfe8877 chore: refresh fal model registry 2026-09-03 08:59:31 +00:00
github-actions[bot] 27afeac7f6 chore: refresh fal model registry 2026-09-02 08:50:21 +00:00
github-actions[bot] 9cb787423e chore: refresh fal model registry 2026-09-01 09:28:25 +00:00
github-actions[bot] a64fd856ae chore: refresh fal model registry 2026-08-31 10:57:07 +00:00
github-actions[bot] f5959e25ec chore: refresh fal model registry 2026-08-30 10:00:07 +00:00
github-actions[bot] 58d3cd1d48 chore: refresh fal model registry 2026-08-29 11:14:07 +00:00
github-actions[bot] 1780dde1bc chore: refresh fal model registry 2026-08-28 16:45:15 +00:00
github-actions[bot] bf9be6fb19 chore: refresh fal model registry 2026-08-27 15:19:15 +00:00
github-actions[bot] 546a6f4923 chore: refresh fal model registry 2026-08-26 04:59:39 +00:00
github-actions[bot] c7adc92c36 chore: refresh fal model registry 2026-08-25 04:57:42 +00:00
github-actions[bot] 46a029acbc chore: refresh fal model registry 2026-08-24 05:08:56 +00:00
github-actions[bot] e8b175d202 chore: refresh fal model registry 2026-08-23 04:58:31 +00:00
github-actions[bot] e8be1b0521 chore: refresh fal model registry 2026-08-22 04:53:27 +00:00
github-actions[bot] c6248be9ec chore: refresh fal model registry 2026-08-21 04:57:46 +00:00
github-actions[bot] f8e89865e9 chore: refresh fal model registry 2026-08-20 04:57:29 +00:00
github-actions[bot] 48b823f3ce chore: refresh fal model registry 2026-08-19 04:55:55 +00:00
github-actions[bot] a408dc63f4 chore: refresh fal model registry 2026-08-18 04:55:51 +00:00
github-actions[bot] 1f95a7b274 chore: refresh fal model registry 2026-08-17 05:01:55 +00:00
github-actions[bot] a1524a7b94 chore: refresh fal model registry 2026-08-16 04:53:30 +00:00
github-actions[bot] 5c87442b98 chore: refresh fal model registry 2026-08-15 04:50:44 +00:00
github-actions[bot] e11eb0fc9f chore: refresh fal model registry 2026-08-14 05:49:21 +00:00
github-actions[bot] 7c34feb913 chore: refresh fal model registry 2026-08-13 05:52:10 +00:00
github-actions[bot] 8747372fc3 chore: refresh fal model registry 2026-08-12 05:49:30 +00:00
github-actions[bot] 9aea820995 chore: refresh fal model registry 2026-08-11 05:32:22 +00:00
github-actions[bot] 486a4e87be chore: refresh fal model registry 2026-08-10 05:48:18 +00:00
github-actions[bot] 4793af4269 chore: refresh fal model registry 2026-08-09 05:26:58 +00:00
github-actions[bot] dd08b8f490 chore: refresh fal model registry 2026-08-08 05:12:09 +00:00
github-actions[bot] b5228cd7a3 chore: refresh fal model registry 2026-08-07 05:55:06 +00:00
github-actions[bot] 3478187364 chore: refresh fal model registry 2026-08-06 07:06:58 +00:00
github-actions[bot] 33f116853b chore: refresh fal model registry 2026-08-05 06:46:51 +00:00
github-actions[bot] 87d59b969f chore: refresh fal model registry 2026-08-04 06:45:03 +00:00
github-actions[bot] 4e42fb7ba1 chore: refresh fal model registry 2026-08-03 07:54:07 +00:00
github-actions[bot] a22efa1bf0 chore: refresh fal model registry 2026-08-02 06:49:00 +00:00
gokayfem f20a036e24 fix: harden media URLs and portable setup 2026-08-01 18:08:25 +03:00
github-actions[bot] 950b500f1b chore: refresh fal model registry 2026-08-01 14:31:00 +00:00
gokayfem 3185a36053 fix: stabilize CI test imports 2026-08-01 17:28:17 +03:00
gokayfem b12aa12e45 feat: automate fal registry updates 2026-08-01 17:21:39 +03:00
gokayfem 6f95834a13 docs: add citation metadata 2026-08-01 03:18:57 +03:00
Gökay Aydoğan 8a47f0598b feat: v2.5.0 — async execution, typed builders, featured tier, registry freshness (#80)
- Async node execution: on ComfyUI with native async support (detected
  via comfy_execution.utils, added in the same commit as async nodes),
  all dynamic nodes and Fal Any Endpoint run as coroutines — independent
  graph branches execute fal calls concurrently with no Submit/Collect
  required. Uploads/downloads/preflight run off-loop; older ComfyUI
  versions keep byte-identical sync behavior. Live-verified: two
  concurrent generations in 2.5s total.
- Typed builder nodes (FAL/Utils/Builders): 8 chainable builders
  (LoRA, embedding, ControlNet, IP-Adapter, reference image/element,
  multi-prompt shot, key-value, JSON merge) replacing JSON-by-hand for
  the 467 object-typed inputs across the catalog; shapes validated
  against live OpenAPI schemas.
- Discovery: FAL/Featured tier (data/featured_models.json, 26 flagship
  endpoints with display-name overrides), 434 models flagged as
  superseded within their family in node help, thumbnails in the
  endpoint picker.
- Registry freshness: startup delta check against the live catalog
  (logs how many models are newer than the snapshot), sidebar Registry
  section with one-click refresh (atomic registry write; restart note).
- Docs: README 1,946 → 327 lines; model tables moved to MODELS.md
  (generator retargeted; weekly refresh workflow now regenerates it);
  CONTRIBUTING.md redirects hand-written-node PRs to the registry and
  featured-list workflow.

Review fixes: spend-guard preflight moved off the event loop in the
async path; registry writes atomically via temp+rename; freshness
daemon gated off in tests; non-finite numbers rejected in FalKeyValue;
sidebar poll budget aligned with the server timeout.
2026-07-02 20:37:15 +03:00
Gökay Aydoğan 31a2c0a35e v2.4.1 — fix provenance for async-collected results (live-smoke finding) (#79)
* fix: provenance lookup for async-collected results (live-smoke finding)

The URL→request_id lookup only searched cached results, so outputs
fetched via Submit→Collect or request-id recovery had no provenance
(sidecars were written with null endpoint/request). Add an explicit
request_urls table populated on every successful fetch path; the lookup
checks it first and falls back to the result-JSON scan.

Verified against the live fal API: full 11-check smoke suite passes,
including the Save→Provenance-from-File round trip on a real CDN file.
Also confirmed live: /account/billing returns 403 for non-admin keys
(handled gracefully with a hint) and the cache/dedup/free-recovery
ledger semantics hold against real requests.

* chore: remove smoke-test artifacts, gitignore output/

* chore: bump to 2.4.1
2026-07-02 16:13:36 +03:00
Gökay Aydoğan ca4251efbe feat: v2.4.0 — 20 utility nodes (dataset prep, media I/O, video/image/data toolkits) (#78)
Everything between local assets and fal endpoints, under FAL/Utils:

- Dataset: Images → Training ZIP URL (LoRA caption layout), Folder → ZIP
  URL, Video → Frame Dataset ZIP URL, Batch Caption Images (parallel VLM,
  wires straight into the trainers). trainer_node's zip helper refactored
  onto shared ArchiveUtils.
- Load: Load Image from URL (multi-URL batching), Load Audio from URL,
  Load Image Folder, Upload Folder as ZIP URL.
- Video (PyAV/cv2, verified against real encoded fixtures): Extract
  Frames (efficient last-frame seek for image-to-video chaining), Trim
  (keyframe remux, no re-encode), Concat (normalizes resolution/fps),
  Mux Audio + Video, Video → Audio.
- Image: Image Grid with Labels, Resize to fal Preset
  (cover/contain/stretch), Image ↔ Base64.
- Data: JSON Extract (dot/bracket paths over result_json outputs),
  Prompt Lines cycler, Text Template.

av added to dependencies (lazy imports keep the pack loading without it).

Review fixes: folder zips log file count/bytes and enforce configurable
[archive] caps before uploading to the CDN; batch captioning re-raises
ComfyUI cancellation instead of recording empty captions; malformed JSON
paths return the default instead of raising; string "false" parses as
boolean False; cold-start balance checks collapse to a single API call.

9 new tests (100 total).
2026-07-02 15:48:55 +03:00
Gökay Aydoğan d6597fb81e feat: v2.3.0 — durable job inbox, provenance, in-canvas cost UI (#77)
- Durable Job Inbox: every async submit is journaled to the shared
  sqlite store and survives ComfyUI restarts; the Fal Job Inbox node
  lists pending/collected jobs and outputs the newest pending
  request_id/endpoint for one-wire recovery via Result by Request ID.
- Provenance: Fal Save Media from URL writes a <file>.fal.json sidecar
  and embeds a fal_provenance PNG text chunk (preserving existing text
  chunks and ICC profiles), returning the request_id as a second output;
  new Fal Provenance from File node reads saved files back into
  endpoint_id + request_id for free re-materialization.
- First web extension (web/ + /fal_api server routes): per-node cost
  pills rendered above every priced fal node (live estimates on
  free-typed endpoint fields), a fal sidebar tab with session spend,
  account balance and a job list with copy/cancel, and in-canvas
  endpoint search with price labels. All features degrade silently on
  older frontends; routes are exception-guarded and cached.

Review fixes: PNG provenance embed preserves source text chunks and ICC
profile; jobs route limit clamped.

15 new tests (91 total).
2026-07-02 15:19:54 +03:00
59 changed files with 8634 additions and 2035 deletions
+15 -5
View File
@@ -9,14 +9,24 @@ on:
- main
jobs:
sidebar-test:
name: Frontend regression tests
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v7
- uses: actions/setup-node@v4
with:
node-version: "22"
- run: node --experimental-vm-modules --test tests/test_*.mjs
lint:
name: Lint (ruff)
runs-on: ubuntu-latest
steps:
- name: Check out code
uses: actions/checkout@v4
uses: actions/checkout@v7
- name: Set up Python
uses: actions/setup-python@v5
uses: actions/setup-python@v7
with:
python-version: "3.11"
- name: Install ruff
@@ -33,9 +43,9 @@ jobs:
python-version: ["3.10", "3.12"]
steps:
- name: Check out code
uses: actions/checkout@v4
uses: actions/checkout@v7
- name: Set up Python
uses: actions/setup-python@v5
uses: actions/setup-python@v7
with:
python-version: ${{ matrix.python-version }}
- name: Install dependencies
@@ -44,7 +54,7 @@ jobs:
pip install pytest
- name: Run tests
run: |
if [ -d tests ]; then pytest tests -x -q; else echo "no tests yet"; fi
if [ -d tests ]; then python -m pytest tests -x -q; else echo "no tests yet"; fi
- name: Registry builder smoke check
run: |
if [ -f scripts/build_registry.py ]; then python scripts/build_registry.py --help; else echo "no registry builder yet"; fi
+86 -30
View File
@@ -1,47 +1,103 @@
name: Refresh fal model registry
on:
push:
branches:
- main
paths:
- ".github/workflows/registry-refresh.yml"
- "data/fal_registry.json"
- "scripts/build_readme.py"
- "scripts/build_registry.py"
- "scripts/validate_registry.py"
- "tests/test_registry_format.py"
- "tests/test_build_registry.py"
- "tests/test_validate_registry.py"
schedule:
# Every Monday at 06:00 UTC
- cron: "0 6 * * 1"
# Daily, off the hour to avoid peak GitHub Actions scheduling delays.
- cron: "17 4 * * *"
workflow_dispatch:
inputs:
allow_large_change:
description: Allow endpoint additions/removals above the automatic safety limits
required: false
default: false
type: boolean
allow_input_removal:
description: Allow reviewed removal of existing input controls or enum choices
required: false
default: false
type: boolean
concurrency:
group: fal-registry-refresh
cancel-in-progress: false
permissions:
contents: write
pull-requests: write
jobs:
refresh:
name: Rebuild registry and open PR
name: Validate and publish registry refresh
runs-on: ubuntu-latest
timeout-minutes: 30
steps:
- name: Check out code
uses: actions/checkout@v4
- name: Check out the default branch
uses: actions/checkout@v7
with:
ref: main
fetch-depth: 0
- name: Set up Python
uses: actions/setup-python@v5
uses: actions/setup-python@v7
with:
python-version: "3.11"
- name: Rebuild registry
run: python scripts/build_registry.py --out data/fal_registry.json
- name: Summarize changes
id: diff
run: |
{
echo "stat<<EOF"
git diff --stat
echo "EOF"
} >> "$GITHUB_OUTPUT"
# create-pull-request skips PR creation when there are no changes.
- name: Create pull request
uses: peter-evans/create-pull-request@v6
with:
branch: chore/registry-refresh
commit-message: "chore: refresh fal model registry"
title: "Refresh fal model registry"
body: |
Automated weekly refresh of `data/fal_registry.json` via `scripts/build_registry.py`.
```
${{ steps.diff.outputs.stat }}
```
delete-branch: true
- name: Build candidate registry
run: >-
python scripts/build_registry.py
--out data/fal_registry.next.json
--preserve-from data/fal_registry.json
- name: Validate candidate against the committed registry
env:
ALLOW_LARGE_CHANGE: ${{ inputs.allow_large_change || 'false' }}
ALLOW_INPUT_REMOVAL: ${{ inputs.allow_input_removal || 'false' }}
run: |
args=(
data/fal_registry.next.json
--baseline data/fal_registry.json
--max-removal-fraction 0.05
--max-addition-fraction 0.25
)
if [[ "$ALLOW_LARGE_CHANGE" == "true" ]]; then
args+=(--allow-large-change)
fi
if [[ "$ALLOW_INPUT_REMOVAL" == "true" ]]; then
args+=(--allow-input-removal)
fi
python scripts/validate_registry.py "${args[@]}"
- name: Promote candidate and regenerate model catalog
run: |
mv data/fal_registry.next.json data/fal_registry.json
python scripts/build_readme.py
- name: Run refresh safety checks
run: |
python -m pip install pytest
python -m pytest tests/test_registry_format.py tests/test_validate_registry.py tests/test_build_registry.py -q
git diff --check
- name: Commit and push changes
run: |
if git diff --quiet -- data/fal_registry.json MODELS.md; then
echo "Registry is already current."
exit 0
fi
git config user.name "github-actions[bot]"
git config user.email "41898282+github-actions[bot]@users.noreply.github.com"
git add data/fal_registry.json MODELS.md
git commit -m "chore: refresh fal model registry"
git pull --rebase origin main
git push origin HEAD:main
+1
View File
@@ -174,3 +174,4 @@ memory-bank/
.DS_Store
.claude/
Node-Docs/
output/
+21
View File
@@ -0,0 +1,21 @@
cff-version: 1.2.0
message: "If you use ComfyUI-fal-API in your work, please cite it using the metadata below."
type: software
title: "ComfyUI-fal-API"
version: "2.5.0"
date-released: 2026-07-02
authors:
- family-names: "Aydoğan"
given-names: "Gökay"
orcid: "https://orcid.org/0000-0002-2343-9433"
abstract: "A ComfyUI integration for the fal model catalog, including curated and generated nodes, workflow utilities, caching, and spend controls."
keywords:
- ComfyUI
- fal
- generative AI
- image generation
- video generation
- API integration
license: Apache-2.0
repository-code: "https://github.com/gokayfem/ComfyUI-fal-API"
url: "https://github.com/gokayfem/ComfyUI-fal-API"
+104
View File
@@ -0,0 +1,104 @@
# Contributing to ComfyUI-fal-API
Thanks for helping! Before you write anything, read this — it will probably save you the PR entirely.
## A new fal model does NOT need code
Historically, adding a model to this pack meant hand-writing a node class. **That is no longer how it works.** Every live public model on fal gets a node automatically, generated at ComfyUI startup from the committed snapshot at `data/fal_registry.json`. No node class, no mapping entry, no code.
The snapshot stays fresh three ways:
- A **daily GitHub Action** (`.github/workflows/registry-refresh.yml`) builds a candidate, validates it against the committed baseline, and commits safe changes automatically.
- Anyone can run it locally: `python scripts/build_registry.py --out data/fal_registry.json` (then `python scripts/build_readme.py` to regenerate [MODELS.md](MODELS.md)).
- The fal sidebar can rebuild the local registry; restart ComfyUI afterward so new node classes register.
The automated validator rejects malformed records, duplicate or unsorted endpoint IDs, suspiciously small catalogs, and unexpectedly large additions or removals. A maintainer can override the change-size thresholds when manually dispatching the workflow; all structural checks still apply.
Refresh validation also rejects removed input controls, lost dropdowns, and removed enum choices on existing endpoints. Inspect these changes against the live API before using the separate `allow_input_removal` workflow option (or `validate_registry.py --allow-input-removal`). The sidebar builds and validates a candidate before replacing the local snapshot, so a failed refresh keeps the previous registry intact.
The generator handles nested references, composed schemas, nullable types and literal choices for every endpoint, and never caps the number of exposed fields. Short string examples become suggested dropdown choices with a **custom value** input; they are not treated as exhaustive enums. Prompts, prose and URLs remain text/media controls. Fix schema patterns in the shared generator and widget/argument translators rather than adding model-specific patches. Add upstream schema fixtures and regressions in `tests/test_build_registry.py`; H3 and the catalog's suggested controls are checked against the shipped registry too.
Missing endpoints are preserved by default with `deprecated: true`, which keeps their node keys available under `FAL/Compatibility` for old workflows. Use `--prune-missing` only for an intentional breaking cleanup. A fresh endpoint record automatically replaces its deprecated copy if it returns to the live catalog.
To check whether a model is already covered:
```bash
grep '"endpoint_id": "fal-ai/your/endpoint"' data/fal_registry.json
# or browse MODELS.md, or search the node browser in ComfyUI
```
If a model is live on [fal.ai/models](https://fal.ai/models) but missing from the snapshot, rerun `scripts/build_registry.py` — if it's *still* missing, open an issue with the endpoint id. And if you need a model **right now**, the **Fal Any Endpoint (fal)** node calls any endpoint by id without any registry entry at all.
So: **please don't open a PR that adds a node class for a new model.** It will be redundant the moment the registry refreshes.
## Want a model promoted or renamed? Edit `data/featured_models.json`
When a model deserves curation — a spot in the **FAL/Featured** menu tier or a friendlier display name — add its endpoint to `data/featured_models.json` (featured tier + display-name override). That's the whole change: one JSON entry, not a new node class.
## When a hand-written node IS justified
A curated node earns its place only when the generated node genuinely can't express the UX:
- **Multi-endpoint orchestration** — one node fanning out to several endpoints (e.g. Combined Video Generation).
- **Special input ergonomics** — first/last-frame image pairing, unified T2V/I2V dispatch, LoRA slots with per-slot scales.
If you're writing one, the rules are non-negotiable:
1. **Import only from the `.fal_utils` facade** (`from .fal_utils import ApiHandler, FalConfig, ImageUtils, ResultProcessor, ...`) — never reach into `nodes/utils/` internals or call `fal_client` directly. The facade gives you the result cache, spend guard, session ledger, and error handling for free.
2. **Raise errors — no silent fallbacks.** Never return blank images or `"Error: ..."` strings; let `ApiHandler` surface fal's actual error message.
3. **Tooltips on every input.** Users should never have to guess a parameter.
4. **Never change existing node keys, input names, or output signatures.** Existing user workflows reference them forever. `tests/legacy_node_keys.json` is the snapshot of keys that must never be removed or renamed, and `tests/test_mappings.py` fails the suite if one disappears. New inputs must be optional with backward-compatible defaults.
5. **Add tests** alongside the existing ones in `tests/`.
## Dev setup
```bash
pip install -r requirements.txt
python -m pytest tests # the suite MUST pass
ruff check . # lint, same as CI
node --experimental-vm-modules --test tests/test_*.mjs
```
CI runs both on every PR (Python 3.10 and 3.12). The most important test to understand is the **compatibility snapshot**: `tests/test_mappings.py` asserts that every node key recorded in `tests/legacy_node_keys.json` still registers. If your change makes it fail, the fix is to restore the key — not to edit the snapshot.
## Architecture map
```
scripts/build_registry.py queries fal's platform APIs → writes the snapshot
scripts/validate_registry.py safety gate for generated registry candidates
data/fal_registry.json committed model catalog (~1,400 models)
data/featured_models.json curation: featured tier + display-name overrides
scripts/build_readme.py renders MODELS.md from the snapshot
nodes/dynamic/ the auto-generated node machinery
registry_loader.py reads the snapshot, applies [dynamic_nodes] config;
never raises — failures degrade to curated-only
factory.py builds one node class per model, in memory
schema_to_inputs.py registry input specs → ComfyUI INPUT_TYPES (+ tooltips)
arguments.py widget/socket values → API arguments (uploads media)
outputs.py API result → IMAGE / VIDEO / AUDIO / result_json
any_endpoint.py the generic "call any endpoint by id" node
nodes/*.py curated hand-written nodes (image, video, llm, vlm,
trainer, upscaler, util_*)
nodes/fal_utils.py import facade — node modules import ONLY from here
nodes/utils/ the implementations behind the facade: api, config,
pricing, result_cache, ledger, billing (spend guard),
job_store, media, archive, errors, logger
nodes/platform_node.py platform nodes (Submit/Collect, costs, request ids)
nodes/inbox_node.py durable job inbox
nodes/billing_node.py account balance
nodes/server_routes.py HTTP endpoints backing the frontend extension
web/ ComfyUI frontend: cost badges, fal sidebar,
endpoint autocomplete
tests/ pytest suite, incl. the legacy_node_keys.json snapshot
```
## PR checklist
- [ ] Not a hand-written node for a single new model (registry covers it — see above)
- [ ] `python -m pytest tests` passes locally
- [ ] `ruff check .` is clean
- [ ] No existing node keys, inputs, or outputs changed
- [ ] New curated node (if truly justified): uses `.fal_utils`, raises errors, has tooltips and tests
- [ ] No secrets, no `config.ini`, no generated artifacts in the diff
+1735
View File
File diff suppressed because it is too large Load Diff
+157 -1855
View File
File diff suppressed because it is too large Load Diff
+6
View File
@@ -11,6 +11,12 @@ node_list = [
"platform_node",
"billing_node",
"inbox_node",
"util_dataset_node",
"util_media_in_node",
"util_video_node",
"util_image_node",
"util_data_node",
"builder_node",
]
NODE_CLASS_MAPPINGS = {}
File diff suppressed because one or more lines are too long
+31
View File
@@ -0,0 +1,31 @@
{
"version": 1,
"featured": [
{"endpoint_id": "fal-ai/kling-video/o3/pro/text-to-video", "display_name": null},
{"endpoint_id": "fal-ai/kling-video/o3/pro/image-to-video", "display_name": null},
{"endpoint_id": "fal-ai/veo3.1", "display_name": "Veo 3.1 Text to Video (fal)"},
{"endpoint_id": "fal-ai/veo3.1/image-to-video", "display_name": "Veo 3.1 Image to Video (fal)"},
{"endpoint_id": "fal-ai/wan/v2.7/text-to-video", "display_name": "Wan 2.7 Text to Video (fal)"},
{"endpoint_id": "fal-ai/wan/v2.7/image-to-video", "display_name": "Wan 2.7 Image to Video (fal)"},
{"endpoint_id": "bytedance/seedance-2.0/text-to-video", "display_name": "Seedance 2.0 Text to Video (fal)"},
{"endpoint_id": "bytedance/seedance-2.0/image-to-video", "display_name": "Seedance 2.0 Image to Video (fal)"},
{"endpoint_id": "fal-ai/sora-2/text-to-video/pro", "display_name": "Sora 2 Pro Text to Video (fal)"},
{"endpoint_id": "fal-ai/sora-2/image-to-video/pro", "display_name": "Sora 2 Pro Image to Video (fal)"},
{"endpoint_id": "fal-ai/minimax/hailuo-2.3/pro/image-to-video", "display_name": null},
{"endpoint_id": "fal-ai/flux-2-max", "display_name": null},
{"endpoint_id": "fal-ai/flux-2-max/edit", "display_name": "Flux 2 Max Edit (fal)"},
{"endpoint_id": "fal-ai/nano-banana-2", "display_name": null},
{"endpoint_id": "fal-ai/nano-banana-2/edit", "display_name": "Nano Banana 2 Edit (fal)"},
{"endpoint_id": "openai/gpt-image-2", "display_name": "GPT Image 2 (fal)"},
{"endpoint_id": "openai/gpt-image-2/edit", "display_name": "GPT Image 2 Edit (fal)"},
{"endpoint_id": "fal-ai/bytedance/seedream/v4.5/text-to-image", "display_name": "Seedream 4.5 Text to Image (fal)"},
{"endpoint_id": "fal-ai/bytedance/seedream/v4.5/edit", "display_name": "Seedream 4.5 Edit (fal)"},
{"endpoint_id": "fal-ai/recraft/v4.1/pro/text-to-image", "display_name": null},
{"endpoint_id": "ideogram/v4", "display_name": null},
{"endpoint_id": "fal-ai/elevenlabs/tts/eleven-v3", "display_name": "ElevenLabs TTS Eleven v3 (fal)"},
{"endpoint_id": "fal-ai/elevenlabs/speech-to-text/scribe-v2", "display_name": "ElevenLabs Scribe v2 (fal)"},
{"endpoint_id": "fal-ai/hunyuan-3d/v3.1/pro/image-to-3d", "display_name": null},
{"endpoint_id": "fal-ai/topaz/upscale/image", "display_name": "Topaz Image Upscale (fal)"},
{"endpoint_id": "fal-ai/topaz/upscale/video", "display_name": null}
]
}
+743
View File
@@ -0,0 +1,743 @@
"""Chainable typed builder nodes for JSON inputs on auto-generated fal nodes.
Auto-generated endpoint nodes render complex object/array inputs (registry
type "json") as raw JSON string widgets. The builders here emit exactly the
JSON those fields expect, and each accepts an optional ``chain`` input so N
builders can be daisy-chained to produce an N-element array (or a merged
object for ``FalKeyValue``).
Shapes were validated against the live OpenAPI schemas
(https://fal.ai/api/openapi/queue/openapi.json?endpoint_id=<id>):
- ``LoraWeight`` {path, scale[, weight_name]} fal-ai/flux-lora,
fal-ai/wan/v2.2-a14b/text-to-video/lora (126 "loras" inputs in registry)
- ``Embedding`` {path, tokens[]} fal-ai/fast-lightning-sdxl
- ``ControlNet`` {path, control_image_url, conditioning_scale,
start_percentage, end_percentage[, variant]} fal-ai/flux-general
- ``IPAdapter`` {path, image_encoder_path, image_url, scale
[, weight_name]} fal-ai/flux-general
- ``ElementInput`` {frontal_image_url, reference_image_urls[]}
fal-ai/kling-image/o1, fal-ai/kling-image/o3/*
- ``KlingV3MultiPromptElement`` {prompt, duration("1".."15")}
fal-ai/kling-video/o3/*/image-to-video
"""
from __future__ import annotations
import json
import math
from typing import Any
from .fal_utils import FalApiError, ImageUtils, logger
_CATEGORY = "FAL/Utils/Builders"
_CHAIN_TOOLTIP = (
"Optional: wire the json output of another builder of the same kind here "
"to append this entry after its entries (chain N builders for N items)."
)
def _parse_chain(node_name: str, chain: str, container: type) -> Any:
"""Parse a prior chain string into ``container`` (list or dict).
An empty/blank chain yields a fresh empty container. Anything that is not
valid JSON of the right container type raises a clear FalApiError.
"""
text = (chain or "").strip()
if not text:
return container()
try:
parsed = json.loads(text)
except ValueError as err:
logger.error("%s: invalid chain JSON: %s", node_name, err)
raise FalApiError(node_name, f"'chain' is not valid JSON: {err}") from err
if not isinstance(parsed, container):
wanted = "array" if container is list else "object"
if isinstance(parsed, dict):
got = "object"
elif isinstance(parsed, list):
got = "array"
else:
got = type(parsed).__name__
raise FalApiError(
node_name,
f"'chain' must be a JSON {wanted} (got {got}). "
f"Only chain {node_name}-compatible builders together.",
)
return parsed
def _append_entry(node_name: str, chain: str, entry: dict[str, Any]) -> str:
"""New JSON array string: entries from ``chain`` plus ``entry`` (no mutation)."""
prior = _parse_chain(node_name, chain, list)
return json.dumps([*prior, entry])
def _require(node_name: str, field: str, value: str) -> str:
"""Strip a required string field, raising when it is blank."""
text = (value or "").strip()
if not text:
raise FalApiError(node_name, f"'{field}' is required and cannot be empty")
return text
def _resolve_image_url(node_name: str, field: str, image: Any, url: str, required: bool) -> str:
"""A connected IMAGE wins (uploaded via fal storage); else the URL string."""
if image is not None:
return ImageUtils.upload_image(image)
text = (url or "").strip()
if not text and required:
raise FalApiError(
node_name,
f"Connect an image or fill '{field}': the schema requires an image URL",
)
return text
class FalLoRAConfig:
"""Append one LoraWeight ({path, scale}) entry to a JSON array."""
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("json",)
FUNCTION = "build"
CATEGORY = _CATEGORY
DESCRIPTION = (
"Build a `loras` JSON array entry ({path, scale}) without hand-writing "
"JSON. Chain several to stack LoRAs. Wire the json output into the "
"`loras` field of 126+ fal nodes (fal-ai/flux-lora, "
"fal-ai/wan/v2.2-a14b/text-to-video/lora, fal-ai/qwen-image, ...)."
)
@classmethod
def INPUT_TYPES(cls) -> dict[str, Any]:
return {
"required": {
"path": (
"STRING",
{
"default": "",
"tooltip": (
"URL or Hugging Face id of the LoRA weights, e.g. "
"https://.../lora.safetensors. Feeds the `loras` field of "
"fal-ai/flux-lora, fal-ai/wan/v2.2-a14b/text-to-video/lora, "
"fal-ai/chrono-edit-lora and 120+ more."
),
},
),
"scale": (
"FLOAT",
{
"default": 1.0,
"min": 0.0,
"max": 4.0,
"step": 0.01,
"tooltip": "LoRA strength merged into the base model (LoraWeight.scale, 0-4).",
},
),
},
"optional": {
"weight_name": (
"STRING",
{
"default": "",
"tooltip": (
"Optional safetensors file name when `path` is a Hugging Face "
"repo with several files (e.g. Wan/Qwen LoRA endpoints). "
"Leave empty otherwise."
),
},
),
"chain": ("STRING", {"forceInput": True, "tooltip": _CHAIN_TOOLTIP}),
},
}
def build(self, path: str, scale: float, weight_name: str = "", chain: str = "") -> tuple[str]:
entry: dict[str, Any] = {
"path": _require("FalLoRAConfig", "path", path),
"scale": float(scale),
}
if (weight_name or "").strip():
entry = {**entry, "weight_name": weight_name.strip()}
return (_append_entry("FalLoRAConfig", chain, entry),)
class FalEmbeddingConfig:
"""Append one Embedding ({path, tokens}) entry to a JSON array."""
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("json",)
FUNCTION = "build"
CATEGORY = _CATEGORY
DESCRIPTION = (
"Build an `embeddings` JSON array entry ({path, tokens}) for SD/SDXL "
"endpoints such as fal-ai/fast-lightning-sdxl, fal-ai/dreamshaper and "
"fal-ai/fast-fooocus-sdxl. Chain several to load multiple embeddings."
)
@classmethod
def INPUT_TYPES(cls) -> dict[str, Any]:
return {
"required": {
"path": (
"STRING",
{
"default": "",
"tooltip": (
"URL or path to the textual-inversion embedding weights, e.g. "
"https://civitai.com/api/download/models/135931. Feeds the "
"`embeddings` field of fal-ai/fast-lightning-sdxl, "
"fal-ai/dreamshaper, fal-ai/fast-fooocus-sdxl."
),
},
),
},
"optional": {
"tokens": (
"STRING",
{
"default": "<s0>, <s1>",
"tooltip": (
"Comma-separated trigger tokens for the embedding "
"(Embedding.tokens). Leave empty to use the endpoint default."
),
},
),
"chain": ("STRING", {"forceInput": True, "tooltip": _CHAIN_TOOLTIP}),
},
}
def build(self, path: str, tokens: str = "<s0>, <s1>", chain: str = "") -> tuple[str]:
entry: dict[str, Any] = {"path": _require("FalEmbeddingConfig", "path", path)}
token_list = [part.strip() for part in (tokens or "").split(",") if part.strip()]
if token_list:
entry = {**entry, "tokens": token_list}
return (_append_entry("FalEmbeddingConfig", chain, entry),)
class FalControlNetConfig:
"""Append one ControlNet conditioning entry to a JSON array."""
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("json",)
FUNCTION = "build"
CATEGORY = _CATEGORY
DESCRIPTION = (
"Build a `controlnets` JSON array entry ({path, control_image_url, "
"conditioning_scale, start/end_percentage}) for fal-ai/flux-general and "
"its variants (image-to-image, inpainting, differential-diffusion). "
"Connect an IMAGE (auto-uploaded) or paste a control image URL."
)
@classmethod
def INPUT_TYPES(cls) -> dict[str, Any]:
return {
"required": {
"path": (
"STRING",
{
"default": "",
"tooltip": (
"URL or Hugging Face path to the ControlNet weights. Feeds the "
"`controlnets` field of fal-ai/flux-general, "
"fal-ai/flux-general/image-to-image, fal-ai/flux-general/inpainting."
),
},
),
"conditioning_scale": (
"FLOAT",
{
"default": 1.0,
"min": 0.0,
"max": 2.0,
"step": 0.01,
"tooltip": "Strength of the ControlNet guidance (ControlNet.conditioning_scale).",
},
),
"start_percentage": (
"FLOAT",
{
"default": 0.0,
"min": 0.0,
"max": 1.0,
"step": 0.01,
"tooltip": "Fraction of total timesteps at which the ControlNet starts applying (0-1).",
},
),
"end_percentage": (
"FLOAT",
{
"default": 1.0,
"min": 0.0,
"max": 1.0,
"step": 0.01,
"tooltip": "Fraction of total timesteps at which the ControlNet stops applying (0-1).",
},
),
},
"optional": {
"control_image": (
"IMAGE",
{
"tooltip": (
"Control image (canny/depth/pose map, ...). Uploaded to fal "
"storage and sent as `control_image_url`. Overrides the URL widget."
),
},
),
"control_image_url": (
"STRING",
{
"default": "",
"tooltip": (
"Direct URL for the control image; used when no IMAGE is connected. "
"The schema requires one of the two."
),
},
),
"variant": (
"STRING",
{
"default": "",
"tooltip": "Optional variant when `path` is a Hugging Face repo key. Leave empty otherwise.",
},
),
"chain": ("STRING", {"forceInput": True, "tooltip": _CHAIN_TOOLTIP}),
},
}
def build(
self,
path: str,
conditioning_scale: float,
start_percentage: float,
end_percentage: float,
control_image: Any = None,
control_image_url: str = "",
variant: str = "",
chain: str = "",
) -> tuple[str]:
node = "FalControlNetConfig"
entry: dict[str, Any] = {
"path": _require(node, "path", path),
"control_image_url": _resolve_image_url(
node, "control_image_url", control_image, control_image_url, required=True
),
"conditioning_scale": float(conditioning_scale),
"start_percentage": float(start_percentage),
"end_percentage": float(end_percentage),
}
if (variant or "").strip():
entry = {**entry, "variant": variant.strip()}
return (_append_entry(node, chain, entry),)
class FalIPAdapterConfig:
"""Append one IP-Adapter entry to a JSON array."""
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("json",)
FUNCTION = "build"
CATEGORY = _CATEGORY
DESCRIPTION = (
"Build an `ip_adapters` JSON array entry ({path, image_encoder_path, "
"image_url, scale}) for fal-ai/flux-general and its variants. Connect "
"an IMAGE (auto-uploaded) or paste a reference image URL. For the older "
"fal-ai/lora `ip_adapter` field (different keys) use FalKeyValue."
)
@classmethod
def INPUT_TYPES(cls) -> dict[str, Any]:
return {
"required": {
"path": (
"STRING",
{
"default": "",
"tooltip": (
"Hugging Face path to the IP-Adapter weights. Feeds the "
"`ip_adapters` field of fal-ai/flux-general, "
"fal-ai/flux-general/image-to-image, fal-ai/flux-general/rf-inversion."
),
},
),
"image_encoder_path": (
"STRING",
{
"default": "openai/clip-vit-large-patch14",
"tooltip": "Path to the image encoder for the IP-Adapter (IPAdapter.image_encoder_path).",
},
),
"scale": (
"FLOAT",
{
"default": 1.0,
"min": 0.0,
"max": 4.0,
"step": 0.01,
"tooltip": "Strength of the IP-Adapter conditioning (IPAdapter.scale).",
},
),
},
"optional": {
"image": (
"IMAGE",
{
"tooltip": (
"Reference image for the IP-Adapter conditioning. Uploaded to fal "
"storage and sent as `image_url`. Overrides the URL widget."
),
},
),
"image_url": (
"STRING",
{
"default": "",
"tooltip": (
"Direct URL for the reference image; used when no IMAGE is connected. "
"The schema requires one of the two."
),
},
),
"weight_name": (
"STRING",
{
"default": "",
"tooltip": (
"Optional safetensors file name containing the IP-Adapter weights "
"(IPAdapter.weight_name). Leave empty otherwise."
),
},
),
"chain": ("STRING", {"forceInput": True, "tooltip": _CHAIN_TOOLTIP}),
},
}
def build(
self,
path: str,
image_encoder_path: str,
scale: float,
image: Any = None,
image_url: str = "",
weight_name: str = "",
chain: str = "",
) -> tuple[str]:
node = "FalIPAdapterConfig"
entry: dict[str, Any] = {
"path": _require(node, "path", path),
"image_encoder_path": _require(node, "image_encoder_path", image_encoder_path),
"image_url": _resolve_image_url(node, "image_url", image, image_url, required=True),
"scale": float(scale),
}
if (weight_name or "").strip():
entry = {**entry, "weight_name": weight_name.strip()}
return (_append_entry(node, chain, entry),)
class FalReferenceImage:
"""Append one Kling ElementInput (reference character/object) to a JSON array."""
_MAX_REFERENCE_IMAGES = 3
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("json",)
FUNCTION = "build"
CATEGORY = _CATEGORY
DESCRIPTION = (
"Build an `elements` JSON array entry ({frontal_image_url, "
"reference_image_urls}) for Kling Omni image endpoints "
"(fal-ai/kling-image/o1, fal-ai/kling-image/o3/text-to-image, "
"fal-ai/kling-image/o3/image-to-image). Images are auto-uploaded. "
"Chain one builder per character/object element."
)
@classmethod
def INPUT_TYPES(cls) -> dict[str, Any]:
return {
"required": {
"frontal_image": (
"IMAGE",
{
"tooltip": (
"Frontal view of the character/object. Uploaded to fal storage and "
"sent as `frontal_image_url` inside the `elements` field of "
"fal-ai/kling-image/o1 and fal-ai/kling-image/o3 endpoints."
),
},
),
},
"optional": {
"reference_images": (
"IMAGE",
{
"tooltip": (
"Optional batch of up to 3 additional views from different angles "
"(sent as `reference_image_urls`)."
),
},
),
"chain": ("STRING", {"forceInput": True, "tooltip": _CHAIN_TOOLTIP}),
},
}
def build(self, frontal_image: Any, reference_images: Any = None, chain: str = "") -> tuple[str]:
node = "FalReferenceImage"
entry: dict[str, Any] = {"frontal_image_url": ImageUtils.upload_image(frontal_image)}
if reference_images is not None:
urls = ImageUtils.prepare_images(reference_images)
if len(urls) > self._MAX_REFERENCE_IMAGES:
raise FalApiError(
node,
f"'reference_images' supports at most {self._MAX_REFERENCE_IMAGES} "
f"images per element (got {len(urls)})",
)
if urls:
entry = {**entry, "reference_image_urls": urls}
return (_append_entry(node, chain, entry),)
class FalMultiPromptShot:
"""Append one Kling multi-prompt shot ({prompt, duration}) to a JSON array."""
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("json",)
FUNCTION = "build"
CATEGORY = _CATEGORY
DESCRIPTION = (
"Build a `multi_prompt` JSON array entry ({prompt, duration}) for Kling "
"O3 video endpoints (fal-ai/kling-video/o3/standard/image-to-video, "
"fal-ai/kling-video/o3/pro/text-to-video, .../4k variants). Chain one "
"builder per shot to script a multi-shot video."
)
@classmethod
def INPUT_TYPES(cls) -> dict[str, Any]:
return {
"required": {
"prompt": (
"STRING",
{
"default": "",
"multiline": True,
"tooltip": (
"The prompt for this shot. Feeds the `multi_prompt` field of "
"fal-ai/kling-video/o3 image-to-video / text-to-video / "
"reference-to-video endpoints."
),
},
),
"duration": (
"INT",
{
"default": 5,
"min": 1,
"max": 15,
"tooltip": "Duration of this shot in seconds (1-15, sent as a string per the schema).",
},
),
},
"optional": {
"chain": ("STRING", {"forceInput": True, "tooltip": _CHAIN_TOOLTIP}),
},
}
def build(self, prompt: str, duration: int, chain: str = "") -> tuple[str]:
node = "FalMultiPromptShot"
entry = {
"prompt": _require(node, "prompt", prompt),
"duration": str(int(duration)),
}
return (_append_entry(node, chain, entry),)
def _typed_value(node: str, value: str, value_type: str) -> Any:
"""Coerce the FalKeyValue string widget into the selected JSON type."""
if value_type == "string":
return value
text = value.strip()
if value_type == "number":
try:
number = float(text)
except ValueError as err:
raise FalApiError(node, f"'value' is not a number: {text!r}") from err
if not math.isfinite(number):
raise FalApiError(node, f"'value' must be a finite number, got: {text!r}")
return int(number) if number.is_integer() else number
if value_type == "boolean":
lowered = text.lower()
if lowered in ("true", "1", "yes"):
return True
if lowered in ("false", "0", "no"):
return False
raise FalApiError(node, f"'value' is not a boolean (use true/false): {text!r}")
# value_type == "json": nested arrays/objects/null, e.g. from another builder
try:
return json.loads(text)
except ValueError as err:
raise FalApiError(node, f"'value' is not valid JSON: {err}") from err
class FalKeyValue:
"""Merge one typed key/value pair into a JSON object (chainable)."""
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("json",)
FUNCTION = "build"
CATEGORY = _CATEGORY
DESCRIPTION = (
"Generic escape hatch: build a JSON OBJECT one typed key at a time. "
"Chain several to fill object fields like `audio_setting` / "
"`voice_setting` (fal-ai/minimax-music/v2, fal-ai/minimax/speech-02-hd) "
"or `validation` (fal-ai/ltx23-trainer-v2). Set value_type to `json` to "
"nest arrays/objects, including outputs of the array builders."
)
@classmethod
def INPUT_TYPES(cls) -> dict[str, Any]:
return {
"required": {
"key": (
"STRING",
{
"default": "",
"tooltip": (
"Object key to set, e.g. sample_rate for `audio_setting` on "
"fal-ai/minimax-music/v2 or speed for `voice_setting` on "
"fal-ai/minimax/speech-02-hd."
),
},
),
"value": (
"STRING",
{
"default": "",
"multiline": True,
"tooltip": "Value for the key, interpreted according to value_type.",
},
),
"value_type": (
["string", "number", "boolean", "json"],
{
"default": "string",
"tooltip": (
"How to encode the value: string as-is, number/boolean parsed, "
"json for nested objects/arrays (e.g. a builder output)."
),
},
),
},
"optional": {
"chain": (
"STRING",
{
"forceInput": True,
"tooltip": (
"Optional: wire another FalKeyValue json output here to merge this "
"key into that object (later keys win)."
),
},
),
},
}
def build(self, key: str, value: str, value_type: str, chain: str = "") -> tuple[str]:
node = "FalKeyValue"
prior = _parse_chain(node, chain, dict)
merged = {**prior, _require(node, "key", key): _typed_value(node, value, value_type)}
return (json.dumps(merged),)
class FalJSONMerge:
"""Merge two builder outputs: arrays concatenate, objects merge (b wins)."""
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("json",)
FUNCTION = "merge"
CATEGORY = _CATEGORY
DESCRIPTION = (
"Merge two JSON strings: two arrays concatenate (a then b), two objects "
"merge with b overriding a. Useful to combine separately built chains "
"before wiring them into one json field (e.g. two `loras` chains, or "
"FalKeyValue objects for `audio_setting` / `validation`)."
)
@classmethod
def INPUT_TYPES(cls) -> dict[str, Any]:
return {
"required": {
"a": (
"STRING",
{
"forceInput": True,
"tooltip": "First JSON array or object (a builder json output). Empty is allowed.",
},
),
"b": (
"STRING",
{
"forceInput": True,
"tooltip": (
"Second JSON array or object. Must be the same container type as "
"'a'; object keys in 'b' override 'a'."
),
},
),
},
}
@staticmethod
def _parse(side: str, text: str) -> Any:
stripped = (text or "").strip()
if not stripped:
return None
try:
parsed = json.loads(stripped)
except ValueError as err:
raise FalApiError("FalJSONMerge", f"'{side}' is not valid JSON: {err}") from err
if not isinstance(parsed, (list, dict)):
raise FalApiError(
"FalJSONMerge",
f"'{side}' must be a JSON array or object, got {type(parsed).__name__}",
)
return parsed
def merge(self, a: str, b: str) -> tuple[str]:
parsed_a = self._parse("a", a)
parsed_b = self._parse("b", b)
if parsed_a is None and parsed_b is None:
raise FalApiError("FalJSONMerge", "Both 'a' and 'b' are empty; nothing to merge")
if parsed_a is None or parsed_b is None:
return (json.dumps(parsed_b if parsed_a is None else parsed_a),)
if isinstance(parsed_a, list) and isinstance(parsed_b, list):
return (json.dumps([*parsed_a, *parsed_b]),)
if isinstance(parsed_a, dict) and isinstance(parsed_b, dict):
return (json.dumps({**parsed_a, **parsed_b}),)
raise FalApiError(
"FalJSONMerge",
"'a' and 'b' must both be arrays or both be objects "
f"(got {type(parsed_a).__name__} and {type(parsed_b).__name__})",
)
NODE_CLASS_MAPPINGS = {
"FalLoRAConfig_fal": FalLoRAConfig,
"FalEmbeddingConfig_fal": FalEmbeddingConfig,
"FalControlNetConfig_fal": FalControlNetConfig,
"FalIPAdapterConfig_fal": FalIPAdapterConfig,
"FalReferenceImage_fal": FalReferenceImage,
"FalMultiPromptShot_fal": FalMultiPromptShot,
"FalKeyValue_fal": FalKeyValue,
"FalJSONMerge_fal": FalJSONMerge,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"FalLoRAConfig_fal": "LoRA Config (fal)",
"FalEmbeddingConfig_fal": "Embedding Config (fal)",
"FalControlNetConfig_fal": "ControlNet Config (fal)",
"FalIPAdapterConfig_fal": "IP-Adapter Config (fal)",
"FalReferenceImage_fal": "Reference Image Element (fal)",
"FalMultiPromptShot_fal": "Multi-Prompt Shot (fal)",
"FalKeyValue_fal": "Key/Value JSON (fal)",
"FalJSONMerge_fal": "JSON Merge (fal)",
}
+48 -5
View File
@@ -2,6 +2,7 @@
from __future__ import annotations
import asyncio
import json
from typing import Any
@@ -13,13 +14,20 @@ from ..fal_utils import (
ResultProcessor,
logger,
)
from .factory import stable_hash
from .factory import _ASYNC_CAPABLE, stable_hash
from .outputs import find_url
ANY_ENDPOINT_KEY = "FalAnyEndpoint_fal"
ANY_ENDPOINT_DISPLAY_NAME = "Fal Any Endpoint (fal)"
def _validated_endpoint(endpoint_id: str) -> str:
endpoint = (endpoint_id or "").strip()
if not endpoint:
raise FalApiError("(any endpoint)", "endpoint_id is required")
return endpoint
def _parse_arguments_json(endpoint_id: str, arguments_json: str) -> dict[str, Any]:
text = (arguments_json or "").strip()
if not text:
@@ -194,7 +202,7 @@ class FalAnyEndpoint:
return float("nan")
return stable_hash(kwargs)
def run(
def _run_sync(
self,
endpoint_id: str,
arguments_json: str = "{}",
@@ -205,9 +213,7 @@ class FalAnyEndpoint:
seed: int = -1,
force_rerun: bool = False,
) -> tuple[Any, Any, Any, str]:
endpoint = (endpoint_id or "").strip()
if not endpoint:
raise FalApiError("(any endpoint)", "endpoint_id is required")
endpoint = _validated_endpoint(endpoint_id)
arguments = build_overlay_arguments(
endpoint, arguments_json, image, image_2, video, audio, seed
@@ -218,3 +224,40 @@ class FalAnyEndpoint:
)
return extract_flexible_outputs(result)
async def _run_async(
self,
endpoint_id: str,
arguments_json: str = "{}",
image: Any = None,
image_2: Any = None,
video: Any = None,
audio: Any = None,
seed: int = -1,
force_rerun: bool = False,
) -> tuple[Any, Any, Any, str]:
endpoint = _validated_endpoint(endpoint_id)
# Media uploads (build_overlay_arguments) and result downloads
# (extract_flexible_outputs) are blocking HTTP, so both run in worker
# threads; the fal call awaits on the loop so other branches proceed.
arguments = await asyncio.to_thread(
build_overlay_arguments,
endpoint,
arguments_json,
image,
image_2,
video,
audio,
seed,
)
result = await ApiHandler.submit_and_get_result_async(
endpoint, arguments, skip_cache=bool(force_rerun)
)
return await asyncio.to_thread(extract_flexible_outputs, result)
# On async-capable ComfyUI the executor awaits the coroutine, running
# other graph branches concurrently; older ComfyUI gets the sync path.
run = _run_async if _ASYNC_CAPABLE else _run_sync
+2 -2
View File
@@ -67,7 +67,7 @@ def _multi_enum_argument(endpoint: str, inp: dict[str, Any], value: Any) -> Any
selected = [part.strip() for part in str(value).split(",") if part.strip()]
if not selected:
return None
allowed = set(inp.get("enum") or [])
allowed = {str(member): member for member in inp.get("enum") or []}
invalid = [part for part in selected if part not in allowed]
if invalid:
raise FalApiError(
@@ -75,7 +75,7 @@ def _multi_enum_argument(endpoint: str, inp: dict[str, Any], value: Any) -> Any
f"Invalid value(s) {invalid} for '{inp['name']}'. "
f"Allowed: {', '.join(sorted(allowed))}",
)
return selected
return [allowed[part] for part in selected]
def _enum_argument(inp: dict[str, Any], value: Any, kwargs: dict[str, Any]) -> Any:
+41 -1
View File
@@ -2,7 +2,9 @@
from __future__ import annotations
import asyncio
import hashlib
import importlib.util
import inspect
import re
from functools import cache
@@ -16,6 +18,25 @@ from .schema_to_inputs import build_input_types
NODE_KEY_PREFIX = "FalAPI_"
def _detect_async_capable() -> bool:
"""True when the host ComfyUI awaits coroutine node FUNCTIONs.
``comfy_execution/utils.py`` was introduced by the exact commit that added
async node support (Comfy-Org/ComfyUI commit 2b653e8c18, PR #8830,
2025-07-10) and has not been touched since, so its presence is a precise
import-time proxy for ``_async_map_node_over_list`` existing in the
executor. Must never raise outside ComfyUI: a missing ``comfy_execution``
package (tests, older ComfyUI) simply selects the sync path.
"""
try:
return importlib.util.find_spec("comfy_execution.utils") is not None
except Exception:
return False
_ASYNC_CAPABLE = _detect_async_capable()
def node_key(model: dict[str, Any]) -> str:
return NODE_KEY_PREFIX + model["endpoint_id"].replace("/", "-")
@@ -77,6 +98,15 @@ def _call_api(endpoint_id: str, arguments: dict[str, Any], skip_cache: bool) ->
return submit(endpoint_id, arguments)
async def _call_api_async(
endpoint_id: str, arguments: dict[str, Any], skip_cache: bool
) -> Any:
submit = ApiHandler.submit_and_get_result_async
if _accepts_skip_cache(submit):
return await submit(endpoint_id, arguments, skip_cache=skip_cache)
return await submit(endpoint_id, arguments)
def _class_name(model: dict[str, Any]) -> str:
return re.sub(r"[^0-9A-Za-z_]", "_", node_key(model))
@@ -109,6 +139,16 @@ def build_node_class(model: dict[str, Any]) -> type:
result = _call_api(endpoint_id, arguments, bool(kwargs.get("force_rerun")))
return process_result(model, result)
async def run_async(self: Any, **kwargs: Any) -> tuple:
# build_arguments uploads media and process_result downloads results —
# blocking HTTP — so both run in worker threads; only the fal call
# itself awaits on the loop, letting other graph branches proceed.
arguments = await asyncio.to_thread(build_arguments, model, kwargs)
result = await _call_api_async(
endpoint_id, arguments, bool(kwargs.get("force_rerun"))
)
return await asyncio.to_thread(process_result, model, result)
attrs = {
"INPUT_TYPES": classmethod(input_types),
"IS_CHANGED": classmethod(is_changed),
@@ -117,7 +157,7 @@ def build_node_class(model: dict[str, Any]) -> type:
"FUNCTION": "run",
"CATEGORY": f"FAL/Models/{category}",
"DESCRIPTION": _description(model),
"run": run,
"run": run_async if _ASYNC_CAPABLE else run,
"_FAL_ENDPOINT_ID": endpoint_id,
}
return type(_class_name(model), (object,), attrs)
+152 -6
View File
@@ -16,6 +16,7 @@ from .factory import build_display_name, build_node_class, node_key
_REGISTRY_FILENAME = "fal_registry.json"
_FIXTURE_FILENAME = "_fixture_registry.json"
_FEATURED_FILENAME = "featured_models.json"
Mappings = tuple[dict[str, type], dict[str, str]]
@@ -28,6 +29,11 @@ def _registry_path() -> Path:
return package_dir / _FIXTURE_FILENAME
def _featured_path() -> Path:
package_dir = Path(__file__).resolve().parent
return package_dir.parents[1] / "data" / _FEATURED_FILENAME
def _truthy(value: Any) -> bool:
if isinstance(value, str):
return value.strip().lower() in ("1", "true", "yes", "on")
@@ -61,6 +67,87 @@ def _read_models() -> list[dict[str, Any]]:
return []
def _read_featured() -> dict[str, str | None]:
"""Curated featured tier: {endpoint_id: display_name_override_or_None}.
Empty dict when the tier is disabled, the file is missing, or unreadable.
"""
if not _truthy(_get_setting("dynamic_nodes", "featured_tier", True)):
logger.info("Featured fal node tier disabled via config")
return {}
path = _featured_path()
try:
with open(path, encoding="utf-8") as handle:
document = json.load(handle)
entries = document.get("featured", [])
if not isinstance(entries, list):
raise ValueError("'featured' is not a list")
featured: dict[str, str | None] = {}
for entry in entries:
if not isinstance(entry, dict) or not entry.get("endpoint_id"):
continue
override = entry.get("display_name")
featured = {
**featured,
str(entry["endpoint_id"]): str(override) if override else None,
}
return featured
except Exception as err:
logger.debug("No featured fal models applied (%s): %s", path, err)
return {}
def _superseded_map(models: list[dict[str, Any]]) -> dict[str, tuple[str, str]]:
"""{endpoint_id: (newest_endpoint_id, newest_published_date)} per family.
Conservative: models are grouped by (family, category) only when the
registry declares a non-empty ``family`` (no fuzzy title matching), and a
model is flagged only when its group has >1 member and its published_at is
strictly older than the group's newest.
"""
groups: dict[tuple[str, str], list[dict[str, Any]]] = {}
for model in models:
if _truthy(model.get("deprecated")):
continue
family = str(model.get("family") or "").strip()
if not family or not model.get("endpoint_id"):
continue
group_key = (family, str(model.get("category") or ""))
groups = {**groups, group_key: groups.get(group_key, []) + [model]}
superseded: dict[str, tuple[str, str]] = {}
for members in groups.values():
if len(members) < 2:
continue
newest = max(members, key=lambda m: str(m.get("published_at") or ""))
newest_date = str(newest.get("published_at") or "")
if not newest_date:
continue
for model in members:
if str(model.get("published_at") or "") < newest_date:
superseded = {
**superseded,
str(model["endpoint_id"]): (str(newest["endpoint_id"]), newest_date[:10]),
}
return superseded
def _apply_superseded_note(node_class: type, newest_id: str, newest_date: str) -> None:
"""Prefix the class DESCRIPTION with a newer-release warning."""
note = f"Superseded: a newer release exists in this family: {newest_id} ({newest_date})"
existing = str(getattr(node_class, "DESCRIPTION", "") or "")
node_class.DESCRIPTION = f"{note}\n\n{existing}".rstrip()
def _apply_deprecated_note(node_class: type, reason: str) -> None:
"""Mark a compatibility-only endpoint without breaking its node key."""
note = "Compatibility node: this endpoint is absent from the latest live fal catalog."
if reason:
note = f"{note} {reason}"
existing = str(getattr(node_class, "DESCRIPTION", "") or "")
node_class.DESCRIPTION = f"{note}\n\n{existing}".rstrip()
def _unique_display_name(name: str, used: set[str]) -> str:
if name not in used:
return name
@@ -71,12 +158,19 @@ def _unique_display_name(name: str, used: set[str]) -> str:
def _build_model_mappings(
models: list[dict[str, Any]], categories: set[str]
) -> tuple[dict[str, type], dict[str, str], int]:
models: list[dict[str, Any]],
categories: set[str],
featured: dict[str, str | None] | None = None,
superseded: dict[str, tuple[str, str]] | None = None,
) -> tuple[dict[str, type], dict[str, str], int, int, int]:
classes: dict[str, type] = {}
display: dict[str, str] = {}
used_names: set[str] = {ANY_ENDPOINT_DISPLAY_NAME}
featured = featured or {}
superseded = superseded or {}
skipped = 0
flagged = 0
deprecated_count = 0
for model in models:
try:
@@ -88,7 +182,28 @@ def _build_model_mappings(
logger.debug("Duplicate dynamic node key skipped: %s", key)
continue
node_class = build_node_class(model)
name = _unique_display_name(build_display_name(model), used_names)
endpoint_id = str(model.get("endpoint_id") or "")
preferred = build_display_name(model)
deprecated = _truthy(model.get("deprecated"))
if deprecated:
category = str(model.get("category") or "other")
node_class.CATEGORY = f"FAL/Compatibility/{category}"
preferred = f"[Unavailable] {preferred}"
_apply_deprecated_note(
node_class, str(model.get("deprecated_reason") or "").strip()
)
deprecated_count += 1
elif endpoint_id in featured:
category = str(model.get("category") or "other")
node_class.CATEGORY = f"FAL/Featured/{category}"
preferred = featured[endpoint_id] or preferred
if not deprecated and endpoint_id in superseded:
newest_id, newest_date = superseded[endpoint_id]
_apply_superseded_note(node_class, newest_id, newest_date)
flagged += 1
name = _unique_display_name(preferred, used_names)
classes = {**classes, key: node_class}
display = {**display, key: name}
used_names.add(name)
@@ -100,7 +215,26 @@ def _build_model_mappings(
err,
)
return classes, display, skipped
return classes, display, skipped, flagged, deprecated_count
def _log_missing_featured(featured: dict[str, str | None], models: list[dict[str, Any]]) -> int:
"""Debug-log featured ids absent from the registry; returns how many matched."""
registry_ids = {str(m.get("endpoint_id") or "") for m in models}
missing = [endpoint_id for endpoint_id in featured if endpoint_id not in registry_ids]
for endpoint_id in missing:
logger.debug("Featured model not in registry, skipped: %s", endpoint_id)
return len(featured) - len(missing)
def _schedule_freshness_check() -> None:
"""Kick off the delayed registry freshness check; never raises."""
try:
from ..utils.freshness import schedule_startup_check
schedule_startup_check()
except Exception as err:
logger.debug("Could not schedule registry freshness check: %s", err)
def load_dynamic_mappings() -> Mappings:
@@ -112,14 +246,26 @@ def load_dynamic_mappings() -> Mappings:
categories = _category_filter()
models = _read_models()
classes, display, skipped = _build_model_mappings(models, categories)
featured = _read_featured()
featured_count = _log_missing_featured(featured, models)
superseded = _superseded_map(models)
classes, display, skipped, flagged, deprecated_count = _build_model_mappings(
models, categories, featured=featured, superseded=superseded
)
all_classes = {ANY_ENDPOINT_KEY: FalAnyEndpoint, **classes}
all_display = {ANY_ENDPOINT_KEY: ANY_ENDPOINT_DISPLAY_NAME, **display}
logger.info(
"Registered %d dynamic fal nodes (skipped %d)", len(all_classes), skipped
"Registered %d dynamic fal nodes (skipped %d, featured %d, "
"%d superseded, %d compatibility-preserved)",
len(all_classes),
skipped,
featured_count,
flagged,
deprecated_count,
)
_schedule_freshness_check()
return all_classes, all_display
except Exception as err:
logger.error("Dynamic fal node loading failed entirely: %s", err)
+5 -1
View File
@@ -71,7 +71,7 @@ def _multi_enum_spec(inp: dict[str, Any]) -> tuple[Any, ...]:
default = inp.get("default")
text = ", ".join(str(v) for v in default) if isinstance(default, list) else ""
description = (inp.get("description") or "").strip()
tooltip = f"{description} Comma-separated. Options: {', '.join(values)}".strip()
tooltip = f"{description} Comma-separated. Options: {', '.join(str(value) for value in values)}".strip()
return ("STRING", {"default": text, "tooltip": tooltip})
@@ -113,6 +113,10 @@ def _string_spec(inp: dict[str, Any]) -> tuple[Any, ...]:
"default": default if isinstance(default, str) else "",
"multiline": bool(inp.get("multiline")),
}
if inp.get("suggestions"):
# Keep STRING at the API/socket boundary. The frontend presents an
# editable dropdown, preserving saved text values and STRING links.
opts.update(fal_suggestions=list(inp["suggestions"]), multiline=False)
return ("STRING", _with_tooltip(opts, inp.get("description")))
+2
View File
@@ -9,6 +9,7 @@ The implementations now live in the ``nodes/utils`` package.
from .utils import (
ApiHandler,
ArchiveUtils,
BillingUtils,
FalApiError,
FalConfig,
@@ -25,6 +26,7 @@ from .utils import (
__all__ = [
"ApiHandler",
"ArchiveUtils",
"BillingUtils",
"FalApiError",
"FalConfig",
+133
View File
@@ -11,6 +11,7 @@ from __future__ import annotations
import json
import os
import threading
import time
from typing import Any, Callable
from .utils.billing import BillingUtils
@@ -188,11 +189,128 @@ def _search_models(
"title": model.get("title") or model["endpoint_id"],
"category": model.get("category"),
"label": (info or {}).get("label"),
"thumbnail": model.get("thumbnail") or None,
}
for model, info in hits[:capped]
]
# -- registry freshness + refresh -----------------------------------------------
_RESTART_NOTE = "Restart ComfyUI and reload the browser after the refresh to load updated controls and new nodes."
_REFRESH_TIMEOUT_S = 1800
_refresh_lock = threading.Lock()
_refresh_state: dict[str, Any] = {
"running": False,
"started_at": None,
"finished_at": None,
"ok": None,
"message": "Registry refresh has not been started.",
}
def _repo_root() -> str:
nodes_dir = os.path.dirname(os.path.abspath(__file__))
return os.path.dirname(nodes_dir)
def _registry_status() -> dict[str, Any]:
"""Cached diff of the live fal catalog vs. the local registry (may fetch)."""
from .utils.freshness import check_for_new_models
return check_for_new_models(timeout_s=20)
def _refresh_status() -> dict[str, Any]:
"""Snapshot of the background registry-refresh state."""
with _refresh_lock:
return {**_refresh_state, "restart_note": _RESTART_NOTE}
def _run_refresh_subprocess() -> tuple[bool, str]:
"""Build and validate a candidate before atomically replacing the registry."""
import subprocess
import sys
import tempfile
root = _repo_root()
baseline = os.path.join(root, "data", "fal_registry.json")
with tempfile.TemporaryDirectory(prefix="fal-registry-", dir=os.path.dirname(baseline)) as workdir:
candidate = os.path.join(workdir, "candidate.json")
commands = [
[sys.executable, os.path.join(root, "scripts", "build_registry.py"),
"--out", candidate, "--preserve-from", baseline],
[sys.executable, os.path.join(root, "scripts", "validate_registry.py"),
candidate, "--baseline", baseline],
]
for command in commands:
completed = subprocess.run(
command, cwd=root, capture_output=True, text=True, timeout=_REFRESH_TIMEOUT_S
)
if completed.returncode != 0:
tail = (completed.stderr or completed.stdout or "").strip()[-500:]
return False, f"{os.path.basename(command[1])} exited with {completed.returncode}: {tail}"
os.replace(candidate, baseline)
return True, f"Registry refreshed. {_RESTART_NOTE}"
def _finish_refresh(ok: bool, message: str) -> None:
global _refresh_state
with _refresh_lock:
_refresh_state = {
**_refresh_state,
"running": False,
"finished_at": time.time(),
"ok": ok,
"message": message,
}
def _refresh_worker(runner: Callable[[], tuple[bool, str]]) -> None:
"""Run the refresh and record the outcome. Never raises."""
try:
ok, message = runner()
except Exception as exc:
logger.warning("server_routes: registry refresh failed: %s", exc)
ok, message = False, f"Registry refresh failed: {exc}"
_finish_refresh(ok, message)
logger.info("server_routes: registry refresh finished (ok=%s): %s", ok, message)
def _start_refresh(
runner: Callable[[], tuple[bool, str]] | None = None,
spawn: Callable[[Callable[[], None]], None] | None = None,
) -> dict[str, Any]:
"""Start a background registry rebuild; no-op when one is already running.
``runner``/``spawn`` are injectable for tests (stub subprocess / run inline).
"""
global _refresh_state
with _refresh_lock:
if _refresh_state["running"]:
return {"started": False, **_refresh_state, "restart_note": _RESTART_NOTE}
_refresh_state = {
**_refresh_state,
"running": True,
"started_at": time.time(),
"finished_at": None,
"ok": None,
"message": "Registry refresh running — rebuilding data/fal_registry.json...",
}
active_runner = runner or _run_refresh_subprocess
def work() -> None:
_refresh_worker(active_runner)
if spawn is not None:
spawn(work)
else:
threading.Thread(target=work, name="fal-registry-refresh", daemon=True).start()
return {"started": True, **_refresh_status()}
def _cancel(endpoint_id: str, request_id: str) -> dict[str, Any]:
"""Best-effort cancel of a queued fal request via fal_client. Never raises."""
endpoint = (endpoint_id or "").strip()
@@ -280,6 +398,18 @@ async def models_route(request: Any) -> Any:
)
async def registry_status_route(request: Any) -> Any:
return _guarded(_registry_status, "/fal_api/registry_status")
async def registry_refresh_start_route(request: Any) -> Any:
return _guarded(_start_refresh, "/fal_api/registry_refresh")
async def registry_refresh_status_route(request: Any) -> Any:
return _guarded(_refresh_status, "/fal_api/registry_refresh")
async def cancel_route(request: Any) -> Any:
try:
body = await request.json()
@@ -299,6 +429,9 @@ ROUTES: tuple[tuple[str, str, Callable[..., Any]], ...] = (
("GET", "/fal_api/jobs", jobs_route),
("GET", "/fal_api/balance", balance_route),
("GET", "/fal_api/models", models_route),
("GET", "/fal_api/registry_status", registry_status_route),
("GET", "/fal_api/registry_refresh", registry_refresh_status_route),
("POST", "/fal_api/registry_refresh", registry_refresh_start_route),
("POST", "/fal_api/cancel", cancel_route),
)
+5 -36
View File
@@ -1,46 +1,15 @@
import os
import tempfile
import zipfile
import torch
from PIL import Image
from .fal_utils import ApiHandler, FalConfig, ImageUtils
from .fal_utils import ApiHandler, ArchiveUtils, FalConfig
# Initialize FalConfig
fal_config = FalConfig()
def create_zip_from_images(images):
"""Create a zip file from a list of images."""
"""Create a zip file from a list of images and upload it (returns the URL)."""
try:
with tempfile.NamedTemporaryFile(suffix=".zip", delete=False) as temp_zip:
with zipfile.ZipFile(temp_zip, "w") as zf:
for idx, img_tensor in enumerate(images):
# Convert tensor to PIL Image
if isinstance(img_tensor, torch.Tensor):
# Convert to numpy and scale to 0-255 range
img_np = (img_tensor.cpu().numpy() * 255).astype("uint8")
# Handle different tensor formats
if img_np.shape[0] == 3: # If in format (C, H, W)
img_np = img_np.transpose(1, 2, 0)
img = Image.fromarray(img_np)
else:
img = img_tensor
# Save image to temporary file
with tempfile.NamedTemporaryFile(
suffix=".png", delete=False
) as temp_img:
img.save(temp_img, format="PNG")
temp_img_path = temp_img.name
# Add to zip file
zf.write(temp_img_path, f"image_{idx}.png")
os.unlink(temp_img_path)
# Upload the zip through the shared utility (raises on failure)
return ImageUtils.upload_file(temp_zip.name)
zip_path = ArchiveUtils.zip_images(images)
# Upload the zip through the shared utility (raises on failure)
return ArchiveUtils.upload_zip(zip_path)
except Exception as e:
return ApiHandler.handle_text_generation_error(
"flux-lora-fast-training", f"Failed to create zip file: {str(e)}"
+259
View File
@@ -0,0 +1,259 @@
"""Data utility nodes: JSON path extraction, prompt line cycling, text templating."""
from __future__ import annotations
import json
import re
from typing import Any
from .fal_utils import FalApiError, logger
_CATEGORY = "FAL/Utils/Data"
_MAX_INDEX = 2**31 - 1
_MISSING = object()
# a path segment is an optional key name followed by zero or more [N] indices
_PATH_SEGMENT = re.compile(r"^([^\[\]]*)((?:\[\d+\])*)$")
_BRACKET_INDEX = re.compile(r"\[(\d+)\]")
def _tokenize_path(path: str) -> list[Any]:
"""Split a dot/bracket path into str keys and int indices. No eval."""
tokens: list[Any] = []
for part in path.split("."):
segment = part.strip()
if not segment:
continue
match = _PATH_SEGMENT.match(segment)
if match is None:
# contract: anything that can't resolve returns the default,
# a malformed segment included — it can never match a key anyway
logger.debug("FalJSONExtract: unparseable path segment %r in %r", segment, path)
return None
name, brackets = match.group(1), match.group(2)
if name:
tokens.append(name)
tokens.extend(int(index) for index in _BRACKET_INDEX.findall(brackets))
return tokens
def _value_to_bool(value: Any) -> bool:
"""Truthiness with JSON-string awareness: "false"/"0"/"no"/"" are False."""
if isinstance(value, str):
return value.strip().lower() not in ("", "false", "0", "no", "none", "null")
return bool(value)
def _walk_path(value: Any, tokens: list[Any]) -> Any:
"""Follow tokens through nested dicts/lists; return _MISSING when absent."""
current = value
for token in tokens:
index = token if isinstance(token, int) else None
if index is None and isinstance(current, list) and str(token).isdigit():
index = int(token) # bare integer segment indexing an array
if index is not None:
if isinstance(current, list) and 0 <= index < len(current):
current = current[index]
else:
return _MISSING
elif isinstance(current, dict) and token in current:
current = current[token]
else:
return _MISSING
return current
def _value_to_text(value: Any) -> str:
"""Strings pass through; everything else is re-serialized as JSON."""
if isinstance(value, str):
return value
return json.dumps(value)
def _value_to_number(value: Any) -> float:
"""Coerce a value to float; anything non-numeric becomes 0.0."""
if isinstance(value, bool):
return 1.0 if value else 0.0
if isinstance(value, (int, float)):
return float(value)
if isinstance(value, str):
try:
return float(value.strip())
except ValueError:
return 0.0
return 0.0
class FalJSONExtract:
"""Pull a value out of a JSON result by dot/bracket path."""
RETURN_TYPES = ("STRING", "FLOAT", "BOOLEAN")
RETURN_NAMES = ("text", "number", "boolean")
FUNCTION = "extract"
CATEGORY = _CATEGORY
DESCRIPTION = (
"Extract a value from JSON text by path (e.g. video.url, "
"images[0].url). Returns it as text, number, and boolean so it can "
"wire straight into other nodes. Missing paths return the default "
"instead of failing the graph."
)
@classmethod
def INPUT_TYPES(cls) -> dict[str, Any]:
return {
"required": {
"json_text": (
"STRING",
{
"forceInput": True,
"multiline": True,
"tooltip": (
"Wire result_json from a Fal Any Endpoint / Fal Collect "
"node here to pick values out of the raw API result."
),
},
),
"path": (
"STRING",
{
"default": "video.url",
"tooltip": (
"Dot/bracket path into the JSON, e.g. video.url, "
"images[0].url, data.items[2].name. Bare integers "
"also index arrays (images.0.url)."
),
},
),
"default": (
"STRING",
{
"default": "",
"tooltip": "Returned as text when the path is missing (not an error)",
},
),
},
}
def extract(self, json_text: str, path: str = "video.url", default: str = "") -> tuple[str, float, bool]:
try:
payload = json.loads(json_text)
except Exception as exc:
logger.error("FalJSONExtract: invalid JSON input: %s", exc)
raise FalApiError("FalJSONExtract", f"Input is not valid JSON: {exc}") from exc
tokens = _tokenize_path(path or "")
value = _walk_path(payload, tokens) if tokens is not None else _MISSING
if value is _MISSING:
logger.debug("FalJSONExtract: path %r missing, returning default", path)
value = default
return (_value_to_text(value), _value_to_number(value), _value_to_bool(value))
class FalPromptLines:
"""Cycle through a multiline prompt list, one line per run."""
RETURN_TYPES = ("STRING", "INT", "INT")
RETURN_NAMES = ("line", "index", "total")
FUNCTION = "pick"
CATEGORY = _CATEGORY
DESCRIPTION = (
"Pick one line from a multiline text by index. The index wraps "
"around (modulo the number of lines), so with control_after_generate "
"set to increment it cycles through your prompt list forever."
)
@classmethod
def INPUT_TYPES(cls) -> dict[str, Any]:
return {
"required": {
"text": (
"STRING",
{
"default": "",
"multiline": True,
"tooltip": "Prompt list, one prompt per line",
},
),
"index": (
"INT",
{
"default": 0,
"min": 0,
"max": _MAX_INDEX,
"control_after_generate": True,
"tooltip": (
"Which line to pick (wraps around). Set the control to "
"'increment' to iterate through your prompt list run by run."
),
},
),
"skip_blank": (
"BOOLEAN",
{"default": True, "tooltip": "Ignore empty/whitespace-only lines"},
),
},
}
def pick(self, text: str, index: int = 0, skip_blank: bool = True) -> tuple[str, int, int]:
lines = (text or "").splitlines()
if skip_blank:
lines = [line for line in lines if line.strip()]
total = len(lines)
if total == 0:
return ("", 0, 0)
effective = int(index) % total
return (lines[effective], effective, total)
class FalTextTemplate:
"""Fill a text template's {a}..{d} placeholders from string inputs."""
RETURN_TYPES = ("STRING",)
FUNCTION = "render"
CATEGORY = _CATEGORY
DESCRIPTION = (
"Substitute {a}, {b}, {c}, {d} placeholders in a template with the "
"connected string inputs — quick prompt assembly without string "
"concatenation chains. Missing inputs become empty text."
)
@classmethod
def INPUT_TYPES(cls) -> dict[str, Any]:
return {
"required": {
"template": (
"STRING",
{
"default": "a photo of {a}, {b} style",
"multiline": True,
"tooltip": "Template text; {a} {b} {c} {d} are replaced with the inputs below",
},
),
},
"optional": {
"a": ("STRING", {"default": "", "tooltip": "Value for {a}"}),
"b": ("STRING", {"default": "", "tooltip": "Value for {b}"}),
"c": ("STRING", {"default": "", "tooltip": "Value for {c}"}),
"d": ("STRING", {"default": "", "tooltip": "Value for {d}"}),
},
}
def render(self, template: str, a: str = "", b: str = "", c: str = "", d: str = "") -> tuple[str]:
result = template or ""
for key, value in (("a", a), ("b", b), ("c", c), ("d", d)):
result = result.replace("{" + key + "}", value or "")
return (result,)
NODE_CLASS_MAPPINGS = {
"FalJSONExtract_fal": FalJSONExtract,
"FalPromptLines_fal": FalPromptLines,
"FalTextTemplate_fal": FalTextTemplate,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"FalJSONExtract_fal": "JSON Extract (fal)",
"FalPromptLines_fal": "Prompt Lines (fal)",
"FalTextTemplate_fal": "Text Template (fal)",
}
+409
View File
@@ -0,0 +1,409 @@
"""Dataset preparation utility nodes (zip building, frame extraction, captioning)."""
from __future__ import annotations
import os
import tempfile
import zipfile
from concurrent.futures import ThreadPoolExecutor
from typing import Any
from .fal_utils import (
ApiHandler,
ArchiveUtils,
FalApiError,
FalConfig,
ImageUtils,
MediaUtils,
logger,
)
# Initialize FalConfig
fal_config = FalConfig()
_CATEGORY = "FAL/Utils/Dataset"
_ARCHIVE_MODEL = "archive"
_VISION_ENDPOINT = "openrouter/router/vision"
_MAX_CAPTION_WORKERS = 8
_STREAM_CHUNK_SIZE = 1 << 20 # 1 MiB
def _safe_unlink(path: str | None) -> None:
"""Delete a temp file, ignoring errors."""
if path is None:
return
try:
os.unlink(path)
except OSError:
pass
def _split_caption_lines(captions: str) -> list[str] | None:
"""Split a multiline caption field into one caption per line (None if empty)."""
if not captions or not captions.strip():
return None
return captions.splitlines()
def _video_to_local_path(video: Any) -> tuple[str, bool]:
"""Resolve a VIDEO input to a local file path. Returns (path, is_temp)."""
source = video.get_stream_source() if hasattr(video, "get_stream_source") else video
if isinstance(source, str):
if source.startswith(("http://", "https://")):
return MediaUtils.download_url_to_temp(source, ".mp4"), True
if not os.path.isfile(source):
raise FalApiError(
_ARCHIVE_MODEL,
f"Video file not found: {source}. Connect a valid VIDEO input.",
)
return source, False
if hasattr(source, "read"):
temp_path: str | None = None
try:
with tempfile.NamedTemporaryFile(suffix=".mp4", delete=False) as temp_file:
temp_path = temp_file.name
while True:
chunk = source.read(_STREAM_CHUNK_SIZE)
if not chunk:
break
temp_file.write(chunk)
return temp_path, True
except Exception as exc:
_safe_unlink(temp_path)
raise FalApiError(
_ARCHIVE_MODEL, f"Failed to buffer video stream to disk: {exc}"
) from exc
raise FalApiError(
_ARCHIVE_MODEL,
"Unsupported VIDEO input: could not resolve a local file, URL, or stream from it.",
)
def _extract_frames_to_zip(video_path: str, every_nth: int, max_frames: int) -> str:
"""Decode a video with cv2, sample every Nth frame as PNG into a zip, return the zip path."""
try:
import cv2
except ImportError as exc:
raise FalApiError(
_ARCHIVE_MODEL,
"opencv-python is required to extract video frames. "
"Install it with 'pip install opencv-python'.",
) from exc
capture = cv2.VideoCapture(video_path)
if not capture.isOpened():
raise FalApiError(
_ARCHIVE_MODEL,
f"Could not open video for decoding: {video_path}. "
"Check that the input is a valid video file.",
)
zip_path: str | None = None
try:
with tempfile.NamedTemporaryFile(suffix=".zip", delete=False) as temp_zip:
zip_path = temp_zip.name
saved = 0
index = 0
with zipfile.ZipFile(zip_path, "w") as zip_file:
while saved < max_frames:
ok, frame = capture.read()
if not ok:
break
if index % every_nth == 0:
encoded, buffer = cv2.imencode(".png", frame)
if not encoded:
raise FalApiError(
_ARCHIVE_MODEL, f"Failed to encode frame {index} as PNG."
)
zip_file.writestr(f"frame_{saved:05d}.png", buffer.tobytes())
saved += 1
index += 1
if saved == 0:
raise FalApiError(
_ARCHIVE_MODEL,
"No frames could be decoded from the video. "
"Check the input video and the every_nth setting.",
)
logger.info("Extracted %d frame(s) from %s", saved, video_path)
return zip_path
except Exception:
_safe_unlink(zip_path)
raise
finally:
capture.release()
class FalImagesToZipURL:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"images": (
"IMAGE",
{
"tooltip": "Images to package as a training dataset zip (image_0.png, image_1.png, ...).",
},
),
},
"optional": {
"captions": (
"STRING",
{
"default": "",
"multiline": True,
"tooltip": "Optional captions, one per line (blank lines allowed). "
"Line count must match the image batch size, or leave empty for no captions. "
"Written as image_0.txt, image_1.txt, ... next to each image.",
},
),
"name_prefix": (
"STRING",
{
"default": "image",
"tooltip": "File name prefix inside the zip (e.g. 'image' -> image_0.png).",
},
),
},
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("zip_url",)
FUNCTION = "create_zip_url"
CATEGORY = _CATEGORY
DESCRIPTION = (
"Zips an IMAGE batch (with optional per-image captions) and uploads it to fal.ai. "
"Feed the URL directly into the LoRA trainer nodes' images_data_url input."
)
def create_zip_url(self, images, captions="", name_prefix="image"):
caption_lines = _split_caption_lines(captions)
zip_path = ArchiveUtils.zip_images(
images, captions=caption_lines, name_prefix=name_prefix
)
return (ArchiveUtils.upload_zip(zip_path),)
class FalFolderToZipURL:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"folder_path": (
"STRING",
{
"default": "",
"tooltip": "Path to a local folder whose files will be zipped and uploaded.",
},
),
"recursive": (
"BOOLEAN",
{
"default": False,
"tooltip": "Also include files from subfolders (hidden entries are always skipped).",
},
),
"extensions": (
"STRING",
{
"default": ".png,.jpg,.jpeg,.webp,.txt",
"tooltip": "Comma-separated list of file extensions to include. Empty includes all files.",
},
),
},
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("zip_url",)
FUNCTION = "create_zip_url"
CATEGORY = _CATEGORY
DESCRIPTION = "Zips a local folder and uploads it to fal.ai, returning the zip URL."
def create_zip_url(self, folder_path, recursive=False, extensions=".png,.jpg,.jpeg,.webp,.txt"):
extension_list = [part.strip() for part in extensions.split(",") if part.strip()]
zip_path = ArchiveUtils.zip_folder(
folder_path,
include_extensions=extension_list or None,
recursive=recursive,
)
return (ArchiveUtils.upload_zip(zip_path),)
class FalVideoToFrameDatasetZip:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"video": (
"VIDEO",
{
"tooltip": "Video to sample frames from for a training dataset.",
},
),
"every_nth": (
"INT",
{
"default": 10,
"min": 1,
"max": 10000,
"step": 1,
"tooltip": "Keep one frame out of every N decoded frames.",
},
),
"max_frames": (
"INT",
{
"default": 200,
"min": 1,
"max": 2000,
"step": 1,
"tooltip": "Stop after this many frames have been saved.",
},
),
},
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("zip_url",)
FUNCTION = "create_zip_url"
CATEGORY = _CATEGORY
DESCRIPTION = (
"Samples frames from a video (every Nth frame, up to max_frames), zips them as PNGs, "
"uploads the zip to fal.ai, and returns the URL."
)
def create_zip_url(self, video, every_nth=10, max_frames=200):
local_path, is_temp = _video_to_local_path(video)
try:
zip_path = _extract_frames_to_zip(local_path, every_nth, max_frames)
finally:
if is_temp:
_safe_unlink(local_path)
return (ArchiveUtils.upload_zip(zip_path),)
class FalBatchCaption:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"images": (
"IMAGE",
{
"tooltip": "Images to caption, one caption per frame.",
},
),
"prompt": (
"STRING",
{
"default": "Describe this image for LoRA training in one dense sentence.",
"multiline": True,
"tooltip": "Instruction sent to the vision model for each image.",
},
),
"model": (
[
"google/gemini-2.5-flash",
"anthropic/claude-sonnet-4.5",
"openai/gpt-4o",
"custom",
],
{
"default": "google/gemini-2.5-flash",
"tooltip": "Vision model to use. Select 'custom' to type any OpenRouter model id "
"in custom_model_name.",
},
),
},
"optional": {
"custom_model_name": (
"STRING",
{
"default": "",
"tooltip": "OpenRouter model id used when model is set to 'custom'.",
},
),
},
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("captions",)
FUNCTION = "caption_images"
CATEGORY = _CATEGORY
DESCRIPTION = (
"Captions each image in a batch with a fal VLM (concurrently, order preserved) and returns "
"one caption per line — wire straight into FalImagesToZipURL's captions input."
)
def caption_images(
self,
images,
prompt="Describe this image for LoRA training in one dense sentence.",
model="google/gemini-2.5-flash",
custom_model_name="",
):
if model == "custom":
if not custom_model_name or not custom_model_name.strip():
raise FalApiError(
_VISION_ENDPOINT,
"custom_model_name is required when model is set to 'custom'.",
)
model = custom_model_name.strip()
image_urls = ImageUtils.prepare_images(images)
if not image_urls:
raise FalApiError(
_VISION_ENDPOINT, "No images provided to caption. Connect an IMAGE batch."
)
def caption_one(image_url: str) -> str:
arguments = {
"model": model,
"prompt": prompt,
"image_urls": [image_url],
"stream": False,
}
result = ApiHandler.submit_and_get_result(_VISION_ENDPOINT, arguments)
# Captions are joined by newline, so flatten any multiline output.
return str(result["output"]).replace("\r", " ").replace("\n", " ").strip()
with ThreadPoolExecutor(max_workers=_MAX_CAPTION_WORKERS) as executor:
futures = [executor.submit(caption_one, url) for url in image_urls]
captions: list[str] = []
failure_count = 0
for index, future in enumerate(futures):
try:
captions = [*captions, future.result()]
except Exception as exc:
# a user Cancel raised inside a worker must stop the node,
# not silently become an empty caption
if exc.__class__.__name__ == "InterruptProcessingException":
raise
logger.warning("Caption for image %d failed: %s", index, exc)
captions = [*captions, ""]
failure_count += 1
if failure_count == len(image_urls):
raise FalApiError(
_VISION_ENDPOINT,
f"All {len(image_urls)} caption request(s) failed. "
"Check the model id, your fal API key, and the queue logs above.",
)
return ("\n".join(captions),)
# Node class mappings
NODE_CLASS_MAPPINGS = {
"FalImagesToZipURL_fal": FalImagesToZipURL,
"FalFolderToZipURL_fal": FalFolderToZipURL,
"FalVideoToFrameDatasetZip_fal": FalVideoToFrameDatasetZip,
"FalBatchCaption_fal": FalBatchCaption,
}
# Node display name mappings
NODE_DISPLAY_NAME_MAPPINGS = {
"FalImagesToZipURL_fal": "Images → Training ZIP URL (fal)",
"FalFolderToZipURL_fal": "Folder → ZIP URL (fal)",
"FalVideoToFrameDatasetZip_fal": "Video → Frame Dataset ZIP URL (fal)",
"FalBatchCaption_fal": "Batch Caption Images (fal VLM)",
}
+423
View File
@@ -0,0 +1,423 @@
"""Image utility nodes: labeled grids, preset resizing, and base64 conversion."""
from __future__ import annotations
import base64
import io
import math
import re
from typing import Any
import numpy as np
import torch
from PIL import Image, ImageDraw, ImageFont
from .fal_utils import FalApiError, ImageUtils, logger
_CATEGORY = "FAL/Utils/Image"
# fal image_size preset -> (width, height)
_PRESET_SIZES = {
"square_hd": (1024, 1024),
"square": (512, 512),
"portrait_4_3": (768, 1024),
"portrait_16_9": (576, 1024),
"landscape_4_3": (1024, 768),
"landscape_16_9": (1024, 576),
}
_CUSTOM_PRESET = "custom"
_LANCZOS = getattr(Image, "Resampling", Image).LANCZOS
_GRID_BG = (24, 24, 24)
_LABEL_BG = (16, 16, 16)
_LABEL_FG = (235, 235, 235)
_LABEL_MARGIN = 4
_ELLIPSIS = "..."
_B64_WHITESPACE = re.compile(r"\s+")
_MIME_BY_FORMAT = {"png": "image/png", "jpeg": "image/jpeg", "webp": "image/webp"}
def _pils_to_tensor(pils: list[Image.Image]) -> torch.Tensor:
"""Stack same-sized PIL images into a float32 (B, H, W, 3) IMAGE tensor."""
arrays = [np.array(pil.convert("RGB")).astype(np.float32) / 255.0 for pil in pils]
return torch.from_numpy(np.stack(arrays, axis=0))
def _image_input_to_pils(images: Any) -> list[Image.Image]:
"""Convert an IMAGE input (batch tensor or list of tensors) to PIL images."""
if isinstance(images, torch.Tensor) and images.ndim == 4:
items: list[Any] = [images[i] for i in range(images.shape[0])]
elif isinstance(images, (list, tuple)):
items = list(images)
else:
items = [images]
if not items:
raise FalApiError("FalImageGrid", "IMAGE input contained no images")
return [ImageUtils.tensor_to_pil(item) for item in items]
def _letterbox(pil: Image.Image, width: int, height: int, fill: tuple[int, int, int]) -> Image.Image:
"""Fit an image inside (width, height) preserving aspect, padded with fill."""
scale = min(width / pil.width, height / pil.height)
new_size = (max(1, round(pil.width * scale)), max(1, round(pil.height * scale)))
resized = pil.convert("RGB").resize(new_size, _LANCZOS)
canvas = Image.new("RGB", (width, height), fill)
offset = ((width - new_size[0]) // 2, (height - new_size[1]) // 2)
canvas.paste(resized, offset)
return canvas
def _truncate_label(draw: ImageDraw.ImageDraw, text: str, font: Any, max_width: int) -> str:
"""Truncate text with an ellipsis so it fits within max_width pixels."""
if draw.textlength(text, font=font) <= max_width:
return text
for end in range(len(text) - 1, 0, -1):
candidate = text[:end].rstrip() + _ELLIPSIS
if draw.textlength(candidate, font=font) <= max_width:
return candidate
return _ELLIPSIS
def _draw_label(
canvas: Image.Image, text: str, x: int, y: int, cell_width: int, label_height: int
) -> None:
"""Draw one centered label line on its dark strip below a cell."""
draw = ImageDraw.Draw(canvas)
draw.rectangle((x, y, x + cell_width - 1, y + label_height - 1), fill=_LABEL_BG)
if not text:
return
font = ImageFont.load_default()
fitted = _truncate_label(draw, text, font, cell_width - 2 * _LABEL_MARGIN)
text_width = draw.textlength(fitted, font=font)
bbox = font.getbbox(fitted)
text_height = bbox[3] - bbox[1]
text_x = x + max(_LABEL_MARGIN, (cell_width - text_width) // 2)
text_y = y + max(0, (label_height - text_height) // 2) - bbox[1]
draw.text((text_x, text_y), fitted, font=font, fill=_LABEL_FG)
def _grid_shape(count: int, columns: int) -> tuple[int, int]:
"""Resolve (columns, rows) for a grid; columns == 0 means auto square-ish."""
cols = columns if columns > 0 else math.ceil(math.sqrt(count))
cols = max(1, min(cols, count))
return cols, math.ceil(count / cols)
def _compose_grid(
pils: list[Image.Image], labels: list[str], columns: int, padding: int, label_height: int
) -> Image.Image:
"""Lay out letterboxed cells (plus optional label strips) on a dark canvas."""
cell_w = max(pil.width for pil in pils)
cell_h = max(pil.height for pil in pils)
strip_h = label_height if labels else 0
cols, rows = _grid_shape(len(pils), columns)
total_w = cols * cell_w + (cols + 1) * padding
total_h = rows * (cell_h + strip_h) + (rows + 1) * padding
canvas = Image.new("RGB", (total_w, total_h), _GRID_BG)
for i, pil in enumerate(pils):
col, row = i % cols, i // cols
x = padding + col * (cell_w + padding)
y = padding + row * (cell_h + strip_h + padding)
canvas.paste(_letterbox(pil, cell_w, cell_h, _GRID_BG), (x, y))
if strip_h:
text = labels[i] if i < len(labels) else ""
_draw_label(canvas, text, x, y + cell_h, cell_w, strip_h)
return canvas
class FalImageGrid:
"""Compose an image batch into a single labeled contact-sheet grid."""
RETURN_TYPES = ("IMAGE",)
FUNCTION = "compose"
CATEGORY = _CATEGORY
DESCRIPTION = (
"Arrange a batch of images into one grid image with optional text "
"labels under each cell. Mixed sizes are letterboxed into uniform "
"cells on a dark background — handy for comparing seeds, prompts, "
"or models side by side."
)
@classmethod
def INPUT_TYPES(cls) -> dict[str, Any]:
return {
"required": {
"images": ("IMAGE", {"tooltip": "Batch of images to arrange into a grid"}),
},
"optional": {
"labels": (
"STRING",
{
"default": "",
"multiline": True,
"tooltip": (
"One label per line, matched to images in batch order. "
"Leave empty for no label strips."
),
},
),
"columns": (
"INT",
{
"default": 0,
"min": 0,
"max": 64,
"tooltip": "Number of grid columns; 0 = auto (roughly square)",
},
),
"cell_padding": (
"INT",
{
"default": 8,
"min": 0,
"max": 64,
"tooltip": "Pixels of dark padding around each cell",
},
),
"label_height": (
"INT",
{
"default": 28,
"min": 12,
"max": 128,
"tooltip": "Height in pixels of the label strip under each cell",
},
),
},
}
def compose(
self,
images: Any,
labels: str = "",
columns: int = 0,
cell_padding: int = 8,
label_height: int = 28,
) -> tuple[torch.Tensor]:
pils = _image_input_to_pils(images)
label_lines = [line.strip() for line in labels.splitlines()] if labels.strip() else []
try:
grid = _compose_grid(pils, label_lines, int(columns), int(cell_padding), int(label_height))
except FalApiError:
raise
except Exception as exc:
logger.error("FalImageGrid: failed to compose grid: %s", exc)
raise FalApiError("FalImageGrid", f"Failed to compose image grid: {exc}") from exc
logger.debug("FalImageGrid: composed %d cells into %dx%d", len(pils), grid.width, grid.height)
return (_pils_to_tensor([grid]),)
def _resize_one(pil: Image.Image, width: int, height: int, mode: str) -> Image.Image:
"""Resize a single PIL image to (width, height) using the given mode."""
source = pil.convert("RGB")
if mode == "stretch":
return source.resize((width, height), _LANCZOS)
if mode == "contain_pad":
return _letterbox(source, width, height, (0, 0, 0))
# cover_crop: scale to fully cover the target, then center-crop
scale = max(width / source.width, height / source.height)
scaled = source.resize(
(max(width, round(source.width * scale)), max(height, round(source.height * scale))),
_LANCZOS,
)
left = (scaled.width - width) // 2
top = (scaled.height - height) // 2
return scaled.crop((left, top, left + width, top + height))
class FalResizeToPreset:
"""Resize images to an exact fal image_size preset (or custom dimensions)."""
RETURN_TYPES = ("IMAGE", "INT", "INT")
RETURN_NAMES = ("image", "width", "height")
FUNCTION = "resize"
CATEGORY = _CATEGORY
DESCRIPTION = (
"Resize images to the exact pixel dimensions of a fal image_size "
"preset (square_hd, portrait_16_9, ...) or custom width/height. "
"Choose cover (crop), contain (letterbox), or stretch. Batch-safe."
)
@classmethod
def INPUT_TYPES(cls) -> dict[str, Any]:
return {
"required": {
"image": ("IMAGE", {"tooltip": "Image (or batch) to resize"}),
"preset": (
[*_PRESET_SIZES.keys(), _CUSTOM_PRESET],
{
"default": "square_hd",
"tooltip": (
"fal image_size preset: square_hd=1024x1024, square=512x512, "
"portrait_4_3=768x1024, portrait_16_9=576x1024, "
"landscape_4_3=1024x768, landscape_16_9=1024x576. "
"'custom' uses the width/height inputs."
),
},
),
"width": (
"INT",
{
"default": 1024,
"min": 8,
"max": 14142,
"step": 8,
"tooltip": "Target width in pixels (used when preset is 'custom')",
},
),
"height": (
"INT",
{
"default": 1024,
"min": 8,
"max": 14142,
"step": 8,
"tooltip": "Target height in pixels (used when preset is 'custom')",
},
),
"mode": (
["cover_crop", "contain_pad", "stretch"],
{
"default": "cover_crop",
"tooltip": (
"cover_crop: fill the frame and center-crop the overflow; "
"contain_pad: fit inside and letterbox with black bars; "
"stretch: ignore aspect ratio"
),
},
),
},
}
def resize(
self,
image: Any,
preset: str = "square_hd",
width: int = 1024,
height: int = 1024,
mode: str = "cover_crop",
) -> tuple[torch.Tensor, int, int]:
if preset == _CUSTOM_PRESET:
target_w, target_h = int(width), int(height)
elif preset in _PRESET_SIZES:
target_w, target_h = _PRESET_SIZES[preset]
else:
raise FalApiError("FalResizeToPreset", f"Unknown preset: {preset!r}")
if target_w < 1 or target_h < 1:
raise FalApiError("FalResizeToPreset", f"Invalid target size: {target_w}x{target_h}")
pils = _image_input_to_pils(image)
try:
resized = [_resize_one(pil, target_w, target_h, mode) for pil in pils]
except Exception as exc:
logger.error("FalResizeToPreset: resize failed: %s", exc)
raise FalApiError("FalResizeToPreset", f"Failed to resize image: {exc}") from exc
return (_pils_to_tensor(resized), target_w, target_h)
class FalImageToBase64:
"""Encode an image as a base64 string (optionally a data: URI)."""
RETURN_TYPES = ("STRING",)
FUNCTION = "encode"
CATEGORY = _CATEGORY
DESCRIPTION = (
"Encode the first image of a batch as base64 text, optionally "
"wrapped in a data: URI — useful for APIs that accept inline "
"base64 images instead of URLs."
)
@classmethod
def INPUT_TYPES(cls) -> dict[str, Any]:
return {
"required": {
"image": ("IMAGE", {"tooltip": "Image to encode (first of the batch is used)"}),
"format": (
["png", "jpeg", "webp"],
{"default": "png", "tooltip": "Encoding format; png and webp are lossless-capable"},
),
"data_uri": (
"BOOLEAN",
{
"default": True,
"tooltip": "Prefix with 'data:image/...;base64,' (most APIs expect this)",
},
),
},
}
def encode(self, image: Any, format: str = "png", data_uri: bool = True) -> tuple[str]:
fmt = (format or "png").lower()
if fmt not in _MIME_BY_FORMAT:
raise FalApiError("FalImageToBase64", f"Unsupported format: {format!r}")
pil = ImageUtils.tensor_to_pil(image).convert("RGB")
try:
buffer = io.BytesIO()
save_kwargs = {"lossless": True} if fmt == "webp" else {}
pil.save(buffer, format=fmt.upper(), **save_kwargs)
encoded = base64.b64encode(buffer.getvalue()).decode("ascii")
except Exception as exc:
logger.error("FalImageToBase64: encoding failed: %s", exc)
raise FalApiError("FalImageToBase64", f"Failed to encode image as {fmt}: {exc}") from exc
if data_uri:
return (f"data:{_MIME_BY_FORMAT[fmt]};base64,{encoded}",)
return (encoded,)
class FalBase64ToImage:
"""Decode a base64 string (raw or data: URI) into an IMAGE tensor."""
RETURN_TYPES = ("IMAGE",)
FUNCTION = "decode"
CATEGORY = _CATEGORY
DESCRIPTION = (
"Decode base64 image data — either a raw base64 string or a full "
"'data:image/...;base64,...' URI — into a ComfyUI IMAGE."
)
@classmethod
def INPUT_TYPES(cls) -> dict[str, Any]:
return {
"required": {
"data": (
"STRING",
{
"default": "",
"multiline": True,
"tooltip": "Raw base64 image data, or a data:image/...;base64,... URI",
},
),
},
}
def decode(self, data: str) -> tuple[torch.Tensor]:
payload = (data or "").strip()
if payload.startswith("data:"):
_, _, payload = payload.partition(",")
payload = _B64_WHITESPACE.sub("", payload)
if not payload:
raise FalApiError("FalBase64ToImage", "No base64 data provided")
try:
raw = base64.b64decode(payload)
pil = Image.open(io.BytesIO(raw)).convert("RGB")
except Exception as exc:
logger.error("FalBase64ToImage: decoding failed: %s", exc)
raise FalApiError("FalBase64ToImage", f"Failed to decode base64 image: {exc}") from exc
return (_pils_to_tensor([pil]),)
NODE_CLASS_MAPPINGS = {
"FalImageGrid_fal": FalImageGrid,
"FalResizeToPreset_fal": FalResizeToPreset,
"FalImageToBase64_fal": FalImageToBase64,
"FalBase64ToImage_fal": FalBase64ToImage,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"FalImageGrid_fal": "Image Grid with Labels (fal)",
"FalResizeToPreset_fal": "Resize to fal Preset (fal)",
"FalImageToBase64_fal": "Image → Base64 (fal)",
"FalBase64ToImage_fal": "Base64 → Image (fal)",
}
+404
View File
@@ -0,0 +1,404 @@
"""Utility loader nodes: bring images/audio/folders into ComfyUI from URLs and disk.
ComfyUI IMAGE convention: float32 tensors in [0, 1] with shape (B, H, W, C).
"""
from __future__ import annotations
import glob
import io
import os
from typing import Any
import numpy as np
import requests
import torch
from PIL import Image
from .fal_utils import FalApiError, MediaUtils, logger
_CATEGORY = "FAL/Utils/Load"
_DOWNLOAD_TIMEOUT = (10, 180)
_DEFAULT_FOLDER_PATTERN = "*.png,*.jpg,*.jpeg,*.webp"
def _split_csv(value: str) -> list[str]:
"""Split a comma-separated string into stripped, non-empty parts."""
return [part.strip() for part in (value or "").split(",") if part.strip()]
def _validate_http_url(node_name: str, url: str) -> str:
"""Validate that a URL is a non-empty http(s) URL and return it stripped."""
stripped = (url or "").strip()
if not stripped:
raise FalApiError(
node_name, "'url' is empty. Provide an http(s) URL to a media file."
)
if not stripped.startswith(("http://", "https://")):
raise FalApiError(
node_name,
f"Invalid URL '{stripped}'. Only http(s) URLs are supported.",
)
return stripped
def _download_pil_image(node_name: str, url: str) -> Image.Image:
"""Download a URL and decode it as an RGB PIL image."""
try:
response = requests.get(url, timeout=_DOWNLOAD_TIMEOUT)
response.raise_for_status()
return Image.open(io.BytesIO(response.content)).convert("RGB")
except FalApiError:
raise
except Exception as exc:
logger.error("%s: failed to download image %s: %s", node_name, url, exc)
raise FalApiError(
node_name,
f"Failed to download or decode image from '{url}': {exc}",
) from exc
def _images_to_batch_tensor(
node_name: str, images: list[Image.Image], labels: list[str]
) -> torch.Tensor:
"""Stack RGB PIL images into a float32 (B, H, W, C) tensor in [0, 1].
Images whose size differs from the first image are resized to match
(with a warning) so the batch stays valid.
"""
first_size = images[0].size # (W, H)
arrays: list[np.ndarray] = []
for img, label in zip(images, labels):
if img.size != first_size:
logger.warning(
"%s: '%s' is %sx%s; resizing to %sx%s to match the first image",
node_name,
label,
img.size[0],
img.size[1],
first_size[0],
first_size[1],
)
img = img.resize(first_size, Image.LANCZOS)
arrays.append(np.array(img).astype(np.float32) / 255.0)
return torch.from_numpy(np.stack(arrays, axis=0))
class FalLoadImageURL:
"""Load one or more images from http(s) URLs into an IMAGE batch."""
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("image",)
FUNCTION = "load"
CATEGORY = _CATEGORY
DESCRIPTION = (
"Download an image from an http(s) URL into a ComfyUI IMAGE tensor. "
"Accepts a comma-separated list of URLs to build a batch; images with "
"differing sizes are resized to match the first."
)
@classmethod
def INPUT_TYPES(cls) -> dict[str, Any]:
return {
"required": {
"url": (
"STRING",
{
"default": "",
"multiline": True,
"tooltip": (
"http(s) URL of the image to load. A comma-separated "
"list of URLs produces a batched IMAGE; mismatched "
"sizes are resized to the first image's size. URLs "
"containing literal commas are not supported in list mode."
),
},
),
},
}
def load(self, url: str) -> tuple[torch.Tensor]:
node_name = "FalLoadImageURL"
urls = [_validate_http_url(node_name, part) for part in _split_csv(url)]
if not urls:
raise FalApiError(
node_name,
"'url' is empty. Provide an http(s) URL (or a comma-separated "
"list of URLs) to image file(s).",
)
images = [_download_pil_image(node_name, u) for u in urls]
return (_images_to_batch_tensor(node_name, images, urls),)
class FalLoadAudioURL:
"""Load audio from an http(s) URL into a ComfyUI AUDIO output."""
RETURN_TYPES = ("AUDIO",)
RETURN_NAMES = ("audio",)
FUNCTION = "load"
CATEGORY = _CATEGORY
DESCRIPTION = (
"Download and decode an audio file from an http(s) URL into a native "
"ComfyUI AUDIO output (waveform + sample rate)."
)
@classmethod
def INPUT_TYPES(cls) -> dict[str, Any]:
return {
"required": {
"url": (
"STRING",
{
"default": "",
"tooltip": (
"http(s) URL of the audio file to download and "
"decode (e.g. the audio_url output of a fal node)."
),
},
),
},
}
def load(self, url: str) -> tuple[dict[str, Any]]:
node_name = "FalLoadAudioURL"
validated = _validate_http_url(node_name, url)
try:
return (MediaUtils.audio_from_url(validated),)
except FalApiError:
raise
except Exception as exc:
logger.error("%s: failed to load audio %s: %s", node_name, validated, exc)
raise FalApiError(
node_name,
f"Failed to load audio from '{validated}': {exc}",
) from exc
def _resolve_folder(node_name: str, folder_path: str) -> str:
"""Expand and validate a folder path, returning its absolute form."""
expanded = os.path.expanduser((folder_path or "").strip())
if not expanded:
raise FalApiError(
node_name, "'folder_path' is empty. Provide a path to a folder."
)
if not os.path.isdir(expanded):
raise FalApiError(
node_name,
f"Folder not found: '{expanded}'. Provide an existing folder path.",
)
return os.path.abspath(expanded)
def _glob_folder_files(folder: str, patterns: list[str]) -> list[str]:
"""Glob a folder with each pattern, deduplicated, unordered."""
matched: set[str] = set()
for pattern in patterns:
for path in glob.glob(os.path.join(folder, pattern)):
if os.path.isfile(path):
matched.add(os.path.abspath(path))
return list(matched)
def _sort_files(files: list[str], sort: str) -> list[str]:
"""Sort file paths deterministically by name or modification time."""
if sort == "modified":
return sorted(files, key=lambda path: (os.path.getmtime(path), path))
return sorted(files)
class FalLoadImageFolder:
"""Load a folder of images from disk into a single IMAGE batch."""
RETURN_TYPES = ("IMAGE", "INT")
RETURN_NAMES = ("images", "count")
FUNCTION = "load"
CATEGORY = _CATEGORY
DESCRIPTION = (
"Load every image matching the pattern(s) in a local folder into one "
"IMAGE batch. Mixed sizes are resized to the first image's dimensions."
)
@classmethod
def INPUT_TYPES(cls) -> dict[str, Any]:
return {
"required": {
"folder_path": (
"STRING",
{
"default": "",
"tooltip": (
"Path to a local folder of images. '~' expands to "
"your home directory."
),
},
),
"pattern": (
"STRING",
{
"default": _DEFAULT_FOLDER_PATTERN,
"tooltip": (
"Comma-separated glob pattern(s) selecting which "
"files to load, e.g. '*.png,*.jpg'."
),
},
),
"max_images": (
"INT",
{
"default": 100,
"min": 1,
"max": 1000,
"step": 1,
"tooltip": "Maximum number of images to load from the folder.",
},
),
"sort": (
["name", "modified"],
{
"default": "name",
"tooltip": (
"Order in which files are loaded: alphabetical by "
"'name' or oldest-first by 'modified' time."
),
},
),
},
}
def load(
self, folder_path: str, pattern: str, max_images: int, sort: str
) -> tuple[torch.Tensor, int]:
node_name = "FalLoadImageFolder"
folder = _resolve_folder(node_name, folder_path)
patterns = _split_csv(pattern) or _split_csv(_DEFAULT_FOLDER_PATTERN)
files = _sort_files(_glob_folder_files(folder, patterns), sort)[:max_images]
if not files:
raise FalApiError(
node_name,
f"No files matching '{', '.join(patterns)}' found in '{folder}'. "
"Adjust 'pattern' or point 'folder_path' at a folder with images.",
)
images = [self._open_image(node_name, path) for path in files]
batch = _images_to_batch_tensor(node_name, images, files)
return (batch, len(files))
@staticmethod
def _open_image(node_name: str, path: str) -> Image.Image:
"""Open a local image file as RGB, normalizing failures."""
try:
with Image.open(path) as img:
return img.convert("RGB")
except Exception as exc:
logger.error("%s: failed to open image %s: %s", node_name, path, exc)
raise FalApiError(
node_name,
f"Failed to open image '{path}': {exc}. Remove or exclude the "
"file via 'pattern' and retry.",
) from exc
def _normalize_extensions(extensions: str) -> list[str] | None:
"""Parse a comma-separated extension filter; empty means no filter."""
parts = [part.lstrip("*").lower() for part in _split_csv(extensions)]
normalized = [part if part.startswith(".") else f".{part}" for part in parts]
cleaned = [part for part in normalized if part != "."]
return cleaned or None
class FalUploadFolderAsZip:
"""Zip a local folder and upload the archive to fal.ai, returning its URL."""
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("zip_url",)
FUNCTION = "upload"
CATEGORY = _CATEGORY
DESCRIPTION = (
"Zip a local folder (optionally recursive / filtered by extension) and "
"upload the archive to fal.ai storage, returning the ZIP's URL — handy "
"for endpoints that take a training-data archive."
)
@classmethod
def INPUT_TYPES(cls) -> dict[str, Any]:
return {
"required": {
"folder_path": (
"STRING",
{
"default": "",
"tooltip": (
"Path to the local folder to zip and upload. '~' "
"expands to your home directory."
),
},
),
"recursive": (
"BOOLEAN",
{
"default": False,
"tooltip": "Include files from subfolders in the ZIP.",
},
),
"extensions": (
"STRING",
{
"default": "",
"tooltip": (
"Comma-separated file extensions to include, e.g. "
"'.png,.jpg'. Leave empty to include all files."
),
},
),
},
}
def upload(
self, folder_path: str, recursive: bool, extensions: str
) -> tuple[str]:
node_name = "FalUploadFolderAsZip"
folder = _resolve_folder(node_name, folder_path)
archive_utils = self._load_archive_utils(node_name)
include_extensions = _normalize_extensions(extensions)
try:
zip_path = archive_utils.zip_folder(
folder, include_extensions=include_extensions, recursive=recursive
)
return (archive_utils.upload_zip(zip_path),)
except FalApiError:
raise
except Exception as exc:
logger.error("%s: failed to zip/upload %s: %s", node_name, folder, exc)
raise FalApiError(
node_name,
f"Failed to zip and upload folder '{folder}': {exc}",
) from exc
@staticmethod
def _load_archive_utils(node_name: str) -> Any:
"""Lazily import ArchiveUtils, degrading with an actionable error."""
try:
from .fal_utils import ArchiveUtils
return ArchiveUtils
except ImportError as exc:
raise FalApiError(
node_name,
"Archive utilities are unavailable in this install "
f"({exc}). Update/reinstall ComfyUI-fal-API so that "
"nodes/utils/archive.py is present.",
) from exc
NODE_CLASS_MAPPINGS = {
"FalLoadImageURL_fal": FalLoadImageURL,
"FalLoadAudioURL_fal": FalLoadAudioURL,
"FalLoadImageFolder_fal": FalLoadImageFolder,
"FalUploadFolderAsZip_fal": FalUploadFolderAsZip,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"FalLoadImageURL_fal": "Load Image from URL (fal)",
"FalLoadAudioURL_fal": "Load Audio from URL (fal)",
"FalLoadImageFolder_fal": "Load Image Folder (fal)",
"FalUploadFolderAsZip_fal": "Upload Folder as ZIP URL (fal)",
}
+800
View File
@@ -0,0 +1,800 @@
"""Local video utility nodes: frame extraction, trim, concat, mux, audio extraction.
These nodes run entirely locally (cv2/PyAV) — no fal.ai API calls — and are
meant to glue video-generation workflows together (e.g. grab the last frame of
a clip and feed it into an image-to-video node to extend the video).
"""
from __future__ import annotations
import os
import tempfile
from fractions import Fraction
from typing import Any
import numpy as np
import torch
from .fal_utils import FalApiError, MediaUtils, logger
_CATEGORY = "FAL/Utils/Video"
_CHUNK_SIZE = 1 << 20 # 1 MiB
_AV_TIME_BASE = 1_000_000 # PyAV container.seek() offset units (microseconds)
_AAC_FRAME_SIZE = 1024 # samples per AAC frame
_TIME_EPS = 1e-6
# ---------------------------------------------------------------------------
# Shared helpers
# ---------------------------------------------------------------------------
def _safe_unlink(path: str | None) -> None:
"""Delete a temp file, ignoring errors."""
if path is None:
return
try:
os.unlink(path)
except OSError:
pass
def _import_cv2(node_name: str) -> Any:
"""Import cv2 lazily with a clear error when missing."""
try:
import cv2
return cv2
except ImportError as exc:
raise FalApiError(
node_name, "OpenCV is required for this node — pip install opencv-python"
) from exc
def _import_av(node_name: str) -> Any:
"""Import PyAV lazily with a clear error when missing."""
try:
import av
return av
except ImportError as exc:
raise FalApiError(node_name, "PyAV is required for this node — pip install av") from exc
def _resolve_video_from_file() -> type | None:
"""Locate ComfyUI's VideoFromFile class across API layouts."""
try:
from comfy_api.input_impl import VideoFromFile
return VideoFromFile
except ImportError:
pass
try:
from comfy_api.latest import input_impl
return getattr(input_impl, "VideoFromFile", None)
except ImportError:
return None
def _wrap_local_video(path: str, node_name: str) -> Any:
"""Wrap a local video file as a ComfyUI VIDEO object."""
video_cls = _resolve_video_from_file()
if video_cls is None:
raise FalApiError(
node_name,
"comfy_api VideoFromFile is unavailable; update ComfyUI to a version "
"that provides comfy_api to use VIDEO outputs.",
)
return video_cls(path)
def _new_temp_path(suffix: str) -> str:
"""Create an empty named temp file and return its path."""
with tempfile.NamedTemporaryFile(suffix=suffix, prefix="fal_util_video_", delete=False) as temp_file:
return temp_file.name
def _spool_stream_to_temp(source: Any, node_name: str) -> str:
"""Write a readable stream to a temp .mp4 file and return its path."""
temp_path: str | None = None
try:
with tempfile.NamedTemporaryFile(suffix=".mp4", prefix="fal_util_video_", delete=False) as temp_file:
temp_path = temp_file.name
while True:
chunk = source.read(_CHUNK_SIZE)
if not chunk:
break
temp_file.write(chunk)
return temp_path
except Exception as exc:
_safe_unlink(temp_path)
raise FalApiError(node_name, f"Failed to read video stream: {exc}") from exc
def _video_input_to_path(video: Any, node_name: str) -> tuple[str, bool]:
"""Resolve a VIDEO input (or path/URL string) to a local file path.
Returns (path, cleanup_needed). cleanup_needed is True when the path is a
temp file created here that the caller must delete when done.
"""
if video is None:
raise FalApiError(node_name, "No video input provided")
source = video.get_stream_source() if hasattr(video, "get_stream_source") else video
if isinstance(source, str):
if source.startswith(("http://", "https://")):
return MediaUtils.download_url_to_temp(source, ".mp4"), True
if os.path.isfile(source):
return source, False
raise FalApiError(node_name, f"Video path does not exist: {source}")
if hasattr(source, "read"):
return _spool_stream_to_temp(source, node_name), True
raise FalApiError(
node_name,
f"Unsupported video input of type {type(video).__name__}; expected a "
"VIDEO object, a local file path, or an http(s) URL string.",
)
def _add_stream_from_template(output: Any, template: Any) -> Any:
"""Add an output stream copying the template's codec parameters."""
if hasattr(output, "add_stream_from_template"):
return output.add_stream_from_template(template)
return output.add_stream(template=template)
def _stream_duration_seconds(container: Any, stream: Any) -> float:
"""Best-effort duration (seconds) of a stream, falling back to container."""
if stream.duration is not None and stream.time_base is not None:
return float(stream.duration * stream.time_base)
if container.duration is not None:
return float(container.duration) / _AV_TIME_BASE
return 0.0
def _encode_audio_array(
av: Any, output: Any, stream: Any, layout: str, sample_rate: int, samples: np.ndarray, start_index: int
) -> int:
"""Encode a planar float32 (C, T) array as AAC frames; returns next sample index."""
total = samples.shape[1]
for offset in range(0, total, _AAC_FRAME_SIZE):
chunk = np.ascontiguousarray(samples[:, offset : offset + _AAC_FRAME_SIZE])
frame = av.AudioFrame.from_ndarray(chunk, format="fltp", layout=layout)
frame.sample_rate = sample_rate
frame.pts = start_index + offset
for packet in stream.encode(frame):
output.mux(packet)
return start_index + total
def _pad_or_truncate(samples: np.ndarray, needed: int) -> np.ndarray:
"""Pad a planar (C, T) array with silence, or truncate, to exactly `needed` samples."""
if samples.shape[1] >= needed:
return samples[:, :needed]
pad = np.zeros((samples.shape[0], needed - samples.shape[1]), dtype=np.float32)
return np.concatenate([samples, pad], axis=1)
def _normalize_audio_frame(array: np.ndarray, channels: int) -> np.ndarray:
"""Normalize a PyAV audio frame array to float32 with shape (C, N)."""
if np.issubdtype(array.dtype, np.integer):
info = np.iinfo(array.dtype)
scale = float(max(abs(info.min), info.max))
array = array.astype(np.float32) / scale
else:
array = array.astype(np.float32)
if array.ndim == 1:
array = array[np.newaxis, :]
if array.shape[0] == 1 and channels > 1:
# Packed/interleaved format: (1, N * C) -> (C, N)
array = array.reshape(-1, channels).T
return array
def _waveform_to_planar(audio: Any, node_name: str) -> tuple[np.ndarray, int]:
"""Convert a ComfyUI AUDIO dict to (planar float32 (C, T) with C in {1, 2}, sample_rate)."""
try:
waveform = audio["waveform"]
sample_rate = int(audio["sample_rate"])
except (KeyError, TypeError) as exc:
raise FalApiError(
node_name, "Expected an AUDIO dict with 'waveform' and 'sample_rate'"
) from exc
if not isinstance(waveform, torch.Tensor):
raise FalApiError(node_name, "AUDIO 'waveform' must be a torch tensor")
tensor = waveform.detach().cpu().to(torch.float32)
if tensor.ndim == 3:
tensor = tensor[0] # (B, C, T) -> (C, T)
if tensor.ndim == 1:
tensor = tensor.unsqueeze(0)
if tensor.ndim != 2:
raise FalApiError(node_name, f"AUDIO waveform has unsupported shape {tuple(waveform.shape)}")
array = tensor.clamp(-1.0, 1.0).numpy()
if array.shape[0] > 2:
logger.warning("%s: waveform has %d channels; keeping the first two", node_name, array.shape[0])
array = array[:2]
return np.ascontiguousarray(array), sample_rate
# ---------------------------------------------------------------------------
# FalExtractFrames
# ---------------------------------------------------------------------------
def _bgr_frames_to_tensor(frames: list[np.ndarray], cv2: Any) -> torch.Tensor:
"""Convert BGR uint8 frames to a (N, H, W, 3) float32 RGB tensor in 0-1."""
rgb = [cv2.cvtColor(frame, cv2.COLOR_BGR2RGB) for frame in frames]
stacked = np.stack(rgb).astype(np.float32) / 255.0
return torch.from_numpy(stacked)
def _frame_via_seek(cap: Any, cv2: Any, index: int) -> np.ndarray | None:
"""Seek to a frame index and read it; returns None when the seek misbehaves."""
cap.set(cv2.CAP_PROP_POS_FRAMES, float(index))
ok, frame = cap.read()
return frame if ok and frame is not None else None
def _scan_frames(cap: Any, cv2: Any, stop_after: int | None = None) -> tuple[np.ndarray | None, int]:
"""Sequentially decode from frame 0; returns (last frame seen, frames read).
Stops after reading `stop_after + 1` frames when `stop_after` is given.
"""
cap.set(cv2.CAP_PROP_POS_FRAMES, 0.0)
last: np.ndarray | None = None
count = 0
while True:
ok, frame = cap.read()
if not ok or frame is None:
break
last = frame
count += 1
if stop_after is not None and count > stop_after:
break
return last, count
class FalExtractFrames:
"""Extract frames from a video as IMAGE outputs (local decode, no API call)."""
RETURN_TYPES = ("IMAGE", "INT")
RETURN_NAMES = ("frames", "frame_count")
FUNCTION = "extract_frames"
CATEGORY = _CATEGORY
DESCRIPTION = (
"Decode a video locally and extract frames. Mode 'last' grabs the final "
"frame — feed it into an image-to-video node to extend/continue a video."
)
@classmethod
def INPUT_TYPES(cls) -> dict[str, Any]:
return {
"required": {
"video": ("VIDEO", {"tooltip": "Video to decode. Tip: use mode 'last' to grab the final frame and feed it into an image-to-video node to extend the video."}),
"mode": (["first", "last", "nth", "every_nth"], {"default": "last", "tooltip": "first/last: single frame. nth: the n-th frame (1-based). every_nth: every n-th frame as a batch, capped at max_frames."}),
"n": ("INT", {"default": 1, "min": 1, "max": 1_000_000, "tooltip": "Frame index (1-based) for mode 'nth'; step size for 'every_nth'. Ignored otherwise."}),
"max_frames": ("INT", {"default": 64, "min": 1, "max": 1024, "tooltip": "Maximum number of frames returned by mode 'every_nth'; ignored for other modes."}),
},
}
def extract_frames(self, video: Any, mode: str, n: int, max_frames: int) -> tuple[torch.Tensor, int]:
node_name = "FalExtractFrames"
cv2 = _import_cv2(node_name)
path, cleanup = _video_input_to_path(video, node_name)
cap = None
try:
cap = cv2.VideoCapture(path)
if not cap.isOpened():
raise FalApiError(node_name, f"OpenCV could not open video: {path}")
reported = int(cap.get(cv2.CAP_PROP_FRAME_COUNT) or 0)
if mode == "every_nth":
frames, frame_count = self._extract_every_nth(cap, cv2, n, max_frames, reported)
else:
frame, frame_count = self._extract_single(cap, cv2, mode, n, reported, node_name)
frames = [frame]
if not frames or frames[0] is None:
raise FalApiError(node_name, f"No frames could be decoded from {path}")
return (_bgr_frames_to_tensor(frames, cv2), frame_count)
except FalApiError:
raise
except Exception as exc:
logger.error("FalExtractFrames failed: %s", exc)
raise FalApiError(node_name, f"Frame extraction failed: {exc}") from exc
finally:
if cap is not None:
cap.release()
if cleanup:
_safe_unlink(path)
@staticmethod
def _extract_every_nth(
cap: Any, cv2: Any, step: int, max_frames: int, reported: int
) -> tuple[list[np.ndarray], int]:
"""Sequentially collect every `step`-th frame, capped at max_frames."""
cap.set(cv2.CAP_PROP_POS_FRAMES, 0.0)
frames: list[np.ndarray] = []
count = 0
reached_eof = False
while True:
ok, frame = cap.read()
if not ok or frame is None:
reached_eof = True
break
if count % step == 0 and len(frames) < max_frames:
frames.append(frame)
count += 1
if len(frames) >= max_frames and reported > 0:
break # early stop: reported count stands in for the true total
frame_count = count if reached_eof or reported <= 0 else reported
return frames, frame_count
@staticmethod
def _extract_single(
cap: Any, cv2: Any, mode: str, n: int, reported: int, node_name: str
) -> tuple[np.ndarray | None, int]:
"""Extract a single frame for modes first/last/nth."""
if mode == "first":
target = 0
elif mode == "nth":
target = n - 1
if reported > 0 and target >= reported:
logger.warning("%s: frame %d beyond end (%d frames); using last frame", node_name, n, reported)
target = reported - 1
elif mode == "last":
target = max(reported - 1, 0)
else:
raise FalApiError(node_name, f"Unknown mode: {mode}")
frame: np.ndarray | None = None
if reported > 0:
# Fast path: direct seek (some codecs mis-seek; fall back below).
frame = _frame_via_seek(cap, cv2, target)
if frame is not None:
return frame, reported
# Sequential fallback: decode from the start.
if mode == "last":
frame, count = _scan_frames(cap, cv2)
return frame, count
frame, read = _scan_frames(cap, cv2, stop_after=target)
if read > target:
# Reached the target; total count comes from metadata or a full scan.
frame_count = reported if reported > 0 else _scan_frames(cap, cv2)[1]
return frame, frame_count
# Hit EOF early: `frame` is the last decodable frame, `read` the true count.
return frame, read
# ---------------------------------------------------------------------------
# FalTrimVideo
# ---------------------------------------------------------------------------
def _find_seek_start(av: Any, container: Any, anchor: Any, start: float) -> float:
"""Seek near `start` and return the timestamp of the first packet (keyframe snap)."""
if start <= 0:
return 0.0
container.seek(int(start * _AV_TIME_BASE), backward=True, any_frame=False)
actual_start = start
for packet in container.demux(anchor):
if packet.pts is None:
continue
actual_start = float(packet.pts * packet.time_base)
break
container.seek(int(start * _AV_TIME_BASE), backward=True, any_frame=False)
return actual_start
def _remux_trim(av: Any, in_path: str, out_path: str, start: float, end: float | None, node_name: str) -> None:
"""Copy packets between timestamps into a new mp4 without re-encoding."""
with av.open(in_path) as container, av.open(out_path, mode="w") as output:
video_in = container.streams.video[0] if container.streams.video else None
audio_in = container.streams.audio[0] if container.streams.audio else None
selected = [stream for stream in (video_in, audio_in) if stream is not None]
if not selected:
raise FalApiError(node_name, "Input has no video or audio streams")
out_streams = {stream.index: _add_stream_from_template(output, stream) for stream in selected}
anchor = video_in if video_in is not None else audio_in
actual_start = _find_seek_start(av, container, anchor, start)
if end is not None and end <= actual_start + _TIME_EPS:
raise FalApiError(
node_name,
f"Trim range is empty: start snapped to keyframe at {actual_start:.3f}s, end is {end:.3f}s",
)
offsets: dict[int, int] = {}
done = dict.fromkeys(out_streams, False)
kept = 0
for packet in container.demux(selected):
if packet.pts is None:
continue
index = packet.stream.index
if done.get(index, True):
continue
time = float(packet.pts * packet.time_base)
if time < actual_start - _TIME_EPS:
continue
if end is not None and time >= end - _TIME_EPS:
done[index] = True
if all(done.values()):
break
continue
if index not in offsets:
offsets[index] = packet.dts if packet.dts is not None else packet.pts
offset = offsets[index]
packet.pts -= offset
if packet.dts is not None:
packet.dts -= offset
packet.stream = out_streams[index]
output.mux(packet)
kept += 1
if kept == 0:
raise FalApiError(node_name, f"Trim produced no packets (start {start:.3f}s may be past the end)")
class FalTrimVideo:
"""Trim a video to [start, end] seconds by remuxing (no re-encode)."""
RETURN_TYPES = ("VIDEO", "STRING")
RETURN_NAMES = ("video", "path")
FUNCTION = "trim_video"
CATEGORY = _CATEGORY
DESCRIPTION = (
"Trim a video without re-encoding by copying packets between timestamps. "
"Fast and lossless, but the start cut snaps to the nearest earlier keyframe."
)
@classmethod
def INPUT_TYPES(cls) -> dict[str, Any]:
return {
"required": {
"video": ("VIDEO", {"tooltip": "Video to trim (video + audio tracks are kept)."}),
"start_seconds": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 100_000.0, "step": 0.1, "tooltip": "Trim start in seconds. Cuts snap to the nearest earlier keyframe (no re-encode)."}),
"end_seconds": ("FLOAT", {"default": 5.0, "min": 0.0, "max": 100_000.0, "step": 0.1, "tooltip": "Trim end in seconds; 0 = to the end of the video."}),
},
}
def trim_video(self, video: Any, start_seconds: float, end_seconds: float) -> tuple[Any, str]:
node_name = "FalTrimVideo"
av = _import_av(node_name)
end = end_seconds if end_seconds > 0 else None
if end is not None and end <= start_seconds:
raise FalApiError(node_name, f"end_seconds ({end:.3f}) must be greater than start_seconds ({start_seconds:.3f})")
path, cleanup = _video_input_to_path(video, node_name)
out_path = _new_temp_path(".mp4")
try:
_remux_trim(av, path, out_path, start_seconds, end, node_name)
# NOTE: out_path is deliberately kept — VideoFromFile reads it lazily.
return (_wrap_local_video(out_path, node_name), out_path)
except FalApiError:
_safe_unlink(out_path)
raise
except Exception as exc:
_safe_unlink(out_path)
logger.error("FalTrimVideo failed: %s", exc)
raise FalApiError(node_name, f"Trim failed: {exc}") from exc
finally:
if cleanup:
_safe_unlink(path)
# ---------------------------------------------------------------------------
# FalConcatVideos
# ---------------------------------------------------------------------------
def _probe_video(av: Any, path: str, node_name: str) -> dict[str, Any]:
"""Probe a clip for resolution, fps, audio presence, and audio rate."""
with av.open(path) as container:
if not container.streams.video:
raise FalApiError(node_name, f"Input has no video stream: {path}")
stream = container.streams.video[0]
fps = stream.average_rate or stream.guessed_rate
if not fps or fps <= 0:
fps = Fraction(30, 1)
audio_rate = None
if container.streams.audio:
audio_rate = int(container.streams.audio[0].rate or 44100)
return {
"width": max(2, stream.width - stream.width % 2),
"height": max(2, stream.height - stream.height % 2),
"fps": Fraction(fps),
"audio_rate": audio_rate,
}
def _encode_video_frame(output: Any, stream: Any, frame: Any, width: int, height: int, time_base: Fraction, index: int) -> None:
"""Scale/convert one decoded frame and encode it at the given frame index."""
scaled = frame.reformat(width=width, height=height, format="yuv420p")
scaled.pts = index
scaled.time_base = time_base
for packet in stream.encode(scaled):
output.mux(packet)
def _append_clip_video(
av: Any, output: Any, stream: Any, path: str, width: int, height: int, fps: Fraction, start_index: int, node_name: str
) -> int:
"""Decode a clip, resample to target fps/size, encode; returns frames emitted."""
time_base = Fraction(1, 1) / fps
step = 1.0 / float(fps)
emitted = 0
with av.open(path) as container:
next_time = 0.0
last = None
for frame in container.decode(container.streams.video[0]):
time = frame.time if frame.time is not None else next_time
while last is not None and time > next_time + _TIME_EPS:
_encode_video_frame(output, stream, last, width, height, time_base, start_index + emitted)
emitted += 1
next_time += step
last = frame
if last is not None:
_encode_video_frame(output, stream, last, width, height, time_base, start_index + emitted)
emitted += 1
if emitted == 0:
raise FalApiError(node_name, f"No video frames decoded from {path}")
return emitted
def _clip_audio_samples(av: Any, path: str, rate: int, needed: int) -> np.ndarray:
"""Decode+resample a clip's audio to stereo float32 (2, needed); silence when absent."""
with av.open(path) as container:
if not container.streams.audio:
return np.zeros((2, needed), dtype=np.float32)
resampler = av.AudioResampler(format="fltp", layout="stereo", rate=rate)
chunks: list[np.ndarray] = []
for frame in container.decode(container.streams.audio[0]):
frame.pts = None # let the resampler track timestamps itself
chunks.extend(out.to_ndarray() for out in resampler.resample(frame))
chunks.extend(out.to_ndarray() for out in resampler.resample(None))
if not chunks:
return np.zeros((2, needed), dtype=np.float32)
samples = np.concatenate(chunks, axis=1).astype(np.float32)
return _pad_or_truncate(samples, needed)
class FalConcatVideos:
"""Concatenate 2-4 videos by re-encoding to the first clip's resolution and fps."""
RETURN_TYPES = ("VIDEO", "STRING")
RETURN_NAMES = ("video", "path")
FUNCTION = "concat_videos"
CATEGORY = _CATEGORY
DESCRIPTION = (
"Concatenate up to 4 videos. All clips are re-encoded (h264 crf 18) and scaled to "
"video_1's resolution and fps, so mismatched codecs/sizes are fine. Audio: the output "
"gets a stereo AAC track when any input has audio; inputs without audio contribute silence."
)
@classmethod
def INPUT_TYPES(cls) -> dict[str, Any]:
return {
"required": {
"video_1": ("VIDEO", {"tooltip": "First clip; its resolution and fps define the output format."}),
"video_2": ("VIDEO", {"tooltip": "Second clip, scaled to match video_1."}),
},
"optional": {
"video_3": ("VIDEO", {"tooltip": "Optional third clip."}),
"video_4": ("VIDEO", {"tooltip": "Optional fourth clip."}),
},
}
def concat_videos(
self, video_1: Any, video_2: Any, video_3: Any = None, video_4: Any = None
) -> tuple[Any, str]:
node_name = "FalConcatVideos"
av = _import_av(node_name)
inputs = [video for video in (video_1, video_2, video_3, video_4) if video is not None]
resolved: list[tuple[str, bool]] = []
out_path = _new_temp_path(".mp4")
try:
resolved = [_video_input_to_path(video, node_name) for video in inputs]
paths = [path for path, _ in resolved]
probes = [_probe_video(av, path, node_name) for path in paths]
target = probes[0]
width, height, fps = target["width"], target["height"], target["fps"]
audio_rates = [probe["audio_rate"] for probe in probes if probe["audio_rate"]]
audio_rate = audio_rates[0] if audio_rates else None
with av.open(out_path, mode="w") as output:
video_out = output.add_stream("libx264", rate=fps, options={"crf": "18", "preset": "veryfast"})
video_out.width = width
video_out.height = height
video_out.pix_fmt = "yuv420p"
audio_out = None
if audio_rate is not None:
audio_out = output.add_stream("aac", rate=audio_rate)
audio_out.layout = "stereo"
frame_index = 0
sample_index = 0
for path in paths:
emitted = _append_clip_video(av, output, video_out, path, width, height, fps, frame_index, node_name)
frame_index += emitted
if audio_out is not None:
needed = round(emitted / float(fps) * audio_rate)
samples = _clip_audio_samples(av, path, audio_rate, needed)
sample_index = _encode_audio_array(av, output, audio_out, "stereo", audio_rate, samples, sample_index)
for packet in video_out.encode(None):
output.mux(packet)
if audio_out is not None:
for packet in audio_out.encode(None):
output.mux(packet)
# NOTE: out_path is deliberately kept — VideoFromFile reads it lazily.
return (_wrap_local_video(out_path, node_name), out_path)
except FalApiError:
_safe_unlink(out_path)
raise
except Exception as exc:
_safe_unlink(out_path)
logger.error("FalConcatVideos failed: %s", exc)
raise FalApiError(node_name, f"Concat failed: {exc}") from exc
finally:
for path, cleanup in resolved:
if cleanup:
_safe_unlink(path)
# ---------------------------------------------------------------------------
# FalMuxAudioVideo
# ---------------------------------------------------------------------------
class FalMuxAudioVideo:
"""Mux an AUDIO waveform onto a video, replacing any existing audio track."""
RETURN_TYPES = ("VIDEO", "STRING")
RETURN_NAMES = ("video", "path")
FUNCTION = "mux_audio_video"
CATEGORY = _CATEGORY
DESCRIPTION = (
"Attach an AUDIO input to a video as an AAC track, replacing any existing audio. "
"Video packets are copied without re-encoding."
)
@classmethod
def INPUT_TYPES(cls) -> dict[str, Any]:
return {
"required": {
"video": ("VIDEO", {"tooltip": "Video track; copied without re-encoding. Existing audio is replaced."}),
"audio": ("AUDIO", {"tooltip": "Audio to attach, encoded as AAC at its own sample rate."}),
"duration_policy": (["video", "shortest"], {"default": "video", "tooltip": "video: keep full video; audio is padded with silence or truncated to match. shortest: cut the output at whichever track ends first."}),
},
}
def mux_audio_video(self, video: Any, audio: Any, duration_policy: str) -> tuple[Any, str]:
node_name = "FalMuxAudioVideo"
av = _import_av(node_name)
samples, sample_rate = _waveform_to_planar(audio, node_name)
layout = "mono" if samples.shape[0] == 1 else "stereo"
audio_duration = samples.shape[1] / float(sample_rate)
path, cleanup = _video_input_to_path(video, node_name)
out_path = _new_temp_path(".mp4")
try:
with av.open(path) as container:
if not container.streams.video:
raise FalApiError(node_name, "Input has no video stream")
video_in = container.streams.video[0]
video_duration = _stream_duration_seconds(container, video_in)
if video_duration <= 0:
raise FalApiError(node_name, "Could not determine the video duration")
target = video_duration if duration_policy == "video" else min(video_duration, audio_duration)
with av.open(out_path, mode="w") as output:
video_out = _add_stream_from_template(output, video_in)
audio_out = output.add_stream("aac", rate=sample_rate)
audio_out.layout = layout
for packet in container.demux(video_in):
if packet.dts is None:
continue
if duration_policy == "shortest" and float(packet.dts * packet.time_base) >= target - _TIME_EPS:
break
packet.stream = video_out
output.mux(packet)
needed = round(target * sample_rate)
_encode_audio_array(
av, output, audio_out, layout, sample_rate, _pad_or_truncate(samples, needed), 0
)
for packet in audio_out.encode(None):
output.mux(packet)
# NOTE: out_path is deliberately kept — VideoFromFile reads it lazily.
return (_wrap_local_video(out_path, node_name), out_path)
except FalApiError:
_safe_unlink(out_path)
raise
except Exception as exc:
_safe_unlink(out_path)
logger.error("FalMuxAudioVideo failed: %s", exc)
raise FalApiError(node_name, f"Mux failed: {exc}") from exc
finally:
if cleanup:
_safe_unlink(path)
# ---------------------------------------------------------------------------
# FalVideoToAudio
# ---------------------------------------------------------------------------
class FalVideoToAudio:
"""Extract the audio track of a video as a ComfyUI AUDIO output."""
RETURN_TYPES = ("AUDIO",)
RETURN_NAMES = ("audio",)
FUNCTION = "video_to_audio"
CATEGORY = _CATEGORY
DESCRIPTION = "Extract a video's audio track as an AUDIO output at its original sample rate."
@classmethod
def INPUT_TYPES(cls) -> dict[str, Any]:
return {
"required": {
"video": ("VIDEO", {"tooltip": "Video whose audio track should be extracted."}),
},
}
def video_to_audio(self, video: Any) -> tuple[dict[str, Any]]:
node_name = "FalVideoToAudio"
av = _import_av(node_name)
path, cleanup = _video_input_to_path(video, node_name)
try:
with av.open(path) as container:
if not container.streams.audio:
raise FalApiError(node_name, "video has no audio track")
stream = container.streams.audio[0]
sample_rate = int(stream.rate or 44100)
channels = int(getattr(stream, "channels", 1) or 1)
chunks = [
_normalize_audio_frame(frame.to_ndarray(), channels)
for frame in container.decode(stream)
]
if not chunks:
raise FalApiError(node_name, "No audio frames could be decoded")
waveform = torch.from_numpy(np.concatenate(chunks, axis=1)).unsqueeze(0)
return ({"waveform": waveform, "sample_rate": sample_rate},)
except FalApiError:
raise
except Exception as exc:
logger.error("FalVideoToAudio failed: %s", exc)
raise FalApiError(node_name, f"Audio extraction failed: {exc}") from exc
finally:
if cleanup:
_safe_unlink(path)
NODE_CLASS_MAPPINGS = {
"FalExtractFrames_fal": FalExtractFrames,
"FalTrimVideo_fal": FalTrimVideo,
"FalConcatVideos_fal": FalConcatVideos,
"FalMuxAudioVideo_fal": FalMuxAudioVideo,
"FalVideoToAudio_fal": FalVideoToAudio,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"FalExtractFrames_fal": "Extract Frames (fal)",
"FalTrimVideo_fal": "Trim Video (fal)",
"FalConcatVideos_fal": "Concat Videos (fal)",
"FalMuxAudioVideo_fal": "Mux Audio + Video (fal)",
"FalVideoToAudio_fal": "Video → Audio (fal)",
}
+2
View File
@@ -1,6 +1,7 @@
"""Core utilities for the ComfyUI-fal-API node pack."""
from .api import ApiHandler
from .archive import ArchiveUtils
from .billing import BillingUtils, SpendGuard
from .config import FalConfig
from .errors import FalApiError, extract_error_message, raise_fal_error
@@ -14,6 +15,7 @@ from .result_cache import ResultCache
__all__ = [
"ApiHandler",
"ArchiveUtils",
"BillingUtils",
"FalApiError",
"FalConfig",
+107 -8
View File
@@ -180,6 +180,43 @@ def _store_result_in_cache(
logger.debug("[%s] result cache store failed: %s", endpoint, exc)
def _remember_result_urls(endpoint: str, request_id: str | None, result: Any) -> None:
"""Best-effort provenance bookkeeping: map result URLs to their request."""
if not request_id:
return
try:
ResultCache().remember_urls(endpoint, request_id, result)
except Exception as exc:
logger.debug("[%s] remember_urls failed: %s", endpoint, exc)
def _finalize_live_call(endpoint: str, request_id: str | None, started: float) -> None:
"""Log the finished call and record it in the session ledger."""
duration_s = time.monotonic() - started
logger.info(
"[%s] call finished in %.1fs (request_id=%s)",
endpoint,
duration_s,
request_id,
)
_record_ledger_entry(endpoint, request_id, duration_s)
async def _close_async_client(client: Any) -> None:
"""Best-effort close of a per-call AsyncClient's underlying httpx client.
fal_client.AsyncClient lazily caches an httpx.AsyncClient per instance
(bound to the current event loop); we create one AsyncClient per call, so
close it here to avoid leaking connections. Resolving ``_client`` does no
network I/O; any failure is swallowed — cleanup must never mask a result
or an error from the call itself.
"""
try:
httpx_client = await client._client
await httpx_client.aclose()
except Exception as exc:
logger.debug("async fal client close failed: %s", exc)
def _raise_generation_error(model_name: str, error: Exception | str) -> NoReturn:
"""Normalize an exception or error string into a raised FalApiError."""
if isinstance(error, BaseException):
@@ -243,16 +280,77 @@ class ApiHandler:
raise
raise_fal_error(endpoint, exc)
finally:
duration_s = time.monotonic() - started
logger.info(
"[%s] call finished in %.1fs (request_id=%s)",
endpoint,
duration_s,
request_id_ref[0],
)
_record_ledger_entry(endpoint, request_id_ref[0], duration_s)
_finalize_live_call(endpoint, request_id_ref[0], started)
_store_result_in_cache(endpoint, arguments, result, request_id_ref[0])
_remember_result_urls(endpoint, request_id_ref[0], result)
return result
@staticmethod
async def submit_and_get_result_async(
endpoint: str,
arguments: dict[str, Any],
skip_cache: bool = False,
) -> Any:
"""Async twin of ``submit_and_get_result`` for async-capable ComfyUI.
Same semantics — spend-guard preflight, persistent result cache,
queue-progress logging, interruption via the queue callback, ledger
recording and cache/provenance bookkeeping — but awaits the fal call
on the event loop so the executor can run other graph branches
concurrently. The AsyncClient is created per call because its cached
httpx client is bound to the current event loop (ComfyUI runs each
prompt in a fresh loop via ``asyncio.run``).
"""
# Cache first: a hit costs nothing, so it must not be blocked by the
# spend guard (which only gates live, billable calls).
if not skip_cache:
cached = ResultCache().get(endpoint, arguments)
if cached is not None:
return cached
# off-loop: preflight may make a blocking balance HTTP call
await asyncio.to_thread(_spend_guard_preflight, endpoint)
from fal_client import AsyncClient
# Validate the key via get_client() first so a missing/placeholder key
# raises the actionable config error instead of a raw auth failure.
FalConfig().get_client()
client = AsyncClient(key=FalConfig().get_key())
callback = _make_queue_callback(endpoint)
request_id_ref: list[str | None] = [None]
def on_enqueue(request_id: str) -> None:
request_id_ref[0] = request_id
# The queue callback checks interruption on every update while the job
# runs; this covers a cancel that landed before submission (and stays
# outside the try so it cannot record a ledger entry for a job that
# was never submitted).
_check_interruption()
started = time.monotonic()
try:
result = await client.subscribe(
endpoint,
arguments=arguments,
with_logs=True,
on_enqueue=on_enqueue,
on_queue_update=callback,
)
except FalApiError:
raise
except Exception as exc:
if _is_interruption(exc):
raise
raise_fal_error(endpoint, exc)
finally:
_finalize_live_call(endpoint, request_id_ref[0], started)
await _close_async_client(client)
_store_result_in_cache(endpoint, arguments, result, request_id_ref[0])
_remember_result_urls(endpoint, request_id_ref[0], result)
return result
@staticmethod
@@ -329,6 +427,7 @@ class ApiHandler:
JobStore().mark_collected(request_id)
except Exception as exc:
logger.debug("[%s] job store mark_collected failed: %s", endpoint, exc)
_remember_result_urls(endpoint, request_id, result)
return result
@staticmethod
+223
View File
@@ -0,0 +1,223 @@
"""Zip archive helpers for dataset preparation (LoRA training uploads)."""
from __future__ import annotations
import io
import os
import tempfile
import zipfile
from typing import Any
import torch
from .config import FalConfig
from .errors import FalApiError
from .images import ImageUtils
from .logger import logger
_MODEL_NAME = "archive"
def _safe_unlink(path: str | None) -> None:
"""Delete a temp file, ignoring errors."""
if path is None:
return
try:
os.unlink(path)
except OSError:
pass
def _split_frames(images: Any) -> list[Any]:
"""Split an IMAGE input (batch tensor, list, or single image) into frames."""
if images is None:
return []
if isinstance(images, torch.Tensor):
if images.ndim == 4:
return [images[i] for i in range(images.shape[0])]
return [images]
if isinstance(images, (list, tuple)):
return list(images)
return [images]
def _normalize_extensions(extensions: Any) -> list[str] | None:
"""Normalize an extension filter to lowercase dot-prefixed suffixes."""
if not extensions:
return None
normalized = []
for ext in extensions:
cleaned = str(ext).strip().lower()
if not cleaned:
continue
normalized.append(cleaned if cleaned.startswith(".") else f".{cleaned}")
return normalized or None
def _matches_filter(file_name: str, extensions: list[str] | None) -> bool:
"""Whether a file passes the hidden-file and extension filters."""
if file_name.startswith("."):
return False
if extensions is None:
return True
return os.path.splitext(file_name)[1].lower() in extensions
def _collect_folder_files(
folder: str, extensions: list[str] | None, recursive: bool
) -> list[str]:
"""List matching files in a folder (sorted, hidden entries skipped)."""
if not recursive:
return [
os.path.join(folder, name)
for name in sorted(os.listdir(folder))
if os.path.isfile(os.path.join(folder, name))
and _matches_filter(name, extensions)
]
matches: list[str] = []
for root, dirs, files in os.walk(folder):
dirs[:] = sorted(d for d in dirs if not d.startswith("."))
for name in sorted(files):
if _matches_filter(name, extensions):
matches.append(os.path.join(root, name))
return matches
def _new_temp_zip_path() -> str:
"""Reserve a temp .zip path and return it."""
with tempfile.NamedTemporaryFile(suffix=".zip", delete=False) as temp_zip:
return temp_zip.name
class ArchiveUtils:
"""Utility functions for building and uploading zip archives."""
@staticmethod
def zip_images(
images: Any,
captions: list[str] | None = None,
name_prefix: str = "image",
) -> str:
"""Zip an IMAGE batch as image_0.png, image_1.png, ... and return the local zip path.
``captions`` (optional, one per image, entries may be "") also writes
image_0.txt, image_1.txt, ... — the standard LoRA-training caption
layout. The caller is responsible for uploading/deleting the zip
(see ``upload_zip``).
"""
frames = _split_frames(images)
if not frames:
raise FalApiError(
_MODEL_NAME,
"No images provided to zip. Connect an IMAGE batch with at least one frame.",
)
if captions is not None and len(captions) != len(frames):
raise FalApiError(
_MODEL_NAME,
f"Caption count ({len(captions)}) does not match image count ({len(frames)}). "
"Provide exactly one caption per image (blank entries are allowed) or none at all.",
)
prefix = (name_prefix or "image").strip() or "image"
zip_path: str | None = None
try:
zip_path = _new_temp_zip_path()
with zipfile.ZipFile(zip_path, "w") as zip_file:
for index, frame in enumerate(frames):
pil_image = ImageUtils.tensor_to_pil(frame)
buffer = io.BytesIO()
pil_image.save(buffer, format="PNG")
zip_file.writestr(f"{prefix}_{index}.png", buffer.getvalue())
if captions is not None:
zip_file.writestr(f"{prefix}_{index}.txt", captions[index])
return zip_path
except FalApiError:
_safe_unlink(zip_path)
raise
except Exception as exc:
_safe_unlink(zip_path)
logger.error("Failed to create image zip: %s", exc)
raise FalApiError(
_MODEL_NAME, f"Failed to create image zip: {exc}"
) from exc
@staticmethod
def zip_folder(
folder_path: str,
include_extensions: list[str] | None = None,
recursive: bool = False,
) -> str:
"""Zip a folder's files and return the local zip path.
``include_extensions`` filters by suffix (e.g. [".png", ".txt"]); None
includes everything. Hidden files/directories are always skipped.
Non-recursive by default; arcnames are relative to the folder.
"""
if not folder_path or not isinstance(folder_path, str) or not folder_path.strip():
raise FalApiError(
_MODEL_NAME,
"folder_path is empty. Provide the path to a folder of dataset files.",
)
folder = os.path.abspath(os.path.expanduser(folder_path.strip()))
if not os.path.isdir(folder):
raise FalApiError(
_MODEL_NAME,
f"Folder not found: {folder}. Provide the path to an existing directory.",
)
extensions = _normalize_extensions(include_extensions)
files = _collect_folder_files(folder, extensions, recursive)
if not files:
suffix_hint = f" matching extensions {extensions}" if extensions else ""
raise FalApiError(
_MODEL_NAME,
f"No files{suffix_hint} found in {folder}. "
"Check the folder contents, the extension filter, and the recursive flag.",
)
# Folder zips get uploaded to fal's CDN: log loudly what is being read
# and cap runaway/hostile selections ([archive] section in config.ini).
total_bytes = sum(os.path.getsize(f) for f in files)
max_files = int(FalConfig().get_setting("archive", "max_files", 5000))
max_mb = float(FalConfig().get_setting("archive", "max_total_mb", 2048))
if len(files) > max_files or total_bytes > max_mb * 1024 * 1024:
raise FalApiError(
_MODEL_NAME,
f"Refusing to zip {len(files)} file(s) / {total_bytes / 1048576:.1f} MiB "
f"from {folder} — over the [archive] limits (max_files={max_files}, "
f"max_total_mb={max_mb:g}). Narrow the folder/extensions or raise the "
"limits in config.ini.",
)
logger.info(
"archive: zipping %d file(s) (%.1f MiB) from %s",
len(files),
total_bytes / 1048576,
folder,
)
zip_path: str | None = None
try:
zip_path = _new_temp_zip_path()
with zipfile.ZipFile(zip_path, "w") as zip_file:
for file_path in files:
zip_file.write(file_path, os.path.relpath(file_path, folder))
return zip_path
except Exception as exc:
_safe_unlink(zip_path)
logger.error("Failed to zip folder %s: %s", folder, exc)
raise FalApiError(
_MODEL_NAME, f"Failed to zip folder {folder}: {exc}"
) from exc
@staticmethod
def upload_zip(zip_path: str) -> str:
"""Upload a local zip to fal.ai and return its URL; the zip is always deleted."""
if not zip_path or not os.path.isfile(zip_path):
raise FalApiError(
_MODEL_NAME,
f"Zip file not found: {zip_path}. Build it with zip_images/zip_folder first.",
)
try:
return ImageUtils.upload_file(zip_path)
finally:
_safe_unlink(zip_path)
+8 -8
View File
@@ -139,16 +139,16 @@ class BillingUtils:
as SpendGuard.preflight do not hammer the API; ``force=True`` bypasses.
"""
global _balance_cache
now = time.time()
if not force:
with _balance_lock:
value, fetched_at = _balance_cache
if fetched_at > 0 and now - fetched_at < _BALANCE_CACHE_TTL_S:
return value
value = _fetch_balance()
# the fetch happens under the lock so a cold-start burst of parallel
# callers (e.g. 8 caption workers hitting SpendGuard at once) collapses
# into a single API call instead of hammering /account/billing
with _balance_lock:
value, fetched_at = _balance_cache
if not force and fetched_at > 0 and time.time() - fetched_at < _BALANCE_CACHE_TTL_S:
return value
value = _fetch_balance()
_balance_cache = [value, time.time()]
return value
return value
@staticmethod
def get_recent_usage(limit: int = 50) -> list[dict[str, Any]] | None:
+217
View File
@@ -0,0 +1,217 @@
"""Checks the live fal.ai catalog for models missing from the local registry.
``check_for_new_models`` diffs the public catalog against the committed
``data/fal_registry.json`` and caches the result module-level (1h TTL) so the
sidebar and the startup check share one fetch. ``schedule_startup_check``
spawns a delayed daemon thread that logs a single INFO line when the local
registry is behind. Nothing in here may break node loading: the startup path
never raises.
"""
from __future__ import annotations
import json
import os
import threading
import time
from typing import Any
from .logger import logger
CATALOG_URL = "https://fal.ai/api/models?page={page}&total={total}"
_USER_AGENT = "ComfyUI-fal-API-freshness/1.0"
_PAGE_SIZE = 100
_MAX_PAGES = 25
_MAX_NEW_LISTED = 25
_CACHE_TTL_S = 3600.0
_STARTUP_DELAY_S = 10.0
_DEFAULT_TIMEOUT_S = 20.0
_lock = threading.Lock()
_cached_result: dict[str, Any] | None = None
_startup_scheduled = False
def _registry_path() -> str:
"""Path to data/fal_registry.json at the repo root."""
utils_dir = os.path.dirname(os.path.abspath(__file__))
repo_root = os.path.dirname(os.path.dirname(utils_dir))
return os.path.join(repo_root, "data", "fal_registry.json")
def _registry_endpoint_ids() -> set[str]:
"""Endpoint ids present in the committed registry; empty set on failure."""
try:
with open(_registry_path(), encoding="utf-8") as handle:
registry = json.load(handle)
models = registry.get("models")
if not isinstance(models, list):
raise ValueError("'models' is not a list")
return {
str(model["endpoint_id"])
for model in models
if isinstance(model, dict) and model.get("endpoint_id")
}
except Exception as err:
logger.debug("freshness: could not read local registry: %s", err)
return set()
def _extract_items(payload: Any) -> list[dict[str, Any]]:
"""Normalize one catalog API page into a list of item dicts."""
if isinstance(payload, list):
raw = payload
elif isinstance(payload, dict):
raw = next(
(
payload[key]
for key in ("items", "models", "data", "results")
if isinstance(payload.get(key), list)
),
[],
)
else:
raw = []
return [item for item in raw if isinstance(item, dict)]
def _fetch_catalog(timeout_s: float) -> list[dict[str, Any]]:
"""Fetch catalog pages until an empty page (hard cap _MAX_PAGES).
Raises RuntimeError when the very first page cannot be fetched; a failure
on a later page returns the partial catalog (better a lower bound than
nothing).
"""
import requests
items: list[dict[str, Any]] = []
for page in range(1, _MAX_PAGES + 1):
url = CATALOG_URL.format(page=page, total=_PAGE_SIZE)
try:
response = requests.get(url, headers={"User-Agent": _USER_AGENT}, timeout=timeout_s)
response.raise_for_status()
page_items = _extract_items(response.json())
except Exception as err:
if page == 1:
raise RuntimeError(f"fal catalog fetch failed: {err}") from err
logger.debug("freshness: catalog page %d failed (%s); using partial catalog", page, err)
break
if not page_items:
break
items = items + page_items
return items
def _is_live_public(item: dict[str, Any]) -> bool:
return bool(
item.get("id")
and item.get("status") == "public"
and not item.get("deprecated")
and not item.get("removed")
)
def _new_model_entry(item: dict[str, Any]) -> dict[str, Any]:
return {
"endpoint_id": str(item.get("id") or ""),
"title": str(item.get("title") or "").strip(),
"category": str(item.get("category") or "").strip(),
"published_at": str(item.get("publishedAt") or item.get("date") or "").strip(),
}
def check_for_new_models(timeout_s: float = _DEFAULT_TIMEOUT_S) -> dict[str, Any]:
"""Diff the live fal catalog against the local registry (cached, 1h TTL).
Returns ``{"new_count", "new_models" (newest first, max 25), "checked_at"}``.
Raises RuntimeError when the catalog cannot be reached at all; failed runs
are never cached.
"""
global _cached_result
with _lock:
if (
_cached_result is not None
and time.time() - float(_cached_result.get("checked_at", 0)) < _CACHE_TTL_S
):
return _cached_result
known_ids = _registry_endpoint_ids()
catalog = _fetch_catalog(timeout_s)
live = [item for item in catalog if _is_live_public(item)]
seen: set[str] = set()
fresh: list[dict[str, Any]] = []
for item in live:
endpoint_id = str(item["id"])
if endpoint_id in known_ids or endpoint_id in seen:
continue
seen.add(endpoint_id)
fresh = fresh + [_new_model_entry(item)]
fresh.sort(key=lambda entry: entry["published_at"], reverse=True)
result = {
"new_count": len(fresh),
"new_models": fresh[:_MAX_NEW_LISTED],
"checked_at": time.time(),
}
with _lock:
_cached_result = result
return result
def _startup_check_enabled() -> bool:
if os.environ.get("FAL_DISABLE_STARTUP_CHECK"):
return False
try:
from .config import FalConfig
value = FalConfig().get_setting("registry", "startup_check", True)
except Exception as err:
logger.debug("freshness: could not read startup_check setting: %s", err)
return True
if isinstance(value, str):
return value.strip().lower() in ("1", "true", "yes", "on")
return bool(value)
def _startup_worker() -> None:
"""Delayed freshness check; logs one INFO line, never raises."""
try:
time.sleep(_STARTUP_DELAY_S)
result = check_for_new_models()
new_count = result.get("new_count", 0)
if new_count:
logger.info(
"fal catalog: %d models newer than the local registry — "
"see the fal sidebar or run scripts/build_registry.py",
new_count,
)
else:
logger.debug("fal catalog: no new model IDs found; existing schemas may still have updates")
except Exception as err:
logger.debug("fal registry freshness check failed: %s", err)
def schedule_startup_check() -> bool:
"""Spawn the delayed startup freshness thread once. Never raises.
Returns True when a thread was started (enabled and not yet scheduled).
"""
global _startup_scheduled
try:
with _lock:
if _startup_scheduled:
return False
_startup_scheduled = True
if not _startup_check_enabled():
logger.debug("freshness: startup check disabled via config")
return False
thread = threading.Thread(
target=_startup_worker, name="fal-registry-freshness", daemon=True
)
thread.start()
return True
except Exception as err:
logger.debug("freshness: could not schedule startup check: %s", err)
return False
+3 -1
View File
@@ -196,7 +196,9 @@ class JobStore:
if status:
query += " WHERE status = ?"
params = (status,)
query += " ORDER BY submitted_at DESC LIMIT ?"
# SQLite's timestamp resolution can tie for back-to-back submits.
# rowid preserves insertion order, so the newest job stays first.
query += " ORDER BY submitted_at DESC, rowid DESC LIMIT ?"
params = (*params, int(limit))
with self._lock:
conn = self._connection()
+27
View File
@@ -33,6 +33,27 @@ def _is_http_url(value: str) -> bool:
return value.startswith(("http://", "https://"))
def _require_http_url(value: Any, operation: str) -> str:
"""Return a normalized HTTP(S) URL or raise an actionable fal error.
``requests`` otherwise turns values such as ``"E"`` (historically the
first character of an error string routed through a list output) into a
cryptic ``MissingSchema`` exception. Validate at the media boundary so
the bad upstream value and the responsible operation remain visible.
"""
url = str(value or "").strip()
parsed = urlparse(url)
if parsed.scheme.lower() not in {"http", "https"} or not parsed.netloc:
preview = repr(url if len(url) <= 120 else f"{url[:117]}...")
raise FalApiError(
operation,
f"Expected an HTTP(S) media URL, got {preview}. "
"The upstream generation may have failed, or a non-URL output "
"may be connected to a media input.",
)
return url
def _suffix_from_url(url: str, default: str) -> str:
"""Derive a file suffix from a URL path, falling back to a default."""
suffix = os.path.splitext(urlparse(url).path)[1]
@@ -164,9 +185,15 @@ class MediaUtils:
http(s) URL inputs pass through untouched.
"""
@staticmethod
def require_http_url(value: Any, operation: str = "media-download") -> str:
"""Public validation hook for legacy nodes that return media URLs."""
return _require_http_url(value, operation)
@staticmethod
def download_url_to_temp(url: str, suffix: str) -> str:
"""Stream a URL to a temp file and return its local path."""
url = _require_http_url(url, "media-download")
temp_path: str | None = None
try:
with requests.get(url, stream=True, timeout=_DOWNLOAD_TIMEOUT) as resp:
+59 -3
View File
@@ -50,6 +50,14 @@ _SCHEMA = (
created REAL
)
""",
"""
CREATE TABLE IF NOT EXISTS request_urls (
url TEXT PRIMARY KEY,
endpoint TEXT,
request_id TEXT,
created REAL
)
""",
)
@@ -301,17 +309,65 @@ class ResultCache:
"misses": self._misses,
}
def remember_urls(self, endpoint: str, request_id: str, result: Any) -> None:
"""Record every media URL in a result → (endpoint, request_id).
Called on every successful fetch (live, async collect, recovery) so
provenance lookups work regardless of which path produced the result.
Best-effort: never raises.
"""
try:
if not request_id or not isinstance(result, dict):
return
urls: list[str] = []
def dig(value: Any) -> None:
if isinstance(value, dict):
candidate = value.get("url")
if isinstance(candidate, str) and candidate.startswith("http"):
urls.append(candidate)
for child in value.values():
dig(child)
elif isinstance(value, list):
for child in value:
dig(child)
dig(result)
if not urls:
return
now = time.time()
with self._lock:
conn = self._connection()
if conn is None:
return
conn.executemany(
"INSERT OR REPLACE INTO request_urls VALUES (?, ?, ?, ?)",
[(u, endpoint, request_id, now) for u in urls[:64]],
)
conn.commit()
except Exception as exc:
logger.debug("result cache remember_urls failed: %s", exc)
def find_request_by_url(self, url: str) -> dict[str, Any] | None:
"""Find the origin of a result URL: {"endpoint_id", "request_id"} or None.
Scans cached results (most recently used first) for one whose JSON
contains ``url`` as an exact substring. Only rows that recorded a
request_id qualify. Best-effort: any failure is a miss.
Checks the explicit request_urls table first (covers async collect and
recovery paths), then falls back to scanning cached results for the
URL as a substring. Best-effort: any failure is a miss.
"""
try:
target = (url or "").strip()
if not target:
return None
with self._lock:
conn = self._connection()
if conn is not None:
row = conn.execute(
"SELECT endpoint, request_id FROM request_urls WHERE url = ?",
(target,),
).fetchone()
if row is not None:
return {"endpoint_id": row[0], "request_id": row[1]}
# LIKE treats %, _ (and our escape char) specially — escape them
# so URLs containing percent-encoding still match literally.
escaped = (
+125 -2
View File
@@ -2770,8 +2770,16 @@ class SeedanceProImageToVideoNode:
"fal-ai/bytedance/seedance/v1/pro/image-to-video", arguments, variations
)
# Return list of video URLs
return ([r["video"]["url"] for r in results],)
# Validate before exposing URLs to downstream loaders. Older error
# paths could leak an error string through this list output; a
# downstream loader would then receive its first character ("E")
# and raise requests.exceptions.MissingSchema.
endpoint = "fal-ai/bytedance/seedance/v1/pro/image-to-video"
video_urls = [
MediaUtils.require_http_url(r["video"]["url"], endpoint)
for r in results
]
return (video_urls,)
except Exception as e:
return ApiHandler.handle_video_generation_error(
@@ -2779,6 +2787,119 @@ class SeedanceProImageToVideoNode:
)
class Seedance25VideoToVideoNode:
"""Curated Seedance 2.5 video editing node."""
ENDPOINT = "bytedance/seedance-2.5/reference-to-video"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"video": (
"VIDEO",
{
"tooltip": "Source video to edit. URL-backed VIDEO inputs are passed through without re-uploading.",
},
),
"prompt": (
"STRING",
{
"default": "",
"multiline": True,
"tooltip": "Describe the edits to apply to the source video.",
},
),
},
"optional": {
"resolution": (
["480p", "720p", "1080p"],
{"default": "720p", "tooltip": "Output video resolution."},
),
"generate_audio": (
"BOOLEAN",
{
"default": True,
"tooltip": "Generate synchronized audio for the edited video.",
},
),
"bitrate_mode": (
["standard", "high"],
{
"default": "standard",
"tooltip": "Use the standard bitrate or request a larger, higher-quality encode.",
},
),
"seed": (
"INT",
{
"default": -1,
"min": -1,
"max": 2147483647,
"tooltip": "Random seed for reproducible results; -1 lets the API choose.",
},
),
"force_rerun": (
"BOOLEAN",
{
"default": False,
"tooltip": "Bypass the persistent result cache and submit a new fal request.",
},
),
},
}
RETURN_TYPES = ("VIDEO", "STRING")
RETURN_NAMES = ("video", "video_url")
FUNCTION = "edit_video"
CATEGORY = "FAL/VideoGeneration"
DESCRIPTION = (
"Edit one video with Seedance 2.5. The endpoint is always called with "
"task='editing' and the source as a single video_urls entry."
)
@classmethod
def IS_CHANGED(cls, force_rerun=False, **_kwargs):
if force_rerun:
return float("nan")
return False
def edit_video(
self,
video,
prompt,
resolution="720p",
generate_audio=True,
bitrate_mode="standard",
seed=-1,
force_rerun=False,
):
try:
uploaded_url = MediaUtils.upload_video(video)
arguments = {
"prompt": prompt,
"task": "editing",
"video_urls": [uploaded_url],
"resolution": resolution,
"generate_audio": generate_audio,
"bitrate_mode": bitrate_mode,
}
if seed != -1:
arguments["seed"] = seed
result = ApiHandler.submit_and_get_result(
self.ENDPOINT,
arguments,
skip_cache=bool(force_rerun),
)
video_url = MediaUtils.require_http_url(
result["video"]["url"], self.ENDPOINT
)
return (MediaUtils.video_from_url(video_url), video_url)
except Exception as e:
return ApiHandler.handle_video_generation_error(self.ENDPOINT, e)
class Veo3Node:
@classmethod
def INPUT_TYPES(cls):
@@ -3725,6 +3846,7 @@ NODE_CLASS_MAPPINGS = {
"DYWanFun22_fal": DYWanFun22Node,
"DYWanUpscaler_fal": DYWanUpscalerNode,
"SeedanceImageToVideo_fal": SeedanceImageToVideoNode,
"Seedance25VideoToVideo_fal": Seedance25VideoToVideoNode,
"SeedanceProImageToVideo_fal": SeedanceProImageToVideoNode,
"SeedanceTextToVideo_fal": SeedanceTextToVideoNode,
"Veo3_fal": Veo3Node,
@@ -3771,6 +3893,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"Veo2ImageToVideo_fal": "Google Veo2 Image-to-Video (fal)",
"WanPro_fal": "Wan Pro Image-to-Video (fal)",
"SeedanceImageToVideo_fal": "Seedance Image-to-Video (fal)",
"Seedance25VideoToVideo_fal": "Seedance 2.5 Video-to-Video (fal)",
"SeedanceProImageToVideo_fal": "Seedance Pro Image-to-Video (fal)",
"SeedanceTextToVideo_fal": "Seedance Text-to-Video (fal)",
"Veo3_fal": "Veo3 Video Generation (fal)",
+2 -1
View File
@@ -1,7 +1,7 @@
[project]
name = "fal-api"
description = "Custom nodes for using fal API with auto-generated full-catalog coverage of fal.ai models. Video generation with Kling, Runway, Luma. Image generation with Flux. LLMs and VLMs OpenAI, Claude, Llama and Gemini."
version = "2.3.0"
version = "2.5.1"
license = {file = "LICENSE"}
requires-python = ">=3.9"
dependencies = [
@@ -11,6 +11,7 @@ dependencies = [
"numpy",
"pillow",
"requests",
"av",
]
[project.urls]
+1
View File
@@ -4,3 +4,4 @@ opencv-python
numpy
pillow
requests
av
+56 -24
View File
@@ -1,10 +1,11 @@
#!/usr/bin/env python3
"""Regenerate the auto-generated model list section of README.md.
"""Regenerate the auto-generated model catalog in MODELS.md.
Reads data/fal_registry.json and rewrites ONLY the section between
`<!-- BEGIN GENERATED MODEL LIST -->` and `<!-- END GENERATED MODEL LIST -->`
in README.md. Everything outside the markers is left untouched, and running
the script twice in a row produces no diff.
in MODELS.md. Everything outside the markers is left untouched, and running
the script twice in a row produces no diff. If MODELS.md does not exist yet,
it is created with a standard header around the markers.
Usage:
python scripts/build_readme.py
@@ -19,11 +20,26 @@ from typing import Any
REPO_ROOT = Path(__file__).resolve().parents[1]
REGISTRY_PATH = REPO_ROOT / "data" / "fal_registry.json"
README_PATH = REPO_ROOT / "README.md"
MODELS_PATH = REPO_ROOT / "MODELS.md"
BEGIN_MARKER = "<!-- BEGIN GENERATED MODEL LIST -->"
END_MARKER = "<!-- END GENERATED MODEL LIST -->"
MODELS_TEMPLATE = f"""# fal Model Catalog — auto-generated
Every live auto-generated model node in [ComfyUI-fal-API](README.md), grouped
by category (largest first). Click a category to expand it. Historical endpoint
schemas retained for workflow compatibility are registered under
`FAL/Compatibility` and intentionally omitted from this live catalog.
Do not edit this file by hand — refresh `data/fal_registry.json` with
`python scripts/build_registry.py`, then regenerate this catalog with
`python scripts/build_readme.py`.
{BEGIN_MARKER}
{END_MARKER}
"""
MODEL_URL_TEMPLATE = "https://fal.ai/models/{endpoint_id}"
@@ -89,48 +105,64 @@ def render_category(category: str, models: list[dict[str, Any]]) -> str:
def render_generated_section(registry: dict[str, Any]) -> str:
models = registry["models"]
model_count = registry.get("model_count", len(models))
published = [str(m.get("published_at", "")) for m in registry.get("models", [])]
live_models = [model for model in models if not model.get("deprecated")]
live_count = registry.get("live_model_count", len(live_models))
deprecated_count = registry.get("deprecated_model_count", len(models) - len(live_models))
published = [str(model.get("published_at", "")) for model in live_models]
generated_date = max(published)[:10] if any(published) else "unknown"
summary = (
f"{model_count} models · newest model {generated_date} · "
f"{live_count} live models · {deprecated_count} compatibility-preserved · "
f"newest model {generated_date} · "
"refresh with `scripts/build_registry.py`"
)
blocks = [
render_category(category, grouped)
for category, grouped in group_by_category(models)
for category, grouped in group_by_category(live_models)
]
return "\n\n".join([summary, *blocks])
def replace_between_markers(readme: str, generated: str) -> str:
begin = readme.find(BEGIN_MARKER)
end = readme.find(END_MARKER)
def replace_between_markers(document: str, generated: str) -> str:
begin = document.find(BEGIN_MARKER)
end = document.find(END_MARKER)
if begin == -1 or end == -1 or end < begin:
raise SystemExit(
f"README.md must contain '{BEGIN_MARKER}' followed by '{END_MARKER}'"
f"MODELS.md must contain '{BEGIN_MARKER}' followed by '{END_MARKER}'"
)
head = readme[: begin + len(BEGIN_MARKER)]
tail = readme[end:]
head = document[: begin + len(BEGIN_MARKER)]
tail = document[end:]
return f"{head}\n\n{generated}\n\n{tail}"
def read_models_document(path: Path) -> str:
if not path.is_file():
return MODELS_TEMPLATE
try:
return path.read_text(encoding="utf-8")
except OSError as err:
raise SystemExit(f"Failed to read {path}: {err}") from err
def main() -> int:
registry = load_registry(REGISTRY_PATH)
try:
readme = README_PATH.read_text(encoding="utf-8")
except OSError as err:
raise SystemExit(f"Failed to read {README_PATH}: {err}") from err
document = read_models_document(MODELS_PATH)
updated = replace_between_markers(readme, render_generated_section(registry))
if updated == readme:
print(f"README.md already up to date ({registry.get('model_count')} models)")
updated = replace_between_markers(document, render_generated_section(registry))
if MODELS_PATH.is_file() and updated == document:
print(
"MODELS.md already up to date "
f"({registry.get('live_model_count', registry.get('model_count'))} live, "
f"{registry.get('deprecated_model_count', 0)} compatibility-preserved)"
)
return 0
README_PATH.write_text(updated, encoding="utf-8")
MODELS_PATH.write_text(updated, encoding="utf-8")
print(
f"README.md model list regenerated: {registry.get('model_count')} models, "
f"{len(group_by_category(registry['models']))} categories"
f"MODELS.md model catalog regenerated: "
f"{registry.get('live_model_count', registry.get('model_count'))} live models, "
f"{registry.get('deprecated_model_count', 0)} compatibility-preserved, "
f"{len(group_by_category([m for m in registry['models'] if not m.get('deprecated')]))} "
"categories"
)
return 0
+169 -40
View File
@@ -17,6 +17,8 @@ Stdlib only. Usage:
import argparse
import json
import logging
import os
import re
import time
import urllib.error
import urllib.request
@@ -30,7 +32,6 @@ USER_AGENT = "ComfyUI-fal-API-registry-builder/1.0"
FETCH_ATTEMPTS = 3
BACKOFF_BASE_SECONDS = 1.5
MAX_INPUT_PROPERTIES = 40
MAX_DESCRIPTION_CHARS = 500
MULTILINE_NAMES = frozenset({"prompt", "negative_prompt", "text", "script", "dialogue"})
MULTILINE_DESCRIPTION_THRESHOLD = 120
@@ -170,7 +171,7 @@ def resolve_ref(schema, components):
ref = schema.get("$ref", "")
if not ref.startswith("#/components/schemas/"):
return schema
name = ref.rsplit("/", 1)[-1]
name = ref.rsplit("/", 1)[-1].replace("~1", "/").replace("~0", "~")
resolved = components.get(name)
if not isinstance(resolved, dict):
return schema
@@ -178,22 +179,6 @@ def resolve_ref(schema, components):
return {**resolved, **siblings}
def non_null_branches(branches, components):
"""Resolve and drop null branches from an anyOf/oneOf list."""
resolved = [resolve_ref(branch, components) for branch in branches if isinstance(branch, dict)]
return [branch for branch in resolved if branch.get("type") != "null"]
def merge_all_of(schema, components):
"""Merge an allOf list (one level), with sibling keys taking precedence."""
merged = {}
for branch in schema.get("allOf", []):
if isinstance(branch, dict):
merged = {**merged, **resolve_ref(branch, components)}
siblings = {key: value for key, value in schema.items() if key != "allOf"}
return {**merged, **siblings}
def is_custom_size_pair(branches):
"""Detect the image_size pattern: [enum-of-presets, width/height object]."""
enum_branch = next((b for b in branches if b.get("enum")), None)
@@ -213,30 +198,92 @@ def is_custom_size_pair(branches):
return None
def normalize_schema(schema, components):
"""Resolve $ref / allOf / anyOf / oneOf one level.
def normalize_schema(schema, components, _seen_refs=frozenset()):
"""Resolve nested references and composition without dropping enum choices.
Returns (resolved_schema, has_custom_size, custom_size_enum_values).
"""
if not isinstance(schema, dict):
return {}, False, None
ref = schema.get("$ref")
if ref and ref in _seen_refs:
return {}, False, None
seen = _seen_refs | {ref} if ref else _seen_refs
resolved = resolve_ref(schema, components)
if "$ref" in resolved and resolved != schema:
return normalize_schema(resolved, components, seen)
has_custom_size, custom_values = False, None
if "allOf" in resolved:
resolved = merge_all_of(resolved, components)
merged = {}
for branch in resolved["allOf"]:
normalized, branch_custom, branch_values = normalize_schema(branch, components, seen)
if branch_custom:
has_custom_size, custom_values = True, branch_values
properties = {**merged.get("properties", {}), **normalized.get("properties", {})}
required = list(dict.fromkeys(merged.get("required", []) + normalized.get("required", [])))
if "enum" in merged and "enum" in normalized:
normalized = {**normalized, "enum": [v for v in merged["enum"] if v in normalized["enum"]]}
merged = {**merged, **normalized}
if properties:
merged["properties"] = properties
if required:
merged["required"] = required
siblings = {key: value for key, value in resolved.items() if key != "allOf"}
if "properties" in siblings:
siblings["properties"] = {**merged.get("properties", {}), **siblings["properties"]}
if "required" in siblings:
siblings["required"] = list(dict.fromkeys(merged.get("required", []) + siblings["required"]))
resolved = {**merged, **siblings}
if "const" in resolved:
literal = resolved["const"]
resolved = {key: value for key, value in resolved.items() if key != "const"}
if literal is None:
resolved = {**resolved, "type": "null"}
else:
resolved = {**resolved, "enum": [literal]}
if isinstance(resolved.get("type"), list):
types = [value for value in resolved["type"] if value != "null"]
if len(types) == 1:
resolved = {**resolved, "type": types[0]}
branches_key = "anyOf" if "anyOf" in resolved else ("oneOf" if "oneOf" in resolved else None)
if branches_key is None:
return resolved, False, None
return resolved, has_custom_size, custom_values
branches = non_null_branches(resolved[branches_key], components)
normalized_branches = [normalize_schema(branch, components, seen) for branch in resolved[branches_key]]
normalized_branches = [entry for entry in normalized_branches if entry[0] and entry[0].get("type") != "null"]
branches = [entry[0] for entry in normalized_branches]
siblings = {key: value for key, value in resolved.items() if key != branches_key}
if not branches:
return siblings, False, None
if len(branches) == 1:
branch, branch_custom, branch_values = normalized_branches[0]
return {**branch, **siblings}, branch_custom, branch_values
custom_enum_branch = is_custom_size_pair(branches)
if custom_enum_branch is not None:
values = list(custom_enum_branch.get("enum", [])) + ["custom_size"]
return {**custom_enum_branch, **siblings}, True, values
if all(branch.get("enum") for branch in branches):
values = []
for branch in branches:
for value in branch["enum"]:
if value is not None and value not in values:
values.append(value)
if "enum" in siblings:
values = [value for value in values if value in siblings["enum"]]
return {**branches[0], **siblings, "enum": values}, False, None
# An enum plus an open string branch is still an open string. Keep the
# literals as suggestions instead of incorrectly restricting API values.
open_string = next((b for b in branches if b.get("type") == "string" and not b.get("enum")), None)
if open_string is not None and all(
b.get("type") == "string" or (b.get("enum") and all(isinstance(value, str) for value in b["enum"]))
for b in branches
):
examples = [value for b in branches for value in b.get("enum", b.get("examples", []))]
return {**open_string, "examples": examples, **siblings}, False, None
enum_branch = next((branch for branch in branches if branch.get("enum")), None)
chosen = enum_branch if enum_branch is not None else branches[0]
return {**chosen, **siblings}, False, None
@@ -284,6 +331,27 @@ def scalar_type_of(schema):
return "json"
def string_suggestions(name, schema):
"""Short example identifiers are suggestions, never strict enum constraints.
Keep prose, prompts and formatted strings as text. The frontend offers
custom values so undocumented languages, voices, model IDs and
future modes remain usable even when examples are incomplete.
"""
if name in MULTILINE_NAMES or name.endswith(("_prompt", "_text")) or schema.get("format"):
return None
examples = schema.get("examples")
if not isinstance(examples, list) or not all(
isinstance(value, str) and re.fullmatch(r"[\w./:+-]{1,80}", value) and "://" not in value
for value in examples
):
return None
values = list(dict.fromkeys(examples))
if len(values) < 2:
return None
return values
def distill_property(name, raw_schema, required_names, components):
"""Distill one input property into a registry input record, or None."""
if name in SKIPPED_PROPERTY_NAMES or name.startswith("_"):
@@ -358,6 +426,10 @@ def distill_property(name, raw_schema, required_names, components):
}
if has_custom_size:
record = {**record, "has_custom_size": True}
if type_name == "string" and not is_list and not media_kind:
suggestions = string_suggestions(name, schema)
if suggestions:
record = {**record, "suggestions": suggestions}
return record
@@ -374,22 +446,11 @@ def ordered_property_names(schema):
def distill_inputs(schema, components, endpoint_id):
"""Distill an Input schema's properties into registry input records."""
schema, _, _ = normalize_schema(schema, components)
properties = schema.get("properties", {})
required_names = set(schema.get("required", []))
names = ordered_property_names(schema)
if len(names) > MAX_INPUT_PROPERTIES:
required_first = [n for n in names if n in required_names]
optional = [n for n in names if n not in required_names]
budget = max(MAX_INPUT_PROPERTIES - len(required_first), 0)
names = required_first + optional[:budget]
logger.info(
"%s: input schema has %d properties, capped to %d",
endpoint_id,
len(properties),
len(names),
)
inputs = []
for name in names:
record = distill_property(name, properties.get(name, {}), required_names, components)
@@ -405,7 +466,7 @@ def distill_inputs(schema, components, endpoint_id):
def ref_name(schema):
"""Extract the local component name from a {'$ref': ...} node."""
ref = schema.get("$ref", "") if isinstance(schema, dict) else ""
return ref.rsplit("/", 1)[-1] if ref.startswith("#/components/schemas/") else None
return ref.rsplit("/", 1)[-1].replace("~1", "/").replace("~0", "~") if ref.startswith("#/components/schemas/") else None
def input_ref_from_paths(doc):
@@ -440,9 +501,21 @@ def output_ref_from_paths(doc):
def select_schema(doc, endpoint_id, suffix, path_lookup):
"""Select the app Input/Output schema from components.schemas."""
components = doc.get("components", {}).get("schemas", {})
# A document can describe several sibling endpoints. Prefer this endpoint's
# operation, including inline schemas, over whichever path happens to be first.
path = "/" + endpoint_id.strip("/")
if suffix == "Input":
operation = doc.get("paths", {}).get(path, {}).get("post", {})
content = operation.get("requestBody", {}).get("content", {})
else:
operation = doc.get("paths", {}).get(path + "/requests/{request_id}", {}).get("get", {})
content = operation.get("responses", {}).get("200", {}).get("content", {})
schema = content.get("application/json", {}).get("schema")
if isinstance(schema, dict):
return normalize_schema(schema, components)[0]
referenced = path_lookup(doc)
if referenced and referenced in components:
return components[referenced]
return normalize_schema(components[referenced], components)[0]
candidates = [name for name in components if name.endswith(suffix)]
if not candidates:
@@ -455,7 +528,7 @@ def select_schema(doc, endpoint_id, suffix, path_lookup):
in normalized_endpoint
]
pool = matching or candidates
return components[max(pool, key=len)]
return normalize_schema(components[max(pool, key=len)], components)[0]
# ---------------------------------------------------------------------------
@@ -547,6 +620,30 @@ def load_json_file(path):
raise RuntimeError(f"Failed to load cache file {path}: {error}") from error
def preserve_missing_records(records, baseline):
"""Retain historical endpoints so saved ComfyUI workflows keep loading.
A missing endpoint may be retired or its OpenAPI schema may be temporarily
unavailable. The old schema remains usable for node registration and is
moved to a compatibility tier at runtime. If the endpoint returns in a
later catalog build, its fresh record automatically replaces this copy.
"""
current_ids = {record["endpoint_id"] for record in records}
preserved = []
for model in baseline.get("models") or []:
if not isinstance(model, dict) or not model.get("endpoint_id"):
continue
if model["endpoint_id"] in current_ids:
continue
compatibility_record = {
**model,
"deprecated": True,
"deprecated_reason": "Endpoint absent from the latest live fal catalog or schema fetch.",
}
preserved.append(compatibility_record)
return records + preserved
def parse_args():
parser = argparse.ArgumentParser(description="Build the fal.ai model registry JSON.")
parser.add_argument("--out", default="data/fal_registry.json", help="Output registry path")
@@ -559,6 +656,16 @@ def parse_args():
parser.add_argument("--catalog-cache", default=None, help="Path to cached catalog JSON")
parser.add_argument("--schemas-cache", default=None, help="Path to cached endpoint_id->OpenAPI JSON")
parser.add_argument("--max-workers", type=int, default=16, help="Concurrent schema fetches")
parser.add_argument(
"--preserve-from",
default=None,
help="Registry whose missing endpoints should be retained for workflow compatibility",
)
parser.add_argument(
"--prune-missing",
action="store_true",
help="Drop endpoints missing from the live build instead of preserving compatibility nodes",
)
return parser.parse_args()
@@ -579,6 +686,11 @@ def main():
logging.basicConfig(level=logging.INFO, format="%(levelname)s %(message)s")
args = parse_args()
baseline = None
baseline_path = args.preserve_from or args.out
if not args.prune_missing and os.path.isfile(baseline_path):
baseline = load_json_file(baseline_path)
now = datetime.now(timezone.utc)
since = now - timedelta(days=args.since_days) if args.since_days > 0 else None
@@ -609,7 +721,11 @@ def main():
continue
records = records + [record]
live_model_count = len(records)
if baseline is not None:
records = preserve_missing_records(records, baseline)
records = sorted(records, key=lambda record: record["endpoint_id"])
deprecated_model_count = sum(bool(record.get("deprecated")) for record in records)
# NOTE: no wall-clock fields (generated_at etc.) — the committed registry
# must be content-deterministic so the weekly refresh workflow only opens a
@@ -618,10 +734,15 @@ def main():
"version": 1,
"window_days": args.since_days,
"model_count": len(records),
"live_model_count": live_model_count,
"deprecated_model_count": deprecated_model_count,
"models": records,
}
with open(args.out, "w", encoding="utf-8") as handle:
# atomic write: the live sidebar refresh runs this inside a running
# ComfyUI — a crash mid-write must not corrupt the tracked registry
tmp_out = args.out + ".tmp"
with open(tmp_out, "w", encoding="utf-8") as handle:
json.dump(
registry,
handle,
@@ -632,8 +753,16 @@ def main():
)
handle.write("\n")
os.replace(tmp_out, args.out)
log_summary(records, skipped)
logger.info("Wrote %d models to %s", len(records), args.out)
logger.info(
"Wrote %d models to %s (%d live, %d compatibility-preserved)",
len(records),
args.out,
live_model_count,
deprecated_model_count,
)
if __name__ == "__main__":
+238
View File
@@ -0,0 +1,238 @@
#!/usr/bin/env python3
"""Validate a generated fal model registry before it replaces the baseline.
The scheduled registry refresh uses this dependency-free gate to reject
truncated catalog responses, duplicate or malformed records, nondeterministic
ordering, and unexpectedly large endpoint changes.
"""
from __future__ import annotations
import argparse
import json
import re
from pathlib import Path
from typing import Any
REQUIRED_TOP_LEVEL = {
"version",
"models",
"model_count",
"live_model_count",
"deprecated_model_count",
}
REQUIRED_MODEL_FIELDS = {
"endpoint_id",
"title",
"category",
"description",
"family",
"lab",
"pricing",
"published_at",
"inputs",
"output_kind",
"output_props",
"thumbnail",
}
OUTPUT_KINDS = {"audio", "file", "image", "images", "json", "text", "video"}
INPUT_TYPES = {"string", "integer", "number", "boolean", "enum", "object", "array", "json"}
ENDPOINT_ID_PATTERN = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._/-]*$")
class RegistryValidationError(ValueError):
"""Raised when a registry cannot safely be promoted."""
def load_registry(path: Path) -> dict[str, Any]:
try:
payload = json.loads(path.read_text(encoding="utf-8"))
except (OSError, json.JSONDecodeError) as exc:
raise RegistryValidationError(f"Could not read {path}: {exc}") from exc
if not isinstance(payload, dict):
raise RegistryValidationError(f"{path} must contain a JSON object")
return payload
def validate_registry(registry: dict[str, Any], *, min_models: int = 500) -> set[str]:
missing_top = REQUIRED_TOP_LEVEL - registry.keys()
if missing_top:
raise RegistryValidationError(f"Registry is missing top-level fields: {sorted(missing_top)}")
if registry["version"] != 1:
raise RegistryValidationError(f"Unsupported registry version: {registry['version']!r}")
models = registry["models"]
if not isinstance(models, list):
raise RegistryValidationError("Registry 'models' must be a list")
if registry["model_count"] != len(models):
raise RegistryValidationError(f"model_count={registry['model_count']} does not match {len(models)} records")
deprecated_count = sum(bool(model.get("deprecated")) for model in models if isinstance(model, dict))
live_count = len(models) - deprecated_count
if registry["live_model_count"] != live_count:
raise RegistryValidationError(
f"live_model_count={registry['live_model_count']} does not match {live_count} records"
)
if registry["deprecated_model_count"] != deprecated_count:
raise RegistryValidationError(
f"deprecated_model_count={registry['deprecated_model_count']} does not match "
f"{deprecated_count} records"
)
if len(models) < min_models:
raise RegistryValidationError(f"Registry has only {len(models)} models; minimum is {min_models}")
endpoint_ids: list[str] = []
for index, model in enumerate(models):
if not isinstance(model, dict):
raise RegistryValidationError(f"Model {index} is not an object")
missing = REQUIRED_MODEL_FIELDS - model.keys()
if missing:
raise RegistryValidationError(f"Model {index} is missing fields: {sorted(missing)}")
endpoint_id = model["endpoint_id"]
if (
not isinstance(endpoint_id, str)
or not endpoint_id.strip()
or not ENDPOINT_ID_PATTERN.fullmatch(endpoint_id)
):
raise RegistryValidationError(f"Model {index} has an invalid endpoint_id")
if not isinstance(model["title"], str) or not model["title"].strip():
raise RegistryValidationError(f"{endpoint_id} has an invalid title")
if not isinstance(model["inputs"], list):
raise RegistryValidationError(f"{endpoint_id} inputs must be a list")
input_names = set()
for inp in model["inputs"]:
if not isinstance(inp, dict) or not isinstance(inp.get("name"), str) or not inp["name"]:
raise RegistryValidationError(f"{endpoint_id} has an invalid input record")
name = inp["name"]
if name in input_names:
raise RegistryValidationError(f"{endpoint_id} has duplicate input {name}")
input_names.add(name)
if inp.get("type") not in INPUT_TYPES:
raise RegistryValidationError(f"{endpoint_id}.{name} has an invalid input type")
if not isinstance(inp.get("required"), bool):
raise RegistryValidationError(f"{endpoint_id}.{name} required must be a boolean")
if inp["type"] == "enum" and (not isinstance(inp.get("enum"), list) or not inp["enum"]):
raise RegistryValidationError(f"{endpoint_id}.{name} has no enum choices")
if "suggestions" in inp and (
inp["type"] != "string" or not isinstance(inp["suggestions"], list)
or not inp["suggestions"] or not all(isinstance(v, str) for v in inp["suggestions"])
):
raise RegistryValidationError(f"{endpoint_id}.{name} has invalid suggestions")
if not isinstance(model["output_props"], list):
raise RegistryValidationError(f"{endpoint_id} output_props must be a list")
if model["output_kind"] not in OUTPUT_KINDS:
raise RegistryValidationError(f"{endpoint_id} has unknown output_kind={model['output_kind']!r}")
if "deprecated" in model and not isinstance(model["deprecated"], bool):
raise RegistryValidationError(f"{endpoint_id} deprecated must be a boolean")
endpoint_ids.append(endpoint_id)
if len(endpoint_ids) != len(set(endpoint_ids)):
duplicates = sorted(endpoint_id for endpoint_id in set(endpoint_ids) if endpoint_ids.count(endpoint_id) > 1)
raise RegistryValidationError(f"Duplicate endpoint IDs: {duplicates[:10]}")
if endpoint_ids != sorted(endpoint_ids):
raise RegistryValidationError("Models must be sorted by endpoint_id")
return set(endpoint_ids)
def compare_registries(
baseline_ids: set[str],
candidate_ids: set[str],
*,
max_removal_fraction: float = 0.05,
max_addition_fraction: float = 0.25,
allow_large_change: bool = False,
) -> tuple[set[str], set[str]]:
if not 0 <= max_removal_fraction <= 1:
raise RegistryValidationError("max_removal_fraction must be between 0 and 1")
if not 0 <= max_addition_fraction <= 1:
raise RegistryValidationError("max_addition_fraction must be between 0 and 1")
added = candidate_ids - baseline_ids
removed = baseline_ids - candidate_ids
removal_fraction = len(removed) / len(baseline_ids) if baseline_ids else 0.0
addition_fraction = len(added) / len(baseline_ids) if baseline_ids else 0.0
if removal_fraction > max_removal_fraction and not allow_large_change:
raise RegistryValidationError(
f"Candidate removes {len(removed)}/{len(baseline_ids)} endpoints "
f"({removal_fraction:.1%}), above the {max_removal_fraction:.1%} limit"
)
if addition_fraction > max_addition_fraction and not allow_large_change:
raise RegistryValidationError(
f"Candidate adds {len(added)}/{len(baseline_ids)} endpoints "
f"({addition_fraction:.1%}), above the {max_addition_fraction:.1%} limit"
)
return added, removed
def compare_model_inputs(baseline: dict[str, Any], candidate: dict[str, Any]) -> list[str]:
"""Find controls or choices that a refresh would remove from existing nodes.
Intentional upstream removals require review; silently accepting a partial
schema can otherwise publish missing duration/resolution controls as valid.
"""
previous = {model["endpoint_id"]: model for model in baseline["models"]}
regressions = []
for model in candidate["models"]:
endpoint_id = model["endpoint_id"]
if endpoint_id not in previous:
continue
current = {inp["name"]: inp for inp in model["inputs"]}
for old in previous[endpoint_id]["inputs"]:
name = old["name"]
new = current.get(name)
if new is None:
regressions.append(f"{endpoint_id}: removed input {name}")
elif old["type"] in ("integer", "number", "boolean") and new["type"] == "json":
regressions.append(f"{endpoint_id}.{name}: lost {old['type']} control")
elif old["type"] == "enum" or old.get("suggestions"):
choices = new.get("enum") if new["type"] == "enum" else new.get("suggestions")
if not choices:
control = "enum control" if old["type"] == "enum" else "suggested choices"
regressions.append(f"{endpoint_id}.{name}: lost {control}")
else:
old_choices = old["enum"] if old["type"] == "enum" else old["suggestions"]
lost = [value for value in old_choices if value not in choices]
if lost:
regressions.append(f"{endpoint_id}.{name}: removed choices {lost}")
return regressions
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("candidate", type=Path, help="Generated registry to validate")
parser.add_argument("--baseline", type=Path, help="Committed registry to compare")
parser.add_argument("--min-models", type=int, default=500)
parser.add_argument("--max-removal-fraction", type=float, default=0.05)
parser.add_argument("--max-addition-fraction", type=float, default=0.25)
parser.add_argument("--allow-large-change", action="store_true")
parser.add_argument("--allow-input-removal", action="store_true",
help="Allow reviewed removals of existing input controls or enum choices")
return parser.parse_args()
def main() -> None:
args = parse_args()
candidate = load_registry(args.candidate)
candidate_ids = validate_registry(candidate, min_models=args.min_models)
added: set[str] = set()
removed: set[str] = set()
if args.baseline:
baseline = load_registry(args.baseline)
baseline_ids = validate_registry(baseline, min_models=args.min_models)
added, removed = compare_registries(
baseline_ids,
candidate_ids,
max_removal_fraction=args.max_removal_fraction,
max_addition_fraction=args.max_addition_fraction,
allow_large_change=args.allow_large_change,
)
regressions = compare_model_inputs(baseline, candidate)
if regressions and not args.allow_input_removal:
raise RegistryValidationError(
"Candidate removes existing controls; review before using --allow-input-removal:\n"
+ "\n".join(regressions[:20])
)
print(f"Registry valid: {len(candidate_ids)} models " f"(+{len(added)} / -{len(removed)} vs baseline)")
if __name__ == "__main__":
main()
+14
View File
@@ -14,6 +14,15 @@ import pytest
ROOT = Path(__file__).resolve().parents[1]
PKG = "ComfyUI_fal_API"
# Keep local helper modules (notably scripts/) ahead of unrelated installed
# packages when pytest is invoked via its console entry point instead of
# ``python -m pytest``.
if str(ROOT) not in sys.path:
sys.path.insert(0, str(ROOT))
# never let the freshness daemon make live network calls during tests
os.environ.setdefault("FAL_DISABLE_STARTUP_CHECK", "1")
# keep the persistent result cache out of the user's real cache dir during tests
os.environ.setdefault(
"COMFYUI_FAL_API_CACHE_DB",
@@ -67,3 +76,8 @@ def factory_mod():
@pytest.fixture(scope="session")
def errors_mod():
return _submodule("nodes.utils.errors")
@pytest.fixture(scope="session")
def media_mod():
return _submodule("nodes.utils.media")
+370
View File
@@ -0,0 +1,370 @@
{
"minimax/h3/image-to-video": {
"source": "https://fal.ai/api/openapi/queue/openapi.json?endpoint_id=minimax/h3/image-to-video",
"input": {
"title": "ImageToVideoHailuo03Input",
"properties": {
"seed": {
"title": "Seed",
"anyOf": [
{
"type": "integer"
},
{
"type": "null"
}
],
"description": "Random seed. A random seed is selected when omitted."
},
"sync_mode": {
"default": false,
"title": "Sync Mode",
"type": "boolean",
"description": "Return the generated video as base64 instead of a CDN URL."
},
"prompt_expansion_mode": {
"default": "balanced",
"examples": [
"disabled",
"fast",
"balanced",
"quality"
],
"title": "Prompt Expansion Mode",
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"description": "How much effort to spend rewriting the prompt before generation. 'disabled' skips prompt expansion. 'fast' returns in about a second. 'balanced' picks per request. 'quality' spends up to ~30s on a richer prompt."
},
"prompt": {
"maxLength": 50000,
"examples": [
"The camera slowly pulls back from the scene, revealing the full landscape as clouds drift overhead and light shifts across the terrain."
],
"minLength": 1,
"title": "Prompt",
"type": "string",
"description": "Text prompt for video generation"
},
"target_audio_url": {
"title": "Target Audio Url",
"anyOf": [
{
"minLength": 1,
"pattern": "\\S",
"type": "string"
},
{
"type": "null"
}
],
"description": "Optional URL of a 2-15 second audio clip (maximum 15 MB) to pin to the generated soundtrack. The original audio replaces the output soundtrack, trimmed or padded with silence to the video duration without changing playback speed. Accepts an HTTP(S) URL or a base64 data URI."
},
"duration": {
"minimum": 5,
"description": "The duration of the video in seconds.",
"title": "Duration",
"default": 5,
"type": "integer",
"maximum": 15
},
"end_image_url": {
"title": "End Image URL",
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"description": "Optional URL of the image to use as the last frame. It may be provided alone for end-only keyframe generation; in that case the output canvas follows this image."
},
"enable_safety_checker": {
"default": true,
"title": "Enable Safety Checker",
"type": "boolean",
"description": "If set to true, the safety checker will be enabled."
},
"image_url": {
"examples": [
"https://storage.googleapis.com/falserverless/example_inputs/hailuo23/pro_i2v_in.jpg"
],
"title": "Image URL",
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"description": "Optional URL of the image to use as the first frame. When provided, the output canvas follows this image. If only end_image_url is provided, the canvas follows that last frame instead. If both images are omitted, the request is handled as text-to-video (16:9 by default)."
},
"resolution": {
"default": "2K",
"type": "string",
"title": "Resolution",
"enum": [
"480P",
"768P",
"2K",
"4K"
],
"description": "The resolution of the generated video. 480P and 768P are native generation modes; 2K and 4K upscale a 768P base result."
}
},
"required": [
"prompt"
],
"type": "object",
"x-fal-order-properties": [
"prompt",
"duration",
"resolution",
"seed",
"enable_safety_checker",
"sync_mode",
"prompt_expansion_mode",
"target_audio_url",
"image_url",
"end_image_url"
]
}
},
"minimax/h3-max/image-to-video": {
"source": "https://fal.ai/api/openapi/queue/openapi.json?endpoint_id=minimax/h3-max/image-to-video",
"input": {
"properties": {
"sync_mode": {
"type": "boolean",
"title": "Sync Mode",
"default": false,
"description": "Return the generated video as base64 instead of a CDN URL."
},
"prompt_expansion_mode": {
"type": "string",
"title": "Prompt Expansion Mode",
"examples": [
"disabled",
"balanced",
"quality"
],
"description": "How much effort to spend rewriting the prompt before generation. 'disabled' skips prompt expansion. 'balanced' returns in about a second. 'quality' spends up to ~30s on a richer prompt.",
"default": "balanced"
},
"end_image_url": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"title": "End Image URL",
"description": "Optional URL of the image to use as the last frame. It may be provided alone for end-only keyframe generation; in that case the output canvas follows this image."
},
"duration": {
"type": "integer",
"maximum": 15,
"title": "Duration",
"default": 5,
"description": "The duration of the video in seconds.",
"minimum": 5
},
"seed": {
"anyOf": [
{
"type": "integer"
},
{
"type": "null"
}
],
"title": "Seed",
"description": "Random seed. A random seed is selected when omitted."
},
"prompt": {
"type": "string",
"title": "Prompt",
"examples": [
"The camera slowly pulls back from the scene, revealing the full landscape as clouds drift overhead and light shifts across the terrain."
],
"minLength": 1,
"maxLength": 50000,
"description": "Text prompt for video generation"
},
"resolution": {
"type": "string",
"title": "Resolution",
"enum": [
"480P",
"768P",
"1080P"
],
"description": "The native generation resolution, or 1080P latent refinement from a native 768P source.",
"default": "768P"
},
"image_url": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"examples": [
"https://storage.googleapis.com/falserverless/example_inputs/hailuo23/pro_i2v_in.jpg"
],
"description": "Optional URL of the image to use as the first frame. When provided, the output canvas follows this image. If only end_image_url is provided, the canvas follows that last frame instead. If both images are omitted, the request is handled as text-to-video (16:9 by default).",
"title": "Image URL"
},
"enable_safety_checker": {
"type": "boolean",
"title": "Enable Safety Checker",
"default": true,
"description": "If set to true, the safety checker will be enabled."
}
},
"type": "object",
"x-fal-order-properties": [
"prompt",
"duration",
"resolution",
"seed",
"enable_safety_checker",
"sync_mode",
"prompt_expansion_mode",
"image_url",
"end_image_url"
],
"title": "TurboImageToVideoHailuo03Input",
"required": [
"prompt",
"prompt_expansion_mode"
]
}
},
"minimax/h3-max-turbo/image-to-video": {
"source": "https://fal.ai/api/openapi/queue/openapi.json?endpoint_id=minimax/h3-max-turbo/image-to-video",
"input": {
"properties": {
"sync_mode": {
"type": "boolean",
"title": "Sync Mode",
"default": false,
"description": "Return the generated video as base64 instead of a CDN URL."
},
"prompt_expansion_mode": {
"type": "string",
"title": "Prompt Expansion Mode",
"examples": [
"disabled",
"balanced",
"quality"
],
"description": "How much effort to spend rewriting the prompt before generation. 'disabled' skips prompt expansion. 'balanced' returns in about a second. 'quality' spends up to ~30s on a richer prompt.",
"default": "balanced"
},
"end_image_url": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"title": "End Image URL",
"description": "Optional URL of the image to use as the last frame. It may be provided alone for end-only keyframe generation; in that case the output canvas follows this image."
},
"duration": {
"type": "integer",
"maximum": 15,
"title": "Duration",
"default": 5,
"description": "The duration of the video in seconds.",
"minimum": 5
},
"seed": {
"anyOf": [
{
"type": "integer"
},
{
"type": "null"
}
],
"title": "Seed",
"description": "Random seed. A random seed is selected when omitted."
},
"prompt": {
"type": "string",
"title": "Prompt",
"examples": [
"The camera slowly pulls back from the scene, revealing the full landscape as clouds drift overhead and light shifts across the terrain."
],
"minLength": 1,
"maxLength": 50000,
"description": "Text prompt for video generation"
},
"resolution": {
"type": "string",
"title": "Resolution",
"enum": [
"480P",
"768P",
"1080P"
],
"description": "The native generation resolution, or 1080P latent refinement from a native 768P source.",
"default": "768P"
},
"image_url": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"examples": [
"https://storage.googleapis.com/falserverless/example_inputs/hailuo23/pro_i2v_in.jpg"
],
"description": "Optional URL of the image to use as the first frame. When provided, the output canvas follows this image. If only end_image_url is provided, the canvas follows that last frame instead. If both images are omitted, the request is handled as text-to-video (16:9 by default).",
"title": "Image URL"
},
"enable_safety_checker": {
"type": "boolean",
"title": "Enable Safety Checker",
"default": true,
"description": "If set to true, the safety checker will be enabled."
}
},
"type": "object",
"x-fal-order-properties": [
"prompt",
"duration",
"resolution",
"seed",
"enable_safety_checker",
"sync_mode",
"prompt_expansion_mode",
"image_url",
"end_image_url"
],
"title": "TurboImageToVideoHailuo03Input",
"required": [
"prompt",
"prompt_expansion_mode"
]
}
}
}
+184
View File
@@ -0,0 +1,184 @@
"""Offline regressions for controls lost or mistyped during schema distillation."""
from __future__ import annotations
import copy
import json
from pathlib import Path
import pytest
from scripts.build_registry import (
build_record,
distill_inputs,
distill_property,
normalize_schema,
)
H3_INPUTS = json.loads((Path(__file__).parent / "fixtures" / "h3_inputs.json").read_text())
def _doc(schema):
return {"components": {"schemas": {"ModelInput": schema}}}
@pytest.mark.parametrize("endpoint_id", H3_INPUTS)
def test_h3_controls_from_upstream_schema(endpoint_id):
inputs = distill_inputs(H3_INPUTS[endpoint_id]["input"], {}, endpoint_id)
by_name = {inp["name"]: inp for inp in inputs}
assert by_name["duration"]["type"] == "integer"
assert (by_name["duration"]["min"], by_name["duration"]["max"]) == (5, 15)
assert by_name["duration"]["default"] == 5
assert by_name["prompt_expansion_mode"]["type"] == "string"
assert by_name["prompt_expansion_mode"]["default"] == "balanced"
required = H3_INPUTS[endpoint_id]["input"]["required"]
assert by_name["prompt_expansion_mode"]["required"] == ("prompt_expansion_mode" in required)
if "h3-max" in endpoint_id:
assert by_name["resolution"]["enum"] == ["480P", "768P", "1080P"]
assert by_name["prompt_expansion_mode"]["suggestions"] == ["disabled", "balanced", "quality"]
else:
assert by_name["resolution"]["enum"] == ["480P", "768P", "2K", "4K"]
assert by_name["prompt_expansion_mode"]["suggestions"] == ["disabled", "fast", "balanced", "quality"]
@pytest.mark.parametrize("name,examples", [
("mode", ["balanced", "quality"]),
("language", ["en", "tr", "ja"]),
("voice", ["Aria", "Rachel"]),
("model", ["vendor/model-a", "vendor/model-b"]),
])
def test_suggestions_are_generic_and_preserve_free_text(name, examples):
raw = {"type": "string", "examples": examples, "default": "new-value"}
original = copy.deepcopy(raw)
inp = distill_property(name, raw, set(), {})
assert inp["type"] == "string"
assert inp["enum"] is None
assert inp["suggestions"] == examples
assert inp["default"] == "new-value"
assert raw == original
@pytest.mark.parametrize("name,raw", [
("prompt", {"type": "string", "examples": ["cat", "dog"]}),
("custom_prompt", {"type": "string", "examples": ["cat", "dog"]}),
("description", {"type": "string", "examples": ["a cat", "a dog"]}),
("image_url", {"type": "string", "examples": ["https://example.com/a", "https://example.com/b"]}),
("contact", {"type": "string", "format": "email", "examples": ["a", "b"]}),
("mode", {"type": "string", "examples": ["only-one"]}),
("mode", {"type": "string", "enum": ["a", "b"], "examples": ["a", "c"]}),
])
def test_free_text_and_true_enums_do_not_get_suggestions(name, raw):
assert "suggestions" not in distill_property(name, raw, set(), {})
def test_large_schema_does_not_silently_drop_optional_controls():
properties = {f"field_{i}": {"type": "boolean"} for i in range(50)}
properties["duration"] = {"type": "integer", "minimum": 5, "maximum": 15}
inputs = distill_inputs({"properties": properties}, {}, "any/future-model")
assert len(inputs) == 51
assert inputs[-1]["name"] == "duration"
def test_nested_refs_and_nullable_composition_keep_controls():
components = {
"Alias": {"$ref": "#/components/schemas/Resolution"},
"Resolution": {"allOf": [{"type": "string", "enum": ["768P", "1080P"]}]},
}
inp = distill_property("resolution", {
"anyOf": [{"$ref": "#/components/schemas/Alias"}, {"type": "null"}],
"default": "1080P",
"description": "Output resolution",
}, set(), components)
assert inp["type"] == "enum"
assert inp["enum"] == ["768P", "1080P"]
assert inp["default"] == "1080P"
assert inp["description"] == "Output resolution"
@pytest.mark.parametrize("keyword", ["anyOf", "oneOf"])
def test_union_of_literals_preserves_all_choices(keyword):
inp = distill_property("resolution", {keyword: [
{"const": "480P"}, {"enum": ["768P", "1080P"]}, {"type": "null"},
]}, set(), {})
assert inp["enum"] == ["480P", "768P", "1080P"]
def test_literal_union_survives_nested_composition_and_repeated_normalization():
raw = {"allOf": [{"anyOf": [{"const": "480P"}, {"const": "1080P"}]}]}
normalized = normalize_schema(raw, {})[0]
assert normalized["enum"] == ["480P", "1080P"]
assert normalize_schema(normalized, {})[0] == normalized
@pytest.mark.parametrize("schema", [
{"allOf": [{"enum": ["480P", "768P"]}, {"enum": ["768P", "1080P"]}]},
{"anyOf": [{"const": "480P"}, {"const": "768P"}], "enum": ["768P"]},
])
def test_composed_enum_respects_intersecting_constraints(schema):
assert distill_property("resolution", schema, set(), {})["enum"] == ["768P"]
def test_all_of_input_objects_keep_inherited_fields_and_required_names():
doc = _doc({"allOf": [
{"$ref": "#/components/schemas/Base"},
{"properties": {"image_url": {"type": "string"}}, "required": ["image_url"]},
], "properties": {"duration": {"type": "integer", "default": 5, "minimum": 5, "maximum": 15}}})
doc["components"]["schemas"]["Base"] = {
"properties": {"prompt": {"type": "string"}}, "required": ["prompt"],
}
record = build_record({"id": "test/model"}, doc)
assert [inp["name"] for inp in record["inputs"]] == ["prompt", "image_url", "duration"]
assert [inp["required"] for inp in record["inputs"]] == [True, True, False]
def test_nullable_type_array_is_a_numeric_control():
inp = distill_property("duration", {"type": ["integer", "null"], "default": 5}, set(), {})
assert inp["type"] == "integer"
def test_reference_cycles_do_not_recurse_forever():
components = {"A": {"$ref": "#/components/schemas/B"}, "B": {"$ref": "#/components/schemas/A"}}
assert normalize_schema({"$ref": "#/components/schemas/A"}, components) == ({}, False, None)
def test_custom_image_size_still_has_preset_and_dimensions():
inp = distill_property("image_size", {"anyOf": [
{"type": "string", "enum": ["square", "landscape"]},
{"type": "object", "properties": {"width": {"type": "integer"}, "height": {"type": "integer"}}},
]}, set(), {})
assert inp["has_custom_size"] is True
assert inp["enum"] == ["square", "landscape", "custom_size"]
def test_nullable_nested_custom_size_preserves_dimension_controls():
inp = distill_property("image_size", {"anyOf": [
{"allOf": [{"anyOf": [
{"enum": ["square", "landscape"]},
{"type": "object", "properties": {"width": {}, "height": {}}},
]}]},
{"const": None},
]}, set(), {})
assert inp["has_custom_size"] is True
assert inp["enum"] == ["square", "landscape", "custom_size"]
def test_open_string_union_keeps_literal_suggestions_without_restricting_values():
inp = distill_property("voice", {"anyOf": [
{"enum": ["Aria", "Rachel"]}, {"type": "string"},
]}, set(), {})
assert inp["type"] == "string"
assert inp["enum"] is None
assert inp["suggestions"] == ["Aria", "Rachel"]
def test_exact_endpoint_path_wins_over_first_path_and_accepts_inline_schema():
def request(schema):
return {"post": {"requestBody": {"content": {"application/json": {"schema": schema}}}}}
doc = _doc({"properties": {"wrong_field": {"type": "string"}}})
doc["paths"] = {
"/test/other": request({"$ref": "#/components/schemas/ModelInput"}),
"/test/model": request({"properties": {"duration": {"type": "integer", "default": 5}}}),
}
record = build_record({"id": "test/model"}, doc)
assert [inp["name"] for inp in record["inputs"]] == ["duration"]
+45
View File
@@ -0,0 +1,45 @@
"""Check shared controls on shipped H3 generation nodes through the API boundary."""
from __future__ import annotations
import json
from pathlib import Path
import pytest
REGISTRY = Path(__file__).resolve().parents[1] / "data" / "fal_registry.json"
H3_SHARED_CONTROL_MODELS = [
model for model in json.loads(REGISTRY.read_text())["models"]
if model["endpoint_id"].startswith(("minimax/h3/", "minimax/h3-max/", "minimax/h3-max-turbo/"))
and not any(
specialty in model["endpoint_id"]
for specialty in ("/trainer", "/lip-sync/", "/styles/")
)
and not model.get("deprecated")
]
@pytest.mark.parametrize("model", H3_SHARED_CONTROL_MODELS, ids=lambda model: model["endpoint_id"])
def test_shipped_h3_shared_control_widgets_and_api_arguments(model, factory_mod, monkeypatch):
monkeypatch.setattr(factory_mod, "_ASYNC_CAPABLE", False)
captured = []
monkeypatch.setattr(factory_mod, "_call_api", lambda *args: captured.append(args) or {})
monkeypatch.setattr(factory_mod, "process_result", lambda *args: ())
cls = factory_mod.build_node_class(model)
inputs = cls.INPUT_TYPES()
widgets = {**inputs["required"], **inputs["optional"]}
assert widgets["duration"][0] == "INT"
assert widgets["duration"][1]["min"] == 5
assert widgets["duration"][1]["max"] == 15
assert widgets["duration"][1]["default"] == 5
bucket = "required" if "h3-max" in model["endpoint_id"] else "optional"
assert "prompt_expansion_mode" in inputs[bucket]
assert widgets["prompt_expansion_mode"][0] == "STRING"
assert "balanced" in widgets["prompt_expansion_mode"][1]["fal_suggestions"]
assert "quality" in widgets["prompt_expansion_mode"][1]["fal_suggestions"]
resolution = "1080P" if "h3-max" in model["endpoint_id"] else "2K"
assert resolution in widgets["resolution"][0]
cls().run(prompt="A camera pans", duration=15, resolution=resolution, prompt_expansion_mode="quality")
assert captured == [(model["endpoint_id"], {
"prompt": "A camera pans", "duration": 15, "resolution": resolution, "prompt_expansion_mode": "quality",
}, False)]
+3 -1
View File
@@ -33,7 +33,9 @@ def test_submit_and_collect_lifecycle(store):
assert pending[0]["request_id"] == "req-b"
def test_entries_newest_first(store):
def test_entries_newest_first(store, monkeypatch):
mod = importlib.import_module(f"{PKG}.nodes.utils.job_store")
monkeypatch.setattr(mod.time, "time", lambda: 1_700_000_000.0)
store.record_submit("fal-ai/a", "req-1")
store.record_submit("fal-ai/b", "req-2")
entries = store.entries()
+73
View File
@@ -0,0 +1,73 @@
"""Regression tests for malformed media URLs and legacy Seedance output."""
from __future__ import annotations
import sys
import pytest
def test_download_rejects_missing_schema_before_requests(monkeypatch, media_mod):
called = False
def unexpected_get(*_args, **_kwargs):
nonlocal called
called = True
raise AssertionError("requests.get must not receive a malformed URL")
monkeypatch.setattr(media_mod.requests, "get", unexpected_get)
with pytest.raises(media_mod.FalApiError, match=r"Expected an HTTP\(S\) media URL"):
media_mod.MediaUtils.download_url_to_temp("E", ".mp4")
assert called is False
def test_url_validator_normalizes_whitespace(media_mod):
assert (
media_mod.MediaUtils.require_http_url(" https://fal.media/video.mp4 ")
== "https://fal.media/video.mp4"
)
def test_seedance_pro_rejects_non_url_result(pack, monkeypatch):
node_cls = pack.NODE_CLASS_MAPPINGS["SeedanceProImageToVideo_fal"]
module = sys.modules[node_cls.__module__]
monkeypatch.setattr(
module.ImageUtils,
"upload_image",
staticmethod(lambda _image: "https://fal.media/input.png"),
)
monkeypatch.setattr(
module.ApiHandler,
"submit_multiple_and_get_results",
staticmethod(lambda *_args, **_kwargs: [{"video": {"url": "E"}}]),
)
with pytest.raises(module.FalApiError, match=r"Expected an HTTP\(S\) media URL"):
node_cls().generate_video("prompt", object(), "5")
def test_seedance_pro_returns_validated_url_list(pack, monkeypatch):
node_cls = pack.NODE_CLASS_MAPPINGS["SeedanceProImageToVideo_fal"]
module = sys.modules[node_cls.__module__]
monkeypatch.setattr(
module.ImageUtils,
"upload_image",
staticmethod(lambda _image: "https://fal.media/input.png"),
)
monkeypatch.setattr(
module.ApiHandler,
"submit_multiple_and_get_results",
staticmethod(
lambda *_args, **_kwargs: [
{"video": {"url": "https://fal.media/output.mp4"}}
]
),
)
assert node_cls().generate_video("prompt", object(), "5") == (
["https://fal.media/output.mp4"],
)
+13
View File
@@ -70,3 +70,16 @@ def test_png_chunk_roundtrip(platform, tmp_path):
endpoint, request_id, _ = node.read(file_path=str(png))
assert endpoint == "fal-ai/flux-2"
assert request_id == "req-png"
def test_remember_urls_covers_async_results(platform):
"""Provenance must work for Submit→Collect results, not just cached calls."""
cache_mod = importlib.import_module(f"{PKG}.nodes.utils.result_cache")
cache = cache_mod.ResultCache()
cache.clear()
url = "https://v3.fal.media/files/x/collected_output.jpg"
# no cache.put() — this simulates the async-collect path
cache.remember_urls("fal-ai/veo3", "req-async", {"images": [{"url": url}]})
hit = cache.find_request_by_url(url)
assert hit == {"endpoint_id": "fal-ai/veo3", "request_id": "req-async"}
cache.clear()
+21
View File
@@ -2,6 +2,7 @@
from __future__ import annotations
import importlib.util
import json
from pathlib import Path
@@ -20,6 +21,9 @@ def test_top_level_shape():
reg = _registry()
assert reg["version"] == 1
assert reg["model_count"] == len(reg["models"])
deprecated = sum(bool(model.get("deprecated")) for model in reg["models"])
assert reg["live_model_count"] == reg["model_count"] - deprecated
assert reg["deprecated_model_count"] == deprecated
assert reg["model_count"] > 500
@@ -76,3 +80,20 @@ def test_enum_defaults_are_members_or_custom_size():
f"{model['endpoint_id']}.{inp['name']}: default "
f"{inp['default']!r} not in enum"
)
def test_every_shipped_input_can_build_a_widget_without_comfyui():
"""The nightly refresh must render every control, without torch/ComfyUI.
The translator has only stdlib dependencies. Load it directly so the
scheduled registry checks don't need to import the whole node pack.
"""
path = REGISTRY.parents[1] / "nodes" / "dynamic" / "schema_to_inputs.py"
spec = importlib.util.spec_from_file_location("registry_widget_check", path)
translator = importlib.util.module_from_spec(spec)
spec.loader.exec_module(translator)
for model in _registry()["models"]:
input_types = translator.build_input_types(model)
widgets = {**input_types["required"], **input_types["optional"]}
for inp in model["inputs"]:
assert inp["name"] in widgets, f"{model['endpoint_id']}: missing {inp['name']} widget"
@@ -0,0 +1,49 @@
from __future__ import annotations
import importlib
from conftest import PKG, _load_package
from helpers import _model
def test_deprecated_endpoint_keeps_key_in_compatibility_tier():
_load_package()
loader = importlib.import_module(f"{PKG}.nodes.dynamic.registry_loader")
model = _model(
[],
endpoint_id="fal-ai/retired",
deprecated=True,
deprecated_reason="Endpoint retired.",
)
classes, display, skipped, flagged, deprecated = loader._build_model_mappings(
[model], set()
)
key = "FalAPI_fal-ai-retired"
assert key in classes
assert classes[key].CATEGORY == "FAL/Compatibility/text-to-image"
assert display[key].startswith("[Unavailable]")
assert "absent from the latest live fal catalog" in classes[key].DESCRIPTION
assert (skipped, flagged, deprecated) == (0, 0, 1)
def test_deprecated_endpoint_is_not_marked_as_superseded():
_load_package()
loader = importlib.import_module(f"{PKG}.nodes.dynamic.registry_loader")
models = [
_model(
[],
endpoint_id="fal-ai/old",
family="family",
published_at="2025-01-01T00:00:00Z",
deprecated=True,
),
_model(
[],
endpoint_id="fal-ai/new",
family="family",
published_at="2026-01-01T00:00:00Z",
),
]
assert loader._superseded_map(models) == {}
+39
View File
@@ -96,6 +96,14 @@ def test_multi_select_enum_is_comma_string(schema_to_inputs):
assert "vocals, drums, bass" in opts["tooltip"]
def test_numeric_multi_select_enum_keeps_api_value_types(schema_to_inputs, arguments_mod):
model = _model([_input("layers", "enum", enum=[1, 2, 4], default=[1, 4], is_list=True)])
typ, opts = schema_to_inputs.build_input_types(model)["optional"]["layers"]
assert typ == "STRING"
assert opts["default"] == "1, 4"
assert arguments_mod.build_arguments(model, {"layers": "1, 4"}) == {"layers": [1, 4]}
def test_json_field_is_multiline_string(schema_to_inputs):
model = _model([_input("loras", "json")])
it = schema_to_inputs.build_input_types(model)
@@ -115,3 +123,34 @@ def test_every_input_has_tooltip_when_description_given(schema_to_inputs):
model = _model([_input("prompt", "string", required=True, description="What to draw")])
it = schema_to_inputs.build_input_types(model)
assert it["required"]["prompt"][1]["tooltip"] == "What to draw"
def test_suggestions_keep_string_socket_and_default(schema_to_inputs, arguments_mod):
model = _model([_input("voice", "string", default="my-voice-id", suggestions=["Aria", "Rachel"], multiline=True)])
spec = schema_to_inputs.build_input_types(model)["optional"]["voice"]
assert spec[0] == "STRING"
assert spec[1]["default"] == "my-voice-id"
assert spec[1]["fal_suggestions"] == ["Aria", "Rachel"]
assert spec[1]["multiline"] is False
assert arguments_mod.build_arguments(model, {"voice": "new-custom-voice"}) == {"voice": "new-custom-voice"}
def test_all_catalog_suggestions_preserve_input_names_and_arbitrary_values(schema_to_inputs, arguments_mod):
import json
from pathlib import Path
registry = json.loads((Path(__file__).resolve().parents[1] / "data" / "fal_registry.json").read_text())
count = 0
for model in registry["models"]:
inputs = schema_to_inputs.build_input_types(model)
widgets = {**inputs["required"], **inputs["optional"]}
for inp in model["inputs"]:
if not inp.get("suggestions"):
continue
count += 1
assert widgets[inp["name"]][0] == "STRING", model["endpoint_id"]
assert widgets[inp["name"]][1]["fal_suggestions"] == inp["suggestions"]
assert arguments_mod.build_arguments(model, {inp["name"]: "arbitrary-future-value"}) == {
inp["name"]: "arbitrary-future-value",
}
assert count > 20
+127
View File
@@ -0,0 +1,127 @@
"""Focused tests for the curated Seedance 2.5 video editing node."""
from __future__ import annotations
import math
import sys
class _URLVideo:
def __init__(self, url: str):
self.url = url
def get_stream_source(self):
return self.url
def _node_and_module(pack):
node_cls = pack.NODE_CLASS_MAPPINGS["Seedance25VideoToVideo_fal"]
return node_cls, sys.modules[node_cls.__module__]
def test_seedance25_video_to_video_schema_and_registration(pack):
node_cls, _module = _node_and_module(pack)
inputs = node_cls.INPUT_TYPES()
assert inputs["required"]["video"][0] == "VIDEO"
assert inputs["required"]["prompt"][0] == "STRING"
assert set(inputs["optional"]) == {
"resolution",
"generate_audio",
"bitrate_mode",
"seed",
"force_rerun",
}
assert node_cls.RETURN_TYPES == ("VIDEO", "STRING")
assert node_cls.RETURN_NAMES == ("video", "video_url")
assert pack.NODE_DISPLAY_NAME_MAPPINGS["Seedance25VideoToVideo_fal"] == (
"Seedance 2.5 Video-to-Video (fal)"
)
def test_seedance25_edit_payload_outputs_and_force_rerun(pack, monkeypatch):
node_cls, module = _node_and_module(pack)
submitted = {}
native_output = object()
monkeypatch.setattr(
module.MediaUtils,
"upload_video",
staticmethod(lambda _video: "https://fal.media/input.mp4"),
)
monkeypatch.setattr(
module.MediaUtils,
"video_from_url",
staticmethod(lambda url: native_output if url.endswith("output.mp4") else None),
)
def submit(endpoint, arguments, *, skip_cache=False):
submitted.update(
endpoint=endpoint, arguments=arguments, skip_cache=skip_cache
)
return {"video": {"url": "https://fal.media/output.mp4"}}
monkeypatch.setattr(
module.ApiHandler, "submit_and_get_result", staticmethod(submit)
)
result = node_cls().edit_video(
object(),
"Turn the daytime scene into night",
resolution="1080p",
generate_audio=False,
bitrate_mode="high",
seed=123,
force_rerun=True,
)
assert submitted == {
"endpoint": "bytedance/seedance-2.5/reference-to-video",
"arguments": {
"prompt": "Turn the daytime scene into night",
"task": "editing",
"video_urls": ["https://fal.media/input.mp4"],
"resolution": "1080p",
"generate_audio": False,
"bitrate_mode": "high",
"seed": 123,
},
"skip_cache": True,
}
assert result == (native_output, "https://fal.media/output.mp4")
assert math.isnan(node_cls.IS_CHANGED(force_rerun=True))
def test_seedance25_url_backed_video_passes_through_without_upload(
pack, monkeypatch
):
node_cls, module = _node_and_module(pack)
source_url = "https://fal.media/already-uploaded.mp4"
captured = {}
def unexpected_upload(_value):
raise AssertionError("URL-backed VIDEO input must not be re-uploaded")
monkeypatch.setattr(
module.ImageUtils, "upload_file", staticmethod(unexpected_upload)
)
monkeypatch.setattr(
module.MediaUtils,
"video_from_url",
staticmethod(lambda _url: object()),
)
def submit(_endpoint, arguments, *, skip_cache=False):
captured.update(arguments)
assert skip_cache is False
return {"video": {"url": "https://fal.media/output.mp4"}}
monkeypatch.setattr(
module.ApiHandler, "submit_and_get_result", staticmethod(submit)
)
node_cls().edit_video(_URLVideo(source_url), "Restyle as watercolor")
assert captured["video_urls"] == [source_url]
assert captured["task"] == "editing"
assert "seed" not in captured
+35
View File
@@ -3,6 +3,8 @@
from __future__ import annotations
import importlib
from pathlib import Path
from types import SimpleNamespace
import pytest
from conftest import PKG, _load_package
@@ -48,3 +50,36 @@ def test_session_shape(routes):
def test_jobs_degrades_gracefully(routes):
payload = routes._jobs(limit=5)
assert "jobs" in payload and "counts" in payload
@pytest.mark.parametrize("failure_step", [None, "build_registry.py", "validate_registry.py"])
def test_refresh_promotes_only_validated_candidates(routes, monkeypatch, tmp_path, failure_step):
import subprocess
data = tmp_path / "data"
data.mkdir()
baseline = data / "fal_registry.json"
baseline.write_text("original registry")
monkeypatch.setattr(routes, "_repo_root", lambda: str(tmp_path))
steps = []
def run(command, **kwargs):
step = Path(command[1]).name
steps.append(step)
assert baseline.read_text() == "original registry"
assert command[-2:] == ["--preserve-from" if step == "build_registry.py" else "--baseline", str(baseline)]
if step == "build_registry.py":
Path(command[3]).write_text("validated candidate")
else:
assert Path(command[2]).read_text() == "validated candidate"
return SimpleNamespace(returncode=int(step == failure_step), stdout="", stderr="schema rejected")
monkeypatch.setattr(subprocess, "run", run)
ok, message = routes._run_refresh_subprocess()
assert ok == (failure_step is None)
assert baseline.read_text() == ("validated candidate" if ok else "original registry")
assert len(list(data.iterdir())) == 1
assert steps[0] == "build_registry.py"
if failure_step != "build_registry.py":
assert steps[1] == "validate_registry.py"
assert "Restart ComfyUI" in message if ok else "schema rejected" in message
+91
View File
@@ -0,0 +1,91 @@
// Exercise the real sidebar module without ComfyUI or third-party DOM libraries.
// Run: node --experimental-vm-modules --test tests/test_sidebar.mjs
import assert from "node:assert/strict";
import { readFile } from "node:fs/promises";
import { setImmediate } from "node:timers/promises";
import test from "node:test";
import vm from "node:vm";
class Element {
constructor(tag) {
this.tag = tag;
this.children = [];
this.listeners = {};
this.isConnected = true;
this.disabled = false;
this.textContent = "";
this.classList = { toggle() {} };
}
append(...children) { this.children.push(...children); }
replaceChildren(...children) { this.children = children; }
addEventListener(event, listener) { this.listeners[event] = listener; }
get lastChild() { return this.children.at(-1); }
find(className) {
if (this.className === className) return this;
return this.children.map((child) => child.find(className)).find(Boolean);
}
}
async function mount(status, { failPost = false, refreshOK = true } = {}) {
const calls = [];
const context = vm.createContext({
document: { createElement: (tag) => new Element(tag), hidden: false },
console: { debug() {} },
setInterval: () => 1,
clearInterval: () => {},
setTimeout: (callback) => { callback(); },
});
const api = new vm.SyntheticModule(["formatUsd", "getJson", "humanAge", "postJson", "shortEndpoint"], function () {
this.setExport("formatUsd", () => "$0");
this.setExport("humanAge", () => "now");
this.setExport("shortEndpoint", (value) => value);
this.setExport("getJson", async (path) => {
if (path === "/registry_status") {
if (status instanceof Error) throw status;
return status;
}
if (path === "/registry_refresh") return { running: false, finished_at: 1, ok: refreshOK, message: "Validation failed" };
return {};
});
this.setExport("postJson", async (path) => {
calls.push(path);
if (failPost) throw new Error("offline");
return { started: true, running: true };
});
}, { context });
const source = await readFile(new URL("../web/fal_sidebar.js", import.meta.url), "utf8");
const sidebar = new vm.SourceTextModule(source, { context });
await sidebar.link(() => api);
await sidebar.evaluate();
const root = new Element("div");
sidebar.namespace.mountPanel(root);
await setImmediate();
return { root, calls };
}
for (const status of [{ new_count: 0, new_models: [] }, new Error("catalog unavailable"), { new_count: 1, new_models: [{ title: "New model" }] }]) {
test(`refresh is available with status ${JSON.stringify(status)}`, async () => {
const { root, calls } = await mount(status);
const button = root.find("fal-registry-refresh");
assert.ok(button, "Existing models need schema refresh even with no new IDs or an unavailable catalog check");
button.listeners.click();
await setImmediate();
assert.deepEqual(calls, ["/registry_refresh"]);
assert.match(root.find("fal-registry-done").textContent, /updated controls/);
assert.equal(button.disabled, false);
});
}
for (const options of [{ failPost: true }, { refreshOK: false }]) {
test(`failed refresh can be retried: ${JSON.stringify(options)}`, async () => {
const { root, calls } = await mount({ new_count: 0 }, options);
const button = root.find("fal-registry-refresh");
button.listeners.click();
await setImmediate();
assert.ok(root.find("fal-registry-error"));
assert.equal(button.disabled, false);
button.listeners.click();
await setImmediate();
assert.equal(calls.length, 2);
});
}
+90
View File
@@ -0,0 +1,90 @@
import assert from "node:assert/strict";
import { readFile } from "node:fs/promises";
import test from "node:test";
import vm from "node:vm";
const source = await readFile(new URL("../web/fal_suggestions.js", import.meta.url), "utf8");
const module = new vm.SourceTextModule(source, { context: vm.createContext({ console }) });
await module.link(() => { throw new Error("Unexpected import"); });
await module.evaluate();
const { setupSuggestedWidgets } = module.namespace;
function setup(name = "mode", defaultValue = "balanced") {
const callbackValues = [];
const promptCalls = [];
class Node {
constructor() {
this.widgets = [
{ name: "prompt", type: "text", value: "test" },
{ name, type: "text", value: defaultValue, options: {}, callback: (value) => callbackValues.push(value) },
{ name: "seed", type: "number", value: 42 },
];
}
onNodeCreated() { this.originalCalled = true; return "original-result"; }
addWidget(type, name, value, callback, options) {
const widget = { type, name, value, callback, options };
this.widgets.push(widget);
return widget;
}
setDirtyCanvas() {}
}
const data = {
name: "FalAPI_any-future-model",
input: { optional: { [name]: ["STRING", { fal_suggestions: ["balanced", "quality"] }] } },
};
setupSuggestedWidgets(Node, data, { canvas: { prompt: (...args) => promptCalls.push(args) } });
const node = new Node();
assert.equal(node.onNodeCreated(), "original-result");
return { node, widget: node.widgets[1], promptCalls, callbackValues, data };
}
for (const name of ["mode", "voice", "language", "model_id"]) {
test(`generic suggested ${name} keeps widget order, default and STRING socket`, () => {
const { node, widget, data } = setup(name, "custom-default");
assert.equal(widget.type, "combo");
assert.equal(widget.value, "custom-default");
assert.deepEqual(node.widgets.map((w) => w.name), ["prompt", name, "seed"]);
assert.equal(node.widgets[2].value, 42);
assert.equal(data.input.optional[name][0], "STRING");
assert.equal(node.originalCalled, true);
});
}
test("saved values missing from examples survive load and serialization", () => {
const { node, widget } = setup();
widget.value = "older-workflow-value";
node.onConfigure({});
assert.ok(widget.options.values.includes("older-workflow-value"));
assert.equal(node.widgets.map((w) => w.value)[1], "older-workflow-value");
});
test("suggestion selection and custom entry preserve original callbacks", () => {
const { widget, promptCalls, callbackValues } = setup();
widget.value = "quality";
widget.callback("quality");
assert.deepEqual(callbackValues, ["quality"]);
const custom = widget.options.values.at(-1);
widget.value = custom;
widget.callback(custom);
assert.equal(widget.value, "quality", "UI-only label must never reach a queued API request");
assert.equal(promptCalls[0][1], "quality");
promptCalls[0][2]("new-mode-from-api");
assert.equal(widget.value, "new-mode-from-api");
assert.ok(widget.options.values.includes("new-mode-from-api"));
assert.deepEqual(callbackValues, ["quality", "new-mode-from-api"]);
});
test("canceling custom entry leaves the current value intact", () => {
const { widget, promptCalls } = setup();
const custom = widget.options.values.at(-1);
widget.value = custom;
widget.callback(custom);
promptCalls[0][2](null);
assert.equal(widget.value, "balanced");
});
test("nodes without suggestion metadata are unchanged", () => {
class Node {}
setupSuggestedWidgets(Node, { name: "FalAPI_plain-model", input: { required: { prompt: ["STRING", {}] } } }, {});
assert.equal(Node.prototype.onNodeCreated, undefined);
});
+120
View File
@@ -0,0 +1,120 @@
"""Tests for the FAL/Utils node layer (dataset, image, data, video basics)."""
from __future__ import annotations
import importlib
import json
import zipfile
import pytest
import torch
from conftest import PKG, _load_package
@pytest.fixture(scope="session")
def archive_mod():
_load_package()
return importlib.import_module(f"{PKG}.nodes.utils.archive")
@pytest.fixture(scope="session")
def image_nodes(pack):
return pack.NODE_CLASS_MAPPINGS
def test_zip_images_with_captions(archive_mod, tmp_path):
images = torch.rand(3, 8, 8, 3)
zip_path = archive_mod.ArchiveUtils.zip_images(images, captions=["a", "", "c"])
try:
with zipfile.ZipFile(zip_path) as zf:
names = sorted(zf.namelist())
assert "image_0.png" in names and "image_2.txt" in names
assert zf.read("image_0.txt").decode() == "a"
finally:
import os
os.unlink(zip_path)
def test_zip_images_caption_mismatch_raises(archive_mod, errors_mod):
with pytest.raises(errors_mod.FalApiError):
archive_mod.ArchiveUtils.zip_images(torch.rand(2, 8, 8, 3), captions=["only one"])
def test_json_extract(pack):
cls = pack.NODE_CLASS_MAPPINGS["FalJSONExtract_fal"]
node = cls()
fn = getattr(node, cls.FUNCTION)
payload = json.dumps({"video": {"url": "https://x/v.mp4"}, "images": [{"url": "https://x/i.png"}], "seed": 42})
assert fn(json_text=payload, path="video.url", default="")[0] == "https://x/v.mp4"
assert fn(json_text=payload, path="images[0].url", default="")[0] == "https://x/i.png"
assert fn(json_text=payload, path="seed", default="")[1] == 42.0
assert fn(json_text=payload, path="missing.path", default="fallback")[0] == "fallback"
def test_prompt_lines_wraps(pack):
cls = pack.NODE_CLASS_MAPPINGS["FalPromptLines_fal"]
node = cls()
fn = getattr(node, cls.FUNCTION)
text = "one\ntwo\nthree"
assert fn(text=text, index=0, skip_blank=True)[0] == "one"
assert fn(text=text, index=4, skip_blank=True)[0] == "two" # wraps modulo 3
def test_resize_to_preset_dims(pack):
cls = pack.NODE_CLASS_MAPPINGS["FalResizeToPreset_fal"]
node = cls()
fn = getattr(node, cls.FUNCTION)
image = torch.rand(1, 300, 500, 3)
out, width, height = fn(image=image, preset="landscape_16_9", width=1024, height=1024, mode="cover_crop")
assert (width, height) == (1024, 576)
assert tuple(out.shape) == (1, 576, 1024, 3)
def test_base64_round_trip(pack):
cm = pack.NODE_CLASS_MAPPINGS
enc_cls, dec_cls = cm["FalImageToBase64_fal"], cm["FalBase64ToImage_fal"]
image = torch.rand(1, 16, 16, 3)
encoded = getattr(enc_cls(), enc_cls.FUNCTION)(image=image, format="png", data_uri=True)[0]
decoded = getattr(dec_cls(), dec_cls.FUNCTION)(data=encoded)[0]
assert tuple(decoded.shape) == (1, 16, 16, 3)
assert torch.allclose(image, decoded, atol=2 / 255)
def test_image_grid_shape(pack):
cls = pack.NODE_CLASS_MAPPINGS["FalImageGrid_fal"]
node = cls()
fn = getattr(node, cls.FUNCTION)
out = fn(images=torch.rand(4, 32, 32, 3), labels="a\nb\nc\nd", columns=2, cell_padding=4, label_height=16)[0]
assert out.ndim == 4 and out.shape[0] == 1 and out.shape[3] == 3
def test_extract_frames_from_real_video(pack, tmp_path):
cv2 = pytest.importorskip("cv2")
import numpy as np
path = str(tmp_path / "clip.mp4")
writer = cv2.VideoWriter(path, cv2.VideoWriter_fourcc(*"mp4v"), 8, (32, 32))
for i in range(16):
frame = np.full((32, 32, 3), 255 if i == 15 else 0, dtype=np.uint8)
writer.write(frame)
writer.release()
cls = pack.NODE_CLASS_MAPPINGS["FalExtractFrames_fal"]
node = cls()
fn = getattr(node, cls.FUNCTION)
frames, count = fn(video=path, mode="last", n=1, max_frames=64)
assert count == 16
assert frames.shape[0] == 1
assert frames.mean().item() > 0.9 # last frame is white
def test_all_util_nodes_have_tooltips(pack):
util_keys = [k for k, c in pack.NODE_CLASS_MAPPINGS.items() if c.CATEGORY.startswith("FAL/Utils")]
assert len(util_keys) == 28 # 20 utility nodes + 8 typed builders
for key in util_keys:
input_types = pack.NODE_CLASS_MAPPINGS[key].INPUT_TYPES()
for bucket in ("required", "optional"):
for name, spec in input_types.get(bucket, {}).items():
if len(spec) > 1 and isinstance(spec[1], dict):
assert "tooltip" in spec[1], f"{key}.{name} missing tooltip"
+220
View File
@@ -0,0 +1,220 @@
from __future__ import annotations
import copy
import pytest
from scripts.build_readme import render_generated_section
from scripts.build_registry import preserve_missing_records
from scripts.validate_registry import (
RegistryValidationError,
compare_model_inputs,
compare_registries,
validate_registry,
)
def _model(endpoint_id: str) -> dict:
return {
"endpoint_id": endpoint_id,
"title": endpoint_id,
"category": "text-to-image",
"description": "",
"family": "",
"lab": "",
"pricing": "",
"published_at": "",
"inputs": [],
"output_kind": "images",
"output_props": [],
"thumbnail": "",
}
def _registry(*endpoint_ids: str) -> dict:
models = [_model(endpoint_id) for endpoint_id in endpoint_ids]
return {
"version": 1,
"model_count": len(models),
"live_model_count": len(models),
"deprecated_model_count": 0,
"models": models,
}
def test_valid_registry_returns_endpoint_ids():
assert validate_registry(_registry("fal-ai/a", "fal-ai/b"), min_models=2) == {
"fal-ai/a",
"fal-ai/b",
}
@pytest.mark.parametrize(
"registry, message",
[
({"version": 1, "models": []}, "top-level"),
(
{
"version": 1,
"model_count": 2,
"live_model_count": 1,
"deprecated_model_count": 0,
"models": [_model("fal-ai/a")],
},
"does not match",
),
(_registry("fal-ai/b", "fal-ai/a"), "sorted"),
(_registry("fal-ai/a", "fal-ai/a"), "Duplicate"),
],
)
def test_invalid_registry_is_rejected(registry, message):
with pytest.raises(RegistryValidationError, match=message):
validate_registry(registry, min_models=1)
def test_unknown_output_kind_is_rejected():
registry = _registry("fal-ai/a")
registry["models"][0]["output_kind"] = "binary"
with pytest.raises(RegistryValidationError, match="unknown output_kind"):
validate_registry(registry, min_models=1)
def test_invalid_endpoint_id_is_rejected():
with pytest.raises(RegistryValidationError, match="invalid endpoint_id"):
validate_registry(_registry("fal-ai/bad endpoint"), min_models=1)
def test_deprecated_counts_must_match_records():
registry = _registry("fal-ai/a")
registry["models"][0]["deprecated"] = True
with pytest.raises(RegistryValidationError, match="live_model_count"):
validate_registry(registry, min_models=1)
def test_large_removal_is_blocked_by_default():
baseline = {f"fal-ai/{index}" for index in range(100)}
candidate = {f"fal-ai/{index}" for index in range(90)}
with pytest.raises(RegistryValidationError, match="above the 5.0% limit"):
compare_registries(baseline, candidate)
def test_large_removal_can_be_explicitly_allowed():
baseline = {f"fal-ai/{index}" for index in range(100)}
candidate = {f"fal-ai/{index}" for index in range(90)}
added, removed = compare_registries(baseline, candidate, allow_large_change=True)
assert added == set()
assert len(removed) == 10
def test_large_addition_is_blocked_by_default():
baseline = {f"fal-ai/{index}" for index in range(100)}
candidate = baseline | {f"new/{index}" for index in range(30)}
with pytest.raises(RegistryValidationError, match="above the 25.0% limit"):
compare_registries(baseline, candidate)
def test_missing_baseline_record_is_preserved_as_deprecated():
current = [_model("fal-ai/current")]
baseline = {"models": [_model("fal-ai/current"), _model("fal-ai/retired")]}
merged = preserve_missing_records(current, baseline)
by_id = {model["endpoint_id"]: model for model in merged}
assert set(by_id) == {"fal-ai/current", "fal-ai/retired"}
assert "deprecated" not in by_id["fal-ai/current"]
assert by_id["fal-ai/retired"]["deprecated"] is True
def test_generated_catalog_lists_live_models_only():
live = _model("fal-ai/live")
retired = {**_model("fal-ai/retired"), "deprecated": True}
registry = {
"models": [live, retired],
"model_count": 2,
"live_model_count": 1,
"deprecated_model_count": 1,
}
rendered = render_generated_section(registry)
assert "1 live models" in rendered
assert "1 compatibility-preserved" in rendered
assert "fal-ai/live" in rendered
assert "fal-ai/retired" not in rendered
def _input_registry():
registry = _registry("fal-ai/video")
registry["models"][0]["inputs"] = [
{"name": "duration", "type": "integer", "default": 5, "required": False},
{"name": "resolution", "type": "enum", "enum": ["768P", "1080P"], "required": False},
]
return registry
@pytest.mark.parametrize("change, message", [
(lambda inputs: inputs.pop(0), "removed input duration"),
(lambda inputs: inputs[1].update(type="string"), "lost enum control"),
(lambda inputs: inputs[1].update(enum=["768P"]), "removed choices"),
(lambda inputs: inputs[0].update(type="json"), "lost integer control"),
])
def test_refresh_detects_disappearing_controls(change, message):
baseline = _input_registry()
candidate = copy.deepcopy(baseline)
change(candidate["models"][0]["inputs"])
assert message in compare_model_inputs(baseline, candidate)[0]
def test_refresh_allows_new_inputs_choices_and_string_to_dropdown():
baseline = _input_registry()
baseline["models"][0]["inputs"].append({"name": "mode", "type": "string", "required": False})
candidate = copy.deepcopy(baseline)
inputs = candidate["models"][0]["inputs"]
inputs.append({"name": "seed", "type": "integer", "required": False})
inputs[1]["enum"].append("4K")
inputs[2].update(type="enum", enum=["balanced", "quality"])
assert compare_model_inputs(baseline, candidate) == []
@pytest.mark.parametrize("change, message", [
(lambda inputs: inputs.append(dict(inputs[0])), "duplicate input"),
(lambda inputs: inputs[0].update(type="unknown"), "invalid input type"),
(lambda inputs: inputs[1].update(enum=[]), "no enum choices"),
(lambda inputs: inputs[0].update(required="false"), "required must be a boolean"),
(lambda inputs: inputs.append(None), "invalid input record"),
])
def test_refresh_rejects_malformed_controls(change, message):
registry = _input_registry()
change(registry["models"][0]["inputs"])
with pytest.raises(RegistryValidationError, match=message):
validate_registry(registry, min_models=1)
def test_refresh_rejects_lost_suggested_choices():
baseline = _input_registry()
baseline["models"][0]["inputs"].append({
"name": "mode", "type": "string", "required": False, "suggestions": ["fast", "quality"],
})
candidate = copy.deepcopy(baseline)
del candidate["models"][0]["inputs"][-1]["suggestions"]
assert "lost suggested choices" in compare_model_inputs(baseline, candidate)[0]
@pytest.mark.parametrize("allow", [False, True])
def test_cli_requires_separate_explicit_override_for_input_removal(monkeypatch, tmp_path, allow):
import json
import sys
from scripts.validate_registry import main
baseline = _input_registry()
candidate = copy.deepcopy(baseline)
candidate["models"][0]["inputs"].pop(0)
for name, registry in [("baseline", baseline), ("candidate", candidate)]:
(tmp_path / f"{name}.json").write_text(json.dumps(registry))
argv = ["validate_registry.py", str(tmp_path / "candidate.json"), "--baseline",
str(tmp_path / "baseline.json"), "--min-models", "1", "--allow-large-change"]
if allow:
argv.append("--allow-input-removal")
monkeypatch.setattr(sys, "argv", argv)
if allow:
main()
else:
with pytest.raises(RegistryValidationError, match="removes existing controls"):
main()
+28 -2
View File
@@ -39,22 +39,48 @@ function findTarget(canvas, value) {
return null;
}
function resultThumb(model) {
if (!model?.thumbnail || typeof model.thumbnail !== "string") return null;
try {
const img = document.createElement("img");
img.className = "fal-suggest-thumb";
img.src = model.thumbnail;
img.loading = "lazy";
img.decoding = "async";
img.alt = "";
img.addEventListener("error", () => {
img.style.display = "none";
});
return img;
} catch (error) {
console.debug("[fal] suggestion thumbnail failed", error);
return null;
}
}
function resultRow(model, apply) {
const row = document.createElement("div");
row.className = "fal-suggest-item";
const thumb = resultThumb(model);
if (thumb) row.append(thumb);
const text = document.createElement("div");
text.className = "fal-suggest-text";
const title = document.createElement("span");
title.className = "fal-suggest-title";
title.textContent = model.title || model.endpoint_id;
const endpoint = document.createElement("span");
endpoint.className = "fal-suggest-endpoint";
endpoint.textContent = model.endpoint_id;
row.append(title, endpoint);
text.append(title, endpoint);
if (model.label) {
const price = document.createElement("span");
price.className = "fal-suggest-price";
price.textContent = model.label;
row.append(price);
text.append(price);
}
row.append(text);
row.addEventListener("mousedown", (event) => {
event.preventDefault();
event.stopPropagation();
+80 -2
View File
@@ -165,12 +165,29 @@
.fal-suggest-item {
display: flex;
flex-direction: column;
gap: 1px;
flex-direction: row;
align-items: center;
gap: 8px;
padding: 6px 10px;
cursor: pointer;
}
.fal-suggest-thumb {
flex: none;
width: 48px;
height: 48px;
object-fit: cover;
border-radius: 6px;
background: rgba(255, 255, 255, 0.05);
}
.fal-suggest-text {
display: flex;
flex-direction: column;
gap: 1px;
min-width: 0;
}
.fal-suggest-item:hover {
background: rgba(167, 139, 250, 0.15);
}
@@ -188,3 +205,64 @@
font-size: 10px;
color: #c4b5fd;
}
/* Registry freshness section in the sidebar panel. */
.fal-registry {
display: flex;
flex-direction: column;
gap: 4px;
}
.fal-registry-news {
display: flex;
flex-direction: column;
gap: 4px;
padding: 8px 10px;
border: 1px solid rgba(167, 139, 250, 0.25);
border-radius: 8px;
background: rgba(255, 255, 255, 0.04);
}
.fal-registry-count {
font-weight: 600;
color: #c4b5fd;
}
.fal-registry-model {
font-size: 11px;
opacity: 0.8;
overflow: hidden;
text-overflow: ellipsis;
white-space: nowrap;
}
.fal-registry-refresh {
margin-top: 4px;
padding: 4px 10px;
border: 1px solid rgba(167, 139, 250, 0.5);
border-radius: 6px;
background: transparent;
color: #ece9fd;
font-size: 11px;
cursor: pointer;
}
.fal-registry-refresh:hover:not(:disabled) {
background: rgba(167, 139, 250, 0.15);
}
.fal-registry-refresh:disabled {
opacity: 0.55;
cursor: default;
}
.fal-registry-done {
font-size: 11px;
color: #86efac;
}
.fal-registry-error {
font-size: 11px;
color: #fca5a5;
}
+6
View File
@@ -5,6 +5,7 @@ import { app } from "../../scripts/app.js";
import { loadPricingMap, setupNodeBadges } from "./fal_badges.js";
import { registerSidebar } from "./fal_sidebar.js";
import { installAutocomplete } from "./fal_autocomplete.js";
import { setupSuggestedWidgets } from "./fal_suggestions.js";
// Start loading the pricing map immediately: node definitions register before
// setup() runs, and the badge drawer looks the map up lazily at draw time.
@@ -30,6 +31,11 @@ app.registerExtension({
} catch (error) {
console.debug("[fal] badge setup failed", error);
}
try {
setupSuggestedWidgets(nodeType, nodeData, app);
} catch (error) {
console.debug("[fal] suggested control setup failed", error);
}
},
async setup() {
+119 -1
View File
@@ -4,6 +4,9 @@ import { formatUsd, getJson, humanAge, postJson, shortEndpoint } from "./fal_api
const REFRESH_MS = 3000;
const JOB_LIMIT = 50;
const REGISTRY_TITLE_LIMIT = 5;
const REGISTRY_POLL_MS = 3000;
const REGISTRY_POLL_MAX = 600;
let refreshTimer = null;
let panelRoot = null;
@@ -87,16 +90,127 @@ function buildPanel() {
const jobsHeader = element("div", "fal-jobs-header", "Jobs");
const jobs = element("div", "fal-jobs");
root.append(stats, jobsHeader, jobs);
const registryHeader = element("div", "fal-jobs-header", "Registry");
const registry = element("div", "fal-registry");
registry.append(element("div", "fal-muted", "checking for new models…"));
root.append(stats, jobsHeader, jobs, registryHeader, registry);
return {
root,
sessionValue: session.lastChild,
balanceValue: balance.lastChild,
jobsHeader,
jobs,
registryHeader,
registry,
};
}
// -- Registry freshness section -------------------------------------------------
function registryDone(view, ok, message) {
const note = element(
"div",
ok ? "fal-registry-done" : "fal-registry-error",
ok ? "done — restart ComfyUI and reload this page to load updated controls and new nodes" : message || "refresh failed"
);
view.registry.append(note);
}
async function pollRefresh(view, button) {
for (let attempt = 0; attempt < REGISTRY_POLL_MAX; attempt += 1) {
await new Promise((resolve) => setTimeout(resolve, REGISTRY_POLL_MS));
if (!view.registry.isConnected) return;
let status = null;
try {
status = await getJson("/registry_refresh");
} catch (error) {
console.debug("[fal] registry refresh poll failed", error);
continue;
}
if (status && status.running === false && status.finished_at) {
registryDone(view, status.ok === true, status.message);
if (button) {
button.disabled = false;
button.textContent = "Refresh registry";
}
return;
}
}
if (button) button.textContent = "Still running \u2014 check back later";
}
async function startRegistryRefresh(view, button) {
try {
button.disabled = true;
button.textContent = "Refreshing…";
const result = await postJson("/registry_refresh", {});
if (!result?.started && result?.running !== true) {
registryDone(view, false, result?.message || "could not start refresh");
button.disabled = false;
button.textContent = "Refresh registry";
return;
}
await pollRefresh(view, button);
} catch (error) {
console.debug("[fal] registry refresh failed", error);
registryDone(view, false, "refresh request failed");
button.disabled = false;
button.textContent = "Refresh registry";
}
}
function renderRegistry(view, status) {
try {
const count = Number(status?.new_count) || 0;
const box = element("div", "fal-registry-news");
box.append(
element("div", "fal-registry-count", status == null
? "Registry status unavailable."
: count > 0 ? `${count} new model${count === 1 ? "" : "s"} on fal` : "No new model IDs found.")
);
box.append(
element("div", "fal-muted", "Refresh to fetch updated controls for existing models too. Restart ComfyUI and reload this page afterward.")
);
const models = Array.isArray(status?.new_models) ? status.new_models : [];
for (const model of models.slice(0, REGISTRY_TITLE_LIMIT)) {
const title = model?.title || model?.endpoint_id || "";
if (!title) continue;
const row = element("div", "fal-registry-model", title);
if (model?.endpoint_id) row.title = model.endpoint_id;
box.append(row);
}
if (count > REGISTRY_TITLE_LIMIT) {
box.append(element("div", "fal-muted", `…and ${count - REGISTRY_TITLE_LIMIT} more`));
}
const button = element("button", "fal-registry-refresh", "Refresh registry");
button.addEventListener("click", () => {
startRegistryRefresh(view, button).catch((error) =>
console.debug("[fal] registry refresh flow failed", error)
);
});
box.append(button);
view.registry.replaceChildren(box);
} catch (error) {
console.debug("[fal] registry render failed", error);
}
}
async function loadRegistrySection(view) {
try {
const status = await getJson("/registry_status");
renderRegistry(view, status);
} catch (error) {
console.debug("[fal] registry status failed", error);
try {
renderRegistry(view, null);
} catch (renderError) {
console.debug("[fal] registry fallback render failed", renderError);
}
}
}
function renderSession(target, data) {
const total = formatUsd(data?.total_usd) ?? "$0";
const calls = data?.calls ?? 0;
@@ -156,6 +270,10 @@ export function mountPanel(container) {
panelRoot = view.root;
container.replaceChildren(view.root);
startRefreshLoop(view);
// Fetched once per panel open (server-side result is cached for an hour).
loadRegistrySection(view).catch((error) =>
console.debug("[fal] registry section load failed", error)
);
}
function mountFloatingFallback() {
+87
View File
@@ -0,0 +1,87 @@
// Suggested string values are an editable dropdown, not an API enum.
// Keep the original input name, STRING socket and serialized widget position.
const CUSTOM = "Enter custom value…";
function replaceStringWidget(node, name, inputOptions, app) {
const suggestions = inputOptions.fal_suggestions;
const index = (node.widgets || []).findIndex((widget) => widget.name === name);
if (index < 0 || !suggestions.length) return;
const original = node.widgets[index];
const originalSize = node.size ? [...node.size] : null;
const values = [...new Set(suggestions)];
let lastValue = original.value ?? "";
let widget;
const choices = (value) => [...new Set([...values, value]), CUSTOM];
const sync = () => {
if (widget.value !== CUSTOM) lastValue = widget.value ?? "";
widget.options.values = choices(lastValue);
};
const apply = (value, ...args) => {
if (value == null) return; // Cancel leaves the previous value intact.
lastValue = String(value);
widget.value = lastValue;
sync();
original.callback?.call(widget, lastValue, ...args);
node.setDirtyCanvas?.(true, true);
};
const onSelect = (value, ...args) => {
if (value !== CUSTOM) {
apply(value, ...args);
return;
}
// Never serialize or submit the UI-only custom-entry label.
widget.value = lastValue;
if (typeof app?.canvas?.prompt === "function") {
app.canvas.prompt(`Custom ${name}`, lastValue, (text) => apply(text, ...args), args.at(-1));
} else {
apply(globalThis.prompt?.(`Custom ${name}`, lastValue), ...args);
}
};
const tooltip = `${inputOptions.tooltip || original.tooltip || ""} Suggested values; choose '${CUSTOM}' to enter any other value.`.trim();
widget = node.addWidget("combo", name, lastValue, onSelect, {
...original.options,
values: choices(lastValue),
tooltip,
});
widget.tooltip = tooltip;
widget._falSyncSuggestions = sync;
widget.value = lastValue;
// addWidget appends. Move it into the old slot so positional workflow values
// and every following widget keep their existing meaning.
const appendedIndex = node.widgets.indexOf(widget);
node.widgets.splice(appendedIndex, 1);
node.widgets.splice(index, 1, widget);
if (original.label != null) widget.label = original.label;
if (original.serializeValue) widget.serializeValue = original.serializeValue.bind(widget);
original.onRemove?.();
if (originalSize) node.setSize?.(originalSize);
}
export function setupSuggestedWidgets(nodeType, nodeData, app) {
if (!nodeData?.name?.startsWith("FalAPI_")) return;
const fields = Object.entries({ ...nodeData.input?.required, ...nodeData.input?.optional })
.filter(([, spec]) => spec?.[0] === "STRING" && Array.isArray(spec?.[1]?.fal_suggestions))
.map(([name, spec]) => [name, spec[1]]);
if (!fields.length) return;
const originalCreated = nodeType.prototype.onNodeCreated;
nodeType.prototype.onNodeCreated = function (...args) {
const result = originalCreated?.apply(this, args);
for (const [name, inputOptions] of fields) {
try {
replaceStringWidget(this, name, inputOptions, app);
} catch (error) {
console.debug(`[fal] suggestion widget failed for ${name}`, error);
}
}
return result;
};
const originalConfigure = nodeType.prototype.onConfigure;
nodeType.prototype.onConfigure = function (...args) {
const result = originalConfigure?.apply(this, args);
// Configuration loads the old positional values after node creation.
for (const widget of this.widgets || []) widget._falSyncSuggestions?.();
return result;
};
}