mirror of
https://github.com/comfyanonymous/ComfyUI.git
synced 2026-07-23 00:08:08 +08:00
Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
0631e8325d | ||
|
|
81e7adf0d5 |
@@ -1,107 +0,0 @@
|
||||
"""
|
||||
Allow case-sensitive tag names.
|
||||
|
||||
Revision ID: 0005_allow_case_sensitive_tags
|
||||
Revises: 0004_drop_tag_type
|
||||
Create Date: 2026-06-16
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision = "0005_allow_case_sensitive_tags"
|
||||
down_revision = "0004_drop_tag_type"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
bind = op.get_bind()
|
||||
if bind.dialect.name == "sqlite":
|
||||
# SQLite cannot ALTER/DROP CHECK constraints. Recreate the small tag
|
||||
# vocabulary table without the lowercase constraint while preserving
|
||||
# existing tag names.
|
||||
op.execute("PRAGMA foreign_keys=OFF")
|
||||
try:
|
||||
op.execute(
|
||||
"CREATE TABLE tags_new ("
|
||||
"name VARCHAR(512) NOT NULL, "
|
||||
"CONSTRAINT pk_tags PRIMARY KEY (name)"
|
||||
")"
|
||||
)
|
||||
op.execute("INSERT INTO tags_new(name) SELECT name FROM tags")
|
||||
op.execute("DROP TABLE tags")
|
||||
op.execute("ALTER TABLE tags_new RENAME TO tags")
|
||||
finally:
|
||||
op.execute("PRAGMA foreign_keys=ON")
|
||||
return
|
||||
|
||||
op.drop_constraint("ck_tags_ck_tags_lowercase", "tags", type_="check")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# Existing mixed-case tags cannot satisfy the old constraint. Lowercase them
|
||||
# before restoring it, merging duplicate vocabulary/link rows that collide.
|
||||
bind = op.get_bind()
|
||||
|
||||
tag_names = [row[0] for row in bind.execute(sa.text("SELECT name FROM tags"))]
|
||||
existing_names = set(tag_names)
|
||||
lowercase_names = sorted({name.lower() for name in tag_names})
|
||||
missing_lowercase_rows = [
|
||||
{"name": name} for name in lowercase_names if name not in existing_names
|
||||
]
|
||||
if missing_lowercase_rows:
|
||||
bind.execute(sa.text("INSERT INTO tags(name) VALUES (:name)"), missing_lowercase_rows)
|
||||
|
||||
link_rows = bind.execute(
|
||||
sa.text(
|
||||
"SELECT asset_reference_id, tag_name, origin, added_at "
|
||||
"FROM asset_reference_tags "
|
||||
"ORDER BY asset_reference_id, tag_name"
|
||||
)
|
||||
).mappings()
|
||||
deduped_links = {}
|
||||
for row in link_rows:
|
||||
key = (row["asset_reference_id"], row["tag_name"].lower())
|
||||
deduped_links.setdefault(
|
||||
key,
|
||||
{
|
||||
"asset_reference_id": row["asset_reference_id"],
|
||||
"tag_name": row["tag_name"].lower(),
|
||||
"origin": row["origin"],
|
||||
"added_at": row["added_at"],
|
||||
},
|
||||
)
|
||||
|
||||
op.execute("DELETE FROM asset_reference_tags")
|
||||
if deduped_links:
|
||||
bind.execute(
|
||||
sa.text(
|
||||
"INSERT INTO asset_reference_tags "
|
||||
"(asset_reference_id, tag_name, origin, added_at) "
|
||||
"VALUES (:asset_reference_id, :tag_name, :origin, :added_at)"
|
||||
),
|
||||
list(deduped_links.values()),
|
||||
)
|
||||
op.execute("DELETE FROM tags WHERE name != lower(name)")
|
||||
|
||||
if bind.dialect.name == "sqlite":
|
||||
op.execute("PRAGMA foreign_keys=OFF")
|
||||
try:
|
||||
op.execute(
|
||||
"CREATE TABLE tags_new ("
|
||||
"name VARCHAR(512) NOT NULL, "
|
||||
"CONSTRAINT pk_tags PRIMARY KEY (name), "
|
||||
"CONSTRAINT ck_tags_lowercase CHECK (name = lower(name))"
|
||||
")"
|
||||
)
|
||||
op.execute("INSERT INTO tags_new(name) SELECT name FROM tags")
|
||||
op.execute("DROP TABLE tags")
|
||||
op.execute("ALTER TABLE tags_new RENAME TO tags")
|
||||
finally:
|
||||
op.execute("PRAGMA foreign_keys=ON")
|
||||
return
|
||||
|
||||
op.create_check_constraint(
|
||||
"ck_tags_ck_tags_lowercase", "tags", "name = lower(name)"
|
||||
)
|
||||
@@ -1,30 +0,0 @@
|
||||
"""
|
||||
Add loader_path column to asset_references.
|
||||
|
||||
Stores the in-root loader path (path relative to the storage root with the
|
||||
top-level model category dropped) derived from file_path at scan/ingest time,
|
||||
so the assets API can return it without re-resolving against every registered
|
||||
model-folder base on every request.
|
||||
|
||||
Revision ID: 0006_add_loader_path
|
||||
Revises: 0005_allow_case_sensitive_tags
|
||||
Create Date: 2026-07-02
|
||||
"""
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
revision = "0006_add_loader_path"
|
||||
down_revision = "0005_allow_case_sensitive_tags"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
with op.batch_alter_table("asset_references") as batch_op:
|
||||
batch_op.add_column(sa.Column("loader_path", sa.Text(), nullable=True))
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
with op.batch_alter_table("asset_references") as batch_op:
|
||||
batch_op.drop_column("loader_path")
|
||||
+12
-10
@@ -40,7 +40,6 @@ from app.assets.services import (
|
||||
upload_from_temp_path,
|
||||
)
|
||||
from app.assets.services.cursor import InvalidCursorError
|
||||
from app.assets.services.path_utils import compute_display_name
|
||||
from app.assets.services.tagging import list_tag_histogram
|
||||
|
||||
ROUTES = web.RouteTableDef()
|
||||
@@ -162,19 +161,11 @@ def _build_asset_response(result: schemas.AssetDetailResult | schemas.UploadResu
|
||||
preview_url = None
|
||||
else:
|
||||
preview_url = _build_preview_url_from_view(result.tags, result.ref.user_metadata)
|
||||
if result.ref.file_path:
|
||||
display_name = compute_display_name(result.ref.file_path)
|
||||
# In-root loader path (model category dropped): what model loaders consume.
|
||||
loader_path = result.ref.loader_path
|
||||
else:
|
||||
display_name, loader_path = None, None
|
||||
asset_content_hash = result.asset.hash if result.asset else None
|
||||
return schemas_out.Asset(
|
||||
id=result.ref.id,
|
||||
name=result.ref.name,
|
||||
hash=asset_content_hash,
|
||||
loader_path=loader_path,
|
||||
display_name=display_name,
|
||||
asset_hash=asset_content_hash,
|
||||
size=int(result.asset.size_bytes) if result.asset else None,
|
||||
mime_type=result.asset.mime_type if result.asset else None,
|
||||
@@ -428,6 +419,17 @@ async def upload_asset(request: web.Request) -> web.Response:
|
||||
400, "INVALID_BODY", f"Validation failed: {ve.json()}"
|
||||
)
|
||||
|
||||
if spec.tags and spec.tags[0] == "models":
|
||||
if (
|
||||
len(spec.tags) < 2
|
||||
or spec.tags[1] not in folder_paths.folder_names_and_paths
|
||||
):
|
||||
delete_temp_file_if_exists(parsed.tmp_path)
|
||||
category = spec.tags[1] if len(spec.tags) >= 2 else ""
|
||||
return _build_error_response(
|
||||
400, "INVALID_BODY", f"unknown models category '{category}'"
|
||||
)
|
||||
|
||||
try:
|
||||
# Fast path: hash exists, create AssetReference without writing anything
|
||||
if spec.hash and parsed.provided_hash_exists is True:
|
||||
@@ -471,7 +473,7 @@ async def upload_asset(request: web.Request) -> web.Response:
|
||||
return _build_error_response(400, e.code, str(e))
|
||||
except ValueError as e:
|
||||
delete_temp_file_if_exists(parsed.tmp_path)
|
||||
return _build_error_response(400, "INVALID_BODY", str(e))
|
||||
return _build_error_response(400, "BAD_REQUEST", str(e))
|
||||
except HashMismatchError as e:
|
||||
delete_temp_file_if_exists(parsed.tmp_path)
|
||||
return _build_error_response(400, "HASH_MISMATCH", str(e))
|
||||
|
||||
@@ -140,7 +140,7 @@ class CreateFromHashBody(BaseModel):
|
||||
if v is None:
|
||||
return []
|
||||
if isinstance(v, list):
|
||||
out = [str(t).strip() for t in v if str(t).strip()]
|
||||
out = [str(t).strip().lower() for t in v if str(t).strip()]
|
||||
seen = set()
|
||||
dedup = []
|
||||
for t in out:
|
||||
@@ -149,7 +149,7 @@ class CreateFromHashBody(BaseModel):
|
||||
dedup.append(t)
|
||||
return dedup
|
||||
if isinstance(v, str):
|
||||
return list(dict.fromkeys(t.strip() for t in v.split(",") if t.strip()))
|
||||
return [t.strip().lower() for t in v.split(",") if t.strip()]
|
||||
return []
|
||||
|
||||
|
||||
@@ -206,7 +206,7 @@ class TagsListQuery(BaseModel):
|
||||
if v is None:
|
||||
return v
|
||||
v = v.strip()
|
||||
return v or None
|
||||
return v.lower() or None
|
||||
|
||||
|
||||
class TagsAdd(BaseModel):
|
||||
@@ -220,7 +220,7 @@ class TagsAdd(BaseModel):
|
||||
for t in v:
|
||||
if not isinstance(t, str):
|
||||
raise TypeError("tags must be strings")
|
||||
tnorm = t.strip()
|
||||
tnorm = t.strip().lower()
|
||||
if tnorm:
|
||||
out.append(tnorm)
|
||||
seen = set()
|
||||
@@ -239,8 +239,8 @@ class TagsRemove(TagsAdd):
|
||||
class UploadAssetSpec(BaseModel):
|
||||
"""Upload Asset operation.
|
||||
|
||||
- tags: labels plus one destination role ('models'|'input'|'output') for new bytes;
|
||||
if role == 'models', exactly one model_type:<folder_name> tag is required
|
||||
- tags: optional list; if provided, first is root ('models'|'input'|'output');
|
||||
if root == 'models', second must be a valid category
|
||||
- name: display name
|
||||
- user_metadata: arbitrary JSON object (optional)
|
||||
- hash: optional canonical 'blake3:<hex>' for validation / fast-path
|
||||
@@ -309,7 +309,7 @@ class UploadAssetSpec(BaseModel):
|
||||
norm = []
|
||||
seen = set()
|
||||
for t in items:
|
||||
tnorm = str(t).strip()
|
||||
tnorm = str(t).strip().lower()
|
||||
if tnorm and tnorm not in seen:
|
||||
seen.add(tnorm)
|
||||
norm.append(tnorm)
|
||||
@@ -335,4 +335,14 @@ class UploadAssetSpec(BaseModel):
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _validate_order(self):
|
||||
if not self.tags:
|
||||
raise ValueError("at least one tag is required for uploads")
|
||||
root = self.tags[0]
|
||||
if root not in {"models", "input", "output"}:
|
||||
raise ValueError("first tag must be one of: models, input, output")
|
||||
if root == "models":
|
||||
if len(self.tags) < 2:
|
||||
raise ValueError(
|
||||
"models uploads require a category tag as the second tag"
|
||||
)
|
||||
return self
|
||||
|
||||
@@ -9,20 +9,8 @@ class Asset(BaseModel):
|
||||
``id`` here is the AssetReference id, not the content-addressed Asset id."""
|
||||
|
||||
id: str
|
||||
name: str = Field(
|
||||
...,
|
||||
deprecated=True,
|
||||
description="Reference label, often caller-provided or derived from the filename. Deprecated for storage path/display semantics; use `loader_path` and `display_name` when present.",
|
||||
)
|
||||
name: str
|
||||
hash: str | None = None
|
||||
loader_path: str | None = Field(
|
||||
default=None,
|
||||
description="The value a loader consumes to load this asset. `None` when no loader can resolve the file.",
|
||||
)
|
||||
display_name: str | None = Field(
|
||||
default=None,
|
||||
description="Human-facing label for the asset. Not unique.",
|
||||
)
|
||||
asset_hash: str | None = None
|
||||
size: int | None = None
|
||||
mime_type: str | None = None
|
||||
|
||||
@@ -140,6 +140,7 @@ async def parse_multipart_upload(
|
||||
provided_mime_type = ((await field.text()) or "").strip() or None
|
||||
elif fname == "preview_id":
|
||||
provided_preview_id = ((await field.text()) or "").strip() or None
|
||||
|
||||
if not file_present and not (provided_hash and provided_hash_exists):
|
||||
raise UploadError(
|
||||
400, "MISSING_FILE", "Form must include a 'file' part or a known 'hash'."
|
||||
|
||||
@@ -76,8 +76,6 @@ class AssetReference(Base):
|
||||
|
||||
# Cache state fields (from former AssetCacheState)
|
||||
file_path: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
# In-root loader path derived from file_path at scan/ingest time.
|
||||
loader_path: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
mtime_ns: Mapped[int | None] = mapped_column(BigInteger, nullable=True)
|
||||
needs_verify: Mapped[bool] = mapped_column(Boolean, nullable=False, default=False)
|
||||
is_missing: Mapped[bool] = mapped_column(Boolean, nullable=False, default=False)
|
||||
|
||||
@@ -650,7 +650,6 @@ def upsert_reference(
|
||||
name: str,
|
||||
mtime_ns: int,
|
||||
owner_id: str = "",
|
||||
loader_path: str | None = None,
|
||||
) -> tuple[bool, bool]:
|
||||
"""Upsert a reference by file_path. Returns (created, updated).
|
||||
|
||||
@@ -660,7 +659,6 @@ def upsert_reference(
|
||||
vals = {
|
||||
"asset_id": asset_id,
|
||||
"file_path": file_path,
|
||||
"loader_path": loader_path,
|
||||
"name": name,
|
||||
"owner_id": owner_id,
|
||||
"mtime_ns": int(mtime_ns),
|
||||
@@ -688,14 +686,13 @@ def upsert_reference(
|
||||
AssetReference.asset_id != asset_id,
|
||||
AssetReference.mtime_ns.is_(None),
|
||||
AssetReference.mtime_ns != int(mtime_ns),
|
||||
AssetReference.loader_path.is_distinct_from(loader_path),
|
||||
AssetReference.is_missing == True, # noqa: E712
|
||||
AssetReference.deleted_at.isnot(None),
|
||||
)
|
||||
)
|
||||
.values(
|
||||
asset_id=asset_id, mtime_ns=int(mtime_ns), loader_path=loader_path,
|
||||
is_missing=False, deleted_at=None, updated_at=now,
|
||||
asset_id=asset_id, mtime_ns=int(mtime_ns), is_missing=False,
|
||||
deleted_at=None, updated_at=now,
|
||||
)
|
||||
)
|
||||
res2 = session.execute(upd)
|
||||
|
||||
@@ -265,8 +265,6 @@ def list_tags_with_usage(
|
||||
order: str = "count_desc",
|
||||
owner_id: str = "",
|
||||
) -> tuple[list[tuple[str, str, int]], int]:
|
||||
prefix_filter = prefix.strip() if prefix else ""
|
||||
|
||||
counts_sq = (
|
||||
select(
|
||||
AssetReferenceTag.tag_name.label("tag_name"),
|
||||
@@ -295,8 +293,9 @@ def list_tags_with_usage(
|
||||
.join(counts_sq, counts_sq.c.tag_name == Tag.name, isouter=True)
|
||||
)
|
||||
|
||||
if prefix_filter:
|
||||
q = q.where(func.substr(Tag.name, 1, len(prefix_filter)) == prefix_filter)
|
||||
if prefix:
|
||||
escaped, esc = escape_sql_like_string(prefix.strip().lower())
|
||||
q = q.where(Tag.name.like(escaped + "%", escape=esc))
|
||||
|
||||
if not include_zero:
|
||||
q = q.where(func.coalesce(counts_sq.c.cnt, 0) > 0)
|
||||
@@ -307,8 +306,9 @@ def list_tags_with_usage(
|
||||
q = q.order_by(func.coalesce(counts_sq.c.cnt, 0).desc(), Tag.name.asc())
|
||||
|
||||
total_q = select(func.count()).select_from(Tag)
|
||||
if prefix_filter:
|
||||
total_q = total_q.where(func.substr(Tag.name, 1, len(prefix_filter)) == prefix_filter)
|
||||
if prefix:
|
||||
escaped, esc = escape_sql_like_string(prefix.strip().lower())
|
||||
total_q = total_q.where(Tag.name.like(escaped + "%", escape=esc))
|
||||
if not include_zero:
|
||||
visible_tags_sq = (
|
||||
select(AssetReferenceTag.tag_name)
|
||||
|
||||
@@ -41,10 +41,10 @@ def get_utc_now() -> datetime:
|
||||
def normalize_tags(tags: list[str] | None) -> list[str]:
|
||||
"""
|
||||
Normalize a list of tags by:
|
||||
- Stripping whitespace.
|
||||
- Removing exact duplicates while preserving order and case.
|
||||
- Stripping whitespace and converting to lowercase.
|
||||
- Removing duplicates.
|
||||
"""
|
||||
return list(dict.fromkeys(t.strip() for t in (tags or []) if (t or "").strip()))
|
||||
return list(dict.fromkeys(t.strip().lower() for t in (tags or []) if (t or "").strip()))
|
||||
|
||||
|
||||
def validate_blake3_hash(s: str) -> str:
|
||||
|
||||
@@ -36,7 +36,7 @@ from app.assets.services.hashing import HashCheckpoint, compute_blake3_hash
|
||||
from app.assets.services.image_dimensions import extract_image_dimensions
|
||||
from app.assets.services.metadata_extract import extract_file_metadata
|
||||
from app.assets.services.path_utils import (
|
||||
compute_loader_path,
|
||||
compute_relative_filename,
|
||||
get_comfy_models_folders,
|
||||
get_name_and_tags_from_asset_path,
|
||||
)
|
||||
@@ -63,7 +63,7 @@ RootType = Literal["models", "input", "output"]
|
||||
def get_prefixes_for_root(root: RootType) -> list[str]:
|
||||
if root == "models":
|
||||
bases: list[str] = []
|
||||
for _bucket, paths, _exts in get_comfy_models_folders():
|
||||
for _bucket, paths in get_comfy_models_folders():
|
||||
bases.extend(paths)
|
||||
return [os.path.abspath(p) for p in bases]
|
||||
if root == "input":
|
||||
@@ -81,7 +81,7 @@ def get_all_known_prefixes() -> list[str]:
|
||||
|
||||
def collect_models_files() -> list[str]:
|
||||
out: list[str] = []
|
||||
for folder_name, bases, _exts in get_comfy_models_folders():
|
||||
for folder_name, bases in get_comfy_models_folders():
|
||||
rel_files = folder_paths.get_filename_list(folder_name) or []
|
||||
for rel_path in rel_files:
|
||||
if not all(is_visible(part) for part in Path(rel_path).parts):
|
||||
@@ -308,7 +308,7 @@ def build_asset_specs(
|
||||
if not stat_p.st_size:
|
||||
continue
|
||||
name, tags = get_name_and_tags_from_asset_path(abs_p)
|
||||
rel_fname = compute_loader_path(abs_p)
|
||||
rel_fname = compute_relative_filename(abs_p)
|
||||
|
||||
# Extract metadata (tier 1: filesystem, tier 2: safetensors header)
|
||||
metadata = None
|
||||
@@ -430,7 +430,7 @@ def enrich_asset(
|
||||
return new_level
|
||||
|
||||
initial_mtime_ns = get_mtime_ns(stat_p)
|
||||
rel_fname = compute_loader_path(file_path)
|
||||
rel_fname = compute_relative_filename(file_path)
|
||||
mime_type: str | None = None
|
||||
metadata = None
|
||||
|
||||
|
||||
@@ -38,7 +38,7 @@ from app.assets.database.queries import (
|
||||
update_reference_updated_at,
|
||||
)
|
||||
from app.assets.helpers import select_best_live_path
|
||||
from app.assets.services.path_utils import compute_loader_path
|
||||
from app.assets.services.path_utils import compute_relative_filename
|
||||
from app.assets.services.schemas import (
|
||||
AssetData,
|
||||
AssetDetailResult,
|
||||
@@ -91,7 +91,7 @@ def update_asset_metadata(
|
||||
update_reference_name(session, reference_id=reference_id, name=name)
|
||||
touched = True
|
||||
|
||||
computed_filename = compute_loader_path(ref.file_path) if ref.file_path else None
|
||||
computed_filename = compute_relative_filename(ref.file_path) if ref.file_path else None
|
||||
|
||||
new_meta: dict | None = None
|
||||
if user_metadata is not None:
|
||||
|
||||
@@ -56,7 +56,6 @@ class ReferenceRow(TypedDict):
|
||||
id: str
|
||||
asset_id: str
|
||||
file_path: str
|
||||
loader_path: str | None
|
||||
mtime_ns: int
|
||||
owner_id: str
|
||||
name: str
|
||||
@@ -135,14 +134,6 @@ def batch_insert_seed_assets(
|
||||
|
||||
for spec in specs:
|
||||
absolute_path = os.path.abspath(spec["abs_path"])
|
||||
existing_asset_id = path_to_asset_id.get(absolute_path)
|
||||
if existing_asset_id is not None:
|
||||
existing_tags = asset_id_to_ref_data[existing_asset_id]["tags"]
|
||||
asset_id_to_ref_data[existing_asset_id]["tags"] = list(
|
||||
dict.fromkeys([*existing_tags, *spec["tags"]])
|
||||
)
|
||||
continue
|
||||
|
||||
asset_id = str(uuid.uuid4())
|
||||
reference_id = str(uuid.uuid4())
|
||||
absolute_path_list.append(absolute_path)
|
||||
@@ -173,8 +164,6 @@ def batch_insert_seed_assets(
|
||||
"id": reference_id,
|
||||
"asset_id": asset_id,
|
||||
"file_path": absolute_path,
|
||||
# spec["fname"] is compute_loader_path(abs_path) from build_asset_specs.
|
||||
"loader_path": spec["fname"],
|
||||
"mtime_ns": spec["mtime_ns"],
|
||||
"owner_id": owner_id,
|
||||
"name": spec["info_name"],
|
||||
|
||||
@@ -33,9 +33,8 @@ from app.assets.services.bulk_ingest import batch_insert_seed_assets
|
||||
from app.assets.services.file_utils import get_size_and_mtime_ns
|
||||
from app.assets.services.image_dimensions import extract_image_dimensions
|
||||
from app.assets.services.path_utils import (
|
||||
compute_loader_path,
|
||||
compute_relative_filename,
|
||||
get_name_and_tags_from_asset_path,
|
||||
get_path_derived_tags_from_path,
|
||||
resolve_destination_from_tags,
|
||||
validate_path_within_base,
|
||||
)
|
||||
@@ -92,7 +91,6 @@ def _ingest_file_from_path(
|
||||
name=info_name or os.path.basename(locator),
|
||||
mtime_ns=mtime_ns,
|
||||
owner_id=owner_id,
|
||||
loader_path=compute_loader_path(locator),
|
||||
)
|
||||
|
||||
# Get the reference we just created/updated
|
||||
@@ -103,32 +101,17 @@ def _ingest_file_from_path(
|
||||
if preview_id and ref.preview_id != preview_id:
|
||||
ref.preview_id = preview_id
|
||||
|
||||
try:
|
||||
backend_tags = get_path_derived_tags_from_path(locator)
|
||||
except ValueError:
|
||||
backend_tags = []
|
||||
caller_tags = normalize_tags(tags)
|
||||
backend_tags = normalize_tags(backend_tags)
|
||||
all_tags = normalize_tags([*caller_tags, *backend_tags])
|
||||
if all_tags:
|
||||
norm = normalize_tags(list(tags))
|
||||
if norm:
|
||||
if require_existing_tags:
|
||||
validate_tags_exist(session, all_tags)
|
||||
if backend_tags:
|
||||
add_tags_to_reference(
|
||||
session,
|
||||
reference_id=reference_id,
|
||||
tags=backend_tags,
|
||||
origin="automatic",
|
||||
create_if_missing=not require_existing_tags,
|
||||
)
|
||||
if caller_tags:
|
||||
add_tags_to_reference(
|
||||
session,
|
||||
reference_id=reference_id,
|
||||
tags=caller_tags,
|
||||
origin=tag_origin,
|
||||
create_if_missing=not require_existing_tags,
|
||||
)
|
||||
validate_tags_exist(session, norm)
|
||||
add_tags_to_reference(
|
||||
session,
|
||||
reference_id=reference_id,
|
||||
tags=norm,
|
||||
origin=tag_origin,
|
||||
create_if_missing=not require_existing_tags,
|
||||
)
|
||||
|
||||
_update_metadata_with_filename(
|
||||
session,
|
||||
@@ -245,7 +228,7 @@ def ingest_existing_file(
|
||||
"mtime_ns": mtime_ns,
|
||||
"info_name": name,
|
||||
"tags": tags,
|
||||
"fname": compute_loader_path(abs_path),
|
||||
"fname": os.path.basename(abs_path),
|
||||
"metadata": None,
|
||||
"hash": None,
|
||||
"mime_type": mime_type,
|
||||
@@ -305,7 +288,7 @@ def _register_existing_asset(
|
||||
return result
|
||||
|
||||
new_meta = dict(user_metadata)
|
||||
computed_filename = compute_loader_path(ref.file_path) if ref.file_path else None
|
||||
computed_filename = compute_relative_filename(ref.file_path) if ref.file_path else None
|
||||
if computed_filename:
|
||||
new_meta["filename"] = computed_filename
|
||||
|
||||
@@ -352,7 +335,7 @@ def _update_metadata_with_filename(
|
||||
current_metadata: dict | None,
|
||||
user_metadata: dict[str, Any],
|
||||
) -> None:
|
||||
computed_filename = compute_loader_path(file_path) if file_path else None
|
||||
computed_filename = compute_relative_filename(file_path) if file_path else None
|
||||
|
||||
current_meta = current_metadata or {}
|
||||
new_meta = dict(current_meta)
|
||||
@@ -491,10 +474,6 @@ def upload_from_temp_path(
|
||||
existing = get_asset_by_hash(session, asset_hash=asset_hash)
|
||||
|
||||
if existing is not None:
|
||||
# Once content is already known, duplicate byte uploads are treated as
|
||||
# reference-only creation. Request tags are labels only here: do not
|
||||
# require upload destination tags, do not move bytes, and do not
|
||||
# synthesize path-derived classification or uploaded provenance.
|
||||
with contextlib.suppress(Exception):
|
||||
if temp_path and os.path.exists(temp_path):
|
||||
os.remove(temp_path)
|
||||
@@ -556,7 +535,7 @@ def upload_from_temp_path(
|
||||
owner_id=owner_id,
|
||||
preview_id=preview_id,
|
||||
user_metadata=user_metadata or {},
|
||||
tags=[*(tags or []), "uploaded"],
|
||||
tags=tags,
|
||||
tag_origin="manual",
|
||||
require_existing_tags=False,
|
||||
)
|
||||
@@ -590,19 +569,15 @@ def register_file_in_place(
|
||||
) -> UploadResult:
|
||||
"""Register an already-saved file in the asset database without moving it.
|
||||
|
||||
This helper is used by upload paths that have already written bytes before
|
||||
registering the file, so it records the same ``uploaded`` tag as the
|
||||
multipart byte-upload path.
|
||||
|
||||
Tags are derived from trusted filesystem classification and merged with any
|
||||
caller-provided tags, matching the behavior of the scanner.
|
||||
Tags are derived from the filesystem path (root category + subfolder names),
|
||||
merged with any caller-provided tags, matching the behavior of the scanner.
|
||||
If the path is not under a known root, only the caller-provided tags are used.
|
||||
"""
|
||||
try:
|
||||
_, path_tags = get_name_and_tags_from_asset_path(abs_path)
|
||||
except ValueError:
|
||||
path_tags = []
|
||||
merged_tags = normalize_tags([*path_tags, *tags, "uploaded"])
|
||||
merged_tags = normalize_tags([*path_tags, *tags])
|
||||
|
||||
try:
|
||||
digest, _ = hashing.compute_blake3_hash(abs_path)
|
||||
|
||||
@@ -3,66 +3,59 @@ from pathlib import Path
|
||||
from typing import Literal
|
||||
|
||||
import folder_paths
|
||||
from app.assets.helpers import normalize_tags
|
||||
|
||||
|
||||
_NON_MODEL_FOLDER_NAMES = frozenset({"configs", "custom_nodes"})
|
||||
_KNOWN_SUBFOLDER_TAGS = frozenset({"3d", "pasted", "painter", "threed", "webcam"})
|
||||
_NON_MODEL_FOLDER_NAMES = frozenset({"custom_nodes"})
|
||||
|
||||
|
||||
def get_comfy_models_folders() -> list[tuple[str, list[str], set[str]]]:
|
||||
"""Build list of (folder_name, base_paths[], extensions) for all model locations.
|
||||
def get_comfy_models_folders() -> list[tuple[str, list[str]]]:
|
||||
"""Build list of (folder_name, base_paths[]) for all model locations.
|
||||
|
||||
Includes every category registered in folder_names_and_paths,
|
||||
regardless of whether its paths are under the main models_dir,
|
||||
but excludes non-model entries like configs and custom_nodes.
|
||||
|
||||
An empty extensions set means the category accepts any extension,
|
||||
matching folder_paths.filter_files_extensions semantics.
|
||||
but excludes non-model entries like custom_nodes.
|
||||
"""
|
||||
targets: list[tuple[str, list[str], set[str]]] = []
|
||||
targets: list[tuple[str, list[str]]] = []
|
||||
for name, values in folder_paths.folder_names_and_paths.items():
|
||||
if name in _NON_MODEL_FOLDER_NAMES:
|
||||
continue
|
||||
paths, exts = values[0], values[1]
|
||||
paths, _exts = values[0], values[1]
|
||||
if paths:
|
||||
targets.append((name, paths, set(exts)))
|
||||
targets.append((name, paths))
|
||||
return targets
|
||||
|
||||
|
||||
def resolve_destination_from_tags(tags: list[str]) -> tuple[str, list[str]]:
|
||||
"""Validates and maps upload routing tags -> (base_dir, subdirs_for_fs).
|
||||
|
||||
The request tags are only used to choose the write destination. Extra tags
|
||||
remain labels; they do not become path components or trusted classification.
|
||||
"""
|
||||
destination_roles = [t for t in tags if t in {"input", "models", "output"}]
|
||||
if len(destination_roles) != 1:
|
||||
raise ValueError("uploads require exactly one destination role: input, models, or output")
|
||||
|
||||
root = destination_roles[0]
|
||||
"""Validates and maps tags -> (base_dir, subdirs_for_fs)"""
|
||||
if not tags:
|
||||
raise ValueError("tags must not be empty")
|
||||
root = tags[0].lower()
|
||||
if root == "models":
|
||||
model_type_tags = [t for t in tags if t.startswith("model_type:")]
|
||||
if len(model_type_tags) != 1:
|
||||
raise ValueError("models uploads require exactly one model_type:<folder_name> tag")
|
||||
folder_name = model_type_tags[0].split(":", 1)[1]
|
||||
if not folder_name:
|
||||
raise ValueError("models uploads require exactly one model_type:<folder_name> tag")
|
||||
model_folder_paths = {
|
||||
name: paths for name, paths, _exts in get_comfy_models_folders()
|
||||
}
|
||||
if len(tags) < 2:
|
||||
raise ValueError("at least two tags required for model asset")
|
||||
try:
|
||||
bases = model_folder_paths[folder_name]
|
||||
bases = folder_paths.folder_names_and_paths[tags[1]][0]
|
||||
except KeyError:
|
||||
raise ValueError(f"unknown model category '{folder_name}'")
|
||||
raise ValueError(f"unknown model category '{tags[1]}'")
|
||||
if not bases:
|
||||
raise ValueError(f"no base path configured for category '{folder_name}'")
|
||||
raise ValueError(f"no base path configured for category '{tags[1]}'")
|
||||
base_dir = os.path.abspath(bases[0])
|
||||
raw_subdirs = tags[2:]
|
||||
elif root == "input":
|
||||
base_dir = os.path.abspath(folder_paths.get_input_directory())
|
||||
else:
|
||||
raw_subdirs = tags[1:]
|
||||
elif root == "output":
|
||||
base_dir = os.path.abspath(folder_paths.get_output_directory())
|
||||
raw_subdirs = tags[1:]
|
||||
else:
|
||||
raise ValueError(f"unknown root tag '{tags[0]}'; expected 'models', 'input', or 'output'")
|
||||
_sep_chars = frozenset(("/", "\\", os.sep))
|
||||
for i in raw_subdirs:
|
||||
if i in (".", "..") or _sep_chars & set(i):
|
||||
raise ValueError("invalid path component in tags")
|
||||
|
||||
return base_dir, []
|
||||
return base_dir, raw_subdirs if raw_subdirs else []
|
||||
|
||||
|
||||
def validate_path_within_base(candidate: str, base: str) -> None:
|
||||
@@ -72,79 +65,14 @@ def validate_path_within_base(candidate: str, base: str) -> None:
|
||||
raise ValueError("destination escapes base directory")
|
||||
|
||||
|
||||
def _compute_relative_path(child: str, parent: str) -> str:
|
||||
rel = os.path.relpath(os.path.abspath(child), os.path.abspath(parent))
|
||||
if rel == ".":
|
||||
return ""
|
||||
return rel.replace(os.sep, "/")
|
||||
|
||||
|
||||
def _is_relative_to(child: str, parent: str) -> bool:
|
||||
return Path(os.path.abspath(child)).is_relative_to(os.path.abspath(parent))
|
||||
|
||||
|
||||
def compute_asset_response_paths(file_path: str) -> tuple[str, str | None] | None:
|
||||
"""Return (logical_path, display_name) for a file path.
|
||||
|
||||
``logical_path`` is the internal namespaced storage locator (e.g.
|
||||
``models/checkpoints/foo/bar.safetensors``); ``display_name`` is the
|
||||
human-facing label below that namespace, served on Asset responses. These
|
||||
are storage locators, not model-loader namespaces. Registered model-folder
|
||||
membership is represented by backend tags such as
|
||||
``model_type:<folder_name>``; these paths only use known storage roots.
|
||||
def compute_relative_filename(file_path: str) -> str | None:
|
||||
"""
|
||||
fp_abs = os.path.abspath(file_path)
|
||||
candidates: list[tuple[int, int, str, str]] = []
|
||||
|
||||
for order, (namespace, base) in enumerate(
|
||||
(
|
||||
("input", folder_paths.get_input_directory()),
|
||||
("output", folder_paths.get_output_directory()),
|
||||
("temp", folder_paths.get_temp_directory()),
|
||||
("models", getattr(folder_paths, "models_dir", "")),
|
||||
)
|
||||
):
|
||||
if not base:
|
||||
continue
|
||||
base_abs = os.path.abspath(base)
|
||||
if _is_relative_to(fp_abs, base_abs):
|
||||
candidates.append((len(base_abs), -order, namespace, base_abs))
|
||||
|
||||
if not candidates:
|
||||
return None
|
||||
|
||||
_base_len, _order, namespace, base = max(candidates)
|
||||
rel = _compute_relative_path(fp_abs, base)
|
||||
public_path = f"{namespace}/{rel}" if rel else namespace
|
||||
return public_path, rel or None
|
||||
|
||||
|
||||
def compute_display_name(file_path: str) -> str | None:
|
||||
"""Return the asset's `display_name`, or None for unknown paths."""
|
||||
result = compute_asset_response_paths(file_path)
|
||||
return result[1] if result else None
|
||||
|
||||
|
||||
def compute_logical_path(file_path: str) -> str | None:
|
||||
"""Return the internal namespaced storage locator, or None for unknown paths."""
|
||||
result = compute_asset_response_paths(file_path)
|
||||
return result[0] if result else None
|
||||
|
||||
|
||||
def compute_loader_path(file_path: str) -> str | None:
|
||||
"""
|
||||
Return the asset's in-root loader path: the path relative to the last
|
||||
well-known folder (the model category), using forward slashes, eg:
|
||||
Return the model's path relative to the last well-known folder (the model category),
|
||||
using forward slashes, eg:
|
||||
/.../models/checkpoints/flux/123/flux.safetensors -> "flux/123/flux.safetensors"
|
||||
/.../models/text_encoders/clip_g.safetensors -> "clip_g.safetensors"
|
||||
|
||||
This is the value model loaders consume (the model category is dropped). It
|
||||
is persisted as ``AssetReference.loader_path`` and served as the public
|
||||
Asset response `loader_path` field. The human-facing `display_name` comes
|
||||
from compute_asset_response_paths().
|
||||
|
||||
For input/output/temp paths the full path relative to that root is returned.
|
||||
For paths outside any known root, returns None.
|
||||
For non-model paths, returns None.
|
||||
"""
|
||||
try:
|
||||
root_category, rel_path = get_asset_category_and_relative_path(file_path)
|
||||
@@ -188,10 +116,9 @@ def get_asset_category_and_relative_path(
|
||||
def _compute_relative(child: str, parent: str) -> str:
|
||||
# Normalize relative path, stripping any leading ".." components
|
||||
# by anchoring to root (os.sep) then computing relpath back from it.
|
||||
rel = os.path.relpath(
|
||||
return os.path.relpath(
|
||||
os.path.join(os.sep, os.path.relpath(child, parent)), os.sep
|
||||
)
|
||||
return "" if rel == "." else rel.replace(os.sep, "/")
|
||||
|
||||
# 1) input
|
||||
input_base = os.path.abspath(folder_paths.get_input_directory())
|
||||
@@ -209,14 +136,8 @@ def get_asset_category_and_relative_path(
|
||||
return "temp", _compute_relative(fp_abs, temp_base)
|
||||
|
||||
# 4) models (check deepest matching base to avoid ambiguity)
|
||||
ext = os.path.splitext(fp_abs)[1].lower()
|
||||
best: tuple[int, str, str] | None = None # (base_len, bucket, rel_inside_bucket)
|
||||
for bucket, bases, extensions in get_comfy_models_folders():
|
||||
# A bucket only lists files within its extension set (empty set
|
||||
# accepts any extension), so a bucket that cannot load the file
|
||||
# must not contribute a loader path.
|
||||
if extensions and ext not in extensions:
|
||||
continue
|
||||
for bucket, bases in get_comfy_models_folders():
|
||||
for b in bases:
|
||||
base_abs = os.path.abspath(b)
|
||||
if not _check_is_within(fp_abs, base_abs):
|
||||
@@ -228,111 +149,25 @@ def get_asset_category_and_relative_path(
|
||||
if best is not None:
|
||||
_, bucket, rel_inside = best
|
||||
combined = os.path.join(bucket, rel_inside)
|
||||
normalized = os.path.relpath(os.path.join(os.sep, combined), os.sep)
|
||||
return "models", normalized.replace(os.sep, "/")
|
||||
return "models", os.path.relpath(os.path.join(os.sep, combined), os.sep)
|
||||
|
||||
raise ValueError(
|
||||
f"Path is not within input, output, temp, or configured model bases: {file_path}"
|
||||
)
|
||||
|
||||
|
||||
def get_backend_system_tags_from_path(path: str) -> list[str]:
|
||||
"""Return trusted backend tags derived from current filesystem facts.
|
||||
|
||||
The returned tags are only the backend-generated system tags: ``models``,
|
||||
``model_type:<folder_name>``, ``input``, ``output``, and ``temp``. Model
|
||||
type tags are based on registered folder names, not path components.
|
||||
|
||||
A ``model_type:<folder_name>`` tag is only emitted when the file's
|
||||
extension is accepted by that folder's registered extension set, so
|
||||
categories sharing a base directory tag only the files they can
|
||||
actually load. Files under a model base whose extension matches no
|
||||
category still get the ``models`` tag.
|
||||
"""
|
||||
fp_abs = os.path.abspath(path)
|
||||
fp_path = Path(fp_abs)
|
||||
tags: list[str] = []
|
||||
|
||||
def _add(tag: str) -> None:
|
||||
if tag not in tags:
|
||||
tags.append(tag)
|
||||
|
||||
for role, base in (
|
||||
("input", folder_paths.get_input_directory()),
|
||||
("output", folder_paths.get_output_directory()),
|
||||
("temp", folder_paths.get_temp_directory()),
|
||||
):
|
||||
if fp_path.is_relative_to(os.path.abspath(base)):
|
||||
_add(role)
|
||||
|
||||
ext = os.path.splitext(fp_abs)[1].lower()
|
||||
model_types: list[str] = []
|
||||
under_models_base = False
|
||||
for folder_name, bases, extensions in get_comfy_models_folders():
|
||||
for base in bases:
|
||||
if fp_path.is_relative_to(os.path.abspath(base)):
|
||||
under_models_base = True
|
||||
# Empty set accepts any extension, matching
|
||||
# folder_paths.filter_files_extensions semantics.
|
||||
if not extensions or ext in extensions:
|
||||
model_types.append(folder_name)
|
||||
break
|
||||
|
||||
if under_models_base:
|
||||
_add("models")
|
||||
for folder_name in model_types:
|
||||
_add(f"model_type:{folder_name}")
|
||||
|
||||
if not tags:
|
||||
raise ValueError(
|
||||
f"Path is not within input, output, temp, or configured model bases: {path}"
|
||||
)
|
||||
return tags
|
||||
|
||||
|
||||
def get_known_subfolder_tags(subfolder: str | None) -> list[str]:
|
||||
"""Return tags for known UI/input subfolder names."""
|
||||
if subfolder in _KNOWN_SUBFOLDER_TAGS:
|
||||
return [subfolder]
|
||||
return []
|
||||
|
||||
|
||||
def get_known_input_subfolder_tags_from_path(path: str) -> list[str]:
|
||||
"""Return known input-layout tags for files in canonical input subfolders.
|
||||
|
||||
These are compatibility tags for current UI-origin input directories such as
|
||||
``pasted`` and ``webcam``. They are intentionally narrow: only files directly
|
||||
inside a known top-level input directory receive the matching tag.
|
||||
"""
|
||||
fp_abs = os.path.abspath(path)
|
||||
input_base = os.path.abspath(folder_paths.get_input_directory())
|
||||
if not Path(fp_abs).is_relative_to(input_base):
|
||||
return []
|
||||
|
||||
rel = os.path.relpath(fp_abs, input_base)
|
||||
parts = Path(rel).parts
|
||||
if len(parts) == 2:
|
||||
return get_known_subfolder_tags(parts[0])
|
||||
return []
|
||||
|
||||
|
||||
def get_path_derived_tags_from_path(path: str) -> list[str]:
|
||||
"""Return all backend-derived tags for an asset path."""
|
||||
tags = get_backend_system_tags_from_path(path)
|
||||
for tag in get_known_input_subfolder_tags_from_path(path):
|
||||
if tag not in tags:
|
||||
tags.append(tag)
|
||||
return tags
|
||||
|
||||
|
||||
def get_name_and_tags_from_asset_path(file_path: str) -> tuple[str, list[str]]:
|
||||
"""Return (name, tags) derived from a filesystem path.
|
||||
|
||||
- name: base filename with extension
|
||||
- tags: backend-derived tags from root/model classification and known input
|
||||
subfolder layout conventions
|
||||
- tags: [root_category] + parent folder names in order
|
||||
|
||||
Raises:
|
||||
ValueError: path does not belong to any known root.
|
||||
"""
|
||||
return Path(file_path).name, get_path_derived_tags_from_path(file_path)
|
||||
root_category, some_path = get_asset_category_and_relative_path(file_path)
|
||||
p = Path(some_path)
|
||||
parent_parts = [
|
||||
part for part in p.parent.parts if part not in (".", "..", p.anchor)
|
||||
]
|
||||
return p.name, list(dict.fromkeys(normalize_tags([root_category, *parent_parts])))
|
||||
|
||||
@@ -25,7 +25,6 @@ class ReferenceData:
|
||||
preview_id: str | None
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
loader_path: str | None = None
|
||||
system_metadata: dict[str, Any] | None = None
|
||||
job_id: str | None = None
|
||||
last_access_time: datetime | None = None
|
||||
@@ -94,7 +93,6 @@ def extract_reference_data(ref: AssetReference) -> ReferenceData:
|
||||
id=ref.id,
|
||||
name=ref.name,
|
||||
file_path=ref.file_path,
|
||||
loader_path=ref.loader_path,
|
||||
user_metadata=ref.user_metadata,
|
||||
preview_id=ref.preview_id,
|
||||
system_metadata=ref.system_metadata,
|
||||
|
||||
@@ -1,46 +0,0 @@
|
||||
"""Runtime config the frontend reads from /features to follow --comfy-api-base.
|
||||
|
||||
For a non-prod comfy.org backend (staging or an ephemeral preview env), "/features" exposes the api and
|
||||
platform base so the frontend talks to it without a rebuild, plus the Firebase environment it should use.
|
||||
Prod bases are left alone and keep their build-time defaults.
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from comfy.cli_args import args
|
||||
|
||||
_STAGING_API_HOST = "stagingapi.comfy.org"
|
||||
_TESTENV_HOST_SUFFIX = ".testenvs.comfy.org"
|
||||
_STAGING_PLATFORM_BASE_URL = "https://stagingplatform.comfy.org"
|
||||
|
||||
|
||||
def _is_staging_tier(host: str) -> bool:
|
||||
return host == _STAGING_API_HOST or host.endswith(_TESTENV_HOST_SUFFIX)
|
||||
|
||||
|
||||
def normalize_comfy_api_base(url: str) -> str:
|
||||
"""Rewrite a testenv's friendly main host to its comfy-api '-registry' sibling."""
|
||||
parsed = urlparse(url)
|
||||
host = parsed.hostname or ""
|
||||
if not host.endswith(_TESTENV_HOST_SUFFIX):
|
||||
return url
|
||||
label = host[: -len(_TESTENV_HOST_SUFFIX)]
|
||||
if label.endswith("-registry"):
|
||||
return url
|
||||
return f"{parsed.scheme or 'https'}://{label}-registry{_TESTENV_HOST_SUFFIX}"
|
||||
|
||||
|
||||
def environment_overrides_for_base(base_url: str) -> dict[str, Any] | None:
|
||||
"""The /features overrides for a staging-tier base, or None for prod."""
|
||||
if not _is_staging_tier(urlparse(base_url).hostname or ""):
|
||||
return None
|
||||
return {
|
||||
"comfy_api_base_url": normalize_comfy_api_base(base_url).rstrip("/"),
|
||||
"comfy_platform_base_url": _STAGING_PLATFORM_BASE_URL,
|
||||
"firebase_env": "dev",
|
||||
}
|
||||
|
||||
|
||||
def get_environment_overrides() -> dict[str, Any] | None:
|
||||
return environment_overrides_for_base(getattr(args, "comfy_api_base", "") or "")
|
||||
@@ -540,6 +540,9 @@ class Wan21(LatentFormat):
|
||||
latents_std = self.latents_std.to(latent.device, latent.dtype)
|
||||
return latent * latents_std / self.scale_factor + latents_mean
|
||||
|
||||
class LingBotVideo(Wan21):
|
||||
pass
|
||||
|
||||
class Wan22(Wan21):
|
||||
latent_channels = 48
|
||||
latent_dimensions = 3
|
||||
|
||||
@@ -0,0 +1,518 @@
|
||||
import math
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from comfy.ldm.flux.math import apply_rope1, rope
|
||||
from comfy.ldm.flux.layers import timestep_embedding
|
||||
from comfy.ldm.modules.attention import optimized_attention
|
||||
|
||||
|
||||
class LingBotVideoRotaryEmbedding(nn.Module):
|
||||
def __init__(self, axes_dims: Tuple[int, ...], axes_lens: Tuple[int, ...], theta: float):
|
||||
super().__init__()
|
||||
self.axes_dims = tuple(axes_dims)
|
||||
self.theta = theta
|
||||
|
||||
def forward(self, position_ids: torch.Tensor) -> torch.Tensor:
|
||||
return torch.cat(
|
||||
[rope(position_ids[None, :, i], self.axes_dims[i], self.theta) for i in range(len(self.axes_dims))],
|
||||
dim=-3,
|
||||
).squeeze(0)
|
||||
|
||||
|
||||
def make_joint_position_ids(
|
||||
text_len: int, grid_t: int, grid_h: int, grid_w: int, device: torch.device, padded_text_len: Optional[int] = None
|
||||
) -> torch.Tensor:
|
||||
"""3D positions in [video; text] order. Text t-axis is 1..text_len; video t-axis starts at text_len+1.
|
||||
|
||||
Matches patchify_and_embed: cap start (1,0,0); vision start (cap_len+1,0,0);
|
||||
freqs ordered with x first and cap second (same order as cat_interleave).
|
||||
"""
|
||||
tt = torch.arange(grid_t, device=device, dtype=torch.int32) + (text_len + 1)
|
||||
hh = torch.arange(grid_h, device=device, dtype=torch.int32)
|
||||
ww = torch.arange(grid_w, device=device, dtype=torch.int32)
|
||||
grid = torch.stack(torch.meshgrid(tt, hh, ww, indexing="ij"), dim=-1).flatten(0, 2)
|
||||
if padded_text_len is None:
|
||||
padded_text_len = text_len
|
||||
text_t = torch.arange(padded_text_len, device=device, dtype=torch.int32) + 1
|
||||
text_pos = torch.stack(
|
||||
[text_t, torch.zeros_like(text_t), torch.zeros_like(text_t)], dim=-1
|
||||
)
|
||||
return torch.cat([grid, text_pos], dim=0) # (Nx + L, 3)
|
||||
|
||||
|
||||
class LingBotVideoTimestepEmbedding(nn.Module):
|
||||
def __init__(self, in_channels, time_embed_dim, bias=True, device=None, dtype=None, operations=None):
|
||||
super().__init__()
|
||||
self.linear_1 = operations.Linear(in_channels, time_embed_dim, bias=bias, device=device, dtype=dtype)
|
||||
self.act = nn.SiLU()
|
||||
self.linear_2 = operations.Linear(time_embed_dim, time_embed_dim, bias=bias, device=device, dtype=dtype)
|
||||
|
||||
def forward(self, sample):
|
||||
return self.linear_2(self.act(self.linear_1(sample)))
|
||||
|
||||
|
||||
class LingBotVideoTextEmbedder(nn.Module):
|
||||
"""Matches CondProjection: RMSNorm(text_dim, eps=1e-6 fixed) -> Linear-SiLU-Linear."""
|
||||
|
||||
def __init__(self, text_dim: int, hidden_size: int, device=None, dtype=None, operations=None):
|
||||
super().__init__()
|
||||
self.norm = operations.RMSNorm(text_dim, eps=1e-6, elementwise_affine=True, device=device, dtype=dtype)
|
||||
self.linear_1 = operations.Linear(text_dim, hidden_size, bias=True, device=device, dtype=dtype)
|
||||
self.linear_2 = operations.Linear(hidden_size, hidden_size, bias=True, device=device, dtype=dtype)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
x = self.norm(x)
|
||||
return self.linear_2(F.silu(self.linear_1(x)))
|
||||
|
||||
|
||||
class LingBotVideoAttention(nn.Module):
|
||||
def __init__(self, hidden_size, num_heads, norm_eps, qkv_bias, out_bias, device=None, dtype=None, operations=None):
|
||||
super().__init__()
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = hidden_size // num_heads
|
||||
self.to_q = operations.Linear(hidden_size, hidden_size, bias=qkv_bias, device=device, dtype=dtype)
|
||||
self.to_k = operations.Linear(hidden_size, hidden_size, bias=qkv_bias, device=device, dtype=dtype)
|
||||
self.to_v = operations.Linear(hidden_size, hidden_size, bias=qkv_bias, device=device, dtype=dtype)
|
||||
self.norm_q = operations.RMSNorm(self.head_dim, eps=norm_eps, elementwise_affine=True, device=device, dtype=dtype)
|
||||
self.norm_k = operations.RMSNorm(self.head_dim, eps=norm_eps, elementwise_affine=True, device=device, dtype=dtype)
|
||||
self.to_out = operations.Linear(hidden_size, hidden_size, bias=out_bias, device=device, dtype=dtype)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x,
|
||||
rotary_emb,
|
||||
attention_mask=None,
|
||||
transformer_options={},
|
||||
):
|
||||
q = self.to_q(x).unflatten(2, (self.num_heads, self.head_dim))
|
||||
k = self.to_k(x).unflatten(2, (self.num_heads, self.head_dim))
|
||||
v = self.to_v(x).unflatten(2, (self.num_heads, self.head_dim))
|
||||
q = apply_rope1(self.norm_q(q), rotary_emb)
|
||||
k = apply_rope1(self.norm_k(k), rotary_emb)
|
||||
out = optimized_attention(
|
||||
q.transpose(1, 2),
|
||||
k.transpose(1, 2),
|
||||
v.transpose(1, 2),
|
||||
heads=self.num_heads,
|
||||
mask=attention_mask,
|
||||
skip_reshape=True,
|
||||
transformer_options=transformer_options,
|
||||
)
|
||||
return self.to_out(out)
|
||||
|
||||
|
||||
class LingBotVideoMLP(nn.Module):
|
||||
def __init__(self, hidden_size, intermediate_size, device=None, dtype=None, operations=None):
|
||||
super().__init__()
|
||||
self.gate_proj = operations.Linear(hidden_size, intermediate_size, bias=False, device=device, dtype=dtype)
|
||||
self.up_proj = operations.Linear(hidden_size, intermediate_size, bias=False, device=device, dtype=dtype)
|
||||
self.down_proj = operations.Linear(intermediate_size, hidden_size, bias=False, device=device, dtype=dtype)
|
||||
|
||||
def forward(self, x):
|
||||
return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x))
|
||||
|
||||
|
||||
class LingBotVideoRouter(nn.Module):
|
||||
"""Matches the TokenChoiceTopKRouter inference path (no capacity/jitter/load stats).
|
||||
|
||||
The asymmetry must be preserved: selection uses the bias-added score, while gating
|
||||
weights gather the bias-free score.
|
||||
"""
|
||||
|
||||
def __init__(self, hidden_size, num_experts, top_k, score_func, norm_topk_prob,
|
||||
n_group, topk_group, route_scale, device=None, dtype=None, operations=None):
|
||||
super().__init__()
|
||||
self.num_experts = num_experts
|
||||
self.top_k = top_k
|
||||
self.score_func = score_func
|
||||
self.norm_topk_prob = norm_topk_prob
|
||||
self.n_group = n_group
|
||||
self.topk_group = topk_group
|
||||
self.route_scale = route_scale
|
||||
self.weight = nn.Parameter(torch.empty(num_experts, hidden_size, device=device, dtype=dtype))
|
||||
self.register_buffer("e_score_correction_bias", torch.zeros(num_experts, device=device, dtype=dtype), persistent=True)
|
||||
|
||||
def _group_limited_topk(self, scores_for_choice):
|
||||
seq_len = scores_for_choice.shape[0]
|
||||
experts_per_group = self.num_experts // self.n_group
|
||||
grouped = scores_for_choice.view(seq_len, self.n_group, experts_per_group)
|
||||
group_scores = grouped.topk(2, dim=-1)[0].sum(dim=-1)
|
||||
group_idx = torch.topk(group_scores, k=self.topk_group, dim=-1, sorted=False)[1]
|
||||
group_mask = torch.zeros_like(group_scores)
|
||||
group_mask.scatter_(1, group_idx, 1)
|
||||
score_mask = (
|
||||
group_mask.unsqueeze(-1)
|
||||
.expand(seq_len, self.n_group, experts_per_group)
|
||||
.reshape(seq_len, -1)
|
||||
)
|
||||
masked = scores_for_choice.masked_fill(~score_mask.bool(), float("-inf"))
|
||||
return torch.topk(masked, k=self.top_k, dim=-1, sorted=False)[1]
|
||||
|
||||
def forward(self, tokens: torch.Tensor):
|
||||
logits = F.linear(tokens, self.weight)
|
||||
if self.score_func == "softmax":
|
||||
scores = F.softmax(logits, dim=-1)
|
||||
else:
|
||||
scores = logits.sigmoid()
|
||||
scores_for_choice = scores + self.e_score_correction_bias.unsqueeze(0)
|
||||
if self.n_group is not None and self.n_group > 1:
|
||||
top_indices = self._group_limited_topk(scores_for_choice)
|
||||
else:
|
||||
top_indices = torch.topk(scores_for_choice, k=self.top_k, dim=-1, sorted=False)[1]
|
||||
top_scores = scores.gather(1, top_indices)
|
||||
if self.top_k > 1 and self.norm_topk_prob:
|
||||
top_scores = top_scores / (top_scores.sum(dim=-1, keepdim=True) + 1e-20)
|
||||
top_scores = top_scores * self.route_scale
|
||||
return top_indices, top_scores.to(tokens.dtype), logits, scores, scores_for_choice
|
||||
|
||||
|
||||
class LingBotVideoGroupedExperts(nn.Module):
|
||||
"""Weight layout matches GroupedExperts: w1 [E,I,H], w2 [E,H,I], w3 [E,I,H]. Eager per-expert compute."""
|
||||
|
||||
def __init__(self, num_experts, hidden_size, intermediate_size, device=None, dtype=None):
|
||||
super().__init__()
|
||||
self.num_experts = num_experts
|
||||
self.w1 = nn.Parameter(torch.empty(num_experts, intermediate_size, hidden_size, device=device, dtype=dtype))
|
||||
self.w2 = nn.Parameter(torch.empty(num_experts, hidden_size, intermediate_size, device=device, dtype=dtype))
|
||||
self.w3 = nn.Parameter(torch.empty(num_experts, intermediate_size, hidden_size, device=device, dtype=dtype))
|
||||
|
||||
|
||||
class LingBotVideoSparseMoeBlock(nn.Module):
|
||||
def __init__(self, hidden_size, intermediate_size, num_experts, top_k,
|
||||
moe_intermediate_size, score_func, norm_topk_prob, n_group, topk_group,
|
||||
routed_scaling_factor, n_shared_experts, device=None, dtype=None, operations=None):
|
||||
super().__init__()
|
||||
self.hidden_size = hidden_size
|
||||
self.num_experts = num_experts
|
||||
self.router = LingBotVideoRouter(
|
||||
hidden_size, num_experts, top_k, score_func, norm_topk_prob,
|
||||
n_group, topk_group, routed_scaling_factor, device=device, dtype=dtype, operations=operations,
|
||||
)
|
||||
self.experts = LingBotVideoGroupedExperts(num_experts, hidden_size, moe_intermediate_size, device=device, dtype=dtype)
|
||||
self.shared_experts = None
|
||||
if n_shared_experts is not None and n_shared_experts > 0:
|
||||
self.shared_experts = LingBotVideoMLP(
|
||||
hidden_size, moe_intermediate_size * n_shared_experts, device=device, dtype=dtype, operations=operations
|
||||
)
|
||||
|
||||
def _run_expert(self, expert_idx: int, tokens: torch.Tensor) -> torch.Tensor:
|
||||
h = F.silu(tokens @ self.experts.w1[expert_idx].transpose(-2, -1))
|
||||
h = h * (tokens @ self.experts.w3[expert_idx].transpose(-2, -1))
|
||||
return h @ self.experts.w2[expert_idx].transpose(-2, -1)
|
||||
|
||||
def _run_selected_experts(
|
||||
self,
|
||||
tokens: torch.Tensor,
|
||||
top_scores: torch.Tensor,
|
||||
top_indices: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
out = tokens.new_zeros(tokens.shape)
|
||||
for expert_idx in range(self.num_experts):
|
||||
selected = top_indices == expert_idx
|
||||
if not bool(selected.any()):
|
||||
continue
|
||||
token_indices, choice_indices = torch.where(selected)
|
||||
expert_tokens = tokens[token_indices]
|
||||
expert_output = self._run_expert(expert_idx, expert_tokens)
|
||||
expert_output = expert_output * top_scores[token_indices, choice_indices].unsqueeze(-1)
|
||||
out.index_add_(0, token_indices, expert_output)
|
||||
return out
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor, padding_mask: Optional[torch.Tensor] = None):
|
||||
# hidden_states: (B, S, H); padding_mask: (B*S,) with 1=valid (only needed when B>1)
|
||||
B = hidden_states.shape[0]
|
||||
tokens = hidden_states.view(-1, self.hidden_size)
|
||||
top_indices, top_scores, logits, scores, scores_for_choice = self.router(tokens)
|
||||
del logits, scores, scores_for_choice
|
||||
if padding_mask is not None:
|
||||
pm = padding_mask.unsqueeze(-1).to(top_scores.dtype)
|
||||
top_scores = top_scores * pm
|
||||
top_scores = top_scores / (top_scores.sum(dim=-1, keepdim=True) + 1e-9)
|
||||
top_scores = top_scores * self.router.route_scale
|
||||
|
||||
out = self._run_selected_experts(tokens, top_scores, top_indices)
|
||||
|
||||
out = out.view(B, -1, self.hidden_size)
|
||||
if self.shared_experts is not None:
|
||||
shared_output = self.shared_experts(hidden_states)
|
||||
out = out + shared_output
|
||||
return out
|
||||
|
||||
|
||||
class LingBotVideoBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size,
|
||||
num_attention_heads,
|
||||
intermediate_size,
|
||||
norm_eps,
|
||||
qkv_bias,
|
||||
out_bias,
|
||||
num_experts,
|
||||
num_experts_per_tok,
|
||||
moe_intermediate_size,
|
||||
decoder_sparse_step,
|
||||
mlp_only_layers,
|
||||
n_shared_experts,
|
||||
score_func,
|
||||
norm_topk_prob,
|
||||
n_group,
|
||||
topk_group,
|
||||
routed_scaling_factor,
|
||||
layer_idx: int,
|
||||
device=None,
|
||||
dtype=None,
|
||||
operations=None,
|
||||
):
|
||||
super().__init__()
|
||||
self.layer_idx = layer_idx
|
||||
h = hidden_size
|
||||
self.scale_shift_table = nn.Parameter(torch.empty(1, 6 * h, device=device, dtype=dtype))
|
||||
self.norm1 = operations.RMSNorm(h, eps=norm_eps, elementwise_affine=True, device=device, dtype=dtype)
|
||||
self.attn = LingBotVideoAttention(
|
||||
h, num_attention_heads, norm_eps, qkv_bias, out_bias, device=device, dtype=dtype, operations=operations
|
||||
)
|
||||
self.norm_post_attn = operations.RMSNorm(h, eps=norm_eps, elementwise_affine=True, device=device, dtype=dtype)
|
||||
self.norm2 = operations.RMSNorm(h, eps=norm_eps, elementwise_affine=True, device=device, dtype=dtype)
|
||||
# Sparsity decision matches MoEBlock: mlp_only_layers + decoder_sparse_step + num_experts
|
||||
if layer_idx not in mlp_only_layers and (
|
||||
num_experts > 0 and (layer_idx + 1) % decoder_sparse_step == 0
|
||||
):
|
||||
self.ffn = LingBotVideoSparseMoeBlock(
|
||||
h, intermediate_size, num_experts, num_experts_per_tok,
|
||||
moe_intermediate_size, score_func, norm_topk_prob,
|
||||
n_group, topk_group, routed_scaling_factor,
|
||||
n_shared_experts, device=device, dtype=dtype, operations=operations,
|
||||
)
|
||||
else:
|
||||
self.ffn = LingBotVideoMLP(h, intermediate_size, device=device, dtype=dtype, operations=operations)
|
||||
self.norm_post_ffn = operations.RMSNorm(h, eps=norm_eps, elementwise_affine=True, device=device, dtype=dtype)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x,
|
||||
temb6,
|
||||
rotary_emb,
|
||||
attention_mask=None,
|
||||
moe_padding_mask=None,
|
||||
transformer_options={},
|
||||
):
|
||||
expected_tokens = x.shape[0] * x.shape[1]
|
||||
if temb6.ndim != 2 or temb6.shape[0] != expected_tokens:
|
||||
raise ValueError(
|
||||
"LingBotVideoBlock expects token-level temb6 with shape "
|
||||
f"(B*S, 6D); got {tuple(temb6.shape)} for hidden states {tuple(x.shape)}."
|
||||
)
|
||||
mod = temb6.view(x.shape[0], x.shape[1], -1) + self.scale_shift_table.unsqueeze(0)
|
||||
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = mod.chunk(6, dim=-1)
|
||||
gate_msa, gate_mlp = gate_msa.tanh(), gate_mlp.tanh()
|
||||
scale_msa, scale_mlp = 1.0 + scale_msa, 1.0 + scale_mlp
|
||||
|
||||
attn_in = self.norm1(x) * scale_msa + shift_msa
|
||||
attn_out = self.attn(
|
||||
attn_in,
|
||||
rotary_emb,
|
||||
attention_mask,
|
||||
transformer_options=transformer_options,
|
||||
)
|
||||
x = x + (gate_msa * self.norm_post_attn(attn_out)).to(x.dtype)
|
||||
|
||||
ffn_in = self.norm2(x) * scale_mlp + shift_mlp
|
||||
if isinstance(self.ffn, LingBotVideoSparseMoeBlock):
|
||||
ffn_out = self.ffn(ffn_in, padding_mask=moe_padding_mask)
|
||||
else:
|
||||
ffn_out = self.ffn(ffn_in)
|
||||
ffn_normed = self.norm_post_ffn(ffn_out)
|
||||
x = x + (gate_mlp * ffn_normed).to(x.dtype)
|
||||
return x
|
||||
|
||||
|
||||
class LingBotVideo(nn.Module):
|
||||
_no_split_modules = ["LingBotVideoBlock"]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
image_model=None,
|
||||
patch_size: Tuple[int, int, int] = (1, 2, 2),
|
||||
in_channels: int = 16,
|
||||
out_channels: int = 16,
|
||||
hidden_size: int = 2048,
|
||||
num_attention_heads: int = 16,
|
||||
depth: int = 24,
|
||||
intermediate_size: int = 6144,
|
||||
text_dim: int = 2560,
|
||||
freq_dim: int = 256,
|
||||
norm_eps: float = 1e-6,
|
||||
rope_theta: float = 256.0,
|
||||
axes_dims: Tuple[int, int, int] = (32, 48, 48),
|
||||
axes_lens: Tuple[int, int, int] = (8192, 1024, 1024),
|
||||
qkv_bias: bool = False,
|
||||
out_bias: bool = True,
|
||||
patch_embed_bias: bool = True,
|
||||
timestep_mlp_bias: bool = True,
|
||||
num_experts: int = 0,
|
||||
num_experts_per_tok: int = 8,
|
||||
moe_intermediate_size: int = 512,
|
||||
decoder_sparse_step: int = 1,
|
||||
mlp_only_layers: Tuple[int, ...] = (),
|
||||
n_shared_experts: Optional[int] = None,
|
||||
score_func: str = "sigmoid",
|
||||
norm_topk_prob: bool = True,
|
||||
n_group: Optional[int] = None,
|
||||
topk_group: Optional[int] = None,
|
||||
routed_scaling_factor: float = 1.0,
|
||||
dtype=None,
|
||||
device=None,
|
||||
operations=None,
|
||||
):
|
||||
super().__init__()
|
||||
self.dtype = dtype
|
||||
self.patch_size = tuple(patch_size)
|
||||
self.out_channels = out_channels
|
||||
head_dim = hidden_size // num_attention_heads
|
||||
assert head_dim == sum(axes_dims), f"head_dim {head_dim} != sum(axes_dims) {sum(axes_dims)}"
|
||||
mlp_only_layers = tuple(mlp_only_layers)
|
||||
|
||||
self.patch_embedder = operations.Linear(
|
||||
in_channels * math.prod(patch_size), hidden_size, bias=patch_embed_bias, device=device, dtype=dtype
|
||||
)
|
||||
self.freq_dim = freq_dim
|
||||
self.time_embedder = LingBotVideoTimestepEmbedding(
|
||||
freq_dim, hidden_size, bias=timestep_mlp_bias, device=device, dtype=dtype, operations=operations
|
||||
)
|
||||
self.time_modulation = nn.Sequential(
|
||||
nn.SiLU(),
|
||||
operations.Linear(hidden_size, 6 * hidden_size, device=device, dtype=dtype),
|
||||
)
|
||||
self.text_embedder = LingBotVideoTextEmbedder(text_dim, hidden_size, device=device, dtype=dtype, operations=operations)
|
||||
self.rope = LingBotVideoRotaryEmbedding(axes_dims, axes_lens, rope_theta)
|
||||
self.blocks = nn.ModuleList(
|
||||
[
|
||||
LingBotVideoBlock(
|
||||
hidden_size=hidden_size,
|
||||
num_attention_heads=num_attention_heads,
|
||||
intermediate_size=intermediate_size,
|
||||
norm_eps=norm_eps,
|
||||
qkv_bias=qkv_bias,
|
||||
out_bias=out_bias,
|
||||
num_experts=num_experts,
|
||||
num_experts_per_tok=num_experts_per_tok,
|
||||
moe_intermediate_size=moe_intermediate_size,
|
||||
decoder_sparse_step=decoder_sparse_step,
|
||||
mlp_only_layers=mlp_only_layers,
|
||||
n_shared_experts=n_shared_experts,
|
||||
score_func=score_func,
|
||||
norm_topk_prob=norm_topk_prob,
|
||||
n_group=n_group,
|
||||
topk_group=topk_group,
|
||||
routed_scaling_factor=routed_scaling_factor,
|
||||
layer_idx=i,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
operations=operations,
|
||||
)
|
||||
for i in range(depth)
|
||||
]
|
||||
)
|
||||
self.norm_out = operations.LayerNorm(hidden_size, elementwise_affine=False, eps=norm_eps, device=device, dtype=dtype)
|
||||
self.norm_out_modulation = nn.Sequential(
|
||||
nn.SiLU(),
|
||||
operations.Linear(hidden_size, 2 * hidden_size, device=device, dtype=dtype),
|
||||
)
|
||||
self.proj_out = operations.Linear(hidden_size, math.prod(patch_size) * out_channels, device=device, dtype=dtype)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor, # (B, C, T, H, W)
|
||||
timestep: torch.Tensor, # (B,) ∈ [0, 1000](= sigma*1000)
|
||||
context: torch.Tensor = None, # (B, L, text_dim)
|
||||
encoder_attention_mask: Optional[torch.Tensor] = None, # (B, L) 1=valid
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
transformer_options={},
|
||||
**kwargs,
|
||||
):
|
||||
encoder_hidden_states = context
|
||||
if encoder_hidden_states is None:
|
||||
raise ValueError("LingBotVideo requires text conditioning.")
|
||||
if encoder_attention_mask is None:
|
||||
encoder_attention_mask = attention_mask
|
||||
B, C, T, H, W = hidden_states.shape
|
||||
pF, pH, pW = self.patch_size
|
||||
gt, gh, gw = T // pF, H // pH, W // pW
|
||||
n_video = gt * gh * gw
|
||||
L = encoder_hidden_states.shape[1]
|
||||
device = hidden_states.device
|
||||
if encoder_attention_mask is not None:
|
||||
text_lens = encoder_attention_mask.sum(dim=-1).long()
|
||||
else:
|
||||
text_lens = torch.full((B,), L, dtype=torch.long, device=device)
|
||||
text_lens_list = [int(v) for v in text_lens.detach().cpu().tolist()]
|
||||
|
||||
# patchify: token order (f h w), feature order (pf ph pw c) -- matches patchify_and_embed
|
||||
patch_tokens = hidden_states.reshape(B, C, gt, pF, gh, pH, gw, pW)
|
||||
patch_tokens = patch_tokens.permute(0, 2, 4, 6, 3, 5, 7, 1).reshape(
|
||||
B,
|
||||
n_video,
|
||||
pF * pH * pW * C,
|
||||
)
|
||||
x = self.patch_embedder(patch_tokens)
|
||||
text = self.text_embedder(encoder_hidden_states)
|
||||
joint = torch.cat([x, text], dim=1) # [video; text]
|
||||
joint_seq_len = joint.shape[1]
|
||||
|
||||
# Per-sample RoPE: video t-axis start = real text length of this sample + 1
|
||||
rotary_parts = [
|
||||
self.rope(make_joint_position_ids(text_lens_list[i], gt, gh, gw, device, L))
|
||||
for i in range(B)
|
||||
]
|
||||
rotary = torch.stack(rotary_parts, dim=0).unsqueeze(2) # (B, S, 1, head_dim/2, 2, 2)
|
||||
|
||||
attention_mask = None
|
||||
moe_padding_mask = None
|
||||
has_padding = encoder_attention_mask is not None and bool((text_lens < L).any())
|
||||
if has_padding:
|
||||
key_mask = torch.cat(
|
||||
[torch.ones(B, n_video, dtype=torch.bool, device=device),
|
||||
encoder_attention_mask.bool()],
|
||||
dim=1,
|
||||
)
|
||||
attention_mask = key_mask[:, None, None, :] # (B,1,1,S) → SDPA broadcast
|
||||
moe_padding_mask = key_mask.reshape(-1) # (B*S,)
|
||||
|
||||
timestep_proj = timestep_embedding(timestep.to(hidden_states.dtype), self.freq_dim, time_factor=1.0)
|
||||
t_emb = self.time_embedder(timestep_proj) # (B, D)
|
||||
temb_input = t_emb.unsqueeze(1).expand(B, joint_seq_len, -1) # (B, S, D)
|
||||
temb6 = self.time_modulation(temb_input.reshape(B * joint_seq_len, -1))
|
||||
temb6 = temb6.reshape(B, joint_seq_len, -1) # (B, S, 6D)
|
||||
|
||||
temb6 = temb6.reshape(temb6.shape[0] * temb6.shape[1], -1)
|
||||
|
||||
for block in self.blocks:
|
||||
joint = block(
|
||||
joint,
|
||||
temb6,
|
||||
rotary,
|
||||
attention_mask,
|
||||
moe_padding_mask,
|
||||
transformer_options=transformer_options,
|
||||
)
|
||||
|
||||
final_mod = self.norm_out_modulation(temb_input.reshape(joint.shape[0] * joint.shape[1], -1))
|
||||
shift, scale = final_mod.reshape(joint.shape[0], joint.shape[1], -1).chunk(2, dim=-1)
|
||||
final_hidden = self.norm_out(joint) * (1.0 + scale) + shift
|
||||
projected = self.proj_out(final_hidden)
|
||||
x = projected[:, :n_video]
|
||||
|
||||
# unpatchify (matches the rearrange in postprocess)
|
||||
Cout = self.out_channels
|
||||
x = x.reshape(B, gt, gh, gw, pF, pH, pW, Cout)
|
||||
x = x.permute(0, 7, 1, 4, 2, 5, 3, 6).reshape(B, Cout, T, H, W)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
LingBotVideoTransformer3DModel = LingBotVideo
|
||||
@@ -63,6 +63,7 @@ import comfy.ldm.kandinsky5.model
|
||||
import comfy.ldm.anima.model
|
||||
import comfy.ldm.ace.ace_step15
|
||||
import comfy.ldm.cogvideo.model
|
||||
import comfy.ldm.lingbot_video.model
|
||||
import comfy.ldm.rt_detr.rtdetr_v4
|
||||
import comfy.ldm.ernie.model
|
||||
import comfy.ldm.sam3.detector
|
||||
@@ -1376,6 +1377,20 @@ class HunyuanVideoSkyreelsI2V(HunyuanVideo):
|
||||
def scale_latent_inpaint(self, latent_image, **kwargs):
|
||||
return super().scale_latent_inpaint(latent_image=latent_image, **kwargs)
|
||||
|
||||
class LingBotVideo(BaseModel):
|
||||
def __init__(self, model_config, model_type=ModelType.FLOW, device=None):
|
||||
super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.lingbot_video.model.LingBotVideo)
|
||||
|
||||
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)
|
||||
return out
|
||||
|
||||
def scale_latent_inpaint(self, latent_image, **kwargs):
|
||||
return latent_image
|
||||
|
||||
class CosmosVideo(BaseModel):
|
||||
def __init__(self, model_config, model_type=ModelType.EDM, image_to_video=False, device=None):
|
||||
super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.cosmos.model.GeneralDIT)
|
||||
|
||||
@@ -232,6 +232,54 @@ def detect_unet_config(state_dict, key_prefix, metadata=None):
|
||||
dit_config["meanflow_sum"] = False
|
||||
return dit_config
|
||||
|
||||
if '{}patch_embedder.weight'.format(key_prefix) in state_dict_keys and '{}text_embedder.norm.weight'.format(key_prefix) in state_dict_keys and '{}blocks.0.attn.to_q.weight'.format(key_prefix) in state_dict_keys: # LingBot Video
|
||||
dit_config = {}
|
||||
patch_size = (1, 2, 2)
|
||||
patch_prod = math.prod(patch_size)
|
||||
patch_embed = state_dict['{}patch_embedder.weight'.format(key_prefix)]
|
||||
proj_out = state_dict['{}proj_out.weight'.format(key_prefix)]
|
||||
hidden_size = patch_embed.shape[0]
|
||||
dit_config["image_model"] = "lingbot_video"
|
||||
dit_config["patch_size"] = patch_size
|
||||
dit_config["in_channels"] = 16
|
||||
dit_config["out_channels"] = proj_out.shape[0] // patch_prod
|
||||
dit_config["hidden_size"] = hidden_size
|
||||
dit_config["num_attention_heads"] = hidden_size // 128
|
||||
dit_config["depth"] = count_blocks(state_dict_keys, '{}blocks.'.format(key_prefix) + '{}.')
|
||||
dit_config["text_dim"] = state_dict['{}text_embedder.norm.weight'.format(key_prefix)].shape[0]
|
||||
dit_config["freq_dim"] = 256
|
||||
dit_config["qkv_bias"] = '{}blocks.0.attn.to_q.bias'.format(key_prefix) in state_dict_keys
|
||||
dit_config["out_bias"] = '{}blocks.0.attn.to_out.bias'.format(key_prefix) in state_dict_keys
|
||||
dit_config["patch_embed_bias"] = '{}patch_embedder.bias'.format(key_prefix) in state_dict_keys
|
||||
dit_config["timestep_mlp_bias"] = '{}time_embedder.linear_1.bias'.format(key_prefix) in state_dict_keys
|
||||
|
||||
moe_key = '{}blocks.0.ffn.experts.w1'.format(key_prefix)
|
||||
if moe_key in state_dict_keys:
|
||||
experts = state_dict[moe_key]
|
||||
dit_config["num_experts"] = experts.shape[0]
|
||||
dit_config["moe_intermediate_size"] = experts.shape[1]
|
||||
if experts.shape[0] == 128:
|
||||
dit_config["n_group"] = 4
|
||||
dit_config["topk_group"] = 2
|
||||
dit_config["routed_scaling_factor"] = 2.5
|
||||
mlp_only_layers = []
|
||||
for i in range(dit_config["depth"]):
|
||||
if '{}blocks.{}.ffn.gate_proj.weight'.format(key_prefix, i) in state_dict_keys:
|
||||
mlp_only_layers.append(i)
|
||||
dit_config["mlp_only_layers"] = tuple(mlp_only_layers)
|
||||
if len(mlp_only_layers) > 0:
|
||||
dit_config["intermediate_size"] = state_dict['{}blocks.{}.ffn.gate_proj.weight'.format(key_prefix, mlp_only_layers[0])].shape[0]
|
||||
if '{}blocks.0.ffn.shared_experts.gate_proj.weight'.format(key_prefix) in state_dict_keys:
|
||||
shared = state_dict['{}blocks.0.ffn.shared_experts.gate_proj.weight'.format(key_prefix)]
|
||||
dit_config["n_shared_experts"] = shared.shape[0] // experts.shape[1]
|
||||
else:
|
||||
dit_config["num_experts"] = 0
|
||||
dit_config["intermediate_size"] = state_dict['{}blocks.0.ffn.gate_proj.weight'.format(key_prefix)].shape[0]
|
||||
|
||||
if metadata is not None and "config" in metadata:
|
||||
dit_config.update(json.loads(metadata["config"]).get("transformer", {}))
|
||||
return dit_config
|
||||
|
||||
if any_suffix_in(state_dict_keys, key_prefix, 'double_blocks.0.img_attn.norm.key_norm.', ["weight", "scale"]) and ('{}img_in.weight'.format(key_prefix) in state_dict_keys or any_suffix_in(state_dict_keys, key_prefix, 'distilled_guidance_layer.norms.0.', ["weight", "scale"])): #Flux, Chroma or Chroma Radiance (has no img_in.weight)
|
||||
dit_config = {}
|
||||
if '{}double_stream_modulation_img.lin.weight'.format(key_prefix) in state_dict_keys:
|
||||
|
||||
@@ -69,6 +69,7 @@ import comfy.text_encoders.ace15
|
||||
import comfy.text_encoders.longcat_image
|
||||
import comfy.text_encoders.qwen35
|
||||
import comfy.text_encoders.qwen3vl
|
||||
import comfy.text_encoders.lingbot_video
|
||||
import comfy.text_encoders.boogu
|
||||
import comfy.text_encoders.ernie
|
||||
import comfy.text_encoders.gemma4
|
||||
@@ -1313,6 +1314,7 @@ class CLIPType(Enum):
|
||||
IDEOGRAM4 = 30
|
||||
BOOGU = 31
|
||||
KREA2 = 32
|
||||
LINGBOT_VIDEO = 33
|
||||
|
||||
|
||||
|
||||
@@ -1646,6 +1648,10 @@ def load_text_encoder_state_dicts(state_dicts=[], embedding_directory=None, clip
|
||||
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)
|
||||
clip_target.tokenizer = comfy.text_encoders.flux.KleinTokenizer8B if te_model == TEModel.QWEN3VL_8B else comfy.text_encoders.flux.KleinTokenizer
|
||||
elif clip_type == CLIPType.LINGBOT_VIDEO and te_model == TEModel.QWEN3VL_4B:
|
||||
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.lingbot_video.te(**llama_detect(clip_data), model_type="qwen3vl_4b")
|
||||
clip_target.tokenizer = comfy.text_encoders.lingbot_video.tokenizer(model_type="qwen3vl_4b")
|
||||
else:
|
||||
clip_data[0] = comfy.utils.state_dict_prefix_replace(clip_data[0], {"model.language_model.": "model.", "model.visual.": "visual.", "lm_head.": "model.lm_head."})
|
||||
qwen3vl_type = {TEModel.QWEN3VL_4B: "qwen3vl_4b", TEModel.QWEN3VL_8B: "qwen3vl_8b"}[te_model]
|
||||
|
||||
@@ -27,6 +27,8 @@ 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.qwen3vl
|
||||
import comfy.text_encoders.lingbot_video
|
||||
import comfy.text_encoders.anima
|
||||
import comfy.text_encoders.ace15
|
||||
import comfy.text_encoders.longcat_image
|
||||
@@ -1023,6 +1025,40 @@ class HunyuanVideoSkyreelsI2V(HunyuanVideo):
|
||||
out = model_base.HunyuanVideoSkyreelsI2V(self, device=device)
|
||||
return out
|
||||
|
||||
class LingBotVideo(supported_models_base.BASE):
|
||||
unet_config = {
|
||||
"image_model": "lingbot_video",
|
||||
}
|
||||
|
||||
sampling_settings = {
|
||||
"shift": 3.0,
|
||||
}
|
||||
|
||||
unet_extra_config = {}
|
||||
latent_format = latent_formats.LingBotVideo
|
||||
|
||||
memory_usage_factor = 1.8
|
||||
|
||||
supported_inference_dtypes = [torch.bfloat16, torch.float32]
|
||||
|
||||
vae_key_prefix = ["vae."]
|
||||
text_encoder_key_prefix = ["text_encoders."]
|
||||
|
||||
def __init__(self, unet_config):
|
||||
super().__init__(unet_config)
|
||||
self.memory_usage_factor = self.memory_usage_factor * (unet_config.get("hidden_size", 2048) / 2048)
|
||||
|
||||
def get_model(self, state_dict, prefix="", device=None):
|
||||
return model_base.LingBotVideo(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_4b.transformer.".format(pref))
|
||||
if len(qwen3vl_detect) > 0:
|
||||
qwen3vl_detect["model_type"] = "qwen3vl_4b"
|
||||
return supported_models_base.ClipTarget(comfy.text_encoders.lingbot_video.tokenizer(model_type="qwen3vl_4b"), comfy.text_encoders.lingbot_video.te(**qwen3vl_detect))
|
||||
return None
|
||||
|
||||
class CosmosT2V(supported_models_base.BASE):
|
||||
unet_config = {
|
||||
"image_model": "cosmos",
|
||||
@@ -2317,6 +2353,7 @@ models = [
|
||||
HunyuanVideoSkyreelsI2V,
|
||||
HunyuanVideoI2V,
|
||||
HunyuanVideo,
|
||||
LingBotVideo,
|
||||
CosmosT2V,
|
||||
CosmosI2V,
|
||||
CosmosT2IPredict2,
|
||||
|
||||
@@ -0,0 +1,149 @@
|
||||
import comfy.text_encoders.qwen3vl
|
||||
from comfy import sd1_clip
|
||||
|
||||
|
||||
PAD_TOKEN = 151643
|
||||
CROP_MARKER = {"type": "lingbot_video_crop_start"}
|
||||
|
||||
PROMPT_TEMPLATE = (
|
||||
"<|im_start|>system\nGiven a user input that may include a text prompt alone, "
|
||||
"a text prompt with an image reference, or a text prompt with a video reference "
|
||||
"or a video reference alone, generate an \"Enhanced prompt\" that provides detailed "
|
||||
"visual descriptions suitable for video generation. Evaluate the level of detail "
|
||||
"in the user's input: if it is simple, enrich it by adding specifics about colors, "
|
||||
"shapes, sizes, textures, lighting, motion dynamics, camera movement, temporal "
|
||||
"progression, and spatial relationships to create vivid, concrete, and temporally "
|
||||
"coherent scenes to create vivid and concrete scenes. Please generate only the "
|
||||
"enhanced description for the prompt below and avoid including any additional "
|
||||
"commentary or evaluations:<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n"
|
||||
"<|im_start|>assistant\n"
|
||||
)
|
||||
IMAGE_PROMPT_TEMPLATE = "<|vision_start|><|image_pad|><|vision_end|>"
|
||||
|
||||
|
||||
def _marker_tuple(example=None):
|
||||
if example is not None and len(example) > 2:
|
||||
return (CROP_MARKER, 1.0, None)
|
||||
return (CROP_MARKER, 1.0)
|
||||
|
||||
|
||||
class LingBotVideoTokenizer(comfy.text_encoders.qwen3vl.Qwen3VLTokenizer):
|
||||
def __init__(self, embedding_directory=None, tokenizer_data={}, model_type="qwen3vl_4b"):
|
||||
super().__init__(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data, model_type=model_type)
|
||||
self._lingbot_crop_start = None
|
||||
|
||||
def lingbot_crop_start(self):
|
||||
if self._lingbot_crop_start is not None:
|
||||
return self._lingbot_crop_start
|
||||
prefix = PROMPT_TEMPLATE.split("{}")[0]
|
||||
tokens = super().tokenize_with_weights(prefix, thinking=True)
|
||||
key = next(iter(tokens))
|
||||
count = 0
|
||||
for token in tokens[key][0]:
|
||||
token_id = token[0]
|
||||
if isinstance(token_id, int) and token_id == PAD_TOKEN:
|
||||
continue
|
||||
count += 1
|
||||
self._lingbot_crop_start = count
|
||||
return count
|
||||
|
||||
def tokenize_with_weights(self, text, return_word_ids=False, images=[], prevent_empty_text=False, thinking=True, **kwargs):
|
||||
image = kwargs.get("image", None)
|
||||
if image is not None and len(images) == 0:
|
||||
images = [image[i:i + 1] for i in range(image.shape[0])]
|
||||
|
||||
prompt_text = text
|
||||
if len(images) > 0 and not prompt_text.startswith("<|vision_start|>"):
|
||||
prompt_text = IMAGE_PROMPT_TEMPLATE + prompt_text
|
||||
|
||||
tokens = super().tokenize_with_weights(
|
||||
prompt_text,
|
||||
return_word_ids=return_word_ids,
|
||||
llama_template=PROMPT_TEMPLATE,
|
||||
images=images,
|
||||
prevent_empty_text=prevent_empty_text,
|
||||
thinking=thinking,
|
||||
**kwargs,
|
||||
)
|
||||
crop_start = self.lingbot_crop_start()
|
||||
for key in tokens:
|
||||
for row in tokens[key]:
|
||||
example = row[0] if len(row) > 0 else None
|
||||
row.insert(crop_start, _marker_tuple(example))
|
||||
return tokens
|
||||
|
||||
|
||||
class LingBotVideoClipModel(comfy.text_encoders.qwen3vl.Qwen3VLClipModel):
|
||||
def __init__(self, device="cpu", dtype=None, attention_mask=True, model_options={}, model_type="qwen3vl_4b"):
|
||||
super().__init__(
|
||||
device=device,
|
||||
layer="last",
|
||||
layer_idx=None,
|
||||
dtype=dtype,
|
||||
attention_mask=attention_mask,
|
||||
model_options=model_options,
|
||||
model_type=model_type,
|
||||
)
|
||||
self.return_attention_masks = False
|
||||
|
||||
@staticmethod
|
||||
def _strip_crop_markers(tokens):
|
||||
clean_tokens = []
|
||||
crop_starts = []
|
||||
for row in tokens:
|
||||
clean_row = []
|
||||
crop_start = None
|
||||
for token in row:
|
||||
token_value = token[0] if isinstance(token, tuple) else token
|
||||
if isinstance(token_value, dict) and token_value.get("type") == CROP_MARKER["type"]:
|
||||
crop_start = len(clean_row)
|
||||
continue
|
||||
clean_row.append(token)
|
||||
clean_tokens.append(clean_row)
|
||||
crop_starts.append(crop_start)
|
||||
return clean_tokens, crop_starts
|
||||
|
||||
def forward(self, tokens):
|
||||
clean_tokens, crop_starts = self._strip_crop_markers(tokens)
|
||||
out = super().forward(clean_tokens)
|
||||
crop_starts = [c for c in crop_starts if c is not None]
|
||||
if len(crop_starts) == 0:
|
||||
return out
|
||||
|
||||
crop_start = min(crop_starts)
|
||||
z, pooled_output = out[:2]
|
||||
z = z[:, crop_start:]
|
||||
|
||||
return z, pooled_output
|
||||
|
||||
|
||||
class LingBotVideoTEModel(comfy.text_encoders.qwen3vl.Qwen3VLTEModel):
|
||||
def __init__(self, device="cpu", dtype=None, model_options={}, model_type="qwen3vl_4b"):
|
||||
clip_model = lambda **kw: LingBotVideoClipModel(**kw, model_type=model_type)
|
||||
sd1_clip.SD1ClipModel.__init__(
|
||||
self,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
name=model_type,
|
||||
clip_model=clip_model,
|
||||
model_options=model_options,
|
||||
)
|
||||
|
||||
|
||||
def tokenizer(model_type="qwen3vl_4b"):
|
||||
class LingBotVideoTokenizer_(LingBotVideoTokenizer):
|
||||
def __init__(self, embedding_directory=None, tokenizer_data={}):
|
||||
super().__init__(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data, model_type=model_type)
|
||||
return LingBotVideoTokenizer_
|
||||
|
||||
|
||||
def te(dtype_llama=None, llama_quantization_metadata=None, model_type="qwen3vl_4b"):
|
||||
class LingBotVideoTEModel_(LingBotVideoTEModel):
|
||||
def __init__(self, device="cpu", dtype=None, model_options={}):
|
||||
if dtype_llama is not None:
|
||||
dtype = dtype_llama
|
||||
if llama_quantization_metadata is not None:
|
||||
model_options = model_options.copy()
|
||||
model_options["quantization_metadata"] = llama_quantization_metadata
|
||||
super().__init__(device=device, dtype=dtype, model_options=model_options, model_type=model_type)
|
||||
return LingBotVideoTEModel_
|
||||
@@ -100,7 +100,6 @@ def _parse_cli_feature_flags() -> dict[str, Any]:
|
||||
# Default server capabilities
|
||||
_CORE_FEATURE_FLAGS: dict[str, Any] = {
|
||||
"supports_preview_metadata": True,
|
||||
"supports_model_type_tags": True,
|
||||
"max_upload_size": args.max_upload_size * 1024 * 1024, # Convert MB to bytes
|
||||
"extension": {"manager": {"supports_v4": True}},
|
||||
"node_replacements": True,
|
||||
|
||||
@@ -1261,155 +1261,6 @@ class DynamicSlot(ComfyTypeI):
|
||||
out_dict[input_type][finalized_id] = value
|
||||
out_dict["dynamic_paths"][finalized_id] = finalize_prefix(curr_prefix, curr_prefix[-1])
|
||||
|
||||
@comfytype(io_type="COMFY_DYNAMICGROUP_V3")
|
||||
class DynamicGroup(ComfyTypeI):
|
||||
"""A repeatable group of widget inputs (e.g. lora_name + strength stacked into N rows).
|
||||
|
||||
At execution time the node receives a ``list[dict]`` where each element is a row.
|
||||
|
||||
Example::
|
||||
|
||||
io.DynamicGroup.Input(
|
||||
"loras",
|
||||
template=[
|
||||
io.Combo.Input("lora_name", options=folder_paths.get_filename_list("loras")),
|
||||
io.Float.Input("strength", default=1.0, min=-100, max=100, step=0.01),
|
||||
],
|
||||
min=0,
|
||||
max=50,
|
||||
)
|
||||
# execute receives: loras: list[dict] = [{"lora_name": "x.safetensors", "strength": 1.0}, ...]
|
||||
"""
|
||||
|
||||
Type = list[dict[str, Any]]
|
||||
_MaxRows = 100
|
||||
|
||||
class Input(DynamicInput):
|
||||
def __init__(
|
||||
self,
|
||||
id: str,
|
||||
template: list["Input"],
|
||||
min: int = 0,
|
||||
max: int = 50,
|
||||
display_name: str = None,
|
||||
optional: bool = False,
|
||||
tooltip: str = None,
|
||||
lazy: bool = None,
|
||||
extra_dict=None,
|
||||
group_name: str = "Group",
|
||||
):
|
||||
super().__init__(id, display_name, optional, tooltip, lazy, extra_dict)
|
||||
assert len(template) > 0, "DynamicGroup template must have at least one field."
|
||||
for t in template:
|
||||
assert isinstance(t, WidgetInput), (
|
||||
f"DynamicGroup template field '{t.id}' must be a WidgetInput subclass "
|
||||
f"(Combo, Float, Int, String, Boolean, Color). Got {type(t).__name__}."
|
||||
)
|
||||
assert not isinstance(t, DynamicInput), (
|
||||
f"DynamicGroup template field '{t.id}' must not be a DynamicInput. "
|
||||
"Nesting dynamic inputs inside DynamicGroup is not supported."
|
||||
)
|
||||
field_ids = [t.id for t in template]
|
||||
assert len(field_ids) == len(set(field_ids)), (
|
||||
f"DynamicGroup template field ids must be unique within a row. Got: {field_ids}"
|
||||
)
|
||||
# Reject "." in group id and template field ids: slot_id encoding uses "." as a
|
||||
# delimiter (<group_id>.<row>.<field_id>), so any "." in these names would cause
|
||||
# path.split(".") to produce the wrong number of segments during decoding.
|
||||
assert "." not in id, (
|
||||
f"DynamicGroup id must not contain '.'. Got: '{id}'"
|
||||
)
|
||||
for t in template:
|
||||
assert "." not in t.id, (
|
||||
f"DynamicGroup template field id must not contain '.'. Got: '{t.id}'"
|
||||
)
|
||||
assert min >= 0, "DynamicGroup min must be >= 0."
|
||||
assert max >= 1, "DynamicGroup max must be >= 1."
|
||||
assert max <= DynamicGroup._MaxRows, f"DynamicGroup max must be <= {DynamicGroup._MaxRows}."
|
||||
assert min <= max, "DynamicGroup min must be <= max."
|
||||
self.template = template
|
||||
self.min = min
|
||||
self.max = max
|
||||
self.group_name = group_name
|
||||
|
||||
def get_all(self) -> list["Input"]:
|
||||
return [self] + list(self.template)
|
||||
|
||||
def as_dict(self):
|
||||
return super().as_dict() | prune_dict({
|
||||
"template": create_input_dict_v1(self.template),
|
||||
"min": self.min,
|
||||
"max": self.max,
|
||||
"group_name": self.group_name,
|
||||
})
|
||||
|
||||
def validate(self):
|
||||
for t in self.template:
|
||||
t.validate()
|
||||
|
||||
@staticmethod
|
||||
def _expand_schema_for_dynamic(
|
||||
out_dict: dict[str, Any],
|
||||
live_inputs: dict[str, Any],
|
||||
value: tuple[str, dict[str, Any]],
|
||||
input_type: str,
|
||||
curr_prefix: list[str] | None,
|
||||
):
|
||||
info = value[1]
|
||||
min_rows: int = info.get("min", 0)
|
||||
max_rows: int = info.get("max", DynamicGroup._MaxRows)
|
||||
template: dict[str, Any] = info.get("template", {})
|
||||
|
||||
# Collect all template field specs across required/optional sections
|
||||
field_specs: list[tuple[str, tuple[str, dict[str, Any]], bool]] = []
|
||||
for field_required_key in ("required", "optional"):
|
||||
section = template.get(field_required_key, {})
|
||||
is_required_field = field_required_key == "required"
|
||||
for field_id, field_value in section.items():
|
||||
field_specs.append((field_id, field_value, is_required_field))
|
||||
|
||||
# Determine how many rows are currently present by scanning live_inputs
|
||||
finalized_prefix = finalize_prefix(curr_prefix)
|
||||
present_rows = 0
|
||||
for live_key in live_inputs:
|
||||
# Keys look like "<prefix>.<row>.<field_id>"
|
||||
if live_key.startswith(finalized_prefix + "."):
|
||||
remainder = live_key[len(finalized_prefix) + 1:]
|
||||
parts = remainder.split(".", 1)
|
||||
if len(parts) >= 1:
|
||||
try:
|
||||
row_idx = int(parts[0])
|
||||
present_rows = max(present_rows, row_idx + 1)
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
if present_rows > max_rows:
|
||||
raise ValueError(
|
||||
f"DynamicGroup input '{finalized_prefix}' received {present_rows} rows but max is {max_rows}."
|
||||
)
|
||||
row_count = max(min_rows, present_rows)
|
||||
|
||||
for row in range(row_count):
|
||||
for field_id, field_value, is_required_field in field_specs:
|
||||
slot_id = f"{finalized_prefix}.{row}.{field_id}"
|
||||
if row < min_rows and is_required_field:
|
||||
out_dict["required"][slot_id] = field_value
|
||||
else:
|
||||
out_dict["optional"][slot_id] = field_value
|
||||
# Register into dynamic_paths so build_nested_inputs places value at the right path
|
||||
out_dict["dynamic_paths"][slot_id] = slot_id
|
||||
|
||||
# Track the list root path so build_nested_inputs can convert the index dict to a list
|
||||
out_dict.setdefault("list_paths", set()).add(finalized_prefix)
|
||||
|
||||
# Handle the empty case (0 rows) – emit an empty-list default for the parent.
|
||||
# This must only fire when there are genuinely no rows; otherwise the parent
|
||||
# path would clobber the per-row dict built from the slot ids above.
|
||||
if row_count == 0:
|
||||
out_dict["dynamic_paths"][finalized_prefix] = finalized_prefix
|
||||
out_dict["dynamic_paths_default_value"][finalized_prefix] = DynamicPathsDefaultValue.EMPTY_LIST
|
||||
|
||||
|
||||
@comfytype(io_type="IMAGECOMPARE")
|
||||
class ImageCompare(ComfyTypeI):
|
||||
Type = dict
|
||||
@@ -1567,8 +1418,6 @@ def setup_dynamic_input_funcs():
|
||||
register_dynamic_input_func(DynamicCombo.io_type, DynamicCombo._expand_schema_for_dynamic)
|
||||
# DynamicSlot.Input
|
||||
register_dynamic_input_func(DynamicSlot.io_type, DynamicSlot._expand_schema_for_dynamic)
|
||||
# DynamicGroup.Input
|
||||
register_dynamic_input_func(DynamicGroup.io_type, DynamicGroup._expand_schema_for_dynamic)
|
||||
|
||||
if len(DYNAMIC_INPUT_LOOKUP) == 0:
|
||||
setup_dynamic_input_funcs()
|
||||
@@ -1580,8 +1429,6 @@ class V3Data(TypedDict):
|
||||
'Dictionary where the keys are the input ids and the values dictate how to turn the inputs into a nested dictionary.'
|
||||
dynamic_paths_default_value: dict[str, Any]
|
||||
'Dictionary where the keys are the input ids and the values are a string from DynamicPathsDefaultValue for the inputs if value is None.'
|
||||
list_paths: set[str]
|
||||
'Set of top-level keys whose index-keyed dict values should be converted to a sorted list[dict] after build_nested_inputs runs.'
|
||||
create_dynamic_tuple: bool
|
||||
'When True, the value of the dynamic input will be in the format (value, path_key).'
|
||||
|
||||
@@ -1923,7 +1770,6 @@ def get_finalized_class_inputs(d: dict[str, Any], live_inputs: dict[str, Any], i
|
||||
"optional": {},
|
||||
"dynamic_paths": {},
|
||||
"dynamic_paths_default_value": {},
|
||||
"list_paths": set(),
|
||||
}
|
||||
d = d.copy()
|
||||
# ignore hidden for parsing
|
||||
@@ -1939,10 +1785,6 @@ def get_finalized_class_inputs(d: dict[str, Any], live_inputs: dict[str, Any], i
|
||||
dynamic_paths_default_value = out_dict.pop("dynamic_paths_default_value", None)
|
||||
if dynamic_paths_default_value is not None and len(dynamic_paths_default_value) > 0:
|
||||
v3_data["dynamic_paths_default_value"] = dynamic_paths_default_value
|
||||
# list_paths: keys whose nested dict should be post-converted to a sorted list[dict]
|
||||
list_paths = out_dict.pop("list_paths", None)
|
||||
if list_paths:
|
||||
v3_data["list_paths"] = list_paths
|
||||
return out_dict, hidden, v3_data
|
||||
|
||||
def parse_class_inputs(out_dict: dict[str, Any], live_inputs: dict[str, Any], curr_dict: dict[str, Any], curr_prefix: list[str] | None=None) -> None:
|
||||
@@ -1978,12 +1820,10 @@ def add_to_dict_v1(i: Input, d: dict):
|
||||
|
||||
class DynamicPathsDefaultValue:
|
||||
EMPTY_DICT = "empty_dict"
|
||||
EMPTY_LIST = "empty_list"
|
||||
|
||||
def build_nested_inputs(values: dict[str, Any], v3_data: V3Data):
|
||||
paths = v3_data.get("dynamic_paths", None)
|
||||
default_value_dict = v3_data.get("dynamic_paths_default_value", {})
|
||||
list_paths: set[str] = v3_data.get("list_paths", set()) or set()
|
||||
if paths is None:
|
||||
return values
|
||||
values = values.copy()
|
||||
@@ -2006,8 +1846,6 @@ def build_nested_inputs(values: dict[str, Any], v3_data: V3Data):
|
||||
default_option = default_value_dict.get(key, None)
|
||||
if default_option == DynamicPathsDefaultValue.EMPTY_DICT:
|
||||
value = {}
|
||||
elif default_option == DynamicPathsDefaultValue.EMPTY_LIST:
|
||||
value = []
|
||||
if create_tuple:
|
||||
value = (value, key)
|
||||
current[p] = value
|
||||
@@ -2015,34 +1853,6 @@ def build_nested_inputs(values: dict[str, Any], v3_data: V3Data):
|
||||
current = current.setdefault(p, {})
|
||||
|
||||
values.update(result)
|
||||
|
||||
# Post-pass: convert index-keyed dicts to sorted lists for io.DynamicGroup fields
|
||||
for list_path in list_paths:
|
||||
parts = list_path.split(".")
|
||||
# Navigate to the parent container, then convert the leaf
|
||||
container = values
|
||||
for part in parts[:-1]:
|
||||
if not isinstance(container, dict) or part not in container:
|
||||
container = None
|
||||
break
|
||||
container = container[part]
|
||||
if container is None:
|
||||
continue
|
||||
leaf_key = parts[-1]
|
||||
leaf = container.get(leaf_key, None)
|
||||
if isinstance(leaf, dict):
|
||||
try:
|
||||
sorted_rows = [leaf[k] for k in sorted(leaf.keys(), key=int)]
|
||||
container[leaf_key] = sorted_rows
|
||||
except (ValueError, TypeError):
|
||||
# Keys are not all integers; leave as-is
|
||||
pass
|
||||
elif isinstance(leaf, list):
|
||||
# Already a list (e.g. the EMPTY_LIST default was applied above)
|
||||
pass
|
||||
elif leaf is None:
|
||||
container[leaf_key] = []
|
||||
|
||||
return values
|
||||
|
||||
|
||||
@@ -2607,9 +2417,7 @@ __all__ = [
|
||||
# Dynamic Types
|
||||
"MatchType",
|
||||
"DynamicCombo",
|
||||
"DynamicSlot",
|
||||
"Autogrow",
|
||||
"DynamicGroup",
|
||||
# Other classes
|
||||
"HiddenHolder",
|
||||
"Hidden",
|
||||
|
||||
@@ -11,7 +11,6 @@ from io import BytesIO
|
||||
from yarl import URL
|
||||
|
||||
from comfy.cli_args import args
|
||||
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
|
||||
@@ -64,7 +63,7 @@ def get_comfy_api_headers(node_cls: type[IO.ComfyNode]) -> dict[str, str]:
|
||||
|
||||
|
||||
def default_base_url() -> str:
|
||||
return normalize_comfy_api_base(getattr(args, "comfy_api_base", "https://api.comfy.org"))
|
||||
return getattr(args, "comfy_api_base", "https://api.comfy.org")
|
||||
|
||||
|
||||
async def sleep_with_interrupt(
|
||||
|
||||
@@ -0,0 +1,167 @@
|
||||
import math
|
||||
import nodes
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
import comfy.latent_formats
|
||||
import comfy.model_management
|
||||
import comfy.utils
|
||||
from comfy_api.latest import ComfyExtension, io
|
||||
from typing_extensions import override
|
||||
|
||||
|
||||
IMAGE_MIN_TOKEN_NUM = 4
|
||||
IMAGE_MAX_TOKEN_NUM = 16384
|
||||
MAX_RATIO = 200
|
||||
SPATIAL_MERGE_SIZE = 2
|
||||
VISION_PATCH_SIZE = 16
|
||||
|
||||
|
||||
def _crop_image(image, width, height):
|
||||
image = image[:1].movedim(-1, 1)
|
||||
image = comfy.utils.common_upscale(image, width, height, "bilinear", "center")
|
||||
return image.movedim(1, -1)[:, :, :, :3]
|
||||
|
||||
|
||||
def _round_by_factor(number, factor):
|
||||
return round(number / factor) * factor
|
||||
|
||||
|
||||
def _ceil_by_factor(number, factor):
|
||||
return math.ceil(number / factor) * factor
|
||||
|
||||
|
||||
def _floor_by_factor(number, factor):
|
||||
return math.floor(number / factor) * factor
|
||||
|
||||
|
||||
def _smart_resize(height, width, factor, min_pixels=None, max_pixels=None):
|
||||
max_pixels = max_pixels if max_pixels is not None else IMAGE_MAX_TOKEN_NUM * factor ** 2
|
||||
min_pixels = min_pixels if min_pixels is not None else IMAGE_MIN_TOKEN_NUM * factor ** 2
|
||||
if max_pixels < min_pixels:
|
||||
raise ValueError("max_pixels must be greater than or equal to min_pixels.")
|
||||
if max(height, width) / min(height, width) > MAX_RATIO:
|
||||
raise ValueError(f"LingBotVideo image aspect ratio must be smaller than {MAX_RATIO}.")
|
||||
|
||||
resized_height = max(factor, _round_by_factor(height, factor))
|
||||
resized_width = max(factor, _round_by_factor(width, factor))
|
||||
if resized_height * resized_width > max_pixels:
|
||||
beta = math.sqrt((height * width) / max_pixels)
|
||||
resized_height = _floor_by_factor(height / beta, factor)
|
||||
resized_width = _floor_by_factor(width / beta, factor)
|
||||
elif resized_height * resized_width < min_pixels:
|
||||
beta = math.sqrt(min_pixels / (height * width))
|
||||
resized_height = _ceil_by_factor(height * beta, factor)
|
||||
resized_width = _ceil_by_factor(width * beta, factor)
|
||||
return resized_height, resized_width
|
||||
|
||||
|
||||
def _vlm_image(image):
|
||||
factor = VISION_PATCH_SIZE * SPATIAL_MERGE_SIZE
|
||||
height, width = image.shape[1:3]
|
||||
resized_height, resized_width = _smart_resize(height, width, factor)
|
||||
array = image[0].detach().cpu().clamp(0, 1).mul(255).byte().numpy()
|
||||
pil_image = Image.fromarray(array, mode="RGB").resize((resized_width, resized_height))
|
||||
array = np.asarray(pil_image).astype(np.float32) / 255.0
|
||||
return torch.from_numpy(array).unsqueeze(0)
|
||||
|
||||
|
||||
class TextEncodeLingBotVideoI2V(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="TextEncodeLingBotVideoI2V",
|
||||
category="model/conditioning/lingbot_video",
|
||||
inputs=[
|
||||
io.Clip.Input("clip"),
|
||||
io.String.Input("prompt", multiline=True, dynamic_prompts=True),
|
||||
io.String.Input("negative_prompt", multiline=True, dynamic_prompts=True, default=""),
|
||||
io.Int.Input("width", default=480, min=16, max=nodes.MAX_RESOLUTION, step=16),
|
||||
io.Int.Input("height", default=480, min=16, max=nodes.MAX_RESOLUTION, step=16),
|
||||
io.Image.Input("image", optional=True),
|
||||
],
|
||||
outputs=[
|
||||
io.Conditioning.Output(display_name="positive"),
|
||||
io.Conditioning.Output(display_name="negative"),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, clip, prompt, negative_prompt, width, height, image=None) -> io.NodeOutput:
|
||||
if height % 16 != 0 or width % 16 != 0:
|
||||
raise ValueError(f"LingBotVideo width and height must be multiples of 16, got {width}x{height}.")
|
||||
if image is None:
|
||||
images = []
|
||||
else:
|
||||
image = _crop_image(image, width, height)
|
||||
images = [_vlm_image(image)]
|
||||
positive = clip.encode_from_tokens_scheduled(clip.tokenize(prompt, images=images))
|
||||
negative = clip.encode_from_tokens_scheduled(clip.tokenize(negative_prompt, images=images))
|
||||
return io.NodeOutput(positive, negative)
|
||||
|
||||
|
||||
class LingBotImageToVideo(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="LingBotImageToVideo",
|
||||
category="model/conditioning/lingbot_video",
|
||||
inputs=[
|
||||
io.Conditioning.Input("positive"),
|
||||
io.Conditioning.Input("negative"),
|
||||
io.Vae.Input("vae"),
|
||||
io.Int.Input("width", default=480, min=16, max=nodes.MAX_RESOLUTION, step=16),
|
||||
io.Int.Input("height", default=480, min=16, max=nodes.MAX_RESOLUTION, step=16),
|
||||
io.Int.Input("length", default=81, min=1, max=nodes.MAX_RESOLUTION, step=4),
|
||||
io.Int.Input("batch_size", default=1, min=1, max=4096),
|
||||
io.Image.Input("start_image", optional=True),
|
||||
],
|
||||
outputs=[
|
||||
io.Conditioning.Output(display_name="positive"),
|
||||
io.Conditioning.Output(display_name="negative"),
|
||||
io.Latent.Output(display_name="latent"),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, positive, negative, vae, width, height, length, batch_size, start_image=None) -> io.NodeOutput:
|
||||
if length != 1 and (length - 1) % 4 != 0:
|
||||
raise ValueError(f"LingBotVideo length must be 1 or 4n+1, got {length}.")
|
||||
if height % 16 != 0 or width % 16 != 0:
|
||||
raise ValueError(f"LingBotVideo width and height must be multiples of 16, got {width}x{height}.")
|
||||
latent_frames = ((length - 1) // 4) + 1
|
||||
latent = torch.zeros(
|
||||
[batch_size, 16, latent_frames, height // 8, width // 8],
|
||||
device=comfy.model_management.intermediate_device(),
|
||||
)
|
||||
out_latent = {"samples": latent}
|
||||
|
||||
if start_image is not None:
|
||||
start_image = _crop_image(start_image, width, height)
|
||||
cond_latent = comfy.latent_formats.LingBotVideo().process_in(vae.encode(start_image))
|
||||
cond_latent = comfy.utils.resize_to_batch_size(cond_latent, batch_size)
|
||||
cond_t = min(cond_latent.shape[2], latent.shape[2])
|
||||
latent[:, :, :cond_t] = cond_latent[:, :, :cond_t].to(device=latent.device, dtype=latent.dtype)
|
||||
noise_mask = torch.ones(
|
||||
(batch_size, 1, latent.shape[2], latent.shape[3], latent.shape[4]),
|
||||
device=latent.device,
|
||||
dtype=latent.dtype,
|
||||
)
|
||||
noise_mask[:, :, :cond_t] = 0.0
|
||||
out_latent["noise_mask"] = noise_mask
|
||||
|
||||
return io.NodeOutput(positive, negative, out_latent)
|
||||
|
||||
|
||||
class LingBotVideoExtension(ComfyExtension):
|
||||
@override
|
||||
async def get_node_list(self) -> list[type[io.ComfyNode]]:
|
||||
return [
|
||||
TextEncodeLingBotVideoI2V,
|
||||
LingBotImageToVideo,
|
||||
]
|
||||
|
||||
|
||||
async def comfy_entrypoint() -> LingBotVideoExtension:
|
||||
return LingBotVideoExtension()
|
||||
@@ -1,107 +0,0 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing_extensions import override
|
||||
|
||||
import comfy.sd
|
||||
import comfy.utils
|
||||
import folder_paths
|
||||
from comfy_api.latest import ComfyExtension, io
|
||||
|
||||
|
||||
def _load_lora_file(lora_name: str):
|
||||
lora_path = folder_paths.get_full_path_or_raise("loras", lora_name)
|
||||
return comfy.utils.load_torch_file(lora_path, safe_load=True, return_metadata=True)
|
||||
|
||||
|
||||
def _lora_template() -> list[io.Input]:
|
||||
return [
|
||||
io.Combo.Input("lora_name", options=folder_paths.get_filename_list("loras"),
|
||||
tooltip="The name of the LoRA file to apply."),
|
||||
io.Float.Input("strength", default=1.0, min=-100.0, max=100.0, step=0.01,
|
||||
tooltip="How strongly to apply this LoRA. 0 = off, negative inverts the effect."),
|
||||
]
|
||||
|
||||
|
||||
class LoadLoraModel(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="LoadLoraModel",
|
||||
display_name="Load LoRA (Model)",
|
||||
search_aliases=["lora", "load lora", "apply lora", "lora model", "lora stack"],
|
||||
category="model/loaders",
|
||||
description="Apply a stack of LoRAs to a diffusion model. Add one row per LoRA; "
|
||||
"each row picks a LoRA file and its strength.",
|
||||
inputs=[
|
||||
io.Model.Input("model", tooltip="The diffusion model the LoRAs will be applied to."),
|
||||
io.DynamicGroup.Input(
|
||||
"loras",
|
||||
template=_lora_template(),
|
||||
min=1,
|
||||
max=50,
|
||||
tooltip="Each row applies one LoRA to the model.",
|
||||
group_name="LoRA",
|
||||
),
|
||||
],
|
||||
outputs=[io.Model.Output(tooltip="The modified diffusion model.")],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, model, loras: list[dict]) -> io.NodeOutput:
|
||||
for row in loras:
|
||||
lora_name = row.get("lora_name")
|
||||
strength = row.get("strength", 1.0)
|
||||
if not lora_name or lora_name == "none" or strength == 0:
|
||||
continue
|
||||
lora, metadata = _load_lora_file(lora_name)
|
||||
model, _ = comfy.sd.load_lora_for_models(model, None, lora, strength, 0, lora_metadata=metadata)
|
||||
return io.NodeOutput(model)
|
||||
|
||||
|
||||
class LoadLoraTextEncoder(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="LoadLoraTextEncoder",
|
||||
display_name="Load LoRA (Text Encoder)",
|
||||
search_aliases=["lora", "load lora", "apply lora", "clip lora", "lora stack"],
|
||||
category="model/loaders",
|
||||
description="Apply a stack of LoRAs to a CLIP text encoder. Add one row per LoRA; "
|
||||
"each row picks a LoRA file and its strength.",
|
||||
inputs=[
|
||||
io.Clip.Input("clip", tooltip="The CLIP text encoder the LoRAs will be applied to."),
|
||||
io.DynamicGroup.Input(
|
||||
"loras",
|
||||
template=_lora_template(),
|
||||
min=1,
|
||||
max=50,
|
||||
tooltip="Each row applies one LoRA to the text encoder.",
|
||||
group_name="LoRA",
|
||||
),
|
||||
],
|
||||
outputs=[io.Clip.Output(tooltip="The modified CLIP text encoder.")],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, clip, loras: list[dict]) -> io.NodeOutput:
|
||||
for row in loras:
|
||||
lora_name = row.get("lora_name")
|
||||
strength = row.get("strength", 1.0)
|
||||
if not lora_name or lora_name == "none" or strength == 0:
|
||||
continue
|
||||
lora, metadata = _load_lora_file(lora_name)
|
||||
_, clip = comfy.sd.load_lora_for_models(None, clip, lora, 0, strength, lora_metadata=metadata)
|
||||
return io.NodeOutput(clip)
|
||||
|
||||
|
||||
class LoraStackExtension(ComfyExtension):
|
||||
@override
|
||||
async def get_node_list(self) -> list[type[io.ComfyNode]]:
|
||||
return [
|
||||
LoadLoraModel,
|
||||
LoadLoraTextEncoder,
|
||||
]
|
||||
|
||||
|
||||
async def comfy_entrypoint() -> LoraStackExtension:
|
||||
return LoraStackExtension()
|
||||
@@ -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"], ),
|
||||
"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", "lingbot_video"], ),
|
||||
},
|
||||
"optional": {
|
||||
"device": (["default", "cpu"], {"advanced": True}),
|
||||
@@ -2460,6 +2460,7 @@ async def init_builtin_extra_nodes():
|
||||
"nodes_tcfg.py",
|
||||
"nodes_context_windows.py",
|
||||
"nodes_qwen.py",
|
||||
"nodes_lingbot_video.py",
|
||||
"nodes_boogu.py",
|
||||
"nodes_chroma_radiance.py",
|
||||
"nodes_pid.py",
|
||||
@@ -2503,7 +2504,6 @@ async def init_builtin_extra_nodes():
|
||||
"nodes_triposplat.py",
|
||||
"nodes_depth_anything_3.py",
|
||||
"nodes_seed.py",
|
||||
"nodes_lora_stack.py",
|
||||
]
|
||||
|
||||
import_failed = []
|
||||
|
||||
+18
-13
@@ -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,6 +144,14 @@ 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}$
|
||||
@@ -1636,7 +1644,7 @@ paths:
|
||||
format: uuid
|
||||
type: string
|
||||
tags:
|
||||
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.
|
||||
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.
|
||||
type: string
|
||||
user_metadata:
|
||||
description: Custom JSON metadata as a string
|
||||
@@ -1821,7 +1829,7 @@ paths:
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/Asset'
|
||||
$ref: '#/components/schemas/AssetUpdated'
|
||||
description: Asset updated successfully
|
||||
"400":
|
||||
content:
|
||||
@@ -2462,9 +2470,6 @@ paths:
|
||||
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:
|
||||
|
||||
@@ -39,7 +39,6 @@ from comfy.deploy_environment import get_deploy_environment
|
||||
import comfy.utils
|
||||
import comfy.model_management
|
||||
from comfy_api import feature_flags
|
||||
from comfy.comfy_api_env import get_environment_overrides
|
||||
import node_helpers
|
||||
from comfyui_version import __version__
|
||||
from app.frontend_management import FrontendManager, parse_version
|
||||
@@ -47,7 +46,6 @@ from comfy_api.internal import _ComfyNodeInternal
|
||||
from app.assets.seeder import asset_seeder
|
||||
from app.assets.api.routes import register_assets_routes
|
||||
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
|
||||
|
||||
from app.user_manager import UserManager
|
||||
@@ -443,9 +441,7 @@ class PromptServer():
|
||||
if args.enable_assets:
|
||||
try:
|
||||
tag = image_upload_type if image_upload_type in ("input", "output") else "input"
|
||||
tags = [tag]
|
||||
tags.extend(get_known_subfolder_tags(subfolder))
|
||||
result = register_file_in_place(abs_path=filepath, name=filename, tags=tags)
|
||||
result = register_file_in_place(abs_path=filepath, name=filename, tags=[tag])
|
||||
resp["asset"] = {
|
||||
"id": result.ref.id,
|
||||
"name": result.ref.name,
|
||||
@@ -728,11 +724,7 @@ class PromptServer():
|
||||
|
||||
@routes.get("/features")
|
||||
async def get_features(request):
|
||||
features = feature_flags.get_server_features()
|
||||
overrides = get_environment_overrides()
|
||||
if overrides:
|
||||
features.update(overrides)
|
||||
return web.json_response(features)
|
||||
return web.json_response(feature_flags.get_server_features())
|
||||
|
||||
@routes.get("/prompt")
|
||||
async def get_prompt(request):
|
||||
|
||||
@@ -8,7 +8,6 @@ upgrade/downgrade for 0003+.
|
||||
"""
|
||||
|
||||
import os
|
||||
import sqlite3
|
||||
|
||||
import pytest
|
||||
from alembic import command
|
||||
@@ -31,12 +30,6 @@ def _make_config(db_path: str) -> Config:
|
||||
return cfg
|
||||
|
||||
|
||||
def _sqlite_path(cfg: Config) -> str:
|
||||
url = cfg.get_main_option("sqlalchemy.url")
|
||||
assert url is not None and url.startswith("sqlite:///")
|
||||
return url.removeprefix("sqlite:///")
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def migration_db(tmp_path):
|
||||
"""Yield an alembic Config pre-upgraded to the baseline revision."""
|
||||
@@ -62,26 +55,3 @@ def test_upgrade_downgrade_cycle(migration_db):
|
||||
command.upgrade(migration_db, "head")
|
||||
command.downgrade(migration_db, _BASELINE)
|
||||
command.upgrade(migration_db, "head")
|
||||
|
||||
|
||||
def test_case_sensitive_tags_downgrade_normalizes_existing_tags(migration_db):
|
||||
"""Downgrading 0005 folds mixed-case tag vocabulary before restoring CHECK."""
|
||||
command.upgrade(migration_db, "0005_allow_case_sensitive_tags")
|
||||
|
||||
db_path = _sqlite_path(migration_db)
|
||||
with sqlite3.connect(db_path) as conn:
|
||||
conn.execute("INSERT INTO tags(name) VALUES (?)", ("NewTag",))
|
||||
conn.execute("INSERT INTO tags(name) VALUES (?)", ("newtag",))
|
||||
conn.execute("INSERT INTO tags(name) VALUES (?)", ("model_type:LLM",))
|
||||
|
||||
command.downgrade(migration_db, "0004_drop_tag_type")
|
||||
|
||||
with sqlite3.connect(db_path) as conn:
|
||||
tags = {row[0] for row in conn.execute("SELECT name FROM tags")}
|
||||
assert "newtag" in tags
|
||||
assert "model_type:llm" in tags
|
||||
assert "NewTag" not in tags
|
||||
assert "model_type:LLM" not in tags
|
||||
|
||||
with pytest.raises(sqlite3.IntegrityError):
|
||||
conn.execute("INSERT INTO tags(name) VALUES (?)", ("Upper",))
|
||||
|
||||
@@ -234,7 +234,7 @@ def seeded_asset(request: pytest.FixtureRequest, http: requests.Session, api_bas
|
||||
p = getattr(request, "param", {}) or {}
|
||||
tags: Optional[list[str]] = p.get("tags")
|
||||
if tags is None:
|
||||
tags = ["models", "model_type:checkpoints", "unit-tests", "alpha"]
|
||||
tags = ["models", "checkpoints", "unit-tests", "alpha"]
|
||||
meta = {"purpose": "test", "epoch": 1, "flags": ["x", "y"], "nullable": None}
|
||||
# Unique content per test so the seed always creates a fresh asset (201).
|
||||
# Delete is now always a soft delete, so content from a prior test survives
|
||||
|
||||
@@ -133,66 +133,6 @@ class TestListReferencesPage:
|
||||
assert total == 1
|
||||
assert refs[0].name == "tagged"
|
||||
|
||||
def test_include_tags_filter_ands_persisted_model_tags(self, session: Session):
|
||||
asset = _make_asset(session, "hash-model-tags")
|
||||
checkpoint = _make_reference(session, asset, name="checkpoint")
|
||||
lora = _make_reference(session, asset, name="lora")
|
||||
input_ref = _make_reference(session, asset, name="input")
|
||||
ensure_tags_exist(
|
||||
session,
|
||||
["models", "model_type:checkpoints", "model_type:loras", "unit-tests"],
|
||||
)
|
||||
add_tags_to_reference(
|
||||
session,
|
||||
reference_id=checkpoint.id,
|
||||
tags=["models", "model_type:checkpoints", "unit-tests"],
|
||||
origin="automatic",
|
||||
)
|
||||
add_tags_to_reference(
|
||||
session,
|
||||
reference_id=lora.id,
|
||||
tags=["models", "model_type:loras", "unit-tests"],
|
||||
origin="automatic",
|
||||
)
|
||||
add_tags_to_reference(
|
||||
session,
|
||||
reference_id=input_ref.id,
|
||||
tags=["unit-tests"],
|
||||
)
|
||||
session.commit()
|
||||
|
||||
refs, _, total = list_references_page(
|
||||
session,
|
||||
include_tags=["models", "model_type:checkpoints", "unit-tests"],
|
||||
)
|
||||
|
||||
assert total == 1
|
||||
assert refs[0].id == checkpoint.id
|
||||
|
||||
def test_include_tags_filter_preserves_model_type_case(self, session: Session):
|
||||
asset = _make_asset(session, "hash-model-case")
|
||||
ref = _make_reference(session, asset, name="llm")
|
||||
ensure_tags_exist(session, ["models", "model_type:LLM"])
|
||||
add_tags_to_reference(
|
||||
session,
|
||||
reference_id=ref.id,
|
||||
tags=["models", "model_type:LLM"],
|
||||
origin="automatic",
|
||||
)
|
||||
session.commit()
|
||||
|
||||
refs, _, total = list_references_page(
|
||||
session, include_tags=["models", "model_type:LLM"]
|
||||
)
|
||||
refs_lower, _, total_lower = list_references_page(
|
||||
session, include_tags=["models", "model_type:llm"]
|
||||
)
|
||||
|
||||
assert total == 1
|
||||
assert refs[0].id == ref.id
|
||||
assert total_lower == 0
|
||||
assert refs_lower == []
|
||||
|
||||
def test_exclude_tags_filter(self, session: Session):
|
||||
asset = _make_asset(session, "hash1")
|
||||
_make_reference(session, asset, name="keep")
|
||||
|
||||
@@ -176,39 +176,6 @@ class TestUpsertReference:
|
||||
ref = session.query(AssetReference).filter_by(file_path=file_path).one()
|
||||
assert ref.mtime_ns == final_mtime
|
||||
|
||||
def test_upsert_refreshes_loader_path_on_existing_reference(self, session: Session):
|
||||
"""Re-ingesting an existing reference writes the loader_path computed
|
||||
by that ingest, healing NULL or stale values even when nothing else
|
||||
about the row changed."""
|
||||
asset = _make_asset(session, "hash1")
|
||||
file_path = "/models/checkpoints/sub/model.safetensors"
|
||||
|
||||
upsert_reference(
|
||||
session, asset_id=asset.id, file_path=file_path, name="model",
|
||||
mtime_ns=100, loader_path=None,
|
||||
)
|
||||
session.commit()
|
||||
|
||||
created, updated = upsert_reference(
|
||||
session, asset_id=asset.id, file_path=file_path, name="model",
|
||||
mtime_ns=100, loader_path="sub/model.safetensors",
|
||||
)
|
||||
session.commit()
|
||||
|
||||
assert created is False
|
||||
assert updated is True
|
||||
ref = session.query(AssetReference).filter_by(file_path=file_path).one()
|
||||
assert ref.loader_path == "sub/model.safetensors"
|
||||
|
||||
# Identical loader_path is a no-op, not a spurious update.
|
||||
created, updated = upsert_reference(
|
||||
session, asset_id=asset.id, file_path=file_path, name="model",
|
||||
mtime_ns=100, loader_path="sub/model.safetensors",
|
||||
)
|
||||
session.commit()
|
||||
assert created is False
|
||||
assert updated is False
|
||||
|
||||
def test_upsert_restores_missing_reference(self, session: Session):
|
||||
"""Upserting a reference that was marked missing should restore it."""
|
||||
asset = _make_asset(session, "hash1")
|
||||
|
||||
@@ -58,7 +58,7 @@ class TestEnsureTagsExist:
|
||||
session.commit()
|
||||
|
||||
tags = session.query(Tag).all()
|
||||
assert {t.name for t in tags} == {"ALPHA", "Beta", "alpha"}
|
||||
assert {t.name for t in tags} == {"alpha", "beta"}
|
||||
|
||||
def test_empty_list_is_noop(self, session: Session):
|
||||
ensure_tags_exist(session, [])
|
||||
@@ -258,16 +258,6 @@ class TestListTagsWithUsage:
|
||||
tag_names = {name for name, _ in rows}
|
||||
assert tag_names == {"alpha", "alphabet"}
|
||||
|
||||
def test_prefix_filter_is_case_sensitive(self, session: Session):
|
||||
ensure_tags_exist(session, ["model_type:LLM", "model_type:llm"])
|
||||
session.commit()
|
||||
|
||||
rows, total = list_tags_with_usage(session, prefix="model_type:L")
|
||||
|
||||
tag_names = {name for name, _ in rows}
|
||||
assert tag_names == {"model_type:LLM"}
|
||||
assert total == 1
|
||||
|
||||
def test_order_by_name(self, session: Session):
|
||||
ensure_tags_exist(session, ["zebra", "alpha", "middle"])
|
||||
session.commit()
|
||||
|
||||
@@ -1,83 +0,0 @@
|
||||
"""Tests for how _build_asset_response derives the response `loader_path`.
|
||||
|
||||
Guards the persist-and-read contract: the response reads the stored
|
||||
`loader_path` verbatim, with no read-time recomputation. Like tags, the
|
||||
value is a seed-time derivative healed by the scan lifecycle.
|
||||
"""
|
||||
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
from app.assets.api.routes import _build_asset_response
|
||||
from app.assets.services.schemas import AssetDetailResult, ReferenceData
|
||||
|
||||
_TS = datetime(2024, 1, 1, 0, 0, 0)
|
||||
|
||||
|
||||
def _make_result(
|
||||
*, file_path: str | None, loader_path: str | None
|
||||
) -> AssetDetailResult:
|
||||
ref = ReferenceData(
|
||||
id="ref-1",
|
||||
name="model.safetensors",
|
||||
file_path=file_path,
|
||||
loader_path=loader_path,
|
||||
user_metadata=None,
|
||||
preview_id=None,
|
||||
created_at=_TS,
|
||||
updated_at=_TS,
|
||||
last_access_time=_TS,
|
||||
)
|
||||
return AssetDetailResult(ref=ref, asset=None, tags=[])
|
||||
|
||||
|
||||
def test_uses_persisted_loader_path_without_recomputing():
|
||||
"""A stored loader_path is returned verbatim, not re-derived from file_path.
|
||||
|
||||
The sentinel value could never be produced by compute_loader_path for this
|
||||
file_path, so seeing it in the response proves the stored column is read.
|
||||
"""
|
||||
result = _make_result(
|
||||
file_path="/unmatched/root/model.safetensors",
|
||||
loader_path="SENTINEL/stored.safetensors",
|
||||
)
|
||||
|
||||
resp = _build_asset_response(result)
|
||||
|
||||
assert resp.loader_path == "SENTINEL/stored.safetensors"
|
||||
|
||||
|
||||
def test_null_stored_loader_path_is_served_as_null(tmp_path: Path):
|
||||
"""No read-time recomputation: a NULL column is served as null even when
|
||||
the path would resolve."""
|
||||
models = tmp_path / "models"
|
||||
ckpt = models / "checkpoints"
|
||||
ckpt.mkdir(parents=True)
|
||||
f = ckpt / "bar.safetensors"
|
||||
f.touch()
|
||||
|
||||
with patch("app.assets.services.path_utils.folder_paths") as mock_fp, patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[("checkpoints", [str(ckpt)], {".safetensors"})],
|
||||
):
|
||||
mock_fp.get_input_directory.return_value = str(tmp_path / "in")
|
||||
mock_fp.get_output_directory.return_value = str(tmp_path / "out")
|
||||
mock_fp.get_temp_directory.return_value = str(tmp_path / "tmp")
|
||||
mock_fp.models_dir = str(models)
|
||||
|
||||
result = _make_result(file_path=str(f), loader_path=None)
|
||||
resp = _build_asset_response(result)
|
||||
|
||||
assert resp.loader_path is None
|
||||
assert resp.display_name == "checkpoints/bar.safetensors"
|
||||
|
||||
|
||||
def test_all_path_fields_null_without_file_path():
|
||||
"""API-created / hash-only references (no file_path) expose no paths."""
|
||||
result = _make_result(file_path=None, loader_path=None)
|
||||
|
||||
resp = _build_asset_response(result)
|
||||
|
||||
assert resp.loader_path is None
|
||||
assert resp.display_name is None
|
||||
@@ -1,14 +1,10 @@
|
||||
"""Tests for bulk ingest services."""
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.assets.database.models import Asset, AssetReference
|
||||
from app.assets.database.queries import get_reference_tags
|
||||
from app.assets.scanner import build_asset_specs
|
||||
from app.assets.services.bulk_ingest import SeedAssetSpec, batch_insert_seed_assets
|
||||
|
||||
|
||||
@@ -105,184 +101,6 @@ class TestBatchInsertSeedAssets:
|
||||
asset = session.query(Asset).filter_by(id=ref.asset_id).first()
|
||||
assert asset.mime_type == expected_mime, f"Expected {expected_mime} for {filename}, got {asset.mime_type}"
|
||||
|
||||
def test_duplicate_paths_merge_tags_before_insert(
|
||||
self, session: Session, temp_dir: Path
|
||||
):
|
||||
"""Overlapping model-folder registrations can emit the same path twice."""
|
||||
file_path = temp_dir / "shared.safetensors"
|
||||
file_path.write_bytes(b"shared model")
|
||||
|
||||
specs: list[SeedAssetSpec] = [
|
||||
{
|
||||
"abs_path": str(file_path),
|
||||
"size_bytes": 12,
|
||||
"mtime_ns": 1234567890000000000,
|
||||
"info_name": "Shared Model",
|
||||
"tags": ["models", "model_type:checkpoints"],
|
||||
"fname": "shared.safetensors",
|
||||
"metadata": None,
|
||||
"hash": None,
|
||||
"mime_type": "application/safetensors",
|
||||
},
|
||||
{
|
||||
"abs_path": str(file_path),
|
||||
"size_bytes": 12,
|
||||
"mtime_ns": 1234567890000000000,
|
||||
"info_name": "Shared Model",
|
||||
"tags": ["models", "model_type:diffusion_models"],
|
||||
"fname": "shared.safetensors",
|
||||
"metadata": None,
|
||||
"hash": None,
|
||||
"mime_type": "application/safetensors",
|
||||
},
|
||||
]
|
||||
|
||||
result = batch_insert_seed_assets(session, specs=specs, owner_id="")
|
||||
|
||||
assert result.inserted_refs == 1
|
||||
assert result.won_paths == 1
|
||||
refs = session.query(AssetReference).all()
|
||||
assert len(refs) == 1
|
||||
assert set(get_reference_tags(session, reference_id=refs[0].id)) == {
|
||||
"models",
|
||||
"model_type:checkpoints",
|
||||
"model_type:diffusion_models",
|
||||
}
|
||||
|
||||
def test_duplicate_paths_are_merged_after_abspath_normalization(
|
||||
self, session: Session, temp_dir: Path, monkeypatch
|
||||
):
|
||||
"""The scanner may emit equivalent paths with different spelling."""
|
||||
file_path = temp_dir / "same-file.safetensors"
|
||||
file_path.write_bytes(b"shared model")
|
||||
monkeypatch.chdir(temp_dir)
|
||||
relative_path = file_path.name
|
||||
absolute_path = os.path.abspath(relative_path)
|
||||
|
||||
specs: list[SeedAssetSpec] = [
|
||||
{
|
||||
"abs_path": relative_path,
|
||||
"size_bytes": 12,
|
||||
"mtime_ns": 1234567890000000000,
|
||||
"info_name": "Shared Model",
|
||||
"tags": ["models", "model_type:checkpoints"],
|
||||
"fname": "same-file.safetensors",
|
||||
"metadata": None,
|
||||
"hash": None,
|
||||
"mime_type": "application/safetensors",
|
||||
},
|
||||
{
|
||||
"abs_path": absolute_path,
|
||||
"size_bytes": 12,
|
||||
"mtime_ns": 1234567890000000000,
|
||||
"info_name": "Shared Model",
|
||||
"tags": ["models", "model_type:diffusion_models"],
|
||||
"fname": "same-file.safetensors",
|
||||
"metadata": None,
|
||||
"hash": None,
|
||||
"mime_type": "application/safetensors",
|
||||
},
|
||||
]
|
||||
|
||||
result = batch_insert_seed_assets(session, specs=specs, owner_id="")
|
||||
|
||||
assert result.inserted_refs == 1
|
||||
assert result.won_paths == 1
|
||||
refs = session.query(AssetReference).all()
|
||||
assert len(refs) == 1
|
||||
assert refs[0].file_path == absolute_path
|
||||
# loader_path is persisted from the spec's fname (compute_loader_path).
|
||||
assert refs[0].loader_path == "same-file.safetensors"
|
||||
assert set(get_reference_tags(session, reference_id=refs[0].id)) == {
|
||||
"models",
|
||||
"model_type:checkpoints",
|
||||
"model_type:diffusion_models",
|
||||
}
|
||||
|
||||
def test_scanner_duplicate_shared_model_paths_keep_all_model_type_tags(
|
||||
self, session: Session, temp_dir: Path
|
||||
):
|
||||
"""Shared extra model roots make scanner collection emit duplicate paths."""
|
||||
shared_root = temp_dir / "shared"
|
||||
input_dir = temp_dir / "input"
|
||||
output_dir = temp_dir / "output"
|
||||
temp_root = temp_dir / "temp"
|
||||
for directory in (shared_root, input_dir, output_dir, temp_root):
|
||||
directory.mkdir()
|
||||
file_path = shared_root / "dual_use_model.safetensors"
|
||||
file_path.write_bytes(b"shared model")
|
||||
|
||||
with (
|
||||
patch("app.assets.services.path_utils.folder_paths") as mock_fp,
|
||||
patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[
|
||||
("checkpoints", [str(shared_root)], {".safetensors"}),
|
||||
("diffusion_models", [str(shared_root)], {".safetensors"}),
|
||||
],
|
||||
),
|
||||
):
|
||||
mock_fp.get_input_directory.return_value = str(input_dir)
|
||||
mock_fp.get_output_directory.return_value = str(output_dir)
|
||||
mock_fp.get_temp_directory.return_value = str(temp_root)
|
||||
|
||||
specs, tag_pool, skipped = build_asset_specs(
|
||||
paths=[str(file_path), str(file_path)],
|
||||
existing_paths=set(),
|
||||
enable_metadata_extraction=False,
|
||||
compute_hashes=False,
|
||||
)
|
||||
|
||||
assert skipped == 0
|
||||
assert len(specs) == 2
|
||||
assert tag_pool == {
|
||||
"models",
|
||||
"model_type:checkpoints",
|
||||
"model_type:diffusion_models",
|
||||
}
|
||||
|
||||
result = batch_insert_seed_assets(session, specs=specs, owner_id="")
|
||||
|
||||
assert result.inserted_refs == 1
|
||||
assert result.won_paths == 1
|
||||
refs = session.query(AssetReference).all()
|
||||
assert len(refs) == 1
|
||||
assert set(get_reference_tags(session, reference_id=refs[0].id)) == {
|
||||
"models",
|
||||
"model_type:checkpoints",
|
||||
"model_type:diffusion_models",
|
||||
}
|
||||
|
||||
def test_loader_path_persisted_as_null_when_fname_is_none(
|
||||
self, session: Session, temp_dir: Path
|
||||
):
|
||||
"""A file with no in-root loader path (fname=None, e.g. an orphan under
|
||||
models_root) persists loader_path as NULL rather than a synthesized value."""
|
||||
file_path = temp_dir / "orphan.bin"
|
||||
file_path.write_bytes(b"x")
|
||||
|
||||
specs: list[SeedAssetSpec] = [
|
||||
{
|
||||
"abs_path": str(file_path),
|
||||
"size_bytes": 1,
|
||||
"mtime_ns": 1234567890000000000,
|
||||
"info_name": "orphan.bin",
|
||||
"tags": [],
|
||||
"fname": None,
|
||||
"metadata": None,
|
||||
"hash": None,
|
||||
"mime_type": None,
|
||||
}
|
||||
]
|
||||
|
||||
result = batch_insert_seed_assets(session, specs=specs, owner_id="")
|
||||
|
||||
assert result.inserted_refs == 1
|
||||
refs = session.query(AssetReference).all()
|
||||
assert len(refs) == 1
|
||||
assert refs[0].file_path == str(file_path)
|
||||
assert refs[0].loader_path is None
|
||||
|
||||
|
||||
class TestMetadataExtraction:
|
||||
def test_extracts_mime_type_for_model_files(self, temp_dir: Path):
|
||||
|
||||
@@ -94,47 +94,6 @@ class TestIngestFileFromPath:
|
||||
ref_tags = get_reference_tags(session, reference_id=result.reference_id)
|
||||
assert set(ref_tags) == {"models", "checkpoints"}
|
||||
|
||||
def test_path_derived_tags_use_automatic_origin(
|
||||
self, mock_create_session, temp_dir: Path, session: Session
|
||||
):
|
||||
input_dir = temp_dir / "input"
|
||||
output_dir = temp_dir / "output"
|
||||
temp_root = temp_dir / "temp"
|
||||
for directory in (input_dir, output_dir, temp_root):
|
||||
directory.mkdir()
|
||||
file_path = input_dir / "pasted" / "tagged.png"
|
||||
file_path.parent.mkdir()
|
||||
file_path.write_bytes(b"data")
|
||||
|
||||
with (
|
||||
patch("app.assets.services.path_utils.folder_paths") as mock_fp,
|
||||
patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[],
|
||||
),
|
||||
):
|
||||
mock_fp.get_input_directory.return_value = str(input_dir)
|
||||
mock_fp.get_output_directory.return_value = str(output_dir)
|
||||
mock_fp.get_temp_directory.return_value = str(temp_root)
|
||||
|
||||
result = _ingest_file_from_path(
|
||||
abs_path=str(file_path),
|
||||
asset_hash="blake3:pathorigin",
|
||||
size_bytes=4,
|
||||
mtime_ns=1234567890000000000,
|
||||
info_name="Tagged Asset",
|
||||
tags=["input", "manual-label"],
|
||||
)
|
||||
|
||||
assert result.reference_id is not None
|
||||
links = session.query(AssetReferenceTag).filter_by(
|
||||
asset_reference_id=result.reference_id
|
||||
)
|
||||
origin_by_tag = {link.tag_name: link.origin for link in links}
|
||||
assert origin_by_tag["input"] == "automatic"
|
||||
assert origin_by_tag["pasted"] == "automatic"
|
||||
assert origin_by_tag["manual-label"] == "manual"
|
||||
|
||||
def test_idempotent_upsert(self, mock_create_session, temp_dir: Path, session: Session):
|
||||
file_path = temp_dir / "dup.bin"
|
||||
file_path.write_bytes(b"content")
|
||||
@@ -329,45 +288,6 @@ class TestIngestExistingFileTagFK:
|
||||
assert "output" in ref_tag_names
|
||||
|
||||
|
||||
class TestIngestExistingFileLoaderPath:
|
||||
"""Outputs saved into a subfolder must persist the subfolder-qualified
|
||||
loader path, not the bare basename (regression: spec["fname"] was
|
||||
os.path.basename)."""
|
||||
|
||||
def test_subfoldered_output_persists_relative_loader_path(
|
||||
self, mock_create_session, temp_dir: Path, session: Session
|
||||
):
|
||||
input_dir = temp_dir / "input"
|
||||
output_dir = temp_dir / "output"
|
||||
temp_root = temp_dir / "temp"
|
||||
for directory in (input_dir, output_dir, temp_root):
|
||||
directory.mkdir()
|
||||
file_path = output_dir / "sub" / "img_00001_.png"
|
||||
file_path.parent.mkdir()
|
||||
file_path.write_bytes(b"image data")
|
||||
|
||||
with (
|
||||
patch("app.assets.services.path_utils.folder_paths") as mock_fp,
|
||||
patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[],
|
||||
),
|
||||
):
|
||||
mock_fp.get_input_directory.return_value = str(input_dir)
|
||||
mock_fp.get_output_directory.return_value = str(output_dir)
|
||||
mock_fp.get_temp_directory.return_value = str(temp_root)
|
||||
|
||||
assert ingest_existing_file(abs_path=str(file_path)) is True
|
||||
|
||||
ref = (
|
||||
session.query(AssetReference)
|
||||
.filter_by(file_path=str(file_path))
|
||||
.one()
|
||||
)
|
||||
assert ref.loader_path == "sub/img_00001_.png"
|
||||
assert (ref.user_metadata or {}).get("filename") == "sub/img_00001_.png"
|
||||
|
||||
|
||||
class TestIngestImageDimensions:
|
||||
"""system_metadata should carry {kind, width, height} for image assets."""
|
||||
|
||||
|
||||
@@ -6,16 +6,7 @@ from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
from app.assets.services.path_utils import (
|
||||
compute_display_name,
|
||||
compute_loader_path,
|
||||
compute_logical_path,
|
||||
get_asset_category_and_relative_path,
|
||||
get_known_input_subfolder_tags_from_path,
|
||||
get_known_subfolder_tags,
|
||||
get_name_and_tags_from_asset_path,
|
||||
resolve_destination_from_tags,
|
||||
)
|
||||
from app.assets.services.path_utils import get_asset_category_and_relative_path
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
@@ -26,8 +17,7 @@ def fake_dirs():
|
||||
input_dir = root_path / "input"
|
||||
output_dir = root_path / "output"
|
||||
temp_dir = root_path / "temp"
|
||||
models_root = root_path / "models"
|
||||
models_dir = models_root / "checkpoints"
|
||||
models_dir = root_path / "models" / "checkpoints"
|
||||
for d in (input_dir, output_dir, temp_dir, models_dir):
|
||||
d.mkdir(parents=True)
|
||||
|
||||
@@ -35,17 +25,15 @@ def fake_dirs():
|
||||
mock_fp.get_input_directory.return_value = str(input_dir)
|
||||
mock_fp.get_output_directory.return_value = str(output_dir)
|
||||
mock_fp.get_temp_directory.return_value = str(temp_dir)
|
||||
mock_fp.models_dir = str(models_root)
|
||||
|
||||
with patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[("checkpoints", [str(models_dir)], {".safetensors"})],
|
||||
return_value=[("checkpoints", [str(models_dir)])],
|
||||
):
|
||||
yield {
|
||||
"input": input_dir,
|
||||
"output": output_dir,
|
||||
"temp": temp_dir,
|
||||
"models_root": models_root,
|
||||
"models": models_dir,
|
||||
}
|
||||
|
||||
@@ -88,538 +76,6 @@ class TestGetAssetCategoryAndRelativePath:
|
||||
cat, rel = get_asset_category_and_relative_path(str(f))
|
||||
assert cat == "models"
|
||||
|
||||
def test_model_path_tags_include_registered_model_type_only(self, fake_dirs):
|
||||
f = fake_dirs["models"] / "subdir" / "model.safetensors"
|
||||
f.parent.mkdir()
|
||||
f.touch()
|
||||
|
||||
_name, tags = get_name_and_tags_from_asset_path(str(f))
|
||||
|
||||
assert "models" in tags
|
||||
assert "model_type:checkpoints" in tags
|
||||
assert "checkpoints" not in tags
|
||||
assert "subdir" not in tags
|
||||
|
||||
def test_model_type_preserves_registered_folder_case(self, fake_dirs):
|
||||
llm_dir = fake_dirs["models"].parent / "LLM"
|
||||
llm_dir.mkdir()
|
||||
f = llm_dir / "model.safetensors"
|
||||
f.touch()
|
||||
|
||||
with patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[("LLM", [str(llm_dir)], {".safetensors"})],
|
||||
):
|
||||
_name, tags = get_name_and_tags_from_asset_path(str(f))
|
||||
|
||||
assert "models" in tags
|
||||
assert "model_type:LLM" in tags
|
||||
assert "model_type:llm" not in tags
|
||||
|
||||
def test_path_components_do_not_create_model_type_tags(self, fake_dirs):
|
||||
f = fake_dirs["models"] / "loras" / "model.safetensors"
|
||||
f.parent.mkdir()
|
||||
f.touch()
|
||||
|
||||
_name, tags = get_name_and_tags_from_asset_path(str(f))
|
||||
|
||||
assert "models" in tags
|
||||
assert "model_type:checkpoints" in tags
|
||||
assert "loras" not in tags
|
||||
assert "model_type:loras" not in tags
|
||||
|
||||
def test_shared_root_returns_all_matching_model_type_tags(self, fake_dirs):
|
||||
shared_root = fake_dirs["models"].parent / "shared"
|
||||
shared_root.mkdir()
|
||||
f = shared_root / "foo.safetensors"
|
||||
f.touch()
|
||||
|
||||
with patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[
|
||||
("checkpoints", [str(shared_root)], {".safetensors"}),
|
||||
("loras", [str(shared_root)], {".safetensors"}),
|
||||
],
|
||||
):
|
||||
_name, tags = get_name_and_tags_from_asset_path(str(f))
|
||||
|
||||
assert "models" in tags
|
||||
assert "model_type:checkpoints" in tags
|
||||
assert "model_type:loras" in tags
|
||||
|
||||
def test_shared_root_model_type_tags_respect_bucket_extensions(self, fake_dirs):
|
||||
"""Buckets sharing a base dir only tag files matching their extensions."""
|
||||
shared_root = fake_dirs["models"].parent / "unet"
|
||||
shared_root.mkdir()
|
||||
safetensors_file = shared_root / "wan.safetensors"
|
||||
gguf_file = shared_root / "wan.gguf"
|
||||
safetensors_file.touch()
|
||||
gguf_file.touch()
|
||||
|
||||
with patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[
|
||||
("diffusion_models", [str(shared_root)], {".safetensors"}),
|
||||
("unet_gguf", [str(shared_root)], {".gguf"}),
|
||||
],
|
||||
):
|
||||
_name, safetensors_tags = get_name_and_tags_from_asset_path(str(safetensors_file))
|
||||
_name, gguf_tags = get_name_and_tags_from_asset_path(str(gguf_file))
|
||||
|
||||
assert "model_type:diffusion_models" in safetensors_tags
|
||||
assert "model_type:unet_gguf" not in safetensors_tags
|
||||
assert "model_type:unet_gguf" in gguf_tags
|
||||
assert "model_type:diffusion_models" not in gguf_tags
|
||||
|
||||
def test_empty_extension_set_tags_any_extension(self, fake_dirs):
|
||||
"""Custom buckets registered without extensions accept every file."""
|
||||
custom_root = fake_dirs["models"].parent / "custom_bucket"
|
||||
custom_root.mkdir()
|
||||
f = custom_root / "weights.bin"
|
||||
f.touch()
|
||||
|
||||
with patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[("custom_bucket", [str(custom_root)], set())],
|
||||
):
|
||||
_name, tags = get_name_and_tags_from_asset_path(str(f))
|
||||
|
||||
assert "models" in tags
|
||||
assert "model_type:custom_bucket" in tags
|
||||
|
||||
def test_no_extension_match_keeps_models_tag_without_model_type(self, fake_dirs):
|
||||
f = fake_dirs["models"] / "notes.txt"
|
||||
f.touch()
|
||||
|
||||
_name, tags = get_name_and_tags_from_asset_path(str(f))
|
||||
|
||||
assert "models" in tags
|
||||
assert not any(tag.startswith("model_type:") for tag in tags)
|
||||
|
||||
def test_output_backed_registered_folder_gets_model_and_output_tags(self, fake_dirs):
|
||||
output_checkpoints_dir = fake_dirs["output"] / "checkpoints"
|
||||
output_checkpoints_dir.mkdir()
|
||||
f = output_checkpoints_dir / "saved.safetensors"
|
||||
f.touch()
|
||||
|
||||
with patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[("checkpoints", [str(output_checkpoints_dir)], {".safetensors"})],
|
||||
):
|
||||
_name, tags = get_name_and_tags_from_asset_path(str(f))
|
||||
|
||||
assert "models" in tags
|
||||
assert "model_type:checkpoints" in tags
|
||||
assert "output" in tags
|
||||
|
||||
def test_temp_path_tags_include_temp_not_output_or_preview(self, fake_dirs):
|
||||
f = fake_dirs["temp"] / "preview.png"
|
||||
f.touch()
|
||||
|
||||
_name, tags = get_name_and_tags_from_asset_path(str(f))
|
||||
|
||||
assert "temp" in tags
|
||||
assert "output" not in tags
|
||||
assert "preview:true" not in tags
|
||||
|
||||
def test_known_subfolder_tags_are_centralized(self):
|
||||
assert get_known_subfolder_tags("pasted") == ["pasted"]
|
||||
assert get_known_subfolder_tags("arbitrary") == []
|
||||
|
||||
def test_known_input_subfolder_tags_are_path_derived_for_direct_children(self, fake_dirs):
|
||||
f = fake_dirs["input"] / "pasted" / "image.png"
|
||||
f.parent.mkdir()
|
||||
f.touch()
|
||||
|
||||
assert get_known_input_subfolder_tags_from_path(str(f)) == ["pasted"]
|
||||
|
||||
_name, tags = get_name_and_tags_from_asset_path(str(f))
|
||||
assert "input" in tags
|
||||
assert "pasted" in tags
|
||||
|
||||
def test_known_input_subfolder_tags_do_not_apply_to_nested_or_other_roots(self, fake_dirs):
|
||||
nested = fake_dirs["input"] / "pasted" / "session" / "image.png"
|
||||
output = fake_dirs["output"] / "pasted" / "image.png"
|
||||
for path in (nested, output):
|
||||
path.parent.mkdir(parents=True)
|
||||
path.touch()
|
||||
|
||||
assert get_known_input_subfolder_tags_from_path(str(nested)) == []
|
||||
assert get_known_input_subfolder_tags_from_path(str(output)) == []
|
||||
|
||||
def test_unknown_path_raises(self, fake_dirs):
|
||||
with pytest.raises(ValueError, match="not within"):
|
||||
get_asset_category_and_relative_path("/some/random/path.png")
|
||||
|
||||
|
||||
class TestResponseStoragePaths:
|
||||
def test_input_file_path_and_display_name_include_subfolder(self, fake_dirs):
|
||||
sub = fake_dirs["input"] / "some" / "folder"
|
||||
sub.mkdir(parents=True)
|
||||
f = sub / "image.png"
|
||||
f.touch()
|
||||
|
||||
assert compute_logical_path(str(f)) == "input/some/folder/image.png"
|
||||
assert compute_display_name(str(f)) == "some/folder/image.png"
|
||||
|
||||
def test_output_file_path_and_display_name_include_subfolder(self, fake_dirs):
|
||||
sub = fake_dirs["output"] / "renders"
|
||||
sub.mkdir()
|
||||
f = sub / "ComfyUI_00001_.png"
|
||||
f.touch()
|
||||
|
||||
assert compute_logical_path(str(f)) == "output/renders/ComfyUI_00001_.png"
|
||||
assert compute_display_name(str(f)) == "renders/ComfyUI_00001_.png"
|
||||
|
||||
def test_temp_file_path_and_display_name(self, fake_dirs):
|
||||
f = fake_dirs["temp"] / "preview.png"
|
||||
f.touch()
|
||||
|
||||
assert compute_logical_path(str(f)) == "temp/preview.png"
|
||||
assert compute_display_name(str(f)) == "preview.png"
|
||||
|
||||
def test_exact_storage_root_has_no_display_name(self, fake_dirs):
|
||||
assert compute_logical_path(str(fake_dirs["input"])) == "input"
|
||||
assert compute_display_name(str(fake_dirs["input"])) is None
|
||||
|
||||
def test_longest_matching_builtin_root_wins(self, fake_dirs, tmp_path: Path):
|
||||
nested_output = fake_dirs["input"] / "nested-output"
|
||||
nested_output.mkdir()
|
||||
f = nested_output / "image.png"
|
||||
f.touch()
|
||||
|
||||
with patch("app.assets.services.path_utils.folder_paths") as mock_fp:
|
||||
mock_fp.get_input_directory.return_value = str(fake_dirs["input"])
|
||||
mock_fp.get_output_directory.return_value = str(nested_output)
|
||||
mock_fp.get_temp_directory.return_value = str(tmp_path / "temp")
|
||||
mock_fp.models_dir = str(fake_dirs["models_root"])
|
||||
|
||||
assert compute_logical_path(str(f)) == "output/image.png"
|
||||
assert compute_display_name(str(f)) == "image.png"
|
||||
|
||||
def test_model_file_path_is_relative_to_physical_models_root(self, fake_dirs):
|
||||
sub = fake_dirs["models"] / "flux"
|
||||
sub.mkdir()
|
||||
f = sub / "model.safetensors"
|
||||
f.touch()
|
||||
|
||||
assert compute_logical_path(str(f)) == "models/checkpoints/flux/model.safetensors"
|
||||
assert compute_display_name(str(f)) == "checkpoints/flux/model.safetensors"
|
||||
|
||||
name, tags = get_name_and_tags_from_asset_path(str(f))
|
||||
assert name == "model.safetensors"
|
||||
assert "models" in tags
|
||||
assert "model_type:checkpoints" in tags
|
||||
assert "checkpoints" not in tags
|
||||
assert "flux" not in tags
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"folder_name",
|
||||
["checkpoints", "clip", "vae", "diffusion_models", "loras"],
|
||||
)
|
||||
def test_output_model_folder_uses_output_storage_file_path(self, fake_dirs, folder_name):
|
||||
output_model_dir = fake_dirs["output"] / folder_name
|
||||
output_model_dir.mkdir(exist_ok=True)
|
||||
default_model_dir = fake_dirs["models_root"] / folder_name
|
||||
default_model_dir.mkdir(exist_ok=True)
|
||||
f = output_model_dir / "saved.safetensors"
|
||||
f.touch()
|
||||
|
||||
with patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[
|
||||
(folder_name, [str(default_model_dir), str(output_model_dir)], {".safetensors"})
|
||||
],
|
||||
):
|
||||
assert compute_logical_path(str(f)) == f"output/{folder_name}/saved.safetensors"
|
||||
assert compute_display_name(str(f)) == f"{folder_name}/saved.safetensors"
|
||||
|
||||
name, tags = get_name_and_tags_from_asset_path(str(f))
|
||||
assert name == "saved.safetensors"
|
||||
assert "output" in tags
|
||||
assert "models" in tags
|
||||
assert f"model_type:{folder_name}" in tags
|
||||
assert folder_name not in tags
|
||||
|
||||
def test_output_model_subfolder_uses_output_storage_file_path(self, fake_dirs):
|
||||
folder_name = "loras"
|
||||
output_model_dir = fake_dirs["output"] / folder_name
|
||||
subdir = output_model_dir / "experiments"
|
||||
subdir.mkdir(parents=True)
|
||||
f = subdir / "my_lora.safetensors"
|
||||
f.touch()
|
||||
|
||||
with patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[(folder_name, [str(output_model_dir)], {".safetensors"})],
|
||||
):
|
||||
assert (
|
||||
compute_logical_path(str(f))
|
||||
== "output/loras/experiments/my_lora.safetensors"
|
||||
)
|
||||
assert compute_display_name(str(f)) == "loras/experiments/my_lora.safetensors"
|
||||
|
||||
name, tags = get_name_and_tags_from_asset_path(str(f))
|
||||
assert name == "my_lora.safetensors"
|
||||
assert "output" in tags
|
||||
assert "models" in tags
|
||||
assert "model_type:loras" in tags
|
||||
assert "loras" not in tags
|
||||
assert "experiments" not in tags
|
||||
|
||||
def test_external_model_folder_without_provenance_has_no_file_path(self, tmp_path: Path):
|
||||
external_checkpoints_dir = tmp_path / "external" / "not_named_like_category"
|
||||
external_checkpoints_dir.mkdir(parents=True)
|
||||
f = external_checkpoints_dir / "external.safetensors"
|
||||
f.touch()
|
||||
|
||||
with patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[("checkpoints", [str(external_checkpoints_dir)], {".safetensors"})],
|
||||
):
|
||||
assert compute_logical_path(str(f)) is None
|
||||
assert compute_display_name(str(f)) is None
|
||||
|
||||
name, tags = get_name_and_tags_from_asset_path(str(f))
|
||||
assert name == "external.safetensors"
|
||||
assert "models" in tags
|
||||
assert "model_type:checkpoints" in tags
|
||||
|
||||
def test_same_relative_model_file_under_multiple_external_roots_has_no_storage_file_path(
|
||||
self, tmp_path: Path
|
||||
):
|
||||
foo_dir = tmp_path / "foo"
|
||||
bar_dir = tmp_path / "bar"
|
||||
foo_dir.mkdir()
|
||||
bar_dir.mkdir()
|
||||
foo_file = foo_dir / "baz.safetensors"
|
||||
bar_file = bar_dir / "baz.safetensors"
|
||||
foo_file.touch()
|
||||
bar_file.touch()
|
||||
|
||||
with patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[("checkpoints", [str(foo_dir), str(bar_dir)], {".safetensors"})],
|
||||
):
|
||||
assert compute_logical_path(str(foo_file)) is None
|
||||
assert compute_logical_path(str(bar_file)) is None
|
||||
assert compute_display_name(str(foo_file)) is None
|
||||
assert compute_display_name(str(bar_file)) is None
|
||||
|
||||
def test_output_clip_folder_uses_output_storage_and_text_encoder_tag(self, fake_dirs):
|
||||
output_clip_dir = fake_dirs["output"] / "clip"
|
||||
output_clip_dir.mkdir()
|
||||
f = output_clip_dir / "clip_l.safetensors"
|
||||
f.touch()
|
||||
|
||||
with patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[("text_encoders", [str(output_clip_dir)], {".safetensors"})],
|
||||
):
|
||||
assert compute_logical_path(str(f)) == "output/clip/clip_l.safetensors"
|
||||
assert compute_display_name(str(f)) == "clip/clip_l.safetensors"
|
||||
|
||||
name, tags = get_name_and_tags_from_asset_path(str(f))
|
||||
assert name == "clip_l.safetensors"
|
||||
assert "output" in tags
|
||||
assert "models" in tags
|
||||
assert "model_type:text_encoders" in tags
|
||||
assert "clip" not in tags
|
||||
|
||||
def test_physical_unet_folder_uses_storage_path_and_diffusion_models_tag(self, fake_dirs):
|
||||
unet_dir = fake_dirs["models_root"] / "unet"
|
||||
diffusion_models_dir = fake_dirs["models_root"] / "diffusion_models"
|
||||
unet_dir.mkdir()
|
||||
diffusion_models_dir.mkdir()
|
||||
f = unet_dir / "wan.safetensors"
|
||||
f.touch()
|
||||
|
||||
with patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[
|
||||
("diffusion_models", [str(unet_dir), str(diffusion_models_dir)], {".safetensors"})
|
||||
],
|
||||
):
|
||||
assert compute_logical_path(str(f)) == "models/unet/wan.safetensors"
|
||||
assert compute_display_name(str(f)) == "unet/wan.safetensors"
|
||||
|
||||
name, tags = get_name_and_tags_from_asset_path(str(f))
|
||||
assert name == "wan.safetensors"
|
||||
assert "models" in tags
|
||||
assert "model_type:diffusion_models" in tags
|
||||
assert "unet" not in tags
|
||||
|
||||
def test_unregistered_file_under_physical_models_root_still_has_storage_file_path(self, fake_dirs):
|
||||
f = fake_dirs["models_root"] / "not_registered" / "orphan.bin"
|
||||
f.parent.mkdir()
|
||||
f.touch()
|
||||
|
||||
assert compute_logical_path(str(f)) == "models/not_registered/orphan.bin"
|
||||
assert compute_display_name(str(f)) == "not_registered/orphan.bin"
|
||||
|
||||
def test_output_checkpoint_folder_without_registration_has_only_output_tag(self, fake_dirs):
|
||||
f = fake_dirs["output"] / "checkpoints" / "saved.safetensors"
|
||||
f.parent.mkdir(exist_ok=True)
|
||||
f.touch()
|
||||
|
||||
with patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[],
|
||||
):
|
||||
assert compute_logical_path(str(f)) == "output/checkpoints/saved.safetensors"
|
||||
assert compute_display_name(str(f)) == "checkpoints/saved.safetensors"
|
||||
|
||||
name, tags = get_name_and_tags_from_asset_path(str(f))
|
||||
assert name == "saved.safetensors"
|
||||
assert "output" in tags
|
||||
assert "models" not in tags
|
||||
assert not any(tag.startswith("model_type:") for tag in tags)
|
||||
|
||||
def test_unknown_path_returns_none(self):
|
||||
assert compute_logical_path("/some/random/path.png") is None
|
||||
assert compute_display_name("/some/random/path.png") is None
|
||||
|
||||
|
||||
class TestLoaderPath:
|
||||
"""In-root loader path: relative to the storage root, model category dropped."""
|
||||
|
||||
def test_model_loader_path_drops_category(self, fake_dirs):
|
||||
sub = fake_dirs["models"] / "flux"
|
||||
sub.mkdir()
|
||||
f = sub / "model.safetensors"
|
||||
f.touch()
|
||||
|
||||
# logical_path keeps the category, file_path (loader) drops it
|
||||
assert compute_logical_path(str(f)) == "models/checkpoints/flux/model.safetensors"
|
||||
assert compute_loader_path(str(f)) == "flux/model.safetensors"
|
||||
|
||||
def test_model_loader_path_flat_file(self, fake_dirs):
|
||||
f = fake_dirs["models"] / "model.safetensors"
|
||||
f.touch()
|
||||
|
||||
assert compute_loader_path(str(f)) == "model.safetensors"
|
||||
|
||||
def test_input_loader_path_keeps_subfolders(self, fake_dirs):
|
||||
sub = fake_dirs["input"] / "some" / "folder"
|
||||
sub.mkdir(parents=True)
|
||||
f = sub / "image.png"
|
||||
f.touch()
|
||||
|
||||
assert compute_loader_path(str(f)) == "some/folder/image.png"
|
||||
|
||||
def test_temp_loader_path(self, fake_dirs):
|
||||
f = fake_dirs["temp"] / "preview.png"
|
||||
f.touch()
|
||||
|
||||
assert compute_loader_path(str(f)) == "preview.png"
|
||||
|
||||
def test_unregistered_file_under_models_root_has_no_loader_path(self, fake_dirs):
|
||||
# Under models_root but not within any registered category base.
|
||||
f = fake_dirs["models_root"] / "not_registered" / "orphan.bin"
|
||||
f.parent.mkdir()
|
||||
f.touch()
|
||||
|
||||
# It still has a namespaced logical_path, but no loader path.
|
||||
assert compute_logical_path(str(f)) == "models/not_registered/orphan.bin"
|
||||
assert compute_loader_path(str(f)) is None
|
||||
|
||||
def test_extension_mismatch_in_registered_bucket_has_no_loader_path(self, fake_dirs):
|
||||
# Inside a registered bucket, but the bucket's extension set cannot
|
||||
# load it: no model_type tag, and no loader path either.
|
||||
f = fake_dirs["models"] / "notes.txt"
|
||||
f.touch()
|
||||
|
||||
assert compute_logical_path(str(f)) == "models/checkpoints/notes.txt"
|
||||
assert compute_loader_path(str(f)) is None
|
||||
|
||||
def test_shared_base_loader_path_uses_extension_matching_bucket(self, fake_dirs):
|
||||
shared_root = fake_dirs["models"].parent / "unet"
|
||||
shared_root.mkdir()
|
||||
f = shared_root / "wan.gguf"
|
||||
f.touch()
|
||||
|
||||
with patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[
|
||||
("diffusion_models", [str(shared_root)], {".safetensors"}),
|
||||
("unet_gguf", [str(shared_root)], {".gguf"}),
|
||||
],
|
||||
):
|
||||
assert compute_loader_path(str(f)) == "wan.gguf"
|
||||
|
||||
def test_match_all_bucket_provides_loader_path_for_any_extension(self, fake_dirs):
|
||||
custom_root = fake_dirs["models"].parent / "custom_bucket"
|
||||
custom_root.mkdir()
|
||||
f = custom_root / "weights.bin"
|
||||
f.touch()
|
||||
|
||||
with patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[("custom_bucket", [str(custom_root)], set())],
|
||||
):
|
||||
assert compute_loader_path(str(f)) == "weights.bin"
|
||||
|
||||
def test_extra_path_model_has_loader_path_but_no_logical_path(self, tmp_path: Path):
|
||||
"""Registered category base outside models_dir (extra_model_paths style).
|
||||
|
||||
Loadable, so loader_path resolves; but it is not under any canonical
|
||||
storage root, so logical_path/display_name are None. This asymmetry is
|
||||
intentional: loader_path resolves every registered model-folder base,
|
||||
logical_path only resolves the canonical storage roots.
|
||||
"""
|
||||
extra = tmp_path / "extra_ckpts"
|
||||
extra.mkdir()
|
||||
f = extra / "foo.safetensors"
|
||||
f.touch()
|
||||
|
||||
with patch("app.assets.services.path_utils.folder_paths") as mock_fp, patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[("checkpoints", [str(extra)], {".safetensors"})],
|
||||
):
|
||||
mock_fp.get_input_directory.return_value = str(tmp_path / "in")
|
||||
mock_fp.get_output_directory.return_value = str(tmp_path / "out")
|
||||
mock_fp.get_temp_directory.return_value = str(tmp_path / "tmp")
|
||||
mock_fp.models_dir = str(tmp_path / "models") # extra is NOT under this
|
||||
|
||||
assert compute_loader_path(str(f)) == "foo.safetensors"
|
||||
assert compute_logical_path(str(f)) is None
|
||||
assert compute_display_name(str(f)) is None
|
||||
|
||||
def test_unknown_path_returns_none(self):
|
||||
assert compute_loader_path("/some/random/path.png") is None
|
||||
|
||||
|
||||
class TestResolveDestinationFromTags:
|
||||
def test_extra_tags_are_not_path_components(self, fake_dirs):
|
||||
base_dir, subdirs = resolve_destination_from_tags(["input", "unit-tests", "foo"])
|
||||
|
||||
assert base_dir == os.path.abspath(fake_dirs["input"])
|
||||
assert subdirs == []
|
||||
|
||||
def test_model_upload_rejects_non_writable_registered_folders(self):
|
||||
with tempfile.TemporaryDirectory() as root:
|
||||
root_path = Path(root)
|
||||
checkpoints_dir = root_path / "models" / "checkpoints"
|
||||
configs_dir = root_path / "models" / "configs"
|
||||
custom_nodes_dir = root_path / "custom_nodes"
|
||||
for path in (checkpoints_dir, configs_dir, custom_nodes_dir):
|
||||
path.mkdir(parents=True)
|
||||
|
||||
with patch("app.assets.services.path_utils.folder_paths") as mock_fp:
|
||||
mock_fp.folder_names_and_paths = {
|
||||
"checkpoints": ([str(checkpoints_dir)], set()),
|
||||
"configs": ([str(configs_dir)], set()),
|
||||
"custom_nodes": ([str(custom_nodes_dir)], set()),
|
||||
}
|
||||
|
||||
base_dir, subdirs = resolve_destination_from_tags(
|
||||
["models", "model_type:checkpoints"]
|
||||
)
|
||||
assert base_dir == os.path.abspath(checkpoints_dir)
|
||||
assert subdirs == []
|
||||
|
||||
for folder_name in ("configs", "custom_nodes"):
|
||||
with pytest.raises(ValueError, match="unknown model category"):
|
||||
resolve_destination_from_tags(
|
||||
["models", f"model_type:{folder_name}"]
|
||||
)
|
||||
|
||||
@@ -19,8 +19,7 @@ def test_seed_asset_removed_when_file_is_deleted(
|
||||
"""Asset without hash (seed) whose file disappears:
|
||||
after triggering sync_seed_assets, Asset + AssetInfo disappear.
|
||||
"""
|
||||
# Create a file directly under input/unit-tests/<case>. Backend tags only
|
||||
# classify the root; nested path components are not exposed as tags.
|
||||
# Create a file directly under input/unit-tests/<case> so tags include "unit-tests"
|
||||
case_dir = comfy_tmp_base_dir / root / "unit-tests" / "syncseed"
|
||||
case_dir.mkdir(parents=True, exist_ok=True)
|
||||
name = f"seed_{uuid.uuid4().hex[:8]}.bin"
|
||||
@@ -33,7 +32,7 @@ def test_seed_asset_removed_when_file_is_deleted(
|
||||
# Verify it is visible via API and carries no hash (seed)
|
||||
r1 = http.get(
|
||||
api_base + "/api/assets",
|
||||
params={"include_tags": root, "name_contains": name},
|
||||
params={"include_tags": "unit-tests,syncseed", "name_contains": name},
|
||||
timeout=120,
|
||||
)
|
||||
body1 = r1.json()
|
||||
@@ -55,7 +54,7 @@ def test_seed_asset_removed_when_file_is_deleted(
|
||||
# It should disappear (AssetInfo and seed Asset gone)
|
||||
r2 = http.get(
|
||||
api_base + "/api/assets",
|
||||
params={"include_tags": root, "name_contains": name},
|
||||
params={"include_tags": "unit-tests,syncseed", "name_contains": name},
|
||||
timeout=120,
|
||||
)
|
||||
body2 = r2.json()
|
||||
@@ -133,7 +132,7 @@ def test_hashed_asset_two_asset_infos_both_get_missing(
|
||||
second_id = b2["id"]
|
||||
|
||||
# Remove the single underlying file
|
||||
p = comfy_tmp_base_dir / "input" / get_asset_filename(created["asset_hash"], ".png")
|
||||
p = comfy_tmp_base_dir / "input" / "unit-tests" / "multiinfo" / get_asset_filename(b2["asset_hash"], ".png")
|
||||
assert p.exists()
|
||||
p.unlink()
|
||||
|
||||
@@ -251,7 +250,8 @@ def test_missing_tag_clears_on_fastpass_when_mtime_and_size_match(
|
||||
|
||||
a = asset_factory(name, [root, "unit-tests", scope], {}, data)
|
||||
aid = a["id"]
|
||||
p = comfy_tmp_base_dir / root / get_asset_filename(a["asset_hash"], ".bin")
|
||||
base = comfy_tmp_base_dir / root / "unit-tests" / scope
|
||||
p = base / get_asset_filename(a["asset_hash"], ".bin")
|
||||
st0 = p.stat()
|
||||
orig_mtime_ns = getattr(st0, "st_mtime_ns", int(st0.st_mtime * 1_000_000_000))
|
||||
|
||||
|
||||
@@ -290,7 +290,7 @@ def test_metadata_filename_is_set_for_seed_asset_without_hash(
|
||||
|
||||
r1 = http.get(
|
||||
api_base + "/api/assets",
|
||||
params={"include_tags": root, "name_contains": name},
|
||||
params={"include_tags": f"unit-tests,{scope}", "name_contains": name},
|
||||
timeout=120,
|
||||
)
|
||||
body = r1.json()
|
||||
|
||||
@@ -23,7 +23,7 @@ def test_download_svg_forced_to_attachment(http: requests.Session, api_base: str
|
||||
svg = b'<svg xmlns="http://www.w3.org/2000/svg"><script>alert(1)</script></svg>'
|
||||
files = {"file": ("evil.svg", svg, "image/svg+xml")}
|
||||
form_data = {
|
||||
"tags": json.dumps(["models", "model_type:checkpoints", "unit-tests", "svgxss"]),
|
||||
"tags": json.dumps(["models", "checkpoints", "unit-tests", "svgxss"]),
|
||||
"name": "evil.svg",
|
||||
}
|
||||
up = http.post(api_base + "/api/assets", files=files, data=form_data, timeout=120)
|
||||
@@ -131,7 +131,7 @@ def test_download_chooses_existing_state_and_updates_access_time(
|
||||
assert t1 > t0
|
||||
|
||||
|
||||
@pytest.mark.parametrize("seeded_asset", [{"tags": ["models", "model_type:checkpoints"]}], indirect=True)
|
||||
@pytest.mark.parametrize("seeded_asset", [{"tags": ["models", "checkpoints"]}], indirect=True)
|
||||
def test_download_missing_file_returns_404(
|
||||
http: requests.Session, api_base: str, comfy_tmp_base_dir: Path, seeded_asset: dict
|
||||
):
|
||||
|
||||
@@ -13,7 +13,7 @@ def _seed(asset_factory, make_asset_bytes, count: int, tag: str) -> list[str]:
|
||||
for n in names:
|
||||
asset_factory(
|
||||
n,
|
||||
["models", "model_type:checkpoints", "unit-tests", tag],
|
||||
["models", "checkpoints", "unit-tests", tag],
|
||||
{},
|
||||
make_asset_bytes(n, size=2048),
|
||||
)
|
||||
@@ -208,7 +208,7 @@ def test_cursor_walks_for_non_name_sorts(sort_field, http: requests.Session, api
|
||||
names = []
|
||||
for i in range(4):
|
||||
n = f"cursor_{sort_field}_{i:02d}.safetensors"
|
||||
asset_factory(n, ["models", "model_type:checkpoints", "unit-tests", f"cursor-{sort_field}"], {}, make_asset_bytes(n, size=2048 + i))
|
||||
asset_factory(n, ["models", "checkpoints", "unit-tests", f"cursor-{sort_field}"], {}, make_asset_bytes(n, size=2048 + i))
|
||||
names.append(n)
|
||||
|
||||
params = {
|
||||
|
||||
@@ -11,7 +11,7 @@ def test_list_assets_paging_and_sort(http: requests.Session, api_base: str, asse
|
||||
for n in names:
|
||||
asset_factory(
|
||||
n,
|
||||
["models", "model_type:checkpoints", "unit-tests", "paging"],
|
||||
["models", "checkpoints", "unit-tests", "paging"],
|
||||
{"epoch": 1},
|
||||
make_asset_bytes(n, size=2048),
|
||||
)
|
||||
@@ -45,8 +45,8 @@ def test_list_assets_paging_and_sort(http: requests.Session, api_base: str, asse
|
||||
|
||||
|
||||
def test_list_assets_include_exclude_and_name_contains(http: requests.Session, api_base: str, asset_factory):
|
||||
a = asset_factory("inc_a.safetensors", ["models", "model_type:checkpoints", "unit-tests", "alpha"], {}, b"X" * 1024)
|
||||
b = asset_factory("inc_b.safetensors", ["models", "model_type:checkpoints", "unit-tests", "beta"], {}, b"Y" * 1024)
|
||||
a = asset_factory("inc_a.safetensors", ["models", "checkpoints", "unit-tests", "alpha"], {}, b"X" * 1024)
|
||||
b = asset_factory("inc_b.safetensors", ["models", "checkpoints", "unit-tests", "beta"], {}, b"Y" * 1024)
|
||||
|
||||
r = http.get(
|
||||
api_base + "/api/assets",
|
||||
@@ -81,7 +81,7 @@ def test_list_assets_include_exclude_and_name_contains(http: requests.Session, a
|
||||
|
||||
|
||||
def test_list_assets_sort_by_size_both_orders(http, api_base, asset_factory, make_asset_bytes):
|
||||
t = ["models", "model_type:checkpoints", "unit-tests", "lf-size"]
|
||||
t = ["models", "checkpoints", "unit-tests", "lf-size"]
|
||||
n1, n2, n3 = "sz1.safetensors", "sz2.safetensors", "sz3.safetensors"
|
||||
asset_factory(n1, t, {}, make_asset_bytes(n1, 1024))
|
||||
asset_factory(n2, t, {}, make_asset_bytes(n2, 2048))
|
||||
@@ -108,7 +108,7 @@ def test_list_assets_sort_by_size_both_orders(http, api_base, asset_factory, mak
|
||||
|
||||
|
||||
def test_list_assets_sort_by_updated_at_desc(http, api_base, asset_factory, make_asset_bytes):
|
||||
t = ["models", "model_type:checkpoints", "unit-tests", "lf-upd"]
|
||||
t = ["models", "checkpoints", "unit-tests", "lf-upd"]
|
||||
a1 = asset_factory("upd_a.safetensors", t, {}, make_asset_bytes("upd_a", 1200))
|
||||
a2 = asset_factory("upd_b.safetensors", t, {}, make_asset_bytes("upd_b", 1200))
|
||||
|
||||
@@ -131,7 +131,7 @@ def test_list_assets_sort_by_updated_at_desc(http, api_base, asset_factory, make
|
||||
|
||||
|
||||
def test_list_assets_sort_by_last_access_time_desc(http, api_base, asset_factory, make_asset_bytes):
|
||||
t = ["models", "model_type:checkpoints", "unit-tests", "lf-access"]
|
||||
t = ["models", "checkpoints", "unit-tests", "lf-access"]
|
||||
asset_factory("acc_a.safetensors", t, {}, make_asset_bytes("acc_a", 1100))
|
||||
time.sleep(0.02)
|
||||
a2 = asset_factory("acc_b.safetensors", t, {}, make_asset_bytes("acc_b", 1100))
|
||||
@@ -154,14 +154,14 @@ def test_list_assets_sort_by_last_access_time_desc(http, api_base, asset_factory
|
||||
|
||||
|
||||
def test_list_assets_include_tags_variants_and_case(http, api_base, asset_factory, make_asset_bytes):
|
||||
t = ["models", "model_type:checkpoints", "unit-tests", "lf-include"]
|
||||
t = ["models", "checkpoints", "unit-tests", "lf-include"]
|
||||
a = asset_factory("incvar_alpha.safetensors", [*t, "alpha"], {}, make_asset_bytes("iva"))
|
||||
asset_factory("incvar_beta.safetensors", [*t, "beta"], {}, make_asset_bytes("ivb"))
|
||||
|
||||
# CSV tag filters are whitespace-trimmed and case-sensitive.
|
||||
# CSV + case-insensitive
|
||||
r1 = http.get(
|
||||
api_base + "/api/assets",
|
||||
params={"include_tags": "unit-tests,lf-include,alpha"},
|
||||
params={"include_tags": "UNIT-TESTS,LF-INCLUDE,alpha"},
|
||||
timeout=120,
|
||||
)
|
||||
b1 = r1.json()
|
||||
@@ -196,14 +196,14 @@ def test_list_assets_include_tags_variants_and_case(http, api_base, asset_factor
|
||||
|
||||
|
||||
def test_list_assets_exclude_tags_dedup_and_case(http, api_base, asset_factory, make_asset_bytes):
|
||||
t = ["models", "model_type:checkpoints", "unit-tests", "lf-exclude"]
|
||||
t = ["models", "checkpoints", "unit-tests", "lf-exclude"]
|
||||
a = asset_factory("ex_a_alpha.safetensors", [*t, "alpha"], {}, make_asset_bytes("exa", 900))
|
||||
asset_factory("ex_b_beta.safetensors", [*t, "beta"], {}, make_asset_bytes("exb", 900))
|
||||
|
||||
# Exclude filters are case-sensitive.
|
||||
# Exclude uppercase should work
|
||||
r1 = http.get(
|
||||
api_base + "/api/assets",
|
||||
params={"include_tags": "unit-tests,lf-exclude", "exclude_tags": "beta"},
|
||||
params={"include_tags": "unit-tests,lf-exclude", "exclude_tags": "BETA"},
|
||||
timeout=120,
|
||||
)
|
||||
b1 = r1.json()
|
||||
@@ -225,7 +225,7 @@ def test_list_assets_exclude_tags_dedup_and_case(http, api_base, asset_factory,
|
||||
|
||||
|
||||
def test_list_assets_name_contains_case_and_specials(http, api_base, asset_factory, make_asset_bytes):
|
||||
t = ["models", "model_type:checkpoints", "unit-tests", "lf-name"]
|
||||
t = ["models", "checkpoints", "unit-tests", "lf-name"]
|
||||
a1 = asset_factory("CaseMix.SAFE", t, {}, make_asset_bytes("cm", 800))
|
||||
a2 = asset_factory("case-other.safetensors", t, {}, make_asset_bytes("co", 800))
|
||||
|
||||
@@ -261,7 +261,7 @@ def test_list_assets_name_contains_case_and_specials(http, api_base, asset_facto
|
||||
|
||||
|
||||
def test_list_assets_offset_beyond_total_and_limit_boundary(http, api_base, asset_factory, make_asset_bytes):
|
||||
t = ["models", "model_type:checkpoints", "unit-tests", "lf-pagelimits"]
|
||||
t = ["models", "checkpoints", "unit-tests", "lf-pagelimits"]
|
||||
asset_factory("pl1.safetensors", t, {}, make_asset_bytes("pl1", 600))
|
||||
asset_factory("pl2.safetensors", t, {}, make_asset_bytes("pl2", 600))
|
||||
asset_factory("pl3.safetensors", t, {}, make_asset_bytes("pl3", 600))
|
||||
@@ -319,7 +319,7 @@ def test_list_assets_name_contains_literal_underscore(
|
||||
- foobar.safetensors (must NOT match)
|
||||
"""
|
||||
scope = f"lf-underscore-{uuid.uuid4().hex[:6]}"
|
||||
tags = ["models", "model_type:checkpoints", "unit-tests", scope]
|
||||
tags = ["models", "checkpoints", "unit-tests", scope]
|
||||
|
||||
a = asset_factory("foo_bar.safetensors", tags, {}, make_asset_bytes("a", 700))
|
||||
b = asset_factory("fooxbar.safetensors", tags, {}, make_asset_bytes("b", 700))
|
||||
|
||||
@@ -5,7 +5,7 @@ def test_meta_and_across_keys_and_types(
|
||||
http, api_base: str, asset_factory, make_asset_bytes
|
||||
):
|
||||
name = "mf_and_mix.safetensors"
|
||||
tags = ["models", "model_type:checkpoints", "unit-tests", "mf-and"]
|
||||
tags = ["models", "checkpoints", "unit-tests", "mf-and"]
|
||||
meta = {"purpose": "mix", "epoch": 1, "active": True, "score": 1.23}
|
||||
asset_factory(name, tags, meta, make_asset_bytes(name, 4096))
|
||||
|
||||
@@ -41,7 +41,7 @@ def test_meta_and_across_keys_and_types(
|
||||
|
||||
def test_meta_type_strictness_int_vs_str_and_bool(http, api_base, asset_factory, make_asset_bytes):
|
||||
name = "mf_types.safetensors"
|
||||
tags = ["models", "model_type:checkpoints", "unit-tests", "mf-types"]
|
||||
tags = ["models", "checkpoints", "unit-tests", "mf-types"]
|
||||
meta = {"epoch": 1, "active": True}
|
||||
asset_factory(name, tags, meta, make_asset_bytes(name))
|
||||
|
||||
@@ -95,7 +95,7 @@ def test_meta_type_strictness_int_vs_str_and_bool(http, api_base, asset_factory,
|
||||
|
||||
def test_meta_any_of_list_of_scalars(http, api_base, asset_factory, make_asset_bytes):
|
||||
name = "mf_list_scalars.safetensors"
|
||||
tags = ["models", "model_type:checkpoints", "unit-tests", "mf-list"]
|
||||
tags = ["models", "checkpoints", "unit-tests", "mf-list"]
|
||||
meta = {"flags": ["red", "green"]}
|
||||
asset_factory(name, tags, meta, make_asset_bytes(name, 3000))
|
||||
|
||||
@@ -134,7 +134,7 @@ def test_meta_none_semantics_missing_or_null_and_any_of_with_none(
|
||||
http, api_base, asset_factory, make_asset_bytes
|
||||
):
|
||||
# a1: key missing; a2: explicit null; a3: concrete value
|
||||
t = ["models", "model_type:checkpoints", "unit-tests", "mf-none"]
|
||||
t = ["models", "checkpoints", "unit-tests", "mf-none"]
|
||||
a1 = asset_factory("mf_none_missing.safetensors", t, {"x": 1}, make_asset_bytes("a1"))
|
||||
a2 = asset_factory("mf_none_null.safetensors", t, {"maybe": None}, make_asset_bytes("a2"))
|
||||
a3 = asset_factory("mf_none_value.safetensors", t, {"maybe": "x"}, make_asset_bytes("a3"))
|
||||
@@ -166,7 +166,7 @@ def test_meta_none_semantics_missing_or_null_and_any_of_with_none(
|
||||
|
||||
def test_meta_nested_json_object_equality(http, api_base, asset_factory, make_asset_bytes):
|
||||
name = "mf_nested_json.safetensors"
|
||||
tags = ["models", "model_type:checkpoints", "unit-tests", "mf-nested"]
|
||||
tags = ["models", "checkpoints", "unit-tests", "mf-nested"]
|
||||
cfg = {"optimizer": "adam", "lr": 0.001, "schedule": {"type": "cosine", "warmup": 100}}
|
||||
asset_factory(name, tags, {"config": cfg}, make_asset_bytes(name, 2200))
|
||||
|
||||
@@ -197,7 +197,7 @@ def test_meta_nested_json_object_equality(http, api_base, asset_factory, make_as
|
||||
|
||||
def test_meta_list_of_objects_any_of(http, api_base, asset_factory, make_asset_bytes):
|
||||
name = "mf_list_objects.safetensors"
|
||||
tags = ["models", "model_type:checkpoints", "unit-tests", "mf-objlist"]
|
||||
tags = ["models", "checkpoints", "unit-tests", "mf-objlist"]
|
||||
transforms = [{"type": "crop", "size": 128}, {"type": "flip", "p": 0.5}]
|
||||
asset_factory(name, tags, {"transforms": transforms}, make_asset_bytes(name, 2048))
|
||||
|
||||
@@ -228,7 +228,7 @@ def test_meta_list_of_objects_any_of(http, api_base, asset_factory, make_asset_b
|
||||
|
||||
def test_meta_with_special_and_unicode_keys(http, api_base, asset_factory, make_asset_bytes):
|
||||
name = "mf_keys_unicode.safetensors"
|
||||
tags = ["models", "model_type:checkpoints", "unit-tests", "mf-keys"]
|
||||
tags = ["models", "checkpoints", "unit-tests", "mf-keys"]
|
||||
meta = {
|
||||
"weird.key": "v1",
|
||||
"path/like": 7,
|
||||
@@ -259,7 +259,7 @@ def test_meta_with_special_and_unicode_keys(http, api_base, asset_factory, make_
|
||||
|
||||
|
||||
def test_meta_with_zero_and_boolean_lists(http, api_base, asset_factory, make_asset_bytes):
|
||||
t = ["models", "model_type:checkpoints", "unit-tests", "mf-zero-bool"]
|
||||
t = ["models", "checkpoints", "unit-tests", "mf-zero-bool"]
|
||||
a0 = asset_factory("mf_zero_count.safetensors", t, {"count": 0}, make_asset_bytes("z", 1025))
|
||||
a1 = asset_factory("mf_bool_list.safetensors", t, {"choices": [True, False]}, make_asset_bytes("b", 1026))
|
||||
|
||||
@@ -286,7 +286,7 @@ def test_meta_with_zero_and_boolean_lists(http, api_base, asset_factory, make_as
|
||||
|
||||
def test_meta_mixed_list_types_and_strictness(http, api_base, asset_factory, make_asset_bytes):
|
||||
name = "mf_mixed_list.safetensors"
|
||||
tags = ["models", "model_type:checkpoints", "unit-tests", "mf-mixed"]
|
||||
tags = ["models", "checkpoints", "unit-tests", "mf-mixed"]
|
||||
meta = {"mix": ["1", 1, True, None]}
|
||||
asset_factory(name, tags, meta, make_asset_bytes(name, 1999))
|
||||
|
||||
@@ -311,7 +311,7 @@ def test_meta_mixed_list_types_and_strictness(http, api_base, asset_factory, mak
|
||||
|
||||
def test_meta_unknown_key_and_none_behavior_with_scope_tags(http, api_base, asset_factory, make_asset_bytes):
|
||||
# Use a unique scope tag to avoid interference
|
||||
t = ["models", "model_type:checkpoints", "unit-tests", "mf-unknown-scope"]
|
||||
t = ["models", "checkpoints", "unit-tests", "mf-unknown-scope"]
|
||||
x = asset_factory("mf_unknown_a.safetensors", t, {"k1": 1}, make_asset_bytes("ua"))
|
||||
y = asset_factory("mf_unknown_b.safetensors", t, {"k2": 2}, make_asset_bytes("ub"))
|
||||
|
||||
@@ -340,13 +340,13 @@ def test_meta_with_tags_include_exclude_and_name_contains(http, api_base, asset_
|
||||
# alpha matches epoch=1; beta has epoch=2
|
||||
a = asset_factory(
|
||||
"mf_tag_alpha.safetensors",
|
||||
["models", "model_type:checkpoints", "unit-tests", "mf-tag", "alpha"],
|
||||
["models", "checkpoints", "unit-tests", "mf-tag", "alpha"],
|
||||
{"epoch": 1},
|
||||
make_asset_bytes("alpha"),
|
||||
)
|
||||
b = asset_factory(
|
||||
"mf_tag_beta.safetensors",
|
||||
["models", "model_type:checkpoints", "unit-tests", "mf-tag", "beta"],
|
||||
["models", "checkpoints", "unit-tests", "mf-tag", "beta"],
|
||||
{"epoch": 2},
|
||||
make_asset_bytes("beta"),
|
||||
)
|
||||
@@ -367,7 +367,7 @@ def test_meta_with_tags_include_exclude_and_name_contains(http, api_base, asset_
|
||||
|
||||
def test_meta_sort_and_paging_under_filter(http, api_base, asset_factory, make_asset_bytes):
|
||||
# Three assets in same scope with different sizes and a common filter key
|
||||
t = ["models", "model_type:checkpoints", "unit-tests", "mf-sort"]
|
||||
t = ["models", "checkpoints", "unit-tests", "mf-sort"]
|
||||
n1, n2, n3 = "mf_sort_1.safetensors", "mf_sort_2.safetensors", "mf_sort_3.safetensors"
|
||||
asset_factory(n1, t, {"group": "g"}, make_asset_bytes(n1, 1024))
|
||||
asset_factory(n2, t, {"group": "g"}, make_asset_bytes(n2, 2048))
|
||||
|
||||
@@ -29,7 +29,7 @@ def create_seed_file(comfy_tmp_base_dir: Path):
|
||||
def find_asset(http: requests.Session, api_base: str):
|
||||
"""Query API for assets matching scope and optional name."""
|
||||
def _find(scope: str, name: str | None = None) -> list[dict]:
|
||||
params = {"limit": "500"}
|
||||
params = {"include_tags": f"unit-tests,{scope}"}
|
||||
if name:
|
||||
params["name_contains"] = name
|
||||
r = http.get(f"{api_base}/api/assets", params=params, timeout=120)
|
||||
@@ -91,7 +91,7 @@ def test_hashed_asset_not_pruned_when_file_missing(
|
||||
data = make_asset_bytes("test", 2048)
|
||||
a = asset_factory("test.bin", ["input", "unit-tests", scope], {}, data)
|
||||
|
||||
path = comfy_tmp_base_dir / "input" / get_asset_filename(a["asset_hash"], ".bin")
|
||||
path = comfy_tmp_base_dir / "input" / "unit-tests" / scope / get_asset_filename(a["asset_hash"], ".bin")
|
||||
path.unlink()
|
||||
|
||||
trigger_sync_seed_assets(http, api_base)
|
||||
@@ -108,20 +108,18 @@ def test_prune_across_multiple_roots(
|
||||
):
|
||||
"""Prune correctly handles assets across input and output roots."""
|
||||
scope = f"multi-{uuid.uuid4().hex[:6]}"
|
||||
input_name = f"{scope}-input.bin"
|
||||
output_name = f"{scope}-output.bin"
|
||||
input_fp = create_seed_file("input", scope, input_name)
|
||||
create_seed_file("output", scope, output_name)
|
||||
input_fp = create_seed_file("input", scope, "input.bin")
|
||||
create_seed_file("output", scope, "output.bin")
|
||||
|
||||
trigger_sync_seed_assets(http, api_base)
|
||||
assert find_asset(scope, input_name)
|
||||
assert find_asset(scope, output_name)
|
||||
assert len(find_asset(scope)) == 2
|
||||
|
||||
input_fp.unlink()
|
||||
trigger_sync_seed_assets(http, api_base)
|
||||
|
||||
assert not find_asset(scope, input_name)
|
||||
assert find_asset(scope, output_name)
|
||||
remaining = find_asset(scope)
|
||||
assert len(remaining) == 1
|
||||
assert remaining[0]["name"] == "output.bin"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("dirname", ["100%_done", "my_folder_name", "has spaces"])
|
||||
|
||||
@@ -10,9 +10,9 @@ def test_tags_present(http: requests.Session, api_base: str, seeded_asset: dict)
|
||||
body1 = r1.json()
|
||||
assert r1.status_code == 200
|
||||
names = [t["name"] for t in body1["tags"]]
|
||||
# A few selected contract tags should exist.
|
||||
# A few system tags from migration should exist:
|
||||
assert "models" in names
|
||||
assert "model_type:checkpoints" in names
|
||||
assert "checkpoints" in names
|
||||
|
||||
# Only used tags before we add anything new from this test cycle
|
||||
r2 = http.get(api_base + "/api/tags", params={"include_zero": "false"}, timeout=120)
|
||||
@@ -21,7 +21,7 @@ def test_tags_present(http: requests.Session, api_base: str, seeded_asset: dict)
|
||||
# We already seeded one asset via fixture, so used tags must be non-empty
|
||||
used_names = [t["name"] for t in body2["tags"]]
|
||||
assert "models" in used_names
|
||||
assert "model_type:checkpoints" in used_names
|
||||
assert "checkpoints" in used_names
|
||||
|
||||
# Prefix filter should refine the list
|
||||
r3 = http.get(api_base + "/api/tags", params={"include_zero": "false", "prefix": "uni"}, timeout=120)
|
||||
@@ -45,7 +45,7 @@ def test_tags_empty_usage(http: requests.Session, api_base: str, asset_factory,
|
||||
body1 = r1.json()
|
||||
assert r1.status_code == 200
|
||||
names = [t["name"] for t in body1["tags"]]
|
||||
assert "models" in names and "model_type:checkpoints" in names
|
||||
assert "models" in names and "checkpoints" in names
|
||||
|
||||
# Create a short-lived asset under input with a unique custom tag
|
||||
scope = f"tags-empty-usage-{uuid.uuid4().hex[:6]}"
|
||||
@@ -89,28 +89,28 @@ def test_tags_empty_usage(http: requests.Session, api_base: str, asset_factory,
|
||||
def test_add_and_remove_tags(http: requests.Session, api_base: str, seeded_asset: dict):
|
||||
aid = seeded_asset["id"]
|
||||
|
||||
# Add tags with duplicates while preserving source case.
|
||||
payload_add = {"tags": ["NewTag", "unit-tests", "NewTag", "BETA"]}
|
||||
# Add tags with duplicates and mixed case
|
||||
payload_add = {"tags": ["NewTag", "unit-tests", "newtag", "BETA"]}
|
||||
r1 = http.post(f"{api_base}/api/assets/{aid}/tags", json=payload_add, timeout=120)
|
||||
b1 = r1.json()
|
||||
assert r1.status_code == 200, b1
|
||||
# stripped, deduplicated; 'unit-tests' was already present from the seed
|
||||
assert set(b1["added"]) == {"NewTag", "BETA"}
|
||||
# normalized, deduplicated; 'unit-tests' was already present from the seed
|
||||
assert set(b1["added"]) == {"newtag", "beta"}
|
||||
assert set(b1["already_present"]) == {"unit-tests"}
|
||||
assert "NewTag" in b1["total_tags"] and "BETA" in b1["total_tags"]
|
||||
assert "newtag" in b1["total_tags"] and "beta" in b1["total_tags"]
|
||||
|
||||
rg = http.get(f"{api_base}/api/assets/{aid}", timeout=120)
|
||||
g = rg.json()
|
||||
assert rg.status_code == 200
|
||||
tags_now = set(g["tags"])
|
||||
assert {"NewTag", "BETA"}.issubset(tags_now)
|
||||
assert {"newtag", "beta"}.issubset(tags_now)
|
||||
|
||||
# Remove a tag and a non-existent tag
|
||||
payload_del = {"tags": ["NewTag", "does-not-exist"]}
|
||||
payload_del = {"tags": ["newtag", "does-not-exist"]}
|
||||
r2 = http.delete(f"{api_base}/api/assets/{aid}/tags", json=payload_del, timeout=120)
|
||||
b2 = r2.json()
|
||||
assert r2.status_code == 200
|
||||
assert set(b2["removed"]) == {"NewTag"}
|
||||
assert set(b2["removed"]) == {"newtag"}
|
||||
assert set(b2["not_present"]) == {"does-not-exist"}
|
||||
|
||||
# Verify remaining tags after deletion
|
||||
@@ -118,44 +118,8 @@ def test_add_and_remove_tags(http: requests.Session, api_base: str, seeded_asset
|
||||
g2 = rg2.json()
|
||||
assert rg2.status_code == 200
|
||||
tags_later = set(g2["tags"])
|
||||
assert "NewTag" not in tags_later
|
||||
assert "BETA" in tags_later # still present
|
||||
|
||||
|
||||
def test_add_system_looking_tags_allowed_as_labels(
|
||||
http: requests.Session, api_base: str, seeded_asset: dict
|
||||
):
|
||||
aid = seeded_asset["id"]
|
||||
|
||||
response = http.post(
|
||||
f"{api_base}/api/assets/{aid}/tags",
|
||||
json={
|
||||
"tags": [
|
||||
"models",
|
||||
"model_type:manual",
|
||||
"model:true",
|
||||
"models:foo",
|
||||
"input:true",
|
||||
"output:true",
|
||||
"uploaded:true",
|
||||
"temp:true",
|
||||
"temporary",
|
||||
]
|
||||
},
|
||||
timeout=120,
|
||||
)
|
||||
body = response.json()
|
||||
|
||||
assert response.status_code == 200, body
|
||||
assert "models" in body["total_tags"]
|
||||
assert "model_type:manual" in body["total_tags"]
|
||||
assert "model:true" in body["total_tags"]
|
||||
assert "models:foo" in body["total_tags"]
|
||||
assert "input:true" in body["total_tags"]
|
||||
assert "output:true" in body["total_tags"]
|
||||
assert "uploaded:true" in body["total_tags"]
|
||||
assert "temp:true" in body["total_tags"]
|
||||
assert "temporary" in body["total_tags"]
|
||||
assert "newtag" not in tags_later
|
||||
assert "beta" in tags_later # still present
|
||||
|
||||
|
||||
def test_tags_list_order_and_prefix(http: requests.Session, api_base: str, seeded_asset: dict):
|
||||
|
||||
@@ -1,14 +1,11 @@
|
||||
import json
|
||||
import uuid
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from pathlib import Path
|
||||
|
||||
import requests
|
||||
import pytest
|
||||
|
||||
from app.assets.api.schemas_in import UploadAssetSpec
|
||||
from app.assets.api.schemas_out import Asset, AssetCreated
|
||||
from helpers import get_asset_filename
|
||||
|
||||
|
||||
def test_asset_created_inherits_hash_field():
|
||||
@@ -23,18 +20,9 @@ def test_asset_created_inherits_hash_field():
|
||||
assert AssetCreated.model_fields["hash"].annotation == Asset.model_fields["hash"].annotation
|
||||
|
||||
|
||||
def test_upload_asset_spec_ignores_subfolder_field():
|
||||
spec = UploadAssetSpec.model_validate(
|
||||
{"tags": ["input"], "subfolder": "pasted", "name": "image.png"}
|
||||
)
|
||||
|
||||
assert "subfolder" not in UploadAssetSpec.model_fields
|
||||
assert not hasattr(spec, "subfolder")
|
||||
|
||||
|
||||
def test_upload_ok_duplicate_reference(http: requests.Session, api_base: str, make_asset_bytes):
|
||||
name = "dup_a.safetensors"
|
||||
tags = ["models", "model_type:checkpoints", "unit-tests", "alpha"]
|
||||
tags = ["models", "checkpoints", "unit-tests", "alpha"]
|
||||
meta = {"purpose": "dup"}
|
||||
data = make_asset_bytes(name)
|
||||
files = {"file": (name, data, "application/octet-stream")}
|
||||
@@ -55,8 +43,6 @@ def test_upload_ok_duplicate_reference(http: requests.Session, api_base: str, ma
|
||||
assert a2["asset_hash"] == a1["asset_hash"]
|
||||
assert a2["hash"] == a1["hash"]
|
||||
assert a2["id"] != a1["id"] # new reference with same content
|
||||
assert a2.get("loader_path") is None
|
||||
assert a2.get("display_name") is None
|
||||
|
||||
# Third upload with the same data but different name also creates new AssetReference
|
||||
files = {"file": (name, data, "application/octet-stream")}
|
||||
@@ -67,14 +53,12 @@ def test_upload_ok_duplicate_reference(http: requests.Session, api_base: str, ma
|
||||
assert a3["asset_hash"] == a1["asset_hash"]
|
||||
assert a3["id"] != a1["id"]
|
||||
assert a3["id"] != a2["id"]
|
||||
assert a3.get("loader_path") is None
|
||||
assert a3.get("display_name") is None
|
||||
|
||||
|
||||
def test_upload_fastpath_from_existing_hash_no_file(http: requests.Session, api_base: str):
|
||||
# Seed a small file first
|
||||
name = "fastpath_seed.safetensors"
|
||||
tags = ["input", "unit-tests"]
|
||||
tags = ["models", "checkpoints", "unit-tests"]
|
||||
meta = {}
|
||||
files = {"file": (name, b"B" * 1024, "application/octet-stream")}
|
||||
form = {"tags": json.dumps(tags), "name": name, "user_metadata": json.dumps(meta)}
|
||||
@@ -85,10 +69,9 @@ def test_upload_fastpath_from_existing_hash_no_file(http: requests.Session, api_
|
||||
assert b1["hash"] == h
|
||||
|
||||
# Now POST /api/assets with only hash and no file
|
||||
hash_only_tags = ["models", "checkpoints", "unit-tests", "hash-labels"]
|
||||
files = [
|
||||
("hash", (None, h)),
|
||||
("tags", (None, json.dumps(hash_only_tags))),
|
||||
("tags", (None, json.dumps(tags))),
|
||||
("name", (None, "fastpath_copy.safetensors")),
|
||||
("user_metadata", (None, json.dumps({"purpose": "copy"}))),
|
||||
]
|
||||
@@ -98,53 +81,6 @@ def test_upload_fastpath_from_existing_hash_no_file(http: requests.Session, api_
|
||||
assert b2["created_new"] is False
|
||||
assert b2["asset_hash"] == h
|
||||
assert b2["hash"] == h
|
||||
assert "models" in b2["tags"]
|
||||
assert "checkpoints" in b2["tags"]
|
||||
assert "uploaded" not in b2["tags"]
|
||||
assert not any(tag.startswith("model_type:") for tag in b2["tags"])
|
||||
assert b2.get("loader_path") is None
|
||||
assert b2.get("display_name") is None
|
||||
|
||||
rg = http.get(f"{api_base}/api/assets/{b2['id']}", timeout=120)
|
||||
detail = rg.json()
|
||||
assert rg.status_code == 200, detail
|
||||
assert detail.get("loader_path") is None
|
||||
assert detail.get("display_name") is None
|
||||
|
||||
|
||||
def test_create_from_hash_with_model_tags_does_not_synthesize_loader_path(
|
||||
http: requests.Session, api_base: str
|
||||
):
|
||||
seed_name = "from_hash_seed.safetensors"
|
||||
seed_tags = ["models", "model_type:checkpoints", "unit-tests"]
|
||||
files = {"file": (seed_name, b"D" * 1024, "application/octet-stream")}
|
||||
form = {
|
||||
"tags": json.dumps(seed_tags),
|
||||
"name": seed_name,
|
||||
"user_metadata": json.dumps({}),
|
||||
}
|
||||
seed_r = http.post(api_base + "/api/assets", data=form, files=files, timeout=120)
|
||||
seed = seed_r.json()
|
||||
assert seed_r.status_code == 201, seed
|
||||
|
||||
payload = {
|
||||
"hash": seed["asset_hash"],
|
||||
"name": "from_hash_copy.safetensors",
|
||||
"tags": ["models", "model_type:checkpoints", "unit-tests", "spoofed"],
|
||||
}
|
||||
created_r = http.post(api_base + "/api/assets/from-hash", json=payload, timeout=120)
|
||||
created = created_r.json()
|
||||
assert created_r.status_code == 201, created
|
||||
assert created["created_new"] is False
|
||||
assert created["asset_hash"] == seed["asset_hash"]
|
||||
assert created.get("loader_path") is None
|
||||
assert created.get("display_name") is None
|
||||
|
||||
detail_r = http.get(f"{api_base}/api/assets/{created['id']}", timeout=120)
|
||||
detail = detail_r.json()
|
||||
assert detail_r.status_code == 200, detail
|
||||
assert detail.get("loader_path") is None
|
||||
assert detail.get("display_name") is None
|
||||
|
||||
|
||||
def test_upload_fastpath_with_known_hash_and_file(
|
||||
@@ -152,7 +88,7 @@ def test_upload_fastpath_with_known_hash_and_file(
|
||||
):
|
||||
# Seed
|
||||
files = {"file": ("seed.safetensors", b"C" * 128, "application/octet-stream")}
|
||||
form = {"tags": json.dumps(["models", "model_type:checkpoints", "unit-tests", "fp"]), "name": "seed.safetensors", "user_metadata": json.dumps({})}
|
||||
form = {"tags": json.dumps(["models", "checkpoints", "unit-tests", "fp"]), "name": "seed.safetensors", "user_metadata": json.dumps({})}
|
||||
r1 = http.post(api_base + "/api/assets", data=form, files=files, timeout=120)
|
||||
b1 = r1.json()
|
||||
assert r1.status_code == 201, b1
|
||||
@@ -168,49 +104,11 @@ def test_upload_fastpath_with_known_hash_and_file(
|
||||
assert b2["created_new"] is False
|
||||
assert b2["asset_hash"] == h
|
||||
assert b2["hash"] == h
|
||||
assert "checkpoints" in b2["tags"]
|
||||
assert "uploaded" not in b2["tags"]
|
||||
assert not any(tag == "model_type:checkpoints" for tag in b2["tags"])
|
||||
|
||||
|
||||
def test_duplicate_byte_upload_is_reference_only_and_does_not_need_destination(
|
||||
http: requests.Session, api_base: str
|
||||
):
|
||||
data = b"duplicate-reference-only" * 64
|
||||
seed_files = {"file": ("duplicate-seed.bin", data, "application/octet-stream")}
|
||||
seed_form = {
|
||||
"tags": json.dumps(["input", "unit-tests", "duplicate-seed"]),
|
||||
"name": "duplicate-seed.bin",
|
||||
"user_metadata": json.dumps({}),
|
||||
}
|
||||
seed_response = http.post(api_base + "/api/assets", data=seed_form, files=seed_files, timeout=120)
|
||||
seed = seed_response.json()
|
||||
assert seed_response.status_code == 201, seed
|
||||
|
||||
duplicate_files = {"file": ("duplicate-copy.bin", data, "application/octet-stream")}
|
||||
duplicate_form = {
|
||||
"tags": json.dumps(["not-a-destination", "unit-tests", "duplicate-copy"]),
|
||||
"name": "duplicate-copy.bin",
|
||||
"user_metadata": json.dumps({}),
|
||||
}
|
||||
duplicate_response = http.post(
|
||||
api_base + "/api/assets", data=duplicate_form, files=duplicate_files, timeout=120
|
||||
)
|
||||
duplicate = duplicate_response.json()
|
||||
|
||||
assert duplicate_response.status_code == 200, duplicate
|
||||
assert duplicate["created_new"] is False
|
||||
assert duplicate["asset_hash"] == seed["asset_hash"]
|
||||
assert "not-a-destination" in duplicate["tags"]
|
||||
assert "uploaded" not in duplicate["tags"]
|
||||
assert "input" not in duplicate["tags"]
|
||||
assert duplicate.get("loader_path") is None
|
||||
assert duplicate.get("display_name") is None
|
||||
|
||||
|
||||
def test_upload_multiple_tags_fields_are_merged(http: requests.Session, api_base: str):
|
||||
data = [
|
||||
("tags", "models,model_type:checkpoints"),
|
||||
("tags", "models,checkpoints"),
|
||||
("tags", json.dumps(["unit-tests", "alpha"])),
|
||||
("name", "merge.safetensors"),
|
||||
("user_metadata", json.dumps({"u": 1})),
|
||||
@@ -226,71 +124,7 @@ def test_upload_multiple_tags_fields_are_merged(http: requests.Session, api_base
|
||||
detail = rg.json()
|
||||
assert rg.status_code == 200, detail
|
||||
tags = set(detail["tags"])
|
||||
assert {"models", "model_type:checkpoints", "unit-tests", "alpha"}.issubset(tags)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
(
|
||||
"tags",
|
||||
"extension",
|
||||
"expected_display_prefix",
|
||||
),
|
||||
[
|
||||
(["input", "unit-tests"], ".png", ""),
|
||||
(
|
||||
["models", "model_type:checkpoints", "unit-tests"],
|
||||
".safetensors",
|
||||
"checkpoints/",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_upload_response_includes_loader_path_and_display_name(
|
||||
tags: list[str],
|
||||
extension: str,
|
||||
expected_display_prefix: str,
|
||||
http: requests.Session,
|
||||
api_base: str,
|
||||
make_asset_bytes,
|
||||
):
|
||||
scope = f"response-paths-{uuid.uuid4().hex[:6]}"
|
||||
scoped_tags = [*tags, scope]
|
||||
name = f"asset_response_path{extension}"
|
||||
|
||||
files = {"file": (name, make_asset_bytes(name, 1024), "application/octet-stream")}
|
||||
form = {
|
||||
"tags": json.dumps(scoped_tags),
|
||||
"name": name,
|
||||
"user_metadata": json.dumps({}),
|
||||
}
|
||||
created_r = http.post(api_base + "/api/assets", data=form, files=files, timeout=120)
|
||||
created = created_r.json()
|
||||
assert created_r.status_code in (200, 201), created
|
||||
stored_filename = get_asset_filename(created["asset_hash"], extension)
|
||||
expected_suffix = stored_filename
|
||||
expected_display_name = f"{expected_display_prefix}{expected_suffix}"
|
||||
# In-root loader path: model category dropped, no subfolders here -> just the filename.
|
||||
expected_loader_path = expected_suffix
|
||||
|
||||
assert created["loader_path"] == expected_loader_path
|
||||
assert created["display_name"] == expected_display_name
|
||||
assert "logical_path" not in created
|
||||
|
||||
detail_r = http.get(f"{api_base}/api/assets/{created['id']}", timeout=120)
|
||||
detail = detail_r.json()
|
||||
assert detail_r.status_code == 200, detail
|
||||
assert detail["loader_path"] == expected_loader_path
|
||||
assert detail["display_name"] == expected_display_name
|
||||
|
||||
list_r = http.get(
|
||||
api_base + "/api/assets",
|
||||
params={"include_tags": f"unit-tests,{scope}", "limit": "50"},
|
||||
timeout=120,
|
||||
)
|
||||
listed = list_r.json()
|
||||
assert list_r.status_code == 200, listed
|
||||
match = next(a for a in listed["assets"] if a["id"] == created["id"])
|
||||
assert match["loader_path"] == expected_loader_path
|
||||
assert match["display_name"] == expected_display_name
|
||||
assert {"models", "checkpoints", "unit-tests", "alpha"}.issubset(tags)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("root", ["input", "output"])
|
||||
@@ -358,55 +192,16 @@ def test_create_from_hash_endpoint_404(http: requests.Session, api_base: str):
|
||||
assert body["error"]["code"] == "ASSET_NOT_FOUND"
|
||||
|
||||
|
||||
def test_create_from_hash_accepts_arbitrary_system_looking_tags(
|
||||
http: requests.Session, api_base: str
|
||||
):
|
||||
files = {"file": ("hash-seed.bin", b"hash-seed" * 64, "application/octet-stream")}
|
||||
form = {
|
||||
"tags": json.dumps(["input", "unit-tests", "hash-seed"]),
|
||||
"name": "hash-seed.bin",
|
||||
"user_metadata": json.dumps({}),
|
||||
}
|
||||
seed_response = http.post(api_base + "/api/assets", data=form, files=files, timeout=120)
|
||||
seed = seed_response.json()
|
||||
assert seed_response.status_code == 201, seed
|
||||
|
||||
response = http.post(
|
||||
api_base + "/api/assets/from-hash",
|
||||
json={
|
||||
"hash": seed["asset_hash"],
|
||||
"name": "hash-copy.bin",
|
||||
"tags": [
|
||||
"models",
|
||||
"model:true",
|
||||
"models:foo",
|
||||
"temporary:true",
|
||||
"unit-tests",
|
||||
"hash-copy",
|
||||
],
|
||||
},
|
||||
timeout=120,
|
||||
)
|
||||
body = response.json()
|
||||
|
||||
assert response.status_code == 201, body
|
||||
assert "models" in body["tags"]
|
||||
assert "model:true" in body["tags"]
|
||||
assert "models:foo" in body["tags"]
|
||||
assert "temporary:true" in body["tags"]
|
||||
assert "uploaded" not in body["tags"]
|
||||
|
||||
|
||||
def test_upload_zero_byte_rejected(http: requests.Session, api_base: str):
|
||||
files = {"file": ("empty.safetensors", b"", "application/octet-stream")}
|
||||
form = {"tags": json.dumps(["models", "model_type:checkpoints", "unit-tests", "edge"]), "name": "empty.safetensors", "user_metadata": json.dumps({})}
|
||||
form = {"tags": json.dumps(["models", "checkpoints", "unit-tests", "edge"]), "name": "empty.safetensors", "user_metadata": json.dumps({})}
|
||||
r = http.post(api_base + "/api/assets", data=form, files=files, timeout=120)
|
||||
body = r.json()
|
||||
assert r.status_code == 400
|
||||
assert body["error"]["code"] == "EMPTY_UPLOAD"
|
||||
|
||||
|
||||
def test_upload_rejects_arbitrary_labels_without_required_destination_role(http: requests.Session, api_base: str):
|
||||
def test_upload_invalid_root_tag_rejected(http: requests.Session, api_base: str):
|
||||
files = {"file": ("badroot.bin", b"A" * 64, "application/octet-stream")}
|
||||
form = {"tags": json.dumps(["not-a-root", "whatever"]), "name": "badroot.bin", "user_metadata": json.dumps({})}
|
||||
r = http.post(api_base + "/api/assets", data=form, files=files, timeout=120)
|
||||
@@ -417,7 +212,7 @@ def test_upload_rejects_arbitrary_labels_without_required_destination_role(http:
|
||||
|
||||
def test_upload_user_metadata_must_be_json(http: requests.Session, api_base: str):
|
||||
files = {"file": ("badmeta.bin", b"A" * 128, "application/octet-stream")}
|
||||
form = {"tags": json.dumps(["models", "model_type:checkpoints", "unit-tests", "edge"]), "name": "badmeta.bin", "user_metadata": "{not json}"}
|
||||
form = {"tags": json.dumps(["models", "checkpoints", "unit-tests", "edge"]), "name": "badmeta.bin", "user_metadata": "{not json}"}
|
||||
r = http.post(api_base + "/api/assets", data=form, files=files, timeout=120)
|
||||
body = r.json()
|
||||
assert r.status_code == 400
|
||||
@@ -433,7 +228,7 @@ def test_upload_requires_multipart(http: requests.Session, api_base: str):
|
||||
|
||||
def test_upload_missing_file_and_hash(http: requests.Session, api_base: str):
|
||||
files = [
|
||||
("tags", (None, json.dumps(["models", "model_type:checkpoints", "unit-tests"]))),
|
||||
("tags", (None, json.dumps(["models", "checkpoints", "unit-tests"]))),
|
||||
("name", (None, "x.safetensors")),
|
||||
]
|
||||
r = http.post(api_base + "/api/assets", files=files, timeout=120)
|
||||
@@ -442,33 +237,17 @@ def test_upload_missing_file_and_hash(http: requests.Session, api_base: str):
|
||||
assert body["error"]["code"] == "MISSING_FILE"
|
||||
|
||||
|
||||
def test_upload_models_unknown_model_type(http: requests.Session, api_base: str):
|
||||
def test_upload_models_unknown_category(http: requests.Session, api_base: str):
|
||||
files = {"file": ("m.safetensors", b"A" * 128, "application/octet-stream")}
|
||||
form = {"tags": json.dumps(["models", "model_type:no_such_category", "unit-tests"]), "name": "m.safetensors"}
|
||||
form = {"tags": json.dumps(["models", "no_such_category", "unit-tests"]), "name": "m.safetensors"}
|
||||
r = http.post(api_base + "/api/assets", data=form, files=files, timeout=120)
|
||||
body = r.json()
|
||||
assert r.status_code == 400, body
|
||||
assert r.status_code == 400
|
||||
assert body["error"]["code"] == "INVALID_BODY"
|
||||
assert body["error"]["message"].startswith("unknown models category")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model_type", ["configs", "custom_nodes"])
|
||||
def test_upload_models_rejects_non_model_registered_folder(
|
||||
model_type: str, http: requests.Session, api_base: str
|
||||
):
|
||||
files = {"file": ("not-a-model.py", b"A" * 128, "application/octet-stream")}
|
||||
form = {
|
||||
"tags": json.dumps(["models", f"model_type:{model_type}", "unit-tests"]),
|
||||
"name": "not-a-model.py",
|
||||
}
|
||||
|
||||
response = http.post(api_base + "/api/assets", data=form, files=files, timeout=120)
|
||||
body = response.json()
|
||||
|
||||
assert response.status_code == 400, body
|
||||
assert body["error"]["code"] == "INVALID_BODY"
|
||||
|
||||
|
||||
def test_upload_models_requires_model_type(http: requests.Session, api_base: str):
|
||||
def test_upload_models_requires_category(http: requests.Session, api_base: str):
|
||||
files = {"file": ("nocat.safetensors", b"A" * 64, "application/octet-stream")}
|
||||
form = {"tags": json.dumps(["models"]), "name": "nocat.safetensors", "user_metadata": json.dumps({})}
|
||||
r = http.post(api_base + "/api/assets", data=form, files=files, timeout=120)
|
||||
@@ -477,152 +256,13 @@ def test_upload_models_requires_model_type(http: requests.Session, api_base: str
|
||||
assert body["error"]["code"] == "INVALID_BODY"
|
||||
|
||||
|
||||
def test_upload_extra_tags_are_labels_not_path_components(http: requests.Session, api_base: str):
|
||||
def test_upload_tags_traversal_guard(http: requests.Session, api_base: str):
|
||||
files = {"file": ("evil.safetensors", b"A" * 256, "application/octet-stream")}
|
||||
form = {"tags": json.dumps(["models", "model_type:checkpoints", "unit-tests", "..", "zzz"]), "name": "evil.safetensors"}
|
||||
form = {"tags": json.dumps(["models", "checkpoints", "unit-tests", "..", "zzz"]), "name": "evil.safetensors"}
|
||||
r = http.post(api_base + "/api/assets", data=form, files=files, timeout=120)
|
||||
body = r.json()
|
||||
assert r.status_code == 201, body
|
||||
assert ".." in body["tags"]
|
||||
assert "zzz" in body["tags"]
|
||||
assert "models" in body["tags"]
|
||||
assert "model_type:checkpoints" in body["tags"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("subfolder", "expected_tag", "unexpected_tags"),
|
||||
[
|
||||
("custom/session", None, {"custom", "session"}),
|
||||
("pasted", "pasted", set()),
|
||||
],
|
||||
)
|
||||
def test_upload_image_accepts_arbitrary_subfolder_but_only_known_values_become_tags(
|
||||
http: requests.Session,
|
||||
api_base: str,
|
||||
comfy_tmp_base_dir: Path,
|
||||
subfolder: str,
|
||||
expected_tag: str | None,
|
||||
unexpected_tags: set[str],
|
||||
):
|
||||
name = f"upload-image-{uuid.uuid4().hex}.png"
|
||||
files = {"image": (name, b"image-upload" * 64, "image/png")}
|
||||
form = {"type": "input", "subfolder": subfolder}
|
||||
|
||||
response = http.post(api_base + "/upload/image", data=form, files=files, timeout=120)
|
||||
body = response.json()
|
||||
|
||||
assert response.status_code == 200, body
|
||||
assert body["subfolder"] == subfolder
|
||||
assert (comfy_tmp_base_dir / "input" / subfolder / body["name"]).exists()
|
||||
|
||||
asset = body["asset"]
|
||||
tags = set(asset["tags"])
|
||||
assert "input" in tags
|
||||
assert "uploaded" in tags
|
||||
if expected_tag:
|
||||
assert expected_tag in tags
|
||||
assert tags.isdisjoint(unexpected_tags)
|
||||
|
||||
|
||||
def test_multipart_upload_accepts_system_looking_extra_labels(
|
||||
http: requests.Session, api_base: str
|
||||
):
|
||||
files = {"file": ("relaxed-labels.bin", b"relaxed" * 64, "application/octet-stream")}
|
||||
form = {
|
||||
"tags": json.dumps(
|
||||
[
|
||||
"input",
|
||||
"unit-tests",
|
||||
"model:true",
|
||||
"models:foo",
|
||||
"temporary",
|
||||
"uploaded:true",
|
||||
]
|
||||
),
|
||||
"name": "relaxed-labels.bin",
|
||||
"user_metadata": json.dumps({}),
|
||||
}
|
||||
response = http.post(api_base + "/api/assets", data=form, files=files, timeout=120)
|
||||
body = response.json()
|
||||
|
||||
assert response.status_code == 201, body
|
||||
assert "input" in body["tags"]
|
||||
assert "model:true" in body["tags"]
|
||||
assert "models:foo" in body["tags"]
|
||||
assert "temporary" in body["tags"]
|
||||
assert "uploaded:true" in body["tags"]
|
||||
|
||||
|
||||
def test_multipart_upload_rejects_ambiguous_destination_roles(
|
||||
http: requests.Session, api_base: str
|
||||
):
|
||||
files = {"file": ("ambiguous.bin", b"ambiguous" * 64, "application/octet-stream")}
|
||||
form = {
|
||||
"tags": json.dumps(["input", "output", "unit-tests"]),
|
||||
"name": "ambiguous.bin",
|
||||
"user_metadata": json.dumps({}),
|
||||
}
|
||||
response = http.post(api_base + "/api/assets", data=form, files=files, timeout=120)
|
||||
body = response.json()
|
||||
|
||||
assert response.status_code == 400, body
|
||||
assert body["error"]["code"] == "INVALID_BODY"
|
||||
|
||||
|
||||
def test_multipart_upload_rejects_multiple_model_types_for_models_destination(
|
||||
http: requests.Session, api_base: str
|
||||
):
|
||||
files = {"file": ("ambiguous-model.safetensors", b"ambiguous-model" * 64, "application/octet-stream")}
|
||||
form = {
|
||||
"tags": json.dumps(
|
||||
["models", "model_type:checkpoints", "model_type:loras", "unit-tests"]
|
||||
),
|
||||
"name": "ambiguous-model.safetensors",
|
||||
"user_metadata": json.dumps({}),
|
||||
}
|
||||
response = http.post(api_base + "/api/assets", data=form, files=files, timeout=120)
|
||||
body = response.json()
|
||||
|
||||
assert response.status_code == 400, body
|
||||
assert body["error"]["code"] == "INVALID_BODY"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("tags", "expected_root", "extension"),
|
||||
[
|
||||
(["input", "unit-tests", "upload-location-input"], "input", ".bin"),
|
||||
(["output", "unit-tests", "upload-location-output"], "output", ".bin"),
|
||||
(
|
||||
["models", "model_type:checkpoints", "unit-tests", "upload-location-model"],
|
||||
"models/checkpoints",
|
||||
".safetensors",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_multipart_upload_role_selects_write_location(
|
||||
http: requests.Session,
|
||||
api_base: str,
|
||||
comfy_tmp_base_dir: Path,
|
||||
tags: list[str],
|
||||
expected_root: str,
|
||||
extension: str,
|
||||
):
|
||||
role = next(tag for tag in tags if tag in {"input", "models", "output"})
|
||||
name = f"{role}-role-upload{extension}"
|
||||
files = {"file": (name, f"{role}-role-bytes".encode() * 64, "application/octet-stream")}
|
||||
form = {
|
||||
"tags": json.dumps(tags),
|
||||
"name": name,
|
||||
"user_metadata": json.dumps({}),
|
||||
}
|
||||
|
||||
response = http.post(api_base + "/api/assets", data=form, files=files, timeout=120)
|
||||
body = response.json()
|
||||
|
||||
assert response.status_code == 201, body
|
||||
stored_name = get_asset_filename(body["asset_hash"], extension)
|
||||
expected_disk_path = comfy_tmp_base_dir / expected_root / stored_name
|
||||
assert expected_disk_path.exists()
|
||||
assert r.status_code == 400
|
||||
assert body["error"]["code"] in ("BAD_REQUEST", "INVALID_BODY")
|
||||
|
||||
|
||||
def test_upload_empty_tags_rejected(http: requests.Session, api_base: str):
|
||||
|
||||
@@ -1,204 +0,0 @@
|
||||
"""Unit tests for io.DynamicGroup: expansion/reconstruction (0-row and N-row cases)."""
|
||||
import sys
|
||||
import types
|
||||
import pytest
|
||||
|
||||
# Stub torch (type-hint only in _io.py; real torch not available in unit-test env)
|
||||
if "torch" not in sys.modules:
|
||||
_torch_stub = types.ModuleType("torch")
|
||||
_torch_stub.Tensor = object # type: ignore[attr-defined]
|
||||
sys.modules["torch"] = _torch_stub
|
||||
|
||||
from comfy_api.latest._io import ( # noqa: E402
|
||||
DynamicGroup,
|
||||
Float,
|
||||
Int,
|
||||
String,
|
||||
Boolean,
|
||||
get_finalized_class_inputs,
|
||||
build_nested_inputs,
|
||||
create_input_dict_v1,
|
||||
setup_dynamic_input_funcs,
|
||||
)
|
||||
|
||||
# Make sure dynamic input funcs are registered (may already be done at import time)
|
||||
setup_dynamic_input_funcs()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def _make_class_inputs(group_input: DynamicGroup.Input) -> dict:
|
||||
"""Wrap a DynamicGroup.Input into the required/optional dict structure."""
|
||||
return create_input_dict_v1([group_input])
|
||||
|
||||
|
||||
def _run(group_input: DynamicGroup.Input, live_values: dict) -> dict:
|
||||
"""End-to-end helper: expand schema + reconstruct values.
|
||||
|
||||
Mirrors the production split in execution.py:
|
||||
1. get_finalized_class_inputs (schema expansion, line 162)
|
||||
2. build_nested_inputs (value reconstruction, line 281)
|
||||
|
||||
The two steps are separate in production because the engine resolves
|
||||
linked node outputs between them, but in tests we supply values directly.
|
||||
"""
|
||||
class_inputs = _make_class_inputs(group_input)
|
||||
_, _, v3_data = get_finalized_class_inputs(class_inputs, live_values)
|
||||
return build_nested_inputs(dict(live_values), v3_data)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Schema construction
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestDynamicGroupInputConstruction:
|
||||
def test_basic_construction(self):
|
||||
inp = DynamicGroup.Input(
|
||||
"loras",
|
||||
template=[
|
||||
Float.Input("strength", default=1.0),
|
||||
String.Input("name"),
|
||||
],
|
||||
min=0,
|
||||
max=10,
|
||||
)
|
||||
assert inp.id == "loras"
|
||||
assert inp.min == 0
|
||||
assert inp.max == 10
|
||||
assert len(inp.template) == 2
|
||||
|
||||
def test_get_all_includes_self_and_template(self):
|
||||
inp = DynamicGroup.Input(
|
||||
"items",
|
||||
template=[Float.Input("value")],
|
||||
)
|
||||
all_inputs = inp.get_all()
|
||||
assert all_inputs[0] is inp
|
||||
assert all_inputs[1].id == "value"
|
||||
|
||||
def test_as_dict_has_template_min_max(self):
|
||||
inp = DynamicGroup.Input(
|
||||
"items",
|
||||
template=[Float.Input("val", default=0.5)],
|
||||
min=1,
|
||||
max=5,
|
||||
)
|
||||
d = inp.as_dict()
|
||||
assert "template" in d
|
||||
assert d["min"] == 1
|
||||
assert d["max"] == 5
|
||||
|
||||
def test_duplicate_field_ids_raises(self):
|
||||
with pytest.raises(AssertionError):
|
||||
DynamicGroup.Input(
|
||||
"bad",
|
||||
template=[Float.Input("x"), Float.Input("x")],
|
||||
)
|
||||
|
||||
def test_empty_template_raises(self):
|
||||
with pytest.raises(AssertionError):
|
||||
DynamicGroup.Input("bad", template=[])
|
||||
|
||||
def test_min_gt_max_raises(self):
|
||||
with pytest.raises(AssertionError):
|
||||
DynamicGroup.Input("bad", template=[Float.Input("x")], min=5, max=3)
|
||||
|
||||
def test_max_exceeds_limit_raises(self):
|
||||
with pytest.raises(AssertionError):
|
||||
DynamicGroup.Input("bad", template=[Float.Input("x")], max=101)
|
||||
|
||||
def test_dynamic_input_in_template_raises(self):
|
||||
with pytest.raises(AssertionError):
|
||||
DynamicGroup.Input(
|
||||
"bad",
|
||||
template=[DynamicGroup.Input("nested", template=[Float.Input("x")])],
|
||||
)
|
||||
|
||||
def test_validate_calls_through(self):
|
||||
inp = DynamicGroup.Input("items", template=[Float.Input("val", min=-1.0, max=1.0)])
|
||||
inp.validate() # should not raise
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 0-row case
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestZeroRows:
|
||||
def test_empty_live_inputs_produces_empty_list(self):
|
||||
"""With min=0 and no live values, the result should be an empty list."""
|
||||
inp = DynamicGroup.Input("loras", template=[Float.Input("strength", default=1.0)], min=0, max=10)
|
||||
assert _run(inp, {}).get("loras") == []
|
||||
|
||||
def test_min_zero_with_values(self):
|
||||
"""min=0 but 2 rows of live data."""
|
||||
inp = DynamicGroup.Input("loras", template=[Float.Input("strength", default=1.0)], min=0, max=10)
|
||||
result = _run(inp, {"loras.0.strength": 0.8, "loras.1.strength": 0.5})
|
||||
assert result["loras"] == [{"strength": 0.8}, {"strength": 0.5}]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# N-row case
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestNRows:
|
||||
def test_two_rows_two_fields(self):
|
||||
"""Two rows with two fields each produce a list[dict]."""
|
||||
inp = DynamicGroup.Input(
|
||||
"loras",
|
||||
template=[String.Input("lora_name"), Float.Input("strength", default=1.0)],
|
||||
min=0, max=50,
|
||||
)
|
||||
result = _run(inp, {
|
||||
"loras.0.lora_name": "model_a.safetensors", "loras.0.strength": 0.9,
|
||||
"loras.1.lora_name": "model_b.safetensors", "loras.1.strength": 0.4,
|
||||
})
|
||||
assert result["loras"] == [
|
||||
{"lora_name": "model_a.safetensors", "strength": 0.9},
|
||||
{"lora_name": "model_b.safetensors", "strength": 0.4},
|
||||
]
|
||||
|
||||
def test_rows_are_sorted_by_index(self):
|
||||
"""Rows must be in ascending index order even if dict iteration is unordered."""
|
||||
inp = DynamicGroup.Input("items", template=[Int.Input("v", default=0)], min=0, max=10)
|
||||
result = _run(inp, {"items.0.v": 10, "items.2.v": 30, "items.1.v": 20})
|
||||
assert [row["v"] for row in result["items"]] == [10, 20, 30]
|
||||
|
||||
def test_min_rows_schema_slots(self):
|
||||
"""With min=2 and no live data, 2 slots must appear in the expanded schema."""
|
||||
inp = DynamicGroup.Input("items", template=[Float.Input("val", default=0.0)], min=2, max=5)
|
||||
out, _, _ = get_finalized_class_inputs(_make_class_inputs(inp), {})
|
||||
all_slots = {**out.get("required", {}), **out.get("optional", {})}
|
||||
assert "items.0.val" in all_slots
|
||||
assert "items.1.val" in all_slots
|
||||
|
||||
def test_min_rows_reconstructs_when_no_values(self):
|
||||
"""min=2 with NO live values must still yield a 2-element list,
|
||||
not collapse to [] (regression: parent-path clobber)."""
|
||||
inp = DynamicGroup.Input("items", template=[Float.Input("val", default=0.0)], min=2, max=5)
|
||||
result = _run(inp, {})
|
||||
assert len(result["items"]) == 2
|
||||
assert all("val" in row for row in result["items"])
|
||||
|
||||
def test_min_rows_reconstructs_with_partial_values(self):
|
||||
"""min=2 with only the first row's value present still yields 2 rows."""
|
||||
inp = DynamicGroup.Input("items", template=[Float.Input("val", default=0.0)], min=2, max=5)
|
||||
result = _run(inp, {"items.0.val": 0.7})
|
||||
assert len(result["items"]) == 2
|
||||
assert result["items"][0]["val"] == 0.7
|
||||
assert result["items"][1]["val"] is None
|
||||
|
||||
def test_list_paths_in_v3_data(self):
|
||||
"""list_paths must contain the group id so build_nested_inputs knows to convert."""
|
||||
inp = DynamicGroup.Input("things", template=[Boolean.Input("flag")], min=0, max=5)
|
||||
_, _, v3_data = get_finalized_class_inputs(_make_class_inputs(inp), {})
|
||||
assert "things" in v3_data.get("list_paths", set())
|
||||
|
||||
def test_no_leftover_flat_keys(self):
|
||||
"""Flat keys must be consumed; only the reconstructed list remains."""
|
||||
inp = DynamicGroup.Input("rows", template=[Float.Input("x", default=0.0)], min=0, max=5)
|
||||
result = _run(inp, {"rows.0.x": 1.0, "rows.1.x": 2.0})
|
||||
assert "rows.0.x" not in result
|
||||
assert "rows.1.x" not in result
|
||||
assert isinstance(result["rows"], list)
|
||||
@@ -11,11 +11,6 @@ from comfy_api.feature_flags import (
|
||||
_coerce_flag_value,
|
||||
_parse_cli_feature_flags,
|
||||
)
|
||||
from comfy.comfy_api_env import (
|
||||
environment_overrides_for_base,
|
||||
get_environment_overrides,
|
||||
normalize_comfy_api_base,
|
||||
)
|
||||
|
||||
|
||||
class TestFeatureFlags:
|
||||
@@ -34,8 +29,6 @@ class TestFeatureFlags:
|
||||
features = get_server_features()
|
||||
assert "supports_preview_metadata" in features
|
||||
assert features["supports_preview_metadata"] is True
|
||||
assert "supports_model_type_tags" in features
|
||||
assert features["supports_model_type_tags"] is True
|
||||
assert "max_upload_size" in features
|
||||
assert isinstance(features["max_upload_size"], (int, float))
|
||||
|
||||
@@ -188,65 +181,3 @@ class TestCliFeatureFlagRegistry:
|
||||
assert "type" in info, f"{key} missing 'type'"
|
||||
assert "default" in info, f"{key} missing 'default'"
|
||||
assert "description" in info, f"{key} missing 'description'"
|
||||
|
||||
|
||||
class TestComfyApiEnv:
|
||||
"""--comfy-api-base staging-tier detection + testenv main-host -> -registry rewrite."""
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"url, expected",
|
||||
[
|
||||
# testenv friendly main host -> comfy-api -registry sibling (slash trimmed)
|
||||
("https://pr-4398.testenvs.comfy.org", "https://pr-4398-registry.testenvs.comfy.org"),
|
||||
("https://pr-4398.testenvs.comfy.org/", "https://pr-4398-registry.testenvs.comfy.org"),
|
||||
("https://pr-4398-registry.testenvs.comfy.org", "https://pr-4398-registry.testenvs.comfy.org"),
|
||||
# staging + everything else -> unchanged (no -registry split)
|
||||
("https://stagingapi.comfy.org", "https://stagingapi.comfy.org"),
|
||||
("https://api.comfy.org", "https://api.comfy.org"),
|
||||
("https://pr-1.testenvs.comfy.org.evil.com", "https://pr-1.testenvs.comfy.org.evil.com"),
|
||||
("", ""),
|
||||
],
|
||||
)
|
||||
def test_normalize_comfy_api_base(self, url, expected):
|
||||
assert normalize_comfy_api_base(url) == expected
|
||||
|
||||
def test_config_for_staging_tier_else_none(self):
|
||||
# ephemeral testenv: friendly main host -> -registry, staging platform, dev Firebase env
|
||||
eph = environment_overrides_for_base("https://pr-1234.testenvs.comfy.org/")
|
||||
assert eph["comfy_api_base_url"] == "https://pr-1234-registry.testenvs.comfy.org"
|
||||
assert eph["comfy_platform_base_url"] == "https://stagingplatform.comfy.org"
|
||||
assert eph["firebase_env"] == "dev"
|
||||
# staging api host: emitted as-is
|
||||
stg = environment_overrides_for_base("https://stagingapi.comfy.org")
|
||||
assert stg["comfy_api_base_url"] == "https://stagingapi.comfy.org"
|
||||
assert stg["comfy_platform_base_url"] == "https://stagingplatform.comfy.org"
|
||||
assert stg["firebase_env"] == "dev"
|
||||
# prod / unknown: nothing
|
||||
assert environment_overrides_for_base("https://api.comfy.org") is None
|
||||
|
||||
def test_environment_overrides_only_for_staging_tier(self, monkeypatch):
|
||||
def set_base(url):
|
||||
monkeypatch.setattr(
|
||||
"comfy.comfy_api_env.args",
|
||||
type("Args", (), {"comfy_api_base": url})(),
|
||||
)
|
||||
|
||||
# The overrides merged into the HTTP /features response are present for staging-tier bases...
|
||||
set_base("https://stagingapi.comfy.org")
|
||||
assert "comfy_api_base_url" in get_environment_overrides()
|
||||
set_base("https://pr-7.testenvs.comfy.org")
|
||||
assert "comfy_api_base_url" in get_environment_overrides()
|
||||
# ...but never for prod.
|
||||
set_base("https://api.comfy.org")
|
||||
assert get_environment_overrides() is None
|
||||
|
||||
def test_server_features_never_carry_env_overrides(self, monkeypatch):
|
||||
"""The WebSocket capability handshake must stay free of routing keys."""
|
||||
monkeypatch.setattr(
|
||||
"comfy.comfy_api_env.args",
|
||||
type("Args", (), {"comfy_api_base": "https://pr-7.testenvs.comfy.org"})(),
|
||||
)
|
||||
features = get_server_features()
|
||||
assert "comfy_api_base_url" not in features
|
||||
assert "comfy_platform_base_url" not in features
|
||||
assert "firebase_env" not in features
|
||||
|
||||
@@ -12,8 +12,6 @@ class TestWebSocketFeatureFlags:
|
||||
# Check expected server features
|
||||
assert "supports_preview_metadata" in features
|
||||
assert features["supports_preview_metadata"] is True
|
||||
assert "supports_model_type_tags" in features
|
||||
assert features["supports_model_type_tags"] is True
|
||||
assert "max_upload_size" in features
|
||||
assert isinstance(features["max_upload_size"], (int, float))
|
||||
|
||||
@@ -77,5 +75,3 @@ class TestWebSocketFeatureFlags:
|
||||
assert server_message["type"] == "feature_flags"
|
||||
assert "supports_preview_metadata" in server_message["data"]
|
||||
assert server_message["data"]["supports_preview_metadata"] is True
|
||||
assert "supports_model_type_tags" in server_message["data"]
|
||||
assert server_message["data"]["supports_model_type_tags"] is True
|
||||
|
||||
Reference in New Issue
Block a user