Compare commits
122
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
ef7c121415 | ||
|
|
28a045b72c | ||
|
|
4a6c7b934e | ||
|
|
be596c6831 | ||
|
|
991ceedfc6 | ||
|
|
2cff2ef7dd | ||
|
|
7414b7c2bd | ||
|
|
d71da89c09 | ||
|
|
432dfd19a0 | ||
|
|
647a38f265 | ||
|
|
82934d7764 | ||
|
|
7f8af1d013 | ||
|
|
68eba7e44c | ||
|
|
bb00c3cd94 | ||
|
|
e81ac14244 | ||
|
|
d481082a1a | ||
|
|
6747ca2e7a | ||
|
|
356f1159ac | ||
|
|
381c0012e4 | ||
|
|
20bc1e301a | ||
|
|
6870e00ef7 | ||
|
|
39fc5ae8df | ||
|
|
e85a103192 | ||
|
|
c7872ac509 | ||
|
|
70cdfe8877 | ||
|
|
27afeac7f6 | ||
|
|
9cb787423e | ||
|
|
a64fd856ae | ||
|
|
f5959e25ec | ||
|
|
58d3cd1d48 | ||
|
|
1780dde1bc | ||
|
|
bf9be6fb19 | ||
|
|
546a6f4923 | ||
|
|
c7adc92c36 | ||
|
|
46a029acbc | ||
|
|
e8b175d202 | ||
|
|
e8be1b0521 | ||
|
|
c6248be9ec | ||
|
|
f8e89865e9 | ||
|
|
48b823f3ce | ||
|
|
a408dc63f4 | ||
|
|
1f95a7b274 | ||
|
|
a1524a7b94 | ||
|
|
5c87442b98 | ||
|
|
e11eb0fc9f | ||
|
|
7c34feb913 | ||
|
|
8747372fc3 | ||
|
|
9aea820995 | ||
|
|
486a4e87be | ||
|
|
4793af4269 | ||
|
|
dd08b8f490 | ||
|
|
b5228cd7a3 | ||
|
|
3478187364 | ||
|
|
33f116853b | ||
|
|
87d59b969f | ||
|
|
4e42fb7ba1 | ||
|
|
a22efa1bf0 | ||
|
|
f20a036e24 | ||
|
|
950b500f1b | ||
|
|
3185a36053 | ||
|
|
b12aa12e45 | ||
|
|
6f95834a13 | ||
|
|
8a47f0598b | ||
|
|
31a2c0a35e | ||
|
|
ca4251efbe | ||
|
|
d6597fb81e | ||
|
|
75a4ce6324 | ||
|
|
b75f731abf | ||
|
|
1b14ab3164 | ||
|
|
a728d7e3ba | ||
|
|
648d4b5ab2 | ||
|
|
809cf424b4 | ||
|
|
56ac5e9613 | ||
|
|
492a963bd4 | ||
|
|
57b78dcca3 | ||
|
|
aa6e5b9531 | ||
|
|
f6a650d407 | ||
|
|
54b3182c6a | ||
|
|
830a467f54 | ||
|
|
1fb220258f | ||
|
|
5382b69e64 | ||
|
|
95b8a044ec | ||
|
|
6273ea0fb2 | ||
|
|
f65213ceb4 | ||
|
|
1e60cc4a0b | ||
|
|
7cd2900150 | ||
|
|
4200bbcedf | ||
|
|
fedda31284 | ||
|
|
105f6a9083 | ||
|
|
f65b8ea0fa | ||
|
|
3f27dd7887 | ||
|
|
b95ff2c86e | ||
|
|
04f19b26c2 | ||
|
|
13a05b8d6a | ||
|
|
a4f22a114b | ||
|
|
be7f74ebee | ||
|
|
331ed2c058 | ||
|
|
97049f29c8 | ||
|
|
9d8c754e8a | ||
|
|
f9b21a5e93 | ||
|
|
fbee93b5b5 | ||
|
|
845b9d46c5 | ||
|
|
31572e6e45 | ||
|
|
b60c18d8a8 | ||
|
|
58c54acbce | ||
|
|
27580456ed | ||
|
|
a68c56134c | ||
|
|
cd9eb99568 | ||
|
|
ef774a511b | ||
|
|
66d4dcf54d | ||
|
|
a6d061c0eb | ||
|
|
34d3a8396e | ||
|
|
06a30a6f21 | ||
|
|
f4f486edb0 | ||
|
|
1f6f476679 | ||
|
|
1e561ac944 | ||
|
|
cf523888a7 | ||
|
|
93aa2cbc04 | ||
|
|
5be02175f3 | ||
|
|
4ff17aa6ef | ||
|
|
a6d29a2d4c | ||
|
|
4215edebf0 |
@@ -0,0 +1,74 @@
|
||||
name: CI
|
||||
|
||||
on:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
pull_request:
|
||||
branches:
|
||||
- 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@v7
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v7
|
||||
with:
|
||||
python-version: "3.11"
|
||||
- name: Install ruff
|
||||
run: pip install ruff
|
||||
- name: Run ruff
|
||||
run: ruff check .
|
||||
|
||||
test:
|
||||
name: Test (python ${{ matrix.python-version }})
|
||||
runs-on: ubuntu-latest
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
python-version: ["3.10", "3.12"]
|
||||
steps:
|
||||
- name: Check out code
|
||||
uses: actions/checkout@v7
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v7
|
||||
with:
|
||||
python-version: ${{ matrix.python-version }}
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
pip install -r requirements.txt
|
||||
pip install pytest
|
||||
- name: Run tests
|
||||
run: |
|
||||
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
|
||||
- name: Validate model registry
|
||||
run: |
|
||||
python -c "
|
||||
import json, pathlib
|
||||
path = pathlib.Path('data/fal_registry.json')
|
||||
if not path.exists():
|
||||
print('no registry file yet')
|
||||
else:
|
||||
registry = json.loads(path.read_text())
|
||||
assert {'version', 'models', 'model_count'} <= set(registry), 'missing required keys'
|
||||
assert registry['model_count'] == len(registry['models']), 'model_count mismatch'
|
||||
assert registry['model_count'] > 500, 'suspiciously few models'
|
||||
print(f\"registry OK: {registry['model_count']} models\")
|
||||
"
|
||||
@@ -0,0 +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:
|
||||
# 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
|
||||
|
||||
jobs:
|
||||
refresh:
|
||||
name: Validate and publish registry refresh
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 30
|
||||
steps:
|
||||
- name: Check out the default branch
|
||||
uses: actions/checkout@v7
|
||||
with:
|
||||
ref: main
|
||||
fetch-depth: 0
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v7
|
||||
with:
|
||||
python-version: "3.11"
|
||||
|
||||
- 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,3 +1,6 @@
|
||||
# Local configuration (contains API keys) — use config.ini.example as a template
|
||||
config.ini
|
||||
|
||||
# Byte-compiled / optimized / DLL files
|
||||
__pycache__/
|
||||
*.py[cod]
|
||||
@@ -164,7 +167,11 @@ cython_debug/
|
||||
# Cursor and SpecStory
|
||||
.specstory/
|
||||
.cursor/
|
||||
.claude
|
||||
.cursorignore
|
||||
.cursorindexingignore
|
||||
memory-bank/
|
||||
.DS_Store
|
||||
.claude/
|
||||
Node-Docs/
|
||||
output/
|
||||
|
||||
Vendored
+3
@@ -0,0 +1,3 @@
|
||||
{
|
||||
"workbench.colorTheme": "Community Material Theme Ocean High Contrast"
|
||||
}
|
||||
@@ -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
@@ -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
|
||||
@@ -1,144 +1,237 @@
|
||||
# ComfyUI-fal-API
|
||||
|
||||
Custom nodes for using Flux models with fal API in ComfyUI with only one API Key for all.
|
||||
**Every fal model in ComfyUI, one API key.**
|
||||
|
||||
Custom nodes that bring the entire [fal.ai](https://fal.ai) catalog into ComfyUI: ~90 curated hand-written nodes for the most popular models, plus ~1,400 auto-generated nodes covering every live public model on fal — image, video, audio, 3D, LLMs and more. One `FAL_KEY` unlocks all of them. With a persistent result cache (never pay for the same call twice), spend guards, async fan-out, and zero-I/O fal→fal chaining. The exact current catalog lives in [MODELS.md](MODELS.md).
|
||||
|
||||
## Table of Contents
|
||||
|
||||
- [What's New](#whats-new)
|
||||
- [Installation](#installation)
|
||||
- [Configuration](#configuration)
|
||||
- [Usage](#usage)
|
||||
- [Available Nodes](#available-nodes)
|
||||
- [Image Generation](#image-generation)
|
||||
- [Video Generation](#video-generation)
|
||||
- [Language Models (LLMs)](#language-models-llms)
|
||||
- [Vision Language Models (VLMs)](#vision-language-models-vlms)
|
||||
- [Screenshots](#screenshots)
|
||||
- [Curated Nodes](#curated-nodes)
|
||||
- [Auto-Generated Nodes](#auto-generated-nodes)
|
||||
- [Platform Utilities](#platform-utilities)
|
||||
- [Utility Nodes](#utility-nodes)
|
||||
- [Troubleshooting](#troubleshooting)
|
||||
- [Contributing](#contributing)
|
||||
- [License](#license)
|
||||
|
||||
## What's New
|
||||
|
||||
### 2.5.1
|
||||
|
||||
- **Schema controls across the catalog** — nested schema definitions and literal choices are resolved, and inputs are no longer truncated by a 40-field limit.
|
||||
- **Editable suggestions** — modes, voices, languages and model IDs with short example values get dropdowns with an **Enter custom value…** option. Saved text values and `STRING` connections keep working.
|
||||
- **Safer updates** — refresh remains available when only existing models changed. Scheduled and sidebar refreshes reject disappearing controls or choices before replacing the registry.
|
||||
|
||||
### 2.5
|
||||
|
||||
- **Async execution** — fal calls run asynchronously on modern ComfyUI, so the UI stays responsive while jobs are in flight.
|
||||
- **Typed builder nodes** — build LoRA lists and reference-element configs with dedicated, connectable nodes instead of hand-typed JSON.
|
||||
- **Featured tier** — hand-picked models surface under **FAL/Featured** in the node menu, and models superseded by newer versions are flagged.
|
||||
- **Registry freshness** — the pack warns when the committed model registry snapshot is stale, and you can refresh it from the fal sidebar without leaving ComfyUI.
|
||||
- **Automatic daily catalog updates** — a guarded GitHub workflow validates and commits new fal endpoints without waiting for a manual PR merge. Missing historical endpoints remain registered under **FAL/Compatibility** so saved workflows still load.
|
||||
- **MODELS.md** — the full ~1,400-model catalog moved out of this README into [MODELS.md](MODELS.md).
|
||||
|
||||
### 2.4
|
||||
|
||||
Twenty utility nodes under `FAL/Utils` covering everything between your assets and a fal endpoint — dataset prep (images → captioned training ZIP), media loading from URLs/folders, video trim/concat/mux/frame-extract, image grids and resizing, and JSON/prompt/text data helpers. See [Utility Nodes](#utility-nodes).
|
||||
|
||||
### 2.3
|
||||
|
||||
Durable **job inbox** (queued jobs survive ComfyUI restarts), **provenance receipts** (saved outputs carry endpoint + request id in a sidecar/PNG chunk and can be re-materialized for free), and the in-canvas web extension: cost badges above nodes, a fal sidebar with spend/balance/live jobs, and endpoint autocomplete.
|
||||
|
||||
### 2.2
|
||||
|
||||
**Persistent result cache** — identical fal calls are served from disk, free and instant, across restarts; uploads are deduplicated too. **Spend guard** — refuse to submit past a session budget or below a balance floor, before money moves. **URL passthrough** — wire a fal node's URL output into the next node's `*_direct_url` input and intermediate media never touches your machine.
|
||||
|
||||
### 2.1
|
||||
|
||||
Platform utilities under `FAL/Platform`: **Fal Submit + Fal Collect** for parallel fan-out (five video jobs run concurrently on fal, not back-to-back), **Fal Result by Request ID** to re-fetch any past result without re-paying, **Fal Cost Estimator**, and **Fal Session Costs**.
|
||||
|
||||
### 2.0
|
||||
|
||||
Every live public model on fal became a node: ~1,400 auto-generated nodes built at startup from the committed `data/fal_registry.json`, with native `IMAGE`/`VIDEO`/`AUDIO` sockets, schema-derived tooltips, and pricing in the help panel. Plus the generic **Fal Any Endpoint** node, visible errors (failed calls raise fal's actual error message instead of silently returning blanks), progress + cancellation, and full backward compatibility for all ~90 curated nodes.
|
||||
|
||||
## Installation
|
||||
|
||||
1. Navigate to your ComfyUI custom nodes directory:
|
||||
The recommended installation method is **ComfyUI Manager**: search for
|
||||
`ComfyUI-fal-API`, install it, and restart ComfyUI. Manager installs the
|
||||
dependencies into the same Python environment ComfyUI uses.
|
||||
|
||||
For a manual installation:
|
||||
|
||||
1. Navigate to your ComfyUI custom-nodes directory and clone this repository:
|
||||
```
|
||||
cd custom_nodes
|
||||
```
|
||||
|
||||
2. Clone this repository:
|
||||
```
|
||||
git clone https://github.com/gokayfem/ComfyUI-fal-API.git
|
||||
cd ComfyUI-fal-API
|
||||
```
|
||||
|
||||
3. Install the required dependencies:
|
||||
2. Install the dependencies with **ComfyUI's Python**, not an unrelated system
|
||||
`pip`:
|
||||
```
|
||||
pip install -r requirements.txt
|
||||
python -m pip install -r requirements.txt
|
||||
```
|
||||
From the root of **ComfyUI Windows Portable**, use its embedded interpreter:
|
||||
```powershell
|
||||
.\python_embeded\python.exe -m pip install -r .\ComfyUI\custom_nodes\ComfyUI-fal-API\requirements.txt
|
||||
```
|
||||
3. Configure your API key (below) and restart ComfyUI. Curated nodes appear under the **FAL** category, auto-generated nodes under **FAL/Models/<category>** (e.g. `FAL/Models/text-to-image`), and hand-picked models under **FAL/Featured** — or just search for any model by name.
|
||||
|
||||
## Configuration
|
||||
|
||||
1. Get your fal API key from [fal.ai](https://fal.ai/dashboard/keys)
|
||||
|
||||
2. Open the `config.ini` file inside `custom_nodes/ComfyUI-fal-API`
|
||||
|
||||
3. Replace `<your_fal_api_key_here>` with your actual fal API key:
|
||||
```ini
|
||||
[API]
|
||||
FAL_KEY = your_actual_api_key
|
||||
2. Copy `config.ini.example` to `config.ini` inside `custom_nodes/ComfyUI-fal-API` (`config.ini` is gitignored, so your key never ends up in a commit)
|
||||
3. Replace `<your_fal_api_key_here>` with your actual fal API key — or set the `FAL_KEY` environment variable instead:
|
||||
```bash
|
||||
export FAL_KEY=your_actual_api_key
|
||||
```
|
||||
|
||||
## Usage
|
||||
### config.ini reference
|
||||
|
||||
After installation and configuration, restart ComfyUI. The new nodes will be available in the node browser under the "FAL" category.
|
||||
All sections besides `[API]` are optional; defaults shown.
|
||||
|
||||
## Available Nodes
|
||||
```ini
|
||||
[API]
|
||||
FAL_KEY = your_actual_api_key
|
||||
|
||||
### Image Generation
|
||||
[dynamic_nodes]
|
||||
; Set to false to load only the curated hand-written nodes.
|
||||
enabled = true
|
||||
; Comma-separated category filter; leave unset to load everything.
|
||||
; categories = text-to-image,image-to-video
|
||||
|
||||
- **Flux Pro (fal)**: Generate high-quality images using the Flux Pro model
|
||||
- **Flux Dev (fal)**: Use the development version of Flux for image generation
|
||||
- **Flux Schnell (fal)**: Fast image generation with Flux Schnell
|
||||
- **Flux Pro 1.1 (fal)**: Latest version of Flux Pro for image generation
|
||||
- **Flux Ultra (fal)**: Ultra-high quality image generation with advanced controls
|
||||
- **Flux General (fal)**: ControlNets, Ipadapters, Loras for Flux Dev
|
||||
- **Flux LoRA (fal)**: Flux with dual LoRA support for custom styles
|
||||
- **Flux Pro Kontext (fal)**: Context-aware single image-to-image generation with max_quality toggle
|
||||
- **Flux Pro Kontext Multi (fal)**: Multi-image composition (2-4 images) with context awareness and max_quality toggle
|
||||
- **Flux Pro Kontext Text-to-Image (fal)**: Text-to-image with aspect ratio controls and max_quality toggle
|
||||
- **Recraft V3 (fal)**: Professional design generation with multiple style options
|
||||
- **Sana (fal)**: High-quality image synthesis with ultra-high resolution support
|
||||
- **HiDream Full (fal)**: Advanced image generation with comprehensive parameter control
|
||||
- **Ideogram v3 (fal)**: Advanced text-to-image generation with typography support
|
||||
[cache]
|
||||
; Persistent result cache: identical fal calls are served from disk (free)
|
||||
; across ComfyUI restarts. Bypass per node with force_rerun.
|
||||
enabled = true
|
||||
ttl_days = 7
|
||||
max_entries = 5000
|
||||
|
||||
### Video Generation
|
||||
[spend_guard]
|
||||
; Refuse to submit jobs once the session's estimated spend reaches the
|
||||
; budget, or when the account balance falls below the floor. 0 = disabled.
|
||||
session_budget_usd = 0
|
||||
min_balance_usd = 0
|
||||
|
||||
- **Kling Video Generation (fal)**: Generate videos using the Kling model
|
||||
- **Kling Pro v1.0 Video Generation (fal)**: Original version of Kling Pro for video generation
|
||||
- **Kling Pro v1.6 Video Generation (fal)**: Latest version of Kling Pro with improved quality
|
||||
- **Kling Master v2.0 Video Generation (fal)**: Advanced video generation with Kling Master
|
||||
- **Runway Gen3 Image-to-Video (fal)**: Convert images to videos using Runway Gen3
|
||||
- **Luma Dream Machine (fal)**: Create videos with Luma Dream Machine
|
||||
- **MiniMax Video Generation (fal)**: Generate videos using MiniMax model
|
||||
- **MiniMax Text-to-Video (fal)**: Create videos from text prompts using MiniMax
|
||||
- **MiniMax Subject Reference (fal)**: Generate videos with subject reference using MiniMax
|
||||
- **Google Veo2 Image-to-Video (fal)**: Convert images to videos using Google's Veo2 model
|
||||
- **Wan Pro Image-to-Video (fal)**: High-quality video generation with Wan Pro model
|
||||
- **Video Upscaler (fal)**: Upscale video quality using AI
|
||||
- **Combined Video Generation (fal)**: Generate videos using multiple services simultaneously
|
||||
- Supports Kling Pro v1.6, Kling Master v2.0, MiniMax, Luma, Veo2, and Wan Pro
|
||||
- Each service can be individually enabled/disabled
|
||||
- Wan Pro runs with safety checker enabled and automatic seed selection
|
||||
- **Load Video from URL**: Load and process videos from a given URL
|
||||
[archive]
|
||||
; Safety caps for the dataset/folder → ZIP upload utilities.
|
||||
max_files = 5000
|
||||
max_total_mb = 2048
|
||||
|
||||
### Language Models (LLMs)
|
||||
[registry]
|
||||
; On startup a background thread compares the local model registry
|
||||
; against fal's live catalog and logs how many new models are available;
|
||||
; the fal sidebar shows them with a one-click refresh (restart required).
|
||||
; startup_check = true
|
||||
|
||||
- **LLM (fal)**: Large Language Model for text generation and processing
|
||||
- Available models:
|
||||
- google/gemini-flash-1.5-8b
|
||||
- anthropic/claude-3.5-sonnet
|
||||
- anthropic/claude-3-haiku
|
||||
- google/gemini-pro-1.5
|
||||
- google/gemini-flash-1.5
|
||||
- meta-llama/llama-3.2-1b-instruct
|
||||
- meta-llama/llama-3.2-3b-instruct
|
||||
- meta-llama/llama-3.1-8b-instruct
|
||||
- meta-llama/llama-3.1-70b-instruct
|
||||
- openai/gpt-4o-mini
|
||||
- openai/gpt-4o
|
||||
[MODELS.md](MODELS.md).**
|
||||
|
||||
### Vision Language Models (VLMs)
|
||||
Find them in the node browser under `FAL/Models/<category>`, or search by model name. Node keys are `FalAPI_<endpoint-id>` (slashes → dashes), so workflows stay stable across registry refreshes. Each generated node gives you:
|
||||
|
||||
- **VLM (fal)**: Vision Language Model for image understanding and text generation
|
||||
- Available models:
|
||||
- google/gemini-flash-1.5-8b
|
||||
- anthropic/claude-3.5-sonnet
|
||||
- anthropic/claude-3-haiku
|
||||
- google/gemini-pro-1.5
|
||||
- google/gemini-flash-1.5
|
||||
- openai/gpt-4o
|
||||
- Supports various tasks such as image captioning, visual question answering, and more
|
||||
- **Native inputs/outputs** — `IMAGE`/`VIDEO`/`AUDIO` sockets; connected media is uploaded to fal automatically, and video/audio/image models return native ComfyUI types (JSON-ish models return the raw result string).
|
||||
- **Tooltips + pricing** — hover any input for fal's own parameter docs; the help panel shows the model's current pricing.
|
||||
- **Seed semantics** — `seed = -1` means "random / omit seed".
|
||||
- **force_rerun** — bypass result caching to re-roll identical inputs.
|
||||
|
||||
**Fal Any Endpoint (fal)** is the escape hatch: one generic node that calls *any* fal endpoint by id with free-form JSON arguments plus optional image/video/audio inputs (uploaded and merged into the matching keys). Outputs are extracted as `IMAGE`/`VIDEO`/`AUDIO`, with the raw result always available as `result_json`. Even brand-new models work the day they launch.
|
||||
|
||||
**Keeping the catalog fresh:** a daily GitHub Action builds a candidate registry, validates its structure and endpoint changes, regenerates [MODELS.md](MODELS.md), and commits the result when it is safe. Large unexpected catalog changes are blocked rather than published. Endpoints missing from a refresh are retained under **FAL/Compatibility** instead of deleting their node keys and breaking saved workflows. You can also run `python scripts/build_registry.py` yourself (then `python scripts/build_readme.py`), or use the refresh button in the fal sidebar. Install updates through ComfyUI Manager to receive the latest committed snapshot; locally refreshed nodes appear after restarting ComfyUI. If ~1,400 extra nodes is more than you want, the `[dynamic_nodes]` config section disables or filters them — see [Configuration](#configuration).
|
||||
|
||||
## Platform Utilities
|
||||
|
||||
Nodes built on fal's platform primitives (queue, request ids, per-model pricing), under `FAL/Platform`:
|
||||
|
||||
| Node | What it does |
|
||||
| --- | --- |
|
||||
| **Fal Submit** / **Fal Collect** | Queue a job on any endpoint and collect it later — wire N Submits into N Collects and all N jobs run **in parallel** on fal, so the graph takes as long as the slowest one, not the sum. |
|
||||
| **Fal Result by Request ID** | Paste any past request id (console log, sidebar, or [fal dashboard](https://fal.ai/dashboard/requests)) to re-fetch its result **without re-generating or re-paying**. |
|
||||
| **Fal Cost Estimator** | Endpoint id + run count → cost report and `total_usd` float from fal's published pricing, *before* you queue anything. |
|
||||
| **Fal Session Costs** | Running ledger of every fal call this session (endpoint, duration, request id, estimated cost), with optional reset. |
|
||||
| **Fal Account Balance** | Your live fal balance (needs an admin-scoped key; scoped keys degrade gracefully). |
|
||||
| **Fal Job Inbox** | Every Submit is journaled to disk, so queued jobs **survive ComfyUI restarts** — lists pending/collected jobs and outputs the newest pending `request_id` + `endpoint_id`. |
|
||||
| **Fal Save Media from URL** | fal result URLs eventually expire; this downloads any result into your `output/` directory and writes a provenance receipt (`<file>.fal.json` sidecar + PNG text chunk with endpoint, request id, source URL). |
|
||||
| **Fal Provenance from File** | Read a saved file's receipt back into `endpoint_id` + `request_id` — re-materialize a generation for free, months later. |
|
||||
|
||||
Behind the nodes, three always-on platform features:
|
||||
|
||||
- **Persistent result cache** — identical fal calls are served from a disk cache, free and instant, across restarts; input uploads are deduplicated. Bypass per node with `force_rerun`; tune via `[cache]`.
|
||||
- **Spend guard** — with `[spend_guard]` configured, the pack refuses to submit once estimated session spend hits your budget or your balance drops below the floor — the node errors *before* money moves.
|
||||
- **URL passthrough** — every generated node's media input has a `*_direct_url` twin and image nodes output `image_urls`; chain fal→fal by URL and intermediate media never touches your machine.
|
||||
|
||||
And a web extension (degrades silently on older ComfyUI): **cost badges** above every priced fal node with live estimates on free-typed endpoint fields, a **fal sidebar** (session spend, balance, live job list with per-job Cancel, registry freshness/refresh), and **endpoint autocomplete** with prices when editing any `endpoint_id` field.
|
||||
|
||||
## Utility Nodes
|
||||
|
||||
Twenty nodes under `FAL/Utils` covering everything between your assets and a fal endpoint — no other packs needed:
|
||||
|
||||
| Category | Nodes |
|
||||
| --- | --- |
|
||||
| `FAL/Utils/Dataset` | *Images → Training ZIP URL* (standard LoRA caption layout), *Folder → ZIP URL*, *Video → Frame Dataset ZIP URL*, *Batch Caption Images* (parallel VLM captioning) |
|
||||
| `FAL/Utils/Load` | *Load Image from URL* (multi-URL batching), *Load Audio from URL*, *Load Image Folder*, *Upload Folder as ZIP URL* |
|
||||
| `FAL/Utils/Video` | *Extract Frames* (efficient last-frame seek — chain into image-to-video for endless extension), *Trim* (keyframe remux, no re-encode), *Concat* (auto resolution/fps normalize), *Mux Audio + Video*, *Video → Audio* |
|
||||
| `FAL/Utils/Image` | *Image Grid with Labels*, *Resize to fal Preset* (cover/contain/stretch to standard image_size dims), *Image ↔ Base64* |
|
||||
| `FAL/Utils/Data` | *JSON Extract* (dot/bracket path queries against any `result_json`), *Prompt Lines* (cycling line picker), *Text Template* |
|
||||
|
||||
The full LoRA-training pipeline needs nothing else: Load Image Folder → Batch Caption → Images→ZIP → any trainer node.
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
If you encounter any errors during installation or usage, try the following:
|
||||
|
||||
1. Ensure you have the latest version of ComfyUI installed
|
||||
2. Update this custom node package:
|
||||
```
|
||||
cd custom_nodes/ComfyUI-fal-API
|
||||
git pull
|
||||
pip install -r requirements.txt
|
||||
python -m pip install -r requirements.txt
|
||||
```
|
||||
3. If you're using ComfyUI Windows Portable, you may need to install fal-client manually:
|
||||
```
|
||||
ComfyUI_windows_portable>.\python_embeded\python.exe -m pip install fal-client
|
||||
3. **Windows Portable Python is blocked or no longer starts after installing a node?**
|
||||
Do not keep rerunning `pip`. From the portable root, first check the exact
|
||||
interpreter and dependency state:
|
||||
```powershell
|
||||
.\python_embeded\python.exe -c "import sys; print(sys.executable); print(sys.version)"
|
||||
.\python_embeded\python.exe -m pip check
|
||||
```
|
||||
This project installs Python packages only; it does not replace or modify
|
||||
`python.exe`. If the first command itself is blocked or the executable was
|
||||
quarantined, review Windows Security **Protection history**. Restore a file
|
||||
only when the portable archive came from the official ComfyUI release, or
|
||||
re-extract a clean official portable build and move your `models`, `input`,
|
||||
`output`, and `user` data across. Avoid disabling antivirus globally. Then
|
||||
reinstall this node through ComfyUI Manager, or use the exact embedded-
|
||||
interpreter requirements command from the Installation section.
|
||||
4. **Dynamic nodes not appearing?** Check the ComfyUI console for a line like `Registered N dynamic fal nodes` at startup. If it says the nodes are disabled, remove `enabled = false` from the `[dynamic_nodes]` section of your `config.ini` (and check the `categories` filter isn't excluding what you're looking for). Any registry loading error is also printed there.
|
||||
5. **`VIDEO` output is `None` or video sockets are missing?** Update ComfyUI — native `VIDEO`/`AUDIO` types require a recent ComfyUI version.
|
||||
6. **API calls failing?** Failed fal requests raise visible errors that include fal's actual error message (validation issues, content policy, quota). Read the error text in ComfyUI — it usually tells you exactly which parameter to fix.
|
||||
7. **Missing duration, resolution, or other model controls?** Update the pack in ComfyUI Manager, then use **fal sidebar → Registry → Refresh registry**. The button is available even when no new model IDs are found: existing APIs can add or change controls without adding a model. Restart ComfyUI and reload the browser afterward. If an existing canvas node still shows old widgets, add a fresh instance of the same node. H3 generation nodes expose **duration (5–15 seconds)**; H3 Max and Max Turbo expose **480P / 768P / 1080P** and a **disabled / balanced / quality** prompt expansion dropdown, following the current [H3 Max API](https://fal.ai/models/minimax/h3-max/image-to-video/api). Original H3 uses its own schema options, including the `fast` prompt expansion mode.
|
||||
|
||||
## Contributing
|
||||
|
||||
Contributions are welcome — but note that **a new fal model usually needs no code at all**: it appears automatically via the registry. Read [CONTRIBUTING.md](CONTRIBUTING.md) before opening a PR; it covers when a hand-written node is (and isn't) justified, the compatibility rules, and dev setup.
|
||||
|
||||
<details>
|
||||
<summary><strong>Cite this project</strong></summary>
|
||||
|
||||
If ComfyUI-fal-API supports your work, please cite the software. GitHub also
|
||||
provides ready-to-copy APA and BibTeX entries via **Cite this repository**.
|
||||
|
||||
```bibtex
|
||||
@software{Aydogan_ComfyUI_fal_API_2026,
|
||||
author = {Aydoğan, Gökay},
|
||||
title = {ComfyUI-fal-API},
|
||||
version = {2.5.0},
|
||||
year = {2026},
|
||||
url = {https://github.com/gokayfem/ComfyUI-fal-API}
|
||||
}
|
||||
```
|
||||
|
||||
[ORCID](https://orcid.org/0000-0002-2343-9433) · [Citation metadata](CITATION.cff)
|
||||
|
||||
</details>
|
||||
|
||||
## License
|
||||
|
||||
This project is licensed under the Apache License 2.0. See the [LICENSE](LICENSE) file for details.
|
||||
|
||||
## Contributing
|
||||
|
||||
Contributions are welcome! Please feel free to submit a Pull Request.
|
||||
|
||||
## Support
|
||||
|
||||
If you encounter any issues or have questions, please open an issue on the [GitHub repository](https://github.com/gokayfem/ComfyUI-fal-API/issues).
|
||||
|
||||
+48
-8
@@ -1,12 +1,22 @@
|
||||
import importlib.util
|
||||
import importlib
|
||||
import importlib.util
|
||||
|
||||
node_list = [
|
||||
"image_node",
|
||||
"video_node",
|
||||
"llm_node",
|
||||
"vlm_node",
|
||||
"trainer_node",
|
||||
"image_node",
|
||||
"video_node",
|
||||
"llm_node",
|
||||
"vlm_node",
|
||||
"trainer_node",
|
||||
"upscaler_node",
|
||||
"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 = {}
|
||||
@@ -16,7 +26,37 @@ for module_name in node_list:
|
||||
imported_module = importlib.import_module(f".nodes.{module_name}", __name__)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {**NODE_CLASS_MAPPINGS, **imported_module.NODE_CLASS_MAPPINGS}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {**NODE_DISPLAY_NAME_MAPPINGS, **imported_module.NODE_DISPLAY_NAME_MAPPINGS}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
**NODE_DISPLAY_NAME_MAPPINGS,
|
||||
**imported_module.NODE_DISPLAY_NAME_MAPPINGS,
|
||||
}
|
||||
|
||||
try:
|
||||
from .nodes.dynamic import get_dynamic_mappings
|
||||
|
||||
dyn_classes, dyn_display = get_dynamic_mappings()
|
||||
# static nodes win on any key collision
|
||||
for k, v in dyn_classes.items():
|
||||
NODE_CLASS_MAPPINGS.setdefault(k, v)
|
||||
for k, v in dyn_display.items():
|
||||
NODE_DISPLAY_NAME_MAPPINGS.setdefault(k, v)
|
||||
except Exception as _dynamic_error: # never break static nodes
|
||||
import logging
|
||||
|
||||
logging.getLogger(__name__).error(
|
||||
"Failed to load dynamic fal nodes: %s", _dynamic_error
|
||||
)
|
||||
|
||||
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||
WEB_DIRECTORY = "./web"
|
||||
|
||||
try:
|
||||
from .nodes import server_routes as _server_routes # noqa: F401 registers /fal_api routes
|
||||
except Exception as _routes_error: # never break node loading over HTTP extras
|
||||
import logging
|
||||
|
||||
logging.getLogger(__name__).warning(
|
||||
"fal API server routes not registered: %s", _routes_error
|
||||
)
|
||||
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS", "WEB_DIRECTORY"]
|
||||
|
||||
@@ -0,0 +1,21 @@
|
||||
[API]
|
||||
FAL_KEY = <your_fal_api_key_here>
|
||||
|
||||
; --- Optional sections (defaults shown; uncomment to change) ---
|
||||
|
||||
; [dynamic_nodes]
|
||||
; enabled = true
|
||||
; categories = text-to-image,image-to-video
|
||||
|
||||
; [cache]
|
||||
; Persistent result cache: identical fal calls are served from disk (free)
|
||||
; across ComfyUI restarts. Bypass per node with force_rerun.
|
||||
; enabled = true
|
||||
; ttl_days = 7
|
||||
; max_entries = 5000
|
||||
|
||||
; [spend_guard]
|
||||
; Refuse to submit jobs once the session's estimated spend reaches the
|
||||
; budget, or when the account balance falls below the floor. 0 = disabled.
|
||||
; session_budget_usd = 0
|
||||
; min_balance_usd = 0
|
||||
File diff suppressed because one or more lines are too long
@@ -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}
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,211 @@
|
||||
{
|
||||
"id": "b7d0f93a-df07-4002-8769-ae88c70c403a",
|
||||
"revision": 0,
|
||||
"last_node_id": 4,
|
||||
"last_link_id": 3,
|
||||
"nodes": [
|
||||
{
|
||||
"id": 2,
|
||||
"type": "LoadImage",
|
||||
"pos": [
|
||||
-3530.263671875,
|
||||
-2397.567138671875
|
||||
],
|
||||
"size": [
|
||||
274.080078125,
|
||||
314
|
||||
],
|
||||
"flags": {},
|
||||
"order": 0,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
2
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "MASK",
|
||||
"type": "MASK",
|
||||
"links": null
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "comfy-core",
|
||||
"ver": "0.3.59",
|
||||
"Node name for S&R": "LoadImage"
|
||||
},
|
||||
"widgets_values": [
|
||||
"knight.jpeg",
|
||||
"image"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 3,
|
||||
"type": "LoadImage",
|
||||
"pos": [
|
||||
-3527.47119140625,
|
||||
-2022.4005126953125
|
||||
],
|
||||
"size": [
|
||||
274.080078125,
|
||||
314
|
||||
],
|
||||
"flags": {},
|
||||
"order": 1,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
1
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "MASK",
|
||||
"type": "MASK",
|
||||
"links": null
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "comfy-core",
|
||||
"ver": "0.3.59",
|
||||
"Node name for S&R": "LoadImage"
|
||||
},
|
||||
"widgets_values": [
|
||||
"mask_knight.jpeg",
|
||||
"image"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 1,
|
||||
"type": "FluxPro1Fill_fal",
|
||||
"pos": [
|
||||
-3013.856689453125,
|
||||
-2388.482177734375
|
||||
],
|
||||
"size": [
|
||||
400,
|
||||
276
|
||||
],
|
||||
"flags": {},
|
||||
"order": 2,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "image",
|
||||
"shape": 7,
|
||||
"type": "IMAGE",
|
||||
"link": 2
|
||||
},
|
||||
{
|
||||
"name": "mask_image",
|
||||
"shape": 7,
|
||||
"type": "IMAGE",
|
||||
"link": 1
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
3
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "fal-api",
|
||||
"ver": "88466f23804f2e9e6a905b82bf6693c754a466ba",
|
||||
"Node name for S&R": "FluxPro1Fill_fal"
|
||||
},
|
||||
"widgets_values": [
|
||||
"A big yellow smiley face.",
|
||||
1,
|
||||
"2",
|
||||
"png",
|
||||
1647,
|
||||
"randomize",
|
||||
false,
|
||||
true
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 4,
|
||||
"type": "PreviewImage",
|
||||
"pos": [
|
||||
-2539.11376953125,
|
||||
-2386.515625
|
||||
],
|
||||
"size": [
|
||||
427.3517150878906,
|
||||
400.66522216796875
|
||||
],
|
||||
"flags": {},
|
||||
"order": 3,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "images",
|
||||
"type": "IMAGE",
|
||||
"link": 3
|
||||
}
|
||||
],
|
||||
"outputs": [],
|
||||
"properties": {
|
||||
"cnr_id": "comfy-core",
|
||||
"ver": "0.3.59",
|
||||
"Node name for S&R": "PreviewImage"
|
||||
},
|
||||
"widgets_values": []
|
||||
}
|
||||
],
|
||||
"links": [
|
||||
[
|
||||
1,
|
||||
3,
|
||||
0,
|
||||
1,
|
||||
1,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
2,
|
||||
2,
|
||||
0,
|
||||
1,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
3,
|
||||
1,
|
||||
0,
|
||||
4,
|
||||
0,
|
||||
"IMAGE"
|
||||
]
|
||||
],
|
||||
"groups": [],
|
||||
"config": {},
|
||||
"extra": {
|
||||
"ds": {
|
||||
"scale": 1.0152559799477112,
|
||||
"offset": [
|
||||
4064.0521902493365,
|
||||
2549.035899408263
|
||||
]
|
||||
},
|
||||
"frontendVersion": "1.27.10",
|
||||
"VHS_latentpreview": false,
|
||||
"VHS_latentpreviewrate": 0,
|
||||
"VHS_MetadataImage": true,
|
||||
"VHS_KeepIntermediate": true
|
||||
},
|
||||
"version": 0.4
|
||||
}
|
||||
@@ -0,0 +1,970 @@
|
||||
{
|
||||
"id": "d3437cd7-7301-49a6-b533-dd71877f7816",
|
||||
"revision": 0,
|
||||
"last_node_id": 20,
|
||||
"last_link_id": 19,
|
||||
"nodes": [
|
||||
{
|
||||
"id": 5,
|
||||
"type": "NanoBananaPro_fal",
|
||||
"pos": [
|
||||
2510,
|
||||
-1370
|
||||
],
|
||||
"size": [
|
||||
400,
|
||||
208
|
||||
],
|
||||
"flags": {},
|
||||
"order": 15,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "images",
|
||||
"shape": 7,
|
||||
"type": "IMAGE",
|
||||
"link": 19
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
4
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "fal-api",
|
||||
"ver": "54b3182c6adf925426dc25e28483951350edd477",
|
||||
"Node name for S&R": "NanoBananaPro_fal",
|
||||
"ue_properties": {
|
||||
"widget_ue_connectable": {},
|
||||
"input_ue_unconnectable": {},
|
||||
"version": "7.4.1"
|
||||
}
|
||||
},
|
||||
"widgets_values": [
|
||||
"a group photo of 14 people",
|
||||
1,
|
||||
"21:9",
|
||||
"png",
|
||||
"2K",
|
||||
false
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 6,
|
||||
"type": "ImpactMakeImageBatch",
|
||||
"pos": [
|
||||
2240,
|
||||
-1370
|
||||
],
|
||||
"size": [
|
||||
156.6236328125,
|
||||
306
|
||||
],
|
||||
"flags": {},
|
||||
"order": 14,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "image1",
|
||||
"shape": 7,
|
||||
"type": "IMAGE",
|
||||
"link": 5
|
||||
},
|
||||
{
|
||||
"name": "image2",
|
||||
"type": "IMAGE",
|
||||
"link": 6
|
||||
},
|
||||
{
|
||||
"name": "image3",
|
||||
"type": "IMAGE",
|
||||
"link": 7
|
||||
},
|
||||
{
|
||||
"name": "image4",
|
||||
"type": "IMAGE",
|
||||
"link": 8
|
||||
},
|
||||
{
|
||||
"name": "image5",
|
||||
"type": "IMAGE",
|
||||
"link": 9
|
||||
},
|
||||
{
|
||||
"name": "image6",
|
||||
"type": "IMAGE",
|
||||
"link": 10
|
||||
},
|
||||
{
|
||||
"name": "image7",
|
||||
"type": "IMAGE",
|
||||
"link": 11
|
||||
},
|
||||
{
|
||||
"name": "image8",
|
||||
"type": "IMAGE",
|
||||
"link": 12
|
||||
},
|
||||
{
|
||||
"name": "image9",
|
||||
"type": "IMAGE",
|
||||
"link": 13
|
||||
},
|
||||
{
|
||||
"name": "image10",
|
||||
"type": "IMAGE",
|
||||
"link": 14
|
||||
},
|
||||
{
|
||||
"name": "image11",
|
||||
"type": "IMAGE",
|
||||
"link": 15
|
||||
},
|
||||
{
|
||||
"name": "image12",
|
||||
"type": "IMAGE",
|
||||
"link": 16
|
||||
},
|
||||
{
|
||||
"name": "image13",
|
||||
"type": "IMAGE",
|
||||
"link": 17
|
||||
},
|
||||
{
|
||||
"name": "image14",
|
||||
"type": "IMAGE",
|
||||
"link": 18
|
||||
},
|
||||
{
|
||||
"name": "image15",
|
||||
"type": "IMAGE",
|
||||
"link": null
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
19
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "comfyui-impact-pack",
|
||||
"ver": "8.25.1",
|
||||
"Node name for S&R": "ImpactMakeImageBatch",
|
||||
"ue_properties": {
|
||||
"widget_ue_connectable": {},
|
||||
"input_ue_unconnectable": {},
|
||||
"version": "7.4.1"
|
||||
}
|
||||
},
|
||||
"widgets_values": []
|
||||
},
|
||||
{
|
||||
"id": 13,
|
||||
"type": "LoadImage",
|
||||
"pos": [
|
||||
1460,
|
||||
-2230
|
||||
],
|
||||
"size": [
|
||||
274.080078125,
|
||||
314
|
||||
],
|
||||
"flags": {},
|
||||
"order": 0,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
9
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "MASK",
|
||||
"type": "MASK",
|
||||
"links": null
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "comfy-core",
|
||||
"ver": "0.3.68",
|
||||
"Node name for S&R": "LoadImage",
|
||||
"ue_properties": {
|
||||
"widget_ue_connectable": {},
|
||||
"input_ue_unconnectable": {},
|
||||
"version": "7.4.1"
|
||||
}
|
||||
},
|
||||
"widgets_values": [
|
||||
"Image_3.jpg",
|
||||
"image"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 2,
|
||||
"type": "PreviewImage",
|
||||
"pos": [
|
||||
2930,
|
||||
-1370
|
||||
],
|
||||
"size": [
|
||||
630,
|
||||
360
|
||||
],
|
||||
"flags": {},
|
||||
"order": 16,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "images",
|
||||
"type": "IMAGE",
|
||||
"link": 4
|
||||
}
|
||||
],
|
||||
"outputs": [],
|
||||
"properties": {
|
||||
"cnr_id": "comfy-core",
|
||||
"ver": "0.3.68",
|
||||
"Node name for S&R": "PreviewImage",
|
||||
"ue_properties": {
|
||||
"widget_ue_connectable": {},
|
||||
"input_ue_unconnectable": {},
|
||||
"version": "7.4.1"
|
||||
}
|
||||
},
|
||||
"widgets_values": []
|
||||
},
|
||||
{
|
||||
"id": 10,
|
||||
"type": "LoadImage",
|
||||
"pos": [
|
||||
1180,
|
||||
-2230
|
||||
],
|
||||
"size": [
|
||||
274.080078125,
|
||||
314
|
||||
],
|
||||
"flags": {},
|
||||
"order": 1,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
8
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "MASK",
|
||||
"type": "MASK",
|
||||
"links": null
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "comfy-core",
|
||||
"ver": "0.3.68",
|
||||
"Node name for S&R": "LoadImage",
|
||||
"ue_properties": {
|
||||
"widget_ue_connectable": {},
|
||||
"input_ue_unconnectable": {},
|
||||
"version": "7.4.1"
|
||||
}
|
||||
},
|
||||
"widgets_values": [
|
||||
"Image_1.png",
|
||||
"image"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 12,
|
||||
"type": "LoadImage",
|
||||
"pos": [
|
||||
1180,
|
||||
-1880
|
||||
],
|
||||
"size": [
|
||||
274.080078125,
|
||||
314
|
||||
],
|
||||
"flags": {},
|
||||
"order": 2,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
10
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "MASK",
|
||||
"type": "MASK",
|
||||
"links": null
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "comfy-core",
|
||||
"ver": "0.3.68",
|
||||
"Node name for S&R": "LoadImage",
|
||||
"ue_properties": {
|
||||
"widget_ue_connectable": {},
|
||||
"input_ue_unconnectable": {},
|
||||
"version": "7.4.1"
|
||||
}
|
||||
},
|
||||
"widgets_values": [
|
||||
"Image_2.png",
|
||||
"image"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 11,
|
||||
"type": "LoadImage",
|
||||
"pos": [
|
||||
1460,
|
||||
-1880
|
||||
],
|
||||
"size": [
|
||||
274.080078125,
|
||||
314
|
||||
],
|
||||
"flags": {},
|
||||
"order": 3,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
7
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "MASK",
|
||||
"type": "MASK",
|
||||
"links": null
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "comfy-core",
|
||||
"ver": "0.3.68",
|
||||
"Node name for S&R": "LoadImage",
|
||||
"ue_properties": {
|
||||
"widget_ue_connectable": {},
|
||||
"input_ue_unconnectable": {},
|
||||
"version": "7.4.1"
|
||||
}
|
||||
},
|
||||
"widgets_values": [
|
||||
"Image_4.png",
|
||||
"image"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 16,
|
||||
"type": "LoadImage",
|
||||
"pos": [
|
||||
1740,
|
||||
-1880
|
||||
],
|
||||
"size": [
|
||||
274.080078125,
|
||||
314
|
||||
],
|
||||
"flags": {},
|
||||
"order": 4,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
6
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "MASK",
|
||||
"type": "MASK",
|
||||
"links": null
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "comfy-core",
|
||||
"ver": "0.3.68",
|
||||
"Node name for S&R": "LoadImage",
|
||||
"ue_properties": {
|
||||
"widget_ue_connectable": {},
|
||||
"input_ue_unconnectable": {},
|
||||
"version": "7.4.1"
|
||||
}
|
||||
},
|
||||
"widgets_values": [
|
||||
"Image_6.jpeg",
|
||||
"image"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 19,
|
||||
"type": "LoadImage",
|
||||
"pos": [
|
||||
2300,
|
||||
-2230
|
||||
],
|
||||
"size": [
|
||||
274.080078125,
|
||||
314
|
||||
],
|
||||
"flags": {},
|
||||
"order": 5,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
13
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "MASK",
|
||||
"type": "MASK",
|
||||
"links": null
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "comfy-core",
|
||||
"ver": "0.3.68",
|
||||
"Node name for S&R": "LoadImage",
|
||||
"ue_properties": {
|
||||
"widget_ue_connectable": {},
|
||||
"input_ue_unconnectable": {},
|
||||
"version": "7.4.1"
|
||||
}
|
||||
},
|
||||
"widgets_values": [
|
||||
"Image_9.jpeg",
|
||||
"image"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 15,
|
||||
"type": "LoadImage",
|
||||
"pos": [
|
||||
1740,
|
||||
-2230
|
||||
],
|
||||
"size": [
|
||||
274.080078125,
|
||||
314
|
||||
],
|
||||
"flags": {},
|
||||
"order": 6,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
11
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "MASK",
|
||||
"type": "MASK",
|
||||
"links": null
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "comfy-core",
|
||||
"ver": "0.3.68",
|
||||
"Node name for S&R": "LoadImage",
|
||||
"ue_properties": {
|
||||
"widget_ue_connectable": {},
|
||||
"input_ue_unconnectable": {},
|
||||
"version": "7.4.1"
|
||||
}
|
||||
},
|
||||
"widgets_values": [
|
||||
"Image_5.png",
|
||||
"image"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 20,
|
||||
"type": "LoadImage",
|
||||
"pos": [
|
||||
2020,
|
||||
-2230
|
||||
],
|
||||
"size": [
|
||||
274.080078125,
|
||||
314
|
||||
],
|
||||
"flags": {},
|
||||
"order": 7,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
12
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "MASK",
|
||||
"type": "MASK",
|
||||
"links": null
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "comfy-core",
|
||||
"ver": "0.3.68",
|
||||
"Node name for S&R": "LoadImage",
|
||||
"ue_properties": {
|
||||
"widget_ue_connectable": {},
|
||||
"input_ue_unconnectable": {},
|
||||
"version": "7.4.1"
|
||||
}
|
||||
},
|
||||
"widgets_values": [
|
||||
"Image_7.png",
|
||||
"image"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 14,
|
||||
"type": "LoadImage",
|
||||
"pos": [
|
||||
2020,
|
||||
-1880
|
||||
],
|
||||
"size": [
|
||||
274.080078125,
|
||||
314
|
||||
],
|
||||
"flags": {},
|
||||
"order": 8,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
14
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "MASK",
|
||||
"type": "MASK",
|
||||
"links": null
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "comfy-core",
|
||||
"ver": "0.3.68",
|
||||
"Node name for S&R": "LoadImage",
|
||||
"ue_properties": {
|
||||
"widget_ue_connectable": {},
|
||||
"input_ue_unconnectable": {},
|
||||
"version": "7.4.1"
|
||||
}
|
||||
},
|
||||
"widgets_values": [
|
||||
"Image_8.png",
|
||||
"image"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 17,
|
||||
"type": "LoadImage",
|
||||
"pos": [
|
||||
2300,
|
||||
-1880
|
||||
],
|
||||
"size": [
|
||||
274.080078125,
|
||||
314
|
||||
],
|
||||
"flags": {},
|
||||
"order": 9,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
16
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "MASK",
|
||||
"type": "MASK",
|
||||
"links": null
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "comfy-core",
|
||||
"ver": "0.3.68",
|
||||
"Node name for S&R": "LoadImage",
|
||||
"ue_properties": {
|
||||
"widget_ue_connectable": {},
|
||||
"input_ue_unconnectable": {},
|
||||
"version": "7.4.1"
|
||||
}
|
||||
},
|
||||
"widgets_values": [
|
||||
"Image_10.png",
|
||||
"image"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 7,
|
||||
"type": "LoadImage",
|
||||
"pos": [
|
||||
2580,
|
||||
-1880
|
||||
],
|
||||
"size": [
|
||||
274.080078125,
|
||||
314
|
||||
],
|
||||
"flags": {},
|
||||
"order": 10,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
5
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "MASK",
|
||||
"type": "MASK",
|
||||
"links": null
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "comfy-core",
|
||||
"ver": "0.3.68",
|
||||
"Node name for S&R": "LoadImage",
|
||||
"ue_properties": {
|
||||
"widget_ue_connectable": {},
|
||||
"input_ue_unconnectable": {},
|
||||
"version": "7.4.1"
|
||||
}
|
||||
},
|
||||
"widgets_values": [
|
||||
"Image_12.png",
|
||||
"image"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 18,
|
||||
"type": "LoadImage",
|
||||
"pos": [
|
||||
2580,
|
||||
-2230
|
||||
],
|
||||
"size": [
|
||||
274.080078125,
|
||||
314
|
||||
],
|
||||
"flags": {},
|
||||
"order": 11,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
15
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "MASK",
|
||||
"type": "MASK",
|
||||
"links": null
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "comfy-core",
|
||||
"ver": "0.3.68",
|
||||
"Node name for S&R": "LoadImage",
|
||||
"ue_properties": {
|
||||
"widget_ue_connectable": {},
|
||||
"input_ue_unconnectable": {},
|
||||
"version": "7.4.1"
|
||||
}
|
||||
},
|
||||
"widgets_values": [
|
||||
"Image_11.jpg",
|
||||
"image"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 9,
|
||||
"type": "LoadImage",
|
||||
"pos": [
|
||||
2860,
|
||||
-2230
|
||||
],
|
||||
"size": [
|
||||
274.080078125,
|
||||
314
|
||||
],
|
||||
"flags": {},
|
||||
"order": 12,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
17
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "MASK",
|
||||
"type": "MASK",
|
||||
"links": null
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "comfy-core",
|
||||
"ver": "0.3.68",
|
||||
"Node name for S&R": "LoadImage",
|
||||
"ue_properties": {
|
||||
"widget_ue_connectable": {},
|
||||
"input_ue_unconnectable": {},
|
||||
"version": "7.4.1"
|
||||
}
|
||||
},
|
||||
"widgets_values": [
|
||||
"Image_13.jpg",
|
||||
"image"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 8,
|
||||
"type": "LoadImage",
|
||||
"pos": [
|
||||
2860,
|
||||
-1880
|
||||
],
|
||||
"size": [
|
||||
274.080078125,
|
||||
314
|
||||
],
|
||||
"flags": {},
|
||||
"order": 13,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
18
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "MASK",
|
||||
"type": "MASK",
|
||||
"links": null
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "comfy-core",
|
||||
"ver": "0.3.68",
|
||||
"Node name for S&R": "LoadImage",
|
||||
"ue_properties": {
|
||||
"widget_ue_connectable": {},
|
||||
"input_ue_unconnectable": {},
|
||||
"version": "7.4.1"
|
||||
}
|
||||
},
|
||||
"widgets_values": [
|
||||
"Image_14.jpg",
|
||||
"image"
|
||||
]
|
||||
}
|
||||
],
|
||||
"links": [
|
||||
[
|
||||
4,
|
||||
5,
|
||||
0,
|
||||
2,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
5,
|
||||
7,
|
||||
0,
|
||||
6,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
6,
|
||||
16,
|
||||
0,
|
||||
6,
|
||||
1,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
7,
|
||||
11,
|
||||
0,
|
||||
6,
|
||||
2,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
8,
|
||||
10,
|
||||
0,
|
||||
6,
|
||||
3,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
9,
|
||||
13,
|
||||
0,
|
||||
6,
|
||||
4,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
10,
|
||||
12,
|
||||
0,
|
||||
6,
|
||||
5,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
11,
|
||||
15,
|
||||
0,
|
||||
6,
|
||||
6,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
12,
|
||||
20,
|
||||
0,
|
||||
6,
|
||||
7,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
13,
|
||||
19,
|
||||
0,
|
||||
6,
|
||||
8,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
14,
|
||||
14,
|
||||
0,
|
||||
6,
|
||||
9,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
15,
|
||||
18,
|
||||
0,
|
||||
6,
|
||||
10,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
16,
|
||||
17,
|
||||
0,
|
||||
6,
|
||||
11,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
17,
|
||||
9,
|
||||
0,
|
||||
6,
|
||||
12,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
18,
|
||||
8,
|
||||
0,
|
||||
6,
|
||||
13,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
19,
|
||||
6,
|
||||
0,
|
||||
5,
|
||||
0,
|
||||
"IMAGE"
|
||||
]
|
||||
],
|
||||
"groups": [],
|
||||
"config": {},
|
||||
"extra": {
|
||||
"ds": {
|
||||
"scale": 0.8769226950000027,
|
||||
"offset": [
|
||||
-1077.4046047064173,
|
||||
2181.4808993231786
|
||||
]
|
||||
},
|
||||
"ue_links": [],
|
||||
"links_added_by_ue": [],
|
||||
"frontendVersion": "1.28.8",
|
||||
"VHS_latentpreview": false,
|
||||
"VHS_latentpreviewrate": 0,
|
||||
"VHS_MetadataImage": true,
|
||||
"VHS_KeepIntermediate": true
|
||||
},
|
||||
"version": 0.4
|
||||
}
|
||||
@@ -0,0 +1,371 @@
|
||||
{
|
||||
"id": "8e5d7b17-8bec-4eae-8c78-0bcc142bd0c8",
|
||||
"revision": 0,
|
||||
"last_node_id": 17,
|
||||
"last_link_id": 28,
|
||||
"nodes": [
|
||||
{
|
||||
"id": 3,
|
||||
"type": "LoadVideo",
|
||||
"pos": [
|
||||
-3565.794921875,
|
||||
-2337.682861328125
|
||||
],
|
||||
"size": [
|
||||
274.080078125,
|
||||
232.1231231689453
|
||||
],
|
||||
"flags": {},
|
||||
"order": 0,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "VIDEO",
|
||||
"type": "VIDEO",
|
||||
"links": [
|
||||
9
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "comfy-core",
|
||||
"ver": "0.3.59",
|
||||
"Node name for S&R": "LoadVideo"
|
||||
},
|
||||
"widgets_values": [
|
||||
"AnimateDiff_00039.mp4",
|
||||
"image"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 6,
|
||||
"type": "VHS_VideoInfo",
|
||||
"pos": [
|
||||
-2305.713623046875,
|
||||
-2289.08837890625
|
||||
],
|
||||
"size": [
|
||||
225.59765625,
|
||||
206
|
||||
],
|
||||
"flags": {},
|
||||
"order": 3,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "video_info",
|
||||
"type": "VHS_VIDEOINFO",
|
||||
"link": 4
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "source_fps🟨",
|
||||
"type": "FLOAT",
|
||||
"links": [
|
||||
5
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "source_frame_count🟨",
|
||||
"type": "INT",
|
||||
"links": null
|
||||
},
|
||||
{
|
||||
"name": "source_duration🟨",
|
||||
"type": "FLOAT",
|
||||
"links": null
|
||||
},
|
||||
{
|
||||
"name": "source_width🟨",
|
||||
"type": "INT",
|
||||
"links": null
|
||||
},
|
||||
{
|
||||
"name": "source_height🟨",
|
||||
"type": "INT",
|
||||
"links": null
|
||||
},
|
||||
{
|
||||
"name": "loaded_fps🟦",
|
||||
"type": "FLOAT",
|
||||
"links": null
|
||||
},
|
||||
{
|
||||
"name": "loaded_frame_count🟦",
|
||||
"type": "INT",
|
||||
"links": null
|
||||
},
|
||||
{
|
||||
"name": "loaded_duration🟦",
|
||||
"type": "FLOAT",
|
||||
"links": null
|
||||
},
|
||||
{
|
||||
"name": "loaded_width🟦",
|
||||
"type": "INT",
|
||||
"links": null
|
||||
},
|
||||
{
|
||||
"name": "loaded_height🟦",
|
||||
"type": "INT",
|
||||
"links": null
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "comfyui-videohelpersuite",
|
||||
"ver": "1.7.7",
|
||||
"Node name for S&R": "VHS_VideoInfo"
|
||||
},
|
||||
"widgets_values": {}
|
||||
},
|
||||
{
|
||||
"id": 9,
|
||||
"type": "Bria_Video_Increase_Resolution_fal",
|
||||
"pos": [
|
||||
-3215.152099609375,
|
||||
-2325.123291015625
|
||||
],
|
||||
"size": [
|
||||
319.8667907714844,
|
||||
106
|
||||
],
|
||||
"flags": {},
|
||||
"order": 1,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "video",
|
||||
"shape": 7,
|
||||
"type": "VIDEO",
|
||||
"link": 9
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "video_url",
|
||||
"type": "STRING",
|
||||
"links": [
|
||||
11
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "fal-api",
|
||||
"ver": "1fb220258fdc77e1acabdd7f1cb32cf8194f57a9",
|
||||
"Node name for S&R": "Bria_Video_Increase_Resolution_fal"
|
||||
},
|
||||
"widgets_values": [
|
||||
2,
|
||||
"",
|
||||
"mp4_h264"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 4,
|
||||
"type": "LoadVideoURL",
|
||||
"pos": [
|
||||
-2754.03515625,
|
||||
-2394.18994140625
|
||||
],
|
||||
"size": [
|
||||
270,
|
||||
266
|
||||
],
|
||||
"flags": {},
|
||||
"order": 2,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "url",
|
||||
"type": "STRING",
|
||||
"widget": {
|
||||
"name": "url"
|
||||
},
|
||||
"link": 11
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "frames",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
27
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "frame_count",
|
||||
"type": "INT",
|
||||
"links": null
|
||||
},
|
||||
{
|
||||
"name": "video_info",
|
||||
"type": "VHS_VIDEOINFO",
|
||||
"links": [
|
||||
4
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "fal-api",
|
||||
"ver": "1fb220258fdc77e1acabdd7f1cb32cf8194f57a9",
|
||||
"Node name for S&R": "LoadVideoURL"
|
||||
},
|
||||
"widgets_values": [
|
||||
"https://example.com/video.mp4",
|
||||
0,
|
||||
"Disabled",
|
||||
512,
|
||||
512,
|
||||
0,
|
||||
0,
|
||||
1
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 5,
|
||||
"type": "VHS_VideoCombine",
|
||||
"pos": [
|
||||
-1946.611328125,
|
||||
-2388.9658203125
|
||||
],
|
||||
"size": [
|
||||
214.7587890625,
|
||||
460.36083984375
|
||||
],
|
||||
"flags": {},
|
||||
"order": 4,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "images",
|
||||
"type": "IMAGE",
|
||||
"link": 27
|
||||
},
|
||||
{
|
||||
"name": "audio",
|
||||
"shape": 7,
|
||||
"type": "AUDIO",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "meta_batch",
|
||||
"shape": 7,
|
||||
"type": "VHS_BatchManager",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "vae",
|
||||
"shape": 7,
|
||||
"type": "VAE",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "frame_rate",
|
||||
"type": "FLOAT",
|
||||
"widget": {
|
||||
"name": "frame_rate"
|
||||
},
|
||||
"link": 5
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "Filenames",
|
||||
"type": "VHS_FILENAMES",
|
||||
"links": []
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "comfyui-videohelpersuite",
|
||||
"ver": "1.7.7",
|
||||
"Node name for S&R": "VHS_VideoCombine"
|
||||
},
|
||||
"widgets_values": {
|
||||
"frame_rate": 8,
|
||||
"loop_count": 0,
|
||||
"filename_prefix": "AnimateDiff",
|
||||
"format": "video/h264-mp4",
|
||||
"pix_fmt": "yuv420p",
|
||||
"crf": 19,
|
||||
"save_metadata": true,
|
||||
"trim_to_audio": false,
|
||||
"pingpong": false,
|
||||
"save_output": true,
|
||||
"videopreview": {
|
||||
"hidden": false,
|
||||
"paused": false,
|
||||
"params": {
|
||||
"filename": "AnimateDiff_00049.mp4",
|
||||
"subfolder": "",
|
||||
"type": "output",
|
||||
"format": "video/h264-mp4",
|
||||
"frame_rate": 25,
|
||||
"workflow": "AnimateDiff_00049.png",
|
||||
"fullpath": "/root/ComfyUI/output/AnimateDiff_00049.mp4"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
],
|
||||
"links": [
|
||||
[
|
||||
4,
|
||||
4,
|
||||
2,
|
||||
6,
|
||||
0,
|
||||
"VHS_VIDEOINFO"
|
||||
],
|
||||
[
|
||||
5,
|
||||
6,
|
||||
0,
|
||||
5,
|
||||
4,
|
||||
"FLOAT"
|
||||
],
|
||||
[
|
||||
9,
|
||||
3,
|
||||
0,
|
||||
9,
|
||||
0,
|
||||
"VIDEO"
|
||||
],
|
||||
[
|
||||
11,
|
||||
9,
|
||||
0,
|
||||
4,
|
||||
0,
|
||||
"STRING"
|
||||
],
|
||||
[
|
||||
27,
|
||||
4,
|
||||
0,
|
||||
5,
|
||||
0,
|
||||
"IMAGE"
|
||||
]
|
||||
],
|
||||
"groups": [],
|
||||
"config": {},
|
||||
"extra": {
|
||||
"ds": {
|
||||
"scale": 1.1167815779424788,
|
||||
"offset": [
|
||||
3727.7338670706927,
|
||||
2641.8882544232415
|
||||
]
|
||||
},
|
||||
"frontendVersion": "1.27.10",
|
||||
"VHS_latentpreview": false,
|
||||
"VHS_latentpreviewrate": 0,
|
||||
"VHS_MetadataImage": true,
|
||||
"VHS_KeepIntermediate": true
|
||||
},
|
||||
"version": 0.4
|
||||
}
|
||||
@@ -0,0 +1,376 @@
|
||||
{
|
||||
"id": "1c871709-2293-44c4-9ba3-e1b5be72bbbe",
|
||||
"revision": 0,
|
||||
"last_node_id": 19,
|
||||
"last_link_id": 30,
|
||||
"nodes": [
|
||||
{
|
||||
"id": 6,
|
||||
"type": "VHS_VideoInfo",
|
||||
"pos": [
|
||||
-2305.713623046875,
|
||||
-2289.08837890625
|
||||
],
|
||||
"size": [
|
||||
225.59765625,
|
||||
206
|
||||
],
|
||||
"flags": {},
|
||||
"order": 3,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "video_info",
|
||||
"type": "VHS_VIDEOINFO",
|
||||
"link": 4
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "source_fps🟨",
|
||||
"type": "FLOAT",
|
||||
"links": [
|
||||
5
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "source_frame_count🟨",
|
||||
"type": "INT",
|
||||
"links": null
|
||||
},
|
||||
{
|
||||
"name": "source_duration🟨",
|
||||
"type": "FLOAT",
|
||||
"links": null
|
||||
},
|
||||
{
|
||||
"name": "source_width🟨",
|
||||
"type": "INT",
|
||||
"links": null
|
||||
},
|
||||
{
|
||||
"name": "source_height🟨",
|
||||
"type": "INT",
|
||||
"links": null
|
||||
},
|
||||
{
|
||||
"name": "loaded_fps🟦",
|
||||
"type": "FLOAT",
|
||||
"links": null
|
||||
},
|
||||
{
|
||||
"name": "loaded_frame_count🟦",
|
||||
"type": "INT",
|
||||
"links": null
|
||||
},
|
||||
{
|
||||
"name": "loaded_duration🟦",
|
||||
"type": "FLOAT",
|
||||
"links": null
|
||||
},
|
||||
{
|
||||
"name": "loaded_width🟦",
|
||||
"type": "INT",
|
||||
"links": null
|
||||
},
|
||||
{
|
||||
"name": "loaded_height🟦",
|
||||
"type": "INT",
|
||||
"links": null
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "comfyui-videohelpersuite",
|
||||
"ver": "1.7.7",
|
||||
"Node name for S&R": "VHS_VideoInfo"
|
||||
},
|
||||
"widgets_values": {}
|
||||
},
|
||||
{
|
||||
"id": 4,
|
||||
"type": "LoadVideoURL",
|
||||
"pos": [
|
||||
-2754.03515625,
|
||||
-2394.18994140625
|
||||
],
|
||||
"size": [
|
||||
270,
|
||||
266
|
||||
],
|
||||
"flags": {},
|
||||
"order": 2,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "url",
|
||||
"type": "STRING",
|
||||
"widget": {
|
||||
"name": "url"
|
||||
},
|
||||
"link": 30
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "frames",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
27
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "frame_count",
|
||||
"type": "INT",
|
||||
"links": null
|
||||
},
|
||||
{
|
||||
"name": "video_info",
|
||||
"type": "VHS_VIDEOINFO",
|
||||
"links": [
|
||||
4
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "fal-api",
|
||||
"ver": "1fb220258fdc77e1acabdd7f1cb32cf8194f57a9",
|
||||
"Node name for S&R": "LoadVideoURL"
|
||||
},
|
||||
"widgets_values": [
|
||||
"https://example.com/video.mp4",
|
||||
0,
|
||||
"Disabled",
|
||||
512,
|
||||
512,
|
||||
0,
|
||||
0,
|
||||
1
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 5,
|
||||
"type": "VHS_VideoCombine",
|
||||
"pos": [
|
||||
-1946.611328125,
|
||||
-2388.9658203125
|
||||
],
|
||||
"size": [
|
||||
214.7587890625,
|
||||
460.36083984375
|
||||
],
|
||||
"flags": {},
|
||||
"order": 4,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "images",
|
||||
"type": "IMAGE",
|
||||
"link": 27
|
||||
},
|
||||
{
|
||||
"name": "audio",
|
||||
"shape": 7,
|
||||
"type": "AUDIO",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "meta_batch",
|
||||
"shape": 7,
|
||||
"type": "VHS_BatchManager",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "vae",
|
||||
"shape": 7,
|
||||
"type": "VAE",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "frame_rate",
|
||||
"type": "FLOAT",
|
||||
"widget": {
|
||||
"name": "frame_rate"
|
||||
},
|
||||
"link": 5
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "Filenames",
|
||||
"type": "VHS_FILENAMES",
|
||||
"links": []
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "comfyui-videohelpersuite",
|
||||
"ver": "1.7.7",
|
||||
"Node name for S&R": "VHS_VideoCombine"
|
||||
},
|
||||
"widgets_values": {
|
||||
"frame_rate": 8,
|
||||
"loop_count": 0,
|
||||
"filename_prefix": "AnimateDiff",
|
||||
"format": "video/h264-mp4",
|
||||
"pix_fmt": "yuv420p",
|
||||
"crf": 19,
|
||||
"save_metadata": true,
|
||||
"trim_to_audio": false,
|
||||
"pingpong": false,
|
||||
"save_output": true,
|
||||
"videopreview": {
|
||||
"hidden": false,
|
||||
"paused": false,
|
||||
"params": {
|
||||
"filename": "AnimateDiff_00050.mp4",
|
||||
"subfolder": "",
|
||||
"type": "output",
|
||||
"format": "video/h264-mp4",
|
||||
"frame_rate": 25,
|
||||
"workflow": "AnimateDiff_00050.png",
|
||||
"fullpath": "/root/ComfyUI/output/AnimateDiff_00050.mp4"
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 3,
|
||||
"type": "LoadVideo",
|
||||
"pos": [
|
||||
-3565.794921875,
|
||||
-2337.682861328125
|
||||
],
|
||||
"size": [
|
||||
274.080078125,
|
||||
232.1231231689453
|
||||
],
|
||||
"flags": {},
|
||||
"order": 0,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "VIDEO",
|
||||
"type": "VIDEO",
|
||||
"links": [
|
||||
29
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "comfy-core",
|
||||
"ver": "0.3.59",
|
||||
"Node name for S&R": "LoadVideo"
|
||||
},
|
||||
"widgets_values": [
|
||||
"AnimateDiff_00039.mp4",
|
||||
"image"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 18,
|
||||
"type": "Seedvr_Upscale_Video_fal",
|
||||
"pos": [
|
||||
-3179.730224609375,
|
||||
-2341.025390625
|
||||
],
|
||||
"size": [
|
||||
270,
|
||||
274
|
||||
],
|
||||
"flags": {},
|
||||
"order": 1,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "video",
|
||||
"shape": 7,
|
||||
"type": "VIDEO",
|
||||
"link": 29
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "video_url",
|
||||
"type": "STRING",
|
||||
"links": [
|
||||
30
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "fal-api",
|
||||
"ver": "e650b470da94100c9315922f37156315e2eea42f",
|
||||
"Node name for S&R": "Seedvr_Upscale_Video_fal"
|
||||
},
|
||||
"widgets_values": [
|
||||
2,
|
||||
"",
|
||||
"factor",
|
||||
"1080p",
|
||||
0.1,
|
||||
"high",
|
||||
"balanced",
|
||||
"X264 (.mp4)"
|
||||
]
|
||||
}
|
||||
],
|
||||
"links": [
|
||||
[
|
||||
4,
|
||||
4,
|
||||
2,
|
||||
6,
|
||||
0,
|
||||
"VHS_VIDEOINFO"
|
||||
],
|
||||
[
|
||||
5,
|
||||
6,
|
||||
0,
|
||||
5,
|
||||
4,
|
||||
"FLOAT"
|
||||
],
|
||||
[
|
||||
27,
|
||||
4,
|
||||
0,
|
||||
5,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
29,
|
||||
3,
|
||||
0,
|
||||
18,
|
||||
0,
|
||||
"VIDEO"
|
||||
],
|
||||
[
|
||||
30,
|
||||
18,
|
||||
0,
|
||||
4,
|
||||
0,
|
||||
"STRING"
|
||||
]
|
||||
],
|
||||
"groups": [],
|
||||
"config": {},
|
||||
"extra": {
|
||||
"ds": {
|
||||
"scale": 1.1167815779424797,
|
||||
"offset": [
|
||||
3838.4549661482356,
|
||||
2565.8890490291633
|
||||
]
|
||||
},
|
||||
"frontendVersion": "1.27.10",
|
||||
"VHS_latentpreview": false,
|
||||
"VHS_latentpreviewrate": 0,
|
||||
"VHS_MetadataImage": true,
|
||||
"VHS_KeepIntermediate": true
|
||||
},
|
||||
"version": 0.4
|
||||
}
|
||||
@@ -0,0 +1,95 @@
|
||||
"""Billing nodes: account balance reporting with spend-guard visibility."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from .fal_utils import BillingUtils, SpendGuard, logger
|
||||
|
||||
_CATEGORY = "FAL/Platform"
|
||||
|
||||
_UNAVAILABLE_HINT = (
|
||||
"Hint: the balance API needs an API key with billing access "
|
||||
"(Authorization: Key ...); check your key at https://fal.ai/dashboard/keys."
|
||||
)
|
||||
|
||||
|
||||
def _guard_line(settings: dict[str, float]) -> str:
|
||||
"""One-line summary of the active spend-guard configuration."""
|
||||
budget = settings.get("session_budget_usd") or 0.0
|
||||
floor = settings.get("min_balance_usd") or 0.0
|
||||
if budget <= 0 and floor <= 0:
|
||||
return (
|
||||
"Spend guard: disabled (set [spend_guard] session_budget_usd / "
|
||||
"min_balance_usd in config.ini to enable)"
|
||||
)
|
||||
parts = []
|
||||
if budget > 0:
|
||||
parts.append(f"session budget ${budget:.2f}")
|
||||
if floor > 0:
|
||||
parts.append(f"min balance ${floor:.2f}")
|
||||
return f"Spend guard: {', '.join(parts)}"
|
||||
|
||||
|
||||
def _build_report(balance: float | None, settings: dict[str, float]) -> str:
|
||||
"""Human-readable balance report. Never raises."""
|
||||
if balance is not None:
|
||||
lines = [f"fal account balance: ${balance:,.2f}"]
|
||||
else:
|
||||
lines = ["fal account balance: unavailable", _UNAVAILABLE_HINT]
|
||||
return "\n".join([*lines, _guard_line(settings)])
|
||||
|
||||
|
||||
class FalBalance:
|
||||
"""Report the fal.ai account credit balance and active spend-guard limits."""
|
||||
|
||||
RETURN_TYPES = ("STRING", "FLOAT")
|
||||
RETURN_NAMES = ("report", "balance_usd")
|
||||
FUNCTION = "check"
|
||||
CATEGORY = _CATEGORY
|
||||
OUTPUT_NODE = True
|
||||
DESCRIPTION = (
|
||||
"Fetch your fal.ai account credit balance and show the active "
|
||||
"spend-guard settings. Never fails: balance_usd is -1.0 when the "
|
||||
"balance API is unavailable."
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> dict[str, Any]:
|
||||
return {
|
||||
"required": {},
|
||||
"optional": {
|
||||
"force_refresh": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": False,
|
||||
"tooltip": "Bypass the 60s balance cache and query fal again",
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def IS_CHANGED(cls, **kwargs: Any) -> Any:
|
||||
# The balance changes outside the graph; always re-run.
|
||||
return float("nan")
|
||||
|
||||
def check(self, force_refresh: bool = False) -> tuple[str, float]:
|
||||
try:
|
||||
balance = BillingUtils.get_balance(force=bool(force_refresh))
|
||||
settings = SpendGuard.settings()
|
||||
report = _build_report(balance, settings)
|
||||
except Exception as err: # This node must never fail the graph.
|
||||
logger.warning("FalBalance: could not build balance report: %s", err)
|
||||
balance = None
|
||||
report = f"fal account balance: unavailable ({err})\n{_UNAVAILABLE_HINT}"
|
||||
return (report, float(balance) if balance is not None else -1.0)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"FalBalance_fal": FalBalance,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"FalBalance_fal": "Fal Account Balance (fal)",
|
||||
}
|
||||
@@ -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)",
|
||||
}
|
||||
@@ -0,0 +1,25 @@
|
||||
"""Dynamic fal.ai node package: auto-generated nodes from the model registry."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
def get_dynamic_mappings() -> tuple[dict[str, type], dict[str, str]]:
|
||||
"""Return (NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS) for dynamic nodes.
|
||||
|
||||
Never raises: any failure (missing registry, missing utils facade, bad
|
||||
schema) yields empty mappings so static node loading is never affected.
|
||||
"""
|
||||
try:
|
||||
from .registry_loader import load_dynamic_mappings
|
||||
|
||||
return load_dynamic_mappings()
|
||||
except Exception as err:
|
||||
import logging
|
||||
|
||||
logging.getLogger(__name__).error(
|
||||
"Failed to load dynamic fal nodes: %s", err
|
||||
)
|
||||
return {}, {}
|
||||
|
||||
|
||||
__all__ = ["get_dynamic_mappings"]
|
||||
@@ -0,0 +1,104 @@
|
||||
{
|
||||
"version": 1,
|
||||
"generated_at": "2026-07-02T00:00:00Z",
|
||||
"model_count": 5,
|
||||
"models": [
|
||||
{
|
||||
"endpoint_id": "fal-ai/flux/dev",
|
||||
"title": "FLUX.1 [dev]",
|
||||
"category": "text-to-image",
|
||||
"lab": "Black Forest Labs",
|
||||
"family": "flux",
|
||||
"description": "FLUX.1 [dev] is a 12 billion parameter flow transformer for text-to-image generation.",
|
||||
"pricing": "$0.025 per megapixel",
|
||||
"published_at": "2024-08-01",
|
||||
"thumbnail": null,
|
||||
"inputs": [
|
||||
{"name": "prompt", "type": "string", "required": true, "default": null, "enum": null, "min": null, "max": null, "description": "The prompt to generate an image from", "media_kind": null, "is_list": false, "multiline": true, "has_custom_size": false},
|
||||
{"name": "image_size", "type": "enum", "required": false, "default": "landscape_4_3", "enum": ["square_hd", "square", "portrait_4_3", "portrait_16_9", "landscape_4_3", "landscape_16_9", "custom_size"], "min": null, "max": null, "description": "The size of the generated image", "media_kind": null, "is_list": false, "multiline": false, "has_custom_size": true},
|
||||
{"name": "num_inference_steps", "type": "integer", "required": false, "default": 28, "enum": null, "min": 1, "max": 50, "description": "Number of inference steps", "media_kind": null, "is_list": false, "multiline": false, "has_custom_size": false},
|
||||
{"name": "guidance_scale", "type": "number", "required": false, "default": 3.5, "enum": null, "min": 1, "max": 20, "description": "CFG scale", "media_kind": null, "is_list": false, "multiline": false, "has_custom_size": false},
|
||||
{"name": "seed", "type": "integer", "required": false, "default": null, "enum": null, "min": null, "max": null, "description": "Random seed", "media_kind": null, "is_list": false, "multiline": false, "has_custom_size": false},
|
||||
{"name": "num_images", "type": "integer", "required": false, "default": 1, "enum": null, "min": 1, "max": 4, "description": "Number of images to generate", "media_kind": null, "is_list": false, "multiline": false, "has_custom_size": false},
|
||||
{"name": "enable_safety_checker", "type": "boolean", "required": false, "default": true, "enum": null, "min": null, "max": null, "description": "Enable the safety checker", "media_kind": null, "is_list": false, "multiline": false, "has_custom_size": false},
|
||||
{"name": "loras", "type": "json", "required": false, "default": null, "enum": null, "min": null, "max": null, "description": "LoRA weights to apply", "media_kind": null, "is_list": true, "multiline": false, "has_custom_size": false}
|
||||
],
|
||||
"output_kind": "images",
|
||||
"output_props": ["images"]
|
||||
},
|
||||
{
|
||||
"endpoint_id": "fal-ai/kling-video/v2/master/image-to-video",
|
||||
"title": "Kling 2.0 Master",
|
||||
"category": "image-to-video",
|
||||
"lab": "Kuaishou",
|
||||
"family": "kling-video",
|
||||
"description": "Generate video clips from an image using Kling 2.0 Master.",
|
||||
"pricing": "$1.40 per 5s video",
|
||||
"published_at": "2025-04-15",
|
||||
"thumbnail": null,
|
||||
"inputs": [
|
||||
{"name": "prompt", "type": "string", "required": true, "default": null, "enum": null, "min": null, "max": null, "description": "Motion prompt", "media_kind": null, "is_list": false, "multiline": true, "has_custom_size": false},
|
||||
{"name": "image_url", "type": "string", "required": true, "default": null, "enum": null, "min": null, "max": null, "description": "Start frame image", "media_kind": "image", "is_list": false, "multiline": false, "has_custom_size": false},
|
||||
{"name": "duration", "type": "enum", "required": false, "default": "5", "enum": ["5", "10"], "min": null, "max": null, "description": "Duration of the video in seconds", "media_kind": null, "is_list": false, "multiline": false, "has_custom_size": false},
|
||||
{"name": "negative_prompt", "type": "string", "required": false, "default": "blur, distort, and low quality", "enum": null, "min": null, "max": null, "description": "Negative prompt", "media_kind": null, "is_list": false, "multiline": false, "has_custom_size": false},
|
||||
{"name": "cfg_scale", "type": "number", "required": false, "default": 0.5, "enum": null, "min": 0, "max": 1, "description": "CFG scale", "media_kind": null, "is_list": false, "multiline": false, "has_custom_size": false}
|
||||
],
|
||||
"output_kind": "video",
|
||||
"output_props": ["video"]
|
||||
},
|
||||
{
|
||||
"endpoint_id": "fal-ai/video-upscaler",
|
||||
"title": "Video Upscaler",
|
||||
"category": "video-to-video",
|
||||
"lab": "fal",
|
||||
"family": "video-upscaler",
|
||||
"description": "Upscale videos by a given factor.",
|
||||
"pricing": "$0.02 per video second",
|
||||
"published_at": "2024-11-01",
|
||||
"thumbnail": null,
|
||||
"inputs": [
|
||||
{"name": "video_url", "type": "string", "required": true, "default": null, "enum": null, "min": null, "max": null, "description": "Video to upscale", "media_kind": "video", "is_list": false, "multiline": false, "has_custom_size": false},
|
||||
{"name": "scale", "type": "number", "required": false, "default": 2, "enum": null, "min": 1, "max": 4, "description": "Upscale factor", "media_kind": null, "is_list": false, "multiline": false, "has_custom_size": false}
|
||||
],
|
||||
"output_kind": "video",
|
||||
"output_props": ["video"]
|
||||
},
|
||||
{
|
||||
"endpoint_id": "fal-ai/kokoro/american-english",
|
||||
"title": "Kokoro TTS",
|
||||
"category": "text-to-speech",
|
||||
"lab": "Kokoro",
|
||||
"family": "kokoro",
|
||||
"description": "Fast and expressive American English text-to-speech.",
|
||||
"pricing": "$0.02 per 1000 characters",
|
||||
"published_at": "2025-01-20",
|
||||
"thumbnail": null,
|
||||
"inputs": [
|
||||
{"name": "prompt", "type": "string", "required": true, "default": null, "enum": null, "min": null, "max": null, "description": "Text to convert to speech", "media_kind": null, "is_list": false, "multiline": true, "has_custom_size": false},
|
||||
{"name": "voice", "type": "enum", "required": false, "default": "af_heart", "enum": ["af_heart", "af_bella", "am_adam", "am_echo"], "min": null, "max": null, "description": "Voice to use", "media_kind": null, "is_list": false, "multiline": false, "has_custom_size": false},
|
||||
{"name": "speed", "type": "number", "required": false, "default": 1.0, "enum": null, "min": 0.5, "max": 2.0, "description": "Speech speed", "media_kind": null, "is_list": false, "multiline": false, "has_custom_size": false}
|
||||
],
|
||||
"output_kind": "audio",
|
||||
"output_props": ["audio"]
|
||||
},
|
||||
{
|
||||
"endpoint_id": "tripo3d/tripo/v2.5/image-to-3d",
|
||||
"title": "Tripo3D v2.5",
|
||||
"category": "image-to-3d",
|
||||
"lab": "Tripo",
|
||||
"family": "tripo",
|
||||
"description": "Generate a textured 3D mesh from a single image.",
|
||||
"pricing": "$0.20 per generation",
|
||||
"published_at": "2025-02-10",
|
||||
"thumbnail": null,
|
||||
"inputs": [
|
||||
{"name": "image_url", "type": "string", "required": true, "default": null, "enum": null, "min": null, "max": null, "description": "Input image", "media_kind": "image", "is_list": false, "multiline": false, "has_custom_size": false},
|
||||
{"name": "image_urls", "type": "array", "required": false, "default": null, "enum": null, "min": null, "max": null, "description": "Optional multi-view images", "media_kind": "image", "is_list": true, "multiline": false, "has_custom_size": false},
|
||||
{"name": "texture", "type": "enum", "required": false, "default": "standard", "enum": ["no", "standard", "HD"], "min": null, "max": null, "description": "Texture quality", "media_kind": null, "is_list": false, "multiline": false, "has_custom_size": false},
|
||||
{"name": "seed", "type": "integer", "required": false, "default": null, "enum": null, "min": null, "max": null, "description": "Random seed", "media_kind": null, "is_list": false, "multiline": false, "has_custom_size": false}
|
||||
],
|
||||
"output_kind": "file",
|
||||
"output_props": ["model_mesh"]
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,485 @@
|
||||
"""Self-test for the dynamic node package. Stdlib only; stubs the utils facade.
|
||||
|
||||
Run: python3 nodes/dynamic/_selftest.py
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import sys
|
||||
import types
|
||||
from pathlib import Path
|
||||
|
||||
PACKAGE_DIR = Path(__file__).resolve().parent
|
||||
NODES_DIR = PACKAGE_DIR.parent
|
||||
REPO_ROOT = NODES_DIR.parent
|
||||
REAL_REGISTRY = REPO_ROOT / "data" / "fal_registry.json"
|
||||
|
||||
PKG = "falapi_nodes"
|
||||
|
||||
|
||||
def _install_stub_facade() -> types.ModuleType:
|
||||
"""Install a stub falapi_nodes.fal_utils satisfying the facade contract."""
|
||||
stub = types.ModuleType(f"{PKG}.fal_utils")
|
||||
|
||||
class FalApiError(Exception):
|
||||
def __init__(self, endpoint, message):
|
||||
super().__init__(f"[{endpoint}] {message}")
|
||||
self.endpoint = endpoint
|
||||
self.message = message
|
||||
|
||||
class FalConfig:
|
||||
def get_setting(self, section, name, default=None):
|
||||
return default
|
||||
|
||||
class ImageUtils:
|
||||
@staticmethod
|
||||
def upload_image(tensor):
|
||||
return "https://stub.fal.media/image.png"
|
||||
|
||||
@staticmethod
|
||||
def prepare_images(images):
|
||||
return ["https://stub.fal.media/1.png", "https://stub.fal.media/2.png"]
|
||||
|
||||
class ResultProcessor:
|
||||
@staticmethod
|
||||
def process_image_result(result):
|
||||
return ("IMAGE_TENSOR",)
|
||||
|
||||
@staticmethod
|
||||
def process_single_image_result(result):
|
||||
return ("IMAGE_TENSOR",)
|
||||
|
||||
class ApiHandler:
|
||||
last_call = None
|
||||
|
||||
@staticmethod
|
||||
def submit_and_get_result(endpoint, arguments):
|
||||
ApiHandler.last_call = (endpoint, arguments)
|
||||
return _CANNED_RESULTS.get(endpoint, {"ok": True})
|
||||
|
||||
class MediaUtils:
|
||||
@staticmethod
|
||||
def video_from_url(url):
|
||||
return "VIDEO_OBJ"
|
||||
|
||||
@staticmethod
|
||||
def audio_from_url(url):
|
||||
return {"waveform": None, "sample_rate": 44100}
|
||||
|
||||
@staticmethod
|
||||
def upload_video(video):
|
||||
return "https://stub.fal.media/video.mp4"
|
||||
|
||||
@staticmethod
|
||||
def upload_audio(audio):
|
||||
return "https://stub.fal.media/audio.wav"
|
||||
|
||||
@staticmethod
|
||||
def download_url_to_temp(url, suffix):
|
||||
return "/tmp/stub" + suffix
|
||||
|
||||
stub.FalApiError = FalApiError
|
||||
stub.FalConfig = FalConfig
|
||||
stub.ImageUtils = ImageUtils
|
||||
stub.ResultProcessor = ResultProcessor
|
||||
stub.ApiHandler = ApiHandler
|
||||
stub.MediaUtils = MediaUtils
|
||||
stub.logger = logging.getLogger("fal_stub")
|
||||
sys.modules[stub.__name__] = stub
|
||||
return stub
|
||||
|
||||
|
||||
_CANNED_RESULTS = {
|
||||
"fal-ai/flux/dev": {"images": [{"url": "https://x/i.png"}], "seed": 1},
|
||||
"fal-ai/kling-video/v2/master/image-to-video": {
|
||||
"video": {"url": "https://x/v.mp4"}
|
||||
},
|
||||
"fal-ai/video-upscaler": {"video": {"url": "https://x/up.mp4"}},
|
||||
"fal-ai/kokoro/american-english": {"audio": {"url": "https://x/a.wav"}},
|
||||
"tripo3d/tripo/v2.5/image-to-3d": {"model_mesh": {"url": "https://x/m.glb"}},
|
||||
}
|
||||
|
||||
|
||||
def _install_package() -> None:
|
||||
pkg = types.ModuleType(PKG)
|
||||
pkg.__path__ = [str(NODES_DIR)]
|
||||
sys.modules[PKG] = pkg
|
||||
|
||||
|
||||
def _load_models(path: Path):
|
||||
with open(path, encoding="utf-8") as handle:
|
||||
return json.load(handle)["models"]
|
||||
|
||||
|
||||
def _check_registry(models, factory, outputs, label):
|
||||
keys = set()
|
||||
names = set()
|
||||
skipped = {}
|
||||
built = 0
|
||||
for model in models:
|
||||
try:
|
||||
cls = factory.build_node_class(model)
|
||||
input_types = cls.INPUT_TYPES()
|
||||
assert isinstance(input_types, dict) and "required" in input_types
|
||||
assert "force_rerun" in input_types.get("optional", {})
|
||||
kind = model.get("output_kind", "json")
|
||||
expected = outputs.RETURN_SPECS.get(kind, outputs.RETURN_SPECS["json"])
|
||||
assert cls.RETURN_TYPES == expected[0], (
|
||||
f"RETURN_TYPES mismatch for {model['endpoint_id']}"
|
||||
)
|
||||
assert cls.RETURN_NAMES == expected[1]
|
||||
key = factory.node_key(model)
|
||||
assert key not in keys, f"key collision: {key}"
|
||||
keys.add(key)
|
||||
names.add(factory.build_display_name(model))
|
||||
built += 1
|
||||
except Exception as err:
|
||||
reason = type(err).__name__ + ": " + str(err)[:80]
|
||||
skipped[reason] = skipped.get(reason, 0) + 1
|
||||
print(f"[{label}] built={built} skipped={sum(skipped.values())}")
|
||||
if skipped:
|
||||
print(f"[{label}] skip reasons histogram:")
|
||||
for reason, count in sorted(skipped.items(), key=lambda kv: -kv[1]):
|
||||
print(f" {count:4d} {reason}")
|
||||
return built, skipped
|
||||
|
||||
|
||||
def _test_fixture_behaviour(dyn, stub):
|
||||
from importlib import import_module
|
||||
|
||||
arguments = import_module(f"{PKG}.dynamic.arguments")
|
||||
factory = import_module(f"{PKG}.dynamic.factory")
|
||||
import_module(f"{PKG}.dynamic.outputs")
|
||||
|
||||
models = _load_models(PACKAGE_DIR / "_fixture_registry.json")
|
||||
by_id = {m["endpoint_id"]: m for m in models}
|
||||
|
||||
# --- arguments: custom_size, seed=-1 omitted, empty json skipped ---
|
||||
flux = by_id["fal-ai/flux/dev"]
|
||||
kwargs = {
|
||||
"prompt": "a cat",
|
||||
"image_size": "custom_size",
|
||||
"width": 512,
|
||||
"height": 768,
|
||||
"num_inference_steps": 28,
|
||||
"guidance_scale": 3.5,
|
||||
"seed": -1,
|
||||
"num_images": 1,
|
||||
"enable_safety_checker": True,
|
||||
"loras": "",
|
||||
"force_rerun": False,
|
||||
}
|
||||
kwargs_snapshot = dict(kwargs)
|
||||
args = arguments.build_arguments(flux, kwargs)
|
||||
assert args["image_size"] == {"width": 512, "height": 768}, args
|
||||
assert "seed" not in args and "loras" not in args and "force_rerun" not in args
|
||||
assert kwargs == kwargs_snapshot, "kwargs were mutated"
|
||||
|
||||
# seed forwarded when != -1; enum passthrough
|
||||
args2 = arguments.build_arguments(flux, {**kwargs, "seed": 42, "image_size": "square"})
|
||||
assert args2["seed"] == 42 and args2["image_size"] == "square"
|
||||
|
||||
# invalid json raises FalApiError
|
||||
try:
|
||||
arguments.build_arguments(flux, {**kwargs, "loras": "{not json"})
|
||||
raise AssertionError("expected FalApiError for bad JSON")
|
||||
except stub.FalApiError:
|
||||
pass
|
||||
|
||||
# valid json parsed
|
||||
args3 = arguments.build_arguments(flux, {**kwargs, "loras": '[{"path": "x"}]'})
|
||||
assert args3["loras"] == [{"path": "x"}]
|
||||
|
||||
# --- media uploads ---
|
||||
kling = by_id["fal-ai/kling-video/v2/master/image-to-video"]
|
||||
kargs = arguments.build_arguments(
|
||||
kling, {"prompt": "move", "image_url": "TENSOR", "duration": "5",
|
||||
"negative_prompt": "", "cfg_scale": 0.5}
|
||||
)
|
||||
assert kargs["image_url"] == "https://stub.fal.media/image.png"
|
||||
assert "negative_prompt" not in kargs # optional empty string skipped
|
||||
|
||||
tripo = by_id["tripo3d/tripo/v2.5/image-to-3d"]
|
||||
targs = arguments.build_arguments(
|
||||
tripo, {"image_url": "TENSOR", "image_urls": "BATCH", "texture": "HD", "seed": -1}
|
||||
)
|
||||
assert targs["image_urls"] == [
|
||||
"https://stub.fal.media/1.png",
|
||||
"https://stub.fal.media/2.png",
|
||||
]
|
||||
|
||||
upscaler = by_id["fal-ai/video-upscaler"]
|
||||
uargs = arguments.build_arguments(upscaler, {"video_url": "VIDEO_OBJ", "scale": 2.0})
|
||||
assert uargs["video_url"] == "https://stub.fal.media/video.mp4"
|
||||
|
||||
# --- end-to-end run() per output kind ---
|
||||
flux_node = factory.build_node_class(flux)()
|
||||
assert flux_node.run(**kwargs) == ("IMAGE_TENSOR", "https://x/i.png")
|
||||
|
||||
kling_node = factory.build_node_class(kling)()
|
||||
out = kling_node.run(prompt="move", image_url="TENSOR", duration="5",
|
||||
negative_prompt="", cfg_scale=0.5)
|
||||
assert out == ("VIDEO_OBJ", "https://x/v.mp4"), out
|
||||
|
||||
tts = by_id["fal-ai/kokoro/american-english"]
|
||||
tts_node = factory.build_node_class(tts)()
|
||||
audio_out = tts_node.run(prompt="hello", voice="af_heart", speed=1.0)
|
||||
assert audio_out[1] == "https://x/a.wav" and isinstance(audio_out[0], dict)
|
||||
|
||||
tripo_node = factory.build_node_class(tripo)()
|
||||
file_out = tripo_node.run(image_url="TENSOR", texture="HD", seed=-1)
|
||||
assert file_out == ("https://x/m.glb",), file_out
|
||||
|
||||
# --- IS_CHANGED semantics ---
|
||||
cls = factory.build_node_class(flux)
|
||||
h1 = cls.IS_CHANGED(prompt="a", force_rerun=False)
|
||||
h2 = cls.IS_CHANGED(prompt="a", force_rerun=False)
|
||||
h3 = cls.IS_CHANGED(prompt="b", force_rerun=False)
|
||||
nan = cls.IS_CHANGED(prompt="a", force_rerun=True)
|
||||
assert h1 == h2 and h1 != h3 and nan != nan # nan != nan
|
||||
|
||||
# --- loader end-to-end (fixture fallback path) ---
|
||||
classes, display = dyn.get_dynamic_mappings()
|
||||
assert "FalAnyEndpoint_fal" in classes
|
||||
assert len(classes) == len(display)
|
||||
assert len(set(display.values())) == len(display), "display name collision"
|
||||
if not REAL_REGISTRY.is_file():
|
||||
assert len(classes) == 6, f"expected 5 fixture + any-endpoint, got {len(classes)}"
|
||||
|
||||
# --- any endpoint node ---
|
||||
any_cls = classes["FalAnyEndpoint_fal"]
|
||||
node = any_cls()
|
||||
any_cls.INPUT_TYPES()
|
||||
res = node.run(
|
||||
endpoint_id="fal-ai/flux/dev",
|
||||
arguments_json='{"prompt": "hi", "image_url": "should-be-overridden"}',
|
||||
image="TENSOR",
|
||||
image_2="TENSOR2",
|
||||
seed=7,
|
||||
)
|
||||
endpoint, sent = sys.modules[f"{PKG}.fal_utils"].ApiHandler.last_call
|
||||
assert sent["image_url"] == "https://stub.fal.media/image.png" # media wins
|
||||
assert sent["image_urls"] == [
|
||||
"https://stub.fal.media/image.png",
|
||||
"https://stub.fal.media/image.png",
|
||||
]
|
||||
assert sent["seed"] == 7 and sent["prompt"] == "hi"
|
||||
assert res[0] == "IMAGE_TENSOR" and json.loads(res[3])["seed"] == 1
|
||||
|
||||
print("[fixture] behaviour tests passed")
|
||||
|
||||
|
||||
def _file_input(name, media_kind, required=False, is_list=False, type_="string"):
|
||||
return {
|
||||
"name": name, "type": type_, "required": required, "default": None,
|
||||
"enum": None, "min": None, "max": None, "description": "",
|
||||
"media_kind": media_kind, "is_list": is_list, "multiline": False,
|
||||
"has_custom_size": False,
|
||||
}
|
||||
|
||||
|
||||
def _test_direct_url_passthrough(stub):
|
||||
from importlib import import_module
|
||||
|
||||
arguments = import_module(f"{PKG}.dynamic.arguments")
|
||||
factory = import_module(f"{PKG}.dynamic.factory")
|
||||
schema = import_module(f"{PKG}.dynamic.schema_to_inputs")
|
||||
|
||||
models = _load_models(PACKAGE_DIR / "_fixture_registry.json")
|
||||
by_id = {m["endpoint_id"]: m for m in models}
|
||||
kling = by_id["fal-ai/kling-video/v2/master/image-to-video"]
|
||||
tripo = by_id["tripo3d/tripo/v2.5/image-to-3d"]
|
||||
flux = by_id["fal-ai/flux/dev"]
|
||||
|
||||
# --- twins present and placed: required media → start of optional ---
|
||||
it = schema.build_input_types(kling)
|
||||
assert list(it["optional"])[0] == "image_url_direct_url"
|
||||
assert it["optional"]["image_url_direct_url"][0] == "STRING"
|
||||
assert "image_url_direct_url" not in it["required"]
|
||||
|
||||
uit = schema.build_input_types(by_id["fal-ai/video-upscaler"])
|
||||
assert list(uit["optional"])[0] == "video_url_direct_url"
|
||||
|
||||
# optional is_list media → twin immediately after its media input
|
||||
tit = schema.build_input_types(tripo)
|
||||
tkeys = list(tit["optional"])
|
||||
assert tkeys[0] == "image_url_direct_url"
|
||||
assert tkeys.index("image_urls_direct_url") == tkeys.index("image_urls") + 1
|
||||
assert "Comma" in tit["optional"]["image_urls_direct_url"][1]["tooltip"]
|
||||
|
||||
# no twin for text-only models or media_kind "file"
|
||||
fit = schema.build_input_types(flux)
|
||||
assert not any(k.endswith("_direct_url") for k in {**fit["required"], **fit["optional"]})
|
||||
file_model = {**kling, "inputs": [_file_input("doc_url", "file", required=True)]}
|
||||
ffit = schema.build_input_types(file_model)
|
||||
assert not any(k.endswith("_direct_url") for k in {**ffit["required"], **ffit["optional"]})
|
||||
|
||||
# collision guard: a literal *_direct_url input suppresses the generated twin
|
||||
collide = {**kling, "inputs": [
|
||||
_file_input("image_url", "image", required=True),
|
||||
_file_input("image_url_direct_url", None),
|
||||
]}
|
||||
cit = schema.build_input_types(collide)
|
||||
assert list(cit["optional"]).count("image_url_direct_url") == 1
|
||||
|
||||
# --- passthrough beats tensor upload; twin key never leaks ---
|
||||
kargs = arguments.build_arguments(kling, {
|
||||
"prompt": "move", "image_url": "TENSOR",
|
||||
"image_url_direct_url": " https://cdn.fal.media/start.png ",
|
||||
"duration": "5", "negative_prompt": "", "cfg_scale": 0.5,
|
||||
})
|
||||
assert kargs["image_url"] == "https://cdn.fal.media/start.png"
|
||||
assert not any(k.endswith("_direct_url") for k in kargs)
|
||||
|
||||
# media input None or absent: URL still wins
|
||||
for image_value in ({"image_url": None}, {}):
|
||||
k2 = arguments.build_arguments(
|
||||
kling,
|
||||
{"prompt": "m", "image_url_direct_url": "https://cdn.fal.media/s.png", **image_value},
|
||||
)
|
||||
assert k2["image_url"] == "https://cdn.fal.media/s.png"
|
||||
|
||||
# blank twin falls back to the normal upload path
|
||||
k3 = arguments.build_arguments(
|
||||
kling, {"prompt": "m", "image_url": "TENSOR", "image_url_direct_url": " "}
|
||||
)
|
||||
assert k3["image_url"] == "https://stub.fal.media/image.png"
|
||||
|
||||
# non-http(s) raises
|
||||
try:
|
||||
arguments.build_arguments(kling, {"prompt": "x", "image_url_direct_url": "ftp://nope"})
|
||||
raise AssertionError("expected FalApiError for non-http(s) direct URL")
|
||||
except stub.FalApiError:
|
||||
pass
|
||||
|
||||
# is_list twin: comma-separated string → list of URLs
|
||||
targs = arguments.build_arguments(tripo, {
|
||||
"image_url": "TENSOR", "texture": "HD", "seed": -1,
|
||||
"image_urls_direct_url": "https://a/1.png, https://a/2.png ,https://a/3.png",
|
||||
})
|
||||
assert targs["image_urls"] == ["https://a/1.png", "https://a/2.png", "https://a/3.png"]
|
||||
assert targs["image_url"] == "https://stub.fal.media/image.png"
|
||||
assert not any(k.endswith("_direct_url") for k in targs)
|
||||
|
||||
try:
|
||||
arguments.build_arguments(tripo, {"image_urls_direct_url": "https://a/1.png, nope"})
|
||||
raise AssertionError("expected FalApiError for bad URL in list")
|
||||
except stub.FalApiError:
|
||||
pass
|
||||
|
||||
# collision model: the literal input passes through as a plain string argument
|
||||
cargs = arguments.build_arguments(
|
||||
collide, {"image_url": "TENSOR", "image_url_direct_url": "not-a-url"}
|
||||
)
|
||||
assert cargs["image_url"] == "https://stub.fal.media/image.png"
|
||||
assert cargs["image_url_direct_url"] == "not-a-url"
|
||||
|
||||
# --- skip_cache plumbing: old signature tolerated, new one receives the flag ---
|
||||
node = factory.build_node_class(flux)()
|
||||
node.run(prompt="hi", force_rerun=True)
|
||||
assert len(stub.ApiHandler.last_call) == 2 # legacy stub: called without skip_cache
|
||||
|
||||
def with_skip(endpoint, arguments, timeout=None, skip_cache=False):
|
||||
stub.ApiHandler.last_call = (endpoint, arguments, skip_cache)
|
||||
return _CANNED_RESULTS.get(endpoint, {"ok": True})
|
||||
|
||||
original = stub.ApiHandler.submit_and_get_result
|
||||
stub.ApiHandler.submit_and_get_result = staticmethod(with_skip)
|
||||
try:
|
||||
node.run(prompt="hi", force_rerun=True)
|
||||
assert stub.ApiHandler.last_call[2] is True
|
||||
node.run(prompt="hi", force_rerun=False)
|
||||
assert stub.ApiHandler.last_call[2] is False
|
||||
finally:
|
||||
stub.ApiHandler.submit_and_get_result = original
|
||||
|
||||
print("[fixture] direct-url passthrough tests passed")
|
||||
|
||||
|
||||
def _twin_sweep(models, schema):
|
||||
"""Real-registry sweep: build every INPUT_TYPES, count twin coverage."""
|
||||
gained_nodes = 0
|
||||
twin_count = 0
|
||||
for model in models:
|
||||
input_types = schema.build_input_types(model)
|
||||
names = {inp["name"] for inp in model.get("inputs", [])}
|
||||
overlap = set(input_types["required"]) & set(input_types["optional"])
|
||||
assert not overlap, f"{model['endpoint_id']}: bucket overlap {overlap}"
|
||||
twins = [
|
||||
key for key in input_types["optional"]
|
||||
if key.endswith("_direct_url") and key not in names
|
||||
]
|
||||
if twins:
|
||||
gained_nodes += 1
|
||||
twin_count += len(twins)
|
||||
print(f"[real] direct-url twins: {twin_count} twin inputs across "
|
||||
f"{gained_nodes}/{len(models)} nodes")
|
||||
|
||||
|
||||
def _dump_samples(models, factory):
|
||||
samples = [
|
||||
("flux", "text-to-image"),
|
||||
("kling", "image-to-video"),
|
||||
(None, "text-to-speech"),
|
||||
(None, "image-to-3d"),
|
||||
(None, "video-to-video"),
|
||||
]
|
||||
seen = set()
|
||||
for hint, category in samples:
|
||||
candidates = [
|
||||
m for m in models
|
||||
if m.get("category") == category and m["endpoint_id"] not in seen
|
||||
]
|
||||
model = next(
|
||||
(m for m in candidates if hint and hint in m["endpoint_id"]),
|
||||
candidates[0] if candidates else None,
|
||||
)
|
||||
if model is None:
|
||||
print(f"-- no sample for {hint or category}")
|
||||
continue
|
||||
seen.add(model["endpoint_id"])
|
||||
cls = factory.build_node_class(model)
|
||||
print(f"\n-- INPUT_TYPES for {model['endpoint_id']} "
|
||||
f"({model.get('output_kind')}):")
|
||||
print(json.dumps(cls.INPUT_TYPES(), indent=2, default=str)[:2500])
|
||||
|
||||
|
||||
def main() -> int:
|
||||
logging.basicConfig(level=logging.INFO)
|
||||
_install_package()
|
||||
stub = _install_stub_facade()
|
||||
|
||||
from importlib import import_module
|
||||
|
||||
dyn = import_module(f"{PKG}.dynamic")
|
||||
factory = import_module(f"{PKG}.dynamic.factory")
|
||||
outputs = import_module(f"{PKG}.dynamic.outputs")
|
||||
|
||||
fixture_models = _load_models(PACKAGE_DIR / "_fixture_registry.json")
|
||||
built, _ = _check_registry(fixture_models, factory, outputs, "fixture")
|
||||
assert built == 5
|
||||
|
||||
_test_fixture_behaviour(dyn, stub)
|
||||
_test_direct_url_passthrough(stub)
|
||||
|
||||
if REAL_REGISTRY.is_file():
|
||||
schema = import_module(f"{PKG}.dynamic.schema_to_inputs")
|
||||
real_models = _load_models(REAL_REGISTRY)
|
||||
built, skipped = _check_registry(real_models, factory, outputs, "real")
|
||||
_twin_sweep(real_models, schema)
|
||||
classes, display = dyn.get_dynamic_mappings()
|
||||
print(f"[real] loader registered {len(classes)} nodes "
|
||||
f"(incl. any-endpoint), display names unique: "
|
||||
f"{len(set(display.values())) == len(display)}")
|
||||
_dump_samples(real_models, factory)
|
||||
else:
|
||||
print("[real] data/fal_registry.json not present; skipped real-registry checks")
|
||||
|
||||
print("\nSELFTEST OK")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
@@ -0,0 +1,263 @@
|
||||
"""Generic node that calls any fal.ai endpoint by id with free-form JSON arguments."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
from ..fal_utils import (
|
||||
ApiHandler,
|
||||
FalApiError,
|
||||
ImageUtils,
|
||||
MediaUtils,
|
||||
ResultProcessor,
|
||||
logger,
|
||||
)
|
||||
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:
|
||||
return {}
|
||||
try:
|
||||
parsed = json.loads(text)
|
||||
except ValueError as err:
|
||||
raise FalApiError(endpoint_id, f"Invalid JSON in 'arguments_json': {err}") from err
|
||||
if not isinstance(parsed, dict):
|
||||
raise FalApiError(endpoint_id, "'arguments_json' must be a JSON object")
|
||||
return parsed
|
||||
|
||||
|
||||
def _media_overlay(
|
||||
image: Any, image_2: Any, video: Any, audio: Any, seed: int
|
||||
) -> dict[str, Any]:
|
||||
overlay: dict[str, Any] = {}
|
||||
if image is not None:
|
||||
first_url = ImageUtils.upload_image(image)
|
||||
overlay = {**overlay, "image_url": first_url}
|
||||
if image_2 is not None:
|
||||
second_url = ImageUtils.upload_image(image_2)
|
||||
overlay = {**overlay, "image_urls": [first_url, second_url]}
|
||||
if video is not None:
|
||||
overlay = {**overlay, "video_url": MediaUtils.upload_video(video)}
|
||||
if audio is not None:
|
||||
overlay = {**overlay, "audio_url": MediaUtils.upload_audio(audio)}
|
||||
if int(seed) != -1:
|
||||
overlay = {**overlay, "seed": int(seed)}
|
||||
return overlay
|
||||
|
||||
|
||||
def build_overlay_arguments(
|
||||
endpoint_id: str,
|
||||
arguments_json: str,
|
||||
image: Any = None,
|
||||
image_2: Any = None,
|
||||
video: Any = None,
|
||||
audio: Any = None,
|
||||
seed: int = -1,
|
||||
) -> dict[str, Any]:
|
||||
"""Merge free-form JSON arguments with uploaded media inputs and seed.
|
||||
|
||||
Connected media inputs win over matching keys in the JSON
|
||||
(image_url, image_urls, video_url, audio_url, seed).
|
||||
"""
|
||||
parsed = _parse_arguments_json(endpoint_id, arguments_json)
|
||||
overlay = _media_overlay(image, image_2, video, audio, seed)
|
||||
return {**parsed, **overlay}
|
||||
|
||||
|
||||
def _extract_images(result: dict[str, Any]) -> Any | None:
|
||||
try:
|
||||
images = result.get("images")
|
||||
if isinstance(images, list) and images:
|
||||
return ResultProcessor.process_image_result(result)[0]
|
||||
if isinstance(result.get("image"), dict):
|
||||
return ResultProcessor.process_single_image_result(result)[0]
|
||||
except Exception as err:
|
||||
logger.debug("FalAnyEndpoint: could not extract images: %s", err)
|
||||
return None
|
||||
|
||||
|
||||
def _extract_video(result: dict[str, Any]) -> Any | None:
|
||||
try:
|
||||
url = find_url(result.get("video"))
|
||||
if url is not None:
|
||||
return MediaUtils.video_from_url(url)
|
||||
except Exception as err:
|
||||
logger.debug("FalAnyEndpoint: could not extract video: %s", err)
|
||||
return None
|
||||
|
||||
|
||||
def _extract_audio(result: dict[str, Any]) -> Any | None:
|
||||
try:
|
||||
url = find_url(result.get("audio"))
|
||||
if url is not None:
|
||||
return MediaUtils.audio_from_url(url)
|
||||
except Exception as err:
|
||||
logger.debug("FalAnyEndpoint: could not extract audio: %s", err)
|
||||
return None
|
||||
|
||||
|
||||
def extract_flexible_outputs(result: dict[str, Any]) -> tuple[Any, Any, Any, str]:
|
||||
"""Opportunistically extract (images, video, audio, raw json) from a result.
|
||||
|
||||
Each media slot is None when the result has no matching content; the raw
|
||||
result is always available as a JSON string in the last slot.
|
||||
"""
|
||||
return (
|
||||
_extract_images(result),
|
||||
_extract_video(result),
|
||||
_extract_audio(result),
|
||||
json.dumps(result, default=str),
|
||||
)
|
||||
|
||||
|
||||
class FalAnyEndpoint:
|
||||
"""Call any fal.ai endpoint with raw JSON arguments plus optional media inputs."""
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "VIDEO", "AUDIO", "STRING")
|
||||
RETURN_NAMES = ("images", "video", "audio", "result_json")
|
||||
FUNCTION = "run"
|
||||
CATEGORY = "FAL/Models"
|
||||
DESCRIPTION = (
|
||||
"Call any fal.ai endpoint by id. Provide arguments as a JSON object; "
|
||||
"connected media inputs are uploaded and override matching keys "
|
||||
"(image_url, image_urls, video_url, audio_url, seed) in the JSON. "
|
||||
"Outputs are extracted opportunistically; the raw result is always "
|
||||
"available as JSON."
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> dict[str, Any]:
|
||||
return {
|
||||
"required": {
|
||||
"endpoint_id": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "fal-ai/flux/dev",
|
||||
"tooltip": "fal endpoint id, e.g. fal-ai/flux/dev",
|
||||
},
|
||||
),
|
||||
"arguments_json": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "{}",
|
||||
"multiline": True,
|
||||
"tooltip": (
|
||||
"JSON object of API arguments. Connected media inputs "
|
||||
"and seed override matching keys here."
|
||||
),
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"image": ("IMAGE", {"tooltip": "Uploaded and sent as image_url"}),
|
||||
"image_2": (
|
||||
"IMAGE",
|
||||
{
|
||||
"tooltip": (
|
||||
"Second image; when set together with 'image', both are "
|
||||
"also sent as image_urls [url1, url2]"
|
||||
)
|
||||
},
|
||||
),
|
||||
"video": ("VIDEO", {"tooltip": "Uploaded and sent as video_url"}),
|
||||
"audio": ("AUDIO", {"tooltip": "Uploaded and sent as audio_url"}),
|
||||
"seed": (
|
||||
"INT",
|
||||
{
|
||||
"default": -1,
|
||||
"min": -1,
|
||||
"max": 2**31 - 1,
|
||||
"control_after_generate": True,
|
||||
"tooltip": "-1 = omit seed; any other value is sent to the API",
|
||||
},
|
||||
),
|
||||
"force_rerun": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": False,
|
||||
"tooltip": "Bypass ComfyUI's cache and call the API again",
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def IS_CHANGED(cls, **kwargs: Any) -> Any:
|
||||
if kwargs.get("force_rerun"):
|
||||
return float("nan")
|
||||
return stable_hash(kwargs)
|
||||
|
||||
def _run_sync(
|
||||
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)
|
||||
|
||||
arguments = build_overlay_arguments(
|
||||
endpoint, arguments_json, image, image_2, video, audio, seed
|
||||
)
|
||||
|
||||
result = ApiHandler.submit_and_get_result(
|
||||
endpoint, arguments, skip_cache=bool(force_rerun)
|
||||
)
|
||||
|
||||
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
|
||||
@@ -0,0 +1,153 @@
|
||||
"""Pure translation of ComfyUI node kwargs back into fal API arguments."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
from ..fal_utils import FalApiError, ImageUtils, MediaUtils
|
||||
from .schema_to_inputs import DIRECT_URL_KINDS, DIRECT_URL_SUFFIX
|
||||
|
||||
_DEFAULT_DIMENSION = 1024
|
||||
|
||||
|
||||
def _direct_url_value(inp: dict[str, Any], kwargs: dict[str, Any]) -> str:
|
||||
"""The stripped '<name>_direct_url' kwarg for a media input, or ''."""
|
||||
if inp.get("media_kind") not in DIRECT_URL_KINDS:
|
||||
return ""
|
||||
raw = kwargs.get(inp["name"] + DIRECT_URL_SUFFIX)
|
||||
return str(raw).strip() if isinstance(raw, str) else ""
|
||||
|
||||
|
||||
def _direct_url_argument(endpoint: str, inp: dict[str, Any], text: str) -> Any:
|
||||
"""Validate a passthrough URL string; is_list inputs accept comma-separated URLs."""
|
||||
twin_name = inp["name"] + DIRECT_URL_SUFFIX
|
||||
error = FalApiError(endpoint, f"'{twin_name}' must be an http(s) URL")
|
||||
if not inp.get("is_list"):
|
||||
if not text.startswith(("http://", "https://")):
|
||||
raise error
|
||||
return text
|
||||
parts = [part.strip() for part in text.split(",") if part.strip()]
|
||||
if not parts or any(not part.startswith(("http://", "https://")) for part in parts):
|
||||
raise error
|
||||
return parts
|
||||
|
||||
|
||||
def _upload_image(inp: dict[str, Any], value: Any) -> Any:
|
||||
if inp.get("is_list"):
|
||||
return ImageUtils.prepare_images(value)
|
||||
return ImageUtils.upload_image(value)
|
||||
|
||||
|
||||
def _media_argument(inp: dict[str, Any], value: Any) -> Any | None:
|
||||
media_kind = inp.get("media_kind")
|
||||
if media_kind == "image":
|
||||
return _upload_image(inp, value)
|
||||
if media_kind == "video":
|
||||
return MediaUtils.upload_video(value)
|
||||
if media_kind == "audio":
|
||||
return MediaUtils.upload_audio(value)
|
||||
# media_kind == "file": already a URL string in the widget
|
||||
text = str(value).strip()
|
||||
return text or None
|
||||
|
||||
|
||||
def _json_argument(endpoint: str, name: str, value: Any) -> Any | None:
|
||||
text = str(value).strip()
|
||||
if not text:
|
||||
return None
|
||||
try:
|
||||
return json.loads(text)
|
||||
except ValueError as err:
|
||||
raise FalApiError(endpoint, f"Invalid JSON in '{name}': {err}") from err
|
||||
|
||||
|
||||
def _multi_enum_argument(endpoint: str, inp: dict[str, Any], value: Any) -> Any | None:
|
||||
"""Comma-separated string widget → validated list of enum members."""
|
||||
selected = [part.strip() for part in str(value).split(",") if part.strip()]
|
||||
if not selected:
|
||||
return None
|
||||
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(
|
||||
endpoint,
|
||||
f"Invalid value(s) {invalid} for '{inp['name']}'. "
|
||||
f"Allowed: {', '.join(sorted(allowed))}",
|
||||
)
|
||||
return [allowed[part] for part in selected]
|
||||
|
||||
|
||||
def _enum_argument(inp: dict[str, Any], value: Any, kwargs: dict[str, Any]) -> Any:
|
||||
if inp.get("has_custom_size") and value == "custom_size":
|
||||
return {
|
||||
"width": int(kwargs.get("width", _DEFAULT_DIMENSION)),
|
||||
"height": int(kwargs.get("height", _DEFAULT_DIMENSION)),
|
||||
}
|
||||
return value
|
||||
|
||||
|
||||
def _scalar_argument(
|
||||
endpoint: str, inp: dict[str, Any], value: Any, kwargs: dict[str, Any]
|
||||
) -> Any | None:
|
||||
input_type = inp.get("type")
|
||||
if input_type == "enum":
|
||||
if inp.get("is_list"):
|
||||
return _multi_enum_argument(endpoint, inp, value)
|
||||
return _enum_argument(inp, value, kwargs)
|
||||
if input_type in ("json", "object", "array"):
|
||||
return _json_argument(endpoint, inp["name"], value)
|
||||
if input_type == "integer":
|
||||
return int(value)
|
||||
if input_type == "number":
|
||||
return float(value)
|
||||
if input_type == "boolean":
|
||||
return bool(value)
|
||||
if input_type == "string":
|
||||
if not inp.get("required") and value == "":
|
||||
return None
|
||||
return value
|
||||
return value
|
||||
|
||||
|
||||
def _argument_for(
|
||||
endpoint: str, inp: dict[str, Any], value: Any, kwargs: dict[str, Any]
|
||||
) -> Any | None:
|
||||
if inp.get("media_kind"):
|
||||
return _media_argument(inp, value)
|
||||
return _scalar_argument(endpoint, inp, value, kwargs)
|
||||
|
||||
|
||||
def build_arguments(model: dict[str, Any], kwargs: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Build the fal API argument dict from node kwargs. Never mutates inputs."""
|
||||
endpoint = model["endpoint_id"]
|
||||
arguments: dict[str, Any] = {}
|
||||
inputs = model.get("inputs", [])
|
||||
input_names = {inp["name"] for inp in inputs}
|
||||
|
||||
for inp in inputs:
|
||||
name = inp["name"]
|
||||
# URL passthrough wins over the media input (which may be None or connected);
|
||||
# skip when the twin name is a real model input (no twin was generated then)
|
||||
if name + DIRECT_URL_SUFFIX not in input_names:
|
||||
direct_url = _direct_url_value(inp, kwargs)
|
||||
if direct_url:
|
||||
resolved = _direct_url_argument(endpoint, inp, direct_url)
|
||||
arguments = {**arguments, name: resolved}
|
||||
continue
|
||||
if name not in kwargs:
|
||||
continue
|
||||
value = kwargs[name]
|
||||
if value is None:
|
||||
continue
|
||||
if name == "seed":
|
||||
seed = int(value)
|
||||
if seed != -1:
|
||||
arguments = {**arguments, "seed": seed}
|
||||
continue
|
||||
resolved = _argument_for(endpoint, inp, value, kwargs)
|
||||
if resolved is None:
|
||||
continue
|
||||
arguments = {**arguments, name: resolved}
|
||||
|
||||
return arguments
|
||||
@@ -0,0 +1,163 @@
|
||||
"""Builds concrete ComfyUI node classes from registry model entries."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import importlib.util
|
||||
import inspect
|
||||
import re
|
||||
from functools import cache
|
||||
from typing import Any
|
||||
|
||||
from ..fal_utils import ApiHandler
|
||||
from .arguments import build_arguments
|
||||
from .outputs import RETURN_SPECS, process_result
|
||||
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("/", "-")
|
||||
|
||||
|
||||
def _slug(text: str) -> str:
|
||||
return re.sub(r"[^a-z0-9]+", "", text.lower())
|
||||
|
||||
|
||||
def build_display_name(model: dict[str, Any]) -> str:
|
||||
endpoint_id = model["endpoint_id"]
|
||||
title = model.get("title") or endpoint_id
|
||||
parts = endpoint_id.split("/")
|
||||
remainder = "/".join(parts[1:]) if len(parts) > 1 else endpoint_id
|
||||
if not remainder or _slug(remainder) == _slug(title):
|
||||
return f"{title} (fal)"
|
||||
return f"{title} · {remainder} (fal)"
|
||||
|
||||
|
||||
def _value_fingerprint(value: Any) -> str:
|
||||
# torch tensors: repr() summarizes large tensors (edge elements only), so two
|
||||
# different images could hash identically — fingerprint the raw bytes instead
|
||||
detach = getattr(value, "detach", None)
|
||||
if callable(detach):
|
||||
try:
|
||||
tensor = value.detach().cpu().contiguous()
|
||||
digest = hashlib.sha256(tensor.numpy().tobytes()).hexdigest()
|
||||
return f"tensor:{tuple(tensor.shape)}:{tensor.dtype}:{digest}"
|
||||
except Exception: # non-numpy-compatible tensor; fall through to repr
|
||||
pass
|
||||
if isinstance(value, dict): # e.g. AUDIO dicts carrying a waveform tensor
|
||||
return repr(sorted((k, _value_fingerprint(v)) for k, v in value.items()))
|
||||
return repr(value)
|
||||
|
||||
|
||||
def stable_hash(kwargs: dict[str, Any]) -> str:
|
||||
payload = repr(sorted((key, _value_fingerprint(value)) for key, value in kwargs.items()))
|
||||
return hashlib.sha256(payload.encode("utf-8")).hexdigest()
|
||||
|
||||
|
||||
|
||||
@cache
|
||||
def _accepts_skip_cache(func: Any) -> bool:
|
||||
"""Whether ``submit_and_get_result`` supports the skip_cache keyword.
|
||||
|
||||
Cached per callable so the check runs once, and re-evaluated automatically
|
||||
when tests stub out ApiHandler. On uninspectable callables fall back to
|
||||
False: calling without skip_cache works with both old and new signatures.
|
||||
"""
|
||||
try:
|
||||
return "skip_cache" in inspect.signature(func).parameters
|
||||
except (TypeError, ValueError):
|
||||
return False
|
||||
|
||||
|
||||
def _call_api(endpoint_id: str, arguments: dict[str, Any], skip_cache: bool) -> Any:
|
||||
submit = ApiHandler.submit_and_get_result
|
||||
if _accepts_skip_cache(submit):
|
||||
return submit(endpoint_id, arguments, skip_cache=skip_cache)
|
||||
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))
|
||||
|
||||
|
||||
def _description(model: dict[str, Any]) -> str:
|
||||
description = model.get("description") or ""
|
||||
pricing = model.get("pricing") or ""
|
||||
if pricing:
|
||||
return f"{description}\n\nPricing: {pricing}".strip()
|
||||
return description
|
||||
|
||||
|
||||
def build_node_class(model: dict[str, Any]) -> type:
|
||||
"""Create a ComfyUI node class for a single registry model entry."""
|
||||
endpoint_id = model["endpoint_id"]
|
||||
kind = model.get("output_kind", "json")
|
||||
return_types, return_names = RETURN_SPECS.get(kind, RETURN_SPECS["json"])
|
||||
category = model.get("category") or "other"
|
||||
|
||||
def input_types(cls: type) -> dict[str, Any]:
|
||||
return build_input_types(model)
|
||||
|
||||
def is_changed(cls: type, **kwargs: Any) -> Any:
|
||||
if kwargs.get("force_rerun"):
|
||||
return float("nan")
|
||||
return stable_hash(kwargs)
|
||||
|
||||
def run(self: Any, **kwargs: Any) -> tuple:
|
||||
arguments = build_arguments(model, kwargs)
|
||||
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),
|
||||
"RETURN_TYPES": return_types,
|
||||
"RETURN_NAMES": return_names,
|
||||
"FUNCTION": "run",
|
||||
"CATEGORY": f"FAL/Models/{category}",
|
||||
"DESCRIPTION": _description(model),
|
||||
"run": run_async if _ASYNC_CAPABLE else run,
|
||||
"_FAL_ENDPOINT_ID": endpoint_id,
|
||||
}
|
||||
return type(_class_name(model), (object,), attrs)
|
||||
@@ -0,0 +1,144 @@
|
||||
"""Return-type specs per output kind and result post-processing."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections.abc import Sequence
|
||||
from typing import Any
|
||||
|
||||
from ..fal_utils import FalApiError, MediaUtils, ResultProcessor
|
||||
|
||||
RETURN_SPECS: dict[str, tuple[tuple[str, ...], tuple[str, ...]]] = {
|
||||
"images": (("IMAGE", "STRING"), ("images", "image_urls")),
|
||||
"image": (("IMAGE", "STRING"), ("image", "image_url")),
|
||||
"video": (("VIDEO", "STRING"), ("video", "video_url")),
|
||||
"audio": (("AUDIO", "STRING"), ("audio", "audio_url")),
|
||||
"text": (("STRING",), ("text",)),
|
||||
"file": (("STRING",), ("file_url",)),
|
||||
"json": (("STRING",), ("json",)),
|
||||
}
|
||||
|
||||
_FILE_PROP_CANDIDATES = (
|
||||
"model_glb",
|
||||
"model_mesh",
|
||||
"model_url",
|
||||
"model_urls",
|
||||
"file",
|
||||
"file_url",
|
||||
"output",
|
||||
"outputs",
|
||||
)
|
||||
|
||||
|
||||
def find_url(value: Any) -> str | None:
|
||||
"""Recursively dig a result fragment for a URL string."""
|
||||
if isinstance(value, str):
|
||||
return value if value.startswith(("http://", "https://", "data:")) else None
|
||||
if isinstance(value, dict):
|
||||
direct = value.get("url")
|
||||
if isinstance(direct, str):
|
||||
return direct
|
||||
for nested in value.values():
|
||||
found = find_url(nested)
|
||||
if found is not None:
|
||||
return found
|
||||
return None
|
||||
if isinstance(value, (list, tuple)):
|
||||
for item in value:
|
||||
found = find_url(item)
|
||||
if found is not None:
|
||||
return found
|
||||
return None
|
||||
|
||||
|
||||
def _url_from_props(result: dict[str, Any], props: Sequence[str]) -> str | None:
|
||||
for prop in props:
|
||||
if prop in result:
|
||||
found = find_url(result[prop])
|
||||
if found is not None:
|
||||
return found
|
||||
return None
|
||||
|
||||
|
||||
def _media_url(model: dict[str, Any], result: dict[str, Any], primary: str) -> str:
|
||||
props: list[str] = [primary]
|
||||
for prop in model.get("output_props") or []:
|
||||
if prop not in props:
|
||||
props = [*props, prop]
|
||||
url = _url_from_props(result, props)
|
||||
if url is None:
|
||||
url = find_url(result)
|
||||
if url is None:
|
||||
raise FalApiError(
|
||||
model["endpoint_id"], f"No {primary} URL found in API result"
|
||||
)
|
||||
return url
|
||||
|
||||
|
||||
def _image_urls_csv(result: dict[str, Any]) -> str:
|
||||
"""CDN URLs of an {"images": [...]} result, comma-joined (twin-input format)."""
|
||||
images = result.get("images")
|
||||
if not isinstance(images, (list, tuple)):
|
||||
return ""
|
||||
urls = [url for url in (find_url(item) for item in images) if url]
|
||||
return ",".join(urls)
|
||||
|
||||
|
||||
def _process_images(result: dict[str, Any]) -> tuple[Any, ...]:
|
||||
tensor = ResultProcessor.process_image_result(result)[0]
|
||||
return (tensor, _image_urls_csv(result))
|
||||
|
||||
|
||||
def _process_image(result: dict[str, Any]) -> tuple[Any, ...]:
|
||||
tensor = ResultProcessor.process_single_image_result(result)[0]
|
||||
url = find_url(result.get("image"))
|
||||
return (tensor, url or "")
|
||||
|
||||
|
||||
def _process_video(model: dict[str, Any], result: dict[str, Any]) -> tuple[Any, ...]:
|
||||
url = _media_url(model, result, "video")
|
||||
return (MediaUtils.video_from_url(url), url)
|
||||
|
||||
|
||||
def _process_audio(model: dict[str, Any], result: dict[str, Any]) -> tuple[Any, ...]:
|
||||
url = _media_url(model, result, "audio")
|
||||
return (MediaUtils.audio_from_url(url), url)
|
||||
|
||||
|
||||
def _process_text(model: dict[str, Any], result: dict[str, Any]) -> tuple[Any, ...]:
|
||||
for prop in model.get("output_props") or []:
|
||||
value = result.get(prop)
|
||||
if isinstance(value, str):
|
||||
return (value,)
|
||||
for value in result.values():
|
||||
if isinstance(value, str):
|
||||
return (value,)
|
||||
return (json.dumps(result, default=str),)
|
||||
|
||||
|
||||
def _process_file(model: dict[str, Any], result: dict[str, Any]) -> tuple[Any, ...]:
|
||||
props = [*(model.get("output_props") or []), *_FILE_PROP_CANDIDATES]
|
||||
url = _url_from_props(result, props)
|
||||
if url is None:
|
||||
url = find_url(result)
|
||||
if url is None:
|
||||
raise FalApiError(model["endpoint_id"], "No file URL found in API result")
|
||||
return (url,)
|
||||
|
||||
|
||||
def process_result(model: dict[str, Any], result: dict[str, Any]) -> tuple[Any, ...]:
|
||||
"""Convert a raw fal API result dict into the node's return tuple."""
|
||||
kind = model.get("output_kind", "json")
|
||||
if kind == "images":
|
||||
return _process_images(result)
|
||||
if kind == "image":
|
||||
return _process_image(result)
|
||||
if kind == "video":
|
||||
return _process_video(model, result)
|
||||
if kind == "audio":
|
||||
return _process_audio(model, result)
|
||||
if kind == "text":
|
||||
return _process_text(model, result)
|
||||
if kind == "file":
|
||||
return _process_file(model, result)
|
||||
return (json.dumps(result, default=str),)
|
||||
@@ -0,0 +1,272 @@
|
||||
"""Loads the fal model registry and builds dynamic node mappings.
|
||||
|
||||
Must never raise: any failure results in empty (or partial) mappings so the
|
||||
static nodes keep loading no matter what.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from ..fal_utils import FalConfig, logger
|
||||
from .any_endpoint import ANY_ENDPOINT_DISPLAY_NAME, ANY_ENDPOINT_KEY, FalAnyEndpoint
|
||||
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]]
|
||||
|
||||
|
||||
def _registry_path() -> Path:
|
||||
package_dir = Path(__file__).resolve().parent
|
||||
real = package_dir.parents[1] / "data" / _REGISTRY_FILENAME
|
||||
if real.is_file():
|
||||
return real
|
||||
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")
|
||||
return bool(value)
|
||||
|
||||
|
||||
def _get_setting(section: str, name: str, default: Any) -> Any:
|
||||
config = FalConfig()
|
||||
getter = getattr(config, "get_setting", None)
|
||||
if getter is None:
|
||||
return default
|
||||
return getter(section, name, default)
|
||||
|
||||
|
||||
def _category_filter() -> set[str]:
|
||||
raw = _get_setting("dynamic_nodes", "categories", "") or ""
|
||||
return {part.strip() for part in str(raw).split(",") if part.strip()}
|
||||
|
||||
|
||||
def _read_models() -> list[dict[str, Any]]:
|
||||
path = _registry_path()
|
||||
try:
|
||||
with open(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 models
|
||||
except Exception as err:
|
||||
logger.error("Failed to read fal registry at %s: %s", path, err)
|
||||
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
|
||||
counter = 2
|
||||
while f"{name} #{counter}" in used:
|
||||
counter += 1
|
||||
return f"{name} #{counter}"
|
||||
|
||||
|
||||
def _build_model_mappings(
|
||||
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:
|
||||
if categories and model.get("category") not in categories:
|
||||
continue
|
||||
key = node_key(model)
|
||||
if key in classes or key == ANY_ENDPOINT_KEY:
|
||||
skipped += 1
|
||||
logger.debug("Duplicate dynamic node key skipped: %s", key)
|
||||
continue
|
||||
node_class = build_node_class(model)
|
||||
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)
|
||||
except Exception as err:
|
||||
skipped += 1
|
||||
logger.debug(
|
||||
"Skipped dynamic node for %s: %s",
|
||||
model.get("endpoint_id", "<unknown>"),
|
||||
err,
|
||||
)
|
||||
|
||||
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:
|
||||
"""Build (NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS) for dynamic nodes."""
|
||||
try:
|
||||
if not _truthy(_get_setting("dynamic_nodes", "enabled", True)):
|
||||
logger.info("Dynamic fal nodes disabled via config")
|
||||
return {}, {}
|
||||
|
||||
categories = _category_filter()
|
||||
models = _read_models()
|
||||
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, 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)
|
||||
return {}, {}
|
||||
@@ -0,0 +1,231 @@
|
||||
"""Pure translation of a registry model schema into a ComfyUI INPUT_TYPES dict."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
INT_MIN = -(2**31)
|
||||
INT_MAX = 2**31 - 1
|
||||
|
||||
SEED_SPEC = (
|
||||
"INT",
|
||||
{
|
||||
"default": -1,
|
||||
"min": -1,
|
||||
"max": INT_MAX,
|
||||
"control_after_generate": True,
|
||||
"tooltip": "-1 = random (fal picks); any other value is sent to the API",
|
||||
},
|
||||
)
|
||||
|
||||
FORCE_RERUN_SPEC = (
|
||||
"BOOLEAN",
|
||||
{"default": False, "tooltip": "Bypass ComfyUI's cache and call the API again"},
|
||||
)
|
||||
|
||||
WIDTH_HEIGHT_OPTS = {"default": 1024, "min": 64, "max": 14142, "step": 8}
|
||||
|
||||
_MEDIA_TYPES = {"image": "IMAGE", "video": "VIDEO", "audio": "AUDIO"}
|
||||
|
||||
# media kinds that get a "<name>_direct_url" passthrough companion input
|
||||
# ("file" is excluded: it is already a plain STRING URL widget)
|
||||
DIRECT_URL_SUFFIX = "_direct_url"
|
||||
DIRECT_URL_KINDS = ("image", "video", "audio")
|
||||
|
||||
|
||||
def _clamp(value: float, lo: float, hi: float) -> float:
|
||||
return max(lo, min(hi, value))
|
||||
|
||||
|
||||
def _with_tooltip(opts: dict[str, Any], description: str | None) -> dict[str, Any]:
|
||||
if description:
|
||||
return {**opts, "tooltip": description}
|
||||
return opts
|
||||
|
||||
|
||||
def _int_spec(inp: dict[str, Any]) -> tuple[Any, ...]:
|
||||
lo = int(inp["min"]) if inp.get("min") is not None else INT_MIN
|
||||
hi = int(inp["max"]) if inp.get("max") is not None else INT_MAX
|
||||
raw_default = inp.get("default")
|
||||
default = int(raw_default) if isinstance(raw_default, (int, float)) else 0
|
||||
opts = {"default": int(_clamp(default, lo, hi)), "min": lo, "max": hi}
|
||||
return ("INT", _with_tooltip(opts, inp.get("description")))
|
||||
|
||||
|
||||
def _float_spec(inp: dict[str, Any]) -> tuple[Any, ...]:
|
||||
has_range = inp.get("min") is not None and inp.get("max") is not None
|
||||
lo = float(inp["min"]) if inp.get("min") is not None else -1e10
|
||||
hi = float(inp["max"]) if inp.get("max") is not None else 1e10
|
||||
step = 0.01 if has_range and (hi - lo) <= 10 else 0.1
|
||||
raw_default = inp.get("default")
|
||||
default = float(raw_default) if isinstance(raw_default, (int, float)) else 0.0
|
||||
opts = {"default": _clamp(default, lo, hi), "min": lo, "max": hi, "step": step}
|
||||
return ("FLOAT", _with_tooltip(opts, inp.get("description")))
|
||||
|
||||
|
||||
def _multi_enum_spec(inp: dict[str, Any]) -> tuple[Any, ...]:
|
||||
"""Array-of-enum inputs: ComfyUI has no multi-select widget, so use a
|
||||
comma-separated string validated at call time."""
|
||||
values = list(inp.get("enum") or [])
|
||||
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(str(value) for value in values)}".strip()
|
||||
return ("STRING", {"default": text, "tooltip": tooltip})
|
||||
|
||||
|
||||
def _enum_spec(inp: dict[str, Any]) -> tuple[Any, ...]:
|
||||
if inp.get("is_list"):
|
||||
return _multi_enum_spec(inp)
|
||||
values = list(inp.get("enum") or [])
|
||||
if not values:
|
||||
return _string_spec(inp)
|
||||
default = inp.get("default")
|
||||
if default not in values:
|
||||
# a dict default on a has_custom_size enum means the API defaults to an
|
||||
# explicit {width, height}; represent that as the custom_size preset
|
||||
if isinstance(default, dict) and "custom_size" in values:
|
||||
default = "custom_size"
|
||||
else:
|
||||
default = values[0]
|
||||
opts = _with_tooltip({"default": default}, inp.get("description"))
|
||||
return (values, opts)
|
||||
|
||||
|
||||
def _custom_size_default(inp: dict[str, Any], dimension: str) -> int:
|
||||
default = inp.get("default")
|
||||
if isinstance(default, dict):
|
||||
value = default.get(dimension)
|
||||
if isinstance(value, int) and value > 0:
|
||||
return int(_clamp(value, 64, 14142))
|
||||
return WIDTH_HEIGHT_OPTS["default"]
|
||||
|
||||
|
||||
def _bool_spec(inp: dict[str, Any]) -> tuple[Any, ...]:
|
||||
opts = {"default": bool(inp.get("default"))}
|
||||
return ("BOOLEAN", _with_tooltip(opts, inp.get("description")))
|
||||
|
||||
|
||||
def _string_spec(inp: dict[str, Any]) -> tuple[Any, ...]:
|
||||
default = inp.get("default")
|
||||
opts = {
|
||||
"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")))
|
||||
|
||||
|
||||
def _json_spec(inp: dict[str, Any]) -> tuple[Any, ...]:
|
||||
default = inp.get("default")
|
||||
if default is None:
|
||||
text = ""
|
||||
elif isinstance(default, str):
|
||||
text = default
|
||||
else:
|
||||
text = json.dumps(default)
|
||||
description = (inp.get("description") or "").strip()
|
||||
tooltip = (description + " (JSON)").strip()
|
||||
return ("STRING", {"default": text, "multiline": True, "tooltip": tooltip})
|
||||
|
||||
|
||||
def _media_spec(inp: dict[str, Any]) -> tuple[Any, ...]:
|
||||
media_kind = inp.get("media_kind")
|
||||
comfy_type = _MEDIA_TYPES.get(media_kind)
|
||||
if comfy_type is not None:
|
||||
return (comfy_type,)
|
||||
# media_kind == "file": plain URL string
|
||||
description = (inp.get("description") or "").strip()
|
||||
tooltip = (description + " (URL to file)").strip()
|
||||
return ("STRING", {"default": "", "tooltip": tooltip})
|
||||
|
||||
|
||||
def _direct_url_spec(name: str, media_kind: str, is_list: bool) -> tuple[Any, ...]:
|
||||
tooltip = (
|
||||
f"fal/CDN URL passthrough for '{name}': when set, this URL is sent directly "
|
||||
f"and the {media_kind} input is ignored — no download/re-upload. "
|
||||
"Chain fal nodes' *_url outputs here."
|
||||
)
|
||||
if is_list:
|
||||
tooltip += " Comma-separate multiple URLs."
|
||||
return ("STRING", {"default": "", "tooltip": tooltip})
|
||||
|
||||
|
||||
def _direct_url_twin(
|
||||
inp: dict[str, Any], existing_names: set[str]
|
||||
) -> tuple[str, tuple[Any, ...]] | None:
|
||||
"""The optional passthrough companion for a media input, or None."""
|
||||
media_kind = inp.get("media_kind")
|
||||
if media_kind not in DIRECT_URL_KINDS:
|
||||
return None
|
||||
twin_name = inp["name"] + DIRECT_URL_SUFFIX
|
||||
if twin_name in existing_names:
|
||||
return None
|
||||
spec = _direct_url_spec(inp["name"], media_kind, bool(inp.get("is_list")))
|
||||
return (twin_name, spec)
|
||||
|
||||
|
||||
def _input_spec(inp: dict[str, Any]) -> tuple[Any, ...]:
|
||||
if inp.get("media_kind"):
|
||||
return _media_spec(inp)
|
||||
input_type = inp.get("type")
|
||||
if input_type == "enum":
|
||||
return _enum_spec(inp)
|
||||
if input_type == "integer":
|
||||
return _int_spec(inp)
|
||||
if input_type == "number":
|
||||
return _float_spec(inp)
|
||||
if input_type == "boolean":
|
||||
return _bool_spec(inp)
|
||||
if input_type in ("json", "object", "array"):
|
||||
return _json_spec(inp)
|
||||
return _string_spec(inp)
|
||||
|
||||
|
||||
def build_input_types(model: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Build a ComfyUI INPUT_TYPES dict from a registry model entry."""
|
||||
required: dict[str, Any] = {}
|
||||
optional: dict[str, Any] = {}
|
||||
# twins of *required* media inputs go at the start of the optional bucket
|
||||
leading_optional: dict[str, Any] = {}
|
||||
custom_size_input: dict[str, Any] | None = None
|
||||
|
||||
inputs = model.get("inputs", [])
|
||||
existing_names = {inp["name"] for inp in inputs}
|
||||
|
||||
for inp in inputs:
|
||||
name = inp["name"]
|
||||
if name == "seed":
|
||||
optional[name] = SEED_SPEC
|
||||
continue
|
||||
if inp.get("has_custom_size") and custom_size_input is None:
|
||||
custom_size_input = inp
|
||||
spec = _input_spec(inp)
|
||||
twin = _direct_url_twin(inp, existing_names)
|
||||
if inp.get("required"):
|
||||
required[name] = spec
|
||||
if twin is not None:
|
||||
leading_optional[twin[0]] = twin[1]
|
||||
else:
|
||||
optional[name] = spec
|
||||
if twin is not None:
|
||||
optional[twin[0]] = twin[1]
|
||||
|
||||
optional = {**leading_optional, **optional}
|
||||
|
||||
if custom_size_input is not None:
|
||||
for dimension in ("width", "height"):
|
||||
if dimension not in required and dimension not in optional:
|
||||
opts = {
|
||||
**WIDTH_HEIGHT_OPTS,
|
||||
"default": _custom_size_default(custom_size_input, dimension),
|
||||
}
|
||||
optional[dimension] = ("INT", opts)
|
||||
|
||||
optional["force_rerun"] = FORCE_RERUN_SPEC
|
||||
|
||||
return {"required": required, "optional": optional}
|
||||
@@ -0,0 +1,42 @@
|
||||
"""Backward-compatible facade for the nodes.utils package.
|
||||
|
||||
Existing node modules import from here, e.g.:
|
||||
|
||||
from .fal_utils import FalConfig, ImageUtils, ResultProcessor, ApiHandler
|
||||
|
||||
The implementations now live in the ``nodes/utils`` package.
|
||||
"""
|
||||
|
||||
from .utils import (
|
||||
ApiHandler,
|
||||
ArchiveUtils,
|
||||
BillingUtils,
|
||||
FalApiError,
|
||||
FalConfig,
|
||||
ImageUtils,
|
||||
JobStore,
|
||||
MediaUtils,
|
||||
PricingUtils,
|
||||
ResultCache,
|
||||
ResultProcessor,
|
||||
SessionLedger,
|
||||
SpendGuard,
|
||||
logger,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"ApiHandler",
|
||||
"ArchiveUtils",
|
||||
"BillingUtils",
|
||||
"FalApiError",
|
||||
"FalConfig",
|
||||
"ImageUtils",
|
||||
"JobStore",
|
||||
"MediaUtils",
|
||||
"PricingUtils",
|
||||
"ResultCache",
|
||||
"ResultProcessor",
|
||||
"SessionLedger",
|
||||
"SpendGuard",
|
||||
"logger",
|
||||
]
|
||||
Regular → Executable
+2026
-470
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,97 @@
|
||||
"""Fal Job Inbox: list async fal jobs recorded across ComfyUI sessions."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from .fal_utils import JobStore, logger
|
||||
|
||||
_CATEGORY = "FAL/Platform"
|
||||
|
||||
_STATUS_CHOICES = ("all", "submitted", "collected")
|
||||
|
||||
|
||||
class FalJobInbox:
|
||||
"""List async fal jobs from the persistent store; jobs survive restarts."""
|
||||
|
||||
RETURN_TYPES = ("STRING", "STRING", "STRING")
|
||||
RETURN_NAMES = ("report", "latest_request_id", "latest_endpoint")
|
||||
FUNCTION = "inbox"
|
||||
CATEGORY = _CATEGORY
|
||||
OUTPUT_NODE = True
|
||||
DESCRIPTION = (
|
||||
"Inbox of async fal jobs recorded by Fal Submit. The store is "
|
||||
"persistent, so jobs queued in a previous session survive a ComfyUI "
|
||||
"restart: submit tonight, restart, then wire latest_request_id and "
|
||||
"latest_endpoint into Fal Result by Request ID to collect tomorrow "
|
||||
"without re-paying. Never fails the graph."
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> dict[str, Any]:
|
||||
return {
|
||||
"required": {},
|
||||
"optional": {
|
||||
"status_filter": (
|
||||
list(_STATUS_CHOICES),
|
||||
{
|
||||
"default": "all",
|
||||
"tooltip": (
|
||||
"Which jobs to list: 'submitted' shows jobs still "
|
||||
"waiting to be collected (including ones queued "
|
||||
"before a restart), 'collected' shows finished ones."
|
||||
),
|
||||
},
|
||||
),
|
||||
"limit": (
|
||||
"INT",
|
||||
{
|
||||
"default": 20,
|
||||
"min": 1,
|
||||
"max": 200,
|
||||
"tooltip": "Maximum number of jobs to list, newest first",
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def IS_CHANGED(cls, **kwargs: Any) -> Any:
|
||||
# The job store mutates outside the graph; always re-run.
|
||||
return float("nan")
|
||||
|
||||
@staticmethod
|
||||
def _filtered_report(store: JobStore, status: str, limit: int) -> str:
|
||||
"""Plain listing of jobs with one status (report() covers 'all')."""
|
||||
entries = store.entries(limit=limit, status=status)
|
||||
lines = [f"Fal job inbox ({status}): {len(entries)} shown, newest first"]
|
||||
lines.extend(
|
||||
f" {entry.get('endpoint') or '(unknown endpoint)'} "
|
||||
f"req={entry.get('request_id') or '-'}"
|
||||
for entry in entries
|
||||
)
|
||||
return "\n".join(lines)
|
||||
|
||||
def inbox(self, status_filter: str = "all", limit: int = 20) -> tuple[str, str, str]:
|
||||
try:
|
||||
store = JobStore()
|
||||
if status_filter in _STATUS_CHOICES[1:]:
|
||||
report = self._filtered_report(store, status_filter, int(limit))
|
||||
else:
|
||||
report = store.report(limit=int(limit))
|
||||
latest = next(iter(store.pending(limit=1)), None)
|
||||
latest_request_id = str(latest.get("request_id") or "") if latest else ""
|
||||
latest_endpoint = str(latest.get("endpoint") or "") if latest else ""
|
||||
return (report, latest_request_id, latest_endpoint)
|
||||
except Exception as exc: # This node must never fail the graph.
|
||||
logger.warning("FalJobInbox: could not read job store: %s", exc)
|
||||
return ("Fal job inbox: report unavailable", "", "")
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"FalJobInbox_fal": FalJobInbox,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"FalJobInbox_fal": "Fal Job Inbox (fal)",
|
||||
}
|
||||
+105
-52
@@ -1,71 +1,124 @@
|
||||
import os
|
||||
import configparser
|
||||
from fal_client.client import SyncClient
|
||||
from .fal_utils import ApiHandler, FalConfig
|
||||
|
||||
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
parent_dir = os.path.dirname(current_dir)
|
||||
config_path = os.path.join(parent_dir, "config.ini")
|
||||
# Initialize FalConfig
|
||||
fal_config = FalConfig()
|
||||
|
||||
config = configparser.ConfigParser()
|
||||
config.read(config_path)
|
||||
|
||||
try:
|
||||
if os.environ.get("FAL_KEY") is not None:
|
||||
print("FAL_KEY found in environment variables")
|
||||
fal_key = os.environ["FAL_KEY"]
|
||||
else:
|
||||
print("FAL_KEY not found in environment variables")
|
||||
fal_key = config['API']['FAL_KEY']
|
||||
print("FAL_KEY found in config.ini")
|
||||
os.environ["FAL_KEY"] = fal_key
|
||||
print("FAL_KEY set in environment variables")
|
||||
|
||||
# Check if FAL key is the default placeholder
|
||||
if fal_key == "<your_fal_api_key_here>":
|
||||
print("WARNING: You are using the default FAL API key placeholder!")
|
||||
print("Please set your actual FAL API key in either:")
|
||||
print("1. The config.ini file under [API] section")
|
||||
print("2. Or as an environment variable named FAL_KEY")
|
||||
print("Get your API key from: https://fal.ai/dashboard/keys")
|
||||
except KeyError:
|
||||
print("Error: FAL_KEY not found in config.ini or environment variables")
|
||||
|
||||
# Create the client with API key
|
||||
fal_client = SyncClient(key=fal_key)
|
||||
|
||||
class LLMNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"prompt": ("STRING", {"default": "", "multiline": True}),
|
||||
"model": (["google/gemini-flash-1.5-8b", "anthropic/claude-3.5-sonnet", "anthropic/claude-3-haiku",
|
||||
"google/gemini-pro-1.5", "google/gemini-flash-1.5", "meta-llama/llama-3.2-1b-instruct",
|
||||
"meta-llama/llama-3.2-3b-instruct", "meta-llama/llama-3.1-8b-instruct",
|
||||
"meta-llama/llama-3.1-70b-instruct", "openai/gpt-4o-mini", "openai/gpt-4o"],
|
||||
{"default": "google/gemini-flash-1.5-8b"}),
|
||||
"system_prompt": ("STRING", {"default": "", "multiline": True}),
|
||||
"prompt": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"multiline": True,
|
||||
"tooltip": "User prompt sent to the model.",
|
||||
},
|
||||
),
|
||||
"model": (
|
||||
[
|
||||
"google/gemini-2.5-flash",
|
||||
"anthropic/claude-sonnet-4.5",
|
||||
"openai/gpt-4.1",
|
||||
"openai/gpt-oss-120b",
|
||||
"meta-llama/llama-4-maverick",
|
||||
"Custom",
|
||||
],
|
||||
{
|
||||
"default": "google/gemini-2.5-flash",
|
||||
"tooltip": "Model to use. Select 'Custom' to type any OpenRouter model id in custom_model_name.",
|
||||
},
|
||||
),
|
||||
"system_prompt": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"multiline": True,
|
||||
"tooltip": "Optional system prompt to steer the model's behavior.",
|
||||
},
|
||||
),
|
||||
"temperature": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 1.0,
|
||||
"min": 0.0,
|
||||
"max": 2.0,
|
||||
"step": 0.1,
|
||||
"tooltip": "Sampling temperature. Lower is more deterministic.",
|
||||
},
|
||||
),
|
||||
"reasoning": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": False,
|
||||
"tooltip": "Request the model's reasoning trace (returned on the 'reasoning' output).",
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"max_tokens": (
|
||||
"INT",
|
||||
{
|
||||
"default": 0,
|
||||
"min": 0,
|
||||
"max": 100000,
|
||||
"tooltip": "Maximum output tokens. 0 uses the model default.",
|
||||
},
|
||||
),
|
||||
"custom_model_name": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"multiline": False,
|
||||
"tooltip": "OpenRouter model id used when model is set to 'Custom'.",
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_TYPES = ("STRING", "STRING",)
|
||||
RETURN_NAMES = ("output", "reasoning",)
|
||||
FUNCTION = "generate_text"
|
||||
CATEGORY = "FAL/LLM"
|
||||
|
||||
def generate_text(self, prompt, model, system_prompt):
|
||||
arguments = {
|
||||
"model": model,
|
||||
"prompt": prompt,
|
||||
"system_prompt": system_prompt,
|
||||
}
|
||||
|
||||
def generate_text(self, prompt, model, system_prompt, temperature, reasoning, max_tokens=0, custom_model_name=""):
|
||||
try:
|
||||
handler = fal_client.submit("fal-ai/any-llm", arguments=arguments)
|
||||
result = handler.get()
|
||||
return (result["output"],)
|
||||
# Handle custom model selection
|
||||
if model == "Custom":
|
||||
if not custom_model_name or custom_model_name.strip() == "":
|
||||
# Raises a clear FalApiError
|
||||
ApiHandler.handle_text_generation_error(
|
||||
"Custom", "Custom model name is required when 'Custom' is selected"
|
||||
)
|
||||
model = custom_model_name.strip()
|
||||
|
||||
arguments = {
|
||||
"model": model,
|
||||
"prompt": prompt,
|
||||
"system_prompt": system_prompt,
|
||||
"temperature": temperature,
|
||||
"reasoning": reasoning,
|
||||
"stream": False,
|
||||
}
|
||||
|
||||
# Only include max_tokens if it's greater than 0
|
||||
if max_tokens > 0:
|
||||
arguments["max_tokens"] = max_tokens
|
||||
|
||||
result = ApiHandler.submit_and_get_result("openrouter/router", arguments)
|
||||
|
||||
# Extract output and reasoning
|
||||
output_text = result.get("output", "")
|
||||
reasoning_text = result.get("reasoning", "")
|
||||
|
||||
return (output_text, reasoning_text)
|
||||
except Exception as e:
|
||||
print(f"Error generating text with LLM: {str(e)}")
|
||||
return ("Error: Unable to generate text.",)
|
||||
# Raises a clear FalApiError (passes an existing FalApiError through
|
||||
# unchanged, so the custom-model validation error is not re-wrapped)
|
||||
return ApiHandler.handle_text_generation_error(model, e)
|
||||
|
||||
|
||||
# Node class mappings
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
|
||||
@@ -0,0 +1,622 @@
|
||||
"""Platform nodes: async submit/collect, result recovery, cost tools, media saving."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import shutil
|
||||
import time
|
||||
from typing import Any
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from .dynamic.any_endpoint import build_overlay_arguments, extract_flexible_outputs
|
||||
from .fal_utils import (
|
||||
ApiHandler,
|
||||
FalApiError,
|
||||
MediaUtils,
|
||||
PricingUtils,
|
||||
ResultCache,
|
||||
SessionLedger,
|
||||
logger,
|
||||
)
|
||||
|
||||
_CATEGORY = "FAL/Platform"
|
||||
|
||||
# Custom ComfyUI type carried between FalSubmit and FalCollect.
|
||||
# Shape: {"endpoint_id": str, "request_id": str}
|
||||
FAL_HANDLE_TYPE = "FAL_HANDLE"
|
||||
|
||||
_FLEXIBLE_RETURN_TYPES = ("IMAGE", "VIDEO", "AUDIO", "STRING")
|
||||
_FLEXIBLE_RETURN_NAMES = ("images", "video", "audio", "result_json")
|
||||
|
||||
_MAX_SEED = 2**31 - 1
|
||||
_SAVE_COUNTER_LIMIT = 100_000
|
||||
|
||||
|
||||
def _collect_result(
|
||||
node_name: str, endpoint_id: str, request_id: str, record_cost: bool = True
|
||||
) -> tuple[Any, Any, Any, str]:
|
||||
"""Fetch a queued result by id and extract flexible outputs from it.
|
||||
|
||||
``record_cost=False`` marks a pure recovery of a past request — the fetch
|
||||
is logged for traceability but adds no new spend to the session ledger.
|
||||
"""
|
||||
endpoint = (endpoint_id or "").strip()
|
||||
request = (request_id or "").strip()
|
||||
if not endpoint or not request:
|
||||
raise FalApiError(node_name, "Both endpoint_id and request_id are required")
|
||||
result = ApiHandler.result_from_request_id(endpoint, request, record_cost=record_cost)
|
||||
return extract_flexible_outputs(result)
|
||||
|
||||
|
||||
class FalSubmit:
|
||||
"""Queue a fal.ai job without waiting; pair with Fal Collect for the result."""
|
||||
|
||||
RETURN_TYPES = (FAL_HANDLE_TYPE, "STRING")
|
||||
RETURN_NAMES = ("handle", "request_id")
|
||||
FUNCTION = "submit"
|
||||
CATEGORY = _CATEGORY
|
||||
DESCRIPTION = (
|
||||
"Submit a job to any fal.ai endpoint and return immediately with a "
|
||||
"handle. Wire several Submits into Collects to run generations in "
|
||||
"parallel instead of one at a time."
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> dict[str, Any]:
|
||||
return {
|
||||
"required": {
|
||||
"endpoint_id": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "fal-ai/kling-video/v3/pro/image-to-video",
|
||||
"tooltip": (
|
||||
"fal endpoint id, e.g. fal-ai/kling-video/v3/pro/image-to-video. "
|
||||
"Submitting queues the job instantly, so multiple Fal Submit nodes "
|
||||
"fan out in parallel; wire each handle into a Fal Collect node to "
|
||||
"wait for its result."
|
||||
),
|
||||
},
|
||||
),
|
||||
"arguments_json": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "{}",
|
||||
"multiline": True,
|
||||
"tooltip": (
|
||||
"JSON object of API arguments. Connected media inputs "
|
||||
"and seed override matching keys here."
|
||||
),
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"image": ("IMAGE", {"tooltip": "Uploaded and sent as image_url"}),
|
||||
"image_2": (
|
||||
"IMAGE",
|
||||
{
|
||||
"tooltip": (
|
||||
"Second image; when set together with 'image', both are "
|
||||
"also sent as image_urls [url1, url2]"
|
||||
)
|
||||
},
|
||||
),
|
||||
"video": ("VIDEO", {"tooltip": "Uploaded and sent as video_url"}),
|
||||
"audio": ("AUDIO", {"tooltip": "Uploaded and sent as audio_url"}),
|
||||
"seed": (
|
||||
"INT",
|
||||
{
|
||||
"default": -1,
|
||||
"min": -1,
|
||||
"max": _MAX_SEED,
|
||||
"control_after_generate": True,
|
||||
"tooltip": "-1 = omit seed; any other value is sent to the API",
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
def submit(
|
||||
self,
|
||||
endpoint_id: str,
|
||||
arguments_json: str = "{}",
|
||||
image: Any = None,
|
||||
image_2: Any = None,
|
||||
video: Any = None,
|
||||
audio: Any = None,
|
||||
seed: int = -1,
|
||||
) -> tuple[dict[str, str], str]:
|
||||
endpoint = (endpoint_id or "").strip()
|
||||
if not endpoint:
|
||||
raise FalApiError("FalSubmit", "endpoint_id is required")
|
||||
|
||||
arguments = build_overlay_arguments(
|
||||
endpoint, arguments_json, image, image_2, video, audio, seed
|
||||
)
|
||||
request_id = ApiHandler.submit_only(endpoint, arguments)
|
||||
logger.info("[%s] submitted request %s", endpoint, request_id)
|
||||
handle = {"endpoint_id": endpoint, "request_id": request_id}
|
||||
return (handle, request_id)
|
||||
|
||||
|
||||
class FalCollect:
|
||||
"""Wait for a queued fal.ai job (from Fal Submit) and extract its outputs."""
|
||||
|
||||
RETURN_TYPES = _FLEXIBLE_RETURN_TYPES
|
||||
RETURN_NAMES = _FLEXIBLE_RETURN_NAMES
|
||||
FUNCTION = "collect"
|
||||
CATEGORY = _CATEGORY
|
||||
DESCRIPTION = (
|
||||
"Block until a job submitted with Fal Submit finishes, then extract "
|
||||
"outputs opportunistically. The raw result is always available as JSON."
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> dict[str, Any]:
|
||||
return {
|
||||
"required": {
|
||||
"handle": (
|
||||
FAL_HANDLE_TYPE,
|
||||
{"tooltip": "Handle from a Fal Submit node"},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
def collect(self, handle: Any) -> tuple[Any, Any, Any, str]:
|
||||
if not isinstance(handle, dict) or not handle.get("endpoint_id") or not handle.get("request_id"):
|
||||
raise FalApiError(
|
||||
"FalCollect",
|
||||
"Expected a FAL_HANDLE dict with 'endpoint_id' and 'request_id' "
|
||||
"keys; connect the handle output of a Fal Submit node.",
|
||||
)
|
||||
return _collect_result("FalCollect", handle["endpoint_id"], handle["request_id"])
|
||||
|
||||
|
||||
class FalResultByRequestId:
|
||||
"""Fetch any past fal.ai generation by its request id, without re-paying."""
|
||||
|
||||
RETURN_TYPES = _FLEXIBLE_RETURN_TYPES
|
||||
RETURN_NAMES = _FLEXIBLE_RETURN_NAMES
|
||||
FUNCTION = "fetch"
|
||||
CATEGORY = _CATEGORY
|
||||
DESCRIPTION = (
|
||||
"Recover a past generation by request id. Results are fetched from "
|
||||
"fal's queue by id, so nothing is re-generated or re-billed."
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> dict[str, Any]:
|
||||
return {
|
||||
"required": {
|
||||
"endpoint_id": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"tooltip": "fal endpoint id the request was originally submitted to",
|
||||
},
|
||||
),
|
||||
"request_id": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"tooltip": (
|
||||
"Recover any past generation from your logs, the session cost "
|
||||
"report, or the fal.ai dashboard without re-paying; the result "
|
||||
"is fetched from fal's queue by this id."
|
||||
),
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
def fetch(self, endpoint_id: str, request_id: str) -> tuple[Any, Any, Any, str]:
|
||||
return _collect_result(
|
||||
"FalResultByRequestId", endpoint_id, request_id, record_cost=False
|
||||
)
|
||||
|
||||
|
||||
class FalCostEstimator:
|
||||
"""Estimate the cost of running an endpoint N times, from registry pricing."""
|
||||
|
||||
RETURN_TYPES = ("STRING", "FLOAT")
|
||||
RETURN_NAMES = ("report", "total_usd")
|
||||
FUNCTION = "estimate"
|
||||
CATEGORY = _CATEGORY
|
||||
DESCRIPTION = (
|
||||
"Estimate USD cost for running a fal endpoint a number of times. "
|
||||
"Never fails: unknown endpoints produce a 'pricing unknown' report."
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> dict[str, Any]:
|
||||
return {
|
||||
"required": {
|
||||
"endpoint_id": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "fal-ai/flux/dev",
|
||||
"tooltip": (
|
||||
"fal endpoint id to estimate. See the model list in the "
|
||||
"README for available endpoint ids."
|
||||
),
|
||||
},
|
||||
),
|
||||
"runs": (
|
||||
"INT",
|
||||
{
|
||||
"default": 1,
|
||||
"min": 1,
|
||||
"max": 10000,
|
||||
"tooltip": "Number of runs to estimate the total cost for",
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
def estimate(self, endpoint_id: str, runs: int = 1) -> tuple[str, float]:
|
||||
endpoint = (endpoint_id or "").strip()
|
||||
try:
|
||||
est = PricingUtils.estimate(endpoint, runs=int(runs))
|
||||
report = PricingUtils.format_report(est)
|
||||
total = est.get("total")
|
||||
except Exception as err: # This node must never fail the graph.
|
||||
logger.warning("FalCostEstimator: could not estimate %s: %s", endpoint, err)
|
||||
report = f"{endpoint or '(no endpoint)'}: pricing unknown ({err})"
|
||||
total = None
|
||||
return (report, float(total) if total is not None else 0.0)
|
||||
|
||||
|
||||
class FalSessionCosts:
|
||||
"""Report every fal API call made this ComfyUI session and its total cost."""
|
||||
|
||||
RETURN_TYPES = ("STRING", "FLOAT")
|
||||
RETURN_NAMES = ("report", "total_usd")
|
||||
FUNCTION = "report"
|
||||
CATEGORY = _CATEGORY
|
||||
OUTPUT_NODE = True
|
||||
DESCRIPTION = (
|
||||
"Report all fal API calls recorded this session (endpoints, request "
|
||||
"ids, estimated costs) and the running total in USD."
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> dict[str, Any]:
|
||||
return {
|
||||
"required": {},
|
||||
"optional": {
|
||||
"reset": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": False,
|
||||
"tooltip": "Clear the ledger after reporting",
|
||||
},
|
||||
),
|
||||
"trigger": (
|
||||
"STRING",
|
||||
{
|
||||
"forceInput": True,
|
||||
"tooltip": (
|
||||
"Connect any upstream text output (e.g. result_json) to "
|
||||
"force this node to run after your generations finish"
|
||||
),
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def IS_CHANGED(cls, **kwargs: Any) -> Any:
|
||||
# The ledger mutates outside the graph; always re-run.
|
||||
return float("nan")
|
||||
|
||||
def report(self, reset: bool = False, trigger: str = "") -> tuple[str, float]:
|
||||
del trigger # Only used to order execution in the graph.
|
||||
ledger = SessionLedger()
|
||||
report_text = ledger.report()
|
||||
total = ledger.total_cost()
|
||||
if reset:
|
||||
ledger.reset()
|
||||
return (report_text, float(total))
|
||||
|
||||
|
||||
def _output_directory() -> str:
|
||||
"""ComfyUI's output directory, or ./output when running outside ComfyUI."""
|
||||
try:
|
||||
import folder_paths
|
||||
except ImportError:
|
||||
return os.path.abspath(os.path.join(os.getcwd(), "output"))
|
||||
return folder_paths.get_output_directory()
|
||||
|
||||
|
||||
def _suffix_from_url(url: str) -> str:
|
||||
suffix = os.path.splitext(urlparse(url).path)[1]
|
||||
return suffix if suffix else ".bin"
|
||||
|
||||
|
||||
def _resolve_save_directory(filename_prefix: str) -> tuple[str, str]:
|
||||
"""Split the prefix into a confined save directory and a basename.
|
||||
|
||||
The resolved directory must stay inside the ComfyUI output directory —
|
||||
a shared workflow must not be able to write outside it via '..' segments.
|
||||
"""
|
||||
prefix = (filename_prefix or "").strip().strip("/") or "fal/media"
|
||||
subdir, basename = os.path.split(prefix)
|
||||
basename = basename or "media"
|
||||
output_root = os.path.realpath(_output_directory())
|
||||
directory = os.path.realpath(os.path.join(output_root, subdir))
|
||||
if directory != output_root and not directory.startswith(output_root + os.sep):
|
||||
raise FalApiError(
|
||||
"FalSaveMediaURL",
|
||||
f"filename_prefix escapes the output directory: {filename_prefix!r}",
|
||||
)
|
||||
return directory, basename
|
||||
|
||||
|
||||
def _claim_unique_destination(directory: str, basename: str, suffix: str) -> str:
|
||||
"""Atomically claim the first free '<basename>_00001<suffix>' path.
|
||||
|
||||
O_CREAT|O_EXCL closes the check-then-act race between concurrent saves.
|
||||
"""
|
||||
for counter in range(1, _SAVE_COUNTER_LIMIT):
|
||||
candidate = os.path.join(directory, f"{basename}_{counter:05d}{suffix}")
|
||||
try:
|
||||
os.close(os.open(candidate, os.O_CREAT | os.O_EXCL | os.O_WRONLY))
|
||||
return candidate
|
||||
except FileExistsError:
|
||||
continue
|
||||
raise FalApiError(
|
||||
"FalSaveMediaURL",
|
||||
f"Could not find a free filename for '{basename}{suffix}' in {directory}",
|
||||
)
|
||||
|
||||
|
||||
_PROVENANCE_VERSION = 1
|
||||
_PNG_PROVENANCE_KEY = "fal_provenance"
|
||||
_SIDECAR_SUFFIX = ".fal.json"
|
||||
|
||||
|
||||
def _build_provenance(target_url: str) -> dict[str, Any]:
|
||||
"""Assemble the provenance receipt for a saved URL; lookup misses are None."""
|
||||
endpoint_id: str | None = None
|
||||
request_id: str | None = None
|
||||
try:
|
||||
found = ResultCache().find_request_by_url(target_url)
|
||||
except Exception as err:
|
||||
logger.warning("FalSaveMediaURL: provenance lookup failed for %s: %s", target_url, err)
|
||||
found = None
|
||||
if found:
|
||||
endpoint_id = found.get("endpoint_id")
|
||||
request_id = found.get("request_id")
|
||||
return {
|
||||
"version": _PROVENANCE_VERSION,
|
||||
"endpoint_id": endpoint_id,
|
||||
"request_id": request_id,
|
||||
"source_url": target_url,
|
||||
"saved_at": time.time(),
|
||||
}
|
||||
|
||||
|
||||
def _write_provenance_sidecar(saved_path: str, provenance: dict[str, Any]) -> None:
|
||||
"""Write '<saved_path>.fal.json' next to the file. Best-effort, never raises."""
|
||||
sidecar_path = saved_path + _SIDECAR_SUFFIX
|
||||
try:
|
||||
with open(sidecar_path, "w", encoding="utf-8") as handle:
|
||||
json.dump(provenance, handle, indent=2)
|
||||
except Exception as err:
|
||||
logger.warning("FalSaveMediaURL: could not write sidecar %s: %s", sidecar_path, err)
|
||||
|
||||
|
||||
def _embed_png_provenance(saved_path: str, provenance: dict[str, Any]) -> None:
|
||||
"""Embed provenance as a PNG text chunk (PNG files only). Never raises."""
|
||||
if not saved_path.lower().endswith(".png"):
|
||||
return
|
||||
try:
|
||||
from PIL import Image
|
||||
from PIL.PngImagePlugin import PngInfo
|
||||
|
||||
info = PngInfo()
|
||||
info.add_text(_PNG_PROVENANCE_KEY, json.dumps(provenance))
|
||||
with Image.open(saved_path) as image:
|
||||
image.load() # read fully before overwriting the same path
|
||||
# carry over the source PNG's existing text chunks and color
|
||||
# profile — a fresh PngInfo would otherwise strip them on re-save
|
||||
for key, value in (getattr(image, "text", {}) or {}).items():
|
||||
if key != _PNG_PROVENANCE_KEY and isinstance(value, str):
|
||||
info.add_text(key, value)
|
||||
save_kwargs: dict[str, Any] = {"pnginfo": info}
|
||||
icc_profile = image.info.get("icc_profile")
|
||||
if icc_profile:
|
||||
save_kwargs["icc_profile"] = icc_profile
|
||||
image.save(saved_path, **save_kwargs)
|
||||
except Exception as err:
|
||||
logger.warning(
|
||||
"FalSaveMediaURL: could not embed PNG provenance in %s: %s", saved_path, err
|
||||
)
|
||||
|
||||
|
||||
def _read_provenance_sidecar(path: str) -> dict[str, Any] | None:
|
||||
"""Load '<path>.fal.json' if present and valid; None otherwise."""
|
||||
sidecar_path = path + _SIDECAR_SUFFIX
|
||||
if not os.path.isfile(sidecar_path):
|
||||
return None
|
||||
try:
|
||||
with open(sidecar_path, encoding="utf-8") as handle:
|
||||
data = json.load(handle)
|
||||
return data if isinstance(data, dict) else None
|
||||
except Exception as err:
|
||||
logger.warning("FalProvenanceFromFile: unreadable sidecar %s: %s", sidecar_path, err)
|
||||
return None
|
||||
|
||||
|
||||
def _read_png_provenance(path: str) -> dict[str, Any] | None:
|
||||
"""Read the 'fal_provenance' PNG text chunk (PNG files only); None otherwise."""
|
||||
if not path.lower().endswith(".png"):
|
||||
return None
|
||||
try:
|
||||
from PIL import Image
|
||||
|
||||
with Image.open(path) as image:
|
||||
raw = getattr(image, "text", {}).get(_PNG_PROVENANCE_KEY)
|
||||
if not raw:
|
||||
return None
|
||||
data = json.loads(raw)
|
||||
return data if isinstance(data, dict) else None
|
||||
except Exception as err:
|
||||
logger.warning(
|
||||
"FalProvenanceFromFile: could not read PNG provenance from %s: %s", path, err
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
class FalSaveMediaURL:
|
||||
"""Download a media URL and save it into the ComfyUI output directory."""
|
||||
|
||||
RETURN_TYPES = ("STRING", "STRING")
|
||||
RETURN_NAMES = ("path", "request_id")
|
||||
FUNCTION = "save"
|
||||
CATEGORY = _CATEGORY
|
||||
OUTPUT_NODE = True
|
||||
DESCRIPTION = (
|
||||
"Download a result URL (video, audio, file, ...) and store it under "
|
||||
"the ComfyUI output directory with a unique, never-overwriting name. "
|
||||
"Every save also writes a provenance receipt (a .fal.json sidecar, "
|
||||
"plus an embedded text chunk for PNGs) so the generation can be "
|
||||
"recovered later for free via 'Fal Provenance from File'."
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> dict[str, Any]:
|
||||
return {
|
||||
"required": {
|
||||
"url": (
|
||||
"STRING",
|
||||
{
|
||||
"forceInput": True,
|
||||
"tooltip": "http(s) URL of the media to download and save",
|
||||
},
|
||||
),
|
||||
"filename_prefix": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "fal/media",
|
||||
"tooltip": "Relative to the ComfyUI output directory; subfolders allowed",
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
def save(self, url: str, filename_prefix: str = "fal/media") -> tuple[str, str]:
|
||||
target_url = (url or "").strip()
|
||||
if not target_url.startswith(("http://", "https://")):
|
||||
raise FalApiError(
|
||||
"FalSaveMediaURL",
|
||||
f"Expected an http(s) URL to save, got: {target_url!r}",
|
||||
)
|
||||
|
||||
directory, basename = _resolve_save_directory(filename_prefix)
|
||||
suffix = _suffix_from_url(target_url)
|
||||
|
||||
temp_path = MediaUtils.download_url_to_temp(target_url, suffix)
|
||||
try:
|
||||
os.makedirs(directory, exist_ok=True)
|
||||
destination = _claim_unique_destination(directory, basename, suffix)
|
||||
shutil.move(temp_path, destination)
|
||||
except FalApiError:
|
||||
raise
|
||||
except Exception as err:
|
||||
logger.error("FalSaveMediaURL: failed to save %s: %s", target_url, err)
|
||||
raise FalApiError(
|
||||
"FalSaveMediaURL", f"Failed to save {target_url}: {err}"
|
||||
) from err
|
||||
finally:
|
||||
if os.path.exists(temp_path):
|
||||
try:
|
||||
os.unlink(temp_path)
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
saved = os.path.abspath(destination)
|
||||
logger.info("FalSaveMediaURL: saved %s -> %s", target_url, saved)
|
||||
|
||||
# Provenance receipt: best-effort, never fails the save itself.
|
||||
provenance = _build_provenance(target_url)
|
||||
_write_provenance_sidecar(saved, provenance)
|
||||
_embed_png_provenance(saved, provenance)
|
||||
request_id = str(provenance.get("request_id") or "")
|
||||
return (saved, request_id)
|
||||
|
||||
|
||||
class FalProvenanceFromFile:
|
||||
"""Read the fal provenance receipt of a previously saved output file."""
|
||||
|
||||
RETURN_TYPES = ("STRING", "STRING", "STRING")
|
||||
RETURN_NAMES = ("endpoint_id", "request_id", "provenance_json")
|
||||
FUNCTION = "read"
|
||||
CATEGORY = _CATEGORY
|
||||
DESCRIPTION = (
|
||||
"Recover where a saved file came from: reads the .fal.json sidecar "
|
||||
"written by 'Fal Save Media from URL' (or the provenance chunk "
|
||||
"embedded in PNGs). Wire endpoint_id and request_id into 'Fal Result "
|
||||
"by Request ID' to re-materialize the generation for free."
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> dict[str, Any]:
|
||||
return {
|
||||
"required": {
|
||||
"file_path": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"tooltip": (
|
||||
"Absolute path of a previously saved output — reads the "
|
||||
".fal.json sidecar or embedded PNG chunk."
|
||||
),
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
def read(self, file_path: str) -> tuple[str, str, str]:
|
||||
path = os.path.expanduser((file_path or "").strip())
|
||||
if not path:
|
||||
raise FalApiError("FalProvenanceFromFile", "file_path is required")
|
||||
if not os.path.isfile(path):
|
||||
raise FalApiError("FalProvenanceFromFile", f"File not found: {path}")
|
||||
|
||||
provenance = _read_provenance_sidecar(path)
|
||||
if provenance is None:
|
||||
provenance = _read_png_provenance(path)
|
||||
if provenance is None:
|
||||
raise FalApiError(
|
||||
"FalProvenanceFromFile",
|
||||
f"No fal provenance found for {path}: expected a "
|
||||
f"'{os.path.basename(path)}{_SIDECAR_SUFFIX}' sidecar next to it, or a "
|
||||
f"'{_PNG_PROVENANCE_KEY}' text chunk inside a PNG. Only files saved by "
|
||||
"'Fal Save Media from URL' carry a provenance receipt.",
|
||||
)
|
||||
|
||||
endpoint_id = str(provenance.get("endpoint_id") or "")
|
||||
request_id = str(provenance.get("request_id") or "")
|
||||
return (endpoint_id, request_id, json.dumps(provenance))
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"FalSubmit_fal": FalSubmit,
|
||||
"FalCollect_fal": FalCollect,
|
||||
"FalResultByRequestId_fal": FalResultByRequestId,
|
||||
"FalCostEstimator_fal": FalCostEstimator,
|
||||
"FalSessionCosts_fal": FalSessionCosts,
|
||||
"FalSaveMediaURL_fal": FalSaveMediaURL,
|
||||
"FalProvenanceFromFile_fal": FalProvenanceFromFile,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"FalSubmit_fal": "Fal Submit (async) (fal)",
|
||||
"FalCollect_fal": "Fal Collect (async result) (fal)",
|
||||
"FalResultByRequestId_fal": "Fal Result by Request ID (fal)",
|
||||
"FalCostEstimator_fal": "Fal Cost Estimator (fal)",
|
||||
"FalSessionCosts_fal": "Fal Session Costs (fal)",
|
||||
"FalSaveMediaURL_fal": "Fal Save Media from URL (fal)",
|
||||
"FalProvenanceFromFile_fal": "Fal Provenance from File (fal)",
|
||||
}
|
||||
@@ -0,0 +1,465 @@
|
||||
"""HTTP routes exposing fal pricing, session costs, jobs and balance to the ComfyUI frontend.
|
||||
|
||||
Registered on ComfyUI's PromptServer under ``/fal_api/*``. The module must
|
||||
import cleanly without ComfyUI (headless/tests): ``register()`` is a no-op when
|
||||
``server`` is unavailable, and every handler is a thin wrapper over a pure
|
||||
function so the payload logic is unit-testable without aiohttp.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import threading
|
||||
import time
|
||||
from typing import Any, Callable
|
||||
|
||||
from .utils.billing import BillingUtils
|
||||
from .utils.ledger import SessionLedger
|
||||
from .utils.logger import logger
|
||||
from .utils.pricing import PricingUtils
|
||||
|
||||
_SESSION_TAIL = 20
|
||||
_DEFAULT_JOB_LIMIT = 50
|
||||
_DEFAULT_SEARCH_LIMIT = 25
|
||||
_MAX_SEARCH_LIMIT = 100
|
||||
|
||||
_cache_lock = threading.Lock()
|
||||
# [(model, pricing_info|None), ...] in registry order; None until first build.
|
||||
_catalog_cache: list[tuple[dict[str, Any], dict[str, Any] | None]] | None = None
|
||||
_pricing_map_cache: dict[str, dict[str, Any]] | None = None
|
||||
|
||||
|
||||
# -- registry catalog (lazy, built once) ---------------------------------------
|
||||
|
||||
|
||||
def _registry_path() -> str:
|
||||
"""Path to data/fal_registry.json at the repo root."""
|
||||
nodes_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
return os.path.join(os.path.dirname(nodes_dir), "data", "fal_registry.json")
|
||||
|
||||
|
||||
def _read_models() -> list[dict[str, Any]]:
|
||||
"""Read registry models; empty list on any 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 [m for m in models if isinstance(m, dict) and m.get("endpoint_id")]
|
||||
except Exception as exc:
|
||||
logger.warning("server_routes: could not load fal registry: %s", exc)
|
||||
return []
|
||||
|
||||
|
||||
def _format_amount(value: float) -> str:
|
||||
"""Format a dollar amount compactly (up to 4 decimals, no trailing zeros)."""
|
||||
text = f"{value:,.4f}".rstrip("0").rstrip(".")
|
||||
return text or "0"
|
||||
|
||||
|
||||
def _pricing_info(parsed: dict[str, Any]) -> dict[str, Any] | None:
|
||||
"""Turn a PricingUtils.parse result into {"label", "per_run"}, or None.
|
||||
|
||||
per_run known -> "≈$X/run"; only per_unit -> "$X per <unit>"; else None.
|
||||
"""
|
||||
per_run = parsed.get("per_run")
|
||||
if isinstance(per_run, (int, float)):
|
||||
return {"label": f"≈${_format_amount(float(per_run))}/run", "per_run": float(per_run)}
|
||||
per_unit = parsed.get("per_unit")
|
||||
unit = parsed.get("unit")
|
||||
if isinstance(per_unit, (int, float)) and unit:
|
||||
return {"label": f"${_format_amount(float(per_unit))} per {unit}", "per_run": None}
|
||||
return None
|
||||
|
||||
|
||||
def _node_key_for(model: dict[str, Any]) -> str:
|
||||
"""Dynamic node class key for a registry model (same as factory.node_key)."""
|
||||
try:
|
||||
from .dynamic.factory import node_key
|
||||
|
||||
return node_key(model)
|
||||
except Exception: # stripped env without the dynamic package's deps
|
||||
return "FalAPI_" + str(model.get("endpoint_id", "")).replace("/", "-")
|
||||
|
||||
|
||||
def _catalog() -> list[tuple[dict[str, Any], dict[str, Any] | None]]:
|
||||
"""Registry models paired with parsed pricing info, cached after first build."""
|
||||
global _catalog_cache
|
||||
if _catalog_cache is not None:
|
||||
return _catalog_cache
|
||||
with _cache_lock:
|
||||
if _catalog_cache is not None:
|
||||
return _catalog_cache
|
||||
_catalog_cache = [
|
||||
(model, _pricing_info(PricingUtils.parse(str(model.get("pricing") or ""))))
|
||||
for model in _read_models()
|
||||
]
|
||||
return _catalog_cache
|
||||
|
||||
|
||||
# -- pure payload builders (unit-tested directly) -------------------------------
|
||||
|
||||
|
||||
def _pricing_map() -> dict[str, dict[str, Any]]:
|
||||
"""{node_class_key: {"label", "per_run"}} for every priced dynamic node."""
|
||||
global _pricing_map_cache
|
||||
if _pricing_map_cache is not None:
|
||||
return _pricing_map_cache
|
||||
mapping = {
|
||||
_node_key_for(model): info for model, info in _catalog() if info is not None
|
||||
}
|
||||
with _cache_lock:
|
||||
_pricing_map_cache = mapping
|
||||
return mapping
|
||||
|
||||
|
||||
def _pricing_single(endpoint_id: str) -> dict[str, Any]:
|
||||
"""Live pricing label for one endpoint; {"label": None} when unknown."""
|
||||
endpoint = (endpoint_id or "").strip()
|
||||
if not endpoint:
|
||||
return {"label": None, "per_run": None}
|
||||
estimate = PricingUtils.estimate(endpoint)
|
||||
per_run = estimate.get("per_run")
|
||||
if isinstance(per_run, (int, float)):
|
||||
return {"label": f"≈${_format_amount(float(per_run))}/run", "per_run": float(per_run)}
|
||||
unit_note = estimate.get("unit_note") or ""
|
||||
if unit_note:
|
||||
return {"label": unit_note, "per_run": None}
|
||||
return {"label": None, "per_run": None}
|
||||
|
||||
|
||||
def _session() -> dict[str, Any]:
|
||||
"""Session ledger totals plus the last few call entries."""
|
||||
ledger = SessionLedger()
|
||||
entries = ledger.entries()
|
||||
return {
|
||||
"total_usd": ledger.total_cost(),
|
||||
"calls": len(entries),
|
||||
"entries": entries[-_SESSION_TAIL:],
|
||||
}
|
||||
|
||||
|
||||
def _jobs(limit: int = _DEFAULT_JOB_LIMIT) -> dict[str, Any]:
|
||||
"""Persistent async-job inbox; degrades to empty when the store is missing."""
|
||||
try:
|
||||
from .utils.job_store import JobStore
|
||||
|
||||
store = JobStore()
|
||||
return {"jobs": store.entries(limit=limit), "counts": store.counts()}
|
||||
except Exception as exc:
|
||||
logger.debug("server_routes: job store unavailable: %s", exc)
|
||||
return {"jobs": [], "counts": {}}
|
||||
|
||||
|
||||
def _balance() -> dict[str, Any]:
|
||||
"""Account credit balance (60s-cached inside BillingUtils)."""
|
||||
return {"balance_usd": BillingUtils.get_balance()}
|
||||
|
||||
|
||||
def _search_models(
|
||||
q: str = "",
|
||||
category: str = "",
|
||||
max_price: float | None = None,
|
||||
limit: int = _DEFAULT_SEARCH_LIMIT,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Search the registry; newest first, filtered by text/category/per-run price."""
|
||||
needle = (q or "").strip().lower()
|
||||
wanted_category = (category or "").strip()
|
||||
capped = max(1, min(int(limit), _MAX_SEARCH_LIMIT))
|
||||
|
||||
def matches(model: dict[str, Any], info: dict[str, Any] | None) -> bool:
|
||||
haystack = f"{model.get('endpoint_id', '')} {model.get('title', '')}".lower()
|
||||
if needle and needle not in haystack:
|
||||
return False
|
||||
if wanted_category and model.get("category") != wanted_category:
|
||||
return False
|
||||
if max_price is not None:
|
||||
per_run = (info or {}).get("per_run")
|
||||
if not isinstance(per_run, (int, float)) or per_run > max_price:
|
||||
return False
|
||||
return True
|
||||
|
||||
hits = [(model, info) for model, info in _catalog() if matches(model, info)]
|
||||
hits.sort(key=lambda pair: str(pair[0].get("published_at") or ""), reverse=True)
|
||||
return [
|
||||
{
|
||||
"endpoint_id": model["endpoint_id"],
|
||||
"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()
|
||||
request = (request_id or "").strip()
|
||||
if not endpoint or not request:
|
||||
return {"ok": False, "error": "endpoint_id and request_id are required"}
|
||||
try:
|
||||
from .utils.config import FalConfig
|
||||
|
||||
FalConfig().get_client().cancel(endpoint, request)
|
||||
logger.info("server_routes: cancelled %s request %s", endpoint, request)
|
||||
return {"ok": True}
|
||||
except Exception as exc:
|
||||
logger.debug("server_routes: cancel %s/%s failed: %s", endpoint, request, exc)
|
||||
return {"ok": False, "error": str(exc)}
|
||||
|
||||
|
||||
# -- aiohttp glue ----------------------------------------------------------------
|
||||
|
||||
|
||||
def _json_response(payload: Any, status: int = 200) -> Any:
|
||||
"""aiohttp JSON response; the raw payload when aiohttp is unavailable (tests)."""
|
||||
try:
|
||||
from aiohttp import web
|
||||
except ImportError:
|
||||
return payload
|
||||
return web.json_response(payload, status=status)
|
||||
|
||||
|
||||
def _guarded(build: Callable[[], Any], route: str) -> Any:
|
||||
"""Run a payload builder; any exception becomes a 500 {"error": ...} JSON."""
|
||||
try:
|
||||
return _json_response(build())
|
||||
except Exception as exc:
|
||||
logger.warning("server_routes: %s failed: %s", route, exc)
|
||||
return _json_response({"error": str(exc)}, status=500)
|
||||
|
||||
|
||||
def _query_int(value: Any, default: int) -> int:
|
||||
try:
|
||||
return int(value)
|
||||
except (TypeError, ValueError):
|
||||
return default
|
||||
|
||||
|
||||
def _query_float(value: Any) -> float | None:
|
||||
try:
|
||||
return float(value)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
|
||||
|
||||
async def pricing_map_route(request: Any) -> Any:
|
||||
return _guarded(_pricing_map, "/fal_api/pricing_map")
|
||||
|
||||
|
||||
async def pricing_route(request: Any) -> Any:
|
||||
endpoint_id = request.query.get("endpoint_id", "")
|
||||
return _guarded(lambda: _pricing_single(endpoint_id), "/fal_api/pricing")
|
||||
|
||||
|
||||
async def session_route(request: Any) -> Any:
|
||||
return _guarded(_session, "/fal_api/session")
|
||||
|
||||
|
||||
async def jobs_route(request: Any) -> Any:
|
||||
limit = _query_int(request.query.get("limit"), _DEFAULT_JOB_LIMIT)
|
||||
limit = max(1, min(limit, _MAX_SEARCH_LIMIT))
|
||||
return _guarded(lambda: _jobs(limit=limit), "/fal_api/jobs")
|
||||
|
||||
|
||||
async def balance_route(request: Any) -> Any:
|
||||
return _guarded(_balance, "/fal_api/balance")
|
||||
|
||||
|
||||
async def models_route(request: Any) -> Any:
|
||||
query = request.query
|
||||
q = query.get("q", "")
|
||||
category = query.get("category", "")
|
||||
max_price = _query_float(query.get("max_price"))
|
||||
limit = _query_int(query.get("limit"), _DEFAULT_SEARCH_LIMIT)
|
||||
return _guarded(
|
||||
lambda: _search_models(q=q, category=category, max_price=max_price, limit=limit),
|
||||
"/fal_api/models",
|
||||
)
|
||||
|
||||
|
||||
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()
|
||||
except Exception:
|
||||
body = {}
|
||||
payload = body if isinstance(body, dict) else {}
|
||||
return _guarded(
|
||||
lambda: _cancel(payload.get("endpoint_id", ""), payload.get("request_id", "")),
|
||||
"/fal_api/cancel",
|
||||
)
|
||||
|
||||
|
||||
ROUTES: tuple[tuple[str, str, Callable[..., Any]], ...] = (
|
||||
("GET", "/fal_api/pricing_map", pricing_map_route),
|
||||
("GET", "/fal_api/pricing", pricing_route),
|
||||
("GET", "/fal_api/session", session_route),
|
||||
("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),
|
||||
)
|
||||
|
||||
|
||||
def register() -> bool:
|
||||
"""Attach the /fal_api routes to ComfyUI's PromptServer. Never raises.
|
||||
|
||||
Returns False (with a debug log) when running headless without ComfyUI.
|
||||
"""
|
||||
try:
|
||||
from server import PromptServer
|
||||
except ImportError:
|
||||
logger.debug("server_routes: ComfyUI server not available; routes not registered")
|
||||
return False
|
||||
try:
|
||||
instance = getattr(PromptServer, "instance", None)
|
||||
if instance is None:
|
||||
logger.debug("server_routes: PromptServer has no instance yet; skipping")
|
||||
return False
|
||||
routes = instance.routes
|
||||
for method, path, handler in ROUTES:
|
||||
adder = routes.get if method == "GET" else routes.post
|
||||
adder(path)(handler)
|
||||
logger.info("server_routes: registered %d /fal_api routes", len(ROUTES))
|
||||
return True
|
||||
except Exception as exc:
|
||||
logger.warning("server_routes: could not register /fal_api routes: %s", exc)
|
||||
return False
|
||||
|
||||
|
||||
register()
|
||||
+387
-157
@@ -1,99 +1,112 @@
|
||||
import os
|
||||
import configparser
|
||||
from fal_client.client import SyncClient
|
||||
import tempfile
|
||||
import zipfile
|
||||
import torch
|
||||
from PIL import Image
|
||||
from .fal_utils import ApiHandler, ArchiveUtils, FalConfig
|
||||
|
||||
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
parent_dir = os.path.dirname(current_dir)
|
||||
config_path = os.path.join(parent_dir, "config.ini")
|
||||
# Initialize FalConfig
|
||||
fal_config = FalConfig()
|
||||
|
||||
config = configparser.ConfigParser()
|
||||
config.read(config_path)
|
||||
|
||||
try:
|
||||
if os.environ.get("FAL_KEY") is not None:
|
||||
print("FAL_KEY found in environment variables")
|
||||
fal_key = os.environ["FAL_KEY"]
|
||||
else:
|
||||
print("FAL_KEY not found in environment variables")
|
||||
fal_key = config['API']['FAL_KEY']
|
||||
print("FAL_KEY found in config.ini")
|
||||
os.environ["FAL_KEY"] = fal_key
|
||||
print("FAL_KEY set in environment variables")
|
||||
|
||||
# Check if FAL key is the default placeholder
|
||||
if fal_key == "<your_fal_api_key_here>":
|
||||
print("WARNING: You are using the default FAL API key placeholder!")
|
||||
print("Please set your actual FAL API key in either:")
|
||||
print("1. The config.ini file under [API] section")
|
||||
print("2. Or as an environment variable named FAL_KEY")
|
||||
print("Get your API key from: https://fal.ai/dashboard/keys")
|
||||
except KeyError:
|
||||
print("Error: FAL_KEY not found in config.ini or environment variables")
|
||||
|
||||
# Create the client with API key
|
||||
fal_client = SyncClient(key=fal_key)
|
||||
|
||||
def create_zip_from_images(images):
|
||||
"""Create a zip file from a list of images."""
|
||||
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
|
||||
"""Create a zip file from a list of images and upload it (returns the URL)."""
|
||||
try:
|
||||
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)}"
|
||||
)
|
||||
|
||||
# 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)
|
||||
|
||||
return fal_client.upload_file(temp_zip.name)
|
||||
|
||||
class FluxLoraTrainerNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
"steps": ("INT", {"default": 1000, "min": 100, "max": 10000, "step": 100}),
|
||||
"create_masks": ("BOOLEAN", {"default": True}),
|
||||
"is_style": ("BOOLEAN", {"default": False}),
|
||||
"images": (
|
||||
"IMAGE",
|
||||
{"tooltip": "Training images. Ignored when images_zip_url is set."},
|
||||
),
|
||||
"steps": (
|
||||
"INT",
|
||||
{
|
||||
"default": 1000,
|
||||
"min": 100,
|
||||
"max": 10000,
|
||||
"step": 100,
|
||||
"tooltip": "Number of training steps.",
|
||||
},
|
||||
),
|
||||
"create_masks": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": True,
|
||||
"tooltip": "Automatically create segmentation masks for subject training.",
|
||||
},
|
||||
),
|
||||
"is_style": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": False,
|
||||
"tooltip": "Enable for style LoRAs instead of subject LoRAs.",
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"trigger_word": ("STRING", {"default": ""}),
|
||||
"images_zip_url": ("STRING", {"default": ""}),
|
||||
"is_input_format_already_preprocessed": ("BOOLEAN", {"default": False}),
|
||||
"data_archive_format": ("STRING", {"default": ""}),
|
||||
}
|
||||
"trigger_word": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"tooltip": "Token used to invoke the trained concept in prompts.",
|
||||
},
|
||||
),
|
||||
"images_zip_url": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"tooltip": "URL of a pre-uploaded zip of training images. Overrides the IMAGE input.",
|
||||
},
|
||||
),
|
||||
"is_input_format_already_preprocessed": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": False,
|
||||
"tooltip": "Set when the archive already contains preprocessed data (images + captions).",
|
||||
},
|
||||
),
|
||||
"data_archive_format": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"tooltip": "Archive format hint (e.g. 'zip') when it cannot be inferred from the URL.",
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("lora_file_url",)
|
||||
FUNCTION = "train_lora"
|
||||
CATEGORY = "FAL/Training"
|
||||
|
||||
def train_lora(self, images, steps, create_masks, is_style, trigger_word="", images_zip_url="",
|
||||
is_input_format_already_preprocessed=False, data_archive_format=""):
|
||||
def train_lora(
|
||||
self,
|
||||
images,
|
||||
steps,
|
||||
create_masks,
|
||||
is_style,
|
||||
trigger_word="",
|
||||
images_zip_url="",
|
||||
is_input_format_already_preprocessed=False,
|
||||
data_archive_format="",
|
||||
):
|
||||
try:
|
||||
# Use provided zip URL if available, otherwise create and upload zip file
|
||||
images_url = images_zip_url if images_zip_url else create_zip_from_images(images)
|
||||
images_url = (
|
||||
images_zip_url if images_zip_url else create_zip_from_images(images)
|
||||
)
|
||||
if not images_url:
|
||||
return ("Error: Unable to upload images.", "")
|
||||
return ApiHandler.handle_text_generation_error(
|
||||
"flux-lora-fast-training", "Failed to upload images"
|
||||
)
|
||||
|
||||
# Prepare arguments for the API
|
||||
arguments = {
|
||||
@@ -103,170 +116,388 @@ class FluxLoraTrainerNode:
|
||||
"is_style": is_style,
|
||||
"is_input_format_already_preprocessed": is_input_format_already_preprocessed,
|
||||
}
|
||||
|
||||
|
||||
if trigger_word:
|
||||
arguments["trigger_word"] = trigger_word
|
||||
|
||||
|
||||
if data_archive_format:
|
||||
arguments["data_archive_format"] = data_archive_format
|
||||
|
||||
# Submit training job
|
||||
handler = fal_client.submit("fal-ai/flux-lora-fast-training", arguments=arguments)
|
||||
result = handler.get()
|
||||
|
||||
result = ApiHandler.submit_and_get_result(
|
||||
"fal-ai/flux-lora-fast-training", arguments
|
||||
)
|
||||
lora_url = result["diffusers_lora_file"]["url"]
|
||||
|
||||
return (lora_url, )
|
||||
return (lora_url,)
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error during LoRA training: {str(e)}")
|
||||
return ("Error: Training failed.", "")
|
||||
return ApiHandler.handle_text_generation_error(
|
||||
"flux-lora-fast-training", e
|
||||
)
|
||||
|
||||
|
||||
class HunyuanVideoLoraTrainerNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
"steps": ("INT", {"default": 1000, "min": 100, "max": 10000, "step": 100}),
|
||||
"images": (
|
||||
"IMAGE",
|
||||
{"tooltip": "Training images. Ignored when images_zip_url is set."},
|
||||
),
|
||||
"steps": (
|
||||
"INT",
|
||||
{
|
||||
"default": 1000,
|
||||
"min": 100,
|
||||
"max": 10000,
|
||||
"step": 100,
|
||||
"tooltip": "Number of training steps.",
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"trigger_word": ("STRING", {"default": ""}),
|
||||
"learning_rate": ("FLOAT", {"default": 0.0001, "min": 0.00001, "max": 0.01}),
|
||||
"do_caption": ("BOOLEAN", {"default": True}),
|
||||
"images_zip_url": ("STRING", {"default": ""}),
|
||||
"data_archive_format": ("STRING", {"default": ""}),
|
||||
}
|
||||
"trigger_word": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"tooltip": "Token used to invoke the trained concept in prompts.",
|
||||
},
|
||||
),
|
||||
"learning_rate": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 0.0001,
|
||||
"min": 0.00001,
|
||||
"max": 0.01,
|
||||
"tooltip": "Training learning rate.",
|
||||
},
|
||||
),
|
||||
"do_caption": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": True,
|
||||
"tooltip": "Automatically caption the training images.",
|
||||
},
|
||||
),
|
||||
"images_zip_url": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"tooltip": "URL of a pre-uploaded zip of training images. Overrides the IMAGE input.",
|
||||
},
|
||||
),
|
||||
"data_archive_format": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"tooltip": "Archive format hint (e.g. 'zip') when it cannot be inferred from the URL.",
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("lora_file_url",)
|
||||
FUNCTION = "train_lora"
|
||||
CATEGORY = "FAL/Training"
|
||||
|
||||
def train_lora(self, images, steps, trigger_word="", learning_rate=0.0001, do_caption=True,
|
||||
images_zip_url="", data_archive_format=""):
|
||||
def train_lora(
|
||||
self,
|
||||
images,
|
||||
steps,
|
||||
trigger_word="",
|
||||
learning_rate=0.0001,
|
||||
do_caption=True,
|
||||
images_zip_url="",
|
||||
data_archive_format="",
|
||||
):
|
||||
try:
|
||||
# Use provided zip URL if available, otherwise create and upload zip file
|
||||
images_url = images_zip_url if images_zip_url else create_zip_from_images(images)
|
||||
images_url = (
|
||||
images_zip_url if images_zip_url else create_zip_from_images(images)
|
||||
)
|
||||
if not images_url:
|
||||
return ("Error: Unable to upload images.", "")
|
||||
return ApiHandler.handle_text_generation_error(
|
||||
"hunyuan-video-lora-training", "Failed to upload images"
|
||||
)
|
||||
|
||||
# Prepare arguments for the API
|
||||
arguments = {
|
||||
"images_data_url": images_url,
|
||||
"steps": steps,
|
||||
"learning_rate": learning_rate,
|
||||
"do_caption": do_caption
|
||||
"do_caption": do_caption,
|
||||
}
|
||||
|
||||
|
||||
if trigger_word:
|
||||
arguments["trigger_word"] = trigger_word
|
||||
|
||||
|
||||
if data_archive_format:
|
||||
arguments["data_archive_format"] = data_archive_format
|
||||
|
||||
# Submit training job
|
||||
handler = fal_client.submit("fal-ai/hunyuan-video-lora-training", arguments=arguments)
|
||||
result = handler.get()
|
||||
|
||||
result = ApiHandler.submit_and_get_result(
|
||||
"fal-ai/hunyuan-video-lora-training", arguments
|
||||
)
|
||||
lora_url = result["diffusers_lora_file"]["url"]
|
||||
|
||||
return (lora_url,)
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error during LoRA training: {str(e)}")
|
||||
return ("Error: Training failed.", "")
|
||||
return ApiHandler.handle_text_generation_error(
|
||||
"hunyuan-video-lora-training", e
|
||||
)
|
||||
|
||||
|
||||
class WanLoraTrainerNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"training_data_url": ("STRING", {"default": ""}),
|
||||
"number_of_steps": ("INT", {"default": 400, "min": 5, "max": 10000, "step": 1}),
|
||||
"learning_rate": ("FLOAT", {"default": 0.0002, "min": 0.00001, "max": 0.01}),
|
||||
"training_data_url": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"tooltip": "URL of the training data archive (images/videos with optional captions).",
|
||||
},
|
||||
),
|
||||
"number_of_steps": (
|
||||
"INT",
|
||||
{
|
||||
"default": 400,
|
||||
"min": 5,
|
||||
"max": 10000,
|
||||
"step": 1,
|
||||
"tooltip": "Number of training steps.",
|
||||
},
|
||||
),
|
||||
"learning_rate": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 0.0002,
|
||||
"min": 0.00001,
|
||||
"max": 0.01,
|
||||
"tooltip": "Training learning rate.",
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"trigger_phrase": ("STRING", {"default": ""}),
|
||||
"auto_scale_input": ("BOOLEAN", {"default": True}),
|
||||
}
|
||||
"trigger_phrase": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"tooltip": "Phrase used to invoke the trained concept in prompts.",
|
||||
},
|
||||
),
|
||||
"auto_scale_input": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": True,
|
||||
"tooltip": "Automatically rescale input media to the training resolution.",
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("lora_file_url",)
|
||||
FUNCTION = "train_lora"
|
||||
CATEGORY = "FAL/Training"
|
||||
|
||||
def train_lora(self, training_data_url, number_of_steps, learning_rate, trigger_phrase="", auto_scale_input=True):
|
||||
def train_lora(
|
||||
self,
|
||||
training_data_url,
|
||||
number_of_steps,
|
||||
learning_rate,
|
||||
trigger_phrase="",
|
||||
auto_scale_input=True,
|
||||
):
|
||||
try:
|
||||
if not training_data_url:
|
||||
return ("Error: No training data URL provided.",)
|
||||
return ApiHandler.handle_text_generation_error(
|
||||
"wan-trainer", "No training data URL provided"
|
||||
)
|
||||
|
||||
# Prepare arguments for the API
|
||||
arguments = {
|
||||
"training_data_url": training_data_url,
|
||||
"number_of_steps": number_of_steps,
|
||||
"learning_rate": learning_rate,
|
||||
"auto_scale_input": auto_scale_input
|
||||
"auto_scale_input": auto_scale_input,
|
||||
}
|
||||
|
||||
|
||||
if trigger_phrase:
|
||||
arguments["trigger_phrase"] = trigger_phrase
|
||||
|
||||
# Submit training job
|
||||
handler = fal_client.submit("fal-ai/wan-trainer", arguments=arguments)
|
||||
result = handler.get()
|
||||
|
||||
result = ApiHandler.submit_and_get_result("fal-ai/wan-trainer", arguments)
|
||||
lora_url = result["lora_file"]["url"]
|
||||
|
||||
return (lora_url,)
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error during LoRA training: {str(e)}")
|
||||
return ("Error: Training failed.",)
|
||||
return ApiHandler.handle_text_generation_error("wan-trainer", e)
|
||||
|
||||
|
||||
class LtxVideoTrainerNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"training_data_url": ("STRING", {"default": ""}),
|
||||
"rank": (["8", "16", "32", "64", "128"], {"default": "128"}),
|
||||
"number_of_steps": ("INT", {"default": 1000, "min": 100, "max": 10000, "step": 1}),
|
||||
"number_of_frames": ("INT", {"default": 81, "min": 1, "max": 1000}),
|
||||
"frame_rate": ("INT", {"default": 25, "min": 1, "max": 60}),
|
||||
"resolution": (["low", "medium", "high"], {"default": "medium"}),
|
||||
"aspect_ratio": (["16:9", "1:1", "9:16"], {"default": "1:1"}),
|
||||
"learning_rate": ("FLOAT", {"default": 0.0002, "min": 0.00001, "max": 0.01}),
|
||||
"training_data_url": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"tooltip": "URL of the training data archive (videos/images with optional captions).",
|
||||
},
|
||||
),
|
||||
"rank": (
|
||||
["8", "16", "32", "64", "128"],
|
||||
{
|
||||
"default": "128",
|
||||
"tooltip": "LoRA rank. Higher rank captures more detail but produces larger files.",
|
||||
},
|
||||
),
|
||||
"number_of_steps": (
|
||||
"INT",
|
||||
{
|
||||
"default": 1000,
|
||||
"min": 100,
|
||||
"max": 10000,
|
||||
"step": 1,
|
||||
"tooltip": "Number of training steps.",
|
||||
},
|
||||
),
|
||||
"number_of_frames": (
|
||||
"INT",
|
||||
{
|
||||
"default": 81,
|
||||
"min": 1,
|
||||
"max": 1000,
|
||||
"tooltip": "Frames per training sample.",
|
||||
},
|
||||
),
|
||||
"frame_rate": (
|
||||
"INT",
|
||||
{
|
||||
"default": 25,
|
||||
"min": 1,
|
||||
"max": 60,
|
||||
"tooltip": "Frame rate used for training samples.",
|
||||
},
|
||||
),
|
||||
"resolution": (
|
||||
["low", "medium", "high"],
|
||||
{"default": "medium", "tooltip": "Training resolution."},
|
||||
),
|
||||
"aspect_ratio": (
|
||||
["16:9", "1:1", "9:16"],
|
||||
{"default": "1:1", "tooltip": "Training aspect ratio."},
|
||||
),
|
||||
"learning_rate": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 0.0002,
|
||||
"min": 0.00001,
|
||||
"max": 0.01,
|
||||
"tooltip": "Training learning rate.",
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"trigger_phrase": ("STRING", {"default": ""}),
|
||||
"auto_scale_input": ("BOOLEAN", {"default": False}),
|
||||
"split_input_into_scenes": ("BOOLEAN", {"default": True}),
|
||||
"split_input_duration_threshold": ("FLOAT", {"default": 30.0, "min": 1.0, "max": 300.0}),
|
||||
"validation_negative_prompt": ("STRING", {"default": "blurry, low quality, bad quality, out of focus"}),
|
||||
"validation_number_of_frames": ("INT", {"default": 81, "min": 1, "max": 1000}),
|
||||
"validation_resolution": (["low", "medium", "high"], {"default": "high"}),
|
||||
"validation_aspect_ratio": (["16:9", "1:1", "9:16"], {"default": "1:1"}),
|
||||
"validation_reverse": ("BOOLEAN", {"default": False}),
|
||||
}
|
||||
"trigger_phrase": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"tooltip": "Phrase used to invoke the trained concept in prompts.",
|
||||
},
|
||||
),
|
||||
"auto_scale_input": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": False,
|
||||
"tooltip": "Automatically rescale input media to the training resolution.",
|
||||
},
|
||||
),
|
||||
"split_input_into_scenes": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": True,
|
||||
"tooltip": "Split long input videos into individual scenes before training.",
|
||||
},
|
||||
),
|
||||
"split_input_duration_threshold": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 30.0,
|
||||
"min": 1.0,
|
||||
"max": 300.0,
|
||||
"tooltip": "Videos longer than this many seconds are split into scenes.",
|
||||
},
|
||||
),
|
||||
"validation_negative_prompt": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "blurry, low quality, bad quality, out of focus",
|
||||
"tooltip": "Negative prompt used for validation renders during training.",
|
||||
},
|
||||
),
|
||||
"validation_number_of_frames": (
|
||||
"INT",
|
||||
{
|
||||
"default": 81,
|
||||
"min": 1,
|
||||
"max": 1000,
|
||||
"tooltip": "Frames per validation render.",
|
||||
},
|
||||
),
|
||||
"validation_resolution": (
|
||||
["low", "medium", "high"],
|
||||
{"default": "high", "tooltip": "Resolution of validation renders."},
|
||||
),
|
||||
"validation_aspect_ratio": (
|
||||
["16:9", "1:1", "9:16"],
|
||||
{"default": "1:1", "tooltip": "Aspect ratio of validation renders."},
|
||||
),
|
||||
"validation_reverse": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": False,
|
||||
"tooltip": "Also render reversed validation videos.",
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("lora_file_url",)
|
||||
FUNCTION = "train_lora"
|
||||
CATEGORY = "FAL/Training"
|
||||
|
||||
def train_lora(self, training_data_url, rank, number_of_steps, number_of_frames, frame_rate,
|
||||
resolution, aspect_ratio, learning_rate, trigger_phrase="", auto_scale_input=False,
|
||||
split_input_into_scenes=True, split_input_duration_threshold=30.0,
|
||||
validation_negative_prompt="blurry, low quality, bad quality, out of focus",
|
||||
validation_number_of_frames=81, validation_resolution="high",
|
||||
validation_aspect_ratio="1:1", validation_reverse=False):
|
||||
def train_lora(
|
||||
self,
|
||||
training_data_url,
|
||||
rank,
|
||||
number_of_steps,
|
||||
number_of_frames,
|
||||
frame_rate,
|
||||
resolution,
|
||||
aspect_ratio,
|
||||
learning_rate,
|
||||
trigger_phrase="",
|
||||
auto_scale_input=False,
|
||||
split_input_into_scenes=True,
|
||||
split_input_duration_threshold=30.0,
|
||||
validation_negative_prompt="blurry, low quality, bad quality, out of focus",
|
||||
validation_number_of_frames=81,
|
||||
validation_resolution="high",
|
||||
validation_aspect_ratio="1:1",
|
||||
validation_reverse=False,
|
||||
):
|
||||
try:
|
||||
if not training_data_url:
|
||||
return ("Error: No training data URL provided.",)
|
||||
return ApiHandler.handle_text_generation_error(
|
||||
"ltx-video-trainer", "No training data URL provided"
|
||||
)
|
||||
|
||||
# Prepare arguments for the API
|
||||
arguments = {
|
||||
@@ -285,23 +516,22 @@ class LtxVideoTrainerNode:
|
||||
"validation_number_of_frames": validation_number_of_frames,
|
||||
"validation_resolution": validation_resolution,
|
||||
"validation_aspect_ratio": validation_aspect_ratio,
|
||||
"validation_reverse": validation_reverse
|
||||
"validation_reverse": validation_reverse,
|
||||
}
|
||||
|
||||
|
||||
if trigger_phrase:
|
||||
arguments["trigger_phrase"] = trigger_phrase
|
||||
|
||||
# Submit training job
|
||||
handler = fal_client.submit("fal-ai/ltx-video-trainer", arguments=arguments)
|
||||
result = handler.get()
|
||||
|
||||
result = ApiHandler.submit_and_get_result(
|
||||
"fal-ai/ltx-video-trainer", arguments
|
||||
)
|
||||
lora_url = result["lora_file"]["url"]
|
||||
|
||||
return (lora_url,)
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error during LoRA training: {str(e)}")
|
||||
return ("Error: Training failed.",)
|
||||
return ApiHandler.handle_text_generation_error("ltx-video-trainer", e)
|
||||
|
||||
|
||||
# Node class mappings
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
@@ -317,4 +547,4 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"HunyuanVideoLoraTrainer_fal": "Hunyuan Video LoRA Trainer (fal)",
|
||||
"WanLoraTrainer_fal": "WAN LoRA Trainer (fal)",
|
||||
"LtxVideoTrainer_fal": "LTX Video LoRA Trainer (fal)",
|
||||
}
|
||||
}
|
||||
|
||||
+486
-122
@@ -1,157 +1,521 @@
|
||||
import os
|
||||
import configparser
|
||||
import tempfile
|
||||
import requests
|
||||
from PIL import Image
|
||||
import io
|
||||
import numpy as np
|
||||
import torch
|
||||
from fal_client.client import SyncClient
|
||||
from .fal_utils import ApiHandler, FalConfig, ImageUtils, ResultProcessor
|
||||
|
||||
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
parent_dir = os.path.dirname(current_dir)
|
||||
config_path = os.path.join(parent_dir, "config.ini")
|
||||
# Initialize FalConfig
|
||||
fal_config = FalConfig()
|
||||
|
||||
config = configparser.ConfigParser()
|
||||
config.read(config_path)
|
||||
|
||||
try:
|
||||
if os.environ.get("FAL_KEY") is not None:
|
||||
print("FAL_KEY found in environment variables")
|
||||
fal_key = os.environ["FAL_KEY"]
|
||||
else:
|
||||
print("FAL_KEY not found in environment variables")
|
||||
fal_key = config['API']['FAL_KEY']
|
||||
print("FAL_KEY found in config.ini")
|
||||
os.environ["FAL_KEY"] = fal_key
|
||||
print("FAL_KEY set in environment variables")
|
||||
|
||||
# Check if FAL key is the default placeholder
|
||||
if fal_key == "<your_fal_api_key_here>":
|
||||
print("WARNING: You are using the default FAL API key placeholder!")
|
||||
print("Please set your actual FAL API key in either:")
|
||||
print("1. The config.ini file under [API] section")
|
||||
print("2. Or as an environment variable named FAL_KEY")
|
||||
print("Get your API key from: https://fal.ai/dashboard/keys")
|
||||
except KeyError:
|
||||
print("Error: FAL_KEY not found in config.ini or environment variables")
|
||||
|
||||
# Create the client with API key
|
||||
fal_client = SyncClient(key=fal_key)
|
||||
|
||||
def upload_image(image):
|
||||
try:
|
||||
if isinstance(image, torch.Tensor):
|
||||
image_np = image.cpu().numpy()
|
||||
else:
|
||||
image_np = np.array(image)
|
||||
|
||||
if image_np.ndim == 4:
|
||||
image_np = image_np.squeeze(0)
|
||||
if image_np.ndim == 2:
|
||||
image_np = np.stack([image_np] * 3, axis=-1)
|
||||
elif image_np.shape[0] == 3:
|
||||
image_np = np.transpose(image_np, (1, 2, 0))
|
||||
|
||||
if image_np.dtype == np.float32 or image_np.dtype == np.float64:
|
||||
image_np = (image_np * 255).astype(np.uint8)
|
||||
|
||||
pil_image = Image.fromarray(image_np)
|
||||
|
||||
with tempfile.NamedTemporaryFile(suffix=".png", delete=False) as temp_file:
|
||||
pil_image.save(temp_file, format="PNG")
|
||||
temp_file_path = temp_file.name
|
||||
|
||||
image_url = fal_client.upload_file(temp_file_path)
|
||||
return image_url
|
||||
except Exception as e:
|
||||
print(f"Error uploading image: {str(e)}")
|
||||
return None
|
||||
finally:
|
||||
if 'temp_file_path' in locals():
|
||||
os.unlink(temp_file_path)
|
||||
|
||||
class UpscalerNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
"upscale_factor": ("FLOAT", {"default": 2.0, "min": 1.0, "max": 4.0, "step": 0.5}),
|
||||
"negative_prompt": ("STRING", {"default": "(worst quality, low quality, normal quality:2)", "multiline": True}),
|
||||
"creativity": ("FLOAT", {"default": 0.35, "min": 0.0, "max": 1.0, "step": 0.05}),
|
||||
"resemblance": ("FLOAT", {"default": 0.6, "min": 0.0, "max": 1.0, "step": 0.05}),
|
||||
"guidance_scale": ("FLOAT", {"default": 4.0, "min": 1.0, "max": 20.0, "step": 0.5}),
|
||||
"num_inference_steps": ("INT", {"default": 18, "min": 1, "max": 100}),
|
||||
"enable_safety_checker": ("BOOLEAN", {"default": True}),
|
||||
"image": ("IMAGE", {"tooltip": "Input image to upscale."}),
|
||||
"upscale_factor": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 2.0,
|
||||
"min": 1.0,
|
||||
"max": 4.0,
|
||||
"step": 0.5,
|
||||
"tooltip": "How much to enlarge the image (1x-4x).",
|
||||
},
|
||||
),
|
||||
"negative_prompt": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "(worst quality, low quality, normal quality:2)",
|
||||
"multiline": True,
|
||||
"tooltip": "Concepts to avoid during the creative upscale.",
|
||||
},
|
||||
),
|
||||
"creativity": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 0.35,
|
||||
"min": 0.0,
|
||||
"max": 1.0,
|
||||
"step": 0.05,
|
||||
"tooltip": "Higher values allow the model to invent more detail.",
|
||||
},
|
||||
),
|
||||
"resemblance": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 0.6,
|
||||
"min": 0.0,
|
||||
"max": 1.0,
|
||||
"step": 0.05,
|
||||
"tooltip": "Higher values keep the result closer to the input image.",
|
||||
},
|
||||
),
|
||||
"guidance_scale": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 4.0,
|
||||
"min": 1.0,
|
||||
"max": 20.0,
|
||||
"step": 0.5,
|
||||
"tooltip": "Classifier-free guidance scale for the diffusion pass.",
|
||||
},
|
||||
),
|
||||
"num_inference_steps": (
|
||||
"INT",
|
||||
{
|
||||
"default": 18,
|
||||
"min": 1,
|
||||
"max": 100,
|
||||
"tooltip": "Number of diffusion steps; more steps is slower but can add detail.",
|
||||
},
|
||||
),
|
||||
"enable_safety_checker": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": True,
|
||||
"tooltip": "Filter potentially unsafe output images.",
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"seed": ("INT", {"default": -1}),
|
||||
}
|
||||
"seed": (
|
||||
"INT",
|
||||
{
|
||||
"default": -1,
|
||||
"tooltip": "Random seed for reproducibility. -1 uses a random seed.",
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "generate_upscaled_image"
|
||||
CATEGORY = "FAL/Image"
|
||||
|
||||
def generate_upscaled_image(self, image, upscale_factor, negative_prompt, creativity, resemblance, guidance_scale, num_inference_steps, enable_safety_checker, seed=-1):
|
||||
image_url = upload_image(image)
|
||||
if not image_url:
|
||||
print("Failed to upload image for upscaling.")
|
||||
return self.create_blank_image()
|
||||
def generate_upscaled_image(
|
||||
self,
|
||||
image,
|
||||
upscale_factor,
|
||||
negative_prompt,
|
||||
creativity,
|
||||
resemblance,
|
||||
guidance_scale,
|
||||
num_inference_steps,
|
||||
enable_safety_checker,
|
||||
seed=-1,
|
||||
):
|
||||
try:
|
||||
# Upload the image using ImageUtils (raises on failure)
|
||||
image_url = ImageUtils.upload_image(image)
|
||||
|
||||
arguments = {
|
||||
"image_url": image_url,
|
||||
"prompt": "masterpiece, best quality, highres",
|
||||
"upscale_factor": upscale_factor,
|
||||
"negative_prompt": negative_prompt,
|
||||
"creativity": creativity,
|
||||
"resemblance": resemblance,
|
||||
"guidance_scale": guidance_scale,
|
||||
"num_inference_steps": num_inference_steps,
|
||||
"enable_safety_checker": enable_safety_checker
|
||||
arguments = {
|
||||
"image_url": image_url,
|
||||
"prompt": "masterpiece, best quality, highres",
|
||||
"upscale_factor": upscale_factor,
|
||||
"negative_prompt": negative_prompt,
|
||||
"creativity": creativity,
|
||||
"resemblance": resemblance,
|
||||
"guidance_scale": guidance_scale,
|
||||
"num_inference_steps": num_inference_steps,
|
||||
"enable_safety_checker": enable_safety_checker,
|
||||
}
|
||||
|
||||
if seed != -1:
|
||||
arguments["seed"] = seed
|
||||
|
||||
result = ApiHandler.submit_and_get_result(
|
||||
"fal-ai/clarity-upscaler", arguments
|
||||
)
|
||||
return ResultProcessor.process_image_result(result)
|
||||
except Exception as e:
|
||||
return ApiHandler.handle_image_generation_error("clarity-upscaler", e)
|
||||
|
||||
|
||||
class SeedvrUpscalerNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE", {"tooltip": "Input image to upscale."}),
|
||||
"upscale_factor": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 2.0,
|
||||
"min": 1.0,
|
||||
"max": 4.0,
|
||||
"step": 0.5,
|
||||
"tooltip": "How much to enlarge the image (1x-4x).",
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"seed": (
|
||||
"INT",
|
||||
{
|
||||
"default": -1,
|
||||
"tooltip": "Random seed for reproducibility. -1 uses a random seed.",
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
if seed != -1:
|
||||
arguments["seed"] = seed
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "generate_upscaled_image"
|
||||
CATEGORY = "FAL/Image"
|
||||
|
||||
def generate_upscaled_image(
|
||||
self,
|
||||
image,
|
||||
upscale_factor,
|
||||
seed=-1,
|
||||
):
|
||||
try:
|
||||
handler = fal_client.submit("fal-ai/clarity-upscaler", arguments=arguments)
|
||||
result = handler.get()
|
||||
return self.process_result(result)
|
||||
except Exception as e:
|
||||
print(f"Error generating upscaled image: {str(e)}")
|
||||
return self.create_blank_image()
|
||||
# Upload the image using ImageUtils (raises on failure)
|
||||
image_url = ImageUtils.upload_image(image)
|
||||
|
||||
def process_result(self, result):
|
||||
arguments = {
|
||||
"image_url": image_url,
|
||||
"upscale_factor": upscale_factor,
|
||||
}
|
||||
|
||||
if seed != -1:
|
||||
arguments["seed"] = seed
|
||||
|
||||
result = ApiHandler.submit_and_get_result(
|
||||
"fal-ai/seedvr/upscale/image", arguments
|
||||
)
|
||||
return ResultProcessor.process_single_image_result(result)
|
||||
except Exception as e:
|
||||
return ApiHandler.handle_image_generation_error("seedvr-upscaler", e)
|
||||
|
||||
|
||||
class SeedvrUpscaleVideoNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"upscale_factor": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 2.0,
|
||||
"min": 0.00,
|
||||
"max": 5.0,
|
||||
"step": 0.01,
|
||||
"tooltip": "Upscaling factor applied when upscale_mode is 'factor'.",
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"video": (
|
||||
"VIDEO",
|
||||
{
|
||||
"tooltip": "Video to upscale. Takes precedence over input_video_url when connected.",
|
||||
},
|
||||
),
|
||||
"input_video_url": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"tooltip": "URL of the video to upscale. Used when no VIDEO input is connected.",
|
||||
},
|
||||
),
|
||||
"upscale_mode": (
|
||||
["factor", "target"],
|
||||
{
|
||||
"default": "factor",
|
||||
"tooltip": "'factor' scales by upscale_factor; 'target' scales to target_resolution.",
|
||||
},
|
||||
),
|
||||
"target_resolution": (
|
||||
["720p", "1080p", "1440p", "2160p"],
|
||||
{
|
||||
"default": "1080p",
|
||||
"tooltip": "Output resolution used when upscale_mode is 'target'.",
|
||||
},
|
||||
),
|
||||
"noise_scale": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 0.1,
|
||||
"min": 0.0,
|
||||
"max": 1.0,
|
||||
"step": 0.05,
|
||||
"tooltip": "Amount of noise conditioning; higher can hallucinate more detail.",
|
||||
},
|
||||
),
|
||||
"output_quality": (
|
||||
["low", "medium", "high", "maximum"],
|
||||
{
|
||||
"default": "high",
|
||||
"tooltip": "Encoding quality of the output video.",
|
||||
},
|
||||
),
|
||||
"output_write_mode": (
|
||||
["fast", "balanced", "small"],
|
||||
{
|
||||
"default": "balanced",
|
||||
"tooltip": "Encoder speed/size trade-off for writing the output file.",
|
||||
},
|
||||
),
|
||||
"output_format": (
|
||||
[
|
||||
"X264 (.mp4)",
|
||||
"VP9 (.webm)",
|
||||
"PRORES444 (.mov)",
|
||||
"GIF (.gif)",
|
||||
],
|
||||
{
|
||||
"default": "X264 (.mp4)",
|
||||
"tooltip": "Container and codec of the output video.",
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("video_url",)
|
||||
FUNCTION = "generate_upscaled_video"
|
||||
CATEGORY = "FAL/VideoUpscaling"
|
||||
|
||||
def generate_upscaled_video(
|
||||
self,
|
||||
upscale_factor=2.0,
|
||||
video=None,
|
||||
input_video_url=None,
|
||||
upscale_mode="factor",
|
||||
target_resolution="1080p",
|
||||
noise_scale=0.1,
|
||||
output_format="X264 (.mp4)",
|
||||
output_quality="high",
|
||||
output_write_mode="balanced",
|
||||
):
|
||||
try:
|
||||
img_url = result["image"]["url"]
|
||||
img_response = requests.get(img_url)
|
||||
img = Image.open(io.BytesIO(img_response.content))
|
||||
img_array = np.array(img).astype(np.float32) / 255.0
|
||||
video_url = input_video_url
|
||||
if video is not None:
|
||||
video_url = ImageUtils.upload_file(video.get_stream_source())
|
||||
if not video_url:
|
||||
return ApiHandler.handle_video_generation_error(
|
||||
"seedvr-upscale-video",
|
||||
"No video provided. Connect a VIDEO input or set input_video_url.",
|
||||
)
|
||||
|
||||
# Stack the images along a new first dimension
|
||||
stacked_images = np.stack([img_array], axis=0)
|
||||
|
||||
# Convert to PyTorch tensor
|
||||
img_tensor = torch.from_numpy(stacked_images)
|
||||
return (img_tensor,)
|
||||
# The API enum is "PRORES4444 (.mov)"; the dropdown historically
|
||||
# exposes "PRORES444 (.mov)", so translate at the argument level.
|
||||
api_output_format = (
|
||||
"PRORES4444 (.mov)"
|
||||
if output_format == "PRORES444 (.mov)"
|
||||
else output_format
|
||||
)
|
||||
|
||||
arguments = {
|
||||
"video_url": video_url,
|
||||
"upscale_mode": upscale_mode,
|
||||
"upscale_factor": upscale_factor,
|
||||
"target_resolution": target_resolution,
|
||||
"noise_scale": noise_scale,
|
||||
"output_format": api_output_format,
|
||||
"output_quality": output_quality,
|
||||
"output_write_mode": output_write_mode,
|
||||
}
|
||||
|
||||
result = ApiHandler.submit_and_get_result(
|
||||
"fal-ai/seedvr/upscale/video", arguments
|
||||
)
|
||||
return (result["video"]["url"],)
|
||||
except Exception as e:
|
||||
print(f"Error processing result: {str(e)}")
|
||||
return self.create_blank_image()
|
||||
return ApiHandler.handle_video_generation_error(
|
||||
"seedvr-upscale-video", e
|
||||
)
|
||||
|
||||
|
||||
class BriaVideoIncreaseResolutionNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"upscale_factor": (
|
||||
"INT",
|
||||
{
|
||||
"default": 2,
|
||||
"min": 2,
|
||||
"max": 4,
|
||||
"step": 2,
|
||||
"tooltip": "Resolution increase factor. The API accepts 2 or 4.",
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"video": (
|
||||
"VIDEO",
|
||||
{
|
||||
"tooltip": "Video to upscale. Takes precedence over input_video_url when connected.",
|
||||
},
|
||||
),
|
||||
"input_video_url": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"tooltip": "URL of the video to upscale. Used when no VIDEO input is connected.",
|
||||
},
|
||||
),
|
||||
"output_container_and_codec": (
|
||||
[
|
||||
"mp4_h264",
|
||||
"mp4_h265",
|
||||
"mov_h265",
|
||||
"mov_proresks",
|
||||
"webm_vp9",
|
||||
"mkv_h265",
|
||||
"mkv_vp9",
|
||||
"gif",
|
||||
],
|
||||
{
|
||||
"default": "mp4_h264",
|
||||
"tooltip": "Container and codec of the output video.",
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("video_url",)
|
||||
FUNCTION = "generate_upscaled_video"
|
||||
CATEGORY = "FAL/VideoUpscaling"
|
||||
|
||||
def generate_upscaled_video(
|
||||
self,
|
||||
video=None,
|
||||
input_video_url=None,
|
||||
upscale_factor=2,
|
||||
output_container_and_codec="mp4_h264",
|
||||
):
|
||||
try:
|
||||
video_url = input_video_url
|
||||
if video is not None:
|
||||
video_url = ImageUtils.upload_file(video.get_stream_source())
|
||||
if not video_url:
|
||||
return ApiHandler.handle_video_generation_error(
|
||||
"bria-video-increase-resolution",
|
||||
"No video provided. Connect a VIDEO input or set input_video_url.",
|
||||
)
|
||||
|
||||
arguments = {
|
||||
"video_url": video_url,
|
||||
"desired_increase": str(upscale_factor),
|
||||
"output_container_and_codec": output_container_and_codec,
|
||||
}
|
||||
|
||||
result = ApiHandler.submit_and_get_result(
|
||||
"bria/video/increase-resolution", arguments
|
||||
)
|
||||
return (result["video"]["url"],)
|
||||
except Exception as e:
|
||||
return ApiHandler.handle_video_generation_error(
|
||||
"bria-video-increase-resolution", e
|
||||
)
|
||||
|
||||
|
||||
class TopazUpscaleVideoNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"upscale_factor": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 2.0,
|
||||
"min": 1.0,
|
||||
"max": 5.0,
|
||||
"step": 0.1,
|
||||
"tooltip": "How much to enlarge the video (1x-5x).",
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"video": (
|
||||
"VIDEO",
|
||||
{
|
||||
"tooltip": "Video to upscale. Takes precedence over input_video_url when connected.",
|
||||
},
|
||||
),
|
||||
"input_video_url": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"tooltip": "URL of the video to upscale. Used when no VIDEO input is connected.",
|
||||
},
|
||||
),
|
||||
"use_fps": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": False,
|
||||
"tooltip": "Enable frame interpolation to target_fps.",
|
||||
},
|
||||
),
|
||||
"target_fps": (
|
||||
"INT",
|
||||
{
|
||||
"default": 0,
|
||||
"min": 0,
|
||||
"max": 60,
|
||||
"tooltip": "Target output frame rate. Only used when use_fps is enabled and value is non-zero.",
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("video_url",)
|
||||
FUNCTION = "generate_upscaled_video"
|
||||
CATEGORY = "FAL/VideoUpscaling"
|
||||
|
||||
def generate_upscaled_video(
|
||||
self,
|
||||
video=None,
|
||||
input_video_url=None,
|
||||
upscale_factor=2.0,
|
||||
use_fps=False,
|
||||
target_fps=0,
|
||||
):
|
||||
try:
|
||||
video_url = input_video_url
|
||||
if video is not None:
|
||||
video_url = ImageUtils.upload_file(video.get_stream_source())
|
||||
if not video_url:
|
||||
return ApiHandler.handle_video_generation_error(
|
||||
"fal-ai/topaz/upscale/video",
|
||||
"No video provided. Connect a VIDEO input or set input_video_url.",
|
||||
)
|
||||
|
||||
arguments = {
|
||||
"video_url": video_url,
|
||||
"upscale_factor": upscale_factor,
|
||||
}
|
||||
if target_fps != 0 and use_fps:
|
||||
arguments["target_fps"] = target_fps
|
||||
|
||||
result = ApiHandler.submit_and_get_result(
|
||||
"fal-ai/topaz/upscale/video", arguments
|
||||
)
|
||||
return (result["video"]["url"],)
|
||||
except Exception as e:
|
||||
return ApiHandler.handle_video_generation_error(
|
||||
"fal-ai/topaz/upscale/video", e
|
||||
)
|
||||
|
||||
def create_blank_image(self):
|
||||
blank_img = Image.new('RGB', (512, 512), color='black')
|
||||
img_array = np.array(blank_img).astype(np.float32) / 255.0
|
||||
img_tensor = torch.from_numpy(img_array)[None,]
|
||||
return (img_tensor,)
|
||||
|
||||
# Node class mappings
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"Upscaler_fal": UpscalerNode,
|
||||
"Seedvr_Upscaler_fal": SeedvrUpscalerNode,
|
||||
"Seedvr_Upscale_Video_fal": SeedvrUpscaleVideoNode,
|
||||
"Bria_Video_Increase_Resolution_fal": BriaVideoIncreaseResolutionNode,
|
||||
"Topaz_Upscale_Video_fal": TopazUpscaleVideoNode,
|
||||
}
|
||||
|
||||
# Node display name mappings
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"Upscaler_fal": "Clarity Upscaler (fal)",
|
||||
}
|
||||
"Seedvr_Upscaler_fal": "Seedvr Upscaler (fal)",
|
||||
"Seedvr_Upscale_Video_fal": "Seedvr Upscale Video (fal)",
|
||||
"Bria_Video_Increase_Resolution_fal": "Bria Video Increase Resolution (fal)",
|
||||
"Topaz_Upscale_Video_fal": "Topaz Upscale Video (fal)",
|
||||
}
|
||||
|
||||
@@ -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)",
|
||||
}
|
||||
@@ -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)",
|
||||
}
|
||||
@@ -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)",
|
||||
}
|
||||
@@ -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)",
|
||||
}
|
||||
@@ -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)",
|
||||
}
|
||||
@@ -0,0 +1,33 @@
|
||||
"""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
|
||||
from .images import ImageUtils, ResultProcessor
|
||||
from .job_store import JobStore
|
||||
from .ledger import SessionLedger
|
||||
from .logger import logger
|
||||
from .media import MediaUtils
|
||||
from .pricing import PricingUtils
|
||||
from .result_cache import ResultCache
|
||||
|
||||
__all__ = [
|
||||
"ApiHandler",
|
||||
"ArchiveUtils",
|
||||
"BillingUtils",
|
||||
"FalApiError",
|
||||
"FalConfig",
|
||||
"ImageUtils",
|
||||
"JobStore",
|
||||
"MediaUtils",
|
||||
"PricingUtils",
|
||||
"ResultCache",
|
||||
"ResultProcessor",
|
||||
"SessionLedger",
|
||||
"SpendGuard",
|
||||
"extract_error_message",
|
||||
"logger",
|
||||
"raise_fal_error",
|
||||
]
|
||||
@@ -0,0 +1,487 @@
|
||||
"""fal.ai API submission helpers for ComfyUI-fal-API."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import concurrent.futures
|
||||
import time
|
||||
from typing import Any, Callable, NoReturn
|
||||
|
||||
from .config import FalConfig
|
||||
from .errors import FalApiError, extract_error_message, raise_fal_error
|
||||
from .job_store import JobStore
|
||||
from .ledger import SessionLedger
|
||||
from .logger import logger
|
||||
from .pricing import PricingUtils
|
||||
from .result_cache import ResultCache
|
||||
|
||||
_RECOVERY_POLL_INTERVAL_S = 0.5
|
||||
|
||||
_MAX_QUEUE_LOG_LINES = 10_000
|
||||
|
||||
|
||||
def _check_interruption() -> None:
|
||||
"""Raise ComfyUI's InterruptProcessingException if the user cancelled.
|
||||
|
||||
A no-op when running outside ComfyUI.
|
||||
"""
|
||||
try:
|
||||
import comfy.model_management as model_management
|
||||
except ImportError:
|
||||
return
|
||||
model_management.throw_exception_if_processing_interrupted()
|
||||
|
||||
|
||||
def _is_interruption(exc: BaseException) -> bool:
|
||||
"""Detect ComfyUI's interruption exception without importing comfy."""
|
||||
return exc.__class__.__name__ == "InterruptProcessingException"
|
||||
|
||||
|
||||
def _log_message_from_entry(entry: Any) -> str | None:
|
||||
"""Extract a printable message from a fal queue log entry."""
|
||||
if isinstance(entry, dict):
|
||||
message = entry.get("message")
|
||||
return str(message) if message else None
|
||||
text = str(entry)
|
||||
return text if text else None
|
||||
|
||||
|
||||
def _make_queue_callback(endpoint: str) -> Callable[[Any], None]:
|
||||
"""Build an on_queue_update callback that logs progress and honors cancel."""
|
||||
import fal_client
|
||||
|
||||
seen_lines: set = set()
|
||||
last_position: list[int | None] = [None]
|
||||
|
||||
def on_queue_update(status: Any) -> None:
|
||||
# Anything raised here (interruption) must propagate to the caller.
|
||||
_check_interruption()
|
||||
|
||||
if isinstance(status, fal_client.InProgress):
|
||||
for entry in status.logs or []:
|
||||
message = _log_message_from_entry(entry)
|
||||
if message and message not in seen_lines:
|
||||
if len(seen_lines) < _MAX_QUEUE_LOG_LINES:
|
||||
seen_lines.add(message)
|
||||
logger.info("[%s] %s", endpoint, message)
|
||||
elif isinstance(status, fal_client.Queued):
|
||||
position = getattr(status, "position", None)
|
||||
if position != last_position[0]:
|
||||
last_position[0] = position
|
||||
logger.info("[%s] queued (position %s)", endpoint, position)
|
||||
|
||||
return on_queue_update
|
||||
|
||||
|
||||
async def _submit_multiple_async(
|
||||
endpoint: str, arguments: dict[str, Any], variations: int
|
||||
) -> list[Any]:
|
||||
"""Submit multiple jobs concurrently and gather results (with exceptions).
|
||||
|
||||
Interruption is only observed between the submit and gather phases — the
|
||||
per-request polling here has no queue callback, so a ComfyUI Cancel takes
|
||||
effect once the in-flight variations settle (known limitation).
|
||||
"""
|
||||
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())
|
||||
|
||||
def variation_arguments(index: int) -> dict[str, Any]:
|
||||
if "seed" in arguments:
|
||||
return {**arguments, "seed": arguments.get("seed", 0) + index}
|
||||
return arguments
|
||||
|
||||
async def submit_and_get(index: int) -> Any:
|
||||
handler = await client.submit(endpoint, arguments=variation_arguments(index))
|
||||
return await handler.get()
|
||||
|
||||
# One flow per variation so a single submit failure only loses that
|
||||
# variation instead of failing the whole batch.
|
||||
return await asyncio.gather(
|
||||
*[submit_and_get(i) for i in range(variations)], return_exceptions=True
|
||||
)
|
||||
|
||||
|
||||
def _partition_results(
|
||||
endpoint: str, raw_results: list[Any]
|
||||
) -> tuple[list[Any], list[tuple]]:
|
||||
"""Split gathered results into successes and logged failures."""
|
||||
successes: list[Any] = []
|
||||
failures: list[tuple] = []
|
||||
for index, item in enumerate(raw_results):
|
||||
if isinstance(item, BaseException):
|
||||
message, status_code = extract_error_message(item)
|
||||
logger.error("[%s] variation %d failed: %s", endpoint, index, message)
|
||||
failures.append((index, message, status_code))
|
||||
else:
|
||||
successes.append(item)
|
||||
return successes, failures
|
||||
|
||||
|
||||
def _record_ledger_entry(
|
||||
endpoint: str,
|
||||
request_id: str | None,
|
||||
duration_s: float,
|
||||
est_cost_override: float | None = None,
|
||||
free: bool = False,
|
||||
) -> None:
|
||||
"""Record one fal call in the session ledger.
|
||||
|
||||
``free=True`` marks a recovery/replay that spent no new money (est_cost
|
||||
None). Best-effort bookkeeping: any pricing or ledger error is swallowed
|
||||
so it can never break generation.
|
||||
"""
|
||||
if free:
|
||||
est_cost = est_cost_override
|
||||
else:
|
||||
try:
|
||||
est_cost = PricingUtils.estimate(endpoint, 1)["total"]
|
||||
except Exception as exc:
|
||||
logger.debug("[%s] cost estimation failed: %s", endpoint, exc)
|
||||
est_cost = None
|
||||
try:
|
||||
SessionLedger().record(endpoint, request_id, duration_s, est_cost)
|
||||
except Exception as exc:
|
||||
logger.debug("[%s] ledger record failed: %s", endpoint, exc)
|
||||
|
||||
|
||||
def _spend_guard_preflight(endpoint: str) -> None:
|
||||
"""Enforce the spend budget before submitting a paid call.
|
||||
|
||||
A no-op when the billing module is absent. A FalApiError raised by
|
||||
SpendGuard (over budget) propagates to the caller.
|
||||
"""
|
||||
try:
|
||||
from .billing import SpendGuard
|
||||
except ImportError:
|
||||
return
|
||||
SpendGuard.preflight(endpoint)
|
||||
|
||||
|
||||
def _store_result_in_cache(
|
||||
endpoint: str,
|
||||
arguments: dict[str, Any],
|
||||
result: Any,
|
||||
request_id: str | None,
|
||||
) -> None:
|
||||
"""Persist a successful live result in the persistent cache.
|
||||
|
||||
Best-effort bookkeeping: only dict results are cached and any cache
|
||||
error is swallowed so it can never break generation.
|
||||
"""
|
||||
if not isinstance(result, dict):
|
||||
return
|
||||
try:
|
||||
ResultCache().put(endpoint, arguments, result, request_id)
|
||||
except Exception as exc:
|
||||
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):
|
||||
if _is_interruption(error) or not isinstance(error, Exception):
|
||||
raise error
|
||||
raise_fal_error(model_name, error)
|
||||
raise FalApiError(model_name, str(error))
|
||||
|
||||
|
||||
class ApiHandler:
|
||||
"""Utility functions for fal.ai API interactions."""
|
||||
|
||||
@staticmethod
|
||||
def submit_and_get_result(
|
||||
endpoint: str,
|
||||
arguments: dict[str, Any],
|
||||
timeout: float | None = None,
|
||||
skip_cache: bool = False,
|
||||
) -> Any:
|
||||
"""Submit a job via client.subscribe and return the final result.
|
||||
|
||||
Checks the spend budget first, then the persistent result cache: an
|
||||
identical previous call returns its stored result immediately (no
|
||||
charge, no ledger entry). Pass ``skip_cache=True`` to force a live
|
||||
call (e.g. force_rerun). Logs queue position and in-progress log
|
||||
lines, and checks for ComfyUI interruption on every queue update.
|
||||
``timeout`` is reserved for future use (fal_client 1.0 subscribe does
|
||||
not accept one).
|
||||
"""
|
||||
del timeout # Reserved; not supported by fal_client 1.0 subscribe.
|
||||
|
||||
# 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
|
||||
|
||||
_spend_guard_preflight(endpoint)
|
||||
|
||||
client = FalConfig().get_client()
|
||||
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
|
||||
|
||||
started = time.monotonic()
|
||||
try:
|
||||
result = 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)
|
||||
|
||||
_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
|
||||
def submit_only(endpoint: str, arguments: dict[str, Any]) -> str:
|
||||
"""Submit a job without waiting and return its request id.
|
||||
|
||||
Checks the spend budget first (async fan-out must respect it too).
|
||||
Does not record to the session ledger — the collect side
|
||||
(``result_from_request_id``) records the call.
|
||||
"""
|
||||
_spend_guard_preflight(endpoint)
|
||||
client = FalConfig().get_client()
|
||||
try:
|
||||
handle = client.submit(endpoint, arguments=arguments)
|
||||
except FalApiError:
|
||||
raise
|
||||
except Exception as exc:
|
||||
if _is_interruption(exc):
|
||||
raise
|
||||
raise_fal_error(endpoint, exc)
|
||||
logger.info("[%s] submitted async (request_id=%s)", endpoint, handle.request_id)
|
||||
# Best-effort bookkeeping: the persistent job inbox lets this request
|
||||
# be found and collected even after a ComfyUI restart.
|
||||
try:
|
||||
JobStore().record_submit(endpoint, handle.request_id)
|
||||
except Exception as exc:
|
||||
logger.debug("[%s] job store record_submit failed: %s", endpoint, exc)
|
||||
return handle.request_id
|
||||
|
||||
@staticmethod
|
||||
def result_from_request_id(
|
||||
endpoint: str, request_id: str, record_cost: bool = True
|
||||
) -> dict[str, Any]:
|
||||
"""Wait for and fetch the result of a previously submitted request.
|
||||
|
||||
Reconstructs a queue handle from the request id, polls until the
|
||||
request completes (honoring ComfyUI interruption), and returns the
|
||||
result payload. A request that already completed returns immediately
|
||||
without incurring new charges — the result-recovery path.
|
||||
"""
|
||||
label = f"{endpoint}#{request_id}"
|
||||
client = FalConfig().get_client()
|
||||
started = time.monotonic()
|
||||
try:
|
||||
handle = client.get_handle(endpoint, request_id)
|
||||
for _status in handle.iter_events(
|
||||
with_logs=False, interval=_RECOVERY_POLL_INTERVAL_S
|
||||
):
|
||||
_check_interruption()
|
||||
result = handle.get()
|
||||
except FalApiError:
|
||||
raise
|
||||
except Exception as exc:
|
||||
if _is_interruption(exc):
|
||||
raise
|
||||
raise_fal_error(label, exc)
|
||||
finally:
|
||||
duration_s = time.monotonic() - started
|
||||
logger.info(
|
||||
"[%s] result recovery finished in %.1fs (request_id=%s)",
|
||||
endpoint,
|
||||
duration_s,
|
||||
request_id,
|
||||
)
|
||||
# record_cost=False marks a pure recovery of an old request:
|
||||
# log the fetch for traceability but count no new spend.
|
||||
if record_cost:
|
||||
_record_ledger_entry(endpoint, request_id, duration_s)
|
||||
else:
|
||||
_record_ledger_entry(endpoint, request_id, duration_s, est_cost_override=None, free=True)
|
||||
# Best-effort bookkeeping: mark the job collected in the persistent
|
||||
# inbox (inserting it if it was submitted in another session).
|
||||
try:
|
||||
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
|
||||
def submit_multiple_and_get_results(
|
||||
endpoint: str, arguments: dict[str, Any], variations: int
|
||||
) -> list[Any]:
|
||||
"""Submit multiple variations concurrently and return successful results.
|
||||
|
||||
Failed variations are logged; raises FalApiError only if ALL fail.
|
||||
"""
|
||||
try:
|
||||
# Run the async code in a dedicated thread to avoid event loop
|
||||
# conflicts with ComfyUI's own loop.
|
||||
with concurrent.futures.ThreadPoolExecutor(max_workers=1) as executor:
|
||||
future = executor.submit(
|
||||
asyncio.run,
|
||||
_submit_multiple_async(endpoint, arguments, variations),
|
||||
)
|
||||
raw_results = future.result()
|
||||
except FalApiError:
|
||||
raise
|
||||
except Exception as exc:
|
||||
if _is_interruption(exc):
|
||||
raise
|
||||
raise_fal_error(endpoint, exc)
|
||||
|
||||
successes, failures = _partition_results(endpoint, raw_results)
|
||||
if not successes:
|
||||
first_message = failures[0][1] if failures else "no results returned"
|
||||
first_status = failures[0][2] if failures else None
|
||||
raise FalApiError(
|
||||
endpoint,
|
||||
f"All {variations} variations failed: {first_message}",
|
||||
first_status,
|
||||
)
|
||||
return successes
|
||||
|
||||
@staticmethod
|
||||
def handle_video_generation_error(
|
||||
model_name: str, error: Exception | str
|
||||
) -> NoReturn:
|
||||
"""Raise a normalized FalApiError for a video generation failure."""
|
||||
_raise_generation_error(model_name, error)
|
||||
|
||||
@staticmethod
|
||||
def handle_image_generation_error(
|
||||
model_name: str, error: Exception | str
|
||||
) -> NoReturn:
|
||||
"""Raise a normalized FalApiError for an image generation failure."""
|
||||
_raise_generation_error(model_name, error)
|
||||
|
||||
@staticmethod
|
||||
def handle_text_generation_error(
|
||||
model_name: str, error: Exception | str
|
||||
) -> NoReturn:
|
||||
"""Raise a normalized FalApiError for a text generation failure."""
|
||||
_raise_generation_error(model_name, error)
|
||||
@@ -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)
|
||||
@@ -0,0 +1,273 @@
|
||||
"""fal.ai platform billing: account balance, usage reconciliation, spend guard.
|
||||
|
||||
Uses the fal Platform APIs (base ``https://api.fal.ai/v1``):
|
||||
|
||||
- ``GET /account/billing?expand=credits`` — current credit balance.
|
||||
- ``GET /models/requests/by-endpoint`` — per-request records (request_id,
|
||||
endpoint_id; fal does not publish a per-request billed amount).
|
||||
- ``GET /models/usage?expand=summary`` — aggregated billed cost per endpoint.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
import time
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
import requests
|
||||
|
||||
from .config import FalConfig
|
||||
from .errors import FalApiError
|
||||
from .ledger import SessionLedger
|
||||
from .logger import logger
|
||||
|
||||
_API_BASE = "https://api.fal.ai/v1"
|
||||
_REQUEST_TIMEOUT = (5, 15)
|
||||
_BALANCE_CACHE_TTL_S = 60.0
|
||||
_MAX_ENDPOINT_FILTERS = 50
|
||||
_MAX_REQUEST_LIMIT = 100
|
||||
_BILLING_DASHBOARD_URL = "https://fal.ai/dashboard/billing"
|
||||
|
||||
_balance_lock = threading.Lock()
|
||||
# [cached balance (float|None), fetched_at unix time]; fetched_at 0 = no fetch yet.
|
||||
_balance_cache: list[Any] = [None, 0.0]
|
||||
|
||||
_warn_lock = threading.Lock()
|
||||
_warned_once: frozenset[str] = frozenset()
|
||||
|
||||
|
||||
def _warn_once(topic: str, message: str) -> None:
|
||||
"""Log ``message`` as WARNING the first time per topic, DEBUG afterwards."""
|
||||
global _warned_once
|
||||
with _warn_lock:
|
||||
first_time = topic not in _warned_once
|
||||
_warned_once = _warned_once | {topic}
|
||||
if first_time:
|
||||
logger.warning(message)
|
||||
else:
|
||||
logger.debug(message)
|
||||
|
||||
|
||||
def _get_json(path: str, params: dict[str, Any], topic: str) -> Any | None:
|
||||
"""GET a Platform API path; return parsed JSON or None on any failure."""
|
||||
key = FalConfig().get_key()
|
||||
if not key:
|
||||
_warn_once(topic, f"Cannot call fal Platform API {path}: FAL_KEY is not configured")
|
||||
return None
|
||||
try:
|
||||
response = requests.get(
|
||||
f"{_API_BASE}{path}",
|
||||
params=params,
|
||||
headers={"Authorization": f"Key {key}"},
|
||||
timeout=_REQUEST_TIMEOUT,
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
except Exception as exc:
|
||||
_warn_once(topic, f"fal Platform API {path} unavailable: {exc}")
|
||||
return None
|
||||
|
||||
|
||||
def _fetch_balance() -> float | None:
|
||||
"""Fetch the current credit balance in USD, or None on any failure."""
|
||||
payload = _get_json("/account/billing", {"expand": "credits"}, topic="balance")
|
||||
if not isinstance(payload, dict):
|
||||
return None
|
||||
credits = payload.get("credits")
|
||||
balance = credits.get("current_balance") if isinstance(credits, dict) else None
|
||||
if isinstance(balance, (int, float)):
|
||||
return float(balance)
|
||||
_warn_once("balance", "fal balance response had no credits.current_balance field")
|
||||
return None
|
||||
|
||||
|
||||
def _ledger_endpoints(entries: list[dict[str, Any]]) -> list[str]:
|
||||
"""Unique endpoint ids from ledger entries, capped at the API filter limit."""
|
||||
seen: list[str] = []
|
||||
for entry in entries:
|
||||
endpoint = entry.get("endpoint_id")
|
||||
if isinstance(endpoint, str) and endpoint and endpoint not in seen:
|
||||
seen.append(endpoint)
|
||||
return seen[:_MAX_ENDPOINT_FILTERS]
|
||||
|
||||
|
||||
def _earliest_timestamp_iso(entries: list[dict[str, Any]]) -> str | None:
|
||||
"""ISO8601 UTC timestamp of the earliest ledger entry, or None."""
|
||||
stamps = [e["timestamp"] for e in entries if isinstance(e.get("timestamp"), (int, float))]
|
||||
if not stamps:
|
||||
return None
|
||||
earliest = datetime.fromtimestamp(min(stamps), tz=timezone.utc)
|
||||
return earliest.strftime("%Y-%m-%dT%H:%M:%SZ")
|
||||
|
||||
|
||||
def _billed_total_since(entries: list[dict[str, Any]]) -> float | None:
|
||||
"""Aggregated billed cost for the ledger's endpoints since the session start.
|
||||
|
||||
Best effort via ``GET /models/usage?expand=summary``; the window covers the
|
||||
whole workspace, so concurrent non-session calls may be included.
|
||||
"""
|
||||
endpoints = _ledger_endpoints(entries)
|
||||
start = _earliest_timestamp_iso(entries)
|
||||
if not endpoints or start is None:
|
||||
return None
|
||||
params = {
|
||||
"expand": "summary",
|
||||
"start": start,
|
||||
"endpoint_id": ",".join(endpoints),
|
||||
"bound_to_timeframe": "false",
|
||||
}
|
||||
payload = _get_json("/models/usage", params, topic="usage")
|
||||
if not isinstance(payload, dict) or not isinstance(payload.get("summary"), list):
|
||||
return None
|
||||
costs = [
|
||||
item.get("cost")
|
||||
for item in payload["summary"]
|
||||
if isinstance(item, dict) and isinstance(item.get("cost"), (int, float))
|
||||
]
|
||||
return float(sum(costs)) if costs else None
|
||||
|
||||
|
||||
class BillingUtils:
|
||||
"""Read-only access to fal account balance and billed usage. Never raises."""
|
||||
|
||||
@staticmethod
|
||||
def get_balance(force: bool = False) -> float | None:
|
||||
"""Current account credit balance in USD, or None on any failure.
|
||||
|
||||
Results (including failures) are cached for 60 seconds so callers such
|
||||
as SpendGuard.preflight do not hammer the API; ``force=True`` bypasses.
|
||||
"""
|
||||
global _balance_cache
|
||||
# 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
|
||||
|
||||
@staticmethod
|
||||
def get_recent_usage(limit: int = 50) -> list[dict[str, Any]] | None:
|
||||
"""Recent per-request records for this session's endpoints, or None.
|
||||
|
||||
Backed by ``GET /models/requests/by-endpoint`` filtered to the
|
||||
endpoints recorded in the SessionLedger. fal's Platform APIs do not
|
||||
expose a per-request billed amount (usage is aggregated), so ``amount``
|
||||
is always None. Returns None when the ledger is empty or the API is
|
||||
unavailable.
|
||||
"""
|
||||
try:
|
||||
endpoints = _ledger_endpoints(SessionLedger().entries())
|
||||
if not endpoints:
|
||||
return None
|
||||
params = {
|
||||
"endpoint_id": ",".join(endpoints),
|
||||
"limit": max(1, min(int(limit), _MAX_REQUEST_LIMIT)),
|
||||
}
|
||||
payload = _get_json("/models/requests/by-endpoint", params, topic="requests")
|
||||
if not isinstance(payload, dict) or not isinstance(payload.get("items"), list):
|
||||
return None
|
||||
return [
|
||||
{
|
||||
"request_id": item.get("request_id"),
|
||||
"endpoint": item.get("endpoint_id"),
|
||||
"amount": None,
|
||||
}
|
||||
for item in payload["items"]
|
||||
if isinstance(item, dict)
|
||||
]
|
||||
except Exception as exc:
|
||||
logger.debug("get_recent_usage failed: %s", exc)
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def reconcile_ledger() -> dict[str, Any]:
|
||||
"""Best-effort reconciliation of the SessionLedger against fal's records.
|
||||
|
||||
Returns ``{"matched": n, "billed_total": float|None, "estimated_total": float}``
|
||||
where ``matched`` counts ledger request_ids confirmed by the requests
|
||||
API and ``billed_total`` is the aggregated billed cost (None when the
|
||||
usage API is unavailable). Never raises.
|
||||
"""
|
||||
estimated_total = SessionLedger().total_cost()
|
||||
result: dict[str, Any] = {
|
||||
"matched": 0,
|
||||
"billed_total": None,
|
||||
"estimated_total": estimated_total,
|
||||
}
|
||||
try:
|
||||
entries = SessionLedger().entries()
|
||||
if not entries:
|
||||
return result
|
||||
ledger_ids = {e["request_id"] for e in entries if e.get("request_id")}
|
||||
usage = BillingUtils.get_recent_usage(limit=_MAX_REQUEST_LIMIT) or []
|
||||
matched = sum(1 for record in usage if record.get("request_id") in ledger_ids)
|
||||
return {
|
||||
**result,
|
||||
"matched": matched,
|
||||
"billed_total": _billed_total_since(entries),
|
||||
}
|
||||
except Exception as exc:
|
||||
logger.debug("reconcile_ledger failed: %s", exc)
|
||||
return result
|
||||
|
||||
|
||||
def _setting_float(name: str) -> float:
|
||||
"""Read a [spend_guard] float setting; 0.0 when unset or unparseable."""
|
||||
try:
|
||||
value = FalConfig().get_setting("spend_guard", name, 0)
|
||||
if isinstance(value, bool):
|
||||
return 0.0
|
||||
return float(value)
|
||||
except Exception as exc:
|
||||
logger.debug("Invalid [spend_guard] %s value: %s", name, exc)
|
||||
return 0.0
|
||||
|
||||
|
||||
class SpendGuard:
|
||||
"""Pre-call spend checks configured via config.ini [spend_guard]."""
|
||||
|
||||
@staticmethod
|
||||
def settings() -> dict[str, float]:
|
||||
"""Active spend-guard settings (0.0 means the check is disabled)."""
|
||||
return {
|
||||
"session_budget_usd": _setting_float("session_budget_usd"),
|
||||
"min_balance_usd": _setting_float("min_balance_usd"),
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def preflight(endpoint: str) -> None:
|
||||
"""Raise FalApiError if a configured spend limit blocks this call.
|
||||
|
||||
Frozen contract: called by ApiHandler before every fal request. Checks
|
||||
are skipped when their setting is 0/unset; an unavailable balance API
|
||||
never blocks. Raises nothing except FalApiError.
|
||||
"""
|
||||
budget = _setting_float("session_budget_usd")
|
||||
if budget > 0:
|
||||
spent = SessionLedger().total_cost()
|
||||
if spent >= budget:
|
||||
raise FalApiError(
|
||||
"spend-guard",
|
||||
f"Session budget ${budget:.2f} reached (spent ~${spent:.2f}). "
|
||||
"Raise [spend_guard] session_budget_usd in config.ini or "
|
||||
"reset the session ledger.",
|
||||
)
|
||||
|
||||
floor = _setting_float("min_balance_usd")
|
||||
if floor > 0:
|
||||
balance = BillingUtils.get_balance()
|
||||
if balance is None:
|
||||
logger.debug(
|
||||
"Spend guard: balance unavailable; allowing request to %s", endpoint
|
||||
)
|
||||
elif balance < floor:
|
||||
raise FalApiError(
|
||||
"spend-guard",
|
||||
f"fal balance ${balance:.2f} is below your ${floor:.2f} floor "
|
||||
f"— top up at {_BILLING_DASHBOARD_URL}",
|
||||
)
|
||||
@@ -0,0 +1,114 @@
|
||||
"""fal.ai API key/config resolution for the ComfyUI-fal-API node pack."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import configparser
|
||||
import os
|
||||
import threading
|
||||
from typing import Any
|
||||
|
||||
from .errors import FalApiError
|
||||
from .logger import logger
|
||||
|
||||
_PLACEHOLDER_KEY = "<your_fal_api_key_here>"
|
||||
_MISSING_KEY_MESSAGE = (
|
||||
"FAL_KEY is not configured. Set the FAL_KEY environment variable or add it "
|
||||
"to config.ini under the [API] section. Get your API key from "
|
||||
"https://fal.ai/dashboard/keys"
|
||||
)
|
||||
|
||||
|
||||
def _config_path() -> str:
|
||||
"""Return the path to config.ini at the repo root (one dir above nodes/)."""
|
||||
utils_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
nodes_dir = os.path.dirname(utils_dir)
|
||||
repo_root = os.path.dirname(nodes_dir)
|
||||
return os.path.join(repo_root, "config.ini")
|
||||
|
||||
|
||||
def _read_config() -> configparser.ConfigParser:
|
||||
"""Read config.ini; returns an empty parser if the file is absent."""
|
||||
parser = configparser.ConfigParser()
|
||||
try:
|
||||
parser.read(_config_path())
|
||||
except configparser.Error as exc:
|
||||
logger.warning("Failed to parse config.ini: %s", exc)
|
||||
return parser
|
||||
|
||||
|
||||
def _resolve_key(parser: configparser.ConfigParser) -> str | None:
|
||||
"""Resolve the FAL key: environment first, then config.ini [API] FAL_KEY."""
|
||||
env_key = os.environ.get("FAL_KEY")
|
||||
if env_key:
|
||||
logger.info("Using FAL_KEY from environment")
|
||||
return env_key
|
||||
|
||||
config_key = parser.get("API", "FAL_KEY", fallback=None)
|
||||
if config_key:
|
||||
logger.info("Using FAL_KEY from config.ini")
|
||||
return config_key
|
||||
return None
|
||||
|
||||
|
||||
def _is_valid_key(key: str | None) -> bool:
|
||||
return bool(key) and key != _PLACEHOLDER_KEY
|
||||
|
||||
|
||||
class FalConfig:
|
||||
"""Singleton holding fal.ai configuration and a cached client."""
|
||||
|
||||
_instance: FalConfig | None = None
|
||||
_lock = threading.Lock()
|
||||
|
||||
def __new__(cls) -> FalConfig:
|
||||
if cls._instance is None:
|
||||
with cls._lock:
|
||||
if cls._instance is None:
|
||||
instance = super().__new__(cls)
|
||||
instance._initialize()
|
||||
cls._instance = instance
|
||||
return cls._instance
|
||||
|
||||
def _initialize(self) -> None:
|
||||
"""Resolve the API key once; never raises at import/construction time."""
|
||||
self._parser = _read_config()
|
||||
self._key: str | None = _resolve_key(self._parser)
|
||||
self._client: Any | None = None
|
||||
|
||||
if not _is_valid_key(self._key):
|
||||
logger.warning(_MISSING_KEY_MESSAGE)
|
||||
|
||||
def get_client(self) -> Any:
|
||||
"""Get or create the cached fal_client SyncClient.
|
||||
|
||||
Raises FalApiError if no valid key is configured.
|
||||
"""
|
||||
if self._client is None:
|
||||
if not _is_valid_key(self._key):
|
||||
raise FalApiError("config", _MISSING_KEY_MESSAGE)
|
||||
from fal_client.client import SyncClient
|
||||
|
||||
self._client = SyncClient(key=self._key)
|
||||
return self._client
|
||||
|
||||
def get_key(self) -> str | None:
|
||||
"""Return the resolved FAL API key (may be None or a placeholder)."""
|
||||
return self._key
|
||||
|
||||
def get_setting(self, section: str, name: str, default: Any = None) -> Any:
|
||||
"""Read an arbitrary config.ini setting, with bool parsing.
|
||||
|
||||
Returns ``default`` if the section or option is absent. Values equal to
|
||||
"true"/"false" (case-insensitive) are returned as booleans.
|
||||
"""
|
||||
try:
|
||||
value = self._parser.get(section, name)
|
||||
except (configparser.NoSectionError, configparser.NoOptionError):
|
||||
return default
|
||||
|
||||
lowered = value.strip().lower()
|
||||
if lowered == "true":
|
||||
return True
|
||||
if lowered == "false":
|
||||
return False
|
||||
return value
|
||||
@@ -0,0 +1,89 @@
|
||||
"""Error types and helpers for normalizing fal.ai API failures."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, NoReturn
|
||||
|
||||
|
||||
class FalApiError(Exception):
|
||||
"""Raised when a fal.ai API call (or related processing) fails."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
message: str,
|
||||
status_code: int | None = None,
|
||||
) -> None:
|
||||
self.model_name = model_name
|
||||
self.message = message
|
||||
self.status_code = status_code
|
||||
formatted = f"[{model_name}] {message}"
|
||||
if status_code is not None:
|
||||
formatted = f"{formatted} (HTTP {status_code})"
|
||||
super().__init__(formatted)
|
||||
|
||||
|
||||
def _flatten_validation_detail(detail: list[Any]) -> str:
|
||||
"""Flatten a FastAPI validation-error list into a readable string."""
|
||||
parts: list[str] = []
|
||||
for item in detail:
|
||||
if isinstance(item, dict):
|
||||
loc = ".".join(str(part) for part in (item.get("loc") or []))
|
||||
msg = str(item.get("msg", item))
|
||||
parts.append(f"{loc}: {msg}" if loc else msg)
|
||||
else:
|
||||
parts.append(str(item))
|
||||
return "; ".join(parts)
|
||||
|
||||
|
||||
def _detail_to_message(detail: Any) -> str:
|
||||
"""Convert a response 'detail' payload into a message string."""
|
||||
if isinstance(detail, str):
|
||||
return detail
|
||||
if isinstance(detail, list):
|
||||
return _flatten_validation_detail(detail)
|
||||
return str(detail)
|
||||
|
||||
|
||||
def _message_from_response(response: Any) -> str | None:
|
||||
"""Extract a human-readable message from an httpx-like response."""
|
||||
try:
|
||||
payload = response.json()
|
||||
except Exception:
|
||||
payload = None
|
||||
|
||||
if isinstance(payload, dict) and "detail" in payload:
|
||||
return _detail_to_message(payload["detail"])
|
||||
if payload is not None:
|
||||
return str(payload)
|
||||
|
||||
text = getattr(response, "text", None)
|
||||
if isinstance(text, str) and text.strip():
|
||||
return text.strip()
|
||||
return None
|
||||
|
||||
|
||||
def extract_error_message(exc: BaseException) -> tuple[str, int | None]:
|
||||
"""Extract a readable message and HTTP status code from an exception.
|
||||
|
||||
Duck-types fal_client.FalClientHTTPError (``.status_code`` plus an
|
||||
httpx ``.response``) so this works without importing fal_client.
|
||||
"""
|
||||
raw_status = getattr(exc, "status_code", None)
|
||||
status_code = raw_status if isinstance(raw_status, int) else None
|
||||
|
||||
response = getattr(exc, "response", None)
|
||||
if response is not None:
|
||||
message = _message_from_response(response)
|
||||
if message:
|
||||
return message, status_code
|
||||
|
||||
return str(exc) or exc.__class__.__name__, status_code
|
||||
|
||||
|
||||
def raise_fal_error(model_name: str, exc: Exception) -> NoReturn:
|
||||
"""Normalize any exception into a FalApiError and raise it."""
|
||||
if isinstance(exc, FalApiError):
|
||||
raise exc
|
||||
message, status_code = extract_error_message(exc)
|
||||
raise FalApiError(model_name, message, status_code) from exc
|
||||
@@ -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
|
||||
@@ -0,0 +1,248 @@
|
||||
"""Image tensor helpers and API result processing for ComfyUI-fal-API.
|
||||
|
||||
ComfyUI IMAGE convention: float32 tensors in [0, 1] with shape (B, H, W, C).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import io
|
||||
import os
|
||||
import tempfile
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import requests
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
from .config import FalConfig
|
||||
from .errors import FalApiError, raise_fal_error
|
||||
from .logger import logger
|
||||
|
||||
_DOWNLOAD_TIMEOUT = (10, 180)
|
||||
_MAX_PARALLEL_TRANSFERS = 8
|
||||
_HASH_CHUNK_SIZE = 1 << 20 # 1 MiB
|
||||
|
||||
|
||||
def _hash_file(path: str) -> str | None:
|
||||
"""Return the sha256 hex digest of a file's bytes, or None on any error."""
|
||||
try:
|
||||
digest = hashlib.sha256()
|
||||
with open(path, "rb") as handle:
|
||||
for chunk in iter(lambda: handle.read(_HASH_CHUNK_SIZE), b""):
|
||||
digest.update(chunk)
|
||||
return digest.hexdigest()
|
||||
except Exception as exc:
|
||||
logger.debug("failed to hash %s for upload cache: %s", path, exc)
|
||||
return None
|
||||
|
||||
|
||||
def _upload_cache_lookup(file_path: Any) -> tuple[str | None, str | None]:
|
||||
"""Hash a local file and consult the persistent upload cache.
|
||||
|
||||
Returns (content_hash, cached_url); both None when the file cannot be
|
||||
hashed or caching is unavailable. Never raises.
|
||||
"""
|
||||
try:
|
||||
if not isinstance(file_path, (str, os.PathLike)):
|
||||
return None, None
|
||||
content_hash = _hash_file(os.fspath(file_path))
|
||||
if content_hash is None:
|
||||
return None, None
|
||||
|
||||
from .result_cache import ResultCache
|
||||
|
||||
cached_url = ResultCache().get_upload(content_hash)
|
||||
if cached_url:
|
||||
logger.debug("upload cache hit (sha256=%s...)", content_hash[:12])
|
||||
return content_hash, cached_url
|
||||
except Exception as exc:
|
||||
logger.debug("upload cache lookup failed: %s", exc)
|
||||
return None, None
|
||||
|
||||
|
||||
def _upload_cache_store(content_hash: str | None, url: str) -> None:
|
||||
"""Remember a completed upload in the persistent cache. Never raises."""
|
||||
if content_hash is None:
|
||||
return
|
||||
try:
|
||||
from .result_cache import ResultCache
|
||||
|
||||
ResultCache().put_upload(content_hash, url)
|
||||
except Exception as exc:
|
||||
logger.debug("upload cache store failed: %s", exc)
|
||||
|
||||
|
||||
def _safe_unlink(path: str) -> None:
|
||||
"""Delete a temp file, ignoring errors."""
|
||||
try:
|
||||
os.unlink(path)
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
|
||||
def _download_image_array(url: str) -> np.ndarray:
|
||||
"""Download an image URL and return a float32 (H, W, 3) array in [0, 1]."""
|
||||
response = requests.get(url, timeout=_DOWNLOAD_TIMEOUT)
|
||||
response.raise_for_status()
|
||||
img = Image.open(io.BytesIO(response.content)).convert("RGB")
|
||||
return np.array(img).astype(np.float32) / 255.0
|
||||
|
||||
|
||||
def _download_image_arrays(urls: list[str]) -> list[np.ndarray]:
|
||||
"""Download image URLs (in parallel when multiple), preserving order."""
|
||||
if len(urls) == 1:
|
||||
return [_download_image_array(urls[0])]
|
||||
max_workers = min(len(urls), _MAX_PARALLEL_TRANSFERS)
|
||||
with ThreadPoolExecutor(max_workers=max_workers) as executor:
|
||||
return list(executor.map(_download_image_array, urls))
|
||||
|
||||
|
||||
def _split_image_batch(images: Any) -> list[Any]:
|
||||
"""Split an IMAGE input into a list of single images, preserving order."""
|
||||
if isinstance(images, torch.Tensor):
|
||||
if images.ndim == 4 and images.shape[0] > 1:
|
||||
return [images[i : i + 1] for i in range(images.shape[0])]
|
||||
return [images]
|
||||
if isinstance(images, (list, tuple)):
|
||||
return list(images)
|
||||
return [images]
|
||||
|
||||
|
||||
class ImageUtils:
|
||||
"""Utility functions for image processing and uploads."""
|
||||
|
||||
@staticmethod
|
||||
def tensor_to_pil(image: Any) -> Image.Image:
|
||||
"""Convert an image tensor (or array-like) to a PIL Image."""
|
||||
try:
|
||||
if isinstance(image, torch.Tensor):
|
||||
image_np = image.detach().cpu().numpy()
|
||||
else:
|
||||
image_np = np.array(image)
|
||||
|
||||
if image_np.ndim == 4:
|
||||
image_np = image_np[0] # Drop batch dimension
|
||||
if image_np.ndim == 2:
|
||||
image_np = np.stack([image_np] * 3, axis=-1) # Grayscale -> RGB
|
||||
elif (
|
||||
image_np.ndim == 3
|
||||
and image_np.shape[0] == 3
|
||||
and image_np.shape[2] not in (1, 3, 4)
|
||||
):
|
||||
image_np = np.transpose(image_np, (1, 2, 0)) # (C, H, W) -> (H, W, C)
|
||||
|
||||
if image_np.dtype in (np.float32, np.float64):
|
||||
image_np = np.clip(image_np * 255.0, 0, 255).astype(np.uint8)
|
||||
|
||||
return Image.fromarray(image_np)
|
||||
except Exception as exc:
|
||||
logger.error("Failed to convert tensor to PIL image: %s", exc)
|
||||
raise FalApiError(
|
||||
"image-utils", f"Failed to convert tensor to image: {exc}"
|
||||
) from exc
|
||||
|
||||
@staticmethod
|
||||
def upload_image(image: Any) -> str:
|
||||
"""Upload an image tensor to fal.ai and return its URL."""
|
||||
pil_image = ImageUtils.tensor_to_pil(image)
|
||||
temp_path: str | None = None
|
||||
try:
|
||||
with tempfile.NamedTemporaryFile(suffix=".png", delete=False) as temp_file:
|
||||
temp_path = temp_file.name
|
||||
pil_image.save(temp_file, format="PNG")
|
||||
return ImageUtils.upload_file(temp_path)
|
||||
finally:
|
||||
if temp_path is not None:
|
||||
_safe_unlink(temp_path)
|
||||
|
||||
@staticmethod
|
||||
def upload_file(file_path: Any) -> str:
|
||||
"""Upload a local file to fal.ai and return its URL.
|
||||
|
||||
Identical file contents reuse the previously uploaded URL via the
|
||||
persistent upload cache (keyed by sha256), skipping the transfer.
|
||||
"""
|
||||
content_hash, cached_url = _upload_cache_lookup(file_path)
|
||||
if cached_url:
|
||||
return cached_url
|
||||
try:
|
||||
client = FalConfig().get_client()
|
||||
url = client.upload_file(file_path)
|
||||
except FalApiError:
|
||||
raise
|
||||
except Exception as exc:
|
||||
logger.error("Failed to upload file %s: %s", file_path, exc)
|
||||
raise_fal_error("file-upload", exc)
|
||||
_upload_cache_store(content_hash, url)
|
||||
return url
|
||||
|
||||
@staticmethod
|
||||
def mask_to_image(mask: torch.Tensor) -> torch.Tensor:
|
||||
"""Convert a MASK tensor to an IMAGE tensor (B, H, W, 3)."""
|
||||
return (
|
||||
mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1]))
|
||||
.movedim(1, -1)
|
||||
.expand(-1, -1, -1, 3)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def prepare_images(images: Any) -> list[str]:
|
||||
"""Upload image input(s) to fal.ai in parallel, preserving order."""
|
||||
if images is None:
|
||||
return []
|
||||
singles = _split_image_batch(images)
|
||||
if not singles:
|
||||
return []
|
||||
if len(singles) == 1:
|
||||
return [ImageUtils.upload_image(singles[0])]
|
||||
max_workers = min(len(singles), _MAX_PARALLEL_TRANSFERS)
|
||||
with ThreadPoolExecutor(max_workers=max_workers) as executor:
|
||||
return list(executor.map(ImageUtils.upload_image, singles))
|
||||
|
||||
|
||||
class ResultProcessor:
|
||||
"""Utility functions for turning API results into ComfyUI tensors."""
|
||||
|
||||
@staticmethod
|
||||
def process_image_result(result: dict[str, Any]) -> tuple:
|
||||
"""Process a multi-image result ({"images": [{"url": ...}, ...]})."""
|
||||
try:
|
||||
urls = [img_info["url"] for img_info in result["images"]]
|
||||
if not urls:
|
||||
raise ValueError("API result contained no images")
|
||||
arrays = _download_image_arrays(urls)
|
||||
stacked = np.stack(arrays, axis=0)
|
||||
return (torch.from_numpy(stacked),)
|
||||
except FalApiError:
|
||||
raise
|
||||
except Exception as exc:
|
||||
logger.error("Failed to process image result: %s", exc)
|
||||
raise FalApiError(
|
||||
"image-result", f"Failed to process image result: {exc}"
|
||||
) from exc
|
||||
|
||||
@staticmethod
|
||||
def process_single_image_result(result: dict[str, Any]) -> tuple:
|
||||
"""Process a single-image result ({"image": {"url": ...}})."""
|
||||
try:
|
||||
img_array = _download_image_array(result["image"]["url"])
|
||||
stacked = np.stack([img_array], axis=0)
|
||||
return (torch.from_numpy(stacked),)
|
||||
except FalApiError:
|
||||
raise
|
||||
except Exception as exc:
|
||||
logger.error("Failed to process single image result: %s", exc)
|
||||
raise FalApiError(
|
||||
"image-result", f"Failed to process single image result: {exc}"
|
||||
) from exc
|
||||
|
||||
@staticmethod
|
||||
def create_blank_image() -> tuple:
|
||||
"""Create a blank black 512x512 IMAGE tensor (kept for compatibility)."""
|
||||
blank_img = Image.new("RGB", (512, 512), color="black")
|
||||
img_array = np.array(blank_img).astype(np.float32) / 255.0
|
||||
img_tensor = torch.from_numpy(img_array)[None,]
|
||||
return (img_tensor,)
|
||||
@@ -0,0 +1,250 @@
|
||||
"""Persistent inbox of async fal jobs so they survive ComfyUI restarts.
|
||||
|
||||
Every job queued through Fal Submit is recorded in a ``jobs`` table inside
|
||||
the same sqlite database as the result cache ("submit tonight, collect
|
||||
tomorrow"). When a result is later fetched by request id — even in a fresh
|
||||
session — the job is marked collected.
|
||||
|
||||
Every public method is best-effort: any sqlite failure degrades to a no-op
|
||||
or an empty result — bookkeeping must never break generation.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import sqlite3
|
||||
import threading
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
from .logger import logger
|
||||
from .result_cache import _default_db_path
|
||||
|
||||
_PRUNE_AFTER_DAYS = 30
|
||||
_SECONDS_PER_DAY = 86400.0
|
||||
|
||||
_STATUS_SUBMITTED = "submitted"
|
||||
_STATUS_COLLECTED = "collected"
|
||||
|
||||
_COLUMNS = ("request_id", "endpoint", "status", "submitted_at", "collected_at", "note")
|
||||
|
||||
_SCHEMA = """
|
||||
CREATE TABLE IF NOT EXISTS jobs (
|
||||
request_id TEXT PRIMARY KEY,
|
||||
endpoint TEXT,
|
||||
status TEXT,
|
||||
submitted_at REAL,
|
||||
collected_at REAL,
|
||||
note TEXT
|
||||
)
|
||||
"""
|
||||
|
||||
|
||||
def _humanize_age(seconds: float) -> str:
|
||||
"""Compact age like '45s', '12m', '2h' or '3d'. Clamped at zero."""
|
||||
seconds = max(0.0, seconds)
|
||||
if seconds < 60:
|
||||
return f"{int(seconds)}s"
|
||||
if seconds < 3600:
|
||||
return f"{int(seconds // 60)}m"
|
||||
if seconds < _SECONDS_PER_DAY:
|
||||
return f"{int(seconds // 3600)}h"
|
||||
return f"{int(seconds // _SECONDS_PER_DAY)}d"
|
||||
|
||||
|
||||
def _format_entry(entry: dict[str, Any], now: float) -> str:
|
||||
"""One report line: ' ⏳ 2h ago fal-ai/kling-video/v3 req=abc123'."""
|
||||
collected = entry.get("status") == _STATUS_COLLECTED
|
||||
icon = "✅" if collected else "⏳"
|
||||
reference = entry.get("collected_at") if collected else entry.get("submitted_at")
|
||||
if not isinstance(reference, (int, float)):
|
||||
reference = entry.get("submitted_at")
|
||||
age = (
|
||||
f"{_humanize_age(now - float(reference))} ago"
|
||||
if isinstance(reference, (int, float))
|
||||
else "age unknown"
|
||||
)
|
||||
endpoint = entry.get("endpoint") or "(unknown endpoint)"
|
||||
return f" {icon} {age} {endpoint} req={entry.get('request_id') or '-'}"
|
||||
|
||||
|
||||
class JobStore:
|
||||
"""Thread-safe singleton over the persistent async-job inbox."""
|
||||
|
||||
_instance: JobStore | None = None
|
||||
_instance_lock = threading.Lock()
|
||||
|
||||
def __new__(cls) -> JobStore:
|
||||
if cls._instance is None:
|
||||
with cls._instance_lock:
|
||||
if cls._instance is None:
|
||||
instance = super().__new__(cls)
|
||||
instance._initialize()
|
||||
cls._instance = instance
|
||||
return cls._instance
|
||||
|
||||
def _initialize(self) -> None:
|
||||
self._lock = threading.Lock()
|
||||
self._conn: sqlite3.Connection | None = None
|
||||
self._connect_failed = False
|
||||
self._db_path = _default_db_path()
|
||||
|
||||
# -- connection -----------------------------------------------------------
|
||||
|
||||
def _connection(self) -> sqlite3.Connection | None:
|
||||
"""Open (once) and return the sqlite connection; None if unavailable.
|
||||
|
||||
Must be called with ``self._lock`` held. A corrupted or unwritable
|
||||
database disables the job store for the session instead of raising.
|
||||
"""
|
||||
if self._conn is not None:
|
||||
return self._conn
|
||||
if self._connect_failed:
|
||||
return None
|
||||
conn: sqlite3.Connection | None = None
|
||||
try:
|
||||
os.makedirs(os.path.dirname(self._db_path), exist_ok=True)
|
||||
conn = sqlite3.connect(self._db_path, check_same_thread=False)
|
||||
conn.execute("PRAGMA journal_mode=WAL")
|
||||
conn.execute(_SCHEMA)
|
||||
conn.commit()
|
||||
self._conn = conn
|
||||
return conn
|
||||
except Exception as exc:
|
||||
self._connect_failed = True
|
||||
if conn is not None:
|
||||
try:
|
||||
conn.close()
|
||||
except Exception:
|
||||
pass
|
||||
logger.debug(
|
||||
"fal job store unavailable this session (%s): %s", self._db_path, exc
|
||||
)
|
||||
return None
|
||||
|
||||
# -- writes ---------------------------------------------------------------
|
||||
|
||||
def record_submit(self, endpoint: str, request_id: str, note: str = "") -> None:
|
||||
"""Record a freshly queued job as 'submitted'. Never raises."""
|
||||
try:
|
||||
if not request_id:
|
||||
return
|
||||
with self._lock:
|
||||
conn = self._connection()
|
||||
if conn is None:
|
||||
return
|
||||
conn.execute(
|
||||
"INSERT OR REPLACE INTO jobs "
|
||||
"(request_id, endpoint, status, submitted_at, collected_at, note) "
|
||||
"VALUES (?, ?, ?, ?, NULL, ?)",
|
||||
(request_id, endpoint, _STATUS_SUBMITTED, time.time(), note),
|
||||
)
|
||||
conn.commit()
|
||||
except Exception as exc:
|
||||
logger.debug("job store record_submit failed: %s", exc)
|
||||
self.prune(_PRUNE_AFTER_DAYS)
|
||||
|
||||
def mark_collected(self, request_id: str) -> None:
|
||||
"""Mark a job 'collected'. Unknown ids are inserted silently (recovery
|
||||
of jobs submitted in other sessions). Never raises."""
|
||||
try:
|
||||
if not request_id:
|
||||
return
|
||||
now = time.time()
|
||||
with self._lock:
|
||||
conn = self._connection()
|
||||
if conn is None:
|
||||
return
|
||||
conn.execute(
|
||||
"INSERT OR IGNORE INTO jobs "
|
||||
"(request_id, endpoint, status, submitted_at, collected_at, note) "
|
||||
"VALUES (?, '', ?, ?, ?, '')",
|
||||
(request_id, _STATUS_COLLECTED, now, now),
|
||||
)
|
||||
conn.execute(
|
||||
"UPDATE jobs SET status = ?, collected_at = ? WHERE request_id = ?",
|
||||
(_STATUS_COLLECTED, now, request_id),
|
||||
)
|
||||
conn.commit()
|
||||
except Exception as exc:
|
||||
logger.debug("job store mark_collected failed: %s", exc)
|
||||
|
||||
def prune(self, older_than_days: float = _PRUNE_AFTER_DAYS) -> None:
|
||||
"""Delete jobs submitted more than ``older_than_days`` ago. Never raises.
|
||||
|
||||
fal queue entries expire long before this window, so stale rows are
|
||||
pure noise by then.
|
||||
"""
|
||||
try:
|
||||
cutoff = time.time() - float(older_than_days) * _SECONDS_PER_DAY
|
||||
with self._lock:
|
||||
conn = self._connection()
|
||||
if conn is None:
|
||||
return
|
||||
conn.execute("DELETE FROM jobs WHERE submitted_at < ?", (cutoff,))
|
||||
conn.commit()
|
||||
except Exception as exc:
|
||||
logger.debug("job store prune failed: %s", exc)
|
||||
|
||||
# -- reads ----------------------------------------------------------------
|
||||
|
||||
def entries(self, limit: int = 50, status: str | None = None) -> list[dict[str, Any]]:
|
||||
"""Return jobs newest first as dicts keyed by column name. Never raises."""
|
||||
try:
|
||||
query = f"SELECT {', '.join(_COLUMNS)} FROM jobs"
|
||||
params: tuple[Any, ...] = ()
|
||||
if status:
|
||||
query += " WHERE status = ?"
|
||||
params = (status,)
|
||||
# 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()
|
||||
if conn is None:
|
||||
return []
|
||||
rows = conn.execute(query, params).fetchall()
|
||||
return [dict(zip(_COLUMNS, row)) for row in rows]
|
||||
except Exception as exc:
|
||||
logger.debug("job store entries failed: %s", exc)
|
||||
return []
|
||||
|
||||
def pending(self, limit: int = 50) -> list[dict[str, Any]]:
|
||||
"""Jobs submitted but not yet collected, newest first. Never raises."""
|
||||
return self.entries(limit=limit, status=_STATUS_SUBMITTED)
|
||||
|
||||
def counts(self) -> dict[str, int]:
|
||||
"""Return {'submitted': n, 'collected': n}. Never raises."""
|
||||
result = {_STATUS_SUBMITTED: 0, _STATUS_COLLECTED: 0}
|
||||
try:
|
||||
with self._lock:
|
||||
conn = self._connection()
|
||||
if conn is None:
|
||||
return result
|
||||
rows = conn.execute(
|
||||
"SELECT status, COUNT(*) FROM jobs GROUP BY status"
|
||||
).fetchall()
|
||||
return {**result, **{str(status): int(count) for status, count in rows if status in result}}
|
||||
except Exception as exc:
|
||||
logger.debug("job store counts failed: %s", exc)
|
||||
return result
|
||||
|
||||
def report(self, limit: int = 20) -> str:
|
||||
"""Multi-line human summary of the async job inbox. Never raises."""
|
||||
try:
|
||||
counts = self.counts()
|
||||
pending = counts.get(_STATUS_SUBMITTED, 0)
|
||||
collected = counts.get(_STATUS_COLLECTED, 0)
|
||||
lines = [
|
||||
f"Fal job inbox: {pending} pending, {collected} collected "
|
||||
"(async jobs survive ComfyUI restarts)"
|
||||
]
|
||||
now = time.time()
|
||||
lines.extend(_format_entry(entry, now) for entry in self.entries(limit=limit))
|
||||
if pending == 0 and collected == 0:
|
||||
lines.append(" (empty — queue jobs with Fal Submit to fill the inbox)")
|
||||
return "\n".join(lines)
|
||||
except Exception as exc:
|
||||
logger.debug("job store report failed: %s", exc)
|
||||
return "Fal job inbox: report unavailable"
|
||||
@@ -0,0 +1,126 @@
|
||||
"""Thread-safe in-memory ledger of fal API calls made this ComfyUI session."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
from .logger import logger
|
||||
|
||||
_REPORT_TAIL = 20
|
||||
|
||||
|
||||
def _format_cost(est_cost: Any) -> str:
|
||||
"""Render an estimated cost, or a placeholder when pricing is unknown."""
|
||||
if isinstance(est_cost, (int, float)):
|
||||
return f"~${est_cost:,.4f}".rstrip("0").rstrip(".")
|
||||
return "cost unknown"
|
||||
|
||||
|
||||
def _format_entry(entry: dict[str, Any]) -> str:
|
||||
"""One report line: ' #12 fal-ai/kling.../v3 12.4s ~$0.35 req=abc123'."""
|
||||
duration = entry.get("duration_s")
|
||||
duration_text = f"{duration:.1f}s" if isinstance(duration, (int, float)) else "?s"
|
||||
request_id = entry.get("request_id") or "-"
|
||||
return (
|
||||
f" #{entry.get('index', '?')} {entry.get('endpoint_id', 'unknown')} "
|
||||
f"{duration_text} {_format_cost(entry.get('est_cost'))} req={request_id}"
|
||||
)
|
||||
|
||||
|
||||
class SessionLedger:
|
||||
"""Singleton recording every fal call this session. Methods never raise."""
|
||||
|
||||
_instance: SessionLedger | None = None
|
||||
_instance_lock = threading.Lock()
|
||||
|
||||
def __new__(cls) -> SessionLedger:
|
||||
if cls._instance is None:
|
||||
with cls._instance_lock:
|
||||
if cls._instance is None:
|
||||
instance = super().__new__(cls)
|
||||
instance._initialize()
|
||||
cls._instance = instance
|
||||
return cls._instance
|
||||
|
||||
def _initialize(self) -> None:
|
||||
self._lock = threading.Lock()
|
||||
self._entries: list[dict[str, Any]] = []
|
||||
self._next_index = 1
|
||||
|
||||
def record(
|
||||
self,
|
||||
endpoint_id: str,
|
||||
request_id: str | None,
|
||||
duration_s: float,
|
||||
est_cost: float | None,
|
||||
) -> None:
|
||||
"""Append one call record (indexed, timestamped). Never raises."""
|
||||
try:
|
||||
with self._lock:
|
||||
entry = {
|
||||
"index": self._next_index,
|
||||
"timestamp": time.time(),
|
||||
"endpoint_id": endpoint_id,
|
||||
"request_id": request_id,
|
||||
"duration_s": duration_s,
|
||||
"est_cost": est_cost,
|
||||
}
|
||||
self._entries = [*self._entries, entry]
|
||||
self._next_index += 1
|
||||
except Exception as exc:
|
||||
logger.debug("SessionLedger.record failed: %s", exc)
|
||||
|
||||
def entries(self) -> list[dict[str, Any]]:
|
||||
"""Return a copy of all recorded entries."""
|
||||
try:
|
||||
with self._lock:
|
||||
return [dict(entry) for entry in self._entries]
|
||||
except Exception as exc:
|
||||
logger.debug("SessionLedger.entries failed: %s", exc)
|
||||
return []
|
||||
|
||||
def total_cost(self) -> float:
|
||||
"""Sum of all known estimated costs (unknowns excluded)."""
|
||||
try:
|
||||
return sum(
|
||||
entry["est_cost"]
|
||||
for entry in self.entries()
|
||||
if isinstance(entry.get("est_cost"), (int, float))
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.debug("SessionLedger.total_cost failed: %s", exc)
|
||||
return 0.0
|
||||
|
||||
def unknown_cost_count(self) -> int:
|
||||
"""Number of recorded calls with no pricing estimate."""
|
||||
try:
|
||||
return sum(1 for entry in self.entries() if entry.get("est_cost") is None)
|
||||
except Exception as exc:
|
||||
logger.debug("SessionLedger.unknown_cost_count failed: %s", exc)
|
||||
return 0
|
||||
|
||||
def reset(self) -> None:
|
||||
"""Clear all entries and restart indexing."""
|
||||
try:
|
||||
with self._lock:
|
||||
self._entries = []
|
||||
self._next_index = 1
|
||||
except Exception as exc:
|
||||
logger.debug("SessionLedger.reset failed: %s", exc)
|
||||
|
||||
def report(self) -> str:
|
||||
"""Multi-line human summary of session usage. Never raises."""
|
||||
try:
|
||||
entries = self.entries()
|
||||
lines = [
|
||||
f"Session fal usage: {len(entries)} calls, "
|
||||
f"~${self.total_cost():,.2f} estimated, "
|
||||
f"{self.unknown_cost_count()} with unknown pricing"
|
||||
]
|
||||
lines.extend(_format_entry(entry) for entry in entries[-_REPORT_TAIL:])
|
||||
return "\n".join(lines)
|
||||
except Exception as exc:
|
||||
logger.debug("SessionLedger.report failed: %s", exc)
|
||||
return "Session fal usage: report unavailable"
|
||||
@@ -0,0 +1,22 @@
|
||||
"""Shared logger for the ComfyUI-fal-API node pack."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
_LOGGER_NAME = "ComfyUI-fal-API"
|
||||
_LOG_FORMAT = "[%(name)s] %(levelname)s: %(message)s"
|
||||
|
||||
|
||||
def _configure_logger() -> logging.Logger:
|
||||
"""Configure the package logger exactly once."""
|
||||
log = logging.getLogger(_LOGGER_NAME)
|
||||
if not log.handlers:
|
||||
handler = logging.StreamHandler()
|
||||
handler.setFormatter(logging.Formatter(_LOG_FORMAT))
|
||||
log.addHandler(handler)
|
||||
log.setLevel(logging.INFO)
|
||||
return log
|
||||
|
||||
|
||||
logger = _configure_logger()
|
||||
@@ -0,0 +1,309 @@
|
||||
"""Video/audio helpers for ComfyUI-fal-API (download, decode, upload)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import tempfile
|
||||
import threading
|
||||
from typing import Any
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import numpy as np
|
||||
import requests
|
||||
import torch
|
||||
|
||||
from .errors import FalApiError
|
||||
from .images import ImageUtils
|
||||
from .logger import logger
|
||||
|
||||
_DOWNLOAD_TIMEOUT = (10, 600)
|
||||
_CHUNK_SIZE = 1 << 20 # 1 MiB
|
||||
_video_warning = {"emitted": False, "lock": threading.Lock()}
|
||||
|
||||
|
||||
def _safe_unlink(path: str) -> None:
|
||||
"""Delete a temp file, ignoring errors."""
|
||||
try:
|
||||
os.unlink(path)
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
|
||||
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]
|
||||
return suffix if suffix else default
|
||||
|
||||
|
||||
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 _warn_video_unavailable_once() -> None:
|
||||
"""Warn (once) that ComfyUI VIDEO output support is unavailable."""
|
||||
with _video_warning["lock"]:
|
||||
if not _video_warning["emitted"]:
|
||||
_video_warning["emitted"] = True
|
||||
logger.warning(
|
||||
"comfy_api VideoFromFile is unavailable; VIDEO outputs will be "
|
||||
"None. Update ComfyUI to a version that provides comfy_api."
|
||||
)
|
||||
|
||||
|
||||
def _normalize_av_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 _load_audio_with_av(path: str) -> tuple[torch.Tensor, int]:
|
||||
"""Decode audio with PyAV; returns (waveform (1, C, T) float32, rate)."""
|
||||
import av
|
||||
|
||||
with av.open(path) as container:
|
||||
stream = container.streams.audio[0]
|
||||
sample_rate = int(stream.rate or 44100)
|
||||
channels = int(getattr(stream, "channels", 1) or 1)
|
||||
frames = [
|
||||
_normalize_av_frame(frame.to_ndarray(), channels)
|
||||
for frame in container.decode(stream)
|
||||
]
|
||||
|
||||
if not frames:
|
||||
raise FalApiError("audio-decode", f"No audio frames decoded from {path}")
|
||||
waveform = torch.from_numpy(np.concatenate(frames, axis=1))
|
||||
return waveform.unsqueeze(0), sample_rate
|
||||
|
||||
|
||||
def _load_audio(path: str) -> tuple[torch.Tensor, int]:
|
||||
"""Decode an audio file to (waveform (1, C, T) float32, sample_rate)."""
|
||||
try:
|
||||
import torchaudio
|
||||
|
||||
waveform, sample_rate = torchaudio.load(path)
|
||||
return waveform.to(torch.float32).unsqueeze(0), int(sample_rate)
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
try:
|
||||
return _load_audio_with_av(path)
|
||||
except ImportError as exc:
|
||||
raise FalApiError(
|
||||
"audio-decode",
|
||||
"Decoding audio requires torchaudio or av (PyAV); neither is "
|
||||
"installed. Install one of them (e.g. 'pip install torchaudio').",
|
||||
) from exc
|
||||
|
||||
|
||||
def _save_wav(path: str, waveform: torch.Tensor, sample_rate: int) -> None:
|
||||
"""Save a (C, T) float32 waveform as WAV (torchaudio, else stdlib PCM16)."""
|
||||
try:
|
||||
import torchaudio
|
||||
|
||||
torchaudio.save(path, waveform, sample_rate)
|
||||
return
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
import wave
|
||||
|
||||
clipped = np.clip(waveform.numpy(), -1.0, 1.0)
|
||||
pcm = (clipped * 32767.0).astype(np.int16)
|
||||
with wave.open(path, "wb") as wav_file:
|
||||
wav_file.setnchannels(pcm.shape[0])
|
||||
wav_file.setsampwidth(2)
|
||||
wav_file.setframerate(int(sample_rate))
|
||||
wav_file.writeframes(pcm.T.reshape(-1).tobytes())
|
||||
|
||||
|
||||
def _stream_to_temp_file(source: Any, suffix: str) -> str:
|
||||
"""Write a readable stream to a temp file and return its path."""
|
||||
with tempfile.NamedTemporaryFile(suffix=suffix, 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
|
||||
|
||||
|
||||
class MediaUtils:
|
||||
"""Utility functions for video/audio download, conversion, and upload.
|
||||
|
||||
Local-file uploads (upload_video/upload_audio) route through
|
||||
ImageUtils.upload_file, which consults the persistent upload cache
|
||||
(sha256 of file bytes) so identical content is never uploaded twice.
|
||||
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:
|
||||
resp.raise_for_status()
|
||||
with tempfile.NamedTemporaryFile(
|
||||
suffix=suffix, delete=False
|
||||
) as temp_file:
|
||||
temp_path = temp_file.name
|
||||
for chunk in resp.iter_content(chunk_size=_CHUNK_SIZE):
|
||||
if chunk:
|
||||
temp_file.write(chunk)
|
||||
return temp_path
|
||||
except Exception as exc:
|
||||
if temp_path is not None:
|
||||
_safe_unlink(temp_path)
|
||||
logger.error("Failed to download %s: %s", url, exc)
|
||||
raise FalApiError(
|
||||
"media-download", f"Failed to download {url}: {exc}"
|
||||
) from exc
|
||||
|
||||
@staticmethod
|
||||
def video_from_url(url: str) -> Any | None:
|
||||
"""Download a video URL and wrap it as a ComfyUI VIDEO object."""
|
||||
video_cls = _resolve_video_from_file()
|
||||
if video_cls is None:
|
||||
_warn_video_unavailable_once()
|
||||
return None
|
||||
# NOTE: the temp file is deliberately not unlinked here — VideoFromFile
|
||||
# reads the path lazily (e.g. when a downstream save node consumes it),
|
||||
# so deleting early would break playback. The OS temp dir reclaims it.
|
||||
local_path = MediaUtils.download_url_to_temp(
|
||||
url, _suffix_from_url(url, default=".mp4")
|
||||
)
|
||||
return video_cls(local_path)
|
||||
|
||||
@staticmethod
|
||||
def audio_from_url(url: str) -> dict[str, Any]:
|
||||
"""Download and decode audio into a ComfyUI AUDIO dict.
|
||||
|
||||
Returns {"waveform": float32 tensor (1, C, T), "sample_rate": int}.
|
||||
"""
|
||||
local_path = MediaUtils.download_url_to_temp(
|
||||
url, _suffix_from_url(url, default=".wav")
|
||||
)
|
||||
try:
|
||||
waveform, sample_rate = _load_audio(local_path)
|
||||
return {"waveform": waveform, "sample_rate": sample_rate}
|
||||
except FalApiError:
|
||||
raise
|
||||
except Exception as exc:
|
||||
logger.error("Failed to decode audio from %s: %s", url, exc)
|
||||
raise FalApiError(
|
||||
"audio-decode", f"Failed to decode audio from {url}: {exc}"
|
||||
) from exc
|
||||
finally:
|
||||
_safe_unlink(local_path)
|
||||
|
||||
@staticmethod
|
||||
def upload_video(video: Any) -> str:
|
||||
"""Upload a ComfyUI VIDEO input (or path/url string) and return a URL."""
|
||||
if isinstance(video, str):
|
||||
return video if _is_http_url(video) else ImageUtils.upload_file(video)
|
||||
|
||||
source = (
|
||||
video.get_stream_source()
|
||||
if hasattr(video, "get_stream_source")
|
||||
else video
|
||||
)
|
||||
if isinstance(source, str) and _is_http_url(source):
|
||||
return source
|
||||
if hasattr(source, "read"):
|
||||
temp_path = _stream_to_temp_file(source, suffix=".mp4")
|
||||
try:
|
||||
return ImageUtils.upload_file(temp_path)
|
||||
finally:
|
||||
_safe_unlink(temp_path)
|
||||
return ImageUtils.upload_file(source)
|
||||
|
||||
@staticmethod
|
||||
def upload_audio(audio: Any) -> str:
|
||||
"""Upload a ComfyUI AUDIO dict (or path/url string) and return a URL."""
|
||||
if isinstance(audio, str):
|
||||
return audio if _is_http_url(audio) else ImageUtils.upload_file(audio)
|
||||
|
||||
try:
|
||||
waveform = audio["waveform"]
|
||||
sample_rate = int(audio["sample_rate"])
|
||||
except (KeyError, TypeError) as exc:
|
||||
raise FalApiError(
|
||||
"audio-upload",
|
||||
"Expected an AUDIO dict with 'waveform' and 'sample_rate'",
|
||||
) from exc
|
||||
|
||||
tensor = waveform.detach().cpu().to(torch.float32)
|
||||
if tensor.ndim == 3:
|
||||
tensor = tensor[0] # (1, C, T) -> (C, T)
|
||||
|
||||
temp_path: str | None = None
|
||||
try:
|
||||
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as temp_file:
|
||||
temp_path = temp_file.name
|
||||
_save_wav(temp_path, tensor, sample_rate)
|
||||
return ImageUtils.upload_file(temp_path)
|
||||
except FalApiError:
|
||||
raise
|
||||
except Exception as exc:
|
||||
logger.error("Failed to save/upload audio: %s", exc)
|
||||
raise FalApiError(
|
||||
"audio-upload", f"Failed to save/upload audio: {exc}"
|
||||
) from exc
|
||||
finally:
|
||||
if temp_path is not None:
|
||||
_safe_unlink(temp_path)
|
||||
@@ -0,0 +1,231 @@
|
||||
"""Pricing parsing and cost estimation from the committed fal model registry."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import threading
|
||||
from typing import Any
|
||||
|
||||
from .logger import logger
|
||||
|
||||
# Units that imply one billable output per run (single-output assumption).
|
||||
_RUN_UNITS = frozenset({"image", "video", "generation", "request", "run"})
|
||||
|
||||
# Tokens that terminate a unit phrase ("$0.015 per megapixel in TURBO mode").
|
||||
_UNIT_STOPWORDS = frozenset(
|
||||
{
|
||||
"a",
|
||||
"along",
|
||||
"an",
|
||||
"and",
|
||||
"are",
|
||||
"at",
|
||||
"each",
|
||||
"for",
|
||||
"if",
|
||||
"in",
|
||||
"is",
|
||||
"on",
|
||||
"or",
|
||||
"per",
|
||||
"plus",
|
||||
"rounded",
|
||||
"the",
|
||||
"to",
|
||||
"when",
|
||||
"will",
|
||||
"with",
|
||||
"without",
|
||||
}
|
||||
)
|
||||
|
||||
_MONEY = r"([\d,]+(?:\.\d+)?)"
|
||||
|
||||
# "For $1.00, you can run this model (with approximately) 12 times"
|
||||
_RUNS_RATIO_RE = re.compile(
|
||||
rf"for\s+\${_MONEY},?\s+you\s+can\s+run\s+this\s+model\s+"
|
||||
rf"(?:with\s+)?(?:approximately\s+)?{_MONEY}\s+times",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
|
||||
# "$0.05 per second of video", "$0.025 per 1000 characters", "$0.08 per image"
|
||||
_PER_UNIT_RE = re.compile(
|
||||
rf"\$\s*{_MONEY}\s+per\s+(\w+(?:\s+\w+){{0,3}})",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
|
||||
_registry_lock = threading.Lock()
|
||||
_pricing_map: dict[str, str] | None = None
|
||||
|
||||
|
||||
def _registry_path() -> str:
|
||||
"""Return the path to data/fal_registry.json at the repo root."""
|
||||
utils_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
nodes_dir = os.path.dirname(utils_dir)
|
||||
repo_root = os.path.dirname(nodes_dir)
|
||||
return os.path.join(repo_root, "data", "fal_registry.json")
|
||||
|
||||
|
||||
def _load_pricing_map() -> dict[str, str]:
|
||||
"""Lazily load the endpoint_id -> pricing-text map (cached module-wide)."""
|
||||
global _pricing_map
|
||||
if _pricing_map is not None:
|
||||
return _pricing_map
|
||||
with _registry_lock:
|
||||
if _pricing_map is not None:
|
||||
return _pricing_map
|
||||
mapping: dict[str, str] = {}
|
||||
try:
|
||||
with open(_registry_path(), encoding="utf-8") as fh:
|
||||
registry = json.load(fh)
|
||||
for model in registry.get("models") or []:
|
||||
endpoint_id = model.get("endpoint_id")
|
||||
if endpoint_id:
|
||||
mapping[endpoint_id] = str(model.get("pricing") or "")
|
||||
except Exception as exc:
|
||||
logger.warning("Could not load pricing registry: %s", exc)
|
||||
_pricing_map = mapping
|
||||
return _pricing_map
|
||||
|
||||
|
||||
def _to_float(text: str) -> float:
|
||||
"""Parse a dollar/count figure, tolerating thousands separators."""
|
||||
return float(text.replace(",", ""))
|
||||
|
||||
|
||||
def _clean_unit(phrase: str) -> str | None:
|
||||
"""Trim a raw unit capture to the meaningful phrase, or None if empty."""
|
||||
tokens = phrase.lower().split()
|
||||
kept: list[str] = []
|
||||
for position, token in enumerate(tokens):
|
||||
stripped = token.strip(".,")
|
||||
if stripped in _UNIT_STOPWORDS:
|
||||
break
|
||||
if stripped == "of":
|
||||
has_object = (
|
||||
bool(kept)
|
||||
and position + 1 < len(tokens)
|
||||
and tokens[position + 1].strip(".,") not in _UNIT_STOPWORDS
|
||||
)
|
||||
if not has_object:
|
||||
break
|
||||
kept.append(stripped)
|
||||
cleaned = " ".join(kept)
|
||||
return cleaned or None
|
||||
|
||||
|
||||
def _head_noun(unit: str) -> str:
|
||||
"""Return the singular head noun of a unit phrase.
|
||||
|
||||
"second of video" -> "second"; "generated image" -> "image";
|
||||
"image generated" -> "image" (trailing participles are dropped).
|
||||
"""
|
||||
tokens = unit.split(" of ")[0].split()
|
||||
while len(tokens) > 1 and tokens[-1].endswith("ed"):
|
||||
tokens = tokens[:-1]
|
||||
head = tokens[-1]
|
||||
if head.endswith("s") and not head.endswith("ss"):
|
||||
return head[:-1]
|
||||
return head
|
||||
|
||||
|
||||
def _format_amount(value: float) -> str:
|
||||
"""Format a dollar amount compactly (up to 4 decimals, no trailing zeros)."""
|
||||
text = f"{value:,.4f}".rstrip("0").rstrip(".")
|
||||
return text or "0"
|
||||
|
||||
|
||||
def _empty_parse(raw: str) -> dict[str, Any]:
|
||||
return {"per_run": None, "per_unit": None, "unit": None, "raw": raw}
|
||||
|
||||
|
||||
class PricingUtils:
|
||||
"""Best-effort cost estimation from the registry's human pricing strings."""
|
||||
|
||||
@staticmethod
|
||||
def parse(pricing_text: str) -> dict[str, Any]:
|
||||
"""Parse a human pricing string into structured numbers.
|
||||
|
||||
Precedence for ``per_run``: the explicit "For $A, you can run this
|
||||
model N times" ratio wins over a "$X per <unit>" price; the latter
|
||||
only sets ``per_run`` for single-output units (image, video,
|
||||
generation, request, run). Never raises; unmatched strings return
|
||||
all-None fields with ``raw`` preserved.
|
||||
"""
|
||||
raw = pricing_text or ""
|
||||
result = _empty_parse(raw)
|
||||
try:
|
||||
ratio = _RUNS_RATIO_RE.search(raw)
|
||||
if ratio:
|
||||
runs = _to_float(ratio.group(2))
|
||||
if runs > 0:
|
||||
result = {**result, "per_run": _to_float(ratio.group(1)) / runs}
|
||||
|
||||
per_unit = _PER_UNIT_RE.search(raw)
|
||||
if per_unit:
|
||||
unit = _clean_unit(per_unit.group(2))
|
||||
if unit:
|
||||
result = {
|
||||
**result,
|
||||
"per_unit": _to_float(per_unit.group(1)),
|
||||
"unit": unit,
|
||||
}
|
||||
if result["per_run"] is None and _head_noun(unit) in _RUN_UNITS:
|
||||
result = {**result, "per_run": result["per_unit"]}
|
||||
except Exception as exc:
|
||||
logger.debug("Failed to parse pricing text %r: %s", raw, exc)
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def pricing_for(endpoint_id: str) -> dict[str, Any] | None:
|
||||
"""Parsed pricing for an endpoint, or None if unknown/unpublished."""
|
||||
pricing_text = _load_pricing_map().get(endpoint_id)
|
||||
if not pricing_text:
|
||||
return None
|
||||
return PricingUtils.parse(pricing_text)
|
||||
|
||||
@staticmethod
|
||||
def estimate(endpoint_id: str, runs: int = 1) -> dict[str, Any]:
|
||||
"""Estimate the cost of ``runs`` runs of an endpoint."""
|
||||
parsed = PricingUtils.pricing_for(endpoint_id) or _empty_parse("")
|
||||
per_run = parsed["per_run"]
|
||||
unit_note = ""
|
||||
if parsed["per_unit"] is not None and parsed["unit"]:
|
||||
unit_note = f"${_format_amount(parsed['per_unit'])} per {parsed['unit']}"
|
||||
return {
|
||||
"endpoint_id": endpoint_id,
|
||||
"runs": runs,
|
||||
"per_run": per_run,
|
||||
"unit_note": unit_note,
|
||||
"total": per_run * runs if per_run is not None else None,
|
||||
"raw": parsed["raw"],
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def format_report(estimate: dict[str, Any]) -> str:
|
||||
"""Render an estimate as a compact human string. Never raises."""
|
||||
try:
|
||||
endpoint_id = estimate.get("endpoint_id") or "unknown"
|
||||
runs = estimate.get("runs") or 1
|
||||
per_run = estimate.get("per_run")
|
||||
total = estimate.get("total")
|
||||
unit_note = estimate.get("unit_note") or ""
|
||||
|
||||
if per_run is not None:
|
||||
report = f"{endpoint_id}: ~${_format_amount(per_run)}/run"
|
||||
if runs != 1 and total is not None:
|
||||
report += f" → ~${_format_amount(total)} for {runs} runs"
|
||||
return report
|
||||
if unit_note:
|
||||
depends_on = (
|
||||
"duration"
|
||||
if re.search(r"\b(second|minute|hour)s?\b", unit_note)
|
||||
else "output size"
|
||||
)
|
||||
return f"{endpoint_id}: {unit_note} (per-run total depends on {depends_on})"
|
||||
return f"{endpoint_id}: pricing not published"
|
||||
except Exception as exc:
|
||||
logger.debug("format_report failed: %s", exc)
|
||||
return "pricing not published"
|
||||
@@ -0,0 +1,442 @@
|
||||
"""Persistent two-layer cache so the same fal call is never paid for twice.
|
||||
|
||||
Layer 1 (``results``) caches full API results keyed by endpoint + canonical
|
||||
arguments. Layer 2 (``uploads``) caches fal media URLs keyed by the sha256 of
|
||||
the uploaded file's bytes, so re-uploading identical content reuses the same
|
||||
URL (which in turn keeps the result-cache key stable).
|
||||
|
||||
Both layers live in one sqlite database that survives ComfyUI restarts.
|
||||
Every public method is best-effort: any sqlite/config failure degrades to a
|
||||
cache miss or a no-op — bookkeeping must never break generation.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
import sqlite3
|
||||
import threading
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
from .config import FalConfig
|
||||
from .logger import logger
|
||||
|
||||
_DB_ENV_VAR = "COMFYUI_FAL_API_CACHE_DB"
|
||||
_DB_SUBDIR = "comfyui-fal-api"
|
||||
_DB_FILENAME = "cache.db"
|
||||
|
||||
_DEFAULT_ENABLED = True
|
||||
_DEFAULT_TTL_DAYS = 7 # fal CDN URLs inside cached results can expire; keep modest.
|
||||
_DEFAULT_MAX_ENTRIES = 5000
|
||||
_SECONDS_PER_DAY = 86400.0
|
||||
|
||||
_SCHEMA = (
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS results (
|
||||
key TEXT PRIMARY KEY,
|
||||
endpoint TEXT,
|
||||
request_id TEXT,
|
||||
result_json TEXT,
|
||||
created REAL,
|
||||
last_used REAL
|
||||
)
|
||||
""",
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS uploads (
|
||||
content_hash TEXT PRIMARY KEY,
|
||||
url TEXT,
|
||||
created REAL
|
||||
)
|
||||
""",
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS request_urls (
|
||||
url TEXT PRIMARY KEY,
|
||||
endpoint TEXT,
|
||||
request_id TEXT,
|
||||
created REAL
|
||||
)
|
||||
""",
|
||||
)
|
||||
|
||||
|
||||
def _default_db_path() -> str:
|
||||
"""Resolve the cache database path.
|
||||
|
||||
Precedence: COMFYUI_FAL_API_CACHE_DB env var, then the ComfyUI user
|
||||
directory, then ~/.cache (when running outside ComfyUI).
|
||||
"""
|
||||
override = os.environ.get(_DB_ENV_VAR)
|
||||
if override:
|
||||
return override
|
||||
try:
|
||||
import folder_paths
|
||||
|
||||
base = folder_paths.get_user_directory()
|
||||
except ImportError:
|
||||
base = os.path.join(os.path.expanduser("~"), ".cache")
|
||||
return os.path.join(base, _DB_SUBDIR, _DB_FILENAME)
|
||||
|
||||
|
||||
def _format_amount(value: float) -> str:
|
||||
"""Format a dollar amount compactly (up to 4 decimals, no trailing zeros)."""
|
||||
text = f"{value:,.4f}".rstrip("0").rstrip(".")
|
||||
return text or "0"
|
||||
|
||||
|
||||
class ResultCache:
|
||||
"""Thread-safe singleton over the persistent result/upload cache."""
|
||||
|
||||
_instance: ResultCache | None = None
|
||||
_instance_lock = threading.Lock()
|
||||
|
||||
def __new__(cls) -> ResultCache:
|
||||
if cls._instance is None:
|
||||
with cls._instance_lock:
|
||||
if cls._instance is None:
|
||||
instance = super().__new__(cls)
|
||||
instance._initialize()
|
||||
cls._instance = instance
|
||||
return cls._instance
|
||||
|
||||
def _initialize(self) -> None:
|
||||
self._lock = threading.Lock()
|
||||
self._conn: sqlite3.Connection | None = None
|
||||
self._connect_failed = False
|
||||
self._db_path = _default_db_path()
|
||||
self._hits = 0
|
||||
self._misses = 0
|
||||
|
||||
# -- connection / config ------------------------------------------------
|
||||
|
||||
def _connection(self) -> sqlite3.Connection | None:
|
||||
"""Open (once) and return the sqlite connection; None if unavailable.
|
||||
|
||||
Must be called with ``self._lock`` held. A corrupted or unwritable
|
||||
database disables the cache for the session instead of raising.
|
||||
"""
|
||||
if self._conn is not None:
|
||||
return self._conn
|
||||
if self._connect_failed:
|
||||
return None
|
||||
conn: sqlite3.Connection | None = None
|
||||
try:
|
||||
os.makedirs(os.path.dirname(self._db_path), exist_ok=True)
|
||||
conn = sqlite3.connect(self._db_path, check_same_thread=False)
|
||||
conn.execute("PRAGMA journal_mode=WAL")
|
||||
for statement in _SCHEMA:
|
||||
conn.execute(statement)
|
||||
conn.commit()
|
||||
self._conn = conn
|
||||
return conn
|
||||
except Exception as exc:
|
||||
self._connect_failed = True
|
||||
if conn is not None:
|
||||
try:
|
||||
conn.close()
|
||||
except Exception:
|
||||
pass
|
||||
logger.debug(
|
||||
"fal cache unavailable, caching disabled this session (%s): %s",
|
||||
self._db_path,
|
||||
exc,
|
||||
)
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _enabled() -> bool:
|
||||
try:
|
||||
return bool(FalConfig().get_setting("cache", "enabled", _DEFAULT_ENABLED))
|
||||
except Exception as exc:
|
||||
logger.debug("cache 'enabled' setting read failed: %s", exc)
|
||||
return _DEFAULT_ENABLED
|
||||
|
||||
@staticmethod
|
||||
def _ttl_seconds() -> float:
|
||||
try:
|
||||
days = float(FalConfig().get_setting("cache", "ttl_days", _DEFAULT_TTL_DAYS))
|
||||
except Exception as exc:
|
||||
logger.debug("cache 'ttl_days' setting read failed: %s", exc)
|
||||
days = float(_DEFAULT_TTL_DAYS)
|
||||
return days * _SECONDS_PER_DAY
|
||||
|
||||
@staticmethod
|
||||
def _max_entries() -> int:
|
||||
try:
|
||||
return int(float(FalConfig().get_setting("cache", "max_entries", _DEFAULT_MAX_ENTRIES)))
|
||||
except Exception as exc:
|
||||
logger.debug("cache 'max_entries' setting read failed: %s", exc)
|
||||
return _DEFAULT_MAX_ENTRIES
|
||||
|
||||
# -- layer 1: result cache ----------------------------------------------
|
||||
|
||||
@staticmethod
|
||||
def make_key(endpoint: str, arguments: dict[str, Any]) -> str:
|
||||
"""Deterministic cache key: sha256 of endpoint + canonical arguments."""
|
||||
canonical = json.dumps(arguments, sort_keys=True, separators=(",", ":"), default=str)
|
||||
return hashlib.sha256(f"{endpoint}\x00{canonical}".encode()).hexdigest()
|
||||
|
||||
def get(self, endpoint: str, arguments: dict[str, Any]) -> dict[str, Any] | None:
|
||||
"""Return the cached result dict for this exact call, or None on miss."""
|
||||
try:
|
||||
if not self._enabled():
|
||||
return None
|
||||
result_json = self._fetch_live_result_json(self.make_key(endpoint, arguments))
|
||||
if result_json is not None:
|
||||
result = json.loads(result_json)
|
||||
if isinstance(result, dict):
|
||||
self._hits += 1
|
||||
self._log_hit(endpoint)
|
||||
return result
|
||||
self._misses += 1
|
||||
return None
|
||||
except Exception as exc:
|
||||
logger.debug("[%s] result cache get failed: %s", endpoint, exc)
|
||||
return None
|
||||
|
||||
def _fetch_live_result_json(self, key: str) -> str | None:
|
||||
"""Fetch a non-expired row's JSON, deleting expired rows on the way."""
|
||||
now = time.time()
|
||||
ttl_seconds = self._ttl_seconds()
|
||||
with self._lock:
|
||||
conn = self._connection()
|
||||
if conn is None:
|
||||
return None
|
||||
row = conn.execute(
|
||||
"SELECT result_json, created FROM results WHERE key = ?", (key,)
|
||||
).fetchone()
|
||||
if row is None:
|
||||
return None
|
||||
result_json, created = row
|
||||
if ttl_seconds > 0 and now - float(created or 0.0) > ttl_seconds:
|
||||
conn.execute("DELETE FROM results WHERE key = ?", (key,))
|
||||
conn.commit()
|
||||
return None
|
||||
conn.execute("UPDATE results SET last_used = ? WHERE key = ?", (now, key))
|
||||
conn.commit()
|
||||
return str(result_json)
|
||||
|
||||
@staticmethod
|
||||
def _log_hit(endpoint: str) -> None:
|
||||
"""Log a cache hit, with the estimated cost saved when pricing is known."""
|
||||
saved = ""
|
||||
try:
|
||||
from .pricing import PricingUtils
|
||||
|
||||
total = PricingUtils.estimate(endpoint, 1)["total"]
|
||||
if isinstance(total, (int, float)):
|
||||
saved = f" (saved ~${_format_amount(total)})"
|
||||
except Exception:
|
||||
saved = ""
|
||||
logger.info("[%s] cache HIT%s — returning stored result, no charge", endpoint, saved)
|
||||
|
||||
def put(
|
||||
self,
|
||||
endpoint: str,
|
||||
arguments: dict[str, Any],
|
||||
result: dict[str, Any],
|
||||
request_id: str | None = None,
|
||||
) -> None:
|
||||
"""Store a successful live result; prunes oldest rows beyond max_entries."""
|
||||
try:
|
||||
if not self._enabled() or not isinstance(result, dict):
|
||||
return
|
||||
key = self.make_key(endpoint, arguments)
|
||||
result_json = json.dumps(result, default=str)
|
||||
max_entries = self._max_entries()
|
||||
now = time.time()
|
||||
with self._lock:
|
||||
conn = self._connection()
|
||||
if conn is None:
|
||||
return
|
||||
conn.execute(
|
||||
"INSERT OR REPLACE INTO results "
|
||||
"(key, endpoint, request_id, result_json, created, last_used) "
|
||||
"VALUES (?, ?, ?, ?, ?, ?)",
|
||||
(key, endpoint, request_id, result_json, now, now),
|
||||
)
|
||||
if max_entries > 0:
|
||||
conn.execute(
|
||||
"DELETE FROM results WHERE key IN "
|
||||
"(SELECT key FROM results ORDER BY last_used DESC LIMIT -1 OFFSET ?)",
|
||||
(max_entries,),
|
||||
)
|
||||
conn.commit()
|
||||
except Exception as exc:
|
||||
logger.debug("[%s] result cache put failed: %s", endpoint, exc)
|
||||
|
||||
def invalidate_key(self, key: str) -> None:
|
||||
"""Delete one result row by its cache key. Never raises."""
|
||||
try:
|
||||
with self._lock:
|
||||
conn = self._connection()
|
||||
if conn is None:
|
||||
return
|
||||
conn.execute("DELETE FROM results WHERE key = ?", (key,))
|
||||
conn.commit()
|
||||
except Exception as exc:
|
||||
logger.debug("result cache invalidate failed: %s", exc)
|
||||
|
||||
def clear(self) -> None:
|
||||
"""Delete all cached results and uploads. Never raises."""
|
||||
try:
|
||||
with self._lock:
|
||||
conn = self._connection()
|
||||
if conn is None:
|
||||
return
|
||||
conn.execute("DELETE FROM results")
|
||||
conn.execute("DELETE FROM uploads")
|
||||
conn.commit()
|
||||
except Exception as exc:
|
||||
logger.debug("result cache clear failed: %s", exc)
|
||||
|
||||
def stats(self) -> dict[str, Any]:
|
||||
"""Return {entries, db_path, hits, misses}; hit/miss counts are per session."""
|
||||
entries = 0
|
||||
try:
|
||||
with self._lock:
|
||||
conn = self._connection()
|
||||
if conn is not None:
|
||||
row = conn.execute("SELECT COUNT(*) FROM results").fetchone()
|
||||
entries = int(row[0]) if row else 0
|
||||
except Exception as exc:
|
||||
logger.debug("result cache stats failed: %s", exc)
|
||||
return {
|
||||
"entries": entries,
|
||||
"db_path": self._db_path,
|
||||
"hits": self._hits,
|
||||
"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.
|
||||
|
||||
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 = (
|
||||
target.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_")
|
||||
)
|
||||
with self._lock:
|
||||
conn = self._connection()
|
||||
if conn is None:
|
||||
return None
|
||||
row = conn.execute(
|
||||
"SELECT endpoint, request_id FROM results "
|
||||
"WHERE request_id IS NOT NULL "
|
||||
"AND result_json LIKE '%' || ? || '%' ESCAPE '\\' "
|
||||
"ORDER BY last_used DESC LIMIT 1",
|
||||
(escaped,),
|
||||
).fetchone()
|
||||
if row is None:
|
||||
return None
|
||||
endpoint, request_id = row
|
||||
return {"endpoint_id": endpoint, "request_id": request_id}
|
||||
except Exception as exc:
|
||||
logger.debug("result cache url lookup failed: %s", exc)
|
||||
return None
|
||||
|
||||
# -- layer 2: upload cache ----------------------------------------------
|
||||
|
||||
def get_upload(self, content_hash: str) -> str | None:
|
||||
"""Return the cached fal URL for previously uploaded content, or None."""
|
||||
try:
|
||||
if not self._enabled():
|
||||
return None
|
||||
now = time.time()
|
||||
ttl_seconds = self._ttl_seconds()
|
||||
with self._lock:
|
||||
conn = self._connection()
|
||||
if conn is None:
|
||||
return None
|
||||
row = conn.execute(
|
||||
"SELECT url, created FROM uploads WHERE content_hash = ?",
|
||||
(content_hash,),
|
||||
).fetchone()
|
||||
if row is None:
|
||||
return None
|
||||
url, created = row
|
||||
if ttl_seconds > 0 and now - float(created or 0.0) > ttl_seconds:
|
||||
conn.execute(
|
||||
"DELETE FROM uploads WHERE content_hash = ?", (content_hash,)
|
||||
)
|
||||
conn.commit()
|
||||
return None
|
||||
return str(url) if url else None
|
||||
except Exception as exc:
|
||||
logger.debug("upload cache get failed: %s", exc)
|
||||
return None
|
||||
|
||||
def put_upload(self, content_hash: str, url: str) -> None:
|
||||
"""Remember that this content hash uploaded to ``url``. Never raises."""
|
||||
try:
|
||||
if not self._enabled():
|
||||
return
|
||||
with self._lock:
|
||||
conn = self._connection()
|
||||
if conn is None:
|
||||
return
|
||||
conn.execute(
|
||||
"INSERT OR REPLACE INTO uploads (content_hash, url, created) "
|
||||
"VALUES (?, ?, ?)",
|
||||
(content_hash, url, time.time()),
|
||||
)
|
||||
conn.commit()
|
||||
except Exception as exc:
|
||||
logger.debug("upload cache put failed: %s", exc)
|
||||
+3573
-587
File diff suppressed because it is too large
Load Diff
+103
-80
@@ -1,53 +1,86 @@
|
||||
import os
|
||||
import configparser
|
||||
from fal_client.client import SyncClient
|
||||
import torch
|
||||
from PIL import Image
|
||||
import tempfile
|
||||
import numpy as np
|
||||
from .fal_utils import ApiHandler, FalConfig, ImageUtils
|
||||
|
||||
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
parent_dir = os.path.dirname(current_dir)
|
||||
config_path = os.path.join(parent_dir, "config.ini")
|
||||
# Initialize FalConfig
|
||||
fal_config = FalConfig()
|
||||
|
||||
config = configparser.ConfigParser()
|
||||
config.read(config_path)
|
||||
|
||||
try:
|
||||
if os.environ.get("FAL_KEY") is not None:
|
||||
print("FAL_KEY found in environment variables")
|
||||
fal_key = os.environ["FAL_KEY"]
|
||||
else:
|
||||
print("FAL_KEY not found in environment variables")
|
||||
fal_key = config['API']['FAL_KEY']
|
||||
print("FAL_KEY found in config.ini")
|
||||
os.environ["FAL_KEY"] = fal_key
|
||||
print("FAL_KEY set in environment variables")
|
||||
|
||||
# Check if FAL key is the default placeholder
|
||||
if fal_key == "<your_fal_api_key_here>":
|
||||
print("WARNING: You are using the default FAL API key placeholder!")
|
||||
print("Please set your actual FAL API key in either:")
|
||||
print("1. The config.ini file under [API] section")
|
||||
print("2. Or as an environment variable named FAL_KEY")
|
||||
print("Get your API key from: https://fal.ai/dashboard/keys")
|
||||
except KeyError:
|
||||
print("Error: FAL_KEY not found in config.ini or environment variables")
|
||||
|
||||
# Create the client with API key
|
||||
fal_client = SyncClient(key=fal_key)
|
||||
|
||||
class VLMNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"prompt": ("STRING", {"default": "", "multiline": True}),
|
||||
"model": (["google/gemini-flash-1.5-8b", "anthropic/claude-3.5-sonnet", "anthropic/claude-3-haiku",
|
||||
"google/gemini-pro-1.5", "google/gemini-flash-1.5", "openai/gpt-4o"],
|
||||
{"default": "google/gemini-flash-1.5-8b"}),
|
||||
"system_prompt": ("STRING", {"default": "", "multiline": True}),
|
||||
"image": ("IMAGE",),
|
||||
"prompt": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"multiline": True,
|
||||
"tooltip": "User prompt sent to the model.",
|
||||
},
|
||||
),
|
||||
"model": (
|
||||
[
|
||||
"google/gemini-2.5-flash",
|
||||
"anthropic/claude-sonnet-4.5",
|
||||
"openai/gpt-4o",
|
||||
"qwen/qwen3-vl-235b-a22b-instruct",
|
||||
"x-ai/grok-4-fast",
|
||||
"Custom",
|
||||
],
|
||||
{
|
||||
"default": "google/gemini-2.5-flash",
|
||||
"tooltip": "Vision model to use. Select 'Custom' to type any OpenRouter model id in custom_model_name.",
|
||||
},
|
||||
),
|
||||
"system_prompt": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"multiline": True,
|
||||
"tooltip": "Optional system prompt to steer the model's behavior.",
|
||||
},
|
||||
),
|
||||
"image": (
|
||||
"IMAGE",
|
||||
{
|
||||
"tooltip": "Image(s) for the model to analyze. Batches are sent as multiple images.",
|
||||
},
|
||||
),
|
||||
"temperature": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 1.0,
|
||||
"min": 0.0,
|
||||
"max": 2.0,
|
||||
"step": 0.1,
|
||||
"tooltip": "Sampling temperature. Lower is more deterministic.",
|
||||
},
|
||||
),
|
||||
"reasoning": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": False,
|
||||
"tooltip": "Request reasoning from the model.",
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"max_tokens": (
|
||||
"INT",
|
||||
{
|
||||
"default": 0,
|
||||
"min": 0,
|
||||
"max": 100000,
|
||||
"tooltip": "Maximum output tokens. 0 uses the model default.",
|
||||
},
|
||||
),
|
||||
"custom_model_name": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"multiline": False,
|
||||
"tooltip": "OpenRouter model id used when model is set to 'Custom'.",
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -55,54 +88,44 @@ class VLMNode:
|
||||
FUNCTION = "generate_text"
|
||||
CATEGORY = "FAL/VLM"
|
||||
|
||||
def generate_text(self, prompt, model, system_prompt, image):
|
||||
def generate_text(self, prompt, model, system_prompt, image, temperature, reasoning, max_tokens=0, custom_model_name=""):
|
||||
try:
|
||||
# Convert the image tensor to a numpy array
|
||||
if isinstance(image, torch.Tensor):
|
||||
image_np = image.cpu().numpy()
|
||||
else:
|
||||
image_np = np.array(image)
|
||||
# Handle custom model selection
|
||||
if model == "Custom":
|
||||
if not custom_model_name or custom_model_name.strip() == "":
|
||||
return ApiHandler.handle_text_generation_error(
|
||||
"Custom", "Custom model name is required when 'Custom' is selected"
|
||||
)
|
||||
model = custom_model_name.strip()
|
||||
|
||||
# Ensure the image is in the correct format (H, W, C)
|
||||
if image_np.ndim == 4:
|
||||
image_np = image_np.squeeze(0) # Remove batch dimension if present
|
||||
if image_np.ndim == 2:
|
||||
image_np = np.stack([image_np] * 3, axis=-1) # Convert grayscale to RGB
|
||||
elif image_np.shape[0] == 3:
|
||||
image_np = np.transpose(image_np, (1, 2, 0)) # Change from (C, H, W) to (H, W, C)
|
||||
|
||||
# Normalize the image data to 0-255 range
|
||||
if image_np.dtype == np.float32 or image_np.dtype == np.float64:
|
||||
image_np = (image_np * 255).astype(np.uint8)
|
||||
|
||||
# Convert to PIL Image
|
||||
pil_image = Image.fromarray(image_np)
|
||||
|
||||
# Save the image to a temporary file
|
||||
with tempfile.NamedTemporaryFile(suffix=".png", delete=False) as temp_file:
|
||||
pil_image.save(temp_file, format="PNG")
|
||||
temp_file_path = temp_file.name
|
||||
|
||||
# Upload the temporary file
|
||||
image_url = fal_client.upload_file(temp_file_path)
|
||||
# Upload single image or batch and collect URLs
|
||||
image_urls = ImageUtils.prepare_images(image)
|
||||
if not image_urls:
|
||||
return ApiHandler.handle_text_generation_error(
|
||||
model, "Failed to upload image(s)"
|
||||
)
|
||||
|
||||
arguments = {
|
||||
"model": model,
|
||||
"prompt": prompt,
|
||||
"system_prompt": system_prompt,
|
||||
"image_url": image_url,
|
||||
"image_urls": image_urls,
|
||||
"temperature": temperature,
|
||||
"reasoning": reasoning,
|
||||
"stream": False,
|
||||
}
|
||||
|
||||
handler = fal_client.submit("fal-ai/any-llm/vision", arguments=arguments)
|
||||
result = handler.get()
|
||||
# Only include max_tokens if it's greater than 0
|
||||
if max_tokens > 0:
|
||||
arguments["max_tokens"] = max_tokens
|
||||
|
||||
result = ApiHandler.submit_and_get_result(
|
||||
"openrouter/router/vision", arguments
|
||||
)
|
||||
return (result["output"],)
|
||||
except Exception as e:
|
||||
print(f"Error generating text with VLM: {str(e)}")
|
||||
return ("Error: Unable to generate text.",)
|
||||
finally:
|
||||
# Clean up the temporary file
|
||||
if 'temp_file_path' in locals():
|
||||
os.unlink(temp_file_path)
|
||||
return ApiHandler.handle_text_generation_error(model, e)
|
||||
|
||||
|
||||
# Node class mappings
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
@@ -112,4 +135,4 @@ NODE_CLASS_MAPPINGS = {
|
||||
# Node display name mappings
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"VLM_fal": "VLM (fal)",
|
||||
}
|
||||
}
|
||||
|
||||
+22
-3
@@ -1,9 +1,18 @@
|
||||
[project]
|
||||
name = "fal-api"
|
||||
description = "Custom nodes for using fal API. Video generation with Kling, Runway, Luma. Image generation with Flux. LLMs and VLMs OpenAI, Claude, Llama and Gemini."
|
||||
version = "1.0.1"
|
||||
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.5.1"
|
||||
license = {file = "LICENSE"}
|
||||
dependencies = ["fal-client", "torch"]
|
||||
requires-python = ">=3.9"
|
||||
dependencies = [
|
||||
"fal-client>=1.0,<2",
|
||||
"torch",
|
||||
"opencv-python",
|
||||
"numpy",
|
||||
"pillow",
|
||||
"requests",
|
||||
"av",
|
||||
]
|
||||
|
||||
[project.urls]
|
||||
Repository = "https://github.com/gokayfem/ComfyUI-fal-API"
|
||||
@@ -13,3 +22,13 @@ Repository = "https://github.com/gokayfem/ComfyUI-fal-API"
|
||||
PublisherId = "gokayfem"
|
||||
DisplayName = "ComfyUI-fal-API"
|
||||
Icon = ""
|
||||
|
||||
[tool.ruff]
|
||||
target-version = "py39"
|
||||
line-length = 120
|
||||
exclude = ["example_workflows"]
|
||||
|
||||
[tool.ruff.lint]
|
||||
select = ["E", "F", "W", "I", "B", "UP"]
|
||||
# E501: legacy long lines throughout the codebase; revisit once files are refactored.
|
||||
ignore = ["E501"]
|
||||
|
||||
+7
-2
@@ -1,2 +1,7 @@
|
||||
fal-client
|
||||
torch
|
||||
fal-client>=1.0,<2
|
||||
torch
|
||||
opencv-python
|
||||
numpy
|
||||
pillow
|
||||
requests
|
||||
av
|
||||
|
||||
@@ -0,0 +1,171 @@
|
||||
#!/usr/bin/env python3
|
||||
"""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 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
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[1]
|
||||
REGISTRY_PATH = REPO_ROOT / "data" / "fal_registry.json"
|
||||
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}"
|
||||
|
||||
|
||||
def load_registry(path: Path) -> dict[str, Any]:
|
||||
try:
|
||||
with open(path, encoding="utf-8") as handle:
|
||||
registry = json.load(handle)
|
||||
except (OSError, ValueError) as err:
|
||||
raise SystemExit(f"Failed to read registry at {path}: {err}") from err
|
||||
if not isinstance(registry.get("models"), list):
|
||||
raise SystemExit(f"Registry at {path} has no 'models' list")
|
||||
return registry
|
||||
|
||||
|
||||
def escape_cell(text: str) -> str:
|
||||
"""Make a value safe inside a markdown table cell."""
|
||||
return " ".join(str(text).split()).replace("|", "\\|")
|
||||
|
||||
|
||||
def group_by_category(
|
||||
models: list[dict[str, Any]],
|
||||
) -> list[tuple[str, list[dict[str, Any]]]]:
|
||||
"""Group models by category, categories sorted by size desc then name."""
|
||||
grouped: dict[str, list[dict[str, Any]]] = {}
|
||||
for model in models:
|
||||
category = str(model.get("category") or "other")
|
||||
grouped = {**grouped, category: [*grouped.get(category, []), model]}
|
||||
return sorted(grouped.items(), key=lambda item: (-len(item[1]), item[0]))
|
||||
|
||||
|
||||
def model_sort_key(model: dict[str, Any]) -> tuple[str, str]:
|
||||
title = str(model.get("title") or model.get("endpoint_id") or "")
|
||||
return (title.casefold(), str(model.get("endpoint_id") or ""))
|
||||
|
||||
|
||||
def render_model_row(model: dict[str, Any]) -> str:
|
||||
endpoint_id = str(model.get("endpoint_id") or "")
|
||||
title = escape_cell(model.get("title") or endpoint_id)
|
||||
lab = escape_cell(model.get("lab") or "—") or "—"
|
||||
output = escape_cell(model.get("output_kind") or "json")
|
||||
url = MODEL_URL_TEMPLATE.format(endpoint_id=endpoint_id)
|
||||
endpoint_cell = f"[`{escape_cell(endpoint_id)}`]({url})"
|
||||
return f"| {title} | {endpoint_cell} | {lab} | {output} |"
|
||||
|
||||
|
||||
def render_category(category: str, models: list[dict[str, Any]]) -> str:
|
||||
rows = [render_model_row(m) for m in sorted(models, key=model_sort_key)]
|
||||
count = len(models)
|
||||
noun = "model" if count == 1 else "models"
|
||||
return "\n".join(
|
||||
[
|
||||
"<details>",
|
||||
f"<summary><strong>{category}</strong> — {count} {noun}</summary>",
|
||||
"",
|
||||
"| Model | Endpoint | Lab | Output |",
|
||||
"| --- | --- | --- | --- |",
|
||||
*rows,
|
||||
"",
|
||||
"</details>",
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
def render_generated_section(registry: dict[str, Any]) -> str:
|
||||
models = registry["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"{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(live_models)
|
||||
]
|
||||
return "\n\n".join([summary, *blocks])
|
||||
|
||||
|
||||
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"MODELS.md must contain '{BEGIN_MARKER}' followed by '{END_MARKER}'"
|
||||
)
|
||||
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)
|
||||
document = read_models_document(MODELS_PATH)
|
||||
|
||||
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
|
||||
|
||||
MODELS_PATH.write_text(updated, encoding="utf-8")
|
||||
print(
|
||||
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
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
@@ -0,0 +1,769 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Build a compact registry of fal.ai model endpoints.
|
||||
|
||||
Distills the fal.ai model catalog plus per-endpoint OpenAPI schemas into a
|
||||
single registry JSON (``data/fal_registry.json``) that a node factory can use
|
||||
to auto-generate ComfyUI nodes.
|
||||
|
||||
Stdlib only. Usage:
|
||||
|
||||
python scripts/build_registry.py \
|
||||
--out data/fal_registry.json \
|
||||
--since-days 0 \
|
||||
--catalog-cache /path/to/fal_models_all.json \
|
||||
--schemas-cache /path/to/fal_schemas_recent.json
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
import time
|
||||
import urllib.error
|
||||
import urllib.request
|
||||
from collections import Counter
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
CATALOG_URL = "https://fal.ai/api/models?page={page}&total=100"
|
||||
SCHEMA_URL = "https://fal.ai/api/openapi/queue/openapi.json?endpoint_id={endpoint_id}"
|
||||
USER_AGENT = "ComfyUI-fal-API-registry-builder/1.0"
|
||||
|
||||
FETCH_ATTEMPTS = 3
|
||||
BACKOFF_BASE_SECONDS = 1.5
|
||||
MAX_DESCRIPTION_CHARS = 500
|
||||
MULTILINE_NAMES = frozenset({"prompt", "negative_prompt", "text", "script", "dialogue"})
|
||||
MULTILINE_DESCRIPTION_THRESHOLD = 120
|
||||
SKIPPED_PROPERTY_NAMES = frozenset({"sync_mode"})
|
||||
FILE_OUTPUT_PROPS = frozenset({"model_glb", "model_mesh", "model_url", "model_urls", "mesh"})
|
||||
|
||||
logger = logging.getLogger("build_registry")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fetching
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def fetch_json(url):
|
||||
"""Fetch a URL and parse JSON, with retries and backoff.
|
||||
|
||||
Returns the parsed document, or None for a 404 (skip-and-log).
|
||||
Raises on persistent non-404 failure.
|
||||
"""
|
||||
last_error = None
|
||||
for attempt in range(FETCH_ATTEMPTS):
|
||||
try:
|
||||
request = urllib.request.Request(url, headers={"User-Agent": USER_AGENT})
|
||||
with urllib.request.urlopen(request, timeout=60) as response:
|
||||
return json.loads(response.read().decode("utf-8"))
|
||||
except urllib.error.HTTPError as error:
|
||||
if error.code == 404:
|
||||
logger.warning("404 for %s, skipping", url)
|
||||
return None
|
||||
last_error = error
|
||||
except (urllib.error.URLError, TimeoutError, ValueError) as error:
|
||||
last_error = error
|
||||
time.sleep(BACKOFF_BASE_SECONDS * (2 ** attempt))
|
||||
raise RuntimeError(f"Failed to fetch {url} after {FETCH_ATTEMPTS} attempts: {last_error}")
|
||||
|
||||
|
||||
def extract_catalog_items(payload):
|
||||
"""Normalize a catalog API response page into a list of items."""
|
||||
if isinstance(payload, list):
|
||||
return payload
|
||||
if isinstance(payload, dict):
|
||||
for key in ("items", "models", "data", "results"):
|
||||
value = payload.get(key)
|
||||
if isinstance(value, list):
|
||||
return value
|
||||
return []
|
||||
|
||||
|
||||
def fetch_catalog():
|
||||
"""Fetch all catalog pages until an empty page is returned."""
|
||||
items = []
|
||||
page = 1
|
||||
while True:
|
||||
payload = fetch_json(CATALOG_URL.format(page=page))
|
||||
page_items = extract_catalog_items(payload)
|
||||
if not page_items:
|
||||
break
|
||||
items = items + page_items
|
||||
logger.info("Fetched catalog page %d (%d items)", page, len(page_items))
|
||||
page += 1
|
||||
return items
|
||||
|
||||
|
||||
def fetch_schemas(endpoint_ids, max_workers):
|
||||
"""Fetch OpenAPI docs for endpoint ids concurrently. Returns id -> doc."""
|
||||
schemas = {}
|
||||
with ThreadPoolExecutor(max_workers=max_workers) as executor:
|
||||
futures = {
|
||||
executor.submit(fetch_json, SCHEMA_URL.format(endpoint_id=endpoint_id)): endpoint_id
|
||||
for endpoint_id in endpoint_ids
|
||||
}
|
||||
for future in as_completed(futures):
|
||||
endpoint_id = futures[future]
|
||||
try:
|
||||
doc = future.result()
|
||||
except RuntimeError as error:
|
||||
logger.warning("Schema fetch failed for %s: %s", endpoint_id, error)
|
||||
continue
|
||||
if doc is not None:
|
||||
schemas = {**schemas, endpoint_id: doc}
|
||||
return schemas
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Catalog filtering
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def parse_published_at(item):
|
||||
"""Parse the model's publication timestamp, or None."""
|
||||
raw = item.get("publishedAt") or item.get("date") or ""
|
||||
if not raw:
|
||||
return None
|
||||
try:
|
||||
return datetime.fromisoformat(raw.replace("Z", "+00:00"))
|
||||
except ValueError:
|
||||
return None
|
||||
|
||||
|
||||
def filter_catalog(catalog, since):
|
||||
"""Keep live, public models (optionally within a publish window; since=None keeps all).
|
||||
|
||||
Returns (kept_items, skip_reason_counter).
|
||||
"""
|
||||
kept = []
|
||||
skipped = Counter()
|
||||
seen_ids = set()
|
||||
for item in catalog:
|
||||
endpoint_id = item.get("id") or ""
|
||||
if not endpoint_id or endpoint_id in seen_ids:
|
||||
skipped["duplicate_or_missing_id"] += 1
|
||||
continue
|
||||
seen_ids.add(endpoint_id)
|
||||
if item.get("status") != "public":
|
||||
skipped["not_public"] += 1
|
||||
continue
|
||||
if item.get("deprecated"):
|
||||
skipped["deprecated"] += 1
|
||||
continue
|
||||
if item.get("removed"):
|
||||
skipped["removed"] += 1
|
||||
continue
|
||||
if since is not None:
|
||||
published = parse_published_at(item)
|
||||
if published is None or published < since:
|
||||
skipped["outside_window"] += 1
|
||||
continue
|
||||
kept = kept + [item]
|
||||
return kept, skipped
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Schema resolution helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def resolve_ref(schema, components):
|
||||
"""Resolve a local $ref against components.schemas, one level."""
|
||||
ref = schema.get("$ref", "")
|
||||
if not ref.startswith("#/components/schemas/"):
|
||||
return schema
|
||||
name = ref.rsplit("/", 1)[-1].replace("~1", "/").replace("~0", "~")
|
||||
resolved = components.get(name)
|
||||
if not isinstance(resolved, dict):
|
||||
return schema
|
||||
siblings = {key: value for key, value in schema.items() if key != "$ref"}
|
||||
return {**resolved, **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)
|
||||
object_branch = next(
|
||||
(
|
||||
b
|
||||
for b in branches
|
||||
if b.get("type") == "object" or "properties" in b
|
||||
),
|
||||
None,
|
||||
)
|
||||
if enum_branch is None or object_branch is None:
|
||||
return None
|
||||
properties = object_branch.get("properties", {})
|
||||
if "width" in properties and "height" in properties:
|
||||
return enum_branch
|
||||
return None
|
||||
|
||||
|
||||
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:
|
||||
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, has_custom_size, custom_values
|
||||
|
||||
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
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Input distillation
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def detect_media_kind(name, schema, is_list):
|
||||
"""Heuristic media kind from a property name (string-typed props only)."""
|
||||
lowered = name.lower()
|
||||
description = str(schema.get("description", "")).lower()
|
||||
if "image_url" in lowered or "mask_url" in lowered:
|
||||
return "image"
|
||||
if lowered.endswith("_image"):
|
||||
return "image"
|
||||
if "video_url" in lowered:
|
||||
return "video"
|
||||
if "audio_url" in lowered or "voice_url" in lowered:
|
||||
return "audio"
|
||||
if "_url" in lowered or lowered == "url" or schema.get("format") == "uri":
|
||||
for kind in ("image", "video", "audio"):
|
||||
if kind in description:
|
||||
return kind
|
||||
return "file"
|
||||
del is_list # signature symmetry; list-ness does not change the kind
|
||||
return None
|
||||
|
||||
|
||||
def trim_text(value, limit=MAX_DESCRIPTION_CHARS):
|
||||
"""Trim a description/title string."""
|
||||
return str(value or "").strip()[:limit]
|
||||
|
||||
|
||||
def scalar_type_of(schema):
|
||||
"""Map an OpenAPI scalar type to a registry type."""
|
||||
type_name = schema.get("type")
|
||||
if schema.get("enum"):
|
||||
return "enum"
|
||||
if type_name in ("integer", "number", "boolean", "string"):
|
||||
return type_name
|
||||
if type_name == "object" or "properties" in schema:
|
||||
return "json"
|
||||
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("_"):
|
||||
return None
|
||||
|
||||
schema, has_custom_size, custom_enum = normalize_schema(raw_schema, components)
|
||||
|
||||
is_list = False
|
||||
if schema.get("type") == "array":
|
||||
is_list = True
|
||||
items, _, _ = normalize_schema(schema.get("items", {}), components)
|
||||
item_type = scalar_type_of(items)
|
||||
if item_type == "json":
|
||||
type_name = "json"
|
||||
is_list = False # rendered as a single JSON field
|
||||
else:
|
||||
type_name = item_type
|
||||
item_schema = items
|
||||
else:
|
||||
type_name = scalar_type_of(schema)
|
||||
item_schema = schema
|
||||
|
||||
enum_values = None
|
||||
if has_custom_size:
|
||||
type_name = "enum"
|
||||
enum_values = custom_enum
|
||||
elif type_name == "enum":
|
||||
enum_values = list(item_schema.get("enum", []))
|
||||
|
||||
minimum = schema.get("minimum", schema.get("exclusiveMinimum"))
|
||||
maximum = schema.get("maximum", schema.get("exclusiveMaximum"))
|
||||
if type_name not in ("integer", "number"):
|
||||
minimum = None
|
||||
maximum = None
|
||||
|
||||
default = schema.get("default", raw_schema.get("default") if isinstance(raw_schema, dict) else None)
|
||||
if type_name == "json" and default is not None and not isinstance(default, str):
|
||||
default = json.dumps(default, ensure_ascii=False, sort_keys=True)
|
||||
|
||||
# Some upstream schemas declare enum members and the default with mismatched
|
||||
# types (e.g. enum ["1","2","4","8"] with default 4). Normalize the default
|
||||
# onto the literal enum member it string-matches so widgets get a valid value.
|
||||
if enum_values and default is not None and default not in enum_values:
|
||||
match = next((v for v in enum_values if str(v) == str(default)), None)
|
||||
if match is not None:
|
||||
default = match
|
||||
|
||||
description = trim_text(schema.get("description") or schema.get("title"))
|
||||
|
||||
media_kind = None
|
||||
if type_name == "string" or (is_list and type_name == "string"):
|
||||
media_kind = detect_media_kind(name, schema, is_list)
|
||||
|
||||
multiline = name in MULTILINE_NAMES or (
|
||||
type_name == "string"
|
||||
and not enum_values
|
||||
and len(description) > MULTILINE_DESCRIPTION_THRESHOLD
|
||||
)
|
||||
|
||||
record = {
|
||||
"name": name,
|
||||
"type": type_name,
|
||||
"required": name in required_names,
|
||||
"default": default,
|
||||
"enum": enum_values,
|
||||
"min": minimum,
|
||||
"max": maximum,
|
||||
"description": description,
|
||||
"media_kind": media_kind,
|
||||
"is_list": is_list,
|
||||
"multiline": multiline,
|
||||
}
|
||||
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
|
||||
|
||||
|
||||
def ordered_property_names(schema):
|
||||
"""Property names, preferring fal's declared ordering."""
|
||||
properties = schema.get("properties", {})
|
||||
declared = schema.get("x-fal-order-properties")
|
||||
if isinstance(declared, list):
|
||||
ordered = [name for name in declared if name in properties]
|
||||
remainder = [name for name in properties if name not in ordered]
|
||||
return ordered + remainder
|
||||
return list(properties)
|
||||
|
||||
|
||||
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)
|
||||
|
||||
inputs = []
|
||||
for name in names:
|
||||
record = distill_property(name, properties.get(name, {}), required_names, components)
|
||||
if record is not None:
|
||||
inputs = inputs + [record]
|
||||
return inputs
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Schema selection
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
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].replace("~1", "/").replace("~0", "~") if ref.startswith("#/components/schemas/") else None
|
||||
|
||||
|
||||
def input_ref_from_paths(doc):
|
||||
"""Name of the schema referenced by the app POST requestBody."""
|
||||
for operations in doc.get("paths", {}).values():
|
||||
post = operations.get("post") if isinstance(operations, dict) else None
|
||||
if not isinstance(post, dict):
|
||||
continue
|
||||
content = post.get("requestBody", {}).get("content", {})
|
||||
schema = content.get("application/json", {}).get("schema", {})
|
||||
name = ref_name(schema)
|
||||
if name:
|
||||
return name
|
||||
return None
|
||||
|
||||
|
||||
def output_ref_from_paths(doc):
|
||||
"""Name of the schema referenced by result GET responses."""
|
||||
for operations in doc.get("paths", {}).values():
|
||||
get = operations.get("get") if isinstance(operations, dict) else None
|
||||
if not isinstance(get, dict):
|
||||
continue
|
||||
for response in get.get("responses", {}).values():
|
||||
content = response.get("content", {}) if isinstance(response, dict) else {}
|
||||
schema = content.get("application/json", {}).get("schema", {})
|
||||
name = ref_name(schema)
|
||||
if name and name.endswith("Output"):
|
||||
return name
|
||||
return None
|
||||
|
||||
|
||||
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 normalize_schema(components[referenced], components)[0]
|
||||
|
||||
candidates = [name for name in components if name.endswith(suffix)]
|
||||
if not candidates:
|
||||
return None
|
||||
normalized_endpoint = "".join(ch for ch in endpoint_id.lower() if ch.isalnum())
|
||||
matching = [
|
||||
name
|
||||
for name in candidates
|
||||
if "".join(ch for ch in name.lower() if ch.isalnum()).replace(suffix.lower(), "")
|
||||
in normalized_endpoint
|
||||
]
|
||||
pool = matching or candidates
|
||||
return normalize_schema(components[max(pool, key=len)], components)[0]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Output kind detection
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def detect_output(schema, components):
|
||||
"""Classify an Output schema. Returns (output_kind, output_props)."""
|
||||
if schema is None:
|
||||
return "json", []
|
||||
properties = schema.get("properties", {})
|
||||
prop_names = list(properties)
|
||||
lowered = {name.lower() for name in prop_names}
|
||||
|
||||
def prop_is_array(name):
|
||||
resolved, _, _ = normalize_schema(properties.get(name, {}), components)
|
||||
return resolved.get("type") == "array"
|
||||
|
||||
if "images" in lowered and prop_is_array("images"):
|
||||
return "images", prop_names
|
||||
if "image" in lowered:
|
||||
return "image", prop_names
|
||||
if "video" in lowered or "videos" in lowered:
|
||||
return "video", prop_names
|
||||
if "audio" in lowered or "audios" in lowered:
|
||||
return "audio", prop_names
|
||||
if lowered & FILE_OUTPUT_PROPS:
|
||||
return "file", prop_names
|
||||
|
||||
if prop_names:
|
||||
all_stringlike = True
|
||||
for name in prop_names:
|
||||
resolved, _, _ = normalize_schema(properties.get(name, {}), components)
|
||||
if scalar_type_of(resolved) != "string":
|
||||
all_stringlike = False
|
||||
break
|
||||
if all_stringlike:
|
||||
return "text", prop_names
|
||||
|
||||
return "json", prop_names
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Record assembly
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def build_record(item, doc):
|
||||
"""Build a single registry record from a catalog item + OpenAPI doc."""
|
||||
endpoint_id = item["id"]
|
||||
components = doc.get("components", {}).get("schemas", {})
|
||||
|
||||
input_schema = select_schema(doc, endpoint_id, "Input", input_ref_from_paths)
|
||||
if input_schema is None:
|
||||
logger.warning("%s: no Input schema found, skipping", endpoint_id)
|
||||
return None
|
||||
|
||||
output_schema = select_schema(doc, endpoint_id, "Output", output_ref_from_paths)
|
||||
output_kind, output_props = detect_output(output_schema, components)
|
||||
|
||||
published = parse_published_at(item)
|
||||
pricing = str(item.get("pricingInfoOverride") or "").replace("**", "").strip()
|
||||
|
||||
return {
|
||||
"endpoint_id": endpoint_id,
|
||||
"title": str(item.get("title") or "").strip(),
|
||||
"category": str(item.get("category") or "").strip(),
|
||||
"lab": str(item.get("modelLab") or "").strip(),
|
||||
"family": str(item.get("modelFamily") or "").strip(),
|
||||
"description": trim_text(item.get("shortDescription")),
|
||||
"pricing": pricing,
|
||||
"published_at": published.isoformat() if published else "",
|
||||
"thumbnail": str(item.get("thumbnailUrl") or "").strip(),
|
||||
"inputs": distill_inputs(input_schema, components, endpoint_id),
|
||||
"output_kind": output_kind,
|
||||
"output_props": output_props,
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Main
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def load_json_file(path):
|
||||
"""Load a JSON cache file."""
|
||||
try:
|
||||
with open(path, encoding="utf-8") as handle:
|
||||
return json.load(handle)
|
||||
except (OSError, ValueError) as error:
|
||||
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")
|
||||
parser.add_argument(
|
||||
"--since-days",
|
||||
type=int,
|
||||
default=0,
|
||||
help="Rolling publish window in days; 0 (default) = all live models",
|
||||
)
|
||||
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()
|
||||
|
||||
|
||||
def log_summary(records, skipped):
|
||||
"""Log counts by category / output kind and skip reasons."""
|
||||
category_counts = Counter(record["category"] for record in records)
|
||||
kind_counts = Counter(record["output_kind"] for record in records)
|
||||
logger.info("Models by category:")
|
||||
for category, count in category_counts.most_common():
|
||||
logger.info(" %-28s %d", category or "(none)", count)
|
||||
logger.info("Models by output_kind:")
|
||||
for kind, count in kind_counts.most_common():
|
||||
logger.info(" %-10s %d", kind, count)
|
||||
logger.info("Skipped: %s", dict(skipped) or "none")
|
||||
|
||||
|
||||
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
|
||||
|
||||
catalog = (
|
||||
load_json_file(args.catalog_cache) if args.catalog_cache else fetch_catalog()
|
||||
)
|
||||
logger.info("Catalog: %d items", len(catalog))
|
||||
|
||||
kept, skipped = filter_catalog(catalog, since)
|
||||
logger.info("After filtering: %d live public models in window", len(kept))
|
||||
|
||||
if args.schemas_cache:
|
||||
schemas = load_json_file(args.schemas_cache)
|
||||
else:
|
||||
schemas = fetch_schemas([item["id"] for item in kept], args.max_workers)
|
||||
logger.info("Schemas available: %d", len(schemas))
|
||||
|
||||
records = []
|
||||
for item in kept:
|
||||
doc = schemas.get(item["id"])
|
||||
if doc is None:
|
||||
skipped["no_schema"] += 1
|
||||
logger.warning("%s: no schema available, skipping", item["id"])
|
||||
continue
|
||||
record = build_record(item, doc)
|
||||
if record is None:
|
||||
skipped["no_input_schema"] += 1
|
||||
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
|
||||
# PR when the model set actually changes.
|
||||
registry = {
|
||||
"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,
|
||||
}
|
||||
|
||||
# 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,
|
||||
indent=None,
|
||||
separators=(",", ":"),
|
||||
sort_keys=True,
|
||||
ensure_ascii=False,
|
||||
)
|
||||
handle.write("\n")
|
||||
|
||||
os.replace(tmp_out, args.out)
|
||||
|
||||
log_summary(records, skipped)
|
||||
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__":
|
||||
main()
|
||||
@@ -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()
|
||||
@@ -0,0 +1,83 @@
|
||||
"""Shared fixtures: load the pack exactly like ComfyUI does (hyphenated dir)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
import importlib.util
|
||||
import os
|
||||
import sys
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
|
||||
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",
|
||||
str(Path(tempfile.mkdtemp(prefix="fal-api-test-cache-")) / "cache.db"),
|
||||
)
|
||||
|
||||
|
||||
def _load_package():
|
||||
if PKG in sys.modules:
|
||||
return sys.modules[PKG]
|
||||
spec = importlib.util.spec_from_file_location(
|
||||
PKG, ROOT / "__init__.py", submodule_search_locations=[str(ROOT)]
|
||||
)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
sys.modules[PKG] = module
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def pack():
|
||||
"""The fully loaded node pack (static + dynamic mappings)."""
|
||||
return _load_package()
|
||||
|
||||
|
||||
def _submodule(name: str):
|
||||
_load_package()
|
||||
return importlib.import_module(f"{PKG}.{name}")
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def schema_to_inputs():
|
||||
return _submodule("nodes.dynamic.schema_to_inputs")
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def arguments_mod():
|
||||
return _submodule("nodes.dynamic.arguments")
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def outputs_mod():
|
||||
return _submodule("nodes.dynamic.outputs")
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def factory_mod():
|
||||
return _submodule("nodes.dynamic.factory")
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def errors_mod():
|
||||
return _submodule("nodes.utils.errors")
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def media_mod():
|
||||
return _submodule("nodes.utils.media")
|
||||
Vendored
+370
@@ -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"
|
||||
]
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,38 @@
|
||||
"""Shared model/input fixture builders for the dynamic-node tests."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
def _model(inputs, **overrides):
|
||||
base = {
|
||||
"endpoint_id": "fal-ai/test/model",
|
||||
"title": "Test Model",
|
||||
"category": "text-to-image",
|
||||
"lab": "Test Lab",
|
||||
"family": "",
|
||||
"description": "",
|
||||
"pricing": "",
|
||||
"published_at": "2026-01-01T00:00:00Z",
|
||||
"thumbnail": "",
|
||||
"inputs": inputs,
|
||||
"output_kind": "images",
|
||||
"output_props": ["images"],
|
||||
}
|
||||
return {**base, **overrides}
|
||||
|
||||
|
||||
def _input(name, type_, **kw):
|
||||
base = {
|
||||
"name": name,
|
||||
"type": type_,
|
||||
"required": False,
|
||||
"default": None,
|
||||
"enum": None,
|
||||
"min": None,
|
||||
"max": None,
|
||||
"description": "",
|
||||
"media_kind": None,
|
||||
"is_list": False,
|
||||
"multiline": False,
|
||||
}
|
||||
return {**base, **kw}
|
||||
@@ -0,0 +1,92 @@
|
||||
{
|
||||
"description": "Node keys registered at v1.0.12 (commit 1b14ab3). These must NEVER be removed or renamed - existing user workflows reference them.",
|
||||
"keys": [
|
||||
"Bria_Video_Increase_Resolution_fal",
|
||||
"CombinedVideoGeneration_fal",
|
||||
"DYWanFun22_fal",
|
||||
"DYWanUpscaler_fal",
|
||||
"Dreamina31TextToImage_fal",
|
||||
"FluxDev_fal",
|
||||
"FluxGeneral_fal",
|
||||
"FluxLoraTrainer_fal",
|
||||
"FluxLora_fal",
|
||||
"FluxPro11_fal",
|
||||
"FluxPro1Fill_fal",
|
||||
"FluxProKontextMulti_fal",
|
||||
"FluxProKontextTextToImage_fal",
|
||||
"FluxProKontext_fal",
|
||||
"FluxPro_fal",
|
||||
"FluxSchnell_fal",
|
||||
"FluxUltra_fal",
|
||||
"GPTImage15Edit_fal",
|
||||
"GPTImage15_fal",
|
||||
"Hidreamfull_fal",
|
||||
"HunyuanVideoLoraTrainer_fal",
|
||||
"Ideogramv3_fal",
|
||||
"Imagen4Preview_fal",
|
||||
"InfinityStarTextToVideo_fal",
|
||||
"Kling21Pro_fal",
|
||||
"Kling25TurboPro_fal",
|
||||
"Kling26Pro_fal",
|
||||
"KlingMaster_fal",
|
||||
"KlingO3Pro_fal",
|
||||
"KlingO3Standard_fal",
|
||||
"KlingOmniImageToVideo_fal",
|
||||
"KlingOmniReferenceToVideo_fal",
|
||||
"KlingOmniVideoToVideoEdit_fal",
|
||||
"KlingOmniVideoToVideoReference_fal",
|
||||
"KlingPro10_fal",
|
||||
"KlingPro16_fal",
|
||||
"KlingV3ProMotionControl_fal",
|
||||
"KlingV3Pro_fal",
|
||||
"KlingV3StandardMotionControl_fal",
|
||||
"KlingV3Standard_fal",
|
||||
"Kling_fal",
|
||||
"Krea_Wan14b_VideoToVideo_fal",
|
||||
"LLM_fal",
|
||||
"LoadVideoURL",
|
||||
"LtxVideoTrainer_fal",
|
||||
"LumaDreamMachine_fal",
|
||||
"MiniMaxSubjectReference_fal",
|
||||
"MiniMaxTextToVideo_fal",
|
||||
"MiniMax_fal",
|
||||
"NanoBanana2_fal",
|
||||
"NanoBananaEdit_fal",
|
||||
"NanoBananaPro_fal",
|
||||
"NanoBananaTextToImage_fal",
|
||||
"PixverseSwapNode_fal",
|
||||
"QwenImageEditPlusLoRA_fal",
|
||||
"QwenImageEdit_fal",
|
||||
"Recraft_fal",
|
||||
"ReveTextToImage_fal",
|
||||
"RunwayGen3_fal",
|
||||
"Sana_fal",
|
||||
"SeedEditV3_fal",
|
||||
"SeedanceImageToVideo_fal",
|
||||
"SeedanceProImageToVideo_fal",
|
||||
"SeedanceTextToVideo_fal",
|
||||
"SeedreamV4Edit_fal",
|
||||
"Seedvr_Upscale_Video_fal",
|
||||
"Seedvr_Upscaler_fal",
|
||||
"Sora2Pro_fal",
|
||||
"Topaz_Upscale_Video_fal",
|
||||
"UploadFile_fal",
|
||||
"UploadVideo_fal",
|
||||
"Upscaler_fal",
|
||||
"VLM_fal",
|
||||
"Veo2ImageToVideo_fal",
|
||||
"Veo31Fast_fal",
|
||||
"Veo31_fal",
|
||||
"Veo3_fal",
|
||||
"VideoUpscaler_fal",
|
||||
"Wan2214b_animate_move_character_fal",
|
||||
"Wan2214b_animate_replace_character_fal",
|
||||
"Wan22VACEFun14b_fal",
|
||||
"Wan25_preview_fal",
|
||||
"Wan26ReferenceToVideo_fal",
|
||||
"Wan26_fal",
|
||||
"WanLoraTrainer_fal",
|
||||
"WanPro_fal",
|
||||
"WanVACEVideoEdit_fal"
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,6 @@
|
||||
# Anchors pytest's rootdir here so the ComfyUI pack's root __init__.py
|
||||
# (which makes the repo root look like a package) is never collected/imported
|
||||
# by pytest itself — the pack is loaded properly via conftest.py instead.
|
||||
[pytest]
|
||||
addopts = --import-mode=importlib
|
||||
pythonpath = .
|
||||
@@ -0,0 +1,133 @@
|
||||
"""Unit tests for kwargs→API-arguments translation (uploads stubbed)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
from helpers import _input, _model
|
||||
|
||||
|
||||
class _FakeImageUtils:
|
||||
@staticmethod
|
||||
def upload_image(_value):
|
||||
return "https://fal.media/img.png"
|
||||
|
||||
@staticmethod
|
||||
def prepare_images(_value):
|
||||
return ["https://fal.media/img1.png", "https://fal.media/img2.png"]
|
||||
|
||||
|
||||
class _FakeMediaUtils:
|
||||
@staticmethod
|
||||
def upload_video(_value):
|
||||
return "https://fal.media/vid.mp4"
|
||||
|
||||
@staticmethod
|
||||
def upload_audio(_value):
|
||||
return "https://fal.media/aud.wav"
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _stub_uploads(monkeypatch, arguments_mod):
|
||||
monkeypatch.setattr(arguments_mod, "ImageUtils", _FakeImageUtils)
|
||||
monkeypatch.setattr(arguments_mod, "MediaUtils", _FakeMediaUtils)
|
||||
|
||||
|
||||
def test_seed_minus_one_omitted(arguments_mod):
|
||||
model = _model([_input("seed", "integer")])
|
||||
args = arguments_mod.build_arguments(model, {"seed": -1})
|
||||
assert "seed" not in args
|
||||
|
||||
|
||||
def test_seed_value_sent(arguments_mod):
|
||||
model = _model([_input("seed", "integer")])
|
||||
args = arguments_mod.build_arguments(model, {"seed": 42})
|
||||
assert args["seed"] == 42
|
||||
|
||||
|
||||
def test_custom_size_expands_to_object(arguments_mod):
|
||||
model = _model([
|
||||
_input("image_size", "enum", enum=["square", "custom_size"],
|
||||
default="square", has_custom_size=True)
|
||||
])
|
||||
args = arguments_mod.build_arguments(
|
||||
model, {"image_size": "custom_size", "width": 832, "height": 1216}
|
||||
)
|
||||
assert args["image_size"] == {"width": 832, "height": 1216}
|
||||
|
||||
|
||||
def test_preset_size_passes_through(arguments_mod):
|
||||
model = _model([
|
||||
_input("image_size", "enum", enum=["square", "custom_size"],
|
||||
default="square", has_custom_size=True)
|
||||
])
|
||||
args = arguments_mod.build_arguments(
|
||||
model, {"image_size": "square", "width": 832, "height": 1216}
|
||||
)
|
||||
assert args["image_size"] == "square"
|
||||
assert "width" not in args and "height" not in args
|
||||
|
||||
|
||||
def test_image_upload_single_and_list(arguments_mod):
|
||||
model = _model([
|
||||
_input("image_url", "string", media_kind="image"),
|
||||
_input("image_urls", "array", media_kind="image", is_list=True),
|
||||
])
|
||||
args = arguments_mod.build_arguments(
|
||||
model, {"image_url": object(), "image_urls": object()}
|
||||
)
|
||||
assert args["image_url"] == "https://fal.media/img.png"
|
||||
assert args["image_urls"] == [
|
||||
"https://fal.media/img1.png",
|
||||
"https://fal.media/img2.png",
|
||||
]
|
||||
|
||||
|
||||
def test_video_and_audio_upload(arguments_mod):
|
||||
model = _model([
|
||||
_input("video_url", "string", media_kind="video"),
|
||||
_input("audio_url", "string", media_kind="audio"),
|
||||
])
|
||||
args = arguments_mod.build_arguments(
|
||||
model, {"video_url": object(), "audio_url": object()}
|
||||
)
|
||||
assert args["video_url"] == "https://fal.media/vid.mp4"
|
||||
assert args["audio_url"] == "https://fal.media/aud.wav"
|
||||
|
||||
|
||||
def test_invalid_json_raises_fal_error(arguments_mod, errors_mod):
|
||||
model = _model([_input("loras", "json")])
|
||||
with pytest.raises(errors_mod.FalApiError):
|
||||
arguments_mod.build_arguments(model, {"loras": "{not json"})
|
||||
|
||||
|
||||
def test_valid_json_parsed(arguments_mod):
|
||||
model = _model([_input("loras", "json")])
|
||||
args = arguments_mod.build_arguments(model, {"loras": '[{"path": "x"}]'})
|
||||
assert args["loras"] == [{"path": "x"}]
|
||||
|
||||
|
||||
def test_empty_optional_string_skipped(arguments_mod):
|
||||
model = _model([_input("negative_prompt", "string")])
|
||||
args = arguments_mod.build_arguments(model, {"negative_prompt": ""})
|
||||
assert "negative_prompt" not in args
|
||||
|
||||
|
||||
def test_multi_enum_split_and_validated(arguments_mod, errors_mod):
|
||||
model = _model([
|
||||
_input("stems", "enum", enum=["vocals", "drums", "bass"], is_list=True)
|
||||
])
|
||||
args = arguments_mod.build_arguments(model, {"stems": "vocals, bass"})
|
||||
assert args["stems"] == ["vocals", "bass"]
|
||||
|
||||
assert "stems" not in arguments_mod.build_arguments(model, {"stems": " "})
|
||||
|
||||
with pytest.raises(errors_mod.FalApiError):
|
||||
arguments_mod.build_arguments(model, {"stems": "vocals, kazoo"})
|
||||
|
||||
|
||||
def test_kwargs_not_mutated(arguments_mod):
|
||||
model = _model([_input("seed", "integer"), _input("prompt", "string", required=True)])
|
||||
kwargs = {"seed": -1, "prompt": "hi"}
|
||||
snapshot = dict(kwargs)
|
||||
arguments_mod.build_arguments(model, kwargs)
|
||||
assert kwargs == snapshot
|
||||
@@ -0,0 +1,41 @@
|
||||
"""Unit tests for the spend guard and balance node."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
|
||||
import pytest
|
||||
from conftest import PKG, _load_package
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def billing_mod():
|
||||
_load_package()
|
||||
return importlib.import_module(f"{PKG}.nodes.utils.billing")
|
||||
|
||||
|
||||
def test_preflight_noop_when_unconfigured(billing_mod):
|
||||
# no [spend_guard] section in config → both checks disabled
|
||||
billing_mod.SpendGuard.preflight("fal-ai/anything")
|
||||
|
||||
|
||||
def test_balance_node_never_raises(pack, monkeypatch, billing_mod):
|
||||
monkeypatch.setattr(
|
||||
billing_mod.BillingUtils, "get_balance", staticmethod(lambda force=False: None)
|
||||
)
|
||||
cls = pack.NODE_CLASS_MAPPINGS["FalBalance_fal"]
|
||||
node = cls()
|
||||
out = getattr(node, cls.FUNCTION)(force_refresh=False)
|
||||
report, balance = out[0], out[1]
|
||||
assert isinstance(report, str) and report
|
||||
assert balance == -1.0
|
||||
|
||||
|
||||
def test_balance_node_reports_value(pack, monkeypatch, billing_mod):
|
||||
monkeypatch.setattr(
|
||||
billing_mod.BillingUtils, "get_balance", staticmethod(lambda force=False: 24.5)
|
||||
)
|
||||
cls = pack.NODE_CLASS_MAPPINGS["FalBalance_fal"]
|
||||
node = cls()
|
||||
out = getattr(node, cls.FUNCTION)(force_refresh=True)
|
||||
assert out[1] == pytest.approx(24.5)
|
||||
@@ -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"]
|
||||
@@ -0,0 +1,79 @@
|
||||
"""Unit tests for fal error extraction and FalApiError formatting."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
class _FakeResponse:
|
||||
def __init__(self, payload):
|
||||
self._payload = payload
|
||||
|
||||
def json(self):
|
||||
if isinstance(self._payload, Exception):
|
||||
raise self._payload
|
||||
return self._payload
|
||||
|
||||
|
||||
class _FakeHTTPError(Exception):
|
||||
"""Duck-typed stand-in for fal_client.FalClientHTTPError."""
|
||||
|
||||
def __init__(self, message, status_code, payload):
|
||||
super().__init__(message)
|
||||
self.status_code = status_code
|
||||
self.response = _FakeResponse(payload)
|
||||
|
||||
|
||||
def test_error_message_includes_model_and_status(errors_mod):
|
||||
err = errors_mod.FalApiError("fal-ai/flux/dev", "boom", 422)
|
||||
assert "fal-ai/flux/dev" in str(err)
|
||||
assert "boom" in str(err)
|
||||
assert "422" in str(err)
|
||||
|
||||
|
||||
def test_extract_string_detail(errors_mod):
|
||||
exc = _FakeHTTPError("HTTP 403", 403, {"detail": "Content policy violation"})
|
||||
message, status = errors_mod.extract_error_message(exc)
|
||||
assert message == "Content policy violation"
|
||||
assert status == 403
|
||||
|
||||
|
||||
def test_extract_validation_list(errors_mod):
|
||||
exc = _FakeHTTPError(
|
||||
"HTTP 422", 422,
|
||||
{"detail": [
|
||||
{"loc": ["body", "prompt"], "msg": "field required"},
|
||||
{"loc": ["body", "seed"], "msg": "not an int"},
|
||||
]},
|
||||
)
|
||||
message, status = errors_mod.extract_error_message(exc)
|
||||
assert "prompt: field required" in message
|
||||
assert "seed: not an int" in message
|
||||
assert status == 422
|
||||
|
||||
|
||||
def test_extract_falls_back_to_str(errors_mod):
|
||||
message, status = errors_mod.extract_error_message(RuntimeError("plain failure"))
|
||||
assert message == "plain failure"
|
||||
assert status is None
|
||||
|
||||
|
||||
def test_extract_survives_bad_response_json(errors_mod):
|
||||
exc = _FakeHTTPError("HTTP 500", 500, ValueError("not json"))
|
||||
message, status = errors_mod.extract_error_message(exc)
|
||||
assert message # falls back to str(exc)
|
||||
assert status == 500
|
||||
|
||||
|
||||
def test_raise_fal_error_chains(errors_mod):
|
||||
original = RuntimeError("root cause")
|
||||
with pytest.raises(errors_mod.FalApiError) as excinfo:
|
||||
errors_mod.raise_fal_error("some-model", original)
|
||||
assert excinfo.value.__cause__ is original
|
||||
|
||||
|
||||
def test_raise_fal_error_passthrough(errors_mod):
|
||||
already = errors_mod.FalApiError("m", "msg")
|
||||
with pytest.raises(errors_mod.FalApiError) as excinfo:
|
||||
errors_mod.raise_fal_error("other", already)
|
||||
assert excinfo.value is already
|
||||
@@ -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)]
|
||||
@@ -0,0 +1,66 @@
|
||||
"""Unit tests for the durable job inbox store."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
|
||||
import pytest
|
||||
from conftest import PKG, _load_package
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def store():
|
||||
_load_package()
|
||||
mod = importlib.import_module(f"{PKG}.nodes.utils.job_store")
|
||||
instance = mod.JobStore()
|
||||
instance.prune(older_than_days=0) # clear anything from other tests
|
||||
yield instance
|
||||
instance.prune(older_than_days=0)
|
||||
|
||||
|
||||
def test_submit_and_collect_lifecycle(store):
|
||||
store.record_submit("fal-ai/kling-video/v3/pro/image-to-video", "req-a")
|
||||
store.record_submit("fal-ai/flux-2", "req-b")
|
||||
assert store.counts()["submitted"] == 2
|
||||
|
||||
store.mark_collected("req-a")
|
||||
counts = store.counts()
|
||||
assert counts["submitted"] == 1
|
||||
assert counts["collected"] == 1
|
||||
|
||||
pending = store.pending()
|
||||
assert len(pending) == 1
|
||||
assert pending[0]["request_id"] == "req-b"
|
||||
|
||||
|
||||
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()
|
||||
assert entries[0]["request_id"] == "req-2"
|
||||
|
||||
|
||||
def test_mark_collected_unknown_id_is_silent(store):
|
||||
store.mark_collected("req-from-another-session")
|
||||
entries = store.entries(status="collected")
|
||||
assert any(e["request_id"] == "req-from-another-session" for e in entries)
|
||||
|
||||
|
||||
def test_report_mentions_pending(store):
|
||||
store.record_submit("fal-ai/kling-video/v3/pro/image-to-video", "req-x")
|
||||
report = store.report()
|
||||
assert "req-x" in report
|
||||
assert "pending" in report.lower()
|
||||
|
||||
|
||||
def test_inbox_node_outputs(pack, store):
|
||||
store.record_submit("fal-ai/veo3", "req-latest")
|
||||
cls = pack.NODE_CLASS_MAPPINGS["FalJobInbox_fal"]
|
||||
node = cls()
|
||||
out = getattr(node, cls.FUNCTION)(status_filter="all", limit=20)
|
||||
report, latest_id, latest_endpoint = out[0], out[1], out[2]
|
||||
assert isinstance(report, str)
|
||||
assert latest_id == "req-latest"
|
||||
assert latest_endpoint == "fal-ai/veo3"
|
||||
@@ -0,0 +1,56 @@
|
||||
"""Unit tests for the session cost ledger."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
import threading
|
||||
|
||||
import pytest
|
||||
from conftest import PKG, _load_package
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def ledger(pricing_mod=None):
|
||||
_load_package()
|
||||
mod = importlib.import_module(f"{PKG}.nodes.utils.ledger")
|
||||
instance = mod.SessionLedger()
|
||||
instance.reset()
|
||||
yield instance
|
||||
instance.reset()
|
||||
|
||||
|
||||
def test_record_and_totals(ledger):
|
||||
ledger.record("fal-ai/a", "req-1", 2.5, 0.10)
|
||||
ledger.record("fal-ai/b", "req-2", 1.0, None)
|
||||
entries = ledger.entries()
|
||||
assert len(entries) == 2
|
||||
assert ledger.total_cost() == pytest.approx(0.10)
|
||||
assert ledger.unknown_cost_count() == 1
|
||||
|
||||
|
||||
def test_report_mentions_calls(ledger):
|
||||
ledger.record("fal-ai/kling-video/v3/pro/image-to-video", "abc123", 12.4, 0.35)
|
||||
report = ledger.report()
|
||||
assert "kling" in report
|
||||
assert "abc123" in report
|
||||
|
||||
|
||||
def test_reset(ledger):
|
||||
ledger.record("fal-ai/a", None, 1.0, 0.5)
|
||||
ledger.reset()
|
||||
assert ledger.entries() == []
|
||||
assert ledger.total_cost() == 0.0
|
||||
|
||||
|
||||
def test_thread_safety(ledger):
|
||||
def worker():
|
||||
for _ in range(100):
|
||||
ledger.record("fal-ai/t", None, 0.1, 0.01)
|
||||
|
||||
threads = [threading.Thread(target=worker) for _ in range(10)]
|
||||
for t in threads:
|
||||
t.start()
|
||||
for t in threads:
|
||||
t.join()
|
||||
assert len(ledger.entries()) == 1000
|
||||
assert ledger.total_cost() == pytest.approx(10.0)
|
||||
@@ -0,0 +1,45 @@
|
||||
"""Integration: the loaded pack must keep every legacy key and stay coherent."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
LEGACY_SNAPSHOT = Path(__file__).with_name("legacy_node_keys.json")
|
||||
|
||||
|
||||
def test_all_legacy_keys_present(pack):
|
||||
"""Backward-compat lock: keys registered at v1.0.12 must never disappear."""
|
||||
legacy = set(json.loads(LEGACY_SNAPSHOT.read_text())["keys"])
|
||||
current = set(pack.NODE_CLASS_MAPPINGS)
|
||||
missing = legacy - current
|
||||
assert not missing, f"legacy node keys removed (breaks user workflows): {sorted(missing)}"
|
||||
|
||||
|
||||
def test_display_names_complete(pack):
|
||||
missing = [k for k in pack.NODE_CLASS_MAPPINGS if k not in pack.NODE_DISPLAY_NAME_MAPPINGS]
|
||||
assert not missing
|
||||
|
||||
|
||||
def test_dynamic_nodes_registered(pack):
|
||||
dynamic = [k for k in pack.NODE_CLASS_MAPPINGS if k.startswith("FalAPI_")]
|
||||
assert len(dynamic) > 500, "dynamic registry failed to load"
|
||||
assert "FalAnyEndpoint_fal" in pack.NODE_CLASS_MAPPINGS
|
||||
|
||||
|
||||
def test_every_node_class_is_valid(pack):
|
||||
for key, cls in pack.NODE_CLASS_MAPPINGS.items():
|
||||
input_types = cls.INPUT_TYPES()
|
||||
assert isinstance(input_types, dict), key
|
||||
assert "required" in input_types or "optional" in input_types, key
|
||||
assert isinstance(cls.RETURN_TYPES, tuple), key
|
||||
assert isinstance(cls.FUNCTION, str) and hasattr(cls, cls.FUNCTION), key
|
||||
assert isinstance(cls.CATEGORY, str) and cls.CATEGORY, key
|
||||
|
||||
|
||||
def test_no_bare_video_category_left(pack):
|
||||
bare = [
|
||||
k for k, cls in pack.NODE_CLASS_MAPPINGS.items()
|
||||
if cls.CATEGORY.lower() == "video"
|
||||
]
|
||||
assert not bare, f"nodes escaped the FAL/ menu namespace: {bare}"
|
||||
@@ -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"],
|
||||
)
|
||||
@@ -0,0 +1,72 @@
|
||||
"""Unit tests for result→ComfyUI-outputs mapping (media decode stubbed)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
|
||||
import pytest
|
||||
from helpers import _model
|
||||
|
||||
_VIDEO_SENTINEL = object()
|
||||
_AUDIO_SENTINEL = {"waveform": "stub", "sample_rate": 44100}
|
||||
|
||||
|
||||
class _FakeMediaUtils:
|
||||
@staticmethod
|
||||
def video_from_url(_url):
|
||||
return _VIDEO_SENTINEL
|
||||
|
||||
@staticmethod
|
||||
def audio_from_url(_url):
|
||||
return _AUDIO_SENTINEL
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _stub_media(monkeypatch, outputs_mod):
|
||||
monkeypatch.setattr(outputs_mod, "MediaUtils", _FakeMediaUtils)
|
||||
|
||||
|
||||
def test_return_specs_cover_all_kinds(outputs_mod):
|
||||
assert set(outputs_mod.RETURN_SPECS) >= {
|
||||
"images", "image", "video", "audio", "text", "file", "json",
|
||||
}
|
||||
for types, names in outputs_mod.RETURN_SPECS.values():
|
||||
assert len(types) == len(names)
|
||||
|
||||
|
||||
def test_video_result(outputs_mod):
|
||||
model = _model([], output_kind="video", output_props=["video"])
|
||||
result = {"video": {"url": "https://fal.media/v.mp4"}}
|
||||
out = outputs_mod.process_result(model, result)
|
||||
assert out == (_VIDEO_SENTINEL, "https://fal.media/v.mp4")
|
||||
|
||||
|
||||
def test_audio_result(outputs_mod):
|
||||
model = _model([], output_kind="audio", output_props=["audio"])
|
||||
result = {"audio": {"url": "https://fal.media/a.mp3"}}
|
||||
out = outputs_mod.process_result(model, result)
|
||||
assert out == (_AUDIO_SENTINEL, "https://fal.media/a.mp3")
|
||||
|
||||
|
||||
def test_text_result(outputs_mod):
|
||||
model = _model([], output_kind="text", output_props=["text"])
|
||||
assert outputs_mod.process_result(model, {"text": "hello"}) == ("hello",)
|
||||
|
||||
|
||||
def test_file_result_digs_url(outputs_mod):
|
||||
model = _model([], output_kind="file", output_props=["model_glb"])
|
||||
result = {"model_glb": {"url": "https://fal.media/m.glb"}}
|
||||
assert outputs_mod.process_result(model, result) == ("https://fal.media/m.glb",)
|
||||
|
||||
|
||||
def test_json_fallback(outputs_mod):
|
||||
model = _model([], output_kind="json", output_props=[])
|
||||
result = {"anything": [1, 2, 3]}
|
||||
(payload,) = outputs_mod.process_result(model, result)
|
||||
assert json.loads(payload) == result
|
||||
|
||||
|
||||
def test_find_url_recursive(outputs_mod):
|
||||
nested = {"a": [{"b": {"url": "https://x/y.bin"}}]}
|
||||
assert outputs_mod.find_url(nested) == "https://x/y.bin"
|
||||
assert outputs_mod.find_url({"no": "url here"}) is None
|
||||
@@ -0,0 +1,76 @@
|
||||
"""Integration tests for the FAL/Platform utility nodes."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
PLATFORM_KEYS = [
|
||||
"FalSubmit_fal",
|
||||
"FalCollect_fal",
|
||||
"FalResultByRequestId_fal",
|
||||
"FalCostEstimator_fal",
|
||||
"FalSessionCosts_fal",
|
||||
"FalSaveMediaURL_fal",
|
||||
]
|
||||
|
||||
|
||||
def test_all_platform_nodes_registered(pack):
|
||||
for key in PLATFORM_KEYS:
|
||||
assert key in pack.NODE_CLASS_MAPPINGS, key
|
||||
assert key in pack.NODE_DISPLAY_NAME_MAPPINGS, key
|
||||
cls = pack.NODE_CLASS_MAPPINGS[key]
|
||||
assert cls.CATEGORY == "FAL/Platform", key
|
||||
input_types = cls.INPUT_TYPES()
|
||||
assert "required" in input_types or "optional" in input_types
|
||||
|
||||
|
||||
def test_submit_returns_handle_type(pack):
|
||||
cls = pack.NODE_CLASS_MAPPINGS["FalSubmit_fal"]
|
||||
assert "FAL_HANDLE" in cls.RETURN_TYPES
|
||||
|
||||
|
||||
def test_collect_rejects_bad_handle(pack, errors_mod):
|
||||
cls = pack.NODE_CLASS_MAPPINGS["FalCollect_fal"]
|
||||
node = cls()
|
||||
fn = getattr(node, cls.FUNCTION)
|
||||
with pytest.raises(errors_mod.FalApiError):
|
||||
fn(handle="not a handle")
|
||||
|
||||
|
||||
def test_cost_estimator_never_raises(pack):
|
||||
cls = pack.NODE_CLASS_MAPPINGS["FalCostEstimator_fal"]
|
||||
node = cls()
|
||||
fn = getattr(node, cls.FUNCTION)
|
||||
report, total = fn(endpoint_id="fal-ai/definitely-not-real", runs=5)
|
||||
assert isinstance(report, str)
|
||||
assert isinstance(total, float)
|
||||
|
||||
|
||||
def test_session_costs_reports(pack):
|
||||
cls = pack.NODE_CLASS_MAPPINGS["FalSessionCosts_fal"]
|
||||
node = cls()
|
||||
fn = getattr(node, cls.FUNCTION)
|
||||
out = fn(reset=False)
|
||||
report, total = out[0], out[1]
|
||||
assert isinstance(report, str)
|
||||
assert isinstance(total, float)
|
||||
|
||||
|
||||
def test_save_media_blocks_path_traversal(pack, errors_mod):
|
||||
import importlib
|
||||
|
||||
from conftest import PKG
|
||||
|
||||
platform = importlib.import_module(f"{PKG}.nodes.platform_node")
|
||||
with pytest.raises(errors_mod.FalApiError):
|
||||
platform._resolve_save_directory("../../../../tmp/evil")
|
||||
directory, basename = platform._resolve_save_directory("fal/media")
|
||||
assert basename == "media"
|
||||
|
||||
|
||||
def test_save_media_rejects_empty_url(pack, errors_mod):
|
||||
cls = pack.NODE_CLASS_MAPPINGS["FalSaveMediaURL_fal"]
|
||||
node = cls()
|
||||
fn = getattr(node, cls.FUNCTION)
|
||||
with pytest.raises((errors_mod.FalApiError, ValueError)):
|
||||
fn(url="", filename_prefix="fal/test")
|
||||
@@ -0,0 +1,56 @@
|
||||
"""Unit tests for the registry pricing parser."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
|
||||
import pytest
|
||||
from conftest import PKG, _load_package
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def pricing_mod():
|
||||
_load_package()
|
||||
return importlib.import_module(f"{PKG}.nodes.utils.pricing")
|
||||
|
||||
|
||||
def test_run_ratio_wins_over_per_unit(pricing_mod):
|
||||
text = (
|
||||
"Your request will cost $0.08 per image. For $1.00, you can run "
|
||||
"this model 12 times."
|
||||
)
|
||||
parsed = pricing_mod.PricingUtils.parse(text)
|
||||
assert parsed["per_run"] == pytest.approx(1.0 / 12.0)
|
||||
|
||||
|
||||
def test_per_unit_only(pricing_mod):
|
||||
parsed = pricing_mod.PricingUtils.parse(
|
||||
"Your request will cost $0.05 per second of video."
|
||||
)
|
||||
assert parsed["per_run"] is None
|
||||
assert parsed["per_unit"] == pytest.approx(0.05)
|
||||
assert "second" in parsed["unit"]
|
||||
|
||||
|
||||
def test_per_image_implies_per_run(pricing_mod):
|
||||
parsed = pricing_mod.PricingUtils.parse("Your request will cost $0.04 per image.")
|
||||
assert parsed["per_run"] == pytest.approx(0.04)
|
||||
|
||||
|
||||
def test_junk_never_raises(pricing_mod):
|
||||
for junk in ("", "free during preview!!", "$", "per per per", None or ""):
|
||||
parsed = pricing_mod.PricingUtils.parse(junk)
|
||||
assert parsed["raw"] == junk
|
||||
|
||||
|
||||
def test_estimate_against_real_registry(pricing_mod):
|
||||
est = pricing_mod.PricingUtils.estimate("fal-ai/nano-banana-2/edit", 10)
|
||||
assert est["runs"] == 10
|
||||
report = pricing_mod.PricingUtils.format_report(est)
|
||||
assert "fal-ai/nano-banana-2/edit" in report
|
||||
|
||||
|
||||
def test_unknown_endpoint_safe(pricing_mod):
|
||||
est = pricing_mod.PricingUtils.estimate("fal-ai/does-not-exist", 3)
|
||||
report = pricing_mod.PricingUtils.format_report(est)
|
||||
assert isinstance(report, str) and report
|
||||
@@ -0,0 +1,85 @@
|
||||
"""Unit tests for provenance sidecars and reproduce-from-file."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
import json
|
||||
|
||||
import pytest
|
||||
from conftest import PKG, _load_package
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def platform():
|
||||
_load_package()
|
||||
return importlib.import_module(f"{PKG}.nodes.platform_node")
|
||||
|
||||
|
||||
def test_find_request_by_url_escapes_wildcards(platform):
|
||||
cache_mod = importlib.import_module(f"{PKG}.nodes.utils.result_cache")
|
||||
cache = cache_mod.ResultCache()
|
||||
cache.clear()
|
||||
url = "https://fal.media/files/a%20b/out_1.png"
|
||||
cache.put(
|
||||
"fal-ai/flux-2",
|
||||
{"prompt": "x"},
|
||||
{"images": [{"url": url}]},
|
||||
"req-prov",
|
||||
)
|
||||
hit = cache.find_request_by_url(url)
|
||||
assert hit == {"endpoint_id": "fal-ai/flux-2", "request_id": "req-prov"}
|
||||
# a percent sign must not act as a wildcard
|
||||
assert cache.find_request_by_url("https://fal.media/files/aXb/out_1.png") is None
|
||||
cache.clear()
|
||||
|
||||
|
||||
def test_provenance_from_sidecar(platform, tmp_path):
|
||||
saved = tmp_path / "out_00001.mp4"
|
||||
saved.write_bytes(b"fake video")
|
||||
sidecar = tmp_path / "out_00001.mp4.fal.json"
|
||||
sidecar.write_text(json.dumps({
|
||||
"version": 1,
|
||||
"endpoint_id": "fal-ai/veo3",
|
||||
"request_id": "req-42",
|
||||
"source_url": "https://fal.media/v.mp4",
|
||||
"saved_at": 0,
|
||||
}))
|
||||
node = platform.FalProvenanceFromFile()
|
||||
endpoint, request_id, blob = node.read(file_path=str(saved))
|
||||
assert endpoint == "fal-ai/veo3"
|
||||
assert request_id == "req-42"
|
||||
assert json.loads(blob)["source_url"] == "https://fal.media/v.mp4"
|
||||
|
||||
|
||||
def test_provenance_missing_raises(platform, tmp_path, errors_mod):
|
||||
bare = tmp_path / "no_provenance.bin"
|
||||
bare.write_bytes(b"data")
|
||||
node = platform.FalProvenanceFromFile()
|
||||
with pytest.raises(errors_mod.FalApiError):
|
||||
node.read(file_path=str(bare))
|
||||
|
||||
|
||||
def test_png_chunk_roundtrip(platform, tmp_path):
|
||||
from PIL import Image
|
||||
|
||||
png = tmp_path / "img.png"
|
||||
Image.new("RGB", (4, 4), "red").save(png)
|
||||
payload = {"version": 1, "endpoint_id": "fal-ai/flux-2", "request_id": "req-png"}
|
||||
platform._embed_png_provenance(str(png), payload)
|
||||
node = platform.FalProvenanceFromFile()
|
||||
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()
|
||||
@@ -0,0 +1,99 @@
|
||||
"""Validate the committed model registry — pure JSON, no heavy imports."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib.util
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
REGISTRY = Path(__file__).resolve().parents[1] / "data" / "fal_registry.json"
|
||||
|
||||
VALID_INPUT_TYPES = {"string", "integer", "number", "boolean", "enum", "object", "array", "json"}
|
||||
VALID_OUTPUT_KINDS = {"images", "image", "video", "audio", "text", "file", "json"}
|
||||
VALID_MEDIA_KINDS = {None, "image", "video", "audio", "file"}
|
||||
|
||||
|
||||
def _registry():
|
||||
return json.loads(REGISTRY.read_text(encoding="utf-8"))
|
||||
|
||||
|
||||
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
|
||||
|
||||
|
||||
def test_models_well_formed():
|
||||
reg = _registry()
|
||||
seen_ids = set()
|
||||
for model in reg["models"]:
|
||||
eid = model["endpoint_id"]
|
||||
assert eid and "/" in eid, f"bad endpoint_id: {eid!r}"
|
||||
assert eid not in seen_ids, f"duplicate endpoint_id: {eid}"
|
||||
seen_ids.add(eid)
|
||||
assert model["title"], f"{eid}: missing title"
|
||||
assert model["category"], f"{eid}: missing category"
|
||||
assert model["output_kind"] in VALID_OUTPUT_KINDS, f"{eid}: {model['output_kind']}"
|
||||
assert isinstance(model["inputs"], list)
|
||||
|
||||
|
||||
def test_inputs_well_formed():
|
||||
reg = _registry()
|
||||
for model in reg["models"]:
|
||||
eid = model["endpoint_id"]
|
||||
names = set()
|
||||
for inp in model["inputs"]:
|
||||
name = inp["name"]
|
||||
assert name not in names, f"{eid}: duplicate input {name}"
|
||||
names.add(name)
|
||||
assert inp["type"] in VALID_INPUT_TYPES, f"{eid}.{name}: {inp['type']}"
|
||||
assert inp.get("media_kind") in VALID_MEDIA_KINDS, f"{eid}.{name}"
|
||||
if inp["type"] == "enum":
|
||||
assert inp.get("enum"), f"{eid}.{name}: enum without values"
|
||||
|
||||
|
||||
def test_enum_defaults_are_members_or_custom_size():
|
||||
reg = _registry()
|
||||
for model in reg["models"]:
|
||||
for inp in model["inputs"]:
|
||||
if inp["type"] == "enum" and inp.get("default") is not None:
|
||||
if inp["default"] in inp["enum"]:
|
||||
continue
|
||||
# two legitimate non-member shapes exist in the wild:
|
||||
# 1. has_custom_size enums defaulting to an explicit
|
||||
# {width, height} object (mapped to the custom_size preset)
|
||||
# 2. multi-select enums (is_list) defaulting to a list of
|
||||
# members (mapped to a comma-separated string widget)
|
||||
if inp.get("has_custom_size") and isinstance(inp["default"], dict):
|
||||
continue
|
||||
if inp.get("is_list") and isinstance(inp["default"], list):
|
||||
assert all(v in inp["enum"] for v in inp["default"]), (
|
||||
f"{model['endpoint_id']}.{inp['name']}: list default "
|
||||
f"contains non-members"
|
||||
)
|
||||
continue
|
||||
raise AssertionError(
|
||||
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) == {}
|
||||
@@ -0,0 +1,52 @@
|
||||
"""Unit tests for the persistent result + upload cache."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
|
||||
import pytest
|
||||
from conftest import PKG, _load_package
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def cache():
|
||||
_load_package()
|
||||
mod = importlib.import_module(f"{PKG}.nodes.utils.result_cache")
|
||||
instance = mod.ResultCache()
|
||||
instance.clear()
|
||||
yield instance
|
||||
instance.clear()
|
||||
|
||||
|
||||
def test_round_trip(cache):
|
||||
args = {"prompt": "a cat", "seed": 7}
|
||||
assert cache.get("fal-ai/test", args) is None
|
||||
cache.put("fal-ai/test", args, {"images": [{"url": "https://x/y.png"}]}, "req-1")
|
||||
hit = cache.get("fal-ai/test", args)
|
||||
assert hit == {"images": [{"url": "https://x/y.png"}]}
|
||||
|
||||
|
||||
def test_key_is_argument_order_independent(cache):
|
||||
a = cache.make_key("fal-ai/test", {"a": 1, "b": 2})
|
||||
b = cache.make_key("fal-ai/test", {"b": 2, "a": 1})
|
||||
assert a == b
|
||||
assert a != cache.make_key("fal-ai/other", {"a": 1, "b": 2})
|
||||
|
||||
|
||||
def test_different_args_miss(cache):
|
||||
cache.put("fal-ai/test", {"prompt": "a"}, {"ok": 1})
|
||||
assert cache.get("fal-ai/test", {"prompt": "b"}) is None
|
||||
|
||||
|
||||
def test_upload_cache_round_trip(cache):
|
||||
assert cache.get_upload("hash123") is None
|
||||
cache.put_upload("hash123", "https://fal.media/up.png")
|
||||
assert cache.get_upload("hash123") == "https://fal.media/up.png"
|
||||
|
||||
|
||||
def test_clear_and_stats(cache):
|
||||
cache.put("fal-ai/test", {"p": 1}, {"ok": 1})
|
||||
cache.clear()
|
||||
assert cache.get("fal-ai/test", {"p": 1}) is None
|
||||
stats = cache.stats()
|
||||
assert "entries" in stats and "db_path" in stats
|
||||
@@ -0,0 +1,156 @@
|
||||
"""Unit tests for the schema→INPUT_TYPES converter."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from helpers import _input, _model
|
||||
|
||||
|
||||
def test_required_and_optional_buckets(schema_to_inputs):
|
||||
model = _model([
|
||||
_input("prompt", "string", required=True, multiline=True),
|
||||
_input("guidance", "number", default=3.5, min=1, max=20),
|
||||
])
|
||||
it = schema_to_inputs.build_input_types(model)
|
||||
assert "prompt" in it["required"]
|
||||
assert "guidance" in it["optional"]
|
||||
assert it["required"]["prompt"][0] == "STRING"
|
||||
assert it["required"]["prompt"][1]["multiline"] is True
|
||||
|
||||
|
||||
def test_enum_becomes_dropdown(schema_to_inputs):
|
||||
model = _model([_input("style", "enum", enum=["a", "b"], default="b")])
|
||||
it = schema_to_inputs.build_input_types(model)
|
||||
spec = it["optional"]["style"]
|
||||
assert spec[0] == ["a", "b"]
|
||||
assert spec[1]["default"] == "b"
|
||||
|
||||
|
||||
def test_int_range_and_default_clamp(schema_to_inputs):
|
||||
model = _model([_input("steps", "integer", default=28, min=1, max=50)])
|
||||
it = schema_to_inputs.build_input_types(model)
|
||||
typ, opts = it["optional"]["steps"]
|
||||
assert typ == "INT"
|
||||
assert opts["min"] == 1 and opts["max"] == 50 and opts["default"] == 28
|
||||
|
||||
|
||||
def test_seed_spec(schema_to_inputs):
|
||||
model = _model([_input("seed", "integer", required=True)])
|
||||
it = schema_to_inputs.build_input_types(model)
|
||||
# seed is always optional regardless of the API marking it required
|
||||
typ, opts = it["optional"]["seed"]
|
||||
assert typ == "INT"
|
||||
assert opts["default"] == -1
|
||||
assert opts["min"] == -1
|
||||
assert opts.get("control_after_generate") is True
|
||||
|
||||
|
||||
def test_media_inputs(schema_to_inputs):
|
||||
model = _model([
|
||||
_input("image_url", "string", required=True, media_kind="image"),
|
||||
_input("video_url", "string", media_kind="video"),
|
||||
_input("audio_url", "string", media_kind="audio"),
|
||||
])
|
||||
it = schema_to_inputs.build_input_types(model)
|
||||
assert it["required"]["image_url"][0] == "IMAGE"
|
||||
assert it["optional"]["video_url"][0] == "VIDEO"
|
||||
assert it["optional"]["audio_url"][0] == "AUDIO"
|
||||
|
||||
|
||||
def test_custom_size_companions(schema_to_inputs):
|
||||
model = _model([
|
||||
_input(
|
||||
"image_size", "enum",
|
||||
enum=["square", "landscape_4_3", "custom_size"],
|
||||
default="landscape_4_3", has_custom_size=True,
|
||||
)
|
||||
])
|
||||
it = schema_to_inputs.build_input_types(model)
|
||||
assert "width" in it["optional"] and "height" in it["optional"]
|
||||
assert it["optional"]["width"][0] == "INT"
|
||||
|
||||
|
||||
def test_dict_default_maps_to_custom_size(schema_to_inputs):
|
||||
model = _model([
|
||||
_input(
|
||||
"image_size", "enum",
|
||||
enum=["square", "custom_size"],
|
||||
default={"width": 2048, "height": 1536},
|
||||
has_custom_size=True,
|
||||
)
|
||||
])
|
||||
it = schema_to_inputs.build_input_types(model)
|
||||
assert it["optional"]["image_size"][1]["default"] == "custom_size"
|
||||
assert it["optional"]["width"][1]["default"] == 2048
|
||||
assert it["optional"]["height"][1]["default"] == 1536
|
||||
|
||||
|
||||
def test_multi_select_enum_is_comma_string(schema_to_inputs):
|
||||
model = _model([
|
||||
_input("stems", "enum", enum=["vocals", "drums", "bass"],
|
||||
default=["vocals", "drums"], is_list=True)
|
||||
])
|
||||
it = schema_to_inputs.build_input_types(model)
|
||||
typ, opts = it["optional"]["stems"]
|
||||
assert typ == "STRING"
|
||||
assert opts["default"] == "vocals, drums"
|
||||
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)
|
||||
typ, opts = it["optional"]["loras"]
|
||||
assert typ == "STRING"
|
||||
assert opts["multiline"] is True
|
||||
|
||||
|
||||
def test_force_rerun_always_present(schema_to_inputs):
|
||||
it = schema_to_inputs.build_input_types(_model([]))
|
||||
typ, opts = it["optional"]["force_rerun"]
|
||||
assert typ == "BOOLEAN"
|
||||
assert opts["default"] is False
|
||||
|
||||
|
||||
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
|
||||
@@ -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
|
||||
@@ -0,0 +1,85 @@
|
||||
"""Unit tests for the /fal_api server routes' pure functions."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
from conftest import PKG, _load_package
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def routes():
|
||||
_load_package()
|
||||
return importlib.import_module(f"{PKG}.nodes.server_routes")
|
||||
|
||||
|
||||
def test_import_without_comfy_server_is_safe(routes):
|
||||
# loaded via conftest without ComfyUI's `server` module present
|
||||
assert routes is not None
|
||||
|
||||
|
||||
def test_pricing_map_covers_registry(routes):
|
||||
pricing_map = routes._pricing_map()
|
||||
assert len(pricing_map) > 100
|
||||
sample = next(iter(pricing_map.values()))
|
||||
assert "label" in sample
|
||||
for key in pricing_map:
|
||||
assert key.startswith("FalAPI_")
|
||||
|
||||
|
||||
def test_search_models(routes):
|
||||
results = routes._search_models(q="kling", category="", max_price=None, limit=10)
|
||||
assert results
|
||||
assert all("kling" in r["endpoint_id"].lower() or "kling" in r["title"].lower() for r in results)
|
||||
|
||||
|
||||
def test_search_models_price_filter(routes):
|
||||
unfiltered = routes._search_models(q="", category="", max_price=None, limit=100)
|
||||
cheap = routes._search_models(q="", category="", max_price=0.02, limit=100)
|
||||
assert len(cheap) < len(unfiltered)
|
||||
|
||||
|
||||
def test_session_shape(routes):
|
||||
payload = routes._session()
|
||||
assert set(payload) >= {"total_usd", "calls"}
|
||||
|
||||
|
||||
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
|
||||
@@ -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);
|
||||
});
|
||||
}
|
||||
@@ -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);
|
||||
});
|
||||
@@ -0,0 +1,67 @@
|
||||
"""Unit tests for fal-CDN URL passthrough (twin inputs + URL outputs)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
from helpers import _input, _model
|
||||
|
||||
|
||||
def test_media_inputs_get_direct_url_twins(schema_to_inputs):
|
||||
model = _model([
|
||||
_input("image_url", "string", required=True, media_kind="image"),
|
||||
_input("video_url", "string", media_kind="video"),
|
||||
_input("doc_url", "string", media_kind="file"),
|
||||
])
|
||||
it = schema_to_inputs.build_input_types(model)
|
||||
assert "image_url_direct_url" in it["optional"]
|
||||
assert "video_url_direct_url" in it["optional"]
|
||||
assert "doc_url_direct_url" not in it["optional"] # file kind excluded
|
||||
|
||||
|
||||
def test_direct_url_wins_over_tensor(arguments_mod, monkeypatch):
|
||||
calls = []
|
||||
monkeypatch.setattr(
|
||||
arguments_mod.ImageUtils,
|
||||
"upload_image",
|
||||
staticmethod(lambda v: calls.append(v) or "https://uploaded/x.png"),
|
||||
)
|
||||
model = _model([_input("image_url", "string", media_kind="image")])
|
||||
args = arguments_mod.build_arguments(
|
||||
model,
|
||||
{"image_url": object(), "image_url_direct_url": "https://fal.media/direct.png"},
|
||||
)
|
||||
assert args["image_url"] == "https://fal.media/direct.png"
|
||||
assert not calls # no upload happened
|
||||
|
||||
|
||||
def test_direct_url_works_without_tensor(arguments_mod):
|
||||
model = _model([_input("image_url", "string", media_kind="image")])
|
||||
args = arguments_mod.build_arguments(
|
||||
model, {"image_url_direct_url": "https://fal.media/direct.png"}
|
||||
)
|
||||
assert args["image_url"] == "https://fal.media/direct.png"
|
||||
|
||||
|
||||
def test_invalid_direct_url_raises(arguments_mod, errors_mod):
|
||||
model = _model([_input("image_url", "string", media_kind="image")])
|
||||
with pytest.raises(errors_mod.FalApiError):
|
||||
arguments_mod.build_arguments(
|
||||
model, {"image_url_direct_url": "not-a-url"}
|
||||
)
|
||||
|
||||
|
||||
def test_direct_url_list_splits_commas(arguments_mod):
|
||||
model = _model([
|
||||
_input("image_urls", "array", media_kind="image", is_list=True)
|
||||
])
|
||||
args = arguments_mod.build_arguments(
|
||||
model,
|
||||
{"image_urls_direct_url": "https://a/1.png, https://b/2.png"},
|
||||
)
|
||||
assert args["image_urls"] == ["https://a/1.png", "https://b/2.png"]
|
||||
|
||||
|
||||
def test_images_output_includes_urls(outputs_mod):
|
||||
types, names = outputs_mod.RETURN_SPECS["images"]
|
||||
assert types == ("IMAGE", "STRING")
|
||||
assert names == ("images", "image_urls")
|
||||
@@ -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"
|
||||
@@ -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()
|
||||
@@ -0,0 +1,67 @@
|
||||
// Fetch helpers and tiny utilities for the fal platform extension. No deps.
|
||||
|
||||
let apiRef = null;
|
||||
try {
|
||||
const mod = await import("../../scripts/api.js");
|
||||
apiRef = mod?.api ?? null;
|
||||
} catch (error) {
|
||||
console.debug("[fal] scripts/api.js unavailable, falling back to fetch()", error);
|
||||
}
|
||||
|
||||
function rawFetch(path, options) {
|
||||
if (apiRef && typeof apiRef.fetchApi === "function") {
|
||||
return apiRef.fetchApi(path, options);
|
||||
}
|
||||
return fetch(path, options);
|
||||
}
|
||||
|
||||
export async function getJson(path) {
|
||||
const response = await rawFetch(`/fal_api${path}`);
|
||||
if (!response.ok) {
|
||||
throw new Error(`GET /fal_api${path} -> HTTP ${response.status}`);
|
||||
}
|
||||
return await response.json();
|
||||
}
|
||||
|
||||
export async function postJson(path, body) {
|
||||
const response = await rawFetch(`/fal_api${path}`, {
|
||||
method: "POST",
|
||||
headers: { "Content-Type": "application/json" },
|
||||
body: JSON.stringify(body ?? {}),
|
||||
});
|
||||
if (!response.ok) {
|
||||
throw new Error(`POST /fal_api${path} -> HTTP ${response.status}`);
|
||||
}
|
||||
return await response.json();
|
||||
}
|
||||
|
||||
export function debounce(fn, delayMs) {
|
||||
let timer = null;
|
||||
return (...args) => {
|
||||
if (timer) clearTimeout(timer);
|
||||
timer = setTimeout(() => {
|
||||
timer = null;
|
||||
fn(...args);
|
||||
}, delayMs);
|
||||
};
|
||||
}
|
||||
|
||||
export function humanAge(unixSeconds) {
|
||||
if (typeof unixSeconds !== "number") return "?";
|
||||
const seconds = Math.max(0, Date.now() / 1000 - unixSeconds);
|
||||
if (seconds < 60) return `${Math.floor(seconds)}s`;
|
||||
if (seconds < 3600) return `${Math.floor(seconds / 60)}m`;
|
||||
if (seconds < 86400) return `${Math.floor(seconds / 3600)}h`;
|
||||
return `${Math.floor(seconds / 86400)}d`;
|
||||
}
|
||||
|
||||
export function shortEndpoint(endpoint) {
|
||||
const parts = String(endpoint || "").split("/").filter(Boolean);
|
||||
if (parts.length <= 1) return endpoint || "(unknown)";
|
||||
return parts.slice(-2).join("/");
|
||||
}
|
||||
|
||||
export function formatUsd(value) {
|
||||
if (typeof value !== "number" || !isFinite(value)) return null;
|
||||
return `$${value.toFixed(value < 10 ? 4 : 2).replace(/\.?0+$/, "") || "0"}`;
|
||||
}
|
||||
@@ -0,0 +1,180 @@
|
||||
// Endpoint autocomplete for free-typed fal nodes (Any Endpoint, Submit, ...).
|
||||
|
||||
import { debounce, getJson } from "./fal_api.js";
|
||||
import { FREE_TYPED_NODES, refreshNodeBadge } from "./fal_badges.js";
|
||||
|
||||
const SEARCH_DEBOUNCE_MS = 300;
|
||||
const RESULT_LIMIT = 25;
|
||||
const ENDPOINT_WIDGET = "endpoint_id";
|
||||
|
||||
let activePopup = null;
|
||||
|
||||
function destroyPopup() {
|
||||
if (!activePopup) return;
|
||||
try {
|
||||
activePopup.cleanup?.();
|
||||
activePopup.element.remove();
|
||||
} catch (error) {
|
||||
console.debug("[fal] popup cleanup failed", error);
|
||||
}
|
||||
activePopup = null;
|
||||
}
|
||||
|
||||
function findTarget(canvas, value) {
|
||||
const pair = canvas?.node_widget;
|
||||
if (
|
||||
Array.isArray(pair) &&
|
||||
pair[0] &&
|
||||
pair[1]?.name === ENDPOINT_WIDGET &&
|
||||
FREE_TYPED_NODES.has(pair[0].type)
|
||||
) {
|
||||
return pair;
|
||||
}
|
||||
const selected = Object.values(canvas?.selected_nodes || {});
|
||||
for (const node of selected) {
|
||||
if (!FREE_TYPED_NODES.has(node?.type)) continue;
|
||||
const widget = (node.widgets || []).find((w) => w?.name === ENDPOINT_WIDGET);
|
||||
if (widget && String(widget.value ?? "") === String(value ?? "")) return [node, widget];
|
||||
}
|
||||
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;
|
||||
text.append(title, endpoint);
|
||||
if (model.label) {
|
||||
const price = document.createElement("span");
|
||||
price.className = "fal-suggest-price";
|
||||
price.textContent = model.label;
|
||||
text.append(price);
|
||||
}
|
||||
row.append(text);
|
||||
row.addEventListener("mousedown", (event) => {
|
||||
event.preventDefault();
|
||||
event.stopPropagation();
|
||||
apply(model.endpoint_id);
|
||||
});
|
||||
return row;
|
||||
}
|
||||
|
||||
function positionPopup(popup, dialog) {
|
||||
try {
|
||||
const rect = dialog.getBoundingClientRect();
|
||||
popup.style.left = `${Math.max(4, rect.left)}px`;
|
||||
popup.style.top = `${rect.bottom + 4}px`;
|
||||
} catch (error) {
|
||||
console.debug("[fal] popup positioning failed", error);
|
||||
}
|
||||
}
|
||||
|
||||
function attachAutocomplete(dialog, input, node, widget) {
|
||||
destroyPopup();
|
||||
const popup = document.createElement("div");
|
||||
popup.className = "fal-suggest";
|
||||
document.body.appendChild(popup);
|
||||
positionPopup(popup, dialog);
|
||||
|
||||
const apply = (endpointId) => {
|
||||
try {
|
||||
widget.value = endpointId;
|
||||
input.value = endpointId;
|
||||
widget.callback?.(endpointId);
|
||||
refreshNodeBadge(node, endpointId);
|
||||
node.setDirtyCanvas?.(true, true);
|
||||
} catch (error) {
|
||||
console.debug("[fal] could not apply endpoint suggestion", error);
|
||||
}
|
||||
destroyPopup();
|
||||
};
|
||||
|
||||
const search = async (query) => {
|
||||
try {
|
||||
const models = await getJson(
|
||||
`/models?q=${encodeURIComponent(query || "")}&limit=${RESULT_LIMIT}`
|
||||
);
|
||||
if (activePopup?.element !== popup) return;
|
||||
popup.replaceChildren(...(models || []).map((model) => resultRow(model, apply)));
|
||||
popup.style.display = models?.length ? "block" : "none";
|
||||
} catch (error) {
|
||||
console.debug("[fal] endpoint search failed", error);
|
||||
}
|
||||
};
|
||||
const debouncedSearch = debounce(() => search(input.value), SEARCH_DEBOUNCE_MS);
|
||||
|
||||
const onInput = () => debouncedSearch();
|
||||
const onKeyDown = (event) => {
|
||||
if (event.key === "Escape" || event.key === "Enter") destroyPopup();
|
||||
};
|
||||
const onOutsideDown = (event) => {
|
||||
if (!popup.contains(event.target) && event.target !== input) destroyPopup();
|
||||
};
|
||||
input.addEventListener("input", onInput);
|
||||
input.addEventListener("keydown", onKeyDown);
|
||||
document.addEventListener("mousedown", onOutsideDown, true);
|
||||
const aliveCheck = setInterval(() => {
|
||||
if (!input.isConnected) destroyPopup();
|
||||
}, 500);
|
||||
|
||||
activePopup = {
|
||||
element: popup,
|
||||
cleanup: () => {
|
||||
input.removeEventListener("input", onInput);
|
||||
input.removeEventListener("keydown", onKeyDown);
|
||||
document.removeEventListener("mousedown", onOutsideDown, true);
|
||||
clearInterval(aliveCheck);
|
||||
},
|
||||
};
|
||||
search(input.value);
|
||||
}
|
||||
|
||||
export function installAutocomplete() {
|
||||
const canvasClass = globalThis.LGraphCanvas;
|
||||
if (!canvasClass?.prototype?.prompt) {
|
||||
console.debug("[fal] LGraphCanvas.prompt unavailable; endpoint autocomplete disabled");
|
||||
return;
|
||||
}
|
||||
const originalPrompt = canvasClass.prototype.prompt;
|
||||
canvasClass.prototype.prompt = function (title, value, callback, event, ...rest) {
|
||||
const dialog = originalPrompt.call(this, title, value, callback, event, ...rest);
|
||||
try {
|
||||
const target = findTarget(this, value);
|
||||
const input = dialog?.querySelector?.("input, textarea");
|
||||
if (target && input) attachAutocomplete(dialog, input, target[0], target[1]);
|
||||
} catch (error) {
|
||||
console.debug("[fal] autocomplete attach failed", error);
|
||||
}
|
||||
return dialog;
|
||||
};
|
||||
}
|
||||
@@ -0,0 +1,151 @@
|
||||
// Cost badges: a small price pill floating above every priced fal node.
|
||||
|
||||
import { debounce, getJson } from "./fal_api.js";
|
||||
|
||||
export const FREE_TYPED_NODES = new Set([
|
||||
"FalAnyEndpoint_fal",
|
||||
"FalSubmit_fal",
|
||||
"FalCostEstimator_fal",
|
||||
"FalResultByRequestId_fal",
|
||||
]);
|
||||
|
||||
const DYNAMIC_NODE_PREFIX = "FalAPI_";
|
||||
const ENDPOINT_WIDGET = "endpoint_id";
|
||||
const LIVE_DEBOUNCE_MS = 500;
|
||||
|
||||
// node class key -> {label, per_run}; filled asynchronously.
|
||||
let pricingMap = {};
|
||||
// endpoint_id -> label|null for free-typed endpoint lookups.
|
||||
const liveLabelCache = new Map();
|
||||
|
||||
export async function loadPricingMap() {
|
||||
try {
|
||||
pricingMap = (await getJson("/pricing_map")) || {};
|
||||
console.debug(`[fal] pricing map loaded (${Object.keys(pricingMap).length} nodes)`);
|
||||
} catch (error) {
|
||||
console.debug("[fal] pricing map unavailable", error);
|
||||
}
|
||||
}
|
||||
|
||||
function titleHeight() {
|
||||
const lg = globalThis.LiteGraph;
|
||||
const height = lg && typeof lg.NODE_TITLE_HEIGHT === "number" ? lg.NODE_TITLE_HEIGHT : 30;
|
||||
return height;
|
||||
}
|
||||
|
||||
function drawPill(node, ctx, text) {
|
||||
if (!text || node?.flags?.collapsed) return;
|
||||
ctx.save();
|
||||
try {
|
||||
ctx.font = "10px Inter, 'Segoe UI', sans-serif";
|
||||
const padX = 7;
|
||||
const height = 16;
|
||||
const width = ctx.measureText(text).width + padX * 2;
|
||||
const x = 0;
|
||||
const y = -titleHeight() - height - 6;
|
||||
ctx.beginPath();
|
||||
if (typeof ctx.roundRect === "function") {
|
||||
ctx.roundRect(x, y, width, height, height / 2);
|
||||
} else {
|
||||
ctx.rect(x, y, width, height);
|
||||
}
|
||||
ctx.fillStyle = "rgba(12, 12, 18, 0.88)";
|
||||
ctx.fill();
|
||||
ctx.lineWidth = 1;
|
||||
ctx.strokeStyle = "rgba(167, 139, 250, 0.45)";
|
||||
ctx.stroke();
|
||||
ctx.fillStyle = "#ece9fd";
|
||||
ctx.textAlign = "left";
|
||||
ctx.textBaseline = "middle";
|
||||
ctx.fillText(text, x + padX, y + height / 2 + 0.5);
|
||||
} finally {
|
||||
ctx.restore();
|
||||
}
|
||||
}
|
||||
|
||||
function labelForNode(node, typeName) {
|
||||
if (FREE_TYPED_NODES.has(typeName)) return node._falPriceLabel || null;
|
||||
return pricingMap[typeName]?.label || null;
|
||||
}
|
||||
|
||||
export function refreshNodeBadge(node, endpointValue) {
|
||||
const endpoint = String(endpointValue ?? "").trim();
|
||||
if (!endpoint) {
|
||||
node._falPriceLabel = null;
|
||||
node.setDirtyCanvas?.(true, true);
|
||||
return;
|
||||
}
|
||||
if (liveLabelCache.has(endpoint)) {
|
||||
node._falPriceLabel = liveLabelCache.get(endpoint);
|
||||
node.setDirtyCanvas?.(true, true);
|
||||
return;
|
||||
}
|
||||
getJson(`/pricing?endpoint_id=${encodeURIComponent(endpoint)}`)
|
||||
.then((data) => {
|
||||
const label = data?.label || null;
|
||||
liveLabelCache.set(endpoint, label);
|
||||
node._falPriceLabel = label;
|
||||
node.setDirtyCanvas?.(true, true);
|
||||
})
|
||||
.catch((error) => console.debug("[fal] live pricing lookup failed", error));
|
||||
}
|
||||
|
||||
function watchEndpointWidget(node) {
|
||||
const widget = (node.widgets || []).find((w) => w?.name === ENDPOINT_WIDGET);
|
||||
if (!widget) return;
|
||||
const refresh = debounce(() => refreshNodeBadge(node, widget.value), LIVE_DEBOUNCE_MS);
|
||||
const previousCallback = widget.callback;
|
||||
widget.callback = function (...args) {
|
||||
const result = previousCallback?.apply(this, args);
|
||||
try {
|
||||
refresh();
|
||||
} catch (error) {
|
||||
console.debug("[fal] endpoint widget watch failed", error);
|
||||
}
|
||||
return result;
|
||||
};
|
||||
refreshNodeBadge(node, widget.value);
|
||||
}
|
||||
|
||||
function hookFreeTypedNode(nodeType) {
|
||||
const previousCreated = nodeType.prototype.onNodeCreated;
|
||||
nodeType.prototype.onNodeCreated = function (...args) {
|
||||
const result = previousCreated?.apply(this, args);
|
||||
try {
|
||||
watchEndpointWidget(this);
|
||||
} catch (error) {
|
||||
console.debug("[fal] could not watch endpoint widget", error);
|
||||
}
|
||||
return result;
|
||||
};
|
||||
const previousConfigure = nodeType.prototype.onConfigure;
|
||||
nodeType.prototype.onConfigure = function (...args) {
|
||||
const result = previousConfigure?.apply(this, args);
|
||||
try {
|
||||
const widget = (this.widgets || []).find((w) => w?.name === ENDPOINT_WIDGET);
|
||||
if (widget) refreshNodeBadge(this, widget.value);
|
||||
} catch (error) {
|
||||
console.debug("[fal] badge refresh on configure failed", error);
|
||||
}
|
||||
return result;
|
||||
};
|
||||
}
|
||||
|
||||
export function setupNodeBadges(nodeType, nodeData) {
|
||||
const typeName = nodeData?.name;
|
||||
if (!typeName || !nodeType?.prototype) return;
|
||||
const isFreeTyped = FREE_TYPED_NODES.has(typeName);
|
||||
if (!isFreeTyped && !typeName.startsWith(DYNAMIC_NODE_PREFIX)) return;
|
||||
|
||||
const previousDraw = nodeType.prototype.onDrawForeground;
|
||||
nodeType.prototype.onDrawForeground = function (ctx, ...args) {
|
||||
const result = previousDraw?.apply(this, [ctx, ...args]);
|
||||
try {
|
||||
drawPill(this, ctx, labelForNode(this, typeName));
|
||||
} catch (error) {
|
||||
console.debug("[fal] badge draw failed", error);
|
||||
}
|
||||
return result;
|
||||
};
|
||||
if (isFreeTyped) hookFreeTypedNode(nodeType);
|
||||
}
|
||||
@@ -0,0 +1,268 @@
|
||||
/* fal platform extension styles: sidebar panel + endpoint suggestions. */
|
||||
|
||||
.fal-panel {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 10px;
|
||||
padding: 12px;
|
||||
font-size: 12px;
|
||||
color: #ece9fd;
|
||||
}
|
||||
|
||||
.fal-panel-title {
|
||||
font-size: 13px;
|
||||
font-weight: 600;
|
||||
letter-spacing: 0.04em;
|
||||
}
|
||||
|
||||
.fal-stats {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 6px;
|
||||
padding: 10px;
|
||||
border: 1px solid rgba(167, 139, 250, 0.25);
|
||||
border-radius: 8px;
|
||||
background: rgba(12, 12, 18, 0.6);
|
||||
}
|
||||
|
||||
.fal-stat {
|
||||
display: flex;
|
||||
justify-content: space-between;
|
||||
gap: 8px;
|
||||
}
|
||||
|
||||
.fal-stat-label {
|
||||
opacity: 0.65;
|
||||
}
|
||||
|
||||
.fal-stat-value {
|
||||
font-variant-numeric: tabular-nums;
|
||||
font-weight: 600;
|
||||
}
|
||||
|
||||
.fal-jobs-header {
|
||||
font-weight: 600;
|
||||
opacity: 0.85;
|
||||
}
|
||||
|
||||
.fal-jobs {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 4px;
|
||||
overflow-y: auto;
|
||||
max-height: 60vh;
|
||||
}
|
||||
|
||||
.fal-job {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 8px;
|
||||
padding: 6px 8px;
|
||||
border-radius: 6px;
|
||||
background: rgba(255, 255, 255, 0.04);
|
||||
}
|
||||
|
||||
.fal-job-icon {
|
||||
flex: none;
|
||||
}
|
||||
|
||||
.fal-job-info {
|
||||
min-width: 0;
|
||||
flex: 1;
|
||||
}
|
||||
|
||||
.fal-job-endpoint {
|
||||
overflow: hidden;
|
||||
text-overflow: ellipsis;
|
||||
white-space: nowrap;
|
||||
}
|
||||
|
||||
.fal-job-meta {
|
||||
font-size: 10px;
|
||||
opacity: 0.6;
|
||||
overflow: hidden;
|
||||
text-overflow: ellipsis;
|
||||
white-space: nowrap;
|
||||
cursor: pointer;
|
||||
}
|
||||
|
||||
.fal-job-meta:hover {
|
||||
opacity: 0.9;
|
||||
}
|
||||
|
||||
.fal-job-cancel {
|
||||
flex: none;
|
||||
padding: 2px 8px;
|
||||
border: 1px solid rgba(248, 113, 113, 0.5);
|
||||
border-radius: 5px;
|
||||
background: transparent;
|
||||
color: #fca5a5;
|
||||
font-size: 10px;
|
||||
cursor: pointer;
|
||||
}
|
||||
|
||||
.fal-job-cancel:hover {
|
||||
background: rgba(248, 113, 113, 0.15);
|
||||
}
|
||||
|
||||
.fal-jobs-empty {
|
||||
padding: 6px 2px;
|
||||
}
|
||||
|
||||
.fal-muted {
|
||||
opacity: 0.55;
|
||||
}
|
||||
|
||||
/* Floating fallback when no sidebar API is available. */
|
||||
|
||||
.fal-floating {
|
||||
position: fixed;
|
||||
right: 14px;
|
||||
bottom: 14px;
|
||||
z-index: 10000;
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
align-items: flex-end;
|
||||
gap: 8px;
|
||||
}
|
||||
|
||||
.fal-floating-toggle {
|
||||
padding: 6px 14px;
|
||||
border: 1px solid rgba(167, 139, 250, 0.5);
|
||||
border-radius: 999px;
|
||||
background: rgba(12, 12, 18, 0.92);
|
||||
color: #ece9fd;
|
||||
font-size: 12px;
|
||||
font-weight: 600;
|
||||
cursor: pointer;
|
||||
}
|
||||
|
||||
.fal-floating-panel {
|
||||
width: 300px;
|
||||
max-height: 70vh;
|
||||
overflow-y: auto;
|
||||
border: 1px solid rgba(167, 139, 250, 0.3);
|
||||
border-radius: 10px;
|
||||
background: rgba(12, 12, 18, 0.95);
|
||||
box-shadow: 0 8px 30px rgba(0, 0, 0, 0.45);
|
||||
}
|
||||
|
||||
/* Endpoint autocomplete popup. */
|
||||
|
||||
.fal-suggest {
|
||||
position: fixed;
|
||||
z-index: 10001;
|
||||
width: 380px;
|
||||
max-height: 320px;
|
||||
overflow-y: auto;
|
||||
border: 1px solid rgba(167, 139, 250, 0.35);
|
||||
border-radius: 8px;
|
||||
background: rgba(12, 12, 18, 0.97);
|
||||
box-shadow: 0 8px 30px rgba(0, 0, 0, 0.5);
|
||||
font-size: 12px;
|
||||
color: #ece9fd;
|
||||
}
|
||||
|
||||
.fal-suggest-item {
|
||||
display: flex;
|
||||
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);
|
||||
}
|
||||
|
||||
.fal-suggest-title {
|
||||
font-weight: 600;
|
||||
}
|
||||
|
||||
.fal-suggest-endpoint {
|
||||
font-size: 10px;
|
||||
opacity: 0.65;
|
||||
}
|
||||
|
||||
.fal-suggest-price {
|
||||
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;
|
||||
}
|
||||
@@ -0,0 +1,60 @@
|
||||
// fal platform extension: cost badges, session/balance/jobs sidebar,
|
||||
// and endpoint autocomplete for the ComfyUI canvas.
|
||||
|
||||
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.
|
||||
const pricingReady = loadPricingMap();
|
||||
|
||||
function injectStylesheet() {
|
||||
try {
|
||||
const link = document.createElement("link");
|
||||
link.rel = "stylesheet";
|
||||
link.href = new URL("./fal_platform.css", import.meta.url).href;
|
||||
document.head.appendChild(link);
|
||||
} catch (error) {
|
||||
console.debug("[fal] stylesheet injection failed", error);
|
||||
}
|
||||
}
|
||||
|
||||
app.registerExtension({
|
||||
name: "fal.platform",
|
||||
|
||||
beforeRegisterNodeDef(nodeType, nodeData) {
|
||||
try {
|
||||
setupNodeBadges(nodeType, nodeData);
|
||||
} 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() {
|
||||
injectStylesheet();
|
||||
try {
|
||||
registerSidebar(app);
|
||||
} catch (error) {
|
||||
console.debug("[fal] sidebar registration failed", error);
|
||||
}
|
||||
try {
|
||||
installAutocomplete();
|
||||
} catch (error) {
|
||||
console.debug("[fal] autocomplete install failed", error);
|
||||
}
|
||||
try {
|
||||
await pricingReady;
|
||||
app.graph?.setDirtyCanvas?.(true, true);
|
||||
} catch (error) {
|
||||
console.debug("[fal] pricing map warmup failed", error);
|
||||
}
|
||||
},
|
||||
});
|
||||
@@ -0,0 +1,322 @@
|
||||
// Sidebar panel: session spend, account balance and the async job inbox.
|
||||
|
||||
import { formatUsd, getJson, humanAge, postJson, shortEndpoint } from "./fal_api.js";
|
||||
|
||||
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;
|
||||
|
||||
function element(tag, className, text) {
|
||||
const el = document.createElement(tag);
|
||||
if (className) el.className = className;
|
||||
if (text != null) el.textContent = text;
|
||||
return el;
|
||||
}
|
||||
|
||||
function copyText(text, feedbackEl) {
|
||||
const done = () => {
|
||||
if (!feedbackEl) return;
|
||||
const original = feedbackEl.textContent;
|
||||
feedbackEl.textContent = "copied!";
|
||||
setTimeout(() => {
|
||||
feedbackEl.textContent = original;
|
||||
}, 900);
|
||||
};
|
||||
try {
|
||||
if (navigator.clipboard?.writeText) {
|
||||
navigator.clipboard.writeText(text).then(done, () => {});
|
||||
return;
|
||||
}
|
||||
} catch (error) {
|
||||
console.debug("[fal] clipboard copy failed", error);
|
||||
}
|
||||
done();
|
||||
}
|
||||
|
||||
async function cancelJob(job, refresh) {
|
||||
try {
|
||||
const result = await postJson("/cancel", {
|
||||
endpoint_id: job.endpoint,
|
||||
request_id: job.request_id,
|
||||
});
|
||||
if (!result?.ok) console.debug("[fal] cancel refused", result?.error);
|
||||
} catch (error) {
|
||||
console.debug("[fal] cancel request failed", error);
|
||||
}
|
||||
refresh();
|
||||
}
|
||||
|
||||
function jobRow(job, refresh) {
|
||||
const row = element("div", "fal-job");
|
||||
const pending = job.status === "submitted";
|
||||
row.append(element("span", "fal-job-icon", pending ? "⏳" : "✅"));
|
||||
|
||||
const info = element("div", "fal-job-info");
|
||||
info.append(element("div", "fal-job-endpoint", shortEndpoint(job.endpoint)));
|
||||
const requestId = String(job.request_id || "-");
|
||||
const meta = element(
|
||||
"div",
|
||||
"fal-job-meta",
|
||||
`${humanAge(pending ? job.submitted_at : job.collected_at ?? job.submitted_at)} ago · ${requestId}`
|
||||
);
|
||||
meta.title = "Click to copy request id";
|
||||
meta.addEventListener("click", () => copyText(requestId, meta));
|
||||
info.append(meta);
|
||||
row.append(info);
|
||||
|
||||
if (pending) {
|
||||
const cancel = element("button", "fal-job-cancel", "Cancel");
|
||||
cancel.addEventListener("click", () => cancelJob(job, refresh));
|
||||
row.append(cancel);
|
||||
}
|
||||
return row;
|
||||
}
|
||||
|
||||
function buildPanel() {
|
||||
const root = element("div", "fal-panel");
|
||||
root.append(element("div", "fal-panel-title", "fal platform"));
|
||||
|
||||
const stats = element("div", "fal-stats");
|
||||
const session = element("div", "fal-stat");
|
||||
session.append(element("span", "fal-stat-label", "Session"), element("span", "fal-stat-value", "…"));
|
||||
const balance = element("div", "fal-stat");
|
||||
balance.append(element("span", "fal-stat-label", "Balance"), element("span", "fal-stat-value", "…"));
|
||||
stats.append(session, balance);
|
||||
|
||||
const jobsHeader = element("div", "fal-jobs-header", "Jobs");
|
||||
const jobs = element("div", "fal-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;
|
||||
target.textContent = `${total} · ${calls} call${calls === 1 ? "" : "s"}`;
|
||||
}
|
||||
|
||||
function renderBalance(target, data) {
|
||||
const balance = formatUsd(data?.balance_usd);
|
||||
target.textContent = balance ?? "unavailable";
|
||||
target.classList.toggle("fal-muted", balance == null);
|
||||
}
|
||||
|
||||
function renderJobs(view, data, refresh) {
|
||||
const jobs = Array.isArray(data?.jobs) ? data.jobs : [];
|
||||
const counts = data?.counts || {};
|
||||
view.jobsHeader.textContent = `Jobs · ${counts.submitted ?? 0} pending, ${counts.collected ?? 0} collected`;
|
||||
view.jobs.replaceChildren(
|
||||
...(jobs.length
|
||||
? jobs.map((job) => jobRow(job, refresh))
|
||||
: [element("div", "fal-muted fal-jobs-empty", "No jobs yet — queue one with Fal Submit.")])
|
||||
);
|
||||
}
|
||||
|
||||
async function refreshPanel(view) {
|
||||
const refresh = () => refreshPanel(view).catch(() => {});
|
||||
const [session, balance, jobs] = await Promise.allSettled([
|
||||
getJson("/session"),
|
||||
getJson("/balance"),
|
||||
getJson(`/jobs?limit=${JOB_LIMIT}`),
|
||||
]);
|
||||
try {
|
||||
if (session.status === "fulfilled") renderSession(view.sessionValue, session.value);
|
||||
if (balance.status === "fulfilled") renderBalance(view.balanceValue, balance.value);
|
||||
else renderBalance(view.balanceValue, {});
|
||||
if (jobs.status === "fulfilled") renderJobs(view, jobs.value, refresh);
|
||||
} catch (error) {
|
||||
console.debug("[fal] panel render failed", error);
|
||||
}
|
||||
}
|
||||
|
||||
function panelVisible() {
|
||||
return !!panelRoot && panelRoot.isConnected && panelRoot.offsetParent !== null && !document.hidden;
|
||||
}
|
||||
|
||||
function startRefreshLoop(view) {
|
||||
if (refreshTimer) clearInterval(refreshTimer);
|
||||
const tick = () => {
|
||||
if (!panelVisible()) return;
|
||||
refreshPanel(view).catch((error) => console.debug("[fal] panel refresh failed", error));
|
||||
};
|
||||
refreshTimer = setInterval(tick, REFRESH_MS);
|
||||
refreshPanel(view).catch((error) => console.debug("[fal] initial panel refresh failed", error));
|
||||
}
|
||||
|
||||
export function mountPanel(container) {
|
||||
const view = buildPanel();
|
||||
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() {
|
||||
const wrapper = element("div", "fal-floating");
|
||||
const panelHost = element("div", "fal-floating-panel");
|
||||
panelHost.style.display = "none";
|
||||
const toggle = element("button", "fal-floating-toggle", "fal");
|
||||
toggle.title = "fal: session cost, balance, jobs";
|
||||
toggle.addEventListener("click", () => {
|
||||
const hidden = panelHost.style.display === "none";
|
||||
panelHost.style.display = hidden ? "block" : "none";
|
||||
if (hidden) mountPanel(panelHost);
|
||||
});
|
||||
wrapper.append(panelHost, toggle);
|
||||
document.body.appendChild(wrapper);
|
||||
}
|
||||
|
||||
export function registerSidebar(app) {
|
||||
try {
|
||||
const manager = app?.extensionManager;
|
||||
if (manager && typeof manager.registerSidebarTab === "function") {
|
||||
manager.registerSidebarTab({
|
||||
id: "fal-platform",
|
||||
icon: "pi pi-bolt",
|
||||
title: "fal",
|
||||
tooltip: "fal: session cost, balance, jobs",
|
||||
type: "custom",
|
||||
render: (el) => {
|
||||
try {
|
||||
mountPanel(el);
|
||||
} catch (error) {
|
||||
console.debug("[fal] sidebar mount failed", error);
|
||||
}
|
||||
},
|
||||
});
|
||||
return;
|
||||
}
|
||||
} catch (error) {
|
||||
console.debug("[fal] sidebar tab registration failed", error);
|
||||
}
|
||||
try {
|
||||
mountFloatingFallback();
|
||||
} catch (error) {
|
||||
console.debug("[fal] floating panel fallback failed", error);
|
||||
}
|
||||
}
|
||||
@@ -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;
|
||||
};
|
||||
}
|
||||
Reference in New Issue
Block a user