Compare commits

..
Author SHA1 Message Date
Talmaj Marinc a7b63915dc Improve logging for failed download links.
Build package / Build Test (3.13) (push) Has been cancelled
Build package / Build Test (3.14) (push) Has been cancelled
Python Linting / Run Ruff (push) Has been cancelled
Python Linting / Run Pylint (push) Has been cancelled
Build package / Build Test (3.10) (push) Has been cancelled
Build package / Build Test (3.11) (push) Has been cancelled
Build package / Build Test (3.12) (push) Has been cancelled
2026-07-13 00:24:37 +02:00
Talmaj Marinc 5d8b3690ff Add read-repos to hf oauth scope. 2026-07-13 00:24:37 +02:00
Talmaj Marinc 515d73f6f2 Update tests. 2026-07-13 00:24:37 +02:00
Talmaj Marinc a14cc66712 Add obfuscation for stored creds. 2026-07-13 00:24:37 +02:00
Talmaj Marinc 09d55ba54a Fix the port for oauth server. 2026-07-13 00:24:37 +02:00
Talmaj Marinc c62c727aac Merge migrations. 2026-07-13 00:24:37 +02:00
Talmaj Marinc 71cf5a11f1 Remove api creds and add oauth. 2026-07-13 00:24:37 +02:00
Talmaj Marinc ee06f45ff3 Add small improvements and query optimization. 2026-07-13 00:24:37 +02:00
Talmaj Marinc 08f0b15d60 Allow whitespaces around download urls, strip them in the backend. 2026-07-13 00:24:37 +02:00
Talmaj Marinc cebe350e0e Add distinction in error messaging for gated models. 2026-07-13 00:24:37 +02:00
Talmaj Marinc 9807d2a743 Add extension check on the final resolved url -> fix downloading from civitAI. 2026-07-13 00:24:37 +02:00
Talmaj Marinc 88be6cd111 Add delete and clear all downloads funcitonalities. 2026-07-13 00:24:37 +02:00
Talmaj Marinc 0ea2773122 Disable newline translations on Windows, \r\n -> \n only. 2026-07-13 00:24:37 +02:00
Talmaj Marinc 589eb6e1d9 Add support for ENV based HF_TOKEN. 2026-07-13 00:24:37 +02:00
Talmaj Marinc 317317c98a Fix an issue with windows lacking have os.pwrite 2026-07-13 00:24:37 +02:00
Talmaj Marinc c5d9d39d0c Fix running CI tests. 2026-07-13 00:24:37 +02:00
Talmaj Marinc ea1e5786b8 Fix more AI detected issues., 2026-07-13 00:24:37 +02:00
Talmaj Marinc 78ca58cf29 Fix ruff. 2026-07-13 00:24:37 +02:00
Talmaj Marinc 5d0835deee Update openapi.yaml. 2026-07-13 00:24:37 +02:00
Talmaj Marinc dc582b0fe2 Remove sending url info over websockets for model downloads. 2026-07-13 00:24:37 +02:00
Talmaj Marinc 343afeb6d0 Fix sweep deleting FAILED partials and fix segmented resume path trusted offsets blindly. 2026-07-13 00:24:37 +02:00
Talmaj Marinc fe20a6483e Add _positive_int in cli_args arguments. 2026-07-13 00:24:37 +02:00
Talmaj Marinc e7b60beaf1 Redact urls in logging and fix concurrent enqueue issue that could corrupt the downloaded files. 2026-07-13 00:24:37 +02:00
Talmaj Marinc ce237f9a79 Simplify docstrings. 2026-07-13 00:24:37 +02:00
Talmaj Marinc 1d79c41064 Normalize malformed safetensors headers into StructuralError. 2026-07-13 00:24:36 +02:00
Talmaj Marinc ccadcaf4e3 Prevent redirects to loopback/internal IPs (SSRF) 2026-07-13 00:24:36 +02:00
Talmaj Marinc 2fc66a5e8d Simplify docstrings. 2026-07-13 00:24:36 +02:00
Talmaj Marinc 15aa401146 Fix IPv4-mapped IPv6 addresses 2026-07-13 00:24:36 +02:00
Talmaj Marinc 4c152491f0 Restrict cleartext HTTP redirects to explicit loopback/dev hosts. 2026-07-13 00:24:36 +02:00
Talmaj Marinc 72ce283d32 Simplify docstrings. 2026-07-13 00:24:36 +02:00
Talmaj Marinc bd27ae82a1 Don't echo full URLs and raw exception text from probe failures. 2026-07-13 00:24:36 +02:00
Talmaj Marinc 9f9a9272ce Simplify docstrings. 2026-07-13 00:24:36 +02:00
Talmaj Marinc a0b35f60be Handle short pwrite() results. 2026-07-13 00:24:36 +02:00
Talmaj Marinc d0e71574a1 Switch to asyncio.to_thread for db calls in job.py 2026-07-13 00:24:36 +02:00
Talmaj Marinc 8a32df7108 Improve _finalize checks for downloads. 2026-07-13 00:24:36 +02:00
Talmaj Marinc 03e6b57a86 Truncate file to 0 before restarting. 2026-07-13 00:24:36 +02:00
Talmaj Marinc 391c628b8e Add max-download-size in case the server tries to send larger files than it reports. 2026-07-13 00:24:36 +02:00
Talmaj Marinc 790e6775da Clear error when the job leaves a failure state. 2026-07-13 00:24:36 +02:00
Talmaj Marinc 3b5483a225 Fix resuming of segmented download. 2026-07-13 00:24:36 +02:00
Talmaj Marinc d73483be0e Simplify docstrings. 2026-07-13 00:24:36 +02:00
Talmaj Marinc acf8f95eb5 Fix potential concurrency and db integrity error. 2026-07-13 00:24:36 +02:00
Talmaj Marinc 2bed0a5d73 Fix an issue when a credential is deleted, resuming of download fails. 2026-07-13 00:24:36 +02:00
Talmaj Marinc 96a3e0ccad Simplify docstrings. 2026-07-13 00:24:36 +02:00
Talmaj Marinc ea583ae155 Fix sercret_last4 2026-07-13 00:24:36 +02:00
Talmaj Marinc cbb0ddd87e Improve normalize_host 2026-07-13 00:24:36 +02:00
Talmaj Marinc b107b834f2 Docstring simplification. 2026-07-13 00:24:36 +02:00
Talmaj Marinc 53da5f94e6 Fix url parsing. 2026-07-13 00:24:36 +02:00
Talmaj Marinc cca6687a0a Simplify migration docstring. 2026-07-13 00:24:36 +02:00
Talmaj Marinc 6f4c742bc5 Add initial commit for model downloader. 2026-07-13 00:24:36 +02:00
84 changed files with 6116 additions and 5359 deletions
+2 -2
View File
@@ -23,9 +23,9 @@ jobs:
# SHA-pinned per zizmor `unpinned-uses: hash-pin`. Bump this SHA to pick up
# upstream changes; keep `workflows_ref` matching so prompts/scripts load
# from the same commit as the workflow definition.
uses: Comfy-Org/github-workflows/.github/workflows/cursor-review.yml@964d5aad37cbfb57c5b23961d42c2fd85868bf1d # github-workflows main (964d5aa)
uses: Comfy-Org/github-workflows/.github/workflows/cursor-review.yml@047ca48febe3a6647608ed2e0c4331b491cb9d6a # github-workflows#9
with:
workflows_ref: 964d5aad37cbfb57c5b23961d42c2fd85868bf1d
workflows_ref: 047ca48febe3a6647608ed2e0c4331b491cb9d6a
diff_excludes: >-
:!**/.claude/**
:!**/dist/**
-41
View File
@@ -19,9 +19,6 @@
better to remove a broken feature path than keep a complicated partial fix.
- Preserve existing APIs, node names, model-loading behavior, file layout, and
workflow compatibility unless the change is explicitly about replacing them.
- When compatibility is explicitly out of scope, remove compatibility-only
aliases, duplicate nodes, legacy entry points, and preset wrappers instead of
retaining parallel ways to perform the same operation.
- Code must look hand-written for this repository. Changes that read like
generic AI-generated code will be rejected automatically: unnecessary helper
layers, vague names, boilerplate comments, defensive branches without a real
@@ -99,13 +96,6 @@
unless they are read by current code and change current behavior. Remove
pass-through or stored-but-unused values instead of preserving upstream or
deprecated API baggage.
- Do not add a model-specific option to a shared helper when only one caller
needs it. Keep one-off behavior at the model integration boundary, or extend
the shared helper only when the option is a coherent reusable capability.
- Implementations of shared model interfaces should accept the standard caller
contract without model-specific rejection branches for optional capabilities
they do not consume. Let supported behavior be determined by implementation
paths that actually use those inputs.
- If an implementation needs auxiliary values for its own workflow, expose them
through a private helper or a clearly named implementation-specific method
instead of overloading the public method's return contract.
@@ -164,10 +154,6 @@
`comfy-kitchen` helpers where they already solve the problem.
- Use optimized comfy-kitchen ops in places where they improve performance
without changing the expected dtype, device, memory, or interface behavior.
- Prefer ComfyUI's shared optimized kernels and backend dispatchers over
handwritten implementations of the same operation. Remove duplicate local
kernels and adapt inputs to the shared operation's documented layout while
preserving the model's original math and output contract.
- All models should use the optimized attention function selected by ComfyUI.
Treat optimized backend functions, dispatch helpers, and capability-selected
callables as opaque. Higher-level code must not inspect function identity,
@@ -190,12 +176,6 @@
- Model detection code that inspects linear weight shapes should only use the
first dimension. The second dimension may be half the original size for
NVFP4 or other 4-bit quantized models.
- A model-detection signature must guard every state-dict key it dereferences.
Do not partially match a format and then raise an incidental `KeyError` while
extracting its configuration.
- Order model-detection checks from established or more-specific signatures to
newer or broader signatures. Put a broad new detector near the generic
fallback when giving it higher precedence could steal another model family.
- Avoid adding `einops` usage in core inference code. Use native torch tensor
ops such as `reshape`, `view`, `permute`, `transpose`, `flatten`, `unflatten`,
`unsqueeze`, and `squeeze` instead.
@@ -212,23 +192,11 @@
methods for scalar or structural calculations.
- Avoid unnecessary casts and transfers. Preserve the intended compute dtype,
storage dtype, bias dtype, and original tensor shape metadata.
- Do not cast the result of an optimized backend operation back to its input
dtype unless that backend's documented result contract requires normalization.
In particular, trust the selected optimized-attention implementation to honor
its dtype contract.
- Keep model-native latent layout handling inside the model or latent-format
owner, not in helper nodes. Do not collapse, expand, pack, or unpack latent
dimensions in nodes or other caller-side adapters just to satisfy a model
forward; the model path should consume and return the native latent shape for
that model family.
- DiT models should accept latent dimensions that are not exact patch-size
multiples. Use `comfy.ldm.common_dit.pad_to_patch_size` on every patchified
target or reference input, then crop only the target output back to its
original dimensions.
- Avoid defensive shape and configuration checks that merely replace the clear
failure from the tensor operation immediately below them. Add explicit
validation only when it provides materially better context at a real boundary
or prevents silent incorrect output.
- Assume inputs to the main model forward are already in the compute dtype by
default, except integer inputs such as some model timestep tensors. Do not add
defensive or convenience casts in model code; it is better for invalid dtype
@@ -292,15 +260,6 @@
- Model implementations should add the minimal number of ComfyUI nodes required
to run the model. Reuse existing nodes as much as possible; adapting the model
to work with existing nodes is strongly preferred over creating new nodes.
- Use `io.Autogrow` for a variable number of repeated inputs instead of a fixed
series of numbered optional sockets. Set its minimum to zero when the model
has a valid no-item path, and cap it only when the model has a real limit.
- Mark inputs optional when execution has a valid path that does not read them.
If one optional input is needed only to process another optional input, do not
force users on the path that supplies neither to connect it.
- Conditioning nodes should normally output conditioning only. Do not expose
input or intermediate images as convenience outputs for downstream sizing or
routing; use the existing image path or a dedicated image operation instead.
- Nodes should output only values they own. Do not add pass-through outputs for
workflow convenience unless the node is explicitly an output node. Existing
models, latents, conditioning, or other inputs should flow directly to the
-1
View File
@@ -1,6 +1,5 @@
* @comfyanonymous @kosinkadink @guill @alexisrolland @rattus128 @kijai
/CODEOWNERS @comfyanonymous
/AGENTS.md @comfyanonymous
/.ci/ @comfyanonymous
/.github/ @comfyanonymous
@@ -0,0 +1,84 @@
"""
Download manager schema.
Adds the two tables that back the server-side model download manager:
transient job/queue state (``downloads`` + per-segment ``download_segments``).
Revision ID: 0005_download_manager
Revises: 0004_drop_tag_type
Create Date: 2026-06-27
"""
from alembic import op
import sqlalchemy as sa
revision = "0005_download_manager"
down_revision = "0004_drop_tag_type"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.create_table(
"downloads",
sa.Column("id", sa.String(length=36), primary_key=True),
sa.Column("url", sa.Text(), nullable=False),
sa.Column("final_url", sa.Text(), nullable=True),
sa.Column("model_id", sa.String(length=1024), nullable=False),
sa.Column("dest_path", sa.Text(), nullable=False),
sa.Column("temp_path", sa.Text(), nullable=False),
sa.Column("status", sa.String(length=16), nullable=False),
sa.Column("priority", sa.Integer(), nullable=False, server_default="0"),
sa.Column("total_bytes", sa.BigInteger(), nullable=True),
sa.Column("bytes_done", sa.BigInteger(), nullable=False, server_default="0"),
sa.Column("etag", sa.String(length=512), nullable=True),
sa.Column("last_modified", sa.String(length=128), nullable=True),
sa.Column(
"accept_ranges", sa.Boolean(), nullable=False, server_default=sa.text("false")
),
sa.Column("expected_sha256", sa.String(length=64), nullable=True),
sa.Column(
"allow_any_extension",
sa.Boolean(),
nullable=False,
server_default=sa.text("false"),
),
sa.Column("attempts", sa.Integer(), nullable=False, server_default="0"),
sa.Column("error", sa.Text(), nullable=True),
sa.Column("created_at", sa.BigInteger(), nullable=False),
sa.Column("updated_at", sa.BigInteger(), nullable=False),
sa.CheckConstraint("bytes_done >= 0", name="ck_downloads_bytes_done_nonneg"),
sa.CheckConstraint(
"total_bytes IS NULL OR total_bytes >= 0",
name="ck_downloads_total_bytes_nonneg",
),
)
op.create_index("ix_downloads_status", "downloads", ["status"])
op.create_index("ix_downloads_priority", "downloads", ["priority"])
op.create_index("ix_downloads_model_id", "downloads", ["model_id"])
op.create_table(
"download_segments",
sa.Column(
"download_id",
sa.String(length=36),
sa.ForeignKey("downloads.id", ondelete="CASCADE"),
nullable=False,
),
sa.Column("idx", sa.Integer(), nullable=False),
sa.Column("start_offset", sa.BigInteger(), nullable=False),
sa.Column("end_offset", sa.BigInteger(), nullable=False),
sa.Column("bytes_done", sa.BigInteger(), nullable=False, server_default="0"),
sa.PrimaryKeyConstraint("download_id", "idx", name="pk_download_segments"),
sa.CheckConstraint("bytes_done >= 0", name="ck_segments_bytes_done_nonneg"),
sa.CheckConstraint("end_offset >= start_offset", name="ck_segments_range"),
)
def downgrade() -> None:
op.drop_table("download_segments")
op.drop_index("ix_downloads_model_id", table_name="downloads")
op.drop_index("ix_downloads_priority", table_name="downloads")
op.drop_index("ix_downloads_status", table_name="downloads")
op.drop_table("downloads")
+6 -4
View File
@@ -4,7 +4,8 @@ import shutil
from app.logger import log_startup_warning
from utils.install_util import get_missing_requirements_message
from filelock import FileLock, Timeout
from comfy.cli_args import args
# Import the module so tests that reload comfy.cli_args see the live object.
import comfy.cli_args
_DB_AVAILABLE = False
Session = None
@@ -21,6 +22,7 @@ try:
from app.database.models import Base
import app.assets.database.models # noqa: F401 — register models with Base.metadata
import app.model_downloader.database.models # noqa: F401 — register models with Base.metadata
_DB_AVAILABLE = True
except ImportError as e:
@@ -57,13 +59,13 @@ def get_alembic_config():
config = Config(config_path)
config.set_main_option("script_location", scripts_path)
config.set_main_option("sqlalchemy.url", args.database_url)
config.set_main_option("sqlalchemy.url", comfy.cli_args.args.database_url)
return config
def get_db_path():
url = args.database_url
url = comfy.cli_args.args.database_url
if url.startswith("sqlite:///"):
return url.split("///")[1]
else:
@@ -97,7 +99,7 @@ def _is_memory_db(db_url):
def init_db():
db_url = args.database_url
db_url = comfy.cli_args.args.database_url
logging.debug(f"Database URL: {db_url}")
if _is_memory_db(db_url):
+202
View File
@@ -0,0 +1,202 @@
"""aiohttp routes for the download manager.
Endpoint surface (all under ``/api/download``), mirroring the response
envelope used by ``app/assets/api/routes.py``:
POST /api/download/enqueue
GET /api/download
POST /api/download/availability
POST /api/download/clear
GET /api/download/auth
POST /api/download/auth/{provider}/login
POST /api/download/auth/{provider}/logout
GET /api/download/{id}
DELETE /api/download/{id}
POST /api/download/{id}/pause
POST /api/download/{id}/resume
POST /api/download/{id}/cancel
POST /api/download/{id}/priority
Note on ordering: the static ``auth`` routes are registered before the dynamic
``/api/download/{id}`` route so a request to ``.../auth`` is not captured as
``id == "auth"``.
"""
from __future__ import annotations
import json
from aiohttp import web
from pydantic import BaseModel, ValidationError
from app.model_downloader.api import schemas_in, schemas_out
from app.model_downloader.auth.oauth import LoginInProgress, OAuthNotConfigured
from app.model_downloader.auth.providers import PROVIDERS
from app.model_downloader.auth.store import AUTH_STORE
from app.model_downloader.manager import DOWNLOAD_MANAGER, DownloadError
ROUTES = web.RouteTableDef()
def register_routes(app: web.Application) -> None:
"""Wire the download-manager routes into the running aiohttp app."""
app.add_routes(ROUTES)
# ----- envelope helpers (same shape as app/assets/api/routes.py) -----
def _error(status: int, code: str, message: str, details: dict | None = None) -> web.Response:
return web.json_response(
{"error": {"code": code, "message": message, "details": details or {}}},
status=status,
)
def _ok(payload, status: int = 200) -> web.Response:
return web.json_response(payload, status=status)
async def _parse(request: web.Request, model: type[BaseModel]):
try:
raw = await request.json()
except json.JSONDecodeError:
return _error(400, "INVALID_JSON", "Request body must be valid JSON.")
try:
return model.model_validate(raw)
except ValidationError as ve:
return _error(400, "INVALID_BODY", "Validation failed.", {"errors": json.loads(ve.json())})
def _from_download_error(e: DownloadError) -> web.Response:
return _error(e.http_status, e.code, e.message)
# ----- downloads: collection + enqueue + availability -----
@ROUTES.post("/api/download/enqueue")
async def enqueue(request: web.Request) -> web.Response:
parsed = await _parse(request, schemas_in.EnqueueRequest)
if isinstance(parsed, web.Response):
return parsed
try:
download_id = await DOWNLOAD_MANAGER.enqueue(
parsed.url,
parsed.model_id,
priority=parsed.priority,
expected_sha256=parsed.expected_sha256,
allow_any_extension=parsed.allow_any_extension,
)
except DownloadError as e:
return _from_download_error(e)
return _ok({"download_id": download_id, "accepted": True}, status=202)
@ROUTES.get("/api/download")
async def list_downloads(request: web.Request) -> web.Response:
return _ok({"downloads": await DOWNLOAD_MANAGER.list()})
@ROUTES.post("/api/download/availability")
async def availability(request: web.Request) -> web.Response:
parsed = await _parse(request, schemas_in.AvailabilityRequest)
if isinstance(parsed, web.Response):
return parsed
return _ok({"models": await DOWNLOAD_MANAGER.availability(parsed.models)})
@ROUTES.post("/api/download/clear")
async def clear(request: web.Request) -> web.Response:
deleted = await DOWNLOAD_MANAGER.clear()
return _ok({"deleted": deleted})
# ----- auth (OAuth login + env-key status) — must precede /{id} -----
@ROUTES.get("/api/download/auth")
async def auth_status(request: web.Request) -> web.Response:
return _ok({"providers": schemas_out.auth_status()})
@ROUTES.post("/api/download/auth/{provider}/login")
async def auth_login(request: web.Request) -> web.Response:
provider = PROVIDERS.get(request.match_info["provider"])
if provider is None:
return _error(400, "UNKNOWN_PROVIDER", "No such auth provider.")
try:
authorize_url = await AUTH_STORE.begin_login(provider)
except OAuthNotConfigured as e:
return _error(400, "OAUTH_NOT_CONFIGURED", str(e))
except LoginInProgress as e:
return _error(409, "LOGIN_IN_PROGRESS", str(e))
return _ok({"authorize_url": authorize_url})
@ROUTES.post("/api/download/auth/{provider}/logout")
async def auth_logout(request: web.Request) -> web.Response:
provider = PROVIDERS.get(request.match_info["provider"])
if provider is None:
return _error(400, "UNKNOWN_PROVIDER", "No such auth provider.")
AUTH_STORE.clear(provider.name)
return _ok({"logged_out": True})
# ----- single download by id (dynamic; registered last) -----
@ROUTES.get("/api/download/{id}")
async def get_download(request: web.Request) -> web.Response:
view = await DOWNLOAD_MANAGER.status(request.match_info["id"])
if view is None:
return _error(404, "NOT_FOUND", "No such download.")
return _ok(view)
@ROUTES.delete("/api/download/{id}")
async def delete_download(request: web.Request) -> web.Response:
try:
await DOWNLOAD_MANAGER.delete(request.match_info["id"])
except DownloadError as e:
return _from_download_error(e)
return _ok({"deleted": True})
@ROUTES.post("/api/download/{id}/pause")
async def pause(request: web.Request) -> web.Response:
try:
await DOWNLOAD_MANAGER.pause(request.match_info["id"])
except DownloadError as e:
return _from_download_error(e)
return _ok({"ok": True})
@ROUTES.post("/api/download/{id}/resume")
async def resume(request: web.Request) -> web.Response:
try:
await DOWNLOAD_MANAGER.resume(request.match_info["id"])
except DownloadError as e:
return _from_download_error(e)
return _ok({"ok": True})
@ROUTES.post("/api/download/{id}/cancel")
async def cancel(request: web.Request) -> web.Response:
try:
await DOWNLOAD_MANAGER.cancel(request.match_info["id"])
except DownloadError as e:
return _from_download_error(e)
return _ok({"ok": True})
@ROUTES.post("/api/download/{id}/priority")
async def set_priority(request: web.Request) -> web.Response:
parsed = await _parse(request, schemas_in.PriorityRequest)
if isinstance(parsed, web.Response):
return parsed
try:
await DOWNLOAD_MANAGER.set_priority(request.match_info["id"], parsed.priority)
except DownloadError as e:
return _from_download_error(e)
return _ok({"ok": True})
+46
View File
@@ -0,0 +1,46 @@
"""Request schemas for the download manager API.
Pydantic enforces shape at the boundary; handlers operate only on validated
values past that point.
"""
from __future__ import annotations
from typing import Optional
from pydantic import BaseModel, Field, field_validator
class EnqueueRequest(BaseModel):
url: str
model_id: str
priority: int = 0
expected_sha256: Optional[str] = None
allow_any_extension: bool = False
@field_validator("url")
@classmethod
def _strip_url(cls, v: str) -> str:
return v.strip()
class PriorityRequest(BaseModel):
priority: int
class AvailabilityRequest(BaseModel):
"""``{model_id: url}`` — the URLs declared in the workflow JSON."""
models: dict[str, str] = Field(default_factory=dict)
@field_validator("models")
@classmethod
def _strip_urls(cls, v: dict[str, str]) -> dict[str, str]:
return {k: url.strip() for k, url in v.items()}
__all__ = [
"EnqueueRequest",
"PriorityRequest",
"AvailabilityRequest",
]
+15
View File
@@ -0,0 +1,15 @@
"""Response helpers for the download manager API.
The download/status read models are plain dicts produced by the manager. This
module serializes the per-provider auth status (never a token) for the API.
"""
from __future__ import annotations
from app.model_downloader.auth.providers import PROVIDERS
from app.model_downloader.auth.store import AUTH_STORE
def auth_status() -> list[dict]:
"""Per-provider auth status — never includes a token."""
return [AUTH_STORE.status(p) for p in PROVIDERS.values()]
+269
View File
@@ -0,0 +1,269 @@
"""Generic OAuth 2.0 PKCE engine + transient loopback callback server.
The flow, per provider:
1. :func:`start_login_flow` builds a PKCE challenge, binds a loopback callback
server on ``127.0.0.1:<CALLBACK_PORT>`` at ``/callback/<provider>``, and
returns the provider's authorize URL for the user to open.
2. The provider redirects the browser back to the loopback URL with a ``code``
and the ``state`` we generated. The server validates ``state``, exchanges the
code for a :class:`Token`, hands it to the ``deliver`` sink, and tears down.
3. If no callback arrives within :data:`_LOGIN_TIMEOUT`, the server tears down.
The callback runs on its own bare server, not ComfyUI's main server: the main
server rejects cross-site navigations (``Sec-Fetch-Site: cross-site`` → 403),
and an OAuth redirect from the provider is exactly such a navigation. The port
is fixed because HuggingFace and Civitai require an exact registered
``redirect_uri`` (port included); only one login runs at a time so the port
never contends with itself.
Only public PKCE clients are supported (no client secret). All outbound calls
go to the provider's own authorize/token endpoints, strictly user-initiated.
"""
from __future__ import annotations
import asyncio
import base64
import hashlib
import logging
import os
import secrets
import time
from typing import Callable
from urllib.parse import urlencode
from aiohttp import web
from app.model_downloader.auth.providers import Provider
from app.model_downloader.auth.token_store import Token
from app.model_downloader.net.session import get_session, ssl_context
CALLBACK_HOST = "127.0.0.1"
# Fixed loopback port for the OAuth redirect. Must match the redirect URI
# registered on the provider's OAuth app; override in lockstep if you change it.
CALLBACK_PORT = int(os.environ.get("COMFY_OAUTH_CALLBACK_PORT", "41954"))
_LOGIN_TIMEOUT = 300.0 # seconds to wait for the browser callback
# The auth tab is opened by the frontend via window.open, so window.close() is
# allowed here; the visible text is the fallback when the browser blocks it.
_SUCCESS_HTML = (
"<!doctype html><meta charset=utf-8><title>ComfyUI</title>"
"<p>Login successful. You can close this window and return to ComfyUI.</p>"
"<script>window.close()</script>"
)
# Token sink: called with the provider name and the exchanged Token.
TokenSink = Callable[[str, Token], None]
class OAuthError(Exception):
"""A user-facing OAuth failure."""
class OAuthNotConfigured(OAuthError):
"""The provider has no public client id configured."""
class LoginInProgress(OAuthError):
"""A login flow for this provider is already running."""
def _b64url(data: bytes) -> str:
return base64.urlsafe_b64encode(data).rstrip(b"=").decode("ascii")
def _make_pkce() -> tuple[str, str]:
"""Return ``(verifier, challenge)`` for the S256 PKCE method."""
verifier = _b64url(secrets.token_bytes(32))
challenge = _b64url(hashlib.sha256(verifier.encode("ascii")).digest())
return verifier, challenge
def build_authorize_url(
provider: Provider, challenge: str, state: str, redirect_uri: str
) -> str:
params = {
"response_type": "code",
"client_id": provider.client_id,
"redirect_uri": redirect_uri,
"scope": provider.scope,
"state": state,
"code_challenge": challenge,
"code_challenge_method": "S256",
}
return f"{provider.authorize_url}?{urlencode(params)}"
def _token_from_payload(payload: dict) -> Token:
expires_in = payload.get("expires_in")
expires_at = int(time.time()) + int(expires_in) if expires_in else 0
return Token(
access_token=payload["access_token"],
refresh_token=payload.get("refresh_token"),
expires_at=expires_at,
token_type=payload.get("token_type", "Bearer"),
scope=payload.get("scope"),
)
async def _post_token(provider: Provider, data: dict) -> Token:
session = await get_session()
resp = await session.post(
provider.token_url,
data=data,
headers={"Accept": "application/json"},
ssl=ssl_context(),
)
try:
if resp.status != 200:
body = await resp.text()
raise OAuthError(
f"{provider.name} token endpoint returned HTTP {resp.status}: {body[:200]}"
)
payload = await resp.json()
finally:
await resp.release()
if "access_token" not in payload:
raise OAuthError(f"{provider.name} token response missing access_token")
return _token_from_payload(payload)
async def exchange_code(
provider: Provider, code: str, verifier: str, redirect_uri: str
) -> Token:
return await _post_token(
provider,
{
"grant_type": "authorization_code",
"code": code,
"redirect_uri": redirect_uri,
"client_id": provider.client_id,
"code_verifier": verifier,
},
)
async def refresh_access_token(provider: Provider, token: Token) -> Token:
if not token.refresh_token:
raise OAuthError(f"{provider.name} token is not refreshable")
refreshed = await _post_token(
provider,
{
"grant_type": "refresh_token",
"refresh_token": token.refresh_token,
"client_id": provider.client_id,
},
)
# Some providers omit a new refresh token on refresh; keep the old one.
if refreshed.refresh_token is None:
refreshed.refresh_token = token.refresh_token
return refreshed
class _LoginFlow:
"""A single in-flight login: owns the loopback server and PKCE state."""
def __init__(self, provider: Provider, deliver: TokenSink) -> None:
self.provider = provider
self.deliver = deliver
self.verifier, self.challenge = _make_pkce()
self.state = secrets.token_urlsafe(24)
self.redirect_uri = f"http://{CALLBACK_HOST}:{CALLBACK_PORT}/callback/{provider.name}"
self._runner: web.AppRunner | None = None
self._timeout_handle: asyncio.TimerHandle | None = None
async def start(self) -> str:
app = web.Application()
app.router.add_get("/callback/{provider}", self._handle_callback)
self._runner = web.AppRunner(app)
await self._runner.setup()
site = web.TCPSite(self._runner, CALLBACK_HOST, CALLBACK_PORT, reuse_address=True)
try:
await site.start()
except OSError as e:
await self._runner.cleanup()
self._runner = None
raise OAuthError(f"could not bind callback port {CALLBACK_PORT}: {e}")
loop = asyncio.get_running_loop()
self._timeout_handle = loop.call_later(
_LOGIN_TIMEOUT, lambda: asyncio.ensure_future(self._teardown())
)
return build_authorize_url(
self.provider, self.challenge, self.state, self.redirect_uri
)
async def _handle_callback(self, request: web.Request) -> web.Response:
if request.match_info.get("provider") != self.provider.name:
return web.Response(text="Unknown login.", content_type="text/plain", status=404)
error = request.query.get("error")
if error:
asyncio.ensure_future(self._teardown())
return web.Response(
text=f"Login failed: {error}", content_type="text/plain", status=400
)
if request.query.get("state") != self.state:
return web.Response(
text="Login failed: state mismatch.",
content_type="text/plain",
status=400,
)
code = request.query.get("code")
if not code:
return web.Response(
text="Login failed: no authorization code.",
content_type="text/plain",
status=400,
)
try:
token = await exchange_code(
self.provider, code, self.verifier, self.redirect_uri
)
except OAuthError as e:
logging.warning("[model_downloader] %s login failed: %s", self.provider.name, e)
asyncio.ensure_future(self._teardown())
return web.Response(
text=f"Login failed: {e}", content_type="text/plain", status=502
)
self.deliver(self.provider.name, token)
asyncio.ensure_future(self._teardown())
return web.Response(text=_SUCCESS_HTML, content_type="text/html")
async def _teardown(self) -> None:
_ACTIVE.pop(self.provider.name, None)
if self._timeout_handle is not None:
self._timeout_handle.cancel()
self._timeout_handle = None
if self._runner is not None:
try:
await self._runner.cleanup()
except Exception:
logging.debug("[model_downloader] callback server cleanup error", exc_info=True)
self._runner = None
_ACTIVE: dict[str, _LoginFlow] = {}
def login_in_progress(provider_name: str) -> bool:
return provider_name in _ACTIVE
async def start_login_flow(provider: Provider, deliver: TokenSink) -> str:
"""Begin a login flow and return the authorize URL to open in a browser.
Binds the fixed-port loopback callback server; only one login may run at a
time since that port is shared.
"""
if not provider.client_id:
raise OAuthNotConfigured(
f"OAuth app not configured for {provider.name}; set "
f"{provider.client_id_env} or use an env API key."
)
if _ACTIVE:
active = next(iter(_ACTIVE))
raise LoginInProgress(f"A login for {active} is already in progress.")
flow = _LoginFlow(provider, deliver)
authorize_url = await flow.start()
_ACTIVE[provider.name] = flow
return authorize_url
+87
View File
@@ -0,0 +1,87 @@
"""Provider registry for download authentication.
A :class:`Provider` describes a hub that can authenticate downloads either from
an environment API key or from an OAuth 2.0 access token. Both HuggingFace and
Civitai are public PKCE clients, so no client secret is ever stored; the public
``client_id`` is a placeholder overridable via env.
"""
from __future__ import annotations
import os
from dataclasses import dataclass
from urllib.parse import urlsplit
def normalize_host(host: str) -> str:
"""Lowercase, strip port, IDNA-encode."""
if not host:
return ""
host = host.strip()
if "://" in host: # a full URL was pasted — extract just the host
host = urlsplit(host).hostname or ""
host = host.lower()
if host.startswith("[") and "]" in host: # bracketed IPv6 literal
host = host[1 : host.index("]")]
elif host.count(":") == 1: # host:port (not IPv6)
host = host.split(":", 1)[0]
try:
host = host.encode("idna").decode("ascii")
except (UnicodeError, ValueError):
pass
return host
@dataclass(frozen=True)
class Provider:
name: str
host: str
authorize_url: str
token_url: str
scope: str
# Env vars to try, in order, for a plain API key.
env_keys: tuple[str, ...]
# Env var overriding the public OAuth client id.
client_id_env: str
# Public PKCE client id. Empty means "not configured" until the env sets it.
default_client_id: str = ""
@property
def client_id(self) -> str:
return os.environ.get(self.client_id_env, self.default_client_id) or ""
def env_token(self) -> str | None:
for var in self.env_keys:
token = os.environ.get(var)
if token:
return token
return None
PROVIDERS: dict[str, Provider] = {
"huggingface": Provider(
name="huggingface",
host="huggingface.co",
authorize_url="https://huggingface.co/oauth/authorize",
token_url="https://huggingface.co/oauth/token",
scope="openid read-repos gated-repos",
env_keys=("HF_TOKEN", "HUGGING_FACE_HUB_TOKEN"),
client_id_env="COMFY_HF_OAUTH_CLIENT_ID",
),
"civitai": Provider(
name="civitai",
host="civitai.com",
authorize_url="https://auth.civitai.com/api/auth/oauth/authorize",
token_url="https://auth.civitai.com/api/auth/oauth/token",
scope="4", # ModelsRead; UserRead is auto-granted
env_keys=("CIVITAI_API_TOKEN", "CIVITAI_API_KEY"),
client_id_env="COMFY_CIVITAI_OAUTH_CLIENT_ID",
),
}
_HOST_TO_PROVIDER = {p.host: p for p in PROVIDERS.values()}
def provider_for_host(host: str) -> Provider | None:
"""Return the provider whose host exactly matches ``host`` (normalized)."""
return _HOST_TO_PROVIDER.get(normalize_host(host))
+42
View File
@@ -0,0 +1,42 @@
"""Per-hop auth resolution (https only).
Recomputed from scratch on every redirect hop: a hop only gets a bearer token
when *its own host* matches a configured provider, so a token bound to
``huggingface.co`` is silently dropped when the request is redirected to a
presigned CDN host — which is exactly what these hubs expect.
For a matching hop: env API key first, then the provider's OAuth access token
(refreshed if expired), else no auth.
"""
from __future__ import annotations
from dataclasses import dataclass, field
from app.model_downloader.auth.providers import provider_for_host
from app.model_downloader.auth.store import AUTH_STORE
@dataclass
class RequestAuth:
"""How to modify a single request to carry a bearer token."""
headers: dict[str, str] = field(default_factory=dict)
async def resolve_auth_for_hop(host: str, scheme: str) -> RequestAuth | None:
"""Resolve the bearer token (if any) to attach for one request hop."""
if scheme.lower() != "https":
return None
provider = provider_for_host(host)
if provider is None:
return None
token = provider.env_token()
if token:
return RequestAuth(headers={"Authorization": f"Bearer {token}"})
access = await AUTH_STORE.get_valid_token(provider)
if access:
return RequestAuth(headers={"Authorization": f"Bearer {access}"})
return None
+64
View File
@@ -0,0 +1,64 @@
"""In-memory OAuth token cache over the on-disk token store.
:data:`AUTH_STORE` is the process singleton the resolver and API talk to. It
lazily loads each provider's token from disk, refreshes an expired access token
via its refresh token, and orchestrates the login flow (delegating the loopback
callback server to :mod:`oauth`).
"""
from __future__ import annotations
from app.model_downloader.auth import oauth, token_store
from app.model_downloader.auth.providers import Provider
from app.model_downloader.auth.token_store import Token
class AuthStore:
def __init__(self) -> None:
# provider name -> Token, or None when known to be absent. A missing key
# means "not yet loaded from disk".
self._cache: dict[str, Token | None] = {}
def _load(self, name: str) -> Token | None:
if name not in self._cache:
self._cache[name] = token_store.load(name)
return self._cache[name]
def set_token(self, name: str, token: Token) -> None:
self._cache[name] = token
token_store.save(name, token)
def clear(self, name: str) -> None:
self._cache[name] = None
token_store.delete(name)
async def get_valid_token(self, provider: Provider) -> str | None:
"""Return a valid access token string for ``provider``, or ``None``.
Refreshes an expired token when a refresh token is available.
"""
token = self._load(provider.name)
if token is None or not token.access_token:
return None
if token.is_expired():
if not token.refresh_token:
return None
token = await oauth.refresh_access_token(provider, token)
self.set_token(provider.name, token)
return token.access_token
async def begin_login(self, provider: Provider) -> str:
"""Start a login flow; returns the authorize URL to open in a browser."""
return await oauth.start_login_flow(provider, self.set_token)
def status(self, provider: Provider) -> dict:
token = self._load(provider.name)
return {
"provider": provider.name,
"logged_in": token is not None and bool(token.access_token),
"login_in_progress": oauth.login_in_progress(provider.name),
"env_key_present": provider.env_token() is not None,
}
AUTH_STORE = AuthStore()
+186
View File
@@ -0,0 +1,186 @@
"""On-disk OAuth token persistence — one machine-bound blob per provider.
Tokens live under ``folder_paths.get_system_user_directory("download_auth")``,
never in the SQLite DB. Each provider file is written ``0600`` and holds an
opaque blob, not readable JSON: the token JSON is XORed with an HMAC-SHA256
keystream whose key is derived from stable machine/install attributes plus a
per-install random salt.
This is obfuscation, not confidentiality. It stops a token from being read by a
human browsing files, grepped out of a backup, or lifted from a folder copied to
another machine (the blob won't decrypt off its origin machine). It does not
protect against code running inside this process (custom nodes) or an attacker
who reads this source and recomputes the key.
"""
from __future__ import annotations
import base64
import getpass
import hashlib
import hmac
import json
import os
import platform
import secrets
import time
from dataclasses import asdict, dataclass
import folder_paths
_SALT_FILE = ".salt"
_SALT_LEN = 32
_NONCE_LEN = 16
_PBKDF2_ITERS = 200_000
@dataclass
class Token:
access_token: str
refresh_token: str | None = None
# Epoch seconds when the access token expires; 0 means "unknown / no expiry".
expires_at: int = 0
token_type: str = "Bearer"
scope: str | None = None
def is_expired(self, skew: int = 60) -> bool:
if not self.expires_at:
return False
return time.time() + skew >= self.expires_at
def _auth_dir() -> str:
path = folder_paths.get_system_user_directory("download_auth")
os.makedirs(path, exist_ok=True)
return path
def _token_path(provider: str) -> str:
return os.path.join(_auth_dir(), f"{provider}.bin")
def _machine_id() -> bytes:
"""A stable per-machine identifier, best-effort across platforms."""
for path in ("/etc/machine-id", "/var/lib/dbus/machine-id"):
try:
with open(path, "rb") as f:
return f.read().strip()
except OSError:
pass
if os.name == "nt":
import winreg
try:
key = winreg.OpenKey(
winreg.HKEY_LOCAL_MACHINE, r"SOFTWARE\Microsoft\Cryptography"
)
try:
guid, _ = winreg.QueryValueEx(key, "MachineGuid")
return str(guid).encode("utf-8")
finally:
winreg.CloseKey(key)
except OSError:
pass
return platform.node().encode("utf-8")
def _machine_material(auth_dir: str) -> bytes:
try:
user = getpass.getuser()
except Exception:
user = ""
parts = (_machine_id(), platform.node().encode("utf-8"), user.encode("utf-8"), auth_dir.encode("utf-8"))
return b"\x00".join(parts)
def _load_or_create_salt(auth_dir: str) -> bytes | None:
path = os.path.join(auth_dir, _SALT_FILE)
try:
with open(path, "rb") as f:
salt = f.read()
if len(salt) == _SALT_LEN:
return salt
except FileNotFoundError:
pass
except OSError:
return None
salt = secrets.token_bytes(_SALT_LEN)
fd = os.open(path, os.O_WRONLY | os.O_CREAT | os.O_TRUNC, 0o600)
with os.fdopen(fd, "wb") as f:
f.write(salt)
os.chmod(path, 0o600)
return salt
def _derive_key(salt: bytes, auth_dir: str) -> bytes:
return hashlib.pbkdf2_hmac(
"sha256", _machine_material(auth_dir), salt, _PBKDF2_ITERS, dklen=32
)
def _keystream(key: bytes, nonce: bytes, n: int) -> bytes:
out = bytearray()
counter = 0
while len(out) < n:
out.extend(hmac.new(key, nonce + counter.to_bytes(8, "big"), hashlib.sha256).digest())
counter += 1
return bytes(out[:n])
def _xor(data: bytes, stream: bytes) -> bytes:
return bytes(a ^ b for a, b in zip(data, stream))
def load(provider: str) -> Token | None:
auth_dir = _auth_dir()
try:
with open(_token_path(provider), "rb") as f:
blob = base64.b64decode(f.read())
except FileNotFoundError:
return None
except (ValueError, OSError):
return None
salt = _load_or_create_salt(auth_dir)
if salt is None or len(blob) <= _NONCE_LEN:
return None
nonce, ciphertext = blob[:_NONCE_LEN], blob[_NONCE_LEN:]
key = _derive_key(salt, auth_dir)
plaintext = _xor(ciphertext, _keystream(key, nonce, len(ciphertext)))
# A wrong machine / corrupt file decrypts to garbage; treat as logged out.
try:
data = json.loads(plaintext)
except ValueError:
return None
if not isinstance(data, dict) or "access_token" not in data:
return None
return Token(
access_token=data.get("access_token", ""),
refresh_token=data.get("refresh_token"),
expires_at=int(data.get("expires_at", 0) or 0),
token_type=data.get("token_type", "Bearer"),
scope=data.get("scope"),
)
def save(provider: str, token: Token) -> None:
auth_dir = _auth_dir()
salt = _load_or_create_salt(auth_dir)
if salt is None:
return
key = _derive_key(salt, auth_dir)
nonce = secrets.token_bytes(_NONCE_LEN)
plaintext = json.dumps(asdict(token)).encode("utf-8")
ciphertext = _xor(plaintext, _keystream(key, nonce, len(plaintext)))
blob = base64.b64encode(nonce + ciphertext)
path = _token_path(provider)
fd = os.open(path, os.O_WRONLY | os.O_CREAT | os.O_TRUNC, 0o600)
with os.fdopen(fd, "wb") as f:
f.write(blob)
os.chmod(path, 0o600)
def delete(provider: str) -> None:
try:
os.remove(_token_path(provider))
except FileNotFoundError:
pass
+32
View File
@@ -0,0 +1,32 @@
"""Shared constants for the download manager.
Status values are persisted as TEXT in the ``downloads`` table; keep them
stable. The lifecycle is:
queued -> active -> verifying -> completed
| |-> paused -> (resume) -> active
| |-> failed (network, retryable) -> queued (backoff)
|-> cancelled
"""
from __future__ import annotations
class DownloadStatus:
QUEUED = "queued"
ACTIVE = "active"
PAUSED = "paused"
VERIFYING = "verifying"
COMPLETED = "completed"
FAILED = "failed"
CANCELLED = "cancelled"
#: States from which a worker is doing (or about to do) network I/O.
LIVE = (QUEUED, ACTIVE, VERIFYING)
#: Terminal states — the job will not transition again on its own.
TERMINAL = (COMPLETED, FAILED, CANCELLED)
# Default temp-file suffix. Distinctive so the startup orphan sweep only
# removes files THIS subsystem created, never unrelated *.tmp files.
TMP_SUFFIX = ".comfy-download.part"
+125
View File
@@ -0,0 +1,125 @@
"""SQLAlchemy models for the download manager.
Two tables:
- ``downloads`` one row per requested file (job + queue state).
- ``download_segments`` per-segment byte progress, for segmented resume.
On completion a finished file is registered into the assets catalog;
``downloads`` is kept only as job history.
"""
from __future__ import annotations
import time
import uuid
from sqlalchemy import (
BigInteger,
Boolean,
CheckConstraint,
ForeignKey,
Index,
Integer,
String,
Text,
)
from sqlalchemy.orm import Mapped, mapped_column, relationship
from app.database.models import Base
def _uuid() -> str:
return str(uuid.uuid4())
def _now() -> int:
return int(time.time())
class Download(Base):
__tablename__ = "downloads"
id: Mapped[str] = mapped_column(String(36), primary_key=True, default=_uuid)
# Original requested URL and the final URL after validated redirects.
url: Mapped[str] = mapped_column(Text, nullable=False)
final_url: Mapped[str | None] = mapped_column(Text, nullable=True)
# Canonical "<directory>/<filename>" identifier (resolved via folder_paths).
model_id: Mapped[str] = mapped_column(String(1024), nullable=False)
# Final on-disk location and the .part write target.
dest_path: Mapped[str] = mapped_column(Text, nullable=False)
temp_path: Mapped[str] = mapped_column(Text, nullable=False)
status: Mapped[str] = mapped_column(String(16), nullable=False)
priority: Mapped[int] = mapped_column(Integer, nullable=False, default=0)
total_bytes: Mapped[int | None] = mapped_column(BigInteger, nullable=True)
bytes_done: Mapped[int] = mapped_column(BigInteger, nullable=False, default=0)
etag: Mapped[str | None] = mapped_column(String(512), nullable=True)
last_modified: Mapped[str | None] = mapped_column(String(128), nullable=True)
accept_ranges: Mapped[bool] = mapped_column(Boolean, nullable=False, default=False)
# Optional hub-provided checksum to verify against (NOT the dedup key).
expected_sha256: Mapped[str | None] = mapped_column(String(64), nullable=True)
allow_any_extension: Mapped[bool] = mapped_column(
Boolean, nullable=False, default=False
)
# How many retryable failures we have seen (for backoff capping).
attempts: Mapped[int] = mapped_column(Integer, nullable=False, default=0)
error: Mapped[str | None] = mapped_column(Text, nullable=True)
created_at: Mapped[int] = mapped_column(BigInteger, nullable=False, default=_now)
updated_at: Mapped[int] = mapped_column(
BigInteger, nullable=False, default=_now, onupdate=_now
)
segments: Mapped[list[DownloadSegment]] = relationship(
"DownloadSegment",
back_populates="download",
cascade="all,delete-orphan",
passive_deletes=True,
order_by="DownloadSegment.idx",
)
__table_args__ = (
Index("ix_downloads_status", "status"),
Index("ix_downloads_priority", "priority"),
Index("ix_downloads_model_id", "model_id"),
CheckConstraint("bytes_done >= 0", name="ck_downloads_bytes_done_nonneg"),
CheckConstraint(
"total_bytes IS NULL OR total_bytes >= 0",
name="ck_downloads_total_bytes_nonneg",
),
)
def __repr__(self) -> str:
return f"<Download id={self.id} model_id={self.model_id!r} status={self.status}>"
class DownloadSegment(Base):
__tablename__ = "download_segments"
download_id: Mapped[str] = mapped_column(
String(36),
ForeignKey("downloads.id", ondelete="CASCADE"),
primary_key=True,
)
idx: Mapped[int] = mapped_column(Integer, primary_key=True)
start_offset: Mapped[int] = mapped_column(BigInteger, nullable=False)
end_offset: Mapped[int] = mapped_column(BigInteger, nullable=False)
bytes_done: Mapped[int] = mapped_column(BigInteger, nullable=False, default=0)
download: Mapped[Download] = relationship("Download", back_populates="segments")
__table_args__ = (
CheckConstraint("bytes_done >= 0", name="ck_segments_bytes_done_nonneg"),
CheckConstraint("end_offset >= start_offset", name="ck_segments_range"),
)
def __repr__(self) -> str:
return (
f"<DownloadSegment {self.download_id}#{self.idx} "
f"{self.start_offset}-{self.end_offset} done={self.bytes_done}>"
)
+177
View File
@@ -0,0 +1,177 @@
"""Synchronous DB access for the download manager.
All functions open their own short-lived session via ``create_session`` and
commit before returning, mirroring ``app/assets`` usage. They are blocking
(SQLite) and should be called from async code through ``asyncio.to_thread``.
"""
from __future__ import annotations
import time
from typing import Optional
from sqlalchemy import delete, select
from app.database.db import create_session
from app.model_downloader.constants import DownloadStatus
from app.model_downloader.database.models import Download, DownloadSegment
# ----- downloads -----
def insert_download(values: dict) -> None:
with create_session() as session:
session.add(Download(**values))
session.commit()
def get_download(download_id: str) -> Optional[Download]:
with create_session() as session:
row = session.get(Download, download_id)
if row is not None:
session.expunge_all()
return row
def list_downloads() -> list[Download]:
with create_session() as session:
rows = list(
session.execute(
select(Download).order_by(Download.created_at.desc())
).scalars()
)
session.expunge_all()
return rows
def has_live_download_for_model(
model_id: str, live_statuses: tuple[str, ...], exclude_id: Optional[str] = None
) -> bool:
with create_session() as session:
stmt = select(Download.id).where(
Download.model_id == model_id,
Download.status.in_(live_statuses),
).limit(1)
if exclude_id is not None:
stmt = stmt.where(Download.id != exclude_id)
return session.execute(stmt).first() is not None
def list_segments(download_id: str) -> list[DownloadSegment]:
with create_session() as session:
rows = list(
session.execute(
select(DownloadSegment)
.where(DownloadSegment.download_id == download_id)
.order_by(DownloadSegment.idx)
).scalars()
)
session.expunge_all()
return rows
def update_download(download_id: str, **fields) -> None:
if not fields:
return
fields.setdefault("updated_at", int(time.time()))
with create_session() as session:
row = session.get(Download, download_id)
if row is None:
return
for key, value in fields.items():
setattr(row, key, value)
session.commit()
def delete_download(download_id: str) -> None:
with create_session() as session:
row = session.get(Download, download_id)
if row is not None:
session.delete(row)
session.commit()
def delete_downloads(download_ids: list[str]) -> int:
"""Delete many downloads in one transaction; returns the number removed.
Uses a bulk ``DELETE ... WHERE id IN (...)``. Segment rows are removed by
the ``ON DELETE CASCADE`` foreign key (SQLite ``PRAGMA foreign_keys=ON`` is
set in ``app/database/db.py``), so this stays consistent without loading the
ORM relationship.
"""
if not download_ids:
return 0
with create_session() as session:
result = session.execute(
delete(Download).where(Download.id.in_(download_ids))
)
session.commit()
return result.rowcount or 0
def replace_segments(download_id: str, segments: list[dict]) -> None:
"""Atomically replace the segment plan for a download."""
with create_session() as session:
session.query(DownloadSegment).filter(
DownloadSegment.download_id == download_id
).delete()
for seg in segments:
session.add(DownloadSegment(download_id=download_id, **seg))
session.commit()
def update_segment_progress(download_id: str, idx: int, bytes_done: int) -> None:
with create_session() as session:
row = session.get(DownloadSegment, {"download_id": download_id, "idx": idx})
if row is None:
return
row.bytes_done = bytes_done
session.commit()
def list_queued_downloads() -> list[Download]:
"""Queued rows ordered for admission (priority desc, then FIFO)."""
with create_session() as session:
rows = list(
session.execute(
select(Download)
.where(Download.status == DownloadStatus.QUEUED)
.order_by(Download.priority.desc(), Download.created_at.asc())
).scalars()
)
session.expunge_all()
return rows
def reconcile_live_downloads() -> list[Download]:
"""Reset any ``active``/``verifying`` rows left by a previous run.
On a clean restart there can be no live worker, so anything still marked
live is stale. Move it back to ``queued`` (offsets are preserved on the
segment rows) so the scheduler re-admits it. Returns the rows that should
be re-queued by the scheduler (queued + paused).
"""
with create_session() as session:
stale = list(
session.execute(
select(Download).where(
Download.status.in_([DownloadStatus.ACTIVE, DownloadStatus.VERIFYING])
)
).scalars()
)
now = int(time.time())
for row in stale:
row.status = DownloadStatus.QUEUED
row.updated_at = now
session.commit()
resumable = list(
session.execute(
select(Download)
.where(Download.status == DownloadStatus.QUEUED)
.order_by(Download.priority.desc(), Download.created_at.asc())
).scalars()
)
session.expunge_all()
return resumable
+611
View File
@@ -0,0 +1,611 @@
"""The per-download worker.
One :class:`DownloadJob` drives a single file from probe to verified, cataloged
completion. It supports cooperative pause / resume / cancel, segmented
multi-connection transfer with positioned writes, and a verification gate
(size + structural + optional sha256) before the atomic rename into place.
Control is cooperative: external callers flip ``_control`` via
:meth:`request_pause` / :meth:`request_cancel`; segment loops observe it between
chunks and raise, which unwinds cleanly and persists resume offsets.
"""
from __future__ import annotations
import asyncio
import logging
import os
import time
from dataclasses import dataclass, field
from typing import Callable, Optional
from comfy.cli_args import args
from app.model_downloader.constants import DownloadStatus
from app.model_downloader.database import queries
from app.model_downloader.engine.planner import (
effective_segment_count,
plan_segments,
)
from app.model_downloader.engine.writer import FileWriter
from app.model_downloader.net.http import open_validated, redact_url
from app.model_downloader.net.probe import gated_error_message, probe
from app.model_downloader.verify import checksum, dedup, structural
_RETRYABLE_STATUSES = {408, 429, 500, 502, 503, 504}
_PERSIST_INTERVAL = 2.0 # seconds between throttled progress persists
class Paused(Exception):
pass
class Cancelled(Exception):
pass
class RemoteChanged(Exception):
"""The remote file changed under a resume (got 200 where 206 expected)."""
class RetryableError(Exception):
pass
class FatalError(Exception):
"""Non-retryable: 4xx, checksum mismatch, structural failure, gated, etc."""
@dataclass
class SegmentRuntime:
idx: int
start: int
end: int # inclusive; may be -1 for unknown-size single stream
bytes_done: int = 0
@property
def length(self) -> int:
return self.end - self.start + 1
@dataclass
class RuntimeState:
download_id: str
model_id: str
url: str
priority: int
status: str
total_bytes: Optional[int] = None
bytes_done: int = 0
error: Optional[str] = None
segments: list[SegmentRuntime] = field(default_factory=list)
started_at: float = field(default_factory=time.monotonic)
_last_bytes: int = 0
_last_time: float = field(default_factory=time.monotonic)
speed_bps: float = 0.0
@property
def progress(self) -> Optional[float]:
if not self.total_bytes:
return None
return min(1.0, self.bytes_done / self.total_bytes)
@property
def eta_seconds(self) -> Optional[float]:
if not self.total_bytes or self.speed_bps <= 0:
return None
remaining = max(0, self.total_bytes - self.bytes_done)
return remaining / self.speed_bps
@dataclass
class JobSpec:
download_id: str
url: str
model_id: str
dest_path: str
temp_path: str
priority: int = 0
expected_sha256: Optional[str] = None
allow_any_extension: bool = False
etag: Optional[str] = None
attempts: int = 0
class DownloadJob:
def __init__(
self, spec: JobSpec, notify_cb: Optional[Callable[[str], None]] = None
) -> None:
self.spec = spec
self._notify = notify_cb
self._control = "run" # run | pause | cancel
self.state = RuntimeState(
download_id=spec.download_id,
model_id=spec.model_id,
url=spec.url,
priority=spec.priority,
status=DownloadStatus.QUEUED,
)
self._writer: Optional[FileWriter] = None
self._etag: Optional[str] = spec.etag
self._last_persist = 0.0
# ----- external control -----
def request_pause(self) -> None:
if self._control == "run":
self._control = "pause"
def request_cancel(self) -> None:
self._control = "cancel"
def _check_control(self) -> None:
if self._control == "cancel":
raise Cancelled()
if self._control == "pause":
raise Paused()
# ----- lifecycle -----
async def run(self) -> str:
"""Run to a terminal/paused state; returns the final status string."""
await self._set_status(DownloadStatus.ACTIVE, error=None)
try:
pr = await self._probe_and_plan()
await self._transfer(pr)
await self._finalize()
await self._set_status(DownloadStatus.COMPLETED)
except Paused:
await self._persist_progress(force=True)
await self._set_status(DownloadStatus.PAUSED)
except Cancelled:
await self._close_writer()
self._remove_temp()
await self._set_status(DownloadStatus.CANCELLED)
except RemoteChanged:
await self._reset_for_restart()
await self._set_status(
DownloadStatus.QUEUED, error="remote file changed; restarting"
)
except RetryableError as e:
await self._persist_progress(force=True)
await self._set_status(DownloadStatus.QUEUED, error=str(e))
except FatalError as e:
await self._close_writer()
self._remove_temp()
await self._set_status(DownloadStatus.FAILED, error=str(e))
except Exception as e: # unexpected -> treat as retryable
logging.warning(
"[model_downloader] %s unexpected error: %s",
self.spec.model_id, e, exc_info=True,
)
await self._persist_progress(force=True)
await self._set_status(DownloadStatus.QUEUED, error=f"{type(e).__name__}: {e}")
finally:
await self._close_writer()
return self.state.status
# ----- probe + plan -----
async def _probe_and_plan(self):
pr = await probe(self.spec.url)
if not pr.ok:
if pr.gated:
raise FatalError(gated_error_message(self.spec.url, pr))
if pr.status == 0 or pr.status in _RETRYABLE_STATUSES:
raise RetryableError(pr.error or "probe failed")
raise FatalError(pr.error or f"probe returned HTTP {pr.status}")
max_bytes = self._max_download_bytes()
if max_bytes is not None and pr.total_bytes is not None and pr.total_bytes > max_bytes:
raise FatalError(
f"file size {pr.total_bytes} exceeds the maximum allowed "
f"download size {max_bytes} (--download-max-bytes)"
)
self._etag = pr.etag or self._etag
self.state.total_bytes = pr.total_bytes
await asyncio.to_thread(
queries.update_download,
self.spec.download_id,
final_url=pr.final_url,
total_bytes=pr.total_bytes,
accept_ranges=pr.accept_ranges,
etag=pr.etag,
last_modified=pr.last_modified,
)
seg_count = effective_segment_count(
pr.total_bytes, pr.accept_ranges, max(1, args.download_segments)
)
existing = await asyncio.to_thread(queries.list_segments, self.spec.download_id)
can_resume_segmented = (
seg_count > 1
and existing
and pr.total_bytes is not None
and existing[-1].end_offset == pr.total_bytes - 1
)
if can_resume_segmented and not self._segmented_part_valid(pr.total_bytes):
# The persisted per-segment offsets describe bytes in a preallocated
# .part that is now gone or the wrong size (e.g. the partial of a
# failed download was swept on restart, or removed by a fatal
# error). Trusting them would skip already-"complete" segments and
# leave zero-filled holes. Discard the offsets and re-plan fresh.
logging.info(
"[model_downloader] %s discarding segmented resume offsets "
"(preallocated .part missing or wrong size); restarting",
self.spec.model_id,
)
self._remove_temp()
await asyncio.to_thread(
queries.replace_segments, self.spec.download_id, []
)
await asyncio.to_thread(
queries.update_download, self.spec.download_id, bytes_done=0
)
existing = []
can_resume_segmented = False
if can_resume_segmented:
# Resume an existing segmented plan.
self.state.segments = [
SegmentRuntime(s.idx, s.start_offset, s.end_offset, s.bytes_done)
for s in existing
]
elif seg_count > 1 and pr.total_bytes is not None:
plans = plan_segments(pr.total_bytes, seg_count)
await asyncio.to_thread(
queries.replace_segments,
self.spec.download_id,
[
{"idx": p.idx, "start_offset": p.start, "end_offset": p.end, "bytes_done": 0}
for p in plans
],
)
self.state.segments = [SegmentRuntime(p.idx, p.start, p.end, 0) for p in plans]
else:
# Single-stream: one logical segment; bytes_done tracked on the row.
row = await asyncio.to_thread(queries.get_download, self.spec.download_id)
resume_from = row.bytes_done if row else 0
end = (pr.total_bytes - 1) if pr.total_bytes else -1
# ``row.bytes_done`` may be the SUM of per-segment offsets from a
# prior segmented run (a preallocated, non-contiguous .part). A
# single-stream resume writes a contiguous prefix, so the offset is
# only trustworthy when the on-disk file is exactly that many
# contiguous bytes. This guards the case where a download that ran
# segmented now resolves to one segment (server dropped
# Accept-Ranges, or --download-segments was lowered between runs):
# resuming over non-contiguous data would corrupt the output.
if resume_from > 0 and not self._contiguous_prefix_valid(resume_from):
logging.info(
"[model_downloader] %s discarding untrusted resume offset "
"%d (on-disk .part not a contiguous prefix); restarting",
self.spec.model_id, resume_from,
)
resume_from = 0
self._remove_temp()
if await asyncio.to_thread(queries.list_segments, self.spec.download_id):
await asyncio.to_thread(
queries.replace_segments, self.spec.download_id, []
)
await asyncio.to_thread(
queries.update_download, self.spec.download_id, bytes_done=0
)
self.state.segments = [SegmentRuntime(0, 0, end, resume_from)]
self._recompute_bytes_done()
return pr
# ----- transfer -----
async def _transfer(self, pr) -> None:
self._writer = FileWriter(self.spec.temp_path)
await self._writer.open()
segmented = len(self.state.segments) > 1
if segmented and self.state.total_bytes:
await self._writer.preallocate(self.state.total_bytes)
await self._run_segmented()
else:
await self._run_single()
await self._writer.flush()
async def _run_segmented(self) -> None:
pending = [
asyncio.ensure_future(self._run_segment(seg))
for seg in self.state.segments
if seg.bytes_done < seg.length
]
if not pending:
return
done, not_done = await asyncio.wait(
pending, return_when=asyncio.FIRST_EXCEPTION
)
first_exc: Optional[BaseException] = None
for task in done:
exc = task.exception()
if exc is not None and first_exc is None:
first_exc = exc
if first_exc is not None:
for task in not_done:
task.cancel()
await asyncio.gather(*not_done, return_exceptions=True)
raise first_exc
async def _run_segment(self, seg: SegmentRuntime) -> None:
offset = seg.start + seg.bytes_done
headers = {
"Range": f"bytes={offset}-{seg.end}",
"Accept-Encoding": "identity",
}
if self._etag:
headers["If-Range"] = self._etag
async with open_validated(
"GET", self.spec.url, headers=headers
) as (resp, _final):
if resp.status == 200:
# Server ignored the range -> remote changed / no resume support.
raise RemoteChanged()
if resp.status not in (206,):
self._raise_for_status(resp.status)
async for chunk in resp.content.iter_chunked(args.download_chunk_size):
self._check_control()
# Never write past this segment's planned range: a
# non-conforming 206 that returns more than the requested
# bytes would otherwise overrun adjacent segments and the
# preallocated file. Cap the write and abort on overflow.
remaining = seg.length - seg.bytes_done
if remaining <= 0:
raise FatalError(
f"segment {seg.idx}: server returned more than the "
f"requested {seg.length} bytes"
)
overflow = len(chunk) > remaining
if overflow:
chunk = chunk[:remaining]
await self._writer.write_at(offset, chunk)
offset += len(chunk)
seg.bytes_done += len(chunk)
self._recompute_bytes_done()
await self._persist_progress()
if overflow:
raise FatalError(
f"segment {seg.idx}: server returned more than the "
f"requested {seg.length} bytes"
)
async def _run_single(self) -> None:
seg = self.state.segments[0]
offset = seg.bytes_done # resume from here for single-stream
headers = {"Accept-Encoding": "identity"}
if offset > 0:
headers["Range"] = f"bytes={offset}-"
if self._etag:
headers["If-Range"] = self._etag
async with open_validated(
"GET", self.spec.url, headers=headers
) as (resp, _final):
if offset > 0 and resp.status == 200:
# Resume not honoured -> start over from the beginning. Truncate
# the existing partial so stale trailing bytes from the prior
# attempt cannot survive past the new (possibly shorter) end.
offset = 0
seg.bytes_done = 0
self.state.bytes_done = 0
await self._writer.truncate(0)
elif offset > 0 and resp.status != 206:
self._raise_for_status(resp.status)
elif offset == 0 and resp.status != 200:
self._raise_for_status(resp.status)
# Byte ceiling for this stream: the known total when the server
# reported a size, otherwise the configured maximum download size.
# Without a bound, a non-conforming response or an unknown-length
# stream (end == -1) that never closes could fill the disk (DoS).
limit = (seg.end + 1) if seg.end >= 0 else self._max_download_bytes()
async for chunk in resp.content.iter_chunked(args.download_chunk_size):
self._check_control()
overflow = False
if limit is not None:
remaining = limit - offset
if remaining <= 0:
raise FatalError(
f"download exceeded the maximum size {limit} bytes"
)
if len(chunk) > remaining:
chunk = chunk[:remaining]
overflow = True
await self._writer.write_at(offset, chunk)
offset += len(chunk)
seg.bytes_done = offset
self.state.bytes_done = offset
await self._persist_progress()
if overflow:
raise FatalError(
f"download exceeded the maximum size {limit} bytes"
)
def _max_download_bytes(self) -> Optional[int]:
"""Configured maximum download size in bytes, or ``None`` if disabled."""
cap = getattr(args, "download_max_bytes", 0)
return cap if cap and cap > 0 else None
def _raise_for_status(self, status: int) -> None:
if status in (401, 403):
raise FatalError(
f"{redact_url(self.spec.url)} returned {status}; authenticate this "
f"host via /api/download/auth or set its API key env var."
)
if status in _RETRYABLE_STATUSES:
raise RetryableError(f"HTTP {status}")
raise FatalError(f"unexpected HTTP {status}")
# ----- finalize / verify (PRD section 8.4) -----
async def _finalize(self) -> None:
self._check_control()
await self._close_writer()
await self._set_status(DownloadStatus.VERIFYING)
total = self.state.total_bytes
segmented = len(self.state.segments) > 1
if segmented:
# The .part was preallocated to total_bytes, so its on-disk size is
# not evidence of completeness: a segment that ends short (truncated
# 206 / server closes mid-range) leaves a zero-filled hole while the
# file size still equals total. Verify each segment wrote its full
# planned range, and trust the byte counter (== sum of segments)
# rather than os.path.getsize for the total check.
for seg in self.state.segments:
if seg.bytes_done != seg.length:
raise FatalError(
f"segment {seg.idx} incomplete: wrote {seg.bytes_done} "
f"of {seg.length} bytes"
)
observed = self.state.bytes_done
else:
# Single-stream writes a contiguous prefix, so the on-disk size is
# an independent witness of how much actually landed.
observed = os.path.getsize(self.spec.temp_path)
if total is not None and observed != total:
raise FatalError(
f"size mismatch: wrote {observed} of {total} bytes"
)
# Structural gate (cheap, no full read) then optional sha256 (full read).
# Both failures are non-retryable (a truncated/corrupt or mismatched file
# will not heal on retry), so surface them as FatalError rather than
# letting the plain Exceptions fall through to the retryable handler.
# ``temp_path`` carries the ``.part`` suffix; pass ``dest_path`` so the
# structural check detects the real file format instead of skipping it.
try:
await asyncio.to_thread(
structural.validate, self.spec.temp_path, self.spec.dest_path
)
if self.spec.expected_sha256:
await asyncio.to_thread(
checksum.verify_sha256,
self.spec.temp_path,
self.spec.expected_sha256,
)
except (structural.StructuralError, checksum.ChecksumError) as e:
raise FatalError(str(e)) from e
os.makedirs(os.path.dirname(self.spec.dest_path), exist_ok=True)
os.replace(self.spec.temp_path, self.spec.dest_path)
logging.info(
"[model_downloader] completed %s (%d bytes)",
self.spec.model_id, observed,
)
# Catalog into the assets system (blake3 dedup identity). Best-effort.
await dedup.register_completed(self.spec.dest_path)
# ----- helpers -----
def _recompute_bytes_done(self) -> None:
self.state.bytes_done = sum(s.bytes_done for s in self.state.segments)
now = time.monotonic()
dt = now - self.state._last_time
if dt >= 0.5:
self.state.speed_bps = (self.state.bytes_done - self.state._last_bytes) / dt
self.state._last_bytes = self.state.bytes_done
self.state._last_time = now
async def _persist_progress(self, force: bool = False) -> None:
# Both the DB write and the websocket notify are gated by the same
# throttle: persisting hits SQLite, and notifying broadcasts to every
# client, so doing either per-chunk (small --download-chunk-size or
# many concurrent segments) would overwhelm both. Skip entirely inside
# the window; the next persist (or a forced one) ships the latest bytes.
now = time.monotonic()
if not force and now - self._last_persist < _PERSIST_INTERVAL:
return
self._last_persist = now
# SQLite is blocking; run it off the event loop per the queries module
# contract so progress persists don't stall the web server.
await asyncio.to_thread(self._write_progress)
if self._notify:
self._notify(self.spec.download_id)
def _write_progress(self) -> None:
queries.update_download(self.spec.download_id, bytes_done=self.state.bytes_done)
for seg in self.state.segments:
if seg.end >= seg.start: # skip unknown-size sentinel
queries.update_segment_progress(
self.spec.download_id, seg.idx, seg.bytes_done
)
async def _reset_for_restart(self) -> None:
await self._close_writer()
self._remove_temp()
for seg in self.state.segments:
seg.bytes_done = 0
self.state.bytes_done = 0
await asyncio.to_thread(
queries.update_download, self.spec.download_id, bytes_done=0
)
if await asyncio.to_thread(queries.list_segments, self.spec.download_id):
await asyncio.to_thread(
queries.replace_segments, self.spec.download_id, []
)
async def _close_writer(self) -> None:
if self._writer is not None:
try:
await self._writer.close()
except Exception:
logging.debug("[model_downloader] writer close error", exc_info=True)
self._writer = None
def _segmented_part_valid(self, total_bytes: int) -> bool:
"""True when the temp file is the preallocated segmented ``.part``.
A segmented transfer preallocates the .part to ``total_bytes`` up front
and tracks how much of each range landed via per-segment offsets. Those
offsets are only trustworthy when the file they describe is still on
disk at its full preallocated size. A missing file (swept after a
failure, removed on a fatal error, deleted by hand) or a wrong-sized one
means the persisted offsets no longer correspond to real bytes and must
not be resumed over. Doing so would skip "complete" segments and leave
zero-filled holes that pass the size-only verification gate.
"""
try:
return os.path.getsize(self.spec.temp_path) == total_bytes
except OSError:
return False
def _contiguous_prefix_valid(self, prefix_len: int) -> bool:
"""True when the temp file is exactly ``prefix_len`` contiguous bytes.
Single-stream resume appends sequentially, so a valid resume point
implies the .part size equals the persisted offset. A larger file (e.g.
one preallocated to ``total_bytes`` by a previous segmented run) or a
missing/short file means the persisted offset is not a trustworthy
contiguous prefix and must not be resumed over.
"""
try:
return os.path.getsize(self.spec.temp_path) == prefix_len
except OSError:
return False
def _remove_temp(self) -> None:
try:
os.remove(self.spec.temp_path)
except FileNotFoundError:
pass
except OSError as e:
logging.warning(
"[model_downloader] could not remove %s: %s", self.spec.temp_path, e
)
async def _set_status(self, status: str, error: Optional[str] = None) -> None:
# ``error`` is authoritative: passing None clears any prior failure
# text so transitions out of a failure state (retry/success) don't
# leave stale messages on RuntimeState or in the persisted row.
self.state.status = status
self.state.error = error
fields = {"status": status, "bytes_done": self.state.bytes_done, "error": error}
if status == DownloadStatus.QUEUED:
fields["attempts"] = self.spec.attempts + 1
self.spec.attempts += 1
await asyncio.to_thread(queries.update_download, self.spec.download_id, **fields)
if self._notify:
self._notify(self.spec.download_id)
+51
View File
@@ -0,0 +1,51 @@
"""Segment planning.
Split a known byte range into S roughly-equal segments, each fetched by its
own coroutine with ``Range: bytes=start-end``. Falls back to a single segment
when the server doesn't support ranges or the size is unknown/too small for
segmentation to be worthwhile.
"""
from __future__ import annotations
from dataclasses import dataclass
# Below this size, the per-connection setup cost outweighs any parallelism.
_MIN_SEGMENT_BYTES = 1 * 1024 * 1024
@dataclass(frozen=True)
class SegmentPlan:
idx: int
start: int
end: int # inclusive
@property
def length(self) -> int:
return self.end - self.start + 1
def effective_segment_count(
total_bytes: int | None, accept_ranges: bool, configured: int
) -> int:
"""How many segments to actually use for this file."""
if not accept_ranges or total_bytes is None or total_bytes <= 0:
return 1
by_size = max(1, total_bytes // _MIN_SEGMENT_BYTES)
return max(1, min(configured, by_size))
def plan_segments(total_bytes: int, num_segments: int) -> list[SegmentPlan]:
"""Return ``num_segments`` contiguous, inclusive byte ranges covering [0, total)."""
if total_bytes <= 0 or num_segments <= 1:
return [SegmentPlan(idx=0, start=0, end=max(0, total_bytes - 1))]
base = total_bytes // num_segments
plans: list[SegmentPlan] = []
start = 0
for i in range(num_segments):
# Last segment soaks up the remainder.
length = base if i < num_segments - 1 else total_bytes - start
end = start + length - 1
plans.append(SegmentPlan(idx=i, start=start, end=end))
start = end + 1
return plans
+110
View File
@@ -0,0 +1,110 @@
"""Positioned, off-loop file writes.
Network I/O stays on the event loop; every blocking disk op (preallocate,
positioned write, fsync) is run in a bounded thread pool via
``run_in_executor`` so downloads never stall inference or the web server.
A single file descriptor is opened for the whole download. Segments write to
their own offsets with ``os.pwrite`` — which is offset-addressed and atomic
per call, so concurrent segment writers need no extra locking. Per-chunk
fsync is avoided; we fsync once at completion.
``os.pwrite`` is unavailable on Windows, so there we fall back to
``os.lseek`` + ``os.write`` guarded by a per-writer lock (the seek/write pair
is not atomic, so concurrent segment writers must be serialized).
"""
from __future__ import annotations
import asyncio
import os
import threading
from concurrent.futures import ThreadPoolExecutor
from typing import Optional
# One shared, bounded pool for all download disk I/O.
_EXECUTOR = ThreadPoolExecutor(max_workers=8, thread_name_prefix="dl-writer")
_HAS_PWRITE = hasattr(os, "pwrite")
# On Windows ``os.open`` defaults to text mode, which translates every ``\n``
# byte into ``\r\n`` on write and corrupts binary payloads (the file grows by
# one byte per 0x0A). ``O_BINARY`` disables that translation; it does not exist
# on POSIX, where the default is already binary.
_O_BINARY = getattr(os, "O_BINARY", 0)
class FileWriter:
"""Owns the ``.part`` file descriptor for one download."""
def __init__(self, path: str) -> None:
self.path = path
self._fd: Optional[int] = None
# Serializes lseek+write on platforms without os.pwrite (Windows).
self._seek_lock = threading.Lock()
def _open(self) -> None:
os.makedirs(os.path.dirname(self.path), exist_ok=True)
self._fd = os.open(self.path, os.O_RDWR | os.O_CREAT | _O_BINARY, 0o644)
async def open(self) -> None:
await asyncio.get_running_loop().run_in_executor(_EXECUTOR, self._open)
async def preallocate(self, size: int) -> None:
"""Grow the file to ``size`` so segments write to their offsets."""
if self._fd is None or size <= 0:
return
await asyncio.get_running_loop().run_in_executor(
_EXECUTOR, os.ftruncate, self._fd, size
)
async def truncate(self, size: int = 0) -> None:
"""Truncate the file to ``size`` bytes (default: empty it)."""
if self._fd is None:
return
await asyncio.get_running_loop().run_in_executor(
_EXECUTOR, os.ftruncate, self._fd, size
)
def _pwrite_all(self, data: bytes, offset: int) -> None:
"""A positioned write may write fewer bytes than requested (signal
interruption, near-ENOSPC); loop until every byte lands so we never
leave a gap while the caller advances by the full chunk length.
Uses ``os.pwrite`` where available (offset-addressed, atomic per call).
On Windows it falls back to ``os.lseek`` + ``os.write`` under a lock,
since that pair is not atomic across concurrent segment writers."""
assert self._fd is not None, "writer not opened"
view = memoryview(data)
written = 0
total = len(view)
while written < total:
if _HAS_PWRITE:
n = os.pwrite(self._fd, view[written:], offset + written)
else:
with self._seek_lock:
os.lseek(self._fd, offset + written, os.SEEK_SET)
n = os.write(self._fd, view[written:])
if n == 0:
raise OSError(
f"positioned write wrote 0 bytes at offset {offset + written} "
f"({written}/{total} bytes written)"
)
written += n
async def write_at(self, offset: int, data: bytes) -> None:
assert self._fd is not None, "writer not opened"
await asyncio.get_running_loop().run_in_executor(
_EXECUTOR, self._pwrite_all, data, offset
)
async def flush(self) -> None:
if self._fd is None:
return
await asyncio.get_running_loop().run_in_executor(_EXECUTOR, os.fsync, self._fd)
async def close(self) -> None:
if self._fd is None:
return
fd, self._fd = self._fd, None
await asyncio.get_running_loop().run_in_executor(_EXECUTOR, os.close, fd)
+444
View File
@@ -0,0 +1,444 @@
"""Public facade for the download manager.
This is the only object the server imports. It validates requests, owns the
:class:`Scheduler`, and exposes a small async API plus read models for status.
"""
from __future__ import annotations
import asyncio
import logging
import os
import uuid
from typing import Callable, Optional
from app.model_downloader.constants import DownloadStatus
from app.model_downloader.database import queries
from app.model_downloader.net.probe import gated_error_message, probe
from app.model_downloader.scheduler import SCHEDULER
from app.model_downloader.security import paths
from app.model_downloader.net.http import redact_url
from app.model_downloader.security.allowlist import (
ALLOWED_MODEL_EXTENSIONS,
filename_extension,
is_host_allowed_url,
is_url_downloadable,
url_path_extension,
)
from app.model_downloader.security.paths import InvalidModelId
# Non-terminal statuses: an existing row in one of these blocks a re-enqueue.
_LIVE_STATUSES = (
DownloadStatus.QUEUED,
DownloadStatus.ACTIVE,
DownloadStatus.PAUSED,
DownloadStatus.VERIFYING,
)
class DownloadError(Exception):
"""A user-facing error with a stable machine-readable code."""
def __init__(self, code: str, message: str, status: int = 400) -> None:
super().__init__(message)
self.code = code
self.message = message
self.http_status = status
class DownloadManager:
def __init__(self) -> None:
self._scheduler = SCHEDULER
self._notify_cb: Optional[Callable[[str], None]] = None
# Serializes the "check for a live download, then write" critical section
# per model_id. ``downloads`` has no uniqueness constraint on model_id
# (history rows are kept), so without this two concurrent enqueue/resume
# calls could both pass the live check and admit two jobs sharing one
# temp/dest path. The manager is a process singleton over a local SQLite
# DB, so an in-process lock is sufficient (and avoids a migration).
self._model_locks: dict[str, asyncio.Lock] = {}
def set_notify(self, cb: Optional[Callable[[str], None]]) -> None:
self._notify_cb = cb
self._scheduler.set_notify(cb)
async def start(self) -> None:
await self._scheduler.start()
# ----- enqueue -----
async def enqueue(
self,
url: str,
model_id: str,
*,
priority: int = 0,
expected_sha256: Optional[str] = None,
allow_any_extension: bool = False,
) -> str:
# Coarse gate first: host/scheme must be allowlisted, and any extension
# present in the URL path must be a known model type. A URL whose path
# carries NO extension (e.g. Civitai's ``/api/download/models/<id>``) is
# admitted here and its real extension is resolved from the network
# below before the download is finally accepted.
if allow_any_extension:
if not is_host_allowed_url(url):
raise DownloadError(
"URL_NOT_ALLOWED",
"URL is not on the download allowlist (host/scheme).",
)
elif not is_url_downloadable(url):
raise DownloadError(
"URL_NOT_ALLOWED",
"URL is not on the download allowlist (host/scheme/extension).",
)
# When the URL path has no extension, follow it to where it resolves and
# adopt the real extension from the response, forcing the stored
# filename to match. Skipped when the caller opted into any extension.
if not allow_any_extension and url_path_extension(url) == "":
resolved_ext = await self._resolve_extension(url)
model_id = paths.apply_extension(model_id, resolved_ext)
try:
paths.parse_model_id(model_id, allow_any_extension)
dest_path, temp_path = paths.resolve_destination(model_id, allow_any_extension)
except InvalidModelId as e:
raise DownloadError("INVALID_MODEL_ID", str(e))
if await asyncio.to_thread(
paths.resolve_existing, model_id, allow_any_extension
):
raise DownloadError(
"ALREADY_AVAILABLE",
f"Model already exists on disk: {model_id}",
status=409,
)
download_id = str(uuid.uuid4())
# Hold the per-model lock across the live check and the insert so a
# concurrent enqueue/resume for the same model_id cannot interleave
# between them and create a second job against the same temp/dest path.
async with self._model_lock(model_id):
if await self._has_live_download(model_id):
raise DownloadError(
"ALREADY_DOWNLOADING",
f"A download for {model_id} is already in progress.",
status=409,
)
await asyncio.to_thread(
queries.insert_download,
{
"id": download_id,
"url": url,
"model_id": model_id,
"dest_path": dest_path,
"temp_path": temp_path,
"status": DownloadStatus.QUEUED,
"priority": priority,
"expected_sha256": expected_sha256,
"allow_any_extension": allow_any_extension,
},
)
logging.info("[model_downloader] enqueued %s -> %s", redact_url(url), model_id)
await self._scheduler.pump()
return download_id
async def _resolve_extension(self, url: str) -> str:
"""Follow ``url`` to its final response and return the real extension.
Used for allowlisted URLs whose path has no extension (e.g. Civitai
download endpoints): the filename lives in the ``Content-Disposition``
header or the post-redirect URL. Raises :class:`DownloadError` when the
URL can't be resolved, needs authentication, or resolves to something
that is not a known model file — so we never persist a bogus destination.
"""
pr = await probe(url)
if not pr.ok:
if pr.gated:
raise DownloadError(
"GATED_REPO" if pr.is_gated_repo else "CREDENTIALS_REQUIRED",
gated_error_message(url, pr),
status=401,
)
raise DownloadError(
"URL_RESOLVE_FAILED",
f"Could not resolve {redact_url(url)}: {pr.error or 'unknown error'}",
status=502,
)
ext = filename_extension(pr.filename) if pr.filename else ""
if ext not in ALLOWED_MODEL_EXTENSIONS:
raise DownloadError(
"URL_NOT_ALLOWED",
f"URL resolves to {pr.filename or '<unknown>'!r}, which is not a "
f"known model file type {ALLOWED_MODEL_EXTENSIONS}.",
)
return ext
def _model_lock(self, model_id: str) -> asyncio.Lock:
# Lazily create one lock per model_id. There is no ``await`` between the
# lookup and the insert, so under the single asyncio thread this is
# atomic and cannot hand out two different locks for the same model_id.
lock = self._model_locks.get(model_id)
if lock is None:
lock = asyncio.Lock()
self._model_locks[model_id] = lock
return lock
async def _has_live_download(
self, model_id: str, *, exclude_id: Optional[str] = None
) -> bool:
return await asyncio.to_thread(
queries.has_live_download_for_model, model_id, _LIVE_STATUSES, exclude_id
)
# ----- control -----
async def pause(self, download_id: str) -> None:
job = self._scheduler.get_job(download_id)
if job is not None:
job.request_pause()
return
row = await asyncio.to_thread(queries.get_download, download_id)
if row is None:
raise DownloadError("NOT_FOUND", "No such download.", status=404)
if row.status == DownloadStatus.QUEUED:
await asyncio.to_thread(
queries.update_download, download_id, status=DownloadStatus.PAUSED
)
async def resume(self, download_id: str) -> None:
row = await asyncio.to_thread(queries.get_download, download_id)
if row is None:
raise DownloadError("NOT_FOUND", "No such download.", status=404)
if row.status not in (DownloadStatus.PAUSED, DownloadStatus.FAILED):
return
# Re-queueing a paused/failed row must respect the single-live-per-model
# invariant: another download (e.g. a newer enqueue) may already be live
# for this model_id and would share this row's temp/dest path. Hold the
# per-model lock across the check and the status flip, and exclude this
# row itself (a paused row is already a "live" status).
async with self._model_lock(row.model_id):
if await self._has_live_download(row.model_id, exclude_id=download_id):
raise DownloadError(
"ALREADY_DOWNLOADING",
f"A download for {row.model_id} is already in progress.",
status=409,
)
await asyncio.to_thread(
queries.update_download,
download_id,
status=DownloadStatus.QUEUED,
error=None,
)
await self._scheduler.pump()
async def cancel(self, download_id: str) -> None:
job = self._scheduler.get_job(download_id)
if job is not None:
job.request_cancel()
return
row = await asyncio.to_thread(queries.get_download, download_id)
if row is None:
raise DownloadError("NOT_FOUND", "No such download.", status=404)
if row.status in _LIVE_STATUSES:
try:
os.remove(row.temp_path)
except OSError:
pass
await asyncio.to_thread(
queries.update_download, download_id, status=DownloadStatus.CANCELLED
)
async def set_priority(self, download_id: str, priority: int) -> None:
row = await asyncio.to_thread(queries.get_download, download_id)
if row is None:
raise DownloadError("NOT_FOUND", "No such download.", status=404)
await asyncio.to_thread(
queries.update_download, download_id, priority=priority
)
# Admission-order only; a higher priority is
# picked up the next time a slot frees. Pump in case a slot is free now.
await self._scheduler.pump()
async def delete(self, download_id: str) -> None:
"""Delete a terminal download so it stays gone from history.
Refuses to delete a live download so a record is never removed out from
under a running worker; cancel it first. Any leftover ``.part`` temp
file (e.g. from a failed transfer) is removed, but the finished model
file on disk is never touched.
"""
if self._scheduler.get_job(download_id) is not None:
raise DownloadError(
"DOWNLOAD_ACTIVE",
"Cannot delete a download that is still in progress.",
status=409,
)
row = await asyncio.to_thread(queries.get_download, download_id)
if row is None:
raise DownloadError("NOT_FOUND", "No such download.", status=404)
if row.status in _LIVE_STATUSES:
raise DownloadError(
"DOWNLOAD_ACTIVE",
"Cannot delete a download that is still in progress.",
status=409,
)
try:
os.remove(row.temp_path)
except OSError:
pass
await asyncio.to_thread(queries.delete_download, download_id)
async def clear(self) -> int:
"""Delete all terminal downloads from history in one transaction.
Skips anything still live (queued/active/paused/verifying, or a running
job) so an in-flight download is never removed out from under a worker.
Finished model files on disk are never touched; only leftover ``.part``
temp files from failed/cancelled transfers are removed. Returns the
number of history rows deleted.
"""
rows = await asyncio.to_thread(queries.list_downloads)
deletable = [
r
for r in rows
if r.status not in _LIVE_STATUSES
and self._scheduler.get_job(r.id) is None
]
if not deletable:
return 0
for r in deletable:
try:
os.remove(r.temp_path)
except OSError:
pass
return await asyncio.to_thread(
queries.delete_downloads, [r.id for r in deletable]
)
# ----- read models -----
def _view(self, row) -> dict:
"""Combine the persisted row with live in-memory progress, if running."""
job = self._scheduler.get_job(row.id)
bytes_done = row.bytes_done
total = row.total_bytes
speed = None
eta = None
segments = None
if job is not None:
st = job.state
bytes_done = st.bytes_done
total = st.total_bytes if st.total_bytes is not None else total
speed = st.speed_bps
eta = st.eta_seconds
segments = [
{"idx": s.idx, "bytes_done": s.bytes_done, "length": s.length}
for s in st.segments
if s.end >= s.start
]
progress = (bytes_done / total) if total else None
return {
"download_id": row.id,
"model_id": row.model_id,
"url": redact_url(row.url),
"status": row.status,
"priority": row.priority,
"total_bytes": total,
"bytes_done": bytes_done,
"progress": progress,
"speed_bps": speed,
"eta_seconds": eta,
"segments": segments,
"error": row.error,
"created_at": row.created_at,
"updated_at": row.updated_at,
}
def _view_from_state(self, job) -> dict:
"""Build a view purely from the live in-memory job state (no DB)."""
st = job.state
return {
"download_id": st.download_id,
"model_id": st.model_id,
"url": redact_url(st.url),
"status": st.status,
"priority": st.priority,
"total_bytes": st.total_bytes,
"bytes_done": st.bytes_done,
"progress": st.progress,
"speed_bps": st.speed_bps,
"eta_seconds": st.eta_seconds,
"segments": [
{"idx": s.idx, "bytes_done": s.bytes_done, "length": s.length}
for s in st.segments
if s.end >= s.start
],
"error": st.error,
}
def status_sync(self, download_id: str) -> Optional[dict]:
"""Synchronous status read for the websocket notify path.
Uses live in-memory state when the job is running (no DB round-trip on
the hot path); falls back to a quick DB read otherwise.
"""
job = self._scheduler.get_job(download_id)
if job is not None:
return self._view_from_state(job)
row = queries.get_download(download_id)
return self._view(row) if row is not None else None
async def status(self, download_id: str) -> Optional[dict]:
row = await asyncio.to_thread(queries.get_download, download_id)
return self._view(row) if row is not None else None
async def list(self) -> list[dict]:
rows = await asyncio.to_thread(queries.list_downloads)
return [self._view(r) for r in rows]
async def availability(self, models: dict[str, str]) -> dict[str, dict]:
"""Bulk per-id ``{state, progress, ...}`` for the frontend poll.
``state`` is ``available`` (on disk), ``downloading`` (live row), or
``missing``. Cheap: a path lookup plus an in-memory/DB status check.
"""
rows = await asyncio.to_thread(queries.list_downloads)
by_model: dict[str, object] = {}
for r in rows:
if r.status in _LIVE_STATUSES or r.model_id not in by_model:
by_model[r.model_id] = r
# ``url_allowed`` mirrors the coarse enqueue gate (host/scheme + a
# non-disallowed extension); URLs whose extension is only known after a
# network resolve — e.g. Civitai download endpoints — report allowed.
out: dict[str, dict] = {}
for model_id, url in models.items():
try:
exists = await asyncio.to_thread(paths.resolve_existing, model_id)
except InvalidModelId:
out[model_id] = {"state": "missing", "url_allowed": is_url_downloadable(url)}
continue
if exists:
out[model_id] = {"state": "available", "url_allowed": is_url_downloadable(url)}
continue
row = by_model.get(model_id)
if row is not None and row.status in _LIVE_STATUSES:
view = self._view(row)
out[model_id] = {
"state": "downloading",
"url_allowed": is_url_downloadable(url),
"download_id": view["download_id"],
"progress": view["progress"],
"bytes_done": view["bytes_done"],
"total_bytes": view["total_bytes"],
"speed_bps": view["speed_bps"],
}
else:
out[model_id] = {"state": "missing", "url_allowed": is_url_downloadable(url)}
return out
DOWNLOAD_MANAGER = DownloadManager()
+142
View File
@@ -0,0 +1,142 @@
"""Manual, validated redirect-following request opener.
Automatic redirects are disabled. We follow hops ourselves
so that on *every* hop we (a) re-validate scheme + reject credentials-in-URL,
(b) recompute which auth — if any — applies to that hop's host, and (c) let the
connector's resolver screen the IP. This is the single place that attaches a
token, so it can never ride a redirect to a CDN host.
"""
from __future__ import annotations
import logging
import re
from contextlib import asynccontextmanager
from typing import AsyncIterator, Optional
from urllib.parse import unquote, urljoin, urlsplit, urlunsplit
import aiohttp
from app.model_downloader.auth.resolver import resolve_auth_for_hop
from app.model_downloader.net.session import get_session
from app.model_downloader.security.ssrf import (
MAX_REDIRECTS,
SSRFError,
check_redirect_hop,
)
_REDIRECT_CODES = {301, 302, 303, 307, 308}
DEFAULT_TIMEOUT = aiohttp.ClientTimeout(total=None, sock_connect=30, sock_read=120)
def redact_url(url: str) -> str:
"""Drop the query string so a query-scheme secret is never logged/stored."""
try:
parts = urlsplit(url)
except ValueError:
return "<unparseable-url>"
return urlunsplit(parts._replace(query=""))
_CD_FILENAME_STAR = re.compile(
r"filename\*\s*=\s*[^']*'[^']*'([^;]+)", re.IGNORECASE
)
_CD_FILENAME_QUOTED = re.compile(r'filename\s*=\s*"([^"]+)"', re.IGNORECASE)
_CD_FILENAME_BARE = re.compile(r"filename\s*=\s*([^;]+)", re.IGNORECASE)
def filename_from_content_disposition(value: Optional[str]) -> Optional[str]:
"""Extract the download filename from a ``Content-Disposition`` header.
Prefers the RFC 5987 ``filename*=`` form (percent-decoded) over the plain
``filename=`` form. Any directory components in the value are stripped so a
hostile header can only influence the *name*, never the target directory.
Returns ``None`` when no filename is present.
"""
if not value:
return None
for pat, decode in (
(_CD_FILENAME_STAR, True),
(_CD_FILENAME_QUOTED, False),
(_CD_FILENAME_BARE, False),
):
m = pat.search(value)
if not m:
continue
raw = m.group(1).strip().strip('"')
if decode:
try:
raw = unquote(raw)
except Exception:
pass
name = raw.replace("\\", "/").rsplit("/", 1)[-1].strip()
if name:
return name
return None
async def _resolve_final_response(
method: str,
url: str,
base_headers: dict[str, str],
timeout: aiohttp.ClientTimeout,
) -> tuple[aiohttp.ClientResponse, str]:
"""Follow redirects manually until a non-redirect response.
Each intermediate redirect response is released before the next hop.
Returns the final ``(response, final_url)``; the caller owns releasing it.
"""
session = await get_session()
current = url
hops = 0
while True:
check_redirect_hop(current, is_initial_url=(hops == 0))
parts = urlsplit(current)
auth = await resolve_auth_for_hop(parts.hostname or "", parts.scheme)
req_headers = dict(base_headers)
if auth is not None:
req_headers.update(auth.headers)
resp = await session.request(
method,
current,
allow_redirects=False,
headers=req_headers,
timeout=timeout,
)
if resp.status in _REDIRECT_CODES and resp.headers.get("Location"):
next_url = urljoin(str(resp.url), resp.headers["Location"])
await resp.release()
hops += 1
if hops > MAX_REDIRECTS:
raise SSRFError(
f"too many redirects (> {MAX_REDIRECTS}) for {redact_url(url)}"
)
current = next_url
continue
return resp, redact_url(str(resp.url))
@asynccontextmanager
async def open_validated(
method: str,
url: str,
*,
headers: Optional[dict[str, str]] = None,
timeout: aiohttp.ClientTimeout = DEFAULT_TIMEOUT,
) -> AsyncIterator[tuple[aiohttp.ClientResponse, str]]:
"""Open ``method url`` following redirects manually and validated.
Yields ``(response, final_url)`` where ``final_url`` is redacted of any
query string. The response is released automatically on exit.
"""
resp, final_url = await _resolve_final_response(
method, url, dict(headers or {}), timeout
)
try:
yield resp, final_url
finally:
try:
await resp.release()
except Exception: # pragma: no cover - best-effort cleanup
logging.debug("[model_downloader] response release error", exc_info=True)
+192
View File
@@ -0,0 +1,192 @@
"""Pre-download probe.
Issues a tiny ranged GET (``Range: bytes=0-0``) — which doubles as a
range-support test — to discover ``Content-Length``, ``Accept-Ranges``,
``ETag``/``Last-Modified``, and the final post-redirect URL. For HuggingFace
LFS files the true size also appears in the non-standard ``X-Linked-Size``
header, which we read as a fallback.
"""
from __future__ import annotations
import logging
from dataclasses import dataclass
from typing import Optional
from urllib.parse import urlparse, urlsplit
import aiohttp
from app.model_downloader.net.http import (
filename_from_content_disposition,
open_validated,
redact_url,
)
from app.model_downloader.net.session import parse_int_header
_PROBE_TIMEOUT = aiohttp.ClientTimeout(total=60, sock_connect=30, sock_read=30)
@dataclass
class ProbeResult:
ok: bool
status: int
final_url: Optional[str] = None
total_bytes: Optional[int] = None
accept_ranges: bool = False
etag: Optional[str] = None
last_modified: Optional[str] = None
gated: bool = False # 401/403 — needs (or has wrong) credentials
error: Optional[str] = None
# HuggingFace's ``X-Error-Code`` header (e.g. ``GatedRepo``,
# ``RepoNotFound``) when the host reports one. Lets us tell "this repo is
# gated — request access" apart from "you just need a token".
error_code: Optional[str] = None
# Filename the server intends this response to be saved as: the
# ``Content-Disposition`` name if present, else the post-redirect URL's
# basename. Used to resolve the real extension for URLs (e.g. Civitai's
# ``/api/download`` endpoints) that carry no extension in their path.
filename: Optional[str] = None
@property
def is_gated_repo(self) -> bool:
"""True when the host says the repo is gated (access must be granted).
Distinct from a plain missing/invalid token: even a valid credential
won't help until the user accepts the model's terms on its page.
"""
return (self.error_code or "").lower() == "gatedrepo"
def _error_detail(error_code: Optional[str], error_message: Optional[str]) -> str:
"""Format the host's ``X-Error-Code``/``X-Error-Message`` for logs/messages."""
detail = ": ".join(p.strip() for p in (error_code, error_message) if p and p.strip())
return f" ({detail})" if detail else ""
def _probe_failure_message(
status: int, error_code: Optional[str], error_message: Optional[str]
) -> str:
msg = f"probe returned HTTP {status}{_error_detail(error_code, error_message)}"
if status == 404:
# HuggingFace returns 404 (not 403) for a private repo the current
# credentials cannot see, so it is indistinguishable from a missing
# file without the hint. Name both causes so the user can check the
# URL or their access/token scope.
msg += (
" — the file may not exist, or it is private/gated and the "
"credentials in use lack access to it"
)
return msg
def _total_from_content_range(value: Optional[str]) -> Optional[int]:
# "bytes 0-0/12345" -> 12345 ; "bytes 0-0/*" -> None
if not value or "/" not in value:
return None
total = value.rsplit("/", 1)[1].strip()
return parse_int_header(total)
def _filename_from_response(
content_disposition: Optional[str], final_url: Optional[str]
) -> Optional[str]:
name = filename_from_content_disposition(content_disposition)
if name:
return name
if final_url:
base = urlsplit(final_url).path.rsplit("/", 1)[-1]
if base:
return base
return None
async def probe(url: str) -> ProbeResult:
"""Probe ``url`` and return discovered metadata, failing soft."""
try:
async with open_validated(
"GET",
url,
headers={"Range": "bytes=0-0", "Accept-Encoding": "identity"},
timeout=_PROBE_TIMEOUT,
) as (resp, final_url):
# HuggingFace (and some others) report the real reason in these
# headers on any status, including 404 for a private/missing repo.
error_code = resp.headers.get("X-Error-Code")
error_message = resp.headers.get("X-Error-Message")
if resp.status in (401, 403):
logging.warning(
"[model_downloader] probe %s -> HTTP %d%s",
redact_url(final_url or url), resp.status,
_error_detail(error_code, error_message),
)
return ProbeResult(
ok=False, status=resp.status, final_url=final_url, gated=True,
error_code=error_code,
error=(
error_message
or f"host returned {resp.status} (authentication required)"
),
)
if resp.status not in (200, 206):
logging.warning(
"[model_downloader] probe %s -> HTTP %d%s",
redact_url(final_url or url), resp.status,
_error_detail(error_code, error_message),
)
return ProbeResult(
ok=False, status=resp.status, final_url=final_url,
error_code=error_code,
error=_probe_failure_message(resp.status, error_code, error_message),
)
headers = resp.headers
accept_ranges = False
total: Optional[int] = None
if resp.status == 206:
accept_ranges = True
total = _total_from_content_range(headers.get("Content-Range"))
else: # 200: server ignored the range
accept_ranges = headers.get("Accept-Ranges", "").lower() == "bytes"
total = parse_int_header(headers.get("Content-Length"))
if total is None:
total = parse_int_header(headers.get("X-Linked-Size"))
return ProbeResult(
ok=True,
status=resp.status,
final_url=final_url,
total_bytes=total,
accept_ranges=accept_ranges,
etag=headers.get("ETag"),
last_modified=headers.get("Last-Modified"),
filename=_filename_from_response(
headers.get("Content-Disposition"), final_url
),
)
except Exception as e: # network / SSRF / timeout
host = urlparse(url).netloc or "<unknown>"
logging.debug("[model_downloader] probe failed for %s: %s", host, type(e).__name__)
return ProbeResult(ok=False, status=0, error="probe failed: network error")
def gated_error_message(url: str, pr: ProbeResult) -> str:
"""Build a user-facing message for a gated/auth-required probe result.
Distinguishes a *gated* repo (access must be requested/granted on the model
page — a token alone is not enough) from a plain missing/invalid credential.
"""
redacted = redact_url(url)
if pr.is_gated_repo:
detail = (pr.error or "access is restricted").rstrip()
if detail and not detail.endswith((".", "!", "?")):
detail += "."
return (
f"{redacted} is a gated model — {detail} Request access on the model's "
f"page, authenticate this host via /api/download/auth (or set its API "
f"key env var), and retry."
)
return (
f"{redacted} requires authentication. Authenticate this host via "
f"/api/download/auth or set its API key env var, and retry."
)
+72
View File
@@ -0,0 +1,72 @@
"""Lazily-created shared :class:`aiohttp.ClientSession`.
A single session reuses TLS handshakes and TCP connections across the probe
and the many segment GETs to the same host (HuggingFace is the dominant
case), which is a large speedup on cold connections and exactly the
connection-reuse strategy that lets us match aria2c.
The connector uses :class:`ValidatingResolver` so every connection — initial
or post-redirect — is screened for private/special-use IPs at connect time.
TLS is pinned to certifi's CA bundle because the OS trust store is not wired
up on some Python installs (python.org macOS, slim containers).
"""
from __future__ import annotations
import asyncio
import ssl
from typing import Optional
import aiohttp
try:
import certifi
_CA_FILE = certifi.where()
except Exception: # pragma: no cover - certifi is a transitive dep of aiohttp
_CA_FILE = None
from comfy.cli_args import args
from app.model_downloader.security.ssrf import ValidatingResolver
_session: Optional[aiohttp.ClientSession] = None
_lock = asyncio.Lock()
def ssl_context() -> ssl.SSLContext:
if _CA_FILE is not None:
return ssl.create_default_context(cafile=_CA_FILE)
return ssl.create_default_context()
async def get_session() -> aiohttp.ClientSession:
"""Return the shared session, creating it on first use."""
global _session
if _session is not None and not _session.closed:
return _session
async with _lock:
if _session is None or _session.closed:
connector = aiohttp.TCPConnector(
limit_per_host=max(1, getattr(args, "download_max_connections_per_host", 16)),
ssl=ssl_context(),
resolver=ValidatingResolver(),
)
_session = aiohttp.ClientSession(connector=connector)
return _session
async def close_session() -> None:
global _session
if _session is not None and not _session.closed:
await _session.close()
_session = None
def parse_int_header(value: Optional[str]) -> Optional[int]:
"""Parse a non-negative integer header value, or None if bad/absent."""
if not value:
return None
try:
n = int(value)
except (TypeError, ValueError):
return None
return n if n >= 0 else None
+176
View File
@@ -0,0 +1,176 @@
"""Priority scheduler + lifecycle.
Owns the set of running jobs and admits queued downloads up to a global
concurrency limit (K), highest priority first, FIFO within a priority. Runs
entirely on the existing ComfyUI asyncio loop; blocking work (disk, hashing,
DB) is offloaded by the job/writer layers.
On startup it reconciles DB vs. disk: ``active``/``verifying`` rows left by a
previous run are reset to ``queued`` and resumed from persisted offsets, and
orphaned ``.part`` files with no live download row are swept.
"""
from __future__ import annotations
import asyncio
import logging
import os
import random
import time
from typing import Callable, Optional
from comfy.cli_args import args
from app.model_downloader.constants import DownloadStatus
from app.model_downloader.database import queries
from app.model_downloader.engine.job import DownloadJob, JobSpec
from app.model_downloader.security import paths
# Backoff for retryable failures
_BACKOFF_BASE = 2.0
_BACKOFF_CAP = 300.0
_MAX_ATTEMPTS = 6
class Scheduler:
def __init__(self) -> None:
self._jobs: dict[str, DownloadJob] = {}
self._tasks: dict[str, asyncio.Task] = {}
self._backoff_until: dict[str, float] = {}
self._pump_lock = asyncio.Lock()
self._notify_cb: Optional[Callable[[str], None]] = None
self._started = False
@property
def max_active(self) -> int:
return max(1, getattr(args, "download_max_active", 3))
def set_notify(self, cb: Optional[Callable[[str], None]]) -> None:
self._notify_cb = cb
def get_job(self, download_id: str) -> Optional[DownloadJob]:
return self._jobs.get(download_id)
def is_active(self, download_id: str) -> bool:
return download_id in self._tasks
# ----- startup -----
async def start(self) -> None:
if self._started:
return
self._started = True
try:
await asyncio.to_thread(queries.reconcile_live_downloads)
await asyncio.to_thread(self._sweep_orphan_temp_files)
except Exception as e:
logging.warning("[model_downloader] startup reconcile failed: %s", e)
await self.pump()
@staticmethod
def _sweep_orphan_temp_files() -> None:
"""Remove ``.part`` files not referenced by a resumable download row.
Resumable partials are preserved; only truly orphaned temp files from
crashed runs are deleted. ``FAILED`` is included because
:meth:`DownloadManager.resume` explicitly permits resuming a
retry-exhausted failed row: deleting its partial here while the
per-segment offsets survive in the DB would make the next resume
preallocate a fresh sparse file, skip every "complete" segment, and
leave zero-filled holes that pass the size-only verification gate.
"""
live = {
row.temp_path
for row in queries.list_downloads()
if row.status
in (
DownloadStatus.QUEUED,
DownloadStatus.PAUSED,
DownloadStatus.FAILED,
)
}
for path in paths.iter_all_tmp_paths():
if path in live:
continue
try:
os.remove(path)
logging.info("[model_downloader] removed orphan temp file: %s", path)
except OSError as e:
logging.warning("[model_downloader] could not remove %s: %s", path, e)
# ----- admission -----
async def pump(self) -> None:
async with self._pump_lock:
slots = self.max_active - len(self._tasks)
if slots <= 0:
return
now = time.monotonic()
candidates = await asyncio.to_thread(queries.list_queued_downloads)
for row in candidates:
if slots <= 0:
break
if row.id in self._tasks:
continue
if self._backoff_until.get(row.id, 0.0) > now:
continue
self._admit(row)
slots -= 1
def _admit(self, row) -> None:
spec = JobSpec(
download_id=row.id,
url=row.url,
model_id=row.model_id,
dest_path=row.dest_path,
temp_path=row.temp_path,
priority=row.priority,
expected_sha256=row.expected_sha256,
allow_any_extension=row.allow_any_extension,
etag=row.etag,
attempts=row.attempts,
)
job = DownloadJob(spec, notify_cb=self._notify_cb)
self._jobs[row.id] = job
self._tasks[row.id] = asyncio.ensure_future(self._run_job(job))
async def _run_job(self, job: DownloadJob) -> None:
download_id = job.spec.download_id
status = DownloadStatus.FAILED
try:
status = await job.run()
except Exception as e: # run() is defensive, but never let a task die silently
logging.error("[model_downloader] job %s crashed: %s", download_id, e)
queries.update_download(
download_id,
status=DownloadStatus.FAILED,
error=f"internal error: {e}",
)
if self._notify_cb:
self._notify_cb(download_id)
finally:
self._tasks.pop(download_id, None)
self._jobs.pop(download_id, None)
if status == DownloadStatus.QUEUED:
if job.spec.attempts >= _MAX_ATTEMPTS:
queries.update_download(
download_id,
status=DownloadStatus.FAILED,
error=f"giving up after {job.spec.attempts} attempts",
)
if self._notify_cb:
self._notify_cb(download_id)
else:
delay = min(
_BACKOFF_CAP, _BACKOFF_BASE ** job.spec.attempts
) + random.uniform(0, 1.0)
self._backoff_until[download_id] = time.monotonic() + delay
asyncio.ensure_future(self._delayed_pump(delay))
await self.pump()
async def _delayed_pump(self, delay: float) -> None:
await asyncio.sleep(delay)
await self.pump()
SCHEDULER = Scheduler()
+140
View File
@@ -0,0 +1,140 @@
"""URL allowlist for server-side model fetches.
Default-deny. A URL is downloadable only when its parsed host + scheme are
allowlisted AND (unless explicitly relaxed) its final filename ends in a
known model extension.
The built-in host defaults mirror the frontend's ``isModelDownloadable``
allowlist so the two flows agree on what is eligible; ``--download-allowed-hosts``
extends it for self-hosted mirrors. Matching is done on ``urlparse().hostname``
(never a raw string prefix) so userinfo tricks like
``http://127.0.0.1@169.254.169.254/x.safetensors`` — whose real host is the
metadata IP — cannot slip past.
"""
from __future__ import annotations
from urllib.parse import urlparse
from comfy.cli_args import args
# host -> set of allowed schemes. Frontend parity (HuggingFace / Civitai /
# localhost). Extra hosts from --download-allowed-hosts are https-only.
_DEFAULT_ALLOWED_HOSTS: dict[str, set[str]] = {
"huggingface.co": {"https"},
"civitai.com": {"https"},
"localhost": {"http", "https"},
"127.0.0.1": {"http", "https"},
}
# Hosts for which loopback addresses are intentionally permitted (the localhost
# "download a local model" feature). Every other host's loopback resolution is
# rejected by the SSRF resolver.
LOOPBACK_HOSTS = frozenset({"localhost", "127.0.0.1", "::1"})
# Known model file extensions (frontend parity). Checked on the final filename.
ALLOWED_MODEL_EXTENSIONS = (
".safetensors",
".sft",
".ckpt",
".pth",
".pt",
".gguf",
".bin",
)
def _allowed_hosts() -> dict[str, set[str]]:
hosts = {h: set(s) for h, s in _DEFAULT_ALLOWED_HOSTS.items()}
for extra in getattr(args, "download_allowed_hosts", []) or []:
host = extra.strip().lower()
if host:
hosts.setdefault(host, set()).add("https")
return hosts
def is_host_allowed(host: str | None, scheme: str | None) -> bool:
"""True iff ``host`` is allowlisted for ``scheme``.
Used both for the initial URL and re-checked on every redirect hop,
so a whitelisted URL cannot 30x into an off-list host.
"""
if not host or not scheme:
return False
allowed = _allowed_hosts().get(host.lower())
return allowed is not None and scheme.lower() in allowed
def has_allowed_extension(path: str, allow_any_extension: bool = False) -> bool:
if allow_any_extension:
return True
return path.lower().endswith(ALLOWED_MODEL_EXTENSIONS)
def filename_extension(name: str) -> str:
"""Lowercased extension (including the leading dot) of a bare filename.
Returns ``""`` when there is no extension. A leading-dot name
(``.safetensors``) is treated as having no extension (all stem), matching
``os.path.splitext`` semantics so dotfiles aren't mistaken for typed files.
"""
base = name.replace("\\", "/").rsplit("/", 1)[-1]
dot = base.rfind(".")
if dot <= 0:
return ""
return base[dot:].lower()
def is_allowed_extension_name(name: str) -> bool:
"""True iff ``name`` ends in one of the known model extensions."""
return name.lower().endswith(ALLOWED_MODEL_EXTENSIONS)
def is_host_allowed_url(url: str) -> bool:
"""True iff ``url`` parses and its host+scheme are allowlisted."""
if not isinstance(url, str) or not url:
return False
try:
parsed = urlparse(url)
except ValueError:
return False
return is_host_allowed(parsed.hostname, parsed.scheme)
def url_path_extension(url: str) -> str:
"""Extension of the URL *path* basename (query ignored), or ``""``."""
try:
parsed = urlparse(url)
except ValueError:
return ""
return filename_extension(parsed.path)
def is_url_downloadable(url: str) -> bool:
"""Coarse enqueue gate: host/scheme allowed and extension not disallowed.
Unlike :func:`is_url_allowed` (which demands a known extension *in the URL*),
this also admits URLs whose path carries no extension at all — e.g. a Civitai
``/api/download/models/<id>`` endpoint whose real filename only shows up in
the redirect target / ``Content-Disposition``. The true extension is then
resolved from the network and re-validated before the download is admitted.
A path bearing an explicit *non-model* extension (``.zip``, ``.html``, ...)
is still rejected here.
"""
if not is_host_allowed_url(url):
return False
ext = url_path_extension(url)
return ext == "" or ext in ALLOWED_MODEL_EXTENSIONS
def is_url_allowed(url: str, allow_any_extension: bool = False) -> bool:
"""Check whether ``url`` is permitted as a server-side download source."""
if not isinstance(url, str) or not url:
return False
try:
parsed = urlparse(url)
except ValueError:
return False
if not is_host_allowed(parsed.hostname, parsed.scheme):
return False
return has_allowed_extension(parsed.path, allow_any_extension)
+132
View File
@@ -0,0 +1,132 @@
"""Path resolution + traversal safety for downloads.
A ``model_id`` is a *relative destination path* of the form
``<directory>/<filename>`` (e.g. ``loras/my_lora.safetensors``). This module
turns one into an absolute on-disk path under one of ComfyUI's registered
model folders, rejecting unknown folders, path traversal, and symlink escape.
This is the only thing that composes destination paths, so the engine never
touches user-supplied path strings directly.
"""
from __future__ import annotations
import os
import re
from typing import Iterator, Optional
import folder_paths
from app.model_downloader.constants import TMP_SUFFIX
from app.model_downloader.security.allowlist import ALLOWED_MODEL_EXTENSIONS
# A model_id component is a single path segment of safe characters — no slashes,
# no "..", no leading dots that could escape the target directory.
_SEGMENT_RE = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._-]*$")
class InvalidModelId(ValueError):
"""Raised when a model_id is malformed or names an unknown model folder."""
def parse_model_id(model_id: str, allow_any_extension: bool = False) -> tuple[str, str]:
"""Split ``<directory>/<filename>`` and validate both components.
Returns ``(directory, filename)``. Does not touch the filesystem.
"""
if not isinstance(model_id, str) or "/" not in model_id:
raise InvalidModelId(
f"model_id must be '<directory>/<filename>', got {model_id!r}"
)
directory, _, filename = model_id.partition("/")
if "/" in filename or not directory or not filename:
raise InvalidModelId(
f"model_id must have exactly one '/' separator, got {model_id!r}"
)
if not _SEGMENT_RE.match(directory):
raise InvalidModelId(f"invalid directory segment {directory!r}")
if not _SEGMENT_RE.match(filename):
raise InvalidModelId(f"invalid filename segment {filename!r}")
if not allow_any_extension and not filename.lower().endswith(
ALLOWED_MODEL_EXTENSIONS
):
raise InvalidModelId(
f"filename must end with a known model extension "
f"{ALLOWED_MODEL_EXTENSIONS}, got {filename!r}"
)
if directory not in folder_paths.folder_names_and_paths:
raise InvalidModelId(f"unknown model folder {directory!r}")
return directory, filename
def apply_extension(model_id: str, ext: str) -> str:
"""Return ``model_id`` with its filename forced to end in ``ext``.
``ext`` includes the leading dot (e.g. ``".safetensors"``). If the filename
already ends in a *known model extension* it is replaced; otherwise ``ext``
is appended (so ``loras/mymodel`` -> ``loras/mymodel.safetensors`` and
``loras/mymodel.ckpt`` -> ``loras/mymodel.safetensors``). A filename with a
non-model suffix (``my.model.v2``) is treated as an extensionless stem and
``ext`` is appended. The directory part is left untouched; validation is
still the caller's job via :func:`parse_model_id`.
"""
directory, sep, filename = model_id.partition("/")
if not sep:
return model_id # malformed; parse_model_id will reject it
low = filename.lower()
for known in ALLOWED_MODEL_EXTENSIONS:
if low.endswith(known):
filename = filename[: -len(known)]
break
return f"{directory}{sep}{filename}{ext}"
def resolve_existing(model_id: str, allow_any_extension: bool = False) -> Optional[str]:
"""Return the absolute path of an installed model, or None if missing.
Honours ``extra_model_paths.yaml`` transparently via ``get_full_path``.
"""
directory, filename = parse_model_id(model_id, allow_any_extension)
return folder_paths.get_full_path(directory, filename)
def resolve_destination(
model_id: str, allow_any_extension: bool = False
) -> tuple[str, str]:
"""Return ``(final_path, temp_path)`` for a download.
Downloads land at the first registered path for the model's directory
(the "primary" location). ``temp_path`` is a sibling ``.part`` file that
is atomically renamed onto ``final_path`` on success. The result is
asserted to stay within the registered root (defence in depth on top of
the segment regex).
"""
directory, filename = parse_model_id(model_id, allow_any_extension)
roots = folder_paths.get_folder_paths(directory)
if not roots:
raise InvalidModelId(f"no on-disk path registered for folder {directory!r}")
root = os.path.realpath(roots[0])
final_path = os.path.realpath(os.path.join(root, filename))
if final_path != root and not final_path.startswith(root + os.sep):
raise InvalidModelId(f"resolved path escapes model root: {model_id!r}")
temp_path = f"{final_path}{TMP_SUFFIX}"
return final_path, temp_path
def iter_all_tmp_paths() -> Iterator[str]:
"""Yield this subsystem's temp files under every registered model folder.
Matches only the distinctive ``TMP_SUFFIX`` so the startup orphan sweep
can never delete temp files created by other tools.
"""
seen_roots: set[str] = set()
for directory in list(folder_paths.folder_names_and_paths.keys()):
for root in folder_paths.get_folder_paths(directory):
if root in seen_roots or not os.path.isdir(root):
continue
seen_roots.add(root)
try:
for entry in os.scandir(root):
if entry.is_file() and entry.name.endswith(TMP_SUFFIX):
yield entry.path
except OSError:
continue
+163
View File
@@ -0,0 +1,163 @@
"""SSRF / exfiltration defenses.
Two cooperating layers:
1. :class:`ValidatingResolver` is installed on the shared connector. Every
connection — the initial probe and every segment GET, including ones made
after a redirect — resolves its host through this resolver, which rejects
any address that lands on a private / special-use IP range. Because the
resolve and the connect happen together inside the connector, there is no
check-then-connect window for DNS rebinding to exploit.
2. :func:`check_redirect_hop` re-validates every hop. The host allowlist gates
only the *initial* user-supplied URL (anti-SSRF for arbitrary input);
legitimate downloads from allowlisted origins redirect to presigned CDN
hosts that are deliberately NOT on the allowlist (HF ->
``cdn-lfs*.huggingface.co``, Civitai -> signed Cloudflare/S3), so hops are
instead screened for scheme, embedded credentials, and — via the resolver
above — private IPs. Credentials are only ever attached when a hop's host
exactly matches a stored credential, so they are dropped on the CDN hop.
Loopback (the "download a local model" feature) is exempt from IP filtering
only for the initial URL: a *redirect* may never target a loopback host or
a blocked IP-literal, which the resolver alone can't enforce (it exempts
loopback literals and never sees IP literals through DNS).
"""
from __future__ import annotations
import ipaddress
import socket
from urllib.parse import urlparse
from aiohttp.abc import AbstractResolver
from aiohttp.resolver import DefaultResolver
from app.model_downloader.security.allowlist import LOOPBACK_HOSTS
# Cap the redirect chain length a hop may use.
MAX_REDIRECTS = 5
class SSRFError(Exception):
"""A hop failed an SSRF / allowlist check."""
def is_scheme_allowed(scheme: str | None, host: str | None) -> bool:
"""True iff ``scheme`` is permitted for ``host`` on a download hop.
https is always allowed; plain http only for loopback/approved dev hosts.
"""
if not scheme:
return False
scheme = scheme.lower()
if scheme == "https":
return True
if scheme == "http":
return bool(host) and host.lower() in LOOPBACK_HOSTS
return False
def is_blocked_ip(ip_str: str) -> bool:
"""True for any address we refuse to connect to.
Covers loopback, link-local (incl. 169.254.169.254 cloud metadata),
RFC1918 private ranges, unique-local (ULA), unspecified (0.0.0.0/::),
multicast and other reserved ranges.
"""
try:
ip = ipaddress.ip_address(ip_str)
except ValueError:
return True # unparseable -> refuse
# On CPython before the gh-113171 fix (backported to 3.12.4/3.11.9/
# 3.10.14/3.9.19) the is_* properties don't see through IPv4-mapped IPv6
# (e.g. ::ffff:169.254.169.254), so resolve and re-check the embedded IPv4
# to keep mapped metadata/private addresses from slipping past the filter.
mapped = getattr(ip, "ipv4_mapped", None)
if mapped is not None:
ip = mapped
return (
ip.is_private
or ip.is_loopback
or ip.is_link_local
or ip.is_multicast
or ip.is_reserved
or ip.is_unspecified
)
class ValidatingResolver(AbstractResolver):
"""Delegating resolver that drops blocked IPs from every resolution.
If a hostname resolves only to blocked addresses, the connection fails
closed with an :class:`OSError`, which aiohttp surfaces as a connection
error to the caller.
"""
def __init__(self) -> None:
self._inner = DefaultResolver()
async def resolve(self, host, port=0, family=socket.AF_INET):
infos = await self._inner.resolve(host, port, family)
# localhost/127.0.0.1 are an explicit, opt-in allowlist feature.
if isinstance(host, str) and host.lower() in LOOPBACK_HOSTS:
return infos
safe = [info for info in infos if not is_blocked_ip(info["host"])]
if not safe:
raise OSError(
f"refusing to connect to {host!r}: resolves only to "
f"private/special-use addresses"
)
return safe
async def close(self) -> None:
await self._inner.close()
def check_redirect_hop(url: str, *, is_initial_url: bool = False) -> str:
"""Validate one hop's URL.
Returns the URL unchanged on success; raises :class:`SSRFError` otherwise.
Requires https for external hosts (http only for loopback/approved dev
hosts) and forbids credentials-in-URL. The host is NOT re-checked against
the allowlist (CDN redirect targets are off-list by design); credential
leakage is prevented by exact host matching at attach time, and the landing
filename's extension is gated separately by the caller.
Loopback/blocked-IP screening: the connector's resolver filters resolvable
hostnames but exempts literal loopback hosts (``localhost``/``127.0.0.1``/
``::1``) and never sees IP literals through DNS. That loopback exemption is
legitimate only for the *initial* user-supplied URL (``is_initial_url``);
on a redirect hop we reject loopback hosts and any blocked IP-literal here,
so a 30x can't steer a server-side GET at loopback/internal services.
"""
try:
parsed = urlparse(url)
except ValueError as e:
raise SSRFError(f"unparseable redirect URL {url!r}: {e}") from e
host = parsed.hostname
if not host:
raise SSRFError(f"redirect URL has no host: {url!r}")
if not is_scheme_allowed(parsed.scheme, host):
raise SSRFError(
f"redirect to disallowed scheme {parsed.scheme!r} for host "
f"{host!r} (https required for external hosts)"
)
if parsed.username or parsed.password:
raise SSRFError("credentials-in-URL are not allowed")
host_is_loopback = host.lower() in LOOPBACK_HOSTS
if not is_initial_url and host_is_loopback:
raise SSRFError(f"redirect to loopback host {host!r} is not allowed")
# IP-literal targets never go through DNS, so the connector's resolver can't
# screen them — check them directly. The only blocked IP allowed through is
# a loopback literal on the initial URL (handled by the exemption above).
try:
ipaddress.ip_address(host)
except ValueError:
is_ip_literal = False
else:
is_ip_literal = True
if is_ip_literal and is_blocked_ip(host) and not (
is_initial_url and host_is_loopback
):
raise SSRFError(f"redirect to blocked internal address {host!r}")
return url
+49
View File
@@ -0,0 +1,49 @@
"""Hub-checksum verification = SHA256.
Only used to confirm a download matches a *provided* ``expected_sha256``. It
is NOT the dedup key (that is blake3, owned by the assets system). The full
sequential read happens at most once, here, only when a checksum was supplied.
"""
from __future__ import annotations
import hashlib
from typing import Callable, Optional
_CHUNK = 8 * 1024 * 1024
InterruptCheck = Callable[[], bool]
class ChecksumError(Exception):
"""The computed SHA256 did not match the expected value."""
def sha256_file(path: str, interrupt_check: Optional[InterruptCheck] = None) -> Optional[str]:
"""Stream the file and return its lowercase hex SHA256.
Returns ``None`` if interrupted via ``interrupt_check``.
"""
h = hashlib.sha256()
with open(path, "rb") as f:
while True:
if interrupt_check is not None and interrupt_check():
return None
chunk = f.read(_CHUNK)
if not chunk:
break
h.update(chunk)
return h.hexdigest()
def verify_sha256(
path: str, expected: str, interrupt_check: Optional[InterruptCheck] = None
) -> None:
"""Raise :class:`ChecksumError` unless the file's SHA256 matches ``expected``."""
actual = sha256_file(path, interrupt_check)
if actual is None:
return # interrupted; caller will re-verify on resume
if actual.lower() != expected.lower():
raise ChecksumError(
f"sha256 mismatch: expected {expected.lower()}, got {actual.lower()}"
)
+53
View File
@@ -0,0 +1,53 @@
"""Dedup + catalog handoff — reuse the assets system.
We do NOT build a parallel indexer. "Do I already have it?" is answered by
``resolve_existing`` (path) at enqueue time and, where a hash is known, by the
assets blake3 catalog. After a completed download we register the file
through the assets ingest path so it is cataloged and (eventually) hashed by
the existing enrichment worker.
"""
from __future__ import annotations
import asyncio
import logging
import os
from typing import Optional
def _register_sync(abs_path: str) -> Optional[str]:
"""Register a finished file into the assets catalog. Returns asset hash."""
try:
from app.assets.services.ingest import register_file_in_place
except Exception as e: # assets package import failure — non-fatal
logging.debug("[model_downloader] assets ingest unavailable: %s", e)
return None
try:
result = register_file_in_place(abs_path, name=os.path.basename(abs_path), tags=[])
return result.asset.hash if result and result.asset else None
except Exception as e:
# The file is already safely on disk; cataloging is best-effort.
logging.warning(
"[model_downloader] could not register %s into assets catalog: %s",
abs_path, e,
)
return None
async def register_completed(abs_path: str) -> Optional[str]:
"""Catalog a completed download via the assets system (off the event loop)."""
return await asyncio.to_thread(_register_sync, abs_path)
def _find_by_hash_sync(blake3_hex: str) -> Optional[str]:
try:
from app.assets.services.asset_management import get_asset_by_hash
except Exception:
return None
asset = get_asset_by_hash("blake3:" + blake3_hex)
return asset.hash if asset is not None else None
async def find_existing_by_hash(blake3_hex: str) -> Optional[str]:
"""Pure DB lookup — never triggers hashing on the hot path."""
return await asyncio.to_thread(_find_by_hash_sync, blake3_hex)
+86
View File
@@ -0,0 +1,86 @@
"""Cheap structural validation, no full read.
For ``.safetensors``/``.sft`` we parse the header (first few KB): it carries
the tensor table and the byte length of the data region. We assert
``file_size == 8 + header_len + data_region_len``. This detects truncation
and most corruption for free, before any crypto hashing. Other extensions
have no cheap structural check and pass through.
"""
from __future__ import annotations
import json
import os
import struct
from typing import Optional
_SAFETENSORS_EXTS = (".safetensors", ".sft")
# A sane upper bound so a corrupt header length can't make us read gigabytes.
_MAX_HEADER_BYTES = 100 * 1024 * 1024
class StructuralError(Exception):
"""The file failed its structural integrity check."""
def validate(path: str, name_hint: Optional[str] = None) -> None:
"""Validate the file at ``path``. Raises :class:`StructuralError` on failure.
The file format is detected from ``name_hint`` when provided, otherwise from
``path``. Callers that download into a temp file with an opaque suffix (e.g.
``*.comfy-download.part``) must pass the final destination name as
``name_hint`` so the format check is not silently skipped.
"""
lower = (name_hint or path).lower()
if lower.endswith(_SAFETENSORS_EXTS):
_validate_safetensors(path)
# No structural check for other formats; the size + (optional) checksum
# gates in the engine cover those.
def _validate_safetensors(path: str) -> None:
file_size = os.path.getsize(path)
if file_size < 8:
raise StructuralError(f"file too small to be safetensors ({file_size} bytes)")
with open(path, "rb") as f:
header_len = struct.unpack("<Q", f.read(8))[0]
if header_len <= 0 or header_len > _MAX_HEADER_BYTES:
raise StructuralError(f"implausible safetensors header length {header_len}")
if 8 + header_len > file_size:
raise StructuralError("safetensors header extends past end of file")
try:
header = json.loads(f.read(header_len).decode("utf-8"))
except (UnicodeDecodeError, json.JSONDecodeError) as e:
raise StructuralError(f"safetensors header is not valid JSON: {e}") from e
if not isinstance(header, dict):
raise StructuralError("safetensors header is not a JSON object")
data_len = 0
for name, entry in header.items():
if name == "__metadata__":
continue
if not isinstance(entry, dict) or "data_offsets" not in entry:
raise StructuralError(f"tensor {name!r} missing data_offsets")
offsets = entry["data_offsets"]
if not (isinstance(offsets, list) and len(offsets) == 2):
raise StructuralError(f"tensor {name!r} has malformed data_offsets")
begin, end = offsets
# bool is an int subclass; reject it explicitly to avoid True/False offsets.
if (
not isinstance(begin, int)
or not isinstance(end, int)
or isinstance(begin, bool)
or isinstance(end, bool)
or begin < 0
or end < begin
):
raise StructuralError(f"tensor {name!r} has malformed data_offsets")
data_len = max(data_len, end)
expected = 8 + header_len + data_len
if file_size != expected:
raise StructuralError(
f"size mismatch: file is {file_size} bytes, header implies {expected} "
f"(8 + {header_len} header + {data_len} data)"
)
+31
View File
@@ -33,6 +33,28 @@ class EnumAction(argparse.Action):
setattr(namespace, self.dest, value)
def _positive_int(value: str) -> int:
"""argparse type that rejects zero and negative integers."""
try:
ivalue = int(value)
except ValueError:
raise argparse.ArgumentTypeError(f"{value!r} is not an integer")
if ivalue <= 0:
raise argparse.ArgumentTypeError(f"{value!r} must be a positive integer (> 0)")
return ivalue
def _non_negative_int(value: str) -> int:
"""argparse type that rejects negatives but allows zero (a disable sentinel)."""
try:
ivalue = int(value)
except ValueError:
raise argparse.ArgumentTypeError(f"{value!r} is not an integer")
if ivalue < 0:
raise argparse.ArgumentTypeError(f"{value!r} must be a non-negative integer (>= 0)")
return ivalue
parser = argparse.ArgumentParser()
parser.add_argument("--listen", type=str, default="127.0.0.1", metavar="IP", nargs="?", const="0.0.0.0,::", help="Specify the IP address to listen on (default: 127.0.0.1). You can give a list of ip addresses by separating them with a comma like: 127.2.2.2,127.3.3.3 If --listen is provided without an argument, it defaults to 0.0.0.0,:: (listens on all ipv4 and ipv6)")
@@ -246,6 +268,15 @@ parser.add_argument("--enable-asset-hashing", action="store_true", help="Compute
parser.add_argument("--feature-flag", type=str, action='append', default=[], metavar="KEY[=VALUE]", help="Set a server feature flag. Use KEY=VALUE to set an explicit value, or bare KEY to set it to true. Can be specified multiple times. Boolean values (true/false) and numbers are auto-converted. Examples: --feature-flag show_signin_button=true or --feature-flag show_signin_button")
parser.add_argument("--list-feature-flags", action="store_true", help="Print the registry of known CLI-settable feature flags as JSON and exit.")
# ----- Model download manager (PRD: docs/prd-download-manager.md) -----
parser.add_argument("--download-segments", type=_positive_int, default=8, metavar="N", help="Number of parallel HTTP range segments per file for the model download manager (default: 8).")
parser.add_argument("--download-max-active", type=_positive_int, default=3, metavar="N", help="Maximum number of model downloads running concurrently (default: 3).")
parser.add_argument("--download-max-connections-per-host", type=_positive_int, default=16, metavar="N", help="Maximum simultaneous connections to a single host for the download manager (default: 16).")
parser.add_argument("--download-chunk-size", type=_positive_int, default=4 * 1024 * 1024, metavar="BYTES", help="Read chunk size in bytes for the download manager (default: 4 MiB).")
parser.add_argument("--download-max-bytes", type=_non_negative_int, default=1024 * 1024 * 1024 * 1024, metavar="BYTES", help="Maximum size in bytes of a single download; aborts transfers that exceed it (guards against malicious/non-conforming hosts filling the disk). Set to 0 to disable (default: 1 TiB).")
parser.add_argument("--download-allowed-hosts", type=str, nargs="*", default=[], metavar="HOST", help="Additional hostnames to add to the download manager allowlist (https only). The built-in defaults always include huggingface.co and civitai.com.")
parser.add_argument("--download-allow-any-extension", action="store_true", help="Allow the download manager to fetch files with any extension (default: only known model extensions like .safetensors).")
if comfy.options.args_parsing:
args = parser.parse_args()
else:
-278
View File
@@ -1,278 +0,0 @@
import re
import torch
from torch import nn
import torch.nn.functional as F
import comfy.ops
import comfy.utils
MODULE_PATTERN = re.compile(r"lllite_dit_blocks_(\d+)_(self_attn_[qkv]_proj|cross_attn_q_proj|mlp_layer1)$")
def _group_norm(channels, device=None, dtype=None, operations=None):
groups = 8
while groups > 1 and channels % groups != 0:
groups //= 2
return operations.GroupNorm(groups, channels, device=device, dtype=dtype)
class AnimaLLLiteResBlock(nn.Module):
def __init__(self, channels, device=None, dtype=None, operations=None):
super().__init__()
self.norm1 = _group_norm(channels, device=device, dtype=dtype, operations=operations)
self.conv1 = operations.Conv2d(channels, channels, kernel_size=3, padding=1, device=device, dtype=dtype)
self.norm2 = _group_norm(channels, device=device, dtype=dtype, operations=operations)
self.conv2 = operations.Conv2d(channels, channels, kernel_size=3, padding=1, device=device, dtype=dtype)
def forward(self, x):
h = self.conv1(F.silu(self.norm1(x)))
h = self.conv2(F.silu(self.norm2(h)))
return x + h
class AnimaLLLiteASPP(nn.Module):
def __init__(self, channels, dilations, device=None, dtype=None, operations=None):
super().__init__()
branches = []
for dilation in dilations:
if dilation == 1:
conv = operations.Conv2d(channels, channels, kernel_size=1, device=device, dtype=dtype)
else:
conv = operations.Conv2d(channels, channels, kernel_size=3, padding=dilation, dilation=dilation, device=device, dtype=dtype)
branches.append(nn.Sequential(conv, _group_norm(channels, device=device, dtype=dtype, operations=operations), nn.SiLU()))
self.branches = nn.ModuleList(branches)
self.global_pool = nn.AdaptiveAvgPool2d(1)
self.global_conv = nn.Sequential(
operations.Conv2d(channels, channels, kernel_size=1, device=device, dtype=dtype),
_group_norm(channels, device=device, dtype=dtype, operations=operations),
nn.SiLU(),
)
self.proj = nn.Sequential(
operations.Conv2d(channels * (len(dilations) + 1), channels, kernel_size=1, device=device, dtype=dtype),
_group_norm(channels, device=device, dtype=dtype, operations=operations),
nn.SiLU(),
)
def forward(self, x):
height, width = x.shape[-2:]
outputs = [branch(x) for branch in self.branches]
pooled = self.global_conv(self.global_pool(x))
outputs.append(F.interpolate(pooled, size=(height, width), mode="bilinear", align_corners=False))
return self.proj(torch.cat(outputs, dim=1))
class AnimaLLLiteConditioning(nn.Module):
def __init__(self, cond_in_channels, cond_dim, cond_emb_dim, cond_resblocks, aspp_dilations, device=None, dtype=None, operations=None):
super().__init__()
half_dim = cond_dim // 2
self.conv1 = operations.Conv2d(cond_in_channels, half_dim, kernel_size=4, stride=4, device=device, dtype=dtype)
self.norm1 = _group_norm(half_dim, device=device, dtype=dtype, operations=operations)
self.conv2 = operations.Conv2d(half_dim, half_dim, kernel_size=3, padding=1, device=device, dtype=dtype)
self.norm2 = _group_norm(half_dim, device=device, dtype=dtype, operations=operations)
self.conv3 = operations.Conv2d(half_dim, cond_dim, kernel_size=4, stride=4, device=device, dtype=dtype)
self.norm3 = _group_norm(cond_dim, device=device, dtype=dtype, operations=operations)
self.resblocks = nn.ModuleList([
AnimaLLLiteResBlock(cond_dim, device=device, dtype=dtype, operations=operations)
for _ in range(cond_resblocks)
])
self.aspp = AnimaLLLiteASPP(cond_dim, aspp_dilations, device=device, dtype=dtype, operations=operations) if aspp_dilations else None
self.proj = operations.Conv2d(cond_dim, cond_emb_dim, kernel_size=1, device=device, dtype=dtype)
self.out_norm = operations.LayerNorm(cond_emb_dim, device=device, dtype=dtype)
def forward(self, x):
x = F.silu(self.norm1(self.conv1(x)))
x = F.silu(self.norm2(self.conv2(x)))
x = F.silu(self.norm3(self.conv3(x)))
for block in self.resblocks:
x = block(x)
if self.aspp is not None:
x = self.aspp(x)
x = self.proj(x).flatten(2).transpose(1, 2).contiguous()
return self.out_norm(x)
class AnimaLLLiteModule(nn.Module):
def __init__(self, in_dim, cond_emb_dim, mlp_dim, device=None, dtype=None, operations=None):
super().__init__()
self.down = operations.Linear(in_dim, mlp_dim, device=device, dtype=dtype)
self.mid = operations.Linear(mlp_dim + cond_emb_dim, mlp_dim, device=device, dtype=dtype)
self.cond_to_film = operations.Linear(cond_emb_dim, 2 * mlp_dim, device=device, dtype=dtype)
self.up = operations.Linear(mlp_dim, in_dim, device=device, dtype=dtype)
self.depth_embed = nn.Parameter(torch.empty(cond_emb_dim, device=device, dtype=dtype), requires_grad=False)
def forward(self, x, cond_emb, strength):
original_shape = x.shape
if x.ndim == 5:
x = x.flatten(1, 3)
if x.shape[0] != cond_emb.shape[0]:
if x.shape[0] % cond_emb.shape[0] != 0:
raise ValueError(f"Anima LLLite batch mismatch: model input batch {x.shape[0]}, control batch {cond_emb.shape[0]}")
cond_emb = cond_emb.repeat(x.shape[0] // cond_emb.shape[0], 1, 1)
if x.shape[1] != cond_emb.shape[1]:
raise ValueError(f"Anima LLLite sequence mismatch: model input has {x.shape[1]} tokens, control has {cond_emb.shape[1]}")
cond_local = cond_emb + comfy.ops.cast_to_input(self.depth_embed, cond_emb)
hidden = F.silu(self.down(x))
gamma, beta = self.cond_to_film(cond_local).chunk(2, dim=-1)
hidden = self.mid(torch.cat((cond_local, hidden), dim=-1))
hidden = F.silu(hidden * (1 + gamma) + beta)
x = x + self.up(hidden) * strength
if len(original_shape) == 5:
x = x.reshape(original_shape)
return x
class AnimaLLLite(nn.Module):
def __init__(self, state_dict, metadata, device=None, dtype=None, operations=None):
super().__init__()
metadata = metadata or {}
version = metadata.get("lllite.version", "2")
if version != "2":
raise ValueError(f"Unsupported Anima LLLite version {version!r}; only named-key v2 checkpoints are supported")
module_names = sorted({key.split(".", 1)[0] for key in state_dict if key.startswith("lllite_dit_blocks_")})
if not module_names:
raise ValueError("Anima LLLite checkpoint has no lllite_dit_blocks_* modules")
cond_in_channels = state_dict["lllite_conditioning1.conv1.weight"].shape[1]
cond_dim = state_dict["lllite_conditioning1.conv3.weight"].shape[0]
cond_emb_dim = state_dict["lllite_conditioning1.proj.weight"].shape[0]
resblock_ids = {int(key.split(".")[2]) for key in state_dict if key.startswith("lllite_conditioning1.resblocks.")}
cond_resblocks = max(resblock_ids) + 1 if resblock_ids else 0
use_aspp = any(key.startswith("lllite_conditioning1.aspp.") for key in state_dict)
dilation_string = metadata.get("lllite.aspp_dilations", "1,2,4,8")
aspp_dilations = tuple(int(value) for value in dilation_string.split(",") if value.strip()) if use_aspp else ()
self.cond_in_channels = cond_in_channels
self.inpaint_masked_input = metadata.get("lllite.inpaint_masked_input", "false").lower() == "true"
self.lllite_conditioning1 = AnimaLLLiteConditioning(
cond_in_channels, cond_dim, cond_emb_dim, cond_resblocks, aspp_dilations,
device=device, dtype=dtype, operations=operations,
)
self.module_names = set()
self.block_count = 0
self.model_dim = None
for name in module_names:
match = MODULE_PATTERN.fullmatch(name)
if match is None:
raise ValueError(f"Unsupported Anima LLLite module name: {name}")
down_shape = state_dict[f"{name}.down.weight"].shape
mlp_dim, in_dim = down_shape
module_cond_dim = state_dict[f"{name}.cond_to_film.weight"].shape[1]
if module_cond_dim != cond_emb_dim:
raise ValueError(f"Anima LLLite conditioning dimension mismatch in {name}: {module_cond_dim} != {cond_emb_dim}")
if self.model_dim is None:
self.model_dim = in_dim
elif self.model_dim != in_dim:
raise ValueError(f"Anima LLLite model dimension mismatch in {name}: {in_dim} != {self.model_dim}")
self.add_module(name, AnimaLLLiteModule(in_dim, cond_emb_dim, mlp_dim, device=device, dtype=dtype, operations=operations))
self.module_names.add(name)
self.block_count = max(self.block_count, int(match.group(1)) + 1)
def encode_conditioning(self, image):
return self.lllite_conditioning1(image)
def apply(self, x, cond_emb, block_index, target, strength):
name = f"lllite_dit_blocks_{block_index}_{target}"
if name not in self.module_names:
return x
return self.get_submodule(name)(x, cond_emb, strength)
class AnimaLLLitePatch:
def __init__(self, model_patch, image, mask, strength, sigma_start, sigma_end):
self.model_patch = model_patch
self.image = image
self.mask = mask
self.strength = strength
self.sigma_start = sigma_start
self.sigma_end = sigma_end
def __call__(self, args):
x = args["x"]
transformer_options = args["transformer_options"]
if self.strength == 0.0:
return args
sigmas = transformer_options.get("sigmas")
if sigmas is not None:
sigma = float(sigmas.max().item())
if not self.sigma_end <= sigma <= self.sigma_start:
return args
if x.shape[2] != 1:
raise ValueError(f"Anima LLLite only supports T=1, got T={x.shape[2]}")
target_height = x.shape[-2] * 8
target_width = x.shape[-1] * 8
image = comfy.utils.common_upscale(
self.image.movedim(-1, 1), target_width, target_height, "bicubic", crop="center"
).clamp(0.0, 1.0)
image = image.to(device=x.device, dtype=x.dtype) * 2.0 - 1.0
if self.model_patch.model.cond_in_channels == 4:
mask = self.mask
if mask.ndim == 3:
mask = mask.unsqueeze(1)
if mask.ndim != 4 or mask.shape[1] != 1:
raise ValueError(f"Anima LLLite mask must have one channel, got shape {tuple(mask.shape)}")
mask = comfy.utils.common_upscale(
mask.float(), target_width, target_height, "nearest-exact", crop="center"
)
if mask.shape[0] != image.shape[0]:
if image.shape[0] % mask.shape[0] != 0:
raise ValueError(
f"Anima LLLite mask batch {mask.shape[0]} cannot be broadcast to image batch {image.shape[0]}"
)
mask = mask.repeat(image.shape[0] // mask.shape[0], 1, 1, 1)
mask = (mask >= 0.5).to(device=x.device, dtype=x.dtype)
if self.model_patch.model.inpaint_masked_input:
image = image * (mask < 0.5).to(image.dtype)
image = torch.cat((image, mask * 2.0 - 1.0), dim=1)
cond_emb = self.model_patch.model.encode_conditioning(image)
transformer_options["model_patch_data"][self] = cond_emb
return args
def to(self, device_or_dtype):
return self
def models(self):
return [self.model_patch]
class AnimaLLLiteAttentionPatch:
def __init__(self, patch, targets):
self.patch = patch
self.targets = targets
def __call__(self, q, k, v, pe=None, attn_mask=None, extra_options=None):
cond_emb = extra_options["model_patch_data"].get(self.patch)
if cond_emb is None:
return {"q": q, "k": k, "v": v, "pe": pe, "attn_mask": attn_mask}
block_index = extra_options["block_index"]
values = {"q": q, "k": k, "v": v}
for value_name, target in self.targets.items():
values[value_name] = self.patch.model_patch.model.apply(
values[value_name], cond_emb, block_index, target, self.patch.strength
)
return {"q": values["q"], "k": values["k"], "v": values["v"], "pe": pe, "attn_mask": attn_mask}
class AnimaLLLiteMLPPatch:
def __init__(self, patch):
self.patch = patch
def __call__(self, args):
cond_emb = args["transformer_options"]["model_patch_data"].get(self.patch)
if cond_emb is None:
return args
args["x"] = self.patch.model_patch.model.apply(
args["x"], cond_emb, args["transformer_options"]["block_index"], "mlp_layer1", self.patch.strength
)
return args
+11 -53
View File
@@ -14,7 +14,6 @@ from torchvision import transforms
import comfy.patcher_extension
from comfy.ldm.modules.attention import optimized_attention
import comfy.ldm.common_dit
import comfy.ops
import comfy.quant_ops
@@ -149,29 +148,11 @@ class Attention(nn.Module):
x: torch.Tensor,
context: Optional[torch.Tensor] = None,
rope_emb: Optional[torch.Tensor] = None,
transformer_options: Optional[dict] = {},
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
q = self.q_proj(x)
context = x if context is None else context
q_input = x
k_input = context
v_input = context
transformer_patches = transformer_options.get("patches", {})
patch_name = "attn1_patch" if self.is_selfattn else "attn2_patch"
if patch_name in transformer_patches:
extra_options = transformer_options.copy()
extra_options["n_heads"] = self.n_heads
extra_options["dim_head"] = self.head_dim
for patch in transformer_patches[patch_name]:
out = patch(q_input, k_input, v_input, pe=rope_emb, attn_mask=None, extra_options=extra_options)
q_input = out.get("q", q_input)
k_input = out.get("k", k_input)
v_input = out.get("v", v_input)
rope_emb = out.get("pe", rope_emb)
q = self.q_proj(q_input)
k = self.k_proj(k_input)
v = self.v_proj(v_input)
k = self.k_proj(context)
v = self.v_proj(context)
q, k, v = map(
lambda t: rearrange(t, "b ... (h d) -> b ... h d", h=self.n_heads, d=self.head_dim),
(q, k, v),
@@ -180,16 +161,11 @@ class Attention(nn.Module):
def apply_norm_and_rotary_pos_emb(
q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, rope_emb: Optional[torch.Tensor]
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
q = self.q_norm(q)
k = self.k_norm(k)
v = self.v_norm(v)
if self.is_selfattn and rope_emb is not None: # only apply to self-attention!
q_scale, _, q_offload_stream = comfy.ops.cast_bias_weight(self.q_norm, q, offloadable=True)
k_scale, _, k_offload_stream = comfy.ops.cast_bias_weight(self.k_norm, k, offloadable=True)
q, k = comfy.quant_ops.ck.rms_rope_split_half(q, k, rope_emb, q_scale, k_scale, self.q_norm.eps)
comfy.ops.uncast_bias_weight(self.q_norm, q_scale, None, q_offload_stream)
comfy.ops.uncast_bias_weight(self.k_norm, k_scale, None, k_offload_stream)
else:
q = self.q_norm(q)
k = self.k_norm(k)
q, k = comfy.quant_ops.ck.apply_rope_split_half(q, k, rope_emb)
return q, k, v
q, k, v = apply_norm_and_rotary_pos_emb(q, k, v, rope_emb)
@@ -212,7 +188,7 @@ class Attention(nn.Module):
x (Tensor): The query tensor of shape [B, Mq, K]
context (Optional[Tensor]): The key tensor of shape [B, Mk, K] or use x as context [self attention] if None
"""
q, k, v = self.compute_qkv(x, context, rope_emb=rope_emb, transformer_options=transformer_options)
q, k, v = self.compute_qkv(x, context, rope_emb=rope_emb)
return self.compute_attention(q, k, v, transformer_options=transformer_options)
@@ -579,14 +555,8 @@ class Block(nn.Module):
self.layer_norm_mlp,
scale_mlp_B_T_1_1_D,
shift_mlp_B_T_1_1_D,
).to(compute_dtype)
patches = transformer_options.get("patches", {})
if "mlp_patch" in patches:
args = {"x": normalized_x_B_T_H_W_D, "transformer_options": transformer_options}
for patch in patches["mlp_patch"]:
args = patch(args)
normalized_x_B_T_H_W_D = args["x"]
result_B_T_H_W_D = self.mlp(normalized_x_B_T_H_W_D)
)
result_B_T_H_W_D = self.mlp(normalized_x_B_T_H_W_D.to(compute_dtype))
x_B_T_H_W_D = torch.addcmul(x_B_T_H_W_D, gate_mlp_B_T_1_1_D.to(residual_dtype), result_B_T_H_W_D.to(residual_dtype))
return x_B_T_H_W_D
@@ -893,22 +863,11 @@ class MiniTrainDIT(nn.Module):
x_B_T_H_W_D.shape == extra_pos_emb_B_T_H_W_D_or_T_H_W_B_D.shape
), f"{x_B_T_H_W_D.shape} != {extra_pos_emb_B_T_H_W_D_or_T_H_W_B_D.shape}"
transformer_options = kwargs.get("transformer_options", {})
patches = transformer_options.get("patches", {})
if "post_input" in patches:
transformer_options = transformer_options.copy()
transformer_options["model_patch_data"] = {}
if "post_input" in patches:
for patch in patches["post_input"]:
out = patch({"img": x_B_T_H_W_D, "x": x_B_C_T_H_W, "transformer_options": transformer_options})
x_B_T_H_W_D = out["img"]
block_kwargs = {
"rope_emb_L_1_1_D": rope_emb_L_1_1_D.unsqueeze(1).unsqueeze(0),
"adaln_lora_B_T_3D": adaln_lora_B_T_3D,
"extra_per_block_pos_emb": extra_pos_emb_B_T_H_W_D_or_T_H_W_B_D,
"transformer_options": transformer_options,
"transformer_options": kwargs.get("transformer_options", {}),
}
# The residual stream for this model has large values. To make fp16 compute_dtype work, we keep the residual stream
@@ -918,8 +877,7 @@ class MiniTrainDIT(nn.Module):
if x_B_T_H_W_D.dtype == torch.float16:
x_B_T_H_W_D = x_B_T_H_W_D.float()
for block_index, block in enumerate(self.blocks):
transformer_options["block_index"] = block_index
for block in self.blocks:
x_B_T_H_W_D = block(
x_B_T_H_W_D,
t_embedding_B_T_D,
+6 -6
View File
@@ -15,24 +15,24 @@ def make_two_pass_attention(ar_len: int, transformer_options=None):
The AR pass goes through SDPA directand bypasses wrappers, it is only ~1% of T at typical edit sizes.
"""
def two_pass_attention(q, k, v, heads, enable_gqa=False, **kwargs):
def two_pass_attention(q, k, v, heads, **kwargs):
B, H, T, D = q.shape
if T < k.shape[2]: # KV-cache hot path: Q is shorter than K/V (cached AR prefix is in K/V only), all fresh Q positions are in the gen region, single full-attention call
out = optimized_attention(q, k, v, heads, mask=None, skip_reshape=True, skip_output_reshape=True, transformer_options=transformer_options, enable_gqa=enable_gqa)
out = optimized_attention(q, k, v, heads, mask=None, skip_reshape=True, skip_output_reshape=True, transformer_options=transformer_options)
elif ar_len >= T:
out = comfy.ops.scaled_dot_product_attention(q, k, v, attn_mask=None, dropout_p=0.0, is_causal=True, enable_gqa=enable_gqa)
out = comfy.ops.scaled_dot_product_attention(q, k, v, attn_mask=None, dropout_p=0.0, is_causal=True)
elif ar_len <= 0:
out = optimized_attention(q, k, v, heads, mask=None, skip_reshape=True, skip_output_reshape=True, transformer_options=transformer_options, enable_gqa=enable_gqa)
out = optimized_attention(q, k, v, heads, mask=None, skip_reshape=True, skip_output_reshape=True, transformer_options=transformer_options)
else:
out_ar = comfy.ops.scaled_dot_product_attention(
q[:, :, :ar_len], k[:, :, :ar_len], v[:, :, :ar_len],
attn_mask=None, dropout_p=0.0, is_causal=True, enable_gqa=enable_gqa,
attn_mask=None, dropout_p=0.0, is_causal=True,
)
out_gen = optimized_attention(
q[:, :, ar_len:], k, v, heads,
mask=None, skip_reshape=True, skip_output_reshape=True,
transformer_options=transformer_options, enable_gqa=enable_gqa,
transformer_options=transformer_options,
)
out = torch.cat([out_ar, out_gen], dim=2)
-445
View File
@@ -1,445 +0,0 @@
# https://github.com/jdopensource/JoyAI-Image-Edit (Apache 2.0)
import math
from typing import Optional, Tuple
import comfy_kitchen
import torch
import torch.nn as nn
import comfy.ldm.common_dit
import comfy.ops
import comfy.patcher_extension
from comfy.ldm.lightricks.model import GELU_approx, PixArtAlphaTextProjection, TimestepEmbedding, Timesteps
from comfy.ldm.modules.attention import optimized_attention
class JoyImageModulate(nn.Module):
def __init__(self, hidden_size: int, factor: int, dtype=None, device=None):
super().__init__()
self.factor = factor
self.modulate_table = nn.Parameter(
torch.empty(1, factor, hidden_size, dtype=dtype, device=device)
)
def forward(self, x: torch.Tensor) -> list:
if x.ndim != 3:
x = x.unsqueeze(1)
table = comfy.ops.cast_to_input(self.modulate_table, x)
return [o.squeeze(1) for o in (table + x).chunk(self.factor, dim=1)]
class JoyImageFeedForward(nn.Module):
def __init__(
self,
dim: int,
inner_dim: int,
dtype=None,
device=None,
operations=None,
):
super().__init__()
self.net = nn.ModuleList([
GELU_approx(dim, inner_dim, dtype=dtype, device=device, operations=operations),
nn.Identity(),
operations.Linear(inner_dim, dim, bias=True, dtype=dtype, device=device),
])
def forward(self, x: torch.Tensor) -> torch.Tensor:
for module in self.net:
x = module(x)
return x
class JoyImageAttention(nn.Module):
def __init__(
self,
dim: int,
num_attention_heads: int,
attention_head_dim: int,
eps: float = 1e-6,
dtype=None,
device=None,
operations=None,
):
super().__init__()
self.num_attention_heads = num_attention_heads
inner_dim = num_attention_heads * attention_head_dim
self.img_attn_qkv = operations.Linear(dim, inner_dim * 3, bias=True, dtype=dtype, device=device)
self.img_attn_q_norm = operations.RMSNorm(attention_head_dim, eps=eps, dtype=dtype, device=device)
self.img_attn_k_norm = operations.RMSNorm(attention_head_dim, eps=eps, dtype=dtype, device=device)
self.img_attn_proj = operations.Linear(inner_dim, dim, bias=True, dtype=dtype, device=device)
self.txt_attn_qkv = operations.Linear(dim, inner_dim * 3, bias=True, dtype=dtype, device=device)
self.txt_attn_q_norm = operations.RMSNorm(attention_head_dim, eps=eps, dtype=dtype, device=device)
self.txt_attn_k_norm = operations.RMSNorm(attention_head_dim, eps=eps, dtype=dtype, device=device)
self.txt_attn_proj = operations.Linear(inner_dim, dim, bias=True, dtype=dtype, device=device)
def forward(
self,
img: torch.Tensor,
txt: torch.Tensor,
image_rotary_emb: torch.Tensor,
transformer_options=None,
) -> Tuple[torch.Tensor, torch.Tensor]:
heads = self.num_attention_heads
img_q, img_k, img_v = self.img_attn_qkv(img).chunk(3, dim=-1)
txt_q, txt_k, txt_v = self.txt_attn_qkv(txt).chunk(3, dim=-1)
img_q = img_q.unflatten(-1, (heads, -1))
img_k = img_k.unflatten(-1, (heads, -1))
img_v = img_v.unflatten(-1, (heads, -1))
txt_q = txt_q.unflatten(-1, (heads, -1))
txt_k = txt_k.unflatten(-1, (heads, -1))
txt_v = txt_v.unflatten(-1, (heads, -1))
img_q = self.img_attn_q_norm(img_q)
img_k = self.img_attn_k_norm(img_k)
txt_q = self.txt_attn_q_norm(txt_q)
txt_k = self.txt_attn_k_norm(txt_k)
img_q, img_k = comfy_kitchen.apply_rope(img_q, img_k, image_rotary_emb)
joint_q = torch.cat([img_q, txt_q], dim=1)
joint_k = torch.cat([img_k, txt_k], dim=1)
joint_v = torch.cat([img_v, txt_v], dim=1)
joint_q = joint_q.flatten(2, 3)
joint_k = joint_k.flatten(2, 3)
joint_v = joint_v.flatten(2, 3)
joint_out = optimized_attention(joint_q, joint_k, joint_v, heads=heads, transformer_options=transformer_options)
seq_img = img.shape[1]
img_out = joint_out[:, :seq_img, :]
txt_out = joint_out[:, seq_img:, :]
img_out = self.img_attn_proj(img_out)
txt_out = self.txt_attn_proj(txt_out)
return img_out, txt_out
class JoyImageTransformerBlock(nn.Module):
def __init__(
self,
dim: int,
num_attention_heads: int,
attention_head_dim: int,
mlp_width_ratio: float = 4.0,
eps: float = 1e-6,
dtype=None,
device=None,
operations=None,
):
super().__init__()
mlp_hidden_dim = int(dim * mlp_width_ratio)
self.img_mod = JoyImageModulate(dim, factor=6, dtype=dtype, device=device)
self.img_norm1 = operations.LayerNorm(dim, elementwise_affine=False, eps=eps, dtype=dtype, device=device)
self.img_norm2 = operations.LayerNorm(dim, elementwise_affine=False, eps=eps, dtype=dtype, device=device)
self.img_mlp = JoyImageFeedForward(dim, inner_dim=mlp_hidden_dim, dtype=dtype, device=device, operations=operations)
self.txt_mod = JoyImageModulate(dim, factor=6, dtype=dtype, device=device)
self.txt_norm1 = operations.LayerNorm(dim, elementwise_affine=False, eps=eps, dtype=dtype, device=device)
self.txt_norm2 = operations.LayerNorm(dim, elementwise_affine=False, eps=eps, dtype=dtype, device=device)
self.txt_mlp = JoyImageFeedForward(dim, inner_dim=mlp_hidden_dim, dtype=dtype, device=device, operations=operations)
self.attn = JoyImageAttention(
dim=dim,
num_attention_heads=num_attention_heads,
attention_head_dim=attention_head_dim,
eps=eps,
dtype=dtype,
device=device,
operations=operations,
)
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
temb: torch.Tensor,
image_rotary_emb: torch.Tensor,
transformer_options=None,
) -> Tuple[torch.Tensor, torch.Tensor]:
(
img_mod1_shift,
img_mod1_scale,
img_mod1_gate,
img_mod2_shift,
img_mod2_scale,
img_mod2_gate,
) = self.img_mod(temb)
(
txt_mod1_shift,
txt_mod1_scale,
txt_mod1_gate,
txt_mod2_shift,
txt_mod2_scale,
txt_mod2_gate,
) = self.txt_mod(temb)
img_normed = self.img_norm1(hidden_states)
txt_normed = self.txt_norm1(encoder_hidden_states)
img_modulated = img_normed * (1 + img_mod1_scale.unsqueeze(1)) + img_mod1_shift.unsqueeze(1)
txt_modulated = txt_normed * (1 + txt_mod1_scale.unsqueeze(1)) + txt_mod1_shift.unsqueeze(1)
img_attn, txt_attn = self.attn(img_modulated, txt_modulated, image_rotary_emb, transformer_options=transformer_options)
hidden_states = hidden_states + img_attn * img_mod1_gate.unsqueeze(1)
encoder_hidden_states = encoder_hidden_states + txt_attn * txt_mod1_gate.unsqueeze(1)
img_ffn_normed = self.img_norm2(hidden_states)
txt_ffn_normed = self.txt_norm2(encoder_hidden_states)
img_ffn_input = img_ffn_normed * (1 + img_mod2_scale.unsqueeze(1)) + img_mod2_shift.unsqueeze(1)
txt_ffn_input = txt_ffn_normed * (1 + txt_mod2_scale.unsqueeze(1)) + txt_mod2_shift.unsqueeze(1)
hidden_states = hidden_states + self.img_mlp(img_ffn_input) * img_mod2_gate.unsqueeze(1)
encoder_hidden_states = encoder_hidden_states + self.txt_mlp(txt_ffn_input) * txt_mod2_gate.unsqueeze(1)
return hidden_states, encoder_hidden_states
class JoyImageTimeTextImageEmbedding(nn.Module):
def __init__(
self,
dim: int,
time_freq_dim: int,
time_proj_dim: int,
text_embed_dim: int,
dtype=None,
device=None,
operations=None,
):
super().__init__()
self.timesteps_proj = Timesteps(num_channels=time_freq_dim, flip_sin_to_cos=True, downscale_freq_shift=0)
self.time_embedder = TimestepEmbedding(
in_channels=time_freq_dim,
time_embed_dim=dim,
dtype=dtype,
device=device,
operations=operations,
)
self.act_fn = nn.SiLU()
self.time_proj = operations.Linear(dim, time_proj_dim, bias=True, dtype=dtype, device=device)
self.text_embedder = PixArtAlphaTextProjection(
text_embed_dim, dim, act_fn="gelu_tanh", dtype=dtype, device=device, operations=operations,
)
def forward(self, timestep: torch.Tensor, encoder_hidden_states: torch.Tensor):
timestep = self.timesteps_proj(timestep)
temb = self.time_embedder(timestep.to(dtype=encoder_hidden_states.dtype)).type_as(encoder_hidden_states)
timestep_proj = self.time_proj(self.act_fn(temb))
encoder_hidden_states = self.text_embedder(encoder_hidden_states)
return temb, timestep_proj, encoder_hidden_states
class JoyImageTransformer3DModel(nn.Module):
def __init__(
self,
patch_size: list = [1, 2, 2],
in_channels: int = 16,
out_channels: Optional[int] = None,
hidden_size: int = 3072,
num_attention_heads: int = 24,
text_dim: int = 4096,
mlp_width_ratio: float = 4.0,
num_layers: int = 20,
rope_dim_list: list = [16, 56, 56],
theta: int = 256,
image_model=None,
dtype=None,
device=None,
operations=None,
):
super().__init__()
self.dtype = dtype
self.out_channels = out_channels or in_channels
self.patch_size = list(patch_size)
self.rope_dim_list = list(rope_dim_list)
self.theta = theta
attention_head_dim = hidden_size // num_attention_heads
self.img_in = operations.Conv3d(
in_channels,
hidden_size,
kernel_size=tuple(self.patch_size),
stride=tuple(self.patch_size),
dtype=dtype,
device=device,
)
self.condition_embedder = JoyImageTimeTextImageEmbedding(
dim=hidden_size,
time_freq_dim=256,
time_proj_dim=hidden_size * 6,
text_embed_dim=text_dim,
dtype=dtype,
device=device,
operations=operations,
)
self.double_blocks = nn.ModuleList([
JoyImageTransformerBlock(
dim=hidden_size,
num_attention_heads=num_attention_heads,
attention_head_dim=attention_head_dim,
mlp_width_ratio=mlp_width_ratio,
dtype=dtype,
device=device,
operations=operations,
)
for _ in range(num_layers)
])
self.norm_out = operations.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, dtype=dtype, device=device)
self.proj_out = operations.Linear(
hidden_size,
self.out_channels * math.prod(self.patch_size),
bias=True,
dtype=dtype,
device=device,
)
def _get_rotary_pos_embed_for_range(
self,
start: Tuple[int, int, int],
stop: Tuple[int, int, int],
device=None,
) -> torch.Tensor:
# 3D RoPE for the patch grid range [start, stop) over (t, h, w). Token order after
# reshape(-1) is (t, h, w), matching the img_in Conv3d flatten.
rope_dim_list = self.rope_dim_list
grids = [torch.arange(start[i], stop[i], dtype=torch.float32, device=device) for i in range(3)]
mesh = torch.stack(torch.meshgrid(*grids, indexing="ij"), dim=0)
angles_parts = []
for i, dim in enumerate(rope_dim_list):
pos = mesh[i].reshape(-1)
freqs = 1.0 / (self.theta ** (torch.arange(0, dim, 2, dtype=torch.float32, device=device)[: (dim // 2)] / dim))
angles_parts.append(torch.outer(pos, freqs))
angles = torch.cat(angles_parts, dim=1)
cos = angles.cos()
sin = angles.sin()
return torch.stack((cos, -sin, sin, cos), dim=-1).unflatten(-1, (2, 2))
def get_rotary_pos_embed_for_components(
self,
component_sizes,
device=None,
) -> torch.Tensor:
# Per-component 3D RoPE. component_sizes is a list of (t, h, w) patch grid sizes in
# sequence order [target, ref0, ref1, ...]; h/w restart at 0 for each component while t
# continues from the running offset, giving every image its own temporal position band.
freqs_parts = []
t_offset = 0
for (t, h, w) in component_sizes:
freqs = self._get_rotary_pos_embed_for_range(
start=(t_offset, 0, 0),
stop=(t_offset + t, h, w),
device=device,
)
freqs_parts.append(freqs)
t_offset += t
return torch.cat(freqs_parts, dim=0).unsqueeze(0).unsqueeze(2)
def unpatchify(self, x: torch.Tensor, t: int, h: int, w: int) -> torch.Tensor:
c = self.out_channels
pt, ph, pw = self.patch_size
x = x.reshape(x.shape[0], t, h, w, pt, ph, pw, c)
x = x.permute(0, 7, 1, 4, 2, 5, 3, 6)
return x.reshape(x.shape[0], c, t * pt, h * ph, w * pw)
def forward(
self,
hidden_states: torch.Tensor,
timestep: torch.Tensor,
context: torch.Tensor = None,
ref_latents=None,
control=None,
transformer_options=None,
**kwargs,
) -> torch.Tensor:
transformer_options = {} if transformer_options is None else transformer_options.copy()
return comfy.patcher_extension.WrapperExecutor.new_class_executor(
self._forward,
self,
comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.DIFFUSION_MODEL, transformer_options)
).execute(hidden_states, timestep, context, ref_latents, transformer_options, **kwargs)
def _forward(
self,
hidden_states: torch.Tensor,
timestep: torch.Tensor,
context: torch.Tensor,
ref_latents=None,
transformer_options=None,
**kwargs,
) -> torch.Tensor:
pt, ph, pw = self.patch_size
_, _, ot, oh, ow = hidden_states.shape
components = [hidden_states, *(ref_latents or [])]
component_sizes = []
img_tokens = []
for comp in components:
comp = comfy.ldm.common_dit.pad_to_patch_size(comp, self.patch_size)
_, _, ct, ch, cw = comp.shape
component_sizes.append((ct // pt, ch // ph, cw // pw))
tokens = self.img_in(comp).flatten(2).transpose(1, 2) # (B, n_i, D)
img_tokens.append(tokens)
img = torch.cat(img_tokens, dim=1)
_, vec, txt = self.condition_embedder(timestep, context)
vec = vec.unflatten(1, (6, -1))
image_rotary_emb = self.get_rotary_pos_embed_for_components(
component_sizes,
device=hidden_states.device,
)
patches_replace = transformer_options.get("patches_replace", {})
blocks_replace = patches_replace.get("dit", {})
transformer_options["total_blocks"] = len(self.double_blocks)
transformer_options["block_type"] = "double"
for i, block in enumerate(self.double_blocks):
transformer_options["block_index"] = i
if ("double_block", i) in blocks_replace:
def block_wrap(args):
out = {}
out["img"], out["txt"] = block(
hidden_states=args["img"],
encoder_hidden_states=args["txt"],
temb=args["vec"],
image_rotary_emb=args["pe"],
transformer_options=args.get("transformer_options"),
)
return out
out = blocks_replace[("double_block", i)]({"img": img,
"txt": txt,
"vec": vec,
"pe": image_rotary_emb,
"transformer_options": transformer_options},
{"original_block": block_wrap})
txt = out["txt"]
img = out["img"]
else:
img, txt = block(
hidden_states=img,
encoder_hidden_states=txt,
temb=vec,
image_rotary_emb=image_rotary_emb,
transformer_options=transformer_options,
)
tt, th, tw = component_sizes[0]
target_tokens = tt * th * tw
img = img[:, :target_tokens, :]
img = self.proj_out(self.norm_out(img))
img = self.unpatchify(img, tt, th, tw)
return img[:, :, :ot, :oh, :ow]
+23 -124
View File
@@ -15,7 +15,6 @@ from einops import rearrange
import comfy.model_management
import comfy.patcher_extension
import comfy.ldm.common_dit
import comfy.utils
from comfy.ldm.flux.layers import EmbedND, timestep_embedding
from comfy.ldm.flux.math import apply_rope
from comfy.ldm.modules.attention import optimized_attention_masked
@@ -74,20 +73,11 @@ class Attention(nn.Module):
self.wo = operations.Linear(dim, dim, bias=bias, device=device, dtype=dtype)
def forward(self, x, freqs=None, mask=None, transformer_options={}):
transformer_patches = transformer_options.get("patches", {})
extra_options = transformer_options.copy()
q, k, v, gate = self.wq(x), self.wk(x), self.wv(x), self.gate(x)
q = rearrange(q, "B L (H D) -> B H L D", H=self.heads)
k = rearrange(k, "B L (H D) -> B H L D", H=self.kvheads)
v = rearrange(v, "B L (H D) -> B H L D", H=self.kvheads)
q, k = self.qknorm(q, k)
if "block_index" in transformer_options and "attn1_patch" in transformer_patches:
for p in transformer_patches["attn1_patch"]:
out = p(q, k, v, pe=freqs, attn_mask=mask, extra_options=extra_options)
q, k, v = out.get("q", q), out.get("k", k), out.get("v", v)
freqs, mask = out.get("pe", freqs), out.get("attn_mask", mask)
if freqs is not None:
q, k = apply_rope(q, k, freqs)
if self.kvheads != self.heads:
@@ -96,11 +86,6 @@ class Attention(nn.Module):
v = v.repeat_interleave(rep, dim=1)
out = optimized_attention_masked(q, k, v, self.heads, mask=mask, skip_reshape=True,
transformer_options=transformer_options)
if "block_index" in transformer_options and "attn1_output_patch" in transformer_patches:
for p in transformer_patches["attn1_output_patch"]:
out = p(out, extra_options)
return self.wo(out * F.sigmoid(gate))
@@ -173,44 +158,8 @@ class SingleStreamBlock(nn.Module):
self.attn = Attention(features, heads, kvheads=kvheads, bias=bias, device=device, dtype=dtype, operations=operations)
self.mlp = SwiGLU(features, multiplier, bias, device=device, dtype=dtype, operations=operations)
def forward(self, x, vec, freqs, mask=None, timestep_zero_index=None, transformer_options={}):
def forward(self, x, vec, freqs, mask=None, transformer_options={}):
prescale, preshift, pregate, postscale, postshift, postgate = self.mod(vec)
if timestep_zero_index is not None:
bs = x.shape[0]
ref_prescale = prescale[bs:]
ref_preshift = preshift[bs:]
ref_pregate = pregate[bs:]
ref_postscale = postscale[bs:]
ref_postshift = postshift[bs:]
ref_postgate = postgate[bs:]
prescale = prescale[:bs]
preshift = preshift[:bs]
pregate = pregate[:bs]
postscale = postscale[:bs]
postshift = postshift[:bs]
postgate = postgate[:bs]
pre = self.prenorm(x)
pre[:, :timestep_zero_index].mul_(1 + prescale).add_(preshift)
pre[:, timestep_zero_index:].mul_(1 + ref_prescale).add_(ref_preshift)
attn = self.attn(pre, freqs, mask, transformer_options=transformer_options)
del pre
attn[:, :timestep_zero_index].mul_(pregate)
attn[:, timestep_zero_index:].mul_(ref_pregate)
x = x + attn
del attn
post = self.postnorm(x)
post[:, :timestep_zero_index].mul_(1 + postscale).add_(postshift)
post[:, timestep_zero_index:].mul_(1 + ref_postscale).add_(ref_postshift)
mlp = self.mlp(post)
del post
mlp[:, :timestep_zero_index].mul_(postgate)
mlp[:, timestep_zero_index:].mul_(ref_postgate)
x = x + mlp
del mlp
return x
x = x + pregate * self.attn((1 + prescale) * self.prenorm(x) + preshift, freqs, mask, transformer_options=transformer_options)
x = x + postgate * self.mlp((1 + postscale) * self.postnorm(x) + postshift)
return x
@@ -232,7 +181,7 @@ class LastLayer(nn.Module):
class SingleStreamDiT(nn.Module):
def __init__(self, features=6144, tdim=256, txtdim=2560, heads=48, kvheads=12, multiplier=4,
layers=28, patch=2, channels=16, bias=False, theta=1e3, txtlayers=12,
txtheads=20, txtkvheads=20, default_ref_method=None, image_model=None,
txtheads=20, txtkvheads=20, image_model=None,
device=None, dtype=None, operations=None, **kwargs):
super().__init__()
self.dtype = dtype
@@ -242,7 +191,6 @@ class SingleStreamDiT(nn.Module):
self.heads = heads
self.txtdim = txtdim
self.txtlayers = txtlayers
self.default_ref_method = default_ref_method
headdim = features // heads
axes = [headdim - 12 * (headdim // 16), 6 * (headdim // 16), 6 * (headdim // 16)]
@@ -273,110 +221,61 @@ class SingleStreamDiT(nn.Module):
operations.Linear(features, features * 6, device=device, dtype=dtype),
)
def forward(self, x, timesteps, context, attention_mask=None, ref_latents=None, transformer_options={}, **kwargs):
def forward(self, x, timesteps, context, attention_mask=None, transformer_options={}, **kwargs):
return comfy.patcher_extension.WrapperExecutor.new_class_executor(
self._forward,
self,
comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.DIFFUSION_MODEL, transformer_options),
).execute(x, timesteps, context, attention_mask, ref_latents, transformer_options, **kwargs)
).execute(x, timesteps, context, attention_mask, transformer_options, **kwargs)
def process_img(self, x, index=0):
patch = self.patch
x = comfy.ldm.common_dit.pad_to_patch_size(x, (patch, patch))
h, w = x.shape[-2] // patch, x.shape[-1] // patch
img = rearrange(x, "b c (h ph) (w pw) -> b (h w) (c ph pw)", ph=patch, pw=patch)
img_ids = torch.zeros(h, w, 3, device=x.device, dtype=torch.float32)
img_ids[..., 0] = index
img_ids[..., 1] = torch.arange(h, device=x.device, dtype=torch.float32)[:, None]
img_ids[..., 2] = torch.arange(w, device=x.device, dtype=torch.float32)[None, :]
return img, img_ids.reshape(1, h * w, 3).repeat(x.shape[0], 1, 1), h, w
def _forward(self, x, timesteps, context, attention_mask=None, ref_latents=None, transformer_options={}, **kwargs):
transformer_options = transformer_options.copy()
def _forward(self, x, timesteps, context, attention_mask=None, transformer_options={}, **kwargs):
temporal = x.ndim == 5
if temporal:
b5, c5, t5, h5, w5 = x.shape
x = x.reshape(b5 * t5, c5, h5, w5)
bs, _, h_orig, w_orig = x.shape
bs, c, H_orig, W_orig = x.shape
patch = self.patch
# Pad the latent up to a multiple of patch (as Flux/Lumina/QwenImage do); crop back at the end.
x = comfy.ldm.common_dit.pad_to_patch_size(x, (patch, patch))
H, W = x.shape[-2], x.shape[-1]
h_, w_ = H // patch, W // patch
# context arrives as (B, seq, txtlayers*txtdim); reshape to (B, txtlayers, seq, txtdim).
context = self._unpack_context(context)
img, imgpos, h_, w_ = self.process_img(x)
img_tokens = img.shape[1]
timestep_zero_index = None
ref_method = kwargs.get("ref_latents_method", self.default_ref_method)
if ref_method is not None and ref_latents is not None and len(ref_latents) > 0:
ref_tokens = []
ref_pos = []
ref_num_tokens = []
for index, ref in enumerate(ref_latents, 1):
if ref.ndim == 5:
rb, rc, rt, rh5, rw5 = ref.shape
ref = ref.reshape(rb * rt, rc, rh5, rw5)
ref = comfy.utils.repeat_to_batch_size(ref, bs)
kontext, kontext_ids, _, _ = self.process_img(ref, index=index)
ref_tokens.append(kontext)
ref_pos.append(kontext_ids)
ref_num_tokens.append(kontext.shape[1])
img = torch.cat([img] + ref_tokens, dim=1)
imgpos = torch.cat([imgpos] + ref_pos, dim=1)
del ref_tokens, ref_pos
if ref_method == "index_timestep_zero":
timestep_zero_index = img_tokens
transformer_options["reference_image_num_tokens"] = ref_num_tokens
img = rearrange(x, "b c (h ph) (w pw) -> b (h w) (c ph pw)", ph=patch, pw=patch)
img = self.first(img)
t = self.tmlp(timestep_embedding(timesteps, self.tdim).unsqueeze(1).to(img.dtype))
tvec = self.tproj(t)
if timestep_zero_index is not None:
t0 = self.tmlp(timestep_embedding(torch.zeros_like(timesteps), self.tdim).unsqueeze(1).to(img.dtype))
tvec = torch.cat((tvec, self.tproj(t0)), dim=0)
context = self.txtfusion(context, mask=None, transformer_options=transformer_options)
context = self.txtmlp(context)
txtlen = context.shape[1]
device = context.device
txtpos = torch.zeros(bs, txtlen, 3, device=device, dtype=torch.float32)
patches = transformer_options.get("patches", {})
if "post_input" in patches:
for p in patches["post_input"]:
out = p({"img": img, "txt": context, "img_ids": imgpos, "txt_ids": txtpos, "transformer_options": transformer_options})
img, context = out["img"], out["txt"]
imgpos, txtpos = out["img_ids"], out["txt_ids"]
txtlen, imglen = context.shape[1], img.shape[1]
combined = torch.cat((context, img), dim=1)
del context, img
if timestep_zero_index is not None:
timestep_zero_index += txtlen
# Position ids: text at 0, image at (0, h_idx, w_idx).
device = combined.device
txtpos = torch.zeros(bs, txtlen, 3, device=device, dtype=torch.float32)
imgids = torch.zeros(h_, w_, 3, device=device, dtype=torch.float32)
imgids[..., 1] = torch.arange(h_, device=device, dtype=torch.float32)[:, None]
imgids[..., 2] = torch.arange(w_, device=device, dtype=torch.float32)[None, :]
imgpos = imgids.reshape(1, h_ * w_, 3).repeat(bs, 1, 1)
pos = torch.cat((txtpos, imgpos), dim=1)
del txtpos, imgpos
freqs = self.pe_embedder(pos)
del pos
transformer_options["total_blocks"] = len(self.blocks)
transformer_options["block_type"] = "single"
transformer_options["img_slice"] = [txtlen, combined.shape[1]]
for i, block in enumerate(self.blocks):
transformer_options["block_index"] = i
combined = block(combined, tvec, freqs, None, timestep_zero_index=timestep_zero_index, transformer_options=transformer_options)
for block in self.blocks:
combined = block(combined, tvec, freqs, None, transformer_options=transformer_options)
final = self.last(combined, t)
del combined
out = final[:, txtlen:txtlen + img_tokens, :]
out = final[:, txtlen:txtlen + imglen, :]
out = rearrange(out, "b (h w) (c ph pw) -> b c (h ph) (w pw)",
h=h_, w=w_, ph=patch, pw=patch, c=self.channels)
out = out[:, :, :h_orig, :w_orig] # crop padding back off
out = out[:, :, :H_orig, :W_orig] # crop padding back off
if temporal:
out = out.reshape(b5, t5, self.channels, h_orig, w_orig).movedim(1, 2)
out = out.reshape(b5, t5, self.channels, H_orig, W_orig).movedim(1, 2)
return out
def _unpack_context(self, context):
-50
View File
@@ -552,7 +552,6 @@ class WanModel(torch.nn.Module):
List of denoised video tensors with original input shapes [C_out, F, H / 8, W / 8]
"""
# embeddings
x_input = x
x = self.patch_embedding(x.float()).to(x.dtype)
grid_sizes = x.shape[2:]
transformer_options["grid_sizes"] = grid_sizes
@@ -565,13 +564,11 @@ class WanModel(torch.nn.Module):
e0 = self.time_projection(e).unflatten(2, (6, self.dim))
full_ref = None
img_offset = 0
if self.ref_conv is not None:
full_ref = kwargs.get("reference_latent", None)
if full_ref is not None:
full_ref = self.ref_conv(full_ref).flatten(2).transpose(1, 2)
x = torch.concat((full_ref, x), dim=1)
img_offset = full_ref.shape[1]
# In-context reference (Bernini)
context_latents = kwargs.get("context_latents", None)
@@ -592,7 +589,6 @@ class WanModel(torch.nn.Module):
context_img_len = clip_fea.shape[-2]
patches_replace = transformer_options.get("patches_replace", {})
patches = transformer_options.get("patches", {})
blocks_replace = patches_replace.get("dit", {})
transformer_options["total_blocks"] = len(self.blocks)
transformer_options["block_type"] = "double"
@@ -608,11 +604,6 @@ class WanModel(torch.nn.Module):
else:
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options)
if "double_block" in patches:
for p in patches["double_block"]:
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": img_offset, "transformer_options": transformer_options})
x = out["img"]
# head
x = self.head(x, e)
@@ -786,7 +777,6 @@ class VaceWanModel(WanModel):
**kwargs,
):
# embeddings
x_input = x
x = self.patch_embedding(x.float()).to(x.dtype)
grid_sizes = x.shape[2:]
transformer_options["grid_sizes"] = grid_sizes
@@ -817,7 +807,6 @@ class VaceWanModel(WanModel):
x_orig = x
patches_replace = transformer_options.get("patches_replace", {})
patches = transformer_options.get("patches", {})
blocks_replace = patches_replace.get("dit", {})
transformer_options["total_blocks"] = len(self.blocks)
transformer_options["block_type"] = "double"
@@ -833,11 +822,6 @@ class VaceWanModel(WanModel):
else:
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options)
if "double_block" in patches:
for p in patches["double_block"]:
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": 0, "transformer_options": transformer_options})
x = out["img"]
ii = self.vace_layers_mapping.get(i, None)
if ii is not None:
for iii in range(len(c)):
@@ -903,7 +887,6 @@ class CameraWanModel(WanModel):
**kwargs,
):
# embeddings
x_input = x
x = self.patch_embedding(x.float()).to(x.dtype)
if self.control_adapter is not None and camera_conditions is not None:
x = x + self.control_adapter(camera_conditions).to(x.dtype)
@@ -926,7 +909,6 @@ class CameraWanModel(WanModel):
context_img_len = clip_fea.shape[-2]
patches_replace = transformer_options.get("patches_replace", {})
patches = transformer_options.get("patches", {})
blocks_replace = patches_replace.get("dit", {})
transformer_options["total_blocks"] = len(self.blocks)
transformer_options["block_type"] = "double"
@@ -942,11 +924,6 @@ class CameraWanModel(WanModel):
else:
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options)
if "double_block" in patches:
for p in patches["double_block"]:
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": 0, "transformer_options": transformer_options})
x = out["img"]
# head
x = self.head(x, e)
@@ -1358,7 +1335,6 @@ class WanModel_S2V(WanModel):
# embeddings
bs, _, time, height, width = x.shape
x_input = x
x = self.patch_embedding(x.float()).to(x.dtype)
if control_video is not None:
x = x + self.cond_encoder(control_video)
@@ -1403,7 +1379,6 @@ class WanModel_S2V(WanModel):
context = self.text_embedding(context)
patches_replace = transformer_options.get("patches_replace", {})
patches = transformer_options.get("patches", {})
blocks_replace = patches_replace.get("dit", {})
transformer_options["total_blocks"] = len(self.blocks)
transformer_options["block_type"] = "double"
@@ -1418,12 +1393,6 @@ class WanModel_S2V(WanModel):
x = out["img"]
else:
x = block(x, e=e0, freqs=freqs, context=context, transformer_options=transformer_options)
if "double_block" in patches:
for p in patches["double_block"]:
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": 0, "transformer_options": transformer_options})
x = out["img"]
if audio_emb is not None:
x = self.audio_injector(x, i, audio_emb, audio_emb_global, seq_len)
# head
@@ -1630,7 +1599,6 @@ class HumoWanModel(WanModel):
bs, _, time, height, width = x.shape
# embeddings
x_input = x
x = self.patch_embedding(x.float()).to(x.dtype)
grid_sizes = x.shape[2:]
x = x.flatten(2).transpose(1, 2)
@@ -1662,7 +1630,6 @@ class HumoWanModel(WanModel):
audio = None
patches_replace = transformer_options.get("patches_replace", {})
patches = transformer_options.get("patches", {})
blocks_replace = patches_replace.get("dit", {})
transformer_options["total_blocks"] = len(self.blocks)
transformer_options["block_type"] = "double"
@@ -1678,11 +1645,6 @@ class HumoWanModel(WanModel):
else:
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, audio=audio, transformer_options=transformer_options)
if "double_block" in patches:
for p in patches["double_block"]:
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": 0, "transformer_options": transformer_options})
x = out["img"]
# head
x = self.head(x, e)
@@ -1698,14 +1660,8 @@ class SCAILWanModel(WanModel):
def forward_orig(self, x, t, context, clip_fea=None, freqs=None, transformer_options={}, pose_latents=None, reference_latent=None, ref_mask_latents=None, sam_latents=None, **kwargs):
x_input = x
img_offset = 0
if reference_latent is not None:
x = torch.cat((reference_latent, x), dim=2)
img_offset = (reference_latent.shape[2] // self.patch_size[0]) * \
(reference_latent.shape[3] // self.patch_size[1]) * \
(reference_latent.shape[4] // self.patch_size[2])
# embeddings
x = self.patch_embedding(x.float()).to(x.dtype)
@@ -1741,7 +1697,6 @@ class SCAILWanModel(WanModel):
context_img_len = clip_fea.shape[-2]
patches_replace = transformer_options.get("patches_replace", {})
patches = transformer_options.get("patches", {})
blocks_replace = patches_replace.get("dit", {})
transformer_options["total_blocks"] = len(self.blocks)
transformer_options["block_type"] = "double"
@@ -1757,11 +1712,6 @@ class SCAILWanModel(WanModel):
else:
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options)
if "double_block" in patches:
for p in patches["double_block"]:
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": img_offset, "transformer_options": transformer_options})
x = out["img"]
# head
x = self.head(x, e)
-9
View File
@@ -493,7 +493,6 @@ class AnimateWanModel(WanModel):
**kwargs,
):
# embeddings
x_input = x
x = self.patch_embedding(x.float()).to(x.dtype)
x, motion_vec = self.after_patch_embedding(x, pose_latents, face_pixel_values)
grid_sizes = x.shape[2:]
@@ -506,13 +505,11 @@ class AnimateWanModel(WanModel):
e0 = self.time_projection(e).unflatten(2, (6, self.dim))
full_ref = None
img_offset = 0
if self.ref_conv is not None:
full_ref = kwargs.get("reference_latent", None)
if full_ref is not None:
full_ref = self.ref_conv(full_ref).flatten(2).transpose(1, 2)
x = torch.concat((full_ref, x), dim=1)
img_offset = full_ref.shape[1]
# context
context = self.text_embedding(context)
@@ -525,7 +522,6 @@ class AnimateWanModel(WanModel):
context_img_len = clip_fea.shape[-2]
patches_replace = transformer_options.get("patches_replace", {})
patches = transformer_options.get("patches", {})
blocks_replace = patches_replace.get("dit", {})
transformer_options["total_blocks"] = len(self.blocks)
transformer_options["block_type"] = "double"
@@ -541,11 +537,6 @@ class AnimateWanModel(WanModel):
else:
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options)
if "double_block" in patches:
for p in patches["double_block"]:
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": img_offset, "transformer_options": transformer_options})
x = out["img"]
if i % 5 == 0 and motion_vec is not None:
x = x + self.face_adapter.fuser_blocks[i // 5](x, motion_vec)
-10
View File
@@ -111,7 +111,6 @@ class WanDancerModel(WanModel):
def forward_orig(self, x, t, context, clip_fea=None, clip_fea_ref=None, freqs=None, audio_embed=None, fps=30, audio_inject_scale=1.0, transformer_options={}, **kwargs):
# embeddings
x_input = x
if int(fps + 0.5) != 30:
x = self.patch_embedding_global(x.float()).to(x.dtype)
else:
@@ -129,13 +128,11 @@ class WanDancerModel(WanModel):
e0 = self.time_projection(e).unflatten(2, (6, self.dim))
full_ref = None
img_offset = 0
if self.ref_conv is not None: # model has the weight, but this wasn't used in the original pipeline
full_ref = kwargs.get("reference_latent", None)
if full_ref is not None:
full_ref = self.ref_conv(full_ref).flatten(2).transpose(1, 2)
x = torch.concat((full_ref, x), dim=1)
img_offset = full_ref.shape[1]
# context
context = self.text_embedding(context)
@@ -166,7 +163,6 @@ class WanDancerModel(WanModel):
context_img_len += clip_fea_ref.shape[-2]
patches_replace = transformer_options.get("patches_replace", {})
patches = transformer_options.get("patches", {})
blocks_replace = patches_replace.get("dit", {})
transformer_options["total_blocks"] = len(self.blocks)
transformer_options["block_type"] = "double"
@@ -181,12 +177,6 @@ class WanDancerModel(WanModel):
x = out["img"]
else:
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options)
if "double_block" in patches:
for p in patches["double_block"]:
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": img_offset, "transformer_options": transformer_options})
x = out["img"]
if audio_emb is not None:
x = self.music_injector(x, i, audio_emb, audio_emb_global=None, seq_len=seq_len, scale=audio_inject_scale)
-149
View File
@@ -1,149 +0,0 @@
# Uni3C controlnet for Wan 2.1: https://github.com/ewrfcas/Uni3C
# Converted from the original diffusers based implementation.
import torch
import torch.nn as nn
from comfy.ldm.flux.layers import EmbedND
from .model import WanSelfAttention
class Uni3CLayerNormZero(nn.Module):
def __init__(
self,
conditioning_dim,
embedding_dim,
eps=1e-5,
device=None, dtype=None, operations=None
):
super().__init__()
self.silu = nn.SiLU()
self.linear = operations.Linear(conditioning_dim, 3 * embedding_dim, device=device, dtype=dtype)
self.norm = operations.LayerNorm(embedding_dim, eps=eps, elementwise_affine=True, device=device, dtype=dtype)
def forward(self, x, temb):
shift, scale, gate = self.linear(self.silu(temb)).chunk(3, dim=1)
x = self.norm(x) * (1 + scale)[:, None, :] + shift[:, None, :]
return x, gate[:, None, :]
class Uni3CAttentionBlock(nn.Module):
def __init__(
self,
dim,
ffn_dim,
num_heads,
time_embed_dim=5120,
eps=1e-6,
device=None, dtype=None, operations=None
):
super().__init__()
operation_settings = {"operations": operations, "device": device, "dtype": dtype}
self.norm1 = Uni3CLayerNormZero(time_embed_dim, dim, device=device, dtype=dtype, operations=operations)
self.self_attn = WanSelfAttention(dim, num_heads, qk_norm=True, eps=eps, operation_settings=operation_settings)
self.norm2 = Uni3CLayerNormZero(time_embed_dim, dim, device=device, dtype=dtype, operations=operations)
self.ffn = nn.Sequential(
operations.Linear(dim, ffn_dim, device=device, dtype=dtype), nn.GELU(approximate='tanh'),
operations.Linear(ffn_dim, dim, device=device, dtype=dtype))
def forward(self, x, temb, freqs):
norm_x, gate_msa = self.norm1(x, temb)
x = x + gate_msa * self.self_attn(norm_x, freqs)
norm_x, gate_ff = self.norm2(x, temb)
x = x + gate_ff * self.ffn(norm_x)
return x
class MaskCamEmbed(nn.Module):
def __init__(
self,
add_channels=7,
mid_channels=256,
conv_out_dim=5120,
device=None, dtype=None, operations=None
):
super().__init__()
self.mask_padding = [0, 0, 0, 0, 3, 0] # first frame conditioning
self.mask_proj = nn.Sequential(
operations.Conv3d(add_channels, mid_channels, kernel_size=(4, 8, 8), stride=(4, 8, 8), device=device, dtype=dtype),
operations.GroupNorm(mid_channels // 8, mid_channels, device=device, dtype=dtype),
nn.SiLU())
self.mask_zero_proj = operations.Conv3d(mid_channels, conv_out_dim, kernel_size=(1, 2, 2), stride=(1, 2, 2), device=device, dtype=dtype)
def forward(self, add_inputs):
add_padded = torch.nn.functional.pad(add_inputs, self.mask_padding, mode="constant", value=0)
add_embeds = self.mask_proj(add_padded)
add_embeds = self.mask_zero_proj(add_embeds)
add_embeds = add_embeds.flatten(2).transpose(1, 2)
return add_embeds
class WanUni3CControlnet(nn.Module):
def __init__(
self,
in_channels=36,
conv_out_dim=5120,
dim=1024,
ffn_dim=8192,
num_heads=16,
num_layers=20,
time_embed_dim=5120,
out_proj_dim=5120,
add_channels=7,
mid_channels=256,
device=None, dtype=None, operations=None
):
super().__init__()
patch_size = (1, 2, 2)
self.num_layers = num_layers
self.controlnet_patch_embedding = operations.Conv3d(
in_channels, conv_out_dim, kernel_size=patch_size, stride=patch_size, device=device, dtype=torch.float32)
self.controlnet_mask_embedding = MaskCamEmbed(add_channels, mid_channels, conv_out_dim, device=device, dtype=dtype, operations=operations)
if conv_out_dim != dim:
self.proj_in = operations.Linear(conv_out_dim, dim, device=device, dtype=dtype)
else:
self.proj_in = nn.Identity()
self.controlnet_blocks = nn.ModuleList([
Uni3CAttentionBlock(dim, ffn_dim, num_heads, time_embed_dim, device=device, dtype=dtype, operations=operations)
for _ in range(num_layers)])
self.proj_out = nn.ModuleList([
operations.Linear(dim, out_proj_dim, device=device, dtype=dtype)
for _ in range(num_layers)])
head_dim = dim // num_heads
self.rope_embedder = EmbedND(dim=head_dim, theta=10000.0, axes_dim=[head_dim - 4 * (head_dim // 6), 2 * (head_dim // 6), 2 * (head_dim // 6)])
def rope_encode(self, t_len, h_len, w_len, device=None, dtype=None):
img_ids = torch.zeros((t_len, h_len, w_len, 3), device=device, dtype=dtype)
img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + torch.arange(t_len, device=device, dtype=dtype).reshape(-1, 1, 1)
img_ids[:, :, :, 1] = img_ids[:, :, :, 1] + torch.arange(h_len, device=device, dtype=dtype).reshape(1, -1, 1)
img_ids[:, :, :, 2] = img_ids[:, :, :, 2] + torch.arange(w_len, device=device, dtype=dtype).reshape(1, 1, -1)
img_ids = img_ids.reshape(1, -1, img_ids.shape[-1])
freqs = self.rope_embedder(img_ids).movedim(1, 2)
return freqs
def process_input(self, control_input, render_mask=None, camera_embedding=None):
# render_mask/camera_embedding are the checkpoint's extra conditioning path, not wired up yet
hidden = self.controlnet_patch_embedding(control_input.float()).to(control_input.dtype)
t_len, h_len, w_len = hidden.shape[2:]
freqs = self.rope_encode(t_len, h_len, w_len, device=hidden.device, dtype=hidden.dtype)
hidden = hidden.flatten(2).transpose(1, 2)
add_inputs = None
if camera_embedding is not None and render_mask is not None:
add_inputs = torch.cat([render_mask, camera_embedding], dim=1)
elif render_mask is not None:
add_inputs = render_mask
if add_inputs is not None:
hidden = hidden + self.controlnet_mask_embedding(add_inputs.to(hidden.dtype))
hidden = self.proj_in(hidden)
return hidden, freqs
def forward_block(self, block_index, hidden, temb, freqs):
hidden = self.controlnet_blocks[block_index](hidden, temb, freqs)
residual = self.proj_out[block_index](hidden)
return hidden, residual
+6 -44
View File
@@ -58,7 +58,6 @@ import comfy.ldm.omnigen.omnigen2
import comfy.ldm.seedvr.model
import comfy.ldm.boogu.model
import comfy.ldm.qwen_image.model
import comfy.ldm.joyimage.model
import comfy.ldm.ideogram4.model
import comfy.ldm.krea2.model
import comfy.ldm.kandinsky5.model
@@ -2024,11 +2023,11 @@ class WAN22_WanDancer(WAN21):
fps = kwargs.get("fps", None)
if fps is not None:
out['fps'] = comfy.conds.CONDConstant(fps)
out['fps'] = comfy.conds.CONDRegular(torch.FloatTensor([fps]))
audio_inject_scale = kwargs.get("audio_inject_scale", None)
if audio_inject_scale is not None:
out['audio_inject_scale'] = comfy.conds.CONDConstant(audio_inject_scale)
out['audio_inject_scale'] = comfy.conds.CONDRegular(torch.FloatTensor([audio_inject_scale]))
return out
class Hunyuan3Dv2(BaseModel):
@@ -2227,7 +2226,10 @@ class Omnigen2(BaseModel):
out['c_crossattn'] = comfy.conds.CONDRegular(cross_attn)
ref_latents = kwargs.get("reference_latents", None)
if ref_latents is not None:
out['ref_latents'] = comfy.conds.CONDList([self.process_latent_in(lat) for lat in ref_latents])
latents = []
for lat in ref_latents:
latents.append(self.process_latent_in(lat))
out['ref_latents'] = comfy.conds.CONDList(latents)
return out
def extra_conds_shapes(self, **kwargs):
@@ -2274,28 +2276,6 @@ class QwenImage(BaseModel):
out['ref_latents'] = list([1, 16, sum(map(lambda a: math.prod(a.size()), ref_latents)) // 16])
return out
class JoyImage(BaseModel):
def __init__(self, model_config, model_type=ModelType.FLOW, device=None):
super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.joyimage.model.JoyImageTransformer3DModel)
self.memory_usage_factor_conds = ("ref_latents",)
def extra_conds(self, **kwargs):
out = super().extra_conds(**kwargs)
cross_attn = kwargs.get("cross_attn", None)
if cross_attn is not None:
out['c_crossattn'] = comfy.conds.CONDRegular(cross_attn)
ref_latents = kwargs.get("reference_latents", None)
if ref_latents is not None:
out['ref_latents'] = comfy.conds.CONDList([self.process_latent_in(lat) for lat in ref_latents])
return out
def extra_conds_shapes(self, **kwargs):
out = {}
ref_latents = kwargs.get("reference_latents", None)
if ref_latents is not None:
out['ref_latents'] = list([1, 16, sum(map(lambda a: math.prod(a.size()), ref_latents)) // 16])
return out
class Ideogram4(BaseModel):
def __init__(self, model_config, model_type=ModelType.FLOW, device=None):
super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.ideogram4.model.Ideogram4Transformer2DModel)
@@ -2314,30 +2294,12 @@ class Ideogram4(BaseModel):
class Krea2(BaseModel):
def __init__(self, model_config, model_type=ModelType.FLUX, device=None):
super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.krea2.model.SingleStreamDiT)
self.memory_usage_factor_conds = ("ref_latents",)
def extra_conds(self, **kwargs):
out = super().extra_conds(**kwargs)
cross_attn = kwargs.get("cross_attn", None)
if cross_attn is not None:
out['c_crossattn'] = comfy.conds.CONDRegular(cross_attn)
ref_latents = kwargs.get("reference_latents", None)
if ref_latents is not None:
latents = []
for lat in ref_latents:
latents.append(self.process_latent_in(lat))
out['ref_latents'] = comfy.conds.CONDList(latents)
ref_latents_method = kwargs.get("reference_latents_method", None)
if ref_latents_method is not None:
out['ref_latents_method'] = comfy.conds.CONDConstant(ref_latents_method)
return out
def extra_conds_shapes(self, **kwargs):
out = {}
ref_latents = kwargs.get("reference_latents", None)
if ref_latents is not None:
out['ref_latents'] = list([1, 16, sum(map(lambda a: math.prod(a.size()), ref_latents)) // 16])
return out
class HunyuanImage21(BaseModel):
-19
View File
@@ -1058,25 +1058,6 @@ def detect_unet_config(state_dict, key_prefix, metadata=None):
dit_config["image_model"] = "SAM31"
return dit_config
if (
'{}double_blocks.0.attn.img_attn_qkv.weight'.format(key_prefix) in state_dict_keys
and '{}double_blocks.0.attn.img_attn_q_norm.weight'.format(key_prefix) in state_dict_keys
and '{}condition_embedder.time_embedder.linear_1.weight'.format(key_prefix) in state_dict_keys
and '{}img_in.weight'.format(key_prefix) in state_dict_keys
and len(state_dict['{}img_in.weight'.format(key_prefix)].shape) == 5
):
img_in = state_dict['{}img_in.weight'.format(key_prefix)]
head_dim = state_dict['{}double_blocks.0.attn.img_attn_q_norm.weight'.format(key_prefix)].shape[0]
return {
"image_model": "joyimage",
"in_channels": img_in.shape[1],
"hidden_size": img_in.shape[0],
"patch_size": list(img_in.shape[2:]),
"num_layers": count_blocks(state_dict_keys, '{}double_blocks.'.format(key_prefix) + '{}.'),
"num_attention_heads": img_in.shape[0] // head_dim,
"text_dim": 4096,
}
if '{}input_blocks.0.0.weight'.format(key_prefix) not in state_dict_keys:
return None
+1 -1
View File
@@ -473,7 +473,7 @@ except:
SUPPORT_FP8_OPS = args.supports_fp8_compute
AMD_RDNA2_AND_OLDER_ARCH = ["gfx1030", "gfx1031", "gfx1035", "gfx1010", "gfx1011", "gfx1012", "gfx906", "gfx900", "gfx803"]
AMD_RDNA2_AND_OLDER_ARCH = ["gfx1030", "gfx1031", "gfx1010", "gfx1011", "gfx1012", "gfx906", "gfx900", "gfx803"]
AMD_ENABLE_MIOPEN_ENV = 'COMFYUI_ENABLE_MIOPEN'
try:
+2 -13
View File
@@ -76,7 +76,6 @@ import comfy.text_encoders.gemma4
import comfy.text_encoders.cogvideo
import comfy.text_encoders.sa3
import comfy.text_encoders.gpt_oss
import comfy.text_encoders.joyimage
import comfy.model_patcher
import comfy.lora
@@ -1378,7 +1377,6 @@ class CLIPType(Enum):
IDEOGRAM4 = 30
BOOGU = 31
KREA2 = 32
JOYIMAGE = 33
@@ -1434,7 +1432,6 @@ class TEModel(Enum):
GPT_OSS_20B = 33
QWEN3VL_4B = 34
QWEN3VL_8B = 35
GEMMA_4_12B = 36
def detect_te_model(sd):
@@ -1464,9 +1461,6 @@ def detect_te_model(sd):
if 'model.layers.0.post_feedforward_layernorm.weight' in sd:
if 'model.layers.59.self_attn.q_norm.weight' in sd:
return TEModel.GEMMA_4_31B
# Gemma4 12B Unified: 48 layers, encoder-free; global layers drop v_proj (attention_k_eq_v).
if 'model.layers.47.self_attn.q_norm.weight' in sd and 'model.layers.5.self_attn.v_proj.weight' not in sd:
return TEModel.GEMMA_4_12B
if 'model.layers.41.self_attn.q_norm.weight' in sd and 'model.layers.47.self_attn.q_norm.weight' not in sd:
return TEModel.GEMMA_4_E4B
if 'model.layers.34.self_attn.q_norm.weight' in sd and 'model.layers.41.self_attn.q_norm.weight' not in sd:
@@ -1622,11 +1616,10 @@ def load_text_encoder_state_dicts(state_dicts=[], embedding_directory=None, clip
clip_target.clip = comfy.text_encoders.sa3.SAT5GemmaModel
clip_target.tokenizer = comfy.text_encoders.sa3.SAT5GemmaTokenizer
tokenizer_data["spiece_model"] = clip_data[0].get("spiece_model", None)
elif te_model in (TEModel.GEMMA_4_E4B, TEModel.GEMMA_4_E2B, TEModel.GEMMA_4_31B, TEModel.GEMMA_4_12B):
elif te_model in (TEModel.GEMMA_4_E4B, TEModel.GEMMA_4_E2B, TEModel.GEMMA_4_31B):
variant = {TEModel.GEMMA_4_E4B: comfy.text_encoders.gemma4.Gemma4_E4B,
TEModel.GEMMA_4_E2B: comfy.text_encoders.gemma4.Gemma4_E2B,
TEModel.GEMMA_4_31B: comfy.text_encoders.gemma4.Gemma4_31B,
TEModel.GEMMA_4_12B: comfy.text_encoders.gemma4.Gemma4_12B}[te_model]
TEModel.GEMMA_4_31B: comfy.text_encoders.gemma4.Gemma4_31B}[te_model]
clip_target.clip = comfy.text_encoders.gemma4.gemma4_te(**llama_detect(clip_data), model_class=variant)
clip_target.tokenizer = variant.tokenizer
tokenizer_data["tokenizer_json"] = clip_data[0].get("tokenizer_json", None)
@@ -1713,10 +1706,6 @@ def load_text_encoder_state_dicts(state_dicts=[], embedding_directory=None, clip
clip_data[0] = comfy.utils.state_dict_prefix_replace(clip_data[0], {"model.language_model.": "model.", "model.visual.": "visual.", "lm_head.": "model.lm_head."})
clip_target.clip = comfy.text_encoders.krea2.te(**llama_detect(clip_data))
clip_target.tokenizer = comfy.text_encoders.krea2.Krea2Tokenizer
elif clip_type == CLIPType.JOYIMAGE and te_model == TEModel.QWEN3VL_8B: # JoyImageEdit: full Qwen3-VL-8B, edit-conditioning template + drop_idx.
clip_data[0] = comfy.utils.state_dict_prefix_replace(clip_data[0], {"model.language_model.": "model.", "model.visual.": "visual.", "lm_head.": "model.lm_head."})
clip_target.clip = comfy.text_encoders.joyimage.te(**llama_detect(clip_data))
clip_target.tokenizer = comfy.text_encoders.joyimage.JoyImageTokenizer
elif clip_type in (CLIPType.FLUX, CLIPType.FLUX2): # Flux2 Klein reuses the Qwen3-VL LM (3-layer tap -> 12288); visual unused.
klein_model_type = "qwen3_8b" if te_model == TEModel.QWEN3VL_8B else "qwen3_4b"
clip_target.clip = comfy.text_encoders.flux.klein_te(**llama_detect(clip_data), model_type=klein_model_type)
-34
View File
@@ -27,7 +27,6 @@ import comfy.text_encoders.z_image
import comfy.text_encoders.ideogram4
import comfy.text_encoders.boogu
import comfy.text_encoders.krea2
import comfy.text_encoders.joyimage
import comfy.text_encoders.anima
import comfy.text_encoders.ace15
import comfy.text_encoders.longcat_image
@@ -1912,38 +1911,6 @@ class QwenImage(supported_models_base.BASE):
hunyuan_detect = comfy.text_encoders.hunyuan_video.llama_detect(state_dict, "{}qwen25_7b.transformer.".format(pref))
return supported_models_base.ClipTarget(comfy.text_encoders.qwen_image.QwenImageTokenizer, comfy.text_encoders.qwen_image.te(**hunyuan_detect))
class JoyImage(supported_models_base.BASE):
unet_config = {
"image_model": "joyimage",
}
sampling_settings = {
"multiplier": 1000,
"shift": 1.5,
}
memory_usage_factor = 1.8
unet_extra_config = {
"theta": 10000,
"rope_dim_list": [16, 56, 56],
}
latent_format = latent_formats.Wan21
supported_inference_dtypes = [torch.bfloat16, torch.float32]
vae_key_prefix = ["vae."]
text_encoder_key_prefix = ["text_encoders."]
def get_model(self, state_dict, prefix="", device=None):
return model_base.JoyImage(self, device=device)
def clip_target(self, state_dict={}):
pref = self.text_encoder_key_prefix[0]
qwen3vl_detect = comfy.text_encoders.hunyuan_video.llama_detect(state_dict, "{}qwen3vl.transformer.".format(pref))
return supported_models_base.ClipTarget(comfy.text_encoders.joyimage.JoyImageTokenizer, comfy.text_encoders.joyimage.te(**qwen3vl_detect))
class HunyuanImage21(HunyuanVideo):
unet_config = {
"image_model": "hunyuan_video",
@@ -2422,7 +2389,6 @@ models = [
Omnigen2,
Boogu,
QwenImage,
JoyImage,
Ideogram4,
Krea2,
Flux2,
+44 -267
View File
@@ -1,15 +1,11 @@
import torch
import torch.nn as nn
import torchaudio.functional as AF
import torchvision.transforms.functional as TVF
import numpy as np
from tokenizers import Tokenizer
from dataclasses import dataclass
import math
from comfy import sd1_clip
import comfy.model_management
import comfy.ops
from comfy.ldm.modules.attention import optimized_attention_for_device
from comfy.rmsnorm import rms_norm
from comfy.text_encoders.llama import RMSNorm, MLP, BaseLlama, BaseGenerate, _make_scaled_embedding
@@ -25,10 +21,6 @@ GEMMA4_VISION_CONFIG = {"hidden_size": 768, "image_size": 896, "intermediate_siz
GEMMA4_VISION_31B_CONFIG = {"hidden_size": 1152, "image_size": 896, "intermediate_size": 4304, "num_attention_heads": 16, "num_hidden_layers": 27, "patch_size": 16, "head_dim": 72, "rms_norm_eps": 1e-6, "position_embedding_size": 10240, "pooling_kernel_size": 3}
GEMMA4_AUDIO_CONFIG = {"hidden_size": 1024, "num_hidden_layers": 12, "num_attention_heads": 8, "intermediate_size": 4096, "conv_kernel_size": 5, "attention_chunk_size": 12, "attention_context_left": 13, "attention_context_right": 0, "attention_logit_cap": 50.0, "output_proj_dims": 1536, "rms_norm_eps": 1e-6, "residual_weight": 0.5}
# Encoder-free (gemma4_unified) multimodal embedders: raw patches/waveform projected directly into LM space.
GEMMA4_UNIFIED_VISION_CONFIG = {"model_patch_size": 48, "patch_size": 16, "pooling_kernel_size": 3, "mm_embed_dim": 3840, "mm_posemb_size": 1120, "output_proj_dims": 3840, "rms_norm_eps": 1e-6}
GEMMA4_UNIFIED_AUDIO_CONFIG = {"audio_samples_per_token": 640, "output_proj_dims": 640, "rms_norm_eps": 1e-6}
@dataclass
class Gemma4Config:
vocab_size: int = 262144
@@ -43,9 +35,6 @@ class Gemma4Config:
transformer_type: str = "gemma4"
head_dim = 256
global_head_dim = 512
num_global_key_value_heads = None
attention_k_eq_v = False
vision_bidirectional = False
rms_norm_add = False
mlp_activation = "gelu_pytorch_tanh"
qkv_bias = False
@@ -62,7 +51,6 @@ class Gemma4Config:
num_kv_shared_layers: int = 18
use_double_wide_mlp: bool = False
stop_tokens = [1, 50, 106]
suppress_tokens = []
vision_config = GEMMA4_VISION_CONFIG
audio_config = GEMMA4_AUDIO_CONFIG
mm_tokens_per_image = 280
@@ -84,30 +72,12 @@ class Gemma4_31B_Config(Gemma4Config):
num_hidden_layers: int = 60
num_attention_heads: int = 32
num_key_value_heads: int = 16
vision_bidirectional = True
sliding_attention = [1024, 1024, 1024, 1024, 1024, False]
hidden_size_per_layer_input: int = 0
num_kv_shared_layers: int = 0
audio_config = None
vision_config = GEMMA4_VISION_31B_CONFIG
@dataclass
class Gemma4_12B_Config(Gemma4Config):
hidden_size: int = 3840
intermediate_size: int = 15360
num_hidden_layers: int = 48
num_attention_heads: int = 16
num_key_value_heads: int = 8
num_global_key_value_heads = 1
attention_k_eq_v = True
vision_bidirectional = True
sliding_attention = [1024, 1024, 1024, 1024, 1024, False]
hidden_size_per_layer_input: int = 0
num_kv_shared_layers: int = 0
audio_config = GEMMA4_UNIFIED_AUDIO_CONFIG
vision_config = GEMMA4_UNIFIED_VISION_CONFIG
suppress_tokens = [258883, 258882]
# unfused RoPE as addcmul_ RoPE diverges from reference code
def _apply_rotary_pos_emb(x, freqs_cis):
@@ -119,18 +89,17 @@ def _apply_rotary_pos_emb(x, freqs_cis):
return out
class Gemma4Attention(nn.Module):
def __init__(self, config, head_dim, num_kv_heads=None, k_eq_v=False, device=None, dtype=None, ops=None):
def __init__(self, config, head_dim, device=None, dtype=None, ops=None):
super().__init__()
self.num_heads = config.num_attention_heads
self.num_kv_heads = num_kv_heads if num_kv_heads is not None else config.num_key_value_heads
self.num_kv_heads = config.num_key_value_heads
self.hidden_size = config.hidden_size
self.head_dim = head_dim
self.inner_size = self.num_heads * head_dim
self.q_proj = ops.Linear(config.hidden_size, self.inner_size, bias=config.qkv_bias, device=device, dtype=dtype)
self.k_proj = ops.Linear(config.hidden_size, self.num_kv_heads * head_dim, bias=config.qkv_bias, device=device, dtype=dtype)
# k_eq_v: V reuses the K projection (no separate v_proj weight)
self.v_proj = None if k_eq_v else ops.Linear(config.hidden_size, self.num_kv_heads * head_dim, bias=config.qkv_bias, device=device, dtype=dtype)
self.v_proj = ops.Linear(config.hidden_size, self.num_kv_heads * head_dim, bias=config.qkv_bias, device=device, dtype=dtype)
self.o_proj = ops.Linear(self.inner_size, config.hidden_size, bias=False, device=device, dtype=dtype)
self.q_norm = None
@@ -164,10 +133,7 @@ class Gemma4Attention(nn.Module):
shareable_kv = None
else:
xk = self.k_proj(hidden_states).view(batch_size, seq_length, self.num_kv_heads, self.head_dim)
if self.v_proj is not None:
xv = self.v_proj(hidden_states).view(batch_size, seq_length, self.num_kv_heads, self.head_dim)
else:
xv = xk # k_eq_v: V is the raw K projection (before k_norm/RoPE)
xv = self.v_proj(hidden_states).view(batch_size, seq_length, self.num_kv_heads, self.head_dim)
if self.k_norm is not None:
xk = self.k_norm(xk)
xv = rms_norm(xv)
@@ -220,10 +186,7 @@ class TransformerBlockGemma4(nn.Module):
head_dim = config.head_dim if self.sliding_attention else config.global_head_dim
# k_eq_v only on global layers, which then use num_global_key_value_heads
k_eq_v = config.attention_k_eq_v and not self.sliding_attention
num_kv_heads = config.num_global_key_value_heads if k_eq_v else config.num_key_value_heads
self.self_attn = Gemma4Attention(config, head_dim=head_dim, num_kv_heads=num_kv_heads, k_eq_v=k_eq_v, device=device, dtype=dtype, ops=ops)
self.self_attn = Gemma4Attention(config, head_dim=head_dim, device=device, dtype=dtype, ops=ops)
num_kv_shared = config.num_kv_shared_layers
first_kv_shared = config.num_hidden_layers - num_kv_shared
@@ -240,9 +203,9 @@ class TransformerBlockGemma4(nn.Module):
self.per_layer_input_gate = ops.Linear(config.hidden_size, self.hidden_size_per_layer_input, bias=False, device=device, dtype=dtype)
self.per_layer_projection = ops.Linear(self.hidden_size_per_layer_input, config.hidden_size, bias=False, device=device, dtype=dtype)
self.post_per_layer_input_norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps, device=device, dtype=dtype)
# layer_scalar exists on every gemma4 variant, independent of per-layer input
self.register_buffer("layer_scalar", torch.empty(1, device=device, dtype=dtype))
self.register_buffer("layer_scalar", torch.ones(1, device=device, dtype=dtype))
else:
self.layer_scalar = None
def forward(self, x, attention_mask=None, freqs_cis=None, past_key_value=None, per_layer_input=None, shared_kv=None):
sliding_window = None
@@ -281,7 +244,8 @@ class TransformerBlockGemma4(nn.Module):
x = self.post_per_layer_input_norm(x)
x = residual + x
x = x * comfy.ops.cast_to_input(self.layer_scalar, x)
if self.layer_scalar is not None:
x = x * self.layer_scalar
return x, present_key_value, shareable_kv
@@ -370,19 +334,6 @@ class Gemma4Transformer(nn.Module):
causal_mask.masked_fill_(torch.ones_like(causal_mask, dtype=torch.bool).triu_(1), min_val)
mask = mask + causal_mask if mask is not None else causal_mask
# Bidirectional attention within each image soft-token block (prefill only; text/audio stay causal).
if self.config.vision_bidirectional and past_len == 0 and embeds_info:
block_ids = torch.full((seq_len,), -1, dtype=torch.long, device=x.device)
group = 0
for info in embeds_info:
if info.get("type") == "image":
start = info["index"]
block_ids[start:start + info["size"]] = group
group += 1
if group > 0:
same_block = (block_ids[:, None] == block_ids[None, :]) & (block_ids[:, None] >= 0)
mask = mask.masked_fill(same_block, 0.0)
# Per-layer inputs
per_layer_inputs = None
if self.hidden_size_per_layer_input:
@@ -403,24 +354,8 @@ class Gemma4Transformer(nn.Module):
shared_global_kv = None # KV from last non-shared global layer
intermediate = None
all_intermediate = None
only_layers = None
if intermediate_output is not None:
if isinstance(intermediate_output, list):
all_intermediate = []
only_layers = {len(self.layers) + layer if layer < 0 else layer for layer in intermediate_output}
elif intermediate_output == "all":
all_intermediate = []
intermediate_output = None
elif intermediate_output < 0:
intermediate_output = len(self.layers) + intermediate_output
next_key_values = []
for i, layer in enumerate(self.layers):
if all_intermediate is not None:
if only_layers is None or (i in only_layers):
all_intermediate.append(x.unsqueeze(1).clone())
past_kv = past_key_values[i] if past_key_values is not None and len(past_key_values) > 0 else None
layer_kwargs = {}
@@ -450,18 +385,7 @@ class Gemma4Transformer(nn.Module):
if self.norm is not None:
x = self.norm(x)
if all_intermediate is not None:
if only_layers is None or (len(self.layers) in only_layers):
all_intermediate.append(x.unsqueeze(1).clone())
if len(all_intermediate) > 0:
intermediate = torch.cat(all_intermediate, dim=1)
if intermediate is not None and final_layer_norm_intermediate and self.norm is not None:
intermediate = self.norm(intermediate)
# Only hand back the KV cache when caching was actually requested; SDClipModel reads
# outputs[2] as the pooled output.
if past_key_values is not None and len(next_key_values) > 0:
if len(next_key_values) > 0:
return x, intermediate, next_key_values
return x, intermediate
@@ -480,8 +404,6 @@ class Gemma4Base(BaseLlama, BaseGenerate, torch.nn.Module):
cap = self.model.config.final_logit_softcapping
if cap:
logits = cap * torch.tanh(logits / cap)
if self.model.config.suppress_tokens:
logits[..., self.model.config.suppress_tokens] = torch.finfo(logits.dtype).min
return logits
def init_kv_cache(self, batch, max_cache_len, device, execution_dtype):
@@ -519,28 +441,6 @@ class Gemma4AudioMixin:
return None, None
class Gemma4UnifiedBase(Gemma4Base):
"""Encoder-free multimodal Gemma4 (gemma4_unified, e.g. 12B): raw image patches and audio frames projected directly into LM space."""
def _init_model(self, config, dtype, device, operations):
self.num_layers = config.num_hidden_layers
self.model = Gemma4Transformer(config, device=device, dtype=dtype, ops=operations)
self.dtype = dtype
self.vision_model = Gemma4UnifiedVisionEmbedder(config.vision_config, device=device, dtype=dtype, ops=operations)
self.multi_modal_projector = Gemma4RMSNormProjector(config.vision_config["output_proj_dims"], config.hidden_size, dtype=dtype, device=device, ops=operations)
self.audio_projector = Gemma4RMSNormProjector(config.audio_config["output_proj_dims"], config.hidden_size, dtype=dtype, device=device, ops=operations)
def preprocess_embed(self, embed, device):
if embed["type"] == "image":
pixels = embed.pop("data").movedim(-1, 1).to(device, dtype=self.dtype) # [B, H, W, C] -> [B, C, H, W], [0,1]
patches, positions = self.vision_model.patchify(pixels)
vision_out = self.vision_model(patches, positions)
return self.multi_modal_projector(vision_out), None
if embed["type"] == "audio":
audio = embed.pop("data").to(device, dtype=self.dtype) # [1, T, audio_samples_per_token]
return self.audio_projector(audio), None
return None, None
# Vision Encoder
def _compute_vision_2d_rope(head_dim, pixel_position_ids, theta=100.0, device=None):
@@ -813,73 +713,6 @@ class Gemma4MultiModalProjector(Gemma4RMSNormProjector):
super().__init__(config.vision_config["hidden_size"], config.hidden_size, dtype=dtype, device=device, ops=ops)
# Encoder-free vision (gemma4_unified): raw merged pixel patches projected directly into LM space.
def _patches_merge(patches, positions_xy, length):
patch_size = math.isqrt(patches.shape[-1] // 3)
k = math.isqrt(patches.shape[-2] // length)
batch = patches.shape[:-2]
max_x = positions_xy[..., 0].max(dim=-1, keepdim=True)[0] + 1
kidx = torch.div(positions_xy, k, rounding_mode="floor")
rem = torch.remainder(positions_xy, k)
order = rem[..., 0] + rem[..., 1] * k + k * k * kidx[..., 0] + k * max_x * kidx[..., 1]
perm = order.long().argsort(dim=-1)
merged = patches.gather(-2, perm.unsqueeze(-1).expand_as(patches))
merged = merged.reshape(*batch, length, k, k, patch_size, patch_size, 3)
merged = merged.permute(*range(len(batch)), -6, -5, -3, -4, -2, -1).reshape(*batch, length, (k * patch_size) ** 2 * 3)
pos = positions_xy.gather(-2, perm.unsqueeze(-1).expand_as(positions_xy))
pad = (positions_xy == -1).all(dim=-1, keepdim=True)
pos = torch.where(pad, positions_xy, pos).reshape(*batch, length, k * k, 2)
pos = torch.div(pos, k, rounding_mode="floor").min(dim=-2)[0]
return merged, pos
class Gemma4UnifiedVisionEmbedder(nn.Module):
"""Encoder-free patch embedder (LN -> Dense -> LN -> +2D posemb -> LN); projection to text space is the separate multi_modal_projector."""
def __init__(self, config, device=None, dtype=None, ops=None):
super().__init__()
self.patch_size = config["patch_size"]
self.pooling_kernel_size = config["pooling_kernel_size"]
patch_dim = config["model_patch_size"] ** 2 * 3
mm_embed_dim = config["mm_embed_dim"]
self.patch_ln1 = ops.LayerNorm(patch_dim, device=device, dtype=dtype)
self.patch_dense = ops.Linear(patch_dim, mm_embed_dim, device=device, dtype=dtype)
self.patch_ln2 = ops.LayerNorm(mm_embed_dim, device=device, dtype=dtype)
self.pos_embedding = nn.Parameter(torch.empty(config["mm_posemb_size"], 2, mm_embed_dim, device=device, dtype=dtype))
self.pos_norm = ops.LayerNorm(mm_embed_dim, device=device, dtype=dtype)
def patchify(self, pixels):
"""pixels: [B, C, H, W] in [0,1] -> merged patches [B, N, 6912], positions [B, N, 2]."""
ps, k = self.patch_size, self.pooling_kernel_size
out_patches, out_positions = [], []
for img in pixels:
ph, pw = img.shape[-2] // ps, img.shape[-1] // ps
teacher = img.reshape(img.shape[0], ph, ps, pw, ps).permute(1, 3, 2, 4, 0).reshape(ph * pw, -1)
grid = torch.meshgrid(torch.arange(pw, device=img.device), torch.arange(ph, device=img.device), indexing="xy")
tpos = torch.stack(grid, dim=-1).reshape(teacher.shape[0], 2)
n_model = teacher.shape[0] // (k * k)
mp, mpos = _patches_merge(teacher.unsqueeze(0), tpos.unsqueeze(0), n_model)
out_patches.append(mp.squeeze(0))
out_positions.append(mpos.squeeze(0))
return torch.stack(out_patches), torch.stack(out_positions)
def forward(self, pixel_values, image_position_ids):
x = self.patch_ln1(pixel_values)
x = self.patch_dense(x)
x = self.patch_ln2(x)
clamped = image_position_ids.clamp(min=0).long()
valid = (image_position_ids != -1).to(x.dtype).unsqueeze(-1)
axes = torch.arange(2, device=image_position_ids.device)
pos = comfy.model_management.cast_to_device(self.pos_embedding, x.device, x.dtype)
pos_embs = (pos[clamped, axes] * valid).sum(-2)
x = x + pos_embs
return self.pos_norm(x)
# Audio Encoder
class Gemma4AudioConvSubsampler(nn.Module):
@@ -1157,30 +990,6 @@ class Gemma4AudioProjector(Gemma4RMSNormProjector):
# Tokenizer and Wrappers
def _get_aspect_ratio_preserving_size(height, width, patch_size, max_patches, pooling_kernel_size):
target_px = max_patches * patch_size ** 2
factor = math.sqrt(target_px / (height * width))
side_mult = pooling_kernel_size * patch_size
target_height = math.floor(factor * height / side_mult) * side_mult
target_width = math.floor(factor * width / side_mult) * side_mult
if target_height == 0 and target_width == 0:
raise ValueError(f"Attempting to resize to a 0 x 0 image. Resized height should be divisible by {side_mult}.")
max_side_length = (max_patches // pooling_kernel_size ** 2) * side_mult
if target_height == 0:
target_height = side_mult
target_width = min(math.floor(width / height) * side_mult, max_side_length)
elif target_width == 0:
target_width = side_mult
target_height = min(math.floor(height / width) * side_mult, max_side_length)
if target_height * target_width > target_px:
raise ValueError(f"Resizing [{height}x{width}] to [{target_height}x{target_width}] exceeds the patch budget.")
return target_height, target_width
class Gemma4_Tokenizer():
tokenizer_json_data = None
@@ -1189,35 +998,25 @@ class Gemma4_Tokenizer():
return {"tokenizer_json": self.tokenizer_json_data}
return {}
def _audio_token_count(self, num_samples):
# Default (E2B/E4B): mel frames after two stride-2 conv subsamples.
_fl = 320 # int(round(16000 * 20.0 / 1000.0))
_hl = 160 # int(round(16000 * 10.0 / 1000.0))
_nmel = (num_samples + _fl // 2 - (_fl + 1)) // _hl + 1
_t = _nmel
for _ in range(2):
_t = (_t + 2 - 3) // 2 + 1
return min(_t, 750)
@staticmethod
def _resample_16k(waveform, sample_rate):
"""Mix to mono and resample to 16kHz. Kaiser params reproduce the reference (transformers
load_audio -> librosa/soxr_hq) to ~1e-12 MSE using only torchaudio."""
def _extract_mel_spectrogram(self, waveform, sample_rate):
"""Extract 128-bin log mel spectrogram.
Uses numpy for FFT/matmul/log to produce bit-identical results with reference code.
"""
# Mix to mono first, then resample to 16kHz
if waveform.dim() > 1 and waveform.shape[0] > 1:
waveform = waveform.mean(dim=0, keepdim=True)
if waveform.dim() == 1:
waveform = waveform.unsqueeze(0)
audio = waveform.float()
audio = waveform.squeeze(0).float().numpy()
if sample_rate != 16000:
audio = AF.resample(audio, sample_rate, 16000, resampling_method="sinc_interp_kaiser",
lowpass_filter_width=121, rolloff=0.9568384289091556, beta=21.01531462440614)
return audio.squeeze(0).contiguous()
def _extract_audio_features(self, waveform, sample_rate):
"""Default (E2B/E4B): 128-bin log mel spectrogram for the conformer audio encoder.
Uses numpy for FFT/matmul/log to produce bit-identical results with reference code.
"""
audio = self._resample_16k(waveform, sample_rate).numpy()
# Use scipy's resample_poly with a high-quality FIR filter to get as close as possible to librosa's resampling (while still not full match)
from scipy.signal import resample_poly, firwin
from math import gcd
g = gcd(sample_rate, 16000)
up, down = 16000 // g, sample_rate // g
L = max(up, down)
h = firwin(160 * L + 1, 0.96 / L, window=('kaiser', 6.5))
audio = resample_poly(audio, up, down, window=h).astype(np.float32)
n = len(audio)
# Pad to multiple of 128, build sample-level mask
@@ -1265,8 +1064,8 @@ class Gemma4_Tokenizer():
if audio is not None:
waveform = audio["waveform"].squeeze(0) if hasattr(audio, "__getitem__") else audio
sample_rate = audio.get("sample_rate", 16000) if hasattr(audio, "get") else 16000
feat, feat_mask = self._extract_audio_features(waveform, sample_rate)
audio_features = [(feat.unsqueeze(0), feat_mask.unsqueeze(0))] # ([1, T, D], [1, T])
mel, mel_mask = self._extract_mel_spectrogram(waveform, sample_rate)
audio_features = [(mel.unsqueeze(0), mel_mask.unsqueeze(0))] # ([1, T, 128], [1, T])
# Process image/video frames
is_video = video is not None
@@ -1291,8 +1090,13 @@ class Gemma4_Tokenizer():
pooling_k = 3
max_soft_tokens = kwargs.get("max_soft_tokens", 70 if is_video else 280)
max_patches = max_soft_tokens * pooling_k * pooling_k
target_h, target_w = _get_aspect_ratio_preserving_size(h, w, patch_size, max_patches, pooling_k)
target_px = max_patches * patch_size * patch_size
factor = (target_px / (h * w)) ** 0.5
side_mult = pooling_k * patch_size
target_h = max(int(factor * h // side_mult) * side_mult, side_mult)
target_w = max(int(factor * w // side_mult) * side_mult, side_mult)
import torchvision.transforms.functional as TVF
for i in range(num_frames):
# rescaling to match reference code
s = (samples[i].clamp(0, 1) * 255).to(torch.uint8) # [C, H, W] uint8
@@ -1311,7 +1115,7 @@ class Gemma4_Tokenizer():
llama_text = llama_template.format(text)
else:
# Build template from modalities present
system = "<|turn>system\n<|think|>\n<turn|>\n" if thinking else ""
system = "<|turn>system\n<|think|><turn|>\n" if thinking else ""
media = ""
if len(images) > 0:
if is_video:
@@ -1331,11 +1135,15 @@ class Gemma4_Tokenizer():
if len(audio_features) > 0:
# Compute audio token count (always at 16kHz)
num_samples = int(waveform.shape[-1] * 16000 / sample_rate) if sample_rate != 16000 else waveform.shape[-1]
n_audio_tokens = self._audio_token_count(num_samples)
_fl = 320 # int(round(16000 * 20.0 / 1000.0))
_hl = 160 # int(round(16000 * 10.0 / 1000.0))
_nmel = (num_samples + _fl // 2 - (_fl + 1)) // _hl + 1
_t = _nmel
for _ in range(2):
_t = (_t + 2 - 3) // 2 + 1
n_audio_tokens = min(_t, 750)
media += "<|audio>" + "<|audio|>" * n_audio_tokens + "<audio|>"
# Non-thinking mode primes an empty thought channel so the model answers directly.
model_open = "" if thinking else "<|channel>thought\n<channel|>"
llama_text = f"{system}<|turn>user\n{text}{media}<turn|>\n<|turn>model\n{model_open}"
llama_text = f"{system}<|turn>user\n{media}{text}<turn|>\n<|turn>model\n"
text_tokens = super().tokenize_with_weights(llama_text, return_word_ids)
@@ -1370,6 +1178,7 @@ class Gemma4_Tokenizer():
class _Gemma4Tokenizer:
"""Tokenizer using the tokenizers (Gemma4 doesn't come with sentencepiece model)"""
def __init__(self, tokenizer_json_bytes=None, **kwargs):
from tokenizers import Tokenizer
if isinstance(tokenizer_json_bytes, torch.Tensor):
tokenizer_json_bytes = bytes(tokenizer_json_bytes.tolist())
self.tokenizer = Tokenizer.from_str(tokenizer_json_bytes.decode("utf-8"))
@@ -1415,30 +1224,6 @@ class Gemma4Tokenizer(sd1_clip.SD1Tokenizer):
super().__init__(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data, name="gemma4", tokenizer=self.tokenizer_class)
class Gemma4UnifiedSDTokenizer(Gemma4SDTokenizer):
"""Encoder-free (gemma4_unified) audio: raw 16kHz waveform frames instead of mel spectrogram."""
embedding_size = 3840
def _extract_audio_features(self, waveform, sample_rate):
audio = self._resample_16k(waveform, sample_rate)
spt = 640 # audio_samples_per_token (40ms at 16kHz)
pad = (-audio.shape[0]) % spt
if pad:
audio = torch.nn.functional.pad(audio, (0, pad))
num_tokens = audio.shape[0] // spt
feats = audio[:num_tokens * spt].reshape(num_tokens, spt)
feats = feats[:750] # audio_seq_length cap (matches reference truncation, ~30s)
mask = torch.ones(feats.shape[0], dtype=torch.bool)
return feats, mask
def _audio_token_count(self, num_samples):
return min((num_samples + 639) // 640, 750)
class Gemma4UnifiedTokenizer(Gemma4Tokenizer):
tokenizer_class = Gemma4UnifiedSDTokenizer
# Model wrappers
class Gemma4Model(sd1_clip.SDClipModel):
model_class = None
@@ -1471,7 +1256,7 @@ class Gemma4Model(sd1_clip.SDClipModel):
expanded_idx += 1
initial_token_ids = [ids]
input_ids = torch.tensor(initial_token_ids, device=self.execution_device)
return self.transformer.generate(embeds, do_sample, max_length, temperature, top_k, top_p, min_p, repetition_penalty, seed, initial_tokens=initial_token_ids[0], presence_penalty=presence_penalty, initial_input_ids=input_ids, embeds_info=embeds_info)
return self.transformer.generate(embeds, do_sample, max_length, temperature, top_k, top_p, min_p, repetition_penalty, seed, initial_tokens=initial_token_ids[0], presence_penalty=presence_penalty, initial_input_ids=input_ids)
def gemma4_te(dtype_llama=None, llama_quantization_metadata=None, model_class=None):
@@ -1511,11 +1296,3 @@ def _make_variant(config_cls):
Gemma4_E4B = _make_variant(Gemma4Config)
Gemma4_E2B = _make_variant(Gemma4_E2B_Config)
Gemma4_31B = _make_variant(Gemma4_31B_Config)
# Gemma4 12B Unified: encoder-free multimodal, distinct base/tokenizer (not via _make_variant).
class Gemma4_12B(Gemma4UnifiedBase):
def __init__(self, config_dict, dtype, device, operations):
super().__init__()
self._init_model(Gemma4_12B_Config(**config_dict), dtype, device, operations)
Gemma4_12B.tokenizer = Gemma4UnifiedTokenizer
-97
View File
@@ -1,97 +0,0 @@
import torch
from comfy import sd1_clip
import comfy.text_encoders.qwen_vl
from comfy.text_encoders.qwen3vl import Qwen3VL, Qwen3VLTokenizer
JOYIMAGE_VISION_BLOCK = "<|vision_start|><|image_pad|><|vision_end|>"
JOYIMAGE_TEMPLATE_TEXT = (
"<|im_start|>system\n \\nDescribe the image by detailing the color, shape, size, texture, "
"quantity, text, spatial relationships of the objects and background:<|im_end|>\n"
"<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n"
)
JOYIMAGE_TEMPLATE_IMAGE = (
"<|im_start|>system\n \\nDescribe the image by detailing the color, shape, size, texture, "
"quantity, text, spatial relationships of the objects and background:<|im_end|>\n"
f"<|im_start|>user\n{JOYIMAGE_VISION_BLOCK}{{}}<|im_end|>\n<|im_start|>assistant\n"
)
# The DiT was trained without the leading system-prompt tokens.
JOYIMAGE_DROP_IDX = 34
PAD_TOKEN = 151643
class Qwen3VL8B_JoyImage(Qwen3VL):
model_type = "qwen3vl_8b"
def preprocess_embed(self, embed, device):
if embed["type"] == "image":
image, grid = comfy.text_encoders.qwen_vl.process_qwen2vl_images(
embed["data"], min_pixels=65536, max_pixels=16777216, patch_size=16,
image_mean=[0.5, 0.5, 0.5], image_std=[0.5, 0.5, 0.5],
interpolation="bicubic",
)
merged, deepstack = self.visual(image.to(device, dtype=torch.float32), grid)
return merged, {"grid": grid, "deepstack": deepstack}
return None, None
class JoyImageTokenizer(Qwen3VLTokenizer):
def __init__(self, embedding_directory=None, tokenizer_data={}):
super().__init__(
embedding_directory=embedding_directory, tokenizer_data=tokenizer_data,
model_type="qwen3vl_8b",
)
self.llama_template = JOYIMAGE_TEMPLATE_TEXT
self.llama_template_images = JOYIMAGE_TEMPLATE_IMAGE
def tokenize_with_weights(self, text, return_word_ids=False, llama_template=None, images=None, **kwargs):
kwargs.pop("thinking", None)
return super().tokenize_with_weights(
text, return_word_ids=return_word_ids, llama_template=llama_template,
images=images or [], thinking=True, **kwargs,
)
class _JoyImageClipModel(sd1_clip.SDClipModel):
def __init__(self, device="cpu", layer="hidden", layer_idx=-1, dtype=None,
attention_mask=True, model_options={}):
super().__init__(
device=device, layer=layer, layer_idx=layer_idx, textmodel_json_config={},
# JoyImage conditions on the pre-final-norm output of the last decoder layer.
dtype=dtype, special_tokens={"pad": PAD_TOKEN}, layer_norm_hidden_state=False,
model_class=Qwen3VL8B_JoyImage, enable_attention_masks=attention_mask,
return_attention_masks=attention_mask, model_options=model_options,
)
class JoyImageTEModel(sd1_clip.SD1ClipModel):
def __init__(self, device="cpu", dtype=None, model_options={}):
super().__init__(
device=device, dtype=dtype, name="qwen3vl_8b",
clip_model=_JoyImageClipModel, model_options=model_options,
)
def encode_token_weights(self, token_weight_pairs):
out, pooled, extra = super().encode_token_weights(token_weight_pairs)
if out.shape[1] <= JOYIMAGE_DROP_IDX:
raise ValueError(
f"JoyImageTEModel: encoded sequence length {out.shape[1]} is shorter "
f"than drop_idx={JOYIMAGE_DROP_IDX}; the prompt did not include the "
f"template prefix."
)
out = out[:, JOYIMAGE_DROP_IDX:]
if "attention_mask" in extra:
extra["attention_mask"] = extra["attention_mask"][:, JOYIMAGE_DROP_IDX:]
return out, pooled, extra
def te(dtype_llama=None, llama_quantization_metadata=None):
class JoyImageTEModel_(JoyImageTEModel):
def __init__(self, device="cpu", dtype=None, model_options={}):
if llama_quantization_metadata is not None:
model_options = model_options.copy()
model_options["quantization_metadata"] = llama_quantization_metadata
if dtype_llama is not None:
dtype = dtype_llama
super().__init__(device=device, dtype=dtype, model_options=model_options)
return JoyImageTEModel_
+2 -2
View File
@@ -876,7 +876,7 @@ class BaseGenerate:
torch.empty([batch, model_config.num_key_value_heads, max_cache_len, model_config.head_dim], device=device, dtype=execution_dtype), 0))
return past_key_values
def generate(self, embeds=None, do_sample=True, max_length=256, temperature=1.0, top_k=50, top_p=0.9, min_p=0.0, repetition_penalty=1.0, seed=42, stop_tokens=None, initial_tokens=[], execution_dtype=None, min_tokens=0, presence_penalty=0.0, initial_input_ids=None, position_ids=None, deepstack_embeds=None, visual_pos_masks=None, embeds_info=None):
def generate(self, embeds=None, do_sample=True, max_length=256, temperature=1.0, top_k=50, top_p=0.9, min_p=0.0, repetition_penalty=1.0, seed=42, stop_tokens=None, initial_tokens=[], execution_dtype=None, min_tokens=0, presence_penalty=0.0, initial_input_ids=None, position_ids=None, deepstack_embeds=None, visual_pos_masks=None):
device = embeds.device
if stop_tokens is None:
@@ -911,7 +911,7 @@ class BaseGenerate:
if step == 0 and deepstack_embeds is not None:
extra["deepstack_embeds"] = deepstack_embeds
extra["visual_pos_masks"] = visual_pos_masks
x, _, past_key_values = self.model.forward(None, embeds=embeds, attention_mask=None, past_key_values=past_key_values, input_ids=current_input_ids, position_ids=position_ids, **extra, embeds_info=(embeds_info if step == 0 else None))
x, _, past_key_values = self.model.forward(None, embeds=embeds, attention_mask=None, past_key_values=past_key_values, input_ids=current_input_ids, position_ids=position_ids, **extra)
logits = self.logits(x)[:, -1]
next_token = self.sample_token(logits, temperature, top_k, top_p, min_p, repetition_penalty, initial_tokens + generated_token_ids, generator, do_sample=do_sample, presence_penalty=presence_penalty)
token_id = next_token[0].item()
+2 -2
View File
@@ -15,7 +15,6 @@ def process_qwen2vl_images(
merge_size: int = 2,
image_mean: list = None,
image_std: list = None,
interpolation: str = "bilinear",
):
if image_mean is None:
image_mean = [0.48145466, 0.4578275, 0.40821073]
@@ -48,9 +47,10 @@ def process_qwen2vl_images(
img_resized = F.interpolate(
img.unsqueeze(0),
size=(h_bar, w_bar),
mode=interpolation,
mode='bilinear',
align_corners=False
).squeeze(0)
normalized = img_resized.clone()
for c in range(3):
normalized[c] = (img_resized[c] - image_mean[c]) / image_std[c]
+1
View File
@@ -105,6 +105,7 @@ _CORE_FEATURE_FLAGS: dict[str, Any] = {
"extension": {"manager": {"supports_v4": True}},
"node_replacements": True,
"assets": args.enable_assets,
"server_side_model_downloads": True,
}
# CLI-provided flags cannot overwrite core flags
+51 -343
View File
@@ -1,6 +1,5 @@
from av.container import InputContainer
from av.subtitles.stream import SubtitleStream
from av.video.reformatter import ColorRange
from fractions import Fraction
from typing import Optional
from .._input import AudioInput, VideoInput
@@ -10,7 +9,6 @@ import itertools
import json
import numpy as np
import math
import os
import torch
from .._util import VideoContainer, VideoCodec, VideoComponents
import logging
@@ -60,57 +58,6 @@ def video_stream_bit_depth(stream) -> int:
return max(component.bits for component in stream.format.components)
def last_decodable_audio_stream(container: InputContainer):
"""Streams FFmpeg has no decoder for have no codec context, and decoding their
packets crashes the process (e.g. APAC spatial-audio track in iPhone)."""
stream = next(
(s for s in reversed(container.streams.audio) if s.codec_context is not None),
None,
)
if stream is None and len(container.streams.audio):
logging.warning("No decodable audio stream found in video; ignoring audio.")
return stream
def probe_audio_params(container: InputContainer, audio_stream, max_packets: int = 200):
"""Containers probed only up to a window (mpegts) leave audio codec parameters unset when
audio starts beyond it; learn them by decoding ahead. The caller must seek back afterwards.
Returns (sample_rate, channels), zeros when the stream never yields a decodable frame."""
for i, packet in enumerate(container.demux(audio_stream)):
try:
frames = packet.decode()
except av.error.FFmpegError:
frames = ()
if frames:
return frames[0].sample_rate, frames[0].layout.nb_channels
if i >= max_packets:
break
return 0, 0
def write_output_metadata(container: InputContainer, output, metadata: dict | None):
"""Copy the source container's metadata, then overlay the caller's tags."""
for key, value in container.metadata.items():
if metadata is None or key not in metadata:
output.metadata[key] = value
if metadata is not None:
for key, value in metadata.items():
output.metadata[key] = value if isinstance(value, str) else json.dumps(value)
def mp4_output_open_kwargs(path: str | io.BytesIO, format: VideoContainer, codec: VideoCodec) -> dict:
if format != VideoContainer.AUTO and format != VideoContainer.MP4:
raise ValueError("Only MP4 format is supported for now")
if codec != VideoCodec.AUTO and codec != VideoCodec.H264:
raise ValueError("Only H264 codec is supported for now")
open_kwargs = {"mode": "w", "options": {"movflags": "use_metadata_tags"}}
if isinstance(format, VideoContainer) and format != VideoContainer.AUTO:
open_kwargs["format"] = format.value
elif isinstance(path, io.BytesIO):
open_kwargs["format"] = "mp4" # no file extension to infer the format from
return open_kwargs
class VideoFromFile(VideoInput):
"""
Class representing video input from a file.
@@ -245,10 +192,13 @@ class VideoFromFile(VideoInput):
return estimated_frames
# 3. Last resort: decode frames and count them (streaming)
start_time, duration = self.get_active_trim_window()
if self.__start_time < 0:
start_time = max(self._get_raw_duration() + self.__start_time, 0)
else:
start_time = self.__start_time
frame_count = 1
start_pts = int(start_time / video_stream.time_base)
end_pts = int((start_time + duration) / video_stream.time_base)
end_pts = int((start_time + self.__duration) / video_stream.time_base)
container.seek(start_pts, stream=video_stream)
frame_iterator = (
container.decode(video_stream)
@@ -303,14 +253,17 @@ class VideoFromFile(VideoInput):
def get_components_internal(self, container: InputContainer) -> VideoComponents:
video_stream = self._get_first_video_stream(container)
start_time, duration = self.get_active_trim_window()
if self.__start_time < 0:
start_time = max(self._get_raw_duration() + self.__start_time, 0)
else:
start_time = self.__start_time
# Get video frames
frames = []
audio_frames = []
alphas = None
start_pts = int(start_time / video_stream.time_base)
end_pts = int((start_time + duration) / video_stream.time_base)
end_pts = int((start_time + self.__duration) / video_stream.time_base)
if start_pts != 0:
container.seek(start_pts, stream=video_stream)
@@ -328,11 +281,18 @@ class VideoFromFile(VideoInput):
video_done = False
audio_done = True
audio_stream = last_decodable_audio_stream(container)
# Use the last decodable audio stream. Streams FFmpeg has no decoder for have no codec context,
# and decoding their packets crashes the process. (e.g. APAC spatial-audio track in iPhone)
audio_stream = next(
(s for s in reversed(container.streams.audio) if s.codec_context is not None),
None,
)
if audio_stream is not None:
streams += [audio_stream]
resampler = av.audio.resampler.AudioResampler(format='fltp')
audio_done = False
elif len(container.streams.audio):
logging.warning("No decodable audio stream found in video; ignoring audio.")
for packet in container.demux(*streams):
if video_done and audio_done:
@@ -345,7 +305,7 @@ class VideoFromFile(VideoInput):
for frame in packet.decode():
if frame.pts < start_pts:
continue
if duration and frame.pts >= end_pts:
if self.__duration and frame.pts >= end_pts:
video_done = True
break
@@ -412,7 +372,7 @@ class VideoFromFile(VideoInput):
map(resampler.resample, packet.decode())
)
for frame in aframes:
if duration and frame.time > start_time + duration:
if self.__duration and frame.time > start_time + self.__duration:
audio_done = True
break
@@ -434,8 +394,8 @@ class VideoFromFile(VideoInput):
if len(audio_frames) > 0:
audio_data = np.concatenate(audio_frames, axis=1) # shape: (channels, total_samples)
if duration:
audio_data = audio_data[..., :int(duration * audio_stream.sample_rate)]
if self.__duration:
audio_data = audio_data[..., :int(self.__duration * audio_stream.sample_rate)]
audio_tensor = torch.from_numpy(audio_data).unsqueeze(0) # shape: (1, channels, total_samples)
audio = AudioInput({
@@ -481,14 +441,28 @@ class VideoFromFile(VideoInput):
if not reuse_streams:
if bit_depth is None:
bit_depth = source_bit_depth
return self._save_transcoded(container, path, format=format, codec=codec, metadata=metadata, bit_depth=bit_depth)
components = self.get_components_internal(container)
video = VideoFromComponents(components)
return video.save_to(
path, format=format, codec=codec, metadata=metadata, bit_depth=bit_depth,
)
streams = container.streams
open_kwargs = get_open_write_kwargs(path, container_format, format)
with av.open(path, **open_kwargs) as output_container:
# Add metadata before writing any streams
write_output_metadata(container, output_container, metadata)
# Copy over the original metadata
for key, value in container.metadata.items():
if metadata is None or key not in metadata:
output_container.metadata[key] = value
# Add our new metadata
if metadata is not None:
for key, value in metadata.items():
if isinstance(value, str):
output_container.metadata[key] = value
else:
output_container.metadata[key] = json.dumps(value)
# Add streams to the new container. Streams with no codec context cannot be used as an output template.
stream_map = {}
@@ -506,282 +480,6 @@ class VideoFromFile(VideoInput):
packet.stream = stream_map[packet.stream]
output_container.mux(packet)
def _save_transcoded(
self,
container: InputContainer,
path: str | io.BytesIO,
format: VideoContainer,
codec: VideoCodec,
metadata: dict | None,
bit_depth: int,
):
"""Re-encode to H.264/AAC one frame at a time; peak memory does not scale with video length."""
open_kwargs = mp4_output_open_kwargs(path, format, codec)
video_stream = self._get_first_video_stream(container)
start_time, duration = self.get_active_trim_window()
start_pts = int(start_time / video_stream.time_base)
end_pts = int((start_time + duration) / video_stream.time_base) if duration else None
stream_end_pts = None
if video_stream.duration is not None:
stream_end_pts = (video_stream.start_time or 0) + video_stream.duration
output_end_pts = end_pts
if stream_end_pts is not None and (output_end_pts is None or stream_end_pts < output_end_pts):
output_end_pts = stream_end_pts
if start_pts != 0:
container.seek(start_pts, stream=video_stream)
audio_stream = last_decodable_audio_stream(container)
pix_fmt = "yuv420p10le" if bit_depth >= 10 else "yuv420p"
rate = Fraction(video_stream.average_rate) if video_stream.average_rate else Fraction(1)
resampler = None
sample_rate = 0
audio_time_base = None
duration_cap = None
if audio_stream is not None:
sample_rate = audio_stream.codec_context.sample_rate
channels = audio_stream.codec_context.channels
if not sample_rate:
sample_rate, channels = probe_audio_params(container, audio_stream)
container.seek(start_pts, stream=video_stream)
if sample_rate:
audio_stream.codec_context.flush_buffers()
else:
logging.warning("Audio stream parameters could not be determined; ignoring audio.")
audio_stream = None
if audio_stream is not None:
audio_time_base = Fraction(1, sample_rate)
layout = {1: "mono", 2: "stereo", 6: "5.1"}.get(channels, "stereo")
resampler = av.audio.resampler.AudioResampler(format="fltp", layout=layout, rate=sample_rate)
if duration:
duration_cap = math.ceil(duration * sample_rate)
streams = [video_stream] if audio_stream is None else [video_stream, audio_stream]
pts_step = max(1, int(round((1 / rate) / video_stream.time_base)))
video_done = False
audio_done = audio_stream is None
video_pts_offset = None
last_video_pts = None
last_video_end = None
# rebased pts -> true display duration: the mp4 muxer pads the last sample with 1/rate otherwise
video_frame_durations = {}
source_size = None
rotation_k = 0
rotation_filter = None
audio_started = False
samples_written = 0
pending_audio = []
# The output opens lazily on the first kept frame: it decides the geometry (90/270 rotation swaps dims),
# and never seeking back keeps webm/mkv leading audio intact.
output = None
out_video = None
out_audio = None
def audio_frame_from_ndarray(nd_planar):
frame = av.AudioFrame.from_ndarray(np.ascontiguousarray(nd_planar), format="fltp", layout=layout)
frame.sample_rate = sample_rate
return frame
def drain_audio(final=False):
# Audio may cover the pts span of the video written so far, capped by the requested duration
nonlocal samples_written, audio_done
if last_video_end is None:
cap = 0
else:
cap = math.ceil(last_video_end * video_stream.time_base * sample_rate)
if duration_cap is not None:
cap = min(cap, duration_cap)
while pending_audio and not audio_done:
frame = pending_audio[0]
if samples_written + frame.samples <= cap:
frame.pts = samples_written
frame.time_base = audio_time_base
output.mux(out_audio.encode(frame))
samples_written += frame.samples
pending_audio.pop(0)
continue
if final:
keep = frame.to_ndarray()[..., :cap - samples_written]
if keep.shape[-1] > 0:
tail = audio_frame_from_ndarray(keep)
tail.pts = samples_written
tail.time_base = audio_time_base
output.mux(out_audio.encode(tail))
samples_written += keep.shape[-1]
pending_audio.clear()
break
if duration_cap is not None and samples_written >= duration_cap:
audio_done = True
return cap
try:
for packet in container.demux(*streams):
if video_done and audio_done:
break
if packet.stream == video_stream and not video_done:
try:
frames = packet.decode()
except av.error.InvalidDataError:
logging.info("pyav decode error")
continue
for frame in frames:
if frame.pts is not None and frame.pts < start_pts:
continue
if end_pts is not None and frame.pts is not None and frame.pts >= end_pts:
video_done = True
if last_video_pts is not None:
# the source continues past the window: hold the last kept frame to the window end
end_offset = video_pts_offset if video_pts_offset is not None else start_pts
last_video_end = max(last_video_end, end_pts - end_offset)
break
# the source's true display duration of this frame; average_rate is not a
# frame duration (sparse/VFR sources), so it is only the fallback
frame_duration = frame.duration if frame.duration else pts_step
if end_pts is not None and frame.pts is not None:
frame_duration = min(frame_duration, end_pts - frame.pts)
if output is None:
rotation_k = int(round(frame.rotation // 90)) % 4 if frame.rotation else 0
if rotation_k % 2:
out_width, out_height = frame.height, frame.width
else:
out_width, out_height = frame.width, frame.height
if out_width % 2 or out_height % 2:
raise ValueError(f"H.264 output requires even dimensions, got {out_width}x{out_height}")
source_size = (frame.width, frame.height)
output = av.open(path, **open_kwargs)
# Add metadata before writing any streams
write_output_metadata(container, output, metadata)
out_video = output.add_stream("h264", rate=rate)
# no B-frames: reordering makes mp4 sample durations follow decode order,
# so irregular-VFR spans and trim windows land wrong
out_video.codec_context.max_b_frames = 0
out_video.width = out_width
out_video.height = out_height
out_video.pix_fmt = pix_fmt
# source pts pass through (rebased to 0), so variable frame rate survives
out_video.codec_context.time_base = video_stream.time_base
if audio_stream is not None:
out_audio = output.add_stream("aac", rate=sample_rate, layout=layout)
if (frame.width, frame.height) != source_size:
# encoding would silently rescale the new geometry into the old one
raise ValueError(
f"Video resolution changes mid-stream "
f"({source_size[0]}x{source_size[1]} -> {frame.width}x{frame.height}); cannot transcode"
)
if rotation_k:
if rotation_filter is None:
g = av.filter.Graph()
g_src = g.add_buffer(width=frame.width, height=frame.height,
format=frame.format.name, time_base=video_stream.time_base)
tail = g_src
for filter_name, filter_args in {1: [("transpose", "cclock")],
2: [("hflip", None), ("vflip", None)],
3: [("transpose", "clock")]}[rotation_k]:
step = g.add(filter_name, filter_args)
tail.link_to(step)
tail = step
g_sink = g.add("buffersink")
tail.link_to(g_sink)
g.configure()
rotation_filter = (g_src, g_sink)
rotation_filter[0].push(frame)
frame = rotation_filter[1].pull()
if frame.color_range == ColorRange.JPEG:
# compress full-range sources (yuvj/MJPEG) to limited range
frame = frame.reformat(format=pix_fmt, src_color_range="JPEG", dst_color_range="MPEG")
else:
frame = frame.reformat(format=pix_fmt)
frame_output_end = None
if frame.pts is not None:
if video_pts_offset is None:
video_pts_offset = frame.pts
frame.pts -= video_pts_offset
if output_end_pts is not None:
frame_output_end = output_end_pts - video_pts_offset
if frame.pts + frame_duration > frame_output_end:
clamped_pts = frame_output_end - frame_duration
if clamped_pts >= 0 and (last_video_pts is None or clamped_pts > last_video_pts):
frame.pts = min(frame.pts, clamped_pts)
elif frame.pts < frame_output_end:
frame_duration = frame_output_end - frame.pts
else:
continue
if frame.pts is None or (last_video_pts is not None and frame.pts <= last_video_pts):
# broken sources emit missing/backward timestamps mid-stream, which the
# muxer rejects; nudge them forward by one nominal frame interval
frame.pts = 0 if last_video_pts is None else last_video_pts + pts_step
if frame_output_end is not None and frame.pts + frame_duration > frame_output_end:
if frame.pts >= frame_output_end:
continue
frame_duration = frame_output_end - frame.pts
last_video_pts = frame.pts
last_video_end = frame.pts + frame_duration
video_frame_durations[frame.pts] = frame_duration
# the decoded pict_type would force x264's frame types (intra-only
# sources like MJPEG/ProRes would come out all-keyframe)
frame.pict_type = 0
for out_packet in out_video.encode(frame):
out_packet.duration = video_frame_durations.pop(out_packet.pts, 0)
output.mux(out_packet)
drain_audio()
elif packet.stream == audio_stream and not audio_done:
for resampled in itertools.chain.from_iterable(map(resampler.resample, packet.decode())):
frame_start = None
if resampled.pts is not None:
# passthrough frames keep the source stream's time base
tb = resampled.time_base if resampled.time_base else audio_time_base
frame_start = float(resampled.pts * tb)
if duration and not audio_started and frame_start >= start_time + duration:
audio_done = True
break
if not audio_started:
if frame_start is None:
frame_start = 0.0
to_skip = max(0, int((start_time - frame_start) * sample_rate))
if to_skip >= resampled.samples:
continue
audio_started = True
if duration and frame_start > start_time:
duration_cap = min(duration_cap, math.ceil((start_time + duration - frame_start) * sample_rate))
if to_skip:
pending_audio.append(audio_frame_from_ndarray(resampled.to_ndarray()[..., to_skip:]))
continue
pending_audio.append(resampled)
if video_done:
# the video window is complete so the cap is final, but containers
# that interleave audio behind video (fragmented mp4) still owe most
# of it: stop only once the demuxed audio covers the cap
cap = drain_audio()
if pending_audio or samples_written >= cap:
drain_audio(final=True)
audio_done = True
break
if output is None:
raise ValueError(f"No decodable video frames found in file '{self.__file}'")
if out_audio is not None and not audio_done:
drain_audio(final=True)
window_fill = last_video_end - last_video_pts if video_done and last_video_pts is not None else 0
for out_packet in out_video.encode(None):
duration = video_frame_durations.pop(out_packet.pts, 0)
if out_packet.pts == last_video_pts:
duration = max(duration, window_fill)
out_packet.duration = duration
output.mux(out_packet)
if out_audio is not None:
output.mux(out_audio.encode(None))
except BaseException:
if output is not None:
output.close()
if isinstance(path, (str, os.PathLike)) and os.path.exists(path):
os.remove(path)
raise
else:
if output is not None:
output.close()
def _get_first_video_stream(self, container: InputContainer):
if len(container.streams.video):
return container.streams.video[0]
@@ -829,12 +527,22 @@ class VideoFromComponents(VideoInput):
bit_depth: int | None = None,
):
"""Save the video to a file path or BytesIO buffer."""
open_kwargs = mp4_output_open_kwargs(path, format, codec)
if format != VideoContainer.AUTO and format != VideoContainer.MP4:
raise ValueError("Only MP4 format is supported for now")
if codec != VideoCodec.AUTO and codec != VideoCodec.H264:
raise ValueError("Only H264 codec is supported for now")
# None means "use the depth this video was created with" (CreateVideo's choice).
if bit_depth is None:
bit_depth = self.__bit_depth
is_10bit = bit_depth >= 10
with av.open(path, **open_kwargs) as output:
extra_kwargs = {}
if isinstance(format, VideoContainer) and format != VideoContainer.AUTO:
extra_kwargs["format"] = format.value
elif isinstance(path, io.BytesIO):
# BytesIO has no file extension, so av.open can't infer the format.
# Default to mp4 since that's the only supported format anyway.
extra_kwargs["format"] = "mp4"
with av.open(path, mode='w', options={'movflags': 'use_metadata_tags'}, **extra_kwargs) as output:
# Add metadata before writing any streams
if metadata is not None:
for key, value in metadata.items():
+1 -58
View File
@@ -1,6 +1,6 @@
from datetime import date
from enum import Enum
from typing import Any, Literal
from typing import Any
from pydantic import BaseModel, Field
@@ -242,60 +242,3 @@ class GeminiGenerateContentResponse(BaseModel):
promptFeedback: GeminiPromptFeedback | None = Field(None)
usageMetadata: GeminiUsageMetadata | None = Field(None)
modelVersion: str | None = Field(None)
class GeminiInteractionTextPart(BaseModel):
type: Literal["text"] = "text"
text: str = Field(...)
class GeminiInteractionMediaPart(BaseModel):
type: str = Field(..., description="One of: image, video, audio, document.")
data: str | None = Field(None, description="Base64-encoded media bytes.")
uri: str | None = Field(None, description="URI of the media, as an alternative to inline data.")
mime_type: str | None = Field(None)
class GeminiInteractionGenerationConfig(BaseModel):
temperature: float | None = Field(None, ge=0.0, le=2.0)
top_p: float | None = Field(None, ge=0.0, le=1.0)
class GeminiInteractionRequest(BaseModel):
model: str = Field(...)
input: list[GeminiInteractionTextPart | GeminiInteractionMediaPart] = Field(...)
generation_config: GeminiInteractionGenerationConfig | None = Field(None)
class GeminiInteractionModalityTokens(BaseModel):
modality: str | None = Field(None, description="One of: text, image, audio, video, document.")
tokens: int | None = Field(None)
class GeminiInteractionUsage(BaseModel):
input_tokens_by_modality: list[GeminiInteractionModalityTokens] | None = Field(None)
output_tokens_by_modality: list[GeminiInteractionModalityTokens] | None = Field(None)
total_thought_tokens: int | None = Field(None)
class GeminiInteractionContent(BaseModel):
type: str | None = Field(None)
text: str | None = Field(None)
data: str | None = Field(None)
uri: str | None = Field(None)
mime_type: str | None = Field(None)
class GeminiInteractionStep(BaseModel):
type: str | None = Field(None)
content: list[GeminiInteractionContent] | None = Field(None)
class GeminiInteraction(BaseModel):
id: str | None = Field(None)
status: str | None = Field(
None,
description="One of: in_progress, requires_action, completed, failed, cancelled, incomplete.",
)
steps: list[GeminiInteractionStep] | None = Field(None)
usage: GeminiInteractionUsage | None = Field(None)
-452
View File
@@ -1,452 +0,0 @@
# (label, avatar_id, avatar_type, supported engines)
HEYGEN_AVATAR_LOOKS: list[tuple[str, str, str, tuple[str, ...]]] = [
(
"Annie Lounge Standing Side",
"Annie_Lounge_Standing_Side_public",
"studio_avatar",
("avatar_v", "avatar_iv", "avatar_iii"),
),
(
"Yara Modern Lecture Hall",
"fd6814ecc5e143cd899e615a80eaa2dc",
"photo_avatar",
("avatar_v", "avatar_iv", "avatar_iii"),
),
(
"Brandon Business Sitting Front",
"Brandon_Business_Sitting_Front_public",
"studio_avatar",
("avatar_v", "avatar_iv", "avatar_iii"),
),
(
"Caroline Business Sitting Side",
"Caroline_Business_Sitting_Side_public",
"studio_avatar",
("avatar_v", "avatar_iv", "avatar_iii"),
),
(
"Ursula Lawyer Angle 4",
"f7173d2bb8584c00bfec6905c5e9a492",
"photo_avatar",
("avatar_v", "avatar_iv", "avatar_iii"),
),
(
"Sofia Corporate Presenter 01 Angle 3",
"fe563971fd2d438e957372dac9e2be8c",
"photo_avatar",
("avatar_v", "avatar_iv", "avatar_iii"),
),
(
"Seoyeon Health Nutrition Coach Angle 3",
"fe3c5d5028d941398d064b8fc64a2dea",
"photo_avatar",
("avatar_v", "avatar_iv", "avatar_iii"),
),
(
"Sanne Fitness Coach Angle 4",
"d967f935a8bf4a0c8f0bccfd66c501d2",
"photo_avatar",
("avatar_v", "avatar_iv", "avatar_iii"),
),
("Sander", "f5cd7b94056f495ca0610602d64a9aa3", "photo_avatar", ("avatar_v", "avatar_iv", "avatar_iii")),
(
"Rupert Personal Development Coach Angle 4",
"f57b3e626adb4bc997b38f64884adce4",
"photo_avatar",
("avatar_v", "avatar_iv", "avatar_iii"),
),
(
"Olivier Professor Angle 2",
"f6659bbb094b459c87c967edbb9ee481",
"photo_avatar",
("avatar_v", "avatar_iv", "avatar_iii"),
),
(
"Obi Health Nutrition Coach Angle 5",
"f3dc2c38201d414382f506d2d8e8d029",
"photo_avatar",
("avatar_v", "avatar_iv", "avatar_iii"),
),
(
"Matilda Modern Office Setting",
"fda889ac354a440da8dbecc410981273",
"photo_avatar",
("avatar_v", "avatar_iv", "avatar_iii"),
),
(
"Mateo Traditional Law Office",
"ff172d6c499c4e47ba6fcc5de631e9fc",
"photo_avatar",
("avatar_v", "avatar_iv", "avatar_iii"),
),
(
"Marlon Inviting Armchair Setting",
"f5a57db099ab462daa3e7c604a05dacc",
"photo_avatar",
("avatar_v", "avatar_iv", "avatar_iii"),
),
(
"Margaret Professor Angle 1",
"fb472bc29ab04bcca576e3703978fecb",
"photo_avatar",
("avatar_v", "avatar_iv", "avatar_iii"),
),
(
"Marek Therapy Coach Angle 3",
"e197768703f1463a93dc25ada1f421fb",
"photo_avatar",
("avatar_v", "avatar_iv", "avatar_iii"),
),
(
"Maeve Warm, Professional Setting",
"faf66681d8cc48dc82c4283200b3e782",
"photo_avatar",
("avatar_v", "avatar_iv", "avatar_iii"),
),
(
"Lorenzo Professor Angle 5",
"fc268dc244bb40d7a554663ce723dcf0",
"photo_avatar",
("avatar_v", "avatar_iv", "avatar_iii"),
),
("Luca", "Luca_public", "studio_avatar", ("avatar_iii",)),
("Bruce", "Bruce_public", "studio_avatar", ("avatar_iii",)),
("Nico", "Nico_public", "studio_avatar", ("avatar_iii",)),
("Lisa", "Lisa_public", "studio_avatar", ("avatar_iii",)),
("Sophie", "Sophie_public", "studio_avatar", ("avatar_iii",)),
("Aiko", "Aiko_public", "studio_avatar", ("avatar_iii",)),
("Rebecca (portrait)", "Rebecca_public", "studio_avatar", ("avatar_iii",)),
("Daphne in Grey blazer (portrait)", "Daphne_public_1", "studio_avatar", ("avatar_iii",)),
("Bryce in Black t-shirt", "Bryce_public_5", "studio_avatar", ("avatar_iii",)),
("Diora in White shirt", "Diora_public_3", "studio_avatar", ("avatar_iii",)),
("Freja in White blazer", "Freja_public_1", "studio_avatar", ("avatar_iii",)),
("Albert in Blue blazer", "Albert_public_2", "studio_avatar", ("avatar_iii",)),
("Emery in Red blazer", "Emery_public_1", "studio_avatar", ("avatar_iii",)),
("Minho in Blue shirt", "Minho_public_6", "studio_avatar", ("avatar_iii",)),
("Aditya in Brown blazer", "Aditya_public_4", "studio_avatar", ("avatar_iii",)),
("Nadim in Blue blazer", "Nadim_public_1", "studio_avatar", ("avatar_iii",)),
("Iker in Black blazer", "Iker_public_1", "studio_avatar", ("avatar_iii",)),
("Nour in Black blazer", "Nour_public_1", "studio_avatar", ("avatar_iii",)),
("Saskia in Blue blazer", "Saskia_public_1", "studio_avatar", ("avatar_iii",)),
("Lucien in Blue blazer", "Lucien_public_1", "studio_avatar", ("avatar_iii",)),
("Esmond in Blue suit", "Esmond_public_3", "studio_avatar", ("avatar_iii",)),
("Jinwoo in Blue suit", "Jinwoo_public_5", "studio_avatar", ("avatar_iii",)),
("Annelore in Red sweater (portrait)", "Annelore_public_3", "studio_avatar", ("avatar_iii",)),
("Bastien in Blue shirt", "Bastien_public_4", "studio_avatar", ("avatar_iii",)),
("Zosia in Khaki blazer", "Zosia_public_3", "studio_avatar", ("avatar_iii",)),
("Tahlia in Dark blue suit", "Tahlia_public_4", "studio_avatar", ("avatar_iii",)),
]
HEYGEN_AVATAR_OPTIONS = [x[0] for x in HEYGEN_AVATAR_LOOKS]
HEYGEN_AVATAR_MAP = {x[0]: (x[1], x[2], x[3]) for x in HEYGEN_AVATAR_LOOKS}
# (label, voice_id) — Starfish-compatible voices for the TTS endpoint
HEYGEN_VOICE_TTS: list[tuple[str, str]] = [
("Chill Brian (English, male)", "d2f4f24783d04e22ab49ee8fdc3715e0"),
("Zain (English, female)", "0047732240584155b1588455313e78ec"),
("Narrator Mateo - Excited 🤩 (Spanish, male)", "0077225a877e457db4572ccaf245910b"),
("Aria (English, female)", "007e1378fc454a9f976db570ba6164a7"),
("Caryns (English, female)", "0082e70326864107823605db0d77f5e0"),
("Klara (English, female)", "01209fdcd1c24a109c86dc24ee0f71c0"),
("Bold Kasia - Excited 🤩 (Polish, female)", "015482a78b9a46ebae74bd0beb17765b"),
("Shaun (English, male)", "01c42cddcfdc4665a57b8d89cba8ffc1"),
("Senthil (English, male)", "01d674cfd32b4728a3fddd21b7e7d543"),
("Cody (English, male)", "01f98ed43e6140349f47dbd37a416827"),
("Saffron (English, female)", "0258bbc2cd8648cfa357adfb833f6d7b"),
("Blanka - Lifelike (English, female)", "02880d1c6fd94b7799d91135581ed810"),
("Rami (English, male)", "02d5366a90af4c7a87157808ff352e33"),
("Rhodes (English, female)", "02dce0a169b3460084b6c914d18fb2c8"),
("Michelle - Voice 1 (English, female)", "02X8sHnuxFpsq1caYWN0"),
("Autumn - UGC 3 (English, female)", "03dca9ebfca441dba55fb14afa0791b7"),
("Reassuring Rupert (English, male)", "03fcf8ecb0a94b6b94e9007edb7c35f8"),
("Rose - UGC -2 (English, female)", "0495e14c2bd74eb3aeeef03583e0bce5"),
("Derya - Lifelike - Broadcaster 🎙️ (English, female)", "04d0ae1d0af2489ca7d3bb402a39a890"),
("Dynamic Derek (English, male)", "0516c2d857eb425c94e90b068241914e"),
("Lotte (English, female)", "052fcfb83d1a4c2f8d0368c226fea4b9"),
("Thanos - Broadcaster 🎙️ (English, male)", "054af44a167344d0af2722fdfef08d17"),
("Marcia (English, female)", "05f19352e8f74b0392a8f411eba40de1"),
("Camden (English, male)", "06468055edd4458aa131a1dfd813c1e9"),
("Rumi (English, female)", "06672207805f41a9ad0af6797f8aa14b"),
("Pippa (English, female)", "06b68c4dbb544935b9af984e80efa4fb"),
("William Prescott - Broadcaster 🎙️ (English, male)", "06c816b952f14fa9b3a6c42aa151f731"),
("Sammy (English, female)", "06e6facd99654b9dbb9308f67bf3a31c"),
("Breezy Bagus (Indonesian, male)", "06e81a5d7c8b41818d3f0b38f7cf15a1"),
("Ben (English, male)", "07ca39b243184dbcb82e7e0f0e524b21"),
("Smooth Dev (English, male)", "07d2ba65847541feb97abc9b60181555"),
("Daran inside booth (English, male)", "080d9383c0314056aef392892e009806"),
("Peppy Stella (English, female)", "084760b4922a44599575c770070ec2d7"),
("Silas (English, male)", "08f561403ec846dbbd8c691cc448f45a"),
("Aditya (English, male)", "09c3d65e44e247dd8b78a97a903feb58"),
("Christy (English, female)", "09d88c036bf449fa905900c08b235a37"),
("Elio (English, male)", "0a0b38624ac64ec6afcd5842a977ca10"),
("Luminous Laksh (Hindi, male)", "0adc547b76a5401c856274c379904eb7"),
("Jeff (English, male)", "0add542e349f4ccaba6ecb3b7ced6034"),
("Tahlia Brooks - Excited 🤩 (English, female)", "0b440d1ac2454d69a73302fc806522b1"),
("Riya Mehta (Hindi, female)", "0b464b2f4e2249a4b5a05e60eaf41e7e"),
("Ben Hart (English, male)", "0b47b5a637e944f9bfd49913999b344b"),
("Skylar (English, female)", "0bbfbda5aa924a68a9d1da7b8496052a"),
("Relaxed Reece (English, male)", "0c2151d538844c70a8b096de533f2828"),
("Daniel (English, male)", "0c23804af39a4946ac6fda42bfff2738"),
("Melani (English, female)", "0c54c6399ad64551a304e1a346677723"),
("Clover (English, female)", "0ccb0bea067d4449ad367baeed7ea2e9"),
("Pedro Lima - Serious 😐 (Portuguese, male)", "0d0e23e8170446e38b18a7380b2d30a8"),
("Ana Carvalho (Portuguese, female)", "0d23c5b2f6004e909802a2e8bfcd52c2"),
("Confident Connor - Excited 🤩 (English, male)", "0dd34c3eb79247238219eea35aeb58cd"),
("Vibrant Victor (Spanish, male)", "1062976ea8bf42f4adc27c7e868b8fde"),
("Young Olivier (French, male)", "1c5dc9a8f8cf4de0932f91d75f43a15d"),
("Émile Noir (French, male)", "25a6a67280574d3da78e97b1935ebfc7"),
("Steadfast Stefan (German, male)", "0eb85e6e8710473b82f7e88609ba3053"),
("Deep Dieter (German, male)", "118949676b0a46629d1ad52981c3ef84"),
("Serene Marco (Italian, male)", "72e922488a614041b5ab5f6ee07e3deb"),
("Murmuring Matteo (Italian, male)", "755902b751654f30a6ef49e8bbcacfec"),
("Gail in car (Multilingual, female)", "0214ac51f93e420f8711d568dcfbc50e"),
("Daran outside walking (Multilingual, male)", "0ac81e725f4948dfa9638ceca216bcfa"),
("BOB - Voice 1 (Chinese, unknown)", "dMkR1XwIkarpNqWUJLnX"),
("Hakeem Hassan (Arabic, male)", "61a4359785664d01a59664ceb87ce6d4"),
("Rami Idris (Arabic, male)", "a0bd2e5d41a74643be47ac75ca9171a2"),
("Bold Kasia - Friendly 😊 (Polish, female)", "331624aec8b24a6c9287b8e16bdf54e8"),
("Tranquil Tulin (Turkish, female)", "61646c861eb64e2d9036d8db51385356"),
("Dynamic Derya (Turkish, female)", "664b73058b784aa89ddb2924c141d441"),
("Quiet Dewa (Indonesian, male)", "1fa1193cf1d74f27ba58531c07ef9862"),
("Cuong (Vietnamese, male)", "8af68d7ea38f4e7ca05cf46c3f7a590b"),
]
HEYGEN_VOICE_TTS_OPTIONS = [x[0] for x in HEYGEN_VOICE_TTS]
HEYGEN_VOICE_TTS_MAP = dict(HEYGEN_VOICE_TTS)
# (label, voice_id) — top-ranked voices for video narration (any engine)
HEYGEN_VOICE_GENERAL: list[tuple[str, str]] = [
("Cassidy (English, female)", "16a09e4706f74997ba4ed05ea11470f6"),
("Hope (English, female)", "42d00d4aac5441279d8536cd6b52c53c"),
("Archer (English, male)", "453c20e1525a429080e2ad9e4b26f2cd"),
("Brittney (English, female)", "4754e1ec667544b0bd18cdf4bec7d6a7"),
("Mark (English, male)", "5d8c378ba8c3434586081a52ac368738"),
("Andrew (English, male)", "6be73833ef9a4eb0aeee399b8fe9d62b"),
("Spuds Oxley (English, male)", "76940a9adcd0490a9ce2cfe9a64a2664"),
("Patrick (English, male)", "7e157ec62c9c45f1adca12faae72c86f"),
("David Castlemore (English, male)", "828b59f834fd4c7188da322b6d9b6c75"),
("Michael C (English, male)", "8661cd40d6c44c709e2d0031c0186ada"),
("Adam Stone (English, male)", "88bb9ee1c81b466eb2a08fdde86d3619"),
("Alex (English, male)", "897d6a9b2c844f56aa077238768fe10a"),
("Monika Sogam (English, female)", "97dd67ab8ce242b6a9e7689cb00c6414"),
("Jessica Anne Bogart (English, female)", "b966c31caf124c2a99f19ff1479c964f"),
("John Doe (English, male)", "c4a8ceb7a2954500bc047fb092bcff3f"),
("Ivy (English, female)", "cef3bc4e0a84424cafcde6f2cf466c97"),
("Chill Brian (English, male)", "d2f4f24783d04e22ab49ee8fdc3715e0"),
("Allison (English, female)", "f8c69e517f424cafaecde32dde57096b"),
("Mia Starset (Norwegian, female)", "000466f8ac6d47a49f5743d50b3778de"),
("William Shanks (Spanish, male)", "001248bb63f847888d37b766ee8b3a47"),
("Zain (English, female)", "0047732240584155b1588455313e78ec"),
("Jora Slobod (Romanian, male)", "00631519159a402ab5d8f719e51532bb"),
("Narrator Mateo - Excited 🤩 (Spanish, male)", "0077225a877e457db4572ccaf245910b"),
("Aria (English, female)", "007e1378fc454a9f976db570ba6164a7"),
("Caryns (English, female)", "0082e70326864107823605db0d77f5e0"),
("Klara (English, female)", "01209fdcd1c24a109c86dc24ee0f71c0"),
("Son Tran (Vietnamese, male)", "0132f85950a94d11ba180f885101bf84"),
("Bold Kasia - Excited 🤩 (Polish, female)", "015482a78b9a46ebae74bd0beb17765b"),
("Marc Aurèle (French, male)", "018a94cf15574491a0bab7f6799ac15b"),
("Shaun (English, male)", "01c42cddcfdc4665a57b8d89cba8ffc1"),
("Senthil (English, male)", "01d674cfd32b4728a3fddd21b7e7d543"),
("Cody (English, male)", "01f98ed43e6140349f47dbd37a416827"),
("Saffron (English, female)", "0258bbc2cd8648cfa357adfb833f6d7b"),
("Blanka - Lifelike (English, female)", "02880d1c6fd94b7799d91135581ed810"),
("Rami (English, male)", "02d5366a90af4c7a87157808ff352e33"),
("Rhodes (English, female)", "02dce0a169b3460084b6c914d18fb2c8"),
("Michelle - Voice 1 (English, female)", "02X8sHnuxFpsq1caYWN0"),
("Tuba (, female)", "034ca0c32b6542028748d6d365d90d6a"),
("Autumn - UGC 3 (English, female)", "03dca9ebfca441dba55fb14afa0791b7"),
("Reassuring Rupert (English, male)", "03fcf8ecb0a94b6b94e9007edb7c35f8"),
]
HEYGEN_VOICE_GENERAL_OPTIONS = [x[0] for x in HEYGEN_VOICE_GENERAL]
HEYGEN_VOICE_GENERAL_MAP = dict(HEYGEN_VOICE_GENERAL)
HEYGEN_TRANSLATE_LANGUAGES = [
"English",
"Spanish",
"Spanish (Spain)",
"Spanish (Mexico)",
"French",
"French (France)",
"German",
"German (Germany)",
"Portuguese",
"Portuguese (Brazil)",
"Italian",
"Italian (Italy)",
"Japanese",
"Japanese (Japan)",
"Korean",
"Chinese (Mandarin, Simplified)",
"Arabic",
"Hindi",
"Hindi (India)",
"Russian",
"Russian (Russia)",
"Dutch",
"Polish",
"Turkish",
"Indonesian",
"Vietnamese",
"Ukrainian",
"Afrikaans (South Africa)",
"Albanian (Albania)",
"Amharic (Ethiopia)",
"Arabic (Algeria)",
"Arabic (Bahrain)",
"Arabic (Egypt)",
"Arabic (Iraq)",
"Arabic (Jordan)",
"Arabic (Kuwait)",
"Arabic (Lebanon)",
"Arabic (Libya)",
"Arabic (Morocco)",
"Arabic (Oman)",
"Arabic (Qatar)",
"Arabic (Saudi Arabia)",
"Arabic (Syria)",
"Arabic (Tunisia)",
"Arabic (United Arab Emirates)",
"Arabic (World)",
"Arabic (Yemen)",
"Armenian (Armenia)",
"Azerbaijani (Latin, Azerbaijan)",
"Bangla (Bangladesh)",
"Basque",
"Belarusian (Belarus)",
"Bengali (India)",
"Bosnian (Bosnia and Herzegovina)",
"Bulgarian",
"Bulgarian (Bulgaria)",
"Burmese (Myanmar)",
"Catalan",
"Chinese (Cantonese, Traditional)",
"Chinese (Jilu Mandarin, Simplified)",
"Chinese (Northeastern Mandarin, Simplified)",
"Chinese (Southwestern Mandarin, Simplified)",
"Chinese (Taiwanese Mandarin, Traditional)",
"Chinese (Wu, Simplified)",
"Chinese (Zhongyuan Mandarin Henan, Simplified)",
"Chinese (Zhongyuan Mandarin Shaanxi, Simplified)",
"Croatian",
"Croatian (Croatia)",
"Czech",
"Czech (Czechia)",
"Danish",
"Danish (Denmark)",
"Dutch (Belgium)",
"Dutch (Netherlands)",
"English (Australia)",
"English (Canada)",
"English (Hong Kong SAR)",
"English (India)",
"English (Ireland)",
"English (Kenya)",
"English (New Zealand)",
"English (Nigeria)",
"English (Philippines)",
"English (Singapore)",
"English (South Africa)",
"English (Tanzania)",
"English (UK)",
"English (United States)",
"Estonian (Estonia)",
"Filipino",
"Filipino (Cebuano)",
"Filipino (Philippines)",
"Finnish",
"Finnish (Finland)",
"French (Belgium)",
"French (Canada)",
"French (Switzerland)",
"Galician",
"Georgian (Georgia)",
"German (Austria)",
"German (Switzerland)",
"Greek",
"Greek (Greece)",
"Gujarati (India)",
"Haitian Creole (Haiti)",
"Hebrew (Israel)",
"Hungarian (Hungary)",
"Icelandic (Iceland)",
"Indonesian (Indonesia)",
"Irish (Ireland)",
"Javanese (Latin, Indonesia)",
"Kannada (India)",
"Kazakh (Kazakhstan)",
"Khmer (Cambodia)",
"Konkani (India)",
"Korean (Korea)",
"Lao (Laos)",
"Latin (Vatican City)",
"Latvian (Latvia)",
"Lithuanian (Lithuania)",
"Luxembourgish (Luxembourg)",
"Macedonian (North Macedonia)",
"Maithili (India)",
"Malagasy (Madagascar)",
"Malay",
"Malay (Malaysia)",
"Malayalam (India)",
"Maltese (Malta)",
"Mandarin",
"Marathi (India)",
"Mongolian (Mongolia)",
"Nepali (Nepal)",
"Norwegian Bokmål (Norway)",
"Norwegian Nynorsk (Norway)",
"Odia (India)",
"Pashto (Afghanistan)",
"Persian (Iran)",
"Polish (Poland)",
"Portuguese (Portugal)",
"Punjabi (India)",
"Romanian",
"Romanian (Romania)",
"Serbian (Latin, Serbia)",
"Sindhi (India)",
"Sinhala (Sri Lanka)",
"Slovak",
"Slovak (Slovakia)",
"Slovenian (Slovenia)",
"Somali (Somalia)",
"Spanish (Argentina)",
"Spanish (Bolivia)",
"Spanish (Chile)",
"Spanish (Colombia)",
"Spanish (Costa Rica)",
"Spanish (Cuba)",
"Spanish (Dominican Republic)",
"Spanish (Ecuador)",
"Spanish (El Salvador)",
"Spanish (Equatorial Guinea)",
"Spanish (Guatemala)",
"Spanish (Honduras)",
"Spanish (Latin America)",
"Spanish (Nicaragua)",
"Spanish (Panama)",
"Spanish (Paraguay)",
"Spanish (Peru)",
"Spanish (Puerto Rico)",
"Spanish (United States)",
"Spanish (Uruguay)",
"Spanish (Venezuela)",
"Sundanese (Indonesia)",
"Swahili (Kenya)",
"Swahili (Tanzania)",
"Swedish",
"Swedish (Sweden)",
"Tamil",
"Tamil (India)",
"Tamil (Malaysia)",
"Tamil (Singapore)",
"Tamil (Sri Lanka)",
"Telugu (India)",
"Thai (Thailand)",
"Turkish (Türkiye)",
"Ukrainian (Ukraine)",
"Urdu (India)",
"Urdu (Pakistan)",
"Uzbek (Latin, Uzbekistan)",
"Vietnamese (Vietnam)",
"Welsh (United Kingdom)",
"Zulu (South Africa)",
]
+1 -1
View File
@@ -128,7 +128,7 @@ class OpenAIResponse(ModelResponseProperties, ResponseProperties):
parallel_tool_calls: bool | None = Field(True)
status: str | None = Field(
None,
description="One of `completed`, `failed`, `in_progress`, `incomplete`, `queued`, or `cancelled`.",
description="One of `completed`, `failed`, `in_progress`, or `incomplete`.",
)
usage: ResponseUsage | None = Field(None)
-49
View File
@@ -1,49 +0,0 @@
from pydantic import BaseModel, Field
class SyncInputItem(BaseModel):
type: str = Field(..., description="Input kind: 'video', 'image' or 'audio'.")
url: str = Field(...)
class SyncActiveSpeakerDetection(BaseModel):
auto_detect: bool | None = Field(
None, description="Detect the active speaker automatically. Video input only; rejected for images."
)
frame_number: int | None = Field(
None, description="Frame used for manual speaker selection. Must be 0 for image inputs."
)
coordinates: list[int] | None = Field(
None, description="Pixel [x, y] of the speaker's face in the frame selected by frame_number."
)
class SyncGenerationOptions(BaseModel):
sync_mode: str | None = Field(
None,
description="How to resolve an audio/video duration mismatch: "
"cut_off, bounce, loop, silence or remap. Ignored for image inputs.",
)
i2v_prompt: str | None = Field(
None, description="Motion prompt for image-to-video generation. Image input only."
)
active_speaker_detection: SyncActiveSpeakerDetection | None = Field(None)
class SyncGenerationRequest(BaseModel):
model: str = Field(..., description="Generation model, e.g. 'sync-3'.")
input: list[SyncInputItem] = Field(
..., description="Exactly one visual input (video or image) plus one audio input."
)
options: SyncGenerationOptions | None = Field(None)
class SyncGeneration(BaseModel):
"""Subset of the Generation object returned by POST /v2/generate and GET /v2/generate/{id}."""
id: str = Field(...)
status: str = Field(..., description="PENDING | PROCESSING | COMPLETED | FAILED | REJECTED")
outputUrl: str | None = Field(None)
outputDuration: float | None = Field(None)
error: str | None = Field(None, description="Human-readable failure message.")
errorCode: str | None = Field(None, description="Stable machine-readable code from the GET /v2/errors catalog.")
+48 -122
View File
@@ -24,11 +24,6 @@ from comfy_api_nodes.apis.gemini import (
GeminiImageGenerateContentRequest,
GeminiImageGenerationConfig,
GeminiInlineData,
GeminiInteraction,
GeminiInteractionGenerationConfig,
GeminiInteractionMediaPart,
GeminiInteractionRequest,
GeminiInteractionTextPart,
GeminiMimeType,
GeminiPart,
GeminiRole,
@@ -56,11 +51,9 @@ from comfy_api_nodes.util import (
)
GEMINI_BASE_ENDPOINT = "/proxy/vertexai/gemini"
GEMINI_INTERACTIONS_ENDPOINT = "/proxy/gemini-interactions"
GEMINI_MAX_INPUT_FILE_SIZE = 20 * 1024 * 1024 # 20 MB
GEMINI_URL_INPUT_BUDGET = 10
GEMINI_MAX_INLINE_BYTES = 18 * 1024 * 1024
GEMINI_INTERACTIONS_MAX_INLINE_BYTES = 90 * 1024 * 1024 # the Interactions API rejects requests over ~100MiB
GEMINI_IMAGE_SYS_PROMPT = (
"You are an expert image-generation engine. You must ALWAYS produce an image.\n"
"Interpret all user input—regardless of "
@@ -238,10 +231,29 @@ async def get_image_from_response(response: GeminiGenerateContentResponse, thoug
return torch.cat(image_tensors, dim=0)
async def get_video_from_response(
response: GeminiGenerateContentResponse, cls: type[IO.ComfyNode] | None = None
) -> InputImpl.VideoFromFile:
parts = get_parts_by_type(response, "video/*")
for part in parts:
if part.inlineData and part.inlineData.data:
return InputImpl.VideoFromFile(BytesIO(base64.b64decode(part.inlineData.data)))
if part.fileData and part.fileData.fileUri:
return await download_url_to_video_output(part.fileData.fileUri, cls=cls)
model_message = get_text_from_response(response).strip()
if model_message:
raise ValueError(f"Gemini did not generate a video. Model response: {model_message}")
raise ValueError(
"Gemini did not generate a video. Try rephrasing your prompt, "
"shortening the requested duration, or reducing the number of input images/videos."
)
def calculate_tokens_price(response: GeminiGenerateContentResponse) -> float | None:
if not response.modelVersion:
return None
# Define prices (Cost per 1,000,000 tokens), see https://cloud.google.com/vertex-ai/generative-ai/pricing
output_video_tokens_price = 0.0
if response.modelVersion == "gemini-2.5-pro":
input_tokens_price = 1.25
output_text_tokens_price = 10.0
@@ -262,10 +274,6 @@ def calculate_tokens_price(response: GeminiGenerateContentResponse) -> float | N
input_tokens_price = 0.25
output_text_tokens_price = 1.50
output_image_tokens_price = 0.0
elif response.modelVersion == "gemini-3.5-flash":
input_tokens_price = 1.50
output_text_tokens_price = 9.0
output_image_tokens_price = 0.0
elif response.modelVersion in ("gemini-3-pro-image-preview", "gemini-3-pro-image"):
input_tokens_price = 2
output_text_tokens_price = 12.0
@@ -278,6 +286,11 @@ def calculate_tokens_price(response: GeminiGenerateContentResponse) -> float | N
input_tokens_price = 0.25
output_text_tokens_price = 1.50
output_image_tokens_price = 30.0
elif response.modelVersion == "gemini-omni-flash-preview":
input_tokens_price = 2.145
output_text_tokens_price = 12.87
output_image_tokens_price = 0.0
output_video_tokens_price = 25.025
else:
return None
final_price = response.usageMetadata.promptTokenCount * input_tokens_price
@@ -285,6 +298,8 @@ def calculate_tokens_price(response: GeminiGenerateContentResponse) -> float | N
for i in response.usageMetadata.candidatesTokensDetails:
if i.modality == Modality.IMAGE:
final_price += output_image_tokens_price * i.tokenCount # for Nano Banana models
elif i.modality == Modality.VIDEO:
final_price += output_video_tokens_price * i.tokenCount # for Omni Flash
else:
final_price += output_text_tokens_price * i.tokenCount
if response.usageMetadata.thoughtsTokenCount:
@@ -292,58 +307,6 @@ def calculate_tokens_price(response: GeminiGenerateContentResponse) -> float | N
return final_price / 1_000_000.0
def get_text_from_interaction(interaction: GeminiInteraction) -> str:
"""Extract and concatenate all model output text from an Interactions API response."""
texts = []
for step in interaction.steps or []:
if step.type != "model_output":
continue
for content in step.content or []:
if content.type == "text" and content.text:
texts.append(content.text)
return "\n".join(texts)
async def get_video_from_interaction(
interaction: GeminiInteraction, cls: type[IO.ComfyNode] | None = None
) -> InputImpl.VideoFromFile:
for step in interaction.steps or []:
if step.type != "model_output":
continue
for content in step.content or []:
if content.type != "video":
continue
if content.data:
return InputImpl.VideoFromFile(BytesIO(base64.b64decode(content.data)))
if content.uri:
return await download_url_to_video_output(content.uri, cls=cls)
model_message = get_text_from_interaction(interaction).strip()
if model_message:
raise ValueError(f"Gemini did not generate a video. Model response: {model_message}")
raise ValueError(
"Gemini did not generate a video. Try rephrasing your prompt, "
"shortening the requested duration, or reducing the number of input images/videos."
)
def calculate_interaction_tokens_price(interaction: GeminiInteraction) -> float | None:
if interaction.usage is None:
return None
input_tokens_price = 1.5
output_tokens_prices = {"text": 9.0, "video": 17.5}
thoughts_tokens_price = 9.0
final_price = 0.0
for i in interaction.usage.input_tokens_by_modality or []:
if i.tokens:
final_price += input_tokens_price * i.tokens
for i in interaction.usage.output_tokens_by_modality or []:
if i.tokens and i.modality in output_tokens_prices:
final_price += output_tokens_prices[i.modality] * i.tokens
if interaction.usage.total_thought_tokens:
final_price += thoughts_tokens_price * interaction.usage.total_thought_tokens
return final_price / 1_000_000.0
def create_video_parts(video_input: Input.Video) -> list[GeminiPart]:
"""Convert a single video input to Gemini API compatible parts (inline MP4/H.264)."""
base_64_string = video_to_base64_string(
@@ -470,24 +433,14 @@ async def build_gemini_media_parts(
part, nbytes = _media_inline_part(kind, payload)
inline_bytes += nbytes
if inline_bytes > max_inline_bytes:
detail = f" after the first {url_budget} inputs are uploaded as URLs" if url_budget else ""
raise ValueError(
f"Too much media to send inline (over {max_inline_bytes // (1024 * 1024)}MB{detail}). "
"Reduce the number or size of attached media."
f"Too much media to send inline (over {max_inline_bytes // (1024 * 1024)}MB after the first "
f"{url_budget} inputs are uploaded as URLs). Reduce the number or size of attached media."
)
parts.append(part)
return parts
def to_interaction_media_part(part: GeminiPart) -> GeminiInteractionMediaPart:
"""Convert a fileData/inlineData GeminiPart into an Interactions API media part."""
if part.fileData:
mime = part.fileData.mimeType.value
return GeminiInteractionMediaPart(type=mime.split("/")[0], uri=part.fileData.fileUri, mime_type=mime)
mime = part.inlineData.mimeType.value
return GeminiInteractionMediaPart(type=mime.split("/")[0], data=part.inlineData.data, mime_type=mime)
class GeminiNode(IO.ComfyNode):
"""
Node to generate text responses from a Gemini model.
@@ -666,12 +619,11 @@ class GeminiNode(IO.ComfyNode):
GEMINI_V2_MODELS: dict[str, str] = {
"Gemini 3.1 Pro": "gemini-3.1-pro-preview",
"Gemini 3.5 Flash": "gemini-3.5-flash",
"Gemini 3.1 Flash-Lite": "gemini-3.1-flash-lite-preview",
}
def _gemini_text_model_inputs(thinking_default: str, thinking_options: list[str] | None = None) -> list[Input]:
def _gemini_text_model_inputs(thinking_default: str) -> list[Input]:
"""Per-model inputs revealed by the model DynamicCombo (shared media + sampling controls)."""
return [
IO.Autogrow.Input(
@@ -709,7 +661,7 @@ def _gemini_text_model_inputs(thinking_default: str, thinking_options: list[str]
),
IO.Combo.Input(
"thinking_level",
options=thinking_options or ["LOW", "HIGH"],
options=["LOW", "HIGH"],
default=thinking_default,
tooltip="How hard the model reasons internally before answering. "
"HIGH improves quality on difficult tasks but costs more (thinking) tokens and is slower.",
@@ -767,10 +719,6 @@ class GeminiNodeV2(IO.ComfyNode):
IO.DynamicCombo.Input(
"model",
options=[
IO.DynamicCombo.Option(
"Gemini 3.5 Flash",
_gemini_text_model_inputs("MEDIUM", ["MINIMAL", "LOW", "MEDIUM", "HIGH"]),
),
IO.DynamicCombo.Option("Gemini 3.1 Pro", _gemini_text_model_inputs("HIGH")),
IO.DynamicCombo.Option("Gemini 3.1 Flash-Lite", _gemini_text_model_inputs("LOW")),
],
@@ -811,13 +759,7 @@ class GeminiNodeV2(IO.ComfyNode):
"type": "list_usd",
"usd": [0.00025, 0.0015],
"format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" }
}
: $contains($m, "3.5 flash") ? {
"type": "list_usd",
"usd": [0.0015, 0.009],
"format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" }
}
: {
} : {
"type": "list_usd",
"usd": [0.002, 0.012],
"format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" }
@@ -1191,9 +1133,7 @@ class GeminiImage2(IO.ComfyNode):
) -> IO.NodeOutput:
validate_string(prompt, strip_whitespace=True, min_length=1)
if model == "Nano Banana 2 (Gemini 3.1 Flash Image)":
model = "gemini-3.1-flash-image"
elif model == "gemini-3-pro-image-preview":
model = "gemini-3-pro-image"
model = "gemini-3.1-flash-image-preview"
parts: list[GeminiPart] = [GeminiPart(text=prompt)]
if images is not None:
@@ -1567,7 +1507,7 @@ class GeminiNanoBanana2V2(IO.ComfyNode):
validate_string(prompt, strip_whitespace=True, min_length=1)
model_choice = model["model"]
if model_choice == "Nano Banana 2 (Gemini 3.1 Flash Image)":
model_id = "gemini-3.1-flash-image"
model_id = "gemini-3.1-flash-image-preview"
elif model_choice == "Nano Banana 2 Lite":
model_id = "gemini-3.1-flash-lite-image"
else:
@@ -1719,7 +1659,7 @@ class GeminiVideoOmni(IO.ComfyNode):
],
is_api_node=True,
price_badge=IO.PriceBadge(
expr='{"type":"usd","usd":0.101,"format":{"suffix":"/second","approximate":true}}'
expr='{"type":"usd","usd":0.146,"format":{"suffix":"/second","approximate":true}}'
),
)
@@ -1738,41 +1678,27 @@ class GeminiVideoOmni(IO.ComfyNode):
for video in videos:
validate_video_duration(video, max_duration=10)
parts: list[GeminiInteractionTextPart | GeminiInteractionMediaPart] = []
parts: list[GeminiPart] = []
if images or videos:
# The Interactions API accepts video only inline or as a Files API URI, not as an HTTP URL.
media_parts = await build_gemini_media_parts(
cls, [], [], videos, url_budget=0, max_inline_bytes=GEMINI_INTERACTIONS_MAX_INLINE_BYTES
)
video_inline_bytes = sum(len(p.inlineData.data) for p in media_parts)
media_parts += await build_gemini_media_parts(
cls, images, [], [], max_inline_bytes=GEMINI_INTERACTIONS_MAX_INLINE_BYTES - video_inline_bytes
)
parts.extend(to_interaction_media_part(p) for p in media_parts)
parts.append(GeminiInteractionTextPart(text=prompt))
interaction = await sync_op(
parts.extend(await build_gemini_media_parts(cls, images, [], videos))
parts.append(GeminiPart(text=prompt))
response = await sync_op(
cls,
ApiEndpoint(path=GEMINI_INTERACTIONS_ENDPOINT, method="POST"),
data=GeminiInteractionRequest(
model=model_id,
input=parts,
generation_config=GeminiInteractionGenerationConfig(
ApiEndpoint(path=f"{GEMINI_BASE_ENDPOINT}/{model_id}", method="POST"),
data=GeminiGenerateContentRequest(
contents=[GeminiContent(role=GeminiRole.user, parts=parts)],
generationConfig=GeminiGenerationConfig(
responseModalities=["TEXT", "VIDEO"],
temperature=model.get("temperature", 1.0),
top_p=model.get("top_p", 0.95),
topP=model.get("top_p", 0.95),
),
),
response_model=GeminiInteraction,
price_extractor=calculate_interaction_tokens_price,
response_model=GeminiGenerateContentResponse,
price_extractor=calculate_tokens_price,
)
if interaction.status != "completed":
model_message = get_text_from_interaction(interaction).strip()
raise ValueError(
f"Gemini interaction did not complete (status: {interaction.status})."
+ (f" Model response: {model_message}" if model_message else "")
)
return IO.NodeOutput(
await get_video_from_interaction(interaction, cls=cls),
get_text_from_interaction(interaction),
await get_video_from_response(response, cls=cls),
get_text_from_response(response),
)
-799
View File
@@ -1,799 +0,0 @@
import uuid
import torch
from typing_extensions import override
from comfy_api.latest import IO, ComfyExtension, Input
from comfy_api_nodes.apis.heygen import (
HEYGEN_AVATAR_MAP,
HEYGEN_AVATAR_OPTIONS,
HEYGEN_TRANSLATE_LANGUAGES,
HEYGEN_VOICE_GENERAL_MAP,
HEYGEN_VOICE_GENERAL_OPTIONS,
HEYGEN_VOICE_TTS_MAP,
HEYGEN_VOICE_TTS_OPTIONS,
)
from comfy_api_nodes.util import (
ApiEndpoint,
audio_bytes_to_audio_input,
download_url_as_bytesio,
download_url_to_image_tensor,
download_url_to_video_output,
downscale_image_tensor_by_max_side,
get_number_of_images,
poll_op_raw,
sync_op_raw,
upload_audio_to_comfyapi,
upload_image_to_comfyapi,
upload_images_to_comfyapi,
upload_video_to_comfyapi,
validate_string,
)
from server import PromptServer
_VIDEOS_PATH = "/proxy/heygen/v3/videos"
_TRANSLATIONS_PATH = "/proxy/heygen/v3/video-translations"
_SPEECH_PATH = "/proxy/heygen/v3/voices/speech"
_AVATARS_PATH = "/proxy/heygen/v3/avatars"
_LOOKS_PATH = "/proxy/heygen/v3/avatars/looks"
_DEFAULT_VOICE_OPTION = "(avatar's default voice)"
_AVATARS_BY_ENGINE = {
e: [label for label, (_aid, _atype, engines) in HEYGEN_AVATAR_MAP.items() if e in engines]
for e in ("avatar_iv", "avatar_iii", "avatar_v")
}
async def _apply_speech_source(cls: type[IO.ComfyNode], payload: dict, speech: dict, require_voice: bool) -> None:
"""Fill script/audio speech fields of a /v3/videos payload from the DynamicCombo dict."""
if speech["speech"] == "audio":
payload["audio_url"] = await upload_audio_to_comfyapi(
cls, speech["audio"], container_format="mp3", codec_name="libmp3lame", mime_type="audio/mpeg"
)
elif speech["speech"] == "script":
validate_string(speech["text"], strip_whitespace=True, min_length=1, max_length=5000)
payload["script"] = speech["text"]
voice_id = speech.get("custom_voice_id", "").strip()
if not voice_id and speech["voice"] != _DEFAULT_VOICE_OPTION:
voice_id = HEYGEN_VOICE_GENERAL_MAP[speech["voice"]]
if voice_id:
payload["voice_id"] = voice_id
elif require_voice:
raise ValueError("A voice is required when driving the video with a text script.")
speed = speech.get("voice_speed", 1.0)
if speed != 1.0:
payload["voice_settings"] = {"speed": round(speed, 2)}
async def _create_and_poll_video(cls: type[IO.ComfyNode], payload: dict) -> dict:
"""POST a /v3/videos payload, poll until terminal, and return the final video data."""
created = await sync_op_raw(
cls,
ApiEndpoint(path=_VIDEOS_PATH, method="POST", headers={"Idempotency-Key": uuid.uuid4().hex}),
data=payload,
)
video_id = (created.get("data") or {}).get("video_id")
if not video_id:
raise ValueError(f"HeyGen did not return a video_id: {created}")
final = await poll_op_raw(
cls,
ApiEndpoint(path=f"{_VIDEOS_PATH}/{video_id}"),
status_extractor=lambda r: (r.get("data") or {}).get("status"),
queued_statuses=["pending", "waiting"],
poll_interval=5.0,
)
data = final["data"]
if not data.get("video_url"):
raise ValueError(f"HeyGen returned no video_url for video {video_id}.")
return data
async def _resolve_avatar(
cls: type[IO.ComfyNode], avatar_label: str, custom_avatar_id: str, engine_choice: str
) -> tuple[str, str | None]:
"""Resolve (avatar_id, engine_type) from the combo/override + engine widgets."""
custom_avatar_id = custom_avatar_id.strip()
if custom_avatar_id:
look = (
await sync_op_raw(
cls,
ApiEndpoint(path=f"{_LOOKS_PATH}/{custom_avatar_id}"),
final_label_on_success=None,
)
).get("data") or {}
avatar_id = custom_avatar_id
avatar_label = look.get("name") or custom_avatar_id
supported = look.get("supported_api_engines") or []
else:
avatar_id, avatar_type, supported = HEYGEN_AVATAR_MAP[avatar_label]
if engine_choice == "auto":
engine = next((e for e in ("avatar_iv", "avatar_iii", "avatar_v") if e in supported), None)
else:
engine = engine_choice
if supported and engine not in supported:
raise ValueError(
f"Avatar '{avatar_label}' does not support the {engine} engine "
f"(supported: {', '.join(supported)}). Set engine to 'auto' to pick "
"a compatible engine automatically."
)
return avatar_id, engine
class HeyGenTalkingPhotoNode(IO.ComfyNode):
"""Animate a still image of a person into a lip-synced talking video."""
@classmethod
def define_schema(cls) -> IO.Schema:
return IO.Schema(
node_id="HeyGenTalkingPhotoNode",
display_name="HeyGen Talking Photo",
category="partner/video/HeyGen",
description="Animate any image of a person into a lip-synced talking video "
"(HeyGen Avatar IV). Drive it with a text script or your own audio.",
inputs=[
IO.Image.Input(
"image",
tooltip="Image of a person to animate. Downscaled automatically if larger than 2K.",
),
IO.DynamicCombo.Input(
"speech",
display_name="speech source",
options=[
IO.DynamicCombo.Option(
"script",
[
IO.String.Input(
"text",
multiline=True,
default="",
tooltip="Text for the avatar to speak (up to 5000 characters). "
"The generated speech must be at least 1 second long.",
),
IO.Combo.Input(
"voice",
options=HEYGEN_VOICE_GENERAL_OPTIONS,
tooltip="Voice for the script (HeyGen's most popular voices).",
),
IO.String.Input(
"custom_voice_id",
default="",
optional=True,
tooltip="Optional HeyGen voice ID. When set, overrides the voice selected above. "
"Any voice from HeyGen's library (2000+) can be used.",
),
IO.Float.Input(
"voice_speed",
default=1.0,
min=0.5,
max=1.5,
step=0.05,
optional=True,
tooltip="Speech speed multiplier.",
),
],
),
IO.DynamicCombo.Option(
"audio",
[
IO.Audio.Input(
"audio",
tooltip="Audio for the avatar to lip-sync, up to 10 minutes.",
),
],
),
],
tooltip="Drive the avatar with a text script (HeyGen text-to-speech) or your own audio.",
),
IO.Combo.Input(
"resolution",
options=["720p", "1080p"],
default="1080p",
optional=True,
tooltip="Output video resolution.",
),
IO.Combo.Input(
"aspect_ratio",
options=["auto", "16:9", "9:16", "1:1", "4:5", "5:4"],
default="auto",
optional=True,
tooltip="Output aspect ratio. 'auto' follows the input image.",
),
IO.Combo.Input(
"expressiveness",
options=["low", "medium", "high"],
default="low",
optional=True,
tooltip="How expressive the animated face and gestures are.",
),
IO.Int.Input(
"seed",
default=42,
min=0,
max=2147483647,
control_after_generate=True,
optional=True,
tooltip="Not sent to HeyGen; change it to force a re-run.",
),
],
outputs=[IO.Video.Output()],
hidden=[
IO.Hidden.auth_token_comfy_org,
IO.Hidden.api_key_comfy_org,
IO.Hidden.unique_id,
],
is_api_node=True,
price_badge=IO.PriceBadge(
expr="""{"type":"usd","usd":0.0715,"format":{"suffix":"/second"}}""",
),
)
@classmethod
async def execute(
cls,
image: Input.Image,
speech: dict,
resolution: str = "1080p",
aspect_ratio: str = "auto",
expressiveness: str = "low",
seed: int = 0,
) -> IO.NodeOutput:
image = downscale_image_tensor_by_max_side(image, max_side=2000)
image_url = await upload_image_to_comfyapi(cls, image, mime_type="image/png", total_pixels=None)
payload = {
"type": "image",
"image": {"type": "url", "url": image_url},
"resolution": resolution,
"aspect_ratio": aspect_ratio,
"expressiveness": expressiveness,
"title": "ComfyUI Talking Photo",
}
await _apply_speech_source(cls, payload, speech, require_voice=True)
video = await _create_and_poll_video(cls, payload)
return IO.NodeOutput(await download_url_to_video_output(video["video_url"]))
class HeyGenAvatarVideoNode(IO.ComfyNode):
"""Generate a presenter video from a HeyGen avatar look."""
@classmethod
def define_schema(cls) -> IO.Schema:
return IO.Schema(
node_id="HeyGenAvatarVideoNode",
display_name="HeyGen Avatar Video",
category="partner/video/HeyGen",
description="Generate a talking-presenter video from a HeyGen avatar. "
"Includes HeyGen's most popular public avatars; any look ID can be supplied as an override.",
inputs=[
IO.DynamicCombo.Input(
"engine",
options=[
IO.DynamicCombo.Option(
"auto",
[
IO.Combo.Input(
"avatar",
options=HEYGEN_AVATAR_OPTIONS,
tooltip="Avatar look to present the video (curated from HeyGen's "
"public library). The best engine the look supports is chosen "
"automatically.",
),
],
),
IO.DynamicCombo.Option(
"avatar_iv",
[
IO.Combo.Input(
"avatar",
options=_AVATARS_BY_ENGINE["avatar_iv"],
tooltip="Avatar looks that support the Avatar IV engine.",
),
],
),
IO.DynamicCombo.Option(
"avatar_iii",
[
IO.Combo.Input(
"avatar",
options=_AVATARS_BY_ENGINE["avatar_iii"],
tooltip="Avatar looks that support the Avatar III engine.",
),
],
),
IO.DynamicCombo.Option(
"avatar_v",
[
IO.Combo.Input(
"avatar",
options=_AVATARS_BY_ENGINE["avatar_v"],
tooltip="Avatar looks that support the Avatar V engine.",
),
],
),
],
tooltip="Rendering engine; each choice lists only the avatars that support it. "
"'auto' offers every avatar and picks its best engine (Avatar IV preferred). "
"Avatar V is highest fidelity, Avatar III is the most affordable.",
),
IO.String.Input(
"custom_avatar_id",
default="",
optional=True,
tooltip="Optional HeyGen avatar look ID. When set, overrides the avatar selected above. "
"Any of HeyGen's 3000+ public looks (or your private avatars) can be used.",
),
IO.DynamicCombo.Input(
"speech",
display_name="speech source",
options=[
IO.DynamicCombo.Option(
"script",
[
IO.String.Input(
"text",
multiline=True,
default="",
tooltip="Text for the avatar to speak (up to 5000 characters). "
"The generated speech must be at least 1 second long.",
),
IO.Combo.Input(
"voice",
options=[_DEFAULT_VOICE_OPTION] + HEYGEN_VOICE_GENERAL_OPTIONS,
tooltip="Voice for the script. The default option uses the voice HeyGen assigned to the avatar.",
),
IO.String.Input(
"custom_voice_id",
default="",
optional=True,
tooltip="Optional HeyGen voice ID. When set, overrides the voice selected above. "
"Any voice from HeyGen's library (2000+) can be used.",
),
IO.Float.Input(
"voice_speed",
default=1.0,
min=0.5,
max=1.5,
step=0.05,
optional=True,
tooltip="Speech speed multiplier.",
),
],
),
IO.DynamicCombo.Option(
"audio",
[
IO.Audio.Input(
"audio",
tooltip="Audio for the avatar to lip-sync, up to 10 minutes.",
),
],
),
],
tooltip="Drive the avatar with a text script (HeyGen text-to-speech) or your own audio.",
),
IO.Combo.Input(
"resolution",
options=["720p", "1080p"],
default="1080p",
optional=True,
tooltip="Output video resolution.",
),
IO.Combo.Input(
"aspect_ratio",
options=["auto", "16:9", "9:16", "1:1", "4:5", "5:4"],
default="auto",
optional=True,
tooltip="Output aspect ratio. 'auto' follows the avatar's source footage.",
),
IO.String.Input(
"background_color",
default="",
optional=True,
tooltip="Optional solid background color as a hex code (e.g. '#00ff00'). "
"Leave empty for the avatar's own background.",
),
IO.Int.Input(
"seed",
default=42,
min=0,
max=2147483647,
control_after_generate=True,
optional=True,
tooltip="Not sent to HeyGen; change it to force a re-run.",
),
],
outputs=[IO.Video.Output()],
hidden=[
IO.Hidden.auth_token_comfy_org,
IO.Hidden.api_key_comfy_org,
IO.Hidden.unique_id,
],
is_api_node=True,
price_badge=IO.PriceBadge(
depends_on=IO.PriceBadgeDepends(widgets=["engine"]),
expr="""
widgets.engine = "avatar_iii"
? {"type":"range_usd","min_usd":0.023881,"max_usd":0.061919,"format":{"suffix":"/second"}}
: widgets.engine = "avatar_v"
? {"type":"usd","usd":0.095381,"format":{"suffix":"/second"}}
: widgets.engine = "avatar_iv"
? {"type":"range_usd","min_usd":0.0715,"max_usd":0.095381,"format":{"suffix":"/second"}}
: {"type":"range_usd","min_usd":0.023881,"max_usd":0.095381,"format":{"suffix":"/second"}}
""",
),
)
@classmethod
async def execute(
cls,
engine: dict,
speech: dict,
custom_avatar_id: str = "",
resolution: str = "1080p",
aspect_ratio: str = "auto",
background_color: str = "",
seed: int = 0,
) -> IO.NodeOutput:
avatar_id, engine_type = await _resolve_avatar(cls, engine["avatar"], custom_avatar_id, engine["engine"])
payload = {
"type": "avatar",
"avatar_id": avatar_id,
"resolution": resolution,
"aspect_ratio": aspect_ratio,
"title": "ComfyUI Avatar Video",
}
if engine_type:
payload["engine"] = {"type": engine_type}
background_color = background_color.strip()
if background_color:
if not background_color.startswith("#"):
raise ValueError("background_color must be a hex color code like '#00ff00'.")
payload["background"] = {"type": "color", "value": background_color}
await _apply_speech_source(cls, payload, speech, require_voice=False)
video = await _create_and_poll_video(cls, payload)
return IO.NodeOutput(await download_url_to_video_output(video["video_url"]))
class HeyGenCreateAvatarNode(IO.ComfyNode):
"""Create a reusable HeyGen avatar from a photo or a text prompt."""
@classmethod
def define_schema(cls) -> IO.Schema:
return IO.Schema(
node_id="HeyGenCreateAvatarNode",
display_name="HeyGen Create Avatar",
category="partner/video/HeyGen",
description="Create your own reusable HeyGen avatar from a photo of a person or "
"from a text prompt (a generated character). Feed the resulting avatar_id into "
"HeyGen Avatar Video's custom_avatar_id — and save the ID somewhere to reuse the "
"avatar in future workflows.",
inputs=[
IO.DynamicCombo.Input(
"source",
options=[
IO.DynamicCombo.Option(
"prompt",
[
IO.String.Input(
"prompt",
multiline=True,
default="",
tooltip="Description of the avatar to generate (up to 1000 characters).",
),
IO.Autogrow.Input(
"reference_images",
template=IO.Autogrow.TemplateNames(
IO.Image.Input("ref_image"),
names=[f"ref_image_{i}" for i in range(1, 4)],
min=0,
),
tooltip="Up to 3 reference images guiding the generated look.",
),
],
),
IO.DynamicCombo.Option(
"photo",
[
IO.Image.Input(
"identity_photo",
tooltip="Photo of the person to turn into an avatar. "
"Downscaled automatically if larger than 2K.",
),
],
),
],
tooltip="Generate a new character from a text prompt, or create the avatar "
"from a connected photo of a person.",
),
],
outputs=[
IO.String.Output(
display_name="avatar_id",
tooltip="Avatar look ID. Pass it to HeyGen Avatar Video's custom_avatar_id; "
"save it to reuse the avatar later.",
),
IO.Image.Output(display_name="preview"),
],
hidden=[
IO.Hidden.auth_token_comfy_org,
IO.Hidden.api_key_comfy_org,
IO.Hidden.unique_id,
],
is_api_node=True,
price_badge=IO.PriceBadge(
expr="""{"type":"usd","usd":1.43}""",
),
)
@classmethod
async def execute(
cls,
source: dict,
) -> IO.NodeOutput:
payload: dict = {"name": "ComfyUI Avatar"}
if source["source"] == "photo":
image = downscale_image_tensor_by_max_side(source["identity_photo"], max_side=2000)
image_url = await upload_image_to_comfyapi(cls, image, mime_type="image/png", total_pixels=None)
payload["type"] = "photo"
payload["file"] = {"type": "url", "url": image_url}
else:
validate_string(source["prompt"], strip_whitespace=True, min_length=1, max_length=1000)
payload["type"] = "prompt"
payload["prompt"] = source["prompt"]
ref_tensors = [t for t in (source.get("reference_images") or {}).values() if t is not None]
if ref_tensors:
n_images = sum(get_number_of_images(t) for t in ref_tensors)
if n_images > 3:
raise ValueError(f"HeyGen accepts at most 3 reference images; got {n_images}.")
scaled = [downscale_image_tensor_by_max_side(t, max_side=2000) for t in ref_tensors]
ref_urls = await upload_images_to_comfyapi(
cls, scaled, max_images=3, mime_type="image/png", total_pixels=None
)
payload["reference_images"] = [{"type": "url", "url": u} for u in ref_urls]
created = await sync_op_raw(
cls,
ApiEndpoint(path=_AVATARS_PATH, method="POST"),
data=payload,
)
look_id = ((created.get("data") or {}).get("avatar_item") or {}).get("id")
if not look_id:
raise ValueError(f"HeyGen did not return an avatar: {created}")
final = await poll_op_raw(
cls,
ApiEndpoint(path=f"{_LOOKS_PATH}/{look_id}"),
# A missing status means the look needed no training and is ready.
status_extractor=lambda r: (r.get("data") or {}).get("status") or "completed",
failed_statuses=["failed", "pending_consent"],
poll_interval=5.0,
)
data = final["data"]
if data.get("preview_image_url"):
preview = await download_url_to_image_tensor(data["preview_image_url"])
else:
preview = torch.zeros(1, 64, 64, 3)
PromptServer.instance.send_progress_text(
f"Please save the avatar_id for reuse.\n\navatar_id: {look_id}",
cls.hidden.unique_id,
)
return IO.NodeOutput(look_id, preview)
class HeyGenVideoTranslateNode(IO.ComfyNode):
"""Translate a spoken video into another language with voice cloning and lip sync."""
@classmethod
def define_schema(cls) -> IO.Schema:
return IO.Schema(
node_id="HeyGenVideoTranslateNode",
display_name="HeyGen Video Translate",
category="partner/video/HeyGen",
description="Translate a spoken video into another language. Clones the original "
"speaker's voice and re-animates the mouth to match the translated speech.",
inputs=[
IO.Video.Input(
"video",
tooltip="Video with speech to translate.",
),
IO.Combo.Input(
"output_language",
options=HEYGEN_TRANSLATE_LANGUAGES,
tooltip="Target language for the translated video.",
),
IO.Combo.Input(
"mode",
options=["speed", "precision"],
default="speed",
tooltip="'speed' is faster; 'precision' produces higher-quality lip sync at twice the price.",
),
IO.Boolean.Input(
"translate_audio_only",
default=False,
optional=True,
tooltip="Only swap the audio track, keeping the original mouth movements (no lip sync).",
),
IO.Int.Input(
"speaker_count",
default=0,
min=0,
max=10,
optional=True,
tooltip="Number of speakers in the video. 0 = detect automatically.",
),
IO.Int.Input(
"seed",
default=42,
min=0,
max=2147483647,
control_after_generate=True,
optional=True,
tooltip="Not sent to HeyGen; change it to force a re-run.",
),
],
outputs=[IO.Video.Output()],
hidden=[
IO.Hidden.auth_token_comfy_org,
IO.Hidden.api_key_comfy_org,
IO.Hidden.unique_id,
],
is_api_node=True,
price_badge=IO.PriceBadge(
depends_on=IO.PriceBadgeDepends(widgets=["mode"]),
expr="""{"type":"usd","usd": widgets.mode = "precision" ? 0.095381 : 0.047619,"""
""""format":{"suffix":"/second"}}""",
),
)
@classmethod
async def execute(
cls,
video: Input.Video,
output_language: str,
mode: str,
translate_audio_only: bool = False,
speaker_count: int = 0,
seed: int = 0,
) -> IO.NodeOutput:
video_url = await upload_video_to_comfyapi(cls, video)
payload = {
"video": {"type": "url", "url": video_url},
"output_languages": [output_language],
"mode": mode,
"translate_audio_only": translate_audio_only,
"title": "ComfyUI Video Translate",
}
if speaker_count > 0:
payload["speaker_num"] = speaker_count
created = await sync_op_raw(
cls,
ApiEndpoint(path=_TRANSLATIONS_PATH, method="POST"),
data=payload,
)
translation_ids = (created.get("data") or {}).get("video_translation_ids") or []
if not translation_ids:
raise ValueError(f"HeyGen did not return a translation ID: {created}")
final = await poll_op_raw(
cls,
ApiEndpoint(path=f"{_TRANSLATIONS_PATH}/{translation_ids[0]}"),
status_extractor=lambda r: (r.get("data") or {}).get("status"),
queued_statuses=["pending"],
poll_interval=5.0,
)
data = final["data"]
if not data.get("video_url"):
raise ValueError(f"HeyGen returned no video_url for translation {translation_ids[0]}.")
return IO.NodeOutput(await download_url_to_video_output(data["video_url"]))
class HeyGenTextToSpeechNode(IO.ComfyNode):
"""Synthesize speech audio from text with HeyGen's Starfish TTS engine."""
@classmethod
def define_schema(cls) -> IO.Schema:
return IO.Schema(
node_id="HeyGenTextToSpeechNode",
display_name="HeyGen Text to Speech",
category="partner/audio/HeyGen",
description="Generate speech audio from text using HeyGen's Starfish TTS engine. "
"Includes HeyGen's most popular voices across 17 languages.",
inputs=[
IO.String.Input(
"text",
multiline=True,
default="",
tooltip="Text to synthesize (up to 5000 characters). The generated speech "
"must be at least 1 second long.",
),
IO.Combo.Input(
"voice",
options=HEYGEN_VOICE_TTS_OPTIONS,
tooltip="Voice to use (curated from HeyGen's most popular Starfish-compatible voices).",
),
IO.String.Input(
"custom_voice_id",
default="",
optional=True,
tooltip="Optional HeyGen voice ID. When set, overrides the voice selected above. "
"The voice must support the Starfish engine.",
),
IO.Float.Input(
"speed",
default=1.0,
min=0.5,
max=2.0,
step=0.05,
optional=True,
tooltip="Speech speed multiplier.",
),
IO.Boolean.Input(
"ssml",
default=False,
optional=True,
tooltip="Treat the text as SSML markup (for pauses, emphasis, and pronunciation control).",
),
IO.Int.Input(
"seed",
default=42,
min=0,
max=2147483647,
control_after_generate=True,
optional=True,
tooltip="Not sent to HeyGen; change it to force a re-run.",
),
],
outputs=[IO.Audio.Output()],
hidden=[
IO.Hidden.auth_token_comfy_org,
IO.Hidden.api_key_comfy_org,
IO.Hidden.unique_id,
],
is_api_node=True,
price_badge=IO.PriceBadge(
expr="""{"type":"usd","usd":0.00095381,"format":{"approximate":true,"suffix":"/second"}}""",
),
)
@classmethod
async def execute(
cls,
text: str,
voice: str,
custom_voice_id: str = "",
speed: float = 1.0,
ssml: bool = False,
seed: int = 0,
) -> IO.NodeOutput:
validate_string(text, strip_whitespace=True, min_length=1, max_length=5000)
payload = {
"text": text,
"voice_id": custom_voice_id.strip() or HEYGEN_VOICE_TTS_MAP[voice],
"speed": round(speed, 2),
}
if ssml:
payload["input_type"] = "ssml"
response = await sync_op_raw(
cls,
ApiEndpoint(path=_SPEECH_PATH, method="POST"),
data=payload,
)
audio_url = (response.get("data") or {}).get("audio_url")
if not audio_url:
raise ValueError(f"HeyGen did not return an audio_url: {response}")
audio_bytes = await download_url_as_bytesio(audio_url)
return IO.NodeOutput(audio_bytes_to_audio_input(audio_bytes.getvalue()))
class HeyGenExtension(ComfyExtension):
@override
async def get_node_list(self) -> list[type[IO.ComfyNode]]:
return [
HeyGenTalkingPhotoNode,
HeyGenAvatarVideoNode,
HeyGenCreateAvatarNode,
HeyGenVideoTranslateNode,
HeyGenTextToSpeechNode,
]
async def comfy_entrypoint() -> HeyGenExtension:
return HeyGenExtension()
-18
View File
@@ -41,9 +41,6 @@ STARTING_POINT_ID_PATTERN = r"<starting_point_id:(.*)>"
class SupportedOpenAIModel(str, Enum):
gpt_5_6_sol = "gpt-5.6-sol"
gpt_5_6_terra = "gpt-5.6-terra"
gpt_5_6_luna = "gpt-5.6-luna"
gpt_5_5_pro = "gpt-5.5-pro"
gpt_5_5 = "gpt-5.5"
gpt_5 = "gpt-5"
@@ -1066,21 +1063,6 @@ class OpenAIChatNode(IO.ComfyNode):
"usd": [0.002, 0.008],
"format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" }
}
: $contains($m, "gpt-5.6-terra") ? {
"type": "list_usd",
"usd": [0.0025, 0.015],
"format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" }
}
: $contains($m, "gpt-5.6-luna") ? {
"type": "list_usd",
"usd": [0.001, 0.006],
"format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" }
}
: $contains($m, "gpt-5.6") ? {
"type": "list_usd",
"usd": [0.005, 0.03],
"format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" }
}
: $contains($m, "gpt-5.5-pro") ? {
"type": "list_usd",
"usd": [0.03, 0.18],
-391
View File
@@ -1,391 +0,0 @@
from typing_extensions import override
from comfy_api.latest import IO, ComfyExtension, Input
from comfy_api_nodes.apis.sync_so import (
SyncActiveSpeakerDetection,
SyncGeneration,
SyncGenerationOptions,
SyncGenerationRequest,
SyncInputItem,
)
from comfy_api_nodes.util import (
ApiEndpoint,
download_url_to_video_output,
downscale_image_tensor,
downscale_image_tensor_by_max_side,
get_image_dimensions,
get_number_of_images,
poll_op,
sync_op,
upload_audio_to_comfyapi,
upload_image_to_comfyapi,
upload_video_to_comfyapi,
validate_audio_duration,
)
class SyncLipSyncNode(IO.ComfyNode):
@classmethod
def define_schema(cls) -> IO.Schema:
return IO.Schema(
node_id="SyncLipSyncNode",
display_name="sync.so Lip Sync",
category="partner/video/sync.so",
description=(
"Re-sync mouth movement in a video to new speech audio using sync.so. "
"Handles close-ups, profiles and obstructions automatically while preserving "
"the speaker's expression. Cost scales with output duration."
),
inputs=[
IO.Video.Input(
"video",
tooltip="Footage of the speaker to re-sync. Up to 4K (4096x2160); "
"a constant frame rate of 24/25/30 fps works best.",
),
IO.Audio.Input(
"audio",
tooltip="Speech audio to sync the mouth to.",
),
IO.Int.Input(
"seed",
default=42,
min=0,
max=2147483647,
control_after_generate=True,
tooltip="Seed controls whether the node should re-run; "
"results are non-deterministic regardless of seed.",
),
IO.DynamicCombo.Input(
"model",
options=[
IO.DynamicCombo.Option(
"sync-3",
[
IO.Combo.Input(
"sync_mode",
options=["bounce", "cut_off", "loop", "silence", "remap"],
default="bounce",
tooltip=(
"How to handle a duration mismatch between video and audio; "
"this also sets the output length. "
"bounce: video plays forward then backward until the audio ends "
"(output = audio length). "
"loop: video restarts until the audio ends (output = audio length). "
"remap: video is time-stretched to match the audio (output = audio length). "
"cut_off: the longer track is trimmed (output = shorter length). "
"silence: nothing is trimmed; the shorter track is padded "
"(output = longer length)."
),
),
IO.Combo.Input(
"speaker_selection",
options=["default", "auto-detect", "coordinates"],
default="default",
tooltip=(
"Which face to lipsync when several people are visible. "
"default: let the model decide. "
"auto-detect: detect and follow the active speaker. "
"coordinates: target the face at pixel (speaker_x, speaker_y) "
"in the frame chosen by speaker_frame."
),
),
IO.Int.Input(
"speaker_frame",
default=0,
min=0,
max=1_000_000,
advanced=True,
tooltip="Video frame used to locate the speaker. "
"Only used when speaker_selection is 'coordinates'.",
),
IO.Int.Input(
"speaker_x",
default=0,
min=0,
max=4096,
advanced=True,
tooltip="X pixel coordinate of the speaker's face. "
"Only used when speaker_selection is 'coordinates'.",
),
IO.Int.Input(
"speaker_y",
default=0,
min=0,
max=4096,
advanced=True,
tooltip="Y pixel coordinate of the speaker's face. "
"Only used when speaker_selection is 'coordinates'.",
),
],
)
],
tooltip="sync.so generation model.",
),
],
outputs=[IO.Video.Output()],
hidden=[
IO.Hidden.auth_token_comfy_org,
IO.Hidden.api_key_comfy_org,
IO.Hidden.unique_id,
],
is_api_node=True,
price_badge=IO.PriceBadge(
expr="""{"type":"usd","usd":0.19019,"format":{"approximate":true,"suffix":"/second"}}""",
),
)
@classmethod
async def execute(
cls,
video: Input.Video,
audio: Input.Audio,
seed: int,
model: dict,
) -> IO.NodeOutput:
try:
width, height = video.get_dimensions()
except Exception:
width = height = None
if width and height and (max(width, height) > 4096 or width * height > 4096 * 2160):
raise ValueError(
f"sync.so rejects videos above 4K (4096x2160); got {width}x{height}. Downscale the video first."
)
validate_audio_duration(audio, max_duration=600)
if model["speaker_selection"] == "auto-detect":
speaker_detection = SyncActiveSpeakerDetection(auto_detect=True)
elif model["speaker_selection"] == "coordinates":
speaker_detection = SyncActiveSpeakerDetection(
frame_number=model["speaker_frame"],
coordinates=[model["speaker_x"], model["speaker_y"]],
)
else:
speaker_detection = None
video_url = await upload_video_to_comfyapi(cls, video, max_duration=600)
audio_url = await upload_audio_to_comfyapi(cls, audio)
generation = await sync_op(
cls,
ApiEndpoint(path="/proxy/synclabs/v2/generate", method="POST"),
response_model=SyncGeneration,
data=SyncGenerationRequest(
model=model["model"],
input=[
SyncInputItem(type="video", url=video_url),
SyncInputItem(type="audio", url=audio_url),
],
options=SyncGenerationOptions(
sync_mode=model["sync_mode"],
active_speaker_detection=speaker_detection,
),
),
)
generation = await poll_op(
cls,
ApiEndpoint(path=f"/proxy/synclabs/v2/generate/{generation.id}"),
response_model=SyncGeneration,
status_extractor=lambda g: g.status,
completed_statuses=["COMPLETED", "FAILED", "REJECTED"],
failed_statuses=[],
queued_statuses=["PENDING"],
poll_interval=10.0,
)
if generation.status != "COMPLETED":
code = f" [{generation.errorCode}]" if generation.errorCode else ""
raise ValueError(
f"sync.so generation {generation.status.lower()}{code}: "
f"{generation.error or 'no error details provided'}"
)
if not generation.outputUrl:
raise ValueError("sync.so generation completed but no output URL was returned.")
return IO.NodeOutput(await download_url_to_video_output(generation.outputUrl))
class SyncTalkingImageNode(IO.ComfyNode):
@classmethod
def define_schema(cls) -> IO.Schema:
return IO.Schema(
node_id="SyncTalkingImageNode",
display_name="sync.so Talking Image",
category="partner/video/sync.so",
description=(
"Animate a still portrait into a talking video driven by speech audio, "
"using sync.so's sync-3 model. The output duration matches the audio. "
"Cost scales with output duration."
),
inputs=[
IO.Image.Input(
"image",
tooltip="A single image with a clearly visible face, up to 4K (4096x2160).",
),
IO.Audio.Input(
"audio",
tooltip="Speech audio driving the talking video; the output duration matches it. "
"Chain any TTS node here to drive the animation from text.",
),
IO.String.Input(
"prompt",
multiline=True,
default="",
tooltip="Optional guidance for how the portrait comes to life, e.g. "
"'make the subject smile and look at the camera'. "
"Leave empty for natural talking motion.",
),
IO.Int.Input(
"seed",
default=0,
min=0,
max=2147483647,
control_after_generate=True,
tooltip="Seed controls whether the node should re-run; "
"results are non-deterministic regardless of seed.",
),
IO.DynamicCombo.Input(
"model",
options=[
IO.DynamicCombo.Option(
"sync-3",
[
IO.Combo.Input(
"speaker_selection",
options=["default", "coordinates"],
default="default",
tooltip=(
"Which face to animate when several people are visible. "
"default: let the model decide. "
"coordinates: target the face at pixel (speaker_x, speaker_y) "
"in the image. Auto-detection is not supported for images."
),
),
IO.Int.Input(
"speaker_x",
default=0,
min=0,
max=4096,
advanced=True,
tooltip="X pixel coordinate of the speaker's face. "
"Only used when speaker_selection is 'coordinates'.",
),
IO.Int.Input(
"speaker_y",
default=0,
min=0,
max=4096,
advanced=True,
tooltip="Y pixel coordinate of the speaker's face. "
"Only used when speaker_selection is 'coordinates'.",
),
IO.Boolean.Input(
"auto_downscale",
default=True,
advanced=True,
tooltip="Automatically downscale the image if it exceeds the 4K "
"(4096x2160) input limit; speaker coordinates are scaled to match. "
"When disabled, an oversized image raises an error instead.",
),
],
)
],
tooltip="sync.so generation model. Image input is exclusive to sync-3.",
),
],
outputs=[IO.Video.Output()],
hidden=[
IO.Hidden.auth_token_comfy_org,
IO.Hidden.api_key_comfy_org,
IO.Hidden.unique_id,
],
is_api_node=True,
price_badge=IO.PriceBadge(
expr="""{"type":"usd","usd":0.19019,"format":{"approximate":true,"suffix":"/second"}}""",
),
)
@classmethod
async def execute(
cls,
image: Input.Image,
audio: Input.Audio,
prompt: str,
seed: int,
model: dict,
) -> IO.NodeOutput:
if get_number_of_images(image) != 1:
raise ValueError("Exactly one image is required; got a batch. Pick one frame first.")
validate_audio_duration(audio, max_duration=600)
height, width = get_image_dimensions(image)
speaker_x, speaker_y = model["speaker_x"], model["speaker_y"]
if max(width, height) > 4096 or width * height > 4096 * 2160:
if not model["auto_downscale"]:
raise ValueError(
f"sync.so rejects images above 4K (4096x2160); got {width}x{height}. "
"Downscale the image first or enable auto_downscale."
)
image = downscale_image_tensor(image, total_pixels=4096 * 2160)
image = downscale_image_tensor_by_max_side(image, max_side=4096)
new_height, new_width = get_image_dimensions(image)
# speaker coordinates are given in the original image's pixel space
speaker_x = min(new_width - 1, round(speaker_x * new_width / width))
speaker_y = min(new_height - 1, round(speaker_y * new_height / height))
if model["speaker_selection"] == "coordinates":
speaker_detection = SyncActiveSpeakerDetection(
frame_number=0, # images have a single frame; auto_detect is rejected by the API
coordinates=[speaker_x, speaker_y],
)
else:
speaker_detection = None
image_url = await upload_image_to_comfyapi(cls, image, mime_type="image/png", total_pixels=None)
audio_url = await upload_audio_to_comfyapi(cls, audio)
generation = await sync_op(
cls,
ApiEndpoint(path="/proxy/synclabs/v2/generate", method="POST"),
response_model=SyncGeneration,
data=SyncGenerationRequest(
model=model["model"],
input=[
SyncInputItem(type="image", url=image_url),
SyncInputItem(type="audio", url=audio_url),
],
options=SyncGenerationOptions(
i2v_prompt=prompt.strip() or None,
active_speaker_detection=speaker_detection,
),
),
)
generation = await poll_op(
cls,
ApiEndpoint(path=f"/proxy/synclabs/v2/generate/{generation.id}"),
response_model=SyncGeneration,
status_extractor=lambda g: g.status,
completed_statuses=["COMPLETED", "FAILED", "REJECTED"],
failed_statuses=[],
queued_statuses=["PENDING"],
poll_interval=10.0,
)
if generation.status != "COMPLETED":
code = f" [{generation.errorCode}]" if generation.errorCode else ""
raise ValueError(
f"sync.so generation {generation.status.lower()}{code}: "
f"{generation.error or 'no error details provided'}"
)
if not generation.outputUrl:
raise ValueError("sync.so generation completed but no output URL was returned.")
return IO.NodeOutput(await download_url_to_video_output(generation.outputUrl))
class SyncExtension(ComfyExtension):
@override
async def get_node_list(self) -> list[type[IO.ComfyNode]]:
return [
SyncLipSyncNode,
SyncTalkingImageNode,
]
async def comfy_entrypoint() -> SyncExtension:
return SyncExtension()
+1 -8
View File
@@ -15,8 +15,6 @@ from comfy.comfy_api_env import normalize_comfy_api_base
from comfy.deploy_environment import get_deploy_environment
from comfy.model_management import processing_interrupted
from comfy_api.latest import IO
from comfy_execution.utils import get_executing_context
from comfyui_version import __version__ as comfyui_version
from .common_exceptions import ProcessingInterrupted
@@ -58,16 +56,11 @@ def get_comfy_api_headers(node_cls: type[IO.ComfyNode]) -> dict[str, str]:
relative/cloud URLs resolved against ``default_base_url()``; because the result
includes auth, callers must not attach it to arbitrary absolute/presigned URLs.
"""
headers = {
return {
**get_auth_header(node_cls),
"Comfy-Env": get_deploy_environment(),
"Comfy-Usage-Source": get_usage_source(node_cls),
"Comfy-Core-Version": comfyui_version,
}
ctx = get_executing_context()
if ctx is not None:
headers["Comfy-Job-Id"] = ctx.prompt_id
return headers
def default_base_url() -> str:
+4 -431
View File
@@ -2,7 +2,6 @@ import logging
import os
import json
import av
import numpy as np
import torch
from PIL import Image
@@ -10,7 +9,7 @@ from typing_extensions import override
import folder_paths
import node_helpers
from comfy_api.latest import ComfyExtension, io, Input, InputImpl, Types
from comfy_api.latest import ComfyExtension, io
def load_and_process_images(image_files, input_dir):
@@ -43,38 +42,6 @@ def load_and_process_images(image_files, input_dir):
return output_images
VALID_VIDEO_EXTENSIONS = [".mp4", ".avi", ".mov", ".webm", ".mkv", ".flv"]
def _decode_selected_frames(video: Input.Video, indices: list[int]) -> Input.Video:
"""Decode only the requested frame indices from a video.
Opens the underlying container once, decodes frames in presentation order,
keeps only the ones whose index is in ``indices``, and returns the result
wrapped in a VideoFromComponents so it still satisfies the VideoInput
contract for downstream nodes.
"""
indices_sorted = sorted(set(indices))
max_idx = indices_sorted[-1]
source = video.get_stream_source()
frames_by_idx: dict[int, torch.Tensor] = {}
with av.open(source, mode="r") as container:
stream = container.streams.video[0]
wanted = set(indices_sorted)
for frame_idx, frame in enumerate(container.decode(stream)):
if frame_idx in wanted:
img = frame.to_ndarray(format="rgb24")
frames_by_idx[frame_idx] = torch.from_numpy(img.copy()).float() / 255.0
if frame_idx >= max_idx:
break
stacked = torch.stack([frames_by_idx[i] for i in indices])
return InputImpl.VideoFromComponents(
Types.VideoComponents(images=stacked, frame_rate=video.get_frame_rate())
)
class LoadImageDataSetFromFolderNode(io.ComfyNode):
@classmethod
def define_schema(cls):
@@ -190,116 +157,6 @@ class LoadImageTextDataSetFromFolderNode(io.ComfyNode):
return io.NodeOutput(output_tensor, captions)
class LoadVideoDataSetFromFolderNode(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="LoadVideoDataSetFromFolder",
search_aliases=["load folder", "load from folder", "load dataset", "load videos", "import dataset"],
display_name="Load Video (from Folder)",
category="video",
description="Load a dataset of videos from a specified folder and return a list of videos. Supported formats: MP4, AVI, MOV, WEBM, MKV, FLV.",
is_experimental=True,
inputs=[
io.Combo.Input(
"folder",
options=folder_paths.get_input_subfolders(),
tooltip="The folder containing video files.",
),
],
outputs=[
io.Video.Output(
display_name="videos",
is_output_list=True,
tooltip="Lazy video references; frames are decoded only when needed downstream.",
),
],
)
@classmethod
def execute(cls, folder):
sub_input_dir = os.path.join(folder_paths.get_input_directory(), folder)
video_files = sorted([
f for f in os.listdir(sub_input_dir)
if any(f.lower().endswith(ext) for ext in VALID_VIDEO_EXTENSIONS)
])
if not video_files:
raise ValueError(f"No video files found in {sub_input_dir}")
videos = [InputImpl.VideoFromFile(os.path.join(sub_input_dir, f)) for f in video_files]
logging.info(f"Loaded {len(videos)} lazy video references from {sub_input_dir}")
return io.NodeOutput(videos)
class LoadVideoTextDataSetFromFolderNode(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="LoadVideoTextDataSetFromFolder",
search_aliases=["load folder", "load from folder", "load dataset", "load videos", "import dataset"],
display_name="Load Video-Text (from Folder)",
category="video",
description="Load a dataset of pairs of videos and text captions from a specified folder and return them as a list. Supported formats: MP4, AVI, MOV, WEBM, MKV, FLV.",
is_experimental=True,
inputs=[
io.Combo.Input(
"folder",
options=folder_paths.get_input_subfolders(),
tooltip="The folder containing video files and .txt captions.",
),
],
outputs=[
io.Video.Output(
display_name="videos",
is_output_list=True,
tooltip="Lazy video references; frames are decoded only when needed downstream.",
),
io.String.Output(
display_name="texts",
is_output_list=True,
tooltip="List of text captions.",
),
],
)
@classmethod
def execute(cls, folder):
sub_input_dir = os.path.join(folder_paths.get_input_directory(), folder)
video_files = []
for item in sorted(os.listdir(sub_input_dir)):
path = os.path.join(sub_input_dir, item)
if any(item.lower().endswith(ext) for ext in VALID_VIDEO_EXTENSIONS):
video_files.append(path)
elif os.path.isdir(path):
# Support kohya-ss/sd-scripts folder structure: {repeat}_{desc}/
repeat = 1
if item.split("_")[0].isdigit():
repeat = int(item.split("_")[0])
video_files.extend([
os.path.join(path, f)
for f in sorted(os.listdir(path))
if any(f.lower().endswith(ext) for ext in VALID_VIDEO_EXTENSIONS)
] * repeat)
if not video_files:
raise ValueError(f"No video files found in {sub_input_dir}")
captions = []
for vf in video_files:
caption_path = os.path.splitext(vf)[0] + ".txt"
if os.path.exists(caption_path):
with open(caption_path, "r", encoding="utf-8") as f:
captions.append(f.read().strip())
else:
captions.append("")
videos = [InputImpl.VideoFromFile(vf) for vf in video_files]
logging.info(f"Loaded {len(videos)} lazy video references with captions from {sub_input_dir}")
return io.NodeOutput(videos, captions)
def save_images_to_folder(image_list, output_dir, prefix="image", overwrite=True):
"""Utility function to save a list of image tensors to disk.
@@ -613,15 +470,7 @@ class ImageProcessingNode(io.ComfyNode):
@classmethod
def execute(cls, images, **kwargs):
"""Execute the node. Routes to _process or _group_process based on mode.
For individual processing (_process), automatically handles multi-frame
inputs (video tensors [T, H, W, C]) by applying _process per-frame and
concatenating the results. This allows all spatial transform nodes to
work with video without modification. Nodes that natively handle batched
tensors (e.g. pure tensor math) can set per_frame_process = False to
skip the per-frame loop.
"""
"""Execute the node. Routes to _process or _group_process based on mode."""
is_group = cls._detect_processing_mode()
if is_group:
@@ -640,16 +489,7 @@ class ImageProcessingNode(io.ComfyNode):
result = cls._group_process(images, **params)
else:
# Individual processing: images is single item, call _process
# Auto-loop over frames for multi-frame inputs (video [T, H, W, C])
# so that PIL-based spatial transforms work per-frame automatically.
if images.shape[0] > 1 and getattr(cls, 'per_frame_process', True):
results = []
for i in range(images.shape[0]):
frame_result = cls._process(images[i:i + 1], **params)
results.append(frame_result)
result = torch.cat(results, dim=0)
else:
result = cls._process(images, **params)
result = cls._process(images, **params)
return io.NodeOutput(result)
@@ -963,7 +803,6 @@ class NormalizeImagesNode(ImageProcessingNode):
display_name = "Normalize Image Colors"
category = "image/color"
description = "Normalize images using mean and standard deviation."
per_frame_process = False # Pure tensor math, handles any batch size
extra_inputs = [
io.Float.Input(
"mean",
@@ -994,7 +833,6 @@ class AdjustBrightnessNode(ImageProcessingNode):
display_name = "Adjust Brightness"
category="image/adjustments"
description = "Adjust the brightness of an image."
per_frame_process = False # Pure tensor math, handles any batch size
extra_inputs = [
io.Float.Input(
"factor",
@@ -1016,7 +854,6 @@ class AdjustContrastNode(ImageProcessingNode):
display_name = "Adjust Contrast"
category="image/adjustments"
description = "Adjust the contrast of an image."
per_frame_process = False # Pure tensor math, handles any batch size
extra_inputs = [
io.Float.Input(
"factor",
@@ -1098,261 +935,6 @@ class ShuffleImageTextDatasetNode(io.ComfyNode):
return io.NodeOutput(shuffled_images, shuffled_texts)
# ========== Video Processing Nodes ==========
class VideoFrameSampleNode(io.ComfyNode):
"""Sample a fixed number of frames from a video using various strategies.
For contiguous strategies ("head"/"tail") the result is a fully lazy
VideoInput (no frames decoded). For non-contiguous strategies
("uniform"/"random") only the selected indices are decoded.
"""
@classmethod
def define_schema(cls):
return io.Schema(
node_id="VideoFrameSample",
search_aliases=["sample frames", "extract frames"],
display_name="Sample Video Frame",
category="video",
description="Sample a fixed number of frames from a video using various strategies.",
is_experimental=True,
inputs=[
io.Video.Input("video", tooltip="Input video."),
io.Int.Input(
"num_frames",
default=16,
min=1,
max=9999,
tooltip="Number of frames to sample.",
),
io.Combo.Input(
"strategy",
options=["uniform", "head", "tail", "random"],
default="uniform",
tooltip="uniform: evenly spaced, head: first N, tail: last N, random: random sorted.",
),
io.Int.Input(
"seed",
default=0,
min=0,
max=0xFFFFFFFFFFFFFFFF,
tooltip="Random seed (only used with 'random' strategy).",
),
],
outputs=[
io.Video.Output(display_name="video", tooltip="Sampled video."),
],
)
@classmethod
def execute(cls, video, num_frames, strategy, seed):
total_frames = video.get_frame_count()
num_frames = min(num_frames, total_frames)
fps = float(video.get_frame_rate())
if strategy == "head":
return io.NodeOutput(
video.as_trimmed(0.0, num_frames / fps, strict_duration=False)
)
if strategy == "tail":
start_t = (total_frames - num_frames) / fps
return io.NodeOutput(
video.as_trimmed(start_t, num_frames / fps, strict_duration=False)
)
if strategy == "uniform":
if num_frames == 1:
indices = [total_frames // 2]
else:
indices = [round(i * (total_frames - 1) / (num_frames - 1)) for i in range(num_frames)]
elif strategy == "random":
rng = np.random.RandomState(seed % (2**32 - 1))
indices = sorted(rng.choice(total_frames, size=num_frames, replace=False).tolist())
else:
raise ValueError(f"Unknown strategy: {strategy}")
return io.NodeOutput(_decode_selected_frames(video, indices))
class VideoTemporalCropNode(io.ComfyNode):
"""Crop a continuous range of frames from a video (fully lazy)."""
@classmethod
def define_schema(cls):
return io.Schema(
node_id="VideoTemporalCrop",
search_aliases=["crop", "crop video", "temporal crop", "truncate video"],
display_name="Crop Video (Temporal)",
category="video/transform",
description="Crop a continuous range of frames from a video.",
is_experimental=True,
inputs=[
io.Video.Input("video", tooltip="Input video."),
io.Int.Input(
"start_frame",
default=0,
min=0,
max=99999,
tooltip="Starting frame index.",
),
io.Int.Input(
"length",
default=16,
min=1,
max=99999,
tooltip="Number of frames to keep.",
),
],
outputs=[
io.Video.Output(display_name="video", tooltip="Cropped video (lazy)."),
],
)
@classmethod
def execute(cls, video, start_frame, length):
total_frames = video.get_frame_count()
fps = float(video.get_frame_rate())
start_frame = min(start_frame, max(total_frames - 1, 0))
length = min(length, total_frames - start_frame)
return io.NodeOutput(
video.as_trimmed(start_frame / fps, length / fps, strict_duration=False)
)
class VideoRandomTemporalCropNode(io.ComfyNode):
"""Randomly crop a continuous range of frames from a video (fully lazy)."""
@classmethod
def define_schema(cls):
return io.Schema(
node_id="VideoRandomTemporalCrop",
search_aliases=["crop", "crop video", "temporal crop", "truncate video", "random crop"],
display_name="Crop Video (Temporal Random)",
category="video/transform",
description="Randomly crop a continuous range of frames from a video.",
is_experimental=True,
inputs=[
io.Video.Input("video", tooltip="Input video."),
io.Int.Input(
"length",
default=16,
min=1,
max=99999,
tooltip="Number of frames to keep.",
),
io.Int.Input(
"seed",
default=0,
min=0,
max=0xFFFFFFFFFFFFFFFF,
tooltip="Random seed.",
),
],
outputs=[
io.Video.Output(display_name="video", tooltip="Cropped video (lazy)."),
],
)
@classmethod
def execute(cls, video, length, seed):
total_frames = video.get_frame_count()
fps = float(video.get_frame_rate())
length = min(length, total_frames)
max_start = total_frames - length
rng = np.random.RandomState(seed % (2**32 - 1))
start = rng.randint(0, max_start + 1) if max_start > 0 else 0
return io.NodeOutput(
video.as_trimmed(start / fps, length / fps, strict_duration=False)
)
class ShuffleVideoDatasetNode(io.ComfyNode):
"""Randomly shuffle the order of videos in the dataset."""
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ShuffleVideoDataset",
search_aliases=["shuffle", "randomize", "mix"],
display_name="Shuffle Videos List",
category="video/batch",
description="Randomly shuffle the order of videos in a list.",
is_experimental=True,
is_input_list=True,
inputs=[
io.Video.Input("videos", tooltip="List of videos to shuffle."),
io.Int.Input(
"seed", default=0, min=0, max=0xFFFFFFFFFFFFFFFF, tooltip="Random seed."
),
],
outputs=[
io.Video.Output(
display_name="videos",
is_output_list=True,
tooltip="Shuffled videos",
),
],
)
@classmethod
def execute(cls, videos, seed):
seed = seed[0] if isinstance(seed, list) else seed
np.random.seed(seed % (2**32 - 1))
indices = np.random.permutation(len(videos))
return io.NodeOutput([videos[i] for i in indices])
class ShuffleVideoTextDatasetNode(io.ComfyNode):
"""Shuffle videos and their captions together, preserving pairs."""
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ShuffleVideoTextDataset",
search_aliases=["shuffle", "randomize", "mix"],
display_name="Shuffle Pairs of Video-Text",
category="dataset/video",
description="Randomly shuffle the order of pairs of video-text in a list.",
is_experimental=True,
is_input_list=True,
inputs=[
io.Video.Input("videos", tooltip="List of videos to shuffle."),
io.String.Input("texts", tooltip="List of texts to shuffle."),
io.Int.Input(
"seed",
default=0,
min=0,
max=0xFFFFFFFFFFFFFFFF,
tooltip="Random seed.",
),
],
outputs=[
io.Video.Output(
display_name="videos",
is_output_list=True,
tooltip="Shuffled videos",
),
io.String.Output(
display_name="texts",
is_output_list=True,
tooltip="Shuffled texts",
),
],
)
@classmethod
def execute(cls, videos, texts, seed):
seed = seed[0] if isinstance(seed, list) else seed
np.random.seed(seed % (2**32 - 1))
indices = np.random.permutation(len(videos))
return io.NodeOutput(
[videos[i] for i in indices],
[texts[i] for i in indices],
)
# ========== Text Transform Nodes ==========
@@ -2026,10 +1608,7 @@ class DatasetExtension(ComfyExtension):
LoadImageTextDataSetFromFolderNode,
SaveImageDataSetToFolderNode,
SaveImageTextDataSetToFolderNode,
# Video data loading nodes
LoadVideoDataSetFromFolderNode,
LoadVideoTextDataSetFromFolderNode,
# Image transform nodes (auto-handle video via per-frame processing)
# Image transform nodes
ResizeImagesByShorterEdgeNode,
ResizeImagesByLongerEdgeNode,
CenterCropImagesNode,
@@ -2039,12 +1618,6 @@ class DatasetExtension(ComfyExtension):
AdjustContrastNode,
ShuffleDatasetNode,
ShuffleImageTextDatasetNode,
# Video processing nodes (lazy VideoInput in/out)
VideoFrameSampleNode,
VideoTemporalCropNode,
VideoRandomTemporalCropNode,
ShuffleVideoDatasetNode,
ShuffleVideoTextDatasetNode,
# Text transform nodes
TextToLowercaseNode,
TextToUppercaseNode,
+3 -3
View File
@@ -10,7 +10,7 @@ def Fourier_filter(x, scale_low=1.0, scale_high=1.5, freq_cutoff=20):
Apply frequency-dependent scaling to an image tensor using Fourier transforms.
Parameters:
x: Input tensor of shape (..., H, W)
x: Input tensor of shape (B, C, H, W)
scale_low: Scaling factor for low-frequency components (default: 1.0)
scale_high: Scaling factor for high-frequency components (default: 1.5)
freq_cutoff: Number of frequency indices around center to consider as low-frequency (default: 20)
@@ -31,8 +31,8 @@ def Fourier_filter(x, scale_low=1.0, scale_high=1.5, freq_cutoff=20):
# Initialize mask with high-frequency scaling factor
mask = torch.ones(x_freq.shape, device=device) * scale_high
m = mask
for d in range(2):
dim = len(x_freq.shape) - 2 + d
for d in range(len(x_freq.shape) - 2):
dim = d + 2
cc = x_freq.shape[dim] // 2
f_c = min(freq_cutoff, cc)
m = m.narrow(dim, cc - f_c, f_c * 2)
+13 -24
View File
@@ -844,18 +844,15 @@ class ImageMergeTileList(IO.ComfyNode):
# Format specifications
# ---------------------------------------------------------------------------
# Maps (file_format, bit_depth, num_channels) -> (quantization scale, numpy dtype,
# av frame pix_fmt, stream pix_fmt). Keeps the encode path declarative instead of branchy.
# Maps (file_format, bit_depth, has_alpha) -> (numpy dtype scale, av pixel format,
# stream pix_fmt). Keeps the encode path declarative instead of branchy.
_FORMAT_SPECS = {
("png", "8-bit", 1): {"scale": 255.0, "dtype": np.uint8, "frame_fmt": "gray", "stream_fmt": "gray"},
("png", "8-bit", 3): {"scale": 255.0, "dtype": np.uint8, "frame_fmt": "rgb24", "stream_fmt": "rgb24"},
("png", "8-bit", 4): {"scale": 255.0, "dtype": np.uint8, "frame_fmt": "rgba", "stream_fmt": "rgba"},
("png", "16-bit", 1): {"scale": 65535.0, "dtype": np.uint16, "frame_fmt": "gray16le", "stream_fmt": "gray16be"},
("png", "16-bit", 3): {"scale": 65535.0, "dtype": np.uint16, "frame_fmt": "rgb48le", "stream_fmt": "rgb48be"},
("png", "16-bit", 4): {"scale": 65535.0, "dtype": np.uint16, "frame_fmt": "rgba64le", "stream_fmt": "rgba64be"},
("exr", "32-bit float", 1): {"scale": 1.0, "dtype": np.float32, "frame_fmt": "grayf32le", "stream_fmt": "grayf32le"},
("exr", "32-bit float", 3): {"scale": 1.0, "dtype": np.float32, "frame_fmt": "gbrpf32le", "stream_fmt": "gbrpf32le"},
("exr", "32-bit float", 4): {"scale": 1.0, "dtype": np.float32, "frame_fmt": "gbrapf32le", "stream_fmt": "gbrapf32le"},
("png", "8-bit", False): {"scale": 255.0, "dtype": np.uint8, "frame_fmt": "rgb24", "stream_fmt": "rgb24"},
("png", "8-bit", True): {"scale": 255.0, "dtype": np.uint8, "frame_fmt": "rgba", "stream_fmt": "rgba"},
("png", "16-bit", False): {"scale": 65535.0, "dtype": np.uint16, "frame_fmt": "rgb48le", "stream_fmt": "rgb48be"},
("png", "16-bit", True): {"scale": 65535.0, "dtype": np.uint16, "frame_fmt": "rgba64le", "stream_fmt": "rgba64be"},
("exr", "32-bit float", False): {"scale": 1.0, "dtype": np.float32, "frame_fmt": "gbrpf32le", "stream_fmt": "gbrpf32le"},
("exr", "32-bit float", True): {"scale": 1.0, "dtype": np.float32, "frame_fmt": "gbrapf32le", "stream_fmt": "gbrapf32le"},
}
@@ -894,11 +891,10 @@ def hlg_to_linear(t: torch.Tensor) -> torch.Tensor:
return torch.cat([hlg_to_linear(rgb), alpha], dim=-1)
# Piecewise: sqrt branch below 0.5, log branch above.
# Clamp the log branch at the 0.5 branch point (not above it) so the
# unselected lane stays finite in exp() without altering selected values;
# Clamp inside the log branch so negative / out-of-range values don't blow up;
# values above 1.0 are allowed and extrapolate naturally.
low = (t ** 2) / 3.0
high = (torch.exp((t.clamp(min=0.5) - _HLG_C) / _HLG_A) + _HLG_B) / 12.0
high = (torch.exp((t.clamp(min=_HLG_C) - _HLG_C) / _HLG_A) + _HLG_B) / 12.0
return torch.where(t <= 0.5, low, high)
@@ -1091,8 +1087,7 @@ def _encode_image(
bit_depth: str,
colorspace: str,
) -> bytes:
"""Encode a single HxWxC (or channel-less HxW grayscale) tensor to PNG or
EXR bytes in memory. Grayscale is written as single-channel PNG / Y-only EXR.
"""Encode a single HxWxC tensor to PNG or EXR bytes in memory.
For EXR the input is interpreted according to `colorspace` and converted
to scene-linear (EXR's convention) before writing:
@@ -1106,16 +1101,10 @@ def _encode_image(
For PNG, colorspace selection does not modify pixels PNG is delivered
sRGB-encoded and there is no PNG path for wide-gamut HDR in this node.
"""
if img_tensor.ndim == 2:
img_tensor = img_tensor.unsqueeze(-1) # Some nodes emit grayscale as (H, W) with no channel dim, mask-style.
height, width, num_channels = img_tensor.shape
has_alpha = num_channels == 4
spec = _FORMAT_SPECS.get((file_format, bit_depth, num_channels))
if spec is None:
raise ValueError(
f"No {file_format}/{bit_depth} encoder for {num_channels}-channel images: "
"supported channel counts are 1 (grayscale), 3 (RGB) and 4 (RGBA)."
)
spec = _FORMAT_SPECS[(file_format, bit_depth, has_alpha)]
if spec["dtype"] == np.float32:
# EXR path: preserve full range, no clamp.
-102
View File
@@ -1,102 +0,0 @@
from typing_extensions import override
import comfy.utils
import node_helpers
from comfy_api.latest import ComfyExtension, io
# fmt: off
BUCKETS_1024 = [
(512, 1792), (512, 1856), (512, 1920), (512, 1984), (512, 2048),
(576, 1600), (576, 1664), (576, 1728), (576, 1792),
(640, 1472), (640, 1536), (640, 1600),
(704, 1344), (704, 1408), (704, 1472),
(768, 1216), (768, 1280), (768, 1344),
(832, 1152), (832, 1216),
(896, 1088), (896, 1152),
(960, 1024), (960, 1088),
(1024, 960), (1024, 1024),
(1088, 896), (1088, 960),
(1152, 832), (1152, 896),
(1216, 768), (1216, 832),
(1280, 768),
(1344, 704), (1344, 768),
(1408, 704),
(1472, 640), (1472, 704),
(1536, 640),
(1600, 576), (1600, 640),
(1664, 576),
(1728, 576),
(1792, 512), (1792, 576),
(1856, 512),
(1920, 512),
(1984, 512),
(2048, 512),
]
# fmt: on
def _find_best_bucket(height: int, width: int) -> tuple[int, int]:
target_ratio = height / width
return min(BUCKETS_1024, key=lambda hw: abs(hw[0] / hw[1] - target_ratio))
def _resize_reference(image):
if image.shape[0] != 1:
raise ValueError("JoyImage reference inputs must contain one image each")
samples = image.movedim(-1, 1)
bucket_h, bucket_w = _find_best_bucket(samples.shape[2], samples.shape[3])
resized = comfy.utils.common_upscale(samples, bucket_w, bucket_h, "bilinear", "center")
return resized.movedim(1, -1)[:, :, :, :3]
def _encode(clip, prompt, vae, images):
resized_images = [_resize_reference(image) for image in images]
conditioning = clip.encode_from_tokens_scheduled(clip.tokenize(prompt, images=resized_images))
if vae is not None and resized_images:
ref_latents = [vae.encode(image) for image in resized_images]
conditioning = node_helpers.conditioning_set_values(
conditioning, {"reference_latents": ref_latents}, append=True,
)
return conditioning
class TextEncodeJoyImageEdit(io.ComfyNode):
@classmethod
def define_schema(cls):
image_template = io.Autogrow.TemplatePrefix(
io.Image.Input("image"),
prefix="image",
min=0,
max=6,
)
return io.Schema(
node_id="TextEncodeJoyImageEdit",
category="model/conditioning/joyimage",
inputs=[
io.Clip.Input("clip"),
io.String.Input("prompt", multiline=True, dynamic_prompts=True),
io.Vae.Input("vae", optional=True),
io.Autogrow.Input("images", template=image_template, optional=True),
],
outputs=[
io.Conditioning.Output(),
],
)
@classmethod
def execute(cls, clip, prompt, vae=None, images: io.Autogrow.Type = None) -> io.NodeOutput:
images = images or {}
return io.NodeOutput(_encode(clip, prompt, vae, list(images.values())))
class JoyImageExtension(ComfyExtension):
@override
async def get_node_list(self) -> list[type[io.ComfyNode]]:
return [
TextEncodeJoyImageEdit,
]
async def comfy_entrypoint() -> JoyImageExtension:
return JoyImageExtension()
+4 -8
View File
@@ -174,9 +174,8 @@ class Preview3DAdvanced(IO.ComfyNode):
filename = f"preview3d_advanced_{uuid.uuid4().hex}.{model_3d.format}"
model_3d.save_to(os.path.join(folder_paths.get_temp_directory(), filename))
viewport_state = viewport_state if isinstance(viewport_state, dict) else {}
camera_info_input = kwargs.get("camera_info", None)
camera_info = camera_info_input if camera_info_input is not None else viewport_state.get('camera_info')
camera_info = camera_info_input if camera_info_input is not None else viewport_state['camera_info']
model_3d_info_input = kwargs.get("model_3d_info", None)
model_3d_info = model_3d_info_input if model_3d_info_input is not None else viewport_state.get('model_3d_info', [])
return IO.NodeOutput(
@@ -244,9 +243,8 @@ class PreviewGaussianSplat(IO.ComfyNode):
filename = f"preview_splat_{uuid.uuid4().hex}.{model_3d.format}"
model_3d.save_to(os.path.join(folder_paths.get_temp_directory(), filename))
viewport_state = viewport_state if isinstance(viewport_state, dict) else {}
camera_info_input = kwargs.get("camera_info", None)
camera_info = camera_info_input if camera_info_input is not None else viewport_state.get('camera_info')
camera_info = camera_info_input if camera_info_input is not None else viewport_state['camera_info']
model_3d_info_input = kwargs.get("model_3d_info", None)
model_3d_info = model_3d_info_input if model_3d_info_input is not None else viewport_state.get('model_3d_info', [])
return IO.NodeOutput(
@@ -305,9 +303,8 @@ class PreviewPointCloud(IO.ComfyNode):
filename = f"preview_pointcloud_{uuid.uuid4().hex}.{model_3d.format}"
model_3d.save_to(os.path.join(folder_paths.get_temp_directory(), filename))
viewport_state = viewport_state if isinstance(viewport_state, dict) else {}
camera_info_input = kwargs.get("camera_info", None)
camera_info = camera_info_input if camera_info_input is not None else viewport_state.get('camera_info')
camera_info = camera_info_input if camera_info_input is not None else viewport_state['camera_info']
model_3d_info_input = kwargs.get("model_3d_info", None)
model_3d_info = model_3d_info_input if model_3d_info_input is not None else viewport_state.get('model_3d_info', [])
return IO.NodeOutput(
@@ -378,9 +375,8 @@ class Load3DAdvanced(IO.ComfyNode):
file_3d = None
if model_file and model_file != "none":
file_3d = Types.File3D(folder_paths.get_annotated_filepath(model_file))
viewport_state = viewport_state if isinstance(viewport_state, dict) else {}
model_3d_info = viewport_state.get('model_3d_info', [])
return IO.NodeOutput(file_3d, model_3d_info, viewport_state.get('camera_info'), width, height)
return IO.NodeOutput(file_3d, model_3d_info, viewport_state['camera_info'], width, height)
class Load3DExtension(ComfyExtension):
+2 -229
View File
@@ -8,8 +8,6 @@ import comfy.ldm.common_dit
import comfy.latent_formats
import comfy.ldm.lumina.controlnet
import comfy.ldm.supir.supir_modules
import comfy.ldm.anima.lllite
import comfy.ldm.wan.uni3c
from comfy.ldm.wan.model_multitalk import WanMultiTalkAttentionBlock, MultiTalkAudioProjModel
from comfy_api.latest import io
from comfy.ldm.supir.supir_patch import SUPIRPatch
@@ -238,12 +236,10 @@ class ModelPatchLoader:
def load_model_patch(self, name):
model_patch_path = folder_paths.get_full_path_or_raise("model_patches", name)
sd, metadata = comfy.utils.load_torch_file(model_patch_path, safe_load=True, return_metadata=True)
sd = comfy.utils.load_torch_file(model_patch_path, safe_load=True)
dtype = comfy.utils.weight_dtype(sd)
if 'lllite_conditioning1.conv1.weight' in sd:
model = comfy.ldm.anima.lllite.AnimaLLLite(sd, metadata, device=comfy.model_management.unet_offload_device(), dtype=dtype, operations=comfy.ops.manual_cast)
elif 'controlnet_blocks.0.y_rms.weight' in sd:
if 'controlnet_blocks.0.y_rms.weight' in sd:
additional_in_dim = sd["img_in.weight"].shape[1] - 64
model = QwenImageBlockWiseControlNet(additional_in_dim=additional_in_dim, device=comfy.model_management.unet_offload_device(), dtype=dtype, operations=comfy.ops.manual_cast)
elif 'feature_embedder.mid_layer_norm.bias' in sd:
@@ -265,37 +261,6 @@ class ModelPatchLoader:
if torch.count_nonzero(ref_weight) == 0:
config['broken'] = True
model = comfy.ldm.lumina.controlnet.ZImage_Control(device=comfy.model_management.unet_offload_device(), dtype=dtype, operations=comfy.ops.manual_cast, **config)
elif 'controlnet_patch_embedding.weight' in sd: # Uni3C controlnet for Wan
attn_key_replace = {".self_attn.to_q.": ".self_attn.q.",
".self_attn.to_k.": ".self_attn.k.",
".self_attn.to_v.": ".self_attn.v.",
".self_attn.to_out.0.": ".self_attn.o."}
converted_sd = {}
for k, w in sd.items():
for r, rr in attn_key_replace.items():
k = k.replace(r, rr)
converted_sd[k] = w
sd = converted_sd
num_layers = sum(1 for k in sd if k.startswith("proj_out.") and k.endswith(".weight"))
conv_out_dim = sd["controlnet_patch_embedding.weight"].shape[0]
if "proj_in.weight" in sd:
dim = sd["proj_in.weight"].shape[0]
else:
dim = conv_out_dim
model = comfy.ldm.wan.uni3c.WanUni3CControlnet(
in_channels=sd["controlnet_patch_embedding.weight"].shape[1],
conv_out_dim=conv_out_dim,
dim=dim,
ffn_dim=sd["controlnet_blocks.0.ffn.0.bias"].shape[0],
num_layers=num_layers,
time_embed_dim=sd["controlnet_blocks.0.norm1.linear.weight"].shape[1],
out_proj_dim=sd["proj_out.0.weight"].shape[0],
add_channels=sd["controlnet_mask_embedding.mask_proj.0.weight"].shape[1],
mid_channels=sd["controlnet_mask_embedding.mask_proj.0.weight"].shape[0],
device=comfy.model_management.unet_offload_device(),
dtype=dtype,
operations=comfy.ops.manual_cast)
elif "audio_proj.proj1.weight" in sd:
model = MultiTalkModelPatch(
audio_window=5, context_tokens=32, vae_scale=4,
@@ -331,50 +296,6 @@ class ModelPatchLoader:
return (model_patcher,)
class AnimaLLLiteApply:
@classmethod
def INPUT_TYPES(s):
return {"required": {"model": ("MODEL",),
"model_patch": ("MODEL_PATCH",),
"image": ("IMAGE",),
"strength": ("FLOAT", {"default": 1.0, "min": -10.0, "max": 10.0, "step": 0.01}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}),
},
"optional": {"mask": ("MASK",),
}}
RETURN_TYPES = ("MODEL",)
FUNCTION = "apply_patch"
EXPERIMENTAL = True
CATEGORY = "model_patches/anima"
def apply_patch(self, model, model_patch, image, strength, start_percent, end_percent, mask=None):
image = image[..., :3]
if model_patch.model.cond_in_channels == 4 and mask is None:
mask = torch.zeros_like(image[..., 0])
elif model_patch.model.cond_in_channels != 4:
mask = None
model_sampling = model.get_model_object("model_sampling")
sigma_start = float(model_sampling.percent_to_sigma(start_percent))
sigma_end = float(model_sampling.percent_to_sigma(end_percent))
patch = comfy.ldm.anima.lllite.AnimaLLLitePatch(model_patch, image, mask, strength, sigma_start, sigma_end)
model_patched = model.clone()
model_patched.set_model_post_input_patch(patch)
model_patched.set_model_attn1_patch(comfy.ldm.anima.lllite.AnimaLLLiteAttentionPatch(
patch,
{"q": "self_attn_q_proj", "k": "self_attn_k_proj", "v": "self_attn_v_proj"},
))
model_patched.set_model_attn2_patch(comfy.ldm.anima.lllite.AnimaLLLiteAttentionPatch(
patch,
{"q": "cross_attn_q_proj"},
))
model_patched.set_model_patch(comfy.ldm.anima.lllite.AnimaLLLiteMLPPatch(patch), "mlp_patch")
return (model_patched,)
class DiffSynthCnetPatch:
def __init__(self, model_patch, vae, image, strength, mask=None):
self.model_patch = model_patch
@@ -593,150 +514,6 @@ class ZImageFunControlnet(QwenImageDiffsynthControlnet):
CATEGORY = "model/patch/z-image"
class WanUni3CCnetPatch:
def __init__(self, model_patch, render_video, vae, latent_format, strength, sigma_start, sigma_end):
self.model_patch = model_patch
self.render_video = render_video
self.vae = vae
self.latent_format = latent_format
self.strength = strength
self.sigma_start = sigma_start
self.sigma_end = sigma_end
self.prepared_render = None
self.temp_data = None
def encode_render_video(self, target_latent_shape):
t_len, h_len, w_len = target_latent_shape
temporal_compression = self.vae.temporal_compression_decode() or 1
spatial_compression = self.vae.spacial_compression_encode()
target_frames = (t_len - 1) * temporal_compression + 1
target_height = h_len * spatial_compression
target_width = w_len * spatial_compression
frames = self.render_video
if frames.shape[0] > target_frames:
frames = frames[:target_frames]
elif frames.shape[0] < target_frames:
last_frame = frames[-1:].expand(target_frames - frames.shape[0], -1, -1, -1)
frames = torch.cat([frames, last_frame], dim=0)
if frames.shape[1] != target_height or frames.shape[2] != target_width:
frames = comfy.utils.common_upscale(frames.movedim(-1, 1), target_width, target_height, "bilinear", "center").movedim(1, -1)
loaded_models = comfy.model_management.loaded_models(only_currently_used=True)
render_latent = self.vae.encode(frames)
comfy.model_management.load_models_gpu(loaded_models)
return self.latent_format.process_in(render_latent)
def build_controlnet_input(self, x, dtype, samples_per_cond):
# first 20 channels of the model input: noise latent + I2V mask (zero padded for T2V)
hidden = x[:samples_per_cond, :20].to(dtype)
if hidden.shape[1] < 20:
pad_shape = list(hidden.shape)
pad_shape[1] = 20 - hidden.shape[1]
hidden = torch.cat([hidden, torch.zeros(pad_shape, dtype=hidden.dtype, device=hidden.device)], dim=1)
render = self.prepared_render
if render is None or render.shape[2:] != hidden.shape[2:]:
render = self.encode_render_video(hidden.shape[2:])
render = render.to(device=hidden.device, dtype=dtype)
self.prepared_render = render
if render.shape[0] != hidden.shape[0]:
render = render.expand(hidden.shape[0], -1, -1, -1, -1)
return torch.cat([hidden, render], dim=1)
def __call__(self, kwargs):
img = kwargs.get("img")
block_index = kwargs.get("block_index")
transformer_options = kwargs.get("transformer_options", {})
if block_index == 0:
self.temp_data = None
active = True
sigmas = transformer_options.get("sigmas", None)
if sigmas is not None:
sigma = sigmas[0].item()
if sigma > self.sigma_start or sigma < self.sigma_end:
active = False
if active:
x = kwargs.get("x")
# cond and uncond chunks share latents, so we can reuse residuals
num_conds = len(transformer_options.get("cond_or_uncond", [0]))
samples_per_cond = x.shape[0]
if num_conds > 0 and x.shape[0] % num_conds == 0:
samples_per_cond = x.shape[0] // num_conds
temb = kwargs.get("vec")[:samples_per_cond]
if temb.ndim == 3:
temb = temb[:, 0]
model = self.model_patch.model
controlnet_input = self.build_controlnet_input(x, img.dtype, samples_per_cond)
hidden, freqs = model.process_input(controlnet_input)
self.temp_data = (hidden, temb.to(img.dtype), freqs)
num_layers = self.model_patch.model.num_layers
if self.temp_data is not None and block_index < num_layers:
hidden, temb, freqs = self.temp_data
hidden, residual = self.model_patch.model.forward_block(block_index, hidden, temb, freqs)
residual = residual.to(img.dtype) * self.strength
if residual.shape[0] != img.shape[0]:
residual = residual.repeat(img.shape[0] // residual.shape[0], 1, 1)
img_offset = kwargs.get("img_offset", 0)
img[:, img_offset:img_offset + residual.shape[1]] += residual
if block_index >= num_layers - 1:
self.temp_data = None
else:
self.temp_data = (hidden, temb, freqs)
return kwargs
def to(self, device_or_dtype):
if isinstance(device_or_dtype, torch.device):
if self.prepared_render is not None:
self.prepared_render = self.prepared_render.to(device_or_dtype)
self.temp_data = None
return self
def models(self):
return [self.model_patch]
class WanUni3CControlnetApply:
@classmethod
def INPUT_TYPES(s):
return {"required": { "model": ("MODEL",),
"model_patch": ("MODEL_PATCH",),
"vae": ("VAE",),
"render_video": ("IMAGE", {"tooltip": "The guidance video rendered from the camera trajectory, most commonly warped point cloud renders of the input image."}),
"strength": ("FLOAT", {"default": 1.0, "min": -10.0, "max": 10.0, "step": 0.01}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}),
}}
RETURN_TYPES = ("MODEL",)
FUNCTION = "apply_patch"
EXPERIMENTAL = True
CATEGORY = "model/patch/wan"
def apply_patch(self, model, model_patch, vae, render_video, strength, start_percent, end_percent):
if not isinstance(model_patch.model, comfy.ldm.wan.uni3c.WanUni3CControlnet):
raise ValueError("The connected model patch is not a Uni3C ControlNet.")
cnet_dim = model_patch.model.controlnet_blocks[0].norm1.linear.in_features
model_dim = getattr(model.get_model_object("diffusion_model"), "dim", None)
if model_dim is None:
raise ValueError("The Uni3C ControlNet only works with Wan models.")
if model_dim != cnet_dim:
raise ValueError("This Uni3C ControlNet expects a Wan model with dim {}, the loaded model has dim {}.".format(cnet_dim, model_dim))
model_patched = model.clone()
model_sampling = model.get_model_object("model_sampling")
sigma_start = model_sampling.percent_to_sigma(start_percent)
sigma_end = model_sampling.percent_to_sigma(end_percent)
latent_format = model.get_model_object("latent_format")
patch = WanUni3CCnetPatch(model_patch, render_video[:, :, :, :3], vae, latent_format, strength, sigma_start, sigma_end)
model_patched.set_model_double_block_patch(patch)
return (model_patched,)
class UsoStyleProjectorPatch:
def __init__(self, model_patch, encoded_image):
self.model_patch = model_patch
@@ -895,18 +672,14 @@ NODE_CLASS_MAPPINGS = {
"ModelPatchLoader": ModelPatchLoader,
"QwenImageDiffsynthControlnet": QwenImageDiffsynthControlnet,
"ZImageFunControlnet": ZImageFunControlnet,
"WanUni3CControlnetApply": WanUni3CControlnetApply,
"USOStyleReference": USOStyleReference,
"SUPIRApply": SUPIRApply,
"AnimaLLLiteApply": AnimaLLLiteApply,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"ModelPatchLoader": "Load Model Patch",
"QwenImageDiffsynthControlnet": "Apply Qwen Image DiffSynth ControlNet",
"ZImageFunControlnet": "Apply Z-Image Fun ControlNet",
"WanUni3CControlnetApply": "Apply Wan Uni3C ControlNet",
"USOStyleReference": "Apply USO Style Reference",
"SUPIRApply": "Apply SUPIR Patch",
"AnimaLLLiteApply": "Apply Anima LLLite",
}
+1 -2
View File
@@ -418,9 +418,8 @@ def _save_file3d_to_output(model_3d: Types.File3D, filename_prefix: str) -> str:
def execute_save_3d_advanced(model_3d, viewport_state, width, height, filename_prefix, kwargs) -> IO.NodeOutput:
model_file = _save_file3d_to_output(model_3d, filename_prefix)
viewport_state = viewport_state if isinstance(viewport_state, dict) else {}
camera_info_input = kwargs.get("camera_info", None)
camera_info = camera_info_input if camera_info_input is not None else viewport_state.get('camera_info')
camera_info = camera_info_input if camera_info_input is not None else viewport_state['camera_info']
model_3d_info_input = kwargs.get("model_3d_info", None)
model_3d_info = model_3d_info_input if model_3d_info_input is not None else viewport_state.get('model_3d_info', [])
return IO.NodeOutput(
+2 -3
View File
@@ -920,11 +920,10 @@ def _run_training_loop(
"""
sigmas = torch.tensor(range(num_images))
noise = comfy_extras.nodes_custom_sampler.Noise_RandomNoise(seed)
ndim = latents[0].ndim
if bucket_mode:
# Use first bucket's first latent as dummy for guider
dummy_latent = latents[0][:1].repeat(num_images, *[1]*(ndim-1))
dummy_latent = latents[0][:1].repeat(num_images, 1, 1, 1)
guider.sample(
noise.generate_noise({"samples": dummy_latent}),
dummy_latent,
@@ -934,7 +933,7 @@ def _run_training_loop(
)
elif multi_res:
# use first latent as dummy latent if multi_res
latents = latents[0].repeat(num_images, *[1]*(ndim-1))
latents = latents[0].repeat(num_images, 1, 1, 1)
guider.sample(
noise.generate_noise({"samples": latents}),
latents,
+1 -1
View File
@@ -1,3 +1,3 @@
# This file is automatically generated by the build process when version is
# updated in pyproject.toml.
__version__ = "0.28.0"
__version__ = "0.27.0"
+2 -2
View File
@@ -426,12 +426,12 @@ def _is_intermediate_output(dynprompt, node_id):
def _send_cached_ui(server, node_id, display_node_id, cached, prompt_id, ui_outputs):
if cached.ui is not None:
ui_outputs[node_id] = cached.ui
if server.client_id is None:
return
cached_ui = cached.ui or {}
server.send_sync("executed", { "node": node_id, "display_node": display_node_id, "output": cached_ui.get("output", None), "prompt_id": prompt_id }, server.client_id)
if cached.ui is not None:
ui_outputs[node_id] = cached.ui
async def execute(server, dynprompt, caches, current_item, extra_data, executed, prompt_id, execution_list, pending_subgraph_results, pending_async_nodes, ui_outputs):
unique_id = current_item
+2 -3
View File
@@ -992,7 +992,7 @@ class CLIPLoader:
@classmethod
def INPUT_TYPES(s):
return {"required": { "clip_name": (folder_paths.get_filename_list("text_encoders"), ),
"type": (["stable_diffusion", "stable_cascade", "sd3", "stable_audio", "mochi", "ltxv", "pixart", "cosmos", "lumina2", "wan", "hidream", "chroma", "ace", "omnigen2", "qwen_image", "hunyuan_image", "flux2", "ovis", "longcat_image", "cogvideox", "lens", "pixeldit", "ideogram4", "boogu", "krea2", "joyimage"], ),
"type": (["stable_diffusion", "stable_cascade", "sd3", "stable_audio", "mochi", "ltxv", "pixart", "cosmos", "lumina2", "wan", "hidream", "chroma", "ace", "omnigen2", "qwen_image", "hunyuan_image", "flux2", "ovis", "longcat_image", "cogvideox", "lens", "pixeldit", "ideogram4", "boogu", "krea2"], ),
},
"optional": {
"device": (["default", "cpu"], {"advanced": True}),
@@ -1002,7 +1002,7 @@ class CLIPLoader:
CATEGORY = "model/loaders"
DESCRIPTION = "Recipes:\nsd: clip-l\nstable cascade: clip-g\nsd3: t5 xxl / clip-g / clip-l\nstable audio: t5 base\nmochi: t5 xxl\ncogvideox: t5 xxl (226-token padding)\ncosmos: old t5 xxl\nlumina2: gemma 2 2B\nwan: umt5 xxl\nhidream: llama-3.1 (Recommend) or t5\nomnigen2: qwen vl 2.5 3B\njoyimage: qwen3-vl 8B\nlens: gpt-oss-20b\npixeldit: gemma 2 2B elm"
DESCRIPTION = "Recipes:\nsd: clip-l\nstable cascade: clip-g\nsd3: t5 xxl / clip-g / clip-l\nstable audio: t5 base\nmochi: t5 xxl\ncogvideox: t5 xxl (226-token padding)\ncosmos: old t5 xxl\nlumina2: gemma 2 2B\nwan: umt5 xxl\nhidream: llama-3.1 (Recommend) or t5\nomnigen2: qwen vl 2.5 3B\nlens: gpt-oss-20b\npixeldit: gemma 2 2B elm"
def load_clip(self, clip_name, type="stable_diffusion", device="default"):
clip_type = getattr(comfy.sd.CLIPType, type.upper(), comfy.sd.CLIPType.STABLE_DIFFUSION)
@@ -2462,7 +2462,6 @@ async def init_builtin_extra_nodes():
"nodes_seedvr.py",
"nodes_context_windows.py",
"nodes_qwen.py",
"nodes_joyimage.py",
"nodes_boogu.py",
"nodes_chroma_radiance.py",
"nodes_pid.py",
+499 -74
View File
@@ -7,18 +7,18 @@ components:
description: Timestamp when the asset was created
format: date-time
type: string
display_name:
description: Display name of the asset. Mirrors name for backwards compatibility.
nullable: true
type: string
file_path:
description: Relative path in global-namespace-root form (e.g. "models/checkpoints/flux.safetensors")
nullable: true
type: string
hash:
description: Blake3 hash of the asset content.
pattern: ^blake3:[a-f0-9]{64}$
type: string
loader_path:
description: The value a loader consumes to load this asset. Null when no loader can resolve the file.
nullable: true
type: string
display_name:
description: Human-facing label for the asset. Not unique.
nullable: true
type: string
id:
description: Unique identifier for the asset
format: uuid
@@ -144,14 +144,6 @@ components:
AssetUpdated:
description: Response returned when an existing asset is successfully updated.
properties:
display_name:
description: Display name of the asset. Mirrors name for backwards compatibility.
nullable: true
type: string
file_path:
description: Relative path in global-namespace-root form (e.g. "models/checkpoints/flux.safetensors")
nullable: true
type: string
hash:
description: Blake3 hash of the asset content.
pattern: ^blake3:[a-f0-9]{64}$
@@ -230,6 +222,89 @@ components:
- base_version
- workflow_json
type: object
DownloadEnqueueRequest:
description: Request body for enqueuing a server-side model download.
properties:
allow_any_extension:
default: false
description: Permit a non-model file extension (default only allows known model extensions).
type: boolean
expected_sha256:
description: Optional hub-provided SHA256 to verify the completed file against (fail-closed).
nullable: true
type: string
model_id:
description: Destination as "<directory>/<filename>", resolving to a registered model folder (e.g. "loras/my_lora.safetensors").
type: string
priority:
default: 0
description: Scheduling priority; higher is admitted first.
type: integer
url:
description: Source URL; must be on the allowlist (host + scheme + extension).
type: string
required:
- url
- model_id
type: object
DownloadStatus:
description: Current state and live progress of a single download.
properties:
bytes_done:
type: integer
created_at:
type: integer
download_id:
format: uuid
type: string
error:
nullable: true
type: string
eta_seconds:
nullable: true
type: number
model_id:
type: string
priority:
type: integer
progress:
description: Fraction in [0,1]; null until total size is known.
nullable: true
type: number
segments:
description: Per-segment progress (segmented downloads only).
items:
properties:
bytes_done:
type: integer
idx:
type: integer
length:
type: integer
type: object
nullable: true
type: array
speed_bps:
nullable: true
type: number
status:
enum:
- queued
- active
- paused
- verifying
- completed
- failed
- cancelled
type: string
total_bytes:
nullable: true
type: integer
updated_at:
type: integer
url:
type: string
type: object
ErrorResponse:
description: Standard error response with a machine-readable code and human-readable message.
properties:
@@ -511,6 +586,27 @@ components:
required:
- history
type: object
DownloadAuthProvider:
description: Per-provider download-auth status. Never includes a token.
properties:
env_key_present:
description: Whether an API key for this provider is set via an environment variable.
type: boolean
logged_in:
description: Whether a stored OAuth token exists for this provider.
type: boolean
login_in_progress:
description: Whether an OAuth login flow is currently awaiting the browser callback.
type: boolean
provider:
description: Provider name (e.g. "huggingface", "civitai").
type: string
required:
- provider
- logged_in
- login_in_progress
- env_key_present
type: object
JobCancelResponse:
description: Response for POST /api/jobs/{job_id}/cancel. Returned on both fresh cancels and idempotent no-ops.
properties:
@@ -530,10 +626,6 @@ components:
description: Job creation timestamp (Unix timestamp in milliseconds)
format: int64
type: integer
execution_end_time:
description: Workflow execution completion timestamp (Unix milliseconds, only present for terminal states)
format: int64
type: integer
execution_error:
allOf:
- $ref: '#/components/schemas/ExecutionError'
@@ -542,10 +634,6 @@ components:
additionalProperties: true
description: Node-level execution metadata (only for terminal states)
type: object
execution_start_time:
description: Workflow execution start timestamp (Unix milliseconds, only present once execution has started)
format: int64
type: integer
execution_status:
additionalProperties: true
description: ComfyUI execution status and timeline (only for terminal states)
@@ -578,12 +666,6 @@ components:
description: Last update timestamp (Unix timestamp in milliseconds)
format: int64
type: integer
user_id:
description: |
ID of the user that owns this job (see the `workspace_id`
description above for why this is always the caller's own id
on a successful response).
type: string
workflow:
additionalProperties: true
description: |
@@ -597,18 +679,6 @@ components:
workflow_id:
description: UUID identifying the workflow graph definition
type: string
workspace_id:
description: |
ID of the workspace that owns this job. A successful (200)
response from this operation is only ever returned for the
caller's own job (see this operation's ownership-scoped
query), so this is always the caller's own workspace —
consumers that also need to correlate this job to its
live-progress broadcast channel (workspace+user scoped; see
the internal common/gateways/broadcast package) can use this
value directly rather than resolving their own identity a
second way.
type: string
required:
- id
- status
@@ -809,6 +879,14 @@ components:
ModelFolder:
description: Represents a folder containing models
properties:
extensions:
description: The folder's registered file-extension allowlist. An empty array means the folder accepts any extension (match-all).
example:
- .ckpt
- .safetensors
items:
type: string
type: array
folders:
description: List of paths where models of this type are stored
example:
@@ -1591,13 +1669,7 @@ paths:
schema:
default: true
type: boolean
- description: |
Filter assets by content hash, in the canonical `blake3:<hex>`
form. Matches regardless of which of this asset store's two
internal hash storage formats the matching row was written
under (the canonical form used by from-hash-created references,
or the raw `<hex>.<ext>`/bare `<hex>` storage key used by direct
uploads) — both represent the same content hash.
- description: Filter assets by exact content hash.
in: query
name: hash
schema:
@@ -1676,7 +1748,7 @@ paths:
format: uuid
type: string
tags:
description: JSON-encoded array of freeform tag strings, e.g. '["models","checkpoint"]'. Common types include "models", "input", "output", and "temp", but any tag can be used in any order.
description: JSON-encoded array of tag strings. For new byte uploads, include exactly one destination role (`input`, `output`, or `models`); `models` uploads also require exactly one `model_type:<folder_name>` tag. Extra tags are stored as labels and do not create path components.
type: string
user_metadata:
description: Custom JSON metadata as a string
@@ -1861,7 +1933,7 @@ paths:
content:
application/json:
schema:
$ref: '#/components/schemas/AssetUpdated'
$ref: '#/components/schemas/Asset'
description: Asset updated successfully
"400":
content:
@@ -2382,6 +2454,377 @@ paths:
summary: Get tag histogram for filtered assets
tags:
- file
/api/download:
get:
description: List all known downloads (queued, active, paused, and terminal) with live progress.
operationId: listDownloads
responses:
"200":
content:
application/json:
schema:
properties:
downloads:
items:
$ref: '#/components/schemas/DownloadStatus'
type: array
type: object
description: List of downloads
summary: List downloads
tags:
- download
/api/download/availability:
post:
description: |
Bulk per-id availability for a set of model_ids declared in a workflow.
Returns whether each model is available on disk, currently downloading
(with progress), or missing, plus whether its URL is on the allowlist.
operationId: getModelsAvailability
requestBody:
content:
application/json:
schema:
properties:
models:
additionalProperties:
type: string
description: Map of "<directory>/<filename>" model_id to its declared source URL.
type: object
type: object
responses:
"200":
content:
application/json:
schema:
properties:
models:
additionalProperties: true
type: object
type: object
description: Per-id availability map
summary: Bulk model availability + status
tags:
- download
/api/download/clear:
post:
description: |
Delete all terminal downloads (completed, failed, cancelled) from history
in one transaction, so the cleared history persists across reloads. Live
downloads (queued, active, paused, verifying) are skipped. Finished model
files on disk are never removed; only leftover .part temp files are cleaned up.
operationId: clearDownloads
responses:
"200":
content:
application/json:
schema:
properties:
deleted:
description: Number of history rows removed.
type: integer
type: object
description: History cleared
summary: Clear terminal downloads from history
tags:
- download
/api/download/auth:
get:
description: Per-provider download-auth status (OAuth login state + env-key presence). Never returns a token.
operationId: getDownloadAuth
responses:
"200":
content:
application/json:
schema:
properties:
providers:
items:
$ref: '#/components/schemas/DownloadAuthProvider'
type: array
type: object
description: Per-provider auth status
summary: Get download auth status
tags:
- download
/api/download/auth/{provider}/login:
post:
description: |
Start an OAuth 2.0 PKCE login for a provider. Spawns a transient loopback
callback server and returns the authorize URL to open in a browser.
operationId: loginDownloadAuth
parameters:
- in: path
name: provider
required: true
schema:
type: string
responses:
"200":
content:
application/json:
schema:
properties:
authorize_url:
type: string
type: object
description: Authorize URL to open in a browser
"400":
content:
application/json:
schema:
$ref: '#/components/schemas/ErrorResponse'
description: Unknown provider or OAuth app not configured
"409":
content:
application/json:
schema:
$ref: '#/components/schemas/ErrorResponse'
description: A login for this provider is already in progress
summary: Start OAuth login for a provider
tags:
- download
/api/download/auth/{provider}/logout:
post:
description: Clear the stored OAuth token (memory + on-disk file) for a provider.
operationId: logoutDownloadAuth
parameters:
- in: path
name: provider
required: true
schema:
type: string
responses:
"200":
content:
application/json:
schema:
properties:
logged_out:
type: boolean
type: object
description: Token cleared
"400":
content:
application/json:
schema:
$ref: '#/components/schemas/ErrorResponse'
description: Unknown provider
summary: Log out a provider
tags:
- download
/api/download/enqueue:
post:
description: |
Enqueue a server-side model download. The URL must be on the allowlist
(host + scheme + extension) and the model_id must be "<directory>/<filename>"
resolving to a registered model folder. Returns immediately; track progress
via GET /api/download/{id} or the "download_progress" websocket event.
operationId: enqueueDownload
requestBody:
content:
application/json:
schema:
$ref: '#/components/schemas/DownloadEnqueueRequest'
responses:
"202":
content:
application/json:
schema:
properties:
accepted:
type: boolean
download_id:
format: uuid
type: string
type: object
description: Download accepted and queued
"400":
content:
application/json:
schema:
$ref: '#/components/schemas/ErrorResponse'
description: Invalid request (bad URL, model_id, or not allowlisted)
"409":
content:
application/json:
schema:
$ref: '#/components/schemas/ErrorResponse'
description: Already on disk or already downloading
summary: Enqueue a model download
tags:
- download
/api/download/{id}:
delete:
description: |
Delete a single terminal download from history so it stays gone across
reloads. Refuses (409) to delete a live download (queued, active, paused,
verifying) — cancel it first. The finished model file on disk is never
removed; only a leftover .part temp file is cleaned up.
operationId: deleteDownload
parameters:
- in: path
name: id
required: true
schema:
format: uuid
type: string
responses:
"200":
content:
application/json:
schema:
properties:
deleted:
type: boolean
type: object
description: Deleted
"404":
content:
application/json:
schema:
$ref: '#/components/schemas/ErrorResponse'
description: No such download
"409":
content:
application/json:
schema:
$ref: '#/components/schemas/ErrorResponse'
description: Download is still in progress
summary: Delete a download from history
tags:
- download
get:
description: Get the current status + progress of a single download.
operationId: getDownloadStatus
parameters:
- in: path
name: id
required: true
schema:
format: uuid
type: string
responses:
"200":
content:
application/json:
schema:
$ref: '#/components/schemas/DownloadStatus'
description: Download status
"404":
content:
application/json:
schema:
$ref: '#/components/schemas/ErrorResponse'
description: No such download
summary: Get download status
tags:
- download
/api/download/{id}/cancel:
post:
description: Cancel a download. The partial file is removed.
operationId: cancelDownload
parameters:
- in: path
name: id
required: true
schema:
format: uuid
type: string
responses:
"200":
content:
application/json:
schema:
properties:
ok:
type: boolean
type: object
description: Cancelled
summary: Cancel a download
tags:
- download
/api/download/{id}/pause:
post:
description: Pause a download. The partial file and per-segment offsets are retained for resume.
operationId: pauseDownload
parameters:
- in: path
name: id
required: true
schema:
format: uuid
type: string
responses:
"200":
content:
application/json:
schema:
properties:
ok:
type: boolean
type: object
description: Paused
summary: Pause a download
tags:
- download
/api/download/{id}/priority:
post:
description: Set a download's scheduling priority. Higher priority is admitted first when a slot frees.
operationId: setDownloadPriority
parameters:
- in: path
name: id
required: true
schema:
format: uuid
type: string
requestBody:
content:
application/json:
schema:
properties:
priority:
type: integer
required:
- priority
type: object
responses:
"200":
content:
application/json:
schema:
properties:
ok:
type: boolean
type: object
description: Priority updated
summary: Set download priority
tags:
- download
/api/download/{id}/resume:
post:
description: Resume a paused (or failed) download from its persisted offsets.
operationId: resumeDownload
parameters:
- in: path
name: id
required: true
schema:
format: uuid
type: string
responses:
"200":
content:
application/json:
schema:
properties:
ok:
type: boolean
type: object
description: Resumed
summary: Resume a download
tags:
- download
/api/embeddings:
get:
description: Returns the list of text-encoder embeddings available on disk.
@@ -2496,29 +2939,15 @@ paths:
schema:
additionalProperties: true
properties:
free_tier_balance:
description: Free-tier job allowance for an authenticated non-paid (FREE-tier) user in the rollout. Absent for paid users and unauthenticated requests. Synthesized from config before a grant row exists so a brand-new user still sees their full allowance.
properties:
allowance:
description: Total free jobs granted for the current period
type: integer
remaining:
description: Free jobs remaining (allowance - used, floored at 0)
type: integer
used:
description: Free jobs consumed so far
type: integer
required:
- allowance
- used
- remaining
type: object
max_upload_size:
description: Maximum upload size in bytes
type: integer
supports_preview_metadata:
description: Whether the server supports preview metadata
type: boolean
supports_model_type_tags:
description: Whether the server supports namespaced model type asset tags
type: boolean
type: object
description: Success
headers:
@@ -3346,12 +3775,6 @@ paths:
schema:
$ref: '#/components/schemas/ErrorResponse'
description: Invalid request parameters
"401":
content:
application/json:
schema:
$ref: '#/components/schemas/ErrorResponse'
description: Unauthorized - Authentication required
"500":
content:
application/json:
@@ -5158,3 +5581,5 @@ tags:
name: queue
- description: Job lifecycle queries
name: job
- description: Model download management
name: download
+1 -1
View File
@@ -1,6 +1,6 @@
[project]
name = "ComfyUI"
version = "0.28.0"
version = "0.27.0"
readme = "README.md"
license = { file = "LICENSE" }
requires-python = ">=3.10"
+4 -4
View File
@@ -1,6 +1,6 @@
comfyui-frontend-package==1.45.21
comfyui-workflow-templates==0.11.12
comfyui-embedded-docs==0.5.8
comfyui-frontend-package==1.45.20
comfyui-workflow-templates==0.11.6
comfyui-embedded-docs==0.5.7
torch
torchsde
torchvision
@@ -22,7 +22,7 @@ alembic
SQLAlchemy>=2.0.0
filelock
av>=16.0.0
comfy-kitchen==0.2.22
comfy-kitchen==0.2.18
comfy-aimdo==0.4.10
requests
simpleeval>=1.0.0
+26
View File
@@ -46,6 +46,8 @@ from app.frontend_management import FrontendManager, parse_version
from comfy_api.internal import _ComfyNodeInternal
from app.assets.seeder import asset_seeder
from app.assets.api.routes import register_assets_routes
from app.model_downloader.api.routes import register_routes as register_model_downloader_routes
from app.model_downloader.manager import DOWNLOAD_MANAGER
from app.assets.services.ingest import register_file_in_place
from app.assets.services.path_utils import get_known_subfolder_tags
from app.assets.services.asset_management import resolve_hash_to_path
@@ -259,6 +261,7 @@ class PromptServer():
else:
register_assets_routes(self.app)
asset_seeder.disable()
register_model_downloader_routes(self.app)
routes = web.RouteTableDef()
self.routes = routes
self.last_node_id = None
@@ -1206,6 +1209,29 @@ class PromptServer():
async def setup(self):
timeout = aiohttp.ClientTimeout(total=None) # no timeout
self.client_session = aiohttp.ClientSession(timeout=timeout)
await self._setup_model_downloader()
async def _setup_model_downloader(self):
"""Start the download manager: push progress over the websocket and
resume any downloads interrupted by a previous run."""
def _notify(download_id: str) -> None:
try:
view = DOWNLOAD_MANAGER.status_sync(download_id)
if view is not None:
# Drop the url field before broadcasting: the redacted URL
# (scheme + host + path) should not leak to every connected
# websocket client. download_id / model_id are sufficient to
# correlate progress on the frontend.
broadcast = {k: v for k, v in view.items() if k != "url"}
self.send_sync("download_progress", broadcast)
except Exception:
logging.debug("download progress notify failed", exc_info=True)
DOWNLOAD_MANAGER.set_notify(_notify)
try:
await DOWNLOAD_MANAGER.start()
except Exception as e:
logging.warning("Failed to start model download manager: %s", e)
def add_routes(self):
self.user_manager.add_routes(self.routes)
+1 -525
View File
@@ -2,12 +2,11 @@ import pytest
import torch
import tempfile
import os
import sys
import av
import io
from fractions import Fraction
from comfy_api.input_impl.video_types import VideoFromFile, VideoFromComponents
from comfy_api.util.video_types import VideoComponents, VideoContainer, VideoCodec
from comfy_api.util.video_types import VideoComponents
from comfy_api.input.basic_types import AudioInput
from av.error import InvalidDataError
@@ -238,526 +237,3 @@ def test_duration_consistency(video_components):
manual_duration = float(components.images.shape[0] / components.frame_rate)
assert duration == pytest.approx(manual_duration)
def create_transcode_source(
width=64, height=64, frames=30, fps=30, audio_streams=1, undecodable_audio=0, rotation=False,
container_format="mov", audio_codec="pcm_s16le",
):
"""Create a temp video that save_to must transcode (mpeg4 video, so codec != h264).
``undecodable_audio`` trailing PCM streams get their fourcc corrupted so no decoder exists
(``codec_context is None``), like the APAC track in iPhone spatial-audio recordings.
``rotation`` patches a 90-degree display matrix into the video track header.
"""
buffer = io.BytesIO()
with av.open(buffer, mode="w", format=container_format) as container:
video_stream = container.add_stream("mpeg4", rate=fps)
video_stream.width = width
video_stream.height = height
video_stream.pix_fmt = "yuv420p"
audio = []
for _ in range(audio_streams + undecodable_audio):
stream = container.add_stream(audio_codec, rate=44100)
stream.sample_rate = 44100
audio.append(stream)
for i in range(frames):
frame = av.VideoFrame.from_ndarray(
torch.full((height, width, 3), (i * 7) % 256, dtype=torch.uint8).numpy(),
format="rgb24",
)
container.mux(video_stream.encode(frame.reformat(format="yuv420p")))
# write audio in 1024-sample frames, like real decoders produce, so the
# per-frame skip/cap logic in the transcode path actually runs
for stream in audio:
for offset in range(0, 44100 * frames // fps, 1024):
n = min(1024, 44100 * frames // fps - offset)
audio_frame = av.AudioFrame.from_ndarray(
torch.zeros(1, n, dtype=torch.int16).numpy(), format="s16", layout="mono"
)
audio_frame.sample_rate = 44100
audio_frame.pts = offset
container.mux(stream.encode(audio_frame))
for stream in [video_stream, *audio]:
container.mux(stream.encode(None))
data = bytearray(buffer.getvalue())
end = len(data)
for _ in range(undecodable_audio):
end = data.rindex(b"sowt", 0, end)
data[end:end + 4] = b"Xpac"
if rotation:
# the 3x3 display matrix sits 40 bytes into the version-0 tkhd payload; first tkhd
# inside moov = video track (search from moov so mdat bytes can't false-match)
matrix_offset = data.index(b"tkhd", data.rindex(b"moov")) + 4 + 40
values = [0, 1 << 16, 0, -(1 << 16), 0, 0, 0, 0, 1 << 30]
data[matrix_offset:matrix_offset + 36] = b"".join(v.to_bytes(4, "big", signed=True) for v in values)
tmp = tempfile.NamedTemporaryFile(suffix=f".{container_format}", delete=False)
tmp.write(bytes(data))
tmp.close()
return tmp.name
def transcode_and_probe(video):
buffer = io.BytesIO()
video.save_to(buffer, format=VideoContainer.MP4, codec=VideoCodec.H264)
buffer.seek(0)
with av.open(buffer) as container:
video_stream = container.streams.video[0]
audio_stream = container.streams.audio[0] if container.streams.audio else None
frames = 0
first_pts = None
for packet in container.demux(video_stream):
for frame in packet.decode():
if first_pts is None:
first_pts = frame.pts
frames += 1
return {
"codec": video_stream.codec_context.name,
"width": video_stream.codec_context.width,
"height": video_stream.codec_context.height,
"frames": frames,
"first_pts": first_pts,
"video_seconds": float(video_stream.duration * video_stream.time_base) if video_stream.duration else None,
"audio_seconds": float(audio_stream.duration * audio_stream.time_base)
if audio_stream and audio_stream.duration else None,
"audio_codecs": [s.codec_context.name for s in container.streams.audio],
}
def test_save_to_transcode_streams_without_buffering_frames():
"""Transcoding must not decode the whole video into memory first (~2 GiB for this source)"""
resource = pytest.importorskip("resource") # no getrusage on Windows
rss_scale = 1 if sys.platform == "darwin" else 1024 # ru_maxrss: bytes on macOS, KiB elsewhere
# ru_maxrss is a lifetime peak: a heavier test running earlier would shrink the measured
# delta and quietly defang this canary, so keep this source the biggest thing in the suite
file_path = create_transcode_source(width=640, height=480, frames=300)
try:
rss_before = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss * rss_scale
result = transcode_and_probe(VideoFromFile(file_path))
rss_delta = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss * rss_scale - rss_before
assert result["codec"] == "h264"
assert result["frames"] == 300
assert rss_delta < 500 * 2**20, f"transcode buffered frames in RAM (peak grew {rss_delta / 2**20:.0f} MiB)"
finally:
os.unlink(file_path)
def test_save_to_transcode_honors_trim_window():
"""start_time/duration trim applies to both video and audio on the streaming path"""
file_path = create_transcode_source(frames=90) # 3s @ 30fps
try:
result = transcode_and_probe(VideoFromFile(file_path, start_time=1, duration=1))
assert result["frames"] == pytest.approx(30, abs=2)
assert result["first_pts"] == 0 # trimmed output is rebased to start at zero
assert result["video_seconds"] == pytest.approx(1.0, abs=0.1)
assert result["audio_seconds"] == pytest.approx(1.0, abs=0.1)
finally:
os.unlink(file_path)
def test_save_to_transcode_keeps_audio_of_sparse_video():
"""Audio that runs ahead of a sparse video track (slideshows, timelapses) must be
kept in full it is only clamped to the video's end, never to the video cursor."""
buffer = io.BytesIO()
with av.open(buffer, mode="w", format="mp4") as container:
video_stream = container.add_stream("mpeg4", rate=30)
video_stream.width = video_stream.height = 64
video_stream.pix_fmt = "yuv420p"
audio_stream = container.add_stream("aac", rate=48000, layout="stereo")
for t in (0, 30, 60): # 3 frames spread over 60 seconds
frame = av.VideoFrame.from_ndarray(
torch.full((64, 64, 3), t * 4, dtype=torch.uint8).numpy(), format="rgb24"
).reformat(format="yuv420p")
frame.pts = t * 15360
frame.time_base = Fraction(1, 15360)
container.mux(video_stream.encode(frame))
container.mux(video_stream.encode(None))
for offset in range(0, 48000 * 60, 1024):
n = min(1024, 48000 * 60 - offset)
audio_frame = av.AudioFrame.from_ndarray(
torch.zeros(2, n, dtype=torch.float32).numpy(), format="fltp", layout="stereo"
)
audio_frame.sample_rate = 48000
audio_frame.pts = offset
audio_frame.time_base = Fraction(1, 48000)
container.mux(audio_stream.encode(audio_frame))
container.mux(audio_stream.encode(None))
buffer.seek(0)
result = transcode_and_probe(VideoFromFile(buffer))
assert result["audio_seconds"] == pytest.approx(60.0, abs=1.0)
def test_save_to_transcode_vfr_audio_covers_video_span():
"""A trim window in the sparse region of a VFR file keeps audio for the true pts span
of the kept frames. Deriving the span as frames/average_rate undercuts it badly: the
average is dominated by the dense region (and can be plain wrong on MediaRecorder files)."""
buffer = io.BytesIO()
with av.open(buffer, mode="w", format="mp4") as container:
video_stream = container.add_stream("mpeg4", rate=30)
video_stream.width = video_stream.height = 64
video_stream.pix_fmt = "yuv420p"
audio_stream = container.add_stream("aac", rate=48000, layout="stereo")
# 10 frames inside the first second, then one every 1.25 s
for i, t in enumerate([x / 10 for x in range(10)] + [1.0, 2.25, 3.5, 4.75]):
frame = av.VideoFrame.from_ndarray(
torch.full((64, 64, 3), (i * 16) % 256, dtype=torch.uint8).numpy(), format="rgb24"
).reformat(format="yuv420p")
frame.pts = int(t * 15360)
frame.time_base = Fraction(1, 15360)
container.mux(video_stream.encode(frame))
container.mux(video_stream.encode(None))
for offset in range(0, 48000 * 6, 1024):
n = min(1024, 48000 * 6 - offset)
audio_frame = av.AudioFrame.from_ndarray(
torch.zeros(2, n, dtype=torch.float32).numpy(), format="fltp", layout="stereo"
)
audio_frame.sample_rate = 48000
audio_frame.pts = offset
audio_frame.time_base = Fraction(1, 48000)
container.mux(audio_stream.encode(audio_frame))
container.mux(audio_stream.encode(None))
buffer.seek(0)
result = transcode_and_probe(VideoFromFile(buffer, start_time=1, duration=5))
# kept frames: 1.0/2.25/3.5/4.75 s -> rebased span 3.75 s + one nominal interval
assert result["frames"] == 4
assert result["audio_seconds"] == pytest.approx(4.0, abs=0.45)
def test_save_to_transcode_trims_audio_in_stream_time_base_units():
"""Matroska audio timestamps tick in 1/1000, not 1/sample_rate; trim and audio timing
must convert through the frame's time base instead of assuming sample units. AAC audio,
because it decodes straight to the encoder's format and hits the resampler passthrough
that keeps the source time base on the frames."""
file_path = create_transcode_source(frames=90, container_format="matroska", audio_codec="aac")
try:
result = transcode_and_probe(VideoFromFile(file_path, start_time=1, duration=1))
assert result["audio_codecs"] == ["aac"]
assert result["video_seconds"] == pytest.approx(1.0, abs=0.1)
assert result["audio_seconds"] == pytest.approx(1.0, abs=0.1)
finally:
os.unlink(file_path)
def test_save_to_transcode_learns_unprobed_audio_params():
"""mpegts is only probed a few seconds deep at open, so an audio stream whose first
packet comes later (live captures where audio kicks in late) still has sample_rate 0
when the transcode starts; the parameters must be learned from the stream itself."""
sample_rate, fps, video_seconds, audio_start = 48000, 30, 13, 12
buffer = io.BytesIO()
with av.open(buffer, mode="w", format="mpegts") as container:
video_stream = container.add_stream("mpeg4", rate=fps)
video_stream.width = video_stream.height = 64
video_stream.pix_fmt = "yuv420p"
audio_stream = container.add_stream("aac", rate=sample_rate, layout="mono")
for i in range(video_seconds * fps):
frame = av.VideoFrame.from_ndarray(
torch.full((64, 64, 3), (i * 7) % 256, dtype=torch.uint8).numpy(), format="rgb24"
)
container.mux(video_stream.encode(frame.reformat(format="yuv420p")))
for offset in range(0, (video_seconds - audio_start) * sample_rate, 1024):
n = min(1024, (video_seconds - audio_start) * sample_rate - offset)
audio_frame = av.AudioFrame.from_ndarray(
torch.zeros(1, n, dtype=torch.float32).numpy(), format="fltp", layout="mono"
)
audio_frame.sample_rate = sample_rate
audio_frame.pts = audio_start * sample_rate + offset
container.mux(audio_stream.encode(audio_frame))
for stream in (video_stream, audio_stream):
container.mux(stream.encode(None))
buffer.seek(0)
with av.open(buffer) as container:
# the scenario requires unprobed parameters; if a future FFmpeg probes deeper,
# push audio_start/video_seconds further out to restore it
assert container.streams.audio[0].codec_context.sample_rate == 0
result = transcode_and_probe(VideoFromFile(buffer))
assert result["frames"] == video_seconds * fps
assert result["audio_codecs"] == ["aac"]
assert result["audio_seconds"] == pytest.approx(1.0, abs=0.1)
buffer.seek(0)
trimmed_before_audio = transcode_and_probe(VideoFromFile(buffer, duration=1))
assert trimmed_before_audio["frames"] == fps
assert trimmed_before_audio["audio_codecs"] == []
assert trimmed_before_audio["audio_seconds"] is None
buffer.seek(0)
trimmed_crossing_audio = transcode_and_probe(VideoFromFile(buffer, start_time=11.5, duration=1))
assert trimmed_crossing_audio["frames"] == fps
assert trimmed_crossing_audio["audio_codecs"] == ["aac"]
assert trimmed_crossing_audio["video_seconds"] == pytest.approx(1.0, abs=0.05)
assert trimmed_crossing_audio["audio_seconds"] == pytest.approx(0.5, abs=0.1)
def test_save_to_transcode_trimmed_fragmented_mp4_keeps_audio():
"""Fragmented mp4 (MediaRecorder, DASH/HLS-derived files) delivers audio well behind
video, so when the trim window's last video frame arrives the audio demuxed so far
does not cover the window yet; the transcode must keep demuxing audio until it does
instead of finalizing on the first audio frame it sees afterwards."""
sample_rate, fps, seconds = 48000, 30, 6
buffer = io.BytesIO()
with av.open(buffer, mode="w", format="mp4", options={"movflags": "frag_keyframe+empty_moov"}) as container:
video_stream = container.add_stream("h264", rate=fps)
video_stream.width = video_stream.height = 64
video_stream.pix_fmt = "yuv420p"
audio_stream = container.add_stream("aac", rate=sample_rate, layout="mono")
next_audio_pts = 0
for i in range(seconds * fps):
frame = av.VideoFrame.from_ndarray(
torch.full((64, 64, 3), (i * 7) % 256, dtype=torch.uint8).numpy(), format="rgb24"
)
container.mux(video_stream.encode(frame.reformat(format="yuv420p")))
while next_audio_pts / sample_rate <= i / fps: # feed audio alongside, like a live pipeline
audio_frame = av.AudioFrame.from_ndarray(
torch.zeros(1, 1024, dtype=torch.float32).numpy(), format="fltp", layout="mono"
)
audio_frame.sample_rate = sample_rate
audio_frame.pts = next_audio_pts
container.mux(audio_stream.encode(audio_frame))
next_audio_pts += 1024
for stream in (video_stream, audio_stream):
container.mux(stream.encode(None))
result = transcode_and_probe(VideoFromFile(buffer, start_time=0.5, duration=1.0))
assert result["video_seconds"] == pytest.approx(1.0, abs=0.05)
assert result["audio_seconds"] == pytest.approx(1.0, abs=0.05)
def test_save_to_transcode_sparse_video_keeps_true_duration():
"""average_rate is not a frame duration: a 3-frame video spanning 60 s averages
0.05 fps, and padding the last frame with 1/average_rate used to extend the
output and the audio kept with it about 20 s past the source span."""
sample_rate = 48000
buffer = io.BytesIO()
with av.open(buffer, mode="w", format="mp4") as container:
video_stream = container.add_stream("mpeg4", rate=30)
video_stream.width = video_stream.height = 64
video_stream.pix_fmt = "yuv420p"
audio_stream = container.add_stream("aac", rate=sample_rate, layout="mono")
for i, second in enumerate((0, 30, 60)):
frame = av.VideoFrame.from_ndarray(
torch.full((64, 64, 3), i * 80, dtype=torch.uint8).numpy(), format="rgb24"
).reformat(format="yuv420p")
frame.pts = second * 30
frame.time_base = Fraction(1, 30)
container.mux(video_stream.encode(frame))
for offset in range(0, 90 * sample_rate, 1024):
n = min(1024, 90 * sample_rate - offset)
audio_frame = av.AudioFrame.from_ndarray(
torch.zeros(1, n, dtype=torch.float32).numpy(), format="fltp", layout="mono"
)
audio_frame.sample_rate = sample_rate
audio_frame.pts = offset
container.mux(audio_stream.encode(audio_frame))
for stream in (video_stream, audio_stream):
container.mux(stream.encode(None))
result = transcode_and_probe(VideoFromFile(buffer))
assert result["frames"] == 3
# the last frame keeps its true stts duration (1/30 s), not 1/average_rate (~20 s)
assert result["video_seconds"] == pytest.approx(60.03, abs=0.05)
assert result["audio_seconds"] == pytest.approx(60.03, abs=0.1)
trimmed = transcode_and_probe(VideoFromFile(buffer, duration=45))
assert trimmed["frames"] == 2
# a kept frame whose source duration crosses the window end is clamped to it
assert trimmed["video_seconds"] == pytest.approx(45.0, abs=0.05)
assert trimmed["audio_seconds"] == pytest.approx(45.0, abs=0.1)
def test_save_to_transcode_clamps_final_pts_to_declared_stream_duration():
"""Some iPhone MOVs report a video stream duration that ends before the final
decoded frame's nominal duration. A transcode must not turn that trailing
timestamp quirk into an extra frame interval compared to the source/remux path."""
fps = 30
buffer = io.BytesIO()
with av.open(buffer, mode="w", format="mp4") as container:
video_stream = container.add_stream("mpeg4", rate=fps)
video_stream.width = video_stream.height = 64
video_stream.pix_fmt = "yuv420p"
for i, pts in enumerate([*range(31), 32]):
frame = av.VideoFrame.from_ndarray(
torch.full((64, 64, 3), (i * 7) % 256, dtype=torch.uint8).numpy(), format="rgb24"
).reformat(format="yuv420p")
frame.pts = pts
frame.time_base = Fraction(1, fps)
container.mux(video_stream.encode(frame))
container.mux(video_stream.encode(None))
class _StreamProxy:
def __init__(self, stream, duration):
self._stream = stream
self.duration = duration
def __getattr__(self, name):
return getattr(self._stream, name)
class _StreamsProxy:
def __init__(self, video_stream):
self.video = [video_stream]
self.audio = []
class _PacketProxy:
def __init__(self, packet, stream):
self._packet = packet
self.stream = stream
def __getattr__(self, name):
return getattr(self._packet, name)
class _ContainerProxy:
def __init__(self, container, stream):
self._container = container
self._stream = stream
self.streams = _StreamsProxy(stream)
def __getattr__(self, name):
return getattr(self._container, name)
def demux(self, *streams):
for packet in self._container.demux(self._stream._stream):
yield _PacketProxy(packet, self._stream)
buffer.seek(0)
output = io.BytesIO()
with av.open(buffer) as container:
real_stream = container.streams.video[0]
declared_duration = 32 * int(round((1 / fps) / real_stream.time_base))
stream = _StreamProxy(real_stream, declared_duration)
VideoFromFile(buffer)._save_transcoded(
_ContainerProxy(container, stream), output, VideoContainer.MP4, VideoCodec.H264, None, 8
)
output.seek(0)
with av.open(output) as container:
video_stream = container.streams.video[0]
frames = [f for p in container.demux(video_stream) for f in p.decode()]
assert len(frames) == 32
assert float(video_stream.duration * video_stream.time_base) == pytest.approx(32 / fps, abs=0.01)
assert float(frames[-1].pts * frames[-1].time_base) == pytest.approx(31 / fps, abs=0.01)
def test_save_to_transcode_irregular_vfr_keeps_span():
"""B-frames reorder packets, and mp4 sample durations follow decode order: the dts
timeline ends before the pts timeline, so an irregular-VFR source's tail holds fell
out of the container (this 20.23 s span used to come out as 15.27 s, and the 10 s
trim as 6.03 s). The transcode encodes without B-frames so every sample keeps its
true display duration."""
durations = [1, 1, 60, 1, 1, 120, 1, 180, 1, 1, 150, 90] # 1/30 s ticks, span 20.2333 s
generator = torch.Generator().manual_seed(7)
buffer = io.BytesIO()
with av.open(buffer, mode="w", format="mp4") as container:
video_stream = container.add_stream("mpeg4", rate=30)
video_stream.width = video_stream.height = 64
video_stream.pix_fmt = "yuv420p"
pts = 0
for duration in durations:
# textured frames, so an encoder with default settings has B-frames to gain from
frame = av.VideoFrame.from_ndarray(
torch.randint(0, 255, (64, 64, 3), generator=generator, dtype=torch.uint8).numpy(),
format="rgb24",
).reformat(format="yuv420p")
frame.pts = pts
frame.time_base = Fraction(1, 30)
pts += duration
for packet in video_stream.encode(frame):
packet.duration = duration # exact stts in the source
container.mux(packet)
container.mux(video_stream.encode(None))
result = transcode_and_probe(VideoFromFile(buffer))
assert result["frames"] == len(durations)
assert result["video_seconds"] == pytest.approx(sum(durations) / 30, abs=0.05)
trimmed = transcode_and_probe(VideoFromFile(buffer, duration=10))
assert trimmed["frames"] == 8 # frames at 12.167 s+ fall outside the window
assert trimmed["video_seconds"] == pytest.approx(10.0, abs=0.05)
def test_save_to_transcode_trim_survives_missing_leading_pts():
"""A trim should survive pts-less kept frames followed by a real-pts frame past the window."""
nulled_frames = 0
class _PacketProxy:
def __init__(self, packet):
self._packet = packet
def __getattr__(self, name):
return getattr(self._packet, name)
@property
def stream(self):
return self._packet.stream
def decode(self):
nonlocal nulled_frames
frames = self._packet.decode()
for frame in frames:
if nulled_frames < 2:
frame.pts = None
nulled_frames += 1
return frames
class _ContainerProxy:
def __init__(self, real):
self._real = real
def __getattr__(self, name):
return getattr(self._real, name)
def demux(self, *streams):
for packet in self._real.demux(*streams):
yield _PacketProxy(packet)
file_path = create_transcode_source(frames=10, audio_streams=0)
try:
buffer = io.BytesIO()
with av.open(file_path) as container:
# 0.05 s window: both pts-less frames are kept (synthesized pts 0 and 512),
# and the first real-pts frame (1024 ticks) already lies past end_pts (768)
VideoFromFile(file_path, duration=0.05)._save_transcoded(
_ContainerProxy(container), buffer, VideoContainer.MP4, VideoCodec.H264, None, 8
)
assert nulled_frames == 2
buffer.seek(0)
with av.open(buffer) as container:
video_stream = container.streams.video[0]
frames = [f for p in container.demux(video_stream) for f in p.decode()]
assert len(frames) == 2
assert float(video_stream.duration * video_stream.time_base) == pytest.approx(2 / 30, abs=0.01)
finally:
os.unlink(file_path)
def test_save_to_transcode_bakes_rotation():
"""A 90-degree display-matrix rotation swaps the output dimensions (portrait video)"""
file_path = create_transcode_source(width=64, height=32, rotation=True)
try:
result = transcode_and_probe(VideoFromFile(file_path))
assert (result["width"], result["height"]) == (32, 64)
assert result["frames"] == 30
finally:
os.unlink(file_path)
def test_save_to_transcode_skips_undecodable_audio():
"""Streaming transcode keeps the decodable audio track and drops undecodable ones;
with no decodable audio at all the output is video-only instead of crashing."""
mixed = all_bad = None
try:
mixed = create_transcode_source(audio_streams=1, undecodable_audio=1)
all_bad = create_transcode_source(audio_streams=0, undecodable_audio=2)
result = transcode_and_probe(VideoFromFile(mixed))
assert result["audio_codecs"] == ["aac"]
assert result["audio_seconds"] == pytest.approx(1.0, abs=0.1)
assert transcode_and_probe(VideoFromFile(all_bad))["audio_codecs"] == []
finally:
for path in (mixed, all_bad):
if path:
os.unlink(path)
@@ -112,17 +112,6 @@ def _make_pid_v1_5_sd(latent_proj_channels=16):
return sd
def _make_joyimage_edit_plus_sd():
sd = {
"img_in.weight": torch.empty(4096, 16, 1, 2, 2, device="meta"),
"condition_embedder.time_embedder.linear_1.weight": torch.empty(1, device="meta"),
"double_blocks.0.attn.img_attn_q_norm.weight": torch.empty(128, device="meta"),
}
for i in range(40):
sd[f"double_blocks.{i}.attn.img_attn_qkv.weight"] = torch.empty(1, device="meta")
return sd
def _add_model_diffusion_prefix(sd):
return {f"model.diffusion_model.{k}": v for k, v in sd.items()}
@@ -269,26 +258,6 @@ class TestModelDetection:
assert processed["pixel_blocks.0.adaLN_modulation_msa.bias"].shape == (12288,)
assert processed["pixel_blocks.0.adaLN_modulation_mlp.bias"].shape == (12288,)
def test_joyimage_edit_plus_detection(self):
sd = _make_joyimage_edit_plus_sd()
unet_config = detect_unet_config(sd, "")
assert unet_config == {
"image_model": "joyimage",
"in_channels": 16,
"hidden_size": 4096,
"patch_size": [1, 2, 2],
"num_layers": 40,
"num_attention_heads": 32,
"text_dim": 4096,
}
assert type(model_config_from_unet_config(unet_config, sd)).__name__ == "JoyImage"
def test_incomplete_joyimage_signature_is_not_detected(self):
sd = _make_joyimage_edit_plus_sd()
del sd["double_blocks.0.attn.img_attn_q_norm.weight"]
assert detect_unet_config(sd, "") is None
def test_unet_config_and_required_keys_combination_is_unique(self):
"""Each model in the registry must have a unique combination of
``unet_config`` and ``required_keys``. If two models share the same
@@ -0,0 +1,90 @@
"""Shared fixtures for the model download manager tests.
These run in-process (no ComfyUI subprocess): a file-backed SQLite DB is
initialized once, a temp model folder is registered with ``folder_paths``, and
the shared aiohttp session is reset between tests so each async test gets a
session bound to its own event loop.
"""
from __future__ import annotations
import asyncio
import os
import tempfile
import pytest
def _drain_scheduler_tasks(scheduler) -> None:
"""Cancel and await live scheduler tasks so none outlive the test.
Uses the actual task handles rather than only clearing ``_tasks``: each
per-test event loop is created by ``asyncio.run``, so a task left behind by
a crashed/aborted test would otherwise keep its coroutine alive. We cancel
every live task and, when its loop is still usable, run it to completion to
let the cancellation propagate before dropping the reference.
"""
for task in list(scheduler._tasks.values()):
if task is None:
continue
loop = task.get_loop()
if task.done() or loop.is_closed():
continue
task.cancel()
if not loop.is_running():
try:
loop.run_until_complete(asyncio.gather(task, return_exceptions=True))
except Exception:
pass
scheduler._tasks.clear()
@pytest.fixture(scope="session", autouse=True)
def _init_db():
import app.database.db as db
from comfy.cli_args import args
fd, db_path = tempfile.mkstemp(suffix="-dlmgr-test.sqlite3")
os.close(fd)
args.database_url = f"sqlite:///{db_path}"
db.init_db()
yield
try:
os.remove(db_path)
except OSError:
pass
@pytest.fixture(autouse=True)
def _reset_runtime():
"""Reset module singletons that hold event-loop-bound or cross-test state."""
import app.model_downloader.net.session as ns
from app.model_downloader.scheduler import SCHEDULER
ns._session = None
_drain_scheduler_tasks(SCHEDULER)
SCHEDULER._jobs.clear()
SCHEDULER._backoff_until.clear()
SCHEDULER._started = False
yield
_drain_scheduler_tasks(SCHEDULER)
ns._session = None
@pytest.fixture
def model_root(tmp_path):
"""Register a temp 'loras' model folder and return its absolute path."""
import folder_paths
root = tmp_path / "loras"
root.mkdir(parents=True, exist_ok=True)
saved = folder_paths.folder_names_and_paths.get("loras")
folder_paths.folder_names_and_paths["loras"] = (
[str(root)],
{".safetensors", ".sft", ".ckpt", ".pt", ".pth"},
)
yield str(root)
if saved is not None:
folder_paths.folder_names_and_paths["loras"] = saved
else:
folder_paths.folder_names_and_paths.pop("loras", None)
@@ -0,0 +1,387 @@
"""Unit tests for download authentication.
Covers env-key resolution, OAuth token resolution/refresh/expiry, provider
host matching (including the per-hop drop on a CDN host), and the auth routes.
Async tests are driven via ``asyncio.run`` so no pytest-asyncio plugin is needed.
"""
from __future__ import annotations
import asyncio
import json
import time
from urllib.parse import parse_qs, urlsplit
import pytest
from aiohttp.test_utils import make_mocked_request
from app.model_downloader.api import routes
from app.model_downloader.auth import oauth, token_store
from app.model_downloader.auth.providers import PROVIDERS, provider_for_host
from app.model_downloader.auth.resolver import resolve_auth_for_hop
from app.model_downloader.auth.store import AUTH_STORE
from app.model_downloader.auth.token_store import Token
_HF_ENV = ("HF_TOKEN", "HUGGING_FACE_HUB_TOKEN")
_CIVITAI_ENV = ("CIVITAI_API_TOKEN", "CIVITAI_API_KEY")
@pytest.fixture
def auth_tmp(monkeypatch, tmp_path):
"""Isolate the on-disk token store and clear the in-memory cache."""
d = tmp_path / "download_auth"
d.mkdir()
monkeypatch.setattr(token_store, "_auth_dir", lambda: str(d))
AUTH_STORE._cache.clear()
yield
AUTH_STORE._cache.clear()
def _clear_env(monkeypatch, *names):
for name in names:
monkeypatch.delenv(name, raising=False)
# ----- provider host matching -----
def test_provider_for_host():
assert provider_for_host("HuggingFace.co:443").name == "huggingface"
assert provider_for_host("civitai.com").name == "civitai"
# sibling CDN hosts must not match — this is what drops the token on redirect
assert provider_for_host("cdn-lfs.huggingface.co") is None
assert provider_for_host("cas-bridge.xethub.hf.co") is None
assert provider_for_host("example.com") is None
# ----- env-key resolution -----
def test_env_key_resolution_hf(monkeypatch, auth_tmp):
monkeypatch.setenv("HF_TOKEN", "hf_env")
async def _run():
auth = await resolve_auth_for_hop("huggingface.co", "https")
assert auth is not None
assert auth.headers["Authorization"] == "Bearer hf_env"
# never over http, never on a CDN redirect host
assert await resolve_auth_for_hop("huggingface.co", "http") is None
assert await resolve_auth_for_hop("cdn-lfs.huggingface.co", "https") is None
asyncio.run(_run())
def test_env_key_resolution_civitai_secondary_var(monkeypatch, auth_tmp):
_clear_env(monkeypatch, "CIVITAI_API_TOKEN")
monkeypatch.setenv("CIVITAI_API_KEY", "civ_env")
async def _run():
auth = await resolve_auth_for_hop("civitai.com", "https")
assert auth is not None
assert auth.headers["Authorization"] == "Bearer civ_env"
asyncio.run(_run())
def test_env_key_takes_precedence_over_oauth(monkeypatch, auth_tmp):
monkeypatch.setenv("HF_TOKEN", "hf_env")
async def _run():
AUTH_STORE.set_token("huggingface", Token(access_token="oauth_acc"))
auth = await resolve_auth_for_hop("huggingface.co", "https")
assert auth.headers["Authorization"] == "Bearer hf_env"
asyncio.run(_run())
# ----- OAuth token resolution / refresh / expiry -----
def test_oauth_token_resolution(monkeypatch, auth_tmp):
_clear_env(monkeypatch, *_HF_ENV)
async def _run():
AUTH_STORE.set_token("huggingface", Token(access_token="acc", expires_at=0))
auth = await resolve_auth_for_hop("huggingface.co", "https")
assert auth is not None
assert auth.headers["Authorization"] == "Bearer acc"
asyncio.run(_run())
def test_oauth_refresh_on_expiry(monkeypatch, auth_tmp):
_clear_env(monkeypatch, *_HF_ENV)
async def fake_refresh(provider, tok):
return Token(
access_token="new_acc",
refresh_token="r2",
expires_at=int(time.time()) + 3600,
)
monkeypatch.setattr(oauth, "refresh_access_token", fake_refresh)
async def _run():
AUTH_STORE.set_token(
"huggingface",
Token(access_token="old", refresh_token="r1", expires_at=1),
)
access = await AUTH_STORE.get_valid_token(PROVIDERS["huggingface"])
assert access == "new_acc"
# the refreshed token is persisted (cache + disk)
assert token_store.load("huggingface").access_token == "new_acc"
asyncio.run(_run())
def test_oauth_expired_without_refresh_returns_none(monkeypatch, auth_tmp):
_clear_env(monkeypatch, *_HF_ENV)
async def _run():
AUTH_STORE.set_token(
"huggingface",
Token(access_token="old", refresh_token=None, expires_at=1),
)
assert await AUTH_STORE.get_valid_token(PROVIDERS["huggingface"]) is None
assert await resolve_auth_for_hop("huggingface.co", "https") is None
asyncio.run(_run())
def test_no_auth_when_nothing_configured(monkeypatch, auth_tmp):
_clear_env(monkeypatch, *_HF_ENV, *_CIVITAI_ENV)
async def _run():
assert await resolve_auth_for_hop("huggingface.co", "https") is None
assert await resolve_auth_for_hop("example.com", "https") is None
asyncio.run(_run())
# ----- auth routes -----
def test_auth_status_route(monkeypatch, auth_tmp):
_clear_env(monkeypatch, *_HF_ENV, *_CIVITAI_ENV)
async def _run():
resp = await routes.auth_status(make_mocked_request("GET", "/api/download/auth"))
data = json.loads(resp.body)
by_name = {p["provider"]: p for p in data["providers"]}
assert set(by_name) == {"huggingface", "civitai"}
assert by_name["huggingface"]["logged_in"] is False
assert by_name["huggingface"]["env_key_present"] is False
asyncio.run(_run())
def test_login_unconfigured_returns_400(monkeypatch, auth_tmp):
_clear_env(monkeypatch, "COMFY_HF_OAUTH_CLIENT_ID")
async def _run():
req = make_mocked_request(
"POST", "/api/download/auth/huggingface/login",
match_info={"provider": "huggingface"},
)
resp = await routes.auth_login(req)
assert resp.status == 400
assert json.loads(resp.body)["error"]["code"] == "OAUTH_NOT_CONFIGURED"
asyncio.run(_run())
def test_login_unknown_provider_returns_400(auth_tmp):
async def _run():
req = make_mocked_request(
"POST", "/api/download/auth/nope/login",
match_info={"provider": "nope"},
)
resp = await routes.auth_login(req)
assert resp.status == 400
assert json.loads(resp.body)["error"]["code"] == "UNKNOWN_PROVIDER"
asyncio.run(_run())
def test_login_start_and_in_progress(monkeypatch, auth_tmp):
monkeypatch.setenv("COMFY_HF_OAUTH_CLIENT_ID", "test-client")
async def _run():
try:
req = make_mocked_request(
"POST", "/api/download/auth/huggingface/login",
match_info={"provider": "huggingface"},
)
resp = await routes.auth_login(req)
assert resp.status == 200
url = json.loads(resp.body)["authorize_url"]
assert url.startswith("https://huggingface.co/oauth/authorize?")
assert "code_challenge=" in url and "code_challenge_method=S256" in url
assert "client_id=test-client" in url
# a second concurrent login is rejected
resp2 = await routes.auth_login(
make_mocked_request(
"POST", "/api/download/auth/huggingface/login",
match_info={"provider": "huggingface"},
)
)
assert resp2.status == 409
finally:
flow = oauth._ACTIVE.get("huggingface")
if flow is not None:
await flow._teardown()
asyncio.run(_run())
def _redirect_uri(authorize_url: str) -> str:
return parse_qs(urlsplit(authorize_url).query)["redirect_uri"][0]
def test_callback_uri_is_fixed_loopback_port(monkeypatch, auth_tmp):
"""The redirect targets the fixed loopback port, ignoring the request Host.
In dev the UI is served by Vite on a different port that proxies ``/api``,
so the login request arrives with the frontend port in ``Host``; the
callback URI must stay pinned to the registered loopback port.
"""
monkeypatch.setenv("COMFY_HF_OAUTH_CLIENT_ID", "hf-client")
async def _run():
try:
req = make_mocked_request(
"POST", "/api/download/auth/huggingface/login",
headers={"Host": "localhost:5173"},
match_info={"provider": "huggingface"},
)
resp = await routes.auth_login(req)
redirect = _redirect_uri(json.loads(resp.body)["authorize_url"])
assert redirect == f"http://127.0.0.1:{oauth.CALLBACK_PORT}/callback/huggingface"
finally:
flow = oauth._ACTIVE.get("huggingface")
if flow is not None:
await flow._teardown()
asyncio.run(_run())
def test_second_login_rejected_single_flight(monkeypatch, auth_tmp):
"""Only one login runs at a time — the shared callback port is single-flight."""
monkeypatch.setenv("COMFY_HF_OAUTH_CLIENT_ID", "hf-client")
monkeypatch.setenv("COMFY_CIVITAI_OAUTH_CLIENT_ID", "civ-client")
async def _run():
try:
hf = await routes.auth_login(
make_mocked_request(
"POST", "/api/download/auth/huggingface/login",
match_info={"provider": "huggingface"},
)
)
assert hf.status == 200
# a different provider can't start while HF holds the callback port
civ = await routes.auth_login(
make_mocked_request(
"POST", "/api/download/auth/civitai/login",
match_info={"provider": "civitai"},
)
)
assert civ.status == 409
assert json.loads(civ.body)["error"]["code"] == "LOGIN_IN_PROGRESS"
assert oauth.login_in_progress("huggingface")
assert not oauth.login_in_progress("civitai")
finally:
for name in ("huggingface", "civitai"):
flow = oauth._ACTIVE.get(name)
if flow is not None:
await flow._teardown()
asyncio.run(_run())
def test_callback_completes_login(monkeypatch, auth_tmp):
"""The loopback callback exchanges the code, stores the token, tears down."""
monkeypatch.setenv("COMFY_HF_OAUTH_CLIENT_ID", "hf-client")
_clear_env(monkeypatch, *_HF_ENV)
async def fake_exchange(provider, code, verifier, redirect_uri):
assert code == "the-code"
return Token(access_token="acc_from_callback")
monkeypatch.setattr(oauth, "exchange_code", fake_exchange)
async def _run():
try:
login = await routes.auth_login(
make_mocked_request(
"POST", "/api/download/auth/huggingface/login",
match_info={"provider": "huggingface"},
)
)
assert login.status == 200
flow = oauth._ACTIVE["huggingface"]
cb = await flow._handle_callback(
make_mocked_request(
"GET",
f"/callback/huggingface?state={flow.state}&code=the-code",
match_info={"provider": "huggingface"},
)
)
assert cb.status == 200
assert "window.close" in cb.text
assert AUTH_STORE.status(PROVIDERS["huggingface"])["logged_in"] is True
# the handler schedules its own teardown; let it run
await asyncio.sleep(0.05)
assert not oauth.login_in_progress("huggingface")
finally:
flow = oauth._ACTIVE.get("huggingface")
if flow is not None:
await flow._teardown()
asyncio.run(_run())
def test_callback_rejects_bad_state(monkeypatch, auth_tmp):
monkeypatch.setenv("COMFY_HF_OAUTH_CLIENT_ID", "hf-client")
async def _run():
try:
await routes.auth_login(
make_mocked_request(
"POST", "/api/download/auth/huggingface/login",
match_info={"provider": "huggingface"},
)
)
flow = oauth._ACTIVE["huggingface"]
cb = await flow._handle_callback(
make_mocked_request(
"GET",
"/callback/huggingface?state=wrong&code=x",
match_info={"provider": "huggingface"},
)
)
assert cb.status == 400
# the flow stays pending so a genuine callback can still arrive
assert oauth.login_in_progress("huggingface")
finally:
flow = oauth._ACTIVE.get("huggingface")
if flow is not None:
await flow._teardown()
asyncio.run(_run())
def test_logout_route_clears_token(auth_tmp):
async def _run():
AUTH_STORE.set_token("civitai", Token(access_token="x"))
assert AUTH_STORE.status(PROVIDERS["civitai"])["logged_in"] is True
req = make_mocked_request(
"POST", "/api/download/auth/civitai/logout",
match_info={"provider": "civitai"},
)
resp = await routes.auth_logout(req)
assert resp.status == 200
assert json.loads(resp.body)["logged_out"] is True
assert AUTH_STORE.status(PROVIDERS["civitai"])["logged_in"] is False
asyncio.run(_run())
@@ -0,0 +1,136 @@
"""Unit tests for ``DownloadManager.delete`` and ``DownloadManager.clear``.
Deleting a terminal row must remove it from history for good (so it does not
reappear on the next ``list``), leave live rows untouched, and clean up any
leftover ``.part`` temp file without touching the finished model file.
``clear()`` is the bulk variant: it removes all terminal rows atomically, skips
live ones, and returns the count of rows deleted.
Async methods are driven via ``asyncio.run`` so no pytest-asyncio plugin is
required.
"""
from __future__ import annotations
import asyncio
import os
import pytest
from app.model_downloader.constants import DownloadStatus
from app.model_downloader.database import queries
from app.model_downloader.manager import DOWNLOAD_MANAGER, DownloadError
def _insert(download_id: str, status: str, *, temp_path: str = "/tmp/none.part") -> None:
queries.insert_download(
{
"id": download_id,
"url": "https://huggingface.co/org/model.safetensors",
"model_id": "loras/model.safetensors",
"dest_path": "/tmp/model.safetensors",
"temp_path": temp_path,
"status": status,
"priority": 0,
}
)
def test_delete_removes_terminal_row_from_history():
_insert("done", DownloadStatus.COMPLETED)
asyncio.run(DOWNLOAD_MANAGER.delete("done"))
assert queries.get_download("done") is None
def test_delete_refuses_live_row():
_insert("live", DownloadStatus.QUEUED)
with pytest.raises(DownloadError) as excinfo:
asyncio.run(DOWNLOAD_MANAGER.delete("live"))
assert excinfo.value.code == "DOWNLOAD_ACTIVE"
assert queries.get_download("live") is not None
def test_delete_missing_row_raises_not_found():
with pytest.raises(DownloadError) as excinfo:
asyncio.run(DOWNLOAD_MANAGER.delete("nope"))
assert excinfo.value.code == "NOT_FOUND"
def test_delete_removes_leftover_temp_file(tmp_path):
partial = tmp_path / "model.safetensors.part"
partial.write_bytes(b"partial")
_insert("failed", DownloadStatus.FAILED, temp_path=str(partial))
asyncio.run(DOWNLOAD_MANAGER.delete("failed"))
assert not os.path.exists(partial)
assert queries.get_download("failed") is None
# ----- clear -----
def test_clear_removes_all_terminal_rows():
_insert("c-done", DownloadStatus.COMPLETED)
_insert("c-fail", DownloadStatus.FAILED)
_insert("c-canc", DownloadStatus.CANCELLED)
deleted = asyncio.run(DOWNLOAD_MANAGER.clear())
assert deleted == 3
assert queries.get_download("c-done") is None
assert queries.get_download("c-fail") is None
assert queries.get_download("c-canc") is None
def test_clear_skips_live_rows():
_insert("cl-queued", DownloadStatus.QUEUED)
_insert("cl-paused", DownloadStatus.PAUSED)
_insert("cl-done", DownloadStatus.COMPLETED)
deleted = asyncio.run(DOWNLOAD_MANAGER.clear())
assert deleted == 1
assert queries.get_download("cl-queued") is not None
assert queries.get_download("cl-paused") is not None
assert queries.get_download("cl-done") is None
def test_clear_returns_zero_when_nothing_to_delete():
_insert("cl-only-live", DownloadStatus.QUEUED)
deleted = asyncio.run(DOWNLOAD_MANAGER.clear())
assert deleted == 0
assert queries.get_download("cl-only-live") is not None
def test_clear_removes_leftover_temp_files(tmp_path):
partial = tmp_path / "clear_partial.part"
partial.write_bytes(b"partial data")
finished = tmp_path / "finished.safetensors"
finished.write_bytes(b"real model weights")
_insert("cl-part", DownloadStatus.FAILED, temp_path=str(partial))
# The finished file is not the temp_path; temp_path for a completed download
# no longer exists (already renamed), so use a non-existent path here to
# verify clear() tolerates a missing temp file without raising.
_insert("cl-comp", DownloadStatus.COMPLETED, temp_path=str(tmp_path / "gone.part"))
asyncio.run(DOWNLOAD_MANAGER.clear())
# Leftover .part from the failed download is cleaned up.
assert not partial.exists()
# Finished model file is never touched.
assert finished.exists()
def test_clear_empty_db_returns_zero():
deleted = asyncio.run(DOWNLOAD_MANAGER.clear())
assert deleted == 0
@@ -0,0 +1,637 @@
"""Integration tests for the download engine against a local aiohttp server.
Covers single-stream and segmented transfers, deterministic resume from a
partial file, and cancel rollback. Async tests are driven via ``asyncio.run``
so no pytest-asyncio plugin is required.
"""
from __future__ import annotations
import asyncio
import json
import os
import struct
import uuid
import pytest
from aiohttp import web
from comfy.cli_args import args
from app.model_downloader.constants import DownloadStatus
from app.model_downloader.database import queries
from app.model_downloader.engine.job import DownloadJob, JobSpec
from app.model_downloader.net.session import close_session
from app.model_downloader.security import paths
PAYLOAD_ETAG = '"v1"'
def _payload(n: int) -> bytes:
return bytes((i * 37 + 11) % 256 for i in range(n))
def _safetensors_payload(total: int) -> bytes:
"""A structurally valid ``.safetensors`` blob of exactly ``total`` bytes.
Success-path tests download to ``.safetensors`` destinations, which the
engine now structurally validates before the atomic rename, so their
payloads must parse as real safetensors (header length + JSON header +
data region whose size matches the declared ``data_offsets``).
"""
def _header(data_len: int) -> bytes:
return json.dumps(
{"w": {"dtype": "U8", "shape": [data_len], "data_offsets": [0, data_len]}}
).encode("utf-8")
# The header's byte length depends on the digit count of ``data_len``, so
# iterate until ``total == 8 + len(header) + data_len`` is self-consistent.
data_len = total - 8 - len(_header(total))
for _ in range(8):
header = _header(data_len)
new_data_len = total - 8 - len(header)
if new_data_len == data_len:
break
data_len = new_data_len
assert data_len >= 0, "total too small for a safetensors payload"
header = _header(data_len)
body = bytes((i * 37 + 11) % 256 for i in range(data_len))
return struct.pack("<Q", len(header)) + header + body
def _range_handler(payload: bytes):
async def handler(request: web.Request) -> web.Response:
rng = request.headers.get("Range")
if rng:
spec = rng.split("=", 1)[1]
s, _, e = spec.partition("-")
start = int(s)
end = int(e) if e else len(payload) - 1
chunk = payload[start : end + 1]
return web.Response(
status=206,
body=chunk,
headers={
"Content-Range": f"bytes {start}-{end}/{len(payload)}",
"Accept-Ranges": "bytes",
"ETag": PAYLOAD_ETAG,
},
)
return web.Response(
status=200, body=payload, headers={"Accept-Ranges": "bytes", "ETag": PAYLOAD_ETAG}
)
return handler
def _content_disposition_handler(payload: bytes, filename: str):
"""A range-capable server that only reveals its filename via a header.
Models a Civitai-style ``/api/download/...`` endpoint: the URL path has no
extension, and the real filename (hence extension) lives in the response
``Content-Disposition`` header.
"""
async def handler(request: web.Request) -> web.Response:
headers = {
"Accept-Ranges": "bytes",
"ETag": PAYLOAD_ETAG,
"Content-Disposition": f'attachment; filename="{filename}"',
}
rng = request.headers.get("Range")
if rng:
spec = rng.split("=", 1)[1]
s, _, e = spec.partition("-")
start = int(s)
end = int(e) if e else len(payload) - 1
chunk = payload[start : end + 1]
return web.Response(
status=206,
body=chunk,
headers={**headers, "Content-Range": f"bytes {start}-{end}/{len(payload)}"},
)
return web.Response(status=200, body=payload, headers=headers)
return handler
def _noranges_handler(payload: bytes):
async def handler(request: web.Request) -> web.Response:
# Always full body, never advertises Accept-Ranges -> single-stream.
return web.Response(status=200, body=payload)
return handler
def _slow_handler(payload: bytes, chunk: int = 16384, delay: float = 0.01):
async def handler(request: web.Request) -> web.StreamResponse:
resp = web.StreamResponse(
status=200, headers={"Content-Length": str(len(payload))}
)
await resp.prepare(request)
for i in range(0, len(payload), chunk):
await resp.write(payload[i : i + chunk])
await asyncio.sleep(delay)
await resp.write_eof()
return resp
return handler
def _overflow_range_handler(payload: bytes, extra: int = 256 * 1024):
"""A non-conforming 206 server that returns MORE than the requested range."""
async def handler(request: web.Request) -> web.Response:
rng = request.headers.get("Range")
if rng:
spec = rng.split("=", 1)[1]
s, _, e = spec.partition("-")
start = int(s)
end = int(e) if e else len(payload) - 1
# Maliciously overrun: append extra bytes past the requested end.
body = payload[start : end + 1] + bytes(extra)
return web.Response(
status=206,
body=body,
headers={
"Content-Range": f"bytes {start}-{end}/{len(payload)}",
"Accept-Ranges": "bytes",
"ETag": PAYLOAD_ETAG,
},
)
return web.Response(
status=200, body=payload, headers={"Accept-Ranges": "bytes", "ETag": PAYLOAD_ETAG}
)
return handler
def _short_range_handler(payload: bytes, drop: int = 64 * 1024):
"""A 206 server that returns fewer bytes than requested for later segments.
Simulates a server cleanly closing a range connection early. The response
is internally consistent (Content-Length matches the short body), so the
client sees no error and the segment just ends short, leaving a zero-filled
hole in the preallocated file.
"""
async def handler(request: web.Request) -> web.Response:
rng = request.headers.get("Range")
if rng:
spec = rng.split("=", 1)[1]
s, _, e = spec.partition("-")
start = int(s)
end = int(e) if e else len(payload) - 1
chunk = payload[start : end + 1]
if start > 0 and len(chunk) > drop:
chunk = chunk[:-drop] # truncate a non-first segment
return web.Response(
status=206,
body=chunk,
headers={
"Content-Range": f"bytes {start}-{end}/{len(payload)}",
"Accept-Ranges": "bytes",
"ETag": PAYLOAD_ETAG,
},
)
return web.Response(
status=200, body=payload, headers={"Accept-Ranges": "bytes", "ETag": PAYLOAD_ETAG}
)
return handler
def _unbounded_handler(total: int, chunk: int = 16384):
"""A 200 stream with no Content-Length / Accept-Ranges (unknown length)."""
async def handler(request: web.Request) -> web.StreamResponse:
resp = web.StreamResponse(status=200)
await resp.prepare(request)
sent = 0
while sent < total:
await resp.write(bytes(min(chunk, total - sent)))
sent += chunk
await resp.write_eof()
return resp
return handler
async def _serve(handler):
app = web.Application()
app.router.add_route("*", "/{name:.*}", handler)
runner = web.AppRunner(app)
await runner.setup()
site = web.TCPSite(runner, "127.0.0.1", 0)
await site.start()
port = site._server.sockets[0].getsockname()[1]
return runner, port
def _insert(model_id: str, url: str, status: str = DownloadStatus.QUEUED) -> tuple[str, str, str]:
final_path, temp_path = paths.resolve_destination(model_id)
download_id = str(uuid.uuid4())
queries.insert_download(
{
"id": download_id,
"url": url,
"model_id": model_id,
"dest_path": final_path,
"temp_path": temp_path,
"status": status,
}
)
return download_id, final_path, temp_path
# ----- single-stream -----
def test_single_stream_download(model_root):
payload = _safetensors_payload(300_000)
async def _run():
await close_session()
runner, port = await _serve(_noranges_handler(payload))
try:
url = f"http://127.0.0.1:{port}/model.safetensors"
did, final_path, _temp = _insert("loras/single.safetensors", url)
job = DownloadJob(JobSpec(
download_id=did, url=url, model_id="loras/single.safetensors",
dest_path=final_path, temp_path=_temp,
))
status = await job.run()
assert status == DownloadStatus.COMPLETED, queries.get_download(did).error
assert os.path.exists(final_path)
assert open(final_path, "rb").read() == payload
finally:
await runner.cleanup()
await close_session()
asyncio.run(_run())
# ----- segmented -----
def test_segmented_download(model_root):
payload = _safetensors_payload(4 * 1024 * 1024) # 4 MiB -> multiple segments
async def _run():
await close_session()
runner, port = await _serve(_range_handler(payload))
try:
url = f"http://127.0.0.1:{port}/model.safetensors"
did, final_path, temp = _insert("loras/seg.safetensors", url)
job = DownloadJob(JobSpec(
download_id=did, url=url, model_id="loras/seg.safetensors",
dest_path=final_path, temp_path=temp,
))
status = await job.run()
assert status == DownloadStatus.COMPLETED, queries.get_download(did).error
assert open(final_path, "rb").read() == payload
# More than one segment row was planned.
assert len(queries.list_segments(did)) > 1
finally:
await runner.cleanup()
await close_session()
asyncio.run(_run())
# ----- deterministic resume from a partial file -----
def test_resume_from_partial(model_root):
payload = _safetensors_payload(512 * 1024) # < 1 MiB -> single segment
async def _run():
await close_session()
runner, port = await _serve(_range_handler(payload))
try:
url = f"http://127.0.0.1:{port}/model.safetensors"
did, final_path, temp = _insert("loras/resume.safetensors", url)
# Simulate a prior partial: first 200 KiB already written, offset persisted.
prefix = 200 * 1024
os.makedirs(os.path.dirname(temp), exist_ok=True)
with open(temp, "wb") as f:
f.write(payload[:prefix])
queries.update_download(did, bytes_done=prefix, etag=PAYLOAD_ETAG)
job = DownloadJob(JobSpec(
download_id=did, url=url, model_id="loras/resume.safetensors",
dest_path=final_path, temp_path=temp, etag=PAYLOAD_ETAG,
))
status = await job.run()
assert status == DownloadStatus.COMPLETED, queries.get_download(did).error
assert open(final_path, "rb").read() == payload
finally:
await runner.cleanup()
await close_session()
asyncio.run(_run())
# ----- cancel rollback -----
def test_cancel_rollback(model_root, monkeypatch):
monkeypatch.setattr(args, "download_chunk_size", 16384, raising=False)
payload = _payload(1024 * 1024)
async def _run():
await close_session()
runner, port = await _serve(_slow_handler(payload))
try:
url = f"http://127.0.0.1:{port}/model.safetensors"
did, final_path, temp = _insert("loras/cancel.safetensors", url)
job = DownloadJob(JobSpec(
download_id=did, url=url, model_id="loras/cancel.safetensors",
dest_path=final_path, temp_path=temp,
))
task = asyncio.ensure_future(job.run())
# Wait until some bytes have been written, then cancel.
for _ in range(200):
await asyncio.sleep(0.01)
if job.state.bytes_done > 0:
break
job.request_cancel()
status = await task
assert status == DownloadStatus.CANCELLED
assert not os.path.exists(temp)
assert not os.path.exists(final_path)
finally:
await runner.cleanup()
await close_session()
asyncio.run(_run())
# ----- size-bound enforcement (malicious / non-conforming hosts) -----
def test_segment_overflow_aborts(model_root):
"""A 206 returning more than the requested range must not overrun."""
payload = _payload(4 * 1024 * 1024) # large enough to segment
async def _run():
await close_session()
runner, port = await _serve(_overflow_range_handler(payload))
try:
url = f"http://127.0.0.1:{port}/model.safetensors"
did, final_path, temp = _insert("loras/overflow.safetensors", url)
job = DownloadJob(JobSpec(
download_id=did, url=url, model_id="loras/overflow.safetensors",
dest_path=final_path, temp_path=temp,
))
status = await job.run()
assert status == DownloadStatus.FAILED
assert not os.path.exists(final_path)
assert not os.path.exists(temp)
finally:
await runner.cleanup()
await close_session()
asyncio.run(_run())
def test_short_segment_fails_closed(model_root):
"""A segment that ends short must fail, not be accepted as complete.
The file is preallocated to total_bytes, so the on-disk size still equals
total even with a zero-filled hole; completeness must be judged per-segment.
"""
payload = _safetensors_payload(4 * 1024 * 1024) # large enough to segment
async def _run():
await close_session()
runner, port = await _serve(_short_range_handler(payload))
try:
url = f"http://127.0.0.1:{port}/model.safetensors"
did, final_path, temp = _insert("loras/short.safetensors", url)
job = DownloadJob(JobSpec(
download_id=did, url=url, model_id="loras/short.safetensors",
dest_path=final_path, temp_path=temp,
))
status = await job.run()
assert status == DownloadStatus.FAILED, queries.get_download(did).error
assert "incomplete" in (queries.get_download(did).error or "")
assert not os.path.exists(final_path)
assert not os.path.exists(temp)
finally:
await runner.cleanup()
await close_session()
asyncio.run(_run())
def test_structural_validation_rejects_corrupt(model_root):
"""A correctly sized but structurally invalid file fails closed (not retried).
Regression for the dead structural gate: validation must key off the
destination extension, not the ``.part`` temp suffix.
"""
payload = _payload(300_000) # right size, but not a valid safetensors blob
async def _run():
await close_session()
runner, port = await _serve(_noranges_handler(payload))
try:
url = f"http://127.0.0.1:{port}/model.safetensors"
did, final_path, temp = _insert("loras/corrupt.safetensors", url)
job = DownloadJob(JobSpec(
download_id=did, url=url, model_id="loras/corrupt.safetensors",
dest_path=final_path, temp_path=temp,
))
status = await job.run()
assert status == DownloadStatus.FAILED, queries.get_download(did).error
assert not os.path.exists(final_path)
assert not os.path.exists(temp)
# Failed closed at first attempt, not re-queued as retryable.
assert queries.get_download(did).attempts == 0
finally:
await runner.cleanup()
await close_session()
asyncio.run(_run())
def test_rejects_oversized_known_download(model_root, monkeypatch):
"""A file whose advertised size exceeds the cap is rejected at probe."""
monkeypatch.setattr(args, "download_max_bytes", 100_000, raising=False)
payload = _payload(300_000)
async def _run():
await close_session()
runner, port = await _serve(_noranges_handler(payload))
try:
url = f"http://127.0.0.1:{port}/model.safetensors"
did, final_path, temp = _insert("loras/toobig.safetensors", url)
job = DownloadJob(JobSpec(
download_id=did, url=url, model_id="loras/toobig.safetensors",
dest_path=final_path, temp_path=temp,
))
status = await job.run()
assert status == DownloadStatus.FAILED
assert not os.path.exists(final_path)
finally:
await runner.cleanup()
await close_session()
asyncio.run(_run())
def test_unknown_length_capped_by_max_bytes(model_root, monkeypatch):
"""An unbounded unknown-length stream is capped by --download-max-bytes."""
monkeypatch.setattr(args, "download_max_bytes", 100_000, raising=False)
monkeypatch.setattr(args, "download_chunk_size", 16384, raising=False)
async def _run():
await close_session()
runner, port = await _serve(_unbounded_handler(2 * 1024 * 1024))
try:
url = f"http://127.0.0.1:{port}/model.safetensors"
did, final_path, temp = _insert("loras/unbounded.safetensors", url)
job = DownloadJob(JobSpec(
download_id=did, url=url, model_id="loras/unbounded.safetensors",
dest_path=final_path, temp_path=temp,
))
status = await job.run()
assert status == DownloadStatus.FAILED
assert not os.path.exists(final_path)
assert not os.path.exists(temp)
finally:
await runner.cleanup()
await close_session()
asyncio.run(_run())
# ----- manager + scheduler end-to-end -----
def test_manager_enqueue_to_completion(model_root):
payload = _safetensors_payload(2 * 1024 * 1024)
async def _run():
await close_session()
from app.model_downloader.manager import DOWNLOAD_MANAGER
runner, port = await _serve(_range_handler(payload))
try:
url = f"http://127.0.0.1:{port}/model.safetensors"
did = await DOWNLOAD_MANAGER.enqueue(url, "loras/e2e.safetensors")
# Wait for completion.
final_path, _ = paths.resolve_destination("loras/e2e.safetensors")
for _ in range(500):
await asyncio.sleep(0.02)
row = queries.get_download(did)
if row.status in DownloadStatus.TERMINAL:
break
row = queries.get_download(did)
assert row.status == DownloadStatus.COMPLETED, row.error
assert open(final_path, "rb").read() == payload
finally:
await runner.cleanup()
await close_session()
asyncio.run(_run())
def test_manager_rejects_disallowed_url(model_root):
async def _run():
from app.model_downloader.manager import DOWNLOAD_MANAGER, DownloadError
with pytest.raises(DownloadError) as ei:
await DOWNLOAD_MANAGER.enqueue(
"https://evil.example.com/x.safetensors", "loras/bad.safetensors"
)
assert ei.value.code == "URL_NOT_ALLOWED"
asyncio.run(_run())
def test_manager_resolves_extensionless_url(model_root):
"""An allowlisted URL with no extension in its path is resolved from the
response, and the stored file adopts the resolved extension."""
payload = _safetensors_payload(1 * 1024 * 1024)
async def _run():
await close_session()
from app.model_downloader.manager import DOWNLOAD_MANAGER
runner, port = await _serve(
_content_disposition_handler(payload, "RealModel.safetensors")
)
try:
# No extension in the path (Civitai-style) and none in the model_id.
url = f"http://127.0.0.1:{port}/api/download/models/12345"
did = await DOWNLOAD_MANAGER.enqueue(url, "loras/my_civitai_model")
row = queries.get_download(did)
# The resolved extension was appended to the model_id + destination.
assert row.model_id == "loras/my_civitai_model.safetensors"
assert row.dest_path.endswith("my_civitai_model.safetensors")
final_path, _ = paths.resolve_destination(
"loras/my_civitai_model.safetensors"
)
for _ in range(500):
await asyncio.sleep(0.02)
row = queries.get_download(did)
if row.status in DownloadStatus.TERMINAL:
break
row = queries.get_download(did)
assert row.status == DownloadStatus.COMPLETED, row.error
assert open(final_path, "rb").read() == payload
finally:
await runner.cleanup()
await close_session()
asyncio.run(_run())
def test_manager_overrides_extension_from_resolution(model_root):
"""A model_id carrying a different known extension is corrected to match
the resolved URL's extension."""
payload = _safetensors_payload(256 * 1024)
async def _run():
await close_session()
from app.model_downloader.manager import DOWNLOAD_MANAGER
runner, port = await _serve(
_content_disposition_handler(payload, "weights.safetensors")
)
try:
url = f"http://127.0.0.1:{port}/api/download/models/777"
# Caller guessed .ckpt; resolution says .safetensors -> corrected.
did = await DOWNLOAD_MANAGER.enqueue(url, "loras/guessed.ckpt")
row = queries.get_download(did)
assert row.model_id == "loras/guessed.safetensors"
finally:
await runner.cleanup()
await close_session()
asyncio.run(_run())
def test_manager_rejects_non_model_resolution(model_root):
"""A URL that resolves to a non-model file is rejected, not downloaded."""
async def _run():
await close_session()
from app.model_downloader.manager import DOWNLOAD_MANAGER, DownloadError
runner, port = await _serve(
_content_disposition_handler(b"not a model", "installer.zip")
)
try:
url = f"http://127.0.0.1:{port}/api/download/models/999"
with pytest.raises(DownloadError) as ei:
await DOWNLOAD_MANAGER.enqueue(url, "loras/whatever")
assert ei.value.code == "URL_NOT_ALLOWED"
finally:
await runner.cleanup()
await close_session()
asyncio.run(_run())
@@ -0,0 +1,81 @@
"""Unit tests for the segment planner and structural safetensors validation."""
from __future__ import annotations
import json
import struct
import pytest
from app.model_downloader.engine.planner import (
effective_segment_count,
plan_segments,
)
from app.model_downloader.verify import structural
# ----- planner -----
def test_plan_segments_covers_full_range_contiguously():
total = 1000
plans = plan_segments(total, 4)
assert len(plans) == 4
assert plans[0].start == 0
assert plans[-1].end == total - 1
# contiguous, no gaps/overlaps
for a, b in zip(plans, plans[1:]):
assert b.start == a.end + 1
assert sum(p.length for p in plans) == total
def test_effective_segment_count_falls_back_to_single():
# No range support -> single
assert effective_segment_count(10_000_000, False, 8) == 1
# Unknown size -> single
assert effective_segment_count(None, True, 8) == 1
# Tiny file -> fewer segments than configured
assert effective_segment_count(1024, True, 8) == 1
# Large file with range support -> configured count
assert effective_segment_count(1_000_000_000, True, 8) == 8
# ----- structural -----
def _make_safetensors(tensor_data_len: int, *, corrupt_size: bool = False) -> bytes:
header = {"t": {"dtype": "F32", "shape": [tensor_data_len], "data_offsets": [0, tensor_data_len]}}
header_bytes = json.dumps(header).encode("utf-8")
body = b"\x00" * tensor_data_len
if corrupt_size:
body = body[:-1] # truncate one byte
return struct.pack("<Q", len(header_bytes)) + header_bytes + body
def test_structural_valid_safetensors(tmp_path):
p = tmp_path / "ok.safetensors"
p.write_bytes(_make_safetensors(256))
structural.validate(str(p)) # no raise
def test_structural_detects_truncation(tmp_path):
p = tmp_path / "bad.safetensors"
p.write_bytes(_make_safetensors(256, corrupt_size=True))
with pytest.raises(structural.StructuralError):
structural.validate(str(p))
def test_structural_skips_unknown_extension(tmp_path):
p = tmp_path / "weights.bin"
p.write_bytes(b"anything")
structural.validate(str(p)) # no structural check, no raise
def test_structural_detects_truncation_via_name_hint(tmp_path):
# The downloader validates the opaque temp file (a ``.part`` path) but keys
# the format check off the final destination name via ``name_hint``, so
# truncation must still be detected instead of silently skipped.
p = tmp_path / "bad.comfy-download.part"
p.write_bytes(_make_safetensors(256, corrupt_size=True))
with pytest.raises(structural.StructuralError):
structural.validate(str(p), name_hint="model.safetensors")
@@ -0,0 +1,231 @@
"""Unit tests for the security layer: allowlist, SSRF checks, path safety."""
from __future__ import annotations
import pytest
from app.model_downloader.security import allowlist, paths
from app.model_downloader.security.ssrf import (
SSRFError,
check_redirect_hop,
is_blocked_ip,
)
# ----- allowlist -----
@pytest.mark.parametrize(
"url,allowed",
[
("https://huggingface.co/org/repo/resolve/main/model.safetensors", True),
("https://civitai.com/api/download/x/model.safetensors", True),
("http://localhost/model.safetensors", True),
# off-list host
("https://evil.example.com/model.safetensors", False),
# http to a non-loopback allowlisted host is not permitted (https only)
("http://huggingface.co/org/repo/resolve/main/model.safetensors", False),
# bad extension on an allowed host
("https://huggingface.co/org/repo/resolve/main/config.json", False),
# userinfo trick: real host is the metadata IP, not 127.0.0.1
("http://127.0.0.1@169.254.169.254/x.safetensors", False),
],
)
def test_is_url_allowed(url, allowed):
assert allowlist.is_url_allowed(url) is allowed
def test_allow_any_extension_relaxes_extension_only():
url = "https://huggingface.co/org/repo/resolve/main/weights.bin"
assert allowlist.is_url_allowed(url) is True # .bin is in the known set
odd = "https://huggingface.co/org/repo/resolve/main/weights.zip"
assert allowlist.is_url_allowed(odd) is False
assert allowlist.is_url_allowed(odd, allow_any_extension=True) is True
@pytest.mark.parametrize(
"url,downloadable",
[
# known model extension in the path -> allowed
("https://civitai.com/x/model.safetensors", True),
# no extension in the path (Civitai download API) -> allowed, resolved later
("https://civitai.com/api/download/models/3031464?fileId=2910346", True),
("https://civitai.com/api/download/models/3031464", True),
# explicit non-model extension -> rejected even on an allowed host
("https://civitai.com/api/download/models/thing.zip", False),
("https://huggingface.co/org/repo/resolve/main/config.json", False),
# off-list host is never downloadable
("https://evil.example.com/api/download/models/1", False),
# http to a non-loopback allowlisted host is not permitted
("http://civitai.com/api/download/models/1", False),
],
)
def test_is_url_downloadable(url, downloadable):
assert allowlist.is_url_downloadable(url) is downloadable
@pytest.mark.parametrize(
"name,ext",
[
("model.safetensors", ".safetensors"),
("model.SAFETENSORS", ".safetensors"),
("archive.tar.gz", ".gz"),
("noext", ""),
(".safetensors", ""), # leading-dot dotfile -> no extension
("a/b/c/model.ckpt", ".ckpt"),
],
)
def test_filename_extension(name, ext):
assert allowlist.filename_extension(name) == ext
# ----- SSRF: blocked IPs -----
@pytest.mark.parametrize(
"ip,blocked",
[
("169.254.169.254", True), # cloud metadata / link-local
("127.0.0.1", True),
("10.0.0.5", True),
("192.168.1.1", True),
("172.16.0.1", True),
("::1", True),
("0.0.0.0", True),
# IPv4-mapped IPv6: must see through the mapping even on CPython
# versions predating the gh-113171 is_* property fix.
("::ffff:169.254.169.254", True), # mapped cloud metadata
("::ffff:127.0.0.1", True), # mapped loopback
("::ffff:10.0.0.1", True), # mapped RFC1918
("::ffff:8.8.8.8", False), # mapped public address stays allowed
("8.8.8.8", False),
("1.1.1.1", False),
("not-an-ip", True), # unparseable -> refuse
],
)
def test_is_blocked_ip(ip, blocked):
assert is_blocked_ip(ip) is blocked
# ----- SSRF: redirect hop validation -----
def test_check_redirect_hop_rejects_bad_scheme_and_userinfo():
with pytest.raises(SSRFError):
check_redirect_hop("ftp://huggingface.co/x.safetensors")
with pytest.raises(SSRFError):
check_redirect_hop("https://user:pass@cdn.example.com/x")
# A CDN host that is NOT on the allowlist is allowed as a redirect target
# (private-IP protection is the resolver's job; credential leak is prevented
# by exact host matching).
assert check_redirect_hop("https://cdn-lfs.huggingface.co/abc") is not None
def test_check_redirect_hop_http_only_for_loopback():
# Plain http to an external host is rejected (no plaintext downgrade).
with pytest.raises(SSRFError):
check_redirect_hop("http://cdn-lfs.huggingface.co/abc")
# http is honored for loopback only on the initial user-supplied URL (the
# "download a local model" feature).
assert (
check_redirect_hop("http://localhost/x.safetensors", is_initial_url=True)
is not None
)
assert (
check_redirect_hop("http://127.0.0.1/x.safetensors", is_initial_url=True)
is not None
)
def test_check_redirect_hop_blocks_loopback_and_ip_literals_on_redirect():
# A redirect (is_initial_url=False, the default) must never reach loopback,
# whether by hostname or by IP literal, nor any other internal IP literal.
for target in (
"http://localhost/x.safetensors",
"http://127.0.0.1/x.safetensors",
"https://[::1]/x.safetensors",
"https://169.254.169.254/x.safetensors", # cloud metadata
"https://10.0.0.5/x.safetensors", # RFC1918
):
with pytest.raises(SSRFError):
check_redirect_hop(target)
# Off-allowlist public CDN hosts (hostnames) remain valid redirect targets;
# their resolved IPs are screened by the connector's resolver.
assert check_redirect_hop("https://cdn-lfs.huggingface.co/abc") is not None
# ----- path safety -----
def test_parse_model_id_valid(model_root):
directory, filename = paths.parse_model_id("loras/my_lora.safetensors")
assert directory == "loras"
assert filename == "my_lora.safetensors"
@pytest.mark.parametrize(
"model_id",
[
"loras/../etc/passwd.safetensors", # traversal
"loras/sub/dir.safetensors", # nested
"unknownfolder/x.safetensors", # unknown folder
"loras/model.txt", # bad extension
"noslash.safetensors", # missing directory
"loras/", # empty filename
],
)
def test_parse_model_id_rejects(model_root, model_id):
with pytest.raises(paths.InvalidModelId):
paths.parse_model_id(model_id)
def test_resolve_destination_stays_in_root(model_root):
final_path, temp_path = paths.resolve_destination("loras/x.safetensors")
assert final_path.startswith(model_root)
assert temp_path.startswith(model_root)
assert temp_path != final_path
@pytest.mark.parametrize(
"model_id,ext,expected",
[
# no extension -> append the resolved one
("loras/my_civitai_model", ".safetensors", "loras/my_civitai_model.safetensors"),
# different known extension -> replace it
("loras/mymodel.ckpt", ".safetensors", "loras/mymodel.safetensors"),
# same extension -> unchanged
("loras/mymodel.safetensors", ".safetensors", "loras/mymodel.safetensors"),
# non-model suffix is treated as a stem, extension appended
("loras/my.model.v2", ".safetensors", "loras/my.model.v2.safetensors"),
# malformed (no slash) is returned untouched for parse_model_id to reject
("noslash", ".safetensors", "noslash"),
],
)
def test_apply_extension(model_id, ext, expected):
assert paths.apply_extension(model_id, ext) == expected
# ----- Content-Disposition filename parsing -----
@pytest.mark.parametrize(
"header,expected",
[
('attachment; filename="model.safetensors"', "model.safetensors"),
("attachment; filename=model.ckpt", "model.ckpt"),
# RFC 5987 form is preferred and percent-decoded
(
"attachment; filename=\"fallback.bin\"; filename*=UTF-8''my%20model.safetensors",
"my model.safetensors",
),
# directory components in a hostile header are stripped to the basename
('attachment; filename="../../etc/passwd"', "passwd"),
('attachment; filename="a\\\\b\\\\model.pt"', "model.pt"),
("inline", None),
(None, None),
],
)
def test_filename_from_content_disposition(header, expected):
from app.model_downloader.net.http import filename_from_content_disposition
assert filename_from_content_disposition(header) == expected
-24
View File
@@ -818,30 +818,6 @@ class TestExecution:
except urllib.error.HTTPError:
pass # Expected behavior
def test_cached_outputs_in_job_without_client_id(self, client: ComfyClient, builder: GraphBuilder):
g = builder
image = g.node("StubImage", content="BLACK", height=32, width=32, batch_size=1)
output = g.node("SaveImage", images=image.out(0))
# Prime the cache with a normal run.
client.run(g)
# Resubmit anonymously (no client_id) so output nodes are cache hits with no websocket client.
data = json.dumps({"prompt": g.finalize()}).encode('utf-8')
req = urllib.request.Request(f"http://{client.server_address}/prompt", data=data)
prompt_id = json.loads(urllib.request.urlopen(req).read())['prompt_id']
for _ in range(100):
job = client.get_job(prompt_id)
if job is not None and job['status'] not in ('pending', 'in_progress'):
break
time.sleep(0.1)
else:
raise AssertionError("Prompt did not complete in time")
assert job['status'] == 'completed'
assert output.id in job['outputs'], "Cached outputs must appear in job outputs without a client_id"
def _create_history_item(self, client, builder):
g = GraphBuilder(prefix="offset_test")
input_node = g.node(