diff --git a/.coderabbit.yaml b/.coderabbit.yaml index 0d1e49270..08629ed8e 100644 --- a/.coderabbit.yaml +++ b/.coderabbit.yaml @@ -4,12 +4,12 @@ early_access: false tone_instructions: "Only comment on issues introduced by this PR's changes. Do not flag pre-existing problems in moved, re-indented, or reformatted code." reviews: - profile: "chill" - request_changes_workflow: false + profile: "assertive" + request_changes_workflow: true high_level_summary: false poem: false review_status: false - review_details: false + review_details: true commit_status: true collapse_walkthrough: true changed_files_summary: false @@ -39,6 +39,14 @@ reviews: - path: "**" instructions: | IMPORTANT: Only comment on issues directly introduced by this PR's code changes. + Treat AGENTS.md as mandatory repository policy, not optional style guidance. + Flag PR changes that violate AGENTS.md even when the code is otherwise functional. + In particular, enforce architecture boundaries, dtype/device/memory rules, + interface contracts, import style, no unnecessary try/except blocks, no inline + imports, no outbound internet paths in core ComfyUI, and narrow scoped fixes. + Prefer direct findings over suggestions when a rule is violated. Only ignore + AGENTS.md when it clearly conflicts with a newer explicit maintainer instruction + in the PR. Do NOT flag pre-existing issues in code that was merely moved, re-indented, de-indented, or reformatted without logic changes. If code appears in the diff only due to whitespace or structural reformatting (e.g., removing a `with:` block), @@ -123,5 +131,10 @@ chat: knowledge_base: opt_out: false + code_guidelines: + enabled: true + filePatterns: + - files: "AGENTS.md" + applyTo: "**" learnings: scope: "auto" diff --git a/.github/workflows/cla.yml b/.github/workflows/cla.yml new file mode 100644 index 000000000..b75397e50 --- /dev/null +++ b/.github/workflows/cla.yml @@ -0,0 +1,91 @@ +name: CLA Assistant + +on: + issue_comment: + types: [created] + pull_request_target: + types: [opened, synchronize, closed] + +permissions: + actions: write + contents: read # 'read' is enough because signatures live in a REMOTE repo + pull-requests: write + statuses: write + +jobs: + cla-assistant: + runs-on: ubuntu-latest + steps: + # The CLA action normally requires every commit author in a PR to sign. + # We only want the PR author to sign, so we allowlist all other committers + # by computing them from the PR's commits and excluding the PR author. + - name: Build author-only allowlist + id: allowlist + if: > + github.event_name == 'pull_request_target' || + (github.event_name == 'issue_comment' && github.event.issue.pull_request && ( + github.event.comment.body == 'recheck' || + github.event.comment.body == 'I have read and agree to the Contributor License Agreement' + )) + env: + GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} + PR_NUMBER: ${{ github.event.pull_request.number || github.event.issue.number }} + PR_AUTHOR: ${{ github.event.pull_request.user.login || github.event.issue.user.login }} + BASE_ALLOWLIST: action@github.com,actions-user,ampagent,claude,comfy-pr-bot,GitHub Action,github-actions,github-actions[bot],Glary Bot,Glary-Bot,*[bot] + run: | + others=$(gh api "repos/${{ github.repository }}/pulls/${PR_NUMBER}/commits" --paginate \ + --jq '.[] | (.author.login // empty), (.committer.login // empty)' \ + | sort -u | grep -vix "${PR_AUTHOR}" | paste -sd, -) + if [ -n "$others" ]; then + echo "allowlist=${BASE_ALLOWLIST},${others}" >> "$GITHUB_OUTPUT" + else + echo "allowlist=${BASE_ALLOWLIST}" >> "$GITHUB_OUTPUT" + fi + + - name: CLA Assistant + # Run on PR events, on "recheck" comment, or when someone posts the exact signing phrase. + # IMPORTANT: this phrase must match `custom-pr-sign-comment` below. + if: > + github.event_name == 'pull_request_target' || + (github.event_name == 'issue_comment' && github.event.issue.pull_request && ( + github.event.comment.body == 'recheck' || + github.event.comment.body == 'I have read and agree to the Contributor License Agreement' + )) + uses: contributor-assistant/github-action@ca4a40a7d1004f18d9960b404b97e5f30a505a08 # v2.6.1 + env: + GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} + # PAT required to write to the centralized signatures repo. + PERSONAL_ACCESS_TOKEN: ${{ secrets.PERSONAL_ACCESS_TOKEN }} + with: + # Where the CLA document lives (shown to contributors) + path-to-document: https://github.com/Comfy-Org/comfy-cla/blob/main/comfyui_icla.md + + # Centralized signature storage + remote-organization-name: comfy-org + remote-repository-name: comfy-cla + path-to-signatures: signatures/cla.json + branch: main + + # Only the PR author must sign: bots plus every non-author committer + # are allowlisted via the "Build author-only allowlist" step above. + # *[bot] is a catch-all for any GitHub App bot account. + allowlist: ${{ steps.allowlist.outputs.allowlist }} + + # Custom PR comment messages + custom-notsigned-prcomment: | + 🎉 Thank you for your contribution, we really appreciate it! 🎉 + + Like many open source projects, we require contributors to sign our [Contributor License Agreement (CLA)](https://github.com/Comfy-Org/comfy-cla/blob/main/comfyui_icla.md). A CLA makes the ownership of contributions explicit, so contributors and the project share a clear understanding of how the code can be used. By signing, you: + + - Confirm that you own your contribution. + - Keep the right to reuse your own code. + - Grant us a copyright license to include and share it within our projects. + + CLAs are standard practice across major open source projects including those under the Apache Software Foundation and the Linux Foundation. Ours is based on the Apache Software Foundation's CLA. Most importantly, it would enable us to relicense the project under a more permissive license in the future, giving the project and its community greater flexibility. + + ✍ **To sign, please post a new comment on this PR with exactly the following text:** ✍ + + custom-pr-sign-comment: I have read and agree to the Contributor License Agreement + + custom-allsigned-prcomment: | + ✅ All contributors have signed the CLA. Thank you! This PR is ready to be merged. diff --git a/AGENTS.md b/AGENTS.md index 70dfaa186..05efd834b 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -8,10 +8,13 @@ directly required. - Prefer practical fixes over broad architecture work. Add abstractions only when they remove real repeated logic or match an existing ComfyUI pattern. +- Prefer fewer dependencies. Do not add new dependencies to ComfyUI unless they + are absolutely necessary. - Delete obsolete code aggressively when newer infrastructure makes it useless. Remove dead fallbacks, migration paths, unused options, debug prints, and compatibility branches that are no longer needed. Do not leave dead branches, - unreachable code, or functions that are never called. + unreachable code, or functions that are never called. If code is not + necessary for the current behavior, remove it. - Revert or disable problematic behavior quickly when it breaks users. It is better to remove a broken feature path than keep a complicated partial fix. - Preserve existing APIs, node names, model-loading behavior, file layout, and @@ -85,6 +88,14 @@ not change a shared method to return extra values, alternate shapes, or sentinel wrappers for one implementation unless the shared interface is explicitly updated. +- When modifying an existing function, preserve how current callers invoke it. + Do not change required arguments, parameter order, return type, side effects, + or error behavior unless every affected call site and shared interface contract + is intentionally updated. +- Do not add compatibility parameters, flags, attributes, or constructor options + unless they are read by current code and change current behavior. Remove + pass-through or stored-but-unused values instead of preserving upstream or + deprecated API baggage. - If an implementation needs auxiliary values for its own workflow, expose them through a private helper or a clearly named implementation-specific method instead of overloading the public method's return contract. @@ -102,6 +113,11 @@ - Do not add freeze, unfreeze, or trainability toggles to model classes. ComfyUI models are always treated as frozen for inference, so explicit freeze functionality is redundant and should not be added. +- Remove training-only behavior such as dropout from inference model code, but + preserve checkpoint and state-dict compatibility when doing so. If deleting a + module would change state-dict keys, module ordering, or checkpoint loading + behavior, replace it with a no-op such as `nn.Identity` instead of removing the + slot outright. ## Python Style @@ -111,6 +127,13 @@ - Do not add unnecessary `try`/`except` blocks. Use them for optional dependency, platform, or backend capability detection only when the program has a useful fallback. Prefer specific exception types when changing new code. +- If a library version is pinned in `requirements.txt`, do not add code to + ComfyUI to handle older versions of that library. +- Remove any workarounds for PyTorch versions that ComfyUI no longer officially + supports. Deprecated workarounds include catching an exception and rerunning + the same op with the input cast to float. If a workaround does not have a + comment naming the exact PyTorch version or versions that still need it, + remove it. - Let unsupported model formats, invalid quantization metadata, and bad states fail with clear errors instead of silently producing lower quality output. - Match the existing local style in the file you edit. This codebase tolerates @@ -129,8 +152,101 @@ adding parallel code paths. Use `comfy.quant_ops`, `comfy.model_management`, `comfy.memory_management`, `comfy.pinned_memory`, `comfy_aimdo`, and `comfy-kitchen` helpers where they already solve the problem. +- Use optimized comfy-kitchen ops in places where they improve performance + without changing the expected dtype, device, memory, or interface behavior. +- All models should use the optimized attention function selected by ComfyUI. + Treat optimized backend functions, dispatch helpers, and capability-selected + callables as opaque. Higher-level code must not inspect function identity, + names, modules, or implementation details to decide behavior. +- Apply the same opacity rule to similar patterns beyond attention: callers + should depend on the documented interface and result contract, not on which + backend implementation was selected underneath. +- Do not use custom inference ops that only duplicate an existing op while + upcasting to float32, such as custom RMSNorm variants. Use the generic ComfyUI + ops and/or native torch ops instead. +- If a model class `__init__` has an `operations` parameter, assume + `operations` is never `None`. Do not add fallback branches or default torch + ops for a missing `operations` object. +- Do not add unnecessary parameters to model, model block, or model ops related + classes. Constructor and forward signatures should carry only values that are + actually needed by that object for inference. +- Reuse existing model classes, blocks, ops, and helper modules when appropriate. + Before implementing a new version of a model component, search the existing + model code for a class or helper that already provides the behavior. +- Model detection code that inspects linear weight shapes should only use the + first dimension. The second dimension may be half the original size for + NVFP4 or other 4-bit quantized models. +- Avoid adding `einops` usage in core inference code. Use native torch tensor + ops such as `reshape`, `view`, `permute`, `transpose`, `flatten`, `unflatten`, + `unsqueeze`, and `squeeze` instead. +- Do not use tensors as general-purpose Python data structures. Keep metadata, + bookkeeping, counters, flags, shape math, padding math, index planning, memory + estimates, and control-flow decisions in plain Python values unless the data + must participate directly in tensor computation. Do not create tensors for + structural metadata that is only used for Python-side control flow. Sequence + lengths, cumulative offsets, split indices, window counts, slice boundaries, + and repeat counts should be kept as Python ints/lists from the point they are + computed. Do not build them as CPU/GPU tensors and then cast, move, validate, + or convert them back to Python for `split`, `tensor_split`, indexing plans, + loops, or cache keys. Avoid creating temporary tensors just to use tensor + methods for scalar or structural calculations. - Avoid unnecessary casts and transfers. Preserve the intended compute dtype, storage dtype, bias dtype, and original tensor shape metadata. +- Keep model-native latent layout handling inside the model or latent-format + owner, not in helper nodes. Do not collapse, expand, pack, or unpack latent + dimensions in nodes or other caller-side adapters just to satisfy a model + forward; the model path should consume and return the native latent shape for + that model family. +- Assume inputs to the main model forward are already in the compute dtype by + default, except integer inputs such as some model timestep tensors. Do not add + defensive or convenience casts in model code; it is better for invalid dtype + plumbing to error clearly than to hide it with unnecessary casts. +- Raw model parameters that are not owned by an op and may be initialized in a + dtype different from the compute dtype should be cast at use in forward or + inference code with `comfy.ops.cast_to_input` or + `comfy.model_management.cast_to` to avoid dtype mismatches. +- Model code should not care what dtype it is initialized in, and model + `__init__` methods should not contain workarounds for specific dtypes. Dtype + workaround code, such as making a model work with fp16 compute, belongs in the + execution or model-management layer that owns compute policy. +- Model code should not perform unnecessary device-to-CPU or CPU-to-device + transfers. New allocations must be created on the correct device and dtype; + never allocate on CPU and then move to GPU, or allocate in one dtype and then + convert to another. +- Model code itself should not perform memory management. Loading, unloading, + offloading, device movement, VRAM policy, cache lifetime, and cleanup belong + in the relevant model-management and execution layers, not inside model + implementations. +- Do not add global, module-level, class-level, singleton, or model-owned stores + for tensors or other large memory that persist across executions. Temporary + caches must be scoped to a single execution or forward/encode/decode call: + allocate them in the owning top-level call, pass them explicitly through the + call stack, and let them be discarded when that call returns. +- Follow the Wan VAE temporal cache pattern for temporary caches: create a local + cache such as `feat_map` for the encode/decode operation, pass it into the + blocks that need it, and do not retain it on the model or in global state. +- In model init code, prefer `torch.empty` for parameter/buffer placeholders + that are populated from the model state dict instead of zero-initializing with + `torch.zeros` or similar. If an allocation is not loaded from the state dict + and is useless for inference, do not include it. +- `nn.Parameter` tensors that are stored in and populated from the model state + dict should be initialized with `torch.empty`, not with zero, random, or + otherwise meaningful initialization. +- Model initialization should describe module structure, not fabricate + checkpoint-owned tensor contents. Parameters and buffers that are loaded from + the state dict must not be manually initialized, reassigned, or filled with + fallback values unless that value is actually used when no checkpoint key + exists. +- When slicing large tensors, copy the slice if the sliced tensor's lifetime + exceeds the current function scope. Do not keep a long-lived view into a large + backing tensor when a smaller copy would release memory sooner. +- Use fused or compound torch operations such as `addcmul` when they naturally + match the math. Reducing Python and torch dispatch overhead is a valid + optimization when it does not obscure the code or change dtype/device + behavior. +- Avoid caches that persist across different executions as much as possible. + Persistent caches are acceptable only when they use a very minimal amount of + memory and have a clear ownership and invalidation story. - When optimizing, favor small measurable changes: fewer allocations, fewer device transfers, less peak memory, better batching, or use of a faster existing backend op. @@ -141,6 +257,20 @@ `CATEGORY`, and registration through the local mapping used by that file. - Keep node changes backward compatible by default. Add inputs with sensible defaults and avoid changing output types unless the request requires it. +- Model implementations should add the minimal number of ComfyUI nodes required + to run the model. Reuse existing nodes as much as possible; adapting the model + to work with existing nodes is strongly preferred over creating new nodes. +- Nodes should output only values they own. Do not add pass-through outputs for + workflow convenience unless the node is explicitly an output node. Existing + models, latents, conditioning, or other inputs should flow directly to the + next consumer instead of being re-emitted unchanged. +- Nodes should expose only inputs they actually read to produce current + behavior. Do not add placeholder, pass-through, compatibility, or + workflow-shaping inputs that are ignored or could flow directly to another + node. +- Node-level code must not patch model code directly. Any node behavior that + modifies, wraps, hooks, or changes model behavior must go through the model + patcher class instead of reaching into model internals. - The official mascot of ComfyUI is a very cute anime girl with massive fennec ears, a big fluffy tail, long blonde wavy hair, and blue eyes. Feel free to use her in ComfyUI materials, UI text, examples, tests, generated assets, or diff --git a/README.md b/README.md index bcec86377..14c8d2cb2 100644 --- a/README.md +++ b/README.md @@ -229,7 +229,7 @@ Python 3.14 works but some custom nodes may have issues. The free threaded varia Python 3.13 is very well supported. If you have trouble with some custom node dependencies on 3.13 you can try 3.12 -torch 2.4 and above is supported but some features and optimizations might only work on newer versions. We generally recommend using the latest major version of pytorch with the latest cuda version unless it is less than 2 weeks old. +torch 2.5 is minimally supported but using a newer version is extremely recommended. Some features and optimizations might only work on newer versions. We generally recommend using the latest major version of pytorch with the latest cuda version unless it is less than 2 weeks old. If your pytorch is more than 6 months old, please update it. ### Instructions: diff --git a/alembic_db/versions/0005_allow_case_sensitive_tags.py b/alembic_db/versions/0005_allow_case_sensitive_tags.py new file mode 100644 index 000000000..bd5f864db --- /dev/null +++ b/alembic_db/versions/0005_allow_case_sensitive_tags.py @@ -0,0 +1,107 @@ +""" +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)" + ) diff --git a/alembic_db/versions/0006_add_loader_path.py b/alembic_db/versions/0006_add_loader_path.py new file mode 100644 index 000000000..afa65312d --- /dev/null +++ b/alembic_db/versions/0006_add_loader_path.py @@ -0,0 +1,30 @@ +""" +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") diff --git a/app/assets/api/routes.py b/app/assets/api/routes.py index cde55ec02..af08a99f0 100644 --- a/app/assets/api/routes.py +++ b/app/assets/api/routes.py @@ -40,6 +40,7 @@ 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() @@ -170,11 +171,19 @@ 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, @@ -315,12 +324,15 @@ async def download_asset_content(request: web.Request) -> web.Response: 404, "FILE_NOT_FOUND", "Underlying file not found on disk." ) - _DANGEROUS_MIME_TYPES = { - "text/html", "text/html-sandboxed", "application/xhtml+xml", - "text/javascript", "text/css", - } - if content_type in _DANGEROUS_MIME_TYPES: + # User-controlled asset content must never render inline in the app origin + # (stored XSS via SVG/HTML/XML). Force dangerous types to download and + # override any requested inline disposition. Centralised through + # folder_paths.is_dangerous_content_type so this can't drift from /view and + # /userdata (the previous inline set here omitted image/svg+xml and missed + # the charset/casing/+xml-dialect bypasses). + if folder_paths.is_dangerous_content_type(content_type): content_type = "application/octet-stream" + disposition = "attachment" safe_name = (filename or "").replace("\r", "").replace("\n", "") encoded = urllib.parse.quote(safe_name) @@ -425,17 +437,6 @@ 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: @@ -479,7 +480,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, "BAD_REQUEST", str(e)) + return _build_error_response(400, "INVALID_BODY", str(e)) except HashMismatchError as e: delete_temp_file_if_exists(parsed.tmp_path) return _build_error_response(400, "HASH_MISMATCH", str(e)) diff --git a/app/assets/api/schemas_in.py b/app/assets/api/schemas_in.py index af666746d..38a942b7b 100644 --- a/app/assets/api/schemas_in.py +++ b/app/assets/api/schemas_in.py @@ -140,7 +140,7 @@ class CreateFromHashBody(BaseModel): if v is None: return [] if isinstance(v, list): - out = [str(t).strip().lower() for t in v if str(t).strip()] + out = [str(t).strip() 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 [t.strip().lower() for t in v.split(",") if t.strip()] + return list(dict.fromkeys(t.strip() 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.lower() or None + return v 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().lower() + tnorm = t.strip() if tnorm: out.append(tnorm) seen = set() @@ -239,8 +239,8 @@ class TagsRemove(TagsAdd): class UploadAssetSpec(BaseModel): """Upload Asset operation. - - tags: optional list; if provided, first is root ('models'|'input'|'output'); - if root == 'models', second must be a valid category + - tags: labels plus one destination role ('models'|'input'|'output') for new bytes; + if role == 'models', exactly one model_type: tag is required - name: display name - user_metadata: arbitrary JSON object (optional) - hash: optional canonical 'blake3:' for validation / fast-path @@ -309,7 +309,7 @@ class UploadAssetSpec(BaseModel): norm = [] seen = set() for t in items: - tnorm = str(t).strip().lower() + tnorm = str(t).strip() if tnorm and tnorm not in seen: seen.add(tnorm) norm.append(tnorm) @@ -335,14 +335,4 @@ 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 diff --git a/app/assets/api/schemas_out.py b/app/assets/api/schemas_out.py index 4e38e19d1..da8251499 100644 --- a/app/assets/api/schemas_out.py +++ b/app/assets/api/schemas_out.py @@ -9,8 +9,20 @@ class Asset(BaseModel): ``id`` here is the AssetReference id, not the content-addressed Asset id.""" id: str - name: 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.", + ) 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 diff --git a/app/assets/api/upload.py b/app/assets/api/upload.py index 13d3d372c..2979f0e20 100644 --- a/app/assets/api/upload.py +++ b/app/assets/api/upload.py @@ -140,7 +140,6 @@ 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'." diff --git a/app/assets/database/models.py b/app/assets/database/models.py index 9b61d309a..329cd483d 100644 --- a/app/assets/database/models.py +++ b/app/assets/database/models.py @@ -76,6 +76,8 @@ 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) diff --git a/app/assets/database/queries/asset_reference.py b/app/assets/database/queries/asset_reference.py index 792411800..967b0e43a 100644 --- a/app/assets/database/queries/asset_reference.py +++ b/app/assets/database/queries/asset_reference.py @@ -650,6 +650,7 @@ 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). @@ -659,6 +660,7 @@ 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), @@ -686,13 +688,14 @@ 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), is_missing=False, - deleted_at=None, updated_at=now, + asset_id=asset_id, mtime_ns=int(mtime_ns), loader_path=loader_path, + is_missing=False, deleted_at=None, updated_at=now, ) ) res2 = session.execute(upd) diff --git a/app/assets/database/queries/tags.py b/app/assets/database/queries/tags.py index d41d73a10..148f34801 100644 --- a/app/assets/database/queries/tags.py +++ b/app/assets/database/queries/tags.py @@ -265,6 +265,8 @@ 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"), @@ -293,9 +295,8 @@ def list_tags_with_usage( .join(counts_sq, counts_sq.c.tag_name == Tag.name, isouter=True) ) - if prefix: - escaped, esc = escape_sql_like_string(prefix.strip().lower()) - q = q.where(Tag.name.like(escaped + "%", escape=esc)) + if prefix_filter: + q = q.where(func.substr(Tag.name, 1, len(prefix_filter)) == prefix_filter) if not include_zero: q = q.where(func.coalesce(counts_sq.c.cnt, 0) > 0) @@ -306,9 +307,8 @@ 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: - escaped, esc = escape_sql_like_string(prefix.strip().lower()) - total_q = total_q.where(Tag.name.like(escaped + "%", escape=esc)) + if prefix_filter: + total_q = total_q.where(func.substr(Tag.name, 1, len(prefix_filter)) == prefix_filter) if not include_zero: visible_tags_sq = ( select(AssetReferenceTag.tag_name) diff --git a/app/assets/helpers.py b/app/assets/helpers.py index 3798f3933..87734d0dc 100644 --- a/app/assets/helpers.py +++ b/app/assets/helpers.py @@ -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 and converting to lowercase. - - Removing duplicates. + - Stripping whitespace. + - Removing exact duplicates while preserving order and case. """ - return list(dict.fromkeys(t.strip().lower() for t in (tags or []) if (t or "").strip())) + return list(dict.fromkeys(t.strip() for t in (tags or []) if (t or "").strip())) def validate_blake3_hash(s: str) -> str: diff --git a/app/assets/scanner.py b/app/assets/scanner.py index 2c1e97840..42c4c1e9d 100644 --- a/app/assets/scanner.py +++ b/app/assets/scanner.py @@ -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_relative_filename, + compute_loader_path, 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 in get_comfy_models_folders(): + for _bucket, paths, _exts 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 in get_comfy_models_folders(): + for folder_name, bases, _exts 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_relative_filename(abs_p) + rel_fname = compute_loader_path(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_relative_filename(file_path) + rel_fname = compute_loader_path(file_path) mime_type: str | None = None metadata = None diff --git a/app/assets/services/asset_management.py b/app/assets/services/asset_management.py index d4e4fc61c..a4c8b5a75 100644 --- a/app/assets/services/asset_management.py +++ b/app/assets/services/asset_management.py @@ -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_relative_filename +from app.assets.services.path_utils import compute_loader_path 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_relative_filename(ref.file_path) if ref.file_path else None + computed_filename = compute_loader_path(ref.file_path) if ref.file_path else None new_meta: dict | None = None if user_metadata is not None: diff --git a/app/assets/services/bulk_ingest.py b/app/assets/services/bulk_ingest.py index 67aad838f..c98658bf1 100644 --- a/app/assets/services/bulk_ingest.py +++ b/app/assets/services/bulk_ingest.py @@ -56,6 +56,7 @@ class ReferenceRow(TypedDict): id: str asset_id: str file_path: str + loader_path: str | None mtime_ns: int owner_id: str name: str @@ -134,6 +135,14 @@ 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) @@ -164,6 +173,8 @@ 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"], diff --git a/app/assets/services/ingest.py b/app/assets/services/ingest.py index 3b6dc237c..1ffb3d634 100644 --- a/app/assets/services/ingest.py +++ b/app/assets/services/ingest.py @@ -33,8 +33,9 @@ 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_relative_filename, + compute_loader_path, get_name_and_tags_from_asset_path, + get_path_derived_tags_from_path, resolve_destination_from_tags, validate_path_within_base, ) @@ -91,6 +92,7 @@ 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 @@ -101,17 +103,32 @@ def _ingest_file_from_path( if preview_id and ref.preview_id != preview_id: ref.preview_id = preview_id - norm = normalize_tags(list(tags)) - if norm: + 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: if 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, - ) + 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, + ) _update_metadata_with_filename( session, @@ -228,7 +245,7 @@ def ingest_existing_file( "mtime_ns": mtime_ns, "info_name": name, "tags": tags, - "fname": os.path.basename(abs_path), + "fname": compute_loader_path(abs_path), "metadata": None, "hash": None, "mime_type": mime_type, @@ -288,7 +305,7 @@ def _register_existing_asset( return result new_meta = dict(user_metadata) - computed_filename = compute_relative_filename(ref.file_path) if ref.file_path else None + computed_filename = compute_loader_path(ref.file_path) if ref.file_path else None if computed_filename: new_meta["filename"] = computed_filename @@ -335,7 +352,7 @@ def _update_metadata_with_filename( current_metadata: dict | None, user_metadata: dict[str, Any], ) -> None: - computed_filename = compute_relative_filename(file_path) if file_path else None + computed_filename = compute_loader_path(file_path) if file_path else None current_meta = current_metadata or {} new_meta = dict(current_meta) @@ -474,6 +491,10 @@ 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) @@ -535,7 +556,7 @@ def upload_from_temp_path( owner_id=owner_id, preview_id=preview_id, user_metadata=user_metadata or {}, - tags=tags, + tags=[*(tags or []), "uploaded"], tag_origin="manual", require_existing_tags=False, ) @@ -569,15 +590,19 @@ def register_file_in_place( ) -> UploadResult: """Register an already-saved file in the asset database without moving it. - Tags are derived from the filesystem path (root category + subfolder names), - merged with any caller-provided tags, matching the behavior of the scanner. + 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. 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]) + merged_tags = normalize_tags([*path_tags, *tags, "uploaded"]) try: digest, _ = hashing.compute_blake3_hash(abs_path) diff --git a/app/assets/services/path_utils.py b/app/assets/services/path_utils.py index 892140ffb..7c27c8878 100644 --- a/app/assets/services/path_utils.py +++ b/app/assets/services/path_utils.py @@ -3,59 +3,66 @@ from pathlib import Path from typing import Literal import folder_paths -from app.assets.helpers import normalize_tags -_NON_MODEL_FOLDER_NAMES = frozenset({"custom_nodes"}) +_NON_MODEL_FOLDER_NAMES = frozenset({"configs", "custom_nodes"}) +_KNOWN_SUBFOLDER_TAGS = frozenset({"3d", "pasted", "painter", "threed", "webcam"}) -def get_comfy_models_folders() -> list[tuple[str, list[str]]]: - """Build list of (folder_name, base_paths[]) for all model locations. +def get_comfy_models_folders() -> list[tuple[str, list[str], set[str]]]: + """Build list of (folder_name, base_paths[], extensions) 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 custom_nodes. + 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. """ - targets: list[tuple[str, list[str]]] = [] + targets: list[tuple[str, list[str], set[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)) + targets.append((name, paths, set(exts))) return targets def resolve_destination_from_tags(tags: list[str]) -> tuple[str, list[str]]: - """Validates and maps tags -> (base_dir, subdirs_for_fs)""" - if not tags: - raise ValueError("tags must not be empty") - root = tags[0].lower() + """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] if root == "models": - if len(tags) < 2: - raise ValueError("at least two tags required for model asset") + 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: tag") + folder_name = model_type_tags[0].split(":", 1)[1] + if not folder_name: + raise ValueError("models uploads require exactly one model_type: tag") + model_folder_paths = { + name: paths for name, paths, _exts in get_comfy_models_folders() + } try: - bases = folder_paths.folder_names_and_paths[tags[1]][0] + bases = model_folder_paths[folder_name] except KeyError: - raise ValueError(f"unknown model category '{tags[1]}'") + raise ValueError(f"unknown model category '{folder_name}'") if not bases: - raise ValueError(f"no base path configured for category '{tags[1]}'") + raise ValueError(f"no base path configured for category '{folder_name}'") base_dir = os.path.abspath(bases[0]) - raw_subdirs = tags[2:] elif root == "input": base_dir = os.path.abspath(folder_paths.get_input_directory()) - 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") + base_dir = os.path.abspath(folder_paths.get_output_directory()) - return base_dir, raw_subdirs if raw_subdirs else [] + return base_dir, [] def validate_path_within_base(candidate: str, base: str) -> None: @@ -65,14 +72,79 @@ def validate_path_within_base(candidate: str, base: str) -> None: raise ValueError("destination escapes base directory") -def compute_relative_filename(file_path: str) -> str | None: +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:``; these paths only use known storage roots. """ - Return the model's path relative to the last well-known folder (the model category), - using forward slashes, eg: + 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: /.../models/checkpoints/flux/123/flux.safetensors -> "flux/123/flux.safetensors" /.../models/text_encoders/clip_g.safetensors -> "clip_g.safetensors" - For non-model paths, returns None. + 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. """ try: root_category, rel_path = get_asset_category_and_relative_path(file_path) @@ -116,9 +188,10 @@ 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. - return os.path.relpath( + rel = 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()) @@ -136,8 +209,14 @@ 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 in get_comfy_models_folders(): + 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 b in bases: base_abs = os.path.abspath(b) if not _check_is_within(fp_abs, base_abs): @@ -149,25 +228,111 @@ def get_asset_category_and_relative_path( if best is not None: _, bucket, rel_inside = best combined = os.path.join(bucket, rel_inside) - return "models", os.path.relpath(os.path.join(os.sep, combined), os.sep) + normalized = os.path.relpath(os.path.join(os.sep, combined), os.sep) + return "models", normalized.replace(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:``, ``input``, ``output``, and ``temp``. Model + type tags are based on registered folder names, not path components. + + A ``model_type:`` 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: [root_category] + parent folder names in order + - tags: backend-derived tags from root/model classification and known input + subfolder layout conventions Raises: ValueError: path does not belong to any known root. """ - 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]))) + return Path(file_path).name, get_path_derived_tags_from_path(file_path) diff --git a/app/assets/services/schemas.py b/app/assets/services/schemas.py index 4d2af8a02..0fda6871d 100644 --- a/app/assets/services/schemas.py +++ b/app/assets/services/schemas.py @@ -25,6 +25,7 @@ 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 @@ -93,6 +94,7 @@ 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, diff --git a/app/model_manager.py b/app/model_manager.py index 8f6e34b33..b0329ce17 100644 --- a/app/model_manager.py +++ b/app/model_manager.py @@ -50,21 +50,45 @@ class ModelFileManager: @routes.get("/experiment/models/preview/{folder}/{path_index}/{filename:.*}") async def get_model_preview(request): folder_name = request.match_info.get("folder", None) - path_index = int(request.match_info.get("path_index", None)) filename = request.match_info.get("filename", None) if folder_name not in folder_paths.folder_names_and_paths: return web.Response(status=404) + # The "{filename:.*}" capture also matches the empty string, which + # would resolve to the folder itself; reject it explicitly. + if not filename: + return web.Response(status=400) + + try: + path_index = int(request.match_info.get("path_index", None)) + except (TypeError, ValueError): + return web.Response(status=400) + folders = folder_paths.folder_names_and_paths[folder_name] + if path_index < 0 or path_index >= len(folders[0]): + return web.Response(status=404) folder = folders[0][path_index] - full_filename = os.path.join(folder, filename) + full_filename = os.path.normpath(os.path.join(folder, filename)) + + # Prevent path traversal: the requested file must stay within the + # configured model folder. `filename` is an unrestricted ".*" capture, + # so values like "../../../../etc/passwd" would otherwise escape it. + if not folder_paths.is_within_directory(folder, full_filename): + return web.Response(status=403) previews = self.get_model_previews(full_filename) default_preview = previews[0] if len(previews) > 0 else None if default_preview is None or (isinstance(default_preview, str) and not os.path.isfile(default_preview)): return web.Response(status=404) + # The preview is selected by a glob inside get_model_previews, so a + # companion file (e.g. "model.preview.png") could itself be a symlink + # resolving outside the model folder. Re-validate the file actually + # opened: is_within_directory realpaths it, catching symlink escape. + if isinstance(default_preview, str) and not folder_paths.is_within_directory(folder, default_preview): + return web.Response(status=403) + try: with Image.open(default_preview) as img: img_bytes = BytesIO() diff --git a/app/user_manager.py b/app/user_manager.py index 7b11e381c..de261ad39 100644 --- a/app/user_manager.py +++ b/app/user_manager.py @@ -6,6 +6,7 @@ import glob import shutil import logging import tempfile +import mimetypes from aiohttp import web from urllib import parse from comfy.cli_args import args @@ -336,7 +337,20 @@ class UserManager(): if not isinstance(path, str): return path - return web.FileResponse(path) + # User data files are arbitrary user-supplied content and are never + # meant to render inline. Disable MIME sniffing and force a download + # so uploaded markup/scripts can't execute in the app origin (stored + # XSS). Content-Disposition: attachment is the load-bearing guard; + # the content-type override and nosniff are defence in depth. + content_type = mimetypes.guess_type(path)[0] or 'application/octet-stream' + if folder_paths.is_dangerous_content_type(content_type): + content_type = 'application/octet-stream' + + return web.FileResponse(path, headers={ + "Content-Type": content_type, + "X-Content-Type-Options": "nosniff", + "Content-Disposition": "attachment", + }) @routes.post("/userdata/{file}") async def post_userdata(request): diff --git a/comfy/cli_args.py b/comfy/cli_args.py index 0c82fb8b8..2e35dc679 100644 --- a/comfy/cli_args.py +++ b/comfy/cli_args.py @@ -225,6 +225,7 @@ parser.add_argument( ) parser.add_argument("--user-directory", type=is_valid_directory, default=None, help="Set the ComfyUI user directory with an absolute path. Overrides --base-directory.") +parser.add_argument("--models-directory", type=is_valid_directory, default=None, help="Set the ComfyUI models directory. Overrides the models folder in --base-directory.") parser.add_argument("--enable-compress-response-body", action="store_true", help="Enable compressing response body.") diff --git a/comfy/ldm/ace/ace_step15.py b/comfy/ldm/ace/ace_step15.py index 2ca2d26c4..02182c49f 100644 --- a/comfy/ldm/ace/ace_step15.py +++ b/comfy/ldm/ace/ace_step15.py @@ -217,10 +217,7 @@ class AceStepAttention(nn.Module): cos, sin = position_embeddings query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin) - n_rep = self.num_heads // self.num_kv_heads - if n_rep > 1: - key_states = key_states.repeat_interleave(n_rep, dim=1) - value_states = value_states.repeat_interleave(n_rep, dim=1) + gqa_kwargs = {"enable_gqa": True} if self.num_heads != self.num_kv_heads else {} attn_bias = None if self.sliding_window is not None and not self.is_cross_attention: @@ -244,7 +241,7 @@ class AceStepAttention(nn.Module): else: attn_bias = window_bias - attn_output = optimized_attention(query_states, key_states, value_states, self.num_heads, attn_bias, skip_reshape=True, low_precision_attention=False) + attn_output = optimized_attention(query_states, key_states, value_states, self.num_heads, attn_bias, skip_reshape=True, low_precision_attention=False, **gqa_kwargs) attn_output = self.o_proj(attn_output) return attn_output diff --git a/comfy/ldm/audio/dit.py b/comfy/ldm/audio/dit.py index c28be5b49..b0759a240 100644 --- a/comfy/ldm/audio/dit.py +++ b/comfy/ldm/audio/dit.py @@ -425,19 +425,16 @@ class Attention(nn.Module): if n == 1 and causal: causal = False - if h != kv_h: - # Repeat interleave kv_heads to match q_heads - heads_per_kv_head = h // kv_h - k, v = map(lambda t: t.repeat_interleave(heads_per_kv_head, dim = 1), (k, v)) + gqa_kwargs = {"enable_gqa": True} if h != kv_h else {} if self.differential: q, q_diff = q.unbind(dim=1) k, k_diff = k.unbind(dim=1) - out = optimized_attention(q, k, v, h, skip_reshape=True, low_precision_attention=False, transformer_options=transformer_options) - out_diff = optimized_attention(q_diff, k_diff, v, h, skip_reshape=True, low_precision_attention=False, transformer_options=transformer_options) + out = optimized_attention(q, k, v, h, skip_reshape=True, low_precision_attention=False, transformer_options=transformer_options, **gqa_kwargs) + out_diff = optimized_attention(q_diff, k_diff, v, h, skip_reshape=True, low_precision_attention=False, transformer_options=transformer_options, **gqa_kwargs) out = out - out_diff else: - out = optimized_attention(q, k, v, h, skip_reshape=True, low_precision_attention=False, transformer_options=transformer_options) + out = optimized_attention(q, k, v, h, skip_reshape=True, low_precision_attention=False, transformer_options=transformer_options, **gqa_kwargs) out = self.to_out(out) diff --git a/comfy/ldm/boogu/model.py b/comfy/ldm/boogu/model.py index 966f3c583..ca88bdeb1 100644 --- a/comfy/ldm/boogu/model.py +++ b/comfy/ldm/boogu/model.py @@ -74,11 +74,8 @@ class BooguDoubleStreamProcessor(nn.Module): key = key.transpose(1, 2) value = value.transpose(1, 2) - if attn.kv_heads < attn.heads: - key = key.repeat_interleave(attn.heads // attn.kv_heads, dim=1) - value = value.repeat_interleave(attn.heads // attn.kv_heads, dim=1) - - hidden_states = optimized_attention_masked(query, key, value, attn.heads, attention_mask, skip_reshape=True, transformer_options=transformer_options) + gqa_kwargs = {"enable_gqa": True} if attn.kv_heads < attn.heads else {} + hidden_states = optimized_attention_masked(query, key, value, attn.heads, attention_mask, skip_reshape=True, transformer_options=transformer_options, **gqa_kwargs) # Split back to instruction/image, apply per-stream output projections, recombine. instruct_hidden_states = self.instruct_out(hidden_states[:, :L_instruct]) diff --git a/comfy/ldm/modules/attention.py b/comfy/ldm/modules/attention.py index 55360535a..2411aff5c 100644 --- a/comfy/ldm/modules/attention.py +++ b/comfy/ldm/modules/attention.py @@ -1,5 +1,6 @@ import math import sys +import inspect import torch import torch.nn.functional as F @@ -14,16 +15,16 @@ from .sub_quadratic_attention import efficient_dot_product_attention from comfy import model_management -TORCH_HAS_GQA = model_management.torch_version_numeric >= (2, 5) - if model_management.xformers_enabled(): import xformers import xformers.ops SAGE_ATTENTION_IS_AVAILABLE = False +SAGE_ATTENTION_SUPPORTS_MASK = False try: from sageattention import sageattn SAGE_ATTENTION_IS_AVAILABLE = True + SAGE_ATTENTION_SUPPORTS_MASK = "attn_mask" in inspect.signature(sageattn).parameters except ImportError as e: if model_management.sage_attention_enabled(): if e.name == "sageattention": @@ -89,6 +90,44 @@ def default(val, d): return val return d +def _gqa_repeat_factor(query_heads, key_heads, value_heads): + if key_heads != value_heads: + raise ValueError(f"Key/value head count mismatch for GQA: {key_heads} != {value_heads}") + if query_heads == key_heads: + return 1 + if query_heads % key_heads != 0: + raise ValueError(f"Query heads must be divisible by key/value heads for GQA: {query_heads} vs {key_heads}") + return query_heads // key_heads + +def _repeat_kv_for_gqa(k, v, query_heads, head_dim): + n_rep = _gqa_repeat_factor(query_heads, k.shape[head_dim], v.shape[head_dim]) + if n_rep > 1: + k = k.repeat_interleave(n_rep, dim=head_dim) + v = v.repeat_interleave(n_rep, dim=head_dim) + return k, v + +def _heads_from_dim(tensor, dim_head, name): + inner_dim = tensor.shape[-1] + if inner_dim % dim_head != 0: + raise ValueError(f"{name} inner dimension {inner_dim} is not divisible by head dimension {dim_head}") + return inner_dim // dim_head + +def _reshape_qkv_to_heads(q, k, v, b, heads, dim_head, enable_gqa=False, expand_kv=True): + q = q.unsqueeze(3).reshape(b, -1, heads, dim_head) + if enable_gqa: + key_heads = _heads_from_dim(k, dim_head, "Key") + value_heads = _heads_from_dim(v, dim_head, "Value") + else: + key_heads = heads + value_heads = heads + k = k.unsqueeze(3).reshape(b, -1, key_heads, dim_head) + v = v.unsqueeze(3).reshape(b, -1, value_heads, dim_head) + if enable_gqa: + _gqa_repeat_factor(heads, key_heads, value_heads) + if expand_kv: + k, v = _repeat_kv_for_gqa(k, v, heads, -2) + return q, k, v + # feedforward class GEGLU(nn.Module): @@ -152,28 +191,19 @@ def attention_basic(q, k, v, heads, mask=None, attn_precision=None, skip_reshape b, _, dim_head = q.shape dim_head //= heads - if kwargs.get("enable_gqa", False) and q.shape[-3] != k.shape[-3]: - n_rep = q.shape[-3] // k.shape[-3] - k = k.repeat_interleave(n_rep, dim=-3) - v = v.repeat_interleave(n_rep, dim=-3) - scale = kwargs.get("scale", dim_head ** -0.5) h = heads if skip_reshape: - q, k, v = map( + if kwargs.get("enable_gqa", False): + k, v = _repeat_kv_for_gqa(k, v, q.shape[-3], -3) + q, k, v = map( lambda t: t.reshape(b * heads, -1, dim_head), (q, k, v), ) else: - q, k, v = map( - lambda t: t.unsqueeze(3) - .reshape(b, -1, heads, dim_head) - .permute(0, 2, 1, 3) - .reshape(b * heads, -1, dim_head) - .contiguous(), - (q, k, v), - ) + q, k, v = _reshape_qkv_to_heads(q, k, v, b, heads, dim_head, kwargs.get("enable_gqa", False)) + q, k, v = map(lambda t: t.permute(0, 2, 1, 3).reshape(b * heads, -1, dim_head).contiguous(), (q, k, v)) # force cast to fp32 to avoid overflowing if attn_precision == torch.float32: @@ -231,13 +261,16 @@ def attention_sub_quad(query, key, value, heads, mask=None, attn_precision=None, query = query * (kwargs["scale"] * dim_head ** 0.5) if skip_reshape: + if kwargs.get("enable_gqa", False): + key, value = _repeat_kv_for_gqa(key, value, query.shape[-3], -3) query = query.reshape(b * heads, -1, dim_head) value = value.reshape(b * heads, -1, dim_head) key = key.reshape(b * heads, -1, dim_head).movedim(1, 2) else: - query = query.unsqueeze(3).reshape(b, -1, heads, dim_head).permute(0, 2, 1, 3).reshape(b * heads, -1, dim_head) - value = value.unsqueeze(3).reshape(b, -1, heads, dim_head).permute(0, 2, 1, 3).reshape(b * heads, -1, dim_head) - key = key.unsqueeze(3).reshape(b, -1, heads, dim_head).permute(0, 2, 3, 1).reshape(b * heads, dim_head, -1) + query, key, value = _reshape_qkv_to_heads(query, key, value, b, heads, dim_head, kwargs.get("enable_gqa", False)) + query = query.permute(0, 2, 1, 3).reshape(b * heads, -1, dim_head) + value = value.permute(0, 2, 1, 3).reshape(b * heads, -1, dim_head) + key = key.permute(0, 2, 3, 1).reshape(b * heads, dim_head, -1) dtype = query.dtype @@ -304,19 +337,15 @@ def attention_split(q, k, v, heads, mask=None, attn_precision=None, skip_reshape scale = kwargs.get("scale", dim_head ** -0.5) if skip_reshape: - q, k, v = map( + if kwargs.get("enable_gqa", False): + k, v = _repeat_kv_for_gqa(k, v, q.shape[-3], -3) + q, k, v = map( lambda t: t.reshape(b * heads, -1, dim_head), (q, k, v), ) else: - q, k, v = map( - lambda t: t.unsqueeze(3) - .reshape(b, -1, heads, dim_head) - .permute(0, 2, 1, 3) - .reshape(b * heads, -1, dim_head) - .contiguous(), - (q, k, v), - ) + q, k, v = _reshape_qkv_to_heads(q, k, v, b, heads, dim_head, kwargs.get("enable_gqa", False)) + q, k, v = map(lambda t: t.permute(0, 2, 1, 3).reshape(b * heads, -1, dim_head).contiguous(), (q, k, v)) r1 = torch.zeros(q.shape[0], q.shape[1], v.shape[2], device=q.device, dtype=q.dtype) @@ -438,7 +467,7 @@ def attention_xformers(q, k, v, heads, mask=None, attn_precision=None, skip_resh disabled_xformers = True if disabled_xformers: - return attention_pytorch(q, k, v, heads, mask, skip_reshape=skip_reshape, **kwargs) + return attention_pytorch(q, k, v, heads, mask, skip_reshape=skip_reshape, skip_output_reshape=skip_output_reshape, **kwargs) if skip_reshape: # b h k d -> b k h d @@ -446,13 +475,12 @@ def attention_xformers(q, k, v, heads, mask=None, attn_precision=None, skip_resh lambda t: t.permute(0, 2, 1, 3), (q, k, v), ) + if kwargs.get("enable_gqa", False): + k, v = _repeat_kv_for_gqa(k, v, q.shape[-2], -2) # actually do the reshaping else: dim_head //= heads - q, k, v = map( - lambda t: t.reshape(b, -1, heads, dim_head), - (q, k, v), - ) + q, k, v = _reshape_qkv_to_heads(q, k, v, b, heads, dim_head, kwargs.get("enable_gqa", False)) if mask is not None: # add a singleton batch dimension @@ -474,7 +502,7 @@ def attention_xformers(q, k, v, heads, mask=None, attn_precision=None, skip_resh mask = mask_out[..., :mask.shape[-1]] mask = mask.expand(b, heads, -1, -1) - out = xformers.ops.memory_efficient_attention(q, k, v, attn_bias=mask) + out = xformers.ops.memory_efficient_attention(q, k, v, attn_bias=mask, scale=kwargs.get("scale", None)) if skip_output_reshape: out = out.permute(0, 2, 1, 3) @@ -498,10 +526,8 @@ def attention_pytorch(q, k, v, heads, mask=None, attn_precision=None, skip_resha else: b, _, dim_head = q.shape dim_head //= heads - q, k, v = map( - lambda t: t.view(b, -1, heads, dim_head).transpose(1, 2), - (q, k, v), - ) + q, k, v = _reshape_qkv_to_heads(q, k, v, b, heads, dim_head, kwargs.get("enable_gqa", False), expand_kv=False) + q, k, v = map(lambda t: t.transpose(1, 2), (q, k, v)) if mask is not None: # add a batch dimension if there isn't already one @@ -511,9 +537,7 @@ def attention_pytorch(q, k, v, heads, mask=None, attn_precision=None, skip_resha if mask.ndim == 3: mask = mask.unsqueeze(1) - # Pass through extra SDPA kwargs (scale, enable_gqa) if provided - # enable_gqa requires PyTorch 2.5+; older versions use manual KV expansion above - sdpa_keys = ("scale", "enable_gqa") if TORCH_HAS_GQA else ("scale",) + sdpa_keys = ("scale", "enable_gqa") sdpa_extra = {k: v for k, v in kwargs.items() if k in sdpa_keys} if SDP_BATCH_LIMIT >= b: @@ -541,20 +565,19 @@ def attention_pytorch(q, k, v, heads, mask=None, attn_precision=None, skip_resha @wrap_attn def attention_sage(q, k, v, heads, mask=None, attn_precision=None, skip_reshape=False, skip_output_reshape=False, **kwargs): - if kwargs.get("low_precision_attention", True) is False: + if kwargs.get("low_precision_attention", True) is False or (mask is not None and not SAGE_ATTENTION_SUPPORTS_MASK): return attention_pytorch(q, k, v, heads, mask=mask, skip_reshape=skip_reshape, skip_output_reshape=skip_output_reshape, **kwargs) exception_fallback = False if skip_reshape: b, _, _, dim_head = q.shape tensor_layout = "HND" + if kwargs.get("enable_gqa", False): + k, v = _repeat_kv_for_gqa(k, v, q.shape[-3], -3) else: b, _, dim_head = q.shape dim_head //= heads - q, k, v = map( - lambda t: t.view(b, -1, heads, dim_head), - (q, k, v), - ) + q, k, v = _reshape_qkv_to_heads(q, k, v, b, heads, dim_head, kwargs.get("enable_gqa", False)) tensor_layout = "NHD" if mask is not None: @@ -565,8 +588,12 @@ def attention_sage(q, k, v, heads, mask=None, attn_precision=None, skip_reshape= if mask.ndim == 3: mask = mask.unsqueeze(1) + sage_kwargs = {"is_causal": False, "tensor_layout": tensor_layout, "sm_scale": kwargs.get("scale", None), "smooth_k": False} + if mask is not None: + sage_kwargs["attn_mask"] = mask + try: - out = sageattn(q, k, v, attn_mask=mask, is_causal=False, tensor_layout=tensor_layout) + out = sageattn(q, k, v, **sage_kwargs) except Exception as e: logging.error("Error running sage attention: {}, using pytorch attention instead.".format(e)) exception_fallback = True @@ -616,7 +643,6 @@ def attention3_sage(q, k, v, heads, mask=None, attn_precision=None, skip_reshape skip_output_reshape=skip_output_reshape, **kwargs ) - q_s, k_s, v_s = q, k, v N = q.shape[2] dim_head = D else: @@ -642,11 +668,15 @@ def attention3_sage(q, k, v, heads, mask=None, attn_precision=None, skip_reshape **kwargs ) - if not skip_reshape: - q_s, k_s, v_s = map( - lambda t: t.view(B, -1, heads, dim_head).permute(0, 2, 1, 3).contiguous(), - (q, k, v), - ) + if skip_reshape: + q_s = q + if kwargs.get("enable_gqa", False): + k_s, v_s = _repeat_kv_for_gqa(k, v, H, -3) + else: + k_s, v_s = k, v + else: + q_s, k_s, v_s = _reshape_qkv_to_heads(q, k, v, B, heads, dim_head, kwargs.get("enable_gqa", False)) + q_s, k_s, v_s = map(lambda t: t.permute(0, 2, 1, 3).contiguous(), (q_s, k_s, v_s)) B, H, L, D = q_s.shape try: @@ -662,7 +692,7 @@ def attention3_sage(q, k, v, heads, mask=None, attn_precision=None, skip_reshape q, k, v, heads, mask=mask, attn_precision=attn_precision, - skip_reshape=False, + skip_reshape=skip_reshape, skip_output_reshape=skip_output_reshape, **kwargs ) @@ -681,19 +711,20 @@ def attention3_sage(q, k, v, heads, mask=None, attn_precision=None, skip_reshape try: @torch.library.custom_op("flash_attention::flash_attn", mutates_args=()) def flash_attn_wrapper(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, - dropout_p: float = 0.0, causal: bool = False) -> torch.Tensor: - return flash_attn_func(q, k, v, dropout_p=dropout_p, causal=causal) + dropout_p: float = 0.0, causal: bool = False, softmax_scale: float = -1.0) -> torch.Tensor: + softmax_scale_arg = None if softmax_scale == -1.0 else softmax_scale + return flash_attn_func(q, k, v, dropout_p=dropout_p, causal=causal, softmax_scale=softmax_scale_arg) @flash_attn_wrapper.register_fake - def flash_attn_fake(q, k, v, dropout_p=0.0, causal=False): + def flash_attn_fake(q, k, v, dropout_p=0.0, causal=False, softmax_scale=-1.0): # Output shape is the same as q return q.new_empty(q.shape) except AttributeError as error: FLASH_ATTN_ERROR = error def flash_attn_wrapper(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, - dropout_p: float = 0.0, causal: bool = False) -> torch.Tensor: + dropout_p: float = 0.0, causal: bool = False, softmax_scale: float = -1.0) -> torch.Tensor: assert False, f"Could not define flash_attn_wrapper: {FLASH_ATTN_ERROR}" @wrap_attn @@ -703,10 +734,8 @@ def attention_flash(q, k, v, heads, mask=None, attn_precision=None, skip_reshape else: b, _, dim_head = q.shape dim_head //= heads - q, k, v = map( - lambda t: t.view(b, -1, heads, dim_head).transpose(1, 2), - (q, k, v), - ) + q, k, v = _reshape_qkv_to_heads(q, k, v, b, heads, dim_head, kwargs.get("enable_gqa", False), expand_kv=False) + q, k, v = map(lambda t: t.transpose(1, 2), (q, k, v)) if mask is not None: # add a batch dimension if there isn't already one @@ -725,10 +754,16 @@ def attention_flash(q, k, v, heads, mask=None, attn_precision=None, skip_reshape v.transpose(1, 2), dropout_p=0.0, causal=False, + softmax_scale=kwargs.get("scale", -1.0), ).transpose(1, 2) except Exception as e: logging.warning(f"Flash Attention failed, using default SDPA: {e}") - out = torch.nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=mask, dropout_p=0.0, is_causal=False) + sdpa_extra = {} + if kwargs.get("enable_gqa", False): + sdpa_extra["enable_gqa"] = True + if "scale" in kwargs: + sdpa_extra["scale"] = kwargs["scale"] + out = torch.nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=mask, dropout_p=0.0, is_causal=False, **sdpa_extra) if not skip_output_reshape: out = ( out.transpose(1, 2).reshape(b, -1, heads * dim_head) @@ -1209,5 +1244,3 @@ class SpatialVideoTransformer(SpatialTransformer): x = self.proj_out(x) out = x + x_in return out - - diff --git a/comfy/ldm/omnigen/omnigen2.py b/comfy/ldm/omnigen/omnigen2.py index b8da4cf39..d18a9f6d0 100644 --- a/comfy/ldm/omnigen/omnigen2.py +++ b/comfy/ldm/omnigen/omnigen2.py @@ -141,11 +141,8 @@ class Attention(nn.Module): key = key.transpose(1, 2) value = value.transpose(1, 2) - if self.kv_heads < self.heads: - key = key.repeat_interleave(self.heads // self.kv_heads, dim=1) - value = value.repeat_interleave(self.heads // self.kv_heads, dim=1) - - hidden_states = optimized_attention_masked(query, key, value, self.heads, attention_mask, skip_reshape=True, transformer_options=transformer_options) + gqa_kwargs = {"enable_gqa": True} if self.kv_heads < self.heads else {} + hidden_states = optimized_attention_masked(query, key, value, self.heads, attention_mask, skip_reshape=True, transformer_options=transformer_options, **gqa_kwargs) hidden_states = self.to_out[0](hidden_states) return hidden_states diff --git a/comfy/ops.py b/comfy/ops.py index 69d32e254..35a1ee31e 100644 --- a/comfy/ops.py +++ b/comfy/ops.py @@ -174,6 +174,8 @@ def cast_modules_with_vbar(comfy_modules, dtype, device, bias_dtype, non_blockin elif xfer_dest2 is not None: xfer_source.prepare(xfer_dest2, stream, copy=True, commit=False) return + else: + return comfy.model_management.cast_to_gathered(xfer_source, xfer_dest, non_blocking=non_blocking, stream=stream, r2=xfer_dest2) def handle_pin(m, pin, source, dest, subset="weights", size=None): diff --git a/comfy/sd.py b/comfy/sd.py index 610c4e2b8..071a3102a 100644 --- a/comfy/sd.py +++ b/comfy/sd.py @@ -468,6 +468,9 @@ class CLIP: def decode(self, token_ids, skip_special_tokens=True): return self.tokenizer.decode(token_ids, skip_special_tokens=skip_special_tokens) + def is_dynamic(self): + return self.patcher.is_dynamic() + class VAE: def __init__(self, sd=None, device=None, config=None, dtype=None, metadata=None): if 'decoder.up_blocks.0.resnets.0.norm1.weight' in sd.keys(): #diffusers format @@ -1251,6 +1254,11 @@ class VAE: except: return None + def is_dynamic(self): + # A VAE built from a state dict with no detectable VAE weights returns early + # from __init__ ("No VAE weights detected") before self.patcher is assigned. + patcher = getattr(self, "patcher", None) + return patcher is not None and patcher.is_dynamic() class StyleModel: def __init__(self, model, device="cpu"): diff --git a/comfy/sd1_clip.py b/comfy/sd1_clip.py index 897186bba..f0fdf1aa5 100644 --- a/comfy/sd1_clip.py +++ b/comfy/sd1_clip.py @@ -543,18 +543,24 @@ class SDTokenizer: def _try_get_embedding(self, embedding_name:str): ''' Takes a potential embedding name and tries to retrieve it. - Returns a Tuple consisting of the embedding and any leftover string, embedding can be None. + Returns a Tuple consisting of the embedding, the cleaned embedding name, and any leftover string, embedding can be None. ''' split_embed = embedding_name.split() embedding_name = split_embed[0] leftover = ' '.join(split_embed[1:]) + + match = re.search(r'[<\[]', embedding_name) + if match is not None: + leftover = embedding_name[match.start():] + (" " + leftover if leftover else "") + embedding_name = embedding_name[:match.start()] + embed = load_embed(embedding_name, self.embedding_directory, self.embedding_size, self.embedding_key) if embed is None: stripped = embedding_name.strip(',') if len(stripped) < len(embedding_name): embed = load_embed(stripped, self.embedding_directory, self.embedding_size, self.embedding_key) - return (embed, "{} {}".format(embedding_name[len(stripped):], leftover)) - return (embed, leftover) + return (embed, embedding_name, "{} {}".format(embedding_name[len(stripped):], leftover)) + return (embed, embedding_name, leftover) def pad_tokens(self, tokens, amount): if self.pad_left: @@ -585,7 +591,7 @@ class SDTokenizer: tokens = [] for weighted_segment, weight in parsed_weights: to_tokenize = unescape_important(weighted_segment) - split = re.split(' {0}|\n{0}'.format(self.embedding_identifier), to_tokenize) + split = re.split(r'(?<=\s){}'.format(re.escape(self.embedding_identifier)), to_tokenize) to_tokenize = [split[0]] for i in range(1, len(split)): to_tokenize.append("{}{}".format(self.embedding_identifier, split[i])) @@ -595,7 +601,7 @@ class SDTokenizer: # if we find an embedding, deal with the embedding if word.startswith(self.embedding_identifier) and self.embedding_directory is not None: embedding_name = word[len(self.embedding_identifier):].strip('\n') - embed, leftover = self._try_get_embedding(embedding_name) + embed, embedding_name, leftover = self._try_get_embedding(embedding_name) if embed is None: logging.warning(f"warning, embedding:{embedding_name} does not exist, ignoring") else: diff --git a/comfy/text_encoders/gpt_oss.py b/comfy/text_encoders/gpt_oss.py index d596ef9a0..066796b6a 100644 --- a/comfy/text_encoders/gpt_oss.py +++ b/comfy/text_encoders/gpt_oss.py @@ -12,7 +12,7 @@ import torch.nn.functional as F import comfy.ops from comfy import sd1_clip -from comfy.ldm.modules.attention import TORCH_HAS_GQA, optimized_attention_for_device +from comfy.ldm.modules.attention import optimized_attention_for_device from comfy.text_encoders.llama import RMSNorm, apply_rope @@ -110,10 +110,6 @@ def _attention_with_sinks(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, sin putting the sink logit in the mask at that column. """ - if num_kv_groups > 1 and not TORCH_HAS_GQA: - k = k.repeat_interleave(num_kv_groups, dim=1) - v = v.repeat_interleave(num_kv_groups, dim=1) - B, _, S_q, D = q.shape H_kv = k.shape[1] S_kv = k.shape[-2] diff --git a/comfy/text_encoders/llama.py b/comfy/text_encoders/llama.py index e9f38a9a2..3f98fb0a5 100644 --- a/comfy/text_encoders/llama.py +++ b/comfy/text_encoders/llama.py @@ -550,10 +550,8 @@ class Attention(nn.Module): xv = xv[:, :, -sliding_window:] attention_mask = attention_mask[..., -sliding_window:] if attention_mask is not None else None - xk = xk.repeat_interleave(self.num_heads // self.num_kv_heads, dim=1) - xv = xv.repeat_interleave(self.num_heads // self.num_kv_heads, dim=1) - - output = optimized_attention(xq, xk, xv, self.num_heads, mask=attention_mask, skip_reshape=True) + gqa_kwargs = {"enable_gqa": True} if self.num_heads != self.num_kv_heads else {} + output = optimized_attention(xq, xk, xv, self.num_heads, mask=attention_mask, skip_reshape=True, **gqa_kwargs) return self.o_proj(output), present_key_value class MLP(nn.Module): @@ -937,22 +935,41 @@ class BaseGenerate: return torch.argmax(logits, dim=-1, keepdim=True) # Sampling mode - if repetition_penalty != 1.0: - for i in range(logits.shape[0]): - for token_id in set(token_history): - logits[i, token_id] *= repetition_penalty if logits[i, token_id] < 0 else 1/repetition_penalty - - if presence_penalty is not None and presence_penalty != 0.0: - for i in range(logits.shape[0]): - for token_id in set(token_history): - logits[i, token_id] -= presence_penalty + if len(token_history) > 0 and (repetition_penalty != 1.0 or (presence_penalty is not None and presence_penalty != 0.0)): + token_ids = torch.tensor(list(set(token_history)), device=logits.device) + token_logits = logits[:, token_ids] + if repetition_penalty != 1.0: + token_logits = torch.where(token_logits < 0, token_logits * repetition_penalty, token_logits / repetition_penalty) + if presence_penalty is not None and presence_penalty != 0.0: + token_logits = token_logits - presence_penalty + logits[:, token_ids] = token_logits if temperature != 1.0: logits = logits / temperature if top_k > 0: - indices_to_remove = logits < torch.topk(logits, top_k)[0][..., -1, None] - logits[indices_to_remove] = torch.finfo(logits.dtype).min + top_k = min(top_k, logits.shape[-1]) + logits, top_indices = torch.topk(logits, top_k) + + if min_p > 0.0: + probs_before_filter = torch.nn.functional.softmax(logits, dim=-1) + top_probs, _ = probs_before_filter.max(dim=-1, keepdim=True) + min_threshold = min_p * top_probs + indices_to_remove = probs_before_filter < min_threshold + logits[indices_to_remove] = torch.finfo(logits.dtype).min + + if top_p < 1.0: + sorted_logits, sorted_indices = torch.sort(logits, descending=True) + cumulative_probs = torch.cumsum(torch.nn.functional.softmax(sorted_logits, dim=-1), dim=-1) + sorted_indices_to_remove = cumulative_probs > top_p + sorted_indices_to_remove[..., 0] = False + indices_to_remove = torch.zeros_like(logits, dtype=torch.bool) + indices_to_remove.scatter_(1, sorted_indices, sorted_indices_to_remove) + logits[indices_to_remove] = torch.finfo(logits.dtype).min + + probs = torch.nn.functional.softmax(logits, dim=-1) + next_token = torch.multinomial(probs, num_samples=1, generator=generator) + return top_indices.gather(1, next_token) if min_p > 0.0: probs_before_filter = torch.nn.functional.softmax(logits, dim=-1) diff --git a/comfy/text_encoders/qwen35.py b/comfy/text_encoders/qwen35.py index 71a17990f..304a4357f 100644 --- a/comfy/text_encoders/qwen35.py +++ b/comfy/text_encoders/qwen35.py @@ -366,12 +366,8 @@ class GatedAttention(nn.Module): xv = torch.cat((past_value[:, :, :index], xv), dim=2) present_key_value = (xk, xv, index + num_tokens) - # Expand KV heads for GQA - if self.num_heads != self.num_kv_heads: - xk = xk.repeat_interleave(self.num_heads // self.num_kv_heads, dim=1) - xv = xv.repeat_interleave(self.num_heads // self.num_kv_heads, dim=1) - - output = optimized_attention(xq, xk, xv, self.num_heads, mask=attention_mask, skip_reshape=True) + gqa_kwargs = {"enable_gqa": True} if self.num_heads != self.num_kv_heads else {} + output = optimized_attention(xq, xk, xv, self.num_heads, mask=attention_mask, skip_reshape=True, **gqa_kwargs) output = output * gate.sigmoid() return self.o_proj(output), present_key_value diff --git a/comfy/text_encoders/qwen3vl.py b/comfy/text_encoders/qwen3vl.py index 59c9aae6d..2082c42e7 100644 --- a/comfy/text_encoders/qwen3vl.py +++ b/comfy/text_encoders/qwen3vl.py @@ -167,7 +167,7 @@ class Qwen3VLTokenizer(sd1_clip.SD1Tokenizer): embed_count = 0 for r in tokens[key_name]: for i in range(len(r)): - if r[i][0] == 151655: # <|image_pad|> + if isinstance(r[i][0], (int, float)) and r[i][0] == 151655: # <|image_pad|> if len(images) > embed_count: r[i] = ({"type": "image", "data": images[embed_count], "original_type": "image"},) + r[i][1:] embed_count += 1 diff --git a/comfy_api/feature_flags.py b/comfy_api/feature_flags.py index 4f7365243..931add361 100644 --- a/comfy_api/feature_flags.py +++ b/comfy_api/feature_flags.py @@ -100,6 +100,7 @@ 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, diff --git a/comfy_api/latest/_input_impl/video_types.py b/comfy_api/latest/_input_impl/video_types.py index 6c69256ab..bc95a5b99 100644 --- a/comfy_api/latest/_input_impl/video_types.py +++ b/comfy_api/latest/_input_impl/video_types.py @@ -281,11 +281,18 @@ class VideoFromFile(VideoInput): video_done = False audio_done = True - if len(container.streams.audio): - audio_stream = container.streams.audio[-1] + # Use the last decodable audio stream. Streams FFmpeg has no decoder for have no codec context, + # and decoding their packets crashes the process. (e.g. APAC spatial-audio track in iPhone) + audio_stream = next( + (s for s in reversed(container.streams.audio) if s.codec_context is not None), + None, + ) + if audio_stream is not None: streams += [audio_stream] resampler = av.audio.resampler.AudioResampler(format='fltp') audio_done = False + elif len(container.streams.audio): + logging.warning("No decodable audio stream found in video; ignoring audio.") for packet in container.demux(*streams): if video_done and audio_done: @@ -457,10 +464,13 @@ class VideoFromFile(VideoInput): else: output_container.metadata[key] = json.dumps(value) - # Add streams to the new container + # Add streams to the new container. Streams with no codec context cannot be used as an output template. stream_map = {} for stream in streams: if isinstance(stream, (av.VideoStream, av.AudioStream, SubtitleStream)): + if stream.codec_context is None: + logging.warning("Skipping %s stream %d with unsupported codec", stream.type, stream.index) + continue out_stream = output_container.add_stream_from_template(template=stream, opaque=True) stream_map[stream] = out_stream diff --git a/comfy_api_nodes/apis/bytedance.py b/comfy_api_nodes/apis/bytedance.py index 2d65d8645..76573304b 100644 --- a/comfy_api_nodes/apis/bytedance.py +++ b/comfy_api_nodes/apis/bytedance.py @@ -1,4 +1,4 @@ -from typing import Literal +from typing import Any, Literal from pydantic import BaseModel, Field @@ -24,8 +24,8 @@ class Seedream4TaskCreationRequest(BaseModel): image: list[str] | None = Field(None, description="Image URLs") size: str = Field(...) seed: int = Field(..., ge=0, le=2147483647) - sequential_image_generation: str = Field("disabled") - sequential_image_generation_options: Seedream4Options = Field(Seedream4Options(max_images=15)) + sequential_image_generation: str | None = Field("disabled") + sequential_image_generation_options: Seedream4Options | None = Field(Seedream4Options(max_images=15)) watermark: bool = Field(False) output_format: str | None = None @@ -261,6 +261,19 @@ _PRESETS_SEEDREAM_4K = [ _CUSTOM_PRESET = [("Custom", None, None)] +_PRESETS_SEEDREAM_2K_PRO = [ + ("(2K) 2048x2048 (1:1)", 2048, 2048), + ("(2K) 1728x2304 (3:4)", 1728, 2304), + ("(2K) 2304x1728 (4:3)", 2304, 1728), + # ("(2K) 2848x1600 (16:9)", 2848, 1600), # 4,556,800 px - temporarily unavailable + # ("(2K) 1600x2848 (9:16)", 1600, 2848), # 4,556,800 px - temporarily unavailable + ("(2K) 1664x2496 (2:3)", 1664, 2496), + ("(2K) 2496x1664 (3:2)", 2496, 1664), + # ("(2K) 3136x1344 (21:9)", 3136, 1344), # 4,214,784 px - temporarily unavailable +] +RECOMMENDED_PRESETS_SEEDREAM_5_PRO = ( + _PRESETS_SEEDREAM_1K + _PRESETS_SEEDREAM_2K_PRO + _CUSTOM_PRESET +) RECOMMENDED_PRESETS_SEEDREAM_5_LITE = ( _PRESETS_SEEDREAM_2K + _PRESETS_SEEDREAM_3K + _PRESETS_SEEDREAM_4K + _CUSTOM_PRESET ) @@ -316,3 +329,36 @@ VIDEO_TASKS_EXECUTION_TIME = { "1080p": 150, }, } + + +class SeedAudioConfig(BaseModel): + format: str = Field(default="mp3") + sample_rate: int = Field(default=24000) + speech_rate: int = Field(default=0) + loudness_rate: int = Field(default=0) + pitch_rate: int = Field(default=0) + + +class SeedAudioReference(BaseModel): + speaker: str | None = Field(default=None) + audio_data: str | None = Field(default=None) + audio_url: str | None = Field(default=None) + image_data: str | None = Field(default=None) + image_url: str | None = Field(default=None) + + +class SeedAudioRequest(BaseModel): + model: str = Field(default="seed-audio-1.0") + text_prompt: str = Field(...) + references: list[SeedAudioReference] | None = Field(default=None) + audio_config: SeedAudioConfig = Field(default_factory=SeedAudioConfig) + watermark: dict[str, Any] = Field(default_factory=dict) + + +class SeedAudioResponse(BaseModel): + audio: str | None = Field(default=None) + url: str | None = Field(default=None) + duration: float | None = Field(default=None) + original_duration: float | None = Field(default=None) + code: int | None = Field(default=None) + message: str | None = Field(default=None) diff --git a/comfy_api_nodes/apis/ideogram.py b/comfy_api_nodes/apis/ideogram.py index c5ad9559f..ee3256e96 100644 --- a/comfy_api_nodes/apis/ideogram.py +++ b/comfy_api_nodes/apis/ideogram.py @@ -33,53 +33,6 @@ class IdeogramColorPalette( ) -class ImageRequest(BaseModel): - aspect_ratio: Optional[str] = Field( - None, - description="Optional. The aspect ratio (e.g., 'ASPECT_16_9', 'ASPECT_1_1'). Cannot be used with resolution. Defaults to 'ASPECT_1_1' if unspecified.", - ) - color_palette: Optional[Dict[str, Any]] = Field( - None, description='Optional. Color palette object. Only for V_2, V_2_TURBO.' - ) - magic_prompt_option: Optional[str] = Field( - None, description="Optional. MagicPrompt usage ('AUTO', 'ON', 'OFF')." - ) - model: str = Field(..., description="The model used (e.g., 'V_2', 'V_2A_TURBO')") - negative_prompt: Optional[str] = Field( - None, - description='Optional. Description of what to exclude. Only for V_1, V_1_TURBO, V_2, V_2_TURBO.', - ) - num_images: Optional[int] = Field( - 1, - description='Optional. Number of images to generate (1-8). Defaults to 1.', - ge=1, - le=8, - ) - prompt: str = Field( - ..., description='Required. The prompt to use to generate the image.' - ) - resolution: Optional[str] = Field( - None, - description="Optional. Resolution (e.g., 'RESOLUTION_1024_1024'). Only for model V_2. Cannot be used with aspect_ratio.", - ) - seed: Optional[int] = Field( - None, - description='Optional. A number between 0 and 2147483647.', - ge=0, - le=2147483647, - ) - style_type: Optional[str] = Field( - None, - description="Optional. Style type ('AUTO', 'GENERAL', 'REALISTIC', 'DESIGN', 'RENDER_3D', 'ANIME'). Only for models V_2 and above.", - ) - - -class IdeogramGenerateRequest(BaseModel): - image_request: ImageRequest = Field( - ..., description='The image generation request parameters.' - ) - - class Datum(BaseModel): is_image_safe: Optional[bool] = Field( None, description='Indicates whether the image is considered safe.' @@ -113,20 +66,6 @@ class StyleCode(RootModel[str]): root: str = Field(..., pattern='^[0-9A-Fa-f]{8}$') -class Datum1(BaseModel): - is_image_safe: Optional[bool] = None - prompt: Optional[str] = None - resolution: Optional[str] = None - seed: Optional[int] = None - style_type: Optional[str] = None - url: Optional[str] = None - - -class IdeogramV3IdeogramResponse(BaseModel): - created: Optional[datetime] = None - data: Optional[List[Datum1]] = None - - class RenderingSpeed1(str, Enum): TURBO = 'TURBO' DEFAULT = 'DEFAULT' diff --git a/comfy_api_nodes/apis/stability.py b/comfy_api_nodes/apis/stability.py deleted file mode 100644 index 5b9b5ac7d..000000000 --- a/comfy_api_nodes/apis/stability.py +++ /dev/null @@ -1,147 +0,0 @@ -from enum import Enum -from typing import Optional - -from pydantic import BaseModel, Field, confloat - - -class StabilityFormat(str, Enum): - png = 'png' - jpeg = 'jpeg' - webp = 'webp' - - -class StabilityAspectRatio(str, Enum): - ratio_1_1 = "1:1" - ratio_16_9 = "16:9" - ratio_9_16 = "9:16" - ratio_3_2 = "3:2" - ratio_2_3 = "2:3" - ratio_5_4 = "5:4" - ratio_4_5 = "4:5" - ratio_21_9 = "21:9" - ratio_9_21 = "9:21" - - -def get_stability_style_presets(include_none=True): - presets = [] - if include_none: - presets.append("None") - return presets + [x.value for x in StabilityStylePreset] - - -class StabilityStylePreset(str, Enum): - _3d_model = "3d-model" - analog_film = "analog-film" - anime = "anime" - cinematic = "cinematic" - comic_book = "comic-book" - digital_art = "digital-art" - enhance = "enhance" - fantasy_art = "fantasy-art" - isometric = "isometric" - line_art = "line-art" - low_poly = "low-poly" - modeling_compound = "modeling-compound" - neon_punk = "neon-punk" - origami = "origami" - photographic = "photographic" - pixel_art = "pixel-art" - tile_texture = "tile-texture" - - -class Stability_SD3_5_Model(str, Enum): - sd3_5_large = "sd3.5-large" - # sd3_5_large_turbo = "sd3.5-large-turbo" - sd3_5_medium = "sd3.5-medium" - - -class Stability_SD3_5_GenerationMode(str, Enum): - text_to_image = "text-to-image" - image_to_image = "image-to-image" - - -class StabilityStable3_5Request(BaseModel): - model: str = Field(...) - mode: str = Field(...) - prompt: str = Field(...) - negative_prompt: Optional[str] = Field(None) - aspect_ratio: Optional[str] = Field(None) - seed: Optional[int] = Field(None) - output_format: Optional[str] = Field(StabilityFormat.png.value) - image: Optional[str] = Field(None) - style_preset: Optional[str] = Field(None) - cfg_scale: float = Field(...) - strength: Optional[confloat(ge=0.0, le=1.0)] = Field(None) - - -class StabilityUpscaleConservativeRequest(BaseModel): - prompt: str = Field(...) - negative_prompt: Optional[str] = Field(None) - seed: Optional[int] = Field(None) - output_format: Optional[str] = Field(StabilityFormat.png.value) - image: Optional[str] = Field(None) - creativity: Optional[confloat(ge=0.2, le=0.5)] = Field(None) - - -class StabilityUpscaleCreativeRequest(BaseModel): - prompt: str = Field(...) - negative_prompt: Optional[str] = Field(None) - seed: Optional[int] = Field(None) - output_format: Optional[str] = Field(StabilityFormat.png.value) - image: Optional[str] = Field(None) - creativity: Optional[confloat(ge=0.1, le=0.5)] = Field(None) - style_preset: Optional[str] = Field(None) - - -class StabilityStableUltraRequest(BaseModel): - prompt: str = Field(...) - negative_prompt: Optional[str] = Field(None) - aspect_ratio: Optional[str] = Field(None) - seed: Optional[int] = Field(None) - output_format: Optional[str] = Field(StabilityFormat.png.value) - image: Optional[str] = Field(None) - style_preset: Optional[str] = Field(None) - strength: Optional[confloat(ge=0.0, le=1.0)] = Field(None) - - -class StabilityStableUltraResponse(BaseModel): - image: Optional[str] = Field(None) - finish_reason: Optional[str] = Field(None) - seed: Optional[int] = Field(None) - - -class StabilityResultsGetResponse(BaseModel): - image: Optional[str] = Field(None) - finish_reason: Optional[str] = Field(None) - seed: Optional[int] = Field(None) - id: Optional[str] = Field(None) - name: Optional[str] = Field(None) - errors: Optional[list[str]] = Field(None) - status: Optional[str] = Field(None) - result: Optional[str] = Field(None) - - -class StabilityAsyncResponse(BaseModel): - id: Optional[str] = Field(None) - - -class StabilityTextToAudioRequest(BaseModel): - model: str = Field(...) - prompt: str = Field(...) - duration: int = Field(190, ge=1, le=190) - seed: int = Field(0, ge=0, le=4294967294) - steps: int = Field(8, ge=4, le=8) - output_format: str = Field("wav") - - -class StabilityAudioToAudioRequest(StabilityTextToAudioRequest): - strength: float = Field(0.01, ge=0.01, le=1.0) - - -class StabilityAudioInpaintRequest(StabilityTextToAudioRequest): - mask_start: int = Field(30, ge=0, le=190) - mask_end: int = Field(190, ge=0, le=190) - - -class StabilityAudioResponse(BaseModel): - audio: Optional[str] = Field(None) diff --git a/comfy_api_nodes/nodes_bytedance.py b/comfy_api_nodes/nodes_bytedance.py index f22415abd..043bc9526 100644 --- a/comfy_api_nodes/nodes_bytedance.py +++ b/comfy_api_nodes/nodes_bytedance.py @@ -1,3 +1,4 @@ +import base64 import hashlib import logging import math @@ -15,11 +16,16 @@ from comfy_api_nodes.apis.bytedance import ( RECOMMENDED_PRESETS_SEEDREAM_4_0, RECOMMENDED_PRESETS_SEEDREAM_4_5, RECOMMENDED_PRESETS_SEEDREAM_5_LITE, + RECOMMENDED_PRESETS_SEEDREAM_5_PRO, SEEDANCE2_REF_VIDEO_PIXEL_LIMITS, VIDEO_TASKS_EXECUTION_TIME, GetAssetResponse, Image2VideoTaskCreationRequest, ImageTaskCreationResponse, + SeedAudioConfig, + SeedAudioReference, + SeedAudioRequest, + SeedAudioResponse, Seedance2TaskCreationRequest, SeedanceCreateAssetRequest, SeedanceCreateAssetResponse, @@ -43,6 +49,8 @@ from comfy_api_nodes.apis.bytedance import ( ) from comfy_api_nodes.util import ( ApiEndpoint, + audio_bytes_to_audio_input, + audio_input_to_mp3, download_url_to_image_tensor, download_url_to_video_output, downscale_image_tensor_by_max_side, @@ -51,11 +59,14 @@ from comfy_api_nodes.util import ( image_tensor_pair_to_batch, poll_op, sync_op, + tensor_to_base64_string, upload_audio_to_comfyapi, upload_image_to_comfyapi, upload_images_to_comfyapi, upload_video_to_comfyapi, + upscale_image_tensor_to_min_pixels, upscale_video_to_min_pixels, + validate_audio_duration, validate_image_aspect_ratio, validate_image_dimensions, validate_string, @@ -70,12 +81,14 @@ _VERIFICATION_POLL_TIMEOUT_SEC = 120 _VERIFICATION_POLL_INTERVAL_SEC = 3 SEEDREAM_MODELS = { + "seedream 5.0 pro": "seedream-5-0-pro-260628", "seedream 5.0 lite": "seedream-5-0-260128", "seedream-4-5-251128": "seedream-4-5-251128", "seedream-4-0-250828": "seedream-4-0-250828", } SEEDREAM_PRESETS = { + "seedream-5-0-pro-260628": RECOMMENDED_PRESETS_SEEDREAM_5_PRO, "seedream-5-0-260128": RECOMMENDED_PRESETS_SEEDREAM_5_LITE, "seedream-4-5-251128": RECOMMENDED_PRESETS_SEEDREAM_4_5, "seedream-4-0-250828": RECOMMENDED_PRESETS_SEEDREAM_4_0, @@ -733,8 +746,15 @@ class ByteDanceSeedreamNode(IO.ComfyNode): return IO.NodeOutput(torch.cat([await download_url_to_image_tensor(i) for i in urls])) -def _seedream_model_inputs(*, max_ref_images: int, presets: list): - return [ +def _seedream_model_inputs( + *, + max_ref_images: int, + presets: list, + max_width: int = 6240, + max_height: int = 4992, + supports_batch: bool = True, +): + inputs = [ IO.Combo.Input( "size_preset", options=[label for label, _, _ in presets], @@ -744,7 +764,7 @@ def _seedream_model_inputs(*, max_ref_images: int, presets: list): "width", default=2048, min=1024, - max=6240, + max=max_width, step=2, tooltip="Custom width for image. Value is working only if `size_preset` is set to `Custom`", ), @@ -752,22 +772,27 @@ def _seedream_model_inputs(*, max_ref_images: int, presets: list): "height", default=2048, min=1024, - max=4992, + max=max_height, step=2, tooltip="Custom height for image. Value is working only if `size_preset` is set to `Custom`", ), - IO.Int.Input( - "max_images", - default=1, - min=1, - max=max_ref_images, - step=1, - display_mode=IO.NumberDisplay.number, - tooltip="Maximum number of images to generate. With 1, exactly one image is produced. " - "With >1, the model generates between 1 and max_images related images " - "(e.g., story scenes, character variations). " - "Total images (input + generated) cannot exceed 15.", - ), + ] + if supports_batch: + inputs.append( + IO.Int.Input( + "max_images", + default=1, + min=1, + max=max_ref_images, + step=1, + display_mode=IO.NumberDisplay.number, + tooltip="Maximum number of images to generate. With 1, exactly one image is produced. " + "With >1, the model generates between 1 and max_images related images " + "(e.g., story scenes, character variations). " + "Total images (input + generated) cannot exceed 15.", + ) + ) + inputs.append( IO.Autogrow.Input( "images", template=IO.Autogrow.TemplateNames( @@ -777,14 +802,18 @@ def _seedream_model_inputs(*, max_ref_images: int, presets: list): ), tooltip=f"Optional reference image(s) for image-to-image or multi-reference generation. " f"Up to {max_ref_images} images.", - ), - IO.Boolean.Input( - "fail_on_partial", - default=False, - tooltip="If enabled, abort execution if any requested images are missing or return an error.", - advanced=True, - ), - ] + ) + ) + if supports_batch: + inputs.append( + IO.Boolean.Input( + "fail_on_partial", + default=False, + tooltip="If enabled, abort execution if any requested images are missing or return an error.", + advanced=True, + ) + ) + return inputs class ByteDanceSeedreamNodeV2(IO.ComfyNode): @@ -806,6 +835,16 @@ class ByteDanceSeedreamNodeV2(IO.ComfyNode): IO.DynamicCombo.Input( "model", options=[ + IO.DynamicCombo.Option( + "seedream 5.0 pro", + _seedream_model_inputs( + max_ref_images=10, + presets=RECOMMENDED_PRESETS_SEEDREAM_5_PRO, + max_width=3136, + max_height=2496, + supports_batch=False, + ), + ), IO.DynamicCombo.Option( "seedream 5.0 lite", _seedream_model_inputs(max_ref_images=14, presets=RECOMMENDED_PRESETS_SEEDREAM_5_LITE), @@ -847,15 +886,27 @@ class ByteDanceSeedreamNodeV2(IO.ComfyNode): ], is_api_node=True, price_badge=IO.PriceBadge( - depends_on=IO.PriceBadgeDepends(widgets=["model"]), + depends_on=IO.PriceBadgeDepends( + widgets=["model", "model.size_preset", "model.width", "model.height"] + ), expr=""" ( - $price := $contains(widgets.model, "5.0 lite") ? 0.035 : - $contains(widgets.model, "4-5") ? 0.04 : 0.03; + $sp := $lookup(widgets, "model.size_preset"); + $px := $lookup(widgets, "model.width") * $lookup(widgets, "model.height"); + $isPro := $contains(widgets.model, "5.0 pro"); + $price := $isPro + ? ( + $contains($sp, "custom") + ? ($px <= 2360000 ? 0.045 : 0.09) + : ($contains($sp, "1k") ? 0.045 : 0.09) + ) + : $contains(widgets.model, "5.0 lite") ? 0.035 + : $contains(widgets.model, "4-5") ? 0.04 + : 0.03; { - "type":"usd", + "type": "usd", "usd": $price, - "format": { "suffix":" x images/Run", "approximate": true } + "format": { "suffix": $isPro ? "/Image" : " x images/Run", "approximate": true } } ) """, @@ -873,6 +924,7 @@ class ByteDanceSeedreamNodeV2(IO.ComfyNode): validate_string(prompt, strip_whitespace=True, min_length=1) model_id = SEEDREAM_MODELS[model["model"]] presets = SEEDREAM_PRESETS[model_id] + is_pro = "seedream-5-0-pro" in model_id size_preset = model.get("size_preset", presets[0][0]) width = model.get("width", 2048) @@ -892,19 +944,29 @@ class ByteDanceSeedreamNodeV2(IO.ComfyNode): out_num_pixels = w * h mp_provided = out_num_pixels / 1_000_000.0 - if ("seedream-4-5" in model_id or "seedream-5-0" in model_id) and out_num_pixels < 3686400: - raise ValueError( - f"Minimum image resolution for the selected model is 3.68MP, but {mp_provided:.2f}MP provided." - ) - if "seedream-4-0" in model_id and out_num_pixels < 921600: - raise ValueError( - f"Minimum image resolution that the selected model can generate is 0.92MP, " - f"but {mp_provided:.2f}MP provided." - ) - if out_num_pixels > 16_777_216: - raise ValueError( - f"Maximum image resolution for the selected model is 16.78MP, but {mp_provided:.2f}MP provided." - ) + if is_pro: + if out_num_pixels < 921_600: + raise ValueError( + f"Minimum image resolution for the selected model is 0.92MP, but {mp_provided:.2f}MP provided." + ) + if out_num_pixels > 4_194_304: + raise ValueError( + f"Maximum image resolution for the selected model is 4.19MP, but {mp_provided:.2f}MP provided." + ) + else: + if ("seedream-4-5" in model_id or "seedream-5-0" in model_id) and out_num_pixels < 3_686_400: + raise ValueError( + f"Minimum image resolution for the selected model is 3.68MP, but {mp_provided:.2f}MP provided." + ) + if "seedream-4-0" in model_id and out_num_pixels < 921_600: + raise ValueError( + f"Minimum image resolution that the selected model can generate is 0.92MP, " + f"but {mp_provided:.2f}MP provided." + ) + if out_num_pixels > 16_777_216: + raise ValueError( + f"Maximum image resolution for the selected model is 16.78MP, but {mp_provided:.2f}MP provided." + ) image_tensors: list[Input.Image] = [t for t in images_dict.values() if t is not None] n_input_images = sum(get_number_of_images(t) for t in image_tensors) @@ -940,8 +1002,8 @@ class ByteDanceSeedreamNodeV2(IO.ComfyNode): image=reference_images_urls, size=f"{w}x{h}", seed=seed, - sequential_image_generation=sequential_image_generation, - sequential_image_generation_options=Seedream4Options(max_images=max_images), + sequential_image_generation=None if is_pro else sequential_image_generation, + sequential_image_generation_options=None if is_pro else Seedream4Options(max_images=max_images), watermark=watermark, ), ) @@ -2474,6 +2536,311 @@ class ByteDanceCreateVideoAsset(IO.ComfyNode): return IO.NodeOutput(asset_id, resolved_group) +MODE_TEXT = "text only" +MODE_AUDIO = "audio reference" +MODE_IMAGE = "image reference" +MODE_SPEAKER = "preset voice" + +# (speaker_id, display_label) for built-in TTS 2.0 voices; resolvable ids are account-scoped. +SEED_AUDIO_PRESET_VOICES: list[tuple[str, str]] = [ + ("zh_female_vv_uranus_bigtts", "Vivi (Female, multilingual)"), + ("zh_female_xiaohe_uranus_bigtts", "Mindy (Female, multilingual)"), + ("en_female_stokie_uranus_bigtts", "Stokie (Female, English)"), + ("en_female_dacey_uranus_bigtts", "Dacey (Female, English)"), + ("en_male_tim_uranus_bigtts", "Tim (Male, English)"), + ("zh_male_m191_uranus_bigtts", "Kian (Male, multilingual)"), + ("zh_male_taocheng_uranus_bigtts", "Cedric (Male, multilingual)"), + ("zh_male_sophie_uranus_bigtts", "Sophie (Female, multilingual)"), + ("zh_female_yingyujiaoxue_uranus_bigtts", "Jean (Female, multilingual)"), + ("zh_male_dayi_uranus_bigtts", "Magnus (Male, multilingual)"), + ("zh_female_mizai_uranus_bigtts", "Mabel (Female, multilingual)"), + ("zh_female_jitangnv_uranus_bigtts", "Nadia (Female, multilingual)"), + ("zh_female_meilinvyou_uranus_bigtts", "Opal (Female, multilingual)"), + ("zh_female_liuchangnv_uranus_bigtts", "Pearl (Female, multilingual)"), + ("zh_male_ruyayichen_uranus_bigtts", "Quentin (Male, multilingual)"), + ("zh_female_vivo_uranus_bigtts", "Vienna (Female, multilingual)"), + ("zh_female_xiaoai_uranus_bigtts", "Alina (Female, multilingual)"), + ("zh_female_cancan_uranus_bigtts", "Corinne (Female, multilingual)"), + ("zh_female_tianmeixiaoyuan_uranus_bigtts", "Esther (Female, multilingual)"), + ("zh_female_tianmeitaozi_uranus_bigtts", "Freya (Female, multilingual)"), + ("zh_female_shuangkuaisisi_uranus_bigtts", "Gigi (Female, multilingual)"), + ("zh_female_peiqi_uranus_bigtts", "Holly (Female, multilingual)"), + ("zh_female_xiaoxue_uranus_bigtts", "Lyla (Female, multilingual)"), + ("zh_female_yuanqi_uranus_bigtts", "Daisy (Female, multilingual)"), + ("zh_female_kefunvsheng_uranus_bigtts", "Tracy (Female, multilingual)"), + ("zh_male_shaonianzixin_uranus_bigtts", "Jess (Male, multilingual)"), + ("zh_female_linjianvhai_uranus_bigtts", "Pinky (Female, multilingual)"), + ("zh_female_kiwi_uranus_bigtts", "Sweety (Female, multilingual)"), + ("zh_female_sajiaoxuemei_uranus_bigtts", "Sandy (Female, multilingual)"), + ("de_male_seven_uranus_bigtts", "Sven (Male, German)"), + ("jp_female_minimi_uranus_bigtts", "Minimi (Female, Japanese)"), + ("fr_male_usseau_uranus_bigtts", "Usseau (Male, French)"), + ("es_male_felipe_uranus_bigtts", "Felipe (Male, Spanish)"), + ("id_male_han_uranus_bigtts", "Han (Male, Indonesian)"), + ("pt_male_martins_uranus_bigtts", "Martins (Male, Portuguese)"), + ("it_male_enzo_uranus_bigtts", "Enzo (Male, Italian)"), + ("kr_male_shane_uranus_bigtts", "Shane (Male, Korean)"), + ("zh_male_liufei_uranus_bigtts", "Felix (Male, Chinese)"), + ("zh_female_qingxinnvsheng_uranus_bigtts", "Celeste (Female, Chinese)"), + ("zh_male_sunwukong_uranus_bigtts", "Monkey King (Male, Chinese)"), +] +SEED_AUDIO_VOICE_OPTIONS = [label for _, label in SEED_AUDIO_PRESET_VOICES] +SEED_AUDIO_VOICE_MAP = {label: speaker_id for speaker_id, label in SEED_AUDIO_PRESET_VOICES} + +_AUDIO_TAG_RE = re.compile(r"@Audio(\d+)", re.IGNORECASE) + + +def max_audio_tag(prompt: str) -> int: + """Highest N referenced as @AudioN in the prompt (0 if none).""" + nums = [int(m) for m in _AUDIO_TAG_RE.findall(prompt or "")] + return max(nums) if nums else 0 + + +def connected_audio_indices(reference_mode: dict) -> list[int]: + """Indices (1-based) of connected reference_audio sockets, in order.""" + return [ + i + for i in range(1, 3 + 1) + if reference_mode.get(f"reference_audio_{i}") is not None + ] + + +def validate_seed_audio_inputs( + text_prompt: str, + mode: str, + audio_indices: list[int], + has_image: bool, + preset_voice: str | None = None, +) -> None: + validate_string(text_prompt, field_name="text_prompt", min_length=1, max_length=3000) + max_tag = max_audio_tag(text_prompt) + + if mode == MODE_TEXT: + if max_tag: + raise ValueError( + f"The prompt references @Audio{max_tag}, but reference mode is '{MODE_TEXT}'. " + f"Switch to '{MODE_AUDIO}' and connect the reference clip(s)." + ) + elif mode == MODE_AUDIO: + if not audio_indices: + raise ValueError( + f"Reference mode '{MODE_AUDIO}' requires at least one reference_audio input " + f"(or switch to '{MODE_TEXT}')." + ) + if audio_indices != list(range(1, len(audio_indices) + 1)): + raise ValueError( + "Connect reference_audio inputs in order without gaps: reference_audio_1, then _2, then _3." + ) + if max_tag > len(audio_indices): + raise ValueError( + f"The prompt references @Audio{max_tag}, but only {len(audio_indices)} " + f"reference audio(s) are connected." + ) + elif mode == MODE_IMAGE: + if not has_image: + raise ValueError(f"Reference mode '{MODE_IMAGE}' requires a reference_image input.") + if max_tag: + raise ValueError( + f"@AudioN tags are not used in '{MODE_IMAGE}' mode; the prompt should contain " + f"only the text to synthesize." + ) + elif mode == MODE_SPEAKER: + if not preset_voice or preset_voice not in SEED_AUDIO_VOICE_MAP: + raise ValueError(f"Reference mode '{MODE_SPEAKER}' requires selecting a preset voice.") + if max_tag > 1: + raise ValueError( + f"'{MODE_SPEAKER}' mode uses a single voice, so @Audio{max_tag} is out of range. " + f"Remove the @AudioN tags — the whole prompt is read in the selected voice." + ) + else: + raise ValueError(f"Unknown reference mode: {mode!r}") + + +class ByteDanceSeedAudioNode(IO.ComfyNode): + + @classmethod + def define_schema(cls) -> IO.Schema: + return IO.Schema( + node_id="ByteDanceSeedAudio", + display_name="ByteDance Seed Audio 1.0", + category="partner/audio/ByteDance", + description=( + "Generate speech, music, sound effects and multi-speaker dialogue from a single prompt " + "with ByteDance Seed Audio 1.0. Describe the voice(s), emotion, ambience, background music " + "and sound effects in the prompt, and include the lines to speak. Optionally pick a built-in " + "preset voice, clone voices from up to 3 reference clips (tagged @Audio1-3 in the prompt), " + "or derive a voice from a character image. Up to 2 minutes of audio per run." + ), + inputs=[ + IO.String.Input( + "text_prompt", + multiline=True, + default="", + tooltip=( + "Describe the voice(s), emotion, pacing, ambience, background music and sound " + "effects, and include the lines to speak (name characters inline for dialogue). " + "In 'audio reference' mode, refer to connected clips by order as @Audio1, @Audio2, " + "@Audio3. Maximum 3000 characters." + ), + ), + IO.DynamicCombo.Input( + "reference_mode", + options=[ + IO.DynamicCombo.Option(MODE_TEXT, []), + IO.DynamicCombo.Option( + MODE_AUDIO, + [ + IO.Audio.Input( + "reference_audio_1", + optional=True, + tooltip="Reference clip for voice cloning, tagged @Audio1 in the prompt. " + "Up to 30s.", + ), + IO.Audio.Input( + "reference_audio_2", + optional=True, + tooltip="Reference clip tagged @Audio2 in the prompt. Up to 30s.", + ), + IO.Audio.Input( + "reference_audio_3", + optional=True, + tooltip="Reference clip tagged @Audio3 in the prompt. Up to 30s.", + ), + ], + ), + IO.DynamicCombo.Option( + MODE_IMAGE, + [ + IO.Image.Input( + "reference_image", + optional=True, + tooltip="A single character image; the model derives a voice from it. " + "Cannot be combined with reference audio.", + ), + ], + ), + IO.DynamicCombo.Option( + MODE_SPEAKER, + [ + IO.Combo.Input( + "preset_voice", + options=SEED_AUDIO_VOICE_OPTIONS, + default=SEED_AUDIO_VOICE_OPTIONS[0], + tooltip="A built-in TTS 2.0 voice that reads the prompt. No reference " + "clip needed, and @AudioN tags are not used in this mode.", + ), + ], + ), + ], + tooltip=( + "How to condition the voice: 'text only' (describe everything in the prompt), " + "'audio reference' (clone up to 3 voices, tagged @Audio1-3), 'image reference' " + "(derive a voice from one character image), or 'preset voice' (pick a built-in " + "named voice that reads the prompt)." + ), + ), + IO.Combo.Input( + "sample_rate", + options=["8000", "16000", "24000", "32000", "44100", "48000"], + default="24000", + tooltip="Output sample rate in Hz.", + ), + IO.Int.Input( + "speech_rate", + default=0, + min=-50, + max=100, + tooltip="Speaking speed. 0 = normal, 100 = 2.0x, -50 = 0.5x.", + ), + IO.Int.Input( + "loudness_rate", + default=0, + min=-50, + max=100, + tooltip="Loudness. 0 = normal, 100 = 2.0x, -50 = 0.5x.", + ), + IO.Int.Input( + "pitch_rate", + default=0, + min=-12, + max=12, + tooltip="Pitch shift in semitones (-12 to 12).", + ), + IO.Int.Input( + "seed", + default=42, + min=0, + max=2147483647, + control_after_generate=True, + tooltip="Seed controls whether the node should re-run; " + "results are non-deterministic regardless of seed.", + ), + ], + outputs=[IO.Audio.Output()], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + expr="""{"type":"usd","usd": 0.2145, "format":{"suffix":"/minute","approximate":true}}""", + ), + ) + + @classmethod + async def execute( + cls, + text_prompt: str, + reference_mode: dict, + sample_rate: str, + speech_rate: int, + loudness_rate: int, + pitch_rate: int, + seed: int, + ) -> IO.NodeOutput: + mode = reference_mode["reference_mode"] + audio_indices = connected_audio_indices(reference_mode) + image = reference_mode.get("reference_image") + preset_voice = reference_mode.get("preset_voice") + validate_seed_audio_inputs(text_prompt, mode, audio_indices, image is not None, preset_voice) + + references: list[SeedAudioReference] | None = None + if mode == MODE_AUDIO: + references = [] + for i in audio_indices: + clip = reference_mode[f"reference_audio_{i}"] + validate_audio_duration(clip, max_duration=30.0) + mp3_bytes = audio_input_to_mp3(clip).getvalue() + references.append(SeedAudioReference(audio_data=base64.b64encode(mp3_bytes).decode("utf-8"))) + elif mode == MODE_IMAGE: + image = upscale_image_tensor_to_min_pixels(image, 160_000) + references = [SeedAudioReference(image_data=tensor_to_base64_string(image, mime_type="image/png"))] + elif mode == MODE_SPEAKER: + references = [SeedAudioReference(speaker=SEED_AUDIO_VOICE_MAP[preset_voice])] + + response = await sync_op( + cls, + ApiEndpoint(path="/proxy/byteplus/api/v3/tts/create", method="POST"), + response_model=SeedAudioResponse, + data=SeedAudioRequest( + text_prompt=text_prompt, + references=references, + audio_config=SeedAudioConfig( + sample_rate=int(sample_rate), + speech_rate=speech_rate, + loudness_rate=loudness_rate, + pitch_rate=pitch_rate, + ), + ), + ) + if not response.audio: + raise Exception( + f"Seed Audio returned no audio (code={response.code}): {response.message}" + ) + return IO.NodeOutput(audio_bytes_to_audio_input(base64.b64decode(response.audio))) + + class ByteDanceExtension(ComfyExtension): @override async def get_node_list(self) -> list[type[IO.ComfyNode]]: @@ -2490,6 +2857,7 @@ class ByteDanceExtension(ComfyExtension): ByteDance2ReferenceNode, ByteDanceCreateImageAsset, ByteDanceCreateVideoAsset, + ByteDanceSeedAudioNode, ] diff --git a/comfy_api_nodes/nodes_ideogram.py b/comfy_api_nodes/nodes_ideogram.py index 3b914a850..cc0467987 100644 --- a/comfy_api_nodes/nodes_ideogram.py +++ b/comfy_api_nodes/nodes_ideogram.py @@ -5,9 +5,7 @@ from PIL import Image import numpy as np import torch from comfy_api_nodes.apis.ideogram import ( - IdeogramGenerateRequest, IdeogramGenerateResponse, - ImageRequest, IdeogramV3Request, IdeogramV3EditRequest, IdeogramV4Request, @@ -21,101 +19,6 @@ from comfy_api_nodes.util import ( validate_string, ) -V1_V1_RES_MAP = { - "Auto":"AUTO", - "512 x 1536":"RESOLUTION_512_1536", - "576 x 1408":"RESOLUTION_576_1408", - "576 x 1472":"RESOLUTION_576_1472", - "576 x 1536":"RESOLUTION_576_1536", - "640 x 1024":"RESOLUTION_640_1024", - "640 x 1344":"RESOLUTION_640_1344", - "640 x 1408":"RESOLUTION_640_1408", - "640 x 1472":"RESOLUTION_640_1472", - "640 x 1536":"RESOLUTION_640_1536", - "704 x 1152":"RESOLUTION_704_1152", - "704 x 1216":"RESOLUTION_704_1216", - "704 x 1280":"RESOLUTION_704_1280", - "704 x 1344":"RESOLUTION_704_1344", - "704 x 1408":"RESOLUTION_704_1408", - "704 x 1472":"RESOLUTION_704_1472", - "720 x 1280":"RESOLUTION_720_1280", - "736 x 1312":"RESOLUTION_736_1312", - "768 x 1024":"RESOLUTION_768_1024", - "768 x 1088":"RESOLUTION_768_1088", - "768 x 1152":"RESOLUTION_768_1152", - "768 x 1216":"RESOLUTION_768_1216", - "768 x 1232":"RESOLUTION_768_1232", - "768 x 1280":"RESOLUTION_768_1280", - "768 x 1344":"RESOLUTION_768_1344", - "832 x 960":"RESOLUTION_832_960", - "832 x 1024":"RESOLUTION_832_1024", - "832 x 1088":"RESOLUTION_832_1088", - "832 x 1152":"RESOLUTION_832_1152", - "832 x 1216":"RESOLUTION_832_1216", - "832 x 1248":"RESOLUTION_832_1248", - "864 x 1152":"RESOLUTION_864_1152", - "896 x 960":"RESOLUTION_896_960", - "896 x 1024":"RESOLUTION_896_1024", - "896 x 1088":"RESOLUTION_896_1088", - "896 x 1120":"RESOLUTION_896_1120", - "896 x 1152":"RESOLUTION_896_1152", - "960 x 832":"RESOLUTION_960_832", - "960 x 896":"RESOLUTION_960_896", - "960 x 1024":"RESOLUTION_960_1024", - "960 x 1088":"RESOLUTION_960_1088", - "1024 x 640":"RESOLUTION_1024_640", - "1024 x 768":"RESOLUTION_1024_768", - "1024 x 832":"RESOLUTION_1024_832", - "1024 x 896":"RESOLUTION_1024_896", - "1024 x 960":"RESOLUTION_1024_960", - "1024 x 1024":"RESOLUTION_1024_1024", - "1088 x 768":"RESOLUTION_1088_768", - "1088 x 832":"RESOLUTION_1088_832", - "1088 x 896":"RESOLUTION_1088_896", - "1088 x 960":"RESOLUTION_1088_960", - "1120 x 896":"RESOLUTION_1120_896", - "1152 x 704":"RESOLUTION_1152_704", - "1152 x 768":"RESOLUTION_1152_768", - "1152 x 832":"RESOLUTION_1152_832", - "1152 x 864":"RESOLUTION_1152_864", - "1152 x 896":"RESOLUTION_1152_896", - "1216 x 704":"RESOLUTION_1216_704", - "1216 x 768":"RESOLUTION_1216_768", - "1216 x 832":"RESOLUTION_1216_832", - "1232 x 768":"RESOLUTION_1232_768", - "1248 x 832":"RESOLUTION_1248_832", - "1280 x 704":"RESOLUTION_1280_704", - "1280 x 720":"RESOLUTION_1280_720", - "1280 x 768":"RESOLUTION_1280_768", - "1280 x 800":"RESOLUTION_1280_800", - "1312 x 736":"RESOLUTION_1312_736", - "1344 x 640":"RESOLUTION_1344_640", - "1344 x 704":"RESOLUTION_1344_704", - "1344 x 768":"RESOLUTION_1344_768", - "1408 x 576":"RESOLUTION_1408_576", - "1408 x 640":"RESOLUTION_1408_640", - "1408 x 704":"RESOLUTION_1408_704", - "1472 x 576":"RESOLUTION_1472_576", - "1472 x 640":"RESOLUTION_1472_640", - "1472 x 704":"RESOLUTION_1472_704", - "1536 x 512":"RESOLUTION_1536_512", - "1536 x 576":"RESOLUTION_1536_576", - "1536 x 640":"RESOLUTION_1536_640", -} - -V1_V2_RATIO_MAP = { - "1:1":"ASPECT_1_1", - "4:3":"ASPECT_4_3", - "3:4":"ASPECT_3_4", - "16:9":"ASPECT_16_9", - "9:16":"ASPECT_9_16", - "2:1":"ASPECT_2_1", - "1:2":"ASPECT_1_2", - "3:2":"ASPECT_3_2", - "2:3":"ASPECT_2_3", - "4:5":"ASPECT_4_5", - "5:4":"ASPECT_5_4", -} V3_RATIO_MAP = { "1:3":"1x3", @@ -229,298 +132,6 @@ async def download_and_process_images(image_urls): return stacked_tensors -class IdeogramV1(IO.ComfyNode): - - @classmethod - def define_schema(cls): - return IO.Schema( - node_id="IdeogramV1", - display_name="Ideogram V1", - category="partner/image/Ideogram", - description="Generates images using the Ideogram V1 model.", - inputs=[ - IO.String.Input( - "prompt", - multiline=True, - default="", - tooltip="Prompt for the image generation", - ), - IO.Boolean.Input( - "turbo", - default=False, - tooltip="Whether to use turbo mode (faster generation, potentially lower quality)", - ), - IO.Combo.Input( - "aspect_ratio", - options=list(V1_V2_RATIO_MAP.keys()), - default="1:1", - tooltip="The aspect ratio for image generation.", - optional=True, - ), - IO.Combo.Input( - "magic_prompt_option", - options=["AUTO", "ON", "OFF"], - default="AUTO", - tooltip="Determine if MagicPrompt should be used in generation", - optional=True, - advanced=True, - ), - IO.Int.Input( - "seed", - default=0, - min=0, - max=2147483647, - step=1, - control_after_generate=True, - display_mode=IO.NumberDisplay.number, - optional=True, - ), - IO.String.Input( - "negative_prompt", - multiline=True, - default="", - tooltip="Description of what to exclude from the image", - optional=True, - ), - IO.Int.Input( - "num_images", - default=1, - min=1, - max=8, - step=1, - display_mode=IO.NumberDisplay.number, - optional=True, - ), - ], - outputs=[ - IO.Image.Output(), - ], - hidden=[ - IO.Hidden.auth_token_comfy_org, - IO.Hidden.api_key_comfy_org, - IO.Hidden.unique_id, - ], - is_api_node=True, - price_badge=IO.PriceBadge( - depends_on=IO.PriceBadgeDepends(widgets=["num_images", "turbo"]), - expr=""" - ( - $n := widgets.num_images; - $base := (widgets.turbo = true) ? 0.0286 : 0.0858; - {"type":"usd","usd": $round($base * $n, 2)} - ) - """, - ), - ) - - @classmethod - async def execute( - cls, - prompt, - turbo=False, - aspect_ratio="1:1", - magic_prompt_option="AUTO", - seed=0, - negative_prompt="", - num_images=1, - ): - # Determine the model based on turbo setting - aspect_ratio = V1_V2_RATIO_MAP.get(aspect_ratio, None) - model = "V_1_TURBO" if turbo else "V_1" - - response = await sync_op( - cls, - ApiEndpoint(path="/proxy/ideogram/generate", method="POST"), - response_model=IdeogramGenerateResponse, - data=IdeogramGenerateRequest( - image_request=ImageRequest( - prompt=prompt, - model=model, - num_images=num_images, - seed=seed, - aspect_ratio=aspect_ratio if aspect_ratio != "ASPECT_1_1" else None, - magic_prompt_option=(magic_prompt_option if magic_prompt_option != "AUTO" else None), - negative_prompt=negative_prompt if negative_prompt else None, - ) - ), - max_retries=1, - ) - - if not response.data or len(response.data) == 0: - raise Exception("No images were generated in the response") - - image_urls = [image_data.url for image_data in response.data if image_data.url] - if not image_urls: - raise Exception("No image URLs were generated in the response") - return IO.NodeOutput(await download_and_process_images(image_urls)) - - -class IdeogramV2(IO.ComfyNode): - - @classmethod - def define_schema(cls): - return IO.Schema( - node_id="IdeogramV2", - display_name="Ideogram V2", - category="partner/image/Ideogram", - description="Generates images using the Ideogram V2 model.", - inputs=[ - IO.String.Input( - "prompt", - multiline=True, - default="", - tooltip="Prompt for the image generation", - ), - IO.Boolean.Input( - "turbo", - default=False, - tooltip="Whether to use turbo mode (faster generation, potentially lower quality)", - ), - IO.Combo.Input( - "aspect_ratio", - options=list(V1_V2_RATIO_MAP.keys()), - default="1:1", - tooltip="The aspect ratio for image generation. Ignored if resolution is not set to AUTO.", - optional=True, - ), - IO.Combo.Input( - "resolution", - options=list(V1_V1_RES_MAP.keys()), - default="Auto", - tooltip="The resolution for image generation. " - "If not set to AUTO, this overrides the aspect_ratio setting.", - optional=True, - ), - IO.Combo.Input( - "magic_prompt_option", - options=["AUTO", "ON", "OFF"], - default="AUTO", - tooltip="Determine if MagicPrompt should be used in generation", - optional=True, - advanced=True, - ), - IO.Int.Input( - "seed", - default=0, - min=0, - max=2147483647, - step=1, - control_after_generate=True, - display_mode=IO.NumberDisplay.number, - optional=True, - ), - IO.Combo.Input( - "style_type", - options=["AUTO", "GENERAL", "REALISTIC", "DESIGN", "RENDER_3D", "ANIME"], - default="NONE", - tooltip="Style type for generation (V2 only)", - optional=True, - advanced=True, - ), - IO.String.Input( - "negative_prompt", - multiline=True, - default="", - tooltip="Description of what to exclude from the image", - optional=True, - ), - IO.Int.Input( - "num_images", - default=1, - min=1, - max=8, - step=1, - display_mode=IO.NumberDisplay.number, - optional=True, - ), - #"color_palette": ( - # IO.STRING, - # { - # "multiline": False, - # "default": "", - # "tooltip": "Color palette preset name or hex colors with weights", - # }, - #), - ], - outputs=[ - IO.Image.Output(), - ], - hidden=[ - IO.Hidden.auth_token_comfy_org, - IO.Hidden.api_key_comfy_org, - IO.Hidden.unique_id, - ], - is_api_node=True, - price_badge=IO.PriceBadge( - depends_on=IO.PriceBadgeDepends(widgets=["num_images", "turbo"]), - expr=""" - ( - $n := widgets.num_images; - $base := (widgets.turbo = true) ? 0.0715 : 0.1144; - {"type":"usd","usd": $round($base * $n, 2)} - ) - """, - ), - ) - - @classmethod - async def execute( - cls, - prompt, - turbo=False, - aspect_ratio="1:1", - resolution="Auto", - magic_prompt_option="AUTO", - seed=0, - style_type="NONE", - negative_prompt="", - num_images=1, - color_palette="", - ): - aspect_ratio = V1_V2_RATIO_MAP.get(aspect_ratio, None) - resolution = V1_V1_RES_MAP.get(resolution, None) - # Determine the model based on turbo setting - model = "V_2_TURBO" if turbo else "V_2" - - # Handle resolution vs aspect_ratio logic - # If resolution is not AUTO, it overrides aspect_ratio - final_resolution = None - final_aspect_ratio = None - - if resolution != "AUTO": - final_resolution = resolution - else: - final_aspect_ratio = aspect_ratio if aspect_ratio != "ASPECT_1_1" else None - - response = await sync_op( - cls, - endpoint=ApiEndpoint(path="/proxy/ideogram/generate", method="POST"), - response_model=IdeogramGenerateResponse, - data=IdeogramGenerateRequest( - image_request=ImageRequest( - prompt=prompt, - model=model, - num_images=num_images, - seed=seed, - aspect_ratio=final_aspect_ratio, - resolution=final_resolution, - magic_prompt_option=(magic_prompt_option if magic_prompt_option != "AUTO" else None), - style_type=style_type if style_type != "NONE" else None, - negative_prompt=negative_prompt if negative_prompt else None, - color_palette=color_palette if color_palette else None, - ) - ), - max_retries=1, - ) - if not response.data or len(response.data) == 0: - raise Exception("No images were generated in the response") - - image_urls = [image_data.url for image_data in response.data if image_data.url] - if not image_urls: - raise Exception("No image URLs were generated in the response") - return IO.NodeOutput(await download_and_process_images(image_urls)) - - class IdeogramV3(IO.ComfyNode): @classmethod @@ -917,8 +528,6 @@ class IdeogramExtension(ComfyExtension): @override async def get_node_list(self) -> list[type[IO.ComfyNode]]: return [ - IdeogramV1, - IdeogramV2, IdeogramV3, IdeogramV4, ] diff --git a/comfy_api_nodes/nodes_stability.py b/comfy_api_nodes/nodes_stability.py deleted file mode 100644 index 9eaba173b..000000000 --- a/comfy_api_nodes/nodes_stability.py +++ /dev/null @@ -1,932 +0,0 @@ -from inspect import cleandoc -from typing import Optional -from typing_extensions import override - -from comfy_api.latest import ComfyExtension, Input, IO -from comfy_api_nodes.apis.stability import ( - StabilityUpscaleConservativeRequest, - StabilityUpscaleCreativeRequest, - StabilityAsyncResponse, - StabilityResultsGetResponse, - StabilityStable3_5Request, - StabilityStableUltraRequest, - StabilityStableUltraResponse, - StabilityAspectRatio, - Stability_SD3_5_Model, - Stability_SD3_5_GenerationMode, - get_stability_style_presets, - StabilityTextToAudioRequest, - StabilityAudioToAudioRequest, - StabilityAudioInpaintRequest, - StabilityAudioResponse, -) -from comfy_api_nodes.util import ( - validate_audio_duration, - validate_string, - audio_input_to_mp3, - bytesio_to_image_tensor, - tensor_to_bytesio, - audio_bytes_to_audio_input, - sync_op, - poll_op, - ApiEndpoint, -) - -import torch -import base64 -from io import BytesIO -from enum import Enum - - -class StabilityPollStatus(str, Enum): - finished = "finished" - in_progress = "in_progress" - failed = "failed" - - -def get_async_dummy_status(x: StabilityResultsGetResponse): - if x.name is not None or x.errors is not None: - return StabilityPollStatus.failed - elif x.finish_reason is not None: - return StabilityPollStatus.finished - return StabilityPollStatus.in_progress - - -class StabilityStableImageUltraNode(IO.ComfyNode): - """ - Generates images synchronously based on prompt and resolution. - """ - - @classmethod - def define_schema(cls): - return IO.Schema( - node_id="StabilityStableImageUltraNode", - display_name="Stability AI Stable Image Ultra", - category="partner/image/Stability AI", - description=cleandoc(cls.__doc__ or ""), - inputs=[ - IO.String.Input( - "prompt", - multiline=True, - default="", - tooltip="What you wish to see in the output image. A strong, descriptive prompt that clearly defines" + - "elements, colors, and subjects will lead to better results. " + - "To control the weight of a given word use the format `(word:weight)`," + - "where `word` is the word you'd like to control the weight of and `weight`" + - "is a value between 0 and 1. For example: `The sky was a crisp (blue:0.3) and (green:0.8)`" + - "would convey a sky that was blue and green, but more green than blue.", - ), - IO.Combo.Input( - "aspect_ratio", - options=StabilityAspectRatio, - default=StabilityAspectRatio.ratio_1_1, - tooltip="Aspect ratio of generated image.", - ), - IO.Combo.Input( - "style_preset", - options=get_stability_style_presets(), - tooltip="Optional desired style of generated image.", - advanced=True, - ), - IO.Int.Input( - "seed", - default=0, - min=0, - max=4294967294, - step=1, - display_mode=IO.NumberDisplay.number, - control_after_generate=True, - tooltip="The random seed used for creating the noise.", - ), - IO.Image.Input( - "image", - optional=True, - ), - IO.String.Input( - "negative_prompt", - default="", - tooltip="A blurb of text describing what you do not wish to see in the output image. This is an advanced feature.", - force_input=True, - optional=True, - advanced=True, - ), - IO.Float.Input( - "image_denoise", - default=0.5, - min=0.0, - max=1.0, - step=0.01, - tooltip="Denoise of input image; 0.0 yields image identical to input, 1.0 is as if no image was provided at all.", - optional=True, - ), - ], - outputs=[ - IO.Image.Output(), - ], - hidden=[ - IO.Hidden.auth_token_comfy_org, - IO.Hidden.api_key_comfy_org, - IO.Hidden.unique_id, - ], - is_api_node=True, - price_badge=IO.PriceBadge( - expr="""{"type":"usd","usd":0.08}""", - ), - ) - - @classmethod - async def execute( - cls, - prompt: str, - aspect_ratio: str, - style_preset: str, - seed: int, - image: Optional[torch.Tensor] = None, - negative_prompt: str = "", - image_denoise: Optional[float] = 0.5, - ) -> IO.NodeOutput: - validate_string(prompt, strip_whitespace=False) - # prepare image binary if image present - image_binary = None - if image is not None: - image_binary = tensor_to_bytesio(image, total_pixels=1504*1504).read() - else: - image_denoise = None - - if not negative_prompt: - negative_prompt = None - if style_preset == "None": - style_preset = None - - files = { - "image": image_binary - } - - response_api = await sync_op( - cls, - ApiEndpoint(path="/proxy/stability/v2beta/stable-image/generate/ultra", method="POST"), - response_model=StabilityStableUltraResponse, - data=StabilityStableUltraRequest( - prompt=prompt, - negative_prompt=negative_prompt, - aspect_ratio=aspect_ratio, - seed=seed, - strength=image_denoise, - style_preset=style_preset, - ), - files=files, - content_type="multipart/form-data", - ) - - if response_api.finish_reason != "SUCCESS": - raise Exception(f"Stable Image Ultra generation failed: {response_api.finish_reason}.") - - image_data = base64.b64decode(response_api.image) - returned_image = bytesio_to_image_tensor(BytesIO(image_data)) - - return IO.NodeOutput(returned_image) - - -class StabilityStableImageSD_3_5Node(IO.ComfyNode): - """ - Generates images synchronously based on prompt and resolution. - """ - - @classmethod - def define_schema(cls): - return IO.Schema( - node_id="StabilityStableImageSD_3_5Node", - display_name="Stability AI Stable Diffusion 3.5 Image", - category="partner/image/Stability AI", - description=cleandoc(cls.__doc__ or ""), - inputs=[ - IO.String.Input( - "prompt", - multiline=True, - default="", - tooltip="What you wish to see in the output image. A strong, descriptive prompt that clearly defines elements, colors, and subjects will lead to better results.", - ), - IO.Combo.Input( - "model", - options=Stability_SD3_5_Model, - ), - IO.Combo.Input( - "aspect_ratio", - options=StabilityAspectRatio, - default=StabilityAspectRatio.ratio_1_1, - tooltip="Aspect ratio of generated image.", - ), - IO.Combo.Input( - "style_preset", - options=get_stability_style_presets(), - tooltip="Optional desired style of generated image.", - advanced=True, - ), - IO.Float.Input( - "cfg_scale", - default=4.0, - min=1.0, - max=10.0, - step=0.1, - tooltip="How strictly the diffusion process adheres to the prompt text (higher values keep your image closer to your prompt)", - ), - IO.Int.Input( - "seed", - default=0, - min=0, - max=4294967294, - step=1, - display_mode=IO.NumberDisplay.number, - control_after_generate=True, - tooltip="The random seed used for creating the noise.", - ), - IO.Image.Input( - "image", - optional=True, - ), - IO.String.Input( - "negative_prompt", - default="", - tooltip="Keywords of what you do not wish to see in the output image. This is an advanced feature.", - force_input=True, - optional=True, - advanced=True, - ), - IO.Float.Input( - "image_denoise", - default=0.5, - min=0.0, - max=1.0, - step=0.01, - tooltip="Denoise of input image; 0.0 yields image identical to input, 1.0 is as if no image was provided at all.", - optional=True, - ), - ], - outputs=[ - IO.Image.Output(), - ], - hidden=[ - IO.Hidden.auth_token_comfy_org, - IO.Hidden.api_key_comfy_org, - IO.Hidden.unique_id, - ], - is_api_node=True, - price_badge=IO.PriceBadge( - depends_on=IO.PriceBadgeDepends(widgets=["model"]), - expr=""" - ( - $contains(widgets.model,"large") - ? {"type":"usd","usd":0.065} - : {"type":"usd","usd":0.035} - ) - """, - ), - ) - - @classmethod - async def execute( - cls, - model: str, - prompt: str, - aspect_ratio: str, - style_preset: str, - seed: int, - cfg_scale: float, - image: Optional[torch.Tensor] = None, - negative_prompt: str = "", - image_denoise: Optional[float] = 0.5, - ) -> IO.NodeOutput: - validate_string(prompt, strip_whitespace=False) - # prepare image binary if image present - image_binary = None - mode = Stability_SD3_5_GenerationMode.text_to_image - if image is not None: - image_binary = tensor_to_bytesio(image, total_pixels=1504*1504).read() - mode = Stability_SD3_5_GenerationMode.image_to_image - aspect_ratio = None - else: - image_denoise = None - - if not negative_prompt: - negative_prompt = None - if style_preset == "None": - style_preset = None - - files = { - "image": image_binary - } - - response_api = await sync_op( - cls, - ApiEndpoint(path="/proxy/stability/v2beta/stable-image/generate/sd3", method="POST"), - response_model=StabilityStableUltraResponse, - data=StabilityStable3_5Request( - prompt=prompt, - negative_prompt=negative_prompt, - aspect_ratio=aspect_ratio, - seed=seed, - strength=image_denoise, - style_preset=style_preset, - cfg_scale=cfg_scale, - model=model, - mode=mode, - ), - files=files, - content_type="multipart/form-data", - ) - - if response_api.finish_reason != "SUCCESS": - raise Exception(f"Stable Diffusion 3.5 Image generation failed: {response_api.finish_reason}.") - - image_data = base64.b64decode(response_api.image) - returned_image = bytesio_to_image_tensor(BytesIO(image_data)) - - return IO.NodeOutput(returned_image) - - -class StabilityUpscaleConservativeNode(IO.ComfyNode): - """ - Upscale image with minimal alterations to 4K resolution. - """ - - @classmethod - def define_schema(cls): - return IO.Schema( - node_id="StabilityUpscaleConservativeNode", - display_name="Stability AI Upscale Conservative", - category="partner/image/Stability AI", - description=cleandoc(cls.__doc__ or ""), - inputs=[ - IO.Image.Input("image"), - IO.String.Input( - "prompt", - multiline=True, - default="", - tooltip="What you wish to see in the output image. A strong, descriptive prompt that clearly defines elements, colors, and subjects will lead to better results.", - ), - IO.Float.Input( - "creativity", - default=0.35, - min=0.2, - max=0.5, - step=0.01, - tooltip="Controls the likelihood of creating additional details not heavily conditioned by the init image.", - ), - IO.Int.Input( - "seed", - default=0, - min=0, - max=4294967294, - step=1, - display_mode=IO.NumberDisplay.number, - control_after_generate=True, - tooltip="The random seed used for creating the noise.", - ), - IO.String.Input( - "negative_prompt", - default="", - tooltip="Keywords of what you do not wish to see in the output image. This is an advanced feature.", - force_input=True, - optional=True, - advanced=True, - ), - ], - outputs=[ - IO.Image.Output(), - ], - hidden=[ - IO.Hidden.auth_token_comfy_org, - IO.Hidden.api_key_comfy_org, - IO.Hidden.unique_id, - ], - is_api_node=True, - price_badge=IO.PriceBadge( - expr="""{"type":"usd","usd":0.4}""", - ), - ) - - @classmethod - async def execute( - cls, - image: torch.Tensor, - prompt: str, - creativity: float, - seed: int, - negative_prompt: str = "", - ) -> IO.NodeOutput: - validate_string(prompt, strip_whitespace=False) - image_binary = tensor_to_bytesio(image, total_pixels=1024*1024).read() - - if not negative_prompt: - negative_prompt = None - - files = { - "image": image_binary - } - - response_api = await sync_op( - cls, - ApiEndpoint(path="/proxy/stability/v2beta/stable-image/upscale/conservative", method="POST"), - response_model=StabilityStableUltraResponse, - data=StabilityUpscaleConservativeRequest( - prompt=prompt, - negative_prompt=negative_prompt, - creativity=round(creativity,2), - seed=seed, - ), - files=files, - content_type="multipart/form-data", - ) - - if response_api.finish_reason != "SUCCESS": - raise Exception(f"Stability Upscale Conservative generation failed: {response_api.finish_reason}.") - - image_data = base64.b64decode(response_api.image) - returned_image = bytesio_to_image_tensor(BytesIO(image_data)) - - return IO.NodeOutput(returned_image) - - -class StabilityUpscaleCreativeNode(IO.ComfyNode): - """ - Upscale image with minimal alterations to 4K resolution. - """ - - @classmethod - def define_schema(cls): - return IO.Schema( - node_id="StabilityUpscaleCreativeNode", - display_name="Stability AI Upscale Creative", - category="partner/image/Stability AI", - description=cleandoc(cls.__doc__ or ""), - inputs=[ - IO.Image.Input("image"), - IO.String.Input( - "prompt", - multiline=True, - default="", - tooltip="What you wish to see in the output image. A strong, descriptive prompt that clearly defines elements, colors, and subjects will lead to better results.", - ), - IO.Float.Input( - "creativity", - default=0.3, - min=0.1, - max=0.5, - step=0.01, - tooltip="Controls the likelihood of creating additional details not heavily conditioned by the init image.", - ), - IO.Combo.Input( - "style_preset", - options=get_stability_style_presets(), - tooltip="Optional desired style of generated image.", - advanced=True, - ), - IO.Int.Input( - "seed", - default=0, - min=0, - max=4294967294, - step=1, - display_mode=IO.NumberDisplay.number, - control_after_generate=True, - tooltip="The random seed used for creating the noise.", - ), - IO.String.Input( - "negative_prompt", - default="", - tooltip="Keywords of what you do not wish to see in the output image. This is an advanced feature.", - force_input=True, - optional=True, - advanced=True, - ), - ], - outputs=[ - IO.Image.Output(), - ], - hidden=[ - IO.Hidden.auth_token_comfy_org, - IO.Hidden.api_key_comfy_org, - IO.Hidden.unique_id, - ], - is_api_node=True, - price_badge=IO.PriceBadge( - expr="""{"type":"usd","usd":0.6}""", - ), - ) - - @classmethod - async def execute( - cls, - image: torch.Tensor, - prompt: str, - creativity: float, - style_preset: str, - seed: int, - negative_prompt: str = "", - ) -> IO.NodeOutput: - validate_string(prompt, strip_whitespace=False) - image_binary = tensor_to_bytesio(image, total_pixels=1024*1024).read() - - if not negative_prompt: - negative_prompt = None - if style_preset == "None": - style_preset = None - - files = { - "image": image_binary - } - - response_api = await sync_op( - cls, - ApiEndpoint(path="/proxy/stability/v2beta/stable-image/upscale/creative", method="POST"), - response_model=StabilityAsyncResponse, - data=StabilityUpscaleCreativeRequest( - prompt=prompt, - negative_prompt=negative_prompt, - creativity=round(creativity,2), - style_preset=style_preset, - seed=seed, - ), - files=files, - content_type="multipart/form-data", - ) - - response_poll = await poll_op( - cls, - ApiEndpoint(path=f"/proxy/stability/v2beta/results/{response_api.id}"), - response_model=StabilityResultsGetResponse, - poll_interval=3, - status_extractor=lambda x: get_async_dummy_status(x), - ) - - if response_poll.finish_reason != "SUCCESS": - raise Exception(f"Stability Upscale Creative generation failed: {response_poll.finish_reason}.") - - image_data = base64.b64decode(response_poll.result) - returned_image = bytesio_to_image_tensor(BytesIO(image_data)) - - return IO.NodeOutput(returned_image) - - -class StabilityUpscaleFastNode(IO.ComfyNode): - """ - Quickly upscales an image via Stability API call to 4x its original size; intended for upscaling low-quality/compressed images. - """ - - @classmethod - def define_schema(cls): - return IO.Schema( - node_id="StabilityUpscaleFastNode", - display_name="Stability AI Upscale Fast", - category="partner/image/Stability AI", - description=cleandoc(cls.__doc__ or ""), - inputs=[ - IO.Image.Input("image"), - ], - outputs=[ - IO.Image.Output(), - ], - hidden=[ - IO.Hidden.auth_token_comfy_org, - IO.Hidden.api_key_comfy_org, - IO.Hidden.unique_id, - ], - is_api_node=True, - price_badge=IO.PriceBadge( - expr="""{"type":"usd","usd":0.02}""", - ), - ) - - @classmethod - async def execute(cls, image: torch.Tensor) -> IO.NodeOutput: - image_binary = tensor_to_bytesio(image, total_pixels=4096*4096).read() - - files = { - "image": image_binary - } - - response_api = await sync_op( - cls, - ApiEndpoint(path="/proxy/stability/v2beta/stable-image/upscale/fast", method="POST"), - response_model=StabilityStableUltraResponse, - files=files, - content_type="multipart/form-data", - ) - - if response_api.finish_reason != "SUCCESS": - raise Exception(f"Stability Upscale Fast failed: {response_api.finish_reason}.") - - image_data = base64.b64decode(response_api.image) - returned_image = bytesio_to_image_tensor(BytesIO(image_data)) - - return IO.NodeOutput(returned_image) - - -class StabilityTextToAudio(IO.ComfyNode): - """Generates high-quality music and sound effects from text descriptions.""" - - @classmethod - def define_schema(cls): - return IO.Schema( - node_id="StabilityTextToAudio", - display_name="Stability AI Text To Audio", - category="partner/audio/Stability AI", - essentials_category="Audio", - description=cleandoc(cls.__doc__ or ""), - inputs=[ - IO.Combo.Input( - "model", - options=["stable-audio-2.5"], - ), - IO.String.Input("prompt", multiline=True, default=""), - IO.Int.Input( - "duration", - default=190, - min=1, - max=190, - step=1, - tooltip="Controls the duration in seconds of the generated audio.", - optional=True, - ), - IO.Int.Input( - "seed", - default=0, - min=0, - max=4294967294, - step=1, - display_mode=IO.NumberDisplay.number, - control_after_generate=True, - tooltip="The random seed used for generation.", - optional=True, - ), - IO.Int.Input( - "steps", - default=8, - min=4, - max=8, - step=1, - tooltip="Controls the number of sampling steps.", - optional=True, - advanced=True, - ), - ], - outputs=[ - IO.Audio.Output(), - ], - hidden=[ - IO.Hidden.auth_token_comfy_org, - IO.Hidden.api_key_comfy_org, - IO.Hidden.unique_id, - ], - is_api_node=True, - price_badge=IO.PriceBadge( - expr="""{"type":"usd","usd":0.2}""", - ), - ) - - @classmethod - async def execute(cls, model: str, prompt: str, duration: int, seed: int, steps: int) -> IO.NodeOutput: - validate_string(prompt, max_length=10000) - payload = StabilityTextToAudioRequest(prompt=prompt, model=model, duration=duration, seed=seed, steps=steps) - response_api = await sync_op( - cls, - ApiEndpoint(path="/proxy/stability/v2beta/audio/stable-audio-2/text-to-audio", method="POST"), - response_model=StabilityAudioResponse, - data=payload, - content_type="multipart/form-data", - ) - if not response_api.audio: - raise ValueError("No audio file was received in response.") - return IO.NodeOutput(audio_bytes_to_audio_input(base64.b64decode(response_api.audio))) - - -class StabilityAudioToAudio(IO.ComfyNode): - """Transforms existing audio samples into new high-quality compositions using text instructions.""" - - @classmethod - def define_schema(cls): - return IO.Schema( - node_id="StabilityAudioToAudio", - display_name="Stability AI Audio To Audio", - category="partner/audio/Stability AI", - description=cleandoc(cls.__doc__ or ""), - inputs=[ - IO.Combo.Input( - "model", - options=["stable-audio-2.5"], - ), - IO.String.Input("prompt", multiline=True, default=""), - IO.Audio.Input("audio", tooltip="Audio must be between 6 and 190 seconds long."), - IO.Int.Input( - "duration", - default=190, - min=1, - max=190, - step=1, - tooltip="Controls the duration in seconds of the generated audio.", - optional=True, - ), - IO.Int.Input( - "seed", - default=0, - min=0, - max=4294967294, - step=1, - display_mode=IO.NumberDisplay.number, - control_after_generate=True, - tooltip="The random seed used for generation.", - optional=True, - ), - IO.Int.Input( - "steps", - default=8, - min=4, - max=8, - step=1, - tooltip="Controls the number of sampling steps.", - optional=True, - advanced=True, - ), - IO.Float.Input( - "strength", - default=1, - min=0.01, - max=1.0, - step=0.01, - display_mode=IO.NumberDisplay.slider, - tooltip="Parameter controls how much influence the audio parameter has on the generated audio.", - optional=True, - ), - ], - outputs=[ - IO.Audio.Output(), - ], - hidden=[ - IO.Hidden.auth_token_comfy_org, - IO.Hidden.api_key_comfy_org, - IO.Hidden.unique_id, - ], - is_api_node=True, - price_badge=IO.PriceBadge( - expr="""{"type":"usd","usd":0.2}""", - ), - ) - - @classmethod - async def execute( - cls, model: str, prompt: str, audio: Input.Audio, duration: int, seed: int, steps: int, strength: float - ) -> IO.NodeOutput: - validate_string(prompt, max_length=10000) - validate_audio_duration(audio, 6, 190) - payload = StabilityAudioToAudioRequest( - prompt=prompt, model=model, duration=duration, seed=seed, steps=steps, strength=strength - ) - response_api = await sync_op( - cls, - ApiEndpoint(path="/proxy/stability/v2beta/audio/stable-audio-2/audio-to-audio", method="POST"), - response_model=StabilityAudioResponse, - data=payload, - content_type="multipart/form-data", - files={"audio": audio_input_to_mp3(audio)}, - ) - if not response_api.audio: - raise ValueError("No audio file was received in response.") - return IO.NodeOutput(audio_bytes_to_audio_input(base64.b64decode(response_api.audio))) - - -class StabilityAudioInpaint(IO.ComfyNode): - """Transforms part of existing audio sample using text instructions.""" - - @classmethod - def define_schema(cls): - return IO.Schema( - node_id="StabilityAudioInpaint", - display_name="Stability AI Audio Inpaint", - category="partner/audio/Stability AI", - description=cleandoc(cls.__doc__ or ""), - inputs=[ - IO.Combo.Input( - "model", - options=["stable-audio-2.5"], - ), - IO.String.Input("prompt", multiline=True, default=""), - IO.Audio.Input("audio", tooltip="Audio must be between 6 and 190 seconds long."), - IO.Int.Input( - "duration", - default=190, - min=1, - max=190, - step=1, - tooltip="Controls the duration in seconds of the generated audio.", - optional=True, - ), - IO.Int.Input( - "seed", - default=0, - min=0, - max=4294967294, - step=1, - display_mode=IO.NumberDisplay.number, - control_after_generate=True, - tooltip="The random seed used for generation.", - optional=True, - ), - IO.Int.Input( - "steps", - default=8, - min=4, - max=8, - step=1, - tooltip="Controls the number of sampling steps.", - optional=True, - advanced=True, - ), - IO.Int.Input( - "mask_start", - default=30, - min=0, - max=190, - step=1, - optional=True, - advanced=True, - ), - IO.Int.Input( - "mask_end", - default=190, - min=0, - max=190, - step=1, - optional=True, - advanced=True, - ), - ], - outputs=[ - IO.Audio.Output(), - ], - hidden=[ - IO.Hidden.auth_token_comfy_org, - IO.Hidden.api_key_comfy_org, - IO.Hidden.unique_id, - ], - is_api_node=True, - price_badge=IO.PriceBadge( - expr="""{"type":"usd","usd":0.2}""", - ), - ) - - @classmethod - async def execute( - cls, - model: str, - prompt: str, - audio: Input.Audio, - duration: int, - seed: int, - steps: int, - mask_start: int, - mask_end: int, - ) -> IO.NodeOutput: - validate_string(prompt, max_length=10000) - if mask_end <= mask_start: - raise ValueError(f"Value of mask_end({mask_end}) should be greater then mask_start({mask_start})") - validate_audio_duration(audio, 6, 190) - - payload = StabilityAudioInpaintRequest( - prompt=prompt, - model=model, - duration=duration, - seed=seed, - steps=steps, - mask_start=mask_start, - mask_end=mask_end, - ) - response_api = await sync_op( - cls, - endpoint=ApiEndpoint(path="/proxy/stability/v2beta/audio/stable-audio-2/inpaint", method="POST"), - response_model=StabilityAudioResponse, - data=payload, - content_type="multipart/form-data", - files={"audio": audio_input_to_mp3(audio)}, - ) - if not response_api.audio: - raise ValueError("No audio file was received in response.") - return IO.NodeOutput(audio_bytes_to_audio_input(base64.b64decode(response_api.audio))) - - -class StabilityExtension(ComfyExtension): - @override - async def get_node_list(self) -> list[type[IO.ComfyNode]]: - return [ - StabilityStableImageUltraNode, - StabilityStableImageSD_3_5Node, - StabilityUpscaleConservativeNode, - StabilityUpscaleCreativeNode, - StabilityUpscaleFastNode, - StabilityTextToAudio, - StabilityAudioToAudio, - StabilityAudioInpaint, - ] - - -async def comfy_entrypoint() -> StabilityExtension: - return StabilityExtension() diff --git a/comfy_api_nodes/util/__init__.py b/comfy_api_nodes/util/__init__.py index 25cb88869..1fb6b96cf 100644 --- a/comfy_api_nodes/util/__init__.py +++ b/comfy_api_nodes/util/__init__.py @@ -26,6 +26,7 @@ from .conversions import ( text_filepath_to_base64_string, text_filepath_to_data_uri, trim_video, + upscale_image_tensor_to_min_pixels, upscale_video_to_min_pixels, video_to_base64_string, ) @@ -99,6 +100,7 @@ __all__ = [ "text_filepath_to_base64_string", "text_filepath_to_data_uri", "trim_video", + "upscale_image_tensor_to_min_pixels", "upscale_video_to_min_pixels", "video_to_base64_string", # Validation utilities diff --git a/comfy_api_nodes/util/conversions.py b/comfy_api_nodes/util/conversions.py index a1b5d599c..9cd644fc0 100644 --- a/comfy_api_nodes/util/conversions.py +++ b/comfy_api_nodes/util/conversions.py @@ -448,6 +448,15 @@ def _compute_upscale_dims(src_w: int, src_h: int, total_pixels: int) -> tuple[in return new_w, new_h +def upscale_image_tensor_to_min_pixels(image: torch.Tensor, total_pixels: int) -> torch.Tensor: + samples = image.movedim(-1, 1) + dims = _compute_upscale_dims(samples.shape[3], samples.shape[2], int(total_pixels)) + if dims is None: + return image + new_w, new_h = dims + return common_upscale(samples, new_w, new_h, "lanczos", "disabled").movedim(1, -1) + + def upscale_video_to_min_pixels(video: Input.Video, min_pixels: int) -> Input.Video: """Upscale a video to meet at least ``min_pixels`` (w * h), preserving aspect ratio. diff --git a/comfy_api_nodes/util/request_logger.py b/comfy_api_nodes/util/request_logger.py index fe0543d9b..70ecaf41a 100644 --- a/comfy_api_nodes/util/request_logger.py +++ b/comfy_api_nodes/util/request_logger.py @@ -9,6 +9,7 @@ from typing import Any import folder_paths logger = logging.getLogger(__name__) +_SENSITIVE_HEADERS = {"authorization", "x-api-key"} def get_log_directory(): @@ -73,6 +74,10 @@ def _format_data_for_logging(data: Any) -> str: return str(data) +def _redact_headers(headers: dict) -> dict: + return {k: ("***" if k.lower() in _SENSITIVE_HEADERS else v) for k, v in headers.items()} + + def log_request_response( operation_id: str, request_method: str, @@ -101,7 +106,7 @@ def log_request_response( log_content.append(f"Method: {request_method}") log_content.append(f"URL: {request_url}") if request_headers: - log_content.append(f"Headers:\n{_format_data_for_logging(request_headers)}") + log_content.append(f"Headers:\n{_format_data_for_logging(_redact_headers(request_headers))}") if request_params: log_content.append(f"Params:\n{_format_data_for_logging(request_params)}") if request_data is not None: diff --git a/comfy_api_nodes/util/upload_helpers.py b/comfy_api_nodes/util/upload_helpers.py index 6d1d107a1..f7029ee78 100644 --- a/comfy_api_nodes/util/upload_helpers.py +++ b/comfy_api_nodes/util/upload_helpers.py @@ -158,7 +158,14 @@ async def upload_video_to_comfyapi( # Convert VideoInput to BytesIO using specified container/codec video_bytes_io = BytesIO() - video.save_to(video_bytes_io, format=container, codec=codec) + try: + video.save_to(video_bytes_io, format=container, codec=codec) + except Exception as e: + raise ValueError( + f"Could not convert the input video to {container.value.upper()} for upload; " + f"the file may be corrupted or use an unsupported codec. " + f"Try re-exporting it as MP4 (H.264). Original error: {e}" + ) from e video_bytes_io.seek(0) return await upload_file_to_comfyapi(cls, video_bytes_io, filename, upload_mime_type, wait_label) diff --git a/comfy_execution/caching.py b/comfy_execution/caching.py index ba1e8bc84..ad75a0e50 100644 --- a/comfy_execution/caching.py +++ b/comfy_execution/caching.py @@ -503,6 +503,21 @@ RAM_CACHE_DEFAULT_RAM_USAGE = 0.05 RAM_CACHE_OLD_WORKFLOW_OOM_MULTIPLIER = 1.3 + +def all_outputs_dynamic(outputs): + if outputs is None: + return False + + for output in outputs: + if isinstance(output, (list, tuple)): + if not all_outputs_dynamic(output): + return False + elif not hasattr(output, "is_dynamic") or not output.is_dynamic(): + return False + + return True + + class RAMPressureCache(LRUCache): def __init__(self, key_class, enable_providers=False): @@ -533,7 +548,11 @@ class RAMPressureCache(LRUCache): for key, cache_entry in self.cache.items(): if not free_active and self.used_generation[key] == self.generation: continue - oom_score = RAM_CACHE_OLD_WORKFLOW_OOM_MULTIPLIER ** (self.generation - self.used_generation[key]) + + if all_outputs_dynamic(cache_entry.outputs) and self.used_generation[key] == self.generation: + continue + + oom_score = RAM_CACHE_OLD_WORKFLOW_OOM_MULTIPLIER ** (self.generation - self.used_generation[key]) ram_usage = RAM_CACHE_DEFAULT_RAM_USAGE def scan_list_for_ram_usage(outputs): diff --git a/comfy_extras/nodes_color.py b/comfy_extras/nodes_color.py index f58e51bff..6d10b26f4 100644 --- a/comfy_extras/nodes_color.py +++ b/comfy_extras/nodes_color.py @@ -16,23 +16,30 @@ class ColorToRGBInt(io.ComfyNode): ], outputs=[ io.Int.Output(display_name="rgb_int"), - io.Color.Output(display_name="hex") + io.Color.Output(display_name="hex"), + io.Float.Output(display_name="alpha"), ], ) @classmethod def execute(cls, color: str) -> io.NodeOutput: - # expect format #RRGGBB - if len(color) != 7 or color[0] != "#": - raise ValueError("Color must be in format #RRGGBB") + # expect format #RRGGBB or #RRGGBBAA + if len(color) not in (7, 9) or color[0] != "#": + raise ValueError("Color must be in format #RRGGBB or #RRGGBBAA") try: int(color[1:], 16) except ValueError: - raise ValueError("Color must be in format #RRGGBB") from None + raise ValueError("Color must be in format #RRGGBB or #RRGGBBAA") from None + + alpha = 1.0 + if len(color) == 9: + alpha = int(color[7:9], 16) / 255.0 + color = color[:7] + r, g, b = hex_to_rgb(color) rgb_int = r * 256 * 256 + g * 256 + b - return io.NodeOutput(rgb_int, color) + return io.NodeOutput(rgb_int, color, alpha) class ColorExtension(ComfyExtension): diff --git a/comfy_extras/nodes_text_overlay.py b/comfy_extras/nodes_text_overlay.py new file mode 100644 index 000000000..4c5cdae60 --- /dev/null +++ b/comfy_extras/nodes_text_overlay.py @@ -0,0 +1,150 @@ +import numpy as np +import torch +from PIL import Image as PILImage, ImageColor, ImageDraw, ImageFont +from typing_extensions import override + +from comfy_api.latest import ComfyExtension, IO + + +class TextOverlay(IO.ComfyNode): + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="TextOverlay", + display_name="Draw Text Overlay", + category="text", + description="Draw text overlay on an image or batch of images.", + search_aliases=["text", "label", "caption", "subtitle", "watermark", "title", "addlabel", "overlay"], + inputs=[ + IO.Image.Input("images"), + IO.String.Input("text", multiline=True, default=""), + IO.Float.Input("font_size", default=5.0, min=0.5, max=50.0, step=0.5, tooltip="Font size as a percentage of the image height."), + IO.Color.Input("color", default="#ffffff", tooltip="Color of the text."), + IO.Combo.Input("position", options=["top", "bottom"], default="top"), + IO.Combo.Input("align", options=["left", "center", "right"], default="left"), + IO.Boolean.Input("outline", default=True, tooltip="Draw a black outline around the text."), + ], + outputs=[IO.Image.Output(display_name="images")], + ) + + @classmethod + def execute(cls, images, text, font_size, color, position, align, outline) -> IO.NodeOutput: + if text.strip() == "": + return IO.NodeOutput(images) + + text = text.replace("\\n", "\n").replace("\\t", "\t") + + text_rgba = cls.parse_color_to_rgba(color) + outline_rgba = (0, 0, 0, 255) if outline else (0, 0, 0, 0) + + # Render the overlay once and composite it across all frames in the batch + height = images.shape[1] + width = images.shape[2] + overlay_rgb, overlay_alpha = cls.render_overlay_text(width, height, text, position, align, font_size, text_rgba, outline_rgba) + overlay_rgb = overlay_rgb.to(device=images.device, dtype=images.dtype) + overlay_alpha = overlay_alpha.to(device=images.device, dtype=images.dtype) + + result = images * (1.0 - overlay_alpha) + overlay_rgb * overlay_alpha + return IO.NodeOutput(result) + + @staticmethod + def parse_color_to_rgba(color_string): + parsed = ImageColor.getrgb(color_string) + + if len(parsed) == 3: + return (*parsed, 255) + + return parsed + + @classmethod + def render_overlay_text(cls, width, height, text, position, align, font_size, text_rgba, outline_rgba): + line_spacing = 1.2 + margin_percent = 1.0 + min_font_percent = 2.0 + min_font_pixels = 10 + outline_thickness_factor = 0.04 + + # Draw onto a transparent layer so the result can be alpha-composited over any frame. + layer = PILImage.new("RGBA", (width, height), (0, 0, 0, 0)) + draw = ImageDraw.Draw(layer) + + margin = int(round(margin_percent / 100.0 * min(width, height))) + max_width = max(1, width - 2 * margin) + max_height = max(1, height - 2 * margin) + + # Font scales with resolution, then shrinks to fit the height. + size = max(1, int(round(font_size / 100.0 * height))) + floor = min(size, max(min_font_pixels, int(round(min_font_percent / 100.0 * height)))) + + while True: + font = ImageFont.load_default(size=size) + stroke = max(1, int(round(size * outline_thickness_factor))) if outline_rgba[3] > 0 else 0 + block = "\n".join(cls.wrap_text(text, font, max_width)) + # convert line spacing to pixel spacing + single = draw.textbbox((0, 0), "Ay", font=font, stroke_width=stroke) + double = draw.multiline_textbbox((0, 0), "Ay\nAy", font=font, spacing=0, stroke_width=stroke) + natural_advance = (double[3] - double[1]) - (single[3] - single[1]) + pixel_spacing = int(round(size * line_spacing - natural_advance)) + box = draw.multiline_textbbox((0, 0), block, font=font, spacing=pixel_spacing, stroke_width=stroke) + block_height = box[3] - box[1] + + if block_height <= max_height or size <= floor: + break + + size = max(floor, int(size * 0.9)) + + anchor_h, x = {"left": ("l", margin), "center": ("m", width / 2), "right": ("r", width - margin)}[align] + + # Offset y so the rendered text sits flush against the margin + if position == "bottom": + y = height - margin - box[3] + else: + y = margin - box[1] + + draw.multiline_text((x, y), block, font=font, fill=text_rgba, anchor=anchor_h + "a", + align=align, spacing=pixel_spacing, stroke_width=stroke, stroke_fill=outline_rgba) + + overlay = np.array(layer).astype(np.float32) / 255.0 + overlay_rgb = torch.from_numpy(overlay[:, :, :3]) + overlay_alpha = torch.from_numpy(overlay[:, :, 3:4]) + return overlay_rgb, overlay_alpha + + @staticmethod + def wrap_text(text, font, max_width): + lines = [] + for raw_line in text.split("\n"): + words = raw_line.split() + if not words: + lines.append("") + continue + current = "" + # Break the line into words and split words that are too long + for word in words: + while font.getlength(word) > max_width and len(word) > 1: + cut = 1 + while cut < len(word) and font.getlength(word[:cut + 1]) <= max_width: + cut += 1 + if current: + lines.append(current) + current = "" + lines.append(word[:cut]) + word = word[cut:] + candidate = word if not current else current + " " + word + if not current or font.getlength(candidate) <= max_width: + current = candidate + else: + lines.append(current) + current = word + if current: + lines.append(current) + return lines + + +class TextOverlayExtension(ComfyExtension): + @override + async def get_node_list(self) -> list[type[IO.ComfyNode]]: + return [TextOverlay] + + +async def comfy_entrypoint() -> TextOverlayExtension: + return TextOverlayExtension() diff --git a/folder_paths.py b/folder_paths.py index 7304e1b73..937428c18 100644 --- a/folder_paths.py +++ b/folder_paths.py @@ -17,7 +17,11 @@ if args.base_directory: else: base_path = os.path.dirname(os.path.realpath(__file__)) -models_dir = os.path.join(base_path, "models") +if args.models_directory: + models_dir = os.path.abspath(args.models_directory) +else: + models_dir = os.path.join(base_path, "models") + folder_names_and_paths["checkpoints"] = ([os.path.join(models_dir, "checkpoints")], supported_pt_extensions) folder_names_and_paths["configs"] = ([os.path.join(models_dir, "configs")], [".yaml"]) @@ -264,6 +268,59 @@ def annotated_filepath(name: str) -> tuple[str, str | None]: return name, base_dir +# Content types a browser may execute or render inline. File endpoints that +# serve user-controlled content must force these to download (and ideally set +# Content-Disposition: attachment) to avoid stored XSS. Centralised here so the +# /view and /userdata handlers can't drift apart. mimetypes.guess_type may +# return either the text/* or application/* spelling depending on platform, so +# both are listed. +DANGEROUS_CONTENT_TYPES = { + 'text/html', 'text/html-sandboxed', 'application/xhtml+xml', + 'text/javascript', 'application/javascript', 'application/x-javascript', + 'application/ecmascript', 'text/css', + 'image/svg+xml', 'application/xml', 'text/xml', + # message/rfc822 (.mht/.mhtml) can carry script in some browsers. + 'message/rfc822', +} + + +def is_dangerous_content_type(content_type: str | None) -> bool: + """Return True if a browser may execute or render `content_type` inline. + + Normalises before matching so the check can't be slipped past with a + charset/boundary parameter (``text/html; charset=utf-8``) or casing + (``TEXT/HTML``). Any XML dialect (``*+xml`` or ``*/xml``) is treated as + dangerous because XML can carry inline script via stylesheet/entity tricks, + which also covers the ``application/{xslt,rss,atom,rdf}+xml`` family without + enumerating each one. Endpoints serving user-controlled content should route + a dangerous type to ``application/octet-stream`` + ``Content-Disposition: + attachment`` + ``X-Content-Type-Options: nosniff``. + """ + if not content_type: + return False + normalized = content_type.split(';', 1)[0].strip().lower() + if normalized in DANGEROUS_CONTENT_TYPES: + return True + return normalized.endswith('+xml') or normalized.endswith('/xml') + + +def is_within_directory(directory: str, target: str) -> bool: + """Return True if `target` resolves to a path inside `directory`. + + Uses realpath on both operands so that a symlink placed inside `directory` + that points elsewhere cannot escape the containment check at open time. + """ + try: + directory = os.path.realpath(directory) + target = os.path.realpath(target) + return os.path.commonpath((directory, target)) == directory + except ValueError: + # ValueError is raised by realpath() on a path with an embedded null + # byte, and by commonpath() on Windows when the paths are on different + # drives. In either case the target is not safely within the directory. + return False + + def get_annotated_filepath(name: str, default_dir: str | None=None) -> str: name, base_dir = annotated_filepath(name) @@ -273,7 +330,12 @@ def get_annotated_filepath(name: str, default_dir: str | None=None) -> str: else: base_dir = get_input_directory() # fallback path - return os.path.join(base_dir, name) + filepath = os.path.abspath(os.path.join(base_dir, name)) + # Prevent path traversal: the resolved path must stay within base_dir. + # repr() the name in the message so a crafted value can't inject log lines. + if not is_within_directory(base_dir, filepath): + raise ValueError("Invalid file path: {!r}".format(name)) + return filepath def exists_annotated_filepath(name) -> bool: @@ -282,7 +344,10 @@ def exists_annotated_filepath(name) -> bool: if base_dir is None: base_dir = get_input_directory() # fallback path - filepath = os.path.join(base_dir, name) + filepath = os.path.abspath(os.path.join(base_dir, name)) + # Treat traversal attempts as non-existent rather than probing the filesystem. + if not is_within_directory(base_dir, filepath): + return False return os.path.exists(filepath) diff --git a/main.py b/main.py index d3f79687e..24a75cf78 100644 --- a/main.py +++ b/main.py @@ -132,6 +132,10 @@ def apply_custom_paths(): if args.base_directory: logging.info(f"Setting base directory to: {folder_paths.base_path}") + # --models-directory + if args.models_directory: + logging.info(f"Setting models directory to: {folder_paths.models_dir}") + # --output-directory, --input-directory, --user-directory if args.output_directory: output_dir = os.path.abspath(args.output_directory) diff --git a/nodes.py b/nodes.py index 9043a8d0a..e126576fe 100644 --- a/nodes.py +++ b/nodes.py @@ -2478,6 +2478,7 @@ async def init_builtin_extra_nodes(): "nodes_glsl.py", "nodes_lora_debug.py", "nodes_textgen.py", + "nodes_text_overlay.py", "nodes_color.py", "nodes_toolkit.py", "nodes_replacements.py", diff --git a/openapi.yaml b/openapi.yaml index c6a8621cc..0cf177815 100644 --- a/openapi.yaml +++ b/openapi.yaml @@ -7,18 +7,18 @@ components: description: Timestamp when the asset was created format: date-time type: string - display_name: - description: Display name of the asset. Mirrors name for backwards compatibility. - nullable: true - type: string - file_path: - description: Relative path in global-namespace-root form (e.g. "models/checkpoints/flux.safetensors") - nullable: true - type: string hash: description: Blake3 hash of the asset content. pattern: ^blake3:[a-f0-9]{64}$ type: string + loader_path: + description: The value a loader consumes to load this asset. Null when no loader can resolve the file. + nullable: true + type: string + display_name: + description: Human-facing label for the asset. Not unique. + nullable: true + type: string id: description: Unique identifier for the asset format: uuid @@ -144,14 +144,6 @@ components: AssetUpdated: description: Response returned when an existing asset is successfully updated. properties: - display_name: - description: Display name of the asset. Mirrors name for backwards compatibility. - nullable: true - type: string - file_path: - description: Relative path in global-namespace-root form (e.g. "models/checkpoints/flux.safetensors") - nullable: true - type: string hash: description: Blake3 hash of the asset content. pattern: ^blake3:[a-f0-9]{64}$ @@ -1644,7 +1636,7 @@ paths: format: uuid type: string tags: - description: JSON-encoded array of freeform tag strings, e.g. '["models","checkpoint"]'. Common types include "models", "input", "output", and "temp", but any tag can be used in any order. + description: JSON-encoded array of tag strings. For new byte uploads, include exactly one destination role (`input`, `output`, or `models`); `models` uploads also require exactly one `model_type:` tag. Extra tags are stored as labels and do not create path components. type: string user_metadata: description: Custom JSON metadata as a string @@ -1829,7 +1821,7 @@ paths: content: application/json: schema: - $ref: '#/components/schemas/AssetUpdated' + $ref: '#/components/schemas/Asset' description: Asset updated successfully "400": content: @@ -2470,6 +2462,9 @@ 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: diff --git a/requirements.txt b/requirements.txt index 1d9fe4137..e72f3045b 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,6 +1,6 @@ comfyui-frontend-package==1.45.20 -comfyui-workflow-templates==0.11.1 -comfyui-embedded-docs==0.5.6 +comfyui-workflow-templates==0.11.6 +comfyui-embedded-docs==0.5.7 torch torchsde torchvision diff --git a/server.py b/server.py index 0837c246e..a4c7d7204 100644 --- a/server.py +++ b/server.py @@ -46,6 +46,7 @@ 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 @@ -127,6 +128,7 @@ def create_cors_middleware(allowed_origin: str): return cors_middleware + def is_loopback(host): if host is None: return False @@ -435,7 +437,9 @@ class PromptServer(): try: tag = image_upload_type if image_upload_type in ("input", "output") else "input" - result = register_file_in_place(abs_path=filepath, name=filename, tags=[tag]) + tags = [tag] + tags.extend(get_known_subfolder_tags(subfolder)) + result = register_file_in_place(abs_path=filepath, name=filename, tags=tags) resp["asset"] = { "id": result.ref.id, "name": result.ref.name, @@ -611,15 +615,30 @@ class PromptServer(): or 'application/octet-stream' ) - # For security, force certain mimetypes to download instead of display - if content_type in {'text/html', 'text/html-sandboxed', 'application/xhtml+xml', 'text/javascript', 'text/css'}: - content_type = 'application/octet-stream' # Forces download + # For security, force renderable/active types (HTML, JS, + # CSS, SVG, XML — anything that can carry inline ' + files = {"file": ("evil.svg", svg, "image/svg+xml")} + form_data = { + "tags": json.dumps(["models", "model_type:checkpoints", "unit-tests", "svgxss"]), + "name": "evil.svg", + } + up = http.post(api_base + "/api/assets", files=files, data=form_data, timeout=120) + body = up.json() + assert up.status_code in (200, 201), body + aid = body["id"] + try: + r = http.get(f"{api_base}/api/assets/{aid}/content?disposition=inline", timeout=120) + r.content + assert r.status_code == 200 + ct = r.headers.get("Content-Type", "").lower() + cd = r.headers.get("Content-Disposition", "").lower() + assert "svg" not in ct, f"SVG served with a renderable content type: {ct!r}" + assert ct.startswith("application/octet-stream"), f"expected octet-stream, got {ct!r}" + assert "attachment" in cd, f"inline disposition not overridden to attachment: {cd!r}" + assert r.headers.get("X-Content-Type-Options", "").lower() == "nosniff" + finally: + with contextlib.suppress(Exception): + http.delete(f"{api_base}/api/assets/{aid}", timeout=30) + + def test_download_attachment_and_inline(http: requests.Session, api_base: str, seeded_asset: dict): aid = seeded_asset["id"] @@ -95,7 +131,7 @@ def test_download_chooses_existing_state_and_updates_access_time( assert t1 > t0 -@pytest.mark.parametrize("seeded_asset", [{"tags": ["models", "checkpoints"]}], indirect=True) +@pytest.mark.parametrize("seeded_asset", [{"tags": ["models", "model_type:checkpoints"]}], indirect=True) def test_download_missing_file_returns_404( http: requests.Session, api_base: str, comfy_tmp_base_dir: Path, seeded_asset: dict ): diff --git a/tests-unit/assets_test/test_list_cursor.py b/tests-unit/assets_test/test_list_cursor.py index a37019fd6..8f4cc8251 100644 --- a/tests-unit/assets_test/test_list_cursor.py +++ b/tests-unit/assets_test/test_list_cursor.py @@ -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", "checkpoints", "unit-tests", tag], + ["models", "model_type: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", "checkpoints", "unit-tests", f"cursor-{sort_field}"], {}, make_asset_bytes(n, size=2048 + i)) + asset_factory(n, ["models", "model_type:checkpoints", "unit-tests", f"cursor-{sort_field}"], {}, make_asset_bytes(n, size=2048 + i)) names.append(n) params = { diff --git a/tests-unit/assets_test/test_list_filter.py b/tests-unit/assets_test/test_list_filter.py index 17bbea5c6..d1cba87b3 100644 --- a/tests-unit/assets_test/test_list_filter.py +++ b/tests-unit/assets_test/test_list_filter.py @@ -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", "checkpoints", "unit-tests", "paging"], + ["models", "model_type: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", "checkpoints", "unit-tests", "alpha"], {}, b"X" * 1024) - b = asset_factory("inc_b.safetensors", ["models", "checkpoints", "unit-tests", "beta"], {}, b"Y" * 1024) + 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) 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", "checkpoints", "unit-tests", "lf-size"] + t = ["models", "model_type: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", "checkpoints", "unit-tests", "lf-upd"] + t = ["models", "model_type: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", "checkpoints", "unit-tests", "lf-access"] + t = ["models", "model_type: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", "checkpoints", "unit-tests", "lf-include"] + t = ["models", "model_type: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 + case-insensitive + # CSV tag filters are whitespace-trimmed and case-sensitive. 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", "checkpoints", "unit-tests", "lf-exclude"] + t = ["models", "model_type: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 uppercase should work + # Exclude filters are case-sensitive. 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", "checkpoints", "unit-tests", "lf-name"] + t = ["models", "model_type: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", "checkpoints", "unit-tests", "lf-pagelimits"] + t = ["models", "model_type: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", "checkpoints", "unit-tests", scope] + tags = ["models", "model_type: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)) diff --git a/tests-unit/assets_test/test_metadata_filters.py b/tests-unit/assets_test/test_metadata_filters.py index 20285a3b3..1864b1eef 100644 --- a/tests-unit/assets_test/test_metadata_filters.py +++ b/tests-unit/assets_test/test_metadata_filters.py @@ -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", "checkpoints", "unit-tests", "mf-and"] + tags = ["models", "model_type: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", "checkpoints", "unit-tests", "mf-types"] + tags = ["models", "model_type: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", "checkpoints", "unit-tests", "mf-list"] + tags = ["models", "model_type: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", "checkpoints", "unit-tests", "mf-none"] + t = ["models", "model_type: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", "checkpoints", "unit-tests", "mf-nested"] + tags = ["models", "model_type: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", "checkpoints", "unit-tests", "mf-objlist"] + tags = ["models", "model_type: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", "checkpoints", "unit-tests", "mf-keys"] + tags = ["models", "model_type: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", "checkpoints", "unit-tests", "mf-zero-bool"] + t = ["models", "model_type: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", "checkpoints", "unit-tests", "mf-mixed"] + tags = ["models", "model_type: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", "checkpoints", "unit-tests", "mf-unknown-scope"] + t = ["models", "model_type: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", "checkpoints", "unit-tests", "mf-tag", "alpha"], + ["models", "model_type:checkpoints", "unit-tests", "mf-tag", "alpha"], {"epoch": 1}, make_asset_bytes("alpha"), ) b = asset_factory( "mf_tag_beta.safetensors", - ["models", "checkpoints", "unit-tests", "mf-tag", "beta"], + ["models", "model_type: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", "checkpoints", "unit-tests", "mf-sort"] + t = ["models", "model_type: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)) diff --git a/tests-unit/assets_test/test_prune_orphaned_assets.py b/tests-unit/assets_test/test_prune_orphaned_assets.py index 1fbd4d4e2..618ec6c8d 100644 --- a/tests-unit/assets_test/test_prune_orphaned_assets.py +++ b/tests-unit/assets_test/test_prune_orphaned_assets.py @@ -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 = {"include_tags": f"unit-tests,{scope}"} + params = {"limit": "500"} 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" / "unit-tests" / scope / get_asset_filename(a["asset_hash"], ".bin") + path = comfy_tmp_base_dir / "input" / get_asset_filename(a["asset_hash"], ".bin") path.unlink() trigger_sync_seed_assets(http, api_base) @@ -108,18 +108,20 @@ def test_prune_across_multiple_roots( ): """Prune correctly handles assets across input and output roots.""" scope = f"multi-{uuid.uuid4().hex[:6]}" - input_fp = create_seed_file("input", scope, "input.bin") - create_seed_file("output", scope, "output.bin") + 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) trigger_sync_seed_assets(http, api_base) - assert len(find_asset(scope)) == 2 + assert find_asset(scope, input_name) + assert find_asset(scope, output_name) input_fp.unlink() trigger_sync_seed_assets(http, api_base) - remaining = find_asset(scope) - assert len(remaining) == 1 - assert remaining[0]["name"] == "output.bin" + assert not find_asset(scope, input_name) + assert find_asset(scope, output_name) @pytest.mark.parametrize("dirname", ["100%_done", "my_folder_name", "has spaces"]) diff --git a/tests-unit/assets_test/test_tags_api.py b/tests-unit/assets_test/test_tags_api.py index 9729b7d03..93786696f 100644 --- a/tests-unit/assets_test/test_tags_api.py +++ b/tests-unit/assets_test/test_tags_api.py @@ -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 system tags from migration should exist: + # A few selected contract tags should exist. assert "models" in names - assert "checkpoints" in names + assert "model_type: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 "checkpoints" in used_names + assert "model_type: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 "checkpoints" in names + assert "models" in names and "model_type: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 and mixed case - payload_add = {"tags": ["NewTag", "unit-tests", "newtag", "BETA"]} + # Add tags with duplicates while preserving source 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 - # normalized, deduplicated; 'unit-tests' was already present from the seed - assert set(b1["added"]) == {"newtag", "beta"} + # stripped, 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,8 +118,44 @@ 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 + 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"] def test_tags_list_order_and_prefix(http: requests.Session, api_base: str, seeded_asset: dict): diff --git a/tests-unit/assets_test/test_uploads.py b/tests-unit/assets_test/test_uploads.py index 427a417cc..7be7b0935 100644 --- a/tests-unit/assets_test/test_uploads.py +++ b/tests-unit/assets_test/test_uploads.py @@ -1,11 +1,14 @@ 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(): @@ -20,9 +23,18 @@ 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", "checkpoints", "unit-tests", "alpha"] + tags = ["models", "model_type:checkpoints", "unit-tests", "alpha"] meta = {"purpose": "dup"} data = make_asset_bytes(name) files = {"file": (name, data, "application/octet-stream")} @@ -43,6 +55,8 @@ 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")} @@ -53,12 +67,14 @@ 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 = ["models", "checkpoints", "unit-tests"] + tags = ["input", "unit-tests"] meta = {} files = {"file": (name, b"B" * 1024, "application/octet-stream")} form = {"tags": json.dumps(tags), "name": name, "user_metadata": json.dumps(meta)} @@ -69,9 +85,10 @@ 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(tags))), + ("tags", (None, json.dumps(hash_only_tags))), ("name", (None, "fastpath_copy.safetensors")), ("user_metadata", (None, json.dumps({"purpose": "copy"}))), ] @@ -81,6 +98,53 @@ 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( @@ -88,7 +152,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", "checkpoints", "unit-tests", "fp"]), "name": "seed.safetensors", "user_metadata": json.dumps({})} + form = {"tags": json.dumps(["models", "model_type: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 @@ -104,11 +168,49 @@ 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,checkpoints"), + ("tags", "models,model_type:checkpoints"), ("tags", json.dumps(["unit-tests", "alpha"])), ("name", "merge.safetensors"), ("user_metadata", json.dumps({"u": 1})), @@ -124,7 +226,71 @@ 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", "checkpoints", "unit-tests", "alpha"}.issubset(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 @pytest.mark.parametrize("root", ["input", "output"]) @@ -192,16 +358,55 @@ 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", "checkpoints", "unit-tests", "edge"]), "name": "empty.safetensors", "user_metadata": json.dumps({})} + form = {"tags": json.dumps(["models", "model_type: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_invalid_root_tag_rejected(http: requests.Session, api_base: str): +def test_upload_rejects_arbitrary_labels_without_required_destination_role(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) @@ -212,7 +417,7 @@ def test_upload_invalid_root_tag_rejected(http: requests.Session, api_base: str) 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", "checkpoints", "unit-tests", "edge"]), "name": "badmeta.bin", "user_metadata": "{not json}"} + form = {"tags": json.dumps(["models", "model_type: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 @@ -228,7 +433,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", "checkpoints", "unit-tests"]))), + ("tags", (None, json.dumps(["models", "model_type:checkpoints", "unit-tests"]))), ("name", (None, "x.safetensors")), ] r = http.post(api_base + "/api/assets", files=files, timeout=120) @@ -237,17 +442,33 @@ def test_upload_missing_file_and_hash(http: requests.Session, api_base: str): assert body["error"]["code"] == "MISSING_FILE" -def test_upload_models_unknown_category(http: requests.Session, api_base: str): +def test_upload_models_unknown_model_type(http: requests.Session, api_base: str): files = {"file": ("m.safetensors", b"A" * 128, "application/octet-stream")} - form = {"tags": json.dumps(["models", "no_such_category", "unit-tests"]), "name": "m.safetensors"} + form = {"tags": json.dumps(["models", "model_type: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 + assert r.status_code == 400, body assert body["error"]["code"] == "INVALID_BODY" - assert body["error"]["message"].startswith("unknown models category") -def test_upload_models_requires_category(http: requests.Session, api_base: str): +@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): 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) @@ -256,13 +477,152 @@ def test_upload_models_requires_category(http: requests.Session, api_base: str): assert body["error"]["code"] == "INVALID_BODY" -def test_upload_tags_traversal_guard(http: requests.Session, api_base: str): +def test_upload_extra_tags_are_labels_not_path_components(http: requests.Session, api_base: str): files = {"file": ("evil.safetensors", b"A" * 256, "application/octet-stream")} - form = {"tags": json.dumps(["models", "checkpoints", "unit-tests", "..", "zzz"]), "name": "evil.safetensors"} + form = {"tags": json.dumps(["models", "model_type: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 == 400 - assert body["error"]["code"] in ("BAD_REQUEST", "INVALID_BODY") + 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() def test_upload_empty_tags_rejected(http: requests.Session, api_base: str): diff --git a/tests-unit/comfy_test/folder_path_test.py b/tests-unit/comfy_test/folder_path_test.py index 775e15c36..a0ef17a4c 100644 --- a/tests-unit/comfy_test/folder_path_test.py +++ b/tests-unit/comfy_test/folder_path_test.py @@ -53,8 +53,11 @@ def test_annotated_filepath(): def test_get_annotated_filepath(): default_dir = "/default/dir" - assert folder_paths.get_annotated_filepath("test.txt", default_dir) == os.path.join(default_dir, "test.txt") - assert folder_paths.get_annotated_filepath("test.txt [output]") == os.path.join(folder_paths.get_output_directory(), "test.txt") + # get_annotated_filepath now normalizes with os.path.abspath (part of the + # GHSA-779p traversal hardening), so compare against the normalized form — + # on Windows abspath also prepends the current drive letter. + assert folder_paths.get_annotated_filepath("test.txt", default_dir) == os.path.abspath(os.path.join(default_dir, "test.txt")) + assert folder_paths.get_annotated_filepath("test.txt [output]") == os.path.abspath(os.path.join(folder_paths.get_output_directory(), "test.txt")) def test_add_model_folder_path_append(clear_folder_paths): folder_paths.add_model_folder_path("test_folder", "/default/path", is_default=True) @@ -160,3 +163,20 @@ def test_base_path_change_clears_old(set_base_dir): for name in ["controlnet", "diffusion_models", "text_encoders"]: assert len(folder_paths.get_folder_paths(name)) == 2 + + +def test_models_directory_cli_and_getters(temp_dir): + try: + with patch.object(sys, 'argv', ["main.py", "--models-directory", temp_dir]): + reload(comfy.cli_args) + reload(folder_paths) + + assert folder_paths.models_dir == os.path.abspath(temp_dir) + + with pytest.raises(Exception): + comfy.cli_args.is_valid_directory(os.path.join(temp_dir, "non_existent_folder_path")) + finally: + with patch.object(sys, 'argv', ["main.py"]): + reload(comfy.cli_args) + reload(folder_paths) + diff --git a/tests-unit/feature_flags_test.py b/tests-unit/feature_flags_test.py index 8ec52a124..a436ab1ec 100644 --- a/tests-unit/feature_flags_test.py +++ b/tests-unit/feature_flags_test.py @@ -29,6 +29,8 @@ 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)) diff --git a/tests-unit/security_test/__init__.py b/tests-unit/security_test/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/tests-unit/security_test/test_ghsa_779p_02_preview_traversal.py b/tests-unit/security_test/test_ghsa_779p_02_preview_traversal.py new file mode 100644 index 000000000..f17fd26ea --- /dev/null +++ b/tests-unit/security_test/test_ghsa_779p_02_preview_traversal.py @@ -0,0 +1,192 @@ +"""CI unit tests for FIX #2 of GHSA-779p-m5rp-r4h4. + +Path traversal / hardening in app/model_manager.py get_model_preview +(route /experiment/models/preview/{folder}/{path_index}/{filename:.*}). + +Reference: https://github.com/Comfy-Org/ComfyUI/security/advisories/GHSA-779p-m5rp-r4h4 +""" +import pytest +import yarl +from io import BytesIO +from PIL import Image +from aiohttp import web +from unittest.mock import patch +from app.model_manager import ModelFileManager + +pytestmark = ( + pytest.mark.asyncio +) # This applies the asyncio mark to all test functions in the module + +@pytest.fixture +def model_manager(): + return ModelFileManager() + +@pytest.fixture +def app(model_manager): + app = web.Application() + routes = web.RouteTableDef() + model_manager.add_routes(routes) + app.add_routes(routes) + return app + + +async def test_legit_preview_returns_200(aiohttp_client, app, tmp_path): + """Sanity: a real preview PNG inside the model folder is served as webp 200.""" + img = Image.new('RGB', (16, 16), color=(255, 0, 128)) + img.save(tmp_path / "test_model.png", format='PNG') + + with patch('folder_paths.folder_names_and_paths', { + 'test_folder': ([str(tmp_path)], None) + }): + client = await aiohttp_client(app) + response = await client.get('/experiment/models/preview/test_folder/0/test_model.png') + + assert response.status == 200 + assert response.content_type == 'image/webp' + + img_bytes = BytesIO(await response.read()) + served = Image.open(img_bytes) + assert served.format + assert served.format.lower() == 'webp' + served.close() + + +async def test_non_integer_path_index_returns_400(aiohttp_client, app, tmp_path): + """A non-integer path_index segment must be rejected with 400.""" + with patch('folder_paths.folder_names_and_paths', { + 'test_folder': ([str(tmp_path)], None) + }): + client = await aiohttp_client(app) + response = await client.get('/experiment/models/preview/test_folder/abc/test_model.png') + + assert response.status == 400 + + +async def test_out_of_range_path_index_returns_404(aiohttp_client, app, tmp_path): + """A path_index beyond the configured folder list must return 404.""" + with patch('folder_paths.folder_names_and_paths', { + 'test_folder': ([str(tmp_path)], None) + }): + client = await aiohttp_client(app) + response = await client.get('/experiment/models/preview/test_folder/99/test_model.png') + + assert response.status == 404 + + +async def test_empty_filename_returns_400(aiohttp_client, app, tmp_path): + """The "{filename:.*}" capture also matches the empty string (trailing + slash). It would resolve to the folder itself and must be rejected with 400.""" + with patch('folder_paths.folder_names_and_paths', { + 'test_folder': ([str(tmp_path)], None) + }): + client = await aiohttp_client(app) + response = await client.get('/experiment/models/preview/test_folder/0/') + + assert response.status == 400 + + +async def test_path_traversal_in_filename_returns_403(aiohttp_client, app, tmp_path): + """Path traversal in {filename} must be rejected with 403 and must NOT read + a file outside the configured model directory. + + GOTCHA: aiohttp/yarl collapses literal ``../`` dot-segments out of the URL + path before it reaches the handler, which would make this test vacuously + pass (the request would hit a different/non-existent route). We percent-encode + the dots and slashes (``%2e%2e%2f``) and send the URL with + ``yarl.URL(..., encoded=True)`` so the bytes survive client-side normalization + untouched; aiohttp's router then percent-decodes them into ``match_info``, + delivering the literal ``../`` traversal to the handler's ``{filename:.*}`` + capture. + + Without the fix the handler computes + ``os.path.normpath(os.path.join(folder, "../../../../etc/hosts"))``, which + escapes ``tmp_path`` and would be passed straight to get_model_previews -> + Image.open, serving bytes from outside the model dir (200/served bytes). The + is_within_directory() containment check is the load-bearing fix that turns + that escape into a 403. + """ + # Sanity-anchor: a legit preview exists inside tmp_path, so a 200 path is + # genuinely reachable — proving the 403 below is the containment check + # firing, not an unrelated 404. + img = Image.new('RGB', (16, 16), color=(255, 0, 128)) + img.save(tmp_path / "test_model.png", format='PNG') + + # Percent-encoded "../../../../etc/hosts" so yarl does not collapse the + # dot-segments before the request leaves the client. + encoded_traversal = '%2e%2e%2f' * 4 + 'etc%2fhosts' + raw_path = '/experiment/models/preview/test_folder/0/' + encoded_traversal + url = yarl.URL(raw_path, encoded=True) + + with patch('folder_paths.folder_names_and_paths', { + 'test_folder': ([str(tmp_path)], None) + }): + client = await aiohttp_client(app) + response = await client.get(url) + + # Confirm the traversal actually reached the handler intact: a 200 here + # would mean either normalization stripped the ``../`` (vacuous pass) or + # the containment check failed open and served outside-dir bytes. + assert response.status == 403, ( + f"expected 403 from is_within_directory() containment check, " + f"got {response.status}; traversal may have been normalized away " + f"or the fix failed open" + ) + body = await response.read() + assert body == b"", "403 response must not carry any file bytes" + + +async def test_symlink_companion_preview_returns_403(aiohttp_client, app, tmp_path): + """A companion preview file is selected by a glob inside get_model_previews + and then opened. If that companion is a symlink whose path is in-dir but + whose target escapes the model folder, it must be rejected with 403 — not + served. The requested path itself stays in-dir (so the first containment + check passes); the load-bearing fix is the SECOND is_within_directory check + on the file actually opened. + """ + model_dir = tmp_path / "models" + model_dir.mkdir() + secret_dir = tmp_path / "secret" + secret_dir.mkdir() + # A real image OUTSIDE the model dir — valid, so without the fix Image.open + # would succeed and its bytes would be served (200). + secret = secret_dir / "secret.png" + Image.new('RGB', (8, 8), color=(0, 0, 0)).save(secret, format='PNG') + # Companion preview, in-dir by name but a symlink escaping the model dir. + # (No real model file is needed — get_model_previews globs companions by + # basename, and omitting a .safetensors avoids the metadata-header read.) + companion = model_dir / "model.preview.png" + try: + companion.symlink_to(secret) + except (OSError, NotImplementedError): + pytest.skip("symlinks not supported on this platform/filesystem") + + with patch('folder_paths.folder_names_and_paths', { + 'test_folder': ([str(model_dir)], None) + }): + client = await aiohttp_client(app) + response = await client.get('/experiment/models/preview/test_folder/0/model.safetensors') + + assert response.status == 403, ( + f"expected 403 — the globbed companion preview is a symlink resolving " + f"outside the model dir and must not be served; got {response.status}" + ) + assert await response.read() == b"" + + +async def test_null_byte_in_filename_no_500(aiohttp_client, app, tmp_path): + """A NUL byte in the filename must yield a clean client rejection, not a 500 + from an uncaught ValueError in is_within_directory's realpath() call.""" + raw_path = '/experiment/models/preview/test_folder/0/' + 'a%00b' + url = yarl.URL(raw_path, encoded=True) + + with patch('folder_paths.folder_names_and_paths', { + 'test_folder': ([str(tmp_path)], None) + }): + client = await aiohttp_client(app) + response = await client.get(url) + + assert response.status != 500, ( + f"NUL byte produced a 500 (uncaught ValueError); expected a clean " + f"4xx rejection, got {response.status}" + ) + assert 400 <= response.status < 500 diff --git a/tests-unit/security_test/test_ghsa_779p_03_annotated_traversal.py b/tests-unit/security_test/test_ghsa_779p_03_annotated_traversal.py new file mode 100644 index 000000000..88102760c --- /dev/null +++ b/tests-unit/security_test/test_ghsa_779p_03_annotated_traversal.py @@ -0,0 +1,165 @@ +"""Security tests for GHSA-779p-m5rp-r4h4 — FIX #3. + +Path traversal in folder_paths.get_annotated_filepath / exists_annotated_filepath, +plus the shared is_within_directory() containment helper. + +These are pure-function tests (no running server). The input/output/temp +directories are pointed at tmp_path via the folder_paths setters, so a crafted +name containing `../`, an absolute path, or a symlink that escapes the base +directory must be rejected. + +Reference: https://github.com/Comfy-Org/ComfyUI/security/advisories/GHSA-779p-m5rp-r4h4 +""" +import os + +import pytest + +import folder_paths +from comfy.options import enable_args_parsing +enable_args_parsing() + + +@pytest.fixture +def sandbox(tmp_path): + """Point folder_paths' input/output/temp dirs at a real temp sandbox. + + Yields the realpath'd base, input, output and temp directories. The original + directory values are restored afterward so tests stay isolated. + """ + base = os.path.realpath(str(tmp_path)) + input_dir = os.path.join(base, "input") + output_dir = os.path.join(base, "output") + temp_dir = os.path.join(base, "temp") + for d in (input_dir, output_dir, temp_dir): + os.makedirs(d, exist_ok=True) + + orig_input = folder_paths.get_input_directory() + orig_output = folder_paths.get_output_directory() + orig_temp = folder_paths.get_temp_directory() + + folder_paths.set_input_directory(input_dir) + folder_paths.set_output_directory(output_dir) + folder_paths.set_temp_directory(temp_dir) + + yield { + "base": base, + "input": input_dir, + "output": output_dir, + "temp": temp_dir, + } + + folder_paths.set_input_directory(orig_input) + folder_paths.set_output_directory(orig_output) + folder_paths.set_temp_directory(orig_temp) + + +# --------------------------------------------------------------------------- +# is_within_directory() — the shared containment helper +# --------------------------------------------------------------------------- + +def test_is_within_directory_legit_child(sandbox): + base = sandbox["input"] + child = os.path.join(base, "sub", "image.png") + assert folder_paths.is_within_directory(base, child) is True + + +def test_is_within_directory_dotdot_escape(sandbox): + base = sandbox["input"] + escape = os.path.join(base, "..", "..", "etc", "passwd") + assert folder_paths.is_within_directory(base, escape) is False + + +def test_is_within_directory_symlink_escape(sandbox): + """A symlink created INSIDE base that points OUTSIDE base must not pass. + + This is the key new hardening: is_within_directory realpath()s both operands, + so a symlink planted in the base directory can't be used to read files + elsewhere. We create a real on-disk symlink and a real secret target to + verify the check actually resolves the link. + """ + base = sandbox["input"] + + # A directory living outside the base, holding a secret file. + outside = os.path.join(sandbox["base"], "outside_secret_dir") + os.makedirs(outside, exist_ok=True) + secret = os.path.join(outside, "secret.txt") + with open(secret, "w") as f: + f.write("top secret") + + # Plant a symlink inside base that points at the outside directory. + # symlink creation can require elevated privileges / Developer Mode on + # Windows, so skip cleanly where it isn't available (same guard as the + # sibling test in test_ghsa_779p_02_preview_traversal.py). + link = os.path.join(base, "escape_link") + try: + os.symlink(outside, link) + except (OSError, NotImplementedError): + pytest.skip("symlinks not supported on this platform/filesystem") + + # Accessing the secret "through" the in-base symlink must be rejected. + target_via_link = os.path.join(link, "secret.txt") + assert folder_paths.is_within_directory(base, target_via_link) is False + + +# --------------------------------------------------------------------------- +# get_annotated_filepath() +# --------------------------------------------------------------------------- + +def test_get_annotated_filepath_legit_name(sandbox): + result = folder_paths.get_annotated_filepath("image.png") + assert result == os.path.join(sandbox["input"], "image.png") + assert folder_paths.is_within_directory(sandbox["input"], result) + + +def test_get_annotated_filepath_input_annotation(sandbox): + result = folder_paths.get_annotated_filepath("image.png [input]") + assert result == os.path.join(sandbox["input"], "image.png") + + +def test_get_annotated_filepath_output_annotation(sandbox): + result = folder_paths.get_annotated_filepath("image.png [output]") + assert result == os.path.join(sandbox["output"], "image.png") + + +def test_get_annotated_filepath_temp_annotation(sandbox): + result = folder_paths.get_annotated_filepath("image.png [temp]") + assert result == os.path.join(sandbox["temp"], "image.png") + + +def test_get_annotated_filepath_dotdot_raises(sandbox): + with pytest.raises(ValueError): + folder_paths.get_annotated_filepath("../etc/passwd") + + +def test_get_annotated_filepath_dotdot_with_annotation_raises(sandbox): + with pytest.raises(ValueError): + folder_paths.get_annotated_filepath("../../etc/passwd [output]") + + +def test_get_annotated_filepath_absolute_escape_raises(sandbox): + with pytest.raises(ValueError): + folder_paths.get_annotated_filepath("/etc/passwd") + + +# --------------------------------------------------------------------------- +# exists_annotated_filepath() +# --------------------------------------------------------------------------- + +def test_exists_annotated_filepath_existing_legit_file(sandbox): + real = os.path.join(sandbox["input"], "real.png") + with open(real, "w") as f: + f.write("data") + assert folder_paths.exists_annotated_filepath("real.png") is True + + +def test_exists_annotated_filepath_traversal_returns_false(sandbox): + """A traversal name must return False without raising and without probing + outside the base directory (must never reach os.path.exists for the escape). + """ + # /etc/passwd exists on POSIX; the function must still report False because + # the resolved path escapes the input directory. + assert folder_paths.exists_annotated_filepath("../../../../../../etc/passwd") is False + + +def test_exists_annotated_filepath_absolute_returns_false(sandbox): + assert folder_paths.exists_annotated_filepath("/etc/passwd") is False diff --git a/tests-unit/security_test/test_ghsa_779p_04_userdata_xss.py b/tests-unit/security_test/test_ghsa_779p_04_userdata_xss.py new file mode 100644 index 000000000..aa1250327 --- /dev/null +++ b/tests-unit/security_test/test_ghsa_779p_04_userdata_xss.py @@ -0,0 +1,147 @@ +""" +CI unit tests for FIX #4 of GHSA-779p-m5rp-r4h4. + +Stored-XSS hardening on GET /userdata/{file} in app/user_manager.py. + +User data files are arbitrary user-supplied content and must never render +inline in the app origin. The getuserdata handler: + - forces Content-Type to application/octet-stream for any type in + folder_paths.DANGEROUS_CONTENT_TYPES (text/html, image/svg+xml, + text/javascript, ...), + - sets X-Content-Type-Options: nosniff, + - sets Content-Disposition: attachment. + +These tests pre-create files in tmp_path and GET them back, asserting the +secure response headers. They mirror the aiohttp_client pattern in +tests-unit/prompt_server_test/user_manager_test.py. +""" + +import pytest +import os +from aiohttp import web +from app.user_manager import UserManager + +pytestmark = ( + pytest.mark.asyncio +) # This applies the asyncio mark to all test functions in the module + + +@pytest.fixture +def user_manager(tmp_path): + um = UserManager() + um.get_request_user_filepath = lambda req, file, **kwargs: os.path.join( + tmp_path, file + ) if file else tmp_path + return um + + +@pytest.fixture +def app(user_manager): + app = web.Application() + routes = web.RouteTableDef() + user_manager.add_routes(routes) + app.add_routes(routes) + return app + + +async def test_html_served_as_octet_stream(aiohttp_client, app, tmp_path): + (tmp_path / "evil.html").write_text( + "" + ) + + client = await aiohttp_client(app) + resp = await client.get("/userdata/evil.html") + + assert resp.status == 200 + ct = resp.headers.get("Content-Type", "") + # The load-bearing assertion: a .html file must NOT be served as text/html. + assert "text/html" not in ct.lower(), ( + f"Content-Type {ct!r} would let a browser render/execute the file (stored XSS)." + ) + assert ct == "application/octet-stream" + assert resp.headers.get("X-Content-Type-Options") == "nosniff" + assert "attachment" in resp.headers.get("Content-Disposition", "") + + +async def test_svg_served_as_octet_stream(aiohttp_client, app, tmp_path): + (tmp_path / "evil.svg").write_text( + '' + '' + '' + "" + ) + + client = await aiohttp_client(app) + resp = await client.get("/userdata/evil.svg") + + assert resp.status == 200 + ct = resp.headers.get("Content-Type", "") + # SVG can carry inline