mirror of
https://github.com/comfyanonymous/ComfyUI.git
synced 2026-07-21 23:41:28 +08:00
Compare commits
23
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
255e1b5019 | ||
|
|
7f287b705e | ||
|
|
b7ba504e06 | ||
|
|
6c62ca0b6b | ||
|
|
3fe9f5fecb | ||
|
|
1073a74976 | ||
|
|
de1b8f3e8d | ||
|
|
77917ed3a6 | ||
|
|
a04ebe05c2 | ||
|
|
9764381998 | ||
|
|
1e04ced089 | ||
|
|
96e0e3585b | ||
|
|
35c1470935 | ||
|
|
694815f498 | ||
|
|
92594ca84c | ||
|
|
2c935de1b1 | ||
|
|
dd17debce5 | ||
|
|
50e5270b86 | ||
|
|
c8b3d54cf8 | ||
|
|
eeb58c0a1d | ||
|
|
981551d073 | ||
|
|
351119eb05 | ||
|
|
418d272cfa |
+16
-3
@@ -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"
|
||||
|
||||
@@ -0,0 +1,294 @@
|
||||
## Engineering Style
|
||||
|
||||
- Keep changes small and direct. Most fixes should touch the narrowest code path
|
||||
that explains the bug, performance issue, dtype issue, model-format issue, or
|
||||
user-facing behavior.
|
||||
- Change the least amount of files possible. A change that touches many files is
|
||||
more likely to be a bad change than a good one unless the broader scope is
|
||||
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. 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
|
||||
workflow compatibility unless the change is explicitly about replacing them.
|
||||
- Code must look hand-written for this repository. Changes that read like
|
||||
generic AI-generated code will be rejected automatically: unnecessary helper
|
||||
layers, vague names, boilerplate comments, defensive branches without a real
|
||||
failure mode, broad rewrites, or code that ignores the local style.
|
||||
|
||||
## Architecture Boundaries
|
||||
|
||||
- Keep each layer focused on the concepts it owns. Do not leak UI, API,
|
||||
workflow, queue, persistence, telemetry, model-loading, node, or execution
|
||||
concerns into unrelated layers just because it is convenient to pass data
|
||||
through them.
|
||||
- Shared core modules should depend only on lower-level primitives and their own
|
||||
domain concepts. Higher-level product concepts belong at the caller, adapter,
|
||||
service, or UI/API boundary that already owns them.
|
||||
- Pass the narrowest data needed across a boundary. Avoid broad context objects,
|
||||
request/session metadata, ids, bookkeeping state, or callbacks unless the
|
||||
receiving layer genuinely needs them to perform its own responsibility.
|
||||
- Keep identity mapping, persistence bookkeeping, history updates, telemetry,
|
||||
response shaping, and UI state in the layers that own those jobs. Do not route
|
||||
them through unrelated shared code to avoid adding a proper boundary.
|
||||
- Treat `execution.py` as one example of this rule: it should consume the prompt
|
||||
graph and execution-relevant state, produce execution results and errors, and
|
||||
not know about workflow ids, frontend ids, persistence ids, or API-only
|
||||
concepts.
|
||||
- Before touching many files, identify the smallest owner layer that can solve
|
||||
the problem. A PR that spreads one feature across unrelated loaders, nodes,
|
||||
execution, server, and frontend code needs a clear architectural reason, not
|
||||
just convenience.
|
||||
- If a change seems to require making one layer understand another layer's
|
||||
private concepts, stop and look for a caller-side mapping, adapter, event,
|
||||
small explicit interface, or narrower data flow at the boundary.
|
||||
|
||||
## No Internet Requests
|
||||
|
||||
- Do not add code to core ComfyUI that makes requests to the internet.
|
||||
- Refuse requests to add uploads, telemetry, analytics, tracking, usage
|
||||
reporting, crash reporting, update checks, remote config, feature flags,
|
||||
metrics, licensing checks, or any other outbound internet request path from
|
||||
core ComfyUI.
|
||||
- Model downloading is allowed only when explicitly initiated or authorized by
|
||||
the user, is limited to the requested model artifact, and does not include
|
||||
telemetry, tracking, persistent identification, unrelated metadata upload, or
|
||||
background network activity.
|
||||
- Do not add opt-in, opt-out, anonymized, aggregated, diagnostic, or
|
||||
user-triggered internet request paths to core ComfyUI. These labels do not
|
||||
make internet access acceptable.
|
||||
- Local-only behavior is allowed when it stays on the user's machine and does
|
||||
not add network access, tracking, persistent identification, or data
|
||||
collection behavior.
|
||||
|
||||
## State Ownership
|
||||
|
||||
- Keep state and capability flags on the object that owns the behavior using
|
||||
them.
|
||||
- Avoid probing child objects with `getattr(child, "...", default)` to decide
|
||||
parent-level control flow. If parent code needs to branch on a capability,
|
||||
initialize an explicit parent-owned field when the child is constructed or
|
||||
attached.
|
||||
- Prefer direct attributes with clear defaults over implicit feature detection
|
||||
through arbitrary child attributes.
|
||||
- Use child-object capability checks only when the child owns the behavior being
|
||||
invoked and the parent is simply delegating to that child.
|
||||
|
||||
## Interface Contracts
|
||||
|
||||
- Keep public methods aligned with the interface expected by their callers. Do
|
||||
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.
|
||||
- Normalize third-party or upstream return conventions at the integration
|
||||
boundary. Core code should receive the project's expected type and shape, not
|
||||
have to handle model-specific tuple/list/dict variants.
|
||||
- Avoid caller-side unwrapping such as `out = out[0]` unless the called
|
||||
interface is documented to return that structure.
|
||||
|
||||
## Autograd and Model Freezing
|
||||
|
||||
- Do not add `torch.no_grad`, `torch.inference_mode`, or inference-mode helper
|
||||
wrappers in ComfyUI code. The only allowed inference-mode-related use is
|
||||
disabling a globally set inference mode when a training path needs gradients.
|
||||
- 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
|
||||
|
||||
- Keep imports at module scope. Avoid inline imports unless they are already part
|
||||
of an established optional-backend probe or are needed to avoid an import
|
||||
cycle.
|
||||
- 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.
|
||||
- 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
|
||||
long lines, simple helper functions, module-level state, and direct tensor
|
||||
operations when they make the code easier to follow.
|
||||
- Keep comments sparse and useful. Strip useless comments that restate the code
|
||||
or describe obvious behavior. Short TODOs are fine when they name the concrete
|
||||
missing follow-up.
|
||||
|
||||
## Model, Device, and Memory Behavior
|
||||
|
||||
- Treat dtype, device placement, VRAM usage, and offloading behavior as core
|
||||
correctness concerns. Check CPU, CUDA, ROCm, MPS, DirectML, XPU, NPU, and low
|
||||
VRAM implications when touching shared execution or loading code.
|
||||
- Prefer native ComfyUI formats and existing quantization/offload helpers over
|
||||
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.
|
||||
|
||||
## Nodes and User-Facing Behavior
|
||||
|
||||
- Follow existing node conventions: `INPUT_TYPES`, `RETURN_TYPES`, `FUNCTION`,
|
||||
`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
|
||||
comments, but do not disrespect her.
|
||||
- Warning and info messages should be short and actionable. Remove noisy or
|
||||
misleading messages rather than adding more logging.
|
||||
- Documentation and README edits should be concise, factual, and tied to the
|
||||
changed behavior.
|
||||
|
||||
## Commit and Review Habits
|
||||
|
||||
- If asked to write commit messages, use short direct subjects like the existing
|
||||
history: `Fix ...`, `Add ...`, `Support ...`, `Remove ...`, `Update ...`,
|
||||
`Make ...`, `Use ...`, `Disable ...`, `Bump ...`, or `Revert ...`.
|
||||
- Keep PR descriptions short and reviewable. State the problem, the behavioral
|
||||
change, and the tests run; avoid long narrative explanations, implementation
|
||||
diaries, or exhaustive file-by-file summaries unless the reviewer explicitly
|
||||
needs that context.
|
||||
- Prefer one coherent behavioral change per commit. Dependency pins, tests, and
|
||||
the code that needs them may be in the same commit when they are inseparable.
|
||||
- In reviews, prioritize real user impact: crashes, wrong dtype/device behavior,
|
||||
memory regressions, broken model loading, workflow incompatibility, and noisy
|
||||
or misleading user-facing output.
|
||||
@@ -306,12 +306,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)
|
||||
|
||||
@@ -0,0 +1,51 @@
|
||||
"""URL allowlist for server-side model fetches.
|
||||
|
||||
Mirrors the frontend's ``isModelDownloadable`` allowlist so the two flows
|
||||
agree on which URLs are eligible for download. Server-side allowlisting is
|
||||
the primary SSRF defense for this subsystem — workflow JSON is untrusted
|
||||
input (anyone can hand-craft one), so we never let the server fetch URLs
|
||||
outside this list.
|
||||
"""
|
||||
|
||||
from urllib.parse import urlparse
|
||||
|
||||
# Frontend parity: ``missingModelDownload-*.js`` exports the same triple
|
||||
# (Civitai / HuggingFace / localhost). Keyed by exact hostname → allowed
|
||||
# schemes, and matched against the *parsed* host (not a raw string prefix),
|
||||
# so URL-userinfo tricks can't slip past — see ``is_url_allowed``.
|
||||
_ALLOWED_HOSTS = {
|
||||
"huggingface.co": {"https"},
|
||||
"civitai.com": {"https"},
|
||||
"localhost": {"http"},
|
||||
"127.0.0.1": {"http"},
|
||||
}
|
||||
|
||||
# Frontend parity: same set as ``a = [...]`` in the bundle.
|
||||
_ALLOWED_MODEL_EXTENSIONS = (
|
||||
".safetensors",
|
||||
".sft",
|
||||
".ckpt",
|
||||
".pth",
|
||||
".pt",
|
||||
)
|
||||
|
||||
|
||||
def is_url_allowed(url: str) -> bool:
|
||||
"""Check whether ``url`` is permitted as a server-side download source.
|
||||
|
||||
True only when the parsed host + scheme are allowlisted AND the path ends
|
||||
in a model extension. Matching on ``parsed.hostname`` (not a string prefix)
|
||||
defeats userinfo tricks like ``http://127.0.0.1:80@169.254.169.254/x.safetensors``,
|
||||
whose real host is ``169.254.169.254``; the extension check rejects non-model
|
||||
URLs on allowed hosts (e.g. ``huggingface.co/api/...``).
|
||||
"""
|
||||
if not isinstance(url, str) or not url:
|
||||
return False
|
||||
try:
|
||||
parsed = urlparse(url)
|
||||
except ValueError:
|
||||
return False
|
||||
host = parsed.hostname
|
||||
if host is None or parsed.scheme not in _ALLOWED_HOSTS.get(host, ()):
|
||||
return False
|
||||
return any(parsed.path.endswith(ext) for ext in _ALLOWED_MODEL_EXTENSIONS)
|
||||
@@ -0,0 +1,359 @@
|
||||
"""Aiohttp routes for the server-side model download subsystem.
|
||||
|
||||
Endpoint surface (all under ``/api/``, all kebab-case):
|
||||
|
||||
- ``POST /api/models-availability-status`` — bulk status + metadata query.
|
||||
- ``POST /api/download-models`` — start a batch of downloads.
|
||||
- ``POST /api/cancel-model-download-session`` — cancel a single in-flight one.
|
||||
- ``GET /api/hf-auth-token-status`` — current HF login state.
|
||||
- ``POST /api/hf-auth-login-start`` — begin the HF OAuth flow.
|
||||
- ``POST /api/hf-auth-logout`` — drop the stored HF token.
|
||||
|
||||
The contract is intentionally narrow: only model_ids of the form
|
||||
``<directory>/<filename>`` (validated via ``app.model_downloader.paths``)
|
||||
are accepted, and only URLs on the same allowlist the frontend already
|
||||
uses (HuggingFace, Civitai, localhost) can be fetched. Both are required
|
||||
to keep the server out of the SSRF business for this feature.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
from typing import Any, Literal, Optional
|
||||
|
||||
from aiohttp import web
|
||||
from pydantic import BaseModel, ValidationError
|
||||
|
||||
from app.model_downloader.allowlist import is_url_allowed
|
||||
from app.model_downloader.download_server import (
|
||||
DOWNLOAD_SERVER,
|
||||
DownloadSession,
|
||||
)
|
||||
from app.model_downloader.downloader import schedule_batch
|
||||
from app.model_downloader.gated_detection import probe_url
|
||||
from app.model_downloader.hf_auth.auth_store import HF_AUTH_STORE
|
||||
from app.model_downloader.hf_auth.eligibility import is_hf_auth_eligible
|
||||
from app.model_downloader.hf_auth.oauth import (
|
||||
OAuthCallbackError,
|
||||
OAuthInProgressError,
|
||||
start_login_flow,
|
||||
)
|
||||
from app.model_downloader.paths import (
|
||||
InvalidModelId,
|
||||
parse_model_id,
|
||||
resolve_existing,
|
||||
)
|
||||
from app.model_downloader.api import schemas_in, schemas_out
|
||||
|
||||
ROUTES = web.RouteTableDef()
|
||||
|
||||
|
||||
def register_routes(app: web.Application) -> None:
|
||||
"""Wire the model-downloader routes into the running aiohttp app.
|
||||
|
||||
Called once from ``server.py`` during ``PromptServer`` startup.
|
||||
"""
|
||||
app.add_routes(ROUTES)
|
||||
|
||||
|
||||
# ----- response helpers (same envelope as app/assets/api/routes.py) -----
|
||||
|
||||
|
||||
ErrorCode = Literal[
|
||||
"INVALID_JSON",
|
||||
"INVALID_BODY",
|
||||
"EMPTY_REQUEST",
|
||||
"INVALID_MODEL_ID",
|
||||
"URL_NOT_ALLOWED",
|
||||
"ALREADY_AVAILABLE",
|
||||
"ALREADY_DOWNLOADING",
|
||||
"MODEL_NOT_DOWNLOADABLE",
|
||||
"NOT_DOWNLOADING",
|
||||
"HF_AUTH_NOT_ELIGIBLE",
|
||||
"HF_AUTH_IN_PROGRESS",
|
||||
"HF_AUTH_START_FAILED",
|
||||
]
|
||||
|
||||
|
||||
def _error(status: int, code: ErrorCode, message: str, details: dict | None = None) -> web.Response:
|
||||
return web.json_response(
|
||||
{"error": {"code": code, "message": message, "details": details or {}}},
|
||||
status=status,
|
||||
)
|
||||
|
||||
|
||||
def _validation_error(code: ErrorCode, ve: ValidationError) -> web.Response:
|
||||
return _error(400, code, "Validation failed.", {"errors": json.loads(ve.json())})
|
||||
|
||||
|
||||
def _ok(payload: BaseModel, status: int = 200) -> web.Response:
|
||||
return web.json_response(
|
||||
payload.model_dump(mode="json", exclude_none=False),
|
||||
status=status,
|
||||
)
|
||||
|
||||
|
||||
async def _parse_body(request: web.Request, model: type[BaseModel]) -> Any:
|
||||
"""Parse a JSON body into a pydantic model or raise a 400 response."""
|
||||
try:
|
||||
raw = await request.json()
|
||||
except json.JSONDecodeError:
|
||||
return _error(400, "INVALID_JSON", "Request body must be valid JSON.")
|
||||
try:
|
||||
return model.model_validate(raw)
|
||||
except ValidationError as ve:
|
||||
return _validation_error("INVALID_BODY", ve)
|
||||
|
||||
|
||||
# ----- 1. availability status (unified: state + metadata per id) -----
|
||||
|
||||
|
||||
@ROUTES.post("/api/models-availability-status")
|
||||
async def models_availability_status(request: web.Request) -> web.Response:
|
||||
"""Return per-id ``{state, progress, file_size, is_hf_downloadable}``.
|
||||
|
||||
State (``available`` / ``missing`` / ``downloading``) is cheap to
|
||||
recompute per call. ``file_size`` and ``is_gated`` are cached
|
||||
server-side per URL. ``is_hf_downloadable`` is recomputed every
|
||||
call from the current token state — that's what makes login + license
|
||||
acceptance show up in the UI within one poll cycle without any
|
||||
frontend cache plumbing.
|
||||
"""
|
||||
parsed = await _parse_body(request, schemas_in.AvailabilityStatusRequest)
|
||||
if isinstance(parsed, web.Response):
|
||||
return parsed
|
||||
|
||||
items = list(parsed.models.items())
|
||||
|
||||
# Run all probes concurrently; each is internally cached per URL.
|
||||
probes = await asyncio.gather(*(probe_url(url) for _, url in items))
|
||||
|
||||
response_models: dict[str, schemas_out.ModelStatusEntry] = {}
|
||||
for (model_id, _url), probe in zip(items, probes):
|
||||
try:
|
||||
parse_model_id(model_id)
|
||||
except InvalidModelId:
|
||||
# Ill-formed identifier: report as missing without 400-ing the
|
||||
# whole batch — the workflow author probably typo'd.
|
||||
response_models[model_id] = schemas_out.ModelStatusEntry(
|
||||
state="missing",
|
||||
file_size=probe.file_size,
|
||||
is_hf_downloadable=probe.is_hf_downloadable,
|
||||
)
|
||||
continue
|
||||
|
||||
active = DOWNLOAD_SERVER.get(model_id)
|
||||
if active is not None:
|
||||
response_models[model_id] = schemas_out.ModelStatusEntry(
|
||||
state="downloading",
|
||||
progress=schemas_out.DownloadProgress(
|
||||
bytes_downloaded=active.bytes_downloaded,
|
||||
total_bytes=active.total_bytes,
|
||||
progress=active.progress,
|
||||
),
|
||||
file_size=probe.file_size,
|
||||
is_hf_downloadable=probe.is_hf_downloadable,
|
||||
)
|
||||
continue
|
||||
|
||||
state: schemas_out.ModelState = (
|
||||
"available" if resolve_existing(model_id) is not None else "missing"
|
||||
)
|
||||
response_models[model_id] = schemas_out.ModelStatusEntry(
|
||||
state=state,
|
||||
file_size=probe.file_size,
|
||||
is_hf_downloadable=probe.is_hf_downloadable,
|
||||
)
|
||||
|
||||
return _ok(schemas_out.AvailabilityStatusResponse(
|
||||
models=response_models,
|
||||
hf_auth=schemas_out.HfAuthStatus(
|
||||
token_available=HF_AUTH_STORE.has_token(),
|
||||
eligible=is_hf_auth_eligible(),
|
||||
),
|
||||
))
|
||||
|
||||
|
||||
# ----- 2. start downloads -----
|
||||
|
||||
|
||||
@ROUTES.post("/api/download-models")
|
||||
async def download_models(request: web.Request) -> web.Response:
|
||||
parsed = await _parse_body(request, schemas_in.DownloadModelsRequest)
|
||||
if isinstance(parsed, web.Response):
|
||||
return parsed
|
||||
|
||||
if not parsed.models:
|
||||
return _error(400, "EMPTY_REQUEST", "No models supplied.")
|
||||
|
||||
# ----- precondition pass: validate everything BEFORE registering anything -----
|
||||
# Atomic semantics: if any model fails any precondition (invalid id,
|
||||
# not allow-listed URL, already on disk, already downloading, or gated),
|
||||
# the entire request fails and no state is changed.
|
||||
requested = list(parsed.models.items())
|
||||
|
||||
for model_id, url in requested:
|
||||
try:
|
||||
parse_model_id(model_id)
|
||||
except InvalidModelId as e:
|
||||
return _error(400, "INVALID_MODEL_ID", str(e),
|
||||
{"model_id": model_id})
|
||||
|
||||
if not is_url_allowed(url):
|
||||
return _error(
|
||||
400, "URL_NOT_ALLOWED",
|
||||
"Server-side downloads only accept HuggingFace, Civitai, "
|
||||
"or localhost URLs ending in a known model extension.",
|
||||
{"model_id": model_id, "url": url},
|
||||
)
|
||||
|
||||
if resolve_existing(model_id) is not None:
|
||||
return _error(409, "ALREADY_AVAILABLE",
|
||||
f"Model already exists on disk: {model_id}",
|
||||
{"model_id": model_id})
|
||||
|
||||
if DOWNLOAD_SERVER.is_downloading(model_id):
|
||||
return _error(409, "ALREADY_DOWNLOADING",
|
||||
f"A download for {model_id} is already in progress.",
|
||||
{"model_id": model_id})
|
||||
|
||||
# Reachability check last — it's the only one that talks to the
|
||||
# network. Concurrent probes. For HF URLs ``is_hf_downloadable``
|
||||
# reflects current token access; for non-HF URLs it's None, and we
|
||||
# treat that as "no info, proceed".
|
||||
probes = await asyncio.gather(*(probe_url(url) for _, url in requested))
|
||||
for (model_id, url), probe in zip(requested, probes):
|
||||
if probe.is_hf_downloadable is False:
|
||||
return _error(
|
||||
400, "MODEL_NOT_DOWNLOADABLE",
|
||||
f"Model {model_id} is gated on HuggingFace and the current "
|
||||
f"server token (if any) does not grant access.",
|
||||
{"model_id": model_id, "url": url},
|
||||
)
|
||||
|
||||
# ----- registration pass: try_register is atomic per model_id -----
|
||||
# Defensive: another request might have raced past our pre-check
|
||||
# between the loop above and here. try_register handles that.
|
||||
sessions: list[DownloadSession] = []
|
||||
for model_id, url in requested:
|
||||
session = DOWNLOAD_SERVER.try_register(model_id, url)
|
||||
if session is None:
|
||||
# Race: someone else got in. Roll back what we registered.
|
||||
for s in sessions:
|
||||
DOWNLOAD_SERVER.cancel(s.model_id)
|
||||
return _error(409, "ALREADY_DOWNLOADING",
|
||||
f"A download for {model_id} is already in progress (race).",
|
||||
{"model_id": model_id})
|
||||
sessions.append(session)
|
||||
|
||||
DOWNLOAD_SERVER.sweep_orphan_tmp_files()
|
||||
schedule_batch(sessions)
|
||||
logging.info(
|
||||
"[model_downloader] scheduled %d downloads: %s",
|
||||
len(sessions), [s.model_id for s in sessions],
|
||||
)
|
||||
|
||||
return _ok(schemas_out.DownloadModelsResponse(
|
||||
accepted=True,
|
||||
scheduled=[s.model_id for s in sessions],
|
||||
), status=202)
|
||||
|
||||
|
||||
# ----- 3. cancel a session -----
|
||||
|
||||
|
||||
@ROUTES.post("/api/cancel-model-download-session")
|
||||
async def cancel_model_download_session(request: web.Request) -> web.Response:
|
||||
parsed = await _parse_body(request, schemas_in.CancelDownloadSessionRequest)
|
||||
if isinstance(parsed, web.Response):
|
||||
return parsed
|
||||
|
||||
try:
|
||||
parse_model_id(parsed.model_id)
|
||||
except InvalidModelId as e:
|
||||
return _error(400, "INVALID_MODEL_ID", str(e), {"model_id": parsed.model_id})
|
||||
|
||||
cancelled = DOWNLOAD_SERVER.cancel(parsed.model_id)
|
||||
if not cancelled:
|
||||
return _error(404, "NOT_DOWNLOADING",
|
||||
f"No active download for {parsed.model_id}.",
|
||||
{"model_id": parsed.model_id})
|
||||
|
||||
return _ok(schemas_out.CancelDownloadSessionResponse(cancelled=True))
|
||||
|
||||
|
||||
# ----- 4. HuggingFace OAuth status / login start / logout -----
|
||||
|
||||
|
||||
@ROUTES.get("/api/hf-auth-token-status")
|
||||
async def hf_auth_token_status(request: web.Request) -> web.Response:
|
||||
"""Return whether the server holds a usable HF token + its username.
|
||||
|
||||
Used by the settings UI and (out-of-band) by the frontend on
|
||||
login completion. ``token_available`` is true even if the cached
|
||||
access_token is expired — as long as a refresh_token exists, the
|
||||
user is "logged in" from their perspective.
|
||||
"""
|
||||
token_present = HF_AUTH_STORE.has_token()
|
||||
username: Optional[str] = None
|
||||
if token_present:
|
||||
# Resolve the username via whoami. Done in a worker thread because
|
||||
# huggingface_hub's whoami is synchronous + blocks on a network call.
|
||||
tok = await HF_AUTH_STORE.get_valid_token()
|
||||
if tok is not None:
|
||||
try:
|
||||
username = await asyncio.to_thread(_whoami_username, tok.access_token)
|
||||
except Exception as e:
|
||||
logging.debug("[hf_auth] whoami failed: %s", e)
|
||||
return _ok(schemas_out.HfAuthTokenStatusResponse(
|
||||
token_available=token_present,
|
||||
username=username,
|
||||
))
|
||||
|
||||
|
||||
def _whoami_username(token: str) -> Optional[str]:
|
||||
"""Sync helper: ask HF for the user name attached to a token."""
|
||||
from huggingface_hub import HfApi
|
||||
info = HfApi().whoami(token=token)
|
||||
if isinstance(info, dict):
|
||||
return info.get("name") or info.get("fullname")
|
||||
return None
|
||||
|
||||
|
||||
@ROUTES.post("/api/hf-auth-login-start")
|
||||
async def hf_auth_login_start(request: web.Request) -> web.Response:
|
||||
"""Begin one OAuth attempt: bind the callback port, return the URL.
|
||||
|
||||
Rejected outright if this deployment isn't eligible (we don't
|
||||
surface the option on multi-tenant / public-IP installs).
|
||||
"""
|
||||
if not is_hf_auth_eligible():
|
||||
return _error(
|
||||
403, "HF_AUTH_NOT_ELIGIBLE",
|
||||
"This server is not eligible for interactive HuggingFace login. "
|
||||
"It must be bound to a loopback address and not running in "
|
||||
"--multi-user mode.",
|
||||
)
|
||||
try:
|
||||
url = await start_login_flow()
|
||||
except OAuthInProgressError:
|
||||
return _error(
|
||||
409, "HF_AUTH_IN_PROGRESS",
|
||||
"Another HuggingFace login attempt is in progress. Try again "
|
||||
"after it completes or times out.",
|
||||
)
|
||||
except OAuthCallbackError as e:
|
||||
return _error(
|
||||
503, "HF_AUTH_START_FAILED",
|
||||
f"Could not start the HuggingFace login flow: {e}",
|
||||
)
|
||||
return _ok(schemas_out.HfAuthLoginStartResponse(authorize_url=url))
|
||||
|
||||
|
||||
@ROUTES.post("/api/hf-auth-logout")
|
||||
async def hf_auth_logout(request: web.Request) -> web.Response:
|
||||
"""Drop the in-memory + on-disk HF token."""
|
||||
HF_AUTH_STORE.clear()
|
||||
return _ok(schemas_out.HfAuthLogoutResponse(logged_out=True))
|
||||
@@ -0,0 +1,41 @@
|
||||
"""Request schemas for the model-downloader API.
|
||||
|
||||
Each endpoint accepts a small JSON body. Pydantic enforces the shape at
|
||||
the boundary; route handlers operate only on validated values past that.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class AvailabilityStatusRequest(BaseModel):
|
||||
"""``POST /api/models-availability-status``.
|
||||
|
||||
Sent by the frontend on each poll. Each entry is ``{model_id: url}``;
|
||||
the URL is the one declared in ``properties.models[i].url`` in the
|
||||
workflow JSON and lets the server compute per-id metadata
|
||||
(``file_size`` + ``is_hf_downloadable``) on the same request.
|
||||
"""
|
||||
models: dict[str, str] = Field(default_factory=dict)
|
||||
|
||||
|
||||
class DownloadModelsRequest(BaseModel):
|
||||
"""``POST /api/download-models``.
|
||||
|
||||
Same shape as the metadata request — the URL for each model_id.
|
||||
Returns immediately after validation and scheduling.
|
||||
"""
|
||||
models: dict[str, str] = Field(default_factory=dict)
|
||||
|
||||
|
||||
class CancelDownloadSessionRequest(BaseModel):
|
||||
"""``POST /api/cancel-model-download-session``."""
|
||||
model_id: str
|
||||
|
||||
|
||||
__all__ = [
|
||||
"AvailabilityStatusRequest",
|
||||
"DownloadModelsRequest",
|
||||
"CancelDownloadSessionRequest",
|
||||
]
|
||||
@@ -0,0 +1,81 @@
|
||||
"""Response schemas for the model-downloader API."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Literal, Optional
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
||||
ModelState = Literal["available", "missing", "downloading"]
|
||||
|
||||
|
||||
class DownloadProgress(BaseModel):
|
||||
"""Embedded in a model entry when its state is ``downloading``."""
|
||||
bytes_downloaded: int
|
||||
total_bytes: Optional[int] = None
|
||||
progress: Optional[float] = None # fraction in [0,1]; null until total known
|
||||
|
||||
|
||||
class ModelStatusEntry(BaseModel):
|
||||
"""Everything the UI needs to render one row, in one shot.
|
||||
|
||||
``state`` reflects what the server has on disk + in-flight; ``file_size``
|
||||
and ``is_hf_downloadable`` come from probes (intrinsic; cached).
|
||||
The HF fields are populated for every poll (cached on the server),
|
||||
so license-acceptance flips show up within one poll interval without
|
||||
any frontend cache invalidation.
|
||||
"""
|
||||
state: ModelState
|
||||
progress: Optional[DownloadProgress] = None
|
||||
file_size: Optional[int] = None
|
||||
# HF-only: True iff the server can fetch this URL with current auth
|
||||
# state. False iff gated and lacking access. None for non-HF URLs.
|
||||
is_hf_downloadable: Optional[bool] = None
|
||||
|
||||
|
||||
class HfAuthStatus(BaseModel):
|
||||
"""Snapshot of HF login state, embedded in availability response."""
|
||||
token_available: bool
|
||||
eligible: bool
|
||||
|
||||
|
||||
class AvailabilityStatusResponse(BaseModel):
|
||||
models: dict[str, ModelStatusEntry]
|
||||
hf_auth: HfAuthStatus
|
||||
|
||||
|
||||
class DownloadModelsResponse(BaseModel):
|
||||
accepted: bool
|
||||
scheduled: list[str]
|
||||
|
||||
|
||||
class CancelDownloadSessionResponse(BaseModel):
|
||||
cancelled: bool
|
||||
|
||||
|
||||
class HfAuthTokenStatusResponse(BaseModel):
|
||||
token_available: bool
|
||||
username: Optional[str] = None
|
||||
|
||||
|
||||
class HfAuthLoginStartResponse(BaseModel):
|
||||
authorize_url: str
|
||||
|
||||
|
||||
class HfAuthLogoutResponse(BaseModel):
|
||||
logged_out: bool
|
||||
|
||||
|
||||
__all__ = [
|
||||
"ModelState",
|
||||
"DownloadProgress",
|
||||
"ModelStatusEntry",
|
||||
"HfAuthStatus",
|
||||
"AvailabilityStatusResponse",
|
||||
"DownloadModelsResponse",
|
||||
"CancelDownloadSessionResponse",
|
||||
"HfAuthTokenStatusResponse",
|
||||
"HfAuthLoginStartResponse",
|
||||
"HfAuthLogoutResponse",
|
||||
]
|
||||
@@ -0,0 +1,179 @@
|
||||
"""Process-wide registry of in-flight model downloads.
|
||||
|
||||
A single instance, ``DOWNLOAD_SERVER``, tracks every currently-running
|
||||
server-side model fetch. Designed to be safe with multiple concurrent
|
||||
clients hitting the API: each model_id has at most one active session,
|
||||
and the API rejects requests that conflict with in-flight downloads.
|
||||
|
||||
Cancellation is cooperative — the download loop checks ``is_active`` on
|
||||
its own session between chunks and raises ``DownloadCancelled`` when the
|
||||
session has been removed from the registry. This avoids the complications
|
||||
of ``Task.cancel()`` from outside the loop while still giving deterministic
|
||||
rollback semantics (the worker is responsible for deleting its own
|
||||
``.tmp`` on the cancel path).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
import threading
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional
|
||||
|
||||
from app.model_downloader.paths import iter_all_tmp_paths
|
||||
|
||||
|
||||
class DownloadCancelled(Exception):
|
||||
"""Raised by the streaming loop when its session has been removed
|
||||
from the registry (cancellation request) and the worker should roll
|
||||
back its ``.tmp`` file."""
|
||||
|
||||
|
||||
@dataclass
|
||||
class DownloadSession:
|
||||
"""One in-flight download.
|
||||
|
||||
``progress`` is a fraction in ``[0.0, 1.0]``; ``None`` until the first
|
||||
byte arrives and we know whether the response carries a
|
||||
``Content-Length``. ``total_bytes`` mirrors that header when present.
|
||||
"""
|
||||
model_id: str
|
||||
url: str
|
||||
progress: Optional[float] = None
|
||||
bytes_downloaded: int = 0
|
||||
total_bytes: Optional[int] = None
|
||||
# Sequence number used solely as identity for the cancellation check —
|
||||
# so that "cancel + restart" doesn't get confused by stale workers.
|
||||
epoch: int = field(default_factory=lambda: 0)
|
||||
|
||||
|
||||
class DownloadServer:
|
||||
"""Singleton registry of active downloads.
|
||||
|
||||
All mutation goes through this object so concurrent route handlers
|
||||
see a consistent view. The ``_lock`` is a plain threading lock
|
||||
because the registry is consulted from both the asyncio event-loop
|
||||
thread (route handlers) and from any worker coroutines spawned to
|
||||
perform downloads.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._lock = threading.Lock()
|
||||
self._sessions: dict[str, DownloadSession] = {}
|
||||
self._epoch_counter = 0
|
||||
self._orphan_sweep_done = False
|
||||
|
||||
# ----- lifecycle -----
|
||||
|
||||
def sweep_orphan_tmp_files(self) -> None:
|
||||
"""Idempotently sweep ``*.tmp`` files left by crashed downloads.
|
||||
|
||||
Deferred off the import path so module load doesn't block on
|
||||
filesystem I/O against potentially-slow mounts. Each route handler
|
||||
that might create a new ``.tmp`` runs this exactly once.
|
||||
"""
|
||||
with self._lock:
|
||||
if self._orphan_sweep_done:
|
||||
return
|
||||
self._orphan_sweep_done = True
|
||||
for path in iter_all_tmp_paths():
|
||||
try:
|
||||
os.remove(path)
|
||||
logging.info("[model_downloader] removed orphan tmp file: %s", path)
|
||||
except OSError as e:
|
||||
logging.warning("[model_downloader] could not remove %s: %s", path, e)
|
||||
|
||||
# ----- queries -----
|
||||
|
||||
def is_downloading(self, model_id: str) -> bool:
|
||||
with self._lock:
|
||||
return model_id in self._sessions
|
||||
|
||||
def get(self, model_id: str) -> Optional[DownloadSession]:
|
||||
with self._lock:
|
||||
return self._sessions.get(model_id)
|
||||
|
||||
def snapshot(self) -> dict[str, DownloadSession]:
|
||||
"""Return a shallow copy of the current sessions map."""
|
||||
with self._lock:
|
||||
return dict(self._sessions)
|
||||
|
||||
# ----- mutations -----
|
||||
|
||||
def try_register(self, model_id: str, url: str) -> Optional[DownloadSession]:
|
||||
"""Atomically register a new session iff none exists for ``model_id``.
|
||||
|
||||
Returns the new session on success, ``None`` if a session is already
|
||||
in flight. Callers must check the return value — the caller is the
|
||||
sole owner of the session it gets back.
|
||||
"""
|
||||
with self._lock:
|
||||
if model_id in self._sessions:
|
||||
return None
|
||||
self._epoch_counter += 1
|
||||
session = DownloadSession(
|
||||
model_id=model_id,
|
||||
url=url,
|
||||
epoch=self._epoch_counter,
|
||||
)
|
||||
self._sessions[model_id] = session
|
||||
return session
|
||||
|
||||
def update_progress(
|
||||
self,
|
||||
session: DownloadSession,
|
||||
bytes_downloaded: int,
|
||||
total_bytes: Optional[int],
|
||||
) -> None:
|
||||
"""Update progress on a session. No-op if the session has been
|
||||
removed (cancelled) — caller should check ``is_active`` separately."""
|
||||
with self._lock:
|
||||
current = self._sessions.get(session.model_id)
|
||||
if current is None or current.epoch != session.epoch:
|
||||
return
|
||||
current.bytes_downloaded = bytes_downloaded
|
||||
current.total_bytes = total_bytes
|
||||
if total_bytes and total_bytes > 0:
|
||||
current.progress = min(1.0, bytes_downloaded / total_bytes)
|
||||
|
||||
def is_active(self, session: DownloadSession) -> bool:
|
||||
"""True iff this exact session is still the registered one for
|
||||
its model_id. False after cancellation, after completion, or if
|
||||
another session has replaced it."""
|
||||
with self._lock:
|
||||
current = self._sessions.get(session.model_id)
|
||||
return current is not None and current.epoch == session.epoch
|
||||
|
||||
def finish(self, session: DownloadSession) -> None:
|
||||
"""Remove a completed (or cancelled) session from the registry.
|
||||
|
||||
Safe to call multiple times. Only removes if the epoch matches —
|
||||
we never accidentally evict a *newer* session for the same model_id.
|
||||
"""
|
||||
with self._lock:
|
||||
current = self._sessions.get(session.model_id)
|
||||
if current is not None and current.epoch == session.epoch:
|
||||
del self._sessions[session.model_id]
|
||||
|
||||
def reset_for_tests(self) -> None:
|
||||
"""Clear all sessions and reset the epoch counter. Test-only."""
|
||||
with self._lock:
|
||||
self._sessions.clear()
|
||||
self._epoch_counter = 0
|
||||
|
||||
def cancel(self, model_id: str) -> bool:
|
||||
"""Remove the session registered for ``model_id``.
|
||||
|
||||
Returns True if there was an active session to cancel. The worker
|
||||
will discover the cancellation on its next ``is_active`` check
|
||||
and roll back its ``.tmp`` file.
|
||||
"""
|
||||
with self._lock:
|
||||
if model_id in self._sessions:
|
||||
del self._sessions[model_id]
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
DOWNLOAD_SERVER = DownloadServer()
|
||||
@@ -0,0 +1,216 @@
|
||||
"""Streaming download worker with progress reporting and cancellation.
|
||||
|
||||
Each download writes to ``<final_path>.tmp`` and atomically renames into
|
||||
place on success. Between chunks the worker checks the registry for
|
||||
cancellation (via ``DownloadServer.is_active``) and rolls back its
|
||||
``.tmp`` on cancel or on any error.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
from typing import Optional
|
||||
|
||||
import aiohttp
|
||||
|
||||
from app.model_downloader.download_server import (
|
||||
DOWNLOAD_SERVER,
|
||||
DownloadCancelled,
|
||||
DownloadSession,
|
||||
)
|
||||
from app.model_downloader.hf_auth.auth_store import HF_AUTH_STORE
|
||||
from app.model_downloader.hf_url import is_hf_url
|
||||
from app.model_downloader.http_client import get_session, parse_content_length
|
||||
from app.model_downloader.paths import resolve_destination
|
||||
|
||||
|
||||
CHUNK_SIZE = 64 * 1024 # 64 KiB — same scale as other ComfyUI download paths.
|
||||
REQUEST_TIMEOUT = aiohttp.ClientTimeout(total=None, sock_connect=30, sock_read=120)
|
||||
|
||||
|
||||
async def stream_to_disk(session: DownloadSession) -> str:
|
||||
"""Run a single download to completion or cancellation.
|
||||
|
||||
Returns the final on-disk path on success. Removes the ``.tmp`` and
|
||||
raises on cancellation or failure. The session is finished
|
||||
(removed from the registry) exactly once, here — callers do not
|
||||
need to call ``DOWNLOAD_SERVER.finish`` themselves.
|
||||
"""
|
||||
final_path, tmp_path = resolve_destination(session.model_id, session.epoch)
|
||||
os.makedirs(os.path.dirname(final_path), exist_ok=True)
|
||||
|
||||
# Wipe any stale .tmp from a previous failed attempt before we start —
|
||||
# otherwise a partial body could masquerade as our completed download
|
||||
# when the rename finally happens.
|
||||
_remove_if_exists(tmp_path)
|
||||
|
||||
bytes_seen = 0
|
||||
try:
|
||||
http = await get_session()
|
||||
headers = _auth_headers_for(session.url)
|
||||
logging.info(
|
||||
"[model_downloader] starting GET %s (auth=%s)",
|
||||
session.url, "yes" if "Authorization" in headers else "no",
|
||||
)
|
||||
async with http.get(
|
||||
session.url,
|
||||
allow_redirects=True,
|
||||
timeout=REQUEST_TIMEOUT,
|
||||
headers=headers,
|
||||
) as resp:
|
||||
if resp.status != 200:
|
||||
# Capture a snippet of the response body so 4xx/5xx aren't
|
||||
# opaque in the logs — HF returns JSON or HTML with a
|
||||
# human-readable reason on failures.
|
||||
body_snippet = await _read_short(resp)
|
||||
logging.warning(
|
||||
"[model_downloader] GET %s failed: status=%d final_url=%s body=%s",
|
||||
session.url, resp.status, str(resp.url), body_snippet,
|
||||
)
|
||||
raise DownloadError(
|
||||
f"unexpected HTTP {resp.status} fetching {session.url}: {body_snippet}",
|
||||
status=resp.status,
|
||||
)
|
||||
|
||||
total = parse_content_length(resp.headers.get("Content-Length"))
|
||||
DOWNLOAD_SERVER.update_progress(session, 0, total)
|
||||
|
||||
with open(tmp_path, "wb") as f:
|
||||
async for chunk in resp.content.iter_chunked(CHUNK_SIZE):
|
||||
# Cancellation check between chunks. Cheap and means
|
||||
# cancellation latency is bounded by one chunk plus
|
||||
# one ``write()`` — typically well under a second
|
||||
# even on slow disks.
|
||||
if not DOWNLOAD_SERVER.is_active(session):
|
||||
raise DownloadCancelled()
|
||||
f.write(chunk)
|
||||
bytes_seen += len(chunk)
|
||||
DOWNLOAD_SERVER.update_progress(session, bytes_seen, total)
|
||||
|
||||
# Final cancellation check before we promote the .tmp to the real
|
||||
# filename — avoids the awkward case where cancel arrives during
|
||||
# the very last chunk and we'd otherwise commit anyway.
|
||||
if not DOWNLOAD_SERVER.is_active(session):
|
||||
raise DownloadCancelled()
|
||||
|
||||
# Size verification before commit. aiohttp already raises
|
||||
# ClientPayloadError on a truncated Content-Length/chunked body,
|
||||
# but this also catches the HTTP/1.0-style case (no Content-Length
|
||||
# + Connection: close) where a short read can masquerade as a
|
||||
# complete download.
|
||||
if total is not None and bytes_seen != total:
|
||||
raise DownloadError(
|
||||
f"size mismatch for {session.model_id}: "
|
||||
f"got {bytes_seen} of {total} bytes from {session.url}"
|
||||
)
|
||||
|
||||
# Atomic rename. os.replace is atomic within the same filesystem,
|
||||
# which is guaranteed here because tmp lives alongside final_path.
|
||||
os.replace(tmp_path, final_path)
|
||||
logging.info(
|
||||
"[model_downloader] downloaded %s (%d bytes) from %s",
|
||||
session.model_id, bytes_seen, session.url,
|
||||
)
|
||||
return final_path
|
||||
|
||||
except DownloadCancelled:
|
||||
logging.info("[model_downloader] cancelled: %s", session.model_id)
|
||||
_remove_if_exists(tmp_path)
|
||||
raise
|
||||
except Exception as e:
|
||||
logging.warning(
|
||||
"[model_downloader] failed: %s from %s: %s: %s",
|
||||
session.model_id, session.url, type(e).__name__, e,
|
||||
exc_info=True,
|
||||
)
|
||||
_remove_if_exists(tmp_path)
|
||||
raise
|
||||
finally:
|
||||
# In all terminal states (success / cancel / error) drop the
|
||||
# session from the registry. Idempotent — only removes if we're
|
||||
# still the live epoch for this model_id.
|
||||
DOWNLOAD_SERVER.finish(session)
|
||||
|
||||
|
||||
class DownloadError(Exception):
|
||||
"""Network / protocol error during a download."""
|
||||
|
||||
def __init__(self, message: str, status: Optional[int] = None) -> None:
|
||||
super().__init__(message)
|
||||
self.status = status
|
||||
|
||||
|
||||
async def _read_short(resp: aiohttp.ClientResponse, limit: int = 512) -> str:
|
||||
"""Read up to ``limit`` bytes of a response body for logging.
|
||||
|
||||
Used to surface the JSON/HTML reason from an HF non-2xx response in
|
||||
server logs instead of just the status code. Best-effort: any
|
||||
error here is swallowed.
|
||||
"""
|
||||
try:
|
||||
raw = await resp.content.read(limit)
|
||||
return raw.decode("utf-8", errors="replace").strip()
|
||||
except Exception:
|
||||
return "<unreadable>"
|
||||
|
||||
|
||||
def _auth_headers_for(url: str) -> dict[str, str]:
|
||||
"""Return any auth headers we should add to the GET for ``url``.
|
||||
|
||||
For HuggingFace URLs we inject the user's OAuth access token as a
|
||||
Bearer header — this is HF's documented way to access gated repos
|
||||
(see ``huggingface_hub.hf_hub_download``'s wire format). For every
|
||||
other host we send no extra headers; allowlisted public files
|
||||
don't need them and we don't want to leak tokens to other hosts.
|
||||
"""
|
||||
if not is_hf_url(url):
|
||||
return {}
|
||||
tok = HF_AUTH_STORE.get_token_sync()
|
||||
if tok is None or not tok.access_token:
|
||||
return {}
|
||||
return {"Authorization": f"Bearer {tok.access_token}"}
|
||||
|
||||
|
||||
def _remove_if_exists(path: str) -> None:
|
||||
try:
|
||||
os.remove(path)
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
except OSError as e:
|
||||
logging.warning("[model_downloader] could not remove %s: %s", path, e)
|
||||
|
||||
|
||||
async def run_batch_sequential(sessions: list[DownloadSession]) -> None:
|
||||
"""Run a list of sessions one after the other.
|
||||
|
||||
Each session is independent: a failure or cancellation of one does
|
||||
not abort the rest. Cancellations are observable via the registry
|
||||
*before* a given download starts, so a session that's been
|
||||
pre-cancelled (cancel before the worker reached it) just gets skipped.
|
||||
"""
|
||||
for session in sessions:
|
||||
# If the session got cancelled before its turn, skip without
|
||||
# touching disk. This is what makes the per-request "sequential
|
||||
# but cancellable" semantic work.
|
||||
if not DOWNLOAD_SERVER.is_active(session):
|
||||
DOWNLOAD_SERVER.finish(session)
|
||||
continue
|
||||
try:
|
||||
await stream_to_disk(session)
|
||||
except DownloadCancelled:
|
||||
# Already logged + tmp removed inside stream_to_disk.
|
||||
continue
|
||||
except Exception:
|
||||
# stream_to_disk already logged. Continue with the rest of the batch.
|
||||
continue
|
||||
|
||||
|
||||
def schedule_batch(sessions: list[DownloadSession]) -> asyncio.Task:
|
||||
"""Kick off ``run_batch_sequential`` on the running event loop.
|
||||
|
||||
Returned task is fire-and-forget; the API handler returns immediately
|
||||
after scheduling and clients observe progress via the polling endpoints.
|
||||
"""
|
||||
return asyncio.create_task(run_batch_sequential(sessions))
|
||||
@@ -0,0 +1,245 @@
|
||||
"""Per-URL probes for the unified availability endpoint.
|
||||
|
||||
Three cached/derived facts per URL:
|
||||
|
||||
- ``is_gated`` intrinsic to the model; cached forever once known.
|
||||
Determined by ``auth_check(repo_id, token=None)``:
|
||||
``GatedRepoError`` → True, success → False.
|
||||
|
||||
- ``is_hf_downloadable`` depends on the *current* token; recomputed every
|
||||
call. For non-gated URLs this is trivially True
|
||||
(no HF call needed). For gated URLs we run
|
||||
``auth_check`` with the stored token each call.
|
||||
|
||||
- ``file_size`` intrinsic to the file. Cached forever once
|
||||
determined (including ``None`` on transient
|
||||
failure — we don't retry). We only attempt the
|
||||
HEAD when we already know the URL is downloadable
|
||||
to us; that way a failed-because-gated probe
|
||||
never lands as a cached ``None``.
|
||||
|
||||
Caches are per-process, in-memory; small, no eviction needed for the
|
||||
workflow-scale (~tens of URLs). Concurrent calls for the same URL
|
||||
deduplicate via per-URL ``asyncio.Lock``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
import aiohttp
|
||||
from huggingface_hub import HfApi
|
||||
from huggingface_hub.errors import (
|
||||
GatedRepoError,
|
||||
HfHubHTTPError,
|
||||
RepositoryNotFoundError,
|
||||
)
|
||||
|
||||
from app.model_downloader.allowlist import is_url_allowed
|
||||
from app.model_downloader.hf_auth.auth_store import HF_AUTH_STORE
|
||||
from app.model_downloader.hf_url import is_hf_url, repo_id_from_url
|
||||
from app.model_downloader.http_client import get_session, parse_content_length
|
||||
|
||||
|
||||
_HEAD_TIMEOUT = aiohttp.ClientTimeout(total=15)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ProbeResult:
|
||||
file_size: Optional[int]
|
||||
is_hf_downloadable: Optional[bool]
|
||||
|
||||
|
||||
# --- caches -------------------------------------------------------------- #
|
||||
|
||||
|
||||
# url → bool. Whether this URL's HF repo gates access. Intrinsic to the
|
||||
# model — never changes for a given URL.
|
||||
_is_gated_cache: dict[str, bool] = {}
|
||||
|
||||
# url → Optional[int]. The file's size in bytes, ``None`` if a probe
|
||||
# was attempted and produced no answer. **Only populated when we knew
|
||||
# the URL was downloadable to us at probe time** — so gated-without-
|
||||
# access never lands a ``None`` here that we'd be stuck with after login.
|
||||
_file_size_cache: dict[str, Optional[int]] = {}
|
||||
|
||||
# Per-URL locks for single-flight probes — when multiple polls arrive
|
||||
# in the same tick for the same URL, exactly one of them runs the HF
|
||||
# call and the others wait on the result.
|
||||
_locks: dict[str, asyncio.Lock] = {}
|
||||
|
||||
|
||||
def _lock_for(url: str) -> asyncio.Lock:
|
||||
lock = _locks.get(url)
|
||||
if lock is None:
|
||||
lock = asyncio.Lock()
|
||||
_locks[url] = lock
|
||||
return lock
|
||||
|
||||
|
||||
def clear_caches_for_tests() -> None:
|
||||
"""Test-only: drop everything."""
|
||||
_is_gated_cache.clear()
|
||||
_file_size_cache.clear()
|
||||
_locks.clear()
|
||||
|
||||
|
||||
# --- public entrypoint --------------------------------------------------- #
|
||||
|
||||
|
||||
async def probe_url(url: str) -> ProbeResult:
|
||||
"""Return downloadability + size for one URL, using caches where safe."""
|
||||
if not is_url_allowed(url):
|
||||
return ProbeResult(file_size=None, is_hf_downloadable=None)
|
||||
if not is_hf_url(url):
|
||||
# Non-HF: ``is_hf_downloadable`` is "not applicable" (None).
|
||||
# Size we still cache so we don't HEAD on every poll.
|
||||
size = await _get_or_probe_size(url, token=None)
|
||||
return ProbeResult(file_size=size, is_hf_downloadable=None)
|
||||
|
||||
repo_id = repo_id_from_url(url)
|
||||
if repo_id is None:
|
||||
return ProbeResult(file_size=None, is_hf_downloadable=None)
|
||||
|
||||
# Determine intrinsic gating once.
|
||||
gated = await _resolve_is_gated(url, repo_id)
|
||||
if gated is None:
|
||||
return ProbeResult(file_size=None, is_hf_downloadable=None)
|
||||
|
||||
# Compute current-token downloadability per call.
|
||||
tok = await HF_AUTH_STORE.get_valid_token()
|
||||
token_str: Optional[str] = tok.access_token if tok else None
|
||||
if not gated:
|
||||
is_hf_downloadable: Optional[bool] = True
|
||||
else:
|
||||
is_hf_downloadable = await _auth_check_with_token(repo_id, token_str)
|
||||
|
||||
if is_hf_downloadable is True:
|
||||
size = await _get_or_probe_size(url, token=token_str)
|
||||
else:
|
||||
# Skip the HEAD entirely — would 401 and we'd be stuck with
|
||||
# cached None that survives a later login.
|
||||
size = None
|
||||
|
||||
return ProbeResult(file_size=size, is_hf_downloadable=is_hf_downloadable)
|
||||
|
||||
|
||||
# --- gated/auth probes --------------------------------------------------- #
|
||||
|
||||
|
||||
async def _resolve_is_gated(url: str, repo_id: str) -> Optional[bool]:
|
||||
"""Decide once whether ``repo_id`` is a gated repo."""
|
||||
cached = _is_gated_cache.get(url)
|
||||
if cached is not None:
|
||||
return cached
|
||||
|
||||
async with _lock_for(url):
|
||||
cached = _is_gated_cache.get(url)
|
||||
if cached is not None:
|
||||
return cached
|
||||
# Probe anonymously (token=None) on purpose: an unauthenticated
|
||||
# auth_check is what makes HF raise GatedRepoError for gated repos.
|
||||
# With a token, a gated-but-accepted repo would succeed and look
|
||||
# ungated.
|
||||
try:
|
||||
await asyncio.to_thread(_auth_check_sync, repo_id, None)
|
||||
_is_gated_cache[url] = False
|
||||
return False
|
||||
except GatedRepoError:
|
||||
_is_gated_cache[url] = True
|
||||
return True
|
||||
except RepositoryNotFoundError:
|
||||
# Repo doesn't exist publicly. Treat as gated — we can't
|
||||
# serve it without auth, and an authenticated check might
|
||||
# still succeed if it's a private repo the user can see.
|
||||
_is_gated_cache[url] = True
|
||||
return True
|
||||
except (HfHubHTTPError, Exception) as e:
|
||||
logging.debug(
|
||||
"[hf_auth] is_gated probe failed for %s (will retry): %s",
|
||||
repo_id, e,
|
||||
)
|
||||
return None # don't cache; retry next call
|
||||
|
||||
|
||||
async def _auth_check_with_token(
|
||||
repo_id: str, token: Optional[str]
|
||||
) -> Optional[bool]:
|
||||
"""Auth-check with the supplied token. True/False/None per outcome."""
|
||||
try:
|
||||
await asyncio.to_thread(_auth_check_sync, repo_id, token)
|
||||
return True
|
||||
except GatedRepoError:
|
||||
return False
|
||||
except RepositoryNotFoundError:
|
||||
return False
|
||||
except HfHubHTTPError as e:
|
||||
# 401/403 covers org-SSO-required, revoked tokens, and similar —
|
||||
# all of which mean "can't fetch right now" from the user's POV.
|
||||
status = getattr(getattr(e, "response", None), "status_code", None)
|
||||
if status in (401, 403):
|
||||
return False
|
||||
logging.debug(
|
||||
"[hf_auth] auth_check transient failure for %s: %s", repo_id, e,
|
||||
)
|
||||
return None
|
||||
except Exception as e:
|
||||
logging.warning("[hf_auth] unexpected auth_check error for %s: %s", repo_id, e)
|
||||
return None
|
||||
|
||||
|
||||
def _auth_check_sync(repo_id: str, token: Optional[str]) -> None:
|
||||
"""Thin sync wrapper around ``HfApi.auth_check`` for ``asyncio.to_thread``."""
|
||||
HfApi().auth_check(repo_id, token=token)
|
||||
|
||||
|
||||
# --- size probe ---------------------------------------------------------- #
|
||||
|
||||
|
||||
async def _get_or_probe_size(url: str, token: Optional[str]) -> Optional[int]:
|
||||
"""Return the cached size or HEAD the URL once and cache the result."""
|
||||
if url in _file_size_cache:
|
||||
return _file_size_cache[url]
|
||||
|
||||
async with _lock_for(url):
|
||||
if url in _file_size_cache:
|
||||
return _file_size_cache[url]
|
||||
size = await _probe_size_once(url, token=token)
|
||||
_file_size_cache[url] = size
|
||||
return size
|
||||
|
||||
|
||||
async def _probe_size_once(url: str, token: Optional[str]) -> Optional[int]:
|
||||
"""HEAD the URL and return the file size in bytes, or None on failure.
|
||||
|
||||
HuggingFace serves LFS-tracked files via a 302 to a signed CDN URL.
|
||||
The real file size lives in the non-standard ``X-Linked-Size`` header
|
||||
on that 302 response (``Content-Length`` is the redirect-body length).
|
||||
Disabling redirect-follow lets us read either header on the same
|
||||
response:
|
||||
|
||||
- LFS files: 302 + ``X-Linked-Size``
|
||||
- Small/non-LFS files: 200 + ``Content-Length``
|
||||
"""
|
||||
headers = {"Authorization": f"Bearer {token}"} if token else {}
|
||||
try:
|
||||
session = await get_session()
|
||||
async with session.head(
|
||||
url, allow_redirects=False, timeout=_HEAD_TIMEOUT, headers=headers,
|
||||
) as resp:
|
||||
linked = parse_content_length(resp.headers.get("X-Linked-Size"))
|
||||
if linked is not None:
|
||||
return linked
|
||||
if resp.status == 200:
|
||||
return parse_content_length(resp.headers.get("Content-Length"))
|
||||
return None
|
||||
except (aiohttp.ClientError, TimeoutError, OSError):
|
||||
return None
|
||||
|
||||
|
||||
# Backward-compat shim so consumers that still import the old name keep
|
||||
# building during the refactor; can be removed once routes are updated.
|
||||
MetadataProbeResult = ProbeResult
|
||||
@@ -0,0 +1,121 @@
|
||||
"""In-memory token cache with lazy disk persistence + refresh.
|
||||
|
||||
Public surface is the ``HF_AUTH_STORE`` singleton. Callers ask
|
||||
``get_valid_token()``; the store transparently refreshes from disk
|
||||
on first use, refreshes via the OAuth refresh_token if the cached
|
||||
access_token is expired, and returns ``None`` if neither path works.
|
||||
|
||||
The refresh path imports ``oauth.refresh_access_token`` lazily to
|
||||
avoid an import cycle (oauth needs the store to save tokens it
|
||||
acquires).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import threading
|
||||
from typing import Optional
|
||||
|
||||
from app.model_downloader.hf_auth.token_store import (
|
||||
Token,
|
||||
delete_token,
|
||||
load_token,
|
||||
save_token,
|
||||
)
|
||||
|
||||
|
||||
class HfAuthStore:
|
||||
def __init__(self) -> None:
|
||||
self._lock = threading.Lock()
|
||||
self._token: Optional[Token] = None
|
||||
self._loaded_from_disk = False
|
||||
|
||||
def _ensure_loaded(self) -> None:
|
||||
"""Read the disk token into memory on first access."""
|
||||
if self._loaded_from_disk:
|
||||
return
|
||||
with self._lock:
|
||||
if self._loaded_from_disk:
|
||||
return
|
||||
self._token = load_token()
|
||||
self._loaded_from_disk = True
|
||||
|
||||
def has_token(self) -> bool:
|
||||
"""Cheap check: is there any token in memory?
|
||||
|
||||
Does not attempt refresh; an expired-but-refreshable token still
|
||||
counts as "logged in" from the user's perspective.
|
||||
"""
|
||||
self._ensure_loaded()
|
||||
return self._token is not None
|
||||
|
||||
def _store_token_locked(self, token: Token) -> None:
|
||||
"""Set the in-memory token and persist it to disk.
|
||||
|
||||
Caller must already hold ``self._lock``. Keeping the disk write inside
|
||||
the lock means memory and disk flip together — a concurrent ``clear()``
|
||||
or refresh can't interleave between them.
|
||||
"""
|
||||
self._token = token
|
||||
self._loaded_from_disk = True
|
||||
save_token(token)
|
||||
|
||||
def set_token(self, token: Token) -> None:
|
||||
"""Replace the in-memory token and persist to disk (atomically)."""
|
||||
with self._lock:
|
||||
self._store_token_locked(token)
|
||||
|
||||
def clear(self) -> None:
|
||||
"""Forget the token in memory and on disk (logout)."""
|
||||
with self._lock:
|
||||
self._token = None
|
||||
self._loaded_from_disk = True
|
||||
delete_token()
|
||||
|
||||
def get_token_sync(self) -> Optional[Token]:
|
||||
"""Return the cached token without refreshing.
|
||||
|
||||
Sync callers (e.g. constructing an Authorization header in a
|
||||
non-async path) use this. They accept an expired token over
|
||||
``None``; HF will simply return 401 and the caller can decide
|
||||
what to do.
|
||||
"""
|
||||
self._ensure_loaded()
|
||||
return self._token
|
||||
|
||||
async def get_valid_token(self) -> Optional[Token]:
|
||||
"""Return a fresh token, refreshing via OAuth if necessary.
|
||||
|
||||
Returns ``None`` if there's no cached token at all, or if the
|
||||
cached token is expired and refresh failed. Callers should
|
||||
treat that as "user is not logged in".
|
||||
"""
|
||||
self._ensure_loaded()
|
||||
with self._lock:
|
||||
tok = self._token
|
||||
if tok is None:
|
||||
return None
|
||||
if tok.is_valid():
|
||||
return tok
|
||||
if not tok.refresh_token:
|
||||
return None
|
||||
|
||||
# Lazy import to avoid the oauth ↔ store import cycle.
|
||||
from app.model_downloader.hf_auth.oauth import refresh_access_token
|
||||
|
||||
try:
|
||||
refreshed = await refresh_access_token(tok.refresh_token)
|
||||
except Exception as e:
|
||||
logging.warning("[hf_auth] token refresh failed: %s", e)
|
||||
return None
|
||||
|
||||
with self._lock:
|
||||
# If a logout (clear) or another update replaced the token while we
|
||||
# were awaiting the refresh, don't resurrect the old session.
|
||||
if self._token is not tok:
|
||||
return None
|
||||
self._store_token_locked(refreshed)
|
||||
return refreshed
|
||||
|
||||
|
||||
HF_AUTH_STORE = HfAuthStore()
|
||||
@@ -0,0 +1,55 @@
|
||||
"""Whether this deployment is allowed to do interactive HF OAuth.
|
||||
|
||||
We only let the server hold a HuggingFace token under a strict trust
|
||||
assumption: this is a *single tenant local* install. Concretely:
|
||||
|
||||
- The server is bound to a loopback address. SSH tunneling /
|
||||
reverse-proxies can defeat this, but it's the strongest signal
|
||||
we have without an authentication system.
|
||||
- ``--multi-user`` is off. A shared token used implicitly by multiple
|
||||
declared users would be a footgun — one user's gated downloads
|
||||
would silently authenticate as another.
|
||||
|
||||
Anything else and the frontend hides the HF login UI entirely; gated
|
||||
models continue to show the "acquire it manually" message.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import ipaddress
|
||||
import socket
|
||||
|
||||
from comfy.cli_args import args
|
||||
|
||||
|
||||
def _is_loopback(host: str | None) -> bool:
|
||||
"""Duplicates ``server.is_loopback`` (small, no shared module yet).
|
||||
|
||||
Resolves a host or IP literal to whether it lives on the loopback
|
||||
interface (127.0.0.0/8 for IPv4, ::1 for IPv6). Returns False for
|
||||
``0.0.0.0`` / ``::`` because those are bind-all wildcards, not
|
||||
loopback.
|
||||
"""
|
||||
if host is None:
|
||||
return False
|
||||
try:
|
||||
return ipaddress.ip_address(host).is_loopback
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
loopback = False
|
||||
for family in (socket.AF_INET, socket.AF_INET6):
|
||||
try:
|
||||
r = socket.getaddrinfo(host, None, family, socket.SOCK_STREAM)
|
||||
for _family, _, _, _, sockaddr in r:
|
||||
if not ipaddress.ip_address(sockaddr[0]).is_loopback:
|
||||
return loopback
|
||||
loopback = True
|
||||
except socket.gaierror:
|
||||
pass
|
||||
return loopback
|
||||
|
||||
|
||||
def is_hf_auth_eligible() -> bool:
|
||||
"""True iff this deployment may surface the HF OAuth flow."""
|
||||
return _is_loopback(args.listen) and not args.multi_user
|
||||
@@ -0,0 +1,301 @@
|
||||
"""OAuth 2.0 PKCE flow against HuggingFace's authorization server.
|
||||
|
||||
Wired so that ``POST /api/hf-auth-login-start`` can:
|
||||
1. Generate state + PKCE verifier/challenge in this process.
|
||||
2. Spin up a short-lived loopback HTTP server at port 41954 to
|
||||
receive the redirect callback from HF.
|
||||
3. Return the ``authorize_url`` for the frontend to open in a new tab.
|
||||
|
||||
After the user grants consent on huggingface.co, HF redirects to the
|
||||
local callback URL with ``code`` and ``state``. The callback server
|
||||
validates ``state`` (CSRF), exchanges the code for tokens via PKCE,
|
||||
hands the resulting Token to ``HF_AUTH_STORE.set_token``, and shuts
|
||||
itself down.
|
||||
|
||||
Before this can be exercised end-to-end a maintainer must register a
|
||||
HuggingFace OAuth app and substitute the ``HF_CLIENT_ID`` placeholder
|
||||
below. See the comment above the constant for the exact steps.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import hashlib
|
||||
import logging
|
||||
import secrets
|
||||
import threading
|
||||
import time
|
||||
|
||||
import aiohttp
|
||||
from aiohttp import web
|
||||
|
||||
from app.model_downloader.hf_auth.auth_store import HF_AUTH_STORE
|
||||
from app.model_downloader.hf_auth.token_store import Token
|
||||
from app.model_downloader.http_client import get_session
|
||||
|
||||
|
||||
# --- HF OAuth app registration -------------------------------------------- #
|
||||
# NOTE: The OAuth client_id below is a placeholder. Before this feature can be
|
||||
# exercised end-to-end, a maintainer must register a HuggingFace OAuth app
|
||||
# under a Comfy-Org-controlled HF account and substitute its client_id here.
|
||||
# Detailed walkthrough is in docs/server-side-model-downloads-handover.html
|
||||
# ("HuggingFace OAuth app setup" section). Short version:
|
||||
# 1. huggingface.co → Settings → Connected Apps → "Create app"
|
||||
# 2. Default Scopes: check ``openid`` + ``profile`` (User Info) and
|
||||
# ``gated-repos`` (Repository Access). Leave everything else off.
|
||||
# 3. Redirect URLs: exactly ``http://127.0.0.1:41954/api/auth/huggingface/callback``
|
||||
# — must match ``REDIRECT_URI`` below; change both in lockstep if you
|
||||
# change ``CALLBACK_PORT``.
|
||||
# 4. Save → copy the resulting Client ID into ``HF_CLIENT_ID`` below.
|
||||
# The client_id is not a secret (it travels through the user's browser in
|
||||
# plaintext); HF's "Public app" type means there's no client secret to
|
||||
# manage — PKCE replaces it.
|
||||
HF_CLIENT_ID = "REPLACE_ME_WITH_COMFY_ORG_HF_OAUTH_CLIENT_ID"
|
||||
|
||||
CALLBACK_HOST = "127.0.0.1"
|
||||
CALLBACK_PORT = 41954
|
||||
CALLBACK_PATH = "/api/auth/huggingface/callback"
|
||||
REDIRECT_URI = f"http://{CALLBACK_HOST}:{CALLBACK_PORT}{CALLBACK_PATH}"
|
||||
|
||||
AUTHORIZE_URL = "https://huggingface.co/oauth/authorize"
|
||||
TOKEN_URL = "https://huggingface.co/oauth/token"
|
||||
_TOKEN_REQUEST_TIMEOUT = aiohttp.ClientTimeout(total=30)
|
||||
# Minimal scope set for the feature:
|
||||
# - openid : required by HF when the app uses OIDC at all
|
||||
# - profile : lets ``HfApi.whoami(token=...)`` return a username for the
|
||||
# settings UI; cosmetic but expected
|
||||
# - gated-repos : grants the token enough to call ``auth_check`` and
|
||||
# download files from public gated repos the user has
|
||||
# accepted the license for. The wider ``read-repos`` scope
|
||||
# would also work (it includes ``gated-repos``) but it
|
||||
# additionally grants private-repo read access, which we
|
||||
# don't need and which makes the consent screen scarier
|
||||
# for the user.
|
||||
SCOPE = "openid profile gated-repos"
|
||||
|
||||
# Maximum time the callback server stays up waiting for the user to
|
||||
# complete consent on huggingface.co. Past this, the port closes and
|
||||
# the user has to click "Log in" again.
|
||||
CALLBACK_TIMEOUT_SECS = 300
|
||||
|
||||
|
||||
# Process-wide lock so two simultaneous /api/hf-auth-login-start
|
||||
# requests don't fight over port CALLBACK_PORT.
|
||||
_OAUTH_LOCK = threading.Lock()
|
||||
|
||||
|
||||
class OAuthInProgressError(Exception):
|
||||
"""Another OAuth attempt is already running."""
|
||||
|
||||
|
||||
class OAuthCallbackError(Exception):
|
||||
"""The OAuth callback returned an error (HF denied, port stolen, etc.)."""
|
||||
|
||||
|
||||
# --- PKCE primitives ------------------------------------------------------ #
|
||||
|
||||
|
||||
def _make_pkce() -> tuple[str, str, str]:
|
||||
"""Return ``(verifier, challenge, state)``.
|
||||
|
||||
Verifier never leaves this process. Challenge and state travel
|
||||
through the user's browser. State is checked on the callback to
|
||||
prevent a malicious cross-origin redirect from injecting a token.
|
||||
"""
|
||||
verifier = secrets.token_urlsafe(64)
|
||||
challenge = (
|
||||
base64.urlsafe_b64encode(hashlib.sha256(verifier.encode("ascii")).digest())
|
||||
.rstrip(b"=")
|
||||
.decode("ascii")
|
||||
)
|
||||
state = secrets.token_urlsafe(32)
|
||||
return verifier, challenge, state
|
||||
|
||||
|
||||
def _build_authorize_url(challenge: str, state: str) -> str:
|
||||
from urllib.parse import urlencode
|
||||
|
||||
params = {
|
||||
"client_id": HF_CLIENT_ID,
|
||||
"redirect_uri": REDIRECT_URI,
|
||||
"response_type": "code",
|
||||
"scope": SCOPE,
|
||||
"state": state,
|
||||
"code_challenge": challenge,
|
||||
"code_challenge_method": "S256",
|
||||
}
|
||||
return f"{AUTHORIZE_URL}?{urlencode(params)}"
|
||||
|
||||
|
||||
# --- Token exchange ------------------------------------------------------- #
|
||||
|
||||
|
||||
async def _exchange_code(code: str, verifier: str) -> Token:
|
||||
"""Trade the authorization code for an access+refresh token pair."""
|
||||
data = {
|
||||
"grant_type": "authorization_code",
|
||||
"code": code,
|
||||
"redirect_uri": REDIRECT_URI,
|
||||
"client_id": HF_CLIENT_ID,
|
||||
"code_verifier": verifier,
|
||||
}
|
||||
session = await get_session()
|
||||
async with session.post(TOKEN_URL, data=data, timeout=_TOKEN_REQUEST_TIMEOUT) as resp:
|
||||
resp.raise_for_status()
|
||||
body = await resp.json()
|
||||
return Token(
|
||||
access_token=body["access_token"],
|
||||
refresh_token=body.get("refresh_token"),
|
||||
expires_at=time.time() + float(body.get("expires_in", 3600)),
|
||||
scope=body.get("scope", SCOPE),
|
||||
)
|
||||
|
||||
|
||||
async def refresh_access_token(refresh_token: str) -> Token:
|
||||
"""Trade a refresh_token for a new access (+ possibly refresh) token."""
|
||||
data = {
|
||||
"grant_type": "refresh_token",
|
||||
"refresh_token": refresh_token,
|
||||
"client_id": HF_CLIENT_ID,
|
||||
}
|
||||
session = await get_session()
|
||||
async with session.post(TOKEN_URL, data=data, timeout=_TOKEN_REQUEST_TIMEOUT) as resp:
|
||||
resp.raise_for_status()
|
||||
body = await resp.json()
|
||||
return Token(
|
||||
access_token=body["access_token"],
|
||||
# If HF doesn't rotate refresh tokens, keep using the existing one.
|
||||
refresh_token=body.get("refresh_token", refresh_token),
|
||||
expires_at=time.time() + float(body.get("expires_in", 3600)),
|
||||
scope=body.get("scope", SCOPE),
|
||||
)
|
||||
|
||||
|
||||
# --- Callback server ------------------------------------------------------ #
|
||||
|
||||
|
||||
async def start_login_flow() -> str:
|
||||
"""Begin one OAuth attempt: spawn the callback server, return the URL.
|
||||
|
||||
Returns the URL the frontend should open in a new tab. Raises
|
||||
``OAuthInProgressError`` if another attempt is already running.
|
||||
The callback server runs in the background until the user
|
||||
completes consent or until ``CALLBACK_TIMEOUT_SECS`` elapses;
|
||||
either way the lock + port are released afterward.
|
||||
"""
|
||||
if not _OAUTH_LOCK.acquire(blocking=False):
|
||||
raise OAuthInProgressError()
|
||||
|
||||
try:
|
||||
verifier, challenge, state = _make_pkce()
|
||||
authorize_url = _build_authorize_url(challenge, state)
|
||||
ready: asyncio.Future[None] = asyncio.get_event_loop().create_future()
|
||||
except BaseException:
|
||||
# Failed before handing the lock to the callback-server task: release it
|
||||
# here. (Once the task is spawned, it owns releasing the lock.)
|
||||
_OAUTH_LOCK.release()
|
||||
raise
|
||||
|
||||
asyncio.create_task(_run_callback_server(verifier, state, ready))
|
||||
# Don't return the URL until the callback server is actually bound and
|
||||
# listening — otherwise HF could redirect to a port nothing is serving and
|
||||
# the login would silently dead-end. ``ready`` raises if the bind failed.
|
||||
await ready
|
||||
return authorize_url
|
||||
|
||||
|
||||
async def _run_callback_server(
|
||||
verifier: str, expected_state: str, ready: "asyncio.Future[None]"
|
||||
) -> None:
|
||||
"""Listen for HF's redirect once, capture the token, then shut down.
|
||||
|
||||
Signals ``ready`` once the port is bound (or with an exception if the bind
|
||||
fails), so ``start_login_flow`` only hands back a URL on a live server.
|
||||
"""
|
||||
received: asyncio.Future[Token] = asyncio.get_event_loop().create_future()
|
||||
|
||||
async def handler(request: web.Request) -> web.Response:
|
||||
try:
|
||||
if request.query.get("state") != expected_state:
|
||||
return web.Response(status=400, text="state mismatch")
|
||||
err = request.query.get("error")
|
||||
if err:
|
||||
received.set_exception(OAuthCallbackError(f"HF returned: {err}"))
|
||||
return web.Response(status=400, text=f"OAuth error: {err}")
|
||||
code = request.query.get("code")
|
||||
if not code:
|
||||
return web.Response(status=400, text="missing code")
|
||||
tok = await _exchange_code(code, verifier)
|
||||
if not received.done():
|
||||
received.set_result(tok)
|
||||
return web.Response(
|
||||
content_type="text/html",
|
||||
text=(
|
||||
"<html><body style='font-family:sans-serif;padding:40px'>"
|
||||
"<h2>HuggingFace login successful</h2>"
|
||||
"<p>You can close this tab and return to ComfyUI.</p>"
|
||||
"</body></html>"
|
||||
),
|
||||
)
|
||||
except Exception as exc:
|
||||
if not received.done():
|
||||
received.set_exception(exc)
|
||||
return web.Response(status=500, text=str(exc))
|
||||
|
||||
app = web.Application()
|
||||
app.router.add_get(CALLBACK_PATH, handler)
|
||||
runner = web.AppRunner(app)
|
||||
try:
|
||||
await runner.setup()
|
||||
site = web.TCPSite(runner, CALLBACK_HOST, CALLBACK_PORT, reuse_address=True)
|
||||
await site.start()
|
||||
except Exception as e:
|
||||
# Couldn't bind the callback port (commonly already in use). Tell the
|
||||
# waiting start_login_flow via ``ready`` so it surfaces a clear error
|
||||
# instead of returning a dead URL, and release the lock for next time.
|
||||
logging.warning("[hf_auth] could not start callback server: %s", e)
|
||||
if not ready.done():
|
||||
ready.set_exception(
|
||||
OAuthCallbackError(f"could not bind callback port {CALLBACK_PORT}: {e}")
|
||||
)
|
||||
_OAUTH_LOCK.release()
|
||||
return
|
||||
|
||||
# Bound and listening — now it's safe for start_login_flow to return the URL.
|
||||
if not ready.done():
|
||||
ready.set_result(None)
|
||||
|
||||
try:
|
||||
token = await asyncio.wait_for(received, timeout=CALLBACK_TIMEOUT_SECS)
|
||||
except asyncio.TimeoutError:
|
||||
logging.info("[hf_auth] OAuth login timed out after %ds", CALLBACK_TIMEOUT_SECS)
|
||||
return
|
||||
except OAuthCallbackError as e:
|
||||
logging.warning("[hf_auth] OAuth callback error: %s", e)
|
||||
return
|
||||
except Exception as e:
|
||||
logging.warning("[hf_auth] unexpected OAuth failure: %s", e)
|
||||
return
|
||||
else:
|
||||
HF_AUTH_STORE.set_token(token)
|
||||
logging.info("[hf_auth] OAuth login complete")
|
||||
finally:
|
||||
await runner.cleanup()
|
||||
if _OAUTH_LOCK.locked():
|
||||
_OAUTH_LOCK.release()
|
||||
|
||||
|
||||
def is_login_in_progress() -> bool:
|
||||
"""True iff a callback server is currently bound + waiting."""
|
||||
return _OAUTH_LOCK.locked()
|
||||
|
||||
|
||||
# Re-export for callers that only want the URL builder (e.g. tests).
|
||||
__all__ = [
|
||||
"start_login_flow",
|
||||
"refresh_access_token",
|
||||
"is_login_in_progress",
|
||||
"OAuthInProgressError",
|
||||
"CALLBACK_TIMEOUT_SECS",
|
||||
]
|
||||
@@ -0,0 +1,94 @@
|
||||
"""On-disk persistence for the HuggingFace OAuth token.
|
||||
|
||||
The token shape mirrors what HF returns on the token exchange: an
|
||||
``access_token``, an optional ``refresh_token``, the absolute epoch at
|
||||
which the access token expires, and the granted scope. We persist
|
||||
this so logging in once survives ComfyUI restarts under the internal
|
||||
``__hf_auth`` system-user directory; the file is mode ``0600`` so only
|
||||
the OS user can read it.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import stat
|
||||
import time
|
||||
from dataclasses import asdict, dataclass
|
||||
from typing import Optional
|
||||
|
||||
import folder_paths
|
||||
|
||||
|
||||
# Treat a token as expired this many seconds before its server-reported
|
||||
# ``expires_at`` so we don't try to use a token mid-request only for it
|
||||
# to flip stale between auth_check and the actual GET.
|
||||
EXPIRY_BUFFER_SECS = 60
|
||||
|
||||
TOKEN_FILENAME = "hf_auth_token.json"
|
||||
|
||||
|
||||
@dataclass
|
||||
class Token:
|
||||
"""One OAuth token + the metadata we need to use it."""
|
||||
access_token: str
|
||||
refresh_token: Optional[str]
|
||||
expires_at: float # absolute epoch seconds
|
||||
scope: str = ""
|
||||
|
||||
def is_valid(self) -> bool:
|
||||
"""True iff we can use this token right now."""
|
||||
return (
|
||||
bool(self.access_token)
|
||||
and (self.expires_at - time.time() > EXPIRY_BUFFER_SECS)
|
||||
)
|
||||
|
||||
|
||||
def _token_dir() -> str:
|
||||
return folder_paths.get_system_user_directory("hf_auth")
|
||||
|
||||
|
||||
def _token_path() -> str:
|
||||
return os.path.join(_token_dir(), TOKEN_FILENAME)
|
||||
|
||||
|
||||
def load_token() -> Optional[Token]:
|
||||
"""Read the persisted token, returning ``None`` if absent or corrupt."""
|
||||
path = _token_path()
|
||||
if not os.path.exists(path):
|
||||
return None
|
||||
try:
|
||||
with open(path, "r", encoding="utf-8") as f:
|
||||
data = json.load(f)
|
||||
return Token(**data)
|
||||
except (OSError, json.JSONDecodeError, TypeError) as e:
|
||||
logging.warning("[hf_auth] could not load token at %s: %s", path, e)
|
||||
return None
|
||||
|
||||
|
||||
def save_token(token: Token) -> None:
|
||||
"""Atomically write the token with 0600 permissions."""
|
||||
path = _token_path()
|
||||
os.makedirs(os.path.dirname(path), exist_ok=True)
|
||||
tmp = path + ".tmp"
|
||||
fd = os.open(tmp, os.O_WRONLY | os.O_CREAT | os.O_TRUNC, 0o600)
|
||||
with os.fdopen(fd, "w", encoding="utf-8") as f:
|
||||
json.dump(asdict(token), f)
|
||||
os.replace(tmp, path)
|
||||
try:
|
||||
os.chmod(path, stat.S_IRUSR | stat.S_IWUSR)
|
||||
except OSError as e:
|
||||
# On Windows / weird filesystems chmod may be a no-op; not fatal.
|
||||
logging.debug("[hf_auth] chmod 0600 on %s failed: %s", path, e)
|
||||
|
||||
|
||||
def delete_token() -> None:
|
||||
"""Remove the persisted token; no-op if it doesn't exist."""
|
||||
path = _token_path()
|
||||
try:
|
||||
os.remove(path)
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
except OSError as e:
|
||||
logging.warning("[hf_auth] could not remove token at %s: %s", path, e)
|
||||
@@ -0,0 +1,41 @@
|
||||
"""Parsers for the ``huggingface.co`` URL shape we accept in workflows.
|
||||
|
||||
The download API accepts URLs of the form
|
||||
``https://huggingface.co/<org>/<repo>/resolve/<rev>/<path/to/file>``.
|
||||
We need to recover ``<org>/<repo>`` (the *repo_id*) from such URLs for
|
||||
``huggingface_hub`` API calls (notably ``HfApi.auth_check``).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Optional
|
||||
from urllib.parse import urlparse
|
||||
|
||||
_HF_HOST = "huggingface.co"
|
||||
|
||||
|
||||
def is_hf_url(url: str) -> bool:
|
||||
"""Cheap host check — does this URL point at huggingface.co?"""
|
||||
try:
|
||||
return urlparse(url).hostname == _HF_HOST
|
||||
except ValueError:
|
||||
return False
|
||||
|
||||
|
||||
def repo_id_from_url(url: str) -> Optional[str]:
|
||||
"""Extract ``<org>/<repo>`` from an HF model file URL.
|
||||
|
||||
Returns ``None`` if the URL isn't on huggingface.co or doesn't look
|
||||
like a model-file URL. The expected shape is
|
||||
``/<org>/<repo>/resolve/<rev>/<path>`` — anything else
|
||||
(datasets, spaces, /tree/, /blob/, …) we treat as out of scope here.
|
||||
"""
|
||||
if not is_hf_url(url):
|
||||
return None
|
||||
parts = urlparse(url).path.lstrip("/").split("/")
|
||||
if len(parts) < 4 or parts[2] != "resolve":
|
||||
return None
|
||||
org, repo = parts[0], parts[1]
|
||||
if not org or not repo:
|
||||
return None
|
||||
return f"{org}/{repo}"
|
||||
@@ -0,0 +1,63 @@
|
||||
"""Lazy module-level aiohttp ClientSession.
|
||||
|
||||
A single shared session means TLS handshakes are reused across HEAD probes
|
||||
and the subsequent GETs to the same host (HuggingFace is the dominant
|
||||
case), which is a noticeable speedup on cold connections.
|
||||
|
||||
We deliberately don't close the session at process exit — aiohttp's
|
||||
warning about unclosed sessions is benign at shutdown, and adding atexit
|
||||
plumbing buys nothing because the OS reclaims the sockets anyway. The
|
||||
session lifetime is the lifetime of the Python process.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import ssl
|
||||
from typing import Optional
|
||||
|
||||
import aiohttp
|
||||
import certifi
|
||||
|
||||
|
||||
# Larger per-host pool than aiohttp's default (=100 total / =0 per host)
|
||||
# so concurrent gated probes + a download to the same host don't queue.
|
||||
_CONNECTOR_LIMIT_PER_HOST = 8
|
||||
|
||||
_session: Optional[aiohttp.ClientSession] = None
|
||||
_lock = asyncio.Lock()
|
||||
|
||||
|
||||
def ssl_context() -> ssl.SSLContext:
|
||||
"""TLS context pinned to certifi's CA bundle.
|
||||
aiohttp's default context uses the OS trust store, which isn't wired up
|
||||
on some Python installs (python.org macOS, slim containers) — there TLS
|
||||
to huggingface.co fails with CERTIFICATE_VERIFY_FAILED.
|
||||
"""
|
||||
return ssl.create_default_context(cafile=certifi.where())
|
||||
|
||||
|
||||
async def get_session() -> aiohttp.ClientSession:
|
||||
"""Return the shared session, creating it on first call."""
|
||||
global _session
|
||||
if _session is not None and not _session.closed:
|
||||
return _session
|
||||
async with _lock:
|
||||
if _session is None or _session.closed:
|
||||
connector = aiohttp.TCPConnector(
|
||||
limit_per_host=_CONNECTOR_LIMIT_PER_HOST,
|
||||
ssl=ssl_context(),
|
||||
)
|
||||
_session = aiohttp.ClientSession(connector=connector)
|
||||
return _session
|
||||
|
||||
|
||||
def parse_content_length(value: Optional[str]) -> Optional[int]:
|
||||
"""Parse a byte-count header value, or None if absent/malformed/negative."""
|
||||
if not value:
|
||||
return None
|
||||
try:
|
||||
n = int(value)
|
||||
except ValueError:
|
||||
return None
|
||||
return n if n >= 0 else None
|
||||
@@ -0,0 +1,111 @@
|
||||
"""Path resolution for model downloads.
|
||||
|
||||
Model identifiers used across the download API are *relative destination
|
||||
paths* of the form ``<directory>/<filename>`` (e.g. ``loras/my_lora.safetensors``).
|
||||
This module turns one of those identifiers into an absolute on-disk path
|
||||
under one of ComfyUI's registered model folders, while rejecting unknown
|
||||
folders, path traversal, and other ill-formed inputs.
|
||||
"""
|
||||
|
||||
import os
|
||||
import re
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import folder_paths
|
||||
|
||||
|
||||
# Constrain components so a model_id can never escape its target directory.
|
||||
# - directory: a single path segment of safe chars
|
||||
# - filename: a single path segment of safe chars, must end with a model extension
|
||||
_SEGMENT_RE = re.compile(r"^[A-Za-z0-9._-]+$")
|
||||
|
||||
# Destination filename must name a model file (same set as the URL allowlist),
|
||||
# so a download can't land as e.g. ``foo.txt`` that ComfyUI won't recognise.
|
||||
_MODEL_EXTENSIONS = (".safetensors", ".sft", ".ckpt", ".pth", ".pt")
|
||||
|
||||
# Distinctive temp suffix so the startup orphan-sweep only removes files THIS
|
||||
# subsystem created — never unrelated ``*.tmp`` files in the model dirs.
|
||||
_TMP_SUFFIX = ".comfy-download.tmp"
|
||||
|
||||
|
||||
class InvalidModelId(ValueError):
|
||||
"""Raised when a model_id is syntactically invalid or refers to an
|
||||
unknown model folder."""
|
||||
|
||||
|
||||
def parse_model_id(model_id: str) -> Tuple[str, str]:
|
||||
"""Split ``<directory>/<filename>`` and validate both components.
|
||||
|
||||
Returns ``(directory, filename)``. Raises ``InvalidModelId`` on
|
||||
malformed input. Does NOT touch the filesystem.
|
||||
"""
|
||||
if not isinstance(model_id, str) or "/" not in model_id:
|
||||
raise InvalidModelId(f"model_id must be '<directory>/<filename>', got {model_id!r}")
|
||||
directory, _, filename = model_id.partition("/")
|
||||
if "/" in filename or not directory or not filename:
|
||||
raise InvalidModelId(f"model_id must be exactly one '/' separator, got {model_id!r}")
|
||||
if not _SEGMENT_RE.match(directory):
|
||||
raise InvalidModelId(f"invalid directory segment {directory!r}")
|
||||
if not _SEGMENT_RE.match(filename):
|
||||
raise InvalidModelId(f"invalid filename segment {filename!r}")
|
||||
if not filename.endswith(_MODEL_EXTENSIONS):
|
||||
raise InvalidModelId(
|
||||
f"filename must end with a model extension {_MODEL_EXTENSIONS}, got {filename!r}"
|
||||
)
|
||||
if directory not in folder_paths.folder_names_and_paths:
|
||||
raise InvalidModelId(f"unknown model folder {directory!r}")
|
||||
return directory, filename
|
||||
|
||||
|
||||
def resolve_existing(model_id: str) -> Optional[str]:
|
||||
"""Return the absolute path of an installed model, or None if missing.
|
||||
|
||||
Honours ``extra_model_paths.yaml`` transparently via
|
||||
``folder_paths.get_full_path``.
|
||||
"""
|
||||
directory, filename = parse_model_id(model_id)
|
||||
return folder_paths.get_full_path(directory, filename)
|
||||
|
||||
|
||||
def resolve_destination(model_id: str, epoch: int = 0) -> Tuple[str, str]:
|
||||
"""Return ``(final_path, tmp_path)`` for a download.
|
||||
|
||||
Downloads land at the first registered path for the model's directory
|
||||
(the "primary" location). The temp sibling is the write target, atomically
|
||||
renamed onto ``final_path`` on success.
|
||||
|
||||
``tmp_path`` embeds the session ``epoch`` so a cancel+retry of the same
|
||||
model never shares a temp path between the old (cancelling) worker and the
|
||||
new attempt — otherwise the old worker's rollback could delete the new
|
||||
worker's in-progress file. The distinctive suffix scopes the orphan sweep.
|
||||
"""
|
||||
directory, filename = parse_model_id(model_id)
|
||||
roots = folder_paths.get_folder_paths(directory)
|
||||
if not roots:
|
||||
raise InvalidModelId(f"no on-disk path registered for folder {directory!r}")
|
||||
root = roots[0]
|
||||
final_path = os.path.join(root, filename)
|
||||
tmp_path = f"{final_path}.{epoch}{_TMP_SUFFIX}"
|
||||
return final_path, tmp_path
|
||||
|
||||
|
||||
def iter_all_tmp_paths():
|
||||
"""Yield this subsystem's temp files under every registered model folder.
|
||||
|
||||
Matches only our distinctive ``_TMP_SUFFIX`` (not every ``*.tmp``) so the
|
||||
startup orphan-sweep can't delete temp files created by other tools.
|
||||
"""
|
||||
seen_roots: set[str] = set()
|
||||
for directory in folder_paths.folder_names_and_paths.keys():
|
||||
for root in folder_paths.get_folder_paths(directory):
|
||||
if root in seen_roots or not os.path.isdir(root):
|
||||
continue
|
||||
seen_roots.add(root)
|
||||
try:
|
||||
for entry in os.scandir(root):
|
||||
if entry.is_file() and entry.name.endswith(_TMP_SUFFIX):
|
||||
yield entry.path
|
||||
except OSError:
|
||||
# Folder might be unreadable / missing on certain mounts —
|
||||
# not fatal, just skip it.
|
||||
continue
|
||||
+26
-2
@@ -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()
|
||||
|
||||
+15
-1
@@ -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):
|
||||
|
||||
+11
-5
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -98,12 +98,24 @@ def _parse_cli_feature_flags() -> dict[str, Any]:
|
||||
|
||||
|
||||
# Default server capabilities
|
||||
def _hf_auth_eligible_at_startup() -> bool:
|
||||
"""Snapshot eligibility once at feature-flag init time.
|
||||
|
||||
Imports lazily because the flags module loads very early in the
|
||||
server boot sequence — earlier than the model_downloader package.
|
||||
"""
|
||||
from app.model_downloader.hf_auth.eligibility import is_hf_auth_eligible
|
||||
return is_hf_auth_eligible()
|
||||
|
||||
|
||||
_CORE_FEATURE_FLAGS: dict[str, Any] = {
|
||||
"supports_preview_metadata": True,
|
||||
"max_upload_size": args.max_upload_size * 1024 * 1024, # Convert MB to bytes
|
||||
"extension": {"manager": {"supports_v4": True}},
|
||||
"node_replacements": True,
|
||||
"assets": args.enable_assets,
|
||||
"server_side_model_downloads": True,
|
||||
"hf_auth_eligible": _hf_auth_eligible_at_startup(),
|
||||
}
|
||||
|
||||
# CLI-provided flags cannot overwrite core flags
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
from typing import Literal
|
||||
from typing import Any, Literal
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
@@ -316,3 +316,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)
|
||||
|
||||
@@ -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'
|
||||
|
||||
@@ -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)
|
||||
@@ -1,3 +1,4 @@
|
||||
import base64
|
||||
import hashlib
|
||||
import logging
|
||||
import math
|
||||
@@ -20,6 +21,10 @@ from comfy_api_nodes.apis.bytedance import (
|
||||
GetAssetResponse,
|
||||
Image2VideoTaskCreationRequest,
|
||||
ImageTaskCreationResponse,
|
||||
SeedAudioConfig,
|
||||
SeedAudioReference,
|
||||
SeedAudioRequest,
|
||||
SeedAudioResponse,
|
||||
Seedance2TaskCreationRequest,
|
||||
SeedanceCreateAssetRequest,
|
||||
SeedanceCreateAssetResponse,
|
||||
@@ -43,6 +48,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 +58,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,
|
||||
@@ -2474,6 +2484,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 +2805,7 @@ class ByteDanceExtension(ComfyExtension):
|
||||
ByteDance2ReferenceNode,
|
||||
ByteDanceCreateImageAsset,
|
||||
ByteDanceCreateVideoAsset,
|
||||
ByteDanceSeedAudioNode,
|
||||
]
|
||||
|
||||
|
||||
|
||||
@@ -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,
|
||||
]
|
||||
|
||||
@@ -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()
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
+63
-2
@@ -264,6 +264,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 +326,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 +340,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)
|
||||
|
||||
|
||||
|
||||
+341
@@ -188,6 +188,49 @@ components:
|
||||
- id
|
||||
- updated_at
|
||||
type: object
|
||||
AvailabilityStatusRequest:
|
||||
description: |
|
||||
Models to query — each entry is `model_id → URL`. The URL lets
|
||||
the server compute file_size + is_hf_downloadable on the same
|
||||
request, eliminating the need for a separate metadata endpoint.
|
||||
properties:
|
||||
models:
|
||||
additionalProperties:
|
||||
type: string
|
||||
description: model_id → URL declared in the workflow.
|
||||
type: object
|
||||
required:
|
||||
- models
|
||||
type: object
|
||||
AvailabilityStatusResponse:
|
||||
description: Per-model state + metadata + HF auth snapshot.
|
||||
properties:
|
||||
hf_auth:
|
||||
$ref: '#/components/schemas/HfAuthStatus'
|
||||
models:
|
||||
additionalProperties:
|
||||
$ref: '#/components/schemas/ModelStatusEntry'
|
||||
type: object
|
||||
required:
|
||||
- models
|
||||
- hf_auth
|
||||
type: object
|
||||
CancelDownloadSessionRequest:
|
||||
description: Request to cancel an in-flight download for a given model_id.
|
||||
properties:
|
||||
model_id:
|
||||
type: string
|
||||
required:
|
||||
- model_id
|
||||
type: object
|
||||
CancelDownloadSessionResponse:
|
||||
description: Result of a cancellation request.
|
||||
properties:
|
||||
cancelled:
|
||||
type: boolean
|
||||
required:
|
||||
- cancelled
|
||||
type: object
|
||||
CreateWorkflowRequest:
|
||||
description: Request body for creating a new saved workflow.
|
||||
properties:
|
||||
@@ -230,6 +273,51 @@ components:
|
||||
- base_version
|
||||
- workflow_json
|
||||
type: object
|
||||
DownloadModelsRequest:
|
||||
description: Map of model_id → URL of files to fetch into the model folders.
|
||||
properties:
|
||||
models:
|
||||
additionalProperties:
|
||||
type: string
|
||||
description: model_id → URL of models to download.
|
||||
type: object
|
||||
required:
|
||||
- models
|
||||
type: object
|
||||
DownloadModelsResponse:
|
||||
description: Acknowledgement that downloads have been scheduled.
|
||||
properties:
|
||||
accepted:
|
||||
description: Always true; the request was scheduled.
|
||||
type: boolean
|
||||
scheduled:
|
||||
description: The list of model_ids whose downloads are now in-flight.
|
||||
items:
|
||||
type: string
|
||||
type: array
|
||||
required:
|
||||
- accepted
|
||||
- scheduled
|
||||
type: object
|
||||
DownloadProgress:
|
||||
description: In-flight download progress; embedded in ModelStatusEntry.
|
||||
properties:
|
||||
bytes_downloaded:
|
||||
format: int64
|
||||
type: integer
|
||||
progress:
|
||||
description: Fraction in [0,1]; null until total_bytes is known.
|
||||
format: float
|
||||
nullable: true
|
||||
type: number
|
||||
total_bytes:
|
||||
description: Content-Length when supplied by the source.
|
||||
format: int64
|
||||
nullable: true
|
||||
type: integer
|
||||
required:
|
||||
- bytes_downloaded
|
||||
type: object
|
||||
ErrorResponse:
|
||||
description: Standard error response with a machine-readable code and human-readable message.
|
||||
properties:
|
||||
@@ -394,6 +482,46 @@ components:
|
||||
- name
|
||||
- info
|
||||
type: object
|
||||
HfAuthLoginStartResponse:
|
||||
description: URL the frontend should open in a new tab to complete login.
|
||||
properties:
|
||||
authorize_url:
|
||||
type: string
|
||||
required:
|
||||
- authorize_url
|
||||
type: object
|
||||
HfAuthLogoutResponse:
|
||||
description: Result of the logout call (always logged_out = true).
|
||||
properties:
|
||||
logged_out:
|
||||
type: boolean
|
||||
required:
|
||||
- logged_out
|
||||
type: object
|
||||
HfAuthStatus:
|
||||
description: Inline snapshot of the server's HuggingFace OAuth state.
|
||||
properties:
|
||||
eligible:
|
||||
description: True iff this deployment can surface interactive HF login.
|
||||
type: boolean
|
||||
token_available:
|
||||
description: True iff a token (possibly expired but refreshable) is stored.
|
||||
type: boolean
|
||||
required:
|
||||
- token_available
|
||||
- eligible
|
||||
type: object
|
||||
HfAuthTokenStatusResponse:
|
||||
description: Whether the server holds an HF OAuth token + resolved username.
|
||||
properties:
|
||||
token_available:
|
||||
type: boolean
|
||||
username:
|
||||
nullable: true
|
||||
type: string
|
||||
required:
|
||||
- token_available
|
||||
type: object
|
||||
HistoryDetailEntry:
|
||||
description: History entry with full prompt data
|
||||
properties:
|
||||
@@ -798,6 +926,41 @@ components:
|
||||
- name
|
||||
- folders
|
||||
type: object
|
||||
ModelStatusEntry:
|
||||
description: Everything the UI needs to render one row of the model.
|
||||
properties:
|
||||
file_size:
|
||||
description: Bytes, when known. Cached server-side per URL.
|
||||
format: int64
|
||||
nullable: true
|
||||
type: integer
|
||||
is_hf_downloadable:
|
||||
description: |
|
||||
HuggingFace-only signal. True if the server can fetch this
|
||||
URL with its current auth state (public, or gated-with-access).
|
||||
False if gated and lacking access. Null for non-HF URLs and
|
||||
for HF URLs whose probe failed entirely.
|
||||
nullable: true
|
||||
type: boolean
|
||||
progress:
|
||||
type: object
|
||||
allOf:
|
||||
- $ref: '#/components/schemas/DownloadProgress'
|
||||
description: Present when `state == downloading`.
|
||||
nullable: true
|
||||
state:
|
||||
description: |
|
||||
`available` — file is on disk.
|
||||
`missing` — not on disk and no download in flight.
|
||||
`downloading` — server is currently fetching the file.
|
||||
enum:
|
||||
- available
|
||||
- missing
|
||||
- downloading
|
||||
type: string
|
||||
required:
|
||||
- state
|
||||
type: object
|
||||
NodeInfo:
|
||||
description: Metadata describing a single ComfyUI node type and its inputs/outputs.
|
||||
properties:
|
||||
@@ -2350,6 +2513,72 @@ paths:
|
||||
summary: Get tag histogram for filtered assets
|
||||
tags:
|
||||
- file
|
||||
/api/cancel-model-download-session:
|
||||
post:
|
||||
description: |
|
||||
Cancels the download session for the given model_id. The worker
|
||||
observes the cancellation between chunks, removes its partial `.tmp`
|
||||
file, and exits without writing the destination path.
|
||||
operationId: postCancelModelDownloadSession
|
||||
requestBody:
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/CancelDownloadSessionRequest'
|
||||
required: true
|
||||
responses:
|
||||
"200":
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/CancelDownloadSessionResponse'
|
||||
description: Session cancelled.
|
||||
"404":
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/ErrorResponse'
|
||||
description: No active download for that model_id.
|
||||
summary: Cancel an in-flight server-side model download
|
||||
tags:
|
||||
- model
|
||||
/api/download-models:
|
||||
post:
|
||||
description: |
|
||||
Schedules downloads for every model_id in the request map. Returns
|
||||
immediately after validation; progress is observed via
|
||||
`/api/models-availability-status`. Fails atomically if any model
|
||||
is already on disk, already downloading, gated, or has a URL that
|
||||
is not on the server's allowlist (HuggingFace, Civitai, localhost).
|
||||
operationId: postDownloadModels
|
||||
requestBody:
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/DownloadModelsRequest'
|
||||
required: true
|
||||
responses:
|
||||
"202":
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/DownloadModelsResponse'
|
||||
description: Downloads accepted and scheduled.
|
||||
"400":
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/ErrorResponse'
|
||||
description: One of the requested models is invalid, gated, or has a non-allowed URL.
|
||||
"409":
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/ErrorResponse'
|
||||
description: One of the requested models is already on disk or downloading.
|
||||
summary: Start a server-side download of one or more models
|
||||
tags:
|
||||
- model
|
||||
/api/embeddings:
|
||||
get:
|
||||
description: Returns the list of text-encoder embeddings available on disk.
|
||||
@@ -2655,6 +2884,74 @@ paths:
|
||||
summary: Get a specific subgraph blueprint
|
||||
tags:
|
||||
- workflow
|
||||
/api/hf-auth-login-start:
|
||||
post:
|
||||
description: |
|
||||
Spawns a short-lived loopback callback server (port 41954) and
|
||||
returns the URL the frontend should open in a new tab. After the
|
||||
user grants consent, HF redirects back to the callback URL with
|
||||
an authorization code; the server exchanges that for a token and
|
||||
persists it. Subsequent `/api/hf-auth-token-status` calls will
|
||||
return `token_available: true`. Rejected with 403 if the
|
||||
deployment is not eligible (not loopback or in --multi-user mode);
|
||||
409 if another login attempt is already in progress.
|
||||
operationId: postHfAuthLoginStart
|
||||
responses:
|
||||
"200":
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/HfAuthLoginStartResponse'
|
||||
description: Login flow started; `authorize_url` is ready.
|
||||
"403":
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/ErrorResponse'
|
||||
description: Deployment is not eligible for interactive HF login.
|
||||
"409":
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/ErrorResponse'
|
||||
description: Another login attempt is already in progress.
|
||||
summary: Begin a HuggingFace OAuth login flow
|
||||
tags:
|
||||
- model
|
||||
/api/hf-auth-logout:
|
||||
post:
|
||||
description: Clears the in-memory cache and removes the on-disk token file.
|
||||
operationId: postHfAuthLogout
|
||||
responses:
|
||||
"200":
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/HfAuthLogoutResponse'
|
||||
description: Logged out (idempotent — succeeds even if no token was held).
|
||||
summary: Drop the stored HuggingFace OAuth token
|
||||
tags:
|
||||
- model
|
||||
/api/hf-auth-token-status:
|
||||
get:
|
||||
description: |
|
||||
Returns `token_available: true` when the server has a token
|
||||
in memory (or on disk) for HuggingFace, irrespective of whether
|
||||
the access_token is currently fresh — an expired one with a
|
||||
refresh_token still counts as "logged in" because we'll refresh
|
||||
transparently on next use. If a username is resolvable via
|
||||
`HfApi.whoami` we return that too, for the settings UI.
|
||||
operationId: getHfAuthTokenStatus
|
||||
responses:
|
||||
"200":
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/HfAuthTokenStatusResponse'
|
||||
description: Token status.
|
||||
summary: Whether the server holds a usable HuggingFace OAuth token
|
||||
tags:
|
||||
- model
|
||||
/api/history:
|
||||
post:
|
||||
deprecated: true
|
||||
@@ -3157,6 +3454,48 @@ paths:
|
||||
summary: Cancel multiple jobs
|
||||
tags:
|
||||
- workflow
|
||||
/api/models-availability-status:
|
||||
post:
|
||||
description: |
|
||||
Given a map of `{model_id: url}` (model_id is
|
||||
`<directory>/<filename>`), returns per-id state plus the
|
||||
metadata the UI needs to render the row:
|
||||
|
||||
- `state` — one of `available` / `missing` / `downloading`
|
||||
- `progress` — embedded when `state == downloading`
|
||||
- `file_size` — bytes (when known)
|
||||
- `is_hf_downloadable` — for HF URLs only: true if the
|
||||
server can currently fetch the file with its stored auth
|
||||
state, false if gated and lacking access, null otherwise
|
||||
|
||||
Designed for 1 Hz polling. `file_size` and the intrinsic
|
||||
"is this model gated" check are cached server-side per URL;
|
||||
`is_hf_downloadable` is recomputed per call so license
|
||||
acceptance and login/logout transitions show up within one
|
||||
poll interval without any client-side cache plumbing.
|
||||
operationId: postModelsAvailabilityStatus
|
||||
requestBody:
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/AvailabilityStatusRequest'
|
||||
required: true
|
||||
responses:
|
||||
"200":
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/AvailabilityStatusResponse'
|
||||
description: Per-model status and metadata.
|
||||
"400":
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/ErrorResponse'
|
||||
description: Malformed request body.
|
||||
summary: Unified per-model status + metadata for the polling UI
|
||||
tags:
|
||||
- model
|
||||
/api/node_replacements:
|
||||
get:
|
||||
description: |
|
||||
@@ -5103,3 +5442,5 @@ tags:
|
||||
name: queue
|
||||
- description: Job lifecycle queries
|
||||
name: job
|
||||
- description: Server-side model availability and downloads
|
||||
name: model
|
||||
|
||||
+2
-1
@@ -1,5 +1,5 @@
|
||||
comfyui-frontend-package==1.45.20
|
||||
comfyui-workflow-templates==0.11.1
|
||||
comfyui-workflow-templates==0.11.2
|
||||
comfyui-embedded-docs==0.5.6
|
||||
torch
|
||||
torchsde
|
||||
@@ -9,6 +9,7 @@ numpy>=1.25.0
|
||||
einops
|
||||
transformers>=4.50.3
|
||||
tokenizers>=0.13.3
|
||||
huggingface_hub
|
||||
sentencepiece
|
||||
safetensors>=0.4.2
|
||||
aiohttp>=3.11.8
|
||||
|
||||
@@ -47,6 +47,7 @@ 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.asset_management import resolve_hash_to_path
|
||||
from app.model_downloader.api.routes import register_routes as register_model_downloader_routes
|
||||
|
||||
from app.user_manager import UserManager
|
||||
from app.model_manager import ModelFileManager
|
||||
@@ -127,6 +128,7 @@ def create_cors_middleware(allowed_origin: str):
|
||||
|
||||
return cors_middleware
|
||||
|
||||
|
||||
def is_loopback(host):
|
||||
if host is None:
|
||||
return False
|
||||
@@ -256,6 +258,7 @@ class PromptServer():
|
||||
else:
|
||||
register_assets_routes(self.app)
|
||||
asset_seeder.disable()
|
||||
register_model_downloader_routes(self.app)
|
||||
routes = web.RouteTableDef()
|
||||
self.routes = routes
|
||||
self.last_node_id = None
|
||||
@@ -616,15 +619,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 <script>
|
||||
# and execute in the page origin) to download instead of
|
||||
# displaying inline, preventing stored XSS. The
|
||||
# attachment disposition is the load-bearing guard: a
|
||||
# bare filename= hint does not force a download per
|
||||
# RFC 6266, so we only attach it on the dangerous branch
|
||||
# to avoid breaking inline display of legitimate images.
|
||||
# Escape backslash/quote per RFC 6266 quoted-string so a
|
||||
# filename containing a double quote (which passes the
|
||||
# ".."/leading-slash filter above) can't break out of the
|
||||
# header's quoted-string and malform the disposition.
|
||||
safe_filename = filename.replace("\\", "\\\\").replace('"', '\\"')
|
||||
disposition = f"filename=\"{safe_filename}\""
|
||||
if folder_paths.is_dangerous_content_type(content_type):
|
||||
content_type = 'application/octet-stream'
|
||||
disposition = f"attachment; filename=\"{safe_filename}\""
|
||||
|
||||
return web.FileResponse(
|
||||
file,
|
||||
headers={
|
||||
"Content-Disposition": f"filename=\"{filename}\"",
|
||||
"Content-Type": content_type
|
||||
"Content-Disposition": disposition,
|
||||
"Content-Type": content_type,
|
||||
"X-Content-Type-Options": "nosniff"
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@@ -0,0 +1,708 @@
|
||||
"""Unit tests for the HuggingFace auth subsystem.
|
||||
|
||||
Covers:
|
||||
- token store: save/load roundtrip, chmod 0600, atomic write, delete
|
||||
- eligibility under various CLI-arg combinations
|
||||
- URL parsing (huggingface.co host detection + repo_id extraction)
|
||||
- HF-aware gated_detection.probe_url (mocked auth_check)
|
||||
- HF auth routes (token status, login start with eligibility gate, logout)
|
||||
- PKCE primitives + authorize URL shape
|
||||
|
||||
The OAuth callback server itself isn't exercised end-to-end here — that
|
||||
requires a real HF server. We test the components (state checking,
|
||||
URL building, code-exchange request shape) instead.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import stat
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from aiohttp import web
|
||||
|
||||
from app.model_downloader.api.routes import register_routes
|
||||
from app.model_downloader.hf_auth import oauth
|
||||
from app.model_downloader.hf_auth.auth_store import HF_AUTH_STORE, HfAuthStore
|
||||
from app.model_downloader.hf_auth.token_store import (
|
||||
EXPIRY_BUFFER_SECS,
|
||||
Token,
|
||||
delete_token,
|
||||
load_token,
|
||||
save_token,
|
||||
)
|
||||
from app.model_downloader.hf_url import is_hf_url, repo_id_from_url
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Fixtures
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def patched_user_dir(tmp_path):
|
||||
"""Redirect ``folder_paths.get_user_directory`` so the token file
|
||||
lands in an isolated tmp_path instead of the real user dir."""
|
||||
user_dir = tmp_path / "user"
|
||||
user_dir.mkdir()
|
||||
with patch("folder_paths.get_user_directory", return_value=str(user_dir)):
|
||||
yield user_dir
|
||||
|
||||
|
||||
def _token_file_path(user_dir) -> str:
|
||||
return os.path.join(user_dir, "__hf_auth", "hf_auth_token.json")
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def fresh_auth_store():
|
||||
"""Wipe singleton state between tests: auth + probe caches."""
|
||||
from app.model_downloader import gated_detection
|
||||
|
||||
HF_AUTH_STORE._token = None
|
||||
HF_AUTH_STORE._loaded_from_disk = False
|
||||
gated_detection.clear_caches_for_tests()
|
||||
yield HF_AUTH_STORE
|
||||
HF_AUTH_STORE._token = None
|
||||
HF_AUTH_STORE._loaded_from_disk = False
|
||||
gated_detection.clear_caches_for_tests()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def app(patched_user_dir, fresh_auth_store):
|
||||
app = web.Application()
|
||||
register_routes(app)
|
||||
return app
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# URL parsing
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
|
||||
def test_is_hf_url_recognises_huggingface_co():
|
||||
assert is_hf_url("https://huggingface.co/x/y/resolve/main/z.safetensors")
|
||||
assert is_hf_url("https://huggingface.co/abc")
|
||||
assert not is_hf_url("https://hf-mirror.com/x/y/resolve/main/z.safetensors")
|
||||
assert not is_hf_url("https://civitai.com/x.safetensors")
|
||||
|
||||
|
||||
def test_repo_id_from_url_extracts_org_and_repo():
|
||||
url = "https://huggingface.co/Lightricks/LTX-2.3-22b-IC-LoRA-HDR/resolve/main/x.safetensors"
|
||||
assert repo_id_from_url(url) == "Lightricks/LTX-2.3-22b-IC-LoRA-HDR"
|
||||
|
||||
|
||||
def test_repo_id_from_url_handles_nested_path():
|
||||
url = "https://huggingface.co/Comfy-Org/ltx-2.3/resolve/main/split_files/loras/x.safetensors"
|
||||
assert repo_id_from_url(url) == "Comfy-Org/ltx-2.3"
|
||||
|
||||
|
||||
def test_repo_id_from_url_returns_none_for_non_hf():
|
||||
assert repo_id_from_url("https://civitai.com/x.safetensors") is None
|
||||
|
||||
|
||||
def test_repo_id_from_url_returns_none_for_non_resolve_paths():
|
||||
assert repo_id_from_url("https://huggingface.co/org/repo/blob/main/x.safetensors") is None
|
||||
assert repo_id_from_url("https://huggingface.co/org") is None
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Token store
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
|
||||
def test_token_store_roundtrip(patched_user_dir):
|
||||
tok = Token(
|
||||
access_token="hf_abc",
|
||||
refresh_token="rf_def",
|
||||
expires_at=9999999999.0,
|
||||
scope="openid profile",
|
||||
)
|
||||
save_token(tok)
|
||||
loaded = load_token()
|
||||
assert loaded == tok
|
||||
|
||||
|
||||
def test_token_store_writes_0600(patched_user_dir):
|
||||
tok = Token(access_token="x", refresh_token=None, expires_at=0.0)
|
||||
save_token(tok)
|
||||
path = _token_file_path(patched_user_dir)
|
||||
mode = stat.S_IMODE(os.stat(path).st_mode)
|
||||
# On Windows we silently no-op chmod; allow either the intended
|
||||
# mode or whatever umask the OS gave us.
|
||||
if os.name == "posix":
|
||||
assert mode == 0o600
|
||||
|
||||
|
||||
def test_token_store_delete_removes_file(patched_user_dir):
|
||||
tok = Token(access_token="x", refresh_token=None, expires_at=0.0)
|
||||
save_token(tok)
|
||||
delete_token()
|
||||
path = _token_file_path(patched_user_dir)
|
||||
assert not os.path.exists(path)
|
||||
# Idempotent: second delete is fine.
|
||||
delete_token()
|
||||
|
||||
|
||||
def test_token_store_load_returns_none_for_missing_file(patched_user_dir):
|
||||
assert load_token() is None
|
||||
|
||||
|
||||
def test_token_store_load_returns_none_for_corrupt_file(patched_user_dir):
|
||||
path = _token_file_path(patched_user_dir)
|
||||
os.makedirs(os.path.dirname(path), exist_ok=True)
|
||||
with open(path, "w") as f:
|
||||
f.write("not json {")
|
||||
assert load_token() is None
|
||||
|
||||
|
||||
def test_token_is_valid_uses_buffer(patched_user_dir):
|
||||
import time
|
||||
|
||||
fresh = Token(access_token="x", refresh_token=None, expires_at=time.time() + 3600)
|
||||
nearly_expired = Token(
|
||||
access_token="x",
|
||||
refresh_token=None,
|
||||
expires_at=time.time() + EXPIRY_BUFFER_SECS - 1,
|
||||
)
|
||||
assert fresh.is_valid()
|
||||
assert not nearly_expired.is_valid()
|
||||
|
||||
|
||||
def test_token_is_valid_rejects_empty_access_token():
|
||||
import time
|
||||
|
||||
tok = Token(access_token="", refresh_token=None, expires_at=time.time() + 3600)
|
||||
assert not tok.is_valid()
|
||||
|
||||
|
||||
def test_token_is_valid_rejects_at_exact_buffer_boundary():
|
||||
import time
|
||||
|
||||
tok = Token(
|
||||
access_token="x",
|
||||
refresh_token=None,
|
||||
expires_at=time.time() + EXPIRY_BUFFER_SECS,
|
||||
)
|
||||
assert not tok.is_valid()
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Auth store
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
|
||||
def test_auth_store_loads_lazily(patched_user_dir):
|
||||
tok = Token(access_token="x", refresh_token=None, expires_at=9999999999.0)
|
||||
save_token(tok)
|
||||
store = HfAuthStore()
|
||||
assert store.has_token()
|
||||
assert store.get_token_sync() == tok
|
||||
|
||||
|
||||
def test_auth_store_set_persists(patched_user_dir):
|
||||
store = HfAuthStore()
|
||||
tok = Token(access_token="x", refresh_token=None, expires_at=9999999999.0)
|
||||
store.set_token(tok)
|
||||
# Token is on disk now — a fresh store sees it.
|
||||
assert HfAuthStore().get_token_sync() == tok
|
||||
|
||||
|
||||
def test_auth_store_clear_removes_in_memory_and_on_disk(patched_user_dir):
|
||||
store = HfAuthStore()
|
||||
tok = Token(access_token="x", refresh_token=None, expires_at=9999999999.0)
|
||||
store.set_token(tok)
|
||||
store.clear()
|
||||
assert not store.has_token()
|
||||
assert HfAuthStore().get_token_sync() is None
|
||||
|
||||
|
||||
def test_auth_store_has_token_true_when_expired_but_refreshable(patched_user_dir):
|
||||
import time
|
||||
|
||||
store = HfAuthStore()
|
||||
expired = Token(
|
||||
access_token="old",
|
||||
refresh_token="rf",
|
||||
expires_at=time.time() - 100,
|
||||
)
|
||||
store.set_token(expired)
|
||||
assert store.has_token()
|
||||
assert not expired.is_valid()
|
||||
|
||||
|
||||
def test_auth_store_get_token_sync_returns_expired_without_refresh(patched_user_dir):
|
||||
import time
|
||||
|
||||
store = HfAuthStore()
|
||||
expired = Token(
|
||||
access_token="old",
|
||||
refresh_token=None,
|
||||
expires_at=time.time() - 100,
|
||||
)
|
||||
store.set_token(expired)
|
||||
assert store.get_token_sync() == expired
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_auth_store_get_valid_returns_none_when_expired_without_refresh(
|
||||
patched_user_dir,
|
||||
):
|
||||
import time
|
||||
|
||||
store = HfAuthStore()
|
||||
expired = Token(
|
||||
access_token="old",
|
||||
refresh_token=None,
|
||||
expires_at=time.time() - 100,
|
||||
)
|
||||
store.set_token(expired)
|
||||
with patch(
|
||||
"app.model_downloader.hf_auth.oauth.refresh_access_token",
|
||||
new=AsyncMock(),
|
||||
) as refresh_mock:
|
||||
result = await store.get_valid_token()
|
||||
assert result is None
|
||||
refresh_mock.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_auth_store_get_valid_returns_fresh_token(patched_user_dir):
|
||||
store = HfAuthStore()
|
||||
import time
|
||||
|
||||
tok = Token(access_token="x", refresh_token=None, expires_at=time.time() + 3600)
|
||||
store.set_token(tok)
|
||||
fetched = await store.get_valid_token()
|
||||
assert fetched == tok
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_auth_store_get_valid_refresh_on_expired(patched_user_dir):
|
||||
store = HfAuthStore()
|
||||
import time
|
||||
|
||||
expired = Token(
|
||||
access_token="old",
|
||||
refresh_token="rf",
|
||||
expires_at=time.time() - 100,
|
||||
)
|
||||
store.set_token(expired)
|
||||
refreshed = Token(
|
||||
access_token="new",
|
||||
refresh_token="rf",
|
||||
expires_at=time.time() + 3600,
|
||||
)
|
||||
with patch(
|
||||
"app.model_downloader.hf_auth.oauth.refresh_access_token",
|
||||
new=AsyncMock(return_value=refreshed),
|
||||
):
|
||||
result = await store.get_valid_token()
|
||||
assert result == refreshed
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_auth_store_get_valid_token_does_not_resurrect_after_logout(
|
||||
patched_user_dir,
|
||||
):
|
||||
"""A logout landing *during* an in-flight refresh must not be undone by
|
||||
the refresh writing the token back (the resurrection race)."""
|
||||
store = HfAuthStore()
|
||||
import time
|
||||
|
||||
expired = Token(
|
||||
access_token="old", refresh_token="rf", expires_at=time.time() - 100
|
||||
)
|
||||
store.set_token(expired)
|
||||
refreshed = Token(
|
||||
access_token="new", refresh_token="rf", expires_at=time.time() + 3600
|
||||
)
|
||||
|
||||
async def fake_refresh(_refresh_token):
|
||||
# Simulate the user clicking "Log out" while the refresh is in flight.
|
||||
store.clear()
|
||||
return refreshed
|
||||
|
||||
with patch(
|
||||
"app.model_downloader.hf_auth.oauth.refresh_access_token",
|
||||
new=fake_refresh,
|
||||
):
|
||||
result = await store.get_valid_token()
|
||||
|
||||
# The refresh result is discarded — logout wins, in memory and on disk.
|
||||
assert result is None
|
||||
assert not store.has_token()
|
||||
assert load_token() is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_auth_store_get_valid_returns_none_on_refresh_failure(patched_user_dir):
|
||||
store = HfAuthStore()
|
||||
import time
|
||||
|
||||
expired = Token(
|
||||
access_token="old",
|
||||
refresh_token="rf",
|
||||
expires_at=time.time() - 100,
|
||||
)
|
||||
store.set_token(expired)
|
||||
with patch(
|
||||
"app.model_downloader.hf_auth.oauth.refresh_access_token",
|
||||
new=AsyncMock(side_effect=RuntimeError("HF down")),
|
||||
):
|
||||
result = await store.get_valid_token()
|
||||
assert result is None
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Eligibility
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"listen,multi_user,expected",
|
||||
[
|
||||
("127.0.0.1", False, True),
|
||||
("127.0.0.1", True, False), # multi-user disables it
|
||||
("0.0.0.0", False, False), # bind-all is not loopback
|
||||
("0.0.0.0", True, False),
|
||||
("192.168.1.5", False, False), # LAN address
|
||||
("::1", False, True), # IPv6 loopback
|
||||
],
|
||||
)
|
||||
def test_eligibility(listen, multi_user, expected, monkeypatch):
|
||||
from app.model_downloader.hf_auth import eligibility
|
||||
from comfy.cli_args import args
|
||||
|
||||
monkeypatch.setattr(args, "listen", listen)
|
||||
monkeypatch.setattr(args, "multi_user", multi_user)
|
||||
assert eligibility.is_hf_auth_eligible() is expected
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# gated_detection HF probe
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_probe_url_hf_public(fresh_auth_store):
|
||||
"""auth_check succeeds with no token → is_hf_downloadable = True."""
|
||||
from app.model_downloader.gated_detection import probe_url
|
||||
|
||||
url = "https://huggingface.co/public/repo/resolve/main/x.safetensors"
|
||||
with patch("app.model_downloader.gated_detection._auth_check_sync"), patch(
|
||||
"app.model_downloader.gated_detection._probe_size_once",
|
||||
new=AsyncMock(return_value=1024),
|
||||
):
|
||||
result = await probe_url(url)
|
||||
assert result.is_hf_downloadable is True
|
||||
assert result.file_size == 1024
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_probe_url_hf_gated_no_access(fresh_auth_store):
|
||||
"""auth_check raises GatedRepoError → is_hf_downloadable = False."""
|
||||
from huggingface_hub.errors import GatedRepoError
|
||||
|
||||
from app.model_downloader.gated_detection import probe_url
|
||||
|
||||
url = "https://huggingface.co/gated/repo/resolve/main/x.safetensors"
|
||||
fake_response = MagicMock(status_code=403)
|
||||
with patch(
|
||||
"app.model_downloader.gated_detection._auth_check_sync",
|
||||
side_effect=GatedRepoError("gated", response=fake_response),
|
||||
), patch(
|
||||
"app.model_downloader.gated_detection._probe_size_once",
|
||||
new=AsyncMock(return_value=None),
|
||||
):
|
||||
result = await probe_url(url)
|
||||
assert result.is_hf_downloadable is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_probe_url_non_hf_skips_auth_check():
|
||||
"""Non-HF URLs never call auth_check; is_hf_downloadable stays None."""
|
||||
from app.model_downloader.gated_detection import probe_url
|
||||
|
||||
url = "https://civitai.com/api/download/models/1.safetensors"
|
||||
with patch(
|
||||
"app.model_downloader.gated_detection._auth_check_sync",
|
||||
) as mocked, patch(
|
||||
"app.model_downloader.gated_detection._probe_size_once",
|
||||
new=AsyncMock(return_value=2048),
|
||||
):
|
||||
result = await probe_url(url)
|
||||
assert result.is_hf_downloadable is None
|
||||
assert result.file_size == 2048
|
||||
mocked.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_is_gated_cached_across_calls(fresh_auth_store):
|
||||
"""Intrinsic ``is_gated`` should be determined exactly once per URL.
|
||||
|
||||
Subsequent ``probe_url`` calls for the same URL must not re-issue
|
||||
the null-token auth_check — that's the whole point of the cache."""
|
||||
from app.model_downloader.gated_detection import probe_url
|
||||
|
||||
url = "https://huggingface.co/public/repo/resolve/main/x.safetensors"
|
||||
with patch(
|
||||
"app.model_downloader.gated_detection._auth_check_sync"
|
||||
) as mocked, patch(
|
||||
"app.model_downloader.gated_detection._probe_size_once",
|
||||
new=AsyncMock(return_value=1024),
|
||||
):
|
||||
await probe_url(url)
|
||||
await probe_url(url)
|
||||
await probe_url(url)
|
||||
# Three probe_url calls × public-only-needs-1-auth_check = 1 call total.
|
||||
assert mocked.call_count == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_file_size_cached_across_calls(fresh_auth_store):
|
||||
"""Once a successful HEAD lands, subsequent calls don't re-HEAD."""
|
||||
from app.model_downloader.gated_detection import probe_url
|
||||
|
||||
url = "https://huggingface.co/public/repo/resolve/main/x.safetensors"
|
||||
with patch(
|
||||
"app.model_downloader.gated_detection._auth_check_sync"
|
||||
), patch(
|
||||
"app.model_downloader.gated_detection._probe_size_once",
|
||||
new=AsyncMock(return_value=2048),
|
||||
) as size_probe:
|
||||
r1 = await probe_url(url)
|
||||
r2 = await probe_url(url)
|
||||
assert r1.file_size == 2048
|
||||
assert r2.file_size == 2048
|
||||
assert size_probe.call_count == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_file_size_not_probed_for_gated_no_access(fresh_auth_store):
|
||||
"""When ``is_hf_downloadable`` is False we must NOT HEAD the URL —
|
||||
otherwise a 401-due-to-gating would land as a cached ``None`` that
|
||||
survives a later successful login."""
|
||||
from app.model_downloader.gated_detection import probe_url
|
||||
from huggingface_hub.errors import GatedRepoError
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
url = "https://huggingface.co/gated/repo/resolve/main/x.safetensors"
|
||||
fake_resp = MagicMock(status_code=403)
|
||||
with patch(
|
||||
"app.model_downloader.gated_detection._auth_check_sync",
|
||||
side_effect=GatedRepoError("gated", response=fake_resp),
|
||||
), patch(
|
||||
"app.model_downloader.gated_detection._probe_size_once",
|
||||
new=AsyncMock(return_value=None),
|
||||
) as size_probe:
|
||||
result = await probe_url(url)
|
||||
assert result.is_hf_downloadable is False
|
||||
assert result.file_size is None
|
||||
assert size_probe.call_count == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_probe_url_passes_token_when_available(fresh_auth_store, patched_user_dir):
|
||||
"""For a gated URL, auth_check runs twice: once with token=None to
|
||||
determine the intrinsic ``is_gated`` flag (cached forever), and once
|
||||
with the stored access_token to determine ``is_hf_downloadable`` for
|
||||
the current user."""
|
||||
from app.model_downloader import gated_detection
|
||||
from app.model_downloader.gated_detection import probe_url
|
||||
from huggingface_hub.errors import GatedRepoError
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
gated_detection.clear_caches_for_tests()
|
||||
fresh_auth_store.set_token(Token(
|
||||
access_token="hf_test_token",
|
||||
refresh_token=None,
|
||||
expires_at=9999999999.0,
|
||||
))
|
||||
url = "https://huggingface.co/private/repo/resolve/main/x.safetensors"
|
||||
|
||||
fake_resp = MagicMock(status_code=403)
|
||||
|
||||
def fake_auth_check(repo_id, token):
|
||||
# Null-token call → repo is gated. Subsequent call with the real
|
||||
# token succeeds (user has access).
|
||||
if token is None:
|
||||
raise GatedRepoError("gated", response=fake_resp)
|
||||
|
||||
with patch(
|
||||
"app.model_downloader.gated_detection._auth_check_sync",
|
||||
side_effect=fake_auth_check,
|
||||
) as mocked, patch(
|
||||
"app.model_downloader.gated_detection._probe_size_once",
|
||||
new=AsyncMock(return_value=None),
|
||||
):
|
||||
result = await probe_url(url)
|
||||
|
||||
# is_hf_downloadable should be True (token-authed call succeeded).
|
||||
assert result.is_hf_downloadable is True
|
||||
# Two calls: (repo_id, None) then (repo_id, <token>).
|
||||
assert mocked.call_count == 2
|
||||
assert mocked.call_args_list[0].args == ("private/repo", None)
|
||||
assert mocked.call_args_list[1].args == ("private/repo", "hf_test_token")
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# OAuth primitives
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
|
||||
def test_make_pkce_returns_distinct_high_entropy_values():
|
||||
verifier1, challenge1, state1 = oauth._make_pkce()
|
||||
verifier2, challenge2, state2 = oauth._make_pkce()
|
||||
assert verifier1 != verifier2
|
||||
assert challenge1 != challenge2
|
||||
assert state1 != state2
|
||||
# Verifier should be at least 43 chars per PKCE spec.
|
||||
assert len(verifier1) >= 43
|
||||
|
||||
|
||||
def test_build_authorize_url_includes_pkce_and_state():
|
||||
url = oauth._build_authorize_url("challenge123", "state456")
|
||||
assert url.startswith(oauth.AUTHORIZE_URL)
|
||||
assert "client_id=" + oauth.HF_CLIENT_ID in url
|
||||
assert "code_challenge=challenge123" in url
|
||||
assert "code_challenge_method=S256" in url
|
||||
assert "state=state456" in url
|
||||
assert "response_type=code" in url
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Routes
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_hf_auth_token_status_empty(aiohttp_client, app):
|
||||
"""No token set → token_available=false, username=null."""
|
||||
client = await aiohttp_client(app)
|
||||
resp = await client.get("/api/hf-auth-token-status")
|
||||
assert resp.status == 200
|
||||
data = await resp.json()
|
||||
assert data == {"token_available": False, "username": None}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_hf_auth_token_status_with_token(
|
||||
aiohttp_client, app, fresh_auth_store, patched_user_dir
|
||||
):
|
||||
"""Token present, whoami works → username is returned."""
|
||||
fresh_auth_store.set_token(Token(
|
||||
access_token="x", refresh_token=None, expires_at=9999999999.0,
|
||||
))
|
||||
with patch(
|
||||
"app.model_downloader.api.routes._whoami_username",
|
||||
return_value="alice",
|
||||
):
|
||||
client = await aiohttp_client(app)
|
||||
resp = await client.get("/api/hf-auth-token-status")
|
||||
assert resp.status == 200
|
||||
assert (await resp.json()) == {"token_available": True, "username": "alice"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_hf_auth_login_start_403_when_ineligible(aiohttp_client, app, monkeypatch):
|
||||
"""Not loopback / multi-user → 403."""
|
||||
monkeypatch.setattr(
|
||||
"app.model_downloader.api.routes.is_hf_auth_eligible",
|
||||
lambda: False,
|
||||
)
|
||||
client = await aiohttp_client(app)
|
||||
resp = await client.post("/api/hf-auth-login-start")
|
||||
assert resp.status == 403
|
||||
assert (await resp.json())["error"]["code"] == "HF_AUTH_NOT_ELIGIBLE"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_hf_auth_login_start_returns_authorize_url(aiohttp_client, app, monkeypatch):
|
||||
"""Eligible + first attempt → 200 with authorize_url."""
|
||||
monkeypatch.setattr(
|
||||
"app.model_downloader.api.routes.is_hf_auth_eligible",
|
||||
lambda: True,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"app.model_downloader.api.routes.start_login_flow",
|
||||
AsyncMock(return_value="https://huggingface.co/oauth/authorize?fake=1"),
|
||||
)
|
||||
client = await aiohttp_client(app)
|
||||
resp = await client.post("/api/hf-auth-login-start")
|
||||
assert resp.status == 200
|
||||
assert (await resp.json())["authorize_url"].startswith(
|
||||
"https://huggingface.co/oauth/authorize"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_hf_auth_login_start_409_when_in_progress(aiohttp_client, app, monkeypatch):
|
||||
"""Lock already held → 409."""
|
||||
from app.model_downloader.hf_auth.oauth import OAuthInProgressError
|
||||
|
||||
monkeypatch.setattr(
|
||||
"app.model_downloader.api.routes.is_hf_auth_eligible",
|
||||
lambda: True,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"app.model_downloader.api.routes.start_login_flow",
|
||||
AsyncMock(side_effect=OAuthInProgressError()),
|
||||
)
|
||||
client = await aiohttp_client(app)
|
||||
resp = await client.post("/api/hf-auth-login-start")
|
||||
assert resp.status == 409
|
||||
assert (await resp.json())["error"]["code"] == "HF_AUTH_IN_PROGRESS"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_hf_auth_login_start_503_when_callback_bind_fails(
|
||||
aiohttp_client, app, monkeypatch
|
||||
):
|
||||
"""Callback server failed to bind (e.g. port busy) → 503, not a dead URL."""
|
||||
from app.model_downloader.hf_auth.oauth import OAuthCallbackError
|
||||
|
||||
monkeypatch.setattr(
|
||||
"app.model_downloader.api.routes.is_hf_auth_eligible",
|
||||
lambda: True,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"app.model_downloader.api.routes.start_login_flow",
|
||||
AsyncMock(side_effect=OAuthCallbackError("could not bind callback port")),
|
||||
)
|
||||
client = await aiohttp_client(app)
|
||||
resp = await client.post("/api/hf-auth-login-start")
|
||||
assert resp.status == 503
|
||||
assert (await resp.json())["error"]["code"] == "HF_AUTH_START_FAILED"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_hf_auth_logout_clears_store(
|
||||
aiohttp_client, app, fresh_auth_store, patched_user_dir
|
||||
):
|
||||
fresh_auth_store.set_token(Token(
|
||||
access_token="x", refresh_token=None, expires_at=9999999999.0,
|
||||
))
|
||||
client = await aiohttp_client(app)
|
||||
resp = await client.post("/api/hf-auth-logout")
|
||||
assert resp.status == 200
|
||||
assert (await resp.json()) == {"logged_out": True}
|
||||
assert not fresh_auth_store.has_token()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_availability_includes_hf_auth_snapshot(aiohttp_client, app, monkeypatch):
|
||||
"""The availability response embeds {token_available, eligible}."""
|
||||
monkeypatch.setattr(
|
||||
"app.model_downloader.api.routes.is_hf_auth_eligible",
|
||||
lambda: True,
|
||||
)
|
||||
client = await aiohttp_client(app)
|
||||
resp = await client.post(
|
||||
"/api/models-availability-status",
|
||||
json={"models": {}},
|
||||
)
|
||||
assert resp.status == 200
|
||||
data = await resp.json()
|
||||
assert "hf_auth" in data
|
||||
assert data["hf_auth"] == {"token_available": False, "eligible": True}
|
||||
@@ -0,0 +1,514 @@
|
||||
"""Unit tests for the server-side model download subsystem.
|
||||
|
||||
Covers the pieces that don't require talking to a real network:
|
||||
|
||||
- path parsing & allowlist (pure functions)
|
||||
- DownloadServer registry lifecycle (in-memory state)
|
||||
- API routes via aiohttp_client + folder_paths/probe_url patches
|
||||
|
||||
Streaming downloads themselves are exercised indirectly — the route-level
|
||||
tests stub out the network probe so we can verify the gating logic in
|
||||
``download_models`` without making real HTTP calls.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from unittest.mock import patch, AsyncMock
|
||||
|
||||
import pytest
|
||||
from aiohttp import web
|
||||
|
||||
from app.model_downloader.allowlist import is_url_allowed
|
||||
from app.model_downloader.api.routes import register_routes
|
||||
from app.model_downloader.download_server import DownloadServer
|
||||
from app.model_downloader.gated_detection import MetadataProbeResult
|
||||
from app.model_downloader.paths import (
|
||||
InvalidModelId,
|
||||
parse_model_id,
|
||||
resolve_destination,
|
||||
resolve_existing,
|
||||
)
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Fixtures
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def model_root(tmp_path):
|
||||
"""A fake ``models/`` root with two registered folder types."""
|
||||
loras_dir = tmp_path / "loras"
|
||||
checkpoints_dir = tmp_path / "checkpoints"
|
||||
loras_dir.mkdir()
|
||||
checkpoints_dir.mkdir()
|
||||
return tmp_path, loras_dir, checkpoints_dir
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def patched_folder_paths(model_root):
|
||||
"""Point folder_paths at our fake roots for the duration of one test."""
|
||||
_root, loras_dir, checkpoints_dir = model_root
|
||||
mapping = {
|
||||
"loras": ([str(loras_dir)], {".safetensors"}),
|
||||
"checkpoints": ([str(checkpoints_dir)], {".safetensors"}),
|
||||
}
|
||||
with patch(
|
||||
"folder_paths.folder_names_and_paths", mapping
|
||||
), patch(
|
||||
"folder_paths.get_folder_paths",
|
||||
side_effect=lambda name: mapping.get(name, ([], set()))[0],
|
||||
):
|
||||
yield mapping
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def fresh_download_server():
|
||||
"""Reset the module-level singleton between tests so registry state
|
||||
doesn't leak across tests sharing the singleton."""
|
||||
from app.model_downloader.download_server import DOWNLOAD_SERVER
|
||||
|
||||
DOWNLOAD_SERVER.reset_for_tests()
|
||||
yield DOWNLOAD_SERVER
|
||||
DOWNLOAD_SERVER.reset_for_tests()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def app(patched_folder_paths, fresh_download_server):
|
||||
app = web.Application()
|
||||
register_routes(app)
|
||||
return app
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Pure helpers: allowlist + path parsing
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
|
||||
def test_allowlist_accepts_hf_safetensors():
|
||||
assert is_url_allowed("https://huggingface.co/x/y/resolve/main/z.safetensors")
|
||||
|
||||
|
||||
def test_allowlist_accepts_civitai_pth():
|
||||
assert is_url_allowed("https://civitai.com/api/download/models/123.pth")
|
||||
|
||||
|
||||
def test_allowlist_rejects_unknown_host():
|
||||
assert not is_url_allowed("https://example.com/x.safetensors")
|
||||
|
||||
|
||||
def test_allowlist_rejects_api_path_on_hf():
|
||||
# On an allowlisted host but not pointing at a model file.
|
||||
assert not is_url_allowed("https://huggingface.co/api/models")
|
||||
|
||||
|
||||
def test_allowlist_rejects_non_https_except_localhost():
|
||||
assert not is_url_allowed("http://huggingface.co/x/y.safetensors")
|
||||
assert is_url_allowed("http://localhost:8000/x.safetensors")
|
||||
|
||||
|
||||
def test_parse_model_id_valid(patched_folder_paths):
|
||||
assert parse_model_id("loras/foo.safetensors") == ("loras", "foo.safetensors")
|
||||
|
||||
|
||||
def test_parse_model_id_rejects_traversal(patched_folder_paths):
|
||||
with pytest.raises(InvalidModelId):
|
||||
parse_model_id("../etc/passwd")
|
||||
|
||||
|
||||
def test_parse_model_id_rejects_unknown_folder(patched_folder_paths):
|
||||
with pytest.raises(InvalidModelId):
|
||||
parse_model_id("nope/x.safetensors")
|
||||
|
||||
|
||||
def test_parse_model_id_rejects_double_slash(patched_folder_paths):
|
||||
with pytest.raises(InvalidModelId):
|
||||
parse_model_id("loras/sub/x.safetensors")
|
||||
|
||||
|
||||
def test_resolve_existing_returns_path_when_present(model_root, patched_folder_paths):
|
||||
_root, loras_dir, _ = model_root
|
||||
target = loras_dir / "foo.safetensors"
|
||||
target.write_bytes(b"x")
|
||||
assert resolve_existing("loras/foo.safetensors") == str(target)
|
||||
|
||||
|
||||
def test_resolve_existing_returns_none_when_absent(patched_folder_paths):
|
||||
assert resolve_existing("loras/missing.safetensors") is None
|
||||
|
||||
|
||||
def test_resolve_destination_returns_tmp_pair(model_root, patched_folder_paths):
|
||||
_root, loras_dir, _ = model_root
|
||||
final, tmp = resolve_destination("loras/foo.safetensors", epoch=7)
|
||||
assert final == str(loras_dir / "foo.safetensors")
|
||||
# Temp path embeds the session epoch (so cancel+retry can't collide on it)
|
||||
# and uses the subsystem-specific suffix the startup sweep matches.
|
||||
assert tmp == f"{final}.7.comfy-download.tmp"
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# DownloadServer registry: lifecycle, races, cancellation epoch semantics
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
|
||||
def test_register_is_exclusive():
|
||||
server = DownloadServer()
|
||||
s1 = server.try_register("loras/x.safetensors", "https://huggingface.co/a")
|
||||
s2 = server.try_register("loras/x.safetensors", "https://huggingface.co/b")
|
||||
assert s1 is not None
|
||||
assert s2 is None
|
||||
assert server.is_downloading("loras/x.safetensors")
|
||||
|
||||
|
||||
def test_cancel_removes_session():
|
||||
server = DownloadServer()
|
||||
server.try_register("loras/x.safetensors", "https://huggingface.co/a")
|
||||
assert server.cancel("loras/x.safetensors") is True
|
||||
assert not server.is_downloading("loras/x.safetensors")
|
||||
|
||||
|
||||
def test_cancel_returns_false_when_absent():
|
||||
server = DownloadServer()
|
||||
assert server.cancel("loras/never.safetensors") is False
|
||||
|
||||
|
||||
def test_finish_only_clears_matching_epoch():
|
||||
"""If a session is cancelled and a new one for the same id is
|
||||
registered, ``finish`` from the original worker must not evict the
|
||||
newer session."""
|
||||
server = DownloadServer()
|
||||
s_old = server.try_register("loras/x.safetensors", "u1")
|
||||
server.cancel("loras/x.safetensors")
|
||||
s_new = server.try_register("loras/x.safetensors", "u2")
|
||||
assert s_new is not None and s_new.epoch != s_old.epoch
|
||||
# Old worker's late finish() is a no-op:
|
||||
server.finish(s_old)
|
||||
assert server.is_downloading("loras/x.safetensors")
|
||||
server.finish(s_new)
|
||||
assert not server.is_downloading("loras/x.safetensors")
|
||||
|
||||
|
||||
def test_is_active_follows_cancellation():
|
||||
server = DownloadServer()
|
||||
s = server.try_register("loras/x.safetensors", "u")
|
||||
assert server.is_active(s)
|
||||
server.cancel("loras/x.safetensors")
|
||||
assert not server.is_active(s)
|
||||
|
||||
|
||||
def test_update_progress_tracks_fraction():
|
||||
server = DownloadServer()
|
||||
s = server.try_register("loras/x.safetensors", "u")
|
||||
server.update_progress(s, 50, 100)
|
||||
snap = server.snapshot()["loras/x.safetensors"]
|
||||
assert snap.bytes_downloaded == 50
|
||||
assert snap.total_bytes == 100
|
||||
assert snap.progress == 0.5
|
||||
|
||||
|
||||
def test_update_progress_with_unknown_total_keeps_progress_none():
|
||||
server = DownloadServer()
|
||||
s = server.try_register("loras/x.safetensors", "u")
|
||||
server.update_progress(s, 50, None)
|
||||
assert server.snapshot()["loras/x.safetensors"].progress is None
|
||||
|
||||
|
||||
def test_cleanup_orphan_tmp_files(model_root):
|
||||
"""Orphan temp left by a crashed download must be swept on first use,
|
||||
while unrelated *.tmp files in the model dir are left untouched."""
|
||||
_root, loras_dir, _ = model_root
|
||||
orphan = loras_dir / "stale.safetensors.3.comfy-download.tmp"
|
||||
orphan.write_bytes(b"partial")
|
||||
unrelated = loras_dir / "someothertool.tmp"
|
||||
unrelated.write_bytes(b"not ours")
|
||||
mapping = {"loras": ([str(loras_dir)], {".safetensors"})}
|
||||
with patch("folder_paths.folder_names_and_paths", mapping), patch(
|
||||
"folder_paths.get_folder_paths",
|
||||
side_effect=lambda name: mapping.get(name, ([], set()))[0],
|
||||
):
|
||||
server = DownloadServer()
|
||||
assert orphan.exists(), "sweep must not run at construction time"
|
||||
server.sweep_orphan_tmp_files()
|
||||
assert not orphan.exists()
|
||||
assert unrelated.exists(), "unrelated .tmp must not be swept"
|
||||
# Idempotent — a second call is a cheap no-op.
|
||||
server.sweep_orphan_tmp_files()
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Route: POST /api/models-availability-status
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_availability_partitions_correctly(
|
||||
aiohttp_client, app, model_root, fresh_download_server
|
||||
):
|
||||
_root, loras_dir, _ = model_root
|
||||
(loras_dir / "present.safetensors").write_bytes(b"x")
|
||||
fresh_download_server.try_register(
|
||||
"loras/inflight.safetensors", "http://localhost:8000/x.safetensors"
|
||||
)
|
||||
client = await aiohttp_client(app)
|
||||
|
||||
# Stub probes — we're testing state assignment, not network calls.
|
||||
with patch(
|
||||
"app.model_downloader.api.routes.probe_url",
|
||||
new=AsyncMock(return_value=MetadataProbeResult(
|
||||
file_size=None, is_hf_downloadable=None,
|
||||
)),
|
||||
):
|
||||
body = {
|
||||
"models": {
|
||||
"loras/present.safetensors": "http://localhost:8000/p.safetensors",
|
||||
"loras/missing.safetensors": "http://localhost:8000/m.safetensors",
|
||||
"loras/inflight.safetensors": "http://localhost:8000/x.safetensors",
|
||||
}
|
||||
}
|
||||
resp = await client.post("/api/models-availability-status", json=body)
|
||||
assert resp.status == 200
|
||||
data = await resp.json()
|
||||
models = data["models"]
|
||||
assert models["loras/present.safetensors"]["state"] == "available"
|
||||
assert models["loras/missing.safetensors"]["state"] == "missing"
|
||||
assert models["loras/inflight.safetensors"]["state"] == "downloading"
|
||||
assert "hf_auth" in data
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_availability_invalid_id_classified_as_missing(aiohttp_client, app):
|
||||
client = await aiohttp_client(app)
|
||||
with patch(
|
||||
"app.model_downloader.api.routes.probe_url",
|
||||
new=AsyncMock(return_value=MetadataProbeResult(
|
||||
file_size=None, is_hf_downloadable=None,
|
||||
)),
|
||||
):
|
||||
resp = await client.post(
|
||||
"/api/models-availability-status",
|
||||
json={"models": {"../etc/passwd": "http://localhost:8000/x.safetensors"}},
|
||||
)
|
||||
assert resp.status == 200
|
||||
data = await resp.json()
|
||||
assert data["models"]["../etc/passwd"]["state"] == "missing"
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Route: POST /api/download-models — precondition gating
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_download_rejects_url_not_in_allowlist(aiohttp_client, app):
|
||||
client = await aiohttp_client(app)
|
||||
resp = await client.post(
|
||||
"/api/download-models",
|
||||
json={"models": {"loras/x.safetensors": "https://evil.com/x.safetensors"}},
|
||||
)
|
||||
assert resp.status == 400
|
||||
err = (await resp.json())["error"]
|
||||
assert err["code"] == "URL_NOT_ALLOWED"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_download_rejects_already_available(
|
||||
aiohttp_client, app, model_root
|
||||
):
|
||||
_root, loras_dir, _ = model_root
|
||||
(loras_dir / "x.safetensors").write_bytes(b"x")
|
||||
client = await aiohttp_client(app)
|
||||
resp = await client.post(
|
||||
"/api/download-models",
|
||||
json={"models": {
|
||||
"loras/x.safetensors": "https://huggingface.co/a/b/resolve/main/x.safetensors"
|
||||
}},
|
||||
)
|
||||
assert resp.status == 409
|
||||
assert (await resp.json())["error"]["code"] == "ALREADY_AVAILABLE"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_download_rejects_already_downloading(
|
||||
aiohttp_client, app, fresh_download_server
|
||||
):
|
||||
fresh_download_server.try_register(
|
||||
"loras/x.safetensors", "https://huggingface.co/u.safetensors"
|
||||
)
|
||||
client = await aiohttp_client(app)
|
||||
resp = await client.post(
|
||||
"/api/download-models",
|
||||
json={"models": {
|
||||
"loras/x.safetensors": "https://huggingface.co/a/b/resolve/main/x.safetensors"
|
||||
}},
|
||||
)
|
||||
assert resp.status == 409
|
||||
assert (await resp.json())["error"]["code"] == "ALREADY_DOWNLOADING"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_download_rejects_gated_model(aiohttp_client, app):
|
||||
client = await aiohttp_client(app)
|
||||
with patch(
|
||||
"app.model_downloader.api.routes.probe_url",
|
||||
new=AsyncMock(return_value=MetadataProbeResult(file_size=None, is_hf_downloadable=False)),
|
||||
):
|
||||
resp = await client.post(
|
||||
"/api/download-models",
|
||||
json={"models": {
|
||||
"loras/x.safetensors": "https://huggingface.co/g/r/resolve/main/x.safetensors"
|
||||
}},
|
||||
)
|
||||
assert resp.status == 400
|
||||
assert (await resp.json())["error"]["code"] == "MODEL_NOT_DOWNLOADABLE"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_download_rejects_invalid_model_id(aiohttp_client, app):
|
||||
client = await aiohttp_client(app)
|
||||
resp = await client.post(
|
||||
"/api/download-models",
|
||||
json={"models": {"../etc/passwd": "https://huggingface.co/x.safetensors"}},
|
||||
)
|
||||
assert resp.status == 400
|
||||
assert (await resp.json())["error"]["code"] == "INVALID_MODEL_ID"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_download_atomic_failure_does_not_register_partial(
|
||||
aiohttp_client, app, model_root, fresh_download_server
|
||||
):
|
||||
"""If one model in a batch fails, none get registered."""
|
||||
_root, loras_dir, _ = model_root
|
||||
(loras_dir / "already.safetensors").write_bytes(b"x")
|
||||
client = await aiohttp_client(app)
|
||||
resp = await client.post(
|
||||
"/api/download-models",
|
||||
json={
|
||||
"models": {
|
||||
"loras/already.safetensors":
|
||||
"https://huggingface.co/a/b/resolve/main/already.safetensors",
|
||||
"loras/new.safetensors":
|
||||
"https://huggingface.co/a/b/resolve/main/new.safetensors",
|
||||
}
|
||||
},
|
||||
)
|
||||
assert resp.status == 409
|
||||
# The "new" model should not have been registered as part of the
|
||||
# failed batch.
|
||||
assert not fresh_download_server.is_downloading("loras/new.safetensors")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_download_schedules_when_all_preconditions_pass(
|
||||
aiohttp_client, app, fresh_download_server
|
||||
):
|
||||
"""Verify the precondition pass, registration pass, and async
|
||||
scheduling all wire up correctly. We patch the streamer to avoid
|
||||
real HTTP while still letting the route execute end-to-end."""
|
||||
started = asyncio.Event()
|
||||
finish_signal = asyncio.Event()
|
||||
|
||||
async def fake_stream(session):
|
||||
started.set()
|
||||
await finish_signal.wait()
|
||||
from app.model_downloader.download_server import DOWNLOAD_SERVER
|
||||
DOWNLOAD_SERVER.finish(session)
|
||||
return "/dev/null"
|
||||
|
||||
with patch(
|
||||
"app.model_downloader.api.routes.probe_url",
|
||||
new=AsyncMock(return_value=MetadataProbeResult(file_size=42, is_hf_downloadable=True)),
|
||||
), patch(
|
||||
"app.model_downloader.downloader.stream_to_disk", new=fake_stream
|
||||
):
|
||||
client = await aiohttp_client(app)
|
||||
resp = await client.post(
|
||||
"/api/download-models",
|
||||
json={"models": {
|
||||
"loras/new.safetensors":
|
||||
"https://huggingface.co/a/b/resolve/main/new.safetensors"
|
||||
}},
|
||||
)
|
||||
assert resp.status == 202
|
||||
body = await resp.json()
|
||||
assert body["accepted"] is True
|
||||
assert body["scheduled"] == ["loras/new.safetensors"]
|
||||
# Wait for the worker to actually start.
|
||||
await asyncio.wait_for(started.wait(), timeout=2.0)
|
||||
assert fresh_download_server.is_downloading("loras/new.safetensors")
|
||||
finish_signal.set()
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Route: POST /api/cancel-model-download-session
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cancel_removes_active_session(
|
||||
aiohttp_client, app, fresh_download_server
|
||||
):
|
||||
fresh_download_server.try_register(
|
||||
"loras/x.safetensors", "https://huggingface.co/u.safetensors"
|
||||
)
|
||||
client = await aiohttp_client(app)
|
||||
resp = await client.post(
|
||||
"/api/cancel-model-download-session",
|
||||
json={"model_id": "loras/x.safetensors"},
|
||||
)
|
||||
assert resp.status == 200
|
||||
assert (await resp.json())["cancelled"] is True
|
||||
assert not fresh_download_server.is_downloading("loras/x.safetensors")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cancel_returns_404_when_none(aiohttp_client, app):
|
||||
client = await aiohttp_client(app)
|
||||
resp = await client.post(
|
||||
"/api/cancel-model-download-session",
|
||||
json={"model_id": "loras/nothing.safetensors"},
|
||||
)
|
||||
assert resp.status == 404
|
||||
assert (await resp.json())["error"]["code"] == "NOT_DOWNLOADING"
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Unified availability response embeds metadata per id
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_availability_embeds_metadata(aiohttp_client, app):
|
||||
"""``file_size`` + ``is_hf_downloadable`` come back on the same
|
||||
request as the state — no separate metadata endpoint."""
|
||||
results = {
|
||||
"https://huggingface.co/a/b/resolve/main/free.safetensors":
|
||||
MetadataProbeResult(file_size=1024, is_hf_downloadable=True),
|
||||
"https://huggingface.co/g/r/resolve/main/gated.safetensors":
|
||||
MetadataProbeResult(file_size=None, is_hf_downloadable=False),
|
||||
}
|
||||
|
||||
async def fake_probe(url):
|
||||
return results[url]
|
||||
|
||||
with patch(
|
||||
"app.model_downloader.api.routes.probe_url", new=fake_probe
|
||||
):
|
||||
client = await aiohttp_client(app)
|
||||
resp = await client.post(
|
||||
"/api/models-availability-status",
|
||||
json={
|
||||
"models": {
|
||||
"loras/free.safetensors":
|
||||
"https://huggingface.co/a/b/resolve/main/free.safetensors",
|
||||
"loras/gated.safetensors":
|
||||
"https://huggingface.co/g/r/resolve/main/gated.safetensors",
|
||||
}
|
||||
},
|
||||
)
|
||||
assert resp.status == 200
|
||||
models = (await resp.json())["models"]
|
||||
assert models["loras/free.safetensors"]["file_size"] == 1024
|
||||
assert models["loras/free.safetensors"]["is_hf_downloadable"] is True
|
||||
assert models["loras/gated.safetensors"]["file_size"] is None
|
||||
assert models["loras/gated.safetensors"]["is_hf_downloadable"] is False
|
||||
@@ -1,3 +1,5 @@
|
||||
import contextlib
|
||||
import json
|
||||
import time
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
@@ -9,6 +11,40 @@ import requests
|
||||
from helpers import get_asset_filename, trigger_sync_seed_assets
|
||||
|
||||
|
||||
def test_download_svg_forced_to_attachment(http: requests.Session, api_base: str):
|
||||
"""GHSA-779p-m5rp-r4h4 CISA-5 (sibling route): an uploaded SVG must never be
|
||||
served inline from GET /api/assets/{id}/content, or an inline <script> runs
|
||||
in the app origin (stored XSS). Even with disposition=inline requested, a
|
||||
dangerous content type must be forced to application/octet-stream +
|
||||
Content-Disposition: attachment + nosniff. Regression guard for the stale
|
||||
inline blocklist that previously omitted image/svg+xml and ignored the
|
||||
centralized folder_paths.is_dangerous_content_type check.
|
||||
"""
|
||||
svg = b'<svg xmlns="http://www.w3.org/2000/svg"><script>alert(1)</script></svg>'
|
||||
files = {"file": ("evil.svg", svg, "image/svg+xml")}
|
||||
form_data = {
|
||||
"tags": json.dumps(["models", "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"]
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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(
|
||||
"<script>console.log('xss-marker-ghsa-779p')</script>"
|
||||
)
|
||||
|
||||
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(
|
||||
'<?xml version="1.0"?>'
|
||||
'<svg xmlns="http://www.w3.org/2000/svg">'
|
||||
'<script>console.log("xss-marker-ghsa-779p")</script>'
|
||||
"</svg>"
|
||||
)
|
||||
|
||||
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 <script>; it must not be served as image/svg+xml.
|
||||
assert "svg" not in ct.lower(), (
|
||||
f"Content-Type {ct!r} would let a browser render the SVG and execute embedded scripts."
|
||||
)
|
||||
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_js_served_as_octet_stream(aiohttp_client, app, tmp_path):
|
||||
(tmp_path / "evil.js").write_text("alert('xss-marker-ghsa-779p')")
|
||||
|
||||
client = await aiohttp_client(app)
|
||||
resp = await client.get("/userdata/evil.js")
|
||||
|
||||
assert resp.status == 200
|
||||
ct = resp.headers.get("Content-Type", "").lower()
|
||||
# Must not be served as any executable JavaScript content type.
|
||||
assert "javascript" not in ct, (
|
||||
f"Content-Type {ct!r} is an executable JS type."
|
||||
)
|
||||
assert "ecmascript" not in ct, (
|
||||
f"Content-Type {ct!r} is an executable JS type."
|
||||
)
|
||||
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_xml_dialect_served_as_octet_stream(aiohttp_client, app, tmp_path):
|
||||
"""An XML dialect outside the original blocklist (.xslt -> application/xslt+xml)
|
||||
must still be forced to download. This pins the normalised *+xml family rule
|
||||
in folder_paths.is_dangerous_content_type(); a plain set-membership test would
|
||||
have served this inline."""
|
||||
(tmp_path / "evil.xslt").write_text(
|
||||
'<?xml version="1.0"?>'
|
||||
'<xsl:stylesheet version="1.0" '
|
||||
'xmlns:xsl="http://www.w3.org/1999/XSL/Transform">'
|
||||
"<!-- xss-marker-ghsa-779p -->"
|
||||
"</xsl:stylesheet>"
|
||||
)
|
||||
|
||||
client = await aiohttp_client(app)
|
||||
resp = await client.get("/userdata/evil.xslt")
|
||||
|
||||
assert resp.status == 200
|
||||
ct = resp.headers.get("Content-Type", "")
|
||||
assert ct == "application/octet-stream", (
|
||||
f"Content-Type {ct!r}: an *+xml dialect must be forced to octet-stream "
|
||||
f"(it can carry inline script via stylesheet/entity tricks)."
|
||||
)
|
||||
assert resp.headers.get("X-Content-Type-Options") == "nosniff"
|
||||
assert "attachment" in resp.headers.get("Content-Disposition", "")
|
||||
|
||||
|
||||
async def test_benign_txt_still_served(aiohttp_client, app, tmp_path):
|
||||
(tmp_path / "note.txt").write_text("just a harmless note")
|
||||
|
||||
client = await aiohttp_client(app)
|
||||
resp = await client.get("/userdata/note.txt")
|
||||
|
||||
assert resp.status == 200
|
||||
assert await resp.text() == "just a harmless note"
|
||||
ct = resp.headers.get("Content-Type", "")
|
||||
# text/plain is not in the dangerous set, so it is acceptable here. The
|
||||
# defence-in-depth headers must still be present regardless.
|
||||
assert "text/plain" in ct.lower()
|
||||
assert resp.headers.get("X-Content-Type-Options") == "nosniff"
|
||||
assert "attachment" in resp.headers.get("Content-Disposition", "")
|
||||
@@ -0,0 +1,138 @@
|
||||
"""CI unit guard for FIX #5 of GHSA-779p-m5rp-r4h4 — the /view forced-download set.
|
||||
|
||||
Vuln #5 was stored XSS via SVG upload: the /view endpoint's Content-Type
|
||||
blocklist covered text/html, text/javascript, etc. but was missing
|
||||
image/svg+xml, so an uploaded SVG carrying an inline <script> was served as
|
||||
image/svg+xml and executed in the page origin when rendered.
|
||||
|
||||
The /view forced-download decision lives in the view_image closure registered by
|
||||
server.PromptServer.add_routes (server.py ~line 596), which calls
|
||||
`folder_paths.is_dangerous_content_type(content_type)` — a normalising check that
|
||||
strips charset/boundary parameters and casing and folds in the whole */xml and
|
||||
*+xml dialect family — rather than a bypassable raw
|
||||
`content_type in folder_paths.DANGEROUS_CONTENT_TYPES` membership test. On a match
|
||||
it rewrites the response to application/octet-stream with a
|
||||
Content-Disposition: attachment header. server.py cannot be imported in a unit
|
||||
test (importing it spins up the full PromptServer/aiohttp app and its global side
|
||||
effects), so these tests pin the underlying dangerous-content data
|
||||
(folder_paths.DANGEROUS_CONTENT_TYPES) and the normalising is_dangerous_content_type()
|
||||
helper that the closure actually calls.
|
||||
|
||||
The end-to-end /view assertion (upload an SVG, GET /view, confirm the response
|
||||
is not served as image/svg+xml) lives in the live POC at
|
||||
.security/pocs/test_security_ghsa_779p.py::TestViewSvgContentType, which
|
||||
requires a running server. This file is the fast, server-free CI guard on the
|
||||
set contents so the blocklist can't silently regress.
|
||||
"""
|
||||
|
||||
import folder_paths
|
||||
|
||||
|
||||
# Active/renderable content types that must be forced to download. Each of these
|
||||
# can carry an inline <script> (or otherwise execute) in the page origin if a
|
||||
# browser renders it. image/svg+xml is the original missing item that caused
|
||||
# vuln #5.
|
||||
DANGEROUS = [
|
||||
'image/svg+xml',
|
||||
'application/xml',
|
||||
'text/xml',
|
||||
'text/html',
|
||||
'text/html-sandboxed',
|
||||
'application/xhtml+xml',
|
||||
'text/javascript',
|
||||
'application/javascript',
|
||||
'application/x-javascript',
|
||||
'application/ecmascript',
|
||||
'text/css',
|
||||
]
|
||||
|
||||
# Benign image types that browsers display inline and that must keep rendering;
|
||||
# forcing these to download would break legitimate previews.
|
||||
BENIGN_INLINE_IMAGES = [
|
||||
'image/png',
|
||||
'image/jpeg',
|
||||
'image/webp',
|
||||
'image/gif',
|
||||
]
|
||||
|
||||
|
||||
def test_dangerous_content_types_is_a_set():
|
||||
assert isinstance(folder_paths.DANGEROUS_CONTENT_TYPES, set)
|
||||
|
||||
|
||||
def test_svg_is_in_the_blocklist():
|
||||
"""The specific item whose absence caused vuln #5."""
|
||||
assert 'image/svg+xml' in folder_paths.DANGEROUS_CONTENT_TYPES, (
|
||||
"image/svg+xml missing from DANGEROUS_CONTENT_TYPES — this is exactly "
|
||||
"the regression that reopens GHSA-779p-m5rp-r4h4 vuln #5 (stored XSS "
|
||||
"via SVG upload on /view)."
|
||||
)
|
||||
|
||||
|
||||
def test_all_dangerous_types_present():
|
||||
missing = [ct for ct in DANGEROUS if ct not in folder_paths.DANGEROUS_CONTENT_TYPES]
|
||||
assert not missing, (
|
||||
f"DANGEROUS_CONTENT_TYPES is missing required active/renderable types: "
|
||||
f"{missing}. The /view closure only forces a download for content types "
|
||||
f"in this set; anything missing here is served inline and can execute."
|
||||
)
|
||||
|
||||
|
||||
def test_benign_inline_image_types_absent():
|
||||
leaked = [ct for ct in BENIGN_INLINE_IMAGES if ct in folder_paths.DANGEROUS_CONTENT_TYPES]
|
||||
assert not leaked, (
|
||||
f"Benign inline-displayable image types found in DANGEROUS_CONTENT_TYPES: "
|
||||
f"{leaked}. Forcing these to download would break legitimate image "
|
||||
f"previews in /view — they must keep rendering inline."
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# is_dangerous_content_type() — the normalising check the /view and /userdata
|
||||
# handlers now call instead of a raw `in DANGEROUS_CONTENT_TYPES` membership
|
||||
# test. An exact-string membership test was bypassable with a charset parameter
|
||||
# or odd casing, and missed the wider XML dialect family; these tests pin the
|
||||
# normalisation so that bypass can't reopen.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def test_function_matches_plain_dangerous_types():
|
||||
for ct in DANGEROUS:
|
||||
assert folder_paths.is_dangerous_content_type(ct) is True, ct
|
||||
|
||||
|
||||
def test_function_strips_parameters_and_casing():
|
||||
"""A charset/boundary parameter or casing must not slip a type past the check.
|
||||
|
||||
This is the bypass surfaced by review: the /view blake3 branch can serve an
|
||||
attacker-controlled, unvalidated asset mime_type like 'text/html; charset=utf-8',
|
||||
which an exact-string set test missed.
|
||||
"""
|
||||
for ct in (
|
||||
'text/html; charset=utf-8',
|
||||
'TEXT/HTML',
|
||||
'Text/HTML; charset=UTF-8',
|
||||
'image/svg+xml; charset=utf-8',
|
||||
' text/html ',
|
||||
):
|
||||
assert folder_paths.is_dangerous_content_type(ct) is True, ct
|
||||
|
||||
|
||||
def test_function_covers_xml_dialect_family():
|
||||
"""Any *+xml / */xml dialect is dangerous without enumerating each one."""
|
||||
for ct in (
|
||||
'application/xslt+xml',
|
||||
'application/rss+xml',
|
||||
'application/atom+xml',
|
||||
'application/rdf+xml',
|
||||
'application/mathml+xml',
|
||||
'message/rfc822',
|
||||
):
|
||||
assert folder_paths.is_dangerous_content_type(ct) is True, ct
|
||||
|
||||
|
||||
def test_function_allows_benign_and_empty():
|
||||
for ct in BENIGN_INLINE_IMAGES + ['application/octet-stream', 'text/plain']:
|
||||
assert folder_paths.is_dangerous_content_type(ct) is False, ct
|
||||
# None / empty (mimetypes.guess_type miss) must not be treated as dangerous.
|
||||
assert folder_paths.is_dangerous_content_type(None) is False
|
||||
assert folder_paths.is_dangerous_content_type('') is False
|
||||
Reference in New Issue
Block a user