Compare commits
91
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
dc5b3f764a | ||
|
|
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 | ||
|
|
a5c55fd44b | ||
|
|
63cb3fcf0d | ||
|
|
e1faa49e25 | ||
|
|
c8c437b495 | ||
|
|
81f9031625 | ||
|
|
1c1be8ae31 | ||
|
|
cd5c9ef258 | ||
|
|
ec8880895d | ||
|
|
f8b1efa75d | ||
|
|
975d555e29 | ||
|
|
4988995bf7 | ||
|
|
ee026dd560 | ||
|
|
93d6ad2875 | ||
|
|
c22792a581 | ||
|
|
779e0b1028 | ||
|
|
2f7f43da45 | ||
|
|
a8bdb5bc6d | ||
|
|
c3fae085d6 | ||
|
|
1c67dda258 | ||
|
|
96b0cd0976 | ||
|
|
116bfbd4e0 | ||
|
|
2797366781 | ||
|
|
ebecad477a | ||
|
|
a46b9465e6 | ||
|
|
1d2ecc823e | ||
|
|
a8c202b045 | ||
|
|
e2a41cc5ff | ||
|
|
68328f8526 | ||
|
|
6a4b736773 | ||
|
|
cfd626541b |
@@ -0,0 +1,64 @@
|
||||
name: CI
|
||||
|
||||
on:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
|
||||
jobs:
|
||||
lint:
|
||||
name: Lint (ruff)
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Check out code
|
||||
uses: actions/checkout@v4
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
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@v4
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
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 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,28 @@
|
||||
name: Publish to Comfy registry
|
||||
on:
|
||||
workflow_dispatch:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
- master
|
||||
paths:
|
||||
- "pyproject.toml"
|
||||
|
||||
permissions:
|
||||
issues: write
|
||||
|
||||
jobs:
|
||||
publish-node:
|
||||
name: Publish Custom Node to registry
|
||||
runs-on: ubuntu-latest
|
||||
if: ${{ github.repository_owner == 'gokayfem' }}
|
||||
steps:
|
||||
- name: Check out code
|
||||
uses: actions/checkout@v4
|
||||
with:
|
||||
submodules: true
|
||||
- name: Publish Custom Node
|
||||
uses: Comfy-Org/publish-node-action@v1
|
||||
with:
|
||||
## Add your own personal access token to your Github Repository secrets and reference it here.
|
||||
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
|
||||
@@ -0,0 +1,49 @@
|
||||
name: Refresh fal model registry
|
||||
|
||||
on:
|
||||
schedule:
|
||||
# Every Monday at 06:00 UTC
|
||||
- cron: "0 6 * * 1"
|
||||
workflow_dispatch:
|
||||
|
||||
permissions:
|
||||
contents: write
|
||||
pull-requests: write
|
||||
|
||||
jobs:
|
||||
refresh:
|
||||
name: Rebuild registry and open PR
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Check out code
|
||||
uses: actions/checkout@v4
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.11"
|
||||
- name: Rebuild registry
|
||||
run: python scripts/build_registry.py --out data/fal_registry.json
|
||||
- name: Regenerate MODELS.md
|
||||
run: python scripts/build_readme.py
|
||||
- name: Summarize changes
|
||||
id: diff
|
||||
run: |
|
||||
{
|
||||
echo "stat<<EOF"
|
||||
git diff --stat
|
||||
echo "EOF"
|
||||
} >> "$GITHUB_OUTPUT"
|
||||
# create-pull-request skips PR creation when there are no changes.
|
||||
- name: Create pull request
|
||||
uses: peter-evans/create-pull-request@v6
|
||||
with:
|
||||
branch: chore/registry-refresh
|
||||
commit-message: "chore: refresh fal model registry"
|
||||
title: "Refresh fal model registry"
|
||||
body: |
|
||||
Automated weekly refresh of `data/fal_registry.json` via `scripts/build_registry.py`.
|
||||
|
||||
```
|
||||
${{ steps.diff.outputs.stat }}
|
||||
```
|
||||
delete-branch: true
|
||||
+15
@@ -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]
|
||||
@@ -160,3 +163,15 @@ cython_debug/
|
||||
# and can be added to the global gitignore or merged into this file. For a more nuclear
|
||||
# option (not recommended) you can uncomment the following to ignore the entire idea folder.
|
||||
#.idea/
|
||||
|
||||
# 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,93 @@
|
||||
# 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 two ways:
|
||||
|
||||
- A **weekly GitHub Action** (`.github/workflows/registry-refresh.yml`) rebuilds the registry and opens a PR.
|
||||
- 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)).
|
||||
|
||||
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
|
||||
```
|
||||
|
||||
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
|
||||
data/fal_registry.json committed model catalog (~1,391 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,107 +1,172 @@
|
||||
# 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,391 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 full model 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
|
||||
|
||||
- **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.
|
||||
- **MODELS.md** — the full ~1,391-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,391 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:
|
||||
```
|
||||
cd custom_nodes
|
||||
```
|
||||
|
||||
2. Clone this repository:
|
||||
```
|
||||
git clone https://github.com/gokayfem/ComfyUI-fal-API.git
|
||||
```
|
||||
|
||||
3. Install the required dependencies:
|
||||
```
|
||||
pip install -r requirements.txt
|
||||
```
|
||||
4. 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 General (fal)**: ControlNets, Ipadapters, Loras for Flux Dev
|
||||
[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 Video Generation (fal)**: Advanced video generation with Kling Pro
|
||||
- **Runway Gen3 Image-to-Video (fal)**: Convert images to videos using Runway Gen3
|
||||
- **Luma Dream Machine (fal)**: Create videos with Luma Dream Machine
|
||||
- **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 weekly GitHub Action refreshes `data/fal_registry.json` and opens a PR; you can also run `python scripts/build_registry.py` yourself (then `python scripts/build_readme.py` to regenerate MODELS.md), or use the refresh button in the fal sidebar. If ~1,391 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-FLUX-fal-API
|
||||
cd custom_nodes/ComfyUI-fal-API
|
||||
git pull
|
||||
pip install -r requirements.txt
|
||||
```
|
||||
@@ -109,15 +174,16 @@ If you encounter any errors during installation or usage, try the following:
|
||||
```
|
||||
ComfyUI_windows_portable>.\python_embeded\python.exe -m pip install fal-client
|
||||
```
|
||||
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.
|
||||
|
||||
## 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.
|
||||
|
||||
## 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,216 @@
|
||||
{
|
||||
"id": "80767774-d39b-4f73-a75a-3c1327f92316",
|
||||
"revision": 0,
|
||||
"last_node_id": 94,
|
||||
"last_link_id": 199,
|
||||
"nodes": [
|
||||
{
|
||||
"id": 93,
|
||||
"type": "LoadImage",
|
||||
"pos": [
|
||||
2148.41357421875,
|
||||
-644.6304931640625
|
||||
],
|
||||
"size": [
|
||||
284.00726318359375,
|
||||
497.4951477050781
|
||||
],
|
||||
"flags": {},
|
||||
"order": 0,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
198
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "MASK",
|
||||
"type": "MASK",
|
||||
"links": null
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "LoadImage"
|
||||
},
|
||||
"widgets_values": [
|
||||
"image (24).png",
|
||||
"image"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 92,
|
||||
"type": "LoadImage",
|
||||
"pos": [
|
||||
1846.3321533203125,
|
||||
-654.7257080078125
|
||||
],
|
||||
"size": [
|
||||
270,
|
||||
510.1016845703125
|
||||
],
|
||||
"flags": {},
|
||||
"order": 1,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
197
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "MASK",
|
||||
"type": "MASK",
|
||||
"links": null
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "LoadImage"
|
||||
},
|
||||
"widgets_values": [
|
||||
"image (25).png",
|
||||
"image"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 94,
|
||||
"type": "PreviewImage",
|
||||
"pos": [
|
||||
2907.259765625,
|
||||
-712.0502319335938
|
||||
],
|
||||
"size": [
|
||||
660.7576293945312,
|
||||
723.7353515625
|
||||
],
|
||||
"flags": {},
|
||||
"order": 3,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "images",
|
||||
"type": "IMAGE",
|
||||
"link": 199
|
||||
}
|
||||
],
|
||||
"outputs": [],
|
||||
"properties": {
|
||||
"Node name for S&R": "PreviewImage"
|
||||
},
|
||||
"widgets_values": []
|
||||
},
|
||||
{
|
||||
"id": 91,
|
||||
"type": "FluxProKontextMulti_fal",
|
||||
"pos": [
|
||||
2462.2119140625,
|
||||
-595.9410400390625
|
||||
],
|
||||
"size": [
|
||||
400,
|
||||
364
|
||||
],
|
||||
"flags": {},
|
||||
"order": 2,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "image_1",
|
||||
"type": "IMAGE",
|
||||
"link": 197
|
||||
},
|
||||
{
|
||||
"name": "image_2",
|
||||
"type": "IMAGE",
|
||||
"link": 198
|
||||
},
|
||||
{
|
||||
"name": "image_3",
|
||||
"shape": 7,
|
||||
"type": "IMAGE",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "image_4",
|
||||
"shape": 7,
|
||||
"type": "IMAGE",
|
||||
"link": null
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
199
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "FluxProKontextMulti_fal"
|
||||
},
|
||||
"widgets_values": [
|
||||
"Woman wearing this backpack on her way to jungle",
|
||||
"9:16",
|
||||
false,
|
||||
3.5,
|
||||
1,
|
||||
"2",
|
||||
"jpeg",
|
||||
false,
|
||||
2075416510,
|
||||
"randomize"
|
||||
]
|
||||
}
|
||||
],
|
||||
"links": [
|
||||
[
|
||||
197,
|
||||
92,
|
||||
0,
|
||||
91,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
198,
|
||||
93,
|
||||
0,
|
||||
91,
|
||||
1,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
199,
|
||||
91,
|
||||
0,
|
||||
94,
|
||||
0,
|
||||
"IMAGE"
|
||||
]
|
||||
],
|
||||
"groups": [],
|
||||
"config": {},
|
||||
"extra": {
|
||||
"ds": {
|
||||
"scale": 0.863837598531476,
|
||||
"offset": [
|
||||
-1745.427153953069,
|
||||
795.6355480141049
|
||||
]
|
||||
},
|
||||
"frontendVersion": "1.18.10",
|
||||
"ue_links": [],
|
||||
"VHS_latentpreview": false,
|
||||
"VHS_latentpreviewrate": 0,
|
||||
"VHS_MetadataImage": true,
|
||||
"VHS_KeepIntermediate": true
|
||||
},
|
||||
"version": 0.4
|
||||
}
|
||||
@@ -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,261 @@
|
||||
{
|
||||
"id": "80767774-d39b-4f73-a75a-3c1327f92316",
|
||||
"revision": 0,
|
||||
"last_node_id": 97,
|
||||
"last_link_id": 202,
|
||||
"nodes": [
|
||||
{
|
||||
"id": 92,
|
||||
"type": "LoadImage",
|
||||
"pos": [
|
||||
1846.3321533203125,
|
||||
-654.7257080078125
|
||||
],
|
||||
"size": [
|
||||
270,
|
||||
510.1016845703125
|
||||
],
|
||||
"flags": {},
|
||||
"order": 0,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
200
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "MASK",
|
||||
"type": "MASK",
|
||||
"links": null
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "LoadImage"
|
||||
},
|
||||
"widgets_values": [
|
||||
"image (25).png",
|
||||
"image"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 96,
|
||||
"type": "LoadVideoURL",
|
||||
"pos": [
|
||||
2672.1376953125,
|
||||
-645.724609375
|
||||
],
|
||||
"size": [
|
||||
270,
|
||||
266
|
||||
],
|
||||
"flags": {
|
||||
"collapsed": false
|
||||
},
|
||||
"order": 2,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "url",
|
||||
"type": "STRING",
|
||||
"widget": {
|
||||
"name": "url"
|
||||
},
|
||||
"link": 201
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "frames",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
202
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "frame_count",
|
||||
"type": "INT",
|
||||
"links": null
|
||||
},
|
||||
{
|
||||
"name": "video_info",
|
||||
"type": "VHS_VIDEOINFO",
|
||||
"links": null
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "LoadVideoURL"
|
||||
},
|
||||
"widgets_values": [
|
||||
"https://example.com/video.mp4",
|
||||
0,
|
||||
"Disabled",
|
||||
512,
|
||||
512,
|
||||
0,
|
||||
0,
|
||||
1
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 97,
|
||||
"type": "VHS_VideoCombine",
|
||||
"pos": [
|
||||
2995.218994140625,
|
||||
-646.7775268554688
|
||||
],
|
||||
"size": [
|
||||
215.01171875,
|
||||
670.6875
|
||||
],
|
||||
"flags": {},
|
||||
"order": 3,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "images",
|
||||
"type": "IMAGE",
|
||||
"link": 202
|
||||
},
|
||||
{
|
||||
"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
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "Filenames",
|
||||
"type": "VHS_FILENAMES",
|
||||
"links": null
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "VHS_VideoCombine"
|
||||
},
|
||||
"widgets_values": {
|
||||
"frame_rate": 25,
|
||||
"loop_count": 0,
|
||||
"filename_prefix": "Veo2",
|
||||
"format": "video/h265-mp4",
|
||||
"pix_fmt": "yuv420p10le",
|
||||
"crf": 22,
|
||||
"save_metadata": false,
|
||||
"pingpong": false,
|
||||
"save_output": true,
|
||||
"videopreview": {
|
||||
"hidden": false,
|
||||
"paused": false,
|
||||
"params": {
|
||||
"filename": "Veo2_00001.mp4",
|
||||
"subfolder": "",
|
||||
"type": "output",
|
||||
"format": "video/h265-mp4",
|
||||
"frame_rate": 25,
|
||||
"workflow": "Veo2_00001.png",
|
||||
"fullpath": "D:\\ComfyUI_windows_portable\\ComfyUI\\output\\Veo2_00001.mp4"
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 95,
|
||||
"type": "Veo2ImageToVideo_fal",
|
||||
"pos": [
|
||||
2199.614990234375,
|
||||
-652.0396118164062
|
||||
],
|
||||
"size": [
|
||||
400,
|
||||
200
|
||||
],
|
||||
"flags": {},
|
||||
"order": 1,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "image",
|
||||
"type": "IMAGE",
|
||||
"link": 200
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "STRING",
|
||||
"type": "STRING",
|
||||
"links": [
|
||||
201
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "Veo2ImageToVideo_fal"
|
||||
},
|
||||
"widgets_values": [
|
||||
"Woman slowly turning",
|
||||
"auto",
|
||||
"5s"
|
||||
]
|
||||
}
|
||||
],
|
||||
"links": [
|
||||
[
|
||||
200,
|
||||
92,
|
||||
0,
|
||||
95,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
201,
|
||||
95,
|
||||
0,
|
||||
96,
|
||||
0,
|
||||
"STRING"
|
||||
],
|
||||
[
|
||||
202,
|
||||
96,
|
||||
0,
|
||||
97,
|
||||
0,
|
||||
"IMAGE"
|
||||
]
|
||||
],
|
||||
"groups": [],
|
||||
"config": {},
|
||||
"extra": {
|
||||
"ds": {
|
||||
"scale": 0.863837598531476,
|
||||
"offset": [
|
||||
-1694.469448600796,
|
||||
792.8487236845591
|
||||
]
|
||||
},
|
||||
"frontendVersion": "1.18.10",
|
||||
"ue_links": [],
|
||||
"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 = set(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 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,250 @@
|
||||
"""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:
|
||||
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 _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]:
|
||||
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
|
||||
|
||||
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)
|
||||
if 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 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
|
||||
|
||||
|
||||
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 = _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 flagged as superseded within their family)",
|
||||
len(all_classes),
|
||||
skipped,
|
||||
featured_count,
|
||||
flagged,
|
||||
)
|
||||
_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,227 @@
|
||||
"""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(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")),
|
||||
}
|
||||
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
+2187
-335
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
-37
@@ -1,56 +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:
|
||||
fal_key = config['API']['FAL_KEY']
|
||||
os.environ["FAL_KEY"] = fal_key
|
||||
except KeyError:
|
||||
print("Error: FAL_KEY not found in config.ini")
|
||||
|
||||
# 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,459 @@
|
||||
"""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 after the refresh finishes: new nodes register at import time."
|
||||
_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]:
|
||||
"""Run scripts/build_registry.py; returns (ok, message)."""
|
||||
import subprocess
|
||||
import sys
|
||||
|
||||
root = _repo_root()
|
||||
command = [
|
||||
sys.executable,
|
||||
os.path.join(root, "scripts", "build_registry.py"),
|
||||
"--out",
|
||||
os.path.join("data", "fal_registry.json"),
|
||||
]
|
||||
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"build_registry.py exited with {completed.returncode}: {tail}"
|
||||
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()
|
||||
+465
-91
@@ -1,84 +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:
|
||||
fal_key = config['API']['FAL_KEY']
|
||||
os.environ["FAL_KEY"] = fal_key
|
||||
except KeyError:
|
||||
print("Error: FAL_KEY not found in config.ini")
|
||||
|
||||
# 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 = {
|
||||
@@ -88,89 +116,435 @@ 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": "",
|
||||
"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": "",
|
||||
"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,
|
||||
):
|
||||
try:
|
||||
if not training_data_url:
|
||||
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,
|
||||
}
|
||||
|
||||
if trigger_phrase:
|
||||
arguments["trigger_phrase"] = trigger_phrase
|
||||
|
||||
# Submit training job
|
||||
result = ApiHandler.submit_and_get_result("fal-ai/wan-trainer", arguments)
|
||||
lora_url = result["lora_file"]["url"]
|
||||
return (lora_url,)
|
||||
|
||||
except Exception as e:
|
||||
return ApiHandler.handle_text_generation_error("wan-trainer", e)
|
||||
|
||||
|
||||
class LtxVideoTrainerNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"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": "",
|
||||
"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,
|
||||
):
|
||||
try:
|
||||
if not training_data_url:
|
||||
return ApiHandler.handle_text_generation_error(
|
||||
"ltx-video-trainer", "No training data URL provided"
|
||||
)
|
||||
|
||||
# Prepare arguments for the API
|
||||
arguments = {
|
||||
"training_data_url": training_data_url,
|
||||
"rank": int(rank),
|
||||
"number_of_steps": number_of_steps,
|
||||
"number_of_frames": number_of_frames,
|
||||
"frame_rate": frame_rate,
|
||||
"resolution": resolution,
|
||||
"aspect_ratio": aspect_ratio,
|
||||
"learning_rate": learning_rate,
|
||||
"auto_scale_input": auto_scale_input,
|
||||
"split_input_into_scenes": split_input_into_scenes,
|
||||
"split_input_duration_threshold": split_input_duration_threshold,
|
||||
"validation_negative_prompt": validation_negative_prompt,
|
||||
"validation_number_of_frames": validation_number_of_frames,
|
||||
"validation_resolution": validation_resolution,
|
||||
"validation_aspect_ratio": validation_aspect_ratio,
|
||||
"validation_reverse": validation_reverse,
|
||||
}
|
||||
|
||||
if trigger_phrase:
|
||||
arguments["trigger_phrase"] = trigger_phrase
|
||||
|
||||
# Submit training job
|
||||
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:
|
||||
return ApiHandler.handle_text_generation_error("ltx-video-trainer", e)
|
||||
|
||||
|
||||
# Node class mappings
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"FluxLoraTrainer_fal": FluxLoraTrainerNode,
|
||||
"HunyuanVideoLoraTrainer_fal": HunyuanVideoLoraTrainerNode,
|
||||
"WanLoraTrainer_fal": WanLoraTrainerNode,
|
||||
"LtxVideoTrainer_fal": LtxVideoTrainerNode,
|
||||
}
|
||||
|
||||
# Node display name mappings
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"FluxLoraTrainer_fal": "Flux LoRA Trainer (fal)",
|
||||
"HunyuanVideoLoraTrainer_fal": "Hunyuan Video LoRA Trainer (fal)",
|
||||
}
|
||||
"WanLoraTrainer_fal": "WAN LoRA Trainer (fal)",
|
||||
"LtxVideoTrainer_fal": "LTX Video LoRA Trainer (fal)",
|
||||
}
|
||||
|
||||
+486
-107
@@ -1,142 +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:
|
||||
fal_key = config['API']['FAL_KEY']
|
||||
os.environ["FAL_KEY"] = fal_key
|
||||
except KeyError:
|
||||
print("Error: FAL_KEY not found in config.ini")
|
||||
|
||||
# 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: local registry is up to date")
|
||||
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,248 @@
|
||||
"""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,)
|
||||
query += " ORDER BY submitted_at 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,282 @@
|
||||
"""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 _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 download_url_to_temp(url: str, suffix: str) -> str:
|
||||
"""Stream a URL to a temp file and return its local path."""
|
||||
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)
|
||||
+3669
-400
File diff suppressed because it is too large
Load Diff
+103
-65
@@ -1,38 +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:
|
||||
fal_key = config['API']['FAL_KEY']
|
||||
os.environ["FAL_KEY"] = fal_key
|
||||
except KeyError:
|
||||
print("Error: FAL_KEY not found in config.ini")
|
||||
|
||||
# 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'.",
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -40,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 = {
|
||||
@@ -97,4 +135,4 @@ NODE_CLASS_MAPPINGS = {
|
||||
# Node display name mappings
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"VLM_fal": "VLM (fal)",
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,34 @@
|
||||
[project]
|
||||
name = "fal-api"
|
||||
description = "Custom nodes for using fal API with auto-generated full-catalog coverage of fal.ai models. Video generation with Kling, Runway, Luma. Image generation with Flux. LLMs and VLMs OpenAI, Claude, Llama and Gemini."
|
||||
version = "2.5.0"
|
||||
license = {file = "LICENSE"}
|
||||
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"
|
||||
# Used by Comfy Registry https://comfyregistry.org
|
||||
|
||||
[tool.comfy]
|
||||
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,159 @@
|
||||
#!/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 auto-generated model node in [ComfyUI-fal-API](README.md), grouped by
|
||||
category (largest first). Click a category to expand it.
|
||||
|
||||
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"]
|
||||
model_count = registry.get("model_count", len(models))
|
||||
published = [str(m.get("published_at", "")) for m in registry.get("models", [])]
|
||||
generated_date = max(published)[:10] if any(published) else "unknown"
|
||||
summary = (
|
||||
f"{model_count} models · newest model {generated_date} · "
|
||||
"refresh with `scripts/build_registry.py`"
|
||||
)
|
||||
blocks = [
|
||||
render_category(category, grouped)
|
||||
for category, grouped in group_by_category(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(f"MODELS.md already up to date ({registry.get('model_count')} models)")
|
||||
return 0
|
||||
|
||||
MODELS_PATH.write_text(updated, encoding="utf-8")
|
||||
print(
|
||||
f"MODELS.md model catalog regenerated: {registry.get('model_count')} models, "
|
||||
f"{len(group_by_category(registry['models']))} categories"
|
||||
)
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
@@ -0,0 +1,646 @@
|
||||
#!/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 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_INPUT_PROPERTIES = 40
|
||||
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]
|
||||
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 non_null_branches(branches, components):
|
||||
"""Resolve and drop null branches from an anyOf/oneOf list."""
|
||||
resolved = [resolve_ref(branch, components) for branch in branches if isinstance(branch, dict)]
|
||||
return [branch for branch in resolved if branch.get("type") != "null"]
|
||||
|
||||
|
||||
def merge_all_of(schema, components):
|
||||
"""Merge an allOf list (one level), with sibling keys taking precedence."""
|
||||
merged = {}
|
||||
for branch in schema.get("allOf", []):
|
||||
if isinstance(branch, dict):
|
||||
merged = {**merged, **resolve_ref(branch, components)}
|
||||
siblings = {key: value for key, value in schema.items() if key != "allOf"}
|
||||
return {**merged, **siblings}
|
||||
|
||||
|
||||
def is_custom_size_pair(branches):
|
||||
"""Detect the image_size pattern: [enum-of-presets, width/height object]."""
|
||||
enum_branch = next((b for b in branches if b.get("enum")), None)
|
||||
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):
|
||||
"""Resolve $ref / allOf / anyOf / oneOf one level.
|
||||
|
||||
Returns (resolved_schema, has_custom_size, custom_size_enum_values).
|
||||
"""
|
||||
if not isinstance(schema, dict):
|
||||
return {}, False, None
|
||||
resolved = resolve_ref(schema, components)
|
||||
if "allOf" in resolved:
|
||||
resolved = merge_all_of(resolved, components)
|
||||
branches_key = "anyOf" if "anyOf" in resolved else ("oneOf" if "oneOf" in resolved else None)
|
||||
if branches_key is None:
|
||||
return resolved, False, None
|
||||
|
||||
branches = non_null_branches(resolved[branches_key], components)
|
||||
siblings = {key: value for key, value in resolved.items() if key != branches_key}
|
||||
if not branches:
|
||||
return siblings, False, None
|
||||
|
||||
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
|
||||
|
||||
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 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}
|
||||
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."""
|
||||
properties = schema.get("properties", {})
|
||||
required_names = set(schema.get("required", []))
|
||||
names = ordered_property_names(schema)
|
||||
|
||||
if len(names) > MAX_INPUT_PROPERTIES:
|
||||
required_first = [n for n in names if n in required_names]
|
||||
optional = [n for n in names if n not in required_names]
|
||||
budget = max(MAX_INPUT_PROPERTIES - len(required_first), 0)
|
||||
names = required_first + optional[:budget]
|
||||
logger.info(
|
||||
"%s: input schema has %d properties, capped to %d",
|
||||
endpoint_id,
|
||||
len(properties),
|
||||
len(names),
|
||||
)
|
||||
|
||||
inputs = []
|
||||
for name in names:
|
||||
record = distill_property(name, properties.get(name, {}), required_names, components)
|
||||
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] 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", {})
|
||||
referenced = path_lookup(doc)
|
||||
if referenced and referenced in components:
|
||||
return components[referenced]
|
||||
|
||||
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 components[max(pool, key=len)]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 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 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")
|
||||
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()
|
||||
|
||||
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]
|
||||
|
||||
records = sorted(records, key=lambda record: record["endpoint_id"])
|
||||
|
||||
# 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),
|
||||
"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", len(records), args.out)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,72 @@
|
||||
"""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"
|
||||
|
||||
# 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")
|
||||
@@ -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,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,64 @@
|
||||
"""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):
|
||||
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,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,78 @@
|
||||
"""Validate the committed model registry — pure JSON, no heavy imports."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
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"])
|
||||
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"
|
||||
)
|
||||
@@ -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,117 @@
|
||||
"""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_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"
|
||||
@@ -0,0 +1,50 @@
|
||||
"""Unit tests for the /fal_api server routes' pure functions."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
|
||||
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
|
||||
@@ -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,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,54 @@
|
||||
// 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";
|
||||
|
||||
// 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);
|
||||
}
|
||||
},
|
||||
|
||||
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,313 @@
|
||||
// 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 to load 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);
|
||||
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");
|
||||
return;
|
||||
}
|
||||
await pollRefresh(view, button);
|
||||
} catch (error) {
|
||||
console.debug("[fal] registry refresh failed", error);
|
||||
registryDone(view, false, "refresh request failed");
|
||||
}
|
||||
}
|
||||
|
||||
function renderRegistry(view, status) {
|
||||
try {
|
||||
const count = Number(status?.new_count) || 0;
|
||||
if (count <= 0) {
|
||||
view.registry.replaceChildren(element("div", "fal-muted", "Registry is up to date."));
|
||||
return;
|
||||
}
|
||||
const box = element("div", "fal-registry-news");
|
||||
box.append(
|
||||
element("div", "fal-registry-count", `${count} new model${count === 1 ? "" : "s"} on fal`)
|
||||
);
|
||||
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 {
|
||||
view.registry.replaceChildren(element("div", "fal-muted", "Registry status unavailable."));
|
||||
} 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);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user